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 pub fn forward(&self, e: &Engine, tokens: &[u32]) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
262 if self.is_gemma4_e4b() { return self.gemma4_e4b_forward(e, tokens, false); }
263 if self.cfg.gemma4.is_some() { return self.gemma4_forward(e, tokens, false); }
264 let cfg = &self.cfg;
265 let n_embd = cfg.n_embd as usize;
266 let t = tokens.len();
267 let eps = cfg.rms_eps;
268 let pos: Vec<i32> = (0..t as i32).collect();
269 let pos_d = e.htod_i32(&pos)?;
270
271 let mut x = self.embed(e, tokens)?; for (il, layer) in self.layers.iter().enumerate() {
274 let mut h = e.uninit(t * n_embd)?;
276 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
277
278 let mixed = match &layer.mixer {
279 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t)?,
280 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
281 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
282 };
283
284 let mut x1 = e.uninit(t * n_embd)?;
286 e.add(&x, &mixed, &mut x1, t * n_embd)?;
287
288 let mut z = e.uninit(t * n_embd)?;
290 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
291 let ffn_out = match &layer.ffn {
292 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
293 let n_ff = ffn_gate.out_features();
294 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
295 let up = g2.pop().unwrap();
296 let gate = g2.pop().unwrap();
297 let mut act = e.uninit(t * n_ff)?;
298 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, t * n_ff)?;
299 e.matmul(ffn_down, &act, t)?
300 }
301 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
302 };
303 let mut x2 = e.uninit(t * n_embd)?;
304 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
305 x = x2;
306 }
307
308 let mut hn = e.uninit(t * n_embd)?;
309 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
310 let logits = e.matmul(&self.output, &hn, t)?;
311 Ok(e.dtoh(&logits)?)
312 }
313
314 pub fn forward_last(&self, e: &Engine, tokens: &[u32]) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
320 if self.cfg.gemma4.is_some() { return self.gemma4_forward(e, tokens, true); }
321 let cfg = &self.cfg;
322 let n_embd = cfg.n_embd as usize;
323 let t = tokens.len();
324 let eps = cfg.rms_eps;
325 let pos: Vec<i32> = (0..t as i32).collect();
326 let pos_d = e.htod_i32(&pos)?;
327
328 let mut x = self.embed(e, tokens)?; let probe = std::env::var("MEMRA_LAYER_PROBE").is_ok();
332 for (il, layer) in self.layers.iter().enumerate() {
333 let mut h = e.uninit(t * n_embd)?;
334 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
335 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} norm ok"); }
336 let mixed = match &layer.mixer {
337 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t)?,
338 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
339 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
340 };
341 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} mixer ok"); }
342 let mut x1 = e.uninit(t * n_embd)?;
343 e.add(&x, &mixed, &mut x1, t * n_embd)?;
344 let mut z = e.uninit(t * n_embd)?;
345 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
346 let ffn_out = match &layer.ffn {
347 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
348 let n_ff = ffn_gate.out_features();
349 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
350 let up = g2.pop().unwrap();
351 let gate = g2.pop().unwrap();
352 let mut act = e.uninit(t * n_ff)?;
353 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, t * n_ff)?;
354 e.matmul(ffn_down, &act, t)?
355 }
356 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
357 };
358 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} ffn ok"); }
359 let mut x2 = e.uninit(t * n_embd)?;
360 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
361 x = x2;
362 }
363 let mut hn = e.uninit(t * n_embd)?;
365 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
366 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)?;
369 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
370 let logits = e.matmul(&self.output, &hlast, 1)?; Ok(e.dtoh(&logits)?)
372 }
373
374 pub fn prime_cache(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
391 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
392 let n_embd = self.cfg.n_embd as usize;
393 let t = tokens.len();
394 assert!(t >= PRIME_MIN_T, "prime_cache needs T >= {PRIME_MIN_T} (caller gates)");
398 assert!(cache.pos + t <= cache.max_ctx, "prime_cache: prompt exceeds cache max_ctx");
399
400 if self.is_gemma4_e4b() {
412 return self.gemma4_e4b_prime(e, tokens, cache);
413 }
414 if self.cfg.gemma4.is_some() {
415 return self.gemma4_prime(e, tokens, cache);
417 }
418 let chunk: usize = std::env::var("MEMRA_PRIME_CHUNK").ok()
419 .and_then(|v| v.parse().ok()).unwrap_or(4096);
420 if chunk == 0 || t <= chunk {
421 return self.prime_chunk(e, tokens, cache);
422 }
423 let mut hiddens = e.uninit(t * n_embd)?;
424 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
425 let mut start = 0usize;
426 while start < t {
427 let mut end = (start + chunk).min(t);
429 if t - end > 0 && t - end < PRIME_MIN_T { end = t; }
430 let (l, hs, x) = self.prime_chunk(e, &tokens[start..end], cache)?;
431 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
432 last = Some((l, hs));
433 start = end;
434 }
435 let (logits, h_seed) = last.unwrap();
436 Ok((logits, h_seed, hiddens))
437 }
438
439 fn gdn_hk(e: &Engine, t: usize, num_v: usize, num_k: usize) -> usize {
446 if Engine::gdn_db_on()
447 && Engine::gdn_chunked_enabled() && t >= 16
448 && e.gdn_mma_enabled(Engine::gdn_chunk_size())
449 && num_k * 2 == num_v
450 {
451 num_k
452 } else {
453 num_v
454 }
455 }
456
457 fn f16out_on(e: &Engine, t: usize) -> bool {
462 crate::f16_ffi::pp_f16_enabled() && t >= 16 && !e.verify_exact_on()
463 && std::env::var("MEMRA_F16OUT").as_deref() != Ok("0")
464 }
465
466 pub fn prime_slabs_get(&self, e: &Engine, t: usize, n_embd: usize, n_ff_max: usize)
469 -> Result<std::sync::MutexGuard<'_, Option<PrimeSlabs>>, Box<dyn std::error::Error>> {
470 let mut g = self.prime_slabs.lock().unwrap();
471 let need_new = match g.as_ref() { None => true, Some(sl) => sl.t_cap < t };
472 if need_new {
473 *g = Some(PrimeSlabs {
474 t_cap: t,
475 h: e.uninit(t * n_embd)?,
476 x1: e.uninit(t * n_embd)?,
477 z: e.uninit(t * n_embd)?,
478 act: e.uninit(t * n_ff_max)?,
479 xa: e.uninit(t * n_embd)?,
480 xb: e.uninit(t * n_embd)?,
481 h16: e.alloc_u8_uninit(t * n_embd * 2)?,
482 z16: e.alloc_u8_uninit(t * n_embd * 2)?,
483 gate: e.uninit(t * n_ff_max)?,
484 up: e.uninit(t * n_ff_max)?,
485 ffn_out: e.uninit(t * n_embd)?,
486 seg_glue: Vec::new(),
487 mixed: e.uninit(t * n_embd)?,
488 seg_mid: Vec::new(),
489 seg_t: 0,
490 });
491 }
492 Ok(g)
493 }
494
495 fn prime_chunk(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
496 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
497 let cfg = &self.cfg;
498 let n_embd = cfg.n_embd as usize;
499 let t = tokens.len();
500 let eps = cfg.rms_eps;
501 let base = cache.pos;
502 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
503 let pos_d = e.htod_i32(&pos)?;
504
505 let x_embed = self.embed(e, tokens)?; let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
510 let n_ff_max = self.layers.iter().map(|l| match &l.ffn {
516 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
517 _ => n_embd,
518 }).max().unwrap_or(n_embd).max(n_embd);
519 let use_slabs = std::env::var("MEMRA_PRIME_SLABS").as_deref() != Ok("0");
520 let mut slab_guard = if use_slabs {
521 Some(self.prime_slabs_get(e, t, n_embd, n_ff_max)?)
522 } else {
523 None
524 };
525 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>);
527 let (mut x_cur, mut x_nxt, sl): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, Option<SlabRefs>);
528 let mut seg: Option<(&mut Vec<Option<cudarc::driver::CudaGraph>>, &mut Vec<Option<cudarc::driver::CudaGraph>>, &mut CudaSlice<f32>, &mut usize)> = None;
529 let mut x_own2;
530 match slab_guard.as_mut() {
531 Some(g) => {
532 let slabs = g.as_mut().unwrap();
533 e.copy_into(&mut slabs.xa, 0, &x_embed, t * n_embd)?;
534 let PrimeSlabs { xa, xb, h, x1, z, act, h16, z16, gate, up, ffn_out, seg_glue, mixed, seg_mid, seg_t, .. } = slabs;
535 x_cur = xa;
536 x_nxt = xb;
537 seg = Some((seg_glue, seg_mid, mixed, seg_t));
538 sl = Some((h, x1, z, act, h16, z16, gate, up, ffn_out));
539 }
540 None => {
541 x_own = x_embed;
542 x_own2 = e.uninit(t * n_embd)?;
543 x_cur = &mut x_own;
544 x_nxt = &mut x_own2;
545 sl = None;
546 }
547 }
548 let mut alloc_h; let mut alloc_x1; let mut alloc_z; let mut alloc_act;
549 let mut alloc_h16; let mut alloc_z16;
550 let mut alloc_gate; let mut alloc_up; let mut alloc_fo;
551 let (h, x1, z, act): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
552 let (h16, z16): (&mut CudaSlice<u8>, &mut CudaSlice<u8>);
553 let (sl_gate, sl_up, sl_fo): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
554 match sl {
555 Some((a, b, c, d, e16, f16b, g, u, fo)) => {
556 h = a; x1 = b; z = c; act = d; h16 = e16; z16 = f16b;
557 sl_gate = g; sl_up = u; sl_fo = fo;
558 }
559 None => {
560 alloc_h = e.uninit(t * n_embd)?;
561 alloc_x1 = e.uninit(t * n_embd)?;
562 alloc_z = e.uninit(t * n_embd)?;
563 alloc_act = e.uninit(t * n_ff_max)?;
564 alloc_h16 = e.alloc_u8_uninit(t * n_embd * 2)?;
565 alloc_z16 = e.alloc_u8_uninit(t * n_embd * 2)?;
566 alloc_gate = e.uninit(t * n_ff_max)?;
567 alloc_up = e.uninit(t * n_ff_max)?;
568 alloc_fo = e.uninit(t * n_embd)?;
569 h = &mut alloc_h; x1 = &mut alloc_x1; z = &mut alloc_z; act = &mut alloc_act;
570 h16 = &mut alloc_h16; z16 = &mut alloc_z16;
571 sl_gate = &mut alloc_gate; sl_up = &mut alloc_up; sl_fo = &mut alloc_fo;
572 }
573 }
574 let n_layers = self.layers.len();
579 let use_seg = f16fuse && seg.is_some()
586 && std::env::var("MEMRA_PRIME_SEG").as_deref() == Ok("1");
587 if let Some((sg, sm, _, st)) = seg.as_mut() {
588 if **st != t {
589 sg.clear();
590 sg.extend((0..n_layers).map(|_| None));
591 sm.clear();
592 sm.extend((0..n_layers).map(|_| None));
593 **st = t;
594 }
595 }
596 {
597 let layer0 = &self.layers[0];
598 if f16fuse {
599 e.rms_norm_f16out(x_cur, layer0.attn_norm.float_data(), h, h16, n_embd, t, eps)?;
600 } else {
601 e.rms_norm(x_cur, layer0.attn_norm.float_data(), h, n_embd, t, eps)?;
602 }
603 }
604 for (il, layer) in self.layers.iter().enumerate() {
605 let hx16 = if f16fuse { Some(&*h16) } else { None };
606 if use_seg {
607 let (pre, pre16, w_out) = match &layer.mixer {
610 Mixer::Full(fa) => {
611 let g3 = match hx16 {
612 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
613 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
614 };
615 let (pre, pre16) = self.full_attn_prime_core_inner(e, fa, g3, &pos_d, t, cache, il)?;
616 (pre, pre16, &fa.wo)
617 }
618 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
619 Mixer::Linear(la) => {
620 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
621 let g4 = match hx16 {
622 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
623 None => e.matmul_group(&ws, h, t)?,
624 };
625 let (pre, pre16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, None)?;
626 (pre, pre16, &la.ssm_out)
627 }
628 };
629 {
630 let (_, sm, mslab, _) = seg.as_mut().unwrap();
631 let pre_n = pre.len() / t;
632 let xh_pre = match pre16 {
633 Some(x) => x,
634 None => e.f16_act(&pre, t * pre_n, pre_n)?,
635 };
636 if !e.try_f16_gemm_pre_into(w_out, &xh_pre, t, mslab)? {
637 let y = e.matmul(w_out, &pre, t)?;
638 e.copy_into(mslab, 0, &y, t * n_embd)?;
639 }
640 if sm[il].is_none() {
641 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
642 let w_post = layer.post_attn_norm.float_data();
643 e.stream().synchronize()?;
644 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
645 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
646 e.add(x_cur, mslab, x1, t * n_embd)?;
647 e.rms_norm_f16out(x1, w_post, z, z16, n_embd, t, eps)?;
648 Ok(())
649 })();
650 let g = e.stream().end_capture(
651 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
652 r?;
653 sm[il] = Some(g?.ok_or("S-mid capture produced no graph")?);
654 }
655 sm[il].as_ref().unwrap().launch()?;
656 }
657 } else {
658 let mixed = match &layer.mixer {
659 Mixer::Full(fa) => self.full_attn_prime(e, fa, h, hx16, &pos_d, t, cache, il)?,
660 Mixer::Linear(la) => self.linear_attn_prime(e, la, h, hx16, t, cache, il)?,
661 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
662 };
663 if f16fuse {
664 e.add_rms_norm_f16out(x_cur, &mixed, layer.post_attn_norm.float_data(),
667 x1, z, z16, n_embd, t, eps)?;
668 } else {
669 e.add(x_cur, &mixed, x1, t * n_embd)?;
670 e.rms_norm(x1, layer.post_attn_norm.float_data(), z, n_embd, t, eps)?;
671 }
672 }
673 let zx16 = if f16fuse { Some(&*z16) } else { None };
674 match &layer.ffn {
675 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
676 let n_ff = ffn_gate.out_features();
677 let mut into_ok = false;
680 if let Some(xh) = zx16 {
681 into_ok = e.try_f16_gemm_pre_into(ffn_gate, xh, t, sl_gate)?
682 && e.try_f16_gemm_pre_into(ffn_up, xh, t, sl_up)?;
683 }
684 if !into_ok {
685 let mut g2 = match zx16 {
686 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], z, xh, t)?,
687 None => e.matmul_group(&[ffn_gate, ffn_up], z, t)?,
688 };
689 let up_y = g2.pop().unwrap();
690 let gate_y = g2.pop().unwrap();
691 e.copy_into(sl_gate, 0, &gate_y, t * n_ff)?;
692 e.copy_into(sl_up, 0, &up_y, t * n_ff)?;
693 }
694 let act16 = if Self::f16out_on(e, t) && self.cfg.m3.is_none() {
697 let mut a16 = e.alloc_u8_uninit(t * n_ff * 2)?;
698 e.silu_mul_f16out(sl_gate, sl_up, act, &mut a16, t * n_ff)?;
699 Some(a16)
700 } else {
701 Self::ffn_act(e, &self.cfg, sl_gate, sl_up, act, t * n_ff)?;
702 None
703 };
704 let xh_act = match act16 {
706 Some(x) => x,
707 None => e.f16_act(act, t * n_ff, n_ff)?,
708 };
709 if !e.try_f16_gemm_pre_into(ffn_down, &xh_act, t, sl_fo)? {
710 let y = e.matmul(ffn_down, &*act, t)?;
711 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
712 }
713 }
714 crate::hybrid::Ffn::Moe(m) => {
715 let y = self.moe_ffn_il(e, m, z, t, il as u16)?;
716 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
717 }
718 }
719 if use_seg && il + 1 < n_layers {
720 let w_next = self.layers[il + 1].attn_norm.float_data();
722 let (sg, _, _, _) = seg.as_mut().unwrap();
723 if sg[il].is_none() {
724 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
725 e.stream().synchronize()?;
726 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
727 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
728 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
729 e.rms_norm_f16out(x_nxt, w_next, h, h16, n_embd, t, eps)?;
730 Ok(())
731 })();
732 let g = e.stream().end_capture(
733 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
734 r?;
735 sg[il] = Some(g?.ok_or("S-glue capture produced no graph")?);
736 }
737 sg[il].as_ref().unwrap().launch()?;
738 } else {
739 if il + 1 < n_layers {
740 let w_next = self.layers[il + 1].attn_norm.float_data();
741 if f16fuse {
742 e.add_rms_norm_f16out(x1, sl_fo, w_next, x_nxt, h, h16, n_embd, t, eps)?;
743 } else {
744 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
745 e.rms_norm(x_nxt, w_next, h, n_embd, t, eps)?;
746 }
747 } else {
748 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
749 }
750 }
751 std::mem::swap(&mut x_cur, &mut x_nxt);
752 }
753 let mut x = e.uninit(t * n_embd)?;
755 e.copy_into(&mut x, 0, x_cur, t * n_embd)?;
756 drop(slab_guard);
757
758 let mut h_seed = e.uninit(n_embd)?;
762 if !crate::spec::spec_hpost() {
763 e.copy_view_into(&mut h_seed, 0, &x.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
764 }
765 let mut hn = e.uninit(t * n_embd)?;
767 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
768 if crate::spec::spec_hpost() {
769 e.copy_view_into(&mut h_seed, 0, &hn.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
770 }
771 let last = e.view(&hn, t * n_embd);
772 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
773 let mut hlast = e.uninit(n_embd)?;
774 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
775 let logits = e.matmul(&self.output, &hlast, 1)?;
776 cache.pos += t;
777 Ok((e.dtoh(&logits)?, h_seed, if crate::spec::spec_hpost() { hn } else { x }))
780 }
781
782 pub fn prime_chunk_captured(&self, e: &Engine, x_in: &CudaSlice<f32>, pos_d: &CudaSlice<i32>,
798 t: usize, cache: &mut Cache,
799 len_d: &CudaSlice<i32>,
800 logits_out: &mut CudaSlice<f32>, h_seed_out: &mut CudaSlice<f32>)
801 -> Result<(), Box<dyn std::error::Error>> {
802 let cfg = &self.cfg;
803 let n_embd = cfg.n_embd as usize;
804 let eps = cfg.rms_eps;
805 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
806 let mut x = e.uninit(t * n_embd)?;
807 e.copy_into(&mut x, 0, x_in, t * n_embd)?;
808 for (il, layer) in self.layers.iter().enumerate() {
809 let mut h = e.uninit(t * n_embd)?;
810 let mut hx16: Option<CudaSlice<u8>> = None;
811 if f16fuse {
812 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
813 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut b16, n_embd, t, eps)?;
814 hx16 = Some(b16);
815 } else {
816 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
817 }
818 let mixed = match &layer.mixer {
819 Mixer::Full(fa) => self.full_attn_prime(e, fa, &h, hx16.as_ref(), pos_d, t, cache, il)?,
820 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
821 Mixer::Linear(la) => {
822 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
823 let g4 = match hx16.as_ref() {
824 Some(xh) => e.matmul_group_xh(&ws, &h, xh, t)?,
825 None => e.matmul_group(&ws, &h, t)?,
826 };
827 self.linear_attn_prime_core_pad(e, la, g4, t, cache, il, Some(len_d))?
828 }
829 };
830 let mut x1 = e.uninit(t * n_embd)?;
831 e.add(&x, &mixed, &mut x1, t * n_embd)?;
832 let mut z = e.uninit(t * n_embd)?;
833 let mut zx16: Option<CudaSlice<u8>> = None;
834 if f16fuse {
835 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
836 e.rms_norm_f16out(&x1, layer.post_attn_norm.float_data(), &mut z, &mut b16, n_embd, t, eps)?;
837 zx16 = Some(b16);
838 } else {
839 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
840 }
841 let ffn_out = match &layer.ffn {
842 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
843 let n_ff = ffn_gate.out_features();
844 let mut g2 = match &zx16 {
845 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], &z, xh, t)?,
846 None => e.matmul_group(&[ffn_gate, ffn_up], &z, t)?,
847 };
848 let up = g2.pop().unwrap();
849 let gate = g2.pop().unwrap();
850 let mut act = e.uninit(t * n_ff)?;
851 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, t * n_ff)?;
852 e.matmul(ffn_down, &act, t)?
853 }
854 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
855 };
856 let mut x2 = e.uninit(t * n_embd)?;
857 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
858 x = x2;
859 }
860 if !crate::spec::spec_hpost() {
862 e.row_gather_dev(&x, h_seed_out, len_d, n_embd)?;
863 }
864 let mut hn = e.uninit(t * n_embd)?;
865 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
866 if crate::spec::spec_hpost() {
867 e.row_gather_dev(&hn, h_seed_out, len_d, n_embd)?;
868 }
869 let mut hlast = e.uninit(n_embd)?;
870 e.row_gather_dev(&hn, &mut hlast, len_d, n_embd)?;
871 let logits = e.matmul(&self.output, &hlast, 1)?;
872 let nv = logits.len();
873 e.copy_into(logits_out, 0, &logits, nv)?;
874 Ok(())
875 }
876
877 pub fn prime_cache_batch(&self, e: &Engine, prompts: &[&[u32]], caches: &mut [&mut Cache])
894 -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
895 let cfg = &self.cfg;
896 let n_embd = cfg.n_embd as usize;
897 let eps = cfg.rms_eps;
898 let b = prompts.len();
899 assert!(b >= 1 && b == caches.len());
900 let pos0s: Vec<usize> = caches.iter().map(|c| c.pos).collect();
901 let carried = pos0s.iter().any(|&p| p > 0);
902 if carried && cfg.gemma4.is_some() {
903 return Err("prime_cache_batch: gemma4 has no continuation prime (v0 fresh-only)".into());
904 }
905 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
906 for &t in &ts { assert!(t >= PRIME_MIN_T, "prime_cache_batch needs T >= {PRIME_MIN_T}"); }
907 for (s, c) in caches.iter().enumerate() {
908 assert!(c.pos + ts[s] <= c.max_ctx, "prime_cache_batch: prompt exceeds cache max_ctx");
909 }
910 let total: usize = ts.iter().sum();
911 let offs: Vec<usize> = ts.iter().scan(0usize, |a, &t| { let o = *a; *a += t; Some(o) }).collect();
912 let pos_ds: Vec<CudaSlice<i32>> = ts.iter().zip(&pos0s)
914 .map(|(&t, &p0)| e.htod_i32(&(p0 as i32..(p0 + t) as i32).collect::<Vec<_>>()))
915 .collect::<Result<_, _>>()?;
916 let split = |e: &Engine, y: &CudaSlice<f32>, dim: usize|
918 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
919 let mut out = Vec::with_capacity(b);
920 for s in 0..b {
921 let mut ys = e.uninit(ts[s] * dim)?;
922 e.copy_view_into(&mut ys, 0, &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim), ts[s] * dim)?;
923 out.push(ys);
924 }
925 Ok(out)
926 };
927
928 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
929 let mut x = self.embed(e, &cat_tokens)?; for (il, layer) in self.layers.iter().enumerate() {
931 let mut h = e.uninit(total * n_embd)?;
932 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
933 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut hx16, n_embd, total, eps)?;
934 let mut mixed = e.uninit(total * n_embd)?;
936 match &layer.mixer {
937 Mixer::Full(fa) => {
938 let g3 = e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], &h, &hx16, total)?;
939 let (n_head, n_head_kv, head_dim) =
945 (self.cfg.n_head as usize, self.cfg.n_head_kv as usize, self.cfg.head_dim_k as usize);
946 let fa_scale = 1.0 / (head_dim as f32).sqrt();
947 let use_favl = !carried
948 && (2..=8).contains(&b)
949 && (head_dim == 256 || head_dim == 128)
950 && self.cfg.attn_out_gate()
951 && std::env::var("MEMRA_NOFA").is_err()
952 && std::env::var("MEMRA_FA_FLOOR").is_err()
953 && std::env::var("MEMRA_FA_PP_W2").as_deref() != Ok("1")
954 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0")
955 && std::env::var("MEMRA_FA_VL").as_deref() != Ok("0");
956 if use_favl {
957 let (qf_w, kf_w, vf_w) =
958 (fa.wq.out_features(), fa.wk.out_features(), fa.wv.out_features());
959 struct APre {
960 q: CudaSlice<f32>, gate: Option<CudaSlice<f32>>,
961 qn: CudaSlice<f32>, kn: CudaSlice<f32>,
962 }
963 let mut aps = Vec::with_capacity(b);
964 for &t in ts.iter().take(b) {
965 aps.push(APre {
966 q: e.uninit(t * n_head * head_dim)?,
967 gate: Some(e.uninit(t * n_head * head_dim)?),
968 qn: e.uninit(t * n_head * head_dim)?,
969 kn: e.uninit(t * n_head_kv * head_dim)?,
970 });
971 }
972 let (kv_dim_k, kv_dim_v, ktb, vtb) = {
973 let kvl = caches[0].kv[il].as_ref().unwrap();
974 (kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes)
975 };
976 let pargs: Vec<crate::AttnPreVl> = (0..b).map(|s| {
977 let (o, t) = (offs[s], ts[s]);
978 let kvl = caches[s].kv[il].as_ref().unwrap();
979 assert!(kvl.len == 0 && kvl.len + t <= caches[s].max_ctx,
980 "prime_cache_batch attn vl: fresh + capacity");
981 crate::AttnPreVl {
982 qf: e.addr_f32v(&g3[0].slice(o * qf_w..(o + t) * qf_w)),
983 kf: e.addr_f32v(&g3[1].slice(o * kf_w..(o + t) * kf_w)),
984 vf: e.addr_f32v(&g3[2].slice(o * vf_w..(o + t) * vf_w)),
985 q: e.addr_f32(&aps[s].q),
986 gate: e.addr_f32(aps[s].gate.as_ref().unwrap()),
987 qn: e.addr_f32(&aps[s].qn), kn: e.addr_f32(&aps[s].kn),
988 kc: e.addr_u8(&kvl.k), vc: e.addr_u8(&kvl.v),
989 t: t as i32, pad: 0,
990 }
991 }).collect();
992 e.attn_pre_vl8(&pargs, fa.q_norm.float_data(), fa.k_norm.float_data(),
993 head_dim, self.cfg.rope_dim_count as usize, n_head, n_head_kv,
994 self.cfg.rms_eps, self.cfg.rope_freq_base, 1.0,
995 kv_dim_k, kv_dim_v, ktb, vtb)?;
996 for s in 0..b {
997 let kvl = caches[s].kv[il].as_mut().unwrap();
998 kvl.len += ts[s];
999 let new_len = kvl.len as i32;
1000 e.set_i32_one(&mut kvl.len_d, new_len)?;
1001 }
1002 let mut attns = Vec::with_capacity(b);
1003 let mut mirrors = Vec::with_capacity(b);
1004 for &t in ts.iter().take(b) {
1005 attns.push(e.uninit(t * n_head * head_dim)?);
1006 let n = t * n_head_kv * head_dim;
1007 mirrors.push((e.alloc_u8_uninit(n * 2)?, e.alloc_u8_uninit(n * 2)?));
1008 }
1009 let fa3_on = match std::env::var("MEMRA_FA3").as_deref() {
1012 Ok("0") => false,
1013 Ok("1") => true,
1014 _ => cfg!(memra_hopper_mma),
1015 };
1016 if fa3_on {
1017 let mut q16s = Vec::with_capacity(b);
1018 let mut v16s = Vec::with_capacity(b);
1019 for s in 0..b {
1020 let t = ts[s];
1021 let mut q16 = e.alloc_u8_uninit(t * n_head * head_dim * 2)?;
1022 e.f32_to_bf16_into(&aps[s].qn, &mut q16, t * n_head * head_dim)?;
1023 let mut k16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
1024 e.f32_to_bf16_into(&aps[s].kn, &mut k16, t * n_head_kv * head_dim)?;
1025 let mut v16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
1026 e.f32_to_bf16_v(&g3[2].slice(offs[s] * vf_w..(offs[s] + t) * vf_w),
1027 &mut v16, t * n_head_kv * head_dim)?;
1028 q16s.push(q16);
1029 v16s.push((k16, v16));
1030 }
1031 let mut qp = [core::ptr::null::<core::ffi::c_void>(); 8];
1032 let mut kp = qp;
1033 let mut vp = qp;
1034 let mut op = [core::ptr::null_mut::<f32>(); 8];
1035 let mut tsv = [0i32; 8];
1036 for s in 0..b {
1037 qp[s] = e.addr_u8(&q16s[s]) as *const core::ffi::c_void;
1038 kp[s] = e.addr_u8(&v16s[s].0) as *const core::ffi::c_void;
1039 vp[s] = e.addr_u8(&v16s[s].1) as *const core::ffi::c_void;
1040 op[s] = e.addr_f32(&attns[s]) as *mut f32;
1041 tsv[s] = ts[s] as i32;
1042 }
1043 let rc = unsafe {
1044 crate::fa3_vl_raw(qp.as_ptr(), kp.as_ptr(), vp.as_ptr(), op.as_ptr(),
1045 tsv.as_ptr(), b as i32, n_head as i32,
1046 n_head_kv as i32, head_dim as i32, fa_scale,
1047 e.stream().cu_stream() as *mut core::ffi::c_void)
1048 };
1049 if rc != 0 {
1050 return Err(format!("memra_fa3_vl rc={rc}").into());
1051 }
1052 } else {
1053 let fargs: Vec<crate::FaSeqVl> = (0..b).map(|s| crate::FaSeqVl {
1054 q: e.addr_f32(&aps[s].qn), k16: e.addr_u8(&mirrors[s].0),
1055 v16: e.addr_u8(&mirrors[s].1), o: e.addr_f32(&attns[s]),
1056 kf: e.addr_f32(&aps[s].kn),
1057 vf: e.addr_f32v(&g3[2].slice(offs[s] * vf_w..(offs[s] + ts[s]) * vf_w)),
1058 t: ts[s] as i32, pad: 0,
1059 }).collect();
1060 e.fa_prefill_vl8(&fargs, head_dim, n_head, n_head_kv, fa_scale)?;
1061 }
1062 for (s, attn) in attns.into_iter().enumerate() {
1063 let (attn_g, ag16) = self.full_attn_prime_post_fa(
1064 e, attn, &aps[s].gate, ts[s], n_head, head_dim)?;
1065 let mut done = false;
1066 if let Some(xh) = &ag16 {
1067 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
1068 }
1069 if !done {
1070 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
1071 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
1072 }
1073 }
1074 } else {
1075 let mut parts: Vec<Vec<CudaSlice<f32>>> = (0..b).map(|_| Vec::new()).collect();
1076 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g3) {
1077 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
1078 parts[s].push(ys);
1079 }
1080 }
1081 for (s, g3s) in parts.into_iter().enumerate() {
1082 let (attn_g, ag16) = self.full_attn_prime_core_inner(
1084 e, fa, g3s, &pos_ds[s], ts[s], caches[s], il)?;
1085 let mut done = false;
1086 if let Some(xh) = &ag16 {
1087 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
1088 }
1089 if !done {
1090 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
1091 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
1092 }
1093 }
1094 }
1095 }
1096 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1097 Mixer::Linear(la) => {
1098 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1103 let g4 = e.matmul_group_xh(&ws, &h, &hx16, total)?;
1104 let outs = self.linear_attn_prime_core_batch(e, la, &g4, &offs, &ts, caches, il)?;
1105 for (s, (gn, gn16)) in outs.into_iter().enumerate() {
1106 let (o, t) = (offs[s], ts[s]);
1107 let mut done = false;
1108 if let Some(xh) = &gn16 {
1109 done = e.try_f16_gemm_pre_into_off(&la.ssm_out, xh, t, &mut mixed, o * n_embd)?;
1110 }
1111 if !done {
1112 let m = e.matmul(&la.ssm_out, &gn, t)?;
1113 e.copy_into(&mut mixed, o * n_embd, &m, t * n_embd)?;
1114 }
1115 }
1116 }
1117 }
1118 let mut x1 = e.uninit(total * n_embd)?;
1119 let mut z = e.uninit(total * n_embd)?;
1120 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
1121 e.add_rms_norm_f16out(&x, &mixed, layer.post_attn_norm.float_data(),
1122 &mut x1, &mut z, &mut zx16, n_embd, total, eps)?;
1123 let ffn_out = match &layer.ffn {
1124 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
1125 let n_ff = ffn_gate.out_features();
1126 let mut g2 = e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?;
1127 let up = g2.pop().unwrap();
1128 let gate = g2.pop().unwrap();
1129 let mut act = e.uninit(total * n_ff)?;
1130 if Self::f16out_on(e, total) && self.cfg.m3.is_none() {
1133 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
1134 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
1135 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
1136 Some(y) => y,
1137 None => e.matmul(ffn_down, &act, total)?,
1138 }
1139 } else {
1140 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, total * n_ff)?;
1141 e.matmul(ffn_down, &act, total)?
1142 }
1143 }
1144 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, total, il as u16)?,
1145 };
1146 let mut x2 = e.uninit(total * n_embd)?;
1147 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
1148 x = x2;
1149 }
1150 let mut hn = e.uninit(total * n_embd)?;
1152 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, total, eps)?;
1153 let mut hcat = e.uninit(b * n_embd)?;
1159 for s in 0..b {
1160 let last0 = (offs[s] + ts[s] - 1) * n_embd;
1161 e.copy_view_into(&mut hcat, s * n_embd, &hn.slice(last0..last0 + n_embd), n_embd)?;
1162 }
1163 let logits_cat = if b >= 2 { e.try_f16_gemm(&self.output, &hcat, b)? } else { None };
1164 let logits_host: Option<Vec<f32>> = match &logits_cat {
1165 Some(lc) => Some(e.dtoh(lc)?),
1166 None => None,
1167 };
1168 let n_vocab = self.output.out_features();
1169 let mut hidden_all = if crate::spec::spec_hpost() {
1170 split(e, &hn, n_embd)?
1171 } else {
1172 split(e, &x, n_embd)?
1173 };
1174 let mut out = Vec::with_capacity(b);
1175 for s in 0..b {
1176 let last0 = (offs[s] + ts[s] - 1) * n_embd;
1177 let mut h_seed = e.uninit(n_embd)?;
1178 if !crate::spec::spec_hpost() {
1179 e.copy_view_into(&mut h_seed, 0, &x.slice(last0..last0 + n_embd), n_embd)?;
1180 } else {
1181 e.copy_view_into(&mut h_seed, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
1182 }
1183 let logits = match &logits_host {
1184 Some(lh) => lh[s * n_vocab..(s + 1) * n_vocab].to_vec(),
1185 None => {
1186 let mut hlast = e.uninit(n_embd)?;
1187 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
1188 e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?
1189 }
1190 };
1191 caches[s].pos += ts[s];
1192 out.push((logits, h_seed, hidden_all.remove(0)));
1193 }
1194 Ok(out)
1195 }
1196
1197 fn full_attn_prime(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>,
1203 hx: Option<&CudaSlice<u8>>,
1204 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1205 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1206 let g3 = match hx {
1211 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
1212 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
1213 };
1214 self.full_attn_prime_core(e, fa, g3, pos_d, t, cache, il)
1215 }
1216
1217 fn full_attn_prime_core(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
1221 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1222 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1223 let (attn_g, ag16) = self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
1224 if let Some(xh) = &ag16 {
1225 if let Some(y) = e.try_f16_gemm_pre(&fa.wo, xh, t)? {
1226 return Ok(y);
1227 }
1228 }
1229 Ok(e.matmul(&fa.wo, &attn_g, t)?)
1230 }
1231
1232 fn full_attn_prime_core_inner(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
1233 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1234 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1235 let cfg = &self.cfg;
1236 let n_head = cfg.n_head as usize;
1237 let n_head_kv = cfg.n_head_kv as usize;
1238 let head_dim = cfg.head_dim_k as usize;
1239 let scale = 1.0 / (head_dim as f32).sqrt();
1240 let (pre, base_len) = self.full_attn_prime_pre_fa(e, fa, g3, pos_d, t, cache, il)?;
1241 let AttnPre { q, k, v, gate } = pre;
1242 let mut attn = e.uninit(t * n_head * head_dim)?;
1243 self.full_attn_prime_fa_dispatch(e, &q, &k, &v, &mut attn, base_len, t, cache, il,
1244 head_dim, n_head, n_head_kv, scale)?;
1245 self.full_attn_prime_post_fa(e, attn, &gate, t, n_head, head_dim)
1246 }
1247
1248 #[allow(clippy::type_complexity)]
1252 fn full_attn_prime_pre_fa(&self, e: &Engine, fa: &FullAttnLayer, mut g3: Vec<CudaSlice<f32>>,
1253 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1254 -> Result<(AttnPre, usize), Box<dyn std::error::Error>> {
1255 let cfg = &self.cfg;
1256 let n_head = cfg.n_head as usize;
1257 let n_head_kv = cfg.n_head_kv as usize;
1258 let head_dim = cfg.head_dim_k as usize;
1259 let eps = cfg.rms_eps;
1260
1261 let gated = cfg.attn_out_gate();
1265 let v = g3.pop().unwrap();
1266 let mut k = g3.pop().unwrap();
1267 let qf = g3.pop().unwrap();
1268 let (mut q, gate) = if gated {
1269 let mut q = e.uninit(t * n_head * head_dim)?;
1270 let mut gate = e.uninit(t * n_head * head_dim)?;
1271 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
1272 (q, Some(gate))
1273 } else {
1274 (qf, None)
1275 };
1276
1277 let mut qn = e.uninit(t * n_head * head_dim)?;
1278 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
1279 q = qn;
1280 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
1281 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
1282 k = kn;
1283 let rope_dims = cfg.rope_dim_count as usize;
1284 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, cfg.rope_freq_base, 1.0)?;
1285 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, cfg.rope_freq_base, 1.0)?;
1286
1287 {
1290 let kvl = cache.kv[il].as_mut().unwrap();
1291 assert!(kvl.len + t <= cache.max_ctx, "prime_cache: KV overflow");
1292 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
1293 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
1294 crate::Engine::kv_fp8_on())?;
1295 kvl.len += t;
1296 let new_len = kvl.len as i32;
1297 e.set_i32_one(&mut kvl.len_d, new_len)?;
1298 }
1299
1300 let base_len = {
1301 let kvl = cache.kv[il].as_ref().unwrap();
1302 kvl.len - t };
1304 Ok((AttnPre { q, k, v, gate }, base_len))
1305 }
1306
1307 #[allow(clippy::too_many_arguments)]
1314 fn full_attn_prime_fa_dispatch(&self, e: &Engine, q: &CudaSlice<f32>, k: &CudaSlice<f32>,
1315 v: &CudaSlice<f32>, attn: &mut CudaSlice<f32>, base_len: usize,
1316 t: usize, cache: &mut Cache, il: usize,
1317 head_dim: usize, n_head: usize, n_head_kv: usize, scale: f32)
1318 -> Result<(), Box<dyn std::error::Error>> {
1319 if base_len == 0 {
1320 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
1324 e.sdpa_naive(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1325 } else {
1326 e.fa_prefill(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1327 }
1328 } else {
1329 let kvl = cache.kv[il].as_ref().unwrap();
1330 let t_kv = base_len + t;
1331 let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
1332 let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
1333 let deqw = std::env::var("MEMRA_PRIME_DEQW").map(|v| v != "0").unwrap_or(true);
1341 if deqw {
1342 e.fa_prefill_view_ws(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
1343 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
1344 crate::Engine::kv_fp8_on())?;
1345 } else {
1346 e.fa_prefill_view(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
1347 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
1348 crate::Engine::kv_fp8_on())?;
1349 }
1350 }
1351 Ok(())
1352 }
1353
1354 fn full_attn_prime_post_fa(&self, e: &Engine, attn: CudaSlice<f32>,
1357 gate: &Option<CudaSlice<f32>>, t: usize,
1358 n_head: usize, head_dim: usize)
1359 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1360 let (attn_g, ag16) = match gate {
1361 Some(gate) => {
1362 let n = t * n_head * head_dim;
1363 let mut ag = e.uninit(n)?;
1364 if Self::f16out_on(e, t) {
1365 let mut a16 = e.alloc_u8_uninit(n * 2)?;
1366 e.sig_mul_f16out(&attn, gate, &mut ag, &mut a16, n)?;
1367 (ag, Some(a16))
1368 } else {
1369 let mut gsig = e.uninit(n)?;
1370 e.sigmoid(gate, &mut gsig, n)?;
1371 e.mul(&attn, &gsig, &mut ag, n)?;
1372 (ag, None)
1373 }
1374 }
1375 None => (attn, None),
1376 };
1377 Ok((attn_g, ag16))
1378 }
1379
1380 fn linear_attn_prime(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>,
1387 hx: Option<&CudaSlice<u8>>, t: usize,
1388 cache: &mut Cache, il: usize)
1389 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1390 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1392 let g4 = match hx {
1393 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
1394 None => e.matmul_group(&ws, h, t)?,
1395 };
1396 self.linear_attn_prime_core(e, la, g4, t, cache, il)
1397 }
1398
1399 fn linear_attn_prime_core(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
1401 t: usize, cache: &mut Cache, il: usize)
1402 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1403 self.linear_attn_prime_core_pad(e, la, g4.drain(..).collect(), t, cache, il, None)
1404 }
1405
1406 #[allow(clippy::too_many_arguments)]
1410 fn linear_attn_prime_core_pad_inner(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
1411 t: usize, cache: &mut Cache, il: usize,
1412 pad_len: Option<&CudaSlice<i32>>)
1413 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1414 let ssm = self.cfg.ssm.as_ref().unwrap();
1416 let d_state = ssm.state_size as usize;
1417 let num_k = ssm.group_count as usize;
1418 let num_v = ssm.time_step_rank as usize;
1419 let key_dim = d_state * num_k;
1420 let value_dim = d_state * num_v;
1421 let conv_dim = key_dim * 2 + value_dim;
1422 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(
1427 e, la,
1428 &qkv_mixed.slice(0..t * conv_dim), &z.slice(0..t * value_dim),
1429 &beta_raw.slice(0..t * num_v), &alpha.slice(0..t * num_v),
1430 t, cache, il, pad_len)
1431 }
1432
1433 #[allow(clippy::too_many_arguments)]
1436 fn linear_attn_gdn_prep(&self, e: &Engine, la: &LinearAttnLayer,
1437 qkv_mixed: &cudarc::driver::CudaView<f32>,
1438 beta_raw: &cudarc::driver::CudaView<f32>,
1439 alpha: &cudarc::driver::CudaView<f32>,
1440 t: usize, cache: &mut Cache, il: usize,
1441 pad_len: Option<&CudaSlice<i32>>)
1442 -> Result<GdnPrep, Box<dyn std::error::Error>> {
1443 let cfg = &self.cfg;
1444 let ssm = cfg.ssm.as_ref().unwrap();
1445 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;
1453 debug_assert!(t >= d_conv - 1, "stateful conv needs T >= pad (PRIME_MIN_T gates)");
1454
1455 let rl = cache.recur[il].as_mut().unwrap();
1460 let hk = Self::gdn_hk(e, t, num_v, num_k);
1461 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
1462 let hk = if conv_fuse { hk } else { num_v }; let mut q_g = e.uninit(d_state * hk * t)?;
1464 let mut k_g = e.uninit(d_state * hk * t)?;
1465 let mut v_g = e.uninit(d_state * num_v * t)?;
1466 if conv_fuse {
1467 e.ssm_conv1d_gdn_state_pad(qkv_mixed, &mut rl.conv_state, la.ssm_conv1d.float_data(),
1468 &mut q_g, &mut k_g, &mut v_g,
1469 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim, hk, pad_len)?;
1470 } else {
1471 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(),
1473 &mut conv_out, conv_dim, t, d_conv, pad_len)?;
1474 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)?;
1475 }
1476 let mut q_l2 = e.uninit(d_state * hk * t)?;
1477 let qb16 = if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(32) {
1481 let mut qb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
1482 e.l2_norm_pp(&q_g, &mut q_l2, Some(&mut qb), d_state, hk * t, eps)?;
1483 Some(qb)
1484 } else {
1485 e.l2_norm_pp(&q_g, &mut q_l2, None, d_state, hk * t, eps)?;
1486 None
1487 };
1488 let mut k_l2 = e.uninit(d_state * hk * t)?;
1489 let kb16 = if Engine::l2_v2_on(d_state) {
1491 let mut kb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
1492 e.l2_norm_pp(&k_g, &mut k_l2, Some(&mut kb), d_state, hk * t, eps)?;
1493 Some(kb)
1494 } else {
1495 e.l2_norm_pp(&k_g, &mut k_l2, None, d_state, hk * t, eps)?;
1496 None
1497 };
1498 let mut beta = e.uninit(t * num_v)?;
1499 e.sigmoid_v(beta_raw, &mut beta, t * num_v)?;
1500 let mut g_log = e.uninit(t * num_v)?;
1501 e.gdn_glog_v(alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
1502 if let Some(len_d) = pad_len {
1503 e.gdn_pad_mask(&mut beta, &mut g_log, len_d, num_v, t)?;
1504 }
1505 Ok(GdnPrep { hk, q_l2, k_l2, v_g, beta, g_log, kb16, qb16 })
1506 }
1507
1508 #[allow(clippy::too_many_arguments)]
1513 fn linear_attn_prime_core_batch(&self, e: &Engine, la: &LinearAttnLayer,
1514 g4: &[CudaSlice<f32>], offs: &[usize], ts: &[usize],
1515 caches: &mut [&mut Cache], il: usize)
1516 -> Result<Vec<(CudaSlice<f32>, Option<CudaSlice<u8>>)>, Box<dyn std::error::Error>> {
1517 let ssm = self.cfg.ssm.as_ref().unwrap();
1518 let d_state = ssm.state_size as usize;
1519 let num_k = ssm.group_count as usize;
1520 let num_v = ssm.time_step_rank as usize;
1521 let key_dim = d_state * num_k;
1522 let value_dim = d_state * num_v;
1523 let conv_dim = key_dim * 2 + value_dim;
1524 let eps = self.cfg.rms_eps;
1525 let scale = 1.0 / (d_state as f32).sqrt();
1526 let b = ts.len();
1527 let c = Engine::gdn_chunk_size();
1528 let carried = caches.iter().any(|c| c.pos > 0);
1531 let use_vl = !carried
1532 && (2..=8).contains(&b)
1533 && Engine::gdn_chunked_enabled() && ts.iter().all(|&t| t >= 16)
1534 && e.gdn_mma_enabled(c)
1535 && std::env::var("MEMRA_GDN_VL").as_deref() != Ok("0");
1536 if !use_vl {
1537 return (0..b).map(|s| {
1538 let (o, t) = (offs[s], ts[s]);
1539 self.linear_attn_prime_core_pad_view(
1540 e, la,
1541 &g4[0].slice(o * conv_dim..(o + t) * conv_dim),
1542 &g4[1].slice(o * value_dim..(o + t) * value_dim),
1543 &g4[2].slice(o * num_v..(o + t) * num_v),
1544 &g4[3].slice(o * num_v..(o + t) * num_v),
1545 t, caches[s], il, None)
1546 }).collect();
1547 }
1548 struct SeqBufs {
1552 conv_out: CudaSlice<f32>, q_g: CudaSlice<f32>, k_g: CudaSlice<f32>, v_g: CudaSlice<f32>,
1553 q_l2: CudaSlice<f32>, k_l2: CudaSlice<f32>, beta: CudaSlice<f32>, g_log: CudaSlice<f32>,
1554 gn: CudaSlice<f32>, gn16: CudaSlice<u8>,
1555 }
1556 let d_conv = ssm.conv_kernel as usize;
1557 let f16o = Self::f16out_on(e, 16);
1558 let hk = Self::gdn_hk(e, 16, num_v, num_k); let mut sb = Vec::with_capacity(b);
1560 let mut pres = Vec::with_capacity(b);
1561 for &t in ts.iter().take(b) {
1562 sb.push(SeqBufs {
1563 conv_out: e.uninit(conv_dim * t)?,
1564 q_g: e.uninit(d_state * hk * t)?,
1565 k_g: e.uninit(d_state * hk * t)?,
1566 v_g: e.uninit(d_state * num_v * t)?,
1567 q_l2: e.uninit(d_state * hk * t)?,
1568 k_l2: e.uninit(d_state * hk * t)?,
1569 beta: e.uninit(t * num_v)?,
1570 g_log: e.uninit(t * num_v)?,
1571 gn: e.uninit(d_state * num_v * t)?,
1572 gn16: e.alloc_u8_uninit(d_state * num_v * t * 2)?,
1573 });
1574 pres.push(e.gdn_chunk_alloc(num_v, t, c, hk)?);
1575 }
1576 let prep_args: Vec<crate::GdnPrepVl> = (0..b).map(|s| {
1577 let (o, t) = (offs[s], ts[s]);
1578 let rl = caches[s].recur[il].as_ref().unwrap();
1579 crate::GdnPrepVl {
1580 qkv: e.addr_f32v(&g4[0].slice(o * conv_dim..(o + t) * conv_dim)),
1581 conv_state: e.addr_f32(&rl.conv_state),
1582 conv_out: e.addr_f32(&sb[s].conv_out),
1583 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),
1584 q_l2: e.addr_f32(&sb[s].q_l2), k_l2: e.addr_f32(&sb[s].k_l2),
1585 beta_raw: e.addr_f32v(&g4[2].slice(o * num_v..(o + t) * num_v)),
1586 alpha: e.addr_f32v(&g4[3].slice(o * num_v..(o + t) * num_v)),
1587 beta: e.addr_f32(&sb[s].beta), g_log: e.addr_f32(&sb[s].g_log),
1588 o: e.addr_f32(&pres[s].o),
1589 z: e.addr_f32v(&g4[1].slice(o * value_dim..(o + t) * value_dim)),
1590 gn: e.addr_f32(&sb[s].gn), gn16: e.addr_u8(&sb[s].gn16),
1591 kb16: if Engine::l2_v2_on(d_state) { e.addr_u8(&pres[s].kb16) } else { 0 },
1592 qb16: if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(c) { e.addr_u8(&pres[s].qb16) } else { 0 },
1593 t: t as i32, pad: 0,
1594 }
1595 }).collect();
1596 let args: Vec<crate::GdnSeqVl> = (0..b).map(|s| {
1597 let rl = caches[s].recur[il].as_ref().unwrap();
1598 crate::GdnSeqVl {
1599 kb16: e.addr_u8(&pres[s].kb16), gcum: e.addr_f32(&pres[s].gcum),
1600 beta: e.addr_f32(&sb[s].beta), u: e.addr_f32(&pres[s].u),
1601 wb16: e.addr_u8(&pres[s].wb16), y: e.addr_u8(&pres[s].y16),
1602 ssnap: e.addr_u8(&pres[s].ssnap16),
1603 state_in: e.addr_f32(&rl.ssm_state), state_out: e.addr_f32(&rl.ssm_state_alt),
1604 q: e.addr_f32(&sb[s].q_l2), p: e.addr_f32(&pres[s].p),
1605 o: e.addr_f32(&pres[s].o),
1606 k: e.addr_f32(&sb[s].k_l2), v: e.addr_f32(&sb[s].v_g),
1607 g: e.addr_f32(&sb[s].g_log), a: e.addr_f32(&pres[s].a),
1608 w: e.addr_f32(&pres[s].w),
1609 t: ts[s] as i32, nc: pres[s].nc as i32,
1610 }
1611 }).collect();
1612 e.gdn_prep_vl8(&prep_args, la.ssm_conv1d.float_data(), la.ssm_dt.float_data(),
1613 la.ssm_a.float_data(), conv_dim, d_conv, d_state, num_v, num_k, key_dim, hk, eps)?;
1614 if !Engine::l2_v2_on(d_state) {
1617 e.gdn_mirror_vl8(&args, num_v, 0, hk)?;
1618 }
1619 let wq8: Option<crate::GdnWVl8> = if e.gdn_wgmma_on(c) {
1621 if !Engine::l2_v2_on(d_state) {
1623 for s in 0..b {
1624 e.f32_to_bf16_into(&sb[s].q_l2, &mut pres[s].qb16, d_state * hk * ts[s])?;
1625 }
1626 }
1627 let mut wa = [crate::GdnWVl::default(); 8];
1628 for s in 0..b {
1629 wa[s] = crate::GdnWVl { qb16: e.addr_u8(&pres[s].qb16), pb16: e.addr_u8(&pres[s].pb16) };
1630 }
1631 Some(crate::GdnWVl8(wa))
1632 } else { None };
1633 e.gdn_chunk_k123_vl8(&args, num_v, hk, wq8.as_ref())?;
1634 e.gdn_chunk_vl8(&args, num_v, scale, hk, wq8.as_ref())?;
1635 if f16o {
1636 e.gdn_tail_vl8(&prep_args, la.ssm_norm.float_data(), d_state, num_v, eps)?;
1637 }
1638 let mut out = Vec::with_capacity(b);
1640 for (s, bufs) in sb.into_iter().enumerate() {
1641 let rl = caches[s].recur[il].as_mut().unwrap();
1642 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
1643 let (o, t) = (offs[s], ts[s]);
1644 let SeqBufs { mut gn, gn16, .. } = bufs;
1645 if f16o {
1646 out.push((gn, Some(gn16)));
1647 } else {
1648 let z_v = g4[1].slice(o * value_dim..(o + t) * value_dim);
1649 e.gated_rmsnorm_zv(&pres[s].o, la.ssm_norm.float_data(), &z_v, &mut gn,
1650 d_state, num_v * t, eps)?;
1651 out.push((gn, None));
1652 }
1653 }
1654 Ok(out)
1655 }
1656
1657 #[allow(clippy::too_many_arguments)]
1661 fn linear_attn_prime_core_pad_view(&self, e: &Engine, la: &LinearAttnLayer,
1662 qkv_mixed: &cudarc::driver::CudaView<f32>,
1663 z: &cudarc::driver::CudaView<f32>,
1664 beta_raw: &cudarc::driver::CudaView<f32>,
1665 alpha: &cudarc::driver::CudaView<f32>,
1666 t: usize, cache: &mut Cache, il: usize,
1667 pad_len: Option<&CudaSlice<i32>>)
1668 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1669 let cfg = &self.cfg;
1670 let ssm = cfg.ssm.as_ref().unwrap();
1671 let d_state = ssm.state_size as usize; let num_v = ssm.time_step_rank as usize; let eps = cfg.rms_eps;
1674 let scale = 1.0 / (d_state as f32).sqrt();
1675
1676 let prep = self.linear_attn_gdn_prep(e, la, qkv_mixed, beta_raw, alpha, t, cache, il, pad_len)?;
1677
1678 let mut o = e.uninit(d_state * num_v * t)?;
1684 let rl = cache.recur[il].as_mut().unwrap();
1685 {
1686 let crate::cache::RecurLayer { ssm_state, ssm_state_alt, .. } = rl;
1687 e.gdn_scan_prefill(&prep.q_l2, &prep.k_l2, &prep.v_g, &prep.g_log, &prep.beta,
1688 prep.kb16.as_ref(), prep.qb16.as_ref(), ssm_state, ssm_state_alt, &mut o, num_v, t, scale,
1689 prep.hk)?;
1690 }
1691 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
1692
1693 let mut gn = e.uninit(d_state * num_v * t)?;
1696 let gn16 = if Self::f16out_on(e, t) {
1697 let mut g16 = e.alloc_u8_uninit(d_state * num_v * t * 2)?;
1698 e.gated_rmsnorm_f16out_zv(&o, la.ssm_norm.float_data(), z, &mut gn, &mut g16,
1699 d_state, num_v * t, eps)?;
1700 Some(g16)
1701 } else {
1702 e.gated_rmsnorm_zv(&o, la.ssm_norm.float_data(), z, &mut gn, d_state, num_v * t, eps)?;
1703 None
1704 };
1705 Ok((gn, gn16))
1706 }
1707
1708 #[allow(clippy::too_many_arguments)]
1710 fn linear_attn_prime_core_pad(&self, e: &Engine, la: &LinearAttnLayer, g4: Vec<CudaSlice<f32>>,
1711 t: usize, cache: &mut Cache, il: usize,
1712 pad_len: Option<&CudaSlice<i32>>)
1713 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1714 let (gn, gn16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, pad_len)?;
1715 if let Some(xh) = &gn16 {
1716 if let Some(y) = e.try_f16_gemm_pre(&la.ssm_out, xh, t)? {
1717 return Ok(y);
1718 }
1719 }
1720 Ok(e.matmul(&la.ssm_out, &gn, t)?)
1721 }
1722
1723 pub fn full_attn(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
1725 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1726 let cfg = &self.cfg;
1727 let _n_embd = cfg.n_embd as usize;
1728 let n_head = cfg.n_head as usize;
1729 let n_head_kv = cfg.n_head_kv as usize;
1730 let head_dim = cfg.head_dim_k as usize;
1731 let eps = cfg.rms_eps;
1732 let scale = 1.0 / (head_dim as f32).sqrt();
1733
1734 let gated = cfg.attn_out_gate();
1737 let mut g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
1739 let v = g3.pop().unwrap();
1740 let mut k = g3.pop().unwrap();
1741 let qf = g3.pop().unwrap();
1742 let (mut q, gate) = if gated {
1743 let mut q = e.uninit(t * n_head * head_dim)?;
1744 let mut gate = e.uninit(t * n_head * head_dim)?;
1745 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
1746 (q, Some(gate))
1747 } else {
1748 (qf, None)
1749 };
1750
1751 let mut qn = e.uninit(t * n_head * head_dim)?;
1753 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
1754 q = qn;
1755 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
1756 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
1757 k = kn;
1758 let rope_dims = cfg.rope_dim_count as usize;
1759 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, cfg.rope_freq_base, 1.0)?;
1760 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, cfg.rope_freq_base, 1.0)?;
1761
1762 let mut attn = e.uninit(t * n_head * head_dim)?;
1764 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
1767 e.sdpa_naive(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1769 } else {
1770 e.fa_prefill(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1771 }
1772
1773 let attn_g = match &gate {
1775 Some(gate) => {
1776 let mut gsig = e.uninit(t * n_head * head_dim)?;
1777 e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
1778 let mut ag = e.uninit(t * n_head * head_dim)?;
1779 e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
1780 ag
1781 }
1782 None => attn,
1783 };
1784
1785 let o = e.matmul(&fa.wo, &attn_g, t)?;
1787 Ok(o)
1788 }
1789
1790 pub fn linear_attn(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>, t: usize)
1792 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1793 let cfg = &self.cfg;
1794 let _n_embd = cfg.n_embd as usize;
1795 let ssm = cfg.ssm.as_ref().unwrap();
1796 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;
1801 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;
1805 let scale = 1.0 / (d_state as f32).sqrt();
1806
1807 let mut g4 = e.matmul_group(&[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha], h, t)?;
1810 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);
1822 let mut q_g = e.uninit(d_state * num_v * t)?;
1823 let mut k_g = e.uninit(d_state * num_v * t)?;
1824 let mut v_g = e.uninit(d_state * num_v * t)?;
1825 e.ssm_conv1d_gdn(&qkv_mixed, la.ssm_conv1d.float_data(), &mut q_g, &mut k_g, &mut v_g,
1826 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim)?;
1827 let mut q_l2 = e.uninit(d_state * num_v * t)?;
1829 e.l2_norm(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
1830 let mut k_l2 = e.uninit(d_state * num_v * t)?;
1831 e.l2_norm(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
1832 let v_gd = v_g;
1833
1834 let mut beta = e.uninit(t * num_v)?;
1837 e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
1838 let mut g_log = e.uninit(t * num_v)?;
1840 e.gdn_glog(&alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
1841
1842 let state_in = e.zeros(d_state * d_state * num_v)?; let mut state_out = e.zeros(d_state * d_state * num_v)?;
1845 let mut o = e.uninit(d_state * num_v * t)?;
1846 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)?;
1847
1848 let mut gn = e.uninit(d_state * num_v * t)?;
1853 e.gated_rmsnorm(&o, la.ssm_norm.float_data(), &z, &mut gn, d_state, num_v * t, eps)?;
1854
1855 let out = e.matmul(&la.ssm_out, &gn, t)?;
1859 Ok(out)
1860 }
1861}
1862
1863impl HybridModel {
1864 pub fn moe_ffn_il(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize, il: u16)
1875 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1876 Self::moe_ffn(e, m, z, t, &self.cfg, il, self.max_moe_block())
1877 }
1878
1879 pub fn moe_ffn_il_zq8(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
1883 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, t: usize, il: u16)
1884 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1885 Self::moe_ffn_inner(e, m, z, zq8, t, &self.cfg, il, self.max_moe_block())
1886 }
1887
1888 pub(crate) fn moe_ffn(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
1896 cfg: &ModelConfig, il: u16, max_block: usize)
1897 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1898 Self::moe_ffn_inner(e, m, z, None, t, cfg, il, max_block)
1899 }
1900
1901 #[allow(clippy::too_many_arguments)]
1902 pub(crate) fn moe_ffn_inner(
1903 e: &Engine,
1904 m: &MoeWeights,
1905 z: &CudaSlice<f32>,
1906 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
1907 t: usize,
1908 cfg: &ModelConfig,
1909 il: u16,
1910 max_block: usize,
1911 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1912 let worker_io = crate::spill_pread::worker_enabled();
1913 let epoch_lfu = std::env::var_os("MEMRA_MOE_LFU_DECAY").is_some();
1914 if Engine::moe_cache_enabled() && (worker_io || epoch_lfu) {
1915 e.with_moe_cache(max_block, |cache, _| {
1916 cache.begin_forward_epoch(il, t);
1917 if worker_io {
1918 cache.begin_worker_scope();
1919 }
1920 Ok(())
1921 })?;
1922 }
1923 if t > 1 && std::env::var("MEMRA_MOE_GROUPED").is_ok() {
1925 let grouped_out = Self::moe_ffn_grouped(e, m, z, t, cfg, il, max_block)?;
1926 if std::env::var("MEMRA_MOE_GATE").is_ok() {
1933 let seq_out = Self::moe_ffn_sequential(e, m, z, t, cfg, il, max_block)?;
1934 let g_host = e.dtoh(&grouped_out)?;
1935 let s_host = e.dtoh(&seq_out)?;
1936 let g_bytes: &[u8] = unsafe { std::slice::from_raw_parts(g_host.as_ptr() as *const u8, g_host.len() * 4) };
1937 let s_bytes: &[u8] = unsafe { std::slice::from_raw_parts(s_host.as_ptr() as *const u8, s_host.len() * 4) };
1938 if g_bytes == s_bytes {
1939 if il == 0 { println!("moe-gate il={il} t={t} BYTE-IDENTICAL (first layer only printed)"); }
1940 } else {
1941 let diffs = g_host.iter().zip(s_host.iter()).enumerate()
1942 .filter(|(_, (a, b))| a != b).count();
1943 let maxdiff = g_host.iter().zip(s_host.iter())
1944 .map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
1945 panic!("moe-gate il={il} t={t} MISMATCH: {diffs}/{} elems differ, maxdiff={maxdiff:.6e}", g_host.len());
1946 }
1947 }
1948 return Ok(grouped_out);
1949 }
1950 Self::moe_ffn_sequential_zq8(e, m, z, zq8, t, cfg, il, max_block)
1951 }
1952
1953 pub(crate) fn moe_ffn_sequential(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
1955 cfg: &ModelConfig, il: u16, max_block: usize)
1956 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1957 Self::moe_ffn_sequential_zq8(e, m, z, None, t, cfg, il, max_block)
1958 }
1959
1960 fn trace_moe_routes(il: u16, t: usize, sel_all: &[u32], weights: &[f32])
1964 -> Result<(), Box<dyn std::error::Error>> {
1965 use std::io::Write as _;
1966 if let Ok(path) = std::env::var("MEMRA_MOE_TRACE") {
1967 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
1968 let ids: Vec<String> = sel_all.iter().map(|s| s.to_string()).collect();
1969 writeln!(f, "{} {} {}", il, t, ids.join(","))?;
1970 }
1971 if let Ok(path) = std::env::var("MEMRA_MOE_WEIGHT_TRACE") {
1972 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
1973 let pairs: Vec<String> = sel_all.iter().zip(weights)
1974 .map(|(expert, weight)| format!("{expert}:{weight:.9}"))
1975 .collect();
1976 writeln!(f, "{} {} {}", il, t, pairs.join(","))?;
1977 }
1978 Ok(())
1979 }
1980
1981 fn trace_moe_input(e: &Engine, il: u16, t: usize, n_embd: usize, z: &CudaSlice<f32>)
1986 -> Result<(), Box<dyn std::error::Error>> {
1987 use std::io::Write as _;
1988 let Ok(dir) = std::env::var("MEMRA_MOE_INPUT_TRACE_DIR") else { return Ok(()) };
1989 let host = e.dtoh(z)?;
1990 if host.len() != t * n_embd {
1991 return Err(format!(
1992 "MoE input trace shape mismatch at layer {il}: got {} values, expected {}x{}",
1993 host.len(), t, n_embd
1994 ).into());
1995 }
1996 let bytes = unsafe {
1997 std::slice::from_raw_parts(
1998 host.as_ptr().cast::<u8>(), host.len() * std::mem::size_of::<f32>()
1999 )
2000 };
2001 let state = MOE_INPUT_TRACE_WRITER.get_or_init(|| std::sync::Mutex::new(None));
2002 let mut state = state.lock().map_err(|_| "MoE input trace writer lock is poisoned")?;
2003 if state.is_none() {
2004 let dir = std::path::PathBuf::from(&dir);
2005 std::fs::create_dir_all(&dir)?;
2006 let index = std::fs::OpenOptions::new().create(true).append(true)
2007 .open(dir.join("index.jsonl"))?;
2008 *state = Some(MoeInputTraceWriter {
2009 dir,
2010 index,
2011 payloads: std::collections::HashMap::new(),
2012 });
2013 }
2014 let writer = state.as_mut().unwrap();
2015 if writer.dir != std::path::Path::new(&dir) {
2016 return Err("MEMRA_MOE_INPUT_TRACE_DIR changed after capture started".into());
2017 }
2018 let file_name = format!("layer-{il:03}.f32");
2019 if !writer.payloads.contains_key(&il) {
2020 let payload = std::fs::OpenOptions::new().create(true).append(true)
2021 .open(writer.dir.join(&file_name))?;
2022 let offset = payload.metadata()?.len();
2023 writer.payloads.insert(il, (payload, offset));
2024 }
2025 let (payload, offset) = writer.payloads.get_mut(&il).unwrap();
2026 let row_offset = *offset;
2027 payload.write_all(bytes)?;
2028 *offset += bytes.len() as u64;
2029 writeln!(
2030 writer.index,
2031 "{{\"format\":\"memra-moe-input-trace-v1\",\"layer\":{il},\"tokens\":{t},\
2032 \"hidden_size\":{n_embd},\"file\":\"{file_name}\",\"offset\":{row_offset},\
2033 \"payload_bytes\":{}}}",
2034 bytes.len()
2035 )?;
2036 Ok(())
2037 }
2038
2039 #[allow(clippy::too_many_arguments)]
2040 pub(crate) fn moe_ffn_sequential_zq8(
2041 e: &Engine,
2042 m: &MoeWeights,
2043 z: &CudaSlice<f32>,
2044 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
2045 t: usize,
2046 cfg: &ModelConfig,
2047 il: u16,
2048 max_block: usize,
2049 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2050 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
2051 let moe = cfg.moe.as_ref().unwrap();
2052 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);
2059 debug_assert_eq!(m.gate_exps.out_f, n_ff_exp);
2060 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);
2063
2064 let use_cache = Engine::moe_cache_enabled();
2065 let uniform_experts = m.has_uniform_expert_layout();
2066 let moe_q8 = uniform_experts && moe_q8_enabled()
2067 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
2068 && q8_expert_supported(m.down_exps.qtype);
2069 let cpu_expert_requested = crate::cpu_experts::configured();
2076 if cpu_expert_requested && (cfg.hy3.is_none() || cfg.m3.is_some()) {
2077 return Err(std::io::Error::other(
2078 "MEMRA_CPU_EXPERT_LIB is experimental and currently gated to Hy3",
2079 )
2080 .into());
2081 }
2082 let cpu_hybrid = cpu_expert_requested && t < PRIME_MIN_T && m.dev_exps.is_none();
2083 let freeze_cpu_residency = cpu_expert_requested
2089 && std::env::var("MEMRA_CPU_EXPERT_FREEZE_CACHE").as_deref() == Ok("1");
2090 let caller_warms_before_freeze = std::env::var("MEMRA_CPU_EXPERT_FREEZE_WARMUP_TOKENS")
2091 .ok()
2092 .and_then(|value| value.parse::<usize>().ok())
2093 .is_some_and(|tokens| tokens > 0);
2094 if cpu_hybrid && freeze_cpu_residency && !caller_warms_before_freeze {
2095 e.freeze_moe_cache();
2096 }
2097 let cache_frozen = use_cache && e.moe_cache_frozen();
2098 let cache_dispatch = use_cache && (!cache_frozen || cpu_hybrid);
2099
2100 let logits = if t < PRIME_MIN_T {
2107 if crate::router_kernel_on() {
2111 e.router_gemv(m.gate_inp.float_data(), z, cfg.n_embd as usize,
2114 m.gate_exps.n_expert, t)?
2115 } else {
2116 e.matmul_decode_exact(&m.gate_inp, z, t)?
2117 }
2118 } else if crate::router_prefill_exact_on() && crate::router_kernel_on() {
2119 e.router_gemv(m.gate_inp.float_data(), z, cfg.n_embd as usize,
2134 m.gate_exps.n_expert, t)?
2135 } else {
2136 e.matmul(&m.gate_inp, z, t)?
2137 };
2138
2139 let no_exp_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
2177 && m.down_exps.macros.is_none();
2178 if cfg.sigmoid_router().is_none() && cfg.m3.is_none() && cfg.hy3.is_none()
2179 && no_exp_macros
2180 && t >= PRIME_MIN_T && m.dev_exps.is_some() && moe_q8_enabled()
2181 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
2182 && q8_expert_supported(m.down_exps.qtype)
2183 && std::env::var("MEMRA_MOE_PAIRS").map(|v| v != "0").unwrap_or(true)
2184 && std::env::var("MEMRA_MOE_STATS").is_err() {
2185 return Self::moe_ffn_pairs(e, m, z, &logits, t, cfg);
2186 }
2187
2188 let dev_ok = uniform_experts && cfg.m3.is_none() && cfg.hy3.is_none();
2198 let observe_routes = std::env::var("MEMRA_MOE_STATS").is_ok()
2202 || std::env::var("MEMRA_MOE_TRACE").is_ok()
2203 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
2204 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok();
2205 if dev_ok && t < PRIME_MIN_T && m.dev_exps.is_some() && n_used <= 8 && moe_dev_enabled()
2206 && !observe_routes {
2207 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
2208 }
2209 if dev_ok && use_cache && n_used <= 8 && moe_dev_enabled()
2210 && !observe_routes {
2211 let row_ok = e.with_moe_cache(max_block, |c, eng| {
2212 if moe_prewarm_enabled() { c.prewarm_layer(il, m, eng)?; }
2213 Ok(c.layer_dev_row(il, n_expert, eng)?.is_some())
2214 })?;
2215 if row_ok {
2216 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
2217 }
2218 }
2219
2220 let (sel_all, w_all, routed_cpu_input) = if let Some(sig) = cfg.sigmoid_router() {
2222 if cpu_hybrid {
2223 let (sel, w, input) = Self::moe_route_sigmoid_with_input(
2224 e,
2225 &logits,
2226 z,
2227 t,
2228 n_expert,
2229 n_used,
2230 m.exp_probs_b.as_deref(),
2231 sig,
2232 m.active_experts.as_deref(),
2233 )?;
2234 (sel, w, Some(input))
2235 } else {
2236 let (sel, w) = Self::moe_route_cfg(
2237 e,
2238 &logits,
2239 t,
2240 n_expert,
2241 n_used,
2242 m.exp_probs_b.as_deref(),
2243 Some(sig),
2244 m.active_experts.as_deref(),
2245 )?;
2246 (sel, w, None)
2247 }
2248 } else {
2249 let (sel, w) = Self::moe_route_cfg(
2250 e,
2251 &logits,
2252 t,
2253 n_expert,
2254 n_used,
2255 None,
2256 None,
2257 m.active_experts.as_deref(),
2258 )?;
2259 (sel, w, None)
2260 };
2261
2262 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
2266 Self::trace_moe_input(e, il, t, n_embd, z)?;
2267
2268 let worker_disk_prefetch =
2280 cache_dispatch && crate::spill_pread::worker_enabled() && !cpu_hybrid;
2281 let promote_worker_h2d =
2282 t == 1 && worker_disk_prefetch && crate::spill_pread::copy_h2d_enabled();
2283 if promote_worker_h2d {
2284 let mut selected_blocks = Vec::with_capacity(n_used * 3);
2285 for &ex in sel_all.iter().take(n_used) {
2286 let ex = ex as u16;
2287 selected_blocks.extend([
2288 BlockId::new(il, PROJ_GATE, ex),
2289 BlockId::new(il, PROJ_UP, ex),
2290 BlockId::new(il, PROJ_DOWN, ex),
2291 ]);
2292 }
2293 for &ex in sel_all.iter().take(n_used) {
2294 Self::moe_prefetch_disk_expert(e, il, ex as usize, m, max_block, &selected_blocks)?;
2295 }
2296 e.with_moe_cache(max_block, |cache, eng| {
2297 cache.promote_worker_reads_at_safe_boundary(
2298 &selected_blocks,
2299 &selected_blocks,
2300 eng,
2301 )?;
2302 Ok(())
2303 })?;
2304 }
2305
2306 if t > 1 && std::env::var("MEMRA_MOE_STATS").is_ok() {
2309 let mut cnt = vec![0u32; n_expert];
2310 for &s in sel_all.iter() { cnt[s as usize] += 1; }
2311 let total = sel_all.len() as f64;
2312 let mut h = 0.0f64;
2313 let mut active = 0usize;
2314 for &c in &cnt { if c > 0 { active += 1; let p = c as f64 / total; h -= p * p.log2(); } }
2315 let maxc = cnt.iter().copied().max().unwrap_or(0);
2316 println!("moe-stats il={} t={} assignments={} active={}/{} entropy={:.3}b (max {:.3}b) mean_tok_per_active={:.2} max_tok_per_expert={}",
2317 il, t, sel_all.len(), active, n_expert, h, (n_expert as f64).log2(), total / active.max(1) as f64, maxc);
2318 }
2319
2320 let gdec_may_fire = uniform_experts && use_cache && n_used <= 8 && gdec_enabled();
2329 let mut moe_out = if gdec_may_fire {
2330 e.uninit(t * n_embd)?
2331 } else {
2332 e.zeros(t * n_embd)?
2333 };
2334 let cpu_input = if cpu_hybrid {
2337 Some(routed_cpu_input.ok_or("CPU expert routing did not return the MoE input")?)
2338 } else {
2339 None
2340 };
2341
2342 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;
2350 let mut scratch_u: Option<CudaSlice<u8>> = None;
2351 let mut scratch_d: Option<CudaSlice<u8>> = None;
2352 let page_window = moe_page_prefetch_window();
2360
2361 for tok in 0..t {
2364 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
2365 let w = &w_all[tok * n_used..(tok + 1) * n_used];
2366 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd); let mut tok_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
2368
2369 let no_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
2383 && m.down_exps.macros.is_none();
2384 if gdec_may_fire && moe_q8 && cfg.m3.is_none() && no_macros {
2385 if tok_q8.is_none() {
2386 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
2387 }
2388 let (zq, zd) = tok_q8.as_ref().unwrap();
2389 if Self::moe_gdec_token_q8(e, m, il, max_block, zq, zd, sel, w,
2390 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
2391 continue;
2392 }
2393 } else if gdec_may_fire && cfg.m3.is_none() && no_macros
2394 && Self::moe_gdec_token(e, m, il, max_block, &zt, sel, w,
2395 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
2396 continue;
2397 }
2398
2399 if gdec_may_fire {
2403 let mut row = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2404 e.memset_zeros_view(&mut row)?;
2405 }
2406
2407 let mut cpu_mask = vec![false; sel.len()];
2413 let cpu_worker = if let Some(host_input) = cpu_input.as_ref() {
2414 let gpu_resident = if use_cache {
2415 e.with_moe_cache(max_block, |cache, _| {
2416 Ok(sel
2417 .iter()
2418 .map(|&expert| {
2419 let expert = expert as u16;
2420 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
2421 .into_iter()
2422 .filter(|&projection| {
2423 cache
2424 .resident(BlockId::new(il, projection, expert))
2425 .is_some()
2426 })
2427 .count()
2428 })
2429 .collect::<Vec<_>>())
2430 })?
2431 } else {
2432 vec![0; sel.len()]
2433 };
2434 let mut cpu_selected = Vec::new();
2435 for (index, (&expert, &route_weight)) in sel.iter().zip(w).enumerate() {
2436 if gpu_resident[index] != 3 {
2437 cpu_mask[index] = true;
2438 crate::cpu_experts::record_incomplete_gpu_residency(gpu_resident[index]);
2439 let expert = expert as usize;
2440 cpu_selected.push((expert, route_weight));
2441 }
2442 }
2443 if crate::cpu_experts::predictor_enabled() {
2444 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
2448 crate::cpu_experts::predictor_submit(il, row);
2449 }
2450 if cpu_selected.is_empty() {
2451 None
2452 } else {
2453 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
2454 let job = crate::cpu_experts::prepare_job(m, il, &cpu_selected, row)
2455 .map_err(std::io::Error::other)?;
2456 Some(crate::cpu_experts::submit(job).map_err(std::io::Error::other)?)
2457 }
2458 } else {
2459 None
2460 };
2461
2462 let worker_window = worker_disk_prefetch
2463 .then(worker_prefetch_window)
2464 .unwrap_or(0);
2465 for (j, &ex) in sel.iter().enumerate() {
2466 if cpu_mask[j] {
2467 continue;
2468 }
2469 let ex = ex as usize;
2470 for next in page_prefetch_positions(j, sel.len(), page_window) {
2471 Self::moe_prefetch_host_expert(sel[next] as usize, m);
2472 }
2473 let keep = [
2474 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_GATE, ex as u16),
2475 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_UP, ex as u16),
2476 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_DOWN, ex as u16),
2477 ];
2478 if worker_disk_prefetch && worker_window > 0 {
2479 for next in worker_prefetch_positions(j, sel.len(), worker_window) {
2480 Self::moe_prefetch_disk_expert(
2481 e,
2482 il,
2483 sel[next] as usize,
2484 m,
2485 max_block,
2486 &keep,
2487 )?;
2488 }
2489 } else if cache_dispatch
2490 && !cpu_hybrid
2491 && moe_prefetch_enabled()
2492 && j + 1 < sel.len()
2493 {
2494 let next = sel[j + 1] as usize;
2495 Self::moe_prefetch_expert(e, il, next, m, max_block, &keep)?;
2496 }
2497 let [gate_q8, up_q8, down_q8] = [moe_q8; 3];
2498 if cache_dispatch && (gate_q8 || up_q8 || down_q8) {
2499 if (gate_q8 || up_q8) && tok_q8.is_none() {
2502 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
2503 }
2504 let gate = if gate_q8 {
2505 let (zq, zd) = tok_q8.as_ref().unwrap();
2506 Self::moe_cached_gemm_q8(e, il, PROJ_GATE, ex, m, max_block, zq, zd)?
2507 } else {
2508 Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?
2509 };
2510 let up = if up_q8 {
2511 let (zq, zd) = tok_q8.as_ref().unwrap();
2512 Self::moe_cached_gemm_q8(e, il, PROJ_UP, ex, m, max_block, zq, zd)?
2513 } else {
2514 Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?
2515 };
2516 let mut act = e.uninit(n_ff_exp)?;
2517 Self::ffn_act_scaled(
2518 e,
2519 cfg,
2520 &gate,
2521 &up,
2522 m.gate_exps.macro_scale(ex),
2523 m.up_exps.macro_scale(ex),
2524 &mut act,
2525 n_ff_exp,
2526 )?;
2527 let y = if down_q8 {
2528 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
2529 Self::moe_cached_gemm_q8(e, il, PROJ_DOWN, ex, m, max_block, &aq2, &ad2)?
2530 } else {
2531 let actv = act.slice(0..n_ff_exp);
2532 Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?
2533 };
2534 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2535 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2537 } else if cache_dispatch {
2538 let gate = Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?;
2543 let up = Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?;
2544 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_scaled(e, cfg, &gate, &up,
2546 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, n_ff_exp)?;
2547 let actv = act.slice(0..n_ff_exp);
2548 let y = Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?;
2549 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2550 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2552 } else if cache_frozen {
2553 let gate = Self::moe_frozen_gemm(
2558 e,
2559 il,
2560 PROJ_GATE,
2561 ex,
2562 m,
2563 max_block,
2564 &zt,
2565 &mut scratch_g,
2566 g_len,
2567 )?;
2568 let up = Self::moe_frozen_gemm(
2569 e,
2570 il,
2571 PROJ_UP,
2572 ex,
2573 m,
2574 max_block,
2575 &zt,
2576 &mut scratch_u,
2577 u_len,
2578 )?;
2579 let mut act = e.uninit(n_ff_exp)?;
2580 Self::ffn_act_scaled(
2581 e,
2582 cfg,
2583 &gate,
2584 &up,
2585 m.gate_exps.macro_scale(ex),
2586 m.up_exps.macro_scale(ex),
2587 &mut act,
2588 n_ff_exp,
2589 )?;
2590 let actv = act.slice(0..n_ff_exp);
2591 let y = Self::moe_frozen_gemm(
2592 e,
2593 il,
2594 PROJ_DOWN,
2595 ex,
2596 m,
2597 max_block,
2598 &actv,
2599 &mut scratch_d,
2600 d_len,
2601 )?;
2602 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2603 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2604 } else {
2605 if scratch_g.is_none() {
2609 scratch_g = Some(e.alloc_u8_uninit(g_len)?);
2610 scratch_u = Some(e.alloc_u8_uninit(u_len)?);
2611 scratch_d = Some(e.alloc_u8_uninit(d_len)?);
2612 }
2613 let (sg, su, sd) = (scratch_g.as_mut().unwrap(), scratch_u.as_mut().unwrap(),
2614 scratch_d.as_mut().unwrap());
2615 let gl = m.gate_exps.expert_layout(ex);
2616 let ul = m.up_exps.expert_layout(ex);
2617 let dl = m.down_exps.expert_layout(ex);
2618 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
2619 let gate = e.qmatvec_view(sg, 0..gl.len, &zt, 1,
2620 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)?;
2621
2622 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
2623 let up = e.qmatvec_view(su, 0..ul.len, &zt, 1,
2624 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)?;
2625
2626 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_scaled(e, cfg, &gate, &up,
2628 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, n_ff_exp)?;
2629
2630 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
2631 let actv = act.slice(0..n_ff_exp);
2632 let y = e.qmatvec_view(sd, 0..dl.len, &actv, 1,
2633 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)?;
2634
2635 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2636 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2637 }
2638 }
2639 if let Some(worker) = cpu_worker {
2640 let cpu_output = worker.wait().map_err(std::io::Error::other)?;
2641 let cpu_output = e.htod(&cpu_output)?;
2642 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2643 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
2644 }
2645 if cpu_hybrid && !cache_frozen && cpu_expert_profile_admit_enabled() {
2646 for (j, &ex) in sel.iter().enumerate() {
2647 if cpu_mask[j] {
2648 Self::moe_profile_admit_expert(e, il, ex as usize, m, max_block)?;
2649 }
2650 }
2651 }
2652 }
2653
2654 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
2659 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
2660 {
2661 let n_ff_sh = gate_shexp.out_features(); let verify_t = t > 1 && t < PRIME_MIN_T;
2670 let (sg_gate, sg_up) = if t == 1 {
2671 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
2672 Some(pair) => pair,
2673 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
2674 }
2675 } else if verify_t {
2676 (e.matmul_decode_exact(gate_shexp, z, t)?, e.matmul_decode_exact(up_shexp, z, t)?)
2677 } else {
2678 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?) };
2680 let mut sa = e.uninit(t * n_ff_sh)?; Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
2682 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
2683 else { e.matmul(down_shexp, &sa, t)? }; let g = match &m.gate_inp_shexp {
2697 Some(gate_inp_shexp) => {
2698 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
2699 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
2700 } else {
2701 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
2702 let mut g = e.uninit(t)?; e.sigmoid(&gs, &mut g, t)?;
2704 g
2705 }
2706 }
2707 None => e.htod(&vec![1.0f32; t])?,
2708 };
2709 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
2711 }
2712
2713 Ok(moe_out)
2714 }
2715
2716 pub fn stage1_h2d_per_token(&self) -> u64 {
2719 use crate::hybrid::Ffn;
2720 let n_used = self.cfg.moe.as_ref().map(|m| m.expert_used_count as u64).unwrap_or(0);
2721 let mut bytes = 0u64;
2722 for l in self.layers.iter() {
2723 if let Ffn::Moe(m) = &l.ffn {
2724 bytes += n_used * (m.gate_exps.max_expert_bytes() + m.up_exps.max_expert_bytes()
2725 + m.down_exps.max_expert_bytes()) as u64;
2726 }
2727 }
2728 bytes
2729 }
2730
2731 pub(crate) fn max_moe_block(&self) -> usize {
2735 use crate::hybrid::Ffn;
2736 let mut mx = 0usize;
2737 let mut scan = |ffn: &Ffn| {
2738 if let Ffn::Moe(m) = ffn {
2739 mx = mx.max(m.gate_exps.max_expert_bytes())
2740 .max(m.up_exps.max_expert_bytes())
2741 .max(m.down_exps.max_expert_bytes());
2742 }
2743 };
2744 for l in self.layers.iter() { scan(&l.ffn); }
2745 if let Some(mtp) = self.mtp.as_ref() { scan(&mtp.ffn); }
2746 mx
2747 }
2748
2749 pub(crate) fn moe_cache_block_sizes(&self) -> Vec<usize> {
2752 use crate::hybrid::Ffn;
2753 let mut sizes = Vec::new();
2754 let mut scan = |ffn: &Ffn| {
2755 let Ffn::Moe(m) = ffn else { return };
2756 for ex in 0..m.gate_exps.n_expert {
2757 if m.active_experts.as_ref().is_some_and(|active| !active[ex]) {
2758 continue;
2759 }
2760 for exps in [&m.gate_exps, &m.up_exps, &m.down_exps] {
2761 let len = exps.expert_layout(ex).len;
2762 if len > 0 {
2763 sizes.push(len);
2764 }
2765 }
2766 }
2767 };
2768 for layer in &self.layers {
2769 scan(&layer.ffn);
2770 }
2771 if let Some(mtp) = &self.mtp {
2772 scan(&mtp.ffn);
2773 }
2774 sizes
2775 }
2776
2777 pub fn save_cpu_expert_residency_profile(
2783 &self,
2784 e: &Engine,
2785 path: &std::path::Path,
2786 ) -> Result<(), Box<dyn std::error::Error>> {
2787 let Some(ids) = e.export_moe_residency() else {
2788 return Err("no MoE residency cache to persist".into());
2789 };
2790 let mut body = format!(
2791 "memra-freeze-profile v1 max_block={} blocks={}\n",
2792 self.max_moe_block(),
2793 ids.len()
2794 );
2795 for (layer, proj, ex) in &ids {
2796 body.push_str(&format!("{layer} {proj} {ex}\n"));
2797 }
2798 let tmp = path.with_extension("tmp");
2799 std::fs::write(&tmp, body)?;
2800 std::fs::rename(&tmp, path)?;
2801 println!(
2802 "[moe-cache] freeze profile saved: {} blocks -> {}",
2803 ids.len(),
2804 path.display()
2805 );
2806 Ok(())
2807 }
2808
2809 pub fn restore_cpu_expert_residency_profile(
2813 &self,
2814 e: &Engine,
2815 path: &std::path::Path,
2816 ) -> Result<bool, Box<dyn std::error::Error>> {
2817 use crate::hybrid::Ffn;
2818 use crate::moe_cache::BlockId;
2819 let Ok(content) = std::fs::read_to_string(path) else {
2820 return Ok(false);
2821 };
2822 let mut lines = content.lines();
2823 let Some(header) = lines.next() else { return Ok(false) };
2824 let expected = format!("memra-freeze-profile v1 max_block={}", self.max_moe_block());
2825 if !header.starts_with(&expected) {
2826 println!(
2827 "[moe-cache] freeze profile ignored (geometry mismatch): {}",
2828 path.display()
2829 );
2830 return Ok(false);
2831 }
2832 let mut by_layer: std::collections::HashMap<u16, Vec<BlockId>> =
2833 std::collections::HashMap::new();
2834 for line in lines {
2835 let mut fields = line.split_whitespace();
2836 let (Some(layer), Some(proj), Some(ex)) =
2837 (fields.next(), fields.next(), fields.next())
2838 else {
2839 continue;
2840 };
2841 let (Ok(layer), Ok(proj), Ok(ex)) =
2842 (layer.parse::<u16>(), proj.parse::<u8>(), ex.parse::<u16>())
2843 else {
2844 continue;
2845 };
2846 by_layer
2847 .entry(layer)
2848 .or_default()
2849 .push(BlockId::new(layer, proj, ex));
2850 }
2851 let requested: usize = by_layer.values().map(Vec::len).sum();
2852 if requested == 0 {
2853 return Ok(false);
2854 }
2855 let max_block = self.max_moe_block();
2856 let mut restaged = 0usize;
2857 let mut stage_layer = |layer_index: u16,
2858 ffn: &Ffn|
2859 -> Result<(), Box<dyn std::error::Error>> {
2860 let Ffn::Moe(m) = ffn else { return Ok(()) };
2861 let Some(ids) = by_layer.get(&layer_index) else {
2862 return Ok(());
2863 };
2864 e.with_moe_cache(max_block, |cache, eng| {
2865 for id in ids {
2866 if cache.restage_block(*id, m, eng)? {
2867 restaged += 1;
2868 }
2869 }
2870 Ok(())
2871 })
2872 };
2873 for (index, layer) in self.layers.iter().enumerate() {
2874 stage_layer(index as u16, &layer.ffn)?;
2875 }
2876 if let Some(mtp) = self.mtp.as_ref() {
2877 stage_layer(u16::MAX, &mtp.ffn)?;
2878 }
2879 e.freeze_moe_cache();
2880 println!(
2881 "[moe-cache] freeze profile restored: {restaged}/{requested} blocks restaged from {}",
2882 path.display()
2883 );
2884 Ok(true)
2885 }
2886
2887 pub fn freeze_cpu_expert_residency(
2889 &self,
2890 e: &Engine,
2891 ) -> Result<(), Box<dyn std::error::Error>> {
2892 e.freeze_moe_cache();
2893 Ok(())
2894 }
2895
2896 pub fn ffn_act(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
2900 act: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
2901 Self::ffn_act_scaled(e, cfg, gate, up, 1.0, 1.0, act, n)
2902 }
2903
2904 #[allow(clippy::too_many_arguments)]
2908 pub(crate) fn ffn_act_scaled(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
2909 gs: f32, us: f32, act: &mut CudaSlice<f32>, n: usize)
2910 -> Result<(), Box<dyn std::error::Error>> {
2911 if let Some(m3) = cfg.m3.as_ref() {
2912 return e.swigluoai_mul_scaled(gate, up, gs, us, m3.swiglu_alpha, m3.swiglu_limit, act, n);
2913 }
2914 if gs == 1.0 && us == 1.0 { return e.silu_mul(gate, up, act, n); }
2915 e.silu_mul_scaled(gate, up, gs, us, act, n)
2916 }
2917
2918 fn moe_route(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2924 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2925 Self::moe_route_cfg(e, logits, t, n_expert, n_used, None, None, None)
2926 }
2927
2928 fn moe_route_cfg(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize,
2936 bias: Option<&[f32]>, sig: Option<(f32, bool)>, active: Option<&[bool]>)
2937 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2938 if let Some((sf, route_norm)) = sig {
2939 let lg = e.dtoh(logits)?;
2941 return Self::moe_route_sigmoid_host(
2942 &lg, t, n_expert, n_used, bias, sf, route_norm, active,
2943 );
2944 }
2945 if active.is_none() && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")) {
2949 return e.moe_router_topk_host(logits, t, n_expert, n_used);
2950 }
2951 let lg = e.dtoh(logits)?; let mut sel = vec![0u32; t * n_used];
2954 let mut w_out = vec![0f32; t * n_used];
2955 for tok in 0..t {
2956 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
2957 let maxl = row.iter().enumerate()
2959 .filter(|(i, _)| active.is_none_or(|mask| mask[*i]))
2960 .map(|(_, &x)| x).fold(f32::NEG_INFINITY, f32::max);
2961 let mut probs = vec![0f32; n_expert];
2962 let mut den = 0f32;
2963 for i in 0..n_expert {
2964 if active.is_some_and(|mask| !mask[i]) { continue; }
2965 let x = (row[i] - maxl).exp(); probs[i] = x; den += x;
2966 }
2967 for p in probs.iter_mut() { *p /= den; }
2968 let mut idx: Vec<usize> = (0..n_expert)
2970 .filter(|&i| active.is_none_or(|mask| mask[i])).collect();
2971 idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
2972 let sl = &idx[..n_used];
2973 let mut wv: Vec<f32> = sl.iter().map(|&i| probs[i]).collect();
2974 let mut ws: f32 = wv.iter().sum();
2975 ws = ws.max(6.103515625e-5_f32); for x in wv.iter_mut() { *x /= ws; }
2977 for j in 0..n_used {
2978 sel[tok * n_used + j] = sl[j] as u32;
2979 w_out[tok * n_used + j] = wv[j];
2980 }
2981 }
2982 Ok((sel, w_out))
2983 }
2984
2985 #[allow(clippy::too_many_arguments)]
2986 fn moe_route_sigmoid_with_input(
2987 e: &Engine,
2988 logits: &CudaSlice<f32>,
2989 input: &CudaSlice<f32>,
2990 t: usize,
2991 n_expert: usize,
2992 n_used: usize,
2993 bias: Option<&[f32]>,
2994 (sf, route_norm): (f32, bool),
2995 active: Option<&[bool]>,
2996 ) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
2997 let (lg, input) = e.dtoh_pair(logits, input)?;
2998 let (sel, w) =
2999 Self::moe_route_sigmoid_host(&lg, t, n_expert, n_used, bias, sf, route_norm, active)?;
3000 Ok((sel, w, input))
3001 }
3002
3003 pub fn start_moe_prefetch_predictor(
3008 &self,
3009 e: &Engine,
3010 cfg: &ModelConfig,
3011 ) -> Result<(), Box<dyn std::error::Error>> {
3012 use crate::hybrid::Ffn;
3013 let Some(sig) = cfg.sigmoid_router() else {
3014 return Err("prefetch predictor requires a sigmoid-router arch".into());
3015 };
3016 let resident: std::collections::HashSet<(u16, u8, u16)> = e
3017 .export_moe_residency()
3018 .ok_or("prefetch predictor needs the frozen MoE residency cache")?
3019 .into_iter()
3020 .collect();
3021 let mut layers = Vec::new();
3022 for (index, layer) in self.layers.iter().enumerate() {
3023 let Ffn::Moe(m) = &layer.ffn else { continue };
3024 let crate::model::GpuTensor::Float { data, .. } = &m.gate_inp else { continue };
3025 let router = e.dtoh(data)?;
3026 let n_expert = m.gate_exps.n_expert;
3027 let n_embd = m.gate_exps.in_f;
3028 if router.len() != n_embd * n_expert {
3029 continue;
3030 }
3031 let build = |exps: &crate::model::HostExps| {
3032 (0..n_expert)
3033 .map(|expert| crate::cpu_experts::predictor_projection(exps, expert))
3034 .collect::<Vec<_>>()
3035 };
3036 layers.push((index as u16, crate::cpu_experts::PredictLayerInit {
3037 router,
3038 bias: m.exp_probs_b.clone(),
3039 active: m.active_experts.clone(),
3040 n_embd,
3041 n_used: cfg
3042 .moe
3043 .as_ref()
3044 .map(|moe| moe.expert_used_count as usize)
3045 .ok_or("prefetch predictor requires MoE config")?,
3046 sig,
3047 weights_n_expert: n_expert,
3048 gate: build(&m.gate_exps),
3049 up: build(&m.up_exps),
3050 down: build(&m.down_exps),
3051 }));
3052 }
3053 crate::cpu_experts::start_prefetch_predictor(layers, resident)
3054 .map_err(|error| error.into())
3055 }
3056
3057 #[allow(clippy::too_many_arguments)]
3060 pub(crate) fn moe_route_sigmoid_host_public(
3061 logits: &[f32],
3062 t: usize,
3063 n_expert: usize,
3064 n_used: usize,
3065 bias: Option<&[f32]>,
3066 sf: f32,
3067 route_norm: bool,
3068 active: Option<&[bool]>,
3069 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
3070 Self::moe_route_sigmoid_host(logits, t, n_expert, n_used, bias, sf, route_norm, active)
3071 }
3072
3073 #[allow(clippy::too_many_arguments)]
3074 fn moe_route_sigmoid_host(
3075 lg: &[f32],
3076 t: usize,
3077 n_expert: usize,
3078 n_used: usize,
3079 bias: Option<&[f32]>,
3080 sf: f32,
3081 route_norm: bool,
3082 active: Option<&[bool]>,
3083 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
3084 if lg.len() != t * n_expert {
3085 return Err(format!(
3086 "sigmoid router logits length mismatch: got {}, expected {}",
3087 lg.len(),
3088 t * n_expert,
3089 )
3090 .into());
3091 }
3092 let mut sel = vec![0u32; t * n_used];
3093 let mut w_out = vec![0f32; t * n_used];
3094 for tok in 0..t {
3095 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
3096 let scores: Vec<f32> = row.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect();
3097 let selsc: Vec<f32> = match bias {
3099 Some(b) => scores.iter().zip(b).map(|(s, bb)| s + bb).collect(),
3100 None => scores.clone(),
3101 };
3102 let mut idx: Vec<usize> = (0..n_expert)
3103 .filter(|&i| active.is_none_or(|mask| mask[i]))
3104 .collect();
3105 idx.sort_by(|&a, &b| selsc[b].total_cmp(&selsc[a]).then(a.cmp(&b)));
3106 let sl = &idx[..n_used];
3107 let mut wv: Vec<f32> = sl.iter().map(|&i| scores[i]).collect();
3108 if route_norm {
3109 let ws: f32 = wv.iter().sum::<f32>().max(1e-20);
3110 for x in wv.iter_mut() {
3111 *x = *x / ws * sf;
3112 }
3113 } else {
3114 for x in wv.iter_mut() {
3115 *x *= sf;
3116 }
3117 }
3118 for j in 0..n_used {
3119 sel[tok * n_used + j] = sl[j] as u32;
3120 w_out[tok * n_used + j] = wv[j];
3121 }
3122 }
3123 Ok((sel, w_out))
3124 }
3125
3126 fn moe_ffn_pairs(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, logits: &CudaSlice<f32>,
3135 t: usize, cfg: &ModelConfig)
3136 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3137 let moe = cfg.moe.as_ref().unwrap();
3138 let n_embd = cfg.n_embd as usize;
3139 let n_expert = moe.expert_count as usize;
3140 let n_used = moe.expert_used_count as usize;
3141 let n_ff_exp = moe.expert_ff_length as usize;
3142 let dev = m.dev_exps.as_ref().unwrap();
3143 let (rbg_d, rbu_d) = if dev.gu_il {
3145 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
3146 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
3147
3148 let (sel_all, w_all) = Self::moe_route(e, logits, t, n_expert, n_used)?;
3149 let n_pairs = t * n_used;
3150 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
3153 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
3154 let pair_w: Vec<f32> = w_all.clone();
3155 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
3156 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
3157 let pt = e.htod_i32(&pair_tok)?;
3158 let px = e.htod_i32(&pair_ex)?;
3159 let pw = e.htod(&pair_w)?;
3160 let toff = e.htod_i32(&tok_off)?;
3161 let tids = e.htod_i32(&tok_ids)?;
3162
3163 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
3167 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
3168 let mut ex_ids: Vec<i32> = Vec::new();
3169 let mut ex_off: Vec<i32> = vec![0];
3170 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
3171 for (ex, list) in by_ex.iter().enumerate() {
3172 if list.is_empty() { continue; }
3173 ex_ids.push(ex as i32);
3174 ex_pairs.extend_from_slice(list);
3175 ex_off.push(ex_pairs.len() as i32);
3176 }
3177 let n_active = ex_ids.len();
3178 let exi = e.htod_i32(&ex_ids)?;
3179 let exo = e.htod_i32(&ex_off)?;
3180 let exp_d = e.htod_i32(&ex_pairs)?;
3181 let _ = &px; static MMA_T: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
3202 let mma_t = *MMA_T.get_or_init(|| {
3203 std::env::var("MEMRA_MOE_MMA_T").ok().and_then(|v| v.parse().ok()).unwrap_or(16)
3204 });
3205 let use_mma = std::env::var("MEMRA_MOE_MMA").map(|v| v != "0").unwrap_or(true)
3206 && t >= mma_t
3207 && q8_expert_dec_supported(m.gate_exps.qtype) && q8_expert_dec_supported(m.up_exps.qtype)
3208 && q8_expert_dec_supported(m.down_exps.qtype)
3209 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
3210 let mma_capable = q8_expert_dec_supported(m.gate_exps.qtype)
3226 && q8_expert_dec_supported(m.up_exps.qtype)
3227 && q8_expert_dec_supported(m.down_exps.qtype)
3228 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
3229 let f16g_mode = crate::moe_f16g_mode();
3230 let f16g = f16g_mode != 0 && t >= mma_t
3231 && (f16g_mode != 3 || !mma_capable)
3232 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
3233 && f16g_proj_ok(m.up_exps.qtype, n_embd)
3234 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp);
3235 if use_mma || f16g {
3236 let y_down = if f16g {
3244 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
3248 let csr_tok_d = e.htod_i32(&csr_tok)?;
3249 let (z_f16, z_s) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
3250 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
3251 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
3252 m.gate_exps.qtype, rbg_d)?;
3253 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
3254 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
3255 m.up_exps.qtype, rbu_d)?;
3256 let act_csr = e.moe_pairs_silu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
3257 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
3258 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
3259 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
3260 m.down_exps.qtype, m.down_exps.row_bytes)?;
3261 e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?
3262 } else {
3263 let z_scr = e.mmq_iq_quantize_act(z, n_embd, t)?;
3265 let gate = e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
3266 n_embd, n_ff_exp, n_active, n_pairs, t,
3267 m.gate_exps.qtype, rbg_d)?;
3268 let up = e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
3269 n_embd, n_ff_exp, n_active, n_pairs, t,
3270 m.up_exps.qtype, rbu_d)?;
3271 let a_scr = if crate::moe_fuse_actq_on() {
3277 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 0)?
3278 } else {
3279 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
3280 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
3281 };
3282 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
3283 let pself = e.htod_i32(&pair_self)?;
3284 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
3285 n_ff_exp, n_embd, n_active, n_pairs, n_pairs,
3286 m.down_exps.qtype, m.down_exps.row_bytes)?
3287 };
3288 let mut moe_out = e.uninit(t * n_embd)?;
3289 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
3290 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
3291 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
3292 {
3293 let n_ff_sh = gate_shexp.out_features();
3294 let sg_gate = e.matmul(gate_shexp, z, t)?;
3295 let sg_up = e.matmul(up_shexp, z, t)?;
3296 let mut sa = e.uninit(t * n_ff_sh)?;
3297 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
3298 let sh = e.matmul(down_shexp, &sa, t)?;
3299 let g = match &m.gate_inp_shexp {
3305 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
3306 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
3307 }
3308 Some(gate_inp_shexp) => {
3309 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
3310 let mut g = e.uninit(t)?;
3311 e.sigmoid(&gs, &mut g, t)?;
3312 g
3313 }
3314 None => e.htod(&vec![1.0f32; t])?,
3315 };
3316 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
3317 }
3318 return Ok(moe_out);
3319 }
3320
3321 let dec = std::env::var("MEMRA_MOE_DEC").map(|v| v != "0").unwrap_or(true);
3324 let matvec = |proj, exi: &_, exo: &_, exp_d: &_, pt: &_, aq: &_, ad: &_,
3325 inf, outf, qtype, rb| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3326 let dec = dec && q8_expert_dec_supported(qtype);
3328 if dec { e.moe_pairs_matvec_q8_dec(&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
3329 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
3330 else { e.moe_pairs_matvec_q8_em (&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
3331 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
3332 };
3333 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3334 let gate = matvec(0, &exi, &exo, &exp_d, &pt, &zq, &zd,
3335 n_embd, n_ff_exp, m.gate_exps.qtype, rbg_d)?;
3336 let up = matvec(1, &exi, &exo, &exp_d, &pt, &zq, &zd,
3337 n_embd, n_ff_exp, m.up_exps.qtype, rbu_d)?;
3338 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
3339 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
3340 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
3342 let pself = e.htod_i32(&pair_self)?;
3343 let y_down = matvec(2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
3344 n_ff_exp, n_embd, m.down_exps.qtype, m.down_exps.row_bytes)?;
3345 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
3347
3348 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
3352 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
3353 {
3354 let n_ff_sh = gate_shexp.out_features();
3355 let sg_gate = e.matmul(gate_shexp, z, t)?;
3356 let sg_up = e.matmul(up_shexp, z, t)?;
3357 let mut sa = e.uninit(t * n_ff_sh)?;
3358 e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
3359 let sh = e.matmul(down_shexp, &sa, t)?;
3360 let g = match &m.gate_inp_shexp {
3365 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
3366 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
3367 }
3368 Some(gate_inp_shexp) => {
3369 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
3370 let mut g = e.uninit(t)?;
3371 e.sigmoid(&gs, &mut g, t)?;
3372 g
3373 }
3374 None => e.htod(&vec![1.0f32; t])?,
3375 };
3376 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
3377 }
3378 Ok(moe_out)
3379 }
3380
3381 #[allow(clippy::too_many_arguments)]
3383 #[allow(clippy::too_many_arguments)]
3384 fn moe_ffn_dev(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
3385 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, logits: &CudaSlice<f32>,
3386 t: usize, cfg: &ModelConfig, il: u16, max_block: usize)
3387 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3388 let moe = cfg.moe.as_ref().unwrap();
3389 let n_embd = cfg.n_embd as usize;
3390 let n_expert = moe.expert_count as usize;
3391 let n_used = moe.expert_used_count as usize;
3392 let n_ff_exp = moe.expert_ff_length as usize;
3393
3394 let (sel_d, mut w_d) = e.moe_router_topk(logits, t, n_expert, n_used)?;
3396 if m.has_macros {
3399 e.moe_w_scale_by_expert(&mut w_d, &sel_d, &m.dev_macros, n_expert, t * n_used)?;
3400 }
3401
3402 let mut moe_out = e.uninit(t * n_embd)?;
3404
3405 if let Some(dev) = m.dev_exps.as_ref() {
3408 let (rbg_d, rbu_d) = if dev.gu_il {
3411 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
3412 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
3413 let q8 = moe_q8_enabled()
3414 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3415 && q8_expert_supported(m.down_exps.qtype);
3416 let rows_arm = q8 && t > 1 && crate::spec::spec_m2()
3425 && n_ff_exp == 512 && n_used <= 8
3426 && std::env::var("MEMRA_MOE_DEVQ8_GU").map(|v| v.is_empty() || v == "v").unwrap_or(true)
3427 && std::env::var("MEMRA_MOE_DEVQ8_DOWN").map(|v| v.is_empty() || v == "w8h2v").unwrap_or(true);
3428 let csr_mode = std::env::var("MEMRA_MOE_CSR").ok()
3437 .and_then(|v| v.parse::<i32>().ok()).unwrap_or(1);
3438 let csr_qt = |qt: i32| qt == crate::QT_IQ4_XS || qt == crate::QT_IQ3_S;
3439 let csr_arm = rows_arm && csr_mode > 0 && t <= 10
3440 && csr_qt(m.gate_exps.qtype) && csr_qt(m.up_exps.qtype)
3441 && csr_qt(m.down_exps.qtype);
3442 if csr_arm {
3443 if csr_mode == 2 {
3444 static ENGAGED: std::sync::Once = std::sync::Once::new();
3445 ENGAGED.call_once(|| eprintln!("[memra] moe CSR byte-compare mode ON (t={t})"));
3446 }
3447 let n_pairs = t * n_used;
3448 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3449 let act = e.moe_gate_up_silu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, n_pairs,
3450 n_embd, n_ff_exp, n_used, n_expert,
3451 m.gate_exps.qtype, m.up_exps.qtype,
3452 rbg_d, rbu_d)?;
3453 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
3454 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
3458 t, n_ff_exp, n_embd, n_used, n_expert,
3459 m.down_exps.qtype, m.down_exps.row_bytes)?;
3460 if csr_mode == 2 {
3461 let act_r = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
3463 n_embd, n_ff_exp, n_used, n_expert,
3464 m.gate_exps.qtype, m.up_exps.qtype,
3465 rbg_d, rbu_d, &m.dev_macros)?;
3466 let mut out_r = e.uninit(t * n_embd)?;
3467 let (aq2r, ad2r) = e.quantize_q8_1(&act_r, n_pairs, n_ff_exp)?;
3468 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2r, &ad2r, &mut out_r,
3469 t, n_ff_exp, n_embd, n_used, n_expert,
3470 m.down_exps.qtype, m.down_exps.row_bytes)?;
3471 let (a1, a2) = (e.dtoh(&act)?, e.dtoh(&act_r)?);
3472 let (o1, o2) = (e.dtoh(&moe_out)?, e.dtoh(&out_r)?);
3473 let ba = a1.iter().zip(&a2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
3474 let bo = o1.iter().zip(&o2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
3475 if ba + bo > 0 {
3476 eprintln!("[csr-check] il={il} t={t} ACT diffs={ba}/{} OUT diffs={bo}/{}",
3477 a1.len(), o1.len());
3478 let sel_h = e.dtoh_i32(&sel_d)?;
3480 let mut shown = 0;
3481 for (i, (x, y)) in a1.iter().zip(&a2).enumerate() {
3482 if x.to_bits() != y.to_bits() && shown < 4 {
3483 let (p, o) = (i / n_ff_exp, i % n_ff_exp);
3484 let ex = sel_h[p];
3485 let npx = sel_h.iter().filter(|&&v| v == ex).count();
3486 eprintln!(" ACT p={p} ex={ex} np={npx} o={o} csr={x:e} rows={y:e}");
3487 shown += 1;
3488 }
3489 }
3490 std::process::exit(3);
3491 }
3492 }
3493 } else if rows_arm {
3494 if std::env::var("MEMRA_MOE_OVERLAP").as_deref() == Ok("1") {
3497 use std::sync::atomic::{AtomicU64, Ordering};
3498 static PAIRS: AtomicU64 = AtomicU64::new(0);
3499 static UNIQ: AtomicU64 = AtomicU64::new(0);
3500 static CALLS: AtomicU64 = AtomicU64::new(0);
3501 let sel_h = e.dtoh_i32(&sel_d)?;
3502 let mut u: Vec<i32> = sel_h.clone(); u.sort_unstable(); u.dedup();
3503 PAIRS.fetch_add(sel_h.len() as u64, Ordering::Relaxed);
3504 UNIQ.fetch_add(u.len() as u64, Ordering::Relaxed);
3505 let c = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
3506 if c % 480 == 0 {
3507 let p = PAIRS.load(Ordering::Relaxed); let q = UNIQ.load(Ordering::Relaxed);
3508 eprintln!("[overlap] calls={c} pairs={p} unique={q} ratio={:.3} (t={t})",
3509 q as f64 / p as f64);
3510 }
3511 }
3512 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3513 let act = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
3514 n_embd, n_ff_exp, n_used, n_expert,
3515 m.gate_exps.qtype, m.up_exps.qtype,
3516 rbg_d, rbu_d, &m.dev_macros)?;
3517 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
3518 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
3519 t, n_ff_exp, n_embd, n_used, n_expert,
3520 m.down_exps.qtype, m.down_exps.row_bytes)?;
3521 } else {
3522 for tok in 0..t {
3523 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
3524 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
3525 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
3526 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3527 if q8 {
3528 let (zq, zd) = match (t, zq8) {
3529 (1, Some((q, d))) => (q.clone(), d.clone()),
3530 _ => e.quantize_q8_1_view(&zt, 1, n_embd)?,
3531 };
3532 let act = e.moe_gate_up_silu8_dev_q8(&dev.ptr_row, &selt, &zq, &zd,
3533 n_embd, n_ff_exp, n_used, n_expert,
3534 m.gate_exps.qtype, m.up_exps.qtype,
3535 rbg_d, rbu_d, &m.dev_macros)?;
3536 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
3537 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selt, &wt, &aq2, &ad2, &mut dst,
3538 n_ff_exp, n_embd, n_used, n_expert,
3539 m.down_exps.qtype, m.down_exps.row_bytes)?;
3540 } else {
3541 let act = e.moe_gate_up_silu8_dev(&dev.ptr_row, &selt, &zt, n_embd, n_ff_exp,
3542 n_used, n_expert,
3543 m.gate_exps.qtype, m.up_exps.qtype,
3544 rbg_d, rbu_d, &m.dev_macros)?;
3545 e.moe_down8_fma_dev(&dev.ptr_row, &selt, &wt, &act, &mut dst,
3546 n_ff_exp, n_embd, n_used, n_expert,
3547 m.down_exps.qtype, m.down_exps.row_bytes)?;
3548 }
3549 }
3550 }
3551 } else {
3552 let q8 = moe_q8_enabled()
3559 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3560 && q8_expert_supported(m.down_exps.qtype);
3561 e.with_moe_cache(max_block, |c, eng| {
3562 let row = c.layer_dev_row(il, n_expert, eng)?
3563 .ok_or("moe_ffn_dev: layer row vanished under the lock")?;
3564 for tok in 0..t {
3565 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
3566 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
3567 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
3568 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3569 if q8 {
3570 let (zq, zd) = match (t, zq8) {
3571 (1, Some((q, d))) => (q.clone(), d.clone()),
3572 _ => eng.quantize_q8_1_view(&zt, 1, n_embd)?,
3573 };
3574 let act = eng.moe_gate_up_silu8_dev_q8(row, &selt, &zq, &zd,
3575 n_embd, n_ff_exp, n_used, n_expert,
3576 m.gate_exps.qtype, m.up_exps.qtype,
3577 m.gate_exps.row_bytes, m.up_exps.row_bytes,
3578 &m.dev_macros)?;
3579 let (aq2, ad2) = eng.quantize_q8_1(&act, n_used, n_ff_exp)?;
3580 eng.moe_down8_fma_dev_q8(row, &selt, &wt, &aq2, &ad2, &mut dst,
3581 n_ff_exp, n_embd, n_used, n_expert,
3582 m.down_exps.qtype, m.down_exps.row_bytes)?;
3583 } else {
3584 let act = eng.moe_gate_up_silu8_dev(row, &selt, &zt, n_embd, n_ff_exp,
3585 n_used, n_expert,
3586 m.gate_exps.qtype, m.up_exps.qtype,
3587 m.gate_exps.row_bytes, m.up_exps.row_bytes,
3588 &m.dev_macros)?;
3589 eng.moe_down8_fma_dev(row, &selt, &wt, &act, &mut dst,
3590 n_ff_exp, n_embd, n_used, n_expert,
3591 m.down_exps.qtype, m.down_exps.row_bytes)?;
3592 }
3593 }
3594 c.hits += (t * 3 * n_used) as u64;
3596 Ok(())
3597 })?;
3598 }
3599
3600 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
3605 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
3606 {
3607 let n_ff_sh = gate_shexp.out_features();
3608 let verify_t = t > 1 && t < PRIME_MIN_T;
3611 let (sg_gate, sg_up) = if t == 1 {
3612 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
3613 Some(pair) => pair,
3614 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
3615 }
3616 } else if verify_t {
3617 let mut fused = None;
3621 if crate::spec::spec_fused_t() && (2..=4).contains(&t)
3622 && e.uses_q8_1_fast(gate_shexp) && e.uses_q8_1_fast(up_shexp) {
3623 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3624 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
3625 }
3626 match fused {
3627 Some(pair) => pair,
3628 None => (e.matmul_decode_exact(gate_shexp, z, t)?,
3629 e.matmul_decode_exact(up_shexp, z, t)?),
3630 }
3631 } else {
3632 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
3633 };
3634 let mut sa = e.uninit(t * n_ff_sh)?; e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
3636 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
3637 else { e.matmul(down_shexp, &sa, t)? };
3638 let g = match &m.gate_inp_shexp {
3642 Some(gate_inp_shexp) => {
3643 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
3646 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
3647 } else {
3648 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
3649 let mut g = e.uninit(t)?;
3650 e.sigmoid(&gs, &mut g, t)?;
3651 g
3652 }
3653 }
3654 None => e.htod(&vec![1.0f32; t])?,
3655 };
3656 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
3657 }
3658
3659 Ok(moe_out)
3660 }
3661
3662 #[allow(clippy::too_many_arguments)]
3672 #[allow(clippy::too_many_arguments)]
3675 fn moe_gdec_token_q8(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
3676 zq: &CudaSlice<i8>, zd: &CudaSlice<f32>, sel: &[u32], w: &[f32],
3677 moe_out: &mut CudaSlice<f32>, tok: usize,
3678 n_embd: usize, n_ff_exp: usize, n_used: usize)
3679 -> Result<bool, Box<dyn std::error::Error>> {
3680 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
3681 use cudarc::driver::DevicePtr;
3682 let ptrs = e.with_moe_cache(max_block, |c, eng| {
3683 let mut g = [0u64; 8];
3684 let mut u = [0u64; 8];
3685 let mut d = [0u64; 8];
3686 for (j, &ex) in sel.iter().enumerate() {
3687 let ex = ex as u16;
3688 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
3689 c.resident(BlockId::new(il, PROJ_UP, ex)),
3690 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
3691 else { return Ok(None); };
3692 let __s = eng.stream();
3693 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
3694 let (pu, _e1) = c.slot(su).device_ptr(&__s);
3695 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
3696 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
3697 }
3698 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
3699 for &ex in sel {
3700 let ex = ex as u16;
3701 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
3702 c.note_profile_hit(BlockId::new(il, proj, ex));
3703 }
3704 }
3705 }
3706 c.hits += (3 * n_used) as u64;
3707 Ok(Some((g, u, d)))
3708 })?;
3709 let Some((g, u, d)) = ptrs else { return Ok(false) };
3710 let mut wv = [0f32; 8];
3711 wv[..n_used].copy_from_slice(w);
3712 let act = e.moe_gate_up_silu8_q8(crate::WPtr8(g), crate::WPtr8(u), zq, zd,
3713 n_embd, n_ff_exp, n_used,
3714 m.gate_exps.qtype, m.up_exps.qtype,
3715 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
3716 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
3718 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3719 e.moe_down8_fma_q8(crate::WPtr8(d), crate::F32x8(wv), &aq2, &ad2, &mut dst,
3720 n_ff_exp, n_embd, n_used,
3721 m.down_exps.qtype, m.down_exps.row_bytes)?;
3722 Ok(true)
3723 }
3724
3725 fn moe_gdec_token(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
3726 zt: &cudarc::driver::CudaView<f32>, sel: &[u32], w: &[f32],
3727 moe_out: &mut CudaSlice<f32>, tok: usize,
3728 n_embd: usize, n_ff_exp: usize, n_used: usize)
3729 -> Result<bool, Box<dyn std::error::Error>> {
3730 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
3731 use cudarc::driver::DevicePtr;
3732 let ptrs = e.with_moe_cache(max_block, |c, eng| {
3734 let mut g = [0u64; 8];
3735 let mut u = [0u64; 8];
3736 let mut d = [0u64; 8];
3737 for (j, &ex) in sel.iter().enumerate() {
3738 let ex = ex as u16;
3739 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
3740 c.resident(BlockId::new(il, PROJ_UP, ex)),
3741 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
3742 else { return Ok(None); };
3743 let __s = eng.stream();
3744 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
3745 let (pu, _e1) = c.slot(su).device_ptr(&__s);
3746 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
3747 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
3748 }
3749 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
3750 for &ex in sel {
3751 let ex = ex as u16;
3752 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
3753 c.note_profile_hit(BlockId::new(il, proj, ex));
3754 }
3755 }
3756 }
3757 c.hits += (3 * n_used) as u64; Ok(Some((g, u, d)))
3759 })?;
3760 let Some((g, u, d)) = ptrs else { return Ok(false) };
3761 let mut wv = [0f32; 8];
3762 wv[..n_used].copy_from_slice(w);
3763 let act = e.moe_gate_up_silu8(crate::WPtr8(g), crate::WPtr8(u), zt,
3765 n_embd, n_ff_exp, n_used,
3766 m.gate_exps.qtype, m.up_exps.qtype,
3767 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
3768 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3769 e.moe_down8_fma_into(crate::WPtr8(d), crate::F32x8(wv), &act, &mut dst,
3770 n_ff_exp, n_embd, n_used,
3771 m.down_exps.qtype, m.down_exps.row_bytes)?;
3772 Ok(true)
3773 }
3774
3775 fn moe_cached_gemm_q8(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
3780 max_block: usize, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
3781 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3782 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
3783 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
3784 let layout = exps.expert_layout(ex);
3785 let id = BlockId::new(il, proj, ex as u16);
3786 let source = exps.expert_source(ex);
3787 e.with_moe_cache(max_block, |c, eng| {
3788 let slot = c.dispatch_source(id, source, eng)?;
3789 let DispatchSlot::Resident(sl) = slot;
3790 let buf = c.slot(sl);
3791 eng.qmatvec_expert_q8(buf, 0..layout.len, aq, ad, 1, exps.in_f, exps.out_f,
3792 layout.qtype, layout.row_bytes)
3793 })
3794 }
3795
3796 fn moe_cached_gemm(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
3797 max_block: usize, x: &cudarc::driver::CudaView<f32>)
3798 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3799 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
3800 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
3801 let layout = exps.expert_layout(ex);
3802 let id = BlockId::new(il, proj, ex as u16);
3803 let source = exps.expert_source(ex);
3804 e.with_moe_cache(max_block, |c, eng| {
3806 let slot = c.dispatch_source(id, source, eng)?;
3807 let DispatchSlot::Resident(sl) = slot;
3810 let buf = c.slot(sl);
3811 eng.qmatvec_view(buf, 0..layout.len, x, 1, exps.in_f, exps.out_f,
3812 layout.qtype, layout.row_bytes)
3813 })
3814 }
3815
3816 fn moe_profile_admit_expert(
3820 e: &Engine,
3821 il: u16,
3822 ex: usize,
3823 m: &MoeWeights,
3824 max_block: usize,
3825 ) -> Result<(), Box<dyn std::error::Error>> {
3826 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3827 e.with_moe_cache(max_block, |cache, eng| {
3828 for (proj, exps) in [
3829 (PROJ_GATE, &m.gate_exps),
3830 (PROJ_UP, &m.up_exps),
3831 (PROJ_DOWN, &m.down_exps),
3832 ] {
3833 let id = BlockId::new(il, proj, ex as u16);
3834 let _ = cache.dispatch_source(id, exps.expert_source(ex), eng)?;
3835 }
3836 Ok(())
3837 })
3838 }
3839
3840 #[allow(clippy::too_many_arguments)]
3843 fn moe_frozen_gemm(
3844 e: &Engine,
3845 il: u16,
3846 proj: u8,
3847 ex: usize,
3848 m: &MoeWeights,
3849 max_block: usize,
3850 x: &cudarc::driver::CudaView<f32>,
3851 scratch: &mut Option<CudaSlice<u8>>,
3852 scratch_len: usize,
3853 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3854 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP};
3855 let exps = match proj {
3856 PROJ_GATE => &m.gate_exps,
3857 PROJ_UP => &m.up_exps,
3858 _ => &m.down_exps,
3859 };
3860 let layout = exps.expert_layout(ex);
3861 let id = BlockId::new(il, proj, ex as u16);
3862 if let Some(output) = e.with_moe_cache(max_block, |cache, eng| {
3863 let Some(slot) = cache.resident(id) else {
3864 return Ok(None);
3865 };
3866 let buf = cache.slot(slot);
3867 Ok(Some(eng.qmatvec_view(
3868 buf,
3869 0..layout.len,
3870 x,
3871 1,
3872 exps.in_f,
3873 exps.out_f,
3874 layout.qtype,
3875 layout.row_bytes,
3876 )?))
3877 })? {
3878 return Ok(output);
3879 }
3880 if scratch.is_none() {
3881 *scratch = Some(e.alloc_u8_uninit(scratch_len)?);
3882 }
3883 let scratch = scratch.as_mut().unwrap();
3884 e.stage_expert(exps.expert_bytes(ex), scratch, 0)?;
3885 e.qmatvec_view(
3886 scratch,
3887 0..layout.len,
3888 x,
3889 1,
3890 exps.in_f,
3891 exps.out_f,
3892 layout.qtype,
3893 layout.row_bytes,
3894 )
3895 }
3896
3897 fn moe_prefetch_expert(
3898 e: &Engine,
3899 il: u16,
3900 ex: usize,
3901 m: &MoeWeights,
3902 max_block: usize,
3903 keep: &[crate::moe_cache::BlockId],
3904 ) -> Result<(), Box<dyn std::error::Error>> {
3905 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3906 e.with_moe_cache(max_block, |c, eng| {
3907 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
3908 (PROJ_DOWN, &m.down_exps)] {
3909 let id = BlockId::new(il, proj, ex as u16);
3910 let _ = c.prefetch_source(id, exps.expert_source(ex), keep, eng)?;
3911 }
3912 Ok(())
3913 })
3914 }
3915
3916 fn moe_prefetch_disk_expert(e: &Engine, il: u16, ex: usize, m: &MoeWeights,
3919 max_block: usize, keep: &[crate::moe_cache::BlockId])
3920 -> Result<(), Box<dyn std::error::Error>> {
3921 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3922 e.with_moe_cache(max_block, |c, eng| {
3923 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
3924 (PROJ_DOWN, &m.down_exps)] {
3925 let source = exps.expert_source(ex);
3926 if let crate::model::ExpertSource::Disk { .. } = &source {
3927 let id = BlockId::new(il, proj, ex as u16);
3928 let _ = c.prefetch_source(id, source, keep, eng)?;
3929 }
3930 }
3931 Ok(())
3932 })
3933 }
3934
3935 #[inline]
3936 fn moe_prefetch_host_expert(ex: usize, m: &MoeWeights) {
3937 let _ = m.gate_exps.prefetch_expert_pages(ex);
3938 let _ = m.up_exps.prefetch_expert_pages(ex);
3939 let _ = m.down_exps.prefetch_expert_pages(ex);
3940 }
3941}
3942
3943impl HybridModel {
3960 pub(crate) fn moe_ffn_grouped(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
3963 cfg: &ModelConfig, il: u16, _max_block: usize)
3964 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3965 let moe = cfg.moe.as_ref().unwrap();
3966 let n_embd = cfg.n_embd as usize;
3967 let n_expert = moe.expert_count as usize;
3968 let n_used = moe.expert_used_count as usize;
3969 let n_ff_exp = moe.expert_ff_length as usize;
3970
3971 let logits = e.matmul(&m.gate_inp, z, t)?;
3973 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
3974 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
3975 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
3976 } else {
3977 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
3978 None, None, m.active_experts.as_deref())?
3979 };
3980 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
3981
3982 struct ExpertGroup {
3986 tok_indices: Vec<i32>, slot_indices: Vec<i32>, weights: Vec<f32>, }
3990 let mut groups: Vec<ExpertGroup> = (0..n_expert).map(|_| ExpertGroup {
3991 tok_indices: Vec::new(), slot_indices: Vec::new(), weights: Vec::new(),
3992 }).collect();
3993
3994 for tok in 0..t {
3995 for j in 0..n_used {
3996 let ex = sel_all[tok * n_used + j] as usize;
3997 let w = w_all[tok * n_used + j];
3998 groups[ex].tok_indices.push(tok as i32);
3999 groups[ex].slot_indices.push(j as i32);
4000 groups[ex].weights.push(w);
4001 }
4002 }
4003
4004 let mut slot_buf = e.zeros(t * n_used * n_embd)?;
4007 let mut wbuf = e.zeros(t * n_used)?; let g_len = m.gate_exps.max_expert_bytes();
4011 let u_len = m.up_exps.max_expert_bytes();
4012 let d_len = m.down_exps.max_expert_bytes();
4013 let use_cache = Engine::moe_cache_enabled();
4014 let max_block = _max_block;
4015
4016 let (mut scratch_g, mut scratch_u, mut scratch_d) = if !use_cache {
4018 (Some(e.alloc_u8(g_len)?), Some(e.alloc_u8(u_len)?), Some(e.alloc_u8(d_len)?))
4019 } else {
4020 (None, None, None)
4021 };
4022
4023 let mut order: Vec<usize> =
4034 (0..n_expert).filter(|&ex| !groups[ex].tok_indices.is_empty()).collect();
4035 order.sort_by(|&a, &b| groups[b].tok_indices.len()
4036 .cmp(&groups[a].tok_indices.len()).then(a.cmp(&b)));
4037 let mut m_dist: Vec<usize> = Vec::new(); let page_window = moe_page_prefetch_window();
4039 let worker_disk_prefetch = use_cache && crate::spill_pread::worker_enabled();
4040 if worker_disk_prefetch {
4041 if let Some(first) = grouped_worker_prefetch_position(order.len(), None) {
4042 Self::moe_prefetch_disk_expert(e, il, order[first], m, max_block, &[])?;
4043 }
4044 }
4045 for (order_pos, &ex) in order.iter().enumerate() {
4046 for next in page_prefetch_positions(order_pos, order.len(), page_window) {
4047 Self::moe_prefetch_host_expert(order[next], m);
4048 }
4049 if worker_disk_prefetch {
4050 if let Some(next) = grouped_worker_prefetch_position(order.len(), Some(order_pos)) {
4051 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
4052 let keep = [
4053 BlockId::new(il, PROJ_GATE, ex as u16),
4054 BlockId::new(il, PROJ_UP, ex as u16),
4055 BlockId::new(il, PROJ_DOWN, ex as u16),
4056 ];
4057 Self::moe_prefetch_disk_expert(e, il, order[next], m, max_block, &keep)?;
4058 }
4059 }
4060 let grp = &groups[ex];
4061 let m_e = grp.tok_indices.len();
4062 m_dist.push(m_e);
4063 let gl = m.gate_exps.expert_layout(ex);
4064 let ul = m.up_exps.expert_layout(ex);
4065 let dl = m.down_exps.expert_layout(ex);
4066
4067 let tok_idx_d = e.htod_i32(&grp.tok_indices)?;
4071 let slot_idx_d = e.htod_i32(&grp.slot_indices)?;
4072 let dmac = m.down_exps.macro_scale(ex);
4073 let weight_d = if dmac == 1.0 { e.htod(&grp.weights)? } else {
4074 let scaled: Vec<f32> = grp.weights.iter().map(|&w| w * dmac).collect();
4075 e.htod(&scaled)?
4076 };
4077
4078 let mut gathered = e.zeros(m_e * n_embd)?;
4080 e.gather_rows(z, &tok_idx_d, &mut gathered, n_embd, m_e)?;
4081 let gv = gathered.slice(0..m_e * n_embd);
4082
4083 let y = if use_cache {
4085 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
4086 let gate = e.with_moe_cache(max_block, |c, eng| {
4088 let id = BlockId::new(il, PROJ_GATE, ex as u16);
4089 let slot = c.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
4090 let buf = c.buf(slot);
4091 eng.qmatvec_view(buf, 0..gl.len, &gv, m_e,
4092 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)
4093 })?;
4094 let up = e.with_moe_cache(max_block, |c, eng| {
4095 let id = BlockId::new(il, PROJ_UP, ex as u16);
4096 let slot = c.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
4097 let buf = c.buf(slot);
4098 eng.qmatvec_view(buf, 0..ul.len, &gv, m_e,
4099 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)
4100 })?;
4101 let mut act = e.zeros(m_e * n_ff_exp)?;
4103 Self::ffn_act_scaled(e, cfg, &gate, &up,
4104 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, m_e * n_ff_exp)?;
4105 let actv = act.slice(0..m_e * n_ff_exp);
4106 e.with_moe_cache(max_block, |c, eng| {
4107 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
4108 let slot = c.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
4109 let buf = c.buf(slot);
4110 eng.qmatvec_view(buf, 0..dl.len, &actv, m_e,
4111 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)
4112 })?
4113 } else {
4114 let sg = scratch_g.as_mut().unwrap();
4116 let su = scratch_u.as_mut().unwrap();
4117 let sd = scratch_d.as_mut().unwrap();
4118 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
4119 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
4120 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
4121 let gate = e.qmatvec_view(sg, 0..gl.len, &gv, m_e,
4122 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)?;
4123 let up = e.qmatvec_view(su, 0..ul.len, &gv, m_e,
4124 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)?;
4125 let mut act = e.zeros(m_e * n_ff_exp)?;
4127 Self::ffn_act_scaled(e, cfg, &gate, &up,
4128 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, m_e * n_ff_exp)?;
4129 let actv = act.slice(0..m_e * n_ff_exp);
4130 e.qmatvec_view(sd, 0..dl.len, &actv, m_e,
4131 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)?
4132 };
4133
4134 e.scatter_slot(&y, &tok_idx_d, &slot_idx_d, &weight_d,
4136 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
4137 }
4138
4139 let mut moe_out = e.zeros(t * n_embd)?;
4141 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, t)?;
4142
4143 if std::env::var("MEMRA_MOE_STATS").is_ok() && !m_dist.is_empty() {
4145 m_dist.sort_unstable();
4146 let active = m_dist.len();
4147 let mean = m_dist.iter().sum::<usize>() as f64 / active as f64;
4148 let median = m_dist[active / 2];
4149 let max_m = *m_dist.last().unwrap();
4150 let min_m = m_dist[0];
4151 let above16 = m_dist.iter().filter(|&&x| x >= 16).count();
4152 println!("moe-grouped il={il} t={t} active={active}/{n_expert} \
4153 m_e: min={min_m} median={median} mean={mean:.1} max={max_m} \
4154 above_gemm_threshold(>=16)={above16}/{active}");
4155 }
4156
4157 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4161 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4162 {
4163 let n_ff_sh = gate_shexp.out_features();
4164 let sg_gate = e.matmul(gate_shexp, z, t)?;
4165 let sg_up = e.matmul(up_shexp, z, t)?;
4166 let mut sa = e.zeros(t * n_ff_sh)?;
4167 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
4168 let sh = e.matmul(down_shexp, &sa, t)?;
4169 let g = match &m.gate_inp_shexp {
4173 Some(gate_inp_shexp) => {
4174 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
4177 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
4178 } else {
4179 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
4180 let mut g = e.uninit(t)?;
4181 e.sigmoid(&gs, &mut g, t)?;
4182 g
4183 }
4184 }
4185 None => e.htod(&vec![1.0f32; t])?,
4186 };
4187 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
4188 }
4189
4190 Ok(moe_out)
4191 }
4192
4193 pub(crate) fn moe_ffn_lockstep(
4200 &self,
4201 e: &Engine,
4202 m: &MoeWeights,
4203 zbatch: &CudaSlice<f32>,
4204 mrows: usize,
4205 il: u16,
4206 max_block: usize,
4207 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4208 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
4209 let cfg = &self.cfg;
4210 let moe = cfg.moe.as_ref().unwrap();
4211 let n_embd = cfg.n_embd as usize;
4212 let n_expert = moe.expert_count as usize;
4213 let n_used = moe.expert_used_count as usize;
4214 let n_ff_exp = moe.expert_ff_length as usize;
4215
4216 let logits = e.matmul(&m.gate_inp, zbatch, mrows)?;
4217 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
4218 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
4219 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
4220 } else {
4221 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
4222 None, None, m.active_experts.as_deref())?
4223 };
4224 Self::trace_moe_routes(il, mrows, &sel_all, &w_all)?;
4225
4226 let resident_expert: Vec<bool> = e.with_moe_cache(max_block, |c, _| {
4228 Ok((0..n_expert)
4229 .map(|ex| {
4230 [PROJ_GATE, PROJ_UP, PROJ_DOWN].into_iter().all(|p| {
4231 c.resident(BlockId::new(il, p, ex as u16)).is_some()
4232 })
4233 })
4234 .collect())
4235 })?;
4236
4237 struct Group {
4238 rows: Vec<i32>,
4239 slots: Vec<i32>,
4240 weights: Vec<f32>,
4241 }
4242 let mut groups: std::collections::HashMap<usize, Group> = Default::default();
4243 let mut cpu_rows: Vec<Vec<(usize, f32)>> = vec![Vec::new(); mrows];
4244 let mut cpu_by_expert: std::collections::HashMap<usize, Vec<(usize, f32)>> =
4245 Default::default();
4246 for row in 0..mrows {
4247 for j in 0..n_used {
4248 let ex = sel_all[row * n_used + j] as usize;
4249 let w = w_all[row * n_used + j];
4250 if resident_expert[ex] {
4251 let group = groups.entry(ex).or_insert_with(|| Group {
4252 rows: Vec::new(),
4253 slots: Vec::new(),
4254 weights: Vec::new(),
4255 });
4256 group.rows.push(row as i32);
4257 group.slots.push(j as i32);
4258 group.weights.push(w);
4259 } else {
4260 crate::cpu_experts::record_incomplete_gpu_residency(0);
4261 cpu_rows[row].push((ex, w));
4262 cpu_by_expert.entry(ex).or_default().push((row, w));
4263 }
4264 }
4265 }
4266
4267 let host_rows = e.dtoh(zbatch)?;
4273 let rows_ok = crate::cpu_experts::rows_supported();
4274 enum CpuPart {
4275 Single { row: usize },
4276 Rows { rows: Vec<usize> },
4277 }
4278 let mut tickets: Vec<(CpuPart, crate::cpu_experts::CpuExpertTicket)> = Vec::new();
4279 let mut rows_served: std::collections::HashSet<(usize, usize)> = Default::default();
4280 if rows_ok {
4281 let mut shared: Vec<(usize, Vec<(usize, f32)>)> = cpu_by_expert
4282 .into_iter()
4283 .filter(|(_, rows)| rows.len() >= 2)
4284 .collect();
4285 shared.sort_by_key(|(ex, _)| *ex);
4286 for (ex, mut row_weights) in shared {
4287 row_weights.sort_by_key(|(row, _)| *row);
4288 let inputs: Vec<(&[f32], f32)> = row_weights
4289 .iter()
4290 .map(|&(row, w)| (&host_rows[row * n_embd..(row + 1) * n_embd], w))
4291 .collect();
4292 let job = crate::cpu_experts::prepare_rows_job(m, ex, &inputs)
4293 .map_err(std::io::Error::other)?;
4294 for &(row, _) in &row_weights {
4295 rows_served.insert((row, ex));
4296 }
4297 tickets.push((
4298 CpuPart::Rows {
4299 rows: row_weights.iter().map(|&(row, _)| row).collect(),
4300 },
4301 crate::cpu_experts::submit_rows(job).map_err(std::io::Error::other)?,
4302 ));
4303 }
4304 }
4305 for (row, selected) in cpu_rows.iter().enumerate() {
4306 let leftover: Vec<(usize, f32)> = selected
4307 .iter()
4308 .copied()
4309 .filter(|&(ex, _)| !rows_served.contains(&(row, ex)))
4310 .collect();
4311 if leftover.is_empty() {
4312 continue;
4313 }
4314 let host_row = &host_rows[row * n_embd..(row + 1) * n_embd];
4315 let job = crate::cpu_experts::prepare_job(m, il, &leftover, host_row)
4316 .map_err(std::io::Error::other)?;
4317 tickets.push((
4318 CpuPart::Single { row },
4319 crate::cpu_experts::submit(job).map_err(std::io::Error::other)?,
4320 ));
4321 }
4322
4323 let mut slot_buf = e.zeros(mrows * n_used * n_embd)?;
4324 let mut wbuf = e.zeros(mrows * n_used)?;
4325 let mut order: Vec<usize> = groups.keys().copied().collect();
4326 order.sort_by(|&a, &b| {
4327 groups[&b].rows.len().cmp(&groups[&a].rows.len()).then(a.cmp(&b))
4328 });
4329 for &ex in &order {
4330 let group = &groups[&ex];
4331 let m_e = group.rows.len();
4332 let gl = m.gate_exps.expert_layout(ex);
4333 let ul = m.up_exps.expert_layout(ex);
4334 let dl = m.down_exps.expert_layout(ex);
4335 let row_idx_d = e.htod_i32(&group.rows)?;
4336 let slot_idx_d = e.htod_i32(&group.slots)?;
4337 let dmac = m.down_exps.macro_scale(ex);
4338 let weight_d = if dmac == 1.0 {
4339 e.htod(&group.weights)?
4340 } else {
4341 let scaled: Vec<f32> = group.weights.iter().map(|&w| w * dmac).collect();
4342 e.htod(&scaled)?
4343 };
4344 let mut gathered = e.zeros(m_e * n_embd)?;
4345 e.gather_rows(zbatch, &row_idx_d, &mut gathered, n_embd, m_e)?;
4346 let gv = gathered.slice(0..m_e * n_embd);
4347 let gate = e.with_moe_cache(max_block, |c, eng| {
4348 let slot = c
4349 .resident(BlockId::new(il, PROJ_GATE, ex as u16))
4350 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
4351 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..gl.len, &gv, m_e,
4352 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)
4353 })?;
4354 let up = e.with_moe_cache(max_block, |c, eng| {
4355 let slot = c
4356 .resident(BlockId::new(il, PROJ_UP, ex as u16))
4357 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
4358 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..ul.len, &gv, m_e,
4359 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)
4360 })?;
4361 let mut act = e.zeros(m_e * n_ff_exp)?;
4362 Self::ffn_act_scaled(e, cfg, &gate, &up,
4363 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, m_e * n_ff_exp)?;
4364 let actv = act.slice(0..m_e * n_ff_exp);
4365 let y = e.with_moe_cache(max_block, |c, eng| {
4366 let slot = c
4367 .resident(BlockId::new(il, PROJ_DOWN, ex as u16))
4368 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
4369 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..dl.len, &actv, m_e,
4370 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)
4371 })?;
4372 e.scatter_slot(&y, &row_idx_d, &slot_idx_d, &weight_d,
4373 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
4374 }
4375 let mut moe_out = e.zeros(mrows * n_embd)?;
4376 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, mrows)?;
4377
4378 let mut row_sums: Vec<Option<Vec<f32>>> = vec![None; mrows];
4380 for (part, ticket) in tickets {
4381 let cpu_output = ticket.wait().map_err(std::io::Error::other)?;
4382 let mut add_row = |row: usize, chunk: &[f32]| {
4383 let sum = row_sums[row].get_or_insert_with(|| vec![0.0f32; n_embd]);
4384 for (accumulator, value) in sum.iter_mut().zip(chunk) {
4385 *accumulator += value;
4386 }
4387 };
4388 match part {
4389 CpuPart::Single { row } => add_row(row, &cpu_output),
4390 CpuPart::Rows { rows } => {
4391 for (slot, row) in rows.into_iter().enumerate() {
4392 add_row(row, &cpu_output[slot * n_embd..(slot + 1) * n_embd]);
4393 }
4394 }
4395 }
4396 }
4397 for (row, sum) in row_sums.into_iter().enumerate() {
4398 let Some(sum) = sum else { continue };
4399 let cpu_output = e.htod(&sum)?;
4400 let mut dst = moe_out.slice_mut(row * n_embd..(row + 1) * n_embd);
4401 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
4402 }
4403
4404 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4405 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4406 {
4407 let n_ff_sh = gate_shexp.out_features();
4408 let sg_gate = e.matmul(gate_shexp, zbatch, mrows)?;
4409 let sg_up = e.matmul(up_shexp, zbatch, mrows)?;
4410 let mut sa = e.zeros(mrows * n_ff_sh)?;
4411 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, mrows * n_ff_sh)?;
4412 let sh = e.matmul(down_shexp, &sa, mrows)?;
4413 let g = match &m.gate_inp_shexp {
4416 Some(gate_inp_shexp) => {
4417 e.sigmoid_dot_rows(zbatch, gate_inp_shexp.float_data(), n_embd, mrows)?
4418 }
4419 None => e.htod(&vec![1.0f32; mrows])?,
4420 };
4421 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, mrows)?;
4422 }
4423
4424 Ok(moe_out)
4425 }
4426}
4427
4428impl HybridModel {
4434 pub(crate) fn gemma4_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
4436 let g = self.cfg.gemma4.as_ref().unwrap();
4437 let swa = g.swa_pattern[il];
4438 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
4439 (hd, g.head_count_kv[il] as usize, self.cfg.n_head as usize,
4443 if swa { g.rope_base_swa } else { g.rope_base_global },
4444 1.0, swa)
4445 }
4446
4447 fn gemma4_suppress(&self, e: &Engine, ld: &mut CudaSlice<f32>, t: usize)
4451 -> Result<(), Box<dyn std::error::Error>> {
4452 if let Some((ids, n)) = self.gemma4_aux.as_ref().and_then(|a| a.suppress_d.as_ref()) {
4453 e.mask_ids_rows(ld, ids, *n, self.output.out_features(), t)?;
4454 }
4455 Ok(())
4456 }
4457
4458 fn gemma4_attn_prime(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
4463 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize,
4464 cache: Option<&mut Cache>)
4465 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4466 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
4467 let eps = self.cfg.rms_eps;
4468 let aux = self.gemma4_aux.as_ref().unwrap();
4469
4470 e.mmq_act_begin();
4473 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)? };
4478
4479 let mut q = e.uninit(t * nh * hd)?;
4480 let mut k = e.uninit(t * nkv * hd)?;
4481 let mut v = e.uninit(t * nkv * hd)?;
4483 static EMIT: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4487 let emit = t >= 16 && crate::Engine::qkvnorm_w_on_prefill(nh * t + 2 * nkv * t, hd)
4488 && *EMIT.get_or_init(|| std::env::var("MEMRA_FA_EMIT").map(|s| s != "0").unwrap_or(true));
4489 let mut qb = e.alloc_uninit::<u8>(if emit { t * nh * hd * 2 } else { 1 })?;
4490 let mut kb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
4491 let mut vb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
4492 let v_f16 = emit && crate::fa_f16pv_on() && match hd {
4495 512 => true,
4496 256 => swa && crate::faw_hp_on() && nh % 2 == 0 && (nh / nkv) % 2 == 0,
4497 _ => false,
4498 };
4499 if emit {
4500 e.rms_norm_qkv_w4b(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
4501 &aux.ones, &mut q, &mut k, &mut v, &mut vb,
4502 hd, nh * t, nkv * t, eps, v_f16)?;
4503 } else {
4504 e.rms_norm_qkv(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
4505 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t, eps)?;
4506 }
4507
4508 let ff = if swa { None } else {
4509 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
4510 };
4511 if emit {
4512 e.rope_neox2_bf16e(&mut q, &mut k, &mut qb, &mut kb, pos_d, hd, hd, nh, nkv, t,
4513 base, 1.0, ff)?;
4514 } else {
4515 e.rope_neox2(&mut q, &mut k, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff)?;
4516 }
4517
4518 if let Some(cache) = cache {
4519 let kvl = cache.kv[il].as_mut().unwrap();
4520 assert_eq!(kvl.len, 0, "gemma4 prime is fresh-prompt only (v0)");
4521 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
4522 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()))?;
4523 kvl.len += t;
4524 }
4525 let mut attn = e.zeros(t * nh * hd)?;
4526 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
4530 if swa && t > win {
4531 if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
4532 if emit { e.fa_prefill_w_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
4533 scale, true, win, v_f16)?; }
4534 else { e.fa_prefill_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true,
4535 win)?; }
4536 } else {
4537 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
4538 }
4539 } else if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
4540 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
4541 } else if hd == 512 && std::env::var("MEMRA_NOFA").is_err() {
4542 if emit { e.fa_prefill_hd512_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
4543 scale, true, v_f16)?; }
4544 else { e.fa_prefill_hd512(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?; }
4545 } else {
4546 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
4547 }
4548 Ok(e.matmul(&fa.wo, &attn, t)?)
4549 }
4550
4551 fn gemma4_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
4553 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
4554 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4555 self.gemma4_attn_prime(e, fa, il, h, pos_d, t, None)
4556 }
4557
4558 fn gemma4_moe_q8(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
4563 bits: &crate::hybrid::Gemma4MoeBits,
4564 mq: &(CudaSlice<i8>, CudaSlice<f32>),
4565 router_in: &CudaSlice<f32>, t: usize)
4566 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4567 let cfg = &self.cfg;
4568 let moe = cfg.moe.as_ref().unwrap();
4569 let n_embd = cfg.n_embd as usize;
4570 let n_expert = moe.expert_count as usize;
4571 let n_used = moe.expert_used_count as usize;
4572 let n_ff_exp = moe.expert_ff_length as usize;
4573 let logits = if crate::router_kernel_on() {
4577 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
4578 } else {
4579 e.matmul(&m.gate_inp, router_in, t)?
4580 };
4581 let dev = m.dev_exps.as_ref().unwrap();
4582 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
4583 &bits.per_expert_scale_d)?;
4584 let (zq, zd) = mq;
4585 if t == 1 {
4586 let selv = sel_d.slice(0..n_used);
4587 let wv = w_d.slice(0..n_used);
4588 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, zq, zd,
4589 n_embd, n_ff_exp, n_used, n_expert,
4590 m.gate_exps.qtype, m.up_exps.qtype,
4591 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
4592 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
4593 let mut moe_out = e.uninit(n_embd)?;
4594 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
4595 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
4596 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
4597 return Ok(moe_out);
4598 }
4599 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
4600 let act = if csr {
4601 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, zq, zd, t * n_used,
4602 n_embd, n_ff_exp, n_used, n_expert,
4603 m.gate_exps.qtype, m.up_exps.qtype,
4604 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4605 } else {
4606 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, zq, zd, t,
4607 n_embd, n_ff_exp, n_used, n_expert,
4608 m.gate_exps.qtype, m.up_exps.qtype,
4609 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4610 };
4611 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
4612 let mut moe_out = e.uninit(t * n_embd)?;
4613 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
4616 n_ff_exp, n_embd, n_used, n_expert,
4617 m.down_exps.qtype, m.down_exps.row_bytes)?;
4618 Ok(moe_out)
4619 }
4620
4621 fn gemma4_moe(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
4625 bits: &crate::hybrid::Gemma4MoeBits, moe_in: &CudaSlice<f32>,
4626 router_in: &CudaSlice<f32>, t: usize)
4627 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4628 let cfg = &self.cfg;
4629 let moe = cfg.moe.as_ref().unwrap();
4630 let n_embd = cfg.n_embd as usize;
4631 let n_expert = moe.expert_count as usize;
4632 let n_used = moe.expert_used_count as usize;
4633 let n_ff_exp = moe.expert_ff_length as usize;
4634
4635 let logits = if t < PRIME_MIN_T && crate::router_kernel_on() {
4639 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
4640 } else {
4641 e.matmul(&m.gate_inp, router_in, t)?
4642 };
4643
4644 if t < PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
4649 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
4650 && expert_dp4a_supported(m.down_exps.qtype)
4651 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0") {
4652 let dev = m.dev_exps.as_ref().unwrap();
4653 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
4654 &bits.per_expert_scale_d)?;
4655 if t == 1 {
4656 let (zq, zd) = e.quantize_q8_1(moe_in, 1, n_embd)?;
4657 let selv = sel_d.slice(0..n_used);
4658 let wv = w_d.slice(0..n_used);
4659 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, &zq, &zd,
4660 n_embd, n_ff_exp, n_used, n_expert,
4661 m.gate_exps.qtype, m.up_exps.qtype,
4662 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
4663 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
4664 let mut moe_out = e.uninit(n_embd)?;
4665 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
4666 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
4667 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
4668 return Ok(moe_out);
4669 }
4670 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
4675 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
4676 let act = if csr {
4677 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, t * n_used,
4678 n_embd, n_ff_exp, n_used, n_expert,
4679 m.gate_exps.qtype, m.up_exps.qtype,
4680 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4681 } else {
4682 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
4683 n_embd, n_ff_exp, n_used, n_expert,
4684 m.gate_exps.qtype, m.up_exps.qtype,
4685 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4686 };
4687 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
4688 let mut moe_out = e.uninit(t * n_embd)?;
4689 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
4690 n_ff_exp, n_embd, n_used, n_expert,
4691 m.down_exps.qtype, m.down_exps.row_bytes)?;
4692 return Ok(moe_out);
4693 }
4694
4695 let (sel_all, mut w_all) = Self::moe_route(e, &logits, t, n_expert, n_used)?;
4696 for (i, &sx) in sel_all.iter().enumerate() {
4697 w_all[i] *= bits.per_expert_scale[sx as usize];
4698 }
4699
4700 if t >= PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
4704 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
4705 && expert_dp4a_supported(m.down_exps.qtype)
4706 && std::env::var("MEMRA_GEMMA_MOE_PAIRS").as_deref() != Ok("0") {
4707 let dev = m.dev_exps.as_ref().unwrap();
4708 let n_pairs = t * n_used;
4709 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
4710 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
4711 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
4712 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
4713 let pt = e.htod_i32(&pair_tok)?;
4714 let pw = e.htod(&w_all)?;
4715 let toff = e.htod_i32(&tok_off)?;
4716 let tids = e.htod_i32(&tok_ids)?;
4717 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
4718 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
4719 let mut ex_ids: Vec<i32> = Vec::new();
4720 let mut ex_off: Vec<i32> = vec![0];
4721 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
4722 for (ex, list) in by_ex.iter().enumerate() {
4723 if list.is_empty() { continue; }
4724 ex_ids.push(ex as i32);
4725 ex_pairs.extend_from_slice(list);
4726 ex_off.push(ex_pairs.len() as i32);
4727 }
4728 let n_active = ex_ids.len();
4729 let exi = e.htod_i32(&ex_ids)?;
4730 let exo = e.htod_i32(&ex_off)?;
4731 let exp_d = e.htod_i32(&ex_pairs)?;
4732 if crate::moe_f16g_gemma_on()
4740 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
4741 && f16g_proj_ok(m.up_exps.qtype, n_embd)
4742 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp) {
4743 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
4744 let csr_tok_d = e.htod_i32(&csr_tok)?;
4745 let (z_f16, z_s) = e.moe_f16g_act(moe_in, Some(&csr_tok_d), n_embd, n_pairs)?;
4746 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
4747 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4748 m.gate_exps.qtype, m.gate_exps.row_bytes)?;
4749 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
4750 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4751 m.up_exps.qtype, m.up_exps.row_bytes)?;
4752 let act_csr = e.moe_pairs_gelu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
4753 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
4754 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
4755 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
4756 m.down_exps.qtype, m.down_exps.row_bytes)?;
4757 let y_down = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
4758 let mut moe_out = e.uninit(t * n_embd)?;
4759 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4760 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
4761 let scan = |v: &[f32]| v.iter().filter(|x| !x.is_finite()).count();
4762 let (yd, mo) = (e.dtoh(&y_down)?, e.dtoh(&moe_out)?);
4763 eprintln!("[f16g-debug] post-permute bad={} post-scatter bad={}",
4764 scan(&yd), scan(&mo));
4765 }
4766 return Ok(moe_out);
4767 }
4768 let mma = n_embd % 256 == 0
4771 && std::env::var("MEMRA_GEMMA_MOE_MMA").as_deref() != Ok("0");
4772 let (gate, up) = if mma {
4773 let z_scr = e.mmq_iq_quantize_act(moe_in, n_embd, t)?;
4774 (e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4775 n_embd, n_ff_exp, n_active, n_pairs, t,
4776 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
4777 e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4778 n_embd, n_ff_exp, n_active, n_pairs, t,
4779 m.up_exps.qtype, m.up_exps.row_bytes)?)
4780 } else {
4781 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
4782 (e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 0, &exi, &exo, &exp_d, &pt, &zq, &zd,
4783 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
4784 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
4785 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 1, &exi, &exo, &exp_d, &pt, &zq, &zd,
4786 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
4787 m.up_exps.qtype, m.up_exps.row_bytes)?)
4788 };
4789 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
4790 let pself = e.htod_i32(&pair_self)?;
4791 let y_down = if mma {
4803 let in_pad = n_ff_exp.div_ceil(256) * 256;
4804 let a_scr = if crate::moe_fuse_actq_on() {
4805 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 1)?
4806 } else {
4807 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4808 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
4809 };
4810 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
4811 in_pad, n_embd, n_active, n_pairs, n_pairs,
4812 m.down_exps.qtype, m.down_exps.row_bytes)?
4813 } else {
4814 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4815 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
4816 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
4817 n_ff_exp, n_embd, n_expert, n_active, n_pairs,
4818 m.down_exps.qtype, m.down_exps.row_bytes)?
4819 };
4820 let mut moe_out = e.uninit(t * n_embd)?;
4821 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4822 return Ok(moe_out);
4823 }
4824
4825 let g_len = m.gate_exps.expert_stride;
4826 let u_len = m.up_exps.expert_stride;
4827 let d_len = m.down_exps.expert_stride;
4828 let dev = m.dev_exps.as_ref().filter(|d| !d.gu_il);
4832 let (mut sg, mut su, mut sd) = if dev.is_some() { (None, None, None) } else {
4833 (Some(e.alloc_u8_uninit(g_len)?), Some(e.alloc_u8_uninit(u_len)?), Some(e.alloc_u8_uninit(d_len)?))
4834 };
4835 let mut moe_out = e.zeros(t * n_embd)?;
4836 for tok in 0..t {
4837 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
4838 let w = &w_all[tok * n_used..(tok + 1) * n_used];
4839 let zt = moe_in.slice(tok * n_embd..(tok + 1) * n_embd);
4840 for (j, &ex) in sel.iter().enumerate() {
4841 let ex = ex as usize;
4842 let gate = match dev {
4843 Some(d) => e.qmatvec_view(&d.gate, ex * g_len..(ex + 1) * g_len, &zt, 1,
4844 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?,
4845 None => {
4846 let sg = sg.as_mut().unwrap();
4847 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
4848 e.qmatvec_view(sg, 0..g_len, &zt, 1,
4849 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?
4850 }
4851 };
4852 let up = match dev {
4853 Some(d) => e.qmatvec_view(&d.up, ex * u_len..(ex + 1) * u_len, &zt, 1,
4854 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?,
4855 None => {
4856 let su = su.as_mut().unwrap();
4857 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
4858 e.qmatvec_view(su, 0..u_len, &zt, 1,
4859 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?
4860 }
4861 };
4862 let mut act = e.uninit(n_ff_exp)?;
4863 e.gelu_tanh_mul(&gate, &up, &mut act, n_ff_exp)?;
4864 let actv = act.slice(0..n_ff_exp);
4865 let y = match dev {
4866 Some(d) => e.qmatvec_view(&d.down, ex * d_len..(ex + 1) * d_len, &actv, 1,
4867 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?,
4868 None => {
4869 let sd = sd.as_mut().unwrap();
4870 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
4871 e.qmatvec_view(sd, 0..d_len, &actv, 1,
4872 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?
4873 }
4874 };
4875 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4876 e.axpy_into(&y, w[j], &mut dst, n_embd)?;
4877 }
4878 }
4879 Ok(moe_out)
4880 }
4881
4882 fn gemma4_layer(&self, e: &Engine, il: usize, layer: &crate::hybrid::HybridLayer,
4884 x: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
4885 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4886 let n_embd = self.cfg.n_embd as usize;
4887 let eps = self.cfg.rms_eps;
4888
4889 let mut h = e.zeros(t * n_embd)?;
4890 e.rms_norm(x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
4891 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
4892 let o = self.gemma4_attn(e, fa, il, &h, pos_d, t)?;
4893 let mut cur = e.zeros(t * n_embd)?;
4895 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
4896 self.gemma4_layer_tail_add(e, layer, &cur, x, t)
4897 }
4898
4899 fn gemma4_layer_tail_add(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
4903 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
4904 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4905 Ok(self.gemma4_layer_tail_add_n(e, layer, cur, x, t, None)?.0)
4906 }
4907
4908 fn gemma4_layer_tail_add_n(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
4911 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
4912 next_norm: Option<&CudaSlice<f32>>)
4913 -> Result<(CudaSlice<f32>, Option<CudaSlice<f32>>), Box<dyn std::error::Error>> {
4914 let n_embd = self.cfg.n_embd as usize;
4915 let bits = layer.gemma4.as_ref().unwrap();
4916 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
4917 let mut xn = e.uninit(t * n_embd)?;
4918 match next_norm {
4919 Some(w) => {
4920 let mut hn = e.uninit(t * n_embd)?;
4921 e.add_scale_rms_norm(&sn, &attn_out, bits.layer_scale, w, &mut xn, &mut hn,
4922 n_embd, t, self.cfg.rms_eps)?;
4923 Ok((xn, Some(hn)))
4924 }
4925 None => {
4926 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
4927 Ok((xn, None))
4928 }
4929 }
4930 }
4931
4932 fn gemma4_layer_tail_core(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
4935 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
4936 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4937 self.gemma4_layer_tail_core_pn(e, layer, cur, x, t, None, false)
4938 }
4939
4940 fn gemma4_layer_tail_core_pn(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
4947 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
4948 pre_norm: Option<&CudaSlice<f32>>, defer_post_norm: bool)
4949 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4950 let n_embd = self.cfg.n_embd as usize;
4951 let eps = self.cfg.rms_eps;
4952 let bits = layer.gemma4.as_ref().unwrap();
4953
4954 let Some(mbits) = bits.moe_bits.as_ref() else {
4957 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
4958 else { panic!("gemma4 dense layer without Dense ffn") };
4959 let mut attn_out = e.uninit(t * n_embd)?;
4960 let mut zsh = e.uninit(t * n_embd)?;
4961 let mut zpair: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
4964 match pre_norm {
4965 Some(wa) if t == 1 => {
4966 zpair = Some(e.rms_pre_add_rms_norm_q8z(cur, wa, x,
4967 bits.ffn_norm.float_data(),
4968 &mut attn_out, &mut zsh,
4969 n_embd, t, eps)?);
4970 }
4971 Some(wa) => e.rms_pre_add_rms_norm(cur, wa, x, bits.ffn_norm.float_data(),
4972 &mut attn_out, &mut zsh, n_embd, t, eps)?,
4973 None => e.add_rms_norm(cur, x, bits.ffn_norm.float_data(), &mut attn_out,
4974 &mut zsh, n_embd, t, eps)?,
4975 }
4976 let n_ff = ffn_gate.out_features();
4977 let (gate, up) = if t == 1 {
4983 let (zq, zd) = match zpair {
4984 Some(p) => p,
4985 None => e.quantize_q8_1(&zsh, 1, n_embd)?,
4986 };
4987 match e.matmul_q4_fused2(ffn_gate, ffn_up, &zq, &zd)? {
4988 Some(p) => p,
4989 None => (e.matmul_pre(ffn_gate, &zq, &zd, &zsh, 1)?,
4990 e.matmul_pre(ffn_up, &zq, &zd, &zsh, 1)?),
4991 }
4992 } else {
4993 static F2B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4998 let f2b = *F2B.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
4999 let fused = if f2b {
5000 let (zq, zd) = e.quantize_q8_1(&zsh, t, n_embd)?;
5001 e.matmul_q4_fused2_batched(ffn_gate, ffn_up, &zq, &zd, t)?
5002 } else { None };
5003 match fused {
5004 Some(p) => p,
5005 None => {
5006 e.mmq_act_begin();
5008 (e.matmul(ffn_gate, &zsh, t)?, e.matmul(ffn_up, &zsh, t)?)
5009 }
5010 }
5011 };
5012 let mut act = e.uninit(t * n_ff)?;
5013 let f0 = if e.uses_q8_1_fast(ffn_down) {
5016 let upv = e.view(&up, t * n_ff);
5017 let up_all = upv.slice(0..t * n_ff);
5018 let (aq, ad) = e.gelu_tanh_mul_q8_1(&gate, &up_all, &mut act, n_ff, t)?;
5019 e.matmul_pre(ffn_down, &aq, &ad, &act, t)?
5020 } else {
5021 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
5022 e.matmul(ffn_down, &act, t)?
5023 };
5024 if defer_post_norm { return Ok((f0, attn_out)); }
5025 let mut sn = e.uninit(t * n_embd)?;
5026 e.rms_norm(&f0, bits.post_ffw_norm.float_data(), &mut sn, n_embd, t, eps)?;
5027 return Ok((sn, attn_out));
5028 };
5029
5030 assert!(pre_norm.is_none(), "pre-norm fold is dense-entry only");
5031 let mut attn_out = e.uninit(t * n_embd)?;
5036 let mut router_in = e.uninit(t * n_embd)?;
5037 let fast_moe = match &layer.ffn {
5038 crate::hybrid::Ffn::Moe(m) => m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
5039 && expert_dp4a_supported(m.gate_exps.qtype)
5040 && expert_dp4a_supported(m.up_exps.qtype)
5041 && expert_dp4a_supported(m.down_exps.qtype)
5042 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0"),
5043 _ => false,
5044 };
5045 let q8z = t < PRIME_MIN_T && fast_moe;
5046 let (zsh_f32, zsh_q8, moe_q8) = if q8z {
5047 let (z0, m2) = e.add_rms_norm3_q8z(cur, x, bits.ffn_norm.float_data(),
5048 &mbits.router_scale_pre,
5049 mbits.pre_ffw_norm_2.float_data(),
5050 &mut attn_out, &mut router_in, n_embd, t, eps)?;
5051 (None, Some(z0), Some(m2))
5052 } else {
5053 let mut zsh = e.uninit(t * n_embd)?;
5054 let mut moe_in = e.uninit(t * n_embd)?;
5055 e.add_rms_norm3(cur, x, bits.ffn_norm.float_data(), &mbits.router_scale_pre,
5056 mbits.pre_ffw_norm_2.float_data(), &mut attn_out, &mut zsh,
5057 &mut router_in, &mut moe_in, n_embd, t, eps)?;
5058 (Some((zsh, moe_in)), None, None)
5059 };
5060 let attn_out2 = attn_out;
5061 #[allow(unused_variables)]
5062 let attn_out = &attn_out2;
5063 let n_ff = mbits.shared_gate.out_features();
5064 let (gate, up) = if let Some((zq, zd)) = zsh_q8.as_ref() {
5065 if t == 1 {
5066 match e.matmul_q4_fused2(&mbits.shared_gate, &mbits.shared_up, zq, zd)? {
5067 Some(p) => p,
5068 None => {
5069 let h0 = e.zeros(0)?;
5070 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, 1)?,
5071 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, 1)?)
5072 }
5073 }
5074 } else {
5075 let h0 = e.zeros(0)?;
5077 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, t)?,
5078 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, t)?)
5079 }
5080 } else {
5081 let (zsh, _) = zsh_f32.as_ref().unwrap();
5082 (e.matmul(&mbits.shared_gate, zsh, t)?, e.matmul(&mbits.shared_up, zsh, t)?)
5083 };
5084 let mut act = e.uninit(t * n_ff)?;
5085 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
5086 let mlp0 = e.matmul(&mbits.shared_down, &act, t)?;
5087 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else { panic!("gemma4 layer not MoE") };
5088 let moe0 = match (&moe_q8, &zsh_f32) {
5089 (Some(mq), _) => self.gemma4_moe_q8(e, m, mbits, mq, &router_in, t)?,
5090 (None, Some((_, moe_in))) => self.gemma4_moe(e, m, mbits, moe_in, &router_in, t)?,
5091 _ => unreachable!(),
5092 };
5093 let mut mlp = e.uninit(t * n_embd)?;
5095 let mut moe = e.uninit(t * n_embd)?;
5096 e.rms_norm2x(&mlp0, &moe0, mbits.post_ffw_norm_1.float_data(),
5097 mbits.post_ffw_norm_2.float_data(), &mut mlp, &mut moe, n_embd, t, eps)?;
5098
5099 let mut sum = e.uninit(t * n_embd)?;
5102 let mut sn = e.uninit(t * n_embd)?;
5103 e.add_rms_norm(&mlp, &moe, bits.post_ffw_norm.float_data(), &mut sum, &mut sn,
5104 n_embd, t, eps)?;
5105 Ok((sn, attn_out2))
5106 }
5107
5108 fn gemma4_layer_tail_add_nq(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5110 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
5111 next_norm: Option<&CudaSlice<f32>>)
5112 -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>> {
5113 let n_embd = self.cfg.n_embd as usize;
5114 let bits = layer.gemma4.as_ref().unwrap();
5115 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
5116 let mut xn = e.uninit(t * n_embd)?;
5117 match next_norm {
5118 Some(w) => {
5119 let pair = e.add_scale_rms_norm_q8_1(&sn, &attn_out, bits.layer_scale, w, &mut xn,
5120 n_embd, t, self.cfg.rms_eps)?;
5121 Ok((xn, Some(pair)))
5122 }
5123 None => {
5124 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
5125 Ok((xn, None))
5126 }
5127 }
5128 }
5129
5130 fn gemma4_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
5133 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
5134 if self.is_gemma4_e4b() { return self.gemma4_e4b_forward(e, tokens, last_only); }
5137 let n_embd = self.cfg.n_embd as usize;
5138 let t = tokens.len();
5139 let pos: Vec<i32> = (0..t as i32).collect();
5140 let pos_d = e.htod_i32(&pos)?;
5141
5142 let mut x = self.embed(e, tokens)?;
5143 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
5144 let probe = std::env::var("MEMRA_GEMMA_PROBE").is_ok();
5147 let stat = |e: &Engine, x: &CudaSlice<f32>, tag: &str| -> Result<(), Box<dyn std::error::Error>> {
5148 let h = e.dtoh(x)?;
5149 let bad = h.iter().filter(|v| !v.is_finite()).count();
5150 let mx = h.iter().filter(|v| v.is_finite()).fold(0.0f32, |m, v| m.max(v.abs()));
5151 eprintln!("[gemma-probe] {tag}: tok0_first3={:?} bad={bad} max={mx:.3e}", &h[..3]);
5152 Ok(())
5153 };
5154 if probe { stat(e, &x, "embed")?; }
5155 for (il, layer) in self.layers.iter().enumerate() {
5156 x = self.gemma4_layer(e, il, layer, &x, &pos_d, t)?;
5157 if probe { stat(e, &x, &format!("L{il}"))?; }
5158 }
5159 let mut hn = e.zeros(t * n_embd)?;
5160 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, self.cfg.rms_eps)?;
5161 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
5162 let n_vocab = self.output.out_features();
5163 let logits = if last_only {
5164 let hv = e.view(&hn, t * n_embd);
5165 let last_row = hv.slice((t - 1) * n_embd..t * n_embd);
5166 let mut hlast = e.zeros(n_embd)?;
5167 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
5168 let mut ld = e.matmul(&self.output, &hlast, 1)?;
5169 e.softcap(&mut ld, cap, n_vocab)?;
5170 self.gemma4_suppress(e, &mut ld, 1)?;
5171 e.dtoh(&ld)?
5172 } else {
5173 let mut ld = e.matmul(&self.output, &hn, t)?;
5174 e.softcap(&mut ld, cap, t * n_vocab)?;
5175 self.gemma4_suppress(e, &mut ld, t)?;
5176 e.dtoh(&ld)?
5177 };
5178 Ok(logits)
5179 }
5180
5181 pub(crate) fn gemma4_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
5186 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5187 assert_eq!(cache.pos, 0, "gemma4 prime v0 is fresh-prompt only");
5188 let n_embd = self.cfg.n_embd as usize;
5189 let eps = self.cfg.rms_eps;
5190 let t = tokens.len();
5191 let pos: Vec<i32> = (0..t as i32).collect();
5192 let pos_d = e.htod_i32(&pos)?;
5193 let mut x = self.embed(e, tokens)?;
5194 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
5195 for (il, layer) in self.layers.iter().enumerate() {
5196 let mut h = e.zeros(t * n_embd)?;
5197 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
5198 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer not full-attn") };
5199 let o = self.gemma4_attn_prime(e, fa, il, &h, &pos_d, t, Some(cache))?;
5200 let mut cur = e.zeros(t * n_embd)?;
5201 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
5202 x = self.gemma4_layer_tail_add(e, layer, &cur, &x, t)?;
5203 self.dflash_tap(e, cache, il, &x, t)?;
5204 }
5205 cache.pos += t;
5206 let hiddens = e.clone_dtod(&x)?;
5207 let xv = e.view(&x, t * n_embd);
5208 let last_row = xv.slice((t - 1) * n_embd..t * n_embd);
5209 let mut h_seed = e.zeros(n_embd)?;
5210 e.copy_view_into(&mut h_seed, 0, &last_row, n_embd)?;
5211 let mut hn = e.uninit(n_embd)?;
5212 e.rms_norm(&h_seed, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
5213 let mut ld = e.matmul(&self.output, &hn, 1)?;
5214 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
5215 e.softcap(&mut ld, cap, self.output.out_features())?;
5216 self.gemma4_suppress(e, &mut ld, 1)?;
5217 let logits = e.dtoh(&ld)?;
5218 Ok((logits, h_seed, hiddens))
5219 }
5220
5221 fn gemma4_decode_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
5226 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
5227 pos_d: &CudaSlice<i32>, cache: &mut Cache)
5228 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5229 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
5230 let eps = self.cfg.rms_eps;
5231 let aux = self.gemma4_aux.as_ref().unwrap();
5232 let (hq, hdq) = (hq, hdq);
5233 let h0 = e.zeros(0)?;
5234 let h = &h0;
5235 let (q0, k0, v0) = if swa {
5236 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
5237 Some(t3) => t3,
5238 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
5239 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
5240 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?),
5241 }
5242 } else {
5243 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, &hq, &hdq)? {
5244 Some(p) => p,
5245 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
5246 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?),
5247 };
5248 let v0 = e.clone_dtod(&k0)?;
5249 (q0, k0, v0)
5250 };
5251 let mut q = e.uninit(nh * hd)?;
5252 let mut k = e.uninit(nkv * hd)?;
5253 let mut v = e.uninit(nkv * hd)?;
5254 let ff = if swa { None } else {
5257 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
5258 };
5259 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
5260 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
5261 pos_d, nh, nkv, base, 1.0, ff, eps)?;
5262 let kvl = cache.kv[il].as_mut().unwrap();
5263 e.append_kv_quantized(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len,
5264 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()))?;
5265 kvl.len += 1;
5266 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5270 let mut attn = e.uninit(nh * hd)?;
5271 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
5273 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5274 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5275 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5276 let base = kvl.len as i32;
5278 e.i32_set_k(&mut kvl.len_d, base)?;
5279 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1, scale,
5280 kvl.k_tok_bytes, kvl.v_tok_bytes, Some((&kvl.len_d, -1)), false,
5281 false, None)?;
5282 return Ok(e.matmul(&fa.wo, &attn, 1)?);
5283 }
5284 if swa && kvl.len > win && hd == 256
5286 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5287 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5288 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5289 let base = kvl.len as i32;
5290 e.i32_set_k(&mut kvl.len_d, base)?;
5291 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1, 1, scale,
5292 win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
5293 return Ok(e.matmul(&fa.wo, &attn, 1)?);
5294 }
5295 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) } else { (0, kvl.len) };
5296 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
5297 (off_tok + t_kv) * kvl.k_tok_bytes);
5298 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
5299 (off_tok + t_kv) * kvl.v_tok_bytes);
5300 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
5301 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
5302 Ok(e.matmul(&fa.wo, &attn, 1)?)
5303 }
5304
5305 #[allow(clippy::too_many_arguments)]
5312 pub fn gemma4_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
5313 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
5314 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
5315 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>)
5316 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
5317 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
5318 self.gemma4_decode_step_dc_into(e, token_d, pos_d, embd_gpu, embd_qt, embd_rb, cache,
5319 n_vocab, cap_bucket_max, &mut tok_out)?;
5320 Ok(tok_out)
5321 }
5322
5323 #[allow(clippy::too_many_arguments)]
5326 pub fn gemma4_decode_step_dc_into(&self, e: &Engine, token_d: &CudaSlice<u32>,
5327 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
5328 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
5329 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
5330 tok_out: &mut CudaSlice<u32>)
5331 -> Result<(), Box<dyn std::error::Error>> {
5332 let n_embd = self.cfg.n_embd as usize;
5333 let eps = self.cfg.rms_eps;
5334 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
5335 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
5336 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
5337 let n_layers = self.layers.len();
5338 for (il, layer) in self.layers.iter().enumerate() {
5339 let (hq, hdq) = match h_carry.take() {
5340 Some(p) => p,
5341 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
5342 };
5343 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
5344 let o = self.gemma4_decode_attn_dc(e, fa, il, &hq, &hdq, pos_d, cache, cap_bucket_max)?;
5345 let mut cur = e.uninit(n_embd)?;
5346 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
5347 let next_norm = if il + 1 < n_layers {
5348 Some(self.layers[il + 1].attn_norm.float_data())
5349 } else { None };
5350 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
5351 x = xn;
5352 h_carry = hn;
5353 }
5354 let mut hn = e.uninit(n_embd)?;
5355 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
5356 let mut logits = e.matmul(&self.output, &hn, 1)?;
5357 self.gemma4_suppress(e, &mut logits, 1)?; e.argmax_token_device_into(&logits, tok_out, n_vocab)?;
5359 e.inc_seqlen(pos_d)?;
5360 if cap_bucket_max.is_none() { cache.pos += 1; }
5361 Ok(())
5362 }
5363
5364 pub fn g4_dc_slots(&self, e: &Engine) -> Result<G4DcSlots, Box<dyn std::error::Error>> {
5371 let n_embd = self.cfg.n_embd as usize;
5372 let n_vocab = self.output.out_features();
5373 let n_layers = self.layers.len();
5374 let (mut qmax, mut kvmax, mut ffmax) = (0usize, 0usize, 0usize);
5375 for il in 0..n_layers {
5376 let (hd, nkv, nh, _b, _s, _w) = self.gemma4_geom(il);
5377 qmax = qmax.max(nh * hd);
5378 kvmax = kvmax.max(nkv * hd);
5379 if let crate::hybrid::Ffn::Dense { ffn_gate, .. } = &self.layers[il].ffn {
5380 ffmax = ffmax.max(ffn_gate.out_features());
5381 }
5382 }
5383 Ok(G4DcSlots {
5384 x: e.uninit(n_embd)?, xn: e.uninit(n_embd)?, cur: e.uninit(n_embd)?,
5385 hq: e.alloc_i8_uninit(n_embd)?, hd_: e.uninit(n_embd / 32)?,
5386 q0: e.uninit(qmax)?, k0: e.uninit(kvmax)?, v0: e.uninit(kvmax)?,
5387 q: e.uninit(qmax)?, k: e.uninit(kvmax)?, v: e.uninit(kvmax)?,
5388 attn: e.uninit(qmax)?, o: e.uninit(n_embd)?,
5389 attn_out: e.uninit(n_embd)?, zsh: e.uninit(n_embd)?,
5390 zq: e.alloc_i8_uninit(n_embd.max(qmax))?, zd: e.uninit(n_embd.max(qmax) / 32)?,
5393 gate: e.uninit(ffmax)?, up: e.uninit(ffmax)?,
5394 act: e.uninit(ffmax)?, actq: e.alloc_i8_uninit(ffmax)?, actd: e.uninit(ffmax / 32)?,
5395 f0: e.uninit(n_embd)?, sn: e.uninit(n_embd)?,
5396 hn: e.uninit(n_embd)?, logits: e.uninit(n_vocab)?,
5397 })
5398 }
5399
5400 fn g4_matvec_m1_into(&self, e: &Engine, w: &crate::model::GpuTensor,
5403 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, y: &mut CudaSlice<f32>)
5404 -> Result<(), Box<dyn std::error::Error>> {
5405 use crate::model::GpuTensor;
5406 let (bytes, qtype, row_bytes, scale, rp) = match w {
5407 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
5408 (bytes, *qtype, *row_bytes, *scale, *rp),
5409 _ => return Err("g4_matvec_m1_into: non-quant tensor".into()),
5410 };
5411 let (mbytes, mrp) = match w {
5412 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5413 _ => (bytes, rp),
5414 };
5415 e.qmatvec_mmvq_into(mbytes, aq, ad, 1, w.in_features(), w.out_features(),
5416 qtype, row_bytes, scale, mrp, y)
5417 }
5418
5419 #[allow(clippy::too_many_arguments)]
5423 pub fn gemma4_decode_step_dc_slotted(&self, e: &Engine, token_d: &CudaSlice<u32>,
5424 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
5425 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
5426 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
5427 sl: &mut G4DcSlots, tok_out: &mut CudaSlice<u32>,
5428 ring: Option<(&mut CudaSlice<u32>, usize)>)
5429 -> Result<(), Box<dyn std::error::Error>> {
5430 let n_embd = self.cfg.n_embd as usize;
5431 let eps = self.cfg.rms_eps;
5432 e.embed_gather_device_into(embd_gpu, token_d, &mut sl.x, n_embd, embd_qt, embd_rb)?;
5433 e.scale_inplace(&mut sl.x, (n_embd as f32).sqrt(), n_embd)?;
5434 let n_layers = self.layers.len();
5435 let mut has_carry = false;
5436 for il in 0..n_layers {
5437 if !has_carry {
5438 e.rms_norm_q8_1_into(&sl.x, self.layers[il].attn_norm.float_data(), n_embd, 1,
5439 eps, &mut sl.hq, &mut sl.hd_)?;
5440 }
5441 has_carry = true;
5442 let layer = &self.layers[il];
5443 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
5444 self.gemma4_decode_attn_dc_slotted(e, fa, il, pos_d, cache, cap_bucket_max, sl)?;
5445 e.rms_norm(&sl.o, layer.post_attn_norm.float_data(), &mut sl.cur, n_embd, 1, eps)?;
5446 let next_norm = if il + 1 < n_layers {
5447 Some(self.layers[il + 1].attn_norm.float_data())
5448 } else { None };
5449 self.gemma4_layer_tail_slotted(e, layer, next_norm, sl)?;
5450 std::mem::swap(&mut sl.x, &mut sl.xn);
5451 }
5452 e.rms_norm(&sl.x, self.output_norm.float_data(), &mut sl.hn, n_embd, 1, eps)?;
5453 e.quantize_q8_1_into(&sl.hn, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
5454 {
5456 let (zq, zd) = (&sl.zq, &sl.zd);
5457 let zq = unsafe { &*(zq as *const CudaSlice<i8>) };
5458 let zd = unsafe { &*(zd as *const CudaSlice<f32>) };
5459 self.g4_matvec_m1_into(e, &self.output, zq, zd, &mut sl.logits)?;
5460 }
5461 self.gemma4_suppress(e, &mut sl.logits, 1)?;
5462 e.argmax_token_device_into(&sl.logits, tok_out, n_vocab)?;
5463 if let Some((ring, base)) = ring {
5464 e.plain_tok_ring(tok_out, pos_d, base, ring)?;
5468 }
5469 e.inc_seqlen(pos_d)?;
5470 if cap_bucket_max.is_none() { cache.pos += 1; }
5471 Ok(())
5472 }
5473
5474 #[allow(clippy::too_many_arguments)]
5476 fn gemma4_decode_attn_dc_slotted(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer,
5477 il: usize, pos_d: &CudaSlice<i32>, cache: &mut Cache,
5478 cap_bucket_max: Option<(usize, usize)>, sl: &mut G4DcSlots)
5479 -> Result<(), Box<dyn std::error::Error>> {
5480 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
5481 let eps = self.cfg.rms_eps;
5482 let aux = self.gemma4_aux.as_ref().unwrap();
5483 {
5484 let hq = unsafe { &*(&sl.hq as *const CudaSlice<i8>) };
5485 let hdq = unsafe { &*(&sl.hd_ as *const CudaSlice<f32>) };
5486 if swa {
5487 if !e.matmul_q4_fused3_into(&fa.wq, &fa.wk, &fa.wv, hq, hdq,
5488 &mut sl.q0, &mut sl.k0, &mut sl.v0)? {
5489 return Err("slotted step: fused3 unavailable (non-uniform trunk)".into());
5490 }
5491 } else {
5492 if !e.matmul_q4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)? {
5493 return Err("slotted step: fused2 unavailable".into());
5494 }
5495 let k0r = unsafe { &*(&sl.k0 as *const CudaSlice<f32>) };
5496 e.copy_into(&mut sl.v0, 0, k0r, nkv * hd)?;
5497 }
5498 }
5499 let ff = if swa { None } else {
5502 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
5503 };
5504 let kvl = cache.kv[il].as_mut().unwrap();
5505 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
5506 if crate::Engine::qkv_append_on() {
5507 e.rms_norm_qkv_rope_append_dc(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(),
5509 fa.k_norm.float_data(), &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
5510 pos_d, nh, nkv, base, 1.0, ff, eps,
5511 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
5512 } else {
5513 e.rms_norm_qkv_rope(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
5514 &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
5515 pos_d, nh, nkv, base, 1.0, ff, eps)?;
5516 e.append_kv_quantized_dc(&sl.k, &sl.v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
5517 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
5518 kv_fp8)?;
5519 }
5520 e.inc_seqlen(&mut kvl.len_d)?;
5521 let (b_swa, b_glob) = cap_bucket_max.expect("slotted step is capture-only");
5522 let k_view = e.view_u8(&kvl.k, kvl.k.len());
5523 let v_view = e.view_u8(&kvl.v, kvl.v.len());
5524 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
5525 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5526 let mut fa_q8 = false;
5530 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
5531 e.fa_decode_rows(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, b_glob - 1,
5532 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5533 Some((&kvl.len_d, -1)), false, false,
5534 Some((&mut sl.zq, &mut sl.zd)))?;
5535 fa_q8 = true;
5536 } else if swa && b_swa > win && hd == 256 && rows_on {
5537 e.fa_decode_rows_w(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv,
5538 &kvl.len_d, -1, 1, scale, win,
5539 kvl.k_tok_bytes, kvl.v_tok_bytes,
5540 Some((&mut sl.zq, &mut sl.zd)))?;
5541 fa_q8 = true;
5542 } else {
5543 let b = if swa { b_swa } else { b_glob };
5544 e.fa_decode_dc(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, &kvl.len_d, b,
5545 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5546 swa && crate::Engine::wkv_on())?;
5547 }
5548 if !fa_q8 {
5549 let aq = unsafe { &*(&sl.attn as *const CudaSlice<f32>) };
5550 e.quantize_q8_1_into(aq, 1, nh * hd, &mut sl.zq, &mut sl.zd)?;
5551 }
5552 {
5553 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
5554 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
5555 self.g4_matvec_m1_into(e, &fa.wo, zq, zd, &mut sl.o)?;
5556 }
5557 Ok(())
5558 }
5559
5560 fn gemma4_layer_tail_slotted(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5563 next_norm: Option<&CudaSlice<f32>>, sl: &mut G4DcSlots)
5564 -> Result<(), Box<dyn std::error::Error>> {
5565 let n_embd = self.cfg.n_embd as usize;
5566 let eps = self.cfg.rms_eps;
5567 let bits = layer.gemma4.as_ref().unwrap();
5568 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
5569 else { return Err("slotted tail: dense ffn only".into()) };
5570 e.add_rms_norm(&sl.cur, &sl.x, bits.ffn_norm.float_data(), &mut sl.attn_out,
5571 &mut sl.zsh, n_embd, 1, eps)?;
5572 let n_ff = ffn_gate.out_features();
5573 {
5574 let zshr = unsafe { &*(&sl.zsh as *const CudaSlice<f32>) };
5575 e.quantize_q8_1_into(zshr, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
5576 }
5577 {
5578 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
5579 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
5580 if !e.matmul_q4_fused2_into(ffn_gate, ffn_up, zq, zd, &mut sl.gate, &mut sl.up)? {
5581 return Err("slotted tail: ffn fused2 unavailable".into());
5582 }
5583 }
5584 debug_assert!(e.uses_q8_1_fast(ffn_down));
5585 {
5586 let upr = unsafe { &*(&sl.up as *const CudaSlice<f32>) };
5587 let upv = e.view(upr, n_ff);
5588 let up_all = upv.slice(0..n_ff);
5589 let gr = unsafe { &*(&sl.gate as *const CudaSlice<f32>) };
5590 e.gelu_tanh_mul_q8_1_into(gr, &up_all, &mut sl.act, n_ff, 1,
5591 &mut sl.actq, &mut sl.actd)?;
5592 }
5593 {
5594 let aq = unsafe { &*(&sl.actq as *const CudaSlice<i8>) };
5595 let ad = unsafe { &*(&sl.actd as *const CudaSlice<f32>) };
5596 self.g4_matvec_m1_into(e, ffn_down, aq, ad, &mut sl.f0)?;
5597 }
5598 e.rms_norm(&sl.f0, bits.post_ffw_norm.float_data(), &mut sl.sn, n_embd, 1, eps)?;
5599 match next_norm {
5600 Some(w) => {
5601 e.add_scale_rms_norm_q8_1_into(&sl.sn, &sl.attn_out, bits.layer_scale, w,
5602 &mut sl.xn, n_embd, 1, eps,
5603 &mut sl.hq, &mut sl.hd_)?;
5604 }
5605 None => {
5606 e.add_scale(&sl.sn, &sl.attn_out, bits.layer_scale, &mut sl.xn, n_embd)?;
5607 }
5608 }
5609 Ok(())
5610 }
5611
5612 #[allow(clippy::too_many_arguments)]
5614 fn gemma4_decode_attn_dc(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
5615 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
5616 pos_d: &CudaSlice<i32>, cache: &mut Cache,
5617 cap_bucket_max: Option<(usize, usize)>)
5618 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5619 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
5620 let eps = self.cfg.rms_eps;
5621 let aux = self.gemma4_aux.as_ref().unwrap();
5622 let (q0, k0, v0) = if swa {
5623 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
5624 Some(t3) => t3,
5625 None => {
5626 let h0 = e.zeros(0)?;
5627 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
5628 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
5629 e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?)
5630 }
5631 }
5632 } else {
5633 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
5634 Some(p) => p,
5635 None => {
5636 let h0 = e.zeros(0)?;
5637 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
5638 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?)
5639 }
5640 };
5641 let v0 = e.clone_dtod(&k0)?;
5642 (q0, k0, v0)
5643 };
5644 let mut q = e.uninit(nh * hd)?;
5645 let mut k = e.uninit(nkv * hd)?;
5646 let mut v = e.uninit(nkv * hd)?;
5647 let ff = if swa { None } else {
5649 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
5650 };
5651 let kvl = cache.kv[il].as_mut().unwrap();
5652 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
5653 if crate::Engine::qkv_append_on() {
5654 e.rms_norm_qkv_rope_append_dc(&q0, &k0, &v0, fa.q_norm.float_data(),
5656 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
5657 pos_d, nh, nkv, base, 1.0, ff, eps,
5658 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
5659 } else {
5660 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
5661 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
5662 pos_d, nh, nkv, base, 1.0, ff, eps)?;
5663 e.append_kv_quantized_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
5664 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
5665 }
5666 e.inc_seqlen(&mut kvl.len_d)?;
5667 let mut attn = e.uninit(nh * hd)?;
5668 let mut fa_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
5671 match cap_bucket_max {
5676 None => {
5677 kvl.len += 1;
5681 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5682 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
5683 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5684 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5687 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5688 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5689 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1,
5690 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5691 Some((&kvl.len_d, -1)), false, false,
5692 Some((&mut aq8, &mut ad8)))?;
5693 fa_q8 = Some((aq8, ad8));
5694 } else if swa && kvl.len > win && hd == 256
5695 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5696 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5698 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5699 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5700 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1,
5701 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes,
5702 Some((&mut aq8, &mut ad8)))?;
5703 fa_q8 = Some((aq8, ad8));
5704 } else {
5705 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) }
5706 else { (0, kvl.len) };
5707 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
5708 (off_tok + t_kv) * kvl.k_tok_bytes);
5709 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
5710 (off_tok + t_kv) * kvl.v_tok_bytes);
5711 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
5712 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
5713 }
5714 }
5715 Some((b_swa, b_glob)) => {
5716 let k_view = e.view_u8(&kvl.k, kvl.k.len());
5722 let v_view = e.view_u8(&kvl.v, kvl.v.len());
5723 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
5724 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5725 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
5726 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5727 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, b_glob - 1,
5728 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5729 Some((&kvl.len_d, -1)), false, false,
5730 Some((&mut aq8, &mut ad8)))?;
5731 fa_q8 = Some((aq8, ad8));
5732 } else if swa && b_swa > win && hd == 256 && rows_on {
5733 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5734 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
5735 &kvl.len_d, -1, 1, scale, win,
5736 kvl.k_tok_bytes, kvl.v_tok_bytes,
5737 Some((&mut aq8, &mut ad8)))?;
5738 fa_q8 = Some((aq8, ad8));
5739 } else {
5740 let b = if swa { b_swa } else { b_glob };
5741 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, b,
5742 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5743 swa && crate::Engine::wkv_on())?;
5744 }
5745 }
5746 }
5747 if let Some((aq8, ad8)) = fa_q8 {
5750 let mut y = e.uninit(fa.wo.out_features())?;
5751 self.g4_matvec_m1_into(e, &fa.wo, &aq8, &ad8, &mut y)?;
5752 return Ok(y);
5753 }
5754 Ok(e.matmul(&fa.wo, &attn, 1)?)
5755 }
5756
5757 pub fn gemma4_generate_graph(&self, e: &Engine, prompt_pos: usize, first_token: u32,
5762 cache: &mut Cache, max_new: usize, eos: &[u32],
5763 mut on_token: impl FnMut(u32) -> bool)
5764 -> Result<(Vec<u32>, crate::decode::StopReason), Box<dyn std::error::Error>> {
5765 if self.is_gemma4_e4b() {
5766 return Err("E4B graph serving is unwired (HANDOVER-E4B.md) — dc-eager is the serving arm".into());
5767 }
5768 use crate::decode::StopReason;
5769 let n_vocab = self.output.out_features();
5770 let n_embd = self.cfg.n_embd as usize;
5771 let embd_gpu = self.embd_gpu.get_or_init(|| {
5772 e.upload_u8(&self.embd.raw).expect("embed table upload")
5773 });
5774 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
5775 for kvl in cache.kv.iter_mut().flatten() {
5776 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
5777 }
5778 let mut token_d = e.stream().clone_htod(&[first_token])?;
5779 let mut pos_d = e.htod_i32(&[prompt_pos as i32])?;
5780 let g4 = self.cfg.gemma4.as_ref().unwrap();
5781 let (hd_s, hd_g) = (g4.key_length_swa as usize, g4.key_length_global as usize);
5782 let nkv_s = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
5784 .find(|p| *p.1).map(|p| *p.0 as usize).unwrap_or(8);
5785 let nkv_g = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
5786 .find(|p| !*p.1).map(|p| *p.0 as usize).unwrap_or(2);
5787 let mut graphs: std::collections::HashMap<((bool, usize), (bool, usize), bool, bool),
5788 (cudarc::driver::CudaGraph,
5789 Vec<Box<dyn std::any::Any + Send>>)> = Default::default();
5790 let mut slots = self.g4_dc_slots(e)?;
5793 const RING: usize = 64;
5796 const DRAIN: usize = 1;
5802 let mut ring = e.stream().alloc_zeros::<u32>(RING)?;
5803 let ring_base = prompt_pos;
5804 let mut out = Vec::with_capacity(max_new);
5805 let mut reason = StopReason::MaxNew;
5806 let mut next = first_token;
5807 let mut captures = 0usize;
5808 for _ in 0..max_new {
5809 out.push(next);
5810 if eos.contains(&next) { reason = StopReason::Eos; break; }
5811 if !on_token(next) { reason = StopReason::Callback; break; }
5812 let t_kv = cache.pos + 1;
5813 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5821 let f512 = crate::fa512_min_tkv();
5822 let key_s = if t_kv > win { (true, usize::MAX) }
5823 else { e.fa_bucket_key(t_kv, hd_s, nkv_s, crate::Engine::wkv_on()) };
5824 let (key_g, rung_end) = if t_kv >= f512 {
5825 let end = (t_kv + 1).next_power_of_two().max(f512 * 2);
5828 ((true, end), end)
5829 } else { (e.fa_bucket_key(t_kv, hd_g, nkv_g, false), t_kv) };
5830 let key = (key_s, key_g, t_kv >= f512, t_kv > win);
5831 if !graphs.contains_key(&key) {
5832 let bucket_max = (t_kv, rung_end);
5833 let snap = cache.snapshot(e)?;
5835 let pos_save = e.dtoh_i32_one(&pos_d)?;
5836 let len_save: Vec<Option<i32>> = cache.kv.iter()
5837 .map(|k| k.as_ref().map(|kvl| e.dtoh_i32_one(&kvl.len_d).unwrap())).collect();
5838 let tok_save = e.dtoh_u32_one(&token_d)?;
5839 let graph = {
5844 let tok_ref = &mut token_d;
5845 let pos_ref = &mut pos_d;
5846 let cache_ref = &mut *cache;
5847 let slots_ref = &mut slots;
5848 let ring_ref = &mut ring;
5849 e.capture_graph_retained_flags(
5850 cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
5851 |e| {
5852 let tok_in = unsafe { &*(tok_ref as *const CudaSlice<u32>) };
5854 let sl = unsafe { &mut *(slots_ref as *mut G4DcSlots) };
5855 let rg = unsafe { &mut *(ring_ref as *mut CudaSlice<u32>) };
5856 self.gemma4_decode_step_dc_slotted(e, tok_in, pos_ref, embd_gpu, qt, rb,
5857 cache_ref, n_vocab, Some(bucket_max),
5858 sl, tok_ref, Some((rg, ring_base)))
5859 })?
5860 };
5861 cache.rollback(e, &snap, 0)?;
5862 e.set_i32_one(&mut pos_d, pos_save)?;
5863 for (il, ls) in len_save.iter().enumerate() {
5864 if let (Some(kvl), Some(v)) = (cache.kv[il].as_mut(), ls) {
5865 e.set_i32_one(&mut kvl.len_d, *v)?;
5866 }
5867 }
5868 e.set_u32_one(&mut token_d, tok_save)?;
5869 if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1") {
5870 if let Ok(c) = crate::graph_update::node_census(&graph.0) {
5871 eprintln!("[graph-census] {c:?}");
5872 }
5873 }
5874 graphs.insert(key, graph);
5875 captures += 1;
5876 }
5877 let mut chunk = 1usize;
5882 let drain_cap: usize = std::env::var("MEMRA_GRAPH_DRAIN").ok()
5883 .and_then(|v| v.parse().ok()).unwrap_or(DRAIN);
5884 while chunk < drain_cap && out.len() + chunk < max_new {
5885 let t_next = cache.pos + 1 + chunk;
5886 let key_s2 = if t_next > win { (true, usize::MAX) }
5887 else { e.fa_bucket_key(t_next, hd_s, nkv_s, crate::Engine::wkv_on()) };
5888 let key_g2 = if t_next >= f512 {
5889 (true, (t_next + 1).next_power_of_two().max(f512 * 2))
5890 } else { e.fa_bucket_key(t_next, hd_g, nkv_g, false) };
5891 if (key_s2, key_g2, t_next >= f512, t_next > win) != key { break; }
5892 chunk += 1;
5893 }
5894 let g = &graphs.get(&key).unwrap().0;
5895 for _ in 0..chunk { g.launch()?; }
5896 e.stream().synchronize()?;
5897 let ringh = e.dtoh_u32(&ring)?;
5898 for j in 0..chunk {
5899 let pos_j = cache.pos + j;
5900 let tok_j = ringh[(pos_j - ring_base) % RING];
5901 cache.pos += 0; if j + 1 == chunk { next = tok_j; }
5903 else {
5904 out.push(tok_j);
5905 if eos.contains(&tok_j) || !on_token(tok_j) {
5906 reason = if eos.contains(&tok_j) { StopReason::Eos }
5907 else { StopReason::Callback };
5908 let keep = cache.pos + j + 1;
5910 e.set_i32_one(&mut pos_d, keep as i32)?;
5911 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
5912 e.set_i32_one(&mut kvl.len_d, keep as i32)?;
5913 kvl.len = keep;
5914 }
5915 cache.pos = keep;
5916 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
5917 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
5918 }
5919 return Ok((out, reason));
5920 }
5921 }
5922 }
5923 cache.pos += chunk;
5924 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) { kvl.len += chunk; }
5925 }
5926 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
5927 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
5928 }
5929 Ok((out, reason))
5930 }
5931
5932 pub(crate) fn gemma4_decode_step_t(&self, e: &Engine, tokens: &[u32], pos0: usize,
5938 cache: &mut Cache)
5939 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
5940 Ok(self.gemma4_decode_step_t_h(e, tokens, pos0, cache)?.0)
5941 }
5942
5943 pub(crate) fn gemma4_decode_step_t_am(&self, e: &Engine, tokens: &[u32], pos0: usize,
5947 cache: &mut Cache)
5948 -> Result<(Vec<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5949 let (ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
5950 let t = tokens.len();
5951 let n_vocab = self.output.out_features();
5952 let mut toks = e.stream().alloc_zeros::<u32>(t)?;
5953 for i in 0..t {
5954 e.argmax_token_device_col(&ld, i, n_vocab, &mut toks, i)?;
5955 }
5956 Ok((e.dtoh_u32(&toks)?, hn))
5957 }
5958
5959 pub(crate) fn gemma4_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
5962 pos0: usize, cache: &mut Cache)
5963 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5964 let (ld, hn) = self.gemma4_verify_trunk(e, &vec![0u32; t], pos0, cache, Some(tok_d))?;
5965 let n_vocab = self.output.out_features();
5966 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
5967 for i in 0..t {
5968 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
5969 }
5970 Ok((vam, hn))
5971 }
5972
5973 pub(crate) fn gemma4_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
5976 cache: &mut Cache)
5977 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5978 let (mut ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
5979 let t = tokens.len();
5980 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
5981 e.softcap(&mut ld, cap, t * self.output.out_features())?;
5982 Ok((e.dtoh(&ld)?, hn))
5983 }
5984
5985 pub(crate) fn verify_stream_scratch(&self, e: &Engine, cap: usize)
5988 -> Result<VerifyStreamScratch, Box<dyn std::error::Error>> {
5989 Ok(VerifyStreamScratch {
5990 pos_d: e.htod_i32(&vec![0i32; cap])?,
5991 row_ctrs: (0..cap).map(|_| e.htod_i32(&[0])).collect::<Result<_, _>>()?,
5992 })
5993 }
5994
5995 pub(crate) fn gemma4_verify_t_am_stream(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
6003 ctr: &CudaSlice<i32>, hint: usize,
6004 cache: &mut Cache,
6005 scr: &mut VerifyStreamScratch)
6006 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6007 let n_embd = self.cfg.n_embd as usize;
6008 let eps = self.cfg.rms_eps;
6009 assert!(t <= scr.row_ctrs.len() && t <= 64);
6010 e.i32_iota_from(ctr, &mut scr.pos_d, t)?;
6011 for i in 0..t {
6012 e.i32_copy_add(ctr, &mut scr.row_ctrs[i], (i + 1) as i32)?;
6013 }
6014 let (pos_d, row_ctrs) = (&scr.pos_d, &scr.row_ctrs);
6015 let embd_gpu = self.embd_gpu.get_or_init(|| {
6016 e.upload_u8(&self.embd.raw).expect("embed table upload")
6017 });
6018 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
6019 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
6020 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6021 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6022 let n_layers = self.layers.len();
6023 for (il, layer) in self.layers.iter().enumerate() {
6024 let (hq, hdq) = match h_carry.take() {
6025 Some(p) => p,
6026 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
6027 };
6028 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6029 let o = self.gemma4_verify_attn_stream(e, fa, il, &hq, &hdq, pos_d, t, cache,
6030 hint, row_ctrs)?;
6031 let mut cur = e.uninit(t * n_embd)?;
6032 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
6033 let next_norm = if il + 1 < n_layers {
6034 Some(self.layers[il + 1].attn_norm.float_data())
6035 } else { None };
6036 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
6037 x = xn;
6038 h_carry = hn;
6039 self.dflash_tap(e, cache, il, &x, t)?;
6040 }
6041 let mut hn = e.uninit(t * n_embd)?;
6042 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
6043 let ld = e.matmul(&self.output, &hn, t)?;
6044 let n_vocab = self.output.out_features();
6045 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
6046 for i in 0..t {
6047 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
6048 }
6049 Ok((vam, hn))
6050 }
6051
6052 fn dflash_tap(&self, e: &Engine, cache: &mut Cache, il: usize, x: &CudaSlice<f32>, t: usize)
6059 -> Result<(), Box<dyn std::error::Error>> {
6060 let Some(taps) = cache.dflash_taps.as_mut() else { return Ok(()) };
6061 let Some(slot) = taps.layer_ids.iter().position(|&l| l == il) else { return Ok(()) };
6062 let h = taps.hidden;
6063 let n_taps = taps.layer_ids.len();
6064 debug_assert_eq!(taps.t, t);
6065 let xv = e.view(x, t * h);
6066 for r in 0..t {
6067 let row = xv.slice(r * h..(r + 1) * h);
6068 e.copy_view_into(&mut taps.buf, r * n_taps * h + slot * h, &row, h)?;
6069 }
6070 Ok(())
6071 }
6072
6073 fn gemma4_verify_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
6074 tok_dev: Option<&CudaSlice<u32>>)
6075 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6076 let n_embd = self.cfg.n_embd as usize;
6077 let eps = self.cfg.rms_eps;
6078 let t = tokens.len();
6079 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
6080 let pos_d = e.htod_i32(&pos)?;
6081 let mut x = match tok_dev {
6082 Some(td) => {
6083 let embd_gpu = self.embd_gpu.get_or_init(|| {
6084 e.upload_u8(&self.embd.raw).expect("embed table upload")
6085 });
6086 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
6087 e.embed_gather_device_td(embd_gpu, td, t, n_embd, qt, rb)?
6088 }
6089 None => e.htod(&self.embd.gather(n_embd, tokens))?,
6090 };
6091 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6092 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6093 let n_layers = self.layers.len();
6094 for (il, layer) in self.layers.iter().enumerate() {
6095 let (hq, hdq) = match h_carry.take() {
6096 Some(p) => p,
6097 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
6098 };
6099 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6100 let o = self.gemma4_verify_attn(e, fa, il, &hq, &hdq, &pos_d, t, cache)?;
6101 let mut cur = e.uninit(t * n_embd)?;
6102 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
6103 let next_norm = if il + 1 < n_layers {
6104 Some(self.layers[il + 1].attn_norm.float_data())
6105 } else { None };
6106 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
6107 x = xn;
6108 h_carry = hn;
6109 self.dflash_tap(e, cache, il, &x, t)?;
6110 }
6111 let mut hn = e.uninit(t * n_embd)?;
6112 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
6113 let mut ld = e.matmul(&self.output, &hn, t)?;
6114 self.gemma4_suppress(e, &mut ld, t)?; cache.pos += t;
6116 Ok((ld, hn))
6117 }
6118
6119 #[allow(clippy::too_many_arguments)]
6127 fn gemma4_verify_attn_stream(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6128 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
6129 pos_d: &CudaSlice<i32>, t: usize,
6130 cache: &mut Cache, hint: usize,
6131 row_ctrs: &[CudaSlice<i32>])
6132 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6133 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
6134 let eps = self.cfg.rms_eps;
6135 let aux = self.gemma4_aux.as_ref().unwrap();
6136 let h0 = e.zeros(0)?;
6137 let h = &h0;
6138 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6141 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
6142 let fused_qkv = if f2b {
6143 if swa {
6144 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
6145 .map(|(a, b, c)| (a, b, Some(c)))
6146 } else {
6147 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
6148 .map(|(a, b)| (a, b, None))
6149 }
6150 } else { None };
6151 let (q0, k0, v0) = match fused_qkv {
6152 Some((a, b, cv)) => {
6153 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
6154 (a, b, v)
6155 }
6156 None => {
6157 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
6158 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
6159 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
6160 else { e.clone_dtod(&k0)? };
6161 (q0, k0, v0)
6162 }
6163 };
6164 let mut q = e.uninit(t * nh * hd)?;
6165 let mut k = e.uninit(t * nkv * hd)?;
6166 let mut v = e.uninit(t * nkv * hd)?;
6167 let ff = if swa { None } else {
6170 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
6171 };
6172 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6173 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
6174 pos_d, nh, nkv, base, 1.0, ff, eps)?;
6175 let kvl = cache.kv[il].as_mut().unwrap();
6176 e.append_kv_quantized_rows_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d, t,
6178 kvl.kv_dim_k, kvl.kv_dim_v,
6179 kvl.k_tok_bytes, kvl.v_tok_bytes,
6180 (!swa && crate::Engine::gkv_on())
6181 || (swa && crate::Engine::wkv_on()))?;
6182 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6185 let mut attn = e.uninit(t * nh * hd)?;
6186 let k_view = e.view_u8(&kvl.k, kvl.k.len());
6187 let v_view = e.view_u8(&kvl.v, kvl.v.len());
6188 if swa && hint + 1 >= win {
6191 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6194 &kvl.len_d, 0, t, scale, win,
6195 kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
6196 } else if hd == 512 && hint + t < crate::fa512_min_tkv() {
6197 let bucket = (hint + t + 2).next_power_of_two()
6210 .min(crate::fa512_min_tkv().saturating_sub(1));
6211 let qv = e.view(&q, t * nh * hd);
6212 for i in 0..t {
6213 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
6214 let mut q_one = e.uninit(nh * hd)?;
6215 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
6216 let mut a_one = e.uninit(nh * hd)?;
6217 e.fa_decode_dc(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv,
6218 &row_ctrs[i], bucket, scale,
6219 kvl.k_tok_bytes, kvl.v_tok_bytes, false)?;
6220 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
6221 }
6222 } else if hd == 512 {
6223 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, hint, t, scale,
6226 kvl.k_tok_bytes, kvl.v_tok_bytes,
6227 Some((&kvl.len_d, 0)), false, false, None)?;
6228 } else {
6229 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6231 &kvl.len_d, hint + t, t, scale,
6232 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
6233 swa && crate::Engine::wkv_on())?;
6234 }
6235 Ok(e.matmul(&fa.wo, &attn, t)?)
6236 }
6237
6238 fn gemma4_verify_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6239 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
6240 pos_d: &CudaSlice<i32>, t: usize,
6241 cache: &mut Cache)
6242 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6243 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
6244 let eps = self.cfg.rms_eps;
6245 let aux = self.gemma4_aux.as_ref().unwrap();
6246 let n_embd = self.cfg.n_embd as usize;
6247 let _ = n_embd;
6248
6249 let h0 = e.zeros(0)?;
6250 let h = &h0;
6251 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6254 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
6255 let fused_qkv = if f2b {
6256 if swa {
6257 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
6258 .map(|(a, b, c)| (a, b, Some(c)))
6259 } else {
6260 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
6261 .map(|(a, b)| (a, b, None))
6262 }
6263 } else { None };
6264 let (q0, k0, v0) = match fused_qkv {
6265 Some((a, b, cv)) => {
6266 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
6267 (a, b, v)
6268 }
6269 None => {
6270 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
6271 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
6272 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
6273 else { e.clone_dtod(&k0)? };
6274 (q0, k0, v0)
6275 }
6276 };
6277 let mut q = e.uninit(t * nh * hd)?;
6278 let mut k = e.uninit(t * nkv * hd)?;
6279 let mut v = e.uninit(t * nkv * hd)?;
6280 let ff = if swa { None } else {
6283 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
6284 };
6285 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6286 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
6287 pos_d, nh, nkv, base, 1.0, ff, eps)?;
6288 let kvl = cache.kv[il].as_mut().unwrap();
6289 let base_len = kvl.len;
6290 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, base_len, t,
6291 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()))?;
6292 kvl.len += t;
6293 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6294 let mut attn = e.uninit(t * nh * hd)?;
6295 let rows_ok = (hd == 256 && base_len + 1 >= crate::fa_vec_min_tkv())
6298 || (hd == 512 && !swa && base_len + 1 >= crate::fa512_min_tkv());
6301 if rows_ok && (!swa || base_len + t <= win) {
6302 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
6303 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
6304 if hd == 512 {
6305 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
6307 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, base_len, t,
6308 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
6309 Some((&kvl.len_d, 0)), false,
6310 swa && crate::Engine::wkv_on(), None)?;
6311 } else {
6312 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
6316 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6317 &kvl.len_d, base_len + t, t, scale,
6318 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
6319 swa && crate::Engine::wkv_on())?;
6320 }
6321 return Ok(e.matmul(&fa.wo, &attn, t)?);
6322 }
6323 if hd == 256 && swa && base_len + 1 >= win
6331 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
6332 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
6333 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
6334 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
6335 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, 0,
6336 t, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
6337 return Ok(e.matmul(&fa.wo, &attn, t)?);
6338 }
6339 for i in 0..t {
6340 let avail = base_len + i + 1;
6341 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
6342 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
6343 (off_tok + t_kv) * kvl.k_tok_bytes);
6344 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
6345 (off_tok + t_kv) * kvl.v_tok_bytes);
6346 let qi = e.view(&q, t * nh * hd);
6347 let q_row = qi.slice(i * nh * hd..(i + 1) * nh * hd);
6348 let mut q_one = e.uninit(nh * hd)?;
6349 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
6350 let mut a_one = e.uninit(nh * hd)?;
6351 if swa && avail > win && hd == 256
6355 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
6356 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
6357 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
6358 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
6359 e.fa_decode_rows_w(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, &kvl.len_d, 0,
6360 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
6361 } else if !swa && hd == 512 && avail >= crate::fa512_min_tkv()
6362 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
6363 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
6364 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
6365 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
6366 e.fa_decode_rows(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, avail - 1, 1,
6367 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
6368 Some((&kvl.len_d, 0)), false, false, None)?;
6369 } else {
6370 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
6371 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
6372 }
6373 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
6374 }
6375 Ok(e.matmul(&fa.wo, &attn, t)?)
6376 }
6377
6378 pub(crate) fn gemma4_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
6381 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6382 if let Some(split) = crate::pp::pp2_split(self.layers.len()) {
6387 return self.gemma4_decode_step_h_pp2(e, token, cache, split);
6388 }
6389 if crate::pp::pp_cuts(self.layers.len()).is_some() {
6390 crate::pp::warn_unwired_once("gemma4 eager decode (N>2)");
6391 }
6392 let n_embd = self.cfg.n_embd as usize;
6393 let eps = self.cfg.rms_eps;
6394 let pos_d = e.htod_i32(&[cache.pos as i32])?;
6395 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
6396 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
6397 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6400 let n_layers = self.layers.len();
6401 for (il, layer) in self.layers.iter().enumerate() {
6402 let (hq, hdq) = match h_carry.take() {
6403 Some(p) => p,
6404 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
6405 };
6406 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6407 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, &pos_d, cache)?;
6408 let mut cur = e.uninit(n_embd)?;
6409 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
6410 let next_norm = if il + 1 < n_layers {
6411 Some(self.layers[il + 1].attn_norm.float_data())
6412 } else { None };
6413 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
6414 x = xn;
6415 h_carry = hn;
6416 }
6417 let mut hn = e.uninit(n_embd)?;
6418 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
6419 let h_seed = e.clone_dtod(&x)?;
6420 let mut ld = e.matmul(&self.output, &hn, 1)?;
6421 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6422 e.softcap(&mut ld, cap, self.output.out_features())?; self.gemma4_suppress(e, &mut ld, 1)?;
6424 let logits = e.dtoh(&ld)?;
6425 cache.pos += 1;
6426 Ok((logits, h_seed))
6427 }
6428
6429 fn gemma4_decode_layers(&self, e: &Engine, mut x: CudaSlice<f32>, lo: usize, hi: usize,
6437 pos_d: &CudaSlice<i32>, cache: &mut Cache)
6438 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6439 let n_embd = self.cfg.n_embd as usize;
6440 let eps = self.cfg.rms_eps;
6441 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6442 for il in lo..hi {
6443 let layer = &self.layers[il];
6444 let (hq, hdq) = match h_carry.take() {
6445 Some(p) => p,
6446 None => e.rms_norm_q8_1(&x, self.layers[il].attn_norm.float_data(), n_embd, 1, eps)?,
6448 };
6449 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6450 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, pos_d, cache)?;
6451 let mut cur = e.uninit(n_embd)?;
6452 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
6453 let next_norm = if il + 1 < hi {
6454 Some(self.layers[il + 1].attn_norm.float_data())
6455 } else { None };
6456 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
6457 x = xn;
6458 h_carry = hn;
6459 }
6460 Ok(x)
6461 }
6462
6463 fn gemma4_decode_step_h_pp2(&self, e: &Engine, token: u32, cache: &mut Cache, split: usize)
6470 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6471 if crate::pp::pp2_streams_off() {
6472 return self.gemma4_decode_step_h_pp2_samestream(e, token, cache, split);
6473 }
6474 let rt = crate::pp::Pp2Rt::get(e)?;
6475 let e0 = rt.engine(0, e);
6476 let e1 = rt.engine(1, e);
6477 let n_embd = self.cfg.n_embd as usize;
6478 let eps = self.cfg.rms_eps;
6479
6480 let (pos_d, slot) = {
6482 let _st0 = rt.enter(0);
6483 let pos_d = e0.htod_i32(&[cache.pos as i32])?;
6484 let mut x = e0.htod(&self.embd.gather(n_embd, &[token]))?;
6485 e0.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
6486 let x = self.gemma4_decode_layers(e0, x, 0, split, &pos_d, cache)?;
6487 let slot = rt.tx(0, &x, n_embd)?;
6488 (pos_d, slot)
6489 };
6490
6491 let _st1 = rt.enter(1);
6493 let x = rt.rx(0, slot, n_embd)?;
6494 let x = self.gemma4_decode_layers(e1, x, split, self.layers.len(), &pos_d, cache)?;
6495
6496 let mut hn = e1.uninit(n_embd)?;
6497 e1.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
6498 let h_seed = e1.clone_dtod(&x)?;
6499 let mut ld = e1.matmul(&self.output, &hn, 1)?;
6500 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6501 e1.softcap(&mut ld, cap, self.output.out_features())?;
6502 self.gemma4_suppress(e1, &mut ld, 1)?;
6503 let logits = e1.dtoh(&ld)?;
6504 cache.pos += 1;
6505 Ok((logits, h_seed))
6506 }
6507
6508 fn gemma4_decode_step_h_pp2_samestream(&self, e: &Engine, token: u32, cache: &mut Cache,
6511 split: usize)
6512 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6513 let n_embd = self.cfg.n_embd as usize;
6514 let eps = self.cfg.rms_eps;
6515 let pos_d = e.htod_i32(&[cache.pos as i32])?;
6516
6517 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
6519 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
6520 let x = self.gemma4_decode_layers(e, x, 0, split, &pos_d, cache)?;
6521
6522 let boundary_tx = e.clone_dtod(&x)?;
6524 let boundary_rx = e.clone_dtod(&boundary_tx)?;
6525
6526 let x = self.gemma4_decode_layers(e, boundary_rx, split, self.layers.len(), &pos_d, cache)?;
6528
6529 let mut hn = e.uninit(n_embd)?;
6530 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
6531 let h_seed = e.clone_dtod(&x)?;
6532 let mut ld = e.matmul(&self.output, &hn, 1)?;
6533 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6534 e.softcap(&mut ld, cap, self.output.out_features())?;
6535 self.gemma4_suppress(e, &mut ld, 1)?;
6536 let logits = e.dtoh(&ld)?;
6537 cache.pos += 1;
6538 Ok((logits, h_seed))
6539 }
6540}
6541
6542impl HybridModel {
6551 pub fn is_gemma4_e4b(&self) -> bool {
6552 self.gemma4_aux.as_ref().is_some_and(|a| a.e4b.is_some())
6553 }
6554
6555 fn gemma4_e4b_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
6559 let g = self.cfg.gemma4.as_ref().unwrap();
6560 let swa = g.swa_pattern[il];
6561 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
6562 let Mixer::Full(fa) = &self.layers[il].mixer else { panic!("e4b layer {il} not full-attn") };
6563 let nh = fa.wq.out_features() / hd;
6564 let nkv = fa.wk.out_features() / hd;
6565 (hd, nkv, nh, if swa { g.rope_base_swa } else { g.rope_base_global }, 1.0, swa)
6566 }
6567
6568 fn gemma4_e4b_kv_target(&self, il: usize) -> Option<usize> {
6570 self.layers[il].gemma4.as_ref()
6571 .and_then(|b| b.e4b.as_ref())
6572 .and_then(|e4| e4.kv_share.map(|t| t as usize))
6573 }
6574
6575 fn gemma4_e4b_inp_pl(&self, e: &Engine, tokens: &[u32], x_scaled: &CudaSlice<f32>, t: usize)
6580 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6581 let tok_d = e.stream().clone_htod(&tokens.to_vec())?;
6582 self.gemma4_e4b_inp_pl_dev(e, &tok_d, x_scaled, t)
6583 }
6584
6585 fn gemma4_e4b_inp_pl_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
6587 x_scaled: &CudaSlice<f32>, t: usize)
6588 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6589 let aux = self.gemma4_aux.as_ref().unwrap();
6590 let m = aux.e4b.as_ref().unwrap();
6591 let n_embd = self.cfg.n_embd as usize;
6592 let n_layer = self.layers.len();
6593 let width = m.n_epl * n_layer;
6594 let tbl = m.tok_tbl_gpu.get_or_init(|| {
6595 e.upload_u8(&m.tok_embd_bytes).expect("e4b per-layer token table upload")
6596 });
6597 let mut a = e.embed_gather_device_td(tbl, tok_d, t, width, m.tok_embd_qt,
6598 m.tok_embd_row_bytes)?;
6599 e.scale_inplace(&mut a, (m.n_epl as f32).sqrt(), t * width)?;
6600 let mut p = e.matmul(&m.model_proj, x_scaled, t)?;
6601 e.scale_inplace(&mut p, 1.0 / (n_embd as f32).sqrt(), t * width)?;
6602 let mut pn = e.uninit(t * width)?;
6603 e.rms_norm(&p, m.proj_norm.float_data(), &mut pn, m.n_epl, t * n_layer,
6604 self.cfg.rms_eps)?;
6605 let mut out = e.uninit(t * width)?;
6606 e.add_scale(&a, &pn, 1.0 / 2f32.sqrt(), &mut out, t * width)?;
6607 Ok(out)
6608 }
6609
6610 #[allow(clippy::too_many_arguments)]
6615 fn gemma4_e4b_attn(&self, e: &Engine, il: usize,
6616 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
6617 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
6618 dc_bucket: Option<usize>)
6619 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6620 let (hd, nkv, nh, base, scale, swa) = self.gemma4_e4b_geom(il);
6621 let eps = self.cfg.rms_eps;
6622 let aux = self.gemma4_aux.as_ref().unwrap();
6623 let Mixer::Full(fa) = &self.layers[il].mixer else { unreachable!() };
6624 let h0 = e.zeros(0)?;
6628 let h = &h0;
6629
6630 let ff = if swa { None } else {
6631 Some(aux.rope_freqs.as_ref().expect("e4b global rope needs rope_freqs.weight"))
6632 };
6633 let share = self.gemma4_e4b_kv_target(il);
6634 let mut kv_f32: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
6636 let mut q;
6637 if let Some(_tgt) = share {
6638 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
6639 q = e.uninit(t * nh * hd)?;
6640 let mut kdummy = e.uninit(1)?;
6643 let mut vdummy = e.uninit(1)?;
6644 e.rms_norm_qkv_rope(&q0, &q0, &q0, fa.q_norm.float_data(),
6645 fa.q_norm.float_data(), &aux.ones,
6646 &mut q, &mut kdummy, &mut vdummy, hd, nh * t, 0,
6647 pos_d, nh, 1, base, 1.0, ff, eps)?;
6648 } else {
6649 let e4bits = self.layers[il].gemma4.as_ref().and_then(|g| g.e4b.as_ref());
6653 let cat = e4bits.and_then(|e4| e4.qkv_cat.as_ref());
6654 q = e.uninit(t * nh * hd)?;
6655 let mut k = e.uninit(t * nkv * hd)?;
6656 let mut v = e.uninit(t * nkv * hd)?;
6657 if t == 1 && cat.is_some() {
6658 let qkv0 = e.matmul_pre(cat.unwrap(), hq, hdq, h, 1)?;
6659 e.rms_norm_qkv_rope_cat(&qkv0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6660 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
6661 pos_d, nh, nkv, base, 1.0, ff, eps)?;
6662 } else {
6663 let (q0, k0, v0) = match if t == 1 {
6664 e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)?
6665 } else {
6666 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6669 if *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0")) {
6670 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
6671 } else { None }
6672 } {
6673 Some(triple) => triple,
6674 None => (e.matmul_pre(&fa.wq, hq, hdq, h, t)?,
6675 e.matmul_pre(&fa.wk, hq, hdq, h, t)?,
6676 e.matmul_pre(&fa.wv, hq, hdq, h, t)?), };
6678 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(),
6681 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v,
6682 hd, nh * t, nkv * t, pos_d, nh, nkv, base, 1.0, ff, eps)?;
6683 }
6684 let kvl = cache.kv[il].as_mut().unwrap();
6685 let cls = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
6689 if dc_bucket.is_some() {
6690 debug_assert!(t == 1);
6695 e.append_kv_quantized_row_dc_inc(&k, &v, &mut kvl.k, &mut kvl.v,
6697 &mut kvl.len_d, kvl.kv_dim_k, kvl.kv_dim_v,
6698 kvl.k_tok_bytes, kvl.v_tok_bytes, cls)?;
6699 } else {
6700 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
6701 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes,
6702 kvl.v_tok_bytes, cls)?;
6703 kvl.len += t;
6704 }
6705 kv_f32 = Some((k, v));
6706 }
6707 let kvl_idx = share.unwrap_or(il);
6710 let kvl = cache.kv[kvl_idx].as_ref().unwrap();
6711 let base_len = kvl.len - t; let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6713 let mut attn = e.uninit(t * nh * hd)?;
6714 if t > 1 && base_len == 0 && std::env::var("MEMRA_NOFA").is_err() {
6726 if let Some((kf, vf)) = &kv_f32 {
6727 if hd == 256 && t <= win {
6728 e.fa_prefill(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
6729 return Ok(e.matmul(&fa.wo, &attn, t)?);
6730 }
6731 if hd == 256 && swa && t > win {
6732 e.fa_prefill_w(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true,
6733 win)?;
6734 return Ok(e.matmul(&fa.wo, &attn, t)?);
6735 }
6736 if hd == 512 && !swa {
6737 e.fa_prefill_hd512(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale,
6738 true)?;
6739 return Ok(e.matmul(&fa.wo, &attn, t)?);
6740 }
6741 } else if share.is_some() {
6742 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
6743 let k_view = e.view_u8(&kvl.k, kvl.k.len());
6744 let v_view = e.view_u8(&kvl.v, kvl.v.len());
6745 if hd == 256 && (!swa || t <= win) {
6746 e.fa_prefill_view(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t, t,
6748 scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
6749 return Ok(e.matmul(&fa.wo, &attn, t)?);
6750 }
6751 let kv_dim = nkv * hd;
6754 let mut kf = e.uninit(t * kv_dim)?;
6755 let mut vf = e.uninit(t * kv_dim)?;
6756 e.fa_dequant_kv_view_f32(&k_view, &v_view, &mut kf, &mut vf, kv_dim, kv_dim,
6757 t, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
6758 if hd == 512 {
6759 e.fa_prefill_hd512(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale,
6760 true)?;
6761 } else {
6762 e.fa_prefill_w(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true,
6763 win)?;
6764 }
6765 return Ok(e.matmul(&fa.wo, &attn, t)?);
6766 }
6767 }
6768 if let Some(bucket) = dc_bucket {
6769 assert!(t == 1);
6774 let bucket = if hd == 512 && win <= crate::fa512_min_tkv() {
6780 bucket.min(crate::fa512_min_tkv().saturating_sub(1))
6781 } else { bucket };
6782 let k_view = e.view_u8(&kvl.k, kvl.k.len());
6783 let v_view = e.view_u8(&kvl.v, kvl.v.len());
6784 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
6785 if crate::Engine::wpf_level() >= 1 {
6793 e.prefetch_weight_l2(&fa.wo)?;
6794 }
6795 if e.uses_q8_1_fast(&fa.wo) {
6798 let mut oq = e.alloc_i8_uninit(nh * hd)?;
6799 let mut od = e.zeros(nh * hd / 32)?;
6800 e.fa_decode_dc_q8(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6801 &kvl.len_d, bucket, scale,
6802 kvl.k_tok_bytes, kvl.v_tok_bytes, g,
6803 Some((&mut oq, &mut od)))?;
6804 return Ok(e.matmul_pre(&fa.wo, &oq, &od, &attn, t)?);
6805 }
6806 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6807 &kvl.len_d, bucket, scale,
6808 kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
6809 return Ok(e.matmul(&fa.wo, &attn, t)?);
6810 }
6811 for i in 0..t {
6812 let avail = base_len + i + 1;
6813 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
6814 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
6815 (off_tok + t_kv) * kvl.k_tok_bytes);
6816 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
6817 (off_tok + t_kv) * kvl.v_tok_bytes);
6818 let qv = e.view(&q, t * nh * hd);
6819 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
6820 let mut q_one = e.uninit(nh * hd)?;
6821 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
6822 let mut a_one = e.uninit(nh * hd)?;
6823 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
6827 kvl.k_tok_bytes, kvl.v_tok_bytes,
6828 (!swa && crate::Engine::gkv_on())
6829 || (swa && crate::Engine::wkv_on()))?;
6830 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
6831 }
6832 Ok(e.matmul(&fa.wo, &attn, t)?)
6833 }
6834
6835 fn gemma4_e4b_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
6840 head_last: bool)
6841 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6842 let n_embd = self.cfg.n_embd as usize;
6843 let t = tokens.len();
6844 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
6845 let pos_d = e.htod_i32(&pos)?;
6846 let mut x = e.htod(&self.embd.gather(n_embd, tokens))?;
6847 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6848 let inp_pl = self.gemma4_e4b_inp_pl(e, tokens, &x, t)?;
6849 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, head_last)
6850 }
6851
6852 fn gemma4_e4b_trunk_core(&self, e: &Engine, x_in: CudaSlice<f32>, inp_pl: CudaSlice<f32>,
6856 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
6857 dc_bucket: Option<usize>, cap_logits: bool, head_last: bool)
6858 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6859 let n_embd = self.cfg.n_embd as usize;
6860 let eps = self.cfg.rms_eps;
6861 let n_layer = self.layers.len();
6862 let mut x = x_in;
6863 let aux_e4b = self.gemma4_aux.as_ref().unwrap().e4b.as_ref().unwrap();
6864 let n_epl = aux_e4b.n_epl;
6865
6866 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6872 for il in 0..n_layer {
6873 let layer = &self.layers[il];
6874 let (hq, hdq) = match h_carry.take() {
6875 Some(p) => p,
6876 None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
6877 };
6878 let o = self.gemma4_e4b_attn(e, il, &hq, &hdq, pos_d, t, cache, dc_bucket)?;
6879 let bits = layer.gemma4.as_ref().unwrap();
6882 let e4b = bits.e4b.as_ref().expect("e4b layer bits");
6883 let fuse_exit = e.uses_q8_1_fast(&e4b.inp_gate);
6894 let (sn, attn_out) = self.gemma4_layer_tail_core_pn(
6895 e, layer, &o, &x, t, Some(layer.post_attn_norm.float_data()), fuse_exit)?;
6896 let mut resid = e.uninit(t * n_embd)?;
6897 let g = if fuse_exit {
6903 let (rq, rd) = e.rms_pre_add_q8_1(&sn, bits.post_ffw_norm.float_data(),
6905 &attn_out, &mut resid, n_embd, t,
6906 self.cfg.rms_eps)?;
6907 e.matmul_pre(&e4b.inp_gate, &rq, &rd, &resid, t)?
6908 } else {
6909 e.add(&sn, &attn_out, &mut resid, t * n_embd)?;
6910 e.matmul(&e4b.inp_gate, &resid, t)?
6911 };
6912 let mut act = e.uninit(t * n_epl)?;
6913 let y = if t == 1 && e.uses_q8_1_fast(&e4b.proj) {
6914 let ipv = e.view(&inp_pl, n_epl * n_layer);
6915 let row = ipv.slice(il * n_epl..(il + 1) * n_epl);
6916 let (aq, ad) = e.gelu_tanh_mul_q8_1(&g, &row, &mut act, n_epl, 1)?;
6917 e.matmul_pre(&e4b.proj, &aq, &ad, &act, t)?
6918 } else {
6919 let mut inp_this = e.uninit(t * n_epl)?;
6920 e.copy_rows_strided(&inp_pl, &mut inp_this, n_epl, t, n_epl * n_layer,
6921 il * n_epl)?;
6922 e.gelu_tanh_mul(&g, &inp_this, &mut act, t * n_epl)?;
6923 e.matmul(&e4b.proj, &act, t)?
6924 };
6925 let next_norm = if il + 1 < n_layer {
6928 self.layers[il + 1].attn_norm.float_data()
6929 } else {
6930 self.output_norm.float_data()
6931 };
6932 let mut xn = e.uninit(t * n_embd)?;
6933 let pair = e.rms_pre_add_scale_rms_norm_q8_1(&y, e4b.post_norm.float_data(),
6934 &resid, bits.layer_scale, next_norm,
6935 &mut xn, n_embd, t, eps)?;
6936 h_carry = Some(pair);
6937 x = xn;
6938 }
6939 let (oq, odq) = h_carry.take().unwrap();
6943 let h0 = e.zeros(0)?;
6944 let hm = if head_last { 1 } else { t };
6945 let (hq, hd) = if head_last && t > 1 {
6946 let mut q1 = e.uninit_i8(n_embd)?;
6947 e.dtod_copy_view_i8(&oq.slice((t - 1) * n_embd..t * n_embd), &mut q1)?;
6948 let nb = n_embd / 32;
6949 let mut d1 = e.uninit(nb)?;
6950 e.dtod_copy_view(&odq.slice((t - 1) * nb..t * nb), &mut d1)?;
6951 (q1, d1)
6952 } else {
6953 (oq, odq)
6954 };
6955 let mut ld = e.matmul_pre(&self.output, &hq, &hd, &h0, hm)?;
6956 if cap_logits {
6960 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6961 e.softcap(&mut ld, cap, hm * self.output.out_features())?;
6962 }
6963 self.gemma4_suppress(e, &mut ld, hm)?; Ok((ld, x))
6965 }
6966
6967 pub fn gemma4_e4b_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
6974 t: usize, pos0: usize, cache: &mut Cache)
6975 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6976 let n_embd = self.cfg.n_embd as usize;
6977 let eps = self.cfg.rms_eps;
6978 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
6979 let pos_d = e.htod_i32(&pos)?;
6980 let embd_gpu = self.embd_gpu.get_or_init(|| {
6981 e.upload_u8(&self.embd.raw).expect("embed table upload")
6982 });
6983 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
6984 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
6985 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6986 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, tok_d, &x, t)?;
6987 let (ld, xp) = self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true,
6988 false)?;
6989 let n_vocab = self.output.out_features();
6992 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
6993 for i in 0..t {
6994 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
6995 }
6996 let mut hn = e.uninit(t * n_embd)?;
6997 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
6998 cache.pos += t;
6999 Ok((vam, hn))
7000 }
7001
7002 pub(crate) fn gemma4_e4b_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
7005 cache: &mut Cache)
7006 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7007 let n_embd = self.cfg.n_embd as usize;
7008 let eps = self.cfg.rms_eps;
7009 let t = tokens.len();
7010 let (ld, xp) = self.gemma4_e4b_trunk(e, tokens, pos0, cache, false)?;
7011 let mut hn = e.uninit(t * n_embd)?;
7012 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
7013 cache.pos += t;
7014 Ok((e.dtoh(&ld)?, hn))
7015 }
7016
7017 pub fn gemma4_e4b_decode_step_dcg(&self, e: &Engine, token_d: &mut CudaSlice<u32>,
7023 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7024 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7025 n_vocab: usize, bucket: usize)
7026 -> Result<(), Box<dyn std::error::Error>> {
7027 let n_embd = self.cfg.n_embd as usize;
7028 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
7029 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
7030 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
7031 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, Some(bucket),
7032 false, false)?;
7033 e.argmax_token_device_into(&ld, token_d, n_vocab)?;
7034 e.inc_seqlen(pos_d)?;
7035 Ok(())
7036 }
7037
7038 #[allow(clippy::too_many_arguments)]
7046 pub fn gemma4_e4b_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
7047 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7048 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7049 n_vocab: usize)
7050 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
7051 let n_embd = self.cfg.n_embd as usize;
7052 let eps = self.cfg.rms_eps;
7053 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
7054 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
7055 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
7056 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, None, false,
7057 false)?;
7058 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
7059 e.argmax_token_device_into(&ld, &mut tok_out, n_vocab)?;
7060 e.inc_seqlen(pos_d)?;
7061 cache.pos += 1;
7062 let _ = eps;
7063 Ok(tok_out)
7064 }
7065
7066 pub(crate) fn gemma4_e4b_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
7069 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7070 let (ld, x) = self.gemma4_e4b_trunk(e, &[token], cache.pos, cache, false)?;
7071 let logits = e.dtoh(&ld)?;
7072 cache.pos += 1;
7073 Ok((logits, x))
7074 }
7075
7076 pub(crate) fn gemma4_e4b_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
7080 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7081 assert_eq!(cache.pos, 0, "e4b prime is fresh-prompt only (v0)");
7082 let n_embd = self.cfg.n_embd as usize;
7083 let t = tokens.len();
7084 let (ld, x) = self.gemma4_e4b_trunk(e, tokens, 0, cache, true)?;
7085 cache.pos += t;
7086 let last = e.dtoh(&ld)?; let xv = e.view(&x, t * n_embd);
7088 let row = xv.slice((t - 1) * n_embd..t * n_embd);
7089 let mut h_seed = e.uninit(n_embd)?;
7090 e.copy_view_into(&mut h_seed, 0, &row, n_embd)?;
7091 Ok((last, h_seed, x))
7092 }
7093
7094 pub(crate) fn gemma4_e4b_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
7096 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
7097 let mut cache = Cache::new(e, &self.cfg, tokens.len() + 8)?;
7098 let (ld, _x) = self.gemma4_e4b_trunk(e, tokens, 0, &mut cache, last_only)?;
7099 Ok(e.dtoh(&ld)?) }
7101}
7102
7103#[cfg(test)]
7104mod page_prefetch_tests {
7105 use super::{
7106 grouped_worker_prefetch_position, page_prefetch_positions,
7107 page_prefetch_window_from_values, worker_prefetch_positions,
7108 };
7109
7110 #[test]
7111 fn page_prefetch_window_keeps_existing_opt_in_default() {
7112 assert_eq!(page_prefetch_window_from_values(false, None), 0);
7113 assert_eq!(page_prefetch_window_from_values(false, Some("8")), 0);
7114 assert_eq!(page_prefetch_window_from_values(true, None), 1);
7115 assert_eq!(page_prefetch_window_from_values(true, Some("bad")), 1);
7116 assert_eq!(page_prefetch_window_from_values(true, Some("0")), 0);
7117 assert_eq!(page_prefetch_window_from_values(true, Some("8")), 8);
7118 }
7119
7120 #[test]
7121 fn rolling_page_prefetch_advises_each_future_expert_once() {
7122 let advised: Vec<_> = (0..7)
7123 .flat_map(|position| page_prefetch_positions(position, 7, 3))
7124 .collect();
7125 assert_eq!(advised, vec![1, 2, 3, 4, 5, 6]);
7126
7127 let one_ahead: Vec<_> = (0..4)
7128 .flat_map(|position| page_prefetch_positions(position, 4, 1))
7129 .collect();
7130 assert_eq!(one_ahead, vec![1, 2, 3]);
7131 assert!(page_prefetch_positions(0, 4, 0).is_empty());
7132 }
7133
7134 #[test]
7135 fn grouped_worker_prefetch_primes_first_then_each_known_next_once() {
7136 assert_eq!(grouped_worker_prefetch_position(0, None), None);
7137 let positions: Vec<_> = std::iter::once(grouped_worker_prefetch_position(4, None).unwrap())
7138 .chain((0..4).filter_map(|position| {
7139 grouped_worker_prefetch_position(4, Some(position))
7140 }))
7141 .collect();
7142 assert_eq!(positions, vec![0, 1, 2, 3]);
7143 assert_eq!(grouped_worker_prefetch_position(1, Some(0)), None);
7144 }
7145
7146 #[test]
7147 fn rolling_worker_prefetch_primes_current_and_each_future_expert_once() {
7148 let queued: Vec<_> = (0..8)
7149 .flat_map(|position| worker_prefetch_positions(position, 8, 5))
7150 .collect();
7151 assert_eq!(queued, (0..8).collect::<Vec<_>>());
7152
7153 let one_at_a_time: Vec<_> = (0..4)
7154 .flat_map(|position| worker_prefetch_positions(position, 4, 1))
7155 .collect();
7156 assert_eq!(one_at_a_time, vec![0, 1, 2, 3]);
7157 assert!(worker_prefetch_positions(0, 4, 0).is_empty());
7158 }
7159}
7160
7161pub struct G4DcSlots {
7162 x: CudaSlice<f32>, xn: CudaSlice<f32>, cur: CudaSlice<f32>,
7163 hq: CudaSlice<i8>, hd_: CudaSlice<f32>,
7164 q0: CudaSlice<f32>, k0: CudaSlice<f32>, v0: CudaSlice<f32>,
7165 q: CudaSlice<f32>, k: CudaSlice<f32>, v: CudaSlice<f32>,
7166 attn: CudaSlice<f32>, o: CudaSlice<f32>,
7167 attn_out: CudaSlice<f32>, zsh: CudaSlice<f32>,
7168 zq: CudaSlice<i8>, zd: CudaSlice<f32>,
7169 gate: CudaSlice<f32>, up: CudaSlice<f32>,
7170 act: CudaSlice<f32>, actq: CudaSlice<i8>, actd: CudaSlice<f32>,
7171 f0: CudaSlice<f32>, sn: CudaSlice<f32>,
7172 hn: CudaSlice<f32>, logits: CudaSlice<f32>,
7173}
7174