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 prime_invariant() -> bool {
278 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
279 *E.get_or_init(|| std::env::var("MEMRA_PRIME_INVARIANT").as_deref() == Ok("1"))
280 }
281
282 pub fn prime_grain() -> usize {
289 static G: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
290 *G.get_or_init(|| {
291 std::env::var("MEMRA_PRIME_GRAIN").ok()
292 .and_then(|v| v.parse().ok()).unwrap_or(4096)
293 .max(PRIME_MIN_T)
294 })
295 }
296
297 pub fn forward(&self, e: &Engine, tokens: &[u32]) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
299 if self.is_gemma4_e4b() { return self.gemma4_e4b_forward(e, tokens, false); }
300 if self.cfg.gemma4.is_some() { return self.gemma4_forward(e, tokens, false); }
301 let cfg = &self.cfg;
302 let n_embd = cfg.n_embd as usize;
303 let t = tokens.len();
304 let eps = cfg.rms_eps;
305 let pos: Vec<i32> = (0..t as i32).collect();
306 let pos_d = e.htod_i32(&pos)?;
307
308 let mut x = self.embed(e, tokens)?; for (il, layer) in self.layers.iter().enumerate() {
311 let mut h = e.uninit(t * n_embd)?;
313 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
314
315 let mixed = match &layer.mixer {
316 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t)?,
317 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
318 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
319 };
320
321 let mut x1 = e.uninit(t * n_embd)?;
323 e.add(&x, &mixed, &mut x1, t * n_embd)?;
324
325 let mut z = e.uninit(t * n_embd)?;
327 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
328 let ffn_out = match &layer.ffn {
329 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
330 let n_ff = ffn_gate.out_features();
331 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
332 let up = g2.pop().unwrap();
333 let gate = g2.pop().unwrap();
334 let mut act = e.uninit(t * n_ff)?;
335 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, t * n_ff)?;
336 e.matmul(ffn_down, &act, t)?
337 }
338 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
339 };
340 let mut x2 = e.uninit(t * n_embd)?;
341 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
342 x = x2;
343 }
344
345 let mut hn = e.uninit(t * n_embd)?;
346 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
347 let logits = e.matmul(&self.output, &hn, t)?;
348 Ok(e.dtoh(&logits)?)
349 }
350
351 pub fn forward_last(&self, e: &Engine, tokens: &[u32]) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
357 if self.cfg.gemma4.is_some() { return self.gemma4_forward(e, tokens, true); }
358 let cfg = &self.cfg;
359 let n_embd = cfg.n_embd as usize;
360 let t = tokens.len();
361 let eps = cfg.rms_eps;
362 let pos: Vec<i32> = (0..t as i32).collect();
363 let pos_d = e.htod_i32(&pos)?;
364
365 let mut x = self.embed(e, tokens)?; let probe = std::env::var("MEMRA_LAYER_PROBE").is_ok();
369 for (il, layer) in self.layers.iter().enumerate() {
370 let mut h = e.uninit(t * n_embd)?;
371 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
372 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} norm ok"); }
373 let mixed = match &layer.mixer {
374 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t)?,
375 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
376 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
377 };
378 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} mixer ok"); }
379 let mut x1 = e.uninit(t * n_embd)?;
380 e.add(&x, &mixed, &mut x1, t * n_embd)?;
381 let mut z = e.uninit(t * n_embd)?;
382 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
383 let ffn_out = match &layer.ffn {
384 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
385 let n_ff = ffn_gate.out_features();
386 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
387 let up = g2.pop().unwrap();
388 let gate = g2.pop().unwrap();
389 let mut act = e.uninit(t * n_ff)?;
390 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, t * n_ff)?;
391 e.matmul(ffn_down, &act, t)?
392 }
393 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
394 };
395 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} ffn ok"); }
396 let mut x2 = e.uninit(t * n_embd)?;
397 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
398 x = x2;
399 }
400 let mut hn = e.uninit(t * n_embd)?;
402 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
403 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)?;
406 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
407 let logits = e.matmul(&self.output, &hlast, 1)?; Ok(e.dtoh(&logits)?)
409 }
410
411 pub fn prime_cache(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
428 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
429 let n_embd = self.cfg.n_embd as usize;
430 let t = tokens.len();
431 assert!(t >= PRIME_MIN_T, "prime_cache needs T >= {PRIME_MIN_T} (caller gates)");
435 assert!(cache.pos + t <= cache.max_ctx, "prime_cache: prompt exceeds cache max_ctx");
436
437 if self.is_gemma4_e4b() {
449 return self.gemma4_e4b_prime(e, tokens, cache);
450 }
451 if self.cfg.gemma4.is_some() {
452 return self.gemma4_prime(e, tokens, cache);
454 }
455 let mut chunk: usize = std::env::var("MEMRA_PRIME_CHUNK").ok()
456 .and_then(|v| v.parse().ok()).unwrap_or(4096);
457 if Self::prime_invariant() {
483 chunk = Self::prime_grain();
484 }
485 if chunk == 0 || t <= chunk {
486 return self.prime_chunk(e, tokens, cache);
487 }
488 let mut hiddens = e.uninit(t * n_embd)?;
489 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
490 let mut start = 0usize;
491 while start < t {
492 let mut end = (start + chunk).min(t);
494 if t - end > 0 && t - end < PRIME_MIN_T { end = t; }
495 let (l, hs, x) = self.prime_chunk(e, &tokens[start..end], cache)?;
496 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
497 last = Some((l, hs));
498 start = end;
499 }
500 let (logits, h_seed) = last.unwrap();
501 Ok((logits, h_seed, hiddens))
502 }
503
504 fn gdn_hk(e: &Engine, t: usize, num_v: usize, num_k: usize) -> usize {
511 if Engine::gdn_db_on()
512 && Engine::gdn_chunked_enabled() && t >= 16
513 && e.gdn_mma_enabled(Engine::gdn_chunk_size())
514 && num_k * 2 == num_v
515 {
516 num_k
517 } else {
518 num_v
519 }
520 }
521
522 fn f16out_on(e: &Engine, t: usize) -> bool {
527 crate::f16_ffi::pp_f16_enabled() && t >= 16 && !e.verify_exact_on()
528 && std::env::var("MEMRA_F16OUT").as_deref() != Ok("0")
529 }
530
531 pub fn prime_slabs_get(&self, e: &Engine, t: usize, n_embd: usize, n_ff_max: usize)
534 -> Result<std::sync::MutexGuard<'_, Option<PrimeSlabs>>, Box<dyn std::error::Error>> {
535 let mut g = self.prime_slabs.lock().unwrap();
536 let need_new = match g.as_ref() { None => true, Some(sl) => sl.t_cap < t };
537 if need_new {
538 *g = Some(PrimeSlabs {
539 t_cap: t,
540 h: e.uninit(t * n_embd)?,
541 x1: e.uninit(t * n_embd)?,
542 z: e.uninit(t * n_embd)?,
543 act: e.uninit(t * n_ff_max)?,
544 xa: e.uninit(t * n_embd)?,
545 xb: e.uninit(t * n_embd)?,
546 h16: e.alloc_u8_uninit(t * n_embd * 2)?,
547 z16: e.alloc_u8_uninit(t * n_embd * 2)?,
548 gate: e.uninit(t * n_ff_max)?,
549 up: e.uninit(t * n_ff_max)?,
550 ffn_out: e.uninit(t * n_embd)?,
551 seg_glue: Vec::new(),
552 mixed: e.uninit(t * n_embd)?,
553 seg_mid: Vec::new(),
554 seg_t: 0,
555 });
556 }
557 Ok(g)
558 }
559
560 fn prime_chunk(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
561 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
562 let cfg = &self.cfg;
563 let n_embd = cfg.n_embd as usize;
564 let t = tokens.len();
565 let eps = cfg.rms_eps;
566 let base = cache.pos;
567 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
568 let pos_d = e.htod_i32(&pos)?;
569
570 let x_embed = self.embed(e, tokens)?; let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
575 let n_ff_max = self.layers.iter().map(|l| match &l.ffn {
581 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
582 _ => n_embd,
583 }).max().unwrap_or(n_embd).max(n_embd);
584 let use_slabs = std::env::var("MEMRA_PRIME_SLABS").as_deref() != Ok("0");
585 let mut slab_guard = if use_slabs {
586 Some(self.prime_slabs_get(e, t, n_embd, n_ff_max)?)
587 } else {
588 None
589 };
590 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>);
592 let (mut x_cur, mut x_nxt, sl): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, Option<SlabRefs>);
593 let mut seg: Option<(&mut Vec<Option<cudarc::driver::CudaGraph>>, &mut Vec<Option<cudarc::driver::CudaGraph>>, &mut CudaSlice<f32>, &mut usize)> = None;
594 let mut x_own2;
595 match slab_guard.as_mut() {
596 Some(g) => {
597 let slabs = g.as_mut().unwrap();
598 e.copy_into(&mut slabs.xa, 0, &x_embed, t * n_embd)?;
599 let PrimeSlabs { xa, xb, h, x1, z, act, h16, z16, gate, up, ffn_out, seg_glue, mixed, seg_mid, seg_t, .. } = slabs;
600 x_cur = xa;
601 x_nxt = xb;
602 seg = Some((seg_glue, seg_mid, mixed, seg_t));
603 sl = Some((h, x1, z, act, h16, z16, gate, up, ffn_out));
604 }
605 None => {
606 x_own = x_embed;
607 x_own2 = e.uninit(t * n_embd)?;
608 x_cur = &mut x_own;
609 x_nxt = &mut x_own2;
610 sl = None;
611 }
612 }
613 let mut alloc_h; let mut alloc_x1; let mut alloc_z; let mut alloc_act;
614 let mut alloc_h16; let mut alloc_z16;
615 let mut alloc_gate; let mut alloc_up; let mut alloc_fo;
616 let (h, x1, z, act): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
617 let (h16, z16): (&mut CudaSlice<u8>, &mut CudaSlice<u8>);
618 let (sl_gate, sl_up, sl_fo): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
619 match sl {
620 Some((a, b, c, d, e16, f16b, g, u, fo)) => {
621 h = a; x1 = b; z = c; act = d; h16 = e16; z16 = f16b;
622 sl_gate = g; sl_up = u; sl_fo = fo;
623 }
624 None => {
625 alloc_h = e.uninit(t * n_embd)?;
626 alloc_x1 = e.uninit(t * n_embd)?;
627 alloc_z = e.uninit(t * n_embd)?;
628 alloc_act = e.uninit(t * n_ff_max)?;
629 alloc_h16 = e.alloc_u8_uninit(t * n_embd * 2)?;
630 alloc_z16 = e.alloc_u8_uninit(t * n_embd * 2)?;
631 alloc_gate = e.uninit(t * n_ff_max)?;
632 alloc_up = e.uninit(t * n_ff_max)?;
633 alloc_fo = e.uninit(t * n_embd)?;
634 h = &mut alloc_h; x1 = &mut alloc_x1; z = &mut alloc_z; act = &mut alloc_act;
635 h16 = &mut alloc_h16; z16 = &mut alloc_z16;
636 sl_gate = &mut alloc_gate; sl_up = &mut alloc_up; sl_fo = &mut alloc_fo;
637 }
638 }
639 let n_layers = self.layers.len();
644 let use_seg = f16fuse && seg.is_some()
651 && std::env::var("MEMRA_PRIME_SEG").as_deref() == Ok("1");
652 if let Some((sg, sm, _, st)) = seg.as_mut() {
653 if **st != t {
654 sg.clear();
655 sg.extend((0..n_layers).map(|_| None));
656 sm.clear();
657 sm.extend((0..n_layers).map(|_| None));
658 **st = t;
659 }
660 }
661 {
662 let layer0 = &self.layers[0];
663 if f16fuse {
664 e.rms_norm_f16out(x_cur, layer0.attn_norm.float_data(), h, h16, n_embd, t, eps)?;
665 } else {
666 e.rms_norm(x_cur, layer0.attn_norm.float_data(), h, n_embd, t, eps)?;
667 }
668 }
669 for (il, layer) in self.layers.iter().enumerate() {
670 let hx16 = if f16fuse { Some(&*h16) } else { None };
671 if use_seg {
672 let (pre, pre16, w_out) = match &layer.mixer {
675 Mixer::Full(fa) => {
676 let g3 = match hx16 {
677 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
678 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
679 };
680 let (pre, pre16) = self.full_attn_prime_core_inner(e, fa, g3, &pos_d, t, cache, il)?;
681 (pre, pre16, &fa.wo)
682 }
683 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
684 Mixer::Linear(la) => {
685 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
686 let g4 = match hx16 {
687 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
688 None => e.matmul_group(&ws, h, t)?,
689 };
690 let (pre, pre16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, None)?;
691 (pre, pre16, &la.ssm_out)
692 }
693 };
694 {
695 let (_, sm, mslab, _) = seg.as_mut().unwrap();
696 let pre_n = pre.len() / t;
697 let xh_pre = match pre16 {
698 Some(x) => x,
699 None => e.f16_act(&pre, t * pre_n, pre_n)?,
700 };
701 if !e.try_f16_gemm_pre_into(w_out, &xh_pre, t, mslab)? {
702 let y = e.matmul(w_out, &pre, t)?;
703 e.copy_into(mslab, 0, &y, t * n_embd)?;
704 }
705 if sm[il].is_none() {
706 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
707 let w_post = layer.post_attn_norm.float_data();
708 e.stream().synchronize()?;
709 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
710 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
711 e.add(x_cur, mslab, x1, t * n_embd)?;
712 e.rms_norm_f16out(x1, w_post, z, z16, n_embd, t, eps)?;
713 Ok(())
714 })();
715 let g = e.stream().end_capture(
716 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
717 r?;
718 sm[il] = Some(g?.ok_or("S-mid capture produced no graph")?);
719 }
720 sm[il].as_ref().unwrap().launch()?;
721 }
722 } else {
723 let mixed = match &layer.mixer {
724 Mixer::Full(fa) => self.full_attn_prime(e, fa, h, hx16, &pos_d, t, cache, il)?,
725 Mixer::Linear(la) => self.linear_attn_prime(e, la, h, hx16, t, cache, il)?,
726 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
727 };
728 if f16fuse {
729 e.add_rms_norm_f16out(x_cur, &mixed, layer.post_attn_norm.float_data(),
732 x1, z, z16, n_embd, t, eps)?;
733 } else {
734 e.add(x_cur, &mixed, x1, t * n_embd)?;
735 e.rms_norm(x1, layer.post_attn_norm.float_data(), z, n_embd, t, eps)?;
736 }
737 }
738 let zx16 = if f16fuse { Some(&*z16) } else { None };
739 match &layer.ffn {
740 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
741 let n_ff = ffn_gate.out_features();
742 let mut into_ok = false;
745 if let Some(xh) = zx16 {
746 into_ok = e.try_f16_gemm_pre_into(ffn_gate, xh, t, sl_gate)?
747 && e.try_f16_gemm_pre_into(ffn_up, xh, t, sl_up)?;
748 }
749 if !into_ok {
750 let mut g2 = match zx16 {
751 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], z, xh, t)?,
752 None => e.matmul_group(&[ffn_gate, ffn_up], z, t)?,
753 };
754 let up_y = g2.pop().unwrap();
755 let gate_y = g2.pop().unwrap();
756 e.copy_into(sl_gate, 0, &gate_y, t * n_ff)?;
757 e.copy_into(sl_up, 0, &up_y, t * n_ff)?;
758 }
759 let act16 = if Self::f16out_on(e, t) && self.cfg.m3.is_none() {
762 let mut a16 = e.alloc_u8_uninit(t * n_ff * 2)?;
763 e.silu_mul_f16out(sl_gate, sl_up, act, &mut a16, t * n_ff)?;
764 Some(a16)
765 } else {
766 Self::ffn_act(e, &self.cfg, sl_gate, sl_up, act, t * n_ff)?;
767 None
768 };
769 let xh_act = match act16 {
771 Some(x) => x,
772 None => e.f16_act(act, t * n_ff, n_ff)?,
773 };
774 if !e.try_f16_gemm_pre_into(ffn_down, &xh_act, t, sl_fo)? {
775 let y = e.matmul(ffn_down, &*act, t)?;
776 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
777 }
778 }
779 crate::hybrid::Ffn::Moe(m) => {
780 let y = self.moe_ffn_il(e, m, z, t, il as u16)?;
781 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
782 }
783 }
784 if use_seg && il + 1 < n_layers {
785 let w_next = self.layers[il + 1].attn_norm.float_data();
787 let (sg, _, _, _) = seg.as_mut().unwrap();
788 if sg[il].is_none() {
789 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
790 e.stream().synchronize()?;
791 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
792 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
793 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
794 e.rms_norm_f16out(x_nxt, w_next, h, h16, n_embd, t, eps)?;
795 Ok(())
796 })();
797 let g = e.stream().end_capture(
798 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
799 r?;
800 sg[il] = Some(g?.ok_or("S-glue capture produced no graph")?);
801 }
802 sg[il].as_ref().unwrap().launch()?;
803 } else {
804 if il + 1 < n_layers {
805 let w_next = self.layers[il + 1].attn_norm.float_data();
806 if f16fuse {
807 e.add_rms_norm_f16out(x1, sl_fo, w_next, x_nxt, h, h16, n_embd, t, eps)?;
808 } else {
809 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
810 e.rms_norm(x_nxt, w_next, h, n_embd, t, eps)?;
811 }
812 } else {
813 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
814 }
815 }
816 if let Some(path) = Self::prime_trace_path() {
822 let row = (base + t - 1) as usize;
823 let host = e.dtoh(x_nxt)?;
824 let last = &host[(t - 1) * n_embd..t * n_embd];
825 use std::io::Write as _;
826 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
827 let mut h64: u64 = 0xcbf29ce484222325;
828 for v in last {
829 h64 ^= v.to_bits() as u64;
830 h64 = h64.wrapping_mul(0x100000001b3);
831 }
832 writeln!(f, "{{\"pos\":{row},\"layer\":{il},\"t\":{t},\"base\":{base},\
833 \"hash\":\"{h64:016x}\",\"v0\":{:.9e},\"v1\":{:.9e},\"v2\":{:.9e}}}",
834 last[0], last[1], last[2])?;
835 }
836 std::mem::swap(&mut x_cur, &mut x_nxt);
837 }
838 let mut x = e.uninit(t * n_embd)?;
840 e.copy_into(&mut x, 0, x_cur, t * n_embd)?;
841 drop(slab_guard);
842
843 let mut h_seed = e.uninit(n_embd)?;
847 if !crate::spec::spec_hpost() {
848 e.copy_view_into(&mut h_seed, 0, &x.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
849 }
850 let mut hn = e.uninit(t * n_embd)?;
852 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
853 if crate::spec::spec_hpost() {
854 e.copy_view_into(&mut h_seed, 0, &hn.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
855 }
856 let last = e.view(&hn, t * n_embd);
857 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
858 let mut hlast = e.uninit(n_embd)?;
859 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
860 let logits = e.matmul(&self.output, &hlast, 1)?;
861 cache.pos += t;
862 Ok((e.dtoh(&logits)?, h_seed, if crate::spec::spec_hpost() { hn } else { x }))
865 }
866
867 pub fn prime_chunk_captured(&self, e: &Engine, x_in: &CudaSlice<f32>, pos_d: &CudaSlice<i32>,
883 t: usize, cache: &mut Cache,
884 len_d: &CudaSlice<i32>,
885 logits_out: &mut CudaSlice<f32>, h_seed_out: &mut CudaSlice<f32>)
886 -> Result<(), Box<dyn std::error::Error>> {
887 let cfg = &self.cfg;
888 let n_embd = cfg.n_embd as usize;
889 let eps = cfg.rms_eps;
890 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
891 let mut x = e.uninit(t * n_embd)?;
892 e.copy_into(&mut x, 0, x_in, t * n_embd)?;
893 for (il, layer) in self.layers.iter().enumerate() {
894 let mut h = e.uninit(t * n_embd)?;
895 let mut hx16: Option<CudaSlice<u8>> = None;
896 if f16fuse {
897 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
898 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut b16, n_embd, t, eps)?;
899 hx16 = Some(b16);
900 } else {
901 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
902 }
903 let mixed = match &layer.mixer {
904 Mixer::Full(fa) => self.full_attn_prime(e, fa, &h, hx16.as_ref(), pos_d, t, cache, il)?,
905 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
906 Mixer::Linear(la) => {
907 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
908 let g4 = match hx16.as_ref() {
909 Some(xh) => e.matmul_group_xh(&ws, &h, xh, t)?,
910 None => e.matmul_group(&ws, &h, t)?,
911 };
912 self.linear_attn_prime_core_pad(e, la, g4, t, cache, il, Some(len_d))?
913 }
914 };
915 let mut x1 = e.uninit(t * n_embd)?;
916 e.add(&x, &mixed, &mut x1, t * n_embd)?;
917 let mut z = e.uninit(t * n_embd)?;
918 let mut zx16: Option<CudaSlice<u8>> = None;
919 if f16fuse {
920 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
921 e.rms_norm_f16out(&x1, layer.post_attn_norm.float_data(), &mut z, &mut b16, n_embd, t, eps)?;
922 zx16 = Some(b16);
923 } else {
924 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
925 }
926 let ffn_out = match &layer.ffn {
927 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
928 let n_ff = ffn_gate.out_features();
929 let mut g2 = match &zx16 {
930 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], &z, xh, t)?,
931 None => e.matmul_group(&[ffn_gate, ffn_up], &z, t)?,
932 };
933 let up = g2.pop().unwrap();
934 let gate = g2.pop().unwrap();
935 let mut act = e.uninit(t * n_ff)?;
936 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, t * n_ff)?;
937 e.matmul(ffn_down, &act, t)?
938 }
939 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
940 };
941 let mut x2 = e.uninit(t * n_embd)?;
942 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
943 x = x2;
944 }
945 if !crate::spec::spec_hpost() {
947 e.row_gather_dev(&x, h_seed_out, len_d, n_embd)?;
948 }
949 let mut hn = e.uninit(t * n_embd)?;
950 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
951 if crate::spec::spec_hpost() {
952 e.row_gather_dev(&hn, h_seed_out, len_d, n_embd)?;
953 }
954 let mut hlast = e.uninit(n_embd)?;
955 e.row_gather_dev(&hn, &mut hlast, len_d, n_embd)?;
956 let logits = e.matmul(&self.output, &hlast, 1)?;
957 let nv = logits.len();
958 e.copy_into(logits_out, 0, &logits, nv)?;
959 Ok(())
960 }
961
962 pub fn prime_cache_batch(&self, e: &Engine, prompts: &[&[u32]], caches: &mut [&mut Cache])
979 -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
980 let cfg = &self.cfg;
981 let n_embd = cfg.n_embd as usize;
982 let eps = cfg.rms_eps;
983 let b = prompts.len();
984 assert!(b >= 1 && b == caches.len());
985 let pos0s: Vec<usize> = caches.iter().map(|c| c.pos).collect();
986 let carried = pos0s.iter().any(|&p| p > 0);
987 if carried && cfg.gemma4.is_some() {
988 return Err("prime_cache_batch: gemma4 has no continuation prime (v0 fresh-only)".into());
989 }
990 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
991 for &t in &ts { assert!(t >= PRIME_MIN_T, "prime_cache_batch needs T >= {PRIME_MIN_T}"); }
992 for (s, c) in caches.iter().enumerate() {
993 assert!(c.pos + ts[s] <= c.max_ctx, "prime_cache_batch: prompt exceeds cache max_ctx");
994 }
995 let total: usize = ts.iter().sum();
996 let offs: Vec<usize> = ts.iter().scan(0usize, |a, &t| { let o = *a; *a += t; Some(o) }).collect();
997 let pos_ds: Vec<CudaSlice<i32>> = ts.iter().zip(&pos0s)
999 .map(|(&t, &p0)| e.htod_i32(&(p0 as i32..(p0 + t) as i32).collect::<Vec<_>>()))
1000 .collect::<Result<_, _>>()?;
1001 let split = |e: &Engine, y: &CudaSlice<f32>, dim: usize|
1003 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
1004 let mut out = Vec::with_capacity(b);
1005 for s in 0..b {
1006 let mut ys = e.uninit(ts[s] * dim)?;
1007 e.copy_view_into(&mut ys, 0, &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim), ts[s] * dim)?;
1008 out.push(ys);
1009 }
1010 Ok(out)
1011 };
1012
1013 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
1014 let mut x = self.embed(e, &cat_tokens)?; for (il, layer) in self.layers.iter().enumerate() {
1016 let mut h = e.uninit(total * n_embd)?;
1017 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
1018 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut hx16, n_embd, total, eps)?;
1019 let mut mixed = e.uninit(total * n_embd)?;
1021 match &layer.mixer {
1022 Mixer::Full(fa) => {
1023 let g3 = e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], &h, &hx16, total)?;
1024 let (n_head, n_head_kv, head_dim) =
1030 (self.cfg.n_head as usize, self.cfg.n_head_kv as usize, self.cfg.head_dim_k as usize);
1031 let fa_scale = 1.0 / (head_dim as f32).sqrt();
1032 let use_favl = !carried
1033 && (2..=8).contains(&b)
1034 && (head_dim == 256 || head_dim == 128)
1035 && self.cfg.attn_out_gate()
1036 && std::env::var("MEMRA_NOFA").is_err()
1037 && std::env::var("MEMRA_FA_FLOOR").is_err()
1038 && std::env::var("MEMRA_FA_PP_W2").as_deref() != Ok("1")
1039 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0")
1040 && std::env::var("MEMRA_FA_VL").as_deref() != Ok("0");
1041 if use_favl {
1042 let (qf_w, kf_w, vf_w) =
1043 (fa.wq.out_features(), fa.wk.out_features(), fa.wv.out_features());
1044 struct APre {
1045 q: CudaSlice<f32>, gate: Option<CudaSlice<f32>>,
1046 qn: CudaSlice<f32>, kn: CudaSlice<f32>,
1047 }
1048 let mut aps = Vec::with_capacity(b);
1049 for &t in ts.iter().take(b) {
1050 aps.push(APre {
1051 q: e.uninit(t * n_head * head_dim)?,
1052 gate: Some(e.uninit(t * n_head * head_dim)?),
1053 qn: e.uninit(t * n_head * head_dim)?,
1054 kn: e.uninit(t * n_head_kv * head_dim)?,
1055 });
1056 }
1057 let (kv_dim_k, kv_dim_v, ktb, vtb) = {
1058 let kvl = caches[0].kv[il].as_ref().unwrap();
1059 (kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes)
1060 };
1061 let pargs: Vec<crate::AttnPreVl> = (0..b).map(|s| {
1062 let (o, t) = (offs[s], ts[s]);
1063 let kvl = caches[s].kv[il].as_ref().unwrap();
1064 assert!(kvl.len == 0 && kvl.len + t <= caches[s].max_ctx,
1065 "prime_cache_batch attn vl: fresh + capacity");
1066 crate::AttnPreVl {
1067 qf: e.addr_f32v(&g3[0].slice(o * qf_w..(o + t) * qf_w)),
1068 kf: e.addr_f32v(&g3[1].slice(o * kf_w..(o + t) * kf_w)),
1069 vf: e.addr_f32v(&g3[2].slice(o * vf_w..(o + t) * vf_w)),
1070 q: e.addr_f32(&aps[s].q),
1071 gate: e.addr_f32(aps[s].gate.as_ref().unwrap()),
1072 qn: e.addr_f32(&aps[s].qn), kn: e.addr_f32(&aps[s].kn),
1073 kc: e.addr_u8(&kvl.k), vc: e.addr_u8(&kvl.v),
1074 t: t as i32, pad: 0,
1075 }
1076 }).collect();
1077 e.attn_pre_vl8(&pargs, fa.q_norm.float_data(), fa.k_norm.float_data(),
1078 head_dim, self.cfg.rope_dim_count as usize, n_head, n_head_kv,
1079 self.cfg.rms_eps, self.cfg.rope_freq_base, 1.0,
1080 kv_dim_k, kv_dim_v, ktb, vtb)?;
1081 for s in 0..b {
1082 let kvl = caches[s].kv[il].as_mut().unwrap();
1083 kvl.len += ts[s];
1084 let new_len = kvl.len as i32;
1085 e.set_i32_one(&mut kvl.len_d, new_len)?;
1086 }
1087 let mut attns = Vec::with_capacity(b);
1088 let mut mirrors = Vec::with_capacity(b);
1089 for &t in ts.iter().take(b) {
1090 attns.push(e.uninit(t * n_head * head_dim)?);
1091 let n = t * n_head_kv * head_dim;
1092 mirrors.push((e.alloc_u8_uninit(n * 2)?, e.alloc_u8_uninit(n * 2)?));
1093 }
1094 let fa3_on = match std::env::var("MEMRA_FA3").as_deref() {
1097 Ok("0") => false,
1098 Ok("1") => true,
1099 _ => cfg!(memra_hopper_mma),
1100 };
1101 if fa3_on {
1102 let mut q16s = Vec::with_capacity(b);
1103 let mut v16s = Vec::with_capacity(b);
1104 for s in 0..b {
1105 let t = ts[s];
1106 let mut q16 = e.alloc_u8_uninit(t * n_head * head_dim * 2)?;
1107 e.f32_to_bf16_into(&aps[s].qn, &mut q16, t * n_head * head_dim)?;
1108 let mut k16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
1109 e.f32_to_bf16_into(&aps[s].kn, &mut k16, t * n_head_kv * head_dim)?;
1110 let mut v16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
1111 e.f32_to_bf16_v(&g3[2].slice(offs[s] * vf_w..(offs[s] + t) * vf_w),
1112 &mut v16, t * n_head_kv * head_dim)?;
1113 q16s.push(q16);
1114 v16s.push((k16, v16));
1115 }
1116 let mut qp = [core::ptr::null::<core::ffi::c_void>(); 8];
1117 let mut kp = qp;
1118 let mut vp = qp;
1119 let mut op = [core::ptr::null_mut::<f32>(); 8];
1120 let mut tsv = [0i32; 8];
1121 for s in 0..b {
1122 qp[s] = e.addr_u8(&q16s[s]) as *const core::ffi::c_void;
1123 kp[s] = e.addr_u8(&v16s[s].0) as *const core::ffi::c_void;
1124 vp[s] = e.addr_u8(&v16s[s].1) as *const core::ffi::c_void;
1125 op[s] = e.addr_f32(&attns[s]) as *mut f32;
1126 tsv[s] = ts[s] as i32;
1127 }
1128 let rc = unsafe {
1129 crate::fa3_vl_raw(qp.as_ptr(), kp.as_ptr(), vp.as_ptr(), op.as_ptr(),
1130 tsv.as_ptr(), b as i32, n_head as i32,
1131 n_head_kv as i32, head_dim as i32, fa_scale,
1132 e.stream().cu_stream() as *mut core::ffi::c_void)
1133 };
1134 if rc != 0 {
1135 return Err(format!("memra_fa3_vl rc={rc}").into());
1136 }
1137 } else {
1138 let fargs: Vec<crate::FaSeqVl> = (0..b).map(|s| crate::FaSeqVl {
1139 q: e.addr_f32(&aps[s].qn), k16: e.addr_u8(&mirrors[s].0),
1140 v16: e.addr_u8(&mirrors[s].1), o: e.addr_f32(&attns[s]),
1141 kf: e.addr_f32(&aps[s].kn),
1142 vf: e.addr_f32v(&g3[2].slice(offs[s] * vf_w..(offs[s] + ts[s]) * vf_w)),
1143 t: ts[s] as i32, pad: 0,
1144 }).collect();
1145 e.fa_prefill_vl8(&fargs, head_dim, n_head, n_head_kv, fa_scale)?;
1146 }
1147 for (s, attn) in attns.into_iter().enumerate() {
1148 let (attn_g, ag16) = self.full_attn_prime_post_fa(
1149 e, attn, &aps[s].gate, ts[s], n_head, head_dim)?;
1150 let mut done = false;
1151 if let Some(xh) = &ag16 {
1152 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
1153 }
1154 if !done {
1155 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
1156 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
1157 }
1158 }
1159 } else {
1160 let mut parts: Vec<Vec<CudaSlice<f32>>> = (0..b).map(|_| Vec::new()).collect();
1161 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g3) {
1162 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
1163 parts[s].push(ys);
1164 }
1165 }
1166 for (s, g3s) in parts.into_iter().enumerate() {
1167 let (attn_g, ag16) = self.full_attn_prime_core_inner(
1169 e, fa, g3s, &pos_ds[s], ts[s], caches[s], il)?;
1170 let mut done = false;
1171 if let Some(xh) = &ag16 {
1172 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
1173 }
1174 if !done {
1175 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
1176 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
1177 }
1178 }
1179 }
1180 }
1181 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1182 Mixer::Linear(la) => {
1183 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1188 let g4 = e.matmul_group_xh(&ws, &h, &hx16, total)?;
1189 let outs = self.linear_attn_prime_core_batch(e, la, &g4, &offs, &ts, caches, il)?;
1190 for (s, (gn, gn16)) in outs.into_iter().enumerate() {
1191 let (o, t) = (offs[s], ts[s]);
1192 let mut done = false;
1193 if let Some(xh) = &gn16 {
1194 done = e.try_f16_gemm_pre_into_off(&la.ssm_out, xh, t, &mut mixed, o * n_embd)?;
1195 }
1196 if !done {
1197 let m = e.matmul(&la.ssm_out, &gn, t)?;
1198 e.copy_into(&mut mixed, o * n_embd, &m, t * n_embd)?;
1199 }
1200 }
1201 }
1202 }
1203 let mut x1 = e.uninit(total * n_embd)?;
1204 let mut z = e.uninit(total * n_embd)?;
1205 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
1206 e.add_rms_norm_f16out(&x, &mixed, layer.post_attn_norm.float_data(),
1207 &mut x1, &mut z, &mut zx16, n_embd, total, eps)?;
1208 let ffn_out = match &layer.ffn {
1209 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
1210 let n_ff = ffn_gate.out_features();
1211 let mut g2 = e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?;
1212 let up = g2.pop().unwrap();
1213 let gate = g2.pop().unwrap();
1214 let mut act = e.uninit(total * n_ff)?;
1215 if Self::f16out_on(e, total) && self.cfg.m3.is_none() {
1218 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
1219 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
1220 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
1221 Some(y) => y,
1222 None => e.matmul(ffn_down, &act, total)?,
1223 }
1224 } else {
1225 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, total * n_ff)?;
1226 e.matmul(ffn_down, &act, total)?
1227 }
1228 }
1229 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, total, il as u16)?,
1230 };
1231 let mut x2 = e.uninit(total * n_embd)?;
1232 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
1233 x = x2;
1234 }
1235 let mut hn = e.uninit(total * n_embd)?;
1237 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, total, eps)?;
1238 let mut hcat = e.uninit(b * n_embd)?;
1244 for s in 0..b {
1245 let last0 = (offs[s] + ts[s] - 1) * n_embd;
1246 e.copy_view_into(&mut hcat, s * n_embd, &hn.slice(last0..last0 + n_embd), n_embd)?;
1247 }
1248 let logits_cat = if b >= 2 { e.try_f16_gemm(&self.output, &hcat, b)? } else { None };
1249 let logits_host: Option<Vec<f32>> = match &logits_cat {
1250 Some(lc) => Some(e.dtoh(lc)?),
1251 None => None,
1252 };
1253 let n_vocab = self.output.out_features();
1254 let mut hidden_all = if crate::spec::spec_hpost() {
1255 split(e, &hn, n_embd)?
1256 } else {
1257 split(e, &x, n_embd)?
1258 };
1259 let mut out = Vec::with_capacity(b);
1260 for s in 0..b {
1261 let last0 = (offs[s] + ts[s] - 1) * n_embd;
1262 let mut h_seed = e.uninit(n_embd)?;
1263 if !crate::spec::spec_hpost() {
1264 e.copy_view_into(&mut h_seed, 0, &x.slice(last0..last0 + n_embd), n_embd)?;
1265 } else {
1266 e.copy_view_into(&mut h_seed, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
1267 }
1268 let logits = match &logits_host {
1269 Some(lh) => lh[s * n_vocab..(s + 1) * n_vocab].to_vec(),
1270 None => {
1271 let mut hlast = e.uninit(n_embd)?;
1272 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
1273 e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?
1274 }
1275 };
1276 caches[s].pos += ts[s];
1277 out.push((logits, h_seed, hidden_all.remove(0)));
1278 }
1279 Ok(out)
1280 }
1281
1282 fn full_attn_prime(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>,
1288 hx: Option<&CudaSlice<u8>>,
1289 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1290 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1291 let g3 = match hx {
1296 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
1297 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
1298 };
1299 self.full_attn_prime_core(e, fa, g3, pos_d, t, cache, il)
1300 }
1301
1302 fn full_attn_prime_core(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
1306 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1307 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1308 let (attn_g, ag16) = self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
1309 if let Some(xh) = &ag16 {
1310 if let Some(y) = e.try_f16_gemm_pre(&fa.wo, xh, t)? {
1311 return Ok(y);
1312 }
1313 }
1314 Ok(e.matmul(&fa.wo, &attn_g, t)?)
1315 }
1316
1317 fn full_attn_prime_core_inner(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
1318 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1319 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1320 let cfg = &self.cfg;
1321 let n_head = cfg.n_head as usize;
1322 let n_head_kv = cfg.n_head_kv as usize;
1323 let head_dim = cfg.head_dim_k as usize;
1324 let scale = 1.0 / (head_dim as f32).sqrt();
1325 let (pre, base_len) = self.full_attn_prime_pre_fa(e, fa, g3, pos_d, t, cache, il)?;
1326 let AttnPre { q, k, v, gate } = pre;
1327 let mut attn = e.uninit(t * n_head * head_dim)?;
1328 self.full_attn_prime_fa_dispatch(e, &q, &k, &v, &mut attn, base_len, t, cache, il,
1329 head_dim, n_head, n_head_kv, scale)?;
1330 self.full_attn_prime_post_fa(e, attn, &gate, t, n_head, head_dim)
1331 }
1332
1333 #[allow(clippy::type_complexity)]
1337 fn full_attn_prime_pre_fa(&self, e: &Engine, fa: &FullAttnLayer, mut g3: Vec<CudaSlice<f32>>,
1338 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1339 -> Result<(AttnPre, usize), Box<dyn std::error::Error>> {
1340 let cfg = &self.cfg;
1341 let n_head = cfg.n_head as usize;
1342 let n_head_kv = cfg.n_head_kv as usize;
1343 let head_dim = cfg.head_dim_k as usize;
1344 let eps = cfg.rms_eps;
1345
1346 let gated = cfg.attn_out_gate();
1350 let v = g3.pop().unwrap();
1351 let mut k = g3.pop().unwrap();
1352 let qf = g3.pop().unwrap();
1353 let (mut q, gate) = if gated {
1354 let mut q = e.uninit(t * n_head * head_dim)?;
1355 let mut gate = e.uninit(t * n_head * head_dim)?;
1356 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
1357 (q, Some(gate))
1358 } else {
1359 (qf, None)
1360 };
1361
1362 let mut qn = e.uninit(t * n_head * head_dim)?;
1363 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
1364 q = qn;
1365 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
1366 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
1367 k = kn;
1368 let rope_dims = cfg.rope_dim_count as usize;
1369 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, cfg.rope_freq_base, 1.0)?;
1370 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, cfg.rope_freq_base, 1.0)?;
1371
1372 {
1375 let kvl = cache.kv[il].as_mut().unwrap();
1376 assert!(kvl.len + t <= cache.max_ctx, "prime_cache: KV overflow");
1377 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
1378 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
1379 crate::Engine::kv_fp8_on())?;
1380 kvl.len += t;
1381 let new_len = kvl.len as i32;
1382 e.set_i32_one(&mut kvl.len_d, new_len)?;
1383 }
1384
1385 let base_len = {
1386 let kvl = cache.kv[il].as_ref().unwrap();
1387 kvl.len - t };
1389 Ok((AttnPre { q, k, v, gate }, base_len))
1390 }
1391
1392 #[allow(clippy::too_many_arguments)]
1399 fn full_attn_prime_fa_dispatch(&self, e: &Engine, q: &CudaSlice<f32>, k: &CudaSlice<f32>,
1400 v: &CudaSlice<f32>, attn: &mut CudaSlice<f32>, base_len: usize,
1401 t: usize, cache: &mut Cache, il: usize,
1402 head_dim: usize, n_head: usize, n_head_kv: usize, scale: f32)
1403 -> Result<(), Box<dyn std::error::Error>> {
1404 if base_len == 0 {
1405 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
1409 e.sdpa_naive(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1410 } else {
1411 e.fa_prefill(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1412 }
1413 } else {
1414 let kvl = cache.kv[il].as_ref().unwrap();
1415 let t_kv = base_len + t;
1416 let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
1417 let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
1418 let deqw = std::env::var("MEMRA_PRIME_DEQW").map(|v| v != "0").unwrap_or(true);
1426 if deqw {
1427 e.fa_prefill_view_ws(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
1428 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
1429 crate::Engine::kv_fp8_on())?;
1430 } else {
1431 e.fa_prefill_view(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
1432 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
1433 crate::Engine::kv_fp8_on())?;
1434 }
1435 }
1436 Ok(())
1437 }
1438
1439 fn full_attn_prime_post_fa(&self, e: &Engine, attn: CudaSlice<f32>,
1442 gate: &Option<CudaSlice<f32>>, t: usize,
1443 n_head: usize, head_dim: usize)
1444 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1445 let (attn_g, ag16) = match gate {
1446 Some(gate) => {
1447 let n = t * n_head * head_dim;
1448 let mut ag = e.uninit(n)?;
1449 if Self::f16out_on(e, t) {
1450 let mut a16 = e.alloc_u8_uninit(n * 2)?;
1451 e.sig_mul_f16out(&attn, gate, &mut ag, &mut a16, n)?;
1452 (ag, Some(a16))
1453 } else {
1454 let mut gsig = e.uninit(n)?;
1455 e.sigmoid(gate, &mut gsig, n)?;
1456 e.mul(&attn, &gsig, &mut ag, n)?;
1457 (ag, None)
1458 }
1459 }
1460 None => (attn, None),
1461 };
1462 Ok((attn_g, ag16))
1463 }
1464
1465 fn linear_attn_prime(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>,
1472 hx: Option<&CudaSlice<u8>>, t: usize,
1473 cache: &mut Cache, il: usize)
1474 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1475 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1477 let g4 = match hx {
1478 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
1479 None => e.matmul_group(&ws, h, t)?,
1480 };
1481 self.linear_attn_prime_core(e, la, g4, t, cache, il)
1482 }
1483
1484 fn linear_attn_prime_core(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
1486 t: usize, cache: &mut Cache, il: usize)
1487 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1488 self.linear_attn_prime_core_pad(e, la, g4.drain(..).collect(), t, cache, il, None)
1489 }
1490
1491 #[allow(clippy::too_many_arguments)]
1495 fn linear_attn_prime_core_pad_inner(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
1496 t: usize, cache: &mut Cache, il: usize,
1497 pad_len: Option<&CudaSlice<i32>>)
1498 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1499 let ssm = self.cfg.ssm.as_ref().unwrap();
1501 let d_state = ssm.state_size as usize;
1502 let num_k = ssm.group_count as usize;
1503 let num_v = ssm.time_step_rank as usize;
1504 let key_dim = d_state * num_k;
1505 let value_dim = d_state * num_v;
1506 let conv_dim = key_dim * 2 + value_dim;
1507 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(
1512 e, la,
1513 &qkv_mixed.slice(0..t * conv_dim), &z.slice(0..t * value_dim),
1514 &beta_raw.slice(0..t * num_v), &alpha.slice(0..t * num_v),
1515 t, cache, il, pad_len)
1516 }
1517
1518 #[allow(clippy::too_many_arguments)]
1521 fn linear_attn_gdn_prep(&self, e: &Engine, la: &LinearAttnLayer,
1522 qkv_mixed: &cudarc::driver::CudaView<f32>,
1523 beta_raw: &cudarc::driver::CudaView<f32>,
1524 alpha: &cudarc::driver::CudaView<f32>,
1525 t: usize, cache: &mut Cache, il: usize,
1526 pad_len: Option<&CudaSlice<i32>>)
1527 -> Result<GdnPrep, Box<dyn std::error::Error>> {
1528 let cfg = &self.cfg;
1529 let ssm = cfg.ssm.as_ref().unwrap();
1530 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;
1538 debug_assert!(t >= d_conv - 1, "stateful conv needs T >= pad (PRIME_MIN_T gates)");
1539
1540 let rl = cache.recur[il].as_mut().unwrap();
1545 let hk = Self::gdn_hk(e, t, num_v, num_k);
1546 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
1547 let hk = if conv_fuse { hk } else { num_v }; let mut q_g = e.uninit(d_state * hk * t)?;
1549 let mut k_g = e.uninit(d_state * hk * t)?;
1550 let mut v_g = e.uninit(d_state * num_v * t)?;
1551 if conv_fuse {
1552 e.ssm_conv1d_gdn_state_pad(qkv_mixed, &mut rl.conv_state, la.ssm_conv1d.float_data(),
1553 &mut q_g, &mut k_g, &mut v_g,
1554 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim, hk, pad_len)?;
1555 } else {
1556 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(),
1558 &mut conv_out, conv_dim, t, d_conv, pad_len)?;
1559 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)?;
1560 }
1561 let mut q_l2 = e.uninit(d_state * hk * t)?;
1562 let qb16 = if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(32) {
1566 let mut qb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
1567 e.l2_norm_pp(&q_g, &mut q_l2, Some(&mut qb), d_state, hk * t, eps)?;
1568 Some(qb)
1569 } else {
1570 e.l2_norm_pp(&q_g, &mut q_l2, None, d_state, hk * t, eps)?;
1571 None
1572 };
1573 let mut k_l2 = e.uninit(d_state * hk * t)?;
1574 let kb16 = if Engine::l2_v2_on(d_state) {
1576 let mut kb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
1577 e.l2_norm_pp(&k_g, &mut k_l2, Some(&mut kb), d_state, hk * t, eps)?;
1578 Some(kb)
1579 } else {
1580 e.l2_norm_pp(&k_g, &mut k_l2, None, d_state, hk * t, eps)?;
1581 None
1582 };
1583 let mut beta = e.uninit(t * num_v)?;
1584 e.sigmoid_v(beta_raw, &mut beta, t * num_v)?;
1585 let mut g_log = e.uninit(t * num_v)?;
1586 e.gdn_glog_v(alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
1587 if let Some(len_d) = pad_len {
1588 e.gdn_pad_mask(&mut beta, &mut g_log, len_d, num_v, t)?;
1589 }
1590 Ok(GdnPrep { hk, q_l2, k_l2, v_g, beta, g_log, kb16, qb16 })
1591 }
1592
1593 #[allow(clippy::too_many_arguments)]
1598 fn linear_attn_prime_core_batch(&self, e: &Engine, la: &LinearAttnLayer,
1599 g4: &[CudaSlice<f32>], offs: &[usize], ts: &[usize],
1600 caches: &mut [&mut Cache], il: usize)
1601 -> Result<Vec<(CudaSlice<f32>, Option<CudaSlice<u8>>)>, Box<dyn std::error::Error>> {
1602 let ssm = self.cfg.ssm.as_ref().unwrap();
1603 let d_state = ssm.state_size as usize;
1604 let num_k = ssm.group_count as usize;
1605 let num_v = ssm.time_step_rank as usize;
1606 let key_dim = d_state * num_k;
1607 let value_dim = d_state * num_v;
1608 let conv_dim = key_dim * 2 + value_dim;
1609 let eps = self.cfg.rms_eps;
1610 let scale = 1.0 / (d_state as f32).sqrt();
1611 let b = ts.len();
1612 let c = Engine::gdn_chunk_size();
1613 let carried = caches.iter().any(|c| c.pos > 0);
1616 let use_vl = !carried
1617 && (2..=8).contains(&b)
1618 && Engine::gdn_chunked_enabled() && ts.iter().all(|&t| t >= 16)
1619 && e.gdn_mma_enabled(c)
1620 && std::env::var("MEMRA_GDN_VL").as_deref() != Ok("0");
1621 if !use_vl {
1622 return (0..b).map(|s| {
1623 let (o, t) = (offs[s], ts[s]);
1624 self.linear_attn_prime_core_pad_view(
1625 e, la,
1626 &g4[0].slice(o * conv_dim..(o + t) * conv_dim),
1627 &g4[1].slice(o * value_dim..(o + t) * value_dim),
1628 &g4[2].slice(o * num_v..(o + t) * num_v),
1629 &g4[3].slice(o * num_v..(o + t) * num_v),
1630 t, caches[s], il, None)
1631 }).collect();
1632 }
1633 struct SeqBufs {
1637 conv_out: CudaSlice<f32>, q_g: CudaSlice<f32>, k_g: CudaSlice<f32>, v_g: CudaSlice<f32>,
1638 q_l2: CudaSlice<f32>, k_l2: CudaSlice<f32>, beta: CudaSlice<f32>, g_log: CudaSlice<f32>,
1639 gn: CudaSlice<f32>, gn16: CudaSlice<u8>,
1640 }
1641 let d_conv = ssm.conv_kernel as usize;
1642 let f16o = Self::f16out_on(e, 16);
1643 let hk = Self::gdn_hk(e, 16, num_v, num_k); let mut sb = Vec::with_capacity(b);
1645 let mut pres = Vec::with_capacity(b);
1646 for &t in ts.iter().take(b) {
1647 sb.push(SeqBufs {
1648 conv_out: e.uninit(conv_dim * t)?,
1649 q_g: e.uninit(d_state * hk * t)?,
1650 k_g: e.uninit(d_state * hk * t)?,
1651 v_g: e.uninit(d_state * num_v * t)?,
1652 q_l2: e.uninit(d_state * hk * t)?,
1653 k_l2: e.uninit(d_state * hk * t)?,
1654 beta: e.uninit(t * num_v)?,
1655 g_log: e.uninit(t * num_v)?,
1656 gn: e.uninit(d_state * num_v * t)?,
1657 gn16: e.alloc_u8_uninit(d_state * num_v * t * 2)?,
1658 });
1659 pres.push(e.gdn_chunk_alloc(num_v, t, c, hk)?);
1660 }
1661 let prep_args: Vec<crate::GdnPrepVl> = (0..b).map(|s| {
1662 let (o, t) = (offs[s], ts[s]);
1663 let rl = caches[s].recur[il].as_ref().unwrap();
1664 crate::GdnPrepVl {
1665 qkv: e.addr_f32v(&g4[0].slice(o * conv_dim..(o + t) * conv_dim)),
1666 conv_state: e.addr_f32(&rl.conv_state),
1667 conv_out: e.addr_f32(&sb[s].conv_out),
1668 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),
1669 q_l2: e.addr_f32(&sb[s].q_l2), k_l2: e.addr_f32(&sb[s].k_l2),
1670 beta_raw: e.addr_f32v(&g4[2].slice(o * num_v..(o + t) * num_v)),
1671 alpha: e.addr_f32v(&g4[3].slice(o * num_v..(o + t) * num_v)),
1672 beta: e.addr_f32(&sb[s].beta), g_log: e.addr_f32(&sb[s].g_log),
1673 o: e.addr_f32(&pres[s].o),
1674 z: e.addr_f32v(&g4[1].slice(o * value_dim..(o + t) * value_dim)),
1675 gn: e.addr_f32(&sb[s].gn), gn16: e.addr_u8(&sb[s].gn16),
1676 kb16: if Engine::l2_v2_on(d_state) { e.addr_u8(&pres[s].kb16) } else { 0 },
1677 qb16: if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(c) { e.addr_u8(&pres[s].qb16) } else { 0 },
1678 t: t as i32, pad: 0,
1679 }
1680 }).collect();
1681 let args: Vec<crate::GdnSeqVl> = (0..b).map(|s| {
1682 let rl = caches[s].recur[il].as_ref().unwrap();
1683 crate::GdnSeqVl {
1684 kb16: e.addr_u8(&pres[s].kb16), gcum: e.addr_f32(&pres[s].gcum),
1685 beta: e.addr_f32(&sb[s].beta), u: e.addr_f32(&pres[s].u),
1686 wb16: e.addr_u8(&pres[s].wb16), y: e.addr_u8(&pres[s].y16),
1687 ssnap: e.addr_u8(&pres[s].ssnap16),
1688 state_in: e.addr_f32(&rl.ssm_state), state_out: e.addr_f32(&rl.ssm_state_alt),
1689 q: e.addr_f32(&sb[s].q_l2), p: e.addr_f32(&pres[s].p),
1690 o: e.addr_f32(&pres[s].o),
1691 k: e.addr_f32(&sb[s].k_l2), v: e.addr_f32(&sb[s].v_g),
1692 g: e.addr_f32(&sb[s].g_log), a: e.addr_f32(&pres[s].a),
1693 w: e.addr_f32(&pres[s].w),
1694 t: ts[s] as i32, nc: pres[s].nc as i32,
1695 }
1696 }).collect();
1697 e.gdn_prep_vl8(&prep_args, la.ssm_conv1d.float_data(), la.ssm_dt.float_data(),
1698 la.ssm_a.float_data(), conv_dim, d_conv, d_state, num_v, num_k, key_dim, hk, eps)?;
1699 if !Engine::l2_v2_on(d_state) {
1702 e.gdn_mirror_vl8(&args, num_v, 0, hk)?;
1703 }
1704 let wq8: Option<crate::GdnWVl8> = if e.gdn_wgmma_on(c) {
1706 if !Engine::l2_v2_on(d_state) {
1708 for s in 0..b {
1709 e.f32_to_bf16_into(&sb[s].q_l2, &mut pres[s].qb16, d_state * hk * ts[s])?;
1710 }
1711 }
1712 let mut wa = [crate::GdnWVl::default(); 8];
1713 for s in 0..b {
1714 wa[s] = crate::GdnWVl { qb16: e.addr_u8(&pres[s].qb16), pb16: e.addr_u8(&pres[s].pb16) };
1715 }
1716 Some(crate::GdnWVl8(wa))
1717 } else { None };
1718 e.gdn_chunk_k123_vl8(&args, num_v, hk, wq8.as_ref())?;
1719 e.gdn_chunk_vl8(&args, num_v, scale, hk, wq8.as_ref())?;
1720 if f16o {
1721 e.gdn_tail_vl8(&prep_args, la.ssm_norm.float_data(), d_state, num_v, eps)?;
1722 }
1723 let mut out = Vec::with_capacity(b);
1725 for (s, bufs) in sb.into_iter().enumerate() {
1726 let rl = caches[s].recur[il].as_mut().unwrap();
1727 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
1728 let (o, t) = (offs[s], ts[s]);
1729 let SeqBufs { mut gn, gn16, .. } = bufs;
1730 if f16o {
1731 out.push((gn, Some(gn16)));
1732 } else {
1733 let z_v = g4[1].slice(o * value_dim..(o + t) * value_dim);
1734 e.gated_rmsnorm_zv(&pres[s].o, la.ssm_norm.float_data(), &z_v, &mut gn,
1735 d_state, num_v * t, eps)?;
1736 out.push((gn, None));
1737 }
1738 }
1739 Ok(out)
1740 }
1741
1742 #[allow(clippy::too_many_arguments)]
1746 fn linear_attn_prime_core_pad_view(&self, e: &Engine, la: &LinearAttnLayer,
1747 qkv_mixed: &cudarc::driver::CudaView<f32>,
1748 z: &cudarc::driver::CudaView<f32>,
1749 beta_raw: &cudarc::driver::CudaView<f32>,
1750 alpha: &cudarc::driver::CudaView<f32>,
1751 t: usize, cache: &mut Cache, il: usize,
1752 pad_len: Option<&CudaSlice<i32>>)
1753 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1754 let cfg = &self.cfg;
1755 let ssm = cfg.ssm.as_ref().unwrap();
1756 let d_state = ssm.state_size as usize; let num_v = ssm.time_step_rank as usize; let eps = cfg.rms_eps;
1759 let scale = 1.0 / (d_state as f32).sqrt();
1760
1761 let prep = self.linear_attn_gdn_prep(e, la, qkv_mixed, beta_raw, alpha, t, cache, il, pad_len)?;
1762
1763 let mut o = e.uninit(d_state * num_v * t)?;
1769 let rl = cache.recur[il].as_mut().unwrap();
1770 {
1771 let crate::cache::RecurLayer { ssm_state, ssm_state_alt, .. } = rl;
1772 e.gdn_scan_prefill(&prep.q_l2, &prep.k_l2, &prep.v_g, &prep.g_log, &prep.beta,
1773 prep.kb16.as_ref(), prep.qb16.as_ref(), ssm_state, ssm_state_alt, &mut o, num_v, t, scale,
1774 prep.hk)?;
1775 }
1776 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
1777
1778 let mut gn = e.uninit(d_state * num_v * t)?;
1781 let gn16 = if Self::f16out_on(e, t) {
1782 let mut g16 = e.alloc_u8_uninit(d_state * num_v * t * 2)?;
1783 e.gated_rmsnorm_f16out_zv(&o, la.ssm_norm.float_data(), z, &mut gn, &mut g16,
1784 d_state, num_v * t, eps)?;
1785 Some(g16)
1786 } else {
1787 e.gated_rmsnorm_zv(&o, la.ssm_norm.float_data(), z, &mut gn, d_state, num_v * t, eps)?;
1788 None
1789 };
1790 Ok((gn, gn16))
1791 }
1792
1793 #[allow(clippy::too_many_arguments)]
1795 fn linear_attn_prime_core_pad(&self, e: &Engine, la: &LinearAttnLayer, g4: Vec<CudaSlice<f32>>,
1796 t: usize, cache: &mut Cache, il: usize,
1797 pad_len: Option<&CudaSlice<i32>>)
1798 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1799 let (gn, gn16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, pad_len)?;
1800 if let Some(xh) = &gn16 {
1801 if let Some(y) = e.try_f16_gemm_pre(&la.ssm_out, xh, t)? {
1802 return Ok(y);
1803 }
1804 }
1805 Ok(e.matmul(&la.ssm_out, &gn, t)?)
1806 }
1807
1808 pub fn full_attn(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
1810 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1811 let cfg = &self.cfg;
1812 let _n_embd = cfg.n_embd as usize;
1813 let n_head = cfg.n_head as usize;
1814 let n_head_kv = cfg.n_head_kv as usize;
1815 let head_dim = cfg.head_dim_k as usize;
1816 let eps = cfg.rms_eps;
1817 let scale = 1.0 / (head_dim as f32).sqrt();
1818
1819 let gated = cfg.attn_out_gate();
1822 let mut g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
1824 let v = g3.pop().unwrap();
1825 let mut k = g3.pop().unwrap();
1826 let qf = g3.pop().unwrap();
1827 let (mut q, gate) = if gated {
1828 let mut q = e.uninit(t * n_head * head_dim)?;
1829 let mut gate = e.uninit(t * n_head * head_dim)?;
1830 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
1831 (q, Some(gate))
1832 } else {
1833 (qf, None)
1834 };
1835
1836 let mut qn = e.uninit(t * n_head * head_dim)?;
1838 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
1839 q = qn;
1840 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
1841 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
1842 k = kn;
1843 let rope_dims = cfg.rope_dim_count as usize;
1844 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, cfg.rope_freq_base, 1.0)?;
1845 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, cfg.rope_freq_base, 1.0)?;
1846
1847 let mut attn = e.uninit(t * n_head * head_dim)?;
1849 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
1852 e.sdpa_naive(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1854 } else {
1855 e.fa_prefill(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1856 }
1857
1858 let attn_g = match &gate {
1860 Some(gate) => {
1861 let mut gsig = e.uninit(t * n_head * head_dim)?;
1862 e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
1863 let mut ag = e.uninit(t * n_head * head_dim)?;
1864 e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
1865 ag
1866 }
1867 None => attn,
1868 };
1869
1870 let o = e.matmul(&fa.wo, &attn_g, t)?;
1872 Ok(o)
1873 }
1874
1875 pub fn linear_attn(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>, t: usize)
1877 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1878 let cfg = &self.cfg;
1879 let _n_embd = cfg.n_embd as usize;
1880 let ssm = cfg.ssm.as_ref().unwrap();
1881 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;
1886 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;
1890 let scale = 1.0 / (d_state as f32).sqrt();
1891
1892 let mut g4 = e.matmul_group(&[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha], h, t)?;
1895 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);
1907 let mut q_g = e.uninit(d_state * num_v * t)?;
1908 let mut k_g = e.uninit(d_state * num_v * t)?;
1909 let mut v_g = e.uninit(d_state * num_v * t)?;
1910 e.ssm_conv1d_gdn(&qkv_mixed, la.ssm_conv1d.float_data(), &mut q_g, &mut k_g, &mut v_g,
1911 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim)?;
1912 let mut q_l2 = e.uninit(d_state * num_v * t)?;
1914 e.l2_norm(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
1915 let mut k_l2 = e.uninit(d_state * num_v * t)?;
1916 e.l2_norm(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
1917 let v_gd = v_g;
1918
1919 let mut beta = e.uninit(t * num_v)?;
1922 e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
1923 let mut g_log = e.uninit(t * num_v)?;
1925 e.gdn_glog(&alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
1926
1927 let state_in = e.zeros(d_state * d_state * num_v)?; let mut state_out = e.zeros(d_state * d_state * num_v)?;
1930 let mut o = e.uninit(d_state * num_v * t)?;
1931 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)?;
1932
1933 let mut gn = e.uninit(d_state * num_v * t)?;
1938 e.gated_rmsnorm(&o, la.ssm_norm.float_data(), &z, &mut gn, d_state, num_v * t, eps)?;
1939
1940 let out = e.matmul(&la.ssm_out, &gn, t)?;
1944 Ok(out)
1945 }
1946}
1947
1948impl HybridModel {
1949 pub fn moe_ffn_il(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize, il: u16)
1960 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1961 Self::moe_ffn(e, m, z, t, &self.cfg, il, self.max_moe_block())
1962 }
1963
1964 pub fn moe_ffn_il_zq8(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
1968 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, t: usize, il: u16)
1969 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1970 Self::moe_ffn_inner(e, m, z, zq8, t, &self.cfg, il, self.max_moe_block())
1971 }
1972
1973 pub(crate) fn moe_ffn(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
1981 cfg: &ModelConfig, il: u16, max_block: usize)
1982 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1983 Self::moe_ffn_inner(e, m, z, None, t, cfg, il, max_block)
1984 }
1985
1986 #[allow(clippy::too_many_arguments)]
1987 pub(crate) fn moe_ffn_inner(
1988 e: &Engine,
1989 m: &MoeWeights,
1990 z: &CudaSlice<f32>,
1991 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
1992 t: usize,
1993 cfg: &ModelConfig,
1994 il: u16,
1995 max_block: usize,
1996 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1997 let worker_io = crate::spill_pread::worker_enabled();
1998 let epoch_lfu = std::env::var_os("MEMRA_MOE_LFU_DECAY").is_some();
1999 if Engine::moe_cache_enabled() && (worker_io || epoch_lfu) {
2000 e.with_moe_cache(max_block, |cache, _| {
2001 cache.begin_forward_epoch(il, t);
2002 if worker_io {
2003 cache.begin_worker_scope();
2004 }
2005 Ok(())
2006 })?;
2007 }
2008 if t > 1 && std::env::var("MEMRA_MOE_GROUPED").is_ok() {
2010 let grouped_out = Self::moe_ffn_grouped(e, m, z, t, cfg, il, max_block)?;
2011 if std::env::var("MEMRA_MOE_GATE").is_ok() {
2018 let seq_out = Self::moe_ffn_sequential(e, m, z, t, cfg, il, max_block)?;
2019 let g_host = e.dtoh(&grouped_out)?;
2020 let s_host = e.dtoh(&seq_out)?;
2021 let g_bytes: &[u8] = unsafe { std::slice::from_raw_parts(g_host.as_ptr() as *const u8, g_host.len() * 4) };
2022 let s_bytes: &[u8] = unsafe { std::slice::from_raw_parts(s_host.as_ptr() as *const u8, s_host.len() * 4) };
2023 if g_bytes == s_bytes {
2024 if il == 0 { println!("moe-gate il={il} t={t} BYTE-IDENTICAL (first layer only printed)"); }
2025 } else {
2026 let diffs = g_host.iter().zip(s_host.iter()).enumerate()
2027 .filter(|(_, (a, b))| a != b).count();
2028 let maxdiff = g_host.iter().zip(s_host.iter())
2029 .map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
2030 panic!("moe-gate il={il} t={t} MISMATCH: {diffs}/{} elems differ, maxdiff={maxdiff:.6e}", g_host.len());
2031 }
2032 }
2033 return Ok(grouped_out);
2034 }
2035 Self::moe_ffn_sequential_zq8(e, m, z, zq8, t, cfg, il, max_block)
2036 }
2037
2038 pub(crate) fn moe_ffn_sequential(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
2040 cfg: &ModelConfig, il: u16, max_block: usize)
2041 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2042 Self::moe_ffn_sequential_zq8(e, m, z, None, t, cfg, il, max_block)
2043 }
2044
2045 fn trace_moe_routes(il: u16, t: usize, sel_all: &[u32], weights: &[f32])
2049 -> Result<(), Box<dyn std::error::Error>> {
2050 use std::io::Write as _;
2051 if let Ok(path) = std::env::var("MEMRA_MOE_TRACE") {
2052 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
2053 let ids: Vec<String> = sel_all.iter().map(|s| s.to_string()).collect();
2054 writeln!(f, "{} {} {}", il, t, ids.join(","))?;
2055 }
2056 if let Ok(path) = std::env::var("MEMRA_MOE_WEIGHT_TRACE") {
2057 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
2058 let pairs: Vec<String> = sel_all.iter().zip(weights)
2059 .map(|(expert, weight)| format!("{expert}:{weight:.9}"))
2060 .collect();
2061 writeln!(f, "{} {} {}", il, t, pairs.join(","))?;
2062 }
2063 Ok(())
2064 }
2065
2066 fn trace_moe_input(e: &Engine, il: u16, t: usize, n_embd: usize, z: &CudaSlice<f32>)
2071 -> Result<(), Box<dyn std::error::Error>> {
2072 use std::io::Write as _;
2073 let Ok(dir) = std::env::var("MEMRA_MOE_INPUT_TRACE_DIR") else { return Ok(()) };
2074 let host = e.dtoh(z)?;
2075 if host.len() != t * n_embd {
2076 return Err(format!(
2077 "MoE input trace shape mismatch at layer {il}: got {} values, expected {}x{}",
2078 host.len(), t, n_embd
2079 ).into());
2080 }
2081 let bytes = unsafe {
2082 std::slice::from_raw_parts(
2083 host.as_ptr().cast::<u8>(), host.len() * std::mem::size_of::<f32>()
2084 )
2085 };
2086 let state = MOE_INPUT_TRACE_WRITER.get_or_init(|| std::sync::Mutex::new(None));
2087 let mut state = state.lock().map_err(|_| "MoE input trace writer lock is poisoned")?;
2088 if state.is_none() {
2089 let dir = std::path::PathBuf::from(&dir);
2090 std::fs::create_dir_all(&dir)?;
2091 let index = std::fs::OpenOptions::new().create(true).append(true)
2092 .open(dir.join("index.jsonl"))?;
2093 *state = Some(MoeInputTraceWriter {
2094 dir,
2095 index,
2096 payloads: std::collections::HashMap::new(),
2097 });
2098 }
2099 let writer = state.as_mut().unwrap();
2100 if writer.dir != std::path::Path::new(&dir) {
2101 return Err("MEMRA_MOE_INPUT_TRACE_DIR changed after capture started".into());
2102 }
2103 let file_name = format!("layer-{il:03}.f32");
2104 if !writer.payloads.contains_key(&il) {
2105 let payload = std::fs::OpenOptions::new().create(true).append(true)
2106 .open(writer.dir.join(&file_name))?;
2107 let offset = payload.metadata()?.len();
2108 writer.payloads.insert(il, (payload, offset));
2109 }
2110 let (payload, offset) = writer.payloads.get_mut(&il).unwrap();
2111 let row_offset = *offset;
2112 payload.write_all(bytes)?;
2113 *offset += bytes.len() as u64;
2114 writeln!(
2115 writer.index,
2116 "{{\"format\":\"memra-moe-input-trace-v1\",\"layer\":{il},\"tokens\":{t},\
2117 \"hidden_size\":{n_embd},\"file\":\"{file_name}\",\"offset\":{row_offset},\
2118 \"payload_bytes\":{}}}",
2119 bytes.len()
2120 )?;
2121 Ok(())
2122 }
2123
2124 #[allow(clippy::too_many_arguments)]
2125 pub(crate) fn moe_ffn_sequential_zq8(
2126 e: &Engine,
2127 m: &MoeWeights,
2128 z: &CudaSlice<f32>,
2129 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
2130 t: usize,
2131 cfg: &ModelConfig,
2132 il: u16,
2133 max_block: usize,
2134 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2135 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
2136 let moe = cfg.moe.as_ref().unwrap();
2137 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);
2144 debug_assert_eq!(m.gate_exps.out_f, n_ff_exp);
2145 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);
2148
2149 let use_cache = Engine::moe_cache_enabled();
2150 let uniform_experts = m.has_uniform_expert_layout();
2151 let moe_q8 = uniform_experts && moe_q8_enabled()
2152 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
2153 && q8_expert_supported(m.down_exps.qtype);
2154 let cpu_expert_requested = crate::cpu_experts::configured();
2161 if cpu_expert_requested && (cfg.hy3.is_none() || cfg.m3.is_some()) {
2162 return Err(std::io::Error::other(
2163 "MEMRA_CPU_EXPERT_LIB is experimental and currently gated to Hy3",
2164 )
2165 .into());
2166 }
2167 let cpu_hybrid = cpu_expert_requested && t < PRIME_MIN_T && m.dev_exps.is_none();
2168 let freeze_cpu_residency = cpu_expert_requested
2174 && std::env::var("MEMRA_CPU_EXPERT_FREEZE_CACHE").as_deref() == Ok("1");
2175 let caller_warms_before_freeze = std::env::var("MEMRA_CPU_EXPERT_FREEZE_WARMUP_TOKENS")
2176 .ok()
2177 .and_then(|value| value.parse::<usize>().ok())
2178 .is_some_and(|tokens| tokens > 0);
2179 if cpu_hybrid && freeze_cpu_residency && !caller_warms_before_freeze {
2180 e.freeze_moe_cache();
2181 }
2182 let cache_frozen = use_cache && e.moe_cache_frozen();
2183 let cache_dispatch = use_cache && (!cache_frozen || cpu_hybrid);
2184
2185 let logits = if t < PRIME_MIN_T {
2192 if crate::router_kernel_on() {
2196 e.router_gemv(m.gate_inp.float_data(), z, cfg.n_embd as usize,
2199 m.gate_exps.n_expert, t)?
2200 } else {
2201 e.matmul_decode_exact(&m.gate_inp, z, t)?
2202 }
2203 } else if crate::router_prefill_exact_on() && crate::router_kernel_on() {
2204 e.router_gemv(m.gate_inp.float_data(), z, cfg.n_embd as usize,
2219 m.gate_exps.n_expert, t)?
2220 } else {
2221 e.matmul(&m.gate_inp, z, t)?
2222 };
2223
2224 let no_exp_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
2262 && m.down_exps.macros.is_none();
2263 if cfg.sigmoid_router().is_none() && cfg.m3.is_none() && cfg.hy3.is_none()
2264 && no_exp_macros
2265 && t >= PRIME_MIN_T && m.dev_exps.is_some() && moe_q8_enabled()
2266 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
2267 && q8_expert_supported(m.down_exps.qtype)
2268 && std::env::var("MEMRA_MOE_PAIRS").map(|v| v != "0").unwrap_or(true)
2269 && std::env::var("MEMRA_MOE_STATS").is_err() {
2270 return Self::moe_ffn_pairs(e, m, z, &logits, t, cfg);
2271 }
2272
2273 let dev_ok = uniform_experts && cfg.m3.is_none() && cfg.hy3.is_none();
2283 let observe_routes = std::env::var("MEMRA_MOE_STATS").is_ok()
2287 || std::env::var("MEMRA_MOE_TRACE").is_ok()
2288 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
2289 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok();
2290 if dev_ok && t < PRIME_MIN_T && m.dev_exps.is_some() && n_used <= 8 && moe_dev_enabled()
2291 && !observe_routes {
2292 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
2293 }
2294 if dev_ok && use_cache && n_used <= 8 && moe_dev_enabled()
2295 && !observe_routes {
2296 let row_ok = e.with_moe_cache(max_block, |c, eng| {
2297 if moe_prewarm_enabled() { c.prewarm_layer(il, m, eng)?; }
2298 Ok(c.layer_dev_row(il, n_expert, eng)?.is_some())
2299 })?;
2300 if row_ok {
2301 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
2302 }
2303 }
2304
2305 let (sel_all, w_all, routed_cpu_input) = if let Some(sig) = cfg.sigmoid_router() {
2307 if cpu_hybrid {
2308 let (sel, w, input) = Self::moe_route_sigmoid_with_input(
2309 e,
2310 &logits,
2311 z,
2312 t,
2313 n_expert,
2314 n_used,
2315 m.exp_probs_b.as_deref(),
2316 sig,
2317 m.active_experts.as_deref(),
2318 )?;
2319 (sel, w, Some(input))
2320 } else {
2321 let (sel, w) = Self::moe_route_cfg(
2322 e,
2323 &logits,
2324 t,
2325 n_expert,
2326 n_used,
2327 m.exp_probs_b.as_deref(),
2328 Some(sig),
2329 m.active_experts.as_deref(),
2330 )?;
2331 (sel, w, None)
2332 }
2333 } else {
2334 let (sel, w) = Self::moe_route_cfg(
2335 e,
2336 &logits,
2337 t,
2338 n_expert,
2339 n_used,
2340 None,
2341 None,
2342 m.active_experts.as_deref(),
2343 )?;
2344 (sel, w, None)
2345 };
2346
2347 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
2351 Self::trace_moe_input(e, il, t, n_embd, z)?;
2352
2353 let worker_disk_prefetch =
2365 cache_dispatch && crate::spill_pread::worker_enabled() && !cpu_hybrid;
2366 let promote_worker_h2d =
2367 t == 1 && worker_disk_prefetch && crate::spill_pread::copy_h2d_enabled();
2368 if promote_worker_h2d {
2369 let mut selected_blocks = Vec::with_capacity(n_used * 3);
2370 for &ex in sel_all.iter().take(n_used) {
2371 let ex = ex as u16;
2372 selected_blocks.extend([
2373 BlockId::new(il, PROJ_GATE, ex),
2374 BlockId::new(il, PROJ_UP, ex),
2375 BlockId::new(il, PROJ_DOWN, ex),
2376 ]);
2377 }
2378 for &ex in sel_all.iter().take(n_used) {
2379 Self::moe_prefetch_disk_expert(e, il, ex as usize, m, max_block, &selected_blocks)?;
2380 }
2381 e.with_moe_cache(max_block, |cache, eng| {
2382 cache.promote_worker_reads_at_safe_boundary(
2383 &selected_blocks,
2384 &selected_blocks,
2385 eng,
2386 )?;
2387 Ok(())
2388 })?;
2389 }
2390
2391 if t > 1 && std::env::var("MEMRA_MOE_STATS").is_ok() {
2394 let mut cnt = vec![0u32; n_expert];
2395 for &s in sel_all.iter() { cnt[s as usize] += 1; }
2396 let total = sel_all.len() as f64;
2397 let mut h = 0.0f64;
2398 let mut active = 0usize;
2399 for &c in &cnt { if c > 0 { active += 1; let p = c as f64 / total; h -= p * p.log2(); } }
2400 let maxc = cnt.iter().copied().max().unwrap_or(0);
2401 println!("moe-stats il={} t={} assignments={} active={}/{} entropy={:.3}b (max {:.3}b) mean_tok_per_active={:.2} max_tok_per_expert={}",
2402 il, t, sel_all.len(), active, n_expert, h, (n_expert as f64).log2(), total / active.max(1) as f64, maxc);
2403 }
2404
2405 let gdec_may_fire = uniform_experts && use_cache && n_used <= 8 && gdec_enabled();
2414 let mut moe_out = if gdec_may_fire {
2415 e.uninit(t * n_embd)?
2416 } else {
2417 e.zeros(t * n_embd)?
2418 };
2419 let cpu_input = if cpu_hybrid {
2422 Some(routed_cpu_input.ok_or("CPU expert routing did not return the MoE input")?)
2423 } else {
2424 None
2425 };
2426
2427 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;
2435 let mut scratch_u: Option<CudaSlice<u8>> = None;
2436 let mut scratch_d: Option<CudaSlice<u8>> = None;
2437 let page_window = moe_page_prefetch_window();
2445
2446 for tok in 0..t {
2449 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
2450 let w = &w_all[tok * n_used..(tok + 1) * n_used];
2451 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd); let mut tok_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
2453
2454 let no_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
2468 && m.down_exps.macros.is_none();
2469 if gdec_may_fire && moe_q8 && cfg.m3.is_none() && no_macros {
2470 if tok_q8.is_none() {
2471 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
2472 }
2473 let (zq, zd) = tok_q8.as_ref().unwrap();
2474 if Self::moe_gdec_token_q8(e, m, il, max_block, zq, zd, sel, w,
2475 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
2476 continue;
2477 }
2478 } else if gdec_may_fire && cfg.m3.is_none() && no_macros
2479 && Self::moe_gdec_token(e, m, il, max_block, &zt, sel, w,
2480 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
2481 continue;
2482 }
2483
2484 if gdec_may_fire {
2488 let mut row = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2489 e.memset_zeros_view(&mut row)?;
2490 }
2491
2492 let mut cpu_mask = vec![false; sel.len()];
2498 let cpu_worker = if let Some(host_input) = cpu_input.as_ref() {
2499 let gpu_resident = if use_cache {
2500 e.with_moe_cache(max_block, |cache, _| {
2501 Ok(sel
2502 .iter()
2503 .map(|&expert| {
2504 let expert = expert as u16;
2505 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
2506 .into_iter()
2507 .filter(|&projection| {
2508 cache
2509 .resident(BlockId::new(il, projection, expert))
2510 .is_some()
2511 })
2512 .count()
2513 })
2514 .collect::<Vec<_>>())
2515 })?
2516 } else {
2517 vec![0; sel.len()]
2518 };
2519 let mut cpu_selected = Vec::new();
2520 for (index, (&expert, &route_weight)) in sel.iter().zip(w).enumerate() {
2521 if gpu_resident[index] != 3 {
2522 cpu_mask[index] = true;
2523 crate::cpu_experts::record_incomplete_gpu_residency(gpu_resident[index]);
2524 let expert = expert as usize;
2525 cpu_selected.push((expert, route_weight));
2526 }
2527 }
2528 if crate::cpu_experts::predictor_enabled() {
2529 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
2533 crate::cpu_experts::predictor_submit(il, row);
2534 }
2535 if cpu_selected.is_empty() {
2536 None
2537 } else {
2538 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
2539 let job = crate::cpu_experts::prepare_job(m, il, &cpu_selected, row)
2540 .map_err(std::io::Error::other)?;
2541 Some(crate::cpu_experts::submit(job).map_err(std::io::Error::other)?)
2542 }
2543 } else {
2544 None
2545 };
2546
2547 let worker_window = worker_disk_prefetch
2548 .then(worker_prefetch_window)
2549 .unwrap_or(0);
2550 for (j, &ex) in sel.iter().enumerate() {
2551 if cpu_mask[j] {
2552 continue;
2553 }
2554 let ex = ex as usize;
2555 for next in page_prefetch_positions(j, sel.len(), page_window) {
2556 Self::moe_prefetch_host_expert(sel[next] as usize, m);
2557 }
2558 let keep = [
2559 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_GATE, ex as u16),
2560 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_UP, ex as u16),
2561 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_DOWN, ex as u16),
2562 ];
2563 if worker_disk_prefetch && worker_window > 0 {
2564 for next in worker_prefetch_positions(j, sel.len(), worker_window) {
2565 Self::moe_prefetch_disk_expert(
2566 e,
2567 il,
2568 sel[next] as usize,
2569 m,
2570 max_block,
2571 &keep,
2572 )?;
2573 }
2574 } else if cache_dispatch
2575 && !cpu_hybrid
2576 && moe_prefetch_enabled()
2577 && j + 1 < sel.len()
2578 {
2579 let next = sel[j + 1] as usize;
2580 Self::moe_prefetch_expert(e, il, next, m, max_block, &keep)?;
2581 }
2582 let [gate_q8, up_q8, down_q8] = [moe_q8; 3];
2583 if cache_dispatch && (gate_q8 || up_q8 || down_q8) {
2584 if (gate_q8 || up_q8) && tok_q8.is_none() {
2587 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
2588 }
2589 let gate = if gate_q8 {
2590 let (zq, zd) = tok_q8.as_ref().unwrap();
2591 Self::moe_cached_gemm_q8(e, il, PROJ_GATE, ex, m, max_block, zq, zd)?
2592 } else {
2593 Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?
2594 };
2595 let up = if up_q8 {
2596 let (zq, zd) = tok_q8.as_ref().unwrap();
2597 Self::moe_cached_gemm_q8(e, il, PROJ_UP, ex, m, max_block, zq, zd)?
2598 } else {
2599 Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?
2600 };
2601 let mut act = e.uninit(n_ff_exp)?;
2602 Self::ffn_act_scaled(
2603 e,
2604 cfg,
2605 &gate,
2606 &up,
2607 m.gate_exps.macro_scale(ex),
2608 m.up_exps.macro_scale(ex),
2609 &mut act,
2610 n_ff_exp,
2611 )?;
2612 let y = if down_q8 {
2613 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
2614 Self::moe_cached_gemm_q8(e, il, PROJ_DOWN, ex, m, max_block, &aq2, &ad2)?
2615 } else {
2616 let actv = act.slice(0..n_ff_exp);
2617 Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?
2618 };
2619 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2620 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2622 } else if cache_dispatch {
2623 let gate = Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?;
2628 let up = Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?;
2629 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_scaled(e, cfg, &gate, &up,
2631 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, n_ff_exp)?;
2632 let actv = act.slice(0..n_ff_exp);
2633 let y = Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?;
2634 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2635 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2637 } else if cache_frozen {
2638 let gate = Self::moe_frozen_gemm(
2643 e,
2644 il,
2645 PROJ_GATE,
2646 ex,
2647 m,
2648 max_block,
2649 &zt,
2650 &mut scratch_g,
2651 g_len,
2652 )?;
2653 let up = Self::moe_frozen_gemm(
2654 e,
2655 il,
2656 PROJ_UP,
2657 ex,
2658 m,
2659 max_block,
2660 &zt,
2661 &mut scratch_u,
2662 u_len,
2663 )?;
2664 let mut act = e.uninit(n_ff_exp)?;
2665 Self::ffn_act_scaled(
2666 e,
2667 cfg,
2668 &gate,
2669 &up,
2670 m.gate_exps.macro_scale(ex),
2671 m.up_exps.macro_scale(ex),
2672 &mut act,
2673 n_ff_exp,
2674 )?;
2675 let actv = act.slice(0..n_ff_exp);
2676 let y = Self::moe_frozen_gemm(
2677 e,
2678 il,
2679 PROJ_DOWN,
2680 ex,
2681 m,
2682 max_block,
2683 &actv,
2684 &mut scratch_d,
2685 d_len,
2686 )?;
2687 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2688 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2689 } else {
2690 if scratch_g.is_none() {
2694 scratch_g = Some(e.alloc_u8_uninit(g_len)?);
2695 scratch_u = Some(e.alloc_u8_uninit(u_len)?);
2696 scratch_d = Some(e.alloc_u8_uninit(d_len)?);
2697 }
2698 let (sg, su, sd) = (scratch_g.as_mut().unwrap(), scratch_u.as_mut().unwrap(),
2699 scratch_d.as_mut().unwrap());
2700 let gl = m.gate_exps.expert_layout(ex);
2701 let ul = m.up_exps.expert_layout(ex);
2702 let dl = m.down_exps.expert_layout(ex);
2703 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
2704 let gate = e.qmatvec_view(sg, 0..gl.len, &zt, 1,
2705 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)?;
2706
2707 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
2708 let up = e.qmatvec_view(su, 0..ul.len, &zt, 1,
2709 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)?;
2710
2711 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_scaled(e, cfg, &gate, &up,
2713 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, n_ff_exp)?;
2714
2715 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
2716 let actv = act.slice(0..n_ff_exp);
2717 let y = e.qmatvec_view(sd, 0..dl.len, &actv, 1,
2718 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)?;
2719
2720 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2721 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2722 }
2723 }
2724 if let Some(worker) = cpu_worker {
2725 let cpu_output = worker.wait().map_err(std::io::Error::other)?;
2726 let cpu_output = e.htod(&cpu_output)?;
2727 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2728 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
2729 }
2730 if cpu_hybrid && !cache_frozen && cpu_expert_profile_admit_enabled() {
2731 for (j, &ex) in sel.iter().enumerate() {
2732 if cpu_mask[j] {
2733 Self::moe_profile_admit_expert(e, il, ex as usize, m, max_block)?;
2734 }
2735 }
2736 }
2737 }
2738
2739 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
2744 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
2745 {
2746 let n_ff_sh = gate_shexp.out_features(); let verify_t = t > 1 && t < PRIME_MIN_T;
2755 let (sg_gate, sg_up) = if t == 1 {
2756 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
2757 Some(pair) => pair,
2758 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
2759 }
2760 } else if verify_t {
2761 (e.matmul_decode_exact(gate_shexp, z, t)?, e.matmul_decode_exact(up_shexp, z, t)?)
2762 } else {
2763 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?) };
2765 let mut sa = e.uninit(t * n_ff_sh)?; Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
2767 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
2768 else { e.matmul(down_shexp, &sa, t)? }; let g = match &m.gate_inp_shexp {
2782 Some(gate_inp_shexp) => {
2783 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
2784 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
2785 } else {
2786 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
2787 let mut g = e.uninit(t)?; e.sigmoid(&gs, &mut g, t)?;
2789 g
2790 }
2791 }
2792 None => e.htod(&vec![1.0f32; t])?,
2793 };
2794 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
2796 }
2797
2798 Ok(moe_out)
2799 }
2800
2801 pub fn stage1_h2d_per_token(&self) -> u64 {
2804 use crate::hybrid::Ffn;
2805 let n_used = self.cfg.moe.as_ref().map(|m| m.expert_used_count as u64).unwrap_or(0);
2806 let mut bytes = 0u64;
2807 for l in self.layers.iter() {
2808 if let Ffn::Moe(m) = &l.ffn {
2809 bytes += n_used * (m.gate_exps.max_expert_bytes() + m.up_exps.max_expert_bytes()
2810 + m.down_exps.max_expert_bytes()) as u64;
2811 }
2812 }
2813 bytes
2814 }
2815
2816 pub(crate) fn max_moe_block(&self) -> usize {
2820 use crate::hybrid::Ffn;
2821 let mut mx = 0usize;
2822 let mut scan = |ffn: &Ffn| {
2823 if let Ffn::Moe(m) = ffn {
2824 mx = mx.max(m.gate_exps.max_expert_bytes())
2825 .max(m.up_exps.max_expert_bytes())
2826 .max(m.down_exps.max_expert_bytes());
2827 }
2828 };
2829 for l in self.layers.iter() { scan(&l.ffn); }
2830 if let Some(mtp) = self.mtp.as_ref() { scan(&mtp.ffn); }
2831 mx
2832 }
2833
2834 pub(crate) fn moe_cache_block_sizes(&self) -> Vec<usize> {
2837 use crate::hybrid::Ffn;
2838 let mut sizes = Vec::new();
2839 let mut scan = |ffn: &Ffn| {
2840 let Ffn::Moe(m) = ffn else { return };
2841 for ex in 0..m.gate_exps.n_expert {
2842 if m.active_experts.as_ref().is_some_and(|active| !active[ex]) {
2843 continue;
2844 }
2845 for exps in [&m.gate_exps, &m.up_exps, &m.down_exps] {
2846 let len = exps.expert_layout(ex).len;
2847 if len > 0 {
2848 sizes.push(len);
2849 }
2850 }
2851 }
2852 };
2853 for layer in &self.layers {
2854 scan(&layer.ffn);
2855 }
2856 if let Some(mtp) = &self.mtp {
2857 scan(&mtp.ffn);
2858 }
2859 sizes
2860 }
2861
2862 pub fn save_cpu_expert_residency_profile(
2868 &self,
2869 e: &Engine,
2870 path: &std::path::Path,
2871 ) -> Result<(), Box<dyn std::error::Error>> {
2872 let Some(ids) = e.export_moe_residency() else {
2873 return Err("no MoE residency cache to persist".into());
2874 };
2875 let mut body = format!(
2876 "memra-freeze-profile v1 max_block={} blocks={}\n",
2877 self.max_moe_block(),
2878 ids.len()
2879 );
2880 for (layer, proj, ex) in &ids {
2881 body.push_str(&format!("{layer} {proj} {ex}\n"));
2882 }
2883 let tmp = path.with_extension("tmp");
2884 std::fs::write(&tmp, body)?;
2885 std::fs::rename(&tmp, path)?;
2886 println!(
2887 "[moe-cache] freeze profile saved: {} blocks -> {}",
2888 ids.len(),
2889 path.display()
2890 );
2891 Ok(())
2892 }
2893
2894 pub fn restore_cpu_expert_residency_profile(
2898 &self,
2899 e: &Engine,
2900 path: &std::path::Path,
2901 ) -> Result<bool, Box<dyn std::error::Error>> {
2902 use crate::hybrid::Ffn;
2903 use crate::moe_cache::BlockId;
2904 let Ok(content) = std::fs::read_to_string(path) else {
2905 return Ok(false);
2906 };
2907 let mut lines = content.lines();
2908 let Some(header) = lines.next() else { return Ok(false) };
2909 let expected = format!("memra-freeze-profile v1 max_block={}", self.max_moe_block());
2910 if !header.starts_with(&expected) {
2911 println!(
2912 "[moe-cache] freeze profile ignored (geometry mismatch): {}",
2913 path.display()
2914 );
2915 return Ok(false);
2916 }
2917 let mut by_layer: std::collections::HashMap<u16, Vec<BlockId>> =
2918 std::collections::HashMap::new();
2919 for line in lines {
2920 let mut fields = line.split_whitespace();
2921 let (Some(layer), Some(proj), Some(ex)) =
2922 (fields.next(), fields.next(), fields.next())
2923 else {
2924 continue;
2925 };
2926 let (Ok(layer), Ok(proj), Ok(ex)) =
2927 (layer.parse::<u16>(), proj.parse::<u8>(), ex.parse::<u16>())
2928 else {
2929 continue;
2930 };
2931 by_layer
2932 .entry(layer)
2933 .or_default()
2934 .push(BlockId::new(layer, proj, ex));
2935 }
2936 let requested: usize = by_layer.values().map(Vec::len).sum();
2937 if requested == 0 {
2938 return Ok(false);
2939 }
2940 let max_block = self.max_moe_block();
2941 let mut restaged = 0usize;
2942 let mut stage_layer = |layer_index: u16,
2943 ffn: &Ffn|
2944 -> Result<(), Box<dyn std::error::Error>> {
2945 let Ffn::Moe(m) = ffn else { return Ok(()) };
2946 let Some(ids) = by_layer.get(&layer_index) else {
2947 return Ok(());
2948 };
2949 e.with_moe_cache(max_block, |cache, eng| {
2950 for id in ids {
2951 if cache.restage_block(*id, m, eng)? {
2952 restaged += 1;
2953 }
2954 }
2955 Ok(())
2956 })
2957 };
2958 for (index, layer) in self.layers.iter().enumerate() {
2959 stage_layer(index as u16, &layer.ffn)?;
2960 }
2961 if let Some(mtp) = self.mtp.as_ref() {
2962 stage_layer(u16::MAX, &mtp.ffn)?;
2963 }
2964 e.freeze_moe_cache();
2965 println!(
2966 "[moe-cache] freeze profile restored: {restaged}/{requested} blocks restaged from {}",
2967 path.display()
2968 );
2969 Ok(true)
2970 }
2971
2972 pub fn freeze_cpu_expert_residency(
2974 &self,
2975 e: &Engine,
2976 ) -> Result<(), Box<dyn std::error::Error>> {
2977 e.freeze_moe_cache();
2978 Ok(())
2979 }
2980
2981 pub fn ffn_act(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
2985 act: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
2986 Self::ffn_act_scaled(e, cfg, gate, up, 1.0, 1.0, act, n)
2987 }
2988
2989 #[allow(clippy::too_many_arguments)]
2993 pub(crate) fn ffn_act_scaled(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
2994 gs: f32, us: f32, act: &mut CudaSlice<f32>, n: usize)
2995 -> Result<(), Box<dyn std::error::Error>> {
2996 if let Some(m3) = cfg.m3.as_ref() {
2997 return e.swigluoai_mul_scaled(gate, up, gs, us, m3.swiglu_alpha, m3.swiglu_limit, act, n);
2998 }
2999 if gs == 1.0 && us == 1.0 { return e.silu_mul(gate, up, act, n); }
3000 e.silu_mul_scaled(gate, up, gs, us, act, n)
3001 }
3002
3003 fn moe_route(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
3009 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
3010 Self::moe_route_cfg(e, logits, t, n_expert, n_used, None, None, None)
3011 }
3012
3013 fn moe_route_cfg(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize,
3021 bias: Option<&[f32]>, sig: Option<(f32, bool)>, active: Option<&[bool]>)
3022 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
3023 if let Some((sf, route_norm)) = sig {
3024 let lg = e.dtoh(logits)?;
3026 return Self::moe_route_sigmoid_host(
3027 &lg, t, n_expert, n_used, bias, sf, route_norm, active,
3028 );
3029 }
3030 if active.is_none() && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")) {
3034 return e.moe_router_topk_host(logits, t, n_expert, n_used);
3035 }
3036 let lg = e.dtoh(logits)?; let mut sel = vec![0u32; t * n_used];
3039 let mut w_out = vec![0f32; t * n_used];
3040 for tok in 0..t {
3041 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
3042 let maxl = row.iter().enumerate()
3044 .filter(|(i, _)| active.is_none_or(|mask| mask[*i]))
3045 .map(|(_, &x)| x).fold(f32::NEG_INFINITY, f32::max);
3046 let mut probs = vec![0f32; n_expert];
3047 let mut den = 0f32;
3048 for i in 0..n_expert {
3049 if active.is_some_and(|mask| !mask[i]) { continue; }
3050 let x = (row[i] - maxl).exp(); probs[i] = x; den += x;
3051 }
3052 for p in probs.iter_mut() { *p /= den; }
3053 let mut idx: Vec<usize> = (0..n_expert)
3055 .filter(|&i| active.is_none_or(|mask| mask[i])).collect();
3056 idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
3057 let sl = &idx[..n_used];
3058 let mut wv: Vec<f32> = sl.iter().map(|&i| probs[i]).collect();
3059 let mut ws: f32 = wv.iter().sum();
3060 ws = ws.max(6.103515625e-5_f32); for x in wv.iter_mut() { *x /= ws; }
3062 for j in 0..n_used {
3063 sel[tok * n_used + j] = sl[j] as u32;
3064 w_out[tok * n_used + j] = wv[j];
3065 }
3066 }
3067 Ok((sel, w_out))
3068 }
3069
3070 #[allow(clippy::too_many_arguments)]
3071 fn moe_route_sigmoid_with_input(
3072 e: &Engine,
3073 logits: &CudaSlice<f32>,
3074 input: &CudaSlice<f32>,
3075 t: usize,
3076 n_expert: usize,
3077 n_used: usize,
3078 bias: Option<&[f32]>,
3079 (sf, route_norm): (f32, bool),
3080 active: Option<&[bool]>,
3081 ) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
3082 let (lg, input) = e.dtoh_pair(logits, input)?;
3083 let (sel, w) =
3084 Self::moe_route_sigmoid_host(&lg, t, n_expert, n_used, bias, sf, route_norm, active)?;
3085 Ok((sel, w, input))
3086 }
3087
3088 pub fn start_moe_prefetch_predictor(
3093 &self,
3094 e: &Engine,
3095 cfg: &ModelConfig,
3096 ) -> Result<(), Box<dyn std::error::Error>> {
3097 use crate::hybrid::Ffn;
3098 let Some(sig) = cfg.sigmoid_router() else {
3099 return Err("prefetch predictor requires a sigmoid-router arch".into());
3100 };
3101 let resident: std::collections::HashSet<(u16, u8, u16)> = e
3102 .export_moe_residency()
3103 .ok_or("prefetch predictor needs the frozen MoE residency cache")?
3104 .into_iter()
3105 .collect();
3106 let mut layers = Vec::new();
3107 for (index, layer) in self.layers.iter().enumerate() {
3108 let Ffn::Moe(m) = &layer.ffn else { continue };
3109 let crate::model::GpuTensor::Float { data, .. } = &m.gate_inp else { continue };
3110 let router = e.dtoh(data)?;
3111 let n_expert = m.gate_exps.n_expert;
3112 let n_embd = m.gate_exps.in_f;
3113 if router.len() != n_embd * n_expert {
3114 continue;
3115 }
3116 let build = |exps: &crate::model::HostExps| {
3117 (0..n_expert)
3118 .map(|expert| crate::cpu_experts::predictor_projection(exps, expert))
3119 .collect::<Vec<_>>()
3120 };
3121 layers.push((index as u16, crate::cpu_experts::PredictLayerInit {
3122 router,
3123 bias: m.exp_probs_b.clone(),
3124 active: m.active_experts.clone(),
3125 n_embd,
3126 n_used: cfg
3127 .moe
3128 .as_ref()
3129 .map(|moe| moe.expert_used_count as usize)
3130 .ok_or("prefetch predictor requires MoE config")?,
3131 sig,
3132 weights_n_expert: n_expert,
3133 gate: build(&m.gate_exps),
3134 up: build(&m.up_exps),
3135 down: build(&m.down_exps),
3136 }));
3137 }
3138 crate::cpu_experts::start_prefetch_predictor(layers, resident)
3139 .map_err(|error| error.into())
3140 }
3141
3142 #[allow(clippy::too_many_arguments)]
3145 pub(crate) fn moe_route_sigmoid_host_public(
3146 logits: &[f32],
3147 t: usize,
3148 n_expert: usize,
3149 n_used: usize,
3150 bias: Option<&[f32]>,
3151 sf: f32,
3152 route_norm: bool,
3153 active: Option<&[bool]>,
3154 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
3155 Self::moe_route_sigmoid_host(logits, t, n_expert, n_used, bias, sf, route_norm, active)
3156 }
3157
3158 #[allow(clippy::too_many_arguments)]
3159 fn moe_route_sigmoid_host(
3160 lg: &[f32],
3161 t: usize,
3162 n_expert: usize,
3163 n_used: usize,
3164 bias: Option<&[f32]>,
3165 sf: f32,
3166 route_norm: bool,
3167 active: Option<&[bool]>,
3168 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
3169 if lg.len() != t * n_expert {
3170 return Err(format!(
3171 "sigmoid router logits length mismatch: got {}, expected {}",
3172 lg.len(),
3173 t * n_expert,
3174 )
3175 .into());
3176 }
3177 let mut sel = vec![0u32; t * n_used];
3178 let mut w_out = vec![0f32; t * n_used];
3179 for tok in 0..t {
3180 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
3181 let scores: Vec<f32> = row.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect();
3182 let selsc: Vec<f32> = match bias {
3184 Some(b) => scores.iter().zip(b).map(|(s, bb)| s + bb).collect(),
3185 None => scores.clone(),
3186 };
3187 let mut idx: Vec<usize> = (0..n_expert)
3188 .filter(|&i| active.is_none_or(|mask| mask[i]))
3189 .collect();
3190 idx.sort_by(|&a, &b| selsc[b].total_cmp(&selsc[a]).then(a.cmp(&b)));
3191 let sl = &idx[..n_used];
3192 let mut wv: Vec<f32> = sl.iter().map(|&i| scores[i]).collect();
3193 if route_norm {
3194 let ws: f32 = wv.iter().sum::<f32>().max(1e-20);
3195 for x in wv.iter_mut() {
3196 *x = *x / ws * sf;
3197 }
3198 } else {
3199 for x in wv.iter_mut() {
3200 *x *= sf;
3201 }
3202 }
3203 for j in 0..n_used {
3204 sel[tok * n_used + j] = sl[j] as u32;
3205 w_out[tok * n_used + j] = wv[j];
3206 }
3207 }
3208 Ok((sel, w_out))
3209 }
3210
3211 fn moe_ffn_pairs(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, logits: &CudaSlice<f32>,
3220 t: usize, cfg: &ModelConfig)
3221 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3222 let moe = cfg.moe.as_ref().unwrap();
3223 let n_embd = cfg.n_embd as usize;
3224 let n_expert = moe.expert_count as usize;
3225 let n_used = moe.expert_used_count as usize;
3226 let n_ff_exp = moe.expert_ff_length as usize;
3227 let dev = m.dev_exps.as_ref().unwrap();
3228 let (rbg_d, rbu_d) = if dev.gu_il {
3230 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
3231 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
3232
3233 let (sel_all, w_all) = Self::moe_route(e, logits, t, n_expert, n_used)?;
3234 let n_pairs = t * n_used;
3235 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
3238 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
3239 let pair_w: Vec<f32> = w_all.clone();
3240 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
3241 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
3242 let pt = e.htod_i32(&pair_tok)?;
3243 let px = e.htod_i32(&pair_ex)?;
3244 let pw = e.htod(&pair_w)?;
3245 let toff = e.htod_i32(&tok_off)?;
3246 let tids = e.htod_i32(&tok_ids)?;
3247
3248 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
3252 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
3253 let mut ex_ids: Vec<i32> = Vec::new();
3254 let mut ex_off: Vec<i32> = vec![0];
3255 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
3256 for (ex, list) in by_ex.iter().enumerate() {
3257 if list.is_empty() { continue; }
3258 ex_ids.push(ex as i32);
3259 ex_pairs.extend_from_slice(list);
3260 ex_off.push(ex_pairs.len() as i32);
3261 }
3262 let n_active = ex_ids.len();
3263 let exi = e.htod_i32(&ex_ids)?;
3264 let exo = e.htod_i32(&ex_off)?;
3265 let exp_d = e.htod_i32(&ex_pairs)?;
3266 let _ = &px; static MMA_T: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
3287 let mma_t = *MMA_T.get_or_init(|| {
3288 std::env::var("MEMRA_MOE_MMA_T").ok().and_then(|v| v.parse().ok()).unwrap_or(16)
3289 });
3290 let use_mma = std::env::var("MEMRA_MOE_MMA").map(|v| v != "0").unwrap_or(true)
3291 && t >= mma_t
3292 && q8_expert_dec_supported(m.gate_exps.qtype) && q8_expert_dec_supported(m.up_exps.qtype)
3293 && q8_expert_dec_supported(m.down_exps.qtype)
3294 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
3295 let mma_capable = q8_expert_dec_supported(m.gate_exps.qtype)
3311 && q8_expert_dec_supported(m.up_exps.qtype)
3312 && q8_expert_dec_supported(m.down_exps.qtype)
3313 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
3314 let f16g_mode = crate::moe_f16g_mode();
3315 let f16g = f16g_mode != 0 && t >= mma_t
3316 && (f16g_mode != 3 || !mma_capable)
3317 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
3318 && f16g_proj_ok(m.up_exps.qtype, n_embd)
3319 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp);
3320 if use_mma || f16g {
3321 let y_down = if f16g {
3329 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
3333 let csr_tok_d = e.htod_i32(&csr_tok)?;
3334 let (z_f16, z_s) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
3335 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
3336 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
3337 m.gate_exps.qtype, rbg_d)?;
3338 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
3339 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
3340 m.up_exps.qtype, rbu_d)?;
3341 let act_csr = e.moe_pairs_silu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
3342 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
3343 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
3344 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
3345 m.down_exps.qtype, m.down_exps.row_bytes)?;
3346 e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?
3347 } else {
3348 let z_scr = e.mmq_iq_quantize_act(z, n_embd, t)?;
3350 let gate = e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
3351 n_embd, n_ff_exp, n_active, n_pairs, t,
3352 m.gate_exps.qtype, rbg_d)?;
3353 let up = e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
3354 n_embd, n_ff_exp, n_active, n_pairs, t,
3355 m.up_exps.qtype, rbu_d)?;
3356 let a_scr = if crate::moe_fuse_actq_on() {
3362 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 0)?
3363 } else {
3364 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
3365 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
3366 };
3367 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
3368 let pself = e.htod_i32(&pair_self)?;
3369 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
3370 n_ff_exp, n_embd, n_active, n_pairs, n_pairs,
3371 m.down_exps.qtype, m.down_exps.row_bytes)?
3372 };
3373 let mut moe_out = e.uninit(t * n_embd)?;
3374 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
3375 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
3376 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
3377 {
3378 let n_ff_sh = gate_shexp.out_features();
3379 let sg_gate = e.matmul(gate_shexp, z, t)?;
3380 let sg_up = e.matmul(up_shexp, z, t)?;
3381 let mut sa = e.uninit(t * n_ff_sh)?;
3382 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
3383 let sh = e.matmul(down_shexp, &sa, t)?;
3384 let g = match &m.gate_inp_shexp {
3390 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
3391 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
3392 }
3393 Some(gate_inp_shexp) => {
3394 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
3395 let mut g = e.uninit(t)?;
3396 e.sigmoid(&gs, &mut g, t)?;
3397 g
3398 }
3399 None => e.htod(&vec![1.0f32; t])?,
3400 };
3401 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
3402 }
3403 return Ok(moe_out);
3404 }
3405
3406 let dec = std::env::var("MEMRA_MOE_DEC").map(|v| v != "0").unwrap_or(true);
3409 let matvec = |proj, exi: &_, exo: &_, exp_d: &_, pt: &_, aq: &_, ad: &_,
3410 inf, outf, qtype, rb| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3411 let dec = dec && q8_expert_dec_supported(qtype);
3413 if dec { e.moe_pairs_matvec_q8_dec(&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
3414 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
3415 else { e.moe_pairs_matvec_q8_em (&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
3416 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
3417 };
3418 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3419 let gate = matvec(0, &exi, &exo, &exp_d, &pt, &zq, &zd,
3420 n_embd, n_ff_exp, m.gate_exps.qtype, rbg_d)?;
3421 let up = matvec(1, &exi, &exo, &exp_d, &pt, &zq, &zd,
3422 n_embd, n_ff_exp, m.up_exps.qtype, rbu_d)?;
3423 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
3424 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
3425 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
3427 let pself = e.htod_i32(&pair_self)?;
3428 let y_down = matvec(2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
3429 n_ff_exp, n_embd, m.down_exps.qtype, m.down_exps.row_bytes)?;
3430 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
3432
3433 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
3437 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
3438 {
3439 let n_ff_sh = gate_shexp.out_features();
3440 let sg_gate = e.matmul(gate_shexp, z, t)?;
3441 let sg_up = e.matmul(up_shexp, z, t)?;
3442 let mut sa = e.uninit(t * n_ff_sh)?;
3443 e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
3444 let sh = e.matmul(down_shexp, &sa, t)?;
3445 let g = match &m.gate_inp_shexp {
3450 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
3451 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
3452 }
3453 Some(gate_inp_shexp) => {
3454 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
3455 let mut g = e.uninit(t)?;
3456 e.sigmoid(&gs, &mut g, t)?;
3457 g
3458 }
3459 None => e.htod(&vec![1.0f32; t])?,
3460 };
3461 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
3462 }
3463 Ok(moe_out)
3464 }
3465
3466 #[allow(clippy::too_many_arguments)]
3468 #[allow(clippy::too_many_arguments)]
3469 fn moe_ffn_dev(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
3470 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, logits: &CudaSlice<f32>,
3471 t: usize, cfg: &ModelConfig, il: u16, max_block: usize)
3472 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3473 let moe = cfg.moe.as_ref().unwrap();
3474 let n_embd = cfg.n_embd as usize;
3475 let n_expert = moe.expert_count as usize;
3476 let n_used = moe.expert_used_count as usize;
3477 let n_ff_exp = moe.expert_ff_length as usize;
3478
3479 let (sel_d, mut w_d) = e.moe_router_topk(logits, t, n_expert, n_used)?;
3481 if m.has_macros {
3484 e.moe_w_scale_by_expert(&mut w_d, &sel_d, &m.dev_macros, n_expert, t * n_used)?;
3485 }
3486
3487 let mut moe_out = e.uninit(t * n_embd)?;
3489
3490 if let Some(dev) = m.dev_exps.as_ref() {
3493 let (rbg_d, rbu_d) = if dev.gu_il {
3496 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
3497 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
3498 let q8 = moe_q8_enabled()
3499 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3500 && q8_expert_supported(m.down_exps.qtype);
3501 let rows_arm = q8 && t > 1 && crate::spec::spec_m2()
3510 && n_ff_exp == 512 && n_used <= 8
3511 && std::env::var("MEMRA_MOE_DEVQ8_GU").map(|v| v.is_empty() || v == "v").unwrap_or(true)
3512 && std::env::var("MEMRA_MOE_DEVQ8_DOWN").map(|v| v.is_empty() || v == "w8h2v").unwrap_or(true);
3513 let csr_mode = std::env::var("MEMRA_MOE_CSR").ok()
3522 .and_then(|v| v.parse::<i32>().ok()).unwrap_or(1);
3523 let csr_qt = |qt: i32| qt == crate::QT_IQ4_XS || qt == crate::QT_IQ3_S;
3524 let csr_arm = rows_arm && csr_mode > 0 && t <= 10
3525 && csr_qt(m.gate_exps.qtype) && csr_qt(m.up_exps.qtype)
3526 && csr_qt(m.down_exps.qtype);
3527 if csr_arm {
3528 if csr_mode == 2 {
3529 static ENGAGED: std::sync::Once = std::sync::Once::new();
3530 ENGAGED.call_once(|| eprintln!("[memra] moe CSR byte-compare mode ON (t={t})"));
3531 }
3532 let n_pairs = t * n_used;
3533 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3534 let act = e.moe_gate_up_silu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, n_pairs,
3535 n_embd, n_ff_exp, n_used, n_expert,
3536 m.gate_exps.qtype, m.up_exps.qtype,
3537 rbg_d, rbu_d)?;
3538 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
3539 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
3543 t, n_ff_exp, n_embd, n_used, n_expert,
3544 m.down_exps.qtype, m.down_exps.row_bytes)?;
3545 if csr_mode == 2 {
3546 let act_r = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
3548 n_embd, n_ff_exp, n_used, n_expert,
3549 m.gate_exps.qtype, m.up_exps.qtype,
3550 rbg_d, rbu_d, &m.dev_macros)?;
3551 let mut out_r = e.uninit(t * n_embd)?;
3552 let (aq2r, ad2r) = e.quantize_q8_1(&act_r, n_pairs, n_ff_exp)?;
3553 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2r, &ad2r, &mut out_r,
3554 t, n_ff_exp, n_embd, n_used, n_expert,
3555 m.down_exps.qtype, m.down_exps.row_bytes)?;
3556 let (a1, a2) = (e.dtoh(&act)?, e.dtoh(&act_r)?);
3557 let (o1, o2) = (e.dtoh(&moe_out)?, e.dtoh(&out_r)?);
3558 let ba = a1.iter().zip(&a2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
3559 let bo = o1.iter().zip(&o2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
3560 if ba + bo > 0 {
3561 eprintln!("[csr-check] il={il} t={t} ACT diffs={ba}/{} OUT diffs={bo}/{}",
3562 a1.len(), o1.len());
3563 let sel_h = e.dtoh_i32(&sel_d)?;
3565 let mut shown = 0;
3566 for (i, (x, y)) in a1.iter().zip(&a2).enumerate() {
3567 if x.to_bits() != y.to_bits() && shown < 4 {
3568 let (p, o) = (i / n_ff_exp, i % n_ff_exp);
3569 let ex = sel_h[p];
3570 let npx = sel_h.iter().filter(|&&v| v == ex).count();
3571 eprintln!(" ACT p={p} ex={ex} np={npx} o={o} csr={x:e} rows={y:e}");
3572 shown += 1;
3573 }
3574 }
3575 std::process::exit(3);
3576 }
3577 }
3578 } else if rows_arm {
3579 if std::env::var("MEMRA_MOE_OVERLAP").as_deref() == Ok("1") {
3582 use std::sync::atomic::{AtomicU64, Ordering};
3583 static PAIRS: AtomicU64 = AtomicU64::new(0);
3584 static UNIQ: AtomicU64 = AtomicU64::new(0);
3585 static CALLS: AtomicU64 = AtomicU64::new(0);
3586 let sel_h = e.dtoh_i32(&sel_d)?;
3587 let mut u: Vec<i32> = sel_h.clone(); u.sort_unstable(); u.dedup();
3588 PAIRS.fetch_add(sel_h.len() as u64, Ordering::Relaxed);
3589 UNIQ.fetch_add(u.len() as u64, Ordering::Relaxed);
3590 let c = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
3591 if c % 480 == 0 {
3592 let p = PAIRS.load(Ordering::Relaxed); let q = UNIQ.load(Ordering::Relaxed);
3593 eprintln!("[overlap] calls={c} pairs={p} unique={q} ratio={:.3} (t={t})",
3594 q as f64 / p as f64);
3595 }
3596 }
3597 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3598 let act = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
3599 n_embd, n_ff_exp, n_used, n_expert,
3600 m.gate_exps.qtype, m.up_exps.qtype,
3601 rbg_d, rbu_d, &m.dev_macros)?;
3602 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
3603 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
3604 t, n_ff_exp, n_embd, n_used, n_expert,
3605 m.down_exps.qtype, m.down_exps.row_bytes)?;
3606 } else {
3607 for tok in 0..t {
3608 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
3609 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
3610 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
3611 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3612 if q8 {
3613 let (zq, zd) = match (t, zq8) {
3614 (1, Some((q, d))) => (q.clone(), d.clone()),
3615 _ => e.quantize_q8_1_view(&zt, 1, n_embd)?,
3616 };
3617 let act = e.moe_gate_up_silu8_dev_q8(&dev.ptr_row, &selt, &zq, &zd,
3618 n_embd, n_ff_exp, n_used, n_expert,
3619 m.gate_exps.qtype, m.up_exps.qtype,
3620 rbg_d, rbu_d, &m.dev_macros)?;
3621 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
3622 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selt, &wt, &aq2, &ad2, &mut dst,
3623 n_ff_exp, n_embd, n_used, n_expert,
3624 m.down_exps.qtype, m.down_exps.row_bytes)?;
3625 } else {
3626 let act = e.moe_gate_up_silu8_dev(&dev.ptr_row, &selt, &zt, n_embd, n_ff_exp,
3627 n_used, n_expert,
3628 m.gate_exps.qtype, m.up_exps.qtype,
3629 rbg_d, rbu_d, &m.dev_macros)?;
3630 e.moe_down8_fma_dev(&dev.ptr_row, &selt, &wt, &act, &mut dst,
3631 n_ff_exp, n_embd, n_used, n_expert,
3632 m.down_exps.qtype, m.down_exps.row_bytes)?;
3633 }
3634 }
3635 }
3636 } else {
3637 let q8 = moe_q8_enabled()
3644 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3645 && q8_expert_supported(m.down_exps.qtype);
3646 e.with_moe_cache(max_block, |c, eng| {
3647 let row = c.layer_dev_row(il, n_expert, eng)?
3648 .ok_or("moe_ffn_dev: layer row vanished under the lock")?;
3649 for tok in 0..t {
3650 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
3651 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
3652 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
3653 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3654 if q8 {
3655 let (zq, zd) = match (t, zq8) {
3656 (1, Some((q, d))) => (q.clone(), d.clone()),
3657 _ => eng.quantize_q8_1_view(&zt, 1, n_embd)?,
3658 };
3659 let act = eng.moe_gate_up_silu8_dev_q8(row, &selt, &zq, &zd,
3660 n_embd, n_ff_exp, n_used, n_expert,
3661 m.gate_exps.qtype, m.up_exps.qtype,
3662 m.gate_exps.row_bytes, m.up_exps.row_bytes,
3663 &m.dev_macros)?;
3664 let (aq2, ad2) = eng.quantize_q8_1(&act, n_used, n_ff_exp)?;
3665 eng.moe_down8_fma_dev_q8(row, &selt, &wt, &aq2, &ad2, &mut dst,
3666 n_ff_exp, n_embd, n_used, n_expert,
3667 m.down_exps.qtype, m.down_exps.row_bytes)?;
3668 } else {
3669 let act = eng.moe_gate_up_silu8_dev(row, &selt, &zt, n_embd, n_ff_exp,
3670 n_used, n_expert,
3671 m.gate_exps.qtype, m.up_exps.qtype,
3672 m.gate_exps.row_bytes, m.up_exps.row_bytes,
3673 &m.dev_macros)?;
3674 eng.moe_down8_fma_dev(row, &selt, &wt, &act, &mut dst,
3675 n_ff_exp, n_embd, n_used, n_expert,
3676 m.down_exps.qtype, m.down_exps.row_bytes)?;
3677 }
3678 }
3679 c.hits += (t * 3 * n_used) as u64;
3681 Ok(())
3682 })?;
3683 }
3684
3685 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
3690 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
3691 {
3692 let n_ff_sh = gate_shexp.out_features();
3693 let verify_t = t > 1 && t < PRIME_MIN_T;
3696 let (sg_gate, sg_up) = if t == 1 {
3697 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
3698 Some(pair) => pair,
3699 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
3700 }
3701 } else if verify_t {
3702 let mut fused = None;
3706 if crate::spec::spec_fused_t() && (2..=4).contains(&t)
3707 && e.uses_q8_1_fast(gate_shexp) && e.uses_q8_1_fast(up_shexp) {
3708 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3709 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
3710 }
3711 match fused {
3712 Some(pair) => pair,
3713 None => (e.matmul_decode_exact(gate_shexp, z, t)?,
3714 e.matmul_decode_exact(up_shexp, z, t)?),
3715 }
3716 } else {
3717 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
3718 };
3719 let mut sa = e.uninit(t * n_ff_sh)?; e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
3721 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
3722 else { e.matmul(down_shexp, &sa, t)? };
3723 let g = match &m.gate_inp_shexp {
3727 Some(gate_inp_shexp) => {
3728 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
3731 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
3732 } else {
3733 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
3734 let mut g = e.uninit(t)?;
3735 e.sigmoid(&gs, &mut g, t)?;
3736 g
3737 }
3738 }
3739 None => e.htod(&vec![1.0f32; t])?,
3740 };
3741 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
3742 }
3743
3744 Ok(moe_out)
3745 }
3746
3747 #[allow(clippy::too_many_arguments)]
3757 #[allow(clippy::too_many_arguments)]
3760 fn moe_gdec_token_q8(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
3761 zq: &CudaSlice<i8>, zd: &CudaSlice<f32>, sel: &[u32], w: &[f32],
3762 moe_out: &mut CudaSlice<f32>, tok: usize,
3763 n_embd: usize, n_ff_exp: usize, n_used: usize)
3764 -> Result<bool, Box<dyn std::error::Error>> {
3765 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
3766 use cudarc::driver::DevicePtr;
3767 let ptrs = e.with_moe_cache(max_block, |c, eng| {
3768 let mut g = [0u64; 8];
3769 let mut u = [0u64; 8];
3770 let mut d = [0u64; 8];
3771 for (j, &ex) in sel.iter().enumerate() {
3772 let ex = ex as u16;
3773 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
3774 c.resident(BlockId::new(il, PROJ_UP, ex)),
3775 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
3776 else { return Ok(None); };
3777 let __s = eng.stream();
3778 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
3779 let (pu, _e1) = c.slot(su).device_ptr(&__s);
3780 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
3781 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
3782 }
3783 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
3784 for &ex in sel {
3785 let ex = ex as u16;
3786 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
3787 c.note_profile_hit(BlockId::new(il, proj, ex));
3788 }
3789 }
3790 }
3791 c.hits += (3 * n_used) as u64;
3792 Ok(Some((g, u, d)))
3793 })?;
3794 let Some((g, u, d)) = ptrs else { return Ok(false) };
3795 let mut wv = [0f32; 8];
3796 wv[..n_used].copy_from_slice(w);
3797 let act = e.moe_gate_up_silu8_q8(crate::WPtr8(g), crate::WPtr8(u), zq, zd,
3798 n_embd, n_ff_exp, n_used,
3799 m.gate_exps.qtype, m.up_exps.qtype,
3800 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
3801 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
3803 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3804 e.moe_down8_fma_q8(crate::WPtr8(d), crate::F32x8(wv), &aq2, &ad2, &mut dst,
3805 n_ff_exp, n_embd, n_used,
3806 m.down_exps.qtype, m.down_exps.row_bytes)?;
3807 Ok(true)
3808 }
3809
3810 fn moe_gdec_token(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
3811 zt: &cudarc::driver::CudaView<f32>, sel: &[u32], w: &[f32],
3812 moe_out: &mut CudaSlice<f32>, tok: usize,
3813 n_embd: usize, n_ff_exp: usize, n_used: usize)
3814 -> Result<bool, Box<dyn std::error::Error>> {
3815 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
3816 use cudarc::driver::DevicePtr;
3817 let ptrs = e.with_moe_cache(max_block, |c, eng| {
3819 let mut g = [0u64; 8];
3820 let mut u = [0u64; 8];
3821 let mut d = [0u64; 8];
3822 for (j, &ex) in sel.iter().enumerate() {
3823 let ex = ex as u16;
3824 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
3825 c.resident(BlockId::new(il, PROJ_UP, ex)),
3826 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
3827 else { return Ok(None); };
3828 let __s = eng.stream();
3829 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
3830 let (pu, _e1) = c.slot(su).device_ptr(&__s);
3831 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
3832 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
3833 }
3834 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
3835 for &ex in sel {
3836 let ex = ex as u16;
3837 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
3838 c.note_profile_hit(BlockId::new(il, proj, ex));
3839 }
3840 }
3841 }
3842 c.hits += (3 * n_used) as u64; Ok(Some((g, u, d)))
3844 })?;
3845 let Some((g, u, d)) = ptrs else { return Ok(false) };
3846 let mut wv = [0f32; 8];
3847 wv[..n_used].copy_from_slice(w);
3848 let act = e.moe_gate_up_silu8(crate::WPtr8(g), crate::WPtr8(u), zt,
3850 n_embd, n_ff_exp, n_used,
3851 m.gate_exps.qtype, m.up_exps.qtype,
3852 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
3853 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3854 e.moe_down8_fma_into(crate::WPtr8(d), crate::F32x8(wv), &act, &mut dst,
3855 n_ff_exp, n_embd, n_used,
3856 m.down_exps.qtype, m.down_exps.row_bytes)?;
3857 Ok(true)
3858 }
3859
3860 fn moe_cached_gemm_q8(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
3865 max_block: usize, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
3866 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3867 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
3868 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
3869 let layout = exps.expert_layout(ex);
3870 let id = BlockId::new(il, proj, ex as u16);
3871 let source = exps.expert_source(ex);
3872 e.with_moe_cache(max_block, |c, eng| {
3873 let slot = c.dispatch_source(id, source, eng)?;
3874 let DispatchSlot::Resident(sl) = slot;
3875 let buf = c.slot(sl);
3876 eng.qmatvec_expert_q8(buf, 0..layout.len, aq, ad, 1, exps.in_f, exps.out_f,
3877 layout.qtype, layout.row_bytes)
3878 })
3879 }
3880
3881 fn moe_cached_gemm(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
3882 max_block: usize, x: &cudarc::driver::CudaView<f32>)
3883 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3884 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
3885 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
3886 let layout = exps.expert_layout(ex);
3887 let id = BlockId::new(il, proj, ex as u16);
3888 let source = exps.expert_source(ex);
3889 e.with_moe_cache(max_block, |c, eng| {
3891 let slot = c.dispatch_source(id, source, eng)?;
3892 let DispatchSlot::Resident(sl) = slot;
3895 let buf = c.slot(sl);
3896 eng.qmatvec_view(buf, 0..layout.len, x, 1, exps.in_f, exps.out_f,
3897 layout.qtype, layout.row_bytes)
3898 })
3899 }
3900
3901 fn moe_profile_admit_expert(
3905 e: &Engine,
3906 il: u16,
3907 ex: usize,
3908 m: &MoeWeights,
3909 max_block: usize,
3910 ) -> Result<(), Box<dyn std::error::Error>> {
3911 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3912 e.with_moe_cache(max_block, |cache, eng| {
3913 for (proj, exps) in [
3914 (PROJ_GATE, &m.gate_exps),
3915 (PROJ_UP, &m.up_exps),
3916 (PROJ_DOWN, &m.down_exps),
3917 ] {
3918 let id = BlockId::new(il, proj, ex as u16);
3919 let _ = cache.dispatch_source(id, exps.expert_source(ex), eng)?;
3920 }
3921 Ok(())
3922 })
3923 }
3924
3925 #[allow(clippy::too_many_arguments)]
3928 fn moe_frozen_gemm(
3929 e: &Engine,
3930 il: u16,
3931 proj: u8,
3932 ex: usize,
3933 m: &MoeWeights,
3934 max_block: usize,
3935 x: &cudarc::driver::CudaView<f32>,
3936 scratch: &mut Option<CudaSlice<u8>>,
3937 scratch_len: usize,
3938 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3939 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP};
3940 let exps = match proj {
3941 PROJ_GATE => &m.gate_exps,
3942 PROJ_UP => &m.up_exps,
3943 _ => &m.down_exps,
3944 };
3945 let layout = exps.expert_layout(ex);
3946 let id = BlockId::new(il, proj, ex as u16);
3947 if let Some(output) = e.with_moe_cache(max_block, |cache, eng| {
3948 let Some(slot) = cache.resident(id) else {
3949 return Ok(None);
3950 };
3951 let buf = cache.slot(slot);
3952 Ok(Some(eng.qmatvec_view(
3953 buf,
3954 0..layout.len,
3955 x,
3956 1,
3957 exps.in_f,
3958 exps.out_f,
3959 layout.qtype,
3960 layout.row_bytes,
3961 )?))
3962 })? {
3963 return Ok(output);
3964 }
3965 if scratch.is_none() {
3966 *scratch = Some(e.alloc_u8_uninit(scratch_len)?);
3967 }
3968 let scratch = scratch.as_mut().unwrap();
3969 e.stage_expert(exps.expert_bytes(ex), scratch, 0)?;
3970 e.qmatvec_view(
3971 scratch,
3972 0..layout.len,
3973 x,
3974 1,
3975 exps.in_f,
3976 exps.out_f,
3977 layout.qtype,
3978 layout.row_bytes,
3979 )
3980 }
3981
3982 fn moe_prefetch_expert(
3983 e: &Engine,
3984 il: u16,
3985 ex: usize,
3986 m: &MoeWeights,
3987 max_block: usize,
3988 keep: &[crate::moe_cache::BlockId],
3989 ) -> Result<(), Box<dyn std::error::Error>> {
3990 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3991 e.with_moe_cache(max_block, |c, eng| {
3992 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
3993 (PROJ_DOWN, &m.down_exps)] {
3994 let id = BlockId::new(il, proj, ex as u16);
3995 let _ = c.prefetch_source(id, exps.expert_source(ex), keep, eng)?;
3996 }
3997 Ok(())
3998 })
3999 }
4000
4001 fn moe_prefetch_disk_expert(e: &Engine, il: u16, ex: usize, m: &MoeWeights,
4004 max_block: usize, keep: &[crate::moe_cache::BlockId])
4005 -> Result<(), Box<dyn std::error::Error>> {
4006 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
4007 e.with_moe_cache(max_block, |c, eng| {
4008 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
4009 (PROJ_DOWN, &m.down_exps)] {
4010 let source = exps.expert_source(ex);
4011 if let crate::model::ExpertSource::Disk { .. } = &source {
4012 let id = BlockId::new(il, proj, ex as u16);
4013 let _ = c.prefetch_source(id, source, keep, eng)?;
4014 }
4015 }
4016 Ok(())
4017 })
4018 }
4019
4020 #[inline]
4021 fn moe_prefetch_host_expert(ex: usize, m: &MoeWeights) {
4022 let _ = m.gate_exps.prefetch_expert_pages(ex);
4023 let _ = m.up_exps.prefetch_expert_pages(ex);
4024 let _ = m.down_exps.prefetch_expert_pages(ex);
4025 }
4026}
4027
4028impl HybridModel {
4045 pub(crate) fn moe_ffn_grouped(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
4048 cfg: &ModelConfig, il: u16, _max_block: usize)
4049 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4050 let moe = cfg.moe.as_ref().unwrap();
4051 let n_embd = cfg.n_embd as usize;
4052 let n_expert = moe.expert_count as usize;
4053 let n_used = moe.expert_used_count as usize;
4054 let n_ff_exp = moe.expert_ff_length as usize;
4055
4056 let logits = e.matmul(&m.gate_inp, z, t)?;
4058 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
4059 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
4060 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
4061 } else {
4062 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
4063 None, None, m.active_experts.as_deref())?
4064 };
4065 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
4066
4067 struct ExpertGroup {
4071 tok_indices: Vec<i32>, slot_indices: Vec<i32>, weights: Vec<f32>, }
4075 let mut groups: Vec<ExpertGroup> = (0..n_expert).map(|_| ExpertGroup {
4076 tok_indices: Vec::new(), slot_indices: Vec::new(), weights: Vec::new(),
4077 }).collect();
4078
4079 for tok in 0..t {
4080 for j in 0..n_used {
4081 let ex = sel_all[tok * n_used + j] as usize;
4082 let w = w_all[tok * n_used + j];
4083 groups[ex].tok_indices.push(tok as i32);
4084 groups[ex].slot_indices.push(j as i32);
4085 groups[ex].weights.push(w);
4086 }
4087 }
4088
4089 let mut slot_buf = e.zeros(t * n_used * n_embd)?;
4092 let mut wbuf = e.zeros(t * n_used)?; let g_len = m.gate_exps.max_expert_bytes();
4096 let u_len = m.up_exps.max_expert_bytes();
4097 let d_len = m.down_exps.max_expert_bytes();
4098 let use_cache = Engine::moe_cache_enabled();
4099 let max_block = _max_block;
4100
4101 let (mut scratch_g, mut scratch_u, mut scratch_d) = if !use_cache {
4103 (Some(e.alloc_u8(g_len)?), Some(e.alloc_u8(u_len)?), Some(e.alloc_u8(d_len)?))
4104 } else {
4105 (None, None, None)
4106 };
4107
4108 let mut order: Vec<usize> =
4119 (0..n_expert).filter(|&ex| !groups[ex].tok_indices.is_empty()).collect();
4120 order.sort_by(|&a, &b| groups[b].tok_indices.len()
4121 .cmp(&groups[a].tok_indices.len()).then(a.cmp(&b)));
4122 let mut m_dist: Vec<usize> = Vec::new(); let page_window = moe_page_prefetch_window();
4124 let worker_disk_prefetch = use_cache && crate::spill_pread::worker_enabled();
4125 if worker_disk_prefetch {
4126 if let Some(first) = grouped_worker_prefetch_position(order.len(), None) {
4127 Self::moe_prefetch_disk_expert(e, il, order[first], m, max_block, &[])?;
4128 }
4129 }
4130 for (order_pos, &ex) in order.iter().enumerate() {
4131 for next in page_prefetch_positions(order_pos, order.len(), page_window) {
4132 Self::moe_prefetch_host_expert(order[next], m);
4133 }
4134 if worker_disk_prefetch {
4135 if let Some(next) = grouped_worker_prefetch_position(order.len(), Some(order_pos)) {
4136 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
4137 let keep = [
4138 BlockId::new(il, PROJ_GATE, ex as u16),
4139 BlockId::new(il, PROJ_UP, ex as u16),
4140 BlockId::new(il, PROJ_DOWN, ex as u16),
4141 ];
4142 Self::moe_prefetch_disk_expert(e, il, order[next], m, max_block, &keep)?;
4143 }
4144 }
4145 let grp = &groups[ex];
4146 let m_e = grp.tok_indices.len();
4147 m_dist.push(m_e);
4148 let gl = m.gate_exps.expert_layout(ex);
4149 let ul = m.up_exps.expert_layout(ex);
4150 let dl = m.down_exps.expert_layout(ex);
4151
4152 let tok_idx_d = e.htod_i32(&grp.tok_indices)?;
4156 let slot_idx_d = e.htod_i32(&grp.slot_indices)?;
4157 let dmac = m.down_exps.macro_scale(ex);
4158 let weight_d = if dmac == 1.0 { e.htod(&grp.weights)? } else {
4159 let scaled: Vec<f32> = grp.weights.iter().map(|&w| w * dmac).collect();
4160 e.htod(&scaled)?
4161 };
4162
4163 let mut gathered = e.zeros(m_e * n_embd)?;
4165 e.gather_rows(z, &tok_idx_d, &mut gathered, n_embd, m_e)?;
4166 let gv = gathered.slice(0..m_e * n_embd);
4167
4168 let y = if use_cache {
4170 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
4171 let gate = e.with_moe_cache(max_block, |c, eng| {
4173 let id = BlockId::new(il, PROJ_GATE, ex as u16);
4174 let slot = c.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
4175 let buf = c.buf(slot);
4176 eng.qmatvec_view(buf, 0..gl.len, &gv, m_e,
4177 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)
4178 })?;
4179 let up = e.with_moe_cache(max_block, |c, eng| {
4180 let id = BlockId::new(il, PROJ_UP, ex as u16);
4181 let slot = c.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
4182 let buf = c.buf(slot);
4183 eng.qmatvec_view(buf, 0..ul.len, &gv, m_e,
4184 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)
4185 })?;
4186 let mut act = e.zeros(m_e * n_ff_exp)?;
4188 Self::ffn_act_scaled(e, cfg, &gate, &up,
4189 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, m_e * n_ff_exp)?;
4190 let actv = act.slice(0..m_e * n_ff_exp);
4191 e.with_moe_cache(max_block, |c, eng| {
4192 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
4193 let slot = c.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
4194 let buf = c.buf(slot);
4195 eng.qmatvec_view(buf, 0..dl.len, &actv, m_e,
4196 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)
4197 })?
4198 } else {
4199 let sg = scratch_g.as_mut().unwrap();
4201 let su = scratch_u.as_mut().unwrap();
4202 let sd = scratch_d.as_mut().unwrap();
4203 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
4204 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
4205 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
4206 let gate = e.qmatvec_view(sg, 0..gl.len, &gv, m_e,
4207 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)?;
4208 let up = e.qmatvec_view(su, 0..ul.len, &gv, m_e,
4209 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)?;
4210 let mut act = e.zeros(m_e * n_ff_exp)?;
4212 Self::ffn_act_scaled(e, cfg, &gate, &up,
4213 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, m_e * n_ff_exp)?;
4214 let actv = act.slice(0..m_e * n_ff_exp);
4215 e.qmatvec_view(sd, 0..dl.len, &actv, m_e,
4216 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)?
4217 };
4218
4219 e.scatter_slot(&y, &tok_idx_d, &slot_idx_d, &weight_d,
4221 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
4222 }
4223
4224 let mut moe_out = e.zeros(t * n_embd)?;
4226 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, t)?;
4227
4228 if std::env::var("MEMRA_MOE_STATS").is_ok() && !m_dist.is_empty() {
4230 m_dist.sort_unstable();
4231 let active = m_dist.len();
4232 let mean = m_dist.iter().sum::<usize>() as f64 / active as f64;
4233 let median = m_dist[active / 2];
4234 let max_m = *m_dist.last().unwrap();
4235 let min_m = m_dist[0];
4236 let above16 = m_dist.iter().filter(|&&x| x >= 16).count();
4237 println!("moe-grouped il={il} t={t} active={active}/{n_expert} \
4238 m_e: min={min_m} median={median} mean={mean:.1} max={max_m} \
4239 above_gemm_threshold(>=16)={above16}/{active}");
4240 }
4241
4242 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4246 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4247 {
4248 let n_ff_sh = gate_shexp.out_features();
4249 let sg_gate = e.matmul(gate_shexp, z, t)?;
4250 let sg_up = e.matmul(up_shexp, z, t)?;
4251 let mut sa = e.zeros(t * n_ff_sh)?;
4252 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
4253 let sh = e.matmul(down_shexp, &sa, t)?;
4254 let g = match &m.gate_inp_shexp {
4258 Some(gate_inp_shexp) => {
4259 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
4262 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
4263 } else {
4264 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
4265 let mut g = e.uninit(t)?;
4266 e.sigmoid(&gs, &mut g, t)?;
4267 g
4268 }
4269 }
4270 None => e.htod(&vec![1.0f32; t])?,
4271 };
4272 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
4273 }
4274
4275 Ok(moe_out)
4276 }
4277
4278 pub(crate) fn moe_ffn_lockstep(
4285 &self,
4286 e: &Engine,
4287 m: &MoeWeights,
4288 zbatch: &CudaSlice<f32>,
4289 mrows: usize,
4290 il: u16,
4291 max_block: usize,
4292 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4293 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
4294 let cfg = &self.cfg;
4295 let moe = cfg.moe.as_ref().unwrap();
4296 let n_embd = cfg.n_embd as usize;
4297 let n_expert = moe.expert_count as usize;
4298 let n_used = moe.expert_used_count as usize;
4299 let n_ff_exp = moe.expert_ff_length as usize;
4300
4301 let logits = e.matmul(&m.gate_inp, zbatch, mrows)?;
4302 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
4303 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
4304 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
4305 } else {
4306 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
4307 None, None, m.active_experts.as_deref())?
4308 };
4309 Self::trace_moe_routes(il, mrows, &sel_all, &w_all)?;
4310
4311 let resident_expert: Vec<bool> = e.with_moe_cache(max_block, |c, _| {
4313 Ok((0..n_expert)
4314 .map(|ex| {
4315 [PROJ_GATE, PROJ_UP, PROJ_DOWN].into_iter().all(|p| {
4316 c.resident(BlockId::new(il, p, ex as u16)).is_some()
4317 })
4318 })
4319 .collect())
4320 })?;
4321
4322 struct Group {
4323 rows: Vec<i32>,
4324 slots: Vec<i32>,
4325 weights: Vec<f32>,
4326 }
4327 let mut groups: std::collections::HashMap<usize, Group> = Default::default();
4328 let mut cpu_rows: Vec<Vec<(usize, f32)>> = vec![Vec::new(); mrows];
4329 let mut cpu_by_expert: std::collections::HashMap<usize, Vec<(usize, f32)>> =
4330 Default::default();
4331 for row in 0..mrows {
4332 for j in 0..n_used {
4333 let ex = sel_all[row * n_used + j] as usize;
4334 let w = w_all[row * n_used + j];
4335 if resident_expert[ex] {
4336 let group = groups.entry(ex).or_insert_with(|| Group {
4337 rows: Vec::new(),
4338 slots: Vec::new(),
4339 weights: Vec::new(),
4340 });
4341 group.rows.push(row as i32);
4342 group.slots.push(j as i32);
4343 group.weights.push(w);
4344 } else {
4345 crate::cpu_experts::record_incomplete_gpu_residency(0);
4346 cpu_rows[row].push((ex, w));
4347 cpu_by_expert.entry(ex).or_default().push((row, w));
4348 }
4349 }
4350 }
4351
4352 let host_rows = e.dtoh(zbatch)?;
4358 let rows_ok = crate::cpu_experts::rows_supported();
4359 enum CpuPart {
4360 Single { row: usize },
4361 Rows { rows: Vec<usize> },
4362 }
4363 let mut tickets: Vec<(CpuPart, crate::cpu_experts::CpuExpertTicket)> = Vec::new();
4364 let mut rows_served: std::collections::HashSet<(usize, usize)> = Default::default();
4365 if rows_ok {
4366 let mut shared: Vec<(usize, Vec<(usize, f32)>)> = cpu_by_expert
4367 .into_iter()
4368 .filter(|(_, rows)| rows.len() >= 2)
4369 .collect();
4370 shared.sort_by_key(|(ex, _)| *ex);
4371 for (ex, mut row_weights) in shared {
4372 row_weights.sort_by_key(|(row, _)| *row);
4373 let inputs: Vec<(&[f32], f32)> = row_weights
4374 .iter()
4375 .map(|&(row, w)| (&host_rows[row * n_embd..(row + 1) * n_embd], w))
4376 .collect();
4377 let job = crate::cpu_experts::prepare_rows_job(m, ex, &inputs)
4378 .map_err(std::io::Error::other)?;
4379 for &(row, _) in &row_weights {
4380 rows_served.insert((row, ex));
4381 }
4382 tickets.push((
4383 CpuPart::Rows {
4384 rows: row_weights.iter().map(|&(row, _)| row).collect(),
4385 },
4386 crate::cpu_experts::submit_rows(job).map_err(std::io::Error::other)?,
4387 ));
4388 }
4389 }
4390 for (row, selected) in cpu_rows.iter().enumerate() {
4391 let leftover: Vec<(usize, f32)> = selected
4392 .iter()
4393 .copied()
4394 .filter(|&(ex, _)| !rows_served.contains(&(row, ex)))
4395 .collect();
4396 if leftover.is_empty() {
4397 continue;
4398 }
4399 let host_row = &host_rows[row * n_embd..(row + 1) * n_embd];
4400 let job = crate::cpu_experts::prepare_job(m, il, &leftover, host_row)
4401 .map_err(std::io::Error::other)?;
4402 tickets.push((
4403 CpuPart::Single { row },
4404 crate::cpu_experts::submit(job).map_err(std::io::Error::other)?,
4405 ));
4406 }
4407
4408 let mut slot_buf = e.zeros(mrows * n_used * n_embd)?;
4409 let mut wbuf = e.zeros(mrows * n_used)?;
4410 let mut order: Vec<usize> = groups.keys().copied().collect();
4411 order.sort_by(|&a, &b| {
4412 groups[&b].rows.len().cmp(&groups[&a].rows.len()).then(a.cmp(&b))
4413 });
4414 for &ex in &order {
4415 let group = &groups[&ex];
4416 let m_e = group.rows.len();
4417 let gl = m.gate_exps.expert_layout(ex);
4418 let ul = m.up_exps.expert_layout(ex);
4419 let dl = m.down_exps.expert_layout(ex);
4420 let row_idx_d = e.htod_i32(&group.rows)?;
4421 let slot_idx_d = e.htod_i32(&group.slots)?;
4422 let dmac = m.down_exps.macro_scale(ex);
4423 let weight_d = if dmac == 1.0 {
4424 e.htod(&group.weights)?
4425 } else {
4426 let scaled: Vec<f32> = group.weights.iter().map(|&w| w * dmac).collect();
4427 e.htod(&scaled)?
4428 };
4429 let mut gathered = e.zeros(m_e * n_embd)?;
4430 e.gather_rows(zbatch, &row_idx_d, &mut gathered, n_embd, m_e)?;
4431 let gv = gathered.slice(0..m_e * n_embd);
4432 let gate = e.with_moe_cache(max_block, |c, eng| {
4433 let slot = c
4434 .resident(BlockId::new(il, PROJ_GATE, ex as u16))
4435 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
4436 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..gl.len, &gv, m_e,
4437 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)
4438 })?;
4439 let up = e.with_moe_cache(max_block, |c, eng| {
4440 let slot = c
4441 .resident(BlockId::new(il, PROJ_UP, ex as u16))
4442 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
4443 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..ul.len, &gv, m_e,
4444 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)
4445 })?;
4446 let mut act = e.zeros(m_e * n_ff_exp)?;
4447 Self::ffn_act_scaled(e, cfg, &gate, &up,
4448 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, m_e * n_ff_exp)?;
4449 let actv = act.slice(0..m_e * n_ff_exp);
4450 let y = e.with_moe_cache(max_block, |c, eng| {
4451 let slot = c
4452 .resident(BlockId::new(il, PROJ_DOWN, ex as u16))
4453 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
4454 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..dl.len, &actv, m_e,
4455 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)
4456 })?;
4457 e.scatter_slot(&y, &row_idx_d, &slot_idx_d, &weight_d,
4458 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
4459 }
4460 let mut moe_out = e.zeros(mrows * n_embd)?;
4461 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, mrows)?;
4462
4463 let mut row_sums: Vec<Option<Vec<f32>>> = vec![None; mrows];
4465 for (part, ticket) in tickets {
4466 let cpu_output = ticket.wait().map_err(std::io::Error::other)?;
4467 let mut add_row = |row: usize, chunk: &[f32]| {
4468 let sum = row_sums[row].get_or_insert_with(|| vec![0.0f32; n_embd]);
4469 for (accumulator, value) in sum.iter_mut().zip(chunk) {
4470 *accumulator += value;
4471 }
4472 };
4473 match part {
4474 CpuPart::Single { row } => add_row(row, &cpu_output),
4475 CpuPart::Rows { rows } => {
4476 for (slot, row) in rows.into_iter().enumerate() {
4477 add_row(row, &cpu_output[slot * n_embd..(slot + 1) * n_embd]);
4478 }
4479 }
4480 }
4481 }
4482 for (row, sum) in row_sums.into_iter().enumerate() {
4483 let Some(sum) = sum else { continue };
4484 let cpu_output = e.htod(&sum)?;
4485 let mut dst = moe_out.slice_mut(row * n_embd..(row + 1) * n_embd);
4486 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
4487 }
4488
4489 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4490 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4491 {
4492 let n_ff_sh = gate_shexp.out_features();
4493 let sg_gate = e.matmul(gate_shexp, zbatch, mrows)?;
4494 let sg_up = e.matmul(up_shexp, zbatch, mrows)?;
4495 let mut sa = e.zeros(mrows * n_ff_sh)?;
4496 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, mrows * n_ff_sh)?;
4497 let sh = e.matmul(down_shexp, &sa, mrows)?;
4498 let g = match &m.gate_inp_shexp {
4501 Some(gate_inp_shexp) => {
4502 e.sigmoid_dot_rows(zbatch, gate_inp_shexp.float_data(), n_embd, mrows)?
4503 }
4504 None => e.htod(&vec![1.0f32; mrows])?,
4505 };
4506 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, mrows)?;
4507 }
4508
4509 Ok(moe_out)
4510 }
4511}
4512
4513impl HybridModel {
4519 pub(crate) fn gemma4_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
4521 let g = self.cfg.gemma4.as_ref().unwrap();
4522 let swa = g.swa_pattern[il];
4523 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
4524 (hd, g.head_count_kv[il] as usize, self.cfg.n_head as usize,
4528 if swa { g.rope_base_swa } else { g.rope_base_global },
4529 1.0, swa)
4530 }
4531
4532 fn gemma4_suppress(&self, e: &Engine, ld: &mut CudaSlice<f32>, t: usize)
4536 -> Result<(), Box<dyn std::error::Error>> {
4537 if let Some((ids, n)) = self.gemma4_aux.as_ref().and_then(|a| a.suppress_d.as_ref()) {
4538 e.mask_ids_rows(ld, ids, *n, self.output.out_features(), t)?;
4539 }
4540 Ok(())
4541 }
4542
4543 fn gemma4_attn_prime(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
4548 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize,
4549 cache: Option<&mut Cache>)
4550 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4551 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
4552 let eps = self.cfg.rms_eps;
4553 let aux = self.gemma4_aux.as_ref().unwrap();
4554
4555 e.mmq_act_begin();
4558 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)? };
4563
4564 let mut q = e.uninit(t * nh * hd)?;
4565 let mut k = e.uninit(t * nkv * hd)?;
4566 let mut v = e.uninit(t * nkv * hd)?;
4568 static EMIT: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4572 let emit = t >= 16 && crate::Engine::qkvnorm_w_on_prefill(nh * t + 2 * nkv * t, hd)
4573 && *EMIT.get_or_init(|| std::env::var("MEMRA_FA_EMIT").map(|s| s != "0").unwrap_or(true));
4574 let mut qb = e.alloc_uninit::<u8>(if emit { t * nh * hd * 2 } else { 1 })?;
4575 let mut kb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
4576 let mut vb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
4577 let v_f16 = emit && crate::fa_f16pv_on() && match hd {
4580 512 => true,
4581 256 => swa && crate::faw_hp_on() && nh % 2 == 0 && (nh / nkv) % 2 == 0,
4582 _ => false,
4583 };
4584 if emit {
4585 e.rms_norm_qkv_w4b(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
4586 &aux.ones, &mut q, &mut k, &mut v, &mut vb,
4587 hd, nh * t, nkv * t, eps, v_f16)?;
4588 } else {
4589 e.rms_norm_qkv(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
4590 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t, eps)?;
4591 }
4592
4593 let ff = if swa { None } else {
4594 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
4595 };
4596 if emit {
4597 e.rope_neox2_bf16e(&mut q, &mut k, &mut qb, &mut kb, pos_d, hd, hd, nh, nkv, t,
4598 base, 1.0, ff)?;
4599 } else {
4600 e.rope_neox2(&mut q, &mut k, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff)?;
4601 }
4602
4603 if let Some(cache) = cache {
4604 let kvl = cache.kv[il].as_mut().unwrap();
4605 assert_eq!(kvl.len, 0, "gemma4 prime is fresh-prompt only (v0)");
4606 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
4607 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()))?;
4608 kvl.len += t;
4609 }
4610 let mut attn = e.zeros(t * nh * hd)?;
4611 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
4615 if swa && t > win {
4616 if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
4617 if emit { e.fa_prefill_w_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
4618 scale, true, win, v_f16)?; }
4619 else { e.fa_prefill_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true,
4620 win)?; }
4621 } else {
4622 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
4623 }
4624 } else if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
4625 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
4626 } else if hd == 512 && std::env::var("MEMRA_NOFA").is_err() {
4627 if emit { e.fa_prefill_hd512_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
4628 scale, true, v_f16)?; }
4629 else { e.fa_prefill_hd512(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?; }
4630 } else {
4631 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
4632 }
4633 Ok(e.matmul(&fa.wo, &attn, t)?)
4634 }
4635
4636 fn gemma4_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
4638 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
4639 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4640 self.gemma4_attn_prime(e, fa, il, h, pos_d, t, None)
4641 }
4642
4643 fn gemma4_moe_q8(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
4648 bits: &crate::hybrid::Gemma4MoeBits,
4649 mq: &(CudaSlice<i8>, CudaSlice<f32>),
4650 router_in: &CudaSlice<f32>, t: usize)
4651 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4652 let cfg = &self.cfg;
4653 let moe = cfg.moe.as_ref().unwrap();
4654 let n_embd = cfg.n_embd as usize;
4655 let n_expert = moe.expert_count as usize;
4656 let n_used = moe.expert_used_count as usize;
4657 let n_ff_exp = moe.expert_ff_length as usize;
4658 let logits = if crate::router_kernel_on() {
4662 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
4663 } else {
4664 e.matmul(&m.gate_inp, router_in, t)?
4665 };
4666 let dev = m.dev_exps.as_ref().unwrap();
4667 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
4668 &bits.per_expert_scale_d)?;
4669 let (zq, zd) = mq;
4670 if t == 1 {
4671 let selv = sel_d.slice(0..n_used);
4672 let wv = w_d.slice(0..n_used);
4673 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, zq, zd,
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 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
4678 let mut moe_out = e.uninit(n_embd)?;
4679 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
4680 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
4681 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
4682 return Ok(moe_out);
4683 }
4684 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
4685 let act = if csr {
4686 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, zq, zd, t * n_used,
4687 n_embd, n_ff_exp, n_used, n_expert,
4688 m.gate_exps.qtype, m.up_exps.qtype,
4689 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4690 } else {
4691 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, zq, zd, t,
4692 n_embd, n_ff_exp, n_used, n_expert,
4693 m.gate_exps.qtype, m.up_exps.qtype,
4694 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4695 };
4696 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
4697 let mut moe_out = e.uninit(t * n_embd)?;
4698 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
4701 n_ff_exp, n_embd, n_used, n_expert,
4702 m.down_exps.qtype, m.down_exps.row_bytes)?;
4703 Ok(moe_out)
4704 }
4705
4706 fn gemma4_moe(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
4710 bits: &crate::hybrid::Gemma4MoeBits, moe_in: &CudaSlice<f32>,
4711 router_in: &CudaSlice<f32>, t: usize)
4712 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4713 let cfg = &self.cfg;
4714 let moe = cfg.moe.as_ref().unwrap();
4715 let n_embd = cfg.n_embd as usize;
4716 let n_expert = moe.expert_count as usize;
4717 let n_used = moe.expert_used_count as usize;
4718 let n_ff_exp = moe.expert_ff_length as usize;
4719
4720 let logits = if t < PRIME_MIN_T && crate::router_kernel_on() {
4724 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
4725 } else {
4726 e.matmul(&m.gate_inp, router_in, t)?
4727 };
4728
4729 if t < PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
4734 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
4735 && expert_dp4a_supported(m.down_exps.qtype)
4736 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0") {
4737 let dev = m.dev_exps.as_ref().unwrap();
4738 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
4739 &bits.per_expert_scale_d)?;
4740 if t == 1 {
4741 let (zq, zd) = e.quantize_q8_1(moe_in, 1, n_embd)?;
4742 let selv = sel_d.slice(0..n_used);
4743 let wv = w_d.slice(0..n_used);
4744 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, &zq, &zd,
4745 n_embd, n_ff_exp, n_used, n_expert,
4746 m.gate_exps.qtype, m.up_exps.qtype,
4747 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
4748 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
4749 let mut moe_out = e.uninit(n_embd)?;
4750 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
4751 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
4752 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
4753 return Ok(moe_out);
4754 }
4755 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
4760 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
4761 let act = if csr {
4762 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, t * n_used,
4763 n_embd, n_ff_exp, n_used, n_expert,
4764 m.gate_exps.qtype, m.up_exps.qtype,
4765 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4766 } else {
4767 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
4768 n_embd, n_ff_exp, n_used, n_expert,
4769 m.gate_exps.qtype, m.up_exps.qtype,
4770 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4771 };
4772 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
4773 let mut moe_out = e.uninit(t * n_embd)?;
4774 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
4775 n_ff_exp, n_embd, n_used, n_expert,
4776 m.down_exps.qtype, m.down_exps.row_bytes)?;
4777 return Ok(moe_out);
4778 }
4779
4780 let (sel_all, mut w_all) = Self::moe_route(e, &logits, t, n_expert, n_used)?;
4781 for (i, &sx) in sel_all.iter().enumerate() {
4782 w_all[i] *= bits.per_expert_scale[sx as usize];
4783 }
4784
4785 if t >= PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
4789 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
4790 && expert_dp4a_supported(m.down_exps.qtype)
4791 && std::env::var("MEMRA_GEMMA_MOE_PAIRS").as_deref() != Ok("0") {
4792 let dev = m.dev_exps.as_ref().unwrap();
4793 let n_pairs = t * n_used;
4794 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
4795 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
4796 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
4797 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
4798 let pt = e.htod_i32(&pair_tok)?;
4799 let pw = e.htod(&w_all)?;
4800 let toff = e.htod_i32(&tok_off)?;
4801 let tids = e.htod_i32(&tok_ids)?;
4802 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
4803 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
4804 let mut ex_ids: Vec<i32> = Vec::new();
4805 let mut ex_off: Vec<i32> = vec![0];
4806 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
4807 for (ex, list) in by_ex.iter().enumerate() {
4808 if list.is_empty() { continue; }
4809 ex_ids.push(ex as i32);
4810 ex_pairs.extend_from_slice(list);
4811 ex_off.push(ex_pairs.len() as i32);
4812 }
4813 let n_active = ex_ids.len();
4814 let exi = e.htod_i32(&ex_ids)?;
4815 let exo = e.htod_i32(&ex_off)?;
4816 let exp_d = e.htod_i32(&ex_pairs)?;
4817 if crate::moe_f16g_gemma_on()
4825 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
4826 && f16g_proj_ok(m.up_exps.qtype, n_embd)
4827 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp) {
4828 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
4829 let csr_tok_d = e.htod_i32(&csr_tok)?;
4830 let (z_f16, z_s) = e.moe_f16g_act(moe_in, Some(&csr_tok_d), n_embd, n_pairs)?;
4831 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
4832 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4833 m.gate_exps.qtype, m.gate_exps.row_bytes)?;
4834 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
4835 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4836 m.up_exps.qtype, m.up_exps.row_bytes)?;
4837 let act_csr = e.moe_pairs_gelu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
4838 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
4839 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
4840 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
4841 m.down_exps.qtype, m.down_exps.row_bytes)?;
4842 let y_down = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
4843 let mut moe_out = e.uninit(t * n_embd)?;
4844 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4845 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
4846 let scan = |v: &[f32]| v.iter().filter(|x| !x.is_finite()).count();
4847 let (yd, mo) = (e.dtoh(&y_down)?, e.dtoh(&moe_out)?);
4848 eprintln!("[f16g-debug] post-permute bad={} post-scatter bad={}",
4849 scan(&yd), scan(&mo));
4850 }
4851 return Ok(moe_out);
4852 }
4853 let mma = n_embd % 256 == 0
4856 && std::env::var("MEMRA_GEMMA_MOE_MMA").as_deref() != Ok("0");
4857 let (gate, up) = if mma {
4858 let z_scr = e.mmq_iq_quantize_act(moe_in, n_embd, t)?;
4859 (e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4860 n_embd, n_ff_exp, n_active, n_pairs, t,
4861 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
4862 e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4863 n_embd, n_ff_exp, n_active, n_pairs, t,
4864 m.up_exps.qtype, m.up_exps.row_bytes)?)
4865 } else {
4866 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
4867 (e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 0, &exi, &exo, &exp_d, &pt, &zq, &zd,
4868 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
4869 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
4870 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 1, &exi, &exo, &exp_d, &pt, &zq, &zd,
4871 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
4872 m.up_exps.qtype, m.up_exps.row_bytes)?)
4873 };
4874 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
4875 let pself = e.htod_i32(&pair_self)?;
4876 let y_down = if mma {
4888 let in_pad = n_ff_exp.div_ceil(256) * 256;
4889 let a_scr = if crate::moe_fuse_actq_on() {
4890 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 1)?
4891 } else {
4892 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4893 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
4894 };
4895 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
4896 in_pad, n_embd, n_active, n_pairs, n_pairs,
4897 m.down_exps.qtype, m.down_exps.row_bytes)?
4898 } else {
4899 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4900 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
4901 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
4902 n_ff_exp, n_embd, n_expert, n_active, n_pairs,
4903 m.down_exps.qtype, m.down_exps.row_bytes)?
4904 };
4905 let mut moe_out = e.uninit(t * n_embd)?;
4906 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4907 return Ok(moe_out);
4908 }
4909
4910 let g_len = m.gate_exps.expert_stride;
4911 let u_len = m.up_exps.expert_stride;
4912 let d_len = m.down_exps.expert_stride;
4913 let dev = m.dev_exps.as_ref().filter(|d| !d.gu_il);
4917 let (mut sg, mut su, mut sd) = if dev.is_some() { (None, None, None) } else {
4918 (Some(e.alloc_u8_uninit(g_len)?), Some(e.alloc_u8_uninit(u_len)?), Some(e.alloc_u8_uninit(d_len)?))
4919 };
4920 let mut moe_out = e.zeros(t * n_embd)?;
4921 for tok in 0..t {
4922 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
4923 let w = &w_all[tok * n_used..(tok + 1) * n_used];
4924 let zt = moe_in.slice(tok * n_embd..(tok + 1) * n_embd);
4925 for (j, &ex) in sel.iter().enumerate() {
4926 let ex = ex as usize;
4927 let gate = match dev {
4928 Some(d) => e.qmatvec_view(&d.gate, ex * g_len..(ex + 1) * g_len, &zt, 1,
4929 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?,
4930 None => {
4931 let sg = sg.as_mut().unwrap();
4932 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
4933 e.qmatvec_view(sg, 0..g_len, &zt, 1,
4934 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?
4935 }
4936 };
4937 let up = match dev {
4938 Some(d) => e.qmatvec_view(&d.up, ex * u_len..(ex + 1) * u_len, &zt, 1,
4939 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?,
4940 None => {
4941 let su = su.as_mut().unwrap();
4942 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
4943 e.qmatvec_view(su, 0..u_len, &zt, 1,
4944 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?
4945 }
4946 };
4947 let mut act = e.uninit(n_ff_exp)?;
4948 e.gelu_tanh_mul(&gate, &up, &mut act, n_ff_exp)?;
4949 let actv = act.slice(0..n_ff_exp);
4950 let y = match dev {
4951 Some(d) => e.qmatvec_view(&d.down, ex * d_len..(ex + 1) * d_len, &actv, 1,
4952 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?,
4953 None => {
4954 let sd = sd.as_mut().unwrap();
4955 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
4956 e.qmatvec_view(sd, 0..d_len, &actv, 1,
4957 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?
4958 }
4959 };
4960 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4961 e.axpy_into(&y, w[j], &mut dst, n_embd)?;
4962 }
4963 }
4964 Ok(moe_out)
4965 }
4966
4967 fn gemma4_layer(&self, e: &Engine, il: usize, layer: &crate::hybrid::HybridLayer,
4969 x: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
4970 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4971 let n_embd = self.cfg.n_embd as usize;
4972 let eps = self.cfg.rms_eps;
4973
4974 let mut h = e.zeros(t * n_embd)?;
4975 e.rms_norm(x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
4976 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
4977 let o = self.gemma4_attn(e, fa, il, &h, pos_d, t)?;
4978 let mut cur = e.zeros(t * n_embd)?;
4980 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
4981 self.gemma4_layer_tail_add(e, layer, &cur, x, t)
4982 }
4983
4984 fn gemma4_layer_tail_add(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
4988 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
4989 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4990 Ok(self.gemma4_layer_tail_add_n(e, layer, cur, x, t, None)?.0)
4991 }
4992
4993 fn gemma4_layer_tail_add_n(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
4996 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
4997 next_norm: Option<&CudaSlice<f32>>)
4998 -> Result<(CudaSlice<f32>, Option<CudaSlice<f32>>), Box<dyn std::error::Error>> {
4999 let n_embd = self.cfg.n_embd as usize;
5000 let bits = layer.gemma4.as_ref().unwrap();
5001 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
5002 let mut xn = e.uninit(t * n_embd)?;
5003 match next_norm {
5004 Some(w) => {
5005 let mut hn = e.uninit(t * n_embd)?;
5006 e.add_scale_rms_norm(&sn, &attn_out, bits.layer_scale, w, &mut xn, &mut hn,
5007 n_embd, t, self.cfg.rms_eps)?;
5008 Ok((xn, Some(hn)))
5009 }
5010 None => {
5011 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
5012 Ok((xn, None))
5013 }
5014 }
5015 }
5016
5017 fn gemma4_layer_tail_core(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5020 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
5021 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5022 self.gemma4_layer_tail_core_pn(e, layer, cur, x, t, None, false)
5023 }
5024
5025 fn gemma4_layer_tail_core_pn(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5032 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
5033 pre_norm: Option<&CudaSlice<f32>>, defer_post_norm: bool)
5034 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5035 let n_embd = self.cfg.n_embd as usize;
5036 let eps = self.cfg.rms_eps;
5037 let bits = layer.gemma4.as_ref().unwrap();
5038
5039 let Some(mbits) = bits.moe_bits.as_ref() else {
5042 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
5043 else { panic!("gemma4 dense layer without Dense ffn") };
5044 let mut attn_out = e.uninit(t * n_embd)?;
5045 let mut zsh = e.uninit(t * n_embd)?;
5046 let mut zpair: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
5049 match pre_norm {
5050 Some(wa) if t == 1 => {
5051 zpair = Some(e.rms_pre_add_rms_norm_q8z(cur, wa, x,
5052 bits.ffn_norm.float_data(),
5053 &mut attn_out, &mut zsh,
5054 n_embd, t, eps)?);
5055 }
5056 Some(wa) => e.rms_pre_add_rms_norm(cur, wa, x, bits.ffn_norm.float_data(),
5057 &mut attn_out, &mut zsh, n_embd, t, eps)?,
5058 None => e.add_rms_norm(cur, x, bits.ffn_norm.float_data(), &mut attn_out,
5059 &mut zsh, n_embd, t, eps)?,
5060 }
5061 let n_ff = ffn_gate.out_features();
5062 let (gate, up) = if t == 1 {
5068 let (zq, zd) = match zpair {
5069 Some(p) => p,
5070 None => e.quantize_q8_1(&zsh, 1, n_embd)?,
5071 };
5072 match e.matmul_q4_fused2(ffn_gate, ffn_up, &zq, &zd)? {
5073 Some(p) => p,
5074 None => (e.matmul_pre(ffn_gate, &zq, &zd, &zsh, 1)?,
5075 e.matmul_pre(ffn_up, &zq, &zd, &zsh, 1)?),
5076 }
5077 } else {
5078 static F2B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5083 let f2b = *F2B.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
5084 let fused = if f2b {
5085 let (zq, zd) = e.quantize_q8_1(&zsh, t, n_embd)?;
5086 e.matmul_q4_fused2_batched(ffn_gate, ffn_up, &zq, &zd, t)?
5087 } else { None };
5088 match fused {
5089 Some(p) => p,
5090 None => {
5091 e.mmq_act_begin();
5093 (e.matmul(ffn_gate, &zsh, t)?, e.matmul(ffn_up, &zsh, t)?)
5094 }
5095 }
5096 };
5097 let mut act = e.uninit(t * n_ff)?;
5098 let f0 = if e.uses_q8_1_fast(ffn_down) {
5101 let upv = e.view(&up, t * n_ff);
5102 let up_all = upv.slice(0..t * n_ff);
5103 let (aq, ad) = e.gelu_tanh_mul_q8_1(&gate, &up_all, &mut act, n_ff, t)?;
5104 e.matmul_pre(ffn_down, &aq, &ad, &act, t)?
5105 } else {
5106 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
5107 e.matmul(ffn_down, &act, t)?
5108 };
5109 if defer_post_norm { return Ok((f0, attn_out)); }
5110 let mut sn = e.uninit(t * n_embd)?;
5111 e.rms_norm(&f0, bits.post_ffw_norm.float_data(), &mut sn, n_embd, t, eps)?;
5112 return Ok((sn, attn_out));
5113 };
5114
5115 assert!(pre_norm.is_none(), "pre-norm fold is dense-entry only");
5116 let mut attn_out = e.uninit(t * n_embd)?;
5121 let mut router_in = e.uninit(t * n_embd)?;
5122 let fast_moe = match &layer.ffn {
5123 crate::hybrid::Ffn::Moe(m) => m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
5124 && expert_dp4a_supported(m.gate_exps.qtype)
5125 && expert_dp4a_supported(m.up_exps.qtype)
5126 && expert_dp4a_supported(m.down_exps.qtype)
5127 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0"),
5128 _ => false,
5129 };
5130 let q8z = t < PRIME_MIN_T && fast_moe;
5131 let (zsh_f32, zsh_q8, moe_q8) = if q8z {
5132 let (z0, m2) = e.add_rms_norm3_q8z(cur, x, bits.ffn_norm.float_data(),
5133 &mbits.router_scale_pre,
5134 mbits.pre_ffw_norm_2.float_data(),
5135 &mut attn_out, &mut router_in, n_embd, t, eps)?;
5136 (None, Some(z0), Some(m2))
5137 } else {
5138 let mut zsh = e.uninit(t * n_embd)?;
5139 let mut moe_in = e.uninit(t * n_embd)?;
5140 e.add_rms_norm3(cur, x, bits.ffn_norm.float_data(), &mbits.router_scale_pre,
5141 mbits.pre_ffw_norm_2.float_data(), &mut attn_out, &mut zsh,
5142 &mut router_in, &mut moe_in, n_embd, t, eps)?;
5143 (Some((zsh, moe_in)), None, None)
5144 };
5145 let attn_out2 = attn_out;
5146 #[allow(unused_variables)]
5147 let attn_out = &attn_out2;
5148 let n_ff = mbits.shared_gate.out_features();
5149 let (gate, up) = if let Some((zq, zd)) = zsh_q8.as_ref() {
5150 if t == 1 {
5151 match e.matmul_q4_fused2(&mbits.shared_gate, &mbits.shared_up, zq, zd)? {
5152 Some(p) => p,
5153 None => {
5154 let h0 = e.zeros(0)?;
5155 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, 1)?,
5156 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, 1)?)
5157 }
5158 }
5159 } else {
5160 let h0 = e.zeros(0)?;
5162 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, t)?,
5163 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, t)?)
5164 }
5165 } else {
5166 let (zsh, _) = zsh_f32.as_ref().unwrap();
5167 (e.matmul(&mbits.shared_gate, zsh, t)?, e.matmul(&mbits.shared_up, zsh, t)?)
5168 };
5169 let mut act = e.uninit(t * n_ff)?;
5170 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
5171 let mlp0 = e.matmul(&mbits.shared_down, &act, t)?;
5172 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else { panic!("gemma4 layer not MoE") };
5173 let moe0 = match (&moe_q8, &zsh_f32) {
5174 (Some(mq), _) => self.gemma4_moe_q8(e, m, mbits, mq, &router_in, t)?,
5175 (None, Some((_, moe_in))) => self.gemma4_moe(e, m, mbits, moe_in, &router_in, t)?,
5176 _ => unreachable!(),
5177 };
5178 let mut mlp = e.uninit(t * n_embd)?;
5180 let mut moe = e.uninit(t * n_embd)?;
5181 e.rms_norm2x(&mlp0, &moe0, mbits.post_ffw_norm_1.float_data(),
5182 mbits.post_ffw_norm_2.float_data(), &mut mlp, &mut moe, n_embd, t, eps)?;
5183
5184 let mut sum = e.uninit(t * n_embd)?;
5187 let mut sn = e.uninit(t * n_embd)?;
5188 e.add_rms_norm(&mlp, &moe, bits.post_ffw_norm.float_data(), &mut sum, &mut sn,
5189 n_embd, t, eps)?;
5190 Ok((sn, attn_out2))
5191 }
5192
5193 fn gemma4_layer_tail_add_nq(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5195 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
5196 next_norm: Option<&CudaSlice<f32>>)
5197 -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>> {
5198 let n_embd = self.cfg.n_embd as usize;
5199 let bits = layer.gemma4.as_ref().unwrap();
5200 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
5201 let mut xn = e.uninit(t * n_embd)?;
5202 match next_norm {
5203 Some(w) => {
5204 let pair = e.add_scale_rms_norm_q8_1(&sn, &attn_out, bits.layer_scale, w, &mut xn,
5205 n_embd, t, self.cfg.rms_eps)?;
5206 Ok((xn, Some(pair)))
5207 }
5208 None => {
5209 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
5210 Ok((xn, None))
5211 }
5212 }
5213 }
5214
5215 fn gemma4_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
5218 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
5219 if self.is_gemma4_e4b() { return self.gemma4_e4b_forward(e, tokens, last_only); }
5222 let n_embd = self.cfg.n_embd as usize;
5223 let t = tokens.len();
5224 let pos: Vec<i32> = (0..t as i32).collect();
5225 let pos_d = e.htod_i32(&pos)?;
5226
5227 let mut x = self.embed(e, tokens)?;
5228 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
5229 let probe = std::env::var("MEMRA_GEMMA_PROBE").is_ok();
5232 let stat = |e: &Engine, x: &CudaSlice<f32>, tag: &str| -> Result<(), Box<dyn std::error::Error>> {
5233 let h = e.dtoh(x)?;
5234 let bad = h.iter().filter(|v| !v.is_finite()).count();
5235 let mx = h.iter().filter(|v| v.is_finite()).fold(0.0f32, |m, v| m.max(v.abs()));
5236 eprintln!("[gemma-probe] {tag}: tok0_first3={:?} bad={bad} max={mx:.3e}", &h[..3]);
5237 Ok(())
5238 };
5239 if probe { stat(e, &x, "embed")?; }
5240 for (il, layer) in self.layers.iter().enumerate() {
5241 x = self.gemma4_layer(e, il, layer, &x, &pos_d, t)?;
5242 if probe { stat(e, &x, &format!("L{il}"))?; }
5243 }
5244 let mut hn = e.zeros(t * n_embd)?;
5245 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, self.cfg.rms_eps)?;
5246 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
5247 let n_vocab = self.output.out_features();
5248 let logits = if last_only {
5249 let hv = e.view(&hn, t * n_embd);
5250 let last_row = hv.slice((t - 1) * n_embd..t * n_embd);
5251 let mut hlast = e.zeros(n_embd)?;
5252 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
5253 let mut ld = e.matmul(&self.output, &hlast, 1)?;
5254 e.softcap(&mut ld, cap, n_vocab)?;
5255 self.gemma4_suppress(e, &mut ld, 1)?;
5256 e.dtoh(&ld)?
5257 } else {
5258 let mut ld = e.matmul(&self.output, &hn, t)?;
5259 e.softcap(&mut ld, cap, t * n_vocab)?;
5260 self.gemma4_suppress(e, &mut ld, t)?;
5261 e.dtoh(&ld)?
5262 };
5263 Ok(logits)
5264 }
5265
5266 pub(crate) fn gemma4_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
5271 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5272 assert_eq!(cache.pos, 0, "gemma4 prime v0 is fresh-prompt only");
5273 let n_embd = self.cfg.n_embd as usize;
5274 let eps = self.cfg.rms_eps;
5275 let t = tokens.len();
5276 let pos: Vec<i32> = (0..t as i32).collect();
5277 let pos_d = e.htod_i32(&pos)?;
5278 let mut x = self.embed(e, tokens)?;
5279 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
5280 for (il, layer) in self.layers.iter().enumerate() {
5281 let mut h = e.zeros(t * n_embd)?;
5282 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
5283 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer not full-attn") };
5284 let o = self.gemma4_attn_prime(e, fa, il, &h, &pos_d, t, Some(cache))?;
5285 let mut cur = e.zeros(t * n_embd)?;
5286 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
5287 x = self.gemma4_layer_tail_add(e, layer, &cur, &x, t)?;
5288 self.dflash_tap(e, cache, il, &x, t)?;
5289 }
5290 cache.pos += t;
5291 let hiddens = e.clone_dtod(&x)?;
5292 let xv = e.view(&x, t * n_embd);
5293 let last_row = xv.slice((t - 1) * n_embd..t * n_embd);
5294 let mut h_seed = e.zeros(n_embd)?;
5295 e.copy_view_into(&mut h_seed, 0, &last_row, n_embd)?;
5296 let mut hn = e.uninit(n_embd)?;
5297 e.rms_norm(&h_seed, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
5298 let mut ld = e.matmul(&self.output, &hn, 1)?;
5299 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
5300 e.softcap(&mut ld, cap, self.output.out_features())?;
5301 self.gemma4_suppress(e, &mut ld, 1)?;
5302 let logits = e.dtoh(&ld)?;
5303 Ok((logits, h_seed, hiddens))
5304 }
5305
5306 fn gemma4_decode_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
5311 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
5312 pos_d: &CudaSlice<i32>, cache: &mut Cache)
5313 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5314 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
5315 let eps = self.cfg.rms_eps;
5316 let aux = self.gemma4_aux.as_ref().unwrap();
5317 let (hq, hdq) = (hq, hdq);
5318 let h0 = e.zeros(0)?;
5319 let h = &h0;
5320 let (q0, k0, v0) = if swa {
5321 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
5322 Some(t3) => t3,
5323 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
5324 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
5325 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?),
5326 }
5327 } else {
5328 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, &hq, &hdq)? {
5329 Some(p) => p,
5330 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
5331 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?),
5332 };
5333 let v0 = e.clone_dtod(&k0)?;
5334 (q0, k0, v0)
5335 };
5336 let mut q = e.uninit(nh * hd)?;
5337 let mut k = e.uninit(nkv * hd)?;
5338 let mut v = e.uninit(nkv * hd)?;
5339 let ff = if swa { None } else {
5342 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
5343 };
5344 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
5345 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
5346 pos_d, nh, nkv, base, 1.0, ff, eps)?;
5347 let kvl = cache.kv[il].as_mut().unwrap();
5348 e.append_kv_quantized(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len,
5349 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()))?;
5350 kvl.len += 1;
5351 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5355 let mut attn = e.uninit(nh * hd)?;
5356 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
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;
5363 e.i32_set_k(&mut kvl.len_d, base)?;
5364 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1, scale,
5365 kvl.k_tok_bytes, kvl.v_tok_bytes, Some((&kvl.len_d, -1)), false,
5366 false, None)?;
5367 return Ok(e.matmul(&fa.wo, &attn, 1)?);
5368 }
5369 if swa && kvl.len > win && hd == 256
5371 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5372 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5373 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5374 let base = kvl.len as i32;
5375 e.i32_set_k(&mut kvl.len_d, base)?;
5376 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1, 1, scale,
5377 win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
5378 return Ok(e.matmul(&fa.wo, &attn, 1)?);
5379 }
5380 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) } else { (0, kvl.len) };
5381 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
5382 (off_tok + t_kv) * kvl.k_tok_bytes);
5383 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
5384 (off_tok + t_kv) * kvl.v_tok_bytes);
5385 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
5386 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
5387 Ok(e.matmul(&fa.wo, &attn, 1)?)
5388 }
5389
5390 #[allow(clippy::too_many_arguments)]
5397 pub fn gemma4_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
5398 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
5399 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
5400 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>)
5401 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
5402 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
5403 self.gemma4_decode_step_dc_into(e, token_d, pos_d, embd_gpu, embd_qt, embd_rb, cache,
5404 n_vocab, cap_bucket_max, &mut tok_out)?;
5405 Ok(tok_out)
5406 }
5407
5408 #[allow(clippy::too_many_arguments)]
5411 pub fn gemma4_decode_step_dc_into(&self, e: &Engine, token_d: &CudaSlice<u32>,
5412 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
5413 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
5414 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
5415 tok_out: &mut CudaSlice<u32>)
5416 -> Result<(), Box<dyn std::error::Error>> {
5417 let n_embd = self.cfg.n_embd as usize;
5418 let eps = self.cfg.rms_eps;
5419 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
5420 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
5421 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
5422 let n_layers = self.layers.len();
5423 for (il, layer) in self.layers.iter().enumerate() {
5424 let (hq, hdq) = match h_carry.take() {
5425 Some(p) => p,
5426 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
5427 };
5428 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
5429 let o = self.gemma4_decode_attn_dc(e, fa, il, &hq, &hdq, pos_d, cache, cap_bucket_max)?;
5430 let mut cur = e.uninit(n_embd)?;
5431 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
5432 let next_norm = if il + 1 < n_layers {
5433 Some(self.layers[il + 1].attn_norm.float_data())
5434 } else { None };
5435 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
5436 x = xn;
5437 h_carry = hn;
5438 }
5439 let mut hn = e.uninit(n_embd)?;
5440 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
5441 let mut logits = e.matmul(&self.output, &hn, 1)?;
5442 self.gemma4_suppress(e, &mut logits, 1)?; e.argmax_token_device_into(&logits, tok_out, n_vocab)?;
5444 e.inc_seqlen(pos_d)?;
5445 if cap_bucket_max.is_none() { cache.pos += 1; }
5446 Ok(())
5447 }
5448
5449 pub fn g4_dc_slots(&self, e: &Engine) -> Result<G4DcSlots, Box<dyn std::error::Error>> {
5456 let n_embd = self.cfg.n_embd as usize;
5457 let n_vocab = self.output.out_features();
5458 let n_layers = self.layers.len();
5459 let (mut qmax, mut kvmax, mut ffmax) = (0usize, 0usize, 0usize);
5460 for il in 0..n_layers {
5461 let (hd, nkv, nh, _b, _s, _w) = self.gemma4_geom(il);
5462 qmax = qmax.max(nh * hd);
5463 kvmax = kvmax.max(nkv * hd);
5464 if let crate::hybrid::Ffn::Dense { ffn_gate, .. } = &self.layers[il].ffn {
5465 ffmax = ffmax.max(ffn_gate.out_features());
5466 }
5467 }
5468 Ok(G4DcSlots {
5469 x: e.uninit(n_embd)?, xn: e.uninit(n_embd)?, cur: e.uninit(n_embd)?,
5470 hq: e.alloc_i8_uninit(n_embd)?, hd_: e.uninit(n_embd / 32)?,
5471 q0: e.uninit(qmax)?, k0: e.uninit(kvmax)?, v0: e.uninit(kvmax)?,
5472 q: e.uninit(qmax)?, k: e.uninit(kvmax)?, v: e.uninit(kvmax)?,
5473 attn: e.uninit(qmax)?, o: e.uninit(n_embd)?,
5474 attn_out: e.uninit(n_embd)?, zsh: e.uninit(n_embd)?,
5475 zq: e.alloc_i8_uninit(n_embd.max(qmax))?, zd: e.uninit(n_embd.max(qmax) / 32)?,
5478 gate: e.uninit(ffmax)?, up: e.uninit(ffmax)?,
5479 act: e.uninit(ffmax)?, actq: e.alloc_i8_uninit(ffmax)?, actd: e.uninit(ffmax / 32)?,
5480 f0: e.uninit(n_embd)?, sn: e.uninit(n_embd)?,
5481 hn: e.uninit(n_embd)?, logits: e.uninit(n_vocab)?,
5482 })
5483 }
5484
5485 fn g4_matvec_m1_into(&self, e: &Engine, w: &crate::model::GpuTensor,
5488 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, y: &mut CudaSlice<f32>)
5489 -> Result<(), Box<dyn std::error::Error>> {
5490 use crate::model::GpuTensor;
5491 let (bytes, qtype, row_bytes, scale, rp) = match w {
5492 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
5493 (bytes, *qtype, *row_bytes, *scale, *rp),
5494 _ => return Err("g4_matvec_m1_into: non-quant tensor".into()),
5495 };
5496 let (mbytes, mrp) = match w {
5497 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5498 _ => (bytes, rp),
5499 };
5500 e.qmatvec_mmvq_into(mbytes, aq, ad, 1, w.in_features(), w.out_features(),
5501 qtype, row_bytes, scale, mrp, y)
5502 }
5503
5504 #[allow(clippy::too_many_arguments)]
5508 pub fn gemma4_decode_step_dc_slotted(&self, e: &Engine, token_d: &CudaSlice<u32>,
5509 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
5510 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
5511 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
5512 sl: &mut G4DcSlots, tok_out: &mut CudaSlice<u32>,
5513 ring: Option<(&mut CudaSlice<u32>, usize)>)
5514 -> Result<(), Box<dyn std::error::Error>> {
5515 let n_embd = self.cfg.n_embd as usize;
5516 let eps = self.cfg.rms_eps;
5517 e.embed_gather_device_into(embd_gpu, token_d, &mut sl.x, n_embd, embd_qt, embd_rb)?;
5518 e.scale_inplace(&mut sl.x, (n_embd as f32).sqrt(), n_embd)?;
5519 let n_layers = self.layers.len();
5520 let mut has_carry = false;
5521 for il in 0..n_layers {
5522 if !has_carry {
5523 e.rms_norm_q8_1_into(&sl.x, self.layers[il].attn_norm.float_data(), n_embd, 1,
5524 eps, &mut sl.hq, &mut sl.hd_)?;
5525 }
5526 has_carry = true;
5527 let layer = &self.layers[il];
5528 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
5529 self.gemma4_decode_attn_dc_slotted(e, fa, il, pos_d, cache, cap_bucket_max, sl)?;
5530 e.rms_norm(&sl.o, layer.post_attn_norm.float_data(), &mut sl.cur, n_embd, 1, eps)?;
5531 let next_norm = if il + 1 < n_layers {
5532 Some(self.layers[il + 1].attn_norm.float_data())
5533 } else { None };
5534 self.gemma4_layer_tail_slotted(e, layer, next_norm, sl)?;
5535 std::mem::swap(&mut sl.x, &mut sl.xn);
5536 }
5537 e.rms_norm(&sl.x, self.output_norm.float_data(), &mut sl.hn, n_embd, 1, eps)?;
5538 e.quantize_q8_1_into(&sl.hn, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
5539 {
5541 let (zq, zd) = (&sl.zq, &sl.zd);
5542 let zq = unsafe { &*(zq as *const CudaSlice<i8>) };
5543 let zd = unsafe { &*(zd as *const CudaSlice<f32>) };
5544 self.g4_matvec_m1_into(e, &self.output, zq, zd, &mut sl.logits)?;
5545 }
5546 self.gemma4_suppress(e, &mut sl.logits, 1)?;
5547 e.argmax_token_device_into(&sl.logits, tok_out, n_vocab)?;
5548 if let Some((ring, base)) = ring {
5549 e.plain_tok_ring(tok_out, pos_d, base, ring)?;
5553 }
5554 e.inc_seqlen(pos_d)?;
5555 if cap_bucket_max.is_none() { cache.pos += 1; }
5556 Ok(())
5557 }
5558
5559 #[allow(clippy::too_many_arguments)]
5561 fn gemma4_decode_attn_dc_slotted(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer,
5562 il: usize, pos_d: &CudaSlice<i32>, cache: &mut Cache,
5563 cap_bucket_max: Option<(usize, usize)>, sl: &mut G4DcSlots)
5564 -> Result<(), Box<dyn std::error::Error>> {
5565 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
5566 let eps = self.cfg.rms_eps;
5567 let aux = self.gemma4_aux.as_ref().unwrap();
5568 {
5569 let hq = unsafe { &*(&sl.hq as *const CudaSlice<i8>) };
5570 let hdq = unsafe { &*(&sl.hd_ as *const CudaSlice<f32>) };
5571 if swa {
5572 if !e.matmul_q4_fused3_into(&fa.wq, &fa.wk, &fa.wv, hq, hdq,
5573 &mut sl.q0, &mut sl.k0, &mut sl.v0)? {
5574 return Err("slotted step: fused3 unavailable (non-uniform trunk)".into());
5575 }
5576 } else {
5577 if !e.matmul_q4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)? {
5578 return Err("slotted step: fused2 unavailable".into());
5579 }
5580 let k0r = unsafe { &*(&sl.k0 as *const CudaSlice<f32>) };
5581 e.copy_into(&mut sl.v0, 0, k0r, nkv * hd)?;
5582 }
5583 }
5584 let ff = if swa { None } else {
5587 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
5588 };
5589 let kvl = cache.kv[il].as_mut().unwrap();
5590 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
5591 if crate::Engine::qkv_append_on() {
5592 e.rms_norm_qkv_rope_append_dc(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(),
5594 fa.k_norm.float_data(), &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
5595 pos_d, nh, nkv, base, 1.0, ff, eps,
5596 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
5597 } else {
5598 e.rms_norm_qkv_rope(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
5599 &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
5600 pos_d, nh, nkv, base, 1.0, ff, eps)?;
5601 e.append_kv_quantized_dc(&sl.k, &sl.v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
5602 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
5603 kv_fp8)?;
5604 }
5605 e.inc_seqlen(&mut kvl.len_d)?;
5606 let (b_swa, b_glob) = cap_bucket_max.expect("slotted step is capture-only");
5607 let k_view = e.view_u8(&kvl.k, kvl.k.len());
5608 let v_view = e.view_u8(&kvl.v, kvl.v.len());
5609 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
5610 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5611 let mut fa_q8 = false;
5615 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
5616 e.fa_decode_rows(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, b_glob - 1,
5617 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5618 Some((&kvl.len_d, -1)), false, false,
5619 Some((&mut sl.zq, &mut sl.zd)))?;
5620 fa_q8 = true;
5621 } else if swa && b_swa > win && hd == 256 && rows_on {
5622 e.fa_decode_rows_w(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv,
5623 &kvl.len_d, -1, 1, scale, win,
5624 kvl.k_tok_bytes, kvl.v_tok_bytes,
5625 Some((&mut sl.zq, &mut sl.zd)))?;
5626 fa_q8 = true;
5627 } else {
5628 let b = if swa { b_swa } else { b_glob };
5629 e.fa_decode_dc(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, &kvl.len_d, b,
5630 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5631 swa && crate::Engine::wkv_on())?;
5632 }
5633 if !fa_q8 {
5634 let aq = unsafe { &*(&sl.attn as *const CudaSlice<f32>) };
5635 e.quantize_q8_1_into(aq, 1, nh * hd, &mut sl.zq, &mut sl.zd)?;
5636 }
5637 {
5638 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
5639 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
5640 self.g4_matvec_m1_into(e, &fa.wo, zq, zd, &mut sl.o)?;
5641 }
5642 Ok(())
5643 }
5644
5645 fn gemma4_layer_tail_slotted(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5648 next_norm: Option<&CudaSlice<f32>>, sl: &mut G4DcSlots)
5649 -> Result<(), Box<dyn std::error::Error>> {
5650 let n_embd = self.cfg.n_embd as usize;
5651 let eps = self.cfg.rms_eps;
5652 let bits = layer.gemma4.as_ref().unwrap();
5653 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
5654 else { return Err("slotted tail: dense ffn only".into()) };
5655 e.add_rms_norm(&sl.cur, &sl.x, bits.ffn_norm.float_data(), &mut sl.attn_out,
5656 &mut sl.zsh, n_embd, 1, eps)?;
5657 let n_ff = ffn_gate.out_features();
5658 {
5659 let zshr = unsafe { &*(&sl.zsh as *const CudaSlice<f32>) };
5660 e.quantize_q8_1_into(zshr, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
5661 }
5662 {
5663 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
5664 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
5665 if !e.matmul_q4_fused2_into(ffn_gate, ffn_up, zq, zd, &mut sl.gate, &mut sl.up)? {
5666 return Err("slotted tail: ffn fused2 unavailable".into());
5667 }
5668 }
5669 debug_assert!(e.uses_q8_1_fast(ffn_down));
5670 {
5671 let upr = unsafe { &*(&sl.up as *const CudaSlice<f32>) };
5672 let upv = e.view(upr, n_ff);
5673 let up_all = upv.slice(0..n_ff);
5674 let gr = unsafe { &*(&sl.gate as *const CudaSlice<f32>) };
5675 e.gelu_tanh_mul_q8_1_into(gr, &up_all, &mut sl.act, n_ff, 1,
5676 &mut sl.actq, &mut sl.actd)?;
5677 }
5678 {
5679 let aq = unsafe { &*(&sl.actq as *const CudaSlice<i8>) };
5680 let ad = unsafe { &*(&sl.actd as *const CudaSlice<f32>) };
5681 self.g4_matvec_m1_into(e, ffn_down, aq, ad, &mut sl.f0)?;
5682 }
5683 e.rms_norm(&sl.f0, bits.post_ffw_norm.float_data(), &mut sl.sn, n_embd, 1, eps)?;
5684 match next_norm {
5685 Some(w) => {
5686 e.add_scale_rms_norm_q8_1_into(&sl.sn, &sl.attn_out, bits.layer_scale, w,
5687 &mut sl.xn, n_embd, 1, eps,
5688 &mut sl.hq, &mut sl.hd_)?;
5689 }
5690 None => {
5691 e.add_scale(&sl.sn, &sl.attn_out, bits.layer_scale, &mut sl.xn, n_embd)?;
5692 }
5693 }
5694 Ok(())
5695 }
5696
5697 #[allow(clippy::too_many_arguments)]
5699 fn gemma4_decode_attn_dc(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
5700 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
5701 pos_d: &CudaSlice<i32>, cache: &mut Cache,
5702 cap_bucket_max: Option<(usize, usize)>)
5703 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5704 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
5705 let eps = self.cfg.rms_eps;
5706 let aux = self.gemma4_aux.as_ref().unwrap();
5707 let (q0, k0, v0) = if swa {
5708 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
5709 Some(t3) => t3,
5710 None => {
5711 let h0 = e.zeros(0)?;
5712 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
5713 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
5714 e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?)
5715 }
5716 }
5717 } else {
5718 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
5719 Some(p) => p,
5720 None => {
5721 let h0 = e.zeros(0)?;
5722 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
5723 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?)
5724 }
5725 };
5726 let v0 = e.clone_dtod(&k0)?;
5727 (q0, k0, v0)
5728 };
5729 let mut q = e.uninit(nh * hd)?;
5730 let mut k = e.uninit(nkv * hd)?;
5731 let mut v = e.uninit(nkv * hd)?;
5732 let ff = if swa { None } else {
5734 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
5735 };
5736 let kvl = cache.kv[il].as_mut().unwrap();
5737 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
5738 if crate::Engine::qkv_append_on() {
5739 e.rms_norm_qkv_rope_append_dc(&q0, &k0, &v0, fa.q_norm.float_data(),
5741 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
5742 pos_d, nh, nkv, base, 1.0, ff, eps,
5743 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
5744 } else {
5745 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
5746 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
5747 pos_d, nh, nkv, base, 1.0, ff, eps)?;
5748 e.append_kv_quantized_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
5749 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
5750 }
5751 e.inc_seqlen(&mut kvl.len_d)?;
5752 let mut attn = e.uninit(nh * hd)?;
5753 let mut fa_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
5756 match cap_bucket_max {
5761 None => {
5762 kvl.len += 1;
5766 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5767 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
5768 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5769 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5772 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5773 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5774 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1,
5775 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5776 Some((&kvl.len_d, -1)), false, false,
5777 Some((&mut aq8, &mut ad8)))?;
5778 fa_q8 = Some((aq8, ad8));
5779 } else if swa && kvl.len > win && hd == 256
5780 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5781 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5783 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5784 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5785 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1,
5786 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes,
5787 Some((&mut aq8, &mut ad8)))?;
5788 fa_q8 = Some((aq8, ad8));
5789 } else {
5790 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) }
5791 else { (0, kvl.len) };
5792 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
5793 (off_tok + t_kv) * kvl.k_tok_bytes);
5794 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
5795 (off_tok + t_kv) * kvl.v_tok_bytes);
5796 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
5797 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
5798 }
5799 }
5800 Some((b_swa, b_glob)) => {
5801 let k_view = e.view_u8(&kvl.k, kvl.k.len());
5807 let v_view = e.view_u8(&kvl.v, kvl.v.len());
5808 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
5809 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5810 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
5811 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5812 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, b_glob - 1,
5813 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5814 Some((&kvl.len_d, -1)), false, false,
5815 Some((&mut aq8, &mut ad8)))?;
5816 fa_q8 = Some((aq8, ad8));
5817 } else if swa && b_swa > win && hd == 256 && rows_on {
5818 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5819 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
5820 &kvl.len_d, -1, 1, scale, win,
5821 kvl.k_tok_bytes, kvl.v_tok_bytes,
5822 Some((&mut aq8, &mut ad8)))?;
5823 fa_q8 = Some((aq8, ad8));
5824 } else {
5825 let b = if swa { b_swa } else { b_glob };
5826 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, b,
5827 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5828 swa && crate::Engine::wkv_on())?;
5829 }
5830 }
5831 }
5832 if let Some((aq8, ad8)) = fa_q8 {
5835 let mut y = e.uninit(fa.wo.out_features())?;
5836 self.g4_matvec_m1_into(e, &fa.wo, &aq8, &ad8, &mut y)?;
5837 return Ok(y);
5838 }
5839 Ok(e.matmul(&fa.wo, &attn, 1)?)
5840 }
5841
5842 pub fn gemma4_generate_graph(&self, e: &Engine, prompt_pos: usize, first_token: u32,
5847 cache: &mut Cache, max_new: usize, eos: &[u32],
5848 mut on_token: impl FnMut(u32) -> bool)
5849 -> Result<(Vec<u32>, crate::decode::StopReason), Box<dyn std::error::Error>> {
5850 if self.is_gemma4_e4b() {
5851 return Err("E4B graph serving is unwired (HANDOVER-E4B.md) — dc-eager is the serving arm".into());
5852 }
5853 use crate::decode::StopReason;
5854 let n_vocab = self.output.out_features();
5855 let n_embd = self.cfg.n_embd as usize;
5856 let embd_gpu = self.embd_gpu.get_or_init(|| {
5857 e.upload_u8(&self.embd.raw).expect("embed table upload")
5858 });
5859 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
5860 for kvl in cache.kv.iter_mut().flatten() {
5861 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
5862 }
5863 let mut token_d = e.stream().clone_htod(&[first_token])?;
5864 let mut pos_d = e.htod_i32(&[prompt_pos as i32])?;
5865 let g4 = self.cfg.gemma4.as_ref().unwrap();
5866 let (hd_s, hd_g) = (g4.key_length_swa as usize, g4.key_length_global as usize);
5867 let nkv_s = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
5869 .find(|p| *p.1).map(|p| *p.0 as usize).unwrap_or(8);
5870 let nkv_g = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
5871 .find(|p| !*p.1).map(|p| *p.0 as usize).unwrap_or(2);
5872 let mut graphs: std::collections::HashMap<((bool, usize), (bool, usize), bool, bool),
5873 (cudarc::driver::CudaGraph,
5874 Vec<Box<dyn std::any::Any + Send>>)> = Default::default();
5875 let mut slots = self.g4_dc_slots(e)?;
5878 const RING: usize = 64;
5881 const DRAIN: usize = 1;
5887 let mut ring = e.stream().alloc_zeros::<u32>(RING)?;
5888 let ring_base = prompt_pos;
5889 let mut out = Vec::with_capacity(max_new);
5890 let mut reason = StopReason::MaxNew;
5891 let mut next = first_token;
5892 let mut captures = 0usize;
5893 for _ in 0..max_new {
5894 out.push(next);
5895 if eos.contains(&next) { reason = StopReason::Eos; break; }
5896 if !on_token(next) { reason = StopReason::Callback; break; }
5897 let t_kv = cache.pos + 1;
5898 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5906 let f512 = crate::fa512_min_tkv();
5907 let key_s = if t_kv > win { (true, usize::MAX) }
5908 else { e.fa_bucket_key(t_kv, hd_s, nkv_s, crate::Engine::wkv_on()) };
5909 let (key_g, rung_end) = if t_kv >= f512 {
5910 let end = (t_kv + 1).next_power_of_two().max(f512 * 2);
5913 ((true, end), end)
5914 } else { (e.fa_bucket_key(t_kv, hd_g, nkv_g, false), t_kv) };
5915 let key = (key_s, key_g, t_kv >= f512, t_kv > win);
5916 if !graphs.contains_key(&key) {
5917 let bucket_max = (t_kv, rung_end);
5918 let snap = cache.snapshot(e)?;
5920 let pos_save = e.dtoh_i32_one(&pos_d)?;
5921 let len_save: Vec<Option<i32>> = cache.kv.iter()
5922 .map(|k| k.as_ref().map(|kvl| e.dtoh_i32_one(&kvl.len_d).unwrap())).collect();
5923 let tok_save = e.dtoh_u32_one(&token_d)?;
5924 let graph = {
5929 let tok_ref = &mut token_d;
5930 let pos_ref = &mut pos_d;
5931 let cache_ref = &mut *cache;
5932 let slots_ref = &mut slots;
5933 let ring_ref = &mut ring;
5934 e.capture_graph_retained_flags(
5935 cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
5936 |e| {
5937 let tok_in = unsafe { &*(tok_ref as *const CudaSlice<u32>) };
5939 let sl = unsafe { &mut *(slots_ref as *mut G4DcSlots) };
5940 let rg = unsafe { &mut *(ring_ref as *mut CudaSlice<u32>) };
5941 self.gemma4_decode_step_dc_slotted(e, tok_in, pos_ref, embd_gpu, qt, rb,
5942 cache_ref, n_vocab, Some(bucket_max),
5943 sl, tok_ref, Some((rg, ring_base)))
5944 })?
5945 };
5946 cache.rollback(e, &snap, 0)?;
5947 e.set_i32_one(&mut pos_d, pos_save)?;
5948 for (il, ls) in len_save.iter().enumerate() {
5949 if let (Some(kvl), Some(v)) = (cache.kv[il].as_mut(), ls) {
5950 e.set_i32_one(&mut kvl.len_d, *v)?;
5951 }
5952 }
5953 e.set_u32_one(&mut token_d, tok_save)?;
5954 if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1") {
5955 if let Ok(c) = crate::graph_update::node_census(&graph.0) {
5956 eprintln!("[graph-census] {c:?}");
5957 }
5958 }
5959 graphs.insert(key, graph);
5960 captures += 1;
5961 }
5962 let mut chunk = 1usize;
5967 let drain_cap: usize = std::env::var("MEMRA_GRAPH_DRAIN").ok()
5968 .and_then(|v| v.parse().ok()).unwrap_or(DRAIN);
5969 while chunk < drain_cap && out.len() + chunk < max_new {
5970 let t_next = cache.pos + 1 + chunk;
5971 let key_s2 = if t_next > win { (true, usize::MAX) }
5972 else { e.fa_bucket_key(t_next, hd_s, nkv_s, crate::Engine::wkv_on()) };
5973 let key_g2 = if t_next >= f512 {
5974 (true, (t_next + 1).next_power_of_two().max(f512 * 2))
5975 } else { e.fa_bucket_key(t_next, hd_g, nkv_g, false) };
5976 if (key_s2, key_g2, t_next >= f512, t_next > win) != key { break; }
5977 chunk += 1;
5978 }
5979 let g = &graphs.get(&key).unwrap().0;
5980 for _ in 0..chunk { g.launch()?; }
5981 e.stream().synchronize()?;
5982 let ringh = e.dtoh_u32(&ring)?;
5983 for j in 0..chunk {
5984 let pos_j = cache.pos + j;
5985 let tok_j = ringh[(pos_j - ring_base) % RING];
5986 cache.pos += 0; if j + 1 == chunk { next = tok_j; }
5988 else {
5989 out.push(tok_j);
5990 if eos.contains(&tok_j) || !on_token(tok_j) {
5991 reason = if eos.contains(&tok_j) { StopReason::Eos }
5992 else { StopReason::Callback };
5993 let keep = cache.pos + j + 1;
5995 e.set_i32_one(&mut pos_d, keep as i32)?;
5996 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
5997 e.set_i32_one(&mut kvl.len_d, keep as i32)?;
5998 kvl.len = keep;
5999 }
6000 cache.pos = keep;
6001 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
6002 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
6003 }
6004 return Ok((out, reason));
6005 }
6006 }
6007 }
6008 cache.pos += chunk;
6009 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) { kvl.len += chunk; }
6010 }
6011 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
6012 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
6013 }
6014 Ok((out, reason))
6015 }
6016
6017 pub(crate) fn gemma4_decode_step_t(&self, e: &Engine, tokens: &[u32], pos0: usize,
6023 cache: &mut Cache)
6024 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
6025 Ok(self.gemma4_decode_step_t_h(e, tokens, pos0, cache)?.0)
6026 }
6027
6028 pub(crate) fn gemma4_decode_step_t_am(&self, e: &Engine, tokens: &[u32], pos0: usize,
6032 cache: &mut Cache)
6033 -> Result<(Vec<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6034 let (ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
6035 let t = tokens.len();
6036 let n_vocab = self.output.out_features();
6037 let mut toks = e.stream().alloc_zeros::<u32>(t)?;
6038 for i in 0..t {
6039 e.argmax_token_device_col(&ld, i, n_vocab, &mut toks, i)?;
6040 }
6041 Ok((e.dtoh_u32(&toks)?, hn))
6042 }
6043
6044 pub(crate) fn gemma4_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
6047 pos0: usize, cache: &mut Cache)
6048 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6049 let (ld, hn) = self.gemma4_verify_trunk(e, &vec![0u32; t], pos0, cache, Some(tok_d))?;
6050 let n_vocab = self.output.out_features();
6051 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
6052 for i in 0..t {
6053 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
6054 }
6055 Ok((vam, hn))
6056 }
6057
6058 pub(crate) fn gemma4_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
6061 cache: &mut Cache)
6062 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6063 let (mut ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
6064 let t = tokens.len();
6065 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6066 e.softcap(&mut ld, cap, t * self.output.out_features())?;
6067 Ok((e.dtoh(&ld)?, hn))
6068 }
6069
6070 pub(crate) fn verify_stream_scratch(&self, e: &Engine, cap: usize)
6073 -> Result<VerifyStreamScratch, Box<dyn std::error::Error>> {
6074 Ok(VerifyStreamScratch {
6075 pos_d: e.htod_i32(&vec![0i32; cap])?,
6076 row_ctrs: (0..cap).map(|_| e.htod_i32(&[0])).collect::<Result<_, _>>()?,
6077 })
6078 }
6079
6080 pub(crate) fn gemma4_verify_t_am_stream(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
6088 ctr: &CudaSlice<i32>, hint: usize,
6089 cache: &mut Cache,
6090 scr: &mut VerifyStreamScratch)
6091 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6092 let n_embd = self.cfg.n_embd as usize;
6093 let eps = self.cfg.rms_eps;
6094 assert!(t <= scr.row_ctrs.len() && t <= 64);
6095 e.i32_iota_from(ctr, &mut scr.pos_d, t)?;
6096 for i in 0..t {
6097 e.i32_copy_add(ctr, &mut scr.row_ctrs[i], (i + 1) as i32)?;
6098 }
6099 let (pos_d, row_ctrs) = (&scr.pos_d, &scr.row_ctrs);
6100 let embd_gpu = self.embd_gpu.get_or_init(|| {
6101 e.upload_u8(&self.embd.raw).expect("embed table upload")
6102 });
6103 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
6104 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
6105 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6106 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6107 let n_layers = self.layers.len();
6108 for (il, layer) in self.layers.iter().enumerate() {
6109 let (hq, hdq) = match h_carry.take() {
6110 Some(p) => p,
6111 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
6112 };
6113 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6114 let o = self.gemma4_verify_attn_stream(e, fa, il, &hq, &hdq, pos_d, t, cache,
6115 hint, row_ctrs)?;
6116 let mut cur = e.uninit(t * n_embd)?;
6117 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
6118 let next_norm = if il + 1 < n_layers {
6119 Some(self.layers[il + 1].attn_norm.float_data())
6120 } else { None };
6121 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
6122 x = xn;
6123 h_carry = hn;
6124 self.dflash_tap(e, cache, il, &x, t)?;
6125 }
6126 let mut hn = e.uninit(t * n_embd)?;
6127 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
6128 let ld = e.matmul(&self.output, &hn, t)?;
6129 let n_vocab = self.output.out_features();
6130 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
6131 for i in 0..t {
6132 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
6133 }
6134 Ok((vam, hn))
6135 }
6136
6137 fn dflash_tap(&self, e: &Engine, cache: &mut Cache, il: usize, x: &CudaSlice<f32>, t: usize)
6144 -> Result<(), Box<dyn std::error::Error>> {
6145 let Some(taps) = cache.dflash_taps.as_mut() else { return Ok(()) };
6146 let Some(slot) = taps.layer_ids.iter().position(|&l| l == il) else { return Ok(()) };
6147 let h = taps.hidden;
6148 let n_taps = taps.layer_ids.len();
6149 debug_assert_eq!(taps.t, t);
6150 let xv = e.view(x, t * h);
6151 for r in 0..t {
6152 let row = xv.slice(r * h..(r + 1) * h);
6153 e.copy_view_into(&mut taps.buf, r * n_taps * h + slot * h, &row, h)?;
6154 }
6155 Ok(())
6156 }
6157
6158 fn gemma4_verify_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
6159 tok_dev: Option<&CudaSlice<u32>>)
6160 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6161 let n_embd = self.cfg.n_embd as usize;
6162 let eps = self.cfg.rms_eps;
6163 let t = tokens.len();
6164 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
6165 let pos_d = e.htod_i32(&pos)?;
6166 let mut x = match tok_dev {
6167 Some(td) => {
6168 let embd_gpu = self.embd_gpu.get_or_init(|| {
6169 e.upload_u8(&self.embd.raw).expect("embed table upload")
6170 });
6171 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
6172 e.embed_gather_device_td(embd_gpu, td, t, n_embd, qt, rb)?
6173 }
6174 None => e.htod(&self.embd.gather(n_embd, tokens))?,
6175 };
6176 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6177 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6178 let n_layers = self.layers.len();
6179 for (il, layer) in self.layers.iter().enumerate() {
6180 let (hq, hdq) = match h_carry.take() {
6181 Some(p) => p,
6182 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
6183 };
6184 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6185 let o = self.gemma4_verify_attn(e, fa, il, &hq, &hdq, &pos_d, t, cache)?;
6186 let mut cur = e.uninit(t * n_embd)?;
6187 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
6188 let next_norm = if il + 1 < n_layers {
6189 Some(self.layers[il + 1].attn_norm.float_data())
6190 } else { None };
6191 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
6192 x = xn;
6193 h_carry = hn;
6194 self.dflash_tap(e, cache, il, &x, t)?;
6195 }
6196 let mut hn = e.uninit(t * n_embd)?;
6197 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
6198 let mut ld = e.matmul(&self.output, &hn, t)?;
6199 self.gemma4_suppress(e, &mut ld, t)?; cache.pos += t;
6201 Ok((ld, hn))
6202 }
6203
6204 #[allow(clippy::too_many_arguments)]
6212 fn gemma4_verify_attn_stream(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6213 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
6214 pos_d: &CudaSlice<i32>, t: usize,
6215 cache: &mut Cache, hint: usize,
6216 row_ctrs: &[CudaSlice<i32>])
6217 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6218 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
6219 let eps = self.cfg.rms_eps;
6220 let aux = self.gemma4_aux.as_ref().unwrap();
6221 let h0 = e.zeros(0)?;
6222 let h = &h0;
6223 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6226 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
6227 let fused_qkv = if f2b {
6228 if swa {
6229 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
6230 .map(|(a, b, c)| (a, b, Some(c)))
6231 } else {
6232 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
6233 .map(|(a, b)| (a, b, None))
6234 }
6235 } else { None };
6236 let (q0, k0, v0) = match fused_qkv {
6237 Some((a, b, cv)) => {
6238 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
6239 (a, b, v)
6240 }
6241 None => {
6242 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
6243 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
6244 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
6245 else { e.clone_dtod(&k0)? };
6246 (q0, k0, v0)
6247 }
6248 };
6249 let mut q = e.uninit(t * nh * hd)?;
6250 let mut k = e.uninit(t * nkv * hd)?;
6251 let mut v = e.uninit(t * nkv * hd)?;
6252 let ff = if swa { None } else {
6255 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
6256 };
6257 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6258 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
6259 pos_d, nh, nkv, base, 1.0, ff, eps)?;
6260 let kvl = cache.kv[il].as_mut().unwrap();
6261 e.append_kv_quantized_rows_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d, t,
6263 kvl.kv_dim_k, kvl.kv_dim_v,
6264 kvl.k_tok_bytes, kvl.v_tok_bytes,
6265 (!swa && crate::Engine::gkv_on())
6266 || (swa && crate::Engine::wkv_on()))?;
6267 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6270 let mut attn = e.uninit(t * nh * hd)?;
6271 let k_view = e.view_u8(&kvl.k, kvl.k.len());
6272 let v_view = e.view_u8(&kvl.v, kvl.v.len());
6273 if swa && hint + 1 >= win {
6276 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6279 &kvl.len_d, 0, t, scale, win,
6280 kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
6281 } else if hd == 512 && hint + t < crate::fa512_min_tkv() {
6282 let bucket = (hint + t + 2).next_power_of_two()
6295 .min(crate::fa512_min_tkv().saturating_sub(1));
6296 let qv = e.view(&q, t * nh * hd);
6297 for i in 0..t {
6298 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
6299 let mut q_one = e.uninit(nh * hd)?;
6300 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
6301 let mut a_one = e.uninit(nh * hd)?;
6302 e.fa_decode_dc(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv,
6303 &row_ctrs[i], bucket, scale,
6304 kvl.k_tok_bytes, kvl.v_tok_bytes, false)?;
6305 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
6306 }
6307 } else if hd == 512 {
6308 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, hint, t, scale,
6311 kvl.k_tok_bytes, kvl.v_tok_bytes,
6312 Some((&kvl.len_d, 0)), false, false, None)?;
6313 } else {
6314 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6316 &kvl.len_d, hint + t, t, scale,
6317 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
6318 swa && crate::Engine::wkv_on())?;
6319 }
6320 Ok(e.matmul(&fa.wo, &attn, t)?)
6321 }
6322
6323 fn gemma4_verify_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6324 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
6325 pos_d: &CudaSlice<i32>, t: usize,
6326 cache: &mut Cache)
6327 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6328 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
6329 let eps = self.cfg.rms_eps;
6330 let aux = self.gemma4_aux.as_ref().unwrap();
6331 let n_embd = self.cfg.n_embd as usize;
6332 let _ = n_embd;
6333
6334 let h0 = e.zeros(0)?;
6335 let h = &h0;
6336 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6339 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
6340 let fused_qkv = if f2b {
6341 if swa {
6342 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
6343 .map(|(a, b, c)| (a, b, Some(c)))
6344 } else {
6345 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
6346 .map(|(a, b)| (a, b, None))
6347 }
6348 } else { None };
6349 let (q0, k0, v0) = match fused_qkv {
6350 Some((a, b, cv)) => {
6351 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
6352 (a, b, v)
6353 }
6354 None => {
6355 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
6356 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
6357 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
6358 else { e.clone_dtod(&k0)? };
6359 (q0, k0, v0)
6360 }
6361 };
6362 let mut q = e.uninit(t * nh * hd)?;
6363 let mut k = e.uninit(t * nkv * hd)?;
6364 let mut v = e.uninit(t * nkv * hd)?;
6365 let ff = if swa { None } else {
6368 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
6369 };
6370 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6371 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
6372 pos_d, nh, nkv, base, 1.0, ff, eps)?;
6373 let kvl = cache.kv[il].as_mut().unwrap();
6374 let base_len = kvl.len;
6375 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, base_len, t,
6376 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()))?;
6377 kvl.len += t;
6378 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6379 let mut attn = e.uninit(t * nh * hd)?;
6380 let rows_ok = (hd == 256 && base_len + 1 >= crate::fa_vec_min_tkv())
6383 || (hd == 512 && !swa && base_len + 1 >= crate::fa512_min_tkv());
6386 if rows_ok && (!swa || base_len + t <= win) {
6387 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
6388 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
6389 if hd == 512 {
6390 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
6392 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, base_len, t,
6393 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
6394 Some((&kvl.len_d, 0)), false,
6395 swa && crate::Engine::wkv_on(), None)?;
6396 } else {
6397 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
6401 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6402 &kvl.len_d, base_len + t, t, scale,
6403 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
6404 swa && crate::Engine::wkv_on())?;
6405 }
6406 return Ok(e.matmul(&fa.wo, &attn, t)?);
6407 }
6408 if hd == 256 && swa && base_len + 1 >= win
6416 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
6417 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
6418 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
6419 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
6420 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, 0,
6421 t, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
6422 return Ok(e.matmul(&fa.wo, &attn, t)?);
6423 }
6424 for i in 0..t {
6425 let avail = base_len + i + 1;
6426 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
6427 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
6428 (off_tok + t_kv) * kvl.k_tok_bytes);
6429 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
6430 (off_tok + t_kv) * kvl.v_tok_bytes);
6431 let qi = e.view(&q, t * nh * hd);
6432 let q_row = qi.slice(i * nh * hd..(i + 1) * nh * hd);
6433 let mut q_one = e.uninit(nh * hd)?;
6434 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
6435 let mut a_one = e.uninit(nh * hd)?;
6436 if swa && avail > win && hd == 256
6440 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
6441 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
6442 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
6443 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
6444 e.fa_decode_rows_w(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, &kvl.len_d, 0,
6445 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
6446 } else if !swa && hd == 512 && avail >= crate::fa512_min_tkv()
6447 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
6448 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
6449 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
6450 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
6451 e.fa_decode_rows(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, avail - 1, 1,
6452 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
6453 Some((&kvl.len_d, 0)), false, false, None)?;
6454 } else {
6455 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
6456 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
6457 }
6458 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
6459 }
6460 Ok(e.matmul(&fa.wo, &attn, t)?)
6461 }
6462
6463 pub(crate) fn gemma4_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
6466 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6467 if let Some(split) = crate::pp::pp2_split(self.layers.len()) {
6472 return self.gemma4_decode_step_h_pp2(e, token, cache, split);
6473 }
6474 if crate::pp::pp_cuts(self.layers.len()).is_some() {
6475 crate::pp::warn_unwired_once("gemma4 eager decode (N>2)");
6476 }
6477 let n_embd = self.cfg.n_embd as usize;
6478 let eps = self.cfg.rms_eps;
6479 let pos_d = e.htod_i32(&[cache.pos as i32])?;
6480 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
6481 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
6482 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6485 let n_layers = self.layers.len();
6486 for (il, layer) in self.layers.iter().enumerate() {
6487 let (hq, hdq) = match h_carry.take() {
6488 Some(p) => p,
6489 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
6490 };
6491 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6492 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, &pos_d, cache)?;
6493 let mut cur = e.uninit(n_embd)?;
6494 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
6495 let next_norm = if il + 1 < n_layers {
6496 Some(self.layers[il + 1].attn_norm.float_data())
6497 } else { None };
6498 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
6499 x = xn;
6500 h_carry = hn;
6501 }
6502 let mut hn = e.uninit(n_embd)?;
6503 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
6504 let h_seed = e.clone_dtod(&x)?;
6505 let mut ld = e.matmul(&self.output, &hn, 1)?;
6506 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6507 e.softcap(&mut ld, cap, self.output.out_features())?; self.gemma4_suppress(e, &mut ld, 1)?;
6509 let logits = e.dtoh(&ld)?;
6510 cache.pos += 1;
6511 Ok((logits, h_seed))
6512 }
6513
6514 fn gemma4_decode_layers(&self, e: &Engine, mut x: CudaSlice<f32>, lo: usize, hi: usize,
6522 pos_d: &CudaSlice<i32>, cache: &mut Cache)
6523 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6524 let n_embd = self.cfg.n_embd as usize;
6525 let eps = self.cfg.rms_eps;
6526 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6527 for il in lo..hi {
6528 let layer = &self.layers[il];
6529 let (hq, hdq) = match h_carry.take() {
6530 Some(p) => p,
6531 None => e.rms_norm_q8_1(&x, self.layers[il].attn_norm.float_data(), n_embd, 1, eps)?,
6533 };
6534 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6535 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, pos_d, cache)?;
6536 let mut cur = e.uninit(n_embd)?;
6537 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
6538 let next_norm = if il + 1 < hi {
6539 Some(self.layers[il + 1].attn_norm.float_data())
6540 } else { None };
6541 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
6542 x = xn;
6543 h_carry = hn;
6544 }
6545 Ok(x)
6546 }
6547
6548 fn gemma4_decode_step_h_pp2(&self, e: &Engine, token: u32, cache: &mut Cache, split: usize)
6555 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6556 if crate::pp::pp2_streams_off() {
6557 return self.gemma4_decode_step_h_pp2_samestream(e, token, cache, split);
6558 }
6559 let rt = crate::pp::Pp2Rt::get(e)?;
6560 let e0 = rt.engine(0, e);
6561 let e1 = rt.engine(1, e);
6562 let n_embd = self.cfg.n_embd as usize;
6563 let eps = self.cfg.rms_eps;
6564
6565 let (pos_d, slot) = {
6567 let _st0 = rt.enter(0);
6568 let pos_d = e0.htod_i32(&[cache.pos as i32])?;
6569 let mut x = e0.htod(&self.embd.gather(n_embd, &[token]))?;
6570 e0.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
6571 let x = self.gemma4_decode_layers(e0, x, 0, split, &pos_d, cache)?;
6572 let slot = rt.tx(0, &x, n_embd)?;
6573 (pos_d, slot)
6574 };
6575
6576 let _st1 = rt.enter(1);
6578 let x = rt.rx(0, slot, n_embd)?;
6579 let x = self.gemma4_decode_layers(e1, x, split, self.layers.len(), &pos_d, cache)?;
6580
6581 let mut hn = e1.uninit(n_embd)?;
6582 e1.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
6583 let h_seed = e1.clone_dtod(&x)?;
6584 let mut ld = e1.matmul(&self.output, &hn, 1)?;
6585 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6586 e1.softcap(&mut ld, cap, self.output.out_features())?;
6587 self.gemma4_suppress(e1, &mut ld, 1)?;
6588 let logits = e1.dtoh(&ld)?;
6589 cache.pos += 1;
6590 Ok((logits, h_seed))
6591 }
6592
6593 fn gemma4_decode_step_h_pp2_samestream(&self, e: &Engine, token: u32, cache: &mut Cache,
6596 split: usize)
6597 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6598 let n_embd = self.cfg.n_embd as usize;
6599 let eps = self.cfg.rms_eps;
6600 let pos_d = e.htod_i32(&[cache.pos as i32])?;
6601
6602 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
6604 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
6605 let x = self.gemma4_decode_layers(e, x, 0, split, &pos_d, cache)?;
6606
6607 let boundary_tx = e.clone_dtod(&x)?;
6609 let boundary_rx = e.clone_dtod(&boundary_tx)?;
6610
6611 let x = self.gemma4_decode_layers(e, boundary_rx, split, self.layers.len(), &pos_d, cache)?;
6613
6614 let mut hn = e.uninit(n_embd)?;
6615 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
6616 let h_seed = e.clone_dtod(&x)?;
6617 let mut ld = e.matmul(&self.output, &hn, 1)?;
6618 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6619 e.softcap(&mut ld, cap, self.output.out_features())?;
6620 self.gemma4_suppress(e, &mut ld, 1)?;
6621 let logits = e.dtoh(&ld)?;
6622 cache.pos += 1;
6623 Ok((logits, h_seed))
6624 }
6625}
6626
6627impl HybridModel {
6636 pub fn is_gemma4_e4b(&self) -> bool {
6637 self.gemma4_aux.as_ref().is_some_and(|a| a.e4b.is_some())
6638 }
6639
6640 fn gemma4_e4b_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
6644 let g = self.cfg.gemma4.as_ref().unwrap();
6645 let swa = g.swa_pattern[il];
6646 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
6647 let Mixer::Full(fa) = &self.layers[il].mixer else { panic!("e4b layer {il} not full-attn") };
6648 let nh = fa.wq.out_features() / hd;
6649 let nkv = fa.wk.out_features() / hd;
6650 (hd, nkv, nh, if swa { g.rope_base_swa } else { g.rope_base_global }, 1.0, swa)
6651 }
6652
6653 fn gemma4_e4b_kv_target(&self, il: usize) -> Option<usize> {
6655 self.layers[il].gemma4.as_ref()
6656 .and_then(|b| b.e4b.as_ref())
6657 .and_then(|e4| e4.kv_share.map(|t| t as usize))
6658 }
6659
6660 fn gemma4_e4b_inp_pl(&self, e: &Engine, tokens: &[u32], x_scaled: &CudaSlice<f32>, t: usize)
6665 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6666 let tok_d = e.stream().clone_htod(&tokens.to_vec())?;
6667 self.gemma4_e4b_inp_pl_dev(e, &tok_d, x_scaled, t)
6668 }
6669
6670 fn gemma4_e4b_inp_pl_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
6672 x_scaled: &CudaSlice<f32>, t: usize)
6673 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6674 let aux = self.gemma4_aux.as_ref().unwrap();
6675 let m = aux.e4b.as_ref().unwrap();
6676 let n_embd = self.cfg.n_embd as usize;
6677 let n_layer = self.layers.len();
6678 let width = m.n_epl * n_layer;
6679 let tbl = m.tok_tbl_gpu.get_or_init(|| {
6680 e.upload_u8(&m.tok_embd_bytes).expect("e4b per-layer token table upload")
6681 });
6682 let mut a = e.embed_gather_device_td(tbl, tok_d, t, width, m.tok_embd_qt,
6683 m.tok_embd_row_bytes)?;
6684 e.scale_inplace(&mut a, (m.n_epl as f32).sqrt(), t * width)?;
6685 let mut p = e.matmul(&m.model_proj, x_scaled, t)?;
6686 e.scale_inplace(&mut p, 1.0 / (n_embd as f32).sqrt(), t * width)?;
6687 let mut pn = e.uninit(t * width)?;
6688 e.rms_norm(&p, m.proj_norm.float_data(), &mut pn, m.n_epl, t * n_layer,
6689 self.cfg.rms_eps)?;
6690 let mut out = e.uninit(t * width)?;
6691 e.add_scale(&a, &pn, 1.0 / 2f32.sqrt(), &mut out, t * width)?;
6692 Ok(out)
6693 }
6694
6695 #[allow(clippy::too_many_arguments)]
6700 fn gemma4_e4b_attn(&self, e: &Engine, il: usize,
6701 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
6702 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
6703 dc_bucket: Option<usize>)
6704 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6705 let (hd, nkv, nh, base, scale, swa) = self.gemma4_e4b_geom(il);
6706 let eps = self.cfg.rms_eps;
6707 let aux = self.gemma4_aux.as_ref().unwrap();
6708 let Mixer::Full(fa) = &self.layers[il].mixer else { unreachable!() };
6709 let h0 = e.zeros(0)?;
6713 let h = &h0;
6714
6715 let ff = if swa { None } else {
6716 Some(aux.rope_freqs.as_ref().expect("e4b global rope needs rope_freqs.weight"))
6717 };
6718 let share = self.gemma4_e4b_kv_target(il);
6719 let mut kv_f32: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
6721 let mut q;
6722 if let Some(_tgt) = share {
6723 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
6724 q = e.uninit(t * nh * hd)?;
6725 let mut kdummy = e.uninit(1)?;
6728 let mut vdummy = e.uninit(1)?;
6729 e.rms_norm_qkv_rope(&q0, &q0, &q0, fa.q_norm.float_data(),
6730 fa.q_norm.float_data(), &aux.ones,
6731 &mut q, &mut kdummy, &mut vdummy, hd, nh * t, 0,
6732 pos_d, nh, 1, base, 1.0, ff, eps)?;
6733 } else {
6734 let e4bits = self.layers[il].gemma4.as_ref().and_then(|g| g.e4b.as_ref());
6738 let cat = e4bits.and_then(|e4| e4.qkv_cat.as_ref());
6739 q = e.uninit(t * nh * hd)?;
6740 let mut k = e.uninit(t * nkv * hd)?;
6741 let mut v = e.uninit(t * nkv * hd)?;
6742 if t == 1 && cat.is_some() {
6743 let qkv0 = e.matmul_pre(cat.unwrap(), hq, hdq, h, 1)?;
6744 e.rms_norm_qkv_rope_cat(&qkv0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6745 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
6746 pos_d, nh, nkv, base, 1.0, ff, eps)?;
6747 } else {
6748 let (q0, k0, v0) = match if t == 1 {
6749 e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)?
6750 } else {
6751 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6754 if *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0")) {
6755 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
6756 } else { None }
6757 } {
6758 Some(triple) => triple,
6759 None => (e.matmul_pre(&fa.wq, hq, hdq, h, t)?,
6760 e.matmul_pre(&fa.wk, hq, hdq, h, t)?,
6761 e.matmul_pre(&fa.wv, hq, hdq, h, t)?), };
6763 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(),
6766 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v,
6767 hd, nh * t, nkv * t, pos_d, nh, nkv, base, 1.0, ff, eps)?;
6768 }
6769 let kvl = cache.kv[il].as_mut().unwrap();
6770 let cls = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
6774 if dc_bucket.is_some() {
6775 debug_assert!(t == 1);
6780 e.append_kv_quantized_row_dc_inc(&k, &v, &mut kvl.k, &mut kvl.v,
6782 &mut kvl.len_d, kvl.kv_dim_k, kvl.kv_dim_v,
6783 kvl.k_tok_bytes, kvl.v_tok_bytes, cls)?;
6784 } else {
6785 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
6786 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes,
6787 kvl.v_tok_bytes, cls)?;
6788 kvl.len += t;
6789 }
6790 kv_f32 = Some((k, v));
6791 }
6792 let kvl_idx = share.unwrap_or(il);
6795 let kvl = cache.kv[kvl_idx].as_ref().unwrap();
6796 let base_len = kvl.len - t; let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6798 let mut attn = e.uninit(t * nh * hd)?;
6799 if t > 1 && base_len == 0 && std::env::var("MEMRA_NOFA").is_err() {
6811 if let Some((kf, vf)) = &kv_f32 {
6812 if hd == 256 && t <= win {
6813 e.fa_prefill(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
6814 return Ok(e.matmul(&fa.wo, &attn, t)?);
6815 }
6816 if hd == 256 && swa && t > win {
6817 e.fa_prefill_w(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true,
6818 win)?;
6819 return Ok(e.matmul(&fa.wo, &attn, t)?);
6820 }
6821 if hd == 512 && !swa {
6822 e.fa_prefill_hd512(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale,
6823 true)?;
6824 return Ok(e.matmul(&fa.wo, &attn, t)?);
6825 }
6826 } else if share.is_some() {
6827 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
6828 let k_view = e.view_u8(&kvl.k, kvl.k.len());
6829 let v_view = e.view_u8(&kvl.v, kvl.v.len());
6830 if hd == 256 && (!swa || t <= win) {
6831 e.fa_prefill_view(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t, t,
6833 scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
6834 return Ok(e.matmul(&fa.wo, &attn, t)?);
6835 }
6836 let kv_dim = nkv * hd;
6839 let mut kf = e.uninit(t * kv_dim)?;
6840 let mut vf = e.uninit(t * kv_dim)?;
6841 e.fa_dequant_kv_view_f32(&k_view, &v_view, &mut kf, &mut vf, kv_dim, kv_dim,
6842 t, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
6843 if hd == 512 {
6844 e.fa_prefill_hd512(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale,
6845 true)?;
6846 } else {
6847 e.fa_prefill_w(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true,
6848 win)?;
6849 }
6850 return Ok(e.matmul(&fa.wo, &attn, t)?);
6851 }
6852 }
6853 if let Some(bucket) = dc_bucket {
6854 assert!(t == 1);
6859 let bucket = if hd == 512 && win <= crate::fa512_min_tkv() {
6865 bucket.min(crate::fa512_min_tkv().saturating_sub(1))
6866 } else { bucket };
6867 let k_view = e.view_u8(&kvl.k, kvl.k.len());
6868 let v_view = e.view_u8(&kvl.v, kvl.v.len());
6869 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
6870 if crate::Engine::wpf_level() >= 1 {
6878 e.prefetch_weight_l2(&fa.wo)?;
6879 }
6880 if e.uses_q8_1_fast(&fa.wo) {
6883 let mut oq = e.alloc_i8_uninit(nh * hd)?;
6884 let mut od = e.zeros(nh * hd / 32)?;
6885 e.fa_decode_dc_q8(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6886 &kvl.len_d, bucket, scale,
6887 kvl.k_tok_bytes, kvl.v_tok_bytes, g,
6888 Some((&mut oq, &mut od)))?;
6889 return Ok(e.matmul_pre(&fa.wo, &oq, &od, &attn, t)?);
6890 }
6891 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6892 &kvl.len_d, bucket, scale,
6893 kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
6894 return Ok(e.matmul(&fa.wo, &attn, t)?);
6895 }
6896 for i in 0..t {
6897 let avail = base_len + i + 1;
6898 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
6899 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
6900 (off_tok + t_kv) * kvl.k_tok_bytes);
6901 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
6902 (off_tok + t_kv) * kvl.v_tok_bytes);
6903 let qv = e.view(&q, t * nh * hd);
6904 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
6905 let mut q_one = e.uninit(nh * hd)?;
6906 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
6907 let mut a_one = e.uninit(nh * hd)?;
6908 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
6912 kvl.k_tok_bytes, kvl.v_tok_bytes,
6913 (!swa && crate::Engine::gkv_on())
6914 || (swa && crate::Engine::wkv_on()))?;
6915 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
6916 }
6917 Ok(e.matmul(&fa.wo, &attn, t)?)
6918 }
6919
6920 fn gemma4_e4b_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
6925 head_last: bool)
6926 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6927 let n_embd = self.cfg.n_embd as usize;
6928 let t = tokens.len();
6929 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
6930 let pos_d = e.htod_i32(&pos)?;
6931 let mut x = e.htod(&self.embd.gather(n_embd, tokens))?;
6932 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6933 let inp_pl = self.gemma4_e4b_inp_pl(e, tokens, &x, t)?;
6934 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, head_last)
6935 }
6936
6937 fn gemma4_e4b_trunk_core(&self, e: &Engine, x_in: CudaSlice<f32>, inp_pl: CudaSlice<f32>,
6941 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
6942 dc_bucket: Option<usize>, cap_logits: bool, head_last: bool)
6943 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6944 let n_embd = self.cfg.n_embd as usize;
6945 let eps = self.cfg.rms_eps;
6946 let n_layer = self.layers.len();
6947 let mut x = x_in;
6948 let aux_e4b = self.gemma4_aux.as_ref().unwrap().e4b.as_ref().unwrap();
6949 let n_epl = aux_e4b.n_epl;
6950
6951 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6957 for il in 0..n_layer {
6958 let layer = &self.layers[il];
6959 let (hq, hdq) = match h_carry.take() {
6960 Some(p) => p,
6961 None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
6962 };
6963 let o = self.gemma4_e4b_attn(e, il, &hq, &hdq, pos_d, t, cache, dc_bucket)?;
6964 let bits = layer.gemma4.as_ref().unwrap();
6967 let e4b = bits.e4b.as_ref().expect("e4b layer bits");
6968 let fuse_exit = e.uses_q8_1_fast(&e4b.inp_gate);
6979 let (sn, attn_out) = self.gemma4_layer_tail_core_pn(
6980 e, layer, &o, &x, t, Some(layer.post_attn_norm.float_data()), fuse_exit)?;
6981 let mut resid = e.uninit(t * n_embd)?;
6982 let g = if fuse_exit {
6988 let (rq, rd) = e.rms_pre_add_q8_1(&sn, bits.post_ffw_norm.float_data(),
6990 &attn_out, &mut resid, n_embd, t,
6991 self.cfg.rms_eps)?;
6992 e.matmul_pre(&e4b.inp_gate, &rq, &rd, &resid, t)?
6993 } else {
6994 e.add(&sn, &attn_out, &mut resid, t * n_embd)?;
6995 e.matmul(&e4b.inp_gate, &resid, t)?
6996 };
6997 let mut act = e.uninit(t * n_epl)?;
6998 let y = if t == 1 && e.uses_q8_1_fast(&e4b.proj) {
6999 let ipv = e.view(&inp_pl, n_epl * n_layer);
7000 let row = ipv.slice(il * n_epl..(il + 1) * n_epl);
7001 let (aq, ad) = e.gelu_tanh_mul_q8_1(&g, &row, &mut act, n_epl, 1)?;
7002 e.matmul_pre(&e4b.proj, &aq, &ad, &act, t)?
7003 } else {
7004 let mut inp_this = e.uninit(t * n_epl)?;
7005 e.copy_rows_strided(&inp_pl, &mut inp_this, n_epl, t, n_epl * n_layer,
7006 il * n_epl)?;
7007 e.gelu_tanh_mul(&g, &inp_this, &mut act, t * n_epl)?;
7008 e.matmul(&e4b.proj, &act, t)?
7009 };
7010 let next_norm = if il + 1 < n_layer {
7013 self.layers[il + 1].attn_norm.float_data()
7014 } else {
7015 self.output_norm.float_data()
7016 };
7017 let mut xn = e.uninit(t * n_embd)?;
7018 let pair = e.rms_pre_add_scale_rms_norm_q8_1(&y, e4b.post_norm.float_data(),
7019 &resid, bits.layer_scale, next_norm,
7020 &mut xn, n_embd, t, eps)?;
7021 h_carry = Some(pair);
7022 x = xn;
7023 }
7024 let (oq, odq) = h_carry.take().unwrap();
7028 let h0 = e.zeros(0)?;
7029 let hm = if head_last { 1 } else { t };
7030 let (hq, hd) = if head_last && t > 1 {
7031 let mut q1 = e.uninit_i8(n_embd)?;
7032 e.dtod_copy_view_i8(&oq.slice((t - 1) * n_embd..t * n_embd), &mut q1)?;
7033 let nb = n_embd / 32;
7034 let mut d1 = e.uninit(nb)?;
7035 e.dtod_copy_view(&odq.slice((t - 1) * nb..t * nb), &mut d1)?;
7036 (q1, d1)
7037 } else {
7038 (oq, odq)
7039 };
7040 let mut ld = e.matmul_pre(&self.output, &hq, &hd, &h0, hm)?;
7041 if cap_logits {
7045 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
7046 e.softcap(&mut ld, cap, hm * self.output.out_features())?;
7047 }
7048 self.gemma4_suppress(e, &mut ld, hm)?; Ok((ld, x))
7050 }
7051
7052 pub fn gemma4_e4b_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
7059 t: usize, pos0: usize, cache: &mut Cache)
7060 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7061 let n_embd = self.cfg.n_embd as usize;
7062 let eps = self.cfg.rms_eps;
7063 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
7064 let pos_d = e.htod_i32(&pos)?;
7065 let embd_gpu = self.embd_gpu.get_or_init(|| {
7066 e.upload_u8(&self.embd.raw).expect("embed table upload")
7067 });
7068 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
7069 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
7070 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
7071 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, tok_d, &x, t)?;
7072 let (ld, xp) = self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true,
7073 false)?;
7074 let n_vocab = self.output.out_features();
7077 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
7078 for i in 0..t {
7079 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
7080 }
7081 let mut hn = e.uninit(t * n_embd)?;
7082 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
7083 cache.pos += t;
7084 Ok((vam, hn))
7085 }
7086
7087 pub(crate) fn gemma4_e4b_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
7090 cache: &mut Cache)
7091 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7092 let n_embd = self.cfg.n_embd as usize;
7093 let eps = self.cfg.rms_eps;
7094 let t = tokens.len();
7095 let (ld, xp) = self.gemma4_e4b_trunk(e, tokens, pos0, cache, false)?;
7096 let mut hn = e.uninit(t * n_embd)?;
7097 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
7098 cache.pos += t;
7099 Ok((e.dtoh(&ld)?, hn))
7100 }
7101
7102 pub fn gemma4_e4b_decode_step_dcg(&self, e: &Engine, token_d: &mut CudaSlice<u32>,
7108 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7109 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7110 n_vocab: usize, bucket: usize)
7111 -> Result<(), Box<dyn std::error::Error>> {
7112 let n_embd = self.cfg.n_embd as usize;
7113 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
7114 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
7115 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
7116 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, Some(bucket),
7117 false, false)?;
7118 e.argmax_token_device_into(&ld, token_d, n_vocab)?;
7119 e.inc_seqlen(pos_d)?;
7120 Ok(())
7121 }
7122
7123 #[allow(clippy::too_many_arguments)]
7131 pub fn gemma4_e4b_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
7132 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7133 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7134 n_vocab: usize)
7135 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
7136 let n_embd = self.cfg.n_embd as usize;
7137 let eps = self.cfg.rms_eps;
7138 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
7139 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
7140 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
7141 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, None, false,
7142 false)?;
7143 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
7144 e.argmax_token_device_into(&ld, &mut tok_out, n_vocab)?;
7145 e.inc_seqlen(pos_d)?;
7146 cache.pos += 1;
7147 let _ = eps;
7148 Ok(tok_out)
7149 }
7150
7151 pub(crate) fn gemma4_e4b_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
7154 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7155 let (ld, x) = self.gemma4_e4b_trunk(e, &[token], cache.pos, cache, false)?;
7156 let logits = e.dtoh(&ld)?;
7157 cache.pos += 1;
7158 Ok((logits, x))
7159 }
7160
7161 pub(crate) fn gemma4_e4b_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
7165 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7166 assert_eq!(cache.pos, 0, "e4b prime is fresh-prompt only (v0)");
7167 let n_embd = self.cfg.n_embd as usize;
7168 let t = tokens.len();
7169 let (ld, x) = self.gemma4_e4b_trunk(e, tokens, 0, cache, true)?;
7170 cache.pos += t;
7171 let last = e.dtoh(&ld)?; let xv = e.view(&x, t * n_embd);
7173 let row = xv.slice((t - 1) * n_embd..t * n_embd);
7174 let mut h_seed = e.uninit(n_embd)?;
7175 e.copy_view_into(&mut h_seed, 0, &row, n_embd)?;
7176 Ok((last, h_seed, x))
7177 }
7178
7179 pub(crate) fn gemma4_e4b_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
7181 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
7182 let mut cache = Cache::new(e, &self.cfg, tokens.len() + 8)?;
7183 let (ld, _x) = self.gemma4_e4b_trunk(e, tokens, 0, &mut cache, last_only)?;
7184 Ok(e.dtoh(&ld)?) }
7186}
7187
7188#[cfg(test)]
7189mod page_prefetch_tests {
7190 use super::{
7191 grouped_worker_prefetch_position, page_prefetch_positions,
7192 page_prefetch_window_from_values, worker_prefetch_positions,
7193 };
7194
7195 #[test]
7196 fn page_prefetch_window_keeps_existing_opt_in_default() {
7197 assert_eq!(page_prefetch_window_from_values(false, None), 0);
7198 assert_eq!(page_prefetch_window_from_values(false, Some("8")), 0);
7199 assert_eq!(page_prefetch_window_from_values(true, None), 1);
7200 assert_eq!(page_prefetch_window_from_values(true, Some("bad")), 1);
7201 assert_eq!(page_prefetch_window_from_values(true, Some("0")), 0);
7202 assert_eq!(page_prefetch_window_from_values(true, Some("8")), 8);
7203 }
7204
7205 #[test]
7206 fn rolling_page_prefetch_advises_each_future_expert_once() {
7207 let advised: Vec<_> = (0..7)
7208 .flat_map(|position| page_prefetch_positions(position, 7, 3))
7209 .collect();
7210 assert_eq!(advised, vec![1, 2, 3, 4, 5, 6]);
7211
7212 let one_ahead: Vec<_> = (0..4)
7213 .flat_map(|position| page_prefetch_positions(position, 4, 1))
7214 .collect();
7215 assert_eq!(one_ahead, vec![1, 2, 3]);
7216 assert!(page_prefetch_positions(0, 4, 0).is_empty());
7217 }
7218
7219 #[test]
7220 fn grouped_worker_prefetch_primes_first_then_each_known_next_once() {
7221 assert_eq!(grouped_worker_prefetch_position(0, None), None);
7222 let positions: Vec<_> = std::iter::once(grouped_worker_prefetch_position(4, None).unwrap())
7223 .chain((0..4).filter_map(|position| {
7224 grouped_worker_prefetch_position(4, Some(position))
7225 }))
7226 .collect();
7227 assert_eq!(positions, vec![0, 1, 2, 3]);
7228 assert_eq!(grouped_worker_prefetch_position(1, Some(0)), None);
7229 }
7230
7231 #[test]
7232 fn rolling_worker_prefetch_primes_current_and_each_future_expert_once() {
7233 let queued: Vec<_> = (0..8)
7234 .flat_map(|position| worker_prefetch_positions(position, 8, 5))
7235 .collect();
7236 assert_eq!(queued, (0..8).collect::<Vec<_>>());
7237
7238 let one_at_a_time: Vec<_> = (0..4)
7239 .flat_map(|position| worker_prefetch_positions(position, 4, 1))
7240 .collect();
7241 assert_eq!(one_at_a_time, vec![0, 1, 2, 3]);
7242 assert!(worker_prefetch_positions(0, 4, 0).is_empty());
7243 }
7244}
7245
7246pub struct G4DcSlots {
7247 x: CudaSlice<f32>, xn: CudaSlice<f32>, cur: CudaSlice<f32>,
7248 hq: CudaSlice<i8>, hd_: CudaSlice<f32>,
7249 q0: CudaSlice<f32>, k0: CudaSlice<f32>, v0: CudaSlice<f32>,
7250 q: CudaSlice<f32>, k: CudaSlice<f32>, v: CudaSlice<f32>,
7251 attn: CudaSlice<f32>, o: CudaSlice<f32>,
7252 attn_out: CudaSlice<f32>, zsh: CudaSlice<f32>,
7253 zq: CudaSlice<i8>, zd: CudaSlice<f32>,
7254 gate: CudaSlice<f32>, up: CudaSlice<f32>,
7255 act: CudaSlice<f32>, actq: CudaSlice<i8>, actd: CudaSlice<f32>,
7256 f0: CudaSlice<f32>, sn: CudaSlice<f32>,
7257 hn: CudaSlice<f32>, logits: CudaSlice<f32>,
7258}
7259