1use std::sync::{Arc, Mutex};
4use cudarc::driver::{CudaContext, CudaStream, CudaModule, CudaFunction, CudaSlice, LaunchConfig, PushKernelArg};
5use cudarc::nvrtc::Ptx;
6
7#[cfg(debug_assertions)]
8pub(crate) fn debug_assert_tensor_stream_device<T>(
9 tensor: &CudaSlice<T>,
10 stream: &CudaStream,
11 site: &str,
12) {
13 let tensor_dev = tensor.ordinal();
14 let stream_dev = stream.context().ordinal();
15 assert_eq!(
16 tensor_dev, stream_dev,
17 "PP cross-device tensor read at {site}: tensor on dev{tensor_dev}, stream on dev{stream_dev}"
18 );
19}
20
21pub use memra_gguf;
22pub use memra_runtime;
23
24pub mod model;
25pub mod forward;
26pub mod hybrid;
27pub mod hybrid_forward;
28pub mod sigrouter_contract;
29pub mod cache {
32 pub use memra_kv::*;
33}
34pub mod decode;
35pub mod decode_batch;
36pub mod moesd;
37pub mod mla;
41pub mod pp;
42pub mod spec;
43pub mod gemma_spec;
44pub mod round_stream;
45pub mod graph_update;
46pub mod dflash;
47pub mod eagle;
48pub use memra_sampling as sampler;
49
50pub fn moe_f16g_mode() -> u8 {
94 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
95 *M.get_or_init(|| match std::env::var("MEMRA_MOE_F16G").as_deref() {
96 Ok("0") => 0,
97 Ok("2") => 2,
98 Ok("3") => 3,
99 Ok(_) => 1,
100 Err(_) => 2,
103 })
104}
105pub fn moe_f16g_sk_params() -> (i32, i32) {
119 static P: std::sync::OnceLock<(i32, i32)> = std::sync::OnceLock::new();
120 *P.get_or_init(|| match std::env::var("MEMRA_F16G_SK").as_deref() {
121 Ok("0") => (-1, 0),
122 Ok("32") => (0, i32::MAX),
123 Ok("128") => (0, 1),
124 _ => {
125 let cross = std::env::var("MEMRA_F16G_SK_CROSS").ok()
126 .and_then(|v| v.parse().ok()).unwrap_or(64);
127 (0, cross)
128 }
129 })
130}
131pub fn moe_f16g_direct_on(qtype: i32) -> bool {
142 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
143 let m = *M.get_or_init(|| match std::env::var("MEMRA_F16G_DIRECT").as_deref() {
144 Ok("0") => 0,
145 Ok("kq") => 1,
146 _ => 2,
147 });
148 match m {
149 0 => false,
150 1 => qtype == QT_Q4_K || qtype == QT_Q6_K,
151 _ => true,
152 }
153}
154pub fn moe_f16g_tail_on() -> bool {
163 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
164 *ON.get_or_init(|| std::env::var("MEMRA_F16G_TAIL").as_deref() != Ok("0"))
165}
166
167pub fn moe_f16g_gemma_on() -> bool {
174 static M: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
175 *M.get_or_init(|| !matches!(std::env::var("MEMRA_MOE_F16G").as_deref(), Ok("0") | Err(_)))
176}
177
178pub fn moe_fuse_actq_on() -> bool {
182 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
183 *ON.get_or_init(|| std::env::var("MEMRA_MOE_FUSE_ACTQ").as_deref() != Ok("0"))
184}
185
186pub fn router_prefill_exact_on() -> bool {
196 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
197 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_PREFILL_EXACT").as_deref() != Ok("0"))
198}
199
200pub fn router_kernel_on() -> bool {
201 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
202 *ON.get_or_init(|| {
203 let on = std::env::var("MEMRA_ROUTER_KERNEL").as_deref() != Ok("0");
204 if !on { eprintln!("[memra] router kernel OFF (rollback: per-column cuBLAS gemv)"); }
205 on
206 })
207}
208
209pub const ROUTER_BATCH_MIN_T: usize = 8;
224pub fn router_batch_on() -> bool {
225 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
226 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_BATCH").as_deref() != Ok("0"))
227}
228mod cpu_experts;
229pub mod moe_cache;
230pub mod spill;
231mod spill_pread;
232#[cfg(memra_cutlass)]
233pub mod cutlass_ffi;
234pub mod mmq_ffi;
235pub mod f16_ffi;
236pub mod prime_graph;
237pub mod fp8_ffi;
238
239const FATBIN: &[u8] = include_bytes!(env!("MEMRA_ENGINE_FATBIN"));
246const HYBRID_FATBIN: &[u8] = include_bytes!(env!("MEMRA_HYBRID_FATBIN"));
247const QMATVEC_FATBIN: &[u8] = include_bytes!(env!("MEMRA_QMATVEC_FATBIN"));
248const FLASH_FATBIN: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN"));
249const GEMM_FATBIN: &[u8] = include_bytes!(env!("MEMRA_GEMM_FATBIN"));
250const ROUTER_FATBIN: &[u8] = include_bytes!(env!("MEMRA_ROUTER_FATBIN"));
251const SAMPLE_FATBIN: &[u8] = include_bytes!(env!("MEMRA_SAMPLE_FATBIN"));
253
254fn gemm_fatbin_bytes() -> std::borrow::Cow<'static, [u8]> {
260 assert!(!(portable_mma_gated() && std::env::var_os("MEMRA_GEMM_FATBIN").is_some()),
261 "MEMRA_GEMM_FATBIN overrides are not allowed in the portable CUDA lane");
262 match std::env::var("MEMRA_GEMM_FATBIN") {
263 Ok(path) => std::borrow::Cow::Owned(
264 std::fs::read(&path).unwrap_or_else(|e| panic!("MEMRA_GEMM_FATBIN read {path}: {e}"))),
265 Err(_) => std::borrow::Cow::Borrowed(GEMM_FATBIN),
266 }
267}
268
269pub(crate) const fn portable_mma_gated() -> bool {
276 cfg!(memra_portable_cuda) && !cfg!(memra_hopper_mma)
277}
278
279const fn legacy_quant_gemm_allowed(portable_cuda: bool, hopper_mma: bool, no_gemm: bool) -> bool {
284 (!portable_cuda || hopper_mma) && !no_gemm
285}
286
287const FLASH_FATBIN_VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VQ4"));
295const FLASH_FATBIN_VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VF8"));
296const FLASH_FATBIN_KF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8"));
297const FLASH_FATBIN_KF8VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VQ4"));
298const FLASH_FATBIN_KF8VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VF8"));
299
300pub use memra_kv::{kv_blk_bytes, kv_cache_formats};
303
304fn flash_fatbin_bytes() -> &'static [u8] {
306 match kv_cache_formats() {
307 ("q8_0", "q5_1") => FLASH_FATBIN,
308 ("q8_0", "q4_0") => FLASH_FATBIN_VQ4,
309 ("q8_0", "fp8") => FLASH_FATBIN_VF8,
310 ("fp8", "q5_1") => FLASH_FATBIN_KF8,
311 ("fp8", "q4_0") => FLASH_FATBIN_KF8VQ4,
312 ("fp8", "fp8") => FLASH_FATBIN_KF8VF8,
313 other => unreachable!("kv_cache_formats returned {other:?}"),
314 }
315}
316
317fn k1_launch_override() -> Option<(u32, u32, u32)> {
324 static K1: std::sync::OnceLock<Option<(u32, u32, u32)>> = std::sync::OnceLock::new();
325 *K1.get_or_init(|| {
326 let v = std::env::var("MEMRA_GEMM_K1_LAUNCH").ok()?;
327 let p: Vec<u32> = v.split(',').filter_map(|s| s.trim().parse().ok()).collect();
328 match p.as_slice() { [bm, bn, w] => Some((*bm, *bn, *w)), _ => None }
329 })
330}
331
332pub(crate) fn wgmma_gemm_enabled() -> bool {
339 static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
340 *V.get_or_init(|| std::env::var("MEMRA_WGMMA").as_deref() == Ok("1"))
341}
342
343pub const FA_VEC_MIN_TKV: usize = 96;
358pub fn fa_vec_min_tkv() -> usize {
362 static V: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
363 *V.get_or_init(|| std::env::var("MEMRA_FA_VEC_MIN").ok()
364 .and_then(|v| v.parse().ok())
365 .unwrap_or_else(|| FA_VEC_MIN_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)))
366}
367
368pub fn fa_f16pv_on() -> bool {
379 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
380 *ON.get_or_init(|| std::env::var("MEMRA_FA_F16PV").map(|v| v != "0")
381 .unwrap_or_else(|_| std::env::var("MEMRA_DRAFT").is_err()))
382}
383
384pub fn fa512_hp_on() -> bool {
388 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
389 *ON.get_or_init(|| std::env::var("MEMRA_FA512_HP").as_deref() != Ok("0"))
390}
391
392pub fn faw_hp_on() -> bool {
396 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
397 *ON.get_or_init(|| std::env::var("MEMRA_FAW_HP").as_deref() != Ok("0"))
398}
399
400pub fn fa512_wide_warps() -> usize {
404 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
405 *N.get_or_init(|| match std::env::var("MEMRA_FA512_W4").as_deref() {
406 Ok("1") => 4, _ => 2,
407 })
408}
409
410pub fn fa512_min_tkv() -> usize {
413 static FA512_MIN: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
414 *FA512_MIN.get_or_init(|| std::env::var("MEMRA_FA512_MIN").ok()
415 .and_then(|v| v.parse().ok()).unwrap_or(512))
416}
417pub static FA_VEC_MIN_DEFAULT: std::sync::atomic::AtomicUsize =
421 std::sync::atomic::AtomicUsize::new(FA_VEC_MIN_TKV);
422pub static FA_SPW_DEFAULT: std::sync::atomic::AtomicUsize =
426 std::sync::atomic::AtomicUsize::new(32);
427pub static FUSED_MR1_DEFAULT: std::sync::atomic::AtomicBool =
433 std::sync::atomic::AtomicBool::new(false);
434pub static ROUTER_W8_DEFAULT: std::sync::atomic::AtomicBool =
441 std::sync::atomic::AtomicBool::new(true);
442pub static FA_SP512_DEFAULT: std::sync::atomic::AtomicUsize =
443 std::sync::atomic::AtomicUsize::new(16);
444pub static RMS_BLOCK_DEFAULT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(256);
449pub static FA_SP_GEMMA: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
451pub static MMQ_SK_FORCE: std::sync::atomic::AtomicI8 = std::sync::atomic::AtomicI8::new(-1);
456pub use memra_kv::KV_FP8_FORCE;
459pub(crate) fn rms_block() -> u32 {
460 static V: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
461 *V.get_or_init(|| std::env::var("MEMRA_RMS_BLOCK").ok()
462 .and_then(|v| v.parse().ok())
463 .unwrap_or_else(|| RMS_BLOCK_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)))
464}
465
466pub(crate) fn fa_split_keys(t_kv: usize, n_head_kv: usize) -> usize {
467 static S: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
468 if let Some(forced) = *S.get_or_init(|| {
469 std::env::var("MEMRA_FA_SPLIT").ok().and_then(|v| v.parse().ok())
470 .filter(|&s: &usize| s >= 8 && s % 8 == 0)
471 }) { return forced; }
472 if FA_SP_GEMMA.load(std::sync::atomic::Ordering::Relaxed)
490 && std::env::var("MEMRA_FA_SP16").as_deref() == Ok("1") {
491 return if t_kv <= 8192 { 16 } else if t_kv <= 16384 { 64 } else { 128 };
492 }
493 let big_rig = fa_sm_count() >= 128;
494 if big_rig {
495 let _ = n_head_kv;
496 if t_kv <= 2048 { 16 } else if t_kv <= 16384 { 64 } else { 128 }
497 } else if n_head_kv <= 4 {
498 if t_kv <= 512 { 8 } else if t_kv <= 16384 { 64 } else { 128 }
519 } else {
520 if t_kv <= 8192 { 32 } else if t_kv <= 16384 { 64 } else { 128 }
521 }
522}
523
524fn fa_sm_count() -> i32 {
527 static N: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
528 *N.get_or_init(|| {
529 cudarc::driver::result::init().ok();
530 cudarc::driver::result::device::get(0)
531 .and_then(|d| unsafe { cudarc::driver::result::device::get_attribute(
532 d, cudarc::driver::sys::CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT) })
533 .unwrap_or(82)
534 })
535}
536
537fn fa_hd_suffix(head_dim: usize) -> Result<&'static str, Box<dyn std::error::Error>> {
541 match head_dim {
542 256 => Ok(""),
543 128 => Ok("_hd128"),
544 d => Err(format!("fa_prefill: no kernel stamped for head_dim={d} (only 256/128); \
545 callers must gate to sdpa_naive").into()),
546 }
547}
548
549pub const QT_Q8_0: i32 = 0;
551pub const QT_Q4_K: i32 = 1;
552pub const QT_Q6_K: i32 = 2;
553pub const QT_Q5_K: i32 = 3;
554pub const QT_Q3_K: i32 = 4;
555pub const QT_IQ4_XS: i32 = 5;
556pub const QT_IQ3_S: i32 = 6;
557pub const QT_NVFP4: i32 = 7;
558pub const QT_F8_E4M3: i32 = 10;
564pub const QT_NVFP4_RP: i32 = 9;
567pub const QT_F32: i32 = 8;
569pub const QT_BF16: i32 = 11;
570pub const QT_Q4_0: i32 = 12; pub const QT_Q2_K: i32 = 13;
575pub const QT_F8_E4M3_BLK: i32 = 14;
591
592pub struct Engine {
594 pub gpu: memra_runtime::Gpu,
595 module: Arc<CudaModule>,
596 hybrid: Arc<CudaModule>,
597 qmatvec: Arc<CudaModule>,
598 flash: Arc<CudaModule>,
599 flash_g: std::sync::OnceLock<Arc<CudaModule>>,
603 gemm: Arc<CudaModule>,
604 router: Arc<CudaModule>,
605 sample: Arc<CudaModule>,
607 moe_cache: Mutex<Option<crate::moe_cache::MoeSlotCache>>,
611 moe_cache_layout: Mutex<Option<Vec<usize>>>,
615 capture_keep_on: std::sync::atomic::AtomicBool,
621 verify_exact: std::sync::atomic::AtomicBool,
626 capture_keep: Mutex<Vec<Box<dyn std::any::Any + Send>>>,
627 pub copy_stream: Arc<CudaStream>,
629 #[cfg(memra_cutlass)]
636 cutlass_scratch: Mutex<Option<crate::cutlass_ffi::CutlassScratch>>,
637 fp8_scratch: Mutex<Option<crate::fp8_ffi::Fp8Scratch>>,
641 fa_vf16_scratch: Mutex<Option<CudaSlice<u8>>>,
644 fa_part_pool: Mutex<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
648 fa_part_retired: Mutex<Vec<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
652 fn_cache: Mutex<std::collections::HashMap<String, CudaFunction>>,
654 f16_scratch: Mutex<Option<crate::f16_ffi::F16Scratch>>,
655 argmax_partials: Mutex<Option<(CudaSlice<f32>, CudaSlice<i32>)>>,
660 prime_deqw_ws: Mutex<Option<(CudaSlice<u8>, CudaSlice<u8>)>>,
665 router_stage: Mutex<Option<PinnedStage>>,
669}
670
671fn fa_v2_on() -> bool {
681 std::env::var("MEMRA_FA_V2").map(|v| v != "0").unwrap_or(true)
687}
688
689fn fa_v3_on() -> bool {
697 std::env::var("MEMRA_FA_V3").map(|v| v != "0").unwrap_or(true)
701}
702
703fn fa_v4_mode() -> &'static str {
708 static M: std::sync::OnceLock<String> = std::sync::OnceLock::new();
709 M.get_or_init(|| std::env::var("MEMRA_FA_V4").unwrap_or_default())
710}
711fn fa_v4_on() -> bool { fa_v4_mode() != "0" } pub static FA_SMEM_TKV_DEFAULT: std::sync::atomic::AtomicUsize =
720 std::sync::atomic::AtomicUsize::new(1024);
721pub static FA_V4_MAX_DEFAULT: std::sync::atomic::AtomicUsize =
722 std::sync::atomic::AtomicUsize::new(usize::MAX);
723pub fn fa_v4_at_pub(t_kv: usize) -> bool { fa_v4_at(t_kv) }
724fn fa_v4_at(t_kv: usize) -> bool {
725 static M: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
726 let mx = *M.get_or_init(|| std::env::var("MEMRA_FA_V4_MAX").ok()
727 .and_then(|v| v.parse().ok())
728 .unwrap_or_else(|| FA_V4_MAX_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)));
729 fa_v4_on() && t_kv < mx
730}
731pub const FA_DEEP_MIN_DEFAULT: usize = 0;
745fn fa_deep_at(t_kv: usize) -> bool {
746 if std::env::var("MEMRA_FA_DEEP").as_deref() == Ok("0") { return false; }
747 let min = std::env::var("MEMRA_FA_DEEP_MIN").ok().and_then(|v| v.parse().ok())
748 .unwrap_or(FA_DEEP_MIN_DEFAULT);
749 t_kv >= min
750}
751pub fn fa_deep_at_pub(t_kv: usize) -> bool { fa_deep_at(t_kv) }
753
754fn fa_v3_active(head_dim: usize) -> bool {
755 fa_v3_on() && head_dim % 128 == 0 && kv_cache_formats() == ("q8_0", "q5_1")
758 && !Engine::kv_fp8_on()
759}
760
761pub fn fa_seqs_eligible(t_kv: usize, head_dim: usize) -> bool {
769 std::env::var("MEMRA_NO_FA_VEC").is_err()
770 && t_kv >= fa_vec_min_tkv()
771 && head_dim == 256
772 && fa_v4_at(t_kv)
773 && !matches!(fa_v4_mode(), "noB3" | "stage")
774 && !Engine::kv_fp8_on()
775}
776pub fn fa_split_keys_pub(t_kv: usize, n_head_kv: usize) -> usize { fa_split_keys(t_kv, n_head_kv) }
778
779struct PinnedStage {
784 ptr: *mut u8,
785 cap: usize,
786}
787unsafe impl Send for PinnedStage {}
788impl PinnedStage {
789 fn new(cap: usize) -> Result<Self, Box<dyn std::error::Error>> {
790 let ptr = unsafe { cudarc::driver::result::malloc_host(cap, 0)? } as *mut u8;
791 Ok(PinnedStage { ptr, cap })
792 }
793}
794impl Drop for PinnedStage {
795 fn drop(&mut self) {
796 let _ = unsafe { cudarc::driver::result::free_host(self.ptr as _) };
797 }
798}
799
800pub const ARGMAX_NB: usize = 256;
803
804pub(crate) use memra_fa3_vl as fa3_vl_raw;
806
807unsafe extern "C" {
808 fn memra_fa3_prefill(q16: *const core::ffi::c_void, k16: *const core::ffi::c_void,
810 v16: *const core::ffi::c_void, o: *mut f32,
811 t: i32, h: i32, hkv: i32, d: i32, scale: f32,
812 stream: *mut core::ffi::c_void) -> i32;
813 pub(crate) fn memra_fa3_vl(q16s: *const *const core::ffi::c_void, k16s: *const *const core::ffi::c_void,
815 v16s: *const *const core::ffi::c_void, os: *const *mut f32,
816 ts: *const i32, b: i32, h: i32, hkv: i32, d: i32, scale: f32,
817 stream: *mut core::ffi::c_void) -> i32;
818}
819
820#[repr(C)]
825#[derive(Clone, Copy)]
826pub struct WPtr8(pub [u64; 8]);
827unsafe impl cudarc::driver::DeviceRepr for WPtr8 {}
828
829#[repr(C)]
834#[derive(Clone, Copy, Default)]
835pub struct GdnSeqVl {
836 pub kb16: u64, pub gcum: u64, pub beta: u64, pub u: u64, pub wb16: u64,
837 pub y: u64, pub ssnap: u64, pub state_in: u64, pub state_out: u64,
838 pub q: u64, pub p: u64, pub o: u64,
839 pub k: u64, pub v: u64, pub g: u64, pub a: u64, pub w: u64,
840 pub t: i32, pub nc: i32,
841}
842unsafe impl cudarc::driver::DeviceRepr for GdnSeqVl {}
843#[repr(C)]
844#[derive(Clone, Copy)]
845pub struct GdnVl8(pub [GdnSeqVl; 8]);
846unsafe impl cudarc::driver::DeviceRepr for GdnVl8 {}
847
848#[repr(C)]
851#[derive(Clone, Copy, Default)]
852pub struct GdnWVl { pub qb16: u64, pub pb16: u64 }
853unsafe impl cudarc::driver::DeviceRepr for GdnWVl {}
854#[repr(C)]
855#[derive(Clone, Copy)]
856pub struct GdnWVl8(pub [GdnWVl; 8]);
857unsafe impl cudarc::driver::DeviceRepr for GdnWVl8 {}
858
859#[repr(C)]
861#[derive(Clone, Copy, Default)]
862pub struct GdnPrepVl {
863 pub qkv: u64, pub conv_state: u64, pub conv_out: u64,
864 pub q_g: u64, pub k_g: u64, pub v_g: u64,
865 pub q_l2: u64, pub k_l2: u64,
866 pub beta_raw: u64, pub alpha: u64, pub beta: u64, pub g_log: u64,
867 pub o: u64, pub z: u64, pub gn: u64, pub gn16: u64,
868 pub kb16: u64,
869 pub qb16: u64,
870 pub t: i32, pub pad: i32,
871}
872unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl {}
873#[repr(C)]
874#[derive(Clone, Copy)]
875pub struct GdnPrepVl8(pub [GdnPrepVl; 8]);
876unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl8 {}
877
878#[repr(C)]
880#[derive(Clone, Copy, Default)]
881pub struct FaSeqVl {
882 pub q: u64, pub k16: u64, pub v16: u64, pub o: u64, pub kf: u64, pub vf: u64,
883 pub t: i32, pub pad: i32,
884}
885unsafe impl cudarc::driver::DeviceRepr for FaSeqVl {}
886#[repr(C)]
887#[derive(Clone, Copy)]
888pub struct FaVl8(pub [FaSeqVl; 8]);
889unsafe impl cudarc::driver::DeviceRepr for FaVl8 {}
890
891#[repr(C)]
893#[derive(Clone, Copy, Default)]
894pub struct AttnPreVl {
895 pub qf: u64, pub kf: u64, pub vf: u64,
896 pub q: u64, pub gate: u64, pub qn: u64, pub kn: u64,
897 pub kc: u64, pub vc: u64,
898 pub t: i32, pub pad: i32,
899}
900unsafe impl cudarc::driver::DeviceRepr for AttnPreVl {}
901#[repr(C)]
902#[derive(Clone, Copy)]
903pub struct AttnPreVl8(pub [AttnPreVl; 8]);
904unsafe impl cudarc::driver::DeviceRepr for AttnPreVl8 {}
905
906pub struct GdnChunkBufs {
909 pub gcum: CudaSlice<f32>,
910 pub a: CudaSlice<f32>,
911 pub p: CudaSlice<f32>,
912 pub u: CudaSlice<f32>,
913 pub w: CudaSlice<f32>,
914 pub kb16: CudaSlice<u8>,
915 pub wb16: CudaSlice<u8>,
916 pub y16: CudaSlice<u8>,
917 pub ssnap16: CudaSlice<u8>,
918 pub qb16: CudaSlice<u8>,
919 pub pb16: CudaSlice<u8>,
920 pub o: CudaSlice<f32>,
921 pub t: usize,
922 pub nc: usize,
923}
924
925#[repr(C)]
927#[derive(Clone, Copy)]
928pub struct F32x8(pub [f32; 8]);
929unsafe impl cudarc::driver::DeviceRepr for F32x8 {}
930
931pub static PRIME_NANOS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
935
936impl Engine {
937 pub fn new(ordinal: usize) -> Result<Self, Box<dyn std::error::Error>> {
938 let gpu = memra_runtime::Gpu::new(ordinal)?;
939 if std::env::var("MEMRA_ARCH_CHECK").as_deref() != Ok("0") {
943 use cudarc::driver::sys::CUdevice_attribute_enum as A;
944 let (maj, min) = cudarc::driver::result::device::get(ordinal as i32)
945 .and_then(|d| unsafe { Ok((
946 cudarc::driver::result::device::get_attribute(d, A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR)?,
947 cudarc::driver::result::device::get_attribute(d, A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR)?)) })
948 .unwrap_or((0, 0));
949 let built = env!("MEMRA_BUILT_CUDA_ARCH");
950 let ok = matches!((built, maj, min),
951 ("120a", 12, 0) | ("120a", 12, 1) | ("100a", 10, 0) | ("90a", 9, 0) | ("89", 8, 9));
952 if !ok {
953 return Err(format!(
954 "memra was built for sm_{built} but device {ordinal} reports compute \
955 capability {maj}.{min}. Rebuild on this machine (MEMRA_CUDA_ARCH \
956 auto-detects the GPU) or set MEMRA_ARCH_CHECK=0 to bypass.").into());
957 }
958 }
959 unsafe {
964 use cudarc::driver::sys;
965 let dev: sys::CUdevice = ordinal as sys::CUdevice;
966 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
967 if sys::cuDeviceGetDefaultMemPool(&mut pool, dev) == sys::CUresult::CUDA_SUCCESS {
968 let mut thresh: u64 = u64::MAX;
969 let _ = sys::cuMemPoolSetAttribute(
970 pool,
971 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
972 &mut thresh as *mut u64 as *mut core::ffi::c_void,
973 );
974 }
975 }
976 let module = gpu.ctx.load_module(Ptx::from_binary(FATBIN.to_vec()))?;
977 let hybrid = gpu.ctx.load_module(Ptx::from_binary(HYBRID_FATBIN.to_vec()))?;
978 let qmatvec = gpu.ctx.load_module(Ptx::from_binary(QMATVEC_FATBIN.to_vec()))?;
979 let flash = gpu.ctx.load_module(Ptx::from_binary(flash_fatbin_bytes().to_vec()))?;
980 let gemm = gpu.ctx.load_module(Ptx::from_binary(gemm_fatbin_bytes().into_owned()))?;
981 let router = gpu.ctx.load_module(Ptx::from_binary(ROUTER_FATBIN.to_vec()))?;
982 let sample = gpu.ctx.load_module(Ptx::from_binary(SAMPLE_FATBIN.to_vec()))?;
983 let copy_stream = gpu.ctx.new_stream()?;
984 if std::env::var("MEMRA_EVT").map(|v| v == "1").unwrap_or(false) {
1000 } else {
1002 unsafe { gpu.ctx.disable_event_tracking(); }
1003 }
1004 Ok(Self { gpu, module, hybrid, qmatvec, flash, flash_g: std::sync::OnceLock::new(), gemm, router, sample,
1005 moe_cache: Mutex::new(None),
1006 moe_cache_layout: Mutex::new(None),
1007 copy_stream,
1008 capture_keep_on: std::sync::atomic::AtomicBool::new(false),
1009 verify_exact: std::sync::atomic::AtomicBool::new(false),
1010 capture_keep: Mutex::new(Vec::new()),
1011 argmax_partials: Mutex::new(None),
1012 prime_deqw_ws: Mutex::new(None),
1013 router_stage: Mutex::new(None),
1014 fp8_scratch: Mutex::new(None),
1015 fa_vf16_scratch: Mutex::new(None),
1016 fa_part_pool: Mutex::new(None),
1017 fa_part_retired: Mutex::new(Vec::new()),
1018 fn_cache: Mutex::new(Default::default()),
1019 f16_scratch: Mutex::new(None),
1020 #[cfg(memra_cutlass)]
1021 cutlass_scratch: Mutex::new(None) })
1022 }
1023
1024 pub fn ctx(&self) -> &Arc<CudaContext> { &self.gpu.ctx }
1025
1026 pub fn pool_cached_bytes(&self) -> usize {
1044 let (reserved, used) = self.pool_reserved_used();
1045 reserved.saturating_sub(used)
1046 }
1047
1048 pub fn pool_reserved_used(&self) -> (usize, usize) {
1055 use cudarc::driver::sys;
1056 unsafe {
1057 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
1058 if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
1059 != sys::CUresult::CUDA_SUCCESS
1060 {
1061 return (0, 0);
1062 }
1063 let (mut reserved, mut used) = (0u64, 0u64);
1064 if sys::cuMemPoolGetAttribute(
1065 pool,
1066 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RESERVED_MEM_CURRENT,
1067 &mut reserved as *mut u64 as *mut core::ffi::c_void,
1068 ) != sys::CUresult::CUDA_SUCCESS {
1069 return (0, 0);
1070 }
1071 if sys::cuMemPoolGetAttribute(
1072 pool,
1073 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_USED_MEM_CURRENT,
1074 &mut used as *mut u64 as *mut core::ffi::c_void,
1075 ) != sys::CUresult::CUDA_SUCCESS {
1076 return (0, 0);
1077 }
1078 (reserved as usize, used as usize)
1079 }
1080 }
1081
1082 pub fn stream(&self) -> Arc<CudaStream> { self.gpu.stream() }
1085 pub fn gkv_on() -> bool {
1088 memra_kv::gkv_on()
1089 }
1090
1091 pub fn wkv_on() -> bool {
1103 memra_kv::wkv_on()
1104 }
1105
1106 pub fn kv_fp8_on() -> bool {
1112 memra_kv::kv_fp8_on()
1113 }
1114
1115 fn fa_func(&self, name: &str, head_dim: usize) -> CudaFunction {
1118 if head_dim == 512 && Self::gkv_on() { self.func_g(name) } else { self.func(name) }
1119 }
1120
1121 fn func_g(&self, name: &str) -> CudaFunction {
1125 let m = self.flash_g.get_or_init(|| {
1126 self.gpu.ctx.load_module(cudarc::nvrtc::Ptx::from_binary(FLASH_FATBIN_KF8VF8.to_vec()))
1127 .expect("load kf8vf8 flash fatbin (fp8-globals arm)")
1128 });
1129 let key = format!("g:{name}");
1130 if let Some(f) = self.fn_cache.lock().unwrap().get(&key) { return f.clone(); }
1131 let f = match m.load_function(name) {
1132 Ok(f) => f,
1133 Err(_) => self.func(name),
1134 };
1135 self.fn_cache.lock().unwrap().insert(key, f.clone());
1136 f
1137 }
1138
1139 fn func(&self, name: &str) -> CudaFunction {
1140 if let Some(f) = self.fn_cache.lock().unwrap().get(name) { return f.clone(); }
1143 let f = self.module.load_function(name)
1144 .or_else(|_| self.hybrid.load_function(name))
1145 .or_else(|_| self.qmatvec.load_function(name))
1146 .or_else(|_| self.flash.load_function(name))
1147 .or_else(|_| self.gemm.load_function(name))
1148 .or_else(|_| self.router.load_function(name))
1149 .or_else(|_| self.sample.load_function(name))
1150 .unwrap_or_else(|_| panic!("kernel {name} not in any fatbin"));
1151 self.fn_cache.lock().unwrap().insert(name.to_string(), f.clone());
1152 f
1153 }
1154
1155 pub fn scatter_trim_logits(&self, src: &CudaSlice<f32>, d2t: &CudaSlice<u32>,
1158 dst: &mut CudaSlice<f32>, d_vocab: usize, n_vocab: usize)
1159 -> Result<(), Box<dyn std::error::Error>> {
1160 let f1 = self.func("scatter_trim_logits_f32");
1161 let f2 = self.func("scatter_trim_logits_pass2_f32");
1162 let (dv, nv) = (d_vocab as i32, n_vocab as i32);
1163 let cfg1 = LaunchConfig { grid_dim: (256, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1164 let __s_b1 = self.gpu.stream();
1165 let mut b1 = __s_b1.launch_builder(&f1);
1166 b1.arg(src).arg(d2t).arg(&mut *dst).arg(&dv).arg(&nv);
1167 unsafe { b1.launch(cfg1)?; }
1168 let cfg2 = LaunchConfig { grid_dim: (d_vocab.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1169 let __s_b2 = self.gpu.stream();
1170 let mut b2 = __s_b2.launch_builder(&f2);
1171 b2.arg(src).arg(d2t).arg(&mut *dst).arg(&dv);
1172 unsafe { b2.launch(cfg2)?; }
1173 Ok(())
1174 }
1175
1176 #[allow(clippy::too_many_arguments)]
1182 pub fn filter_stats(&self, x: &CudaSlice<f32>, row_stride: usize, rows: &CudaSlice<i32>,
1183 out_th: &mut CudaSlice<f32>, out_z: &mut CudaSlice<f32>,
1184 out_max: &mut CudaSlice<f32>, n: usize, nrow: usize,
1185 temp: f32, top_k: i32, top_p: f32, min_p: f32)
1186 -> Result<(), Box<dyn std::error::Error>> {
1187 let f = self.func("filter_stats_f32");
1188 let (ni, nr, rs) = (n as i32, nrow as i32, row_stride as i64);
1189 let cfg = LaunchConfig { grid_dim: (nrow as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
1190 let __s_b = self.gpu.stream();
1191 let mut b = __s_b.launch_builder(&f);
1192 b.arg(x).arg(&rs).arg(rows).arg(&mut *out_th).arg(&mut *out_z).arg(&mut *out_max)
1193 .arg(&ni).arg(&nr).arg(&temp).arg(&top_k).arg(&top_p).arg(&min_p);
1194 unsafe { b.launch(cfg)?; }
1195 Ok(())
1196 }
1197
1198 #[allow(clippy::too_many_arguments)]
1200 pub fn softmax_gather_filtered(&self, x: &CudaSlice<f32>, row_stride: usize,
1201 ids: &CudaSlice<u32>, rows: &CudaSlice<i32>,
1202 th: &CudaSlice<f32>, z: &CudaSlice<f32>,
1203 out: &mut CudaSlice<f32>, n: usize, npair: usize, temp: f32)
1204 -> Result<(), Box<dyn std::error::Error>> {
1205 let f = self.func("softmax_gather_filtered_f32");
1206 let (ni, np, rs) = (n as i32, npair as i32, row_stride as i64);
1207 let cfg = LaunchConfig { grid_dim: (npair as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1208 let __s_b = self.gpu.stream();
1209 let mut b = __s_b.launch_builder(&f);
1210 b.arg(x).arg(&rs).arg(ids).arg(rows).arg(th).arg(z).arg(&mut *out).arg(&ni).arg(&np).arg(&temp);
1211 unsafe { b.launch(cfg)?; }
1212 Ok(())
1213 }
1214
1215 #[allow(clippy::too_many_arguments)]
1217 pub fn residual_sample_filtered(&self, p: &CudaSlice<f32>, q: Option<&CudaSlice<f32>>, n: usize,
1218 temp: f32, seed: u64, stream_pos: u32,
1219 p_stats: (f32, f32, f32), q_stats: (f32, f32, f32),
1220 out_tok: &mut CudaSlice<u32>)
1221 -> Result<(), Box<dyn std::error::Error>> {
1222 let f = self.func("residual_sample_filtered_f32");
1223 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1224 let has_q: i32 = q.is_some() as i32;
1225 let qbuf = q.unwrap_or(p);
1226 let (pm, pth, pz) = p_stats; let (qm, qth, qz) = q_stats;
1227 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
1228 let __s_b = self.gpu.stream();
1229 let mut b = __s_b.launch_builder(&f);
1230 b.arg(p).arg(qbuf).arg(&has_q).arg(&ni).arg(&temp).arg(&slo).arg(&shi).arg(&stream_pos)
1231 .arg(&pm).arg(&pth).arg(&pz).arg(&qm).arg(&qth).arg(&qz).arg(&mut *out_tok);
1232 unsafe { b.launch(cfg)?; }
1233 Ok(())
1234 }
1235
1236 #[allow(clippy::too_many_arguments)]
1238 pub fn gumbel_perturb_filtered(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
1239 seed: u64, stream_pos: u32, temp: f32, row_max: f32, th: f32)
1240 -> Result<(), Box<dyn std::error::Error>> {
1241 let f = self.func("gumbel_perturb_filtered_f32");
1242 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1243 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1244 let __s_b = self.gpu.stream();
1245 let mut b = __s_b.launch_builder(&f);
1246 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp).arg(&row_max).arg(&th);
1247 unsafe { b.launch(cfg)?; }
1248 Ok(())
1249 }
1250
1251 #[allow(clippy::too_many_arguments)]
1255 pub fn penalize_logits(&self, x: &mut CudaSlice<f32>, hist: &CudaSlice<u32>, n_hist: usize,
1256 rep: f32, freq: f32, present: f32, n: usize)
1257 -> Result<(), Box<dyn std::error::Error>> {
1258 if n_hist == 0 { return Ok(()); }
1259 let f = self.func("penalize_logits_f32");
1260 let (nh, ni) = (n_hist as i32, n as i32);
1261 let cfg = LaunchConfig { grid_dim: (n_hist.div_ceil(128) as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
1262 let __s_b = self.gpu.stream();
1263 let mut b = __s_b.launch_builder(&f);
1264 b.arg(&mut *x).arg(hist).arg(&nh).arg(&rep).arg(&freq).arg(&present).arg(&ni);
1265 unsafe { b.launch(cfg)?; }
1266 Ok(())
1267 }
1268
1269 #[allow(clippy::too_many_arguments)]
1271 pub fn penalize_logits_rows(&self, x: &mut CudaSlice<f32>, hist: &CudaSlice<u32>, n_hist: usize,
1272 rep: f32, freq: f32, present: f32, n: usize, nrow: usize)
1273 -> Result<(), Box<dyn std::error::Error>> {
1274 if n_hist == 0 || nrow == 0 { return Ok(()); }
1275 let f = self.func("penalize_logits_rows_f32");
1276 let (nh, ni, nr) = (n_hist as i32, n as i32, nrow as i32);
1277 let cfg = LaunchConfig { grid_dim: (n_hist.div_ceil(128) as u32, nrow as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
1278 let __s_b = self.gpu.stream();
1279 let mut b = __s_b.launch_builder(&f);
1280 b.arg(&mut *x).arg(hist).arg(&nh).arg(&rep).arg(&freq).arg(&present).arg(&ni).arg(&nr);
1281 unsafe { b.launch(cfg)?; }
1282 Ok(())
1283 }
1284
1285 pub fn wpf_level() -> u32 {
1293 static ON: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
1294 *ON.get_or_init(|| std::env::var("MEMRA_WPF").ok()
1295 .and_then(|v| v.parse().ok()).unwrap_or(1))
1296 }
1297
1298 pub fn set_verify_exact(&self, on: bool) {
1310 self.verify_exact.store(on, std::sync::atomic::Ordering::Relaxed);
1311 }
1312 pub(crate) fn verify_exact_on(&self) -> bool {
1313 self.verify_exact.load(std::sync::atomic::Ordering::Relaxed)
1314 }
1315
1316 pub fn qkv_append_on() -> bool {
1319 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1320 *ON.get_or_init(|| std::env::var("MEMRA_QKV_APPEND").map(|v| v != "0").unwrap_or(true))
1321 }
1322
1323 pub fn pdl_wb_on() -> bool {
1326 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1327 *ON.get_or_init(|| std::env::var("MEMRA_PDL_WB").map(|v| v != "0").unwrap_or(true))
1328 }
1329
1330 pub fn pdl_mmvq_on() -> bool {
1334 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1335 *ON.get_or_init(|| std::env::var("MEMRA_PDL_MMVQ").map(|v| v != "0").unwrap_or(true))
1336 }
1337
1338 pub fn pdl_on() -> bool {
1339 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1340 *ON.get_or_init(|| std::env::var("MEMRA_PDL").map(|v| v != "0").unwrap_or(true))
1341 }
1342
1343 fn q40_mr1_on() -> bool {
1349 static Q40MR: std::sync::OnceLock<Option<u32>> = std::sync::OnceLock::new();
1350 match *Q40MR.get_or_init(|| std::env::var("MEMRA_Q40_MR").ok()
1351 .and_then(|v| v.parse().ok())) {
1352 Some(v) => v == 1,
1353 None => crate::FUSED_MR1_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
1354 }
1355 }
1356
1357 fn pdl_func_flash(&self, g: bool, name: &'static str)
1362 -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
1363 use cudarc::driver::sys as cu;
1364 static MODS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool), usize>>> =
1371 std::sync::Mutex::new(None);
1372 static FNS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool, &'static str), usize>>> =
1373 std::sync::Mutex::new(None);
1374 let ctx_key = self.ctx().cu_ctx() as usize;
1375 if let Some(&f) = FNS.lock().unwrap().get_or_insert_with(Default::default)
1376 .get(&(ctx_key, g, name)) { return Ok(f as cu::CUfunction); }
1377 let module = {
1378 let mut mods = MODS.lock().unwrap();
1379 let map = mods.get_or_insert_with(Default::default);
1380 match map.get(&(ctx_key, g)) {
1381 Some(&m) => m,
1382 None => {
1383 let m = self.pdl_load_module_in_ctx(
1384 if g { FLASH_FATBIN_KF8VF8 } else { FLASH_FATBIN })?;
1385 map.insert((ctx_key, g), m);
1386 m
1387 }
1388 }
1389 };
1390 let cname = std::ffi::CString::new(name)?;
1391 let mut f: cu::CUfunction = std::ptr::null_mut();
1392 let r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
1393 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("pdl_func_flash {name} (g={g}): {r:?}").into()); }
1394 FNS.lock().unwrap().get_or_insert_with(Default::default)
1395 .insert((ctx_key, g, name), f as usize);
1396 Ok(f)
1397 }
1398
1399 fn pdl_load_module_in_ctx(&self, bytes: &[u8]) -> Result<usize, Box<dyn std::error::Error>> {
1404 use cudarc::driver::sys as cu;
1405 let mut prev: cu::CUcontext = std::ptr::null_mut();
1406 unsafe { cu::cuCtxGetCurrent(&mut prev).result()?; }
1407 self.ctx().bind_to_thread()?;
1408 let mut m: cu::CUmodule = std::ptr::null_mut();
1409 let r = unsafe { cu::cuModuleLoadData(&mut m, bytes.as_ptr() as *const std::ffi::c_void) };
1410 let restore = if prev.is_null() { cu::CUresult::CUDA_SUCCESS }
1411 else { unsafe { cu::cuCtxSetCurrent(prev) } };
1412 if r != cu::CUresult::CUDA_SUCCESS {
1413 return Err(format!("pdl module load: {r:?}").into());
1414 }
1415 if restore != cu::CUresult::CUDA_SUCCESS {
1416 return Err(format!("pdl module load: ctx restore {restore:?}").into());
1417 }
1418 Ok(m as usize)
1419 }
1420
1421 fn pdl_func(&self, name: &'static str) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
1422 use cudarc::driver::sys as cu;
1423 static MODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
1426 std::sync::Mutex::new(None);
1427 static QMODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
1430 std::sync::Mutex::new(None);
1431 static FNS: std::sync::Mutex<Option<std::collections::HashMap<(usize, &'static str), usize>>> =
1432 std::sync::Mutex::new(None);
1433 let ctx_key = self.ctx().cu_ctx() as usize;
1434 if let Some(&f) = FNS.lock().unwrap().get_or_insert_with(Default::default)
1435 .get(&(ctx_key, name)) { return Ok(f as cu::CUfunction); }
1436 let module = {
1437 let mut mods = MODULES.lock().unwrap();
1438 let map = mods.get_or_insert_with(Default::default);
1439 match map.get(&ctx_key) {
1440 Some(&m) => m,
1441 None => {
1442 let m = self.pdl_load_module_in_ctx(FATBIN)?;
1443 map.insert(ctx_key, m);
1444 m
1445 }
1446 }
1447 };
1448 let cname = std::ffi::CString::new(name)?;
1449 let mut f: cu::CUfunction = std::ptr::null_mut();
1450 let mut r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
1451 if r == cu::CUresult::CUDA_ERROR_NOT_FOUND {
1452 let qmodule = {
1453 let mut mods = QMODULES.lock().unwrap();
1454 let map = mods.get_or_insert_with(Default::default);
1455 match map.get(&ctx_key) {
1456 Some(&m) => m,
1457 None => {
1458 let m = self.pdl_load_module_in_ctx(QMATVEC_FATBIN)?;
1459 map.insert(ctx_key, m);
1460 m
1461 }
1462 }
1463 };
1464 r = unsafe { cu::cuModuleGetFunction(&mut f, qmodule as cu::CUmodule, cname.as_ptr()) };
1465 }
1466 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("pdl_func {name}: {r:?}").into()); }
1467 FNS.lock().unwrap().get_or_insert_with(Default::default)
1468 .insert((ctx_key, name), f as usize);
1469 Ok(f)
1470 }
1471
1472 unsafe fn launch_pdl_flash(&self, g: bool, name: &'static str, grid: (u32, u32, u32),
1484 block: (u32, u32, u32), smem: u32,
1485 params: &mut [*mut std::ffi::c_void])
1486 -> Result<(), Box<dyn std::error::Error>> {
1487 use cudarc::driver::sys as cu;
1488 let f = self.pdl_func_flash(g, name)?;
1489 if smem > 0 {
1490 let r = unsafe { cu::cuFuncSetAttribute(f,
1492 cu::CUfunction_attribute_enum::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
1493 smem as i32) };
1494 if r != cu::CUresult::CUDA_SUCCESS {
1495 return Err(format!("pdl smem attr {name}: {r:?}").into());
1496 }
1497 }
1498 let mut attr = cu::CUlaunchAttribute {
1499 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
1500 pad: [0; 4],
1501 value: cu::CUlaunchAttributeValue { programmaticStreamSerializationAllowed: 1 },
1502 };
1503 let cfg = cu::CUlaunchConfig {
1504 gridDimX: grid.0, gridDimY: grid.1, gridDimZ: grid.2,
1505 blockDimX: block.0, blockDimY: block.1, blockDimZ: block.2,
1506 sharedMemBytes: smem, hStream: self.gpu.stream().cu_stream(),
1507 attrs: &mut attr, numAttrs: 1,
1508 };
1509 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
1510 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("launch_pdl_flash {name}: {r:?}").into()); }
1511 Ok(())
1512 }
1513
1514 unsafe fn launch_pdl(&self, name: &'static str, grid: (u32, u32, u32), block: (u32, u32, u32),
1515 params: &mut [*mut std::ffi::c_void])
1516 -> Result<(), Box<dyn std::error::Error>> {
1517 use cudarc::driver::sys as cu;
1518 let f = self.pdl_func(name)?;
1519 let mut attr = cu::CUlaunchAttribute {
1520 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
1521 pad: [0; 4],
1522 value: cu::CUlaunchAttributeValue { programmaticStreamSerializationAllowed: 1 },
1523 };
1524 let cfg = cu::CUlaunchConfig {
1525 gridDimX: grid.0, gridDimY: grid.1, gridDimZ: grid.2,
1526 blockDimX: block.0, blockDimY: block.1, blockDimZ: block.2,
1527 sharedMemBytes: 0, hStream: self.gpu.stream().cu_stream(),
1528 attrs: &mut attr, numAttrs: 1,
1529 };
1530 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
1531 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("launch_pdl {name}: {r:?}").into()); }
1532 Ok(())
1533 }
1534
1535 pub fn prefetch_weight_l2(&self, w: &crate::model::GpuTensor)
1538 -> Result<(), Box<dyn std::error::Error>> {
1539 if let crate::model::GpuTensor::Quant { bytes, rp4, .. } = w {
1540 let p = rp4.as_ref().unwrap_or(bytes);
1541 self.prefetch_l2(p, p.len())?;
1542 }
1543 Ok(())
1544 }
1545
1546 pub fn gather_row_bf16(&self, table: &CudaSlice<u8>, tok: &CudaSlice<u32>, idx: usize,
1549 dst: &mut CudaSlice<f32>, ncols: usize)
1550 -> Result<(), Box<dyn std::error::Error>> {
1551 let f = self.func("gather_row_bf16_f32");
1552 let cfg = LaunchConfig { grid_dim: (ncols.div_ceil(256) as u32, 1, 1),
1553 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1554 let (nc, ix) = (ncols as i32, idx as i32);
1555 let __s_b = self.gpu.stream();
1556 let mut b = __s_b.launch_builder(&f);
1557 b.arg(table).arg(tok).arg(&ix).arg(dst).arg(&nc);
1558 unsafe { b.launch(cfg)?; }
1559 Ok(())
1560 }
1561
1562 pub fn add_row_inplace(&self, logits: &mut CudaSlice<f32>, bias: &CudaSlice<f32>,
1564 n: usize, row_off: usize)
1565 -> Result<(), Box<dyn std::error::Error>> {
1566 let f = self.func("add_row_inplace_f32");
1567 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1),
1568 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1569 let (ni, off) = (n as i32, row_off as i64);
1570 let __s_b = self.gpu.stream();
1571 let mut b = __s_b.launch_builder(&f);
1572 b.arg(logits).arg(bias).arg(&ni).arg(&off);
1573 unsafe { b.launch(cfg)?; }
1574 Ok(())
1575 }
1576
1577 pub fn prefetch_l2(&self, p: &CudaSlice<u8>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
1579 let f = self.func("prefetch_l2_bytes");
1580 let lines = n.div_ceil(128);
1581 let ni = n as i64;
1582 let cfg = LaunchConfig { grid_dim: (lines.div_ceil(256) as u32, 1, 1),
1583 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1584 let __s_b = self.gpu.stream();
1585 let mut b = __s_b.launch_builder(&f);
1586 b.arg(p).arg(&ni);
1587 unsafe { b.launch(cfg)?; }
1588 Ok(())
1589 }
1590
1591 pub fn router_gemv(&self, w: &CudaSlice<f32>, x: &CudaSlice<f32>, n_embd: usize,
1594 n_experts: usize, t: usize)
1595 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1596 let w8 = match std::env::var("MEMRA_ROUTER_V2").as_deref() {
1602 Ok("0") => false,
1603 Ok(_) => true,
1604 Err(_) => ROUTER_W8_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
1605 };
1606 let batch = w8 && t >= ROUTER_BATCH_MIN_T && router_batch_on();
1616 self.router_gemv_form(w, x, n_embd, n_experts, t, w8, batch)
1617 }
1618
1619 pub fn router_gemv_form(&self, w: &CudaSlice<f32>, x: &CudaSlice<f32>, n_embd: usize,
1622 n_experts: usize, t: usize, w8: bool, batch: bool)
1623 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1624 debug_assert!(!batch || w8, "batch twin exists for the w8 form only");
1625 let mut y = self.alloc_uninit::<f32>(t * n_experts)?;
1626 let f = if batch { self.func("router_gemv_f32_w8_batch") }
1627 else if w8 { self.func("router_gemv_f32_w8") }
1628 else { self.func("router_gemv_f32") };
1629 let (ne, nx, ti) = (n_embd as i32, n_experts as i32, t as i32);
1630 let cfg = if batch {
1631 LaunchConfig { grid_dim: (n_experts.div_ceil(8) as u32, t.div_ceil(8) as u32, 1),
1632 block_dim: (32, 8, 1), shared_mem_bytes: 0 }
1633 } else {
1634 LaunchConfig { grid_dim: (n_experts as u32, t as u32, 1),
1635 block_dim: (32, if w8 { 8 } else { 1 }, 1), shared_mem_bytes: 0 }
1636 };
1637 let __s_b = self.gpu.stream();
1638 let mut b = __s_b.launch_builder(&f);
1639 b.arg(w).arg(x).arg(&mut y).arg(&ne).arg(&nx).arg(&ti);
1640 unsafe { b.launch(cfg)?; }
1641 Ok(y)
1642 }
1643
1644 pub fn rows_permute(&self, src: &CudaSlice<f32>, idx: &CudaSlice<i32>, nrows: usize,
1646 ncols: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1647 let mut dst = self.alloc_uninit::<f32>(nrows * ncols)?;
1648 let f = self.func("rows_permute_f32");
1649 let (nc, nr) = (ncols as i32, nrows as i32);
1650 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (256, 1, 1),
1651 shared_mem_bytes: 0 };
1652 let __s_b = self.gpu.stream();
1653 let mut b = __s_b.launch_builder(&f);
1654 b.arg(src).arg(idx).arg(&mut dst).arg(&nc).arg(&nr);
1655 unsafe { b.launch(cfg)?; }
1656 Ok(dst)
1657 }
1658
1659 pub fn sigmoid_dot_rows(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, n_embd: usize,
1664 t: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1665 static OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1668 if *OFF.get_or_init(|| std::env::var("MEMRA_SHEXP_DOT").as_deref() == Ok("0")) {
1669 let gs = self.linear(x, w, t, n_embd, 1)?;
1670 let mut g = self.uninit(t)?;
1671 self.sigmoid(&gs, &mut g, t)?;
1672 return Ok(g);
1673 }
1674 let mut g = self.alloc_uninit::<f32>(t)?;
1680 let f = self.func("sigmoid_dot_rows_f32");
1681 let (ne, ti) = (n_embd as i32, t as i32);
1682 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (32, 8, 1),
1683 shared_mem_bytes: 0 };
1684 let __s_b = self.gpu.stream();
1685 let mut b = __s_b.launch_builder(&f);
1686 b.arg(x).arg(w).arg(&mut g).arg(&ne).arg(&ti);
1687 unsafe { b.launch(cfg)?; }
1688 Ok(g)
1689 }
1690
1691 pub fn spec_rollback_stream(&self, len_ptrs: &CudaSlice<u64>, pos_start: &CudaSlice<i32>,
1693 acc: &CudaSlice<u32>, base: usize, n_rows: usize)
1694 -> Result<(), Box<dyn std::error::Error>> {
1695 let f = self.func("spec_rollback_stream");
1696 let (b, nr) = (base as i32, n_rows as i32);
1697 let cfg = LaunchConfig { grid_dim: (n_rows.div_ceil(64) as u32, 1, 1),
1698 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1699 let __s_bl = self.gpu.stream();
1700 let mut bl = __s_bl.launch_builder(&f);
1701 bl.arg(len_ptrs).arg(pos_start).arg(acc).arg(&b).arg(&nr);
1702 unsafe { bl.launch(cfg)?; }
1703 Ok(())
1704 }
1705
1706 pub fn plain_tok_ring(&self, vam: &CudaSlice<u32>, pos_start: &CudaSlice<i32>,
1708 base: usize, ring: &mut CudaSlice<u32>)
1709 -> Result<(), Box<dyn std::error::Error>> {
1710 let f = self.func("plain_tok_ring");
1711 let (b, cap) = (base as i32, ring.len() as i32);
1712 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1713 let __s_bl = self.gpu.stream();
1714 let mut bl = __s_bl.launch_builder(&f);
1715 bl.arg(vam).arg(pos_start).arg(&b).arg(&mut *ring).arg(&cap);
1716 unsafe { bl.launch(cfg)?; }
1717 Ok(())
1718 }
1719
1720 pub fn spec_ring_commit(&self, vtok: &CudaSlice<u32>, acc: &CudaSlice<u32>,
1722 brk: &CudaSlice<u32>, ring: &mut CudaSlice<u32>,
1723 pend: &mut CudaSlice<u32>)
1724 -> Result<(), Box<dyn std::error::Error>> {
1725 let f = self.func("spec_ring_commit");
1726 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1727 let __s_b = self.gpu.stream();
1728 let mut b = __s_b.launch_builder(&f);
1729 b.arg(vtok).arg(acc).arg(brk).arg(ring).arg(pend);
1730 unsafe { b.launch(cfg)?; }
1731 Ok(())
1732 }
1733 pub fn i32_copy_add(&self, src: &CudaSlice<i32>, dst: &mut CudaSlice<i32>, delta: i32)
1734 -> Result<(), Box<dyn std::error::Error>> {
1735 let f = self.func("i32_copy_add");
1736 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1737 let __s_b = self.gpu.stream();
1738 let mut b = __s_b.launch_builder(&f);
1739 b.arg(src).arg(dst).arg(&delta);
1740 unsafe { b.launch(cfg)?; }
1741 Ok(())
1742 }
1743 pub fn u32_copy(&self, src: &CudaSlice<u32>, dst: &mut CudaSlice<u32>)
1744 -> Result<(), Box<dyn std::error::Error>> {
1745 let f = self.func("u32_copy");
1746 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1747 let __s_b = self.gpu.stream();
1748 let mut b = __s_b.launch_builder(&f);
1749 b.arg(src).arg(dst);
1750 unsafe { b.launch(cfg)?; }
1751 Ok(())
1752 }
1753
1754 pub fn spec_adapt_k(&self, acc: &CudaSlice<u32>, brk: &mut CudaSlice<u32>,
1758 floor: usize, cap: usize)
1759 -> Result<(), Box<dyn std::error::Error>> {
1760 let f = self.func("spec_adapt_k");
1761 let (fl, cp) = (floor as i32, cap as i32);
1762 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1763 let __s_b = self.gpu.stream();
1764 let mut b = __s_b.launch_builder(&f);
1765 b.arg(acc).arg(brk).arg(&fl).arg(&cp);
1766 unsafe { b.launch(cfg)?; }
1767 Ok(())
1768 }
1769
1770 pub fn spec_accept_greedy_dc(&self, preds: &CudaSlice<u32>, vtok: &CudaSlice<u32>,
1772 last_pred: &CudaSlice<u32>, brk: &CudaSlice<u32>,
1773 out: &mut CudaSlice<u32>)
1774 -> Result<(), Box<dyn std::error::Error>> {
1775 let f = self.func("spec_accept_greedy_dc");
1776 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1777 let __s_b = self.gpu.stream();
1778 let mut b = __s_b.launch_builder(&f);
1779 b.arg(preds).arg(vtok).arg(last_pred).arg(brk).arg(out);
1780 unsafe { b.launch(cfg)?; }
1781 Ok(())
1782 }
1783
1784 pub fn pos_iota(&self, pos0: &CudaSlice<i32>, out: &mut CudaSlice<i32>, t: usize)
1786 -> Result<(), Box<dyn std::error::Error>> {
1787 let f = self.func("pos_iota_i32");
1788 let ti = t as i32;
1789 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (t.max(1) as u32, 1, 1),
1790 shared_mem_bytes: 0 };
1791 let __s_b = self.gpu.stream();
1792 let mut b = __s_b.launch_builder(&f);
1793 b.arg(pos0).arg(out).arg(&ti);
1794 unsafe { b.launch(cfg)?; }
1795 Ok(())
1796 }
1797 #[allow(clippy::too_many_arguments)]
1798 pub fn append_kv_quantized_rows_dc(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
1799 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
1800 t0_dev: &CudaSlice<i32>, t: usize,
1801 kv_dim_k: usize, kv_dim_v: usize,
1802 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
1803 -> Result<(), Box<dyn std::error::Error>> {
1804 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_rows_dc") }
1805 else { self.func("append_quantize_kv_q8_0_q5_1_rows_dc") };
1806 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
1807 let cfg = LaunchConfig { grid_dim: (nblk, t as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1808 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
1809 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
1810 let __s_b = self.gpu.stream();
1811 let mut b = __s_b.launch_builder(&f);
1812 b.arg(k_rows).arg(v_rows).arg(kc).arg(vc).arg(t0_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
1813 unsafe { b.launch(cfg)?; }
1814 Ok(())
1815 }
1816
1817 #[allow(clippy::too_many_arguments)]
1820 pub fn append_kv_quantized_row_dc_inc(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
1821 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
1822 t0_dev: &mut CudaSlice<i32>,
1823 kv_dim_k: usize, kv_dim_v: usize,
1824 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
1825 -> Result<(), Box<dyn std::error::Error>> {
1826 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_dc_inc") }
1827 else { self.func("append_quantize_kv_q8_0_q5_1_dc_inc") };
1828 let nthreads = ((kv_dim_k.max(kv_dim_v) / 32) * 32).min(1024) as u32;
1829 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (nthreads, 1, 1),
1830 shared_mem_bytes: 0 };
1831 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
1832 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
1833 let __s_b = self.gpu.stream();
1834 let mut b = __s_b.launch_builder(&f);
1835 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(t0_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
1836 unsafe { b.launch(cfg)?; }
1837 Ok(())
1838 }
1839
1840 pub fn pack_tok_p(&self, tok: &CudaSlice<u32>, p: &CudaSlice<f32>, out: &mut CudaSlice<u32>,
1842 slot: usize) -> Result<(), Box<dyn std::error::Error>> {
1843 let f = self.func("pack_tok_p");
1844 let sl = slot as i32;
1845 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1846 let __s_b = self.gpu.stream();
1847 let mut b = __s_b.launch_builder(&f);
1848 b.arg(tok).arg(p).arg(out).arg(&sl);
1849 unsafe { b.launch(cfg)?; }
1850 Ok(())
1851 }
1852 pub fn tok_map_u32(&self, tok: &mut CudaSlice<u32>, map: &CudaSlice<u32>)
1853 -> Result<(), Box<dyn std::error::Error>> {
1854 let f = self.func("tok_map_u32");
1855 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1856 let __s_b = self.gpu.stream();
1857 let mut b = __s_b.launch_builder(&f);
1858 b.arg(tok).arg(map);
1859 unsafe { b.launch(cfg)?; }
1860 Ok(())
1861 }
1862
1863 #[allow(clippy::too_many_arguments)]
1865 pub fn spec_assemble_verify(&self, tokp: &CudaSlice<u32>, pend: &CudaSlice<u32>,
1866 d2t: Option<&CudaSlice<u32>>, vtok: &mut CudaSlice<u32>,
1867 brk: &mut CudaSlice<u32>, p_min: f32, k: usize, pmin0: bool)
1868 -> Result<(), Box<dyn std::error::Error>> {
1869 let f = self.func("spec_assemble_verify");
1870 let (ki, pm) = (k as i32, if pmin0 { 1i32 } else { 0i32 });
1871 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1872 let __s_b = self.gpu.stream();
1873 let mut b = __s_b.launch_builder(&f);
1874 match d2t {
1875 Some(m) => { b.arg(tokp).arg(pend).arg(m).arg(vtok).arg(brk).arg(&p_min).arg(&ki).arg(&pm);
1876 unsafe { b.launch(cfg)?; } }
1877 None => { let null: u64 = 0;
1878 b.arg(tokp).arg(pend).arg(&null).arg(vtok).arg(brk).arg(&p_min).arg(&ki).arg(&pm);
1879 unsafe { b.launch(cfg)?; } }
1880 }
1881 Ok(())
1882 }
1883
1884 #[allow(clippy::too_many_arguments)]
1886 pub fn ssm_conv_ring_rebuild_dc(&self, qkv_tm: &CudaSlice<f32>, ring_old: &CudaSlice<f32>,
1887 conv_state: &mut CudaSlice<f32>, conv_dim: usize,
1888 acc: &CudaSlice<u32>, base: usize, t_v: usize, d_conv: usize)
1889 -> Result<(), Box<dyn std::error::Error>> {
1890 let f = self.func("ssm_conv_ring_rebuild_f32_dc");
1891 let n = conv_dim * (d_conv - 1);
1892 let cfg = LaunchConfig::for_num_elems(n as u32);
1893 let (cd, b0, tv, dc) = (conv_dim as i32, base as i32, t_v as i32, d_conv as i32);
1894 let __s_b = self.gpu.stream();
1895 let mut b = __s_b.launch_builder(&f);
1896 b.arg(qkv_tm).arg(ring_old).arg(conv_state).arg(&cd).arg(acc).arg(&b0).arg(&tv).arg(&dc);
1897 unsafe { b.launch(cfg)?; }
1898 Ok(())
1899 }
1900 #[allow(clippy::too_many_arguments)]
1901 pub fn gdn_scan_s128_dc(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
1902 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
1903 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
1904 n_head: usize, acc: &CudaSlice<u32>, base: usize, t_v: usize,
1905 scale: f32)
1906 -> Result<(), Box<dyn std::error::Error>> {
1907 let f = self.func("gdn_scan_s128_dc");
1908 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
1909 let cfg = LaunchConfig {
1910 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
1911 block_dim: (WARP, COLS_PER_BLOCK, 1),
1912 shared_mem_bytes: 0,
1913 };
1914 let (h, b0, tv) = (n_head as i32, base as i32, t_v as i32);
1915 let __s_b = self.gpu.stream();
1916 let mut b = __s_b.launch_builder(&f);
1917 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in).arg(state_out).arg(o)
1918 .arg(&h).arg(acc).arg(&b0).arg(&tv).arg(&scale);
1919 unsafe { b.launch(cfg)?; }
1920 Ok(())
1921 }
1922
1923 pub fn spec_rollback_kv(&self, len_ptrs: &CudaSlice<u64>, saved: &CudaSlice<i32>,
1925 acc: &CudaSlice<u32>, base: usize, n_layer: usize)
1926 -> Result<(), Box<dyn std::error::Error>> {
1927 let f = self.func("spec_rollback_kv");
1928 let (b, nl) = (base as i32, n_layer as i32);
1929 let cfg = LaunchConfig { grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
1930 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1931 let __s_bl = self.gpu.stream();
1932 let mut bl = __s_bl.launch_builder(&f);
1933 bl.arg(len_ptrs).arg(saved).arg(acc).arg(&b).arg(&nl);
1934 unsafe { bl.launch(cfg)?; }
1935 Ok(())
1936 }
1937
1938 pub fn spec_fork_valid(&self, acc: &CudaSlice<u32>, optimistic_pending: u32,
1940 valid: &mut CudaSlice<u32>)
1941 -> Result<(), Box<dyn std::error::Error>> {
1942 let f = self.func("spec_fork_valid");
1943 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1),
1944 shared_mem_bytes: 0 };
1945 let __s_bl = self.gpu.stream();
1946 let mut bl = __s_bl.launch_builder(&f);
1947 bl.arg(acc).arg(&optimistic_pending).arg(valid);
1948 unsafe { bl.launch(cfg)?; }
1949 Ok(())
1950 }
1951
1952 pub fn spec_fork_reconcile_kv(&self, len_ptrs: &CudaSlice<u64>, saved: &CudaSlice<i32>,
1954 acc: &CudaSlice<u32>, valid: &CudaSlice<u32>, base: usize,
1955 n_layer: usize)
1956 -> Result<(), Box<dyn std::error::Error>> {
1957 let f = self.func("spec_fork_reconcile_kv");
1958 let (b, nl) = (base as i32, n_layer as i32);
1959 let cfg = LaunchConfig { grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
1960 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1961 let __s_bl = self.gpu.stream();
1962 let mut bl = __s_bl.launch_builder(&f);
1963 bl.arg(len_ptrs).arg(saved).arg(acc).arg(valid).arg(&b).arg(&nl);
1964 unsafe { bl.launch(cfg)?; }
1965 Ok(())
1966 }
1967
1968 pub fn spec_fork_restore_f32(&self, snapshot: &CudaSlice<f32>, state: &mut CudaSlice<f32>,
1970 valid: &CudaSlice<u32>)
1971 -> Result<(), Box<dyn std::error::Error>> {
1972 assert_eq!(snapshot.len(), state.len(), "fork recurrent snapshot shape mismatch");
1973 let f = self.func("spec_fork_restore_f32");
1974 let n = state.len() as i32;
1975 let blocks = state.len().div_ceil(256).min(65535).max(1) as u32;
1976 let cfg = LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1),
1977 shared_mem_bytes: 0 };
1978 let __s_bl = self.gpu.stream();
1979 let mut bl = __s_bl.launch_builder(&f);
1980 bl.arg(snapshot).arg(state).arg(valid).arg(&n);
1981 unsafe { bl.launch(cfg)?; }
1982 Ok(())
1983 }
1984
1985 pub fn spec_seed_gather(&self, vx: &CudaSlice<f32>, fill_prev: &CudaSlice<f32>,
1988 acc: &CudaSlice<u32>, h_seed: &mut CudaSlice<f32>,
1989 base: usize, n_embd: usize)
1990 -> Result<(), Box<dyn std::error::Error>> {
1991 let f = self.func("spec_seed_gather");
1992 let (b, ne) = (base as i32, n_embd as i32);
1993 let cfg = LaunchConfig { grid_dim: (n_embd.div_ceil(256) as u32, 1, 1),
1994 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1995 let __s_bl = self.gpu.stream();
1996 let mut bl = __s_bl.launch_builder(&f);
1997 bl.arg(vx).arg(fill_prev).arg(acc).arg(h_seed).arg(&b).arg(&ne);
1998 unsafe { bl.launch(cfg)?; }
1999 Ok(())
2000 }
2001
2002
2003 pub fn spec_accept_greedy(&self, preds: &CudaSlice<u32>, draft: &CudaSlice<u32>,
2005 last_pred: u32, base: usize, k_round: usize,
2006 out: &mut CudaSlice<u32>)
2007 -> Result<(), Box<dyn std::error::Error>> {
2008 let f = self.func("spec_accept_greedy");
2009 let (b, k) = (base as i32, k_round as i32);
2010 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2011 let __s_bl = self.gpu.stream();
2012 let mut bl = __s_bl.launch_builder(&f);
2013 bl.arg(preds).arg(draft).arg(&last_pred).arg(&b).arg(&k).arg(out);
2014 unsafe { bl.launch(cfg)?; }
2015 Ok(())
2016 }
2017
2018 pub fn gumbel_perturb(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
2025 seed: u64, stream_pos: u32, temp: f32)
2026 -> Result<(), Box<dyn std::error::Error>> {
2027 let f = self.func("gumbel_perturb_f32");
2028 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2029 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2030 let __s_b = self.gpu.stream();
2031 let mut b = __s_b.launch_builder(&f);
2032 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
2033 unsafe { b.launch(cfg)?; }
2034 Ok(())
2035 }
2036
2037 pub fn mask_logits_col(&self, logits: &mut CudaSlice<f32>, mask: &CudaSlice<u32>,
2045 col: usize, n: usize, mask_words: usize)
2046 -> Result<(), Box<dyn std::error::Error>> {
2047 let f = self.func("mask_logits_f32");
2048 let (ci, ni, mw) = (col as i32, n as i32, mask_words as i32);
2049 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256).min(1024) as u32, 1, 1),
2050 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2051 let __s_b = self.gpu.stream();
2052 let mut b = __s_b.launch_builder(&f);
2053 b.arg(&mut *logits).arg(mask).arg(&ci).arg(&ni).arg(&mw);
2054 unsafe { b.launch(cfg)?; }
2055 Ok(())
2056 }
2057
2058 pub fn gumbel_perturb_col(&self, x: &CudaSlice<f32>, col: usize, y: &mut CudaSlice<f32>,
2065 n: usize, seed: u64, stream_pos: u32, temp: f32)
2066 -> Result<(), Box<dyn std::error::Error>> {
2067 let f = self.func("gumbel_perturb_f32");
2068 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2069 let col_view = x.slice(col * n..(col + 1) * n);
2070 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2071 let __s_b = self.gpu.stream();
2072 let mut b = __s_b.launch_builder(&f);
2073 b.arg(&col_view).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
2074 unsafe { b.launch(cfg)?; }
2075 Ok(())
2076 }
2077
2078 pub fn sctr_inc(&self, ctr: &mut CudaSlice<u32>) -> Result<(), Box<dyn std::error::Error>> {
2083 let f = self.func("memra_sctr_inc");
2084 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
2085 let __s_b = self.gpu.stream();
2086 let mut b = __s_b.launch_builder(&f);
2087 b.arg(&mut *ctr);
2088 unsafe { b.launch(cfg)?; }
2089 Ok(())
2090 }
2091
2092 pub fn gumbel_perturb_ctr(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
2097 seed: u64, ctr: &CudaSlice<u32>, temp: f32)
2098 -> Result<(), Box<dyn std::error::Error>> {
2099 let f = self.func("gumbel_perturb_ctr_f32");
2100 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2101 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2102 let __s_b = self.gpu.stream();
2103 let mut b = __s_b.launch_builder(&f);
2104 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(ctr).arg(&temp);
2105 unsafe { b.launch(cfg)?; }
2106 Ok(())
2107 }
2108
2109 pub fn softmax_gather(&self, x: &CudaSlice<f32>, row_stride: usize,
2113 ids: &CudaSlice<u32>, rows: &CudaSlice<i32>,
2114 out: &mut CudaSlice<f32>, n: usize, npair: usize, temp: f32)
2115 -> Result<(), Box<dyn std::error::Error>> {
2116 let f = self.func("softmax_gather_f32");
2117 let (ni, rs) = (n as i32, row_stride as i64);
2118 let np = npair as i32;
2119 let cfg = LaunchConfig { grid_dim: (npair as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2120 let __s_b = self.gpu.stream();
2121 let mut b = __s_b.launch_builder(&f);
2122 b.arg(x).arg(&rs).arg(ids).arg(rows).arg(&mut *out).arg(&ni).arg(&np).arg(&temp);
2123 unsafe { b.launch(cfg)?; }
2124 Ok(())
2125 }
2126
2127 pub fn residual_sample(&self, p: &CudaSlice<f32>, q: Option<&CudaSlice<f32>>, n: usize,
2131 temp: f32, seed: u64, stream_pos: u32,
2132 out_tok: &mut CudaSlice<u32>)
2133 -> Result<(), Box<dyn std::error::Error>> {
2134 let f = self.func("residual_sample_f32");
2135 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2136 let nth = 1024u32;
2137 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (nth, 1, 1), shared_mem_bytes: 0 };
2138 let has_q: i32 = q.is_some() as i32;
2139 let qbuf = q.unwrap_or(p); let __s_b = self.gpu.stream();
2141 let mut b = __s_b.launch_builder(&f);
2142 b.arg(p).arg(qbuf).arg(&has_q).arg(&ni).arg(&temp).arg(&slo).arg(&shi).arg(&stream_pos)
2143 .arg(&mut *out_tok);
2144 unsafe { b.launch(cfg)?; }
2145 Ok(())
2146 }
2147
2148 pub fn with_moe_cache<R>(&self, max_block_bytes: usize,
2153 f: impl FnOnce(&mut crate::moe_cache::MoeSlotCache, &Engine) -> Result<R, Box<dyn std::error::Error>>)
2154 -> Result<R, Box<dyn std::error::Error>> {
2155 let mut guard = self.moe_cache.lock().unwrap();
2156 if guard.is_none() {
2157 *guard = Some(crate::moe_cache::MoeSlotCache::new(self, max_block_bytes)?);
2158 }
2159 let cache = guard.as_mut().unwrap();
2160 f(cache, self)
2161 }
2162
2163 pub fn freeze_moe_cache(&self) {
2166 if let Some(cache) = self.moe_cache.lock().unwrap().as_mut() {
2167 cache.freeze();
2168 }
2169 }
2170
2171 pub fn export_moe_residency(&self) -> Option<Vec<(u16, u8, u16)>> {
2174 self.moe_cache
2175 .lock()
2176 .unwrap()
2177 .as_ref()
2178 .map(crate::moe_cache::MoeSlotCache::export_residency)
2179 }
2180
2181 pub(crate) fn moe_cache_frozen(&self) -> bool {
2182 self.moe_cache
2183 .lock()
2184 .unwrap()
2185 .as_ref()
2186 .is_some_and(crate::moe_cache::MoeSlotCache::is_frozen)
2187 }
2188
2189 pub fn frozen_cpu_experts_prefer_tokenwise_prime(&self) -> bool {
2196 crate::cpu_experts::configured()
2197 && self.moe_cache_frozen()
2198 && std::env::var("MEMRA_CPU_EXPERT_BATCHED_PRIME").as_deref() != Ok("1")
2199 }
2200
2201 pub(crate) fn configure_moe_cache_layout(&self, block_bytes: Vec<usize>) {
2203 assert!(
2204 self.moe_cache.lock().unwrap().is_none(),
2205 "MoE cache layout configured after cache construction"
2206 );
2207 *self.moe_cache_layout.lock().unwrap() = Some(block_bytes);
2208 }
2209
2210 pub(crate) fn moe_cache_layout(&self) -> Option<Vec<usize>> {
2211 self.moe_cache_layout.lock().unwrap().clone()
2212 }
2213
2214 pub fn moe_cache_enabled() -> bool {
2216 std::env::var("MEMRA_MOE_CACHE").as_deref() != Ok("0")
2217 }
2218
2219 pub fn moe_cache_stats(&self) -> Option<(u64, u64, u64, usize)> {
2222 let guard = self.moe_cache.lock().unwrap();
2223 guard.as_ref() .map(|c| (c.hits, c.misses, c.staged_bytes, c.n_slots()))
2224 }
2225
2226 pub fn cpu_expert_stats(
2230 &self,
2231 ) -> Option<(u64, u64, u64, u64, u64, u64, u64, u64, u64, u64, u64)> {
2232 crate::cpu_experts::configured().then(crate::cpu_experts::stats)
2233 }
2234
2235 pub fn cpu_expert_predictor_stats(&self) -> (u64, u64) {
2238 crate::cpu_experts::predictor_stats()
2239 }
2240
2241 pub fn cpu_expert_exposed_wait_ns(&self) -> Option<u64> {
2242 crate::cpu_experts::configured().then(crate::cpu_experts::exposed_wait_ns)
2243 }
2244
2245 pub fn cpu_expert_gpu_residency_stats(&self) -> Option<(u64, u64, u64)> {
2248 crate::cpu_experts::configured().then(crate::cpu_experts::incomplete_gpu_residency_stats)
2249 }
2250
2251 pub fn moe_pread_stats(&self) -> Option<(u64, u64, u64, u64, u64, u64, u64)> {
2254
2255 let guard = self.moe_cache.lock().unwrap();
2256 guard.as_ref().and_then(|cache| cache.pread_stats()).map(|stats| (
2257 stats.reads,
2258 stats.bytes,
2259 stats.read_errors,
2260 stats.short_reads,
2261 stats.fallbacks,
2262 stats.buffer_waits,
2263 stats.ring_full,
2264 ))
2265 }
2266
2267 pub fn spill_config_fallbacks(&self) -> u64 {
2269 crate::spill_pread::config_fallbacks()
2270 }
2271
2272 pub fn moe_cache_reset_counters(&self) {
2274 if let Some(c) = self.moe_cache.lock().unwrap().as_mut() { c.reset_counters(); }
2275 }
2276
2277 pub fn htod_bytes(&self, v: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2278 Ok(self.gpu.stream().clone_htod(v)?)
2279 }
2280
2281 pub fn htod_bytes_padded(&self, v: &[u8], pad: usize)
2285 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2286 let mut d = self.alloc_u8_uninit(v.len() + pad)?;
2287 {
2288 let mut view = d.slice_mut(0..v.len());
2289 self.gpu.stream().memcpy_htod(v, &mut view)?;
2290 }
2291 Ok(d)
2292 }
2293
2294 pub fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
2296 -> Result<(), Box<dyn std::error::Error>> {
2297 let mut view = dst.slice_mut(off..off + len);
2298 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2299 Ok(())
2300 }
2301
2302 pub fn copy_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &CudaSlice<u8>, len: usize)
2305 -> Result<(), Box<dyn std::error::Error>> {
2306 let mut view = dst.slice_mut(off..off + len);
2307 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2308 Ok(())
2309 }
2310
2311 pub fn copy_u8_range_into(
2313 &self,
2314 dst: &mut CudaSlice<u8>,
2315 dst_off: usize,
2316 src: &CudaSlice<u8>,
2317 src_off: usize,
2318 len: usize,
2319 ) -> Result<(), Box<dyn std::error::Error>> {
2320 let mut dst_view = dst.slice_mut(dst_off..dst_off + len);
2321 self.gpu
2322 .stream()
2323 .memcpy_dtod(&src.slice(src_off..src_off + len), &mut dst_view)?;
2324 Ok(())
2325 }
2326
2327 pub fn prepare_kv_append(
2331 &self,
2332 kv: &mut crate::cache::KvLayer,
2333 retain_from: usize,
2334 append_rows: usize,
2335 ) -> Result<usize, Box<dyn std::error::Error>> {
2336 let Some(plan) = kv
2337 .ring
2338 .as_ref()
2339 .map(|ring| ring.append_plan(kv.len, retain_from, append_rows))
2340 .transpose()?
2341 else {
2342 return Ok(kv.len);
2343 };
2344 match plan {
2345 crate::cache::KvRingAppend::Contiguous { write_row } => Ok(write_row),
2346 crate::cache::KvRingAppend::Rebase {
2347 src_row,
2348 keep_rows,
2349 new_base,
2350 write_row,
2351 } => {
2352 if keep_rows > 0 {
2353 let k_len = keep_rows * kv.k_tok_bytes;
2354 let v_len = keep_rows * kv.v_tok_bytes;
2355 let mut k_tmp = self.alloc_u8_uninit(k_len)?;
2356 let mut v_tmp = self.alloc_u8_uninit(v_len)?;
2357 self.copy_u8_range_into(
2358 &mut k_tmp,
2359 0,
2360 &kv.k,
2361 src_row * kv.k_tok_bytes,
2362 k_len,
2363 )?;
2364 self.copy_u8_range_into(
2365 &mut v_tmp,
2366 0,
2367 &kv.v,
2368 src_row * kv.v_tok_bytes,
2369 v_len,
2370 )?;
2371 self.copy_u8_into(&mut kv.k, 0, &k_tmp, k_len)?;
2372 self.copy_u8_into(&mut kv.v, 0, &v_tmp, v_len)?;
2373 }
2374 kv.ring.as_mut().unwrap().apply_rebase(new_base);
2375 Ok(write_row)
2376 }
2377 }
2378 }
2379
2380 pub fn htod_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &[u8])
2383 -> Result<(), Box<dyn std::error::Error>> {
2384 let mut view = dst.slice_mut(off..off + src.len());
2385 self.gpu.stream().memcpy_htod(src, &mut view)?;
2386 Ok(())
2387 }
2388
2389 pub fn view<'a>(&self, b: &'a CudaSlice<f32>, len: usize) -> cudarc::driver::CudaView<'a, f32> {
2390 b.slice(0..len)
2391 }
2392
2393 pub fn view_u8_range<'a>(&self, b: &'a CudaSlice<u8>, start: usize, end: usize)
2396 -> cudarc::driver::CudaView<'a, u8> {
2397 b.slice(start..end)
2398 }
2399 pub fn view_u8<'a>(&self, b: &'a CudaSlice<u8>, len: usize) -> cudarc::driver::CudaView<'a, u8> {
2400 b.slice(0..len)
2401 }
2402
2403 pub fn append_kv_quantized(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2407 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2408 kv_dim_k: usize, kv_dim_v: usize,
2409 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2410 -> Result<(), Box<dyn std::error::Error>> {
2411 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") } else { self.func("append_quantize_kv_q8_0_q5_1") };
2412 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2413 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2414 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2415 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2416 let __s_b = self.gpu.stream();
2417 let mut b = __s_b.launch_builder(&f);
2418 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2419 unsafe { b.launch(cfg)?; }
2420 Ok(())
2421 }
2422
2423 pub fn append_kv_quantized_dc(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2427 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t_dev: &CudaSlice<i32>,
2428 kv_dim_k: usize, kv_dim_v: usize,
2429 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2430 -> Result<(), Box<dyn std::error::Error>> {
2431 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2432 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
2433 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2434 if Self::pdl_on() && Self::pdl_wb_on() {
2436 use cudarc::driver::{DevicePtr, DevicePtrMut};
2437 let s = &self.gpu.stream();
2438 let (pk, _g0) = k_row.device_ptr(s); let (pv, _g1) = v_row.device_ptr(s);
2439 let (pkc, _g2) = kc.device_ptr_mut(s); let (pvc, _g3) = vc.device_ptr_mut(s);
2440 let (pt, _g4) = t_dev.device_ptr(s);
2441 let mut ps = [
2442 &pk as *const _ as *mut std::ffi::c_void, &pv as *const _ as *mut _,
2443 &pkc as *const _ as *mut _, &pvc as *const _ as *mut _,
2444 &pt as *const _ as *mut _, &kdk as *const _ as *mut _,
2445 &kdv as *const _ as *mut _, &ktb as *const _ as *mut _,
2446 &vtb as *const _ as *mut _,
2447 ];
2448 unsafe { self.launch_pdl_flash(g, "append_quantize_kv_q8_0_q5_1_dc",
2449 (nblk, 1, 1), (32, 1, 1), 0, &mut ps)?; }
2450 return Ok(());
2451 }
2452 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_dc") } else { self.func("append_quantize_kv_q8_0_q5_1_dc") };
2453 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2454 let __s_b = self.gpu.stream();
2455 let mut b = __s_b.launch_builder(&f);
2456 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(t_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2457 unsafe { b.launch(cfg)?; }
2458 Ok(())
2459 }
2460
2461 #[allow(clippy::too_many_arguments)]
2468 pub fn append_kv_quantized_rows(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
2469 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
2470 t0: usize, t: usize, kv_dim_k: usize, kv_dim_v: usize,
2471 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2472 -> Result<(), Box<dyn std::error::Error>> {
2473 if std::env::var("MEMRA_PRIME_APPEND_LOOP").is_ok() {
2474 for i in 0..t {
2475 let k_row = k_rows.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
2476 let v_row = v_rows.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
2477 self.append_kv_quantized_view(&k_row, &v_row, kc, vc, t0 + i,
2478 kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes, g)?;
2479 }
2480 return Ok(());
2481 }
2482 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_rows") } else { self.func("append_quantize_kv_q8_0_q5_1_rows") };
2483 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2484 let cfg = LaunchConfig { grid_dim: (nblk, t as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2485 let (t0i, kdk, kdv) = (t0 as i32, kv_dim_k as i32, kv_dim_v as i32);
2486 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2487 let __s_b = self.gpu.stream();
2488 let mut b = __s_b.launch_builder(&f);
2489 b.arg(k_rows).arg(v_rows).arg(kc).arg(vc).arg(&t0i).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2490 unsafe { b.launch(cfg)?; }
2491 Ok(())
2492 }
2493
2494 pub fn inc_seqlen(&self, p: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
2498 let f = self.func("inc_i32");
2499 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
2500 let __s_b = self.gpu.stream();
2501 let mut b = __s_b.launch_builder(&f);
2502 b.arg(p);
2503 unsafe { b.launch(cfg)?; }
2504 Ok(())
2505 }
2506
2507 pub fn append_kv_quantized_view(&self, k_row: &cudarc::driver::CudaView<f32>,
2510 v_row: &cudarc::driver::CudaView<f32>,
2511 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2512 kv_dim_k: usize, kv_dim_v: usize,
2513 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2514 -> Result<(), Box<dyn std::error::Error>> {
2515 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") }
2516 else { self.func("append_quantize_kv_q8_0_q5_1") };
2517 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2518 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2519 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2520 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2521 let __s_b = self.gpu.stream();
2522 let mut b = __s_b.launch_builder(&f);
2523 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2524 unsafe { b.launch(cfg)?; }
2525 Ok(())
2526 }
2527
2528 pub fn copy_view_into(&self, dst: &mut CudaSlice<f32>, off: usize,
2531 src: &cudarc::driver::CudaView<f32>, len: usize)
2532 -> Result<(), Box<dyn std::error::Error>> {
2533 let mut view = dst.slice_mut(off..off + len);
2534 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2535 Ok(())
2536 }
2537
2538 pub fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2542 let mut dst = self.gpu.stream().alloc_zeros::<f32>(src.len())?;
2543 self.gpu.stream().memcpy_dtod(src, &mut dst)?;
2544 Ok(dst)
2545 }
2546
2547 pub fn dtod_copy_view(&self, src: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<f32>)
2550 -> Result<(), Box<dyn std::error::Error>> {
2551 self.gpu.stream().memcpy_dtod(src, dst)?;
2552 Ok(())
2553 }
2554
2555 pub fn dtod_copy_view_i8(&self, src: &cudarc::driver::CudaView<i8>, dst: &mut CudaSlice<i8>)
2557 -> Result<(), Box<dyn std::error::Error>> {
2558 self.gpu.stream().memcpy_dtod(src, dst)?;
2559 Ok(())
2560 }
2561
2562 pub fn dtod_copy_into(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, offset: usize)
2564 -> Result<(), Box<dyn std::error::Error>> {
2565 let n = src.len();
2566 let mut dv = dst.slice_mut(offset..offset + n);
2567 self.gpu.stream().memcpy_dtod(src, &mut dv)?;
2568 Ok(())
2569 }
2570
2571 pub fn uninit_i8(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
2573 self.alloc_uninit::<i8>(n)
2574 }
2575
2576 pub fn qmatvec(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
2578 qtype: i32, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2579 let f = self.func("qmatvec_f32");
2580 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2582 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2583 let __s_b = self.gpu.stream();
2584 let mut b = __s_b.launch_builder(&f);
2585 b.arg(w).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2586 unsafe { b.launch(cfg)?; }
2587 Ok(y)
2588 }
2589
2590 pub fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2592 let s = self.gpu.stream().alloc_zeros::<u8>(n)?;
2593 self.keep_if_capturing(&s);
2594 Ok(s)
2595 }
2596
2597 pub fn alloc_u8_uninit(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2601 let s = unsafe { self.gpu.stream().alloc::<u8>(n)? };
2602 self.keep_if_capturing(&s);
2603 Ok(s)
2604 }
2605
2606 pub fn memset_zeros_view(&self, dst: &mut cudarc::driver::CudaViewMut<f32>)
2609 -> Result<(), Box<dyn std::error::Error>> {
2610 self.gpu.stream().memset_zeros(dst)?;
2611 Ok(())
2612 }
2613
2614 pub fn stage_expert(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2620 -> Result<(), Box<dyn std::error::Error>> {
2621 let mut dst = scratch.slice_mut(off..off + host_bytes.len()); self.gpu.stream().memcpy_htod(host_bytes, &mut dst)?; Ok(())
2624 }
2625
2626 pub fn moe_router_topk(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2632 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2633 let f = self.func("moe_router_topk_f32");
2634 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?; let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?; let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2637 shared_mem_bytes: 0 };
2638 let (ne, nu) = (n_expert as i32, n_used as i32);
2639 let __s_b = self.gpu.stream();
2640 let mut b = __s_b.launch_builder(&f);
2641 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2642 unsafe { b.launch(cfg)?; }
2643 Ok((sel_idx, sel_w))
2644 }
2645
2646 pub fn moe_router_topk_scaled(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize,
2649 n_used: usize, ex_scale: &CudaSlice<f32>)
2650 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2651 let f = self.func("moe_router_topk_scaled_f32");
2656 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
2657 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
2658 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2659 shared_mem_bytes: 0 };
2660 let (ne, nu) = (n_expert as i32, n_used as i32);
2661 let __s_b = self.gpu.stream();
2662 let mut b = __s_b.launch_builder(&f);
2663 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu).arg(ex_scale);
2664 unsafe { b.launch(cfg)?; }
2665 Ok((sel_idx, sel_w))
2666 }
2667
2668 pub fn moe_router_topk_host(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2676 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2677 let f = self.func("moe_router_topk_f32");
2678 let n = t * n_used;
2679 let mut sel_idx = self.alloc_uninit::<i32>(n)?;
2680 let mut sel_w = self.alloc_uninit::<f32>(n)?;
2681 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2682 shared_mem_bytes: 0 };
2683 let (ne, nu) = (n_expert as i32, n_used as i32);
2684 let __s_b = self.gpu.stream();
2685 let mut b = __s_b.launch_builder(&f);
2686 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2687 unsafe { b.launch(cfg)?; }
2688 let bytes = n * 8;
2690 let mut guard = self.router_stage.lock().unwrap();
2691 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
2692 *guard = Some(PinnedStage::new(bytes.max(4096))?);
2693 }
2694 let stage = guard.as_mut().unwrap();
2695 let (si, sw) = unsafe {
2696 (std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
2697 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n))
2698 };
2699 self.gpu.stream().memcpy_dtoh(&sel_idx, si)?; self.gpu.stream().memcpy_dtoh(&sel_w, sw)?; self.gpu.stream().synchronize()?; Ok((si.iter().map(|&i| i as u32).collect(), sw.to_vec()))
2703 }
2704
2705 #[allow(clippy::too_many_arguments)]
2709 pub fn moe_router_sigmoid_topk(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize,
2710 n_used: usize, active_count: usize,
2711 correction_bias: &CudaSlice<f32>,
2712 active: &CudaSlice<u8>, scaling_factor: f32, route_norm: bool)
2713 -> Result<(CudaSlice<i32>, CudaSlice<f32>),
2714 Box<dyn std::error::Error>> {
2715 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
2716 if n_expert == 0 || n_expert > 1024 || n_used == 0 || n_used > n_expert {
2717 return Err(format!(
2718 "sigmoid router shape unsupported: n_expert={n_expert}, n_used={n_used}",
2719 ).into());
2720 }
2721 if logits.len() < t * n_expert || correction_bias.len() != n_expert
2722 || active.len() != n_expert {
2723 return Err(format!(
2724 "sigmoid router buffer mismatch: logits={} bias={} active={} expected logits>={} row={}",
2725 logits.len(), correction_bias.len(), active.len(), t * n_expert, n_expert,
2726 ).into());
2727 }
2728 let f = self.func("moe_router_sigmoid_topk_f32");
2729 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
2730 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
2731 let threads = n_expert.div_ceil(32) * 32;
2732 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (threads as u32, 1, 1),
2733 shared_mem_bytes: 0 };
2734 let (ne, nu, rn) = (n_expert as i32, n_used as i32, i32::from(route_norm));
2735 let __s_b = self.gpu.stream();
2736 let mut b = __s_b.launch_builder(&f);
2737 b.arg(logits).arg(correction_bias).arg(active).arg(&mut sel_idx).arg(&mut sel_w)
2738 .arg(&ne).arg(&nu).arg(&scaling_factor).arg(&rn);
2739 unsafe { b.launch(cfg)?; }
2740 Ok((sel_idx, sel_w))
2741 }
2742
2743 #[allow(clippy::too_many_arguments)]
2746 pub fn moe_router_sigmoid_topk_host(
2747 &self,
2748 logits: &CudaSlice<f32>,
2749 t: usize,
2750 n_expert: usize,
2751 n_used: usize,
2752 active_count: usize,
2753 correction_bias: &CudaSlice<f32>,
2754 active: &CudaSlice<u8>,
2755 scaling_factor: f32,
2756 route_norm: bool,
2757 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2758 let (sel_idx, sel_w) = self.moe_router_sigmoid_topk(
2759 logits, t, n_expert, n_used, active_count, correction_bias, active, scaling_factor,
2760 route_norm,
2761 )?;
2762 let n = t * n_used;
2763 let bytes = n * 8;
2764 let mut guard = self.router_stage.lock().unwrap();
2765 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
2766 *guard = Some(PinnedStage::new(bytes.max(4096))?);
2767 }
2768 let stage = guard.as_mut().unwrap();
2769 let (si, sw) = unsafe {
2770 (std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
2771 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n))
2772 };
2773 self.gpu.stream().memcpy_dtoh(&sel_idx, si)?;
2774 self.gpu.stream().memcpy_dtoh(&sel_w, sw)?;
2775 self.gpu.stream().synchronize()?;
2776 Ok((si.iter().map(|&i| i as u32).collect(), sw.to_vec()))
2777 }
2778
2779 pub fn stage_expert_async(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2783 -> Result<cudarc::driver::CudaEvent, Box<dyn std::error::Error>> {
2784 let mut dst = scratch.slice_mut(off..off + host_bytes.len());
2785 self.copy_stream.memcpy_htod(host_bytes, &mut dst)?;
2786 Ok(self.copy_stream.record_event(None)?)
2787 }
2788
2789 pub fn compute_wait(&self, ev: &cudarc::driver::CudaEvent) -> Result<(), Box<dyn std::error::Error>> {
2791 self.gpu.stream().wait(ev)?;
2792 Ok(())
2793 }
2794
2795 pub fn qmatvec_view(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2800 x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize, out_f: usize,
2801 qtype: i32, row_bytes: usize)
2802 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2803 let f = self.func("qmatvec_f32");
2804 let wv = w.slice(range); let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2807 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2808 let __s_b = self.gpu.stream();
2809 let mut b = __s_b.launch_builder(&f);
2810 b.arg(&wv).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2811 unsafe { b.launch(cfg)?; }
2812 Ok(y)
2813 }
2814
2815 #[allow(clippy::too_many_arguments)]
2822 pub fn moe_gate_up_silu8_q8(&self, gp: WPtr8, up: WPtr8,
2826 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2827 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2828 rb_g: usize, rb_u: usize)
2829 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2830 let f = self.func("moe_gate_up_silu8_q8");
2831 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
2832 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2833 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2834 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2835 let __s_b = self.gpu.stream();
2836 let mut b = __s_b.launch_builder(&f);
2837 b.arg(&gp).arg(&up).arg(aq).arg(ad).arg(&mut act)
2838 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2839 unsafe { b.launch(cfg)?; }
2840 Ok(act)
2841 }
2842
2843 #[allow(clippy::too_many_arguments)]
2844 pub fn moe_down8_fma_q8(&self, dp: WPtr8, w: F32x8,
2845 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
2846 dst: &mut cudarc::driver::CudaViewMut<f32>,
2847 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2848 -> Result<(), Box<dyn std::error::Error>> {
2849 let f = self.func("moe_down8_fma_q8");
2850 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2851 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2852 let (inf, outf, nu, rbi) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2853 let __s_b = self.gpu.stream();
2854 let mut b = __s_b.launch_builder(&f);
2855 b.arg(&dp).arg(&w).arg(aq2).arg(ad2).arg(dst)
2856 .arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbi);
2857 unsafe { b.launch(cfg)?; }
2858 Ok(())
2859 }
2860
2861 pub fn qmatvec_expert_q8(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2863 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
2864 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize)
2865 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2866 let f = self.func("qmatvec_expert_q8");
2867 let wv = w.slice(range);
2868 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
2869 const ROWS: u32 = 4; let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, m as u32, 1),
2871 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2872 let (inf, outf, mi, rbi) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
2873 let __s_b = self.gpu.stream();
2874 let mut b = __s_b.launch_builder(&f);
2875 b.arg(&wv).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qtype).arg(&rbi);
2876 unsafe { b.launch(cfg)?; }
2877 Ok(y)
2878 }
2879
2880 pub fn moe_gate_up_silu8(&self, gp: WPtr8, up: WPtr8, x: &cudarc::driver::CudaView<f32>,
2881 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2882 rb_g: usize, rb_u: usize)
2883 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2884 let f = self.func("moe_gate_up_silu8_f32");
2885 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2887 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2888 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2889 let __s_b = self.gpu.stream();
2890 let mut b = __s_b.launch_builder(&f);
2891 b.arg(&gp).arg(&up).arg(x).arg(&mut act)
2892 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2893 unsafe { b.launch(cfg)?; }
2894 Ok(act)
2895 }
2896
2897 #[allow(clippy::too_many_arguments)]
2903 pub fn moe_down8_fma_into(&self, dp: WPtr8, w: F32x8, act: &CudaSlice<f32>,
2904 dst: &mut cudarc::driver::CudaViewMut<f32>,
2905 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2906 -> Result<(), Box<dyn std::error::Error>> {
2907 let f = self.func("moe_down8_fma_f32");
2908 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2909 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2910 let (inf, outf, nu, rbv) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2911 let __s_b = self.gpu.stream();
2912 let mut b = __s_b.launch_builder(&f);
2913 b.arg(&dp).arg(&w).arg(act).arg(dst).arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbv);
2914 unsafe { b.launch(cfg)?; }
2915 Ok(())
2916 }
2917
2918 #[allow(clippy::too_many_arguments)]
2923 #[allow(clippy::too_many_arguments)]
2938 #[allow(clippy::too_many_arguments)]
2940 pub fn moe_pairs_matvec_q8(&self, table: &CudaSlice<u64>, proj: i32,
2941 pair_tok: &CudaSlice<i32>, pair_ex: &CudaSlice<i32>,
2942 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2943 in_f: usize, out_f: usize, n_expert: usize, n_pairs: usize,
2944 qtype: i32, row_bytes: usize)
2945 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2946 let f = self.func("moe_pairs_matvec_q8");
2947 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2948 const ROWS: u32 = 4;
2949 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_pairs as u32, 1),
2950 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2951 let (inf, outf, ne, np, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2952 n_pairs as i32, row_bytes as i64);
2953 let __s_b = self.gpu.stream();
2954 let mut b = __s_b.launch_builder(&f);
2955 b.arg(table).arg(&proj).arg(pair_tok).arg(pair_ex).arg(aq).arg(ad).arg(&mut y)
2956 .arg(&inf).arg(&outf).arg(&ne).arg(&np).arg(&qtype).arg(&rbi);
2957 unsafe { b.launch(cfg)?; }
2958 Ok(y)
2959 }
2960
2961 #[allow(clippy::too_many_arguments)]
2963 pub fn moe_pairs_matvec_q8_em(&self, table: &CudaSlice<u64>, proj: i32,
2964 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2965 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2966 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2967 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2968 n_pairs: usize, qtype: i32, row_bytes: usize)
2969 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2970 let f = self.func("moe_pairs_matvec_q8_em");
2971 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2972 const ROWS: u32 = 4;
2973 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
2974 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2975 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2976 n_active as i32, row_bytes as i64);
2977 let __s_b = self.gpu.stream();
2978 let mut b = __s_b.launch_builder(&f);
2979 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
2980 .arg(aq).arg(ad).arg(&mut y)
2981 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
2982 unsafe { b.launch(cfg)?; }
2983 Ok(y)
2984 }
2985
2986 #[allow(clippy::too_many_arguments)]
2989 pub fn moe_pairs_matvec_q8_dec(&self, table: &CudaSlice<u64>, proj: i32,
2990 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2991 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2992 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2993 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2994 n_pairs: usize, qtype: i32, row_bytes: usize)
2995 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2996 let f = self.func("moe_pairs_matvec_q8_dec");
2997 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2998 const ROWS: u32 = 4;
2999 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
3000 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
3001 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
3002 n_active as i32, row_bytes as i64);
3003 let __s_b = self.gpu.stream();
3004 let mut b = __s_b.launch_builder(&f);
3005 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
3006 .arg(aq).arg(ad).arg(&mut y)
3007 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
3008 unsafe { b.launch(cfg)?; }
3009 Ok(y)
3010 }
3011
3012 pub fn moe_pairs_gelu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
3013 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3014 let f = self.func("moe_pairs_gelu_mul");
3015 let mut act = self.alloc_uninit::<f32>(n)?;
3016 let cfg = LaunchConfig::for_num_elems(n as u32);
3017 let nl = n as i64;
3018 let __s_b = self.gpu.stream();
3019 let mut b = __s_b.launch_builder(&f);
3020 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
3021 unsafe { b.launch(cfg)?; }
3022 Ok(act)
3023 }
3024
3025 pub fn moe_pairs_silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
3026 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3027 let f = self.func("moe_pairs_silu_mul");
3028 let mut act = self.alloc_uninit::<f32>(n)?;
3029 let cfg = LaunchConfig::for_num_elems(n as u32);
3030 let nl = n as i64;
3031 let __s_b = self.gpu.stream();
3032 let mut b = __s_b.launch_builder(&f);
3033 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
3034 unsafe { b.launch(cfg)?; }
3035 Ok(act)
3036 }
3037
3038 #[allow(clippy::too_many_arguments)]
3039 pub fn moe_pairs_scatter(&self, y_down: &CudaSlice<f32>, pair_w: &CudaSlice<f32>,
3040 tok_pair_off: &CudaSlice<i32>, tok_pair_ids: &CudaSlice<i32>,
3041 moe_out: &mut CudaSlice<f32>, t: usize, n_embd: usize)
3042 -> Result<(), Box<dyn std::error::Error>> {
3043 let f = self.func("moe_pairs_scatter");
3044 let cfg = LaunchConfig { grid_dim: (((n_embd + 255) / 256) as u32, t as u32, 1),
3045 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3046 let ne = n_embd as i32;
3047 let __s_b = self.gpu.stream();
3048 let mut b = __s_b.launch_builder(&f);
3049 b.arg(y_down).arg(pair_w).arg(tok_pair_off).arg(tok_pair_ids).arg(moe_out).arg(&ne);
3050 unsafe { b.launch(cfg)?; }
3051 Ok(())
3052 }
3053
3054 #[allow(clippy::too_many_arguments)]
3058 pub fn moe_gate_up_gelu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3059 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3060 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3061 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3062 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3063 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3064 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3065 rb_g as i64, rb_u as i64);
3066 let f = self.func("moe_gate_up_gelu8_dev_q8");
3067 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3068 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3069 let __s_b = self.gpu.stream();
3070 let mut b = __s_b.launch_builder(&f);
3071 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3072 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
3073 unsafe { b.launch(cfg)?; }
3074 Ok(act)
3075 }
3076
3077 #[allow(clippy::too_many_arguments)]
3079 pub fn moe_gate_up_gelu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3080 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
3081 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3082 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3083 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3084 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
3085 let (inf, nff, ne, rbg, rbu, nu) = (in_f as i32, n_ff as i32, n_expert as i32,
3086 rb_g as i64, rb_u as i64, n_used as i32);
3087 let f = self.func("moe_gate_up_gelu8_dev_q8_rows");
3088 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
3089 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3090 let __s_b = self.gpu.stream();
3091 let mut b = __s_b.launch_builder(&f);
3092 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3093 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu);
3094 unsafe { b.launch(cfg)?; }
3095 Ok(act)
3096 }
3097
3098 #[allow(clippy::too_many_arguments)]
3100 pub fn moe_gate_up_gelu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3101 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, n_pairs: usize,
3102 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3103 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3104 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3105 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
3106 let (inf, nff, ne, rbg, rbu, nu, npi) = (in_f as i32, n_ff as i32, n_expert as i32,
3107 rb_g as i64, rb_u as i64, n_used as i32,
3108 n_pairs as i32);
3109 let f = self.func("moe_gate_up_gelu8_dev_q8_csr");
3110 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
3111 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3112 let __s_b = self.gpu.stream();
3113 let mut b = __s_b.launch_builder(&f);
3114 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3115 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
3116 unsafe { b.launch(cfg)?; }
3117 Ok(act)
3118 }
3119
3120 #[allow(clippy::too_many_arguments)]
3122 pub fn moe_down8_fma_dev_q8_rows_g(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3123 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3124 dst: &mut CudaSlice<f32>, t: usize,
3125 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3126 qt: i32, rb: usize)
3127 -> Result<(), Box<dyn std::error::Error>> {
3128 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3129 n_expert as i32, rb as i64);
3130 let step_b1_w8 = t == 1 && in_f == 1280 && out_f == 4096
3134 && n_used == 8 && qt == QT_IQ4_XS;
3135 let f = self.func(if step_b1_w8 {
3136 "moe_down8_fma_dev_q8_rows_w8"
3137 } else {
3138 "moe_down8_fma_dev_q8_rows_g"
3139 });
3140 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, t as u32),
3141 block_dim: (32, if step_b1_w8 { 8 } else { 1 }, 1),
3142 shared_mem_bytes: 0 };
3143 let __s_b = self.gpu.stream();
3144 let mut b = __s_b.launch_builder(&f);
3145 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3146 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3147 unsafe { b.launch(cfg)?; }
3148 Ok(())
3149 }
3150
3151 pub fn rp_probe_q4(&self, m: usize) -> Result<(f64, f64), Box<dyn std::error::Error>> {
3155 let (out_f, in_f) = (2048usize, 2816usize);
3156 let nblk = in_f / 32;
3157 let mut seed = 0x9E3779B97F4A7C15u64;
3158 let mut rng = move || { seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407); (seed >> 33) as u8 };
3159 let mut w = vec![0u8; out_f * nblk * 18];
3160 for b in w.iter_mut() { *b = rng(); }
3161 for r in 0..out_f {
3162 for g in 0..nblk {
3163 let off = (r * nblk + g) * 18;
3164 w[off] = 0x00; w[off + 1] = 0x2C; }
3166 }
3167 let qplane = out_f * nblk * 16;
3168 let mut wrp = vec![0u8; w.len()];
3169 for r in 0..out_f {
3170 for g in 0..nblk {
3171 let src = &w[(r * nblk + g) * 18..(r * nblk + g) * 18 + 18];
3172 wrp[qplane + (r * nblk + g) * 2..qplane + (r * nblk + g) * 2 + 2]
3173 .copy_from_slice(&src[0..2]);
3174 wrp[(r * nblk + g) * 16..(r * nblk + g) * 16 + 16].copy_from_slice(&src[2..18]);
3175 }
3176 }
3177 let w_d = self.htod_bytes(&w)?;
3178 let wrp_d = self.htod_bytes(&wrp)?;
3179 let mut aq = vec![0i8; m * in_f];
3180 for v in aq.iter_mut() { *v = rng() as i8; }
3181 let aq_d = self.htod_i8(&aq)?;
3182 let ad_d = self.htod(&vec![0.03125f32; m * nblk])?;
3183 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
3184 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
3185 const RPB: u32 = 4;
3186 let cfg = LaunchConfig { grid_dim: ((out_f as u32).div_ceil(RPB), 1, 1),
3187 block_dim: (32, RPB, 1), shared_mem_bytes: 0 };
3188 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
3189 let (rb, qp) = ((nblk * 18) as i64, qplane as i64);
3190 let fb = self.func("qmatvec_q4_0_mmvq_b4");
3191 let fr = self.func("qmatvec_q4_0_mmvq_b4_rp");
3192 {
3193 let __s_b = self.gpu.stream();
3194 let mut b = __s_b.launch_builder(&fb);
3195 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3196 unsafe { b.launch(cfg)?; }
3197 let __s_b = self.gpu.stream();
3198 let mut b = __s_b.launch_builder(&fr);
3199 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1).arg(&inf).arg(&outf).arg(&mi).arg(&qp);
3200 unsafe { b.launch(cfg)?; }
3201 }
3202 self.gpu.stream().synchronize()?;
3203 let (h0, h1) = (self.dtoh(&y0)?, self.dtoh(&y1)?);
3204 let nd = h0.iter().zip(&h1).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
3205 if nd != 0 { return Err(format!("rp twin not bitwise: {nd}/{} diffs", h0.len()).into()); }
3206 let mut time = |rp: bool| -> Result<f64, Box<dyn std::error::Error>> {
3207 self.gpu.stream().synchronize()?;
3208 let t0 = std::time::Instant::now();
3209 for _ in 0..500 {
3210 if rp {
3211 let __s_b = self.gpu.stream();
3212 let mut b = __s_b.launch_builder(&fr);
3213 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1)
3214 .arg(&inf).arg(&outf).arg(&mi).arg(&qp);
3215 unsafe { b.launch(cfg)?; }
3216 } else {
3217 let __s_b = self.gpu.stream();
3218 let mut b = __s_b.launch_builder(&fb);
3219 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0)
3220 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3221 unsafe { b.launch(cfg)?; }
3222 }
3223 }
3224 self.gpu.stream().synchronize()?;
3225 Ok(t0.elapsed().as_secs_f64() * 1e6 / 500.0)
3226 };
3227 let _ = time(false)?; let _ = time(true)?; Ok((time(false)?, time(true)?))
3229 }
3230
3231 pub fn build_q4_rp4(&self, t: &mut crate::model::GpuTensor)
3236 -> Result<(), Box<dyn std::error::Error>> {
3237 use crate::model::GpuTensor;
3238 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3239 if *qtype != QT_Q4_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3240 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3241 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 18 { return Ok(()); }
3242 let nblk = in_f / 32;
3243 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 18)?;
3244 let f = self.func("q4_0_split_rp_build");
3245 let n = (out_f * nblk) as i32;
3246 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3247 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3248 let (of, nb) = (out_f as i32, nblk as i32);
3249 let _ = n;
3250 let __s_b = self.gpu.stream();
3251 let mut b = __s_b.launch_builder(&f);
3252 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3253 unsafe { b.launch(cfg)?; }
3254 *rp4 = Some(dst);
3255 Ok(())
3256 }
3257
3258 pub fn build_q8_rp4(&self, t: &mut crate::model::GpuTensor)
3263 -> Result<(), Box<dyn std::error::Error>> {
3264 use crate::model::GpuTensor;
3265 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3266 if *qtype != QT_Q8_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3267 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3268 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 34 { return Ok(()); }
3269 *rp4 = Some(self.build_q8_rp4_raw(bytes, in_f, out_f)?);
3270 Ok(())
3271 }
3272
3273 pub fn build_q8_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize)
3276 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3277 assert!(in_f % 32 == 0);
3278 let nblk = in_f / 32;
3279 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 34)?;
3280 let f = self.func("q8_0_split_rp_build");
3281 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3282 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3283 let (of, nb) = (out_f as i32, nblk as i32);
3284 let __s_b = self.gpu.stream();
3285 let mut b = __s_b.launch_builder(&f);
3286 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3287 unsafe { b.launch(cfg)?; }
3288 Ok(dst)
3289 }
3290
3291 pub fn build_q4k_rp4(&self, t: &mut crate::model::GpuTensor)
3299 -> Result<(), Box<dyn std::error::Error>> {
3300 use crate::model::GpuTensor;
3301 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3302 if *qtype != QT_Q4_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3303 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3304 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 144 { return Ok(()); }
3305 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q4_K)?);
3306 Ok(())
3307 }
3308
3309 pub fn build_q6k_rp4(&self, t: &mut crate::model::GpuTensor)
3310 -> Result<(), Box<dyn std::error::Error>> {
3311 use crate::model::GpuTensor;
3312 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3313 if *qtype != QT_Q6_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3314 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3315 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 210 { return Ok(()); }
3316 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q6_K)?);
3317 Ok(())
3318 }
3319
3320 pub fn build_kq_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize, qtype: i32)
3322 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3323 assert!(in_f % 256 == 0);
3324 let nsbk = in_f / 256;
3325 let (sb_bytes, kname) = match qtype {
3326 QT_Q4_K => (144usize, "q4_K_split_rp_build"),
3327 QT_Q6_K => (210usize, "q6_K_split_rp_build"),
3328 _ => return Err(format!("build_kq_rp4_raw: qtype {qtype} has no rp mirror").into()),
3329 };
3330 let mut dst = self.alloc_uninit::<u8>(out_f * nsbk * sb_bytes)?;
3331 let f = self.func(kname);
3332 let cfg = LaunchConfig { grid_dim: (((out_f * nsbk) as u32).div_ceil(256), 1, 1),
3333 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3334 let (of, nb) = (out_f as i32, nsbk as i32);
3335 let __s_b = self.gpu.stream();
3336 let mut b = __s_b.launch_builder(&f);
3337 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3338 unsafe { b.launch(cfg)?; }
3339 Ok(dst)
3340 }
3341
3342 pub fn kqrp_enabled() -> bool {
3346 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3347 *ON.get_or_init(|| match std::env::var("MEMRA_KQRP").as_deref() {
3348 Ok("0") => false,
3349 Ok(_) => true,
3350 Err(_) => cfg!(memra_hopper_mma),
3351 })
3352 }
3353
3354 pub fn build_q4_rp_swap(&self, t: &mut crate::model::GpuTensor)
3360 -> Result<bool, Box<dyn std::error::Error>> {
3361 self.build_q4_rp4(t)?;
3362 self.gpu.stream().synchronize()?; use crate::model::GpuTensor;
3364 let GpuTensor::Quant { bytes, rp4, rp, .. } = t else { return Ok(false) };
3365 match rp4.take() {
3366 Some(split) => {
3367 *bytes = split; *rp = true;
3369 Ok(true)
3370 }
3371 None => Ok(false),
3372 }
3373 }
3374
3375 pub fn q4rp_enabled() -> bool {
3377 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3378 *ON.get_or_init(|| std::env::var("MEMRA_Q4RP").map(|v| v != "0").unwrap_or(true))
3379 }
3380
3381 pub fn copy_rows_strided(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
3384 row_elems: usize, n_rows: usize, src_stride: usize, src_off: usize)
3385 -> Result<(), Box<dyn std::error::Error>> {
3386 let f = self.func("copy_rows_strided_f32");
3387 let cfg = LaunchConfig { grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
3388 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3389 let (re, nr) = (row_elems as i32, n_rows as i32);
3390 let (st, off) = (src_stride as i64, src_off as i64);
3391 let __s_b = self.gpu.stream();
3392 let mut b = __s_b.launch_builder(&f);
3393 b.arg(src).arg(&mut *dst).arg(&re).arg(&nr).arg(&st).arg(&off);
3394 unsafe { b.launch(cfg)?; }
3395 Ok(())
3396 }
3397
3398 pub fn u32_set_k(&self, dst: &mut CudaSlice<u32>, v: u32, idx: usize)
3400 -> Result<(), Box<dyn std::error::Error>> {
3401 let f = self.func("u32_set_k");
3402 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3403 let ii = idx as i32;
3404 let __s_b = self.gpu.stream();
3405 let mut b = __s_b.launch_builder(&f);
3406 b.arg(dst).arg(&v).arg(&ii);
3407 unsafe { b.launch(cfg)?; }
3408 Ok(())
3409 }
3410
3411 pub fn i32_add_k(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
3413 let f = self.func("i32_add_k");
3414 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3415 let __s_b = self.gpu.stream();
3416 let mut b = __s_b.launch_builder(&f);
3417 b.arg(d).arg(&v);
3418 unsafe { b.launch(cfg)?; }
3419 Ok(())
3420 }
3421
3422 pub fn i32_iota_from(&self, ctr: &CudaSlice<i32>, dst: &mut CudaSlice<i32>, n: usize)
3424 -> Result<(), Box<dyn std::error::Error>> {
3425 let f = self.func("i32_iota_from");
3426 let cfg = LaunchConfig::for_num_elems(n as u32);
3427 let ni = n as i32;
3428 let __s_b = self.gpu.stream();
3429 let mut b = __s_b.launch_builder(&f);
3430 b.arg(ctr).arg(dst).arg(&ni);
3431 unsafe { b.launch(cfg)?; }
3432 Ok(())
3433 }
3434
3435 pub fn u32_map_k(&self, buf: &mut CudaSlice<u32>, map: &CudaSlice<u32>, idx: usize)
3437 -> Result<(), Box<dyn std::error::Error>> {
3438 let f = self.func("u32_map_k");
3439 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3440 let ii = idx as i32;
3441 let __s_b = self.gpu.stream();
3442 let mut b = __s_b.launch_builder(&f);
3443 b.arg(buf).arg(map).arg(&ii);
3444 unsafe { b.launch(cfg)?; }
3445 Ok(())
3446 }
3447
3448 #[allow(clippy::too_many_arguments)]
3450 pub fn u32_pack2(&self, a: &CudaSlice<u32>, off_a: usize, n1: usize,
3451 b_in: &CudaSlice<u32>, n2: usize, out: &mut CudaSlice<u32>)
3452 -> Result<(), Box<dyn std::error::Error>> {
3453 let f = self.func("u32_pack2");
3454 let cfg = LaunchConfig::for_num_elems((n1 + n2) as u32);
3455 let (oa, i1, i2) = (off_a as i32, n1 as i32, n2 as i32);
3456 let __s_b = self.gpu.stream();
3457 let mut b = __s_b.launch_builder(&f);
3458 b.arg(a).arg(&oa).arg(&i1).arg(b_in).arg(&i2).arg(out);
3459 unsafe { b.launch(cfg)?; }
3460 Ok(())
3461 }
3462
3463 pub fn moe_w_exscale(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3465 s: &CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
3466 let f = self.func("moe_w_exscale");
3467 let cfg = LaunchConfig::for_num_elems(n as u32);
3468 let ni = n as i32;
3469 let __s_b = self.gpu.stream();
3470 let mut b = __s_b.launch_builder(&f);
3471 b.arg(w).arg(sel).arg(s).arg(&ni);
3472 unsafe { b.launch(cfg)?; }
3473 Ok(())
3474 }
3475
3476 pub fn moe_w_scale_by_expert(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3479 macros: &CudaSlice<f32>, n_expert: usize, n: usize)
3480 -> Result<(), Box<dyn std::error::Error>> {
3481 let f = self.func("moe_w_scale_by_expert");
3482 let cfg = LaunchConfig { grid_dim: (n.div_ceil(64) as u32, 1, 1),
3483 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
3484 let (ne, nn) = (n_expert as i32, n as i32);
3485 let __s_b = self.gpu.stream();
3486 let mut b = __s_b.launch_builder(&f);
3487 b.arg(w).arg(sel).arg(macros).arg(&ne).arg(&nn);
3488 unsafe { b.launch(cfg)?; }
3489 Ok(())
3490 }
3491
3492 pub fn moe_gate_up_silu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3493 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3494 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3495 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3496 macros: &CudaSlice<f32>)
3497 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3498 static GU: std::sync::OnceLock<(String, u32)> = std::sync::OnceLock::new();
3499 let (mode, wpb) = GU.get_or_init(|| {
3500 let mode = std::env::var("MEMRA_MOE_DEVQ8_GU").unwrap_or_default();
3501 let wpb = std::env::var("MEMRA_MOE_DEVQ8_WPB").ok()
3502 .and_then(|v| v.parse().ok()).unwrap_or(4u32).clamp(1, 16);
3503 (mode, wpb)
3504 });
3505 let (mode, wpb) = (mode.as_str(), *wpb);
3506 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3507 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3508 rb_g as i64, rb_u as i64);
3509 let (f, cfg) = match mode {
3510 "1" | "2" | "4" => {
3511 let rpw: u32 = mode.parse().unwrap();
3512 let f = self.func(match rpw { 1 => "moe_gate_up_silu8_dev_q8_r1",
3513 2 => "moe_gate_up_silu8_dev_q8_r2",
3514 _ => "moe_gate_up_silu8_dev_q8_r4" });
3515 let rows_per_block = (rpw * wpb) as usize;
3516 let gx = n_ff.div_ceil(rows_per_block) as u32;
3517 (f, LaunchConfig { grid_dim: (gx, n_used as u32, 1),
3518 block_dim: (32, wpb, 1), shared_mem_bytes: 0 })
3519 }
3520 "j8" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8"),
3521 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3522 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3523 "vsm2" => {
3525 let f = self.func("moe_gate_up_silu8_dev_q8_vsm2");
3526 let sh = (rb_g + rb_u) as u32;
3527 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3528 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3529 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3530 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3531 }
3532 "vsm" => {
3533 let f = self.func("moe_gate_up_silu8_dev_q8_vsm");
3534 let sh = (rb_g + rb_u) as u32;
3535 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3536 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3537 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3538 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3539 }
3540 "sg" => (self.func("moe_gate_up_silu8_dev_q8_sg"),
3541 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3542 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3543 "j8sg" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8sg"),
3544 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3545 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3546 "u64" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_u64"),
3547 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3548 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3549 "gs4" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_gs4"),
3550 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3551 block_dim: (32, 4, 1), shared_mem_bytes: 0 }),
3552 "v" | "" => (self.func("moe_gate_up_silu8_dev_q8_v"),
3554 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3555 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3556 "s2" => (self.func("moe_gate_up_silu8_dev_q8_s2"),
3557 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3558 block_dim: (32, 2, 1), shared_mem_bytes: 0 }),
3559 "s2z" => {
3560 let rz = wpb.min(16); (self.func("moe_gate_up_silu8_dev_q8_s2z"),
3562 LaunchConfig { grid_dim: (n_ff.div_ceil(rz as usize) as u32, n_used as u32, 1),
3563 block_dim: (32, 2, rz), shared_mem_bytes: 0 })
3564 }
3565 _ => (self.func("moe_gate_up_silu8_dev_q8"),
3566 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3567 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3568 };
3569 let __s_b = self.gpu.stream();
3570 let mut b = __s_b.launch_builder(&f);
3571 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3572 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3573 unsafe { b.launch(cfg)?; }
3574 Ok(act)
3575 }
3576
3577 #[allow(clippy::too_many_arguments)]
3578 pub fn moe_down8_fma_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3579 w: &cudarc::driver::CudaView<f32>,
3580 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3581 dst: &mut cudarc::driver::CudaViewMut<f32>,
3582 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3583 qt: i32, rb: usize)
3584 -> Result<(), Box<dyn std::error::Error>> {
3585 static DOWN: std::sync::OnceLock<String> = std::sync::OnceLock::new();
3586 let mode = DOWN.get_or_init(|| std::env::var("MEMRA_MOE_DEVQ8_DOWN").unwrap_or_default());
3587 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3588 n_expert as i32, rb as i64);
3589 let (f, cfg) = match mode.as_str() {
3592 m @ ("1" | "2" | "4") if n_used <= 8 => {
3593 let rpw: usize = m.parse().unwrap();
3594 let f = self.func(match rpw { 1 => "moe_down8_fma_dev_q8_w8r1",
3595 2 => "moe_down8_fma_dev_q8_w8r2",
3596 _ => "moe_down8_fma_dev_q8_w8r4" });
3597 (f, LaunchConfig { grid_dim: (out_f.div_ceil(rpw) as u32, 1, 1),
3598 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3599 }
3600 "h2" if in_f == 512 => (self.func("moe_down8_fma_dev_q8_h2"),
3601 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3602 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3603 "" if in_f == 704 && n_used <= 8 =>
3606 (self.func("moe_down8_fma_dev_q8_w8r2"),
3607 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3608 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3609 "w8h2v" | "" if in_f == 512 && n_used <= 8 =>
3613 (self.func("moe_down8_fma_dev_q8_w8h2v"),
3614 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3615 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3616 "w8h2r2v" if in_f == 512 && n_used <= 8 =>
3617 (self.func("moe_down8_fma_dev_q8_w8h2r2v"),
3618 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3619 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3620 "w8h2r2" if in_f == 512 && n_used <= 8 =>
3621 (self.func("moe_down8_fma_dev_q8_w8h2r2"),
3622 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3623 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3624 "w8h2" if in_f == 512 && n_used <= 8 =>
3625 (self.func("moe_down8_fma_dev_q8_w8h2"),
3626 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3627 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3628 _ => (self.func("moe_down8_fma_dev_q8"),
3629 LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3630 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3631 };
3632 let __s_b = self.gpu.stream();
3633 let mut b = __s_b.launch_builder(&f);
3634 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3635 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3636 unsafe { b.launch(cfg)?; }
3637 Ok(())
3638 }
3639
3640 #[allow(clippy::too_many_arguments)]
3647 pub fn moe_gate_up_silu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3648 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
3649 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3650 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3651 macros: &CudaSlice<f32>)
3652 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3653 let f = self.func("moe_gate_up_silu8_dev_q8_v_rows");
3654 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
3655 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
3656 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3657 let (inf, nff, ne, nu, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3658 n_used as i32, rb_g as i64, rb_u as i64);
3659 let __s_b = self.gpu.stream();
3660 let mut b = __s_b.launch_builder(&f);
3661 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3662 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(macros);
3663 unsafe { b.launch(cfg)?; }
3664 Ok(act)
3665 }
3666
3667 #[allow(clippy::too_many_arguments)]
3672 pub fn moe_down8_fma_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3673 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3674 dst: &mut CudaSlice<f32>, t: usize,
3675 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3676 qt: i32, rb: usize)
3677 -> Result<(), Box<dyn std::error::Error>> {
3678 assert!(in_f == 512 && n_used <= 8, "down rows twin is w8h2v shape-gated");
3679 let f = self.func("moe_down8_fma_dev_q8_w8h2v_rows");
3680 let cfg = LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, t as u32),
3681 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 };
3682 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3683 n_expert as i32, rb as i64);
3684 let __s_b = self.gpu.stream();
3685 let mut b = __s_b.launch_builder(&f);
3686 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3687 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3688 unsafe { b.launch(cfg)?; }
3689 Ok(())
3690 }
3691
3692 #[allow(clippy::too_many_arguments)]
3696 pub fn moe_gate_up_silu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3697 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3698 n_pairs: usize, in_f: usize, n_ff: usize, n_used: usize,
3699 n_expert: usize, qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3700 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3701 let f = self.func("moe_gate_up_silu8_dev_q8_csr_iq4");
3702 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
3703 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
3704 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3705 let (inf, nff, ne, nu, npi, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3706 n_used as i32, n_pairs as i32, rb_g as i64, rb_u as i64);
3707 let __s_b = self.gpu.stream();
3708 let mut b = __s_b.launch_builder(&f);
3709 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3710 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
3711 unsafe { b.launch(cfg)?; }
3712 Ok(act)
3713 }
3714
3715
3716 #[allow(clippy::too_many_arguments)]
3720 pub fn moe_down8_fma_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3721 sel: &cudarc::driver::CudaView<i32>,
3722 w: &cudarc::driver::CudaView<f32>,
3723 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3724 dst: &mut cudarc::driver::CudaViewMut<f32>,
3725 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3726 qt: i32, rb: usize)
3727 -> Result<(), Box<dyn std::error::Error>> {
3728 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3729 n_expert as i32, rb as i64);
3730 let (f, cfg) = match variant {
3731 "w8h2" | "w8h2v" => {
3732 (self.func(if variant == "w8h2" { "moe_down8_fma_dev_q8_w8h2" }
3733 else { "moe_down8_fma_dev_q8_w8h2v" }),
3734 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3735 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3736 }
3737 "w8h2r2" | "w8h2r2v" => {
3738 (self.func(if variant == "w8h2r2" { "moe_down8_fma_dev_q8_w8h2r2" }
3739 else { "moe_down8_fma_dev_q8_w8h2r2v" }),
3740 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3741 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3742 }
3743 _ => (self.func("moe_down8_fma_dev_q8"),
3744 LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3745 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3746 };
3747 let __s_b = self.gpu.stream();
3748 let mut b = __s_b.launch_builder(&f);
3749 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3750 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3751 unsafe { b.launch(cfg)?; }
3752 Ok(())
3753 }
3754
3755 #[allow(clippy::too_many_arguments)]
3757 pub fn moe_gate_up_silu8_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3758 sel: &cudarc::driver::CudaView<i32>,
3759 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3760 in_f: usize, n_ff: usize, n_used: usize,
3761 n_expert: usize, qt_g: i32, qt_u: i32,
3762 rb_g: usize, rb_u: usize)
3763 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3764 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3765 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3766 rb_g as i64, rb_u as i64);
3767 let f = self.func(if variant == "v" { "moe_gate_up_silu8_dev_q8_v" }
3768 else { "moe_gate_up_silu8_dev_q8" });
3769 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3770 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3771 let __s_b = self.gpu.stream();
3772 let mut b = __s_b.launch_builder(&f);
3773 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3774 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
3775 unsafe { b.launch(cfg)?; }
3776 Ok(act)
3777 }
3778
3779 pub fn moe_gate_up_silu8_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3780 x: &cudarc::driver::CudaView<f32>,
3781 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3782 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3783 macros: &CudaSlice<f32>)
3784 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3785 let f = self.func("moe_gate_up_silu8_dev");
3786 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3788 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3789 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3790 rb_g as i64, rb_u as i64);
3791 let __s_b = self.gpu.stream();
3792 let mut b = __s_b.launch_builder(&f);
3793 b.arg(table).arg(sel).arg(x).arg(&mut act)
3794 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3795 unsafe { b.launch(cfg)?; }
3796 Ok(act)
3797 }
3798
3799 #[allow(clippy::too_many_arguments)]
3802 pub fn moe_down8_fma_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3803 w: &cudarc::driver::CudaView<f32>, act: &CudaSlice<f32>,
3804 dst: &mut cudarc::driver::CudaViewMut<f32>,
3805 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3806 qt: i32, rb: usize)
3807 -> Result<(), Box<dyn std::error::Error>> {
3808 let f = self.func("moe_down8_fma_dev");
3809 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3810 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3811 let (inf, outf, nu, ne, rbv) = (in_f as i32, out_f as i32, n_used as i32,
3812 n_expert as i32, rb as i64);
3813 let __s_b = self.gpu.stream();
3814 let mut b = __s_b.launch_builder(&f);
3815 b.arg(table).arg(sel).arg(w).arg(act).arg(dst)
3816 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbv);
3817 unsafe { b.launch(cfg)?; }
3818 Ok(())
3819 }
3820
3821 pub fn axpy_into(&self, src: &CudaSlice<f32>, alpha: f32,
3823 dst: &mut cudarc::driver::CudaViewMut<f32>, n: usize)
3824 -> Result<(), Box<dyn std::error::Error>> {
3825 let f = self.func("axpy_f32");
3826 let cfg = LaunchConfig::for_num_elems(n as u32);
3827 let (a, ni) = (alpha, n as i32);
3828 let __s_b = self.gpu.stream();
3829 let mut b = __s_b.launch_builder(&f);
3830 b.arg(src).arg(dst).arg(&a).arg(&ni);
3831 unsafe { b.launch(cfg)?; }
3832 Ok(())
3833 }
3834
3835 pub fn add_scaled_rows(&self, src: &CudaSlice<f32>, scale: &CudaSlice<f32>,
3837 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
3838 -> Result<(), Box<dyn std::error::Error>> {
3839 let f = self.func("add_scaled_rows_f32");
3840 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
3841 let (nc, nr) = (ncols as i32, nrows as i32);
3842 let __s_b = self.gpu.stream();
3843 let mut b = __s_b.launch_builder(&f);
3844 b.arg(src).arg(scale).arg(dst).arg(&nc).arg(&nr);
3845 unsafe { b.launch(cfg)?; }
3846 Ok(())
3847 }
3848
3849 pub fn gather_rows(&self, src: &CudaSlice<f32>, idx: &CudaSlice<i32>,
3853 dst: &mut CudaSlice<f32>, ncols: usize, m_e: usize)
3854 -> Result<(), Box<dyn std::error::Error>> {
3855 let f = self.func("gather_rows_f32");
3856 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3857 let (nc, me) = (ncols as i32, m_e as i32);
3858 let __s_b = self.gpu.stream();
3859 let mut b = __s_b.launch_builder(&f);
3860 b.arg(src).arg(idx).arg(dst).arg(&nc).arg(&me);
3861 unsafe { b.launch(cfg)?; }
3862 Ok(())
3863 }
3864
3865 pub fn scatter_slot(&self, src: &CudaSlice<f32>, tok_idx: &CudaSlice<i32>,
3870 slot_idx: &CudaSlice<i32>, weight: &CudaSlice<f32>,
3871 dst: &mut CudaSlice<f32>, wbuf: &mut CudaSlice<f32>,
3872 ncols: usize, n_used: usize, m_e: usize)
3873 -> Result<(), Box<dyn std::error::Error>> {
3874 let f = self.func("scatter_add_slot_f32");
3875 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3876 let (nc, nu, me) = (ncols as i32, n_used as i32, m_e as i32);
3877 let __s_b = self.gpu.stream();
3878 let mut b = __s_b.launch_builder(&f);
3879 b.arg(src).arg(tok_idx).arg(slot_idx).arg(weight).arg(dst).arg(wbuf).arg(&nc).arg(&nu).arg(&me);
3880 unsafe { b.launch(cfg)?; }
3881 Ok(())
3882 }
3883
3884 pub fn reduce_slots(&self, slots: &CudaSlice<f32>, wbuf: &CudaSlice<f32>,
3888 dst: &mut CudaSlice<f32>, ncols: usize, n_used: usize, t: usize)
3889 -> Result<(), Box<dyn std::error::Error>> {
3890 let f = self.func("reduce_slots_f32");
3891 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
3892 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
3893 let __s_b = self.gpu.stream();
3894 let mut b = __s_b.launch_builder(&f);
3895 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
3896 unsafe { b.launch(cfg)?; }
3897 Ok(())
3898 }
3899
3900 pub fn quantize_q8_1_view(&self, x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize)
3907 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3908 let f = self.func("quantize_q8_1");
3909 let nblk = in_f / 32;
3910 let mut q = self.alloc_uninit::<i8>(m * in_f)?;
3911 let mut d = self.alloc_uninit::<f32>(m * nblk)?;
3912 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
3913 let (inf, mi) = (in_f as i32, m as i32);
3914 let __s_b = self.gpu.stream();
3915 let mut b = __s_b.launch_builder(&f);
3916 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3917 unsafe { b.launch(cfg)?; }
3918 Ok((q, d))
3919 }
3920
3921 pub fn quantize_q8_1(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3922 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3923 let nblk = in_f / 32;
3924 let mut q = self.alloc_uninit::<i8>(m * in_f)?; let mut d = self.alloc_uninit::<f32>(m * nblk)?; let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
3928 let (inf, mi) = (in_f as i32, m as i32);
3929 if Self::pdl_on() && Self::pdl_wb_on() {
3930 {
3931 use cudarc::driver::{DevicePtr, DevicePtrMut};
3932 let s = &self.gpu.stream();
3933 let (px, _g0) = x.device_ptr(s);
3934 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
3935 let mut ps = [
3936 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
3937 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
3938 &mi as *const _ as *mut _,
3939 ];
3940 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
3941 }
3942 return Ok((q, d));
3943 }
3944 let f = self.func("quantize_q8_1");
3945 let __s_b = self.gpu.stream();
3946 let mut b = __s_b.launch_builder(&f);
3947 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3948 unsafe { b.launch(cfg)?; }
3949 Ok((q, d))
3950 }
3951
3952 pub fn quantize_fp4_act(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3956 -> Result<(CudaSlice<u32>, CudaSlice<u8>), Box<dyn std::error::Error>> {
3957 let f = self.func("quantize_fp4_act");
3958 let nb16 = in_f / 16;
3959 let mut aq4 = self.alloc_uninit::<u32>(m * (in_f / 8))?; let mut ad4 = self.alloc_uninit::<u8>(m * nb16)?; let cfg = LaunchConfig::for_num_elems((m * nb16) as u32);
3962 let (inf, mi) = (in_f as i32, m as i32);
3963 let __s_b = self.gpu.stream();
3964 let mut b = __s_b.launch_builder(&f);
3965 b.arg(x).arg(&mut aq4).arg(&mut ad4).arg(&inf).arg(&mi);
3966 unsafe { b.launch(cfg)?; }
3967 Ok((aq4, ad4))
3968 }
3969
3970 pub fn qmatvec_gemm_nvfp4_fp4(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3975 in_f: usize, out_f: usize, row_bytes: usize, scale: f32)
3976 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3977 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
3978 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
3979 let mut y = self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)?;
3980 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
3981 Ok(y)
3982 }
3983
3984 fn fp4_gemm_launch(&self, bytes: &CudaSlice<u8>, aq4: &CudaSlice<u32>, ad4: &CudaSlice<u8>,
3987 m: usize, in_f: usize, out_f: usize, row_bytes: usize)
3988 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3989 let f = self.func("qmatvec_gemm_nvfp4_fp4");
3990 let mut y = self.alloc_uninit::<f32>(m * out_f)?; const BM: u32 = 64; const BN: u32 = 256;
3992 let cfg = LaunchConfig {
3993 grid_dim: ((out_f as u32 + BM - 1) / BM, (m as u32 + BN - 1) / BN, 1),
3994 block_dim: (32, 4, 1), shared_mem_bytes: 0,
3995 };
3996 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3997 let __s_b = self.gpu.stream();
3998 let mut b = __s_b.launch_builder(&f);
3999 b.arg(bytes).arg(aq4).arg(ad4).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
4000 unsafe { b.launch(cfg)?; }
4001 Ok(y)
4002 }
4003
4004 pub fn qmatvec_gemm_nvfp4_fp4_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
4006 in_f: usize, out_f: usize, row_bytes: usize)
4007 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4008 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
4009 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
4010 self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)
4011 }
4012
4013 pub fn qmatvec_q8_0_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4015 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4016 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
4017 let f = self.func("qmatvec_q8_0_dp4a");
4018 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
4020 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
4021 let __s_b = self.gpu.stream();
4022 let mut b = __s_b.launch_builder(&f);
4023 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
4024 unsafe { b.launch(cfg)?; }
4025 Ok(y)
4026 }
4027
4028 #[allow(non_snake_case)] pub fn qmatvec_q4_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4031 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4032 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
4033 let f = self.func("qmatvec_q4_K_dp4a");
4034 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
4036 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
4037 let __s_b = self.gpu.stream();
4038 let mut b = __s_b.launch_builder(&f);
4039 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
4040 unsafe { b.launch(cfg)?; }
4041 Ok(y)
4042 }
4043
4044 #[allow(non_snake_case)] pub fn qmatvec_q6_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4047 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4048 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
4049 let f = self.func("qmatvec_q6_K_dp4a");
4050 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
4052 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
4053 let __s_b = self.gpu.stream();
4054 let mut b = __s_b.launch_builder(&f);
4055 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
4056 unsafe { b.launch(cfg)?; }
4057 Ok(y)
4058 }
4059
4060 #[allow(non_snake_case)] pub fn qmatvec_q5_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4063 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4064 self.qmatvec_dp4a_named("qmatvec_q5_K_dp4a", w, x, m, in_f, out_f, row_bytes)
4065 }
4066 #[allow(non_snake_case)] pub fn qmatvec_q3_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4069 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4070 self.qmatvec_dp4a_named("qmatvec_q3_K_dp4a", w, x, m, in_f, out_f, row_bytes)
4071 }
4072 pub fn qmatvec_nvfp4_fast_rp(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4074 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4075 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
4076 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a_rp", w, x, m, in_f, out_f, row_bytes)
4077 }
4078 pub fn qmatvec_nvfp4_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4080 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4081 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
4084 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a", w, x, m, in_f, out_f, row_bytes)
4085 }
4086 #[allow(non_snake_case)] pub fn qmatvec_iq4_XS_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4089 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4090 self.qmatvec_dp4a_named("qmatvec_iq4_XS_dp4a", w, x, m, in_f, out_f, row_bytes)
4091 }
4092
4093 fn qmatvec_dp4a_named(&self, name: &str, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
4095 in_f: usize, out_f: usize, row_bytes: usize)
4096 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4097 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
4098 let f = self.func(name);
4099 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
4101 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
4102 let __s_b = self.gpu.stream();
4103 let mut b = __s_b.launch_builder(&f);
4104 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
4105 unsafe { b.launch(cfg)?; }
4106 Ok(y)
4107 }
4108
4109 pub fn htod(&self, v: &[f32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4110 Ok(self.gpu.stream().clone_htod(v)?)
4111 }
4112 pub fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
4113 Ok(self.gpu.stream().clone_htod(v)?)
4114 }
4115 pub fn htod_i8(&self, v: &[i8]) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
4117 Ok(self.gpu.stream().clone_htod(v)?)
4118 }
4119 pub fn htod_u64(&self, v: &[u64]) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
4120 Ok(self.gpu.stream().clone_htod(v)?)
4121 }
4122 pub fn dtoh_view(&self, d: &cudarc::driver::CudaView<f32>)
4124 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
4125 let v = self.gpu.stream().clone_dtoh(d)?;
4126 self.gpu.stream().synchronize()?;
4127 Ok(v)
4128 }
4129 pub fn dtoh(&self, d: &CudaSlice<f32>) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
4130 let v = self.gpu.stream().clone_dtoh(d)?;
4131 self.gpu.stream().synchronize()?;
4132 Ok(v)
4133 }
4134 pub fn dtoh_pair(
4138 &self,
4139 a: &CudaSlice<f32>,
4140 b: &CudaSlice<f32>,
4141 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
4142 let av = self.gpu.stream().clone_dtoh(a)?;
4143 let bv = self.gpu.stream().clone_dtoh(b)?;
4144 self.gpu.stream().synchronize()?;
4145 Ok((av, bv))
4146 }
4147 pub fn dtoh_i32(&self, d: &CudaSlice<i32>) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
4149 let v = self.gpu.stream().clone_dtoh(d)?;
4150 self.gpu.stream().synchronize()?;
4151 Ok(v)
4152 }
4153 pub fn dtoh_u8(&self, d: &CudaSlice<u8>) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
4155 let v = self.gpu.stream().clone_dtoh(d)?;
4156 self.gpu.stream().synchronize()?;
4157 Ok(v)
4158 }
4159 pub fn dtoh_u8_view(&self, d: &cudarc::driver::CudaView<u8>)
4160 -> Result<Vec<u8>, Box<dyn std::error::Error>> {
4161 let v = self.gpu.stream().clone_dtoh(d)?;
4162 self.gpu.stream().synchronize()?;
4163 Ok(v)
4164 }
4165 pub fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4166 let s = self.gpu.stream().alloc_zeros::<f32>(n)?;
4167 self.keep_if_capturing(&s);
4168 Ok(s)
4169 }
4170
4171 pub fn prob_of_token_device(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>, n_vocab: usize)
4180 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4181 let nb = ARGMAX_NB;
4182 let mut part = self.alloc_uninit::<f32>(nb)?;
4183 let mut p = self.alloc_uninit::<f32>(1)?;
4184 let f1 = self.func("prob_of_token_partial_f32");
4185 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4186 let nv = n_vocab as i32;
4187 let __s_b1 = self.gpu.stream();
4188 let mut b1 = __s_b1.launch_builder(&f1);
4189 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
4190 unsafe { b1.launch(cfg1)?; }
4191 let f2 = self.func("prob_of_token_final_f32");
4192 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4193 let nbi = nb as i32;
4194 let __s_b2 = self.gpu.stream();
4195 let mut b2 = __s_b2.launch_builder(&f2);
4196 b2.arg(&part).arg(&mut p).arg(&nbi);
4197 unsafe { b2.launch(cfg2)?; }
4198 Ok(p)
4199 }
4200
4201 pub fn prob_of_token_device_col(&self, logits: &CudaSlice<f32>,
4208 tok_all: &CudaSlice<u32>, tok_idx: usize,
4209 p_out: &mut CudaSlice<f32>, p_idx: usize, n_vocab: usize)
4210 -> Result<(), Box<dyn std::error::Error>> {
4211 let tok_v = tok_all.slice(tok_idx..tok_idx + 1);
4212 let mut p_v = p_out.slice_mut(p_idx..p_idx + 1);
4213 let nb = ARGMAX_NB;
4214 let mut part = self.alloc_uninit::<f32>(nb)?;
4215 let f1 = self.func("prob_of_token_partial_f32");
4216 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4217 let nv = n_vocab as i32;
4218 let __s_b1 = self.gpu.stream();
4219 let mut b1 = __s_b1.launch_builder(&f1);
4220 b1.arg(logits).arg(&tok_v).arg(&mut part).arg(&nv);
4221 unsafe { b1.launch(cfg1)?; }
4222 let f2 = self.func("prob_of_token_final_f32");
4223 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4224 let nbi = nb as i32;
4225 let __s_b2 = self.gpu.stream();
4226 let mut b2 = __s_b2.launch_builder(&f2);
4227 b2.arg(&part).arg(&mut p_v).arg(&nbi);
4228 unsafe { b2.launch(cfg2)?; }
4229 Ok(())
4230 }
4231
4232 pub fn prob_of_token_device_into(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>,
4233 p_out: &mut CudaSlice<f32>, n_vocab: usize)
4234 -> Result<(), Box<dyn std::error::Error>> {
4235 let nb = ARGMAX_NB;
4236 let mut part = self.alloc_uninit::<f32>(nb)?;
4237 let f1 = self.func("prob_of_token_partial_f32");
4238 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4239 let nv = n_vocab as i32;
4240 let __s_b1 = self.gpu.stream();
4241 let mut b1 = __s_b1.launch_builder(&f1);
4242 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
4243 unsafe { b1.launch(cfg1)?; }
4244 let f2 = self.func("prob_of_token_final_f32");
4245 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4246 let nbi = nb as i32;
4247 let __s_b2 = self.gpu.stream();
4248 let mut b2 = __s_b2.launch_builder(&f2);
4249 b2.arg(&part).arg(p_out).arg(&nbi);
4250 unsafe { b2.launch(cfg2)?; }
4251 Ok(())
4252 }
4253
4254 pub fn argmax_token_device(&self, logits: &CudaSlice<f32>, n_vocab: usize)
4255 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4256 let mut tok = unsafe { self.gpu.stream().alloc::<u32>(1)? };
4257 self.argmax_token_device_into(logits, &mut tok, n_vocab)?;
4258 Ok(tok)
4259 }
4260 pub fn argmax_token_device_into(&self, logits: &CudaSlice<f32>, tok: &mut CudaSlice<u32>,
4267 n_vocab: usize) -> Result<(), Box<dyn std::error::Error>> {
4268 let nb = ARGMAX_NB;
4269 let f1 = self.func("argmax_partial_f32");
4270 let f2 = self.func("argmax_final_f32");
4271 let mut guard = self.argmax_partials.lock().unwrap();
4272 if guard.is_none() {
4273 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4276 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4277 *guard = Some((pv, pi));
4278 }
4279 let (part_v, part_i) = guard.as_mut().unwrap();
4280 let nv = n_vocab as i32;
4281 let nbi = nb as i32;
4282 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4284 let __s_b1 = self.gpu.stream();
4285 let mut b1 = __s_b1.launch_builder(&f1);
4286 b1.arg(logits).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4287 unsafe { b1.launch(cfg1)?; }
4288 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4290 let __s_b2 = self.gpu.stream();
4291 let mut b2 = __s_b2.launch_builder(&f2);
4292 b2.arg(&*part_v).arg(&*part_i).arg(tok).arg(&nbi);
4293 unsafe { b2.launch(cfg2)?; }
4294 Ok(())
4295 }
4296 pub fn argmax_token_device_col(&self, logits: &CudaSlice<f32>, col: usize, n_vocab: usize,
4302 toks: &mut CudaSlice<u32>, out_idx: usize)
4303 -> Result<(), Box<dyn std::error::Error>> {
4304 let nb = ARGMAX_NB;
4305 let f1 = self.func("argmax_partial_f32");
4306 let f2 = self.func("argmax_final_f32");
4307 let mut guard = self.argmax_partials.lock().unwrap();
4308 if guard.is_none() {
4309 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4310 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4311 *guard = Some((pv, pi));
4312 }
4313 let (part_v, part_i) = guard.as_mut().unwrap();
4314 let col_view = logits.slice(col * n_vocab..(col + 1) * n_vocab);
4315 let nv = n_vocab as i32;
4316 let nbi = nb as i32;
4317 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4318 let __s_b1 = self.gpu.stream();
4319 let mut b1 = __s_b1.launch_builder(&f1);
4320 b1.arg(&col_view).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4321 unsafe { b1.launch(cfg1)?; }
4322 let mut tok_view = toks.slice_mut(out_idx..out_idx + 1);
4323 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4324 let __s_b2 = self.gpu.stream();
4325 let mut b2 = __s_b2.launch_builder(&f2);
4326 b2.arg(&*part_v).arg(&*part_i).arg(&mut tok_view).arg(&nbi);
4327 unsafe { b2.launch(cfg2)?; }
4328 Ok(())
4329 }
4330 pub fn htod_u32_v(&self, v: &[u32]) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4332 Ok(self.gpu.stream().clone_htod(v)?)
4333 }
4334 pub fn dtoh_u32(&self, d: &CudaSlice<u32>) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
4335 let v = self.gpu.stream().clone_dtoh(d)?;
4336 self.gpu.stream().synchronize()?;
4337 Ok(v)
4338 }
4339 pub fn htod_u32_into(&self, dst: &mut CudaSlice<u32>, src: &[u32])
4343 -> Result<(), Box<dyn std::error::Error>> {
4344 let mut view = dst.slice_mut(0..src.len());
4345 self.gpu.stream().memcpy_htod(src, &mut view)?;
4346 Ok(())
4347 }
4348
4349 pub fn htod_i32_into(&self, dst: &mut CudaSlice<i32>, src: &[i32])
4352 -> Result<(), Box<dyn std::error::Error>> {
4353 let mut view = dst.slice_mut(0..src.len());
4354 self.gpu.stream().memcpy_htod(src, &mut view)?;
4355 Ok(())
4356 }
4357
4358 pub fn alloc_u32_zeroed(&self, n: usize) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4359 let s = self.gpu.stream().alloc_zeros::<u32>(n)?;
4360 self.keep_if_capturing(&s);
4361 Ok(s)
4362 }
4363 pub fn embed_gather_device_into(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4366 x_out: &mut CudaSlice<f32>, n_embd: usize, qtype: i32,
4367 row_bytes: usize) -> Result<(), Box<dyn std::error::Error>> {
4368 let f = self.func("embed_gather_u32");
4369 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
4370 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4371 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
4372 let __s_b = self.gpu.stream();
4373 let mut b = __s_b.launch_builder(&f);
4374 b.arg(embd).arg(token_d).arg(x_out).arg(&ne).arg(&qt).arg(&rb);
4375 unsafe { b.launch(cfg)?; }
4376 Ok(())
4377 }
4378 pub fn dtoh_i32_one(&self, d: &CudaSlice<i32>) -> Result<i32, Box<dyn std::error::Error>> {
4380 let v = self.gpu.stream().clone_dtoh(d)?;
4381 self.gpu.stream().synchronize()?;
4382 Ok(v[0])
4383 }
4384 pub fn i32_set_k(&self, dst: &mut CudaSlice<i32>, v: i32)
4391 -> Result<(), Box<dyn std::error::Error>> {
4392 let f = self.func("i32_set_k");
4393 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
4394 let idx = 0i32;
4395 let __s_b = self.gpu.stream();
4396 let mut b = __s_b.launch_builder(&f);
4397 b.arg(dst).arg(&v).arg(&idx);
4398 unsafe { b.launch(cfg)?; }
4399 Ok(())
4400 }
4401
4402 pub fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
4403 self.gpu.stream().memcpy_htod(&[v], d)?;
4404 Ok(())
4405 }
4406 pub fn set_u32_one(&self, d: &mut CudaSlice<u32>, v: u32) -> Result<(), Box<dyn std::error::Error>> {
4409 self.gpu.stream().memcpy_htod(&[v], d)?;
4410 Ok(())
4411 }
4412 pub fn dtoh_u32_one(&self, d: &CudaSlice<u32>) -> Result<u32, Box<dyn std::error::Error>> {
4414 let v = self.gpu.stream().clone_dtoh(d)?;
4415 self.gpu.stream().synchronize()?;
4416 Ok(v[0])
4417 }
4418 pub fn upload_u8(&self, bytes: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
4420 Ok(self.gpu.stream().clone_htod(bytes)?)
4421 }
4422 pub fn embed_gather_device(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4426 n_embd: usize, qtype: i32, row_bytes: usize)
4427 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4428 let f = self.func("embed_gather_u32");
4429 let mut x = self.alloc_uninit::<f32>(n_embd)?;
4430 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
4431 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4432 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
4433 let __s_b = self.gpu.stream();
4434 let mut b = __s_b.launch_builder(&f);
4435 b.arg(embd).arg(token_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb);
4436 unsafe { b.launch(cfg)?; }
4437 Ok(x)
4438 }
4439
4440
4441 pub fn embed_gather_device_t(&self, embd: &CudaSlice<u8>, tokens: &[u32],
4445 n_embd: usize, qtype: i32, row_bytes: usize)
4446 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4447 let t = tokens.len();
4448 let tok_d = self.gpu.stream().clone_htod(tokens)?;
4449 let f = self.func("embed_gather_u32_t");
4450 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4451 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4452 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4453 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4454 let __s_b = self.gpu.stream();
4455 let mut b = __s_b.launch_builder(&f);
4456 b.arg(embd).arg(&tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4457 unsafe { b.launch(cfg)?; }
4458 Ok(x)
4459 }
4460
4461 pub fn embed_gather_device_tv(&self, embd: &CudaSlice<u8>, tok_v: &cudarc::driver::CudaView<u32>,
4466 t: usize, n_embd: usize, qtype: i32, row_bytes: usize)
4467 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4468 let f = self.func("embed_gather_u32_t");
4469 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4470 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4471 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4472 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4473 let __s_b = self.gpu.stream();
4474 let mut b = __s_b.launch_builder(&f);
4475 b.arg(embd).arg(tok_v).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4476 unsafe { b.launch(cfg)?; }
4477 Ok(x)
4478 }
4479
4480 pub fn embed_gather_device_td(&self, embd: &CudaSlice<u8>, tok_d: &CudaSlice<u32>, t: usize,
4481 n_embd: usize, qtype: i32, row_bytes: usize)
4482 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4483 let f = self.func("embed_gather_u32_t");
4484 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4485 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4486 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4487 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4488 let __s_b = self.gpu.stream();
4489 let mut b = __s_b.launch_builder(&f);
4490 b.arg(embd).arg(tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4491 unsafe { b.launch(cfg)?; }
4492 Ok(x)
4493 }
4494
4495 #[inline]
4501 fn keep_if_capturing<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, s: &CudaSlice<T>) {
4503 if self.capture_keep_on.load(std::sync::atomic::Ordering::Relaxed) {
4504 self.capture_keep.lock().unwrap().push(Box::new(s.clone()));
4505 }
4506 }
4507
4508 fn alloc_uninit<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, n: usize)
4509 -> Result<CudaSlice<T>, Box<dyn std::error::Error>> {
4510 let mut s = unsafe { self.gpu.stream().alloc::<T>(n)? };
4511 {
4515 static Z: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4516 if *Z.get_or_init(|| std::env::var("MEMRA_DEBUG_ZERO_ALLOCS").as_deref() == Ok("1")) {
4517 use cudarc::driver::DevicePtrMut;
4519 let n_bytes = s.len() * std::mem::size_of::<T>();
4520 let stream = self.gpu.stream();
4521 let (p_, _g) = s.device_ptr_mut(&stream);
4522 unsafe {
4523 cudarc::driver::sys::cuMemsetD8Async(p_, 0, n_bytes, stream.cu_stream())
4524 .result()?;
4525 }
4526 }
4527 }
4528 self.keep_if_capturing(&s);
4529 Ok(s)
4530 }
4531
4532 pub fn uninit_q8_pair(&self, n: usize)
4537 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4538 Ok((self.alloc_uninit::<i8>(n)?, self.alloc_uninit::<f32>(n / 32)?))
4539 }
4540
4541 pub fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4542 self.alloc_uninit::<f32>(n)
4543 }
4544
4545 pub fn alloc_i8_uninit(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
4547 self.alloc_uninit::<i8>(n)
4548 }
4549
4550 #[allow(clippy::too_many_arguments)]
4554 pub fn rms_norm3(&self, x: &CudaSlice<f32>, w0: &CudaSlice<f32>, w1: &CudaSlice<f32>,
4555 w2: &CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4556 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4557 -> Result<(), Box<dyn std::error::Error>> {
4558 let f = self.func("rms_norm3_f32");
4559 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4560 let (nc, e) = (ncols as i32, eps);
4561 let __s_b = self.gpu.stream();
4562 let mut b = __s_b.launch_builder(&f);
4563 b.arg(x).arg(w0).arg(w1).arg(w2).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e);
4564 unsafe { b.launch(cfg)?; }
4565 Ok(())
4566 }
4567
4568 #[allow(clippy::too_many_arguments)]
4570 pub fn qkvnorm_w_on_prefill(rows: usize, ncols: usize) -> bool {
4573 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4574 *WARP_ON.get_or_init(|| {
4575 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4576 }) && ncols % 4 == 0 && rows >= 64
4577 }
4578
4579 #[allow(clippy::too_many_arguments)]
4582 pub fn rms_norm_qkv_w4b(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4583 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4584 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4585 dvb: &mut CudaSlice<u8>,
4586 ncols: usize, rq: usize, rk: usize, eps: f32, vf16: bool)
4587 -> Result<(), Box<dyn std::error::Error>> {
4588 assert!(ncols % 4 == 0 && rq + 2 * rk >= 64);
4589 let f = self.func("rms_norm_qkv_w4b_f32");
4590 let rows = (rq + 2 * rk) as u32;
4591 let cfg = LaunchConfig {
4592 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4593 };
4594 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4595 let vf = vf16 as i32;
4596 let __s_b = self.gpu.stream();
4597 let mut b = __s_b.launch_builder(&f);
4598 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv).arg(&mut *dvb)
4599 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e).arg(&vf);
4600 unsafe { b.launch(cfg)?; }
4601 Ok(())
4602 }
4603
4604 pub fn rms_norm_qkv(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4605 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4606 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4607 ncols: usize, rq: usize, rk: usize, eps: f32)
4608 -> Result<(), Box<dyn std::error::Error>> {
4609 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4613 let warp_on = *WARP_ON.get_or_init(|| {
4614 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4615 });
4616 if warp_on && ncols % 4 == 0 && rq + 2 * rk >= 64 {
4619 let f = self.func("rms_norm_qkv_w4_f32");
4620 let rows = (rq + 2 * rk) as u32;
4621 let cfg = LaunchConfig {
4622 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4623 };
4624 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4625 let __s_b = self.gpu.stream();
4626 let mut b = __s_b.launch_builder(&f);
4627 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4628 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e);
4629 unsafe { b.launch(cfg)?; }
4630 return Ok(());
4631 }
4632 let f = self.func("rms_norm_qkv_f32");
4633 let grid = (rq + 2 * rk) as u32;
4634 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4635 let (nc, rqi, rki, e) = (ncols as i32, rq as i32, rk as i32, eps);
4636 let __s_b = self.gpu.stream();
4637 let mut b = __s_b.launch_builder(&f);
4638 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4639 .arg(&nc).arg(&rqi).arg(&rki).arg(&e);
4640 unsafe { b.launch(cfg)?; }
4641 Ok(())
4642 }
4643
4644 #[allow(clippy::too_many_arguments)]
4646 pub fn rms_norm2x(&self, a: &CudaSlice<f32>, bb: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4647 wb: &CudaSlice<f32>, da: &mut CudaSlice<f32>, db: &mut CudaSlice<f32>,
4648 ncols: usize, nrows: usize, eps: f32)
4649 -> Result<(), Box<dyn std::error::Error>> {
4650 let f = self.func("rms_norm2x_f32");
4651 let cfg = LaunchConfig { grid_dim: (2 * nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4652 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
4653 let __s_b = self.gpu.stream();
4654 let mut b = __s_b.launch_builder(&f);
4655 b.arg(a).arg(bb).arg(wa).arg(wb).arg(da).arg(db).arg(&nc).arg(&nr).arg(&e);
4656 unsafe { b.launch(cfg)?; }
4657 Ok(())
4658 }
4659
4660 pub fn softcap(&self, y: &mut CudaSlice<f32>, cap: f32, n: usize)
4662 -> Result<(), Box<dyn std::error::Error>> {
4663 let f = self.func("softcap_f32");
4664 let cfg = LaunchConfig::for_num_elems(n as u32);
4665 let ni = n as i32;
4666 let __s_b = self.gpu.stream();
4667 let mut b = __s_b.launch_builder(&f);
4668 b.arg(y).arg(&cap).arg(&ni);
4669 unsafe { b.launch(cfg)?; }
4670 Ok(())
4671 }
4672
4673 pub fn mask_ids_rows(&self, y: &mut CudaSlice<f32>, ids: &CudaSlice<i32>, n_ids: usize,
4676 n_vocab: usize, t: usize)
4677 -> Result<(), Box<dyn std::error::Error>> {
4678 let f = self.func("mask_ids_rows_f32");
4679 let cfg = LaunchConfig::for_num_elems((n_ids * t) as u32);
4680 let (ni, nv, ti) = (n_ids as i32, n_vocab as i32, t as i32);
4681 let __s_b = self.gpu.stream();
4682 let mut b = __s_b.launch_builder(&f);
4683 b.arg(y).arg(ids).arg(&ni).arg(&nv).arg(&ti);
4684 unsafe { b.launch(cfg)?; }
4685 Ok(())
4686 }
4687
4688 #[allow(clippy::too_many_arguments)]
4690 pub fn add_scale_rms_norm(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4691 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4692 ncols: usize, nrows: usize, eps: f32)
4693 -> Result<(), Box<dyn std::error::Error>> {
4694 let f = self.func("add_scale_rms_norm_f32");
4695 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4696 let (nc, e2) = (ncols as i32, eps);
4697 let __s_b = self.gpu.stream();
4698 let mut b = __s_b.launch_builder(&f);
4699 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(dst).arg(&nc).arg(&e2);
4700 unsafe { b.launch(cfg)?; }
4701 Ok(())
4702 }
4703
4704 #[allow(clippy::too_many_arguments)]
4707 pub fn add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4708 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4709 ncols: usize, nrows: usize, eps: f32)
4710 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4711 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4712 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4713 let (nc, e2) = (ncols as i32, eps);
4714 if Self::pdl_on() && Self::pdl_wb_on() {
4715 {
4716 use cudarc::driver::{DevicePtr, DevicePtrMut};
4717 let s = &self.gpu.stream();
4718 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4719 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4720 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4721 let mut ps = [
4722 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4723 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4724 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4725 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4726 &e2 as *const _ as *mut _,
4727 ];
4728 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4729 (rms_block(), 1, 1), &mut ps)?; }
4730 }
4731 return Ok((out_q, out_d));
4732 }
4733 let f = self.func("add_scale_rms_norm_q8_1");
4734 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4735 let __s_b = self.gpu.stream();
4736 let mut b = __s_b.launch_builder(&f);
4737 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e2);
4738 unsafe { b.launch(cfg)?; }
4739 Ok((out_q, out_d))
4740 }
4741
4742 #[allow(clippy::too_many_arguments)]
4744 pub fn add_scale_rms_norm_q8_1_into(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4745 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4746 ncols: usize, nrows: usize, eps: f32,
4747 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4748 -> Result<(), Box<dyn std::error::Error>> {
4749 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4750 let (nc, e2) = (ncols as i32, eps);
4751 if Self::pdl_on() && Self::pdl_wb_on() {
4752 use cudarc::driver::{DevicePtr, DevicePtrMut};
4753 let s = &self.gpu.stream();
4754 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4755 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4756 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4757 let mut ps = [
4758 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4759 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4760 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4761 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4762 &e2 as *const _ as *mut _,
4763 ];
4764 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4765 (rms_block(), 1, 1), &mut ps)?; }
4766 return Ok(());
4767 }
4768 let f = self.func("add_scale_rms_norm_q8_1");
4769 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4770 let __s_b = self.gpu.stream();
4771 let mut b = __s_b.launch_builder(&f);
4772 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut *out_q).arg(&mut *out_d).arg(&nc).arg(&e2);
4773 unsafe { b.launch(cfg)?; }
4774 Ok(())
4775 }
4776
4777 #[allow(clippy::too_many_arguments)]
4780 pub fn rms_pre_add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4781 b_in: &CudaSlice<f32>, c: f32,
4782 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4783 ncols: usize, nrows: usize, eps: f32)
4784 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4785 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4786 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4787 let (nc, e2) = (ncols as i32, eps);
4788 if Self::pdl_on() {
4789 {
4790 use cudarc::driver::{DevicePtr, DevicePtrMut};
4791 let s = &self.gpu.stream();
4792 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
4793 let (pb, _g2) = b_in.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
4794 let (pr, _g4) = res.device_ptr_mut(s);
4795 let (pq, _g5) = out_q.device_ptr_mut(s); let (pd, _g6) = out_d.device_ptr_mut(s);
4796 let mut ps = [
4797 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
4798 &pb as *const _ as *mut _, &c as *const _ as *mut _,
4799 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
4800 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4801 &nc as *const _ as *mut _, &e2 as *const _ as *mut _,
4802 ];
4803 unsafe { self.launch_pdl("rms_pre_add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4804 (rms_block(), 1, 1), &mut ps)?; }
4805 }
4806 return Ok((out_q, out_d));
4807 }
4808 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
4809 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4810 let __s_b = self.gpu.stream();
4811 let mut b = __s_b.launch_builder(&f);
4812 b.arg(a).arg(wa).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e2);
4813 unsafe { b.launch(cfg)?; }
4814 Ok((out_q, out_d))
4815 }
4816
4817 pub fn gelu_tanh_mul_q8_1(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4820 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
4821 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4822 debug_assert!(ncols % 128 == 0);
4823 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4824 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4825 let nc = ncols as i32;
4826 if Self::pdl_on() {
4827 {
4828 use cudarc::driver::{DevicePtr, DevicePtrMut};
4829 let s = &self.gpu.stream();
4830 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4831 let (pact, _g2) = act.device_ptr_mut(s);
4832 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4833 let mut ps = [
4834 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4835 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4836 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4837 ];
4838 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4839 (rms_block(), 1, 1), &mut ps)?; }
4840 }
4841 return Ok((out_q, out_d));
4842 }
4843 let f = self.func("gelu_tanh_mul_q8_1");
4844 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4845 let __s_b = self.gpu.stream();
4846 let mut b = __s_b.launch_builder(&f);
4847 b.arg(gate).arg(up).arg(act).arg(&mut out_q).arg(&mut out_d).arg(&nc);
4848 unsafe { b.launch(cfg)?; }
4849 Ok((out_q, out_d))
4850 }
4851
4852 #[allow(clippy::too_many_arguments)]
4854 pub fn gelu_tanh_mul_q8_1_into(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4855 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
4856 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4857 -> Result<(), Box<dyn std::error::Error>> {
4858 debug_assert!(ncols % 128 == 0);
4859 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4860 let nc = ncols as i32;
4861 if Self::pdl_on() {
4862 use cudarc::driver::{DevicePtr, DevicePtrMut};
4863 let s = &self.gpu.stream();
4864 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4865 let (pact, _g2) = act.device_ptr_mut(s);
4866 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4867 let mut ps = [
4868 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4869 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4870 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4871 ];
4872 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4873 (rms_block(), 1, 1), &mut ps)?; }
4874 return Ok(());
4875 }
4876 let f = self.func("gelu_tanh_mul_q8_1");
4877 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4878 let __s_b = self.gpu.stream();
4879 let mut b = __s_b.launch_builder(&f);
4880 b.arg(gate).arg(up).arg(&mut *act).arg(&mut *out_q).arg(&mut *out_d).arg(&nc);
4881 unsafe { b.launch(cfg)?; }
4882 Ok(())
4883 }
4884
4885 #[allow(clippy::too_many_arguments)]
4887 pub fn add_rms_norm3_q8z(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4888 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4889 res: &mut CudaSlice<f32>, out1: &mut CudaSlice<f32>,
4890 ncols: usize, nrows: usize, eps: f32)
4891 -> Result<((CudaSlice<i8>, CudaSlice<f32>), (CudaSlice<i8>, CudaSlice<f32>)), Box<dyn std::error::Error>> {
4892 let mut q0 = self.alloc_uninit::<i8>(nrows * ncols)?;
4893 let mut d0 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4894 let mut q2 = self.alloc_uninit::<i8>(nrows * ncols)?;
4895 let mut d2 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4896 let f = self.func("add_rms_norm3_q8z_f32");
4897 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4898 let (nc, e2) = (ncols as i32, eps);
4899 let __s_b = self.gpu.stream();
4900 let mut b = __s_b.launch_builder(&f);
4901 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res)
4902 .arg(&mut q0).arg(&mut d0).arg(out1).arg(&mut q2).arg(&mut d2).arg(&nc).arg(&e2);
4903 unsafe { b.launch(cfg)?; }
4904 Ok(((q0, d0), (q2, d2)))
4905 }
4906
4907 #[allow(clippy::too_many_arguments)]
4909 pub fn add_rms_norm3(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4910 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4911 res: &mut CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4912 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4913 -> Result<(), Box<dyn std::error::Error>> {
4914 let f = self.func("add_rms_norm3_f32");
4915 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4916 let (nc, e2) = (ncols as i32, eps);
4917 let __s_b = self.gpu.stream();
4918 let mut b = __s_b.launch_builder(&f);
4919 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e2);
4920 unsafe { b.launch(cfg)?; }
4921 Ok(())
4922 }
4923
4924 pub fn add_scale(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4926 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
4927 let f = self.func("add_scale_f32");
4928 let cfg = LaunchConfig::for_num_elems(n as u32);
4929 let ni = n as i32;
4930 let __s_b = self.gpu.stream();
4931 let mut b = __s_b.launch_builder(&f);
4932 b.arg(a).arg(b_in).arg(&c).arg(dst).arg(&ni);
4933 unsafe { b.launch(cfg)?; }
4934 Ok(())
4935 }
4936
4937 pub fn rms_norm(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4938 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4939 let (nc, e) = (ncols as i32, eps);
4940 if Self::pdl_on() && Self::pdl_wb_on() {
4941 use cudarc::driver::{DevicePtr, DevicePtrMut};
4942 let s = &self.gpu.stream();
4943 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4944 let (pd, _g2) = dst.device_ptr_mut(s);
4945 let mut ps = [
4946 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4947 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4948 &e as *const _ as *mut _,
4949 ];
4950 unsafe { self.launch_pdl("rms_norm_f32", (nrows as u32, 1, 1),
4951 (rms_block(), 1, 1), &mut ps)?; }
4952 return Ok(());
4953 }
4954 let f = self.func("rms_norm_f32");
4955 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4956 let __s_b = self.gpu.stream();
4957 let mut b = __s_b.launch_builder(&f);
4958 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4959 unsafe { b.launch(cfg)?; }
4960 Ok(())
4961 }
4962
4963 pub fn rms_norm_decode(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4971 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4972 let f = self.func("rms_norm_f32");
4973 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4974 let (nc, e) = (ncols as i32, eps);
4975 let __s_b = self.gpu.stream();
4976 let mut b = __s_b.launch_builder(&f);
4977 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4978 unsafe { b.launch(cfg)?; }
4979 Ok(())
4980 }
4981
4982 pub fn rms_norm_q8_1(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize, nrows: usize,
4986 eps: f32) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4987 let nblk = ncols / 32;
4988 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
4989 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
4990 let (nc, e) = (ncols as i32, eps);
4991 if Self::pdl_on() {
4992 {
4993 use cudarc::driver::{DevicePtr, DevicePtrMut};
4994 let s = &self.gpu.stream();
4995 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4996 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
4997 let mut ps = [
4998 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4999 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
5000 &nc as *const _ as *mut _, &e as *const _ as *mut _,
5001 ];
5002 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
5003 &mut ps)?; }
5004 }
5005 return Ok((q, d));
5006 }
5007 let f = self.func("rms_norm_q8_1");
5008 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
5011 let __s_b = self.gpu.stream();
5012 let mut b = __s_b.launch_builder(&f);
5013 b.arg(x).arg(w).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
5014 unsafe { b.launch(cfg)?; }
5015 Ok((q, d))
5016 }
5017
5018 pub fn rms_norm_q8_1_into(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize,
5021 nrows: usize, eps: f32,
5022 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
5023 -> Result<(), Box<dyn std::error::Error>> {
5024 let nblk = ncols / 32;
5025 debug_assert!(q.len() >= nrows * ncols && d.len() >= nrows * nblk);
5026 let (nc, e) = (ncols as i32, eps);
5027 if Self::pdl_on() {
5028 use cudarc::driver::{DevicePtr, DevicePtrMut};
5029 let s = &self.gpu.stream();
5030 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
5031 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
5032 let mut ps = [
5033 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
5034 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
5035 &nc as *const _ as *mut _, &e as *const _ as *mut _,
5036 ];
5037 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
5038 &mut ps)?; }
5039 return Ok(());
5040 }
5041 let f = self.func("rms_norm_q8_1");
5042 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
5043 let __s_b = self.gpu.stream();
5044 let mut b = __s_b.launch_builder(&f);
5045 b.arg(x).arg(w).arg(&mut *q).arg(&mut *d).arg(&nc).arg(&e);
5046 unsafe { b.launch(cfg)?; }
5047 Ok(())
5048 }
5049
5050 pub fn quantize_q8_1_into(&self, x: &CudaSlice<f32>, m: usize, in_f: usize,
5052 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
5053 -> Result<(), Box<dyn std::error::Error>> {
5054 let nblk = in_f / 32;
5055 debug_assert!(q.len() >= m * in_f && d.len() >= m * nblk);
5056 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
5057 let (inf, mi) = (in_f as i32, m as i32);
5058 if Self::pdl_on() && Self::pdl_wb_on() {
5059 use cudarc::driver::{DevicePtr, DevicePtrMut};
5060 let s = &self.gpu.stream();
5061 let (px, _g0) = x.device_ptr(s);
5062 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
5063 let mut ps = [
5064 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
5065 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
5066 &mi as *const _ as *mut _,
5067 ];
5068 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
5069 return Ok(());
5070 }
5071 let f = self.func("quantize_q8_1");
5072 let __s_b = self.gpu.stream();
5073 let mut b = __s_b.launch_builder(&f);
5074 b.arg(x).arg(&mut *q).arg(&mut *d).arg(&inf).arg(&mi);
5075 unsafe { b.launch(cfg)?; }
5076 Ok(())
5077 }
5078
5079 pub fn add_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
5083 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
5084 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5085 let nblk = ncols / 32;
5086 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
5087 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
5088 let f = self.func("add_rms_norm_q8_1");
5089 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
5091 let (nc, e) = (ncols as i32, eps);
5092 let __s_bld = self.gpu.stream();
5093 let mut bld = __s_bld.launch_builder(&f);
5094 bld.arg(a).arg(b_in).arg(w).arg(res).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
5095 unsafe { bld.launch(cfg)?; }
5096 Ok((q, d))
5097 }
5098
5099 pub fn add_rms_norm(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
5103 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
5104 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5105 let (nc, e) = (ncols as i32, eps);
5106 if Self::pdl_on() && Self::pdl_wb_on() {
5107 use cudarc::driver::{DevicePtr, DevicePtrMut};
5108 let s = &self.gpu.stream();
5109 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b.device_ptr(s);
5110 let (pw, _g2) = w.device_ptr(s);
5111 let (pr, _g3) = res.device_ptr_mut(s); let (pd, _g4) = dst.device_ptr_mut(s);
5112 let mut ps = [
5113 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
5114 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
5115 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
5116 &e as *const _ as *mut _,
5117 ];
5118 unsafe { self.launch_pdl("add_rms_norm_f32", (nrows as u32, 1, 1),
5119 (rms_block(), 1, 1), &mut ps)?; }
5120 return Ok(());
5121 }
5122 let f = self.func("add_rms_norm_f32");
5123 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5124 let __s_b2 = self.gpu.stream();
5125 let mut b2 = __s_b2.launch_builder(&f);
5126 b2.arg(a).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
5127 unsafe { b2.launch(cfg)?; }
5128 Ok(())
5129 }
5130
5131 #[allow(clippy::too_many_arguments)]
5134 pub fn rms_pre_add_rms_norm(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
5135 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
5136 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5137 ncols: usize, nrows: usize, eps: f32)
5138 -> Result<(), Box<dyn std::error::Error>> {
5139 let f = self.func("rms_pre_add_rms_norm_f32");
5140 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5141 let (nc, e) = (ncols as i32, eps);
5142 let __s_b2 = self.gpu.stream();
5143 let mut b2 = __s_b2.launch_builder(&f);
5144 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
5145 unsafe { b2.launch(cfg)?; }
5146 Ok(())
5147 }
5148
5149 #[allow(clippy::too_many_arguments)]
5151 pub fn rms_pre_add_rms_norm_q8z(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
5152 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
5153 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5154 ncols: usize, nrows: usize, eps: f32)
5155 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5156 debug_assert!(ncols % 128 == 0);
5157 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5158 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5159 let (nc, e) = (ncols as i32, eps);
5160 if Self::pdl_on() {
5161 {
5162 use cudarc::driver::{DevicePtr, DevicePtrMut};
5163 let s = &self.gpu.stream();
5164 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
5165 let (pb, _g2) = b.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
5166 let (pr, _g4) = res.device_ptr_mut(s); let (pdst, _g5) = dst.device_ptr_mut(s);
5167 let (pq, _g6) = out_q.device_ptr_mut(s); let (pd, _g7) = out_d.device_ptr_mut(s);
5168 let mut ps = [
5169 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
5170 &pb as *const _ as *mut _, &pw as *const _ as *mut _,
5171 &pr as *const _ as *mut _, &pdst as *const _ as *mut _,
5172 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
5173 &nc as *const _ as *mut _, &e as *const _ as *mut _,
5174 ];
5175 unsafe { self.launch_pdl("rms_pre_add_rms_norm_q8z_f32", (nrows as u32, 1, 1),
5176 (rms_block(), 1, 1), &mut ps)?; }
5177 }
5178 return Ok((out_q, out_d));
5179 }
5180 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
5181 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5182 let __s_b2 = self.gpu.stream();
5183 let mut b2 = __s_b2.launch_builder(&f);
5184 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst)
5185 .arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e);
5186 unsafe { b2.launch(cfg)?; }
5187 Ok((out_q, out_d))
5188 }
5189
5190 pub fn build_q4_out_concat3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
5194 w2: &crate::model::GpuTensor)
5195 -> Result<Option<crate::model::GpuTensor>, Box<dyn std::error::Error>> {
5196 use crate::model::GpuTensor;
5197 let part = |w: &GpuTensor| -> Option<(usize, usize)> {
5198 match w {
5199 GpuTensor::Quant { qtype, row_bytes, rp, .. }
5200 if *qtype == QT_Q4_0 && !*rp => Some((*row_bytes, w.out_features())),
5201 _ => None,
5202 }
5203 };
5204 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (part(w0), part(w1), part(w2))
5205 else { return Ok(None) };
5206 if rb0 != rb1 || rb0 != rb2
5207 || w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
5208 return Ok(None);
5209 }
5210 fn bytes_of(w: &crate::model::GpuTensor) -> &CudaSlice<u8> {
5211 match w { crate::model::GpuTensor::Quant { bytes, .. } => bytes, _ => unreachable!() }
5212 }
5213 let (b0, b1, b2) = (bytes_of(w0), bytes_of(w1), bytes_of(w2));
5214 let total = rb0 * (o0 + o1 + o2);
5215 let mut cat = self.alloc_u8(total)?;
5216 self.copy_u8_into(&mut cat, 0, b0, rb0 * o0)?;
5217 self.copy_u8_into(&mut cat, rb0 * o0, b1, rb1 * o1)?;
5218 self.copy_u8_into(&mut cat, rb0 * (o0 + o1), b2, rb2 * o2)?;
5219 Ok(Some(GpuTensor::Quant {
5220 bytes: cat, qtype: QT_Q4_0, row_bytes: rb0,
5221 ne: vec![w0.in_features() as u64, (o0 + o1 + o2) as u64], scale: 1.0, rp: false,
5222 #[cfg(memra_cutlass)]
5223 cutlass: None,
5224 fp8: None, blk: None, rp4: None, f16: None,
5225 }))
5226 }
5227
5228 #[allow(clippy::too_many_arguments)]
5230 pub fn rms_norm_qkv_rope_cat(&self, qkv: &CudaSlice<f32>,
5231 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5232 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5233 head_dim: usize, rq: usize, rk: usize,
5234 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5235 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32)
5236 -> Result<(), Box<dyn std::error::Error>> {
5237 let rows = rq + rk + rk;
5238 let theta_scale = base.powf(-2.0 / head_dim as f32);
5239 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5240 if Self::pdl_on() {
5241 use cudarc::driver::{DevicePtr, DevicePtrMut};
5242 let s = &self.gpu.stream();
5243 let (pqkv, _g0) = qkv.device_ptr(s);
5244 let (pwq, _g1) = wq.device_ptr(s); let (pwk, _g2) = wk.device_ptr(s);
5245 let (pwv, _g3) = wv.device_ptr(s);
5246 let (pq, _g4) = q.device_ptr_mut(s); let (pk, _g5) = k.device_ptr_mut(s);
5247 let (pv, _g6) = v.device_ptr_mut(s);
5248 let (ppos, _g7) = pos.device_ptr(s);
5249 let (pff, _g8) = match ff {
5250 Some(t) => { let (p, g) = t.device_ptr(s); (p, Some(g)) }
5251 None => (0, None),
5252 };
5253 let mut ps = [
5254 &pqkv as *const _ as *mut std::ffi::c_void,
5255 &pwq as *const _ as *mut _, &pwk as *const _ as *mut _,
5256 &pwv as *const _ as *mut _,
5257 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5258 &pv as *const _ as *mut _,
5259 &nc as *const _ as *mut _, &rqi as *const _ as *mut _,
5260 &rki as *const _ as *mut _, &ppos as *const _ as *mut _,
5261 &nhq as *const _ as *mut _, &nhk as *const _ as *mut _,
5262 &theta_scale as *const _ as *mut _, &freq_scale as *const _ as *mut _,
5263 &pff as *const _ as *mut _, &eps as *const _ as *mut _,
5264 ];
5265 unsafe { self.launch_pdl("rms_norm_qkv_rope_cat_f32", (rows as u32, 1, 1),
5266 (rms_block(), 1, 1), &mut ps)?; }
5267 return Ok(());
5268 }
5269 let f = self.func("rms_norm_qkv_rope_cat_f32");
5270 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5271 let __s_b = self.gpu.stream();
5272 let mut b = __s_b.launch_builder(&f);
5273 match ff {
5274 Some(t) => { b.arg(qkv).arg(wq).arg(wk).arg(wv)
5275 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5276 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5277 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5278 unsafe { b.launch(cfg)?; } }
5279 None => { let null: u64 = 0;
5280 b.arg(qkv).arg(wq).arg(wk).arg(wv)
5281 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5282 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5283 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5284 unsafe { b.launch(cfg)?; } }
5285 }
5286 Ok(())
5287 }
5288
5289 #[allow(clippy::too_many_arguments)]
5291 pub fn rms_norm_qkv_rope(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>, v0: &CudaSlice<f32>,
5292 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5293 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5294 head_dim: usize, rq: usize, rk: usize,
5295 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5296 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32)
5297 -> Result<(), Box<dyn std::error::Error>> {
5298 let f = self.func("rms_norm_qkv_rope_f32");
5299 let rows = rq + rk + rk; let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5301 let theta_scale = base.powf(-2.0 / head_dim as f32);
5302 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5303 let __s_b = self.gpu.stream();
5304 let mut b = __s_b.launch_builder(&f);
5305 match ff {
5306 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5307 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5308 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5309 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5310 unsafe { b.launch(cfg)?; } }
5311 None => { let null: u64 = 0;
5312 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5313 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5314 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5315 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5316 unsafe { b.launch(cfg)?; } }
5317 }
5318 Ok(())
5319 }
5320
5321 #[allow(clippy::too_many_arguments)]
5325 pub fn rms_norm_qkv_rope_append_dc(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>,
5326 v0: &CudaSlice<f32>,
5327 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5328 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5329 head_dim: usize, rq: usize, rk: usize,
5330 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5331 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32,
5332 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
5333 t_dev: &CudaSlice<i32>, k_tok_bytes: usize, v_tok_bytes: usize,
5334 g: bool)
5335 -> Result<(), Box<dyn std::error::Error>> {
5336 let rows = rq + rk + rk;
5337 let theta_scale = base.powf(-2.0 / head_dim as f32);
5338 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5339 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
5340 if Self::pdl_on() && Self::pdl_wb_on() {
5341 use cudarc::driver::{DevicePtr, DevicePtrMut};
5342 let s = &self.gpu.stream();
5343 let (p0, _a0) = q0.device_ptr(s); let (p1, _a1) = k0.device_ptr(s);
5344 let (p2, _a2) = v0.device_ptr(s);
5345 let (pwq, _a3) = wq.device_ptr(s); let (pwk, _a4) = wk.device_ptr(s);
5346 let (pwv, _a5) = wv.device_ptr(s);
5347 let (pq, _a6) = q.device_ptr_mut(s); let (pk, _a7) = k.device_ptr_mut(s);
5348 let (pv, _a8) = v.device_ptr_mut(s);
5349 let (pp, _a9) = pos.device_ptr(s);
5350 let pff: u64 = match ff { Some(t) => { let (p, _gg) = t.device_ptr(s); p as u64 }
5351 None => 0 };
5352 let (pkc, _a10) = kc.device_ptr_mut(s); let (pvc, _a11) = vc.device_ptr_mut(s);
5353 let (pt, _a12) = t_dev.device_ptr(s);
5354 let mut ps = [
5355 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
5356 &p2 as *const _ as *mut _, &pwq as *const _ as *mut _,
5357 &pwk as *const _ as *mut _, &pwv as *const _ as *mut _,
5358 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5359 &pv as *const _ as *mut _, &nc as *const _ as *mut _,
5360 &rqi as *const _ as *mut _, &rki as *const _ as *mut _,
5361 &pp as *const _ as *mut _, &nhq as *const _ as *mut _,
5362 &nhk as *const _ as *mut _, &theta_scale as *const _ as *mut _,
5363 &freq_scale as *const _ as *mut _, &pff as *const _ as *mut _,
5364 &eps as *const _ as *mut _, &pkc as *const _ as *mut _,
5365 &pvc as *const _ as *mut _, &pt as *const _ as *mut _,
5366 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
5367 ];
5368 unsafe { self.launch_pdl_flash(g, "rms_norm_qkv_rope_append_dc_f32",
5369 (rows as u32, 1, 1), (rms_block(), 1, 1), 0, &mut ps)?; }
5370 return Ok(());
5371 }
5372 let f = if g { self.func_g("rms_norm_qkv_rope_append_dc_f32") }
5373 else { self.func("rms_norm_qkv_rope_append_dc_f32") };
5374 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5375 let __s_b = self.gpu.stream();
5376 let mut b = __s_b.launch_builder(&f);
5377 match ff {
5378 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5379 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5380 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5381 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps)
5382 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5383 unsafe { b.launch(cfg)?; } }
5384 None => { let null: u64 = 0;
5385 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5386 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5387 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5388 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps)
5389 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5390 unsafe { b.launch(cfg)?; } }
5391 }
5392 Ok(())
5393 }
5394
5395 pub fn add_q8_1(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
5397 ncols: usize, nrows: usize)
5398 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5399 debug_assert!(ncols % 128 == 0);
5400 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5401 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5402 let f = self.func("add_q8_1_f32");
5403 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5404 let nc = ncols as i32;
5405 let __s_b2 = self.gpu.stream();
5406 let mut b2 = __s_b2.launch_builder(&f);
5407 b2.arg(a).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc);
5408 unsafe { b2.launch(cfg)?; }
5409 Ok((out_q, out_d))
5410 }
5411
5412 pub fn rms_pre_add_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>, b: &CudaSlice<f32>,
5416 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
5417 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5418 debug_assert!(ncols % 128 == 0);
5419 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5420 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5421 let f = self.func("rms_pre_add_q8_1_f32");
5422 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1),
5423 shared_mem_bytes: 0 };
5424 let (nc, ep) = (ncols as i32, eps);
5425 let __s_b2 = self.gpu.stream();
5426 let mut b2 = __s_b2.launch_builder(&f);
5427 b2.arg(a).arg(wa).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
5428 unsafe { b2.launch(cfg)?; }
5429 Ok((out_q, out_d))
5430 }
5431
5432 pub fn l2_v2_on(ncols: usize) -> bool {
5436 ncols == 128 && std::env::var("MEMRA_L2_V2").as_deref() != Ok("0")
5437 }
5438
5439 pub fn l2_norm_pp(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5440 dst16: Option<&mut CudaSlice<u8>>, ncols: usize, nrows: usize,
5441 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5442 if Self::l2_v2_on(ncols) {
5443 let f = self.func("l2_norm_pp_v2_f32");
5444 let rows_per_block = 8u32; let cfg = LaunchConfig { grid_dim: ((nrows as u32).div_ceil(rows_per_block), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
5446 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
5447 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
5449 let __s_b = self.gpu.stream();
5450 let mut b = __s_b.launch_builder(&f);
5451 b.arg(x).arg(dst).arg(&d16).arg(&nc).arg(&nr).arg(&e);
5452 unsafe { b.launch(cfg)?; }
5453 return Ok(());
5454 }
5455 self.l2_norm(x, dst, ncols, nrows, eps)
5456 }
5457
5458 pub fn l2_norm(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
5459 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5460 let f = self.func("l2_norm_f32");
5461 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
5462 let (nc, e) = (ncols as i32, eps);
5463 let __s_b = self.gpu.stream();
5464 let mut b = __s_b.launch_builder(&f);
5465 b.arg(x).arg(dst).arg(&nc).arg(&e);
5466 unsafe { b.launch(cfg)?; }
5467 Ok(())
5468 }
5469
5470 pub fn l2_norm_decode(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize,
5476 nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5477 let f = self.func("l2_norm_f32");
5478 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
5479 let (nc, e) = (ncols as i32, eps);
5480 let __s_b = self.gpu.stream();
5481 let mut b = __s_b.launch_builder(&f);
5482 b.arg(x).arg(dst).arg(&nc).arg(&e);
5483 unsafe { b.launch(cfg)?; }
5484 Ok(())
5485 }
5486
5487 pub fn rope_neox(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5489 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32, freq_scale: f32)
5490 -> Result<(), Box<dyn std::error::Error>> {
5491 let f = self.func("rope_neox_f32");
5492 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5493 let grid = (n_heads * n_tokens) as u32;
5494 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5495 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5496 let __s_b = self.gpu.stream();
5497 let mut b = __s_b.launch_builder(&f);
5498 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale);
5499 unsafe { b.launch(cfg)?; }
5500 Ok(())
5501 }
5502
5503 pub fn rope_neox_ff(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5505 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32,
5506 freq_scale: f32, ff: &CudaSlice<f32>)
5507 -> Result<(), Box<dyn std::error::Error>> {
5508 let f = self.func("rope_neox_ff_f32");
5509 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5510 let grid = (n_heads * n_tokens) as u32;
5511 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5512 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5513 let __s_b = self.gpu.stream();
5514 let mut b = __s_b.launch_builder(&f);
5515 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale).arg(ff);
5516 unsafe { b.launch(cfg)?; }
5517 Ok(())
5518 }
5519
5520 #[allow(clippy::too_many_arguments)]
5522 pub fn rope_neox2(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
5523 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
5524 nh_q: usize, nh_k: usize, n_tokens: usize, freq_base: f32,
5525 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
5526 -> Result<(), Box<dyn std::error::Error>> {
5527 let f = self.func("rope_neox2_f32");
5528 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5529 let grid = ((nh_q + nh_k) * n_tokens) as u32;
5530 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5531 let (hd, nd, nq, nk, nt) = (head_dim as i32, n_dims as i32, nh_q as i32, nh_k as i32, n_tokens as i32);
5532 let __s_b = self.gpu.stream();
5533 let mut b = __s_b.launch_builder(&f);
5534 b.arg(q).arg(k).arg(pos).arg(&hd).arg(&nd).arg(&nq).arg(&nk).arg(&nt)
5535 .arg(&theta_scale).arg(&freq_scale);
5536 match ff {
5537 Some(ffv) => { b.arg(ffv); unsafe { b.launch(cfg)?; } }
5538 None => {
5539 let null: u64 = 0;
5540 b.arg(&null);
5541 unsafe { b.launch(cfg)?; }
5542 }
5543 }
5544 Ok(())
5545 }
5546
5547 pub fn gelu_tanh_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5549 -> Result<(), Box<dyn std::error::Error>> {
5550 let f = self.func("gelu_tanh_mul_f32");
5551 let cfg = LaunchConfig::for_num_elems(n as u32);
5552 let ni = n as i32;
5553 let __s_b = self.gpu.stream();
5554 let mut b = __s_b.launch_builder(&f);
5555 b.arg(gate).arg(up).arg(dst).arg(&ni);
5556 unsafe { b.launch(cfg)?; }
5557 Ok(())
5558 }
5559
5560 pub fn silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5561 -> Result<(), Box<dyn std::error::Error>> {
5562 let f = self.func("silu_mul_f32");
5563 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5565 let ni = n as i32;
5566 let __s_b = self.gpu.stream();
5567 let mut b = __s_b.launch_builder(&f);
5568 b.arg(gate).arg(up).arg(dst).arg(&ni);
5569 unsafe { b.launch(cfg)?; }
5570 Ok(())
5571 }
5572
5573 pub fn silu_mul_f16out(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
5576 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
5577 -> Result<(), Box<dyn std::error::Error>> {
5578 let f = self.func("silu_mul_f16out_f32");
5579 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5580 let ni = n as i32;
5581 let __s_b = self.gpu.stream();
5582 let mut b = __s_b.launch_builder(&f);
5583 b.arg(gate).arg(up).arg(dst).arg(dst16).arg(&ni);
5584 unsafe { b.launch(cfg)?; }
5585 Ok(())
5586 }
5587
5588 pub fn silu_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5595 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
5596 let f = self.func("silu_mul_scaled_f32");
5597 let cfg = LaunchConfig::for_num_elems(n as u32);
5598 let ni = n as i32;
5599 let (gsf, usf) = (gs, us);
5600 let __s_b = self.gpu.stream();
5601 let mut b = __s_b.launch_builder(&f);
5602 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(dst).arg(&ni);
5603 unsafe { b.launch(cfg)?; }
5604 Ok(())
5605 }
5606
5607 #[allow(clippy::too_many_arguments)]
5611 pub fn swigluoai_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5612 alpha: f32, limit: f32, dst: &mut CudaSlice<f32>, n: usize)
5613 -> Result<(), Box<dyn std::error::Error>> {
5614 let f = self.func("swigluoai_mul_scaled_f32");
5615 let cfg = LaunchConfig::for_num_elems(n as u32);
5616 let ni = n as i32;
5617 let __s_b = self.gpu.stream();
5618 let mut b = __s_b.launch_builder(&f);
5619 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&alpha).arg(&limit).arg(dst).arg(&ni);
5620 unsafe { b.launch(cfg)?; }
5621 Ok(())
5622 }
5623
5624 pub fn silu_mul_scaled_q8_1(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5632 n: usize)
5633 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5634 let f = self.func("silu_mul_scaled_q8_1");
5635 let nblk = n / 32;
5636 let mut aq = self.alloc_uninit::<i8>(n)?; let mut ad = self.alloc_uninit::<f32>(nblk)?; let cfg = LaunchConfig::for_num_elems(n as u32);
5640 let (gsf, usf, ni) = (gs, us, n as i32);
5641 let __s_b = self.gpu.stream();
5642 let mut b = __s_b.launch_builder(&f);
5643 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(&mut aq).arg(&mut ad).arg(&ni);
5644 unsafe { b.launch(cfg)?; }
5645 Ok((aq, ad))
5646 }
5647
5648 pub fn add(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5649 -> Result<(), Box<dyn std::error::Error>> {
5650 let f = self.func("add_f32");
5651 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5653 let ni = n as i32;
5654 let __s_bld = self.gpu.stream();
5655 let mut bld = __s_bld.launch_builder(&f);
5656 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5657 unsafe { bld.launch(cfg)?; }
5658 Ok(())
5659 }
5660
5661 pub fn mul(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5662 -> Result<(), Box<dyn std::error::Error>> {
5663 let f = self.func("mul_f32");
5664 let cfg = LaunchConfig::for_num_elems(n as u32);
5665 let ni = n as i32;
5666 let __s_bld = self.gpu.stream();
5667 let mut bld = __s_bld.launch_builder(&f);
5668 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5669 unsafe { bld.launch(cfg)?; }
5670 Ok(())
5671 }
5672
5673 pub fn matmul(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
5676 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5677 use crate::model::GpuTensor;
5678 let in_f = w.in_features();
5679 let out_f = w.out_features();
5680 #[allow(non_snake_case)]
5688 let GEMM_M_THRESHOLD = if self.verify_exact_on() { usize::MAX } else { 16usize };
5691
5692 const GEMM_MIN_OUT_F: usize = 128; if m >= GEMM_M_THRESHOLD {
5717 if let Some(y) = self.try_fp8_gemm(w, x, m)? { return Ok(y); }
5718 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(y); }
5725 if let Some(y) = self.try_f16_gemm(w, x, m)? { return Ok(y); }
5728 }
5729 if let GpuTensor::Quant { qtype, .. } = w {
5744 if *qtype == QT_F8_E4M3_BLK {
5745 if m >= GEMM_M_THRESHOLD {
5746 if let Some(y) = self.try_e4m3_blk_prefill(w, x, m)? { return Ok(y); }
5747 }
5748 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5749 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
5750 }
5751 }
5752 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.mmq_supports(w) {
5753 return self.qmatvec_mmq(w, x, m);
5754 }
5755 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.gemm_supports(w) {
5756 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5757 return self.qmatvec_gemm(w, &aq, &ad, m);
5758 }
5759 if m >= GEMM_M_THRESHOLD {
5762 if let Some(y) = self.try_fp4_gemm(w, x, m, in_f, out_f)? { return Ok(y); }
5763 }
5764 let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
5768 if m == 1 && fast {
5773 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, scale, .. } = w {
5774 if self.mmvq_supports(*qtype) {
5775 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5779 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5780 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, rp);
5781 }
5782 }
5783 }
5784 if (2..=16).contains(&m) && fast && std::env::var("MEMRA_NO_BATCHED").is_err()
5800 && (m <= 4 || Self::b8_enabled()) {
5801 let m_ok = m <= 8 || matches!(w, GpuTensor::Quant { qtype, .. }
5811 if *qtype == QT_Q4_0 || *qtype == QT_Q6_K || *qtype == QT_F8_E4M3
5812 || *qtype == QT_NVFP4 || *qtype == QT_Q4_K || *qtype == QT_Q5_K || *qtype == QT_Q8_0);
5813 if m_ok {
5814 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, .. } = w {
5815 if self.batched_supports(*qtype) && self.mmvq_supports(*qtype) {
5816 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5817 let mcols = Self::batched_mcols(m);
5818 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5819 let mut y = self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, mcols, 1.0, rp)?;
5820 if let GpuTensor::Quant { scale, .. } = w {
5821 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5822 }
5823 return Ok(y);
5824 }
5825 }
5826 }
5827 }
5828 if fast {
5834 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. } = w {
5835 if *qtype == QT_F8_E4M3 {
5836 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5837 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes,
5838 *scale, false);
5839 }
5840 }
5841 }
5842 let mut y = match w {
5843 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q8_0 =>
5844 self.qmatvec_q8_0_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5845 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q4_K =>
5846 self.qmatvec_q4_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5847 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q6_K =>
5848 self.qmatvec_q6_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5849 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q5_K =>
5850 self.qmatvec_q5_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5851 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q3_K =>
5852 self.qmatvec_q3_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5853 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } if fast && *qtype == QT_NVFP4 =>
5854 self.qmatvec_dp4a_named(
5855 if *rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
5856 bytes, x, m, in_f, out_f, *row_bytes)?,
5857 GpuTensor::Quant { bytes, qtype, row_bytes, .. }
5861 if fast && *qtype == QT_IQ4_XS && Self::iq_fast_enabled() =>
5862 self.qmatvec_iq4_XS_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5863 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } =>
5868 self.qmatvec(bytes, x, m, in_f, out_f,
5871 if *rp && *qtype == QT_NVFP4 { QT_NVFP4_RP } else { *qtype },
5872 *row_bytes)?,
5873 GpuTensor::Float { data, .. } => self.linear(x, data, m, in_f, out_f)?,
5874 GpuTensor::FloatBf16 { data, .. } =>
5877 self.linear_bf16_chunked(x, data, m, in_f, out_f, false)?,
5878 };
5879 if let GpuTensor::Quant { scale, .. } = w {
5881 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5882 }
5883 Ok(y)
5884 }
5885
5886 pub fn uses_q8_1_fast(&self, w: &crate::model::GpuTensor) -> bool {
5889 use crate::model::GpuTensor;
5890 if std::env::var("MEMRA_FAST").as_deref() == Ok("0") { return false; }
5891 match w {
5892 GpuTensor::Quant { qtype, .. } => matches!(*qtype,
5899 QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q3_K | QT_NVFP4 | QT_F8_E4M3
5900 | QT_F8_E4M3_BLK | QT_Q4_0)
5901 || (*qtype == QT_IQ4_XS && Self::iq_fast_enabled()),
5902 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
5903 }
5904 }
5905
5906 pub fn matmul_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
5911 x_fallback: &CudaSlice<f32>, m: usize)
5912 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5913 use crate::model::GpuTensor;
5914 let x_raw_ok = x_fallback.len() >= m * w.in_features();
5920 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5923 if let Some(y) = self.try_fp8_gemm(w, x_fallback, m)? { return Ok(y); }
5924 if let Some(y) = self.try_fp8_blk_mmq(w, x_fallback, m)? { return Ok(y); }
5927 if let Some(y) = self.try_f16_gemm(w, x_fallback, m)? { return Ok(y); }
5929 }
5930 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5936 if let Some(y) = self.try_e4m3_blk_prefill(w, x_fallback, m)? { return Ok(y); }
5937 }
5938 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
5939 if m >= 16 && w.out_features() >= 128 && self.mmq_supports(w) && !self.verify_exact_on()
5944 && x_raw_ok {
5945 return self.qmatvec_mmq(w, x_fallback, m);
5946 }
5947 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5950 if let Some(y) = self.try_fp4_gemm(w, x_fallback, m, w.in_features(), w.out_features())? {
5951 return Ok(y);
5952 }
5953 }
5954 if m >= 16 && self.gemm_supports(w) && !self.verify_exact_on() {
5957 return self.qmatvec_gemm(w, aq, ad, m);
5958 }
5959 if !self.uses_q8_1_fast(w) { return self.matmul(w, x_fallback, m); }
5960 let in_f = w.in_features();
5961 let out_f = w.out_features();
5962 let (bytes, qtype, row_bytes, scale, rp) = match w {
5963 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
5964 _ => unreachable!("uses_q8_1_fast guaranteed Quant"),
5965 };
5966 let (mbytes, mrp) = match w {
5969 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5970 _ => (bytes, rp),
5971 };
5972 if m == 1 && self.mmvq_supports(qtype) {
5976 return self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, mrp);
5977 }
5978 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5991 && std::env::var("MEMRA_NO_BATCHED").is_err()
5992 && (m <= 4 || Self::b8_enabled())
5993 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_NVFP4
5997 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_F8_E4M3 || qtype == QT_Q8_0) {
5998 let mcols = Self::batched_mcols(m);
5999 return self.qmatvec_mmvq_batched(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, mrp);
6000 }
6001 if qtype == QT_F8_E4M3 || qtype == QT_Q4_0 {
6007 let (b2, r2) = if qtype == QT_Q4_0 { (mbytes, mrp) } else { (bytes, rp) };
6008 return self.qmatvec_mmvq(b2, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, r2);
6009 }
6010 let name = match qtype {
6011 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
6012 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
6013 QT_Q3_K => "qmatvec_q3_K_dp4a",
6014 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
6015 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
6016 _ => unreachable!(),
6017 };
6018 let f = self.func(name);
6019 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
6021 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
6022 let __s_b = self.gpu.stream();
6023 let mut b = __s_b.launch_builder(&f);
6024 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
6025 unsafe { b.launch(cfg)?; }
6026 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
6027 Ok(y)
6028 }
6029
6030 pub fn matmul_decode_exact(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
6038 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6039 use crate::model::GpuTensor;
6040 if let GpuTensor::Float { data, .. } = w {
6048 return self.linear_decode_exact(x, data, m, w.in_features(), w.out_features());
6049 }
6050 if let GpuTensor::FloatBf16 { data, .. } = w {
6053 let (in_f, out_f) = (w.in_features(), w.out_features());
6054 return self.linear_bf16_chunked(x, data, m, in_f, out_f, true);
6055 }
6056 if !self.uses_q8_1_fast(w) { return self.matmul(w, x, m); }
6057 let in_f = w.in_features();
6058 let out_f = w.out_features();
6059 let (bytes, qtype, row_bytes, scale, rp) = match w {
6060 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6061 _ => return self.matmul(w, x, m),
6062 };
6063 let (bytes, rp) = match w {
6066 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
6067 _ => (bytes, rp),
6068 };
6069 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6070 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
6074 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
6083 && std::env::var("MEMRA_NO_BATCHED").is_err()
6084 && (m <= 4 || Self::b8_enabled())
6085 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
6088 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
6089 let mcols = Self::batched_mcols(m);
6090 return self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
6091 }
6092 if self.mmvq_supports(qtype) {
6093 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
6096 }
6097 self.matmul_pre(w, &aq, &ad, x, m)
6100 }
6101
6102 pub fn matmul_decode_exact_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
6112 ad: &CudaSlice<f32>, m: usize)
6113 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6114 use crate::model::GpuTensor;
6115 debug_assert!(self.uses_q8_1_fast(w),
6116 "matmul_decode_exact_pre: caller must guarantee q8_1-fast");
6117 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
6119 let in_f = w.in_features();
6120 let out_f = w.out_features();
6121 let (bytes, qtype, row_bytes, scale, rp) = match w {
6122 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
6123 (bytes, *qtype, *row_bytes, *scale, *rp),
6124 _ => return Err("matmul_decode_exact_pre: Quant tensor required (q8_1-fast contract)".into()),
6125 };
6126 let (bytes, rp) = match w {
6128 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
6129 _ => (bytes, rp),
6130 };
6131 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
6133 && std::env::var("MEMRA_NO_BATCHED").is_err()
6134 && (m <= 4 || Self::b8_enabled())
6135 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
6136 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
6137 let mcols = Self::batched_mcols(m);
6138 return self.qmatvec_mmvq_batched(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
6139 }
6140 if self.mmvq_supports(qtype) {
6141 return self.qmatvec_mmvq(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
6142 }
6143 let x0 = self.zeros(0)?;
6146 self.matmul_pre(w, aq, ad, &x0, m)
6147 }
6148
6149 pub fn matmul_decode_exact_dual_pre(&self, w0: &crate::model::GpuTensor,
6158 w1: &crate::model::GpuTensor,
6159 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6160 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
6161 use crate::model::GpuTensor;
6162 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6163 let on = *ON.get_or_init(|| {
6164 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
6165 });
6166 if !on || !(2..=7).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
6167 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
6168 return Ok(None);
6169 }
6170 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6175 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6176 if w1.in_features() != in_f || w1.out_features() != out_f {
6177 return Ok(None);
6178 }
6179 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
6180 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
6181 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
6182 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
6183 (b0, b1, *rb0, *s0, *s1, *rp0),
6184 _ => return Ok(None),
6185 };
6186 if m > 4 && !(rp && Self::b8_enabled()
6189 && std::env::var("MEMRA_B567").as_deref() != Ok("0")) {
6190 return Ok(None);
6191 }
6192 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, aq, ad, m, in_f, out_f, row_bytes, rp)?;
6193 Ok(Some(((y0, s0), (y1, s1))))
6194 }
6195
6196 pub fn matmul_decode_exact_dual(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6212 x: &CudaSlice<f32>, m: usize)
6213 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6214 use crate::model::GpuTensor;
6215 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6216 let on = *ON.get_or_init(|| {
6217 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
6218 });
6219 if !on || !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
6220 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
6221 return Ok(None);
6222 }
6223 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6228 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6229 if w1.in_features() != in_f || w1.out_features() != out_f {
6230 return Ok(None);
6231 }
6232 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
6233 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
6234 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
6235 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
6236 (b0, b1, *rb0, *s0, *s1, *rp0),
6237 _ => return Ok(None),
6238 };
6239 if std::env::var("MEMRA_DEBUG").is_ok() {
6242 static ONCE: std::sync::Once = std::sync::Once::new();
6243 ONCE.call_once(|| eprintln!("[memra] dual gate+up batched ENGAGED (m={m} rp={rp})"));
6244 }
6245 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6246 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, &aq, &ad, m, in_f, out_f, row_bytes, rp)?;
6247 let mut y0 = y0;
6248 let mut y1 = y1;
6249 if s0 != 1.0 { self.scale_inplace(&mut y0, s0, m * out_f)?; }
6250 if s1 != 1.0 { self.scale_inplace(&mut y1, s1, m * out_f)?; }
6251 Ok(Some((y0, y1)))
6252 }
6253
6254 #[allow(clippy::too_many_arguments)]
6259 pub fn qmatvec_batched_dual_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6260 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6261 m: usize, in_f: usize, out_f: usize, row_bytes: usize, rp: bool)
6262 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6263 const ROWS_PER_BLOCK: u32 = 4;
6264 let mcols = Self::batched_mcols(m);
6265 let tiny_rp1 = rp && mcols == 4 && out_f <= 128
6268 && std::env::var("MEMRA_NVFP4_AUX_DUAL").as_deref() != Ok("0");
6269 let (name, rows_per_block) = if tiny_rp1 {
6270 ("qmatvec_nvfp4_mmvq_dual_b4_rp", ROWS_PER_BLOCK)
6271 } else { match (mcols, rp, m) {
6272 (2, false, _) => ("qmatvec_nvfp4_mmvq_dual_b2", ROWS_PER_BLOCK),
6273 (4, false, _) => ("qmatvec_nvfp4_mmvq_dual_b4_r2", ROWS_PER_BLOCK * 2),
6274 (2, true, _) => ("qmatvec_nvfp4_mmvq_dual_b2_rp", ROWS_PER_BLOCK),
6275 (4, true, _) => ("qmatvec_nvfp4_mmvq_dual_b4_rpr2", ROWS_PER_BLOCK * 2),
6276 (8, true, 5) => ("qmatvec_nvfp4_mmvq_dual_b5_rpr2", ROWS_PER_BLOCK * 2),
6277 (8, true, 6) => ("qmatvec_nvfp4_mmvq_dual_b6_rpr2", ROWS_PER_BLOCK * 2),
6278 (8, true, 7) => ("qmatvec_nvfp4_mmvq_dual_b7_rpr2", ROWS_PER_BLOCK * 2),
6279 _ => return Err(format!("qmatvec_batched_dual_raw: no dual kernel for m {m}").into()),
6280 }};
6281 let f = self.func(name);
6282 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
6283 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
6284 let cfg = LaunchConfig {
6285 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6286 block_dim: (32, ROWS_PER_BLOCK, 1),
6287 shared_mem_bytes: 0,
6288 };
6289 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
6290 let __s_b = self.gpu.stream();
6291 let mut b = __s_b.launch_builder(&f);
6292 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6293 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
6294 unsafe { b.launch(cfg)?; }
6295 Ok((y0, y1))
6296 }
6297
6298 pub fn matmul_pre_dual_noscale(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6310 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6311 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
6312 use crate::model::GpuTensor;
6313 if m != 1 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6314 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6324 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6325 if w1.in_features() != in_f || w1.out_features() != out_f { return Ok(None); }
6326 let no_mirror = |w: &crate::model::GpuTensor| {
6339 !matches!(w, GpuTensor::Quant { rp4: Some(_), .. })
6340 };
6341 if self.q8_ffn_fuse2_on()
6342 && no_mirror(w0) && no_mirror(w1)
6343 && let Some([p0, p1]) = self.q8_fused_params(&[w0, w1])
6344 {
6345 let (y0, y1) = self.q8_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2)?;
6346 return Ok(Some(((y0, 1.0), (y1, 1.0))));
6347 }
6348 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6358 let (y0, y1) = self.e4m3_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2,
6359 1.0, 1.0)?;
6360 return Ok(Some(((y0, p0.3), (y1, p1.3))));
6361 }
6362 let (b0, q0, rb0, s0, rp0) = match w0 {
6363 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6364 _ => return Ok(None),
6365 };
6366 let (b1, q1, rb1, s1, rp1) = match w1 {
6367 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6368 _ => return Ok(None),
6369 };
6370 if q0 != QT_NVFP4 || q1 != QT_NVFP4 || rb0 != rb1 || rp0 != rp1 { return Ok(None); }
6371 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
6373 let rows_per_block = ROWS_PER_BLOCK * RPW;
6374 let f = self.func(if rp0 { "qmatvec_nvfp4_mmvq_dual_mr2_rp" } else { "qmatvec_nvfp4_mmvq_dual_mr2" });
6375 let mut y0 = self.alloc_uninit::<f32>(out_f)?;
6376 let mut y1 = self.alloc_uninit::<f32>(out_f)?;
6377 let cfg = LaunchConfig {
6378 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6379 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0,
6380 };
6381 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, rb0 as i64);
6382 let one = 1.0f32;
6385 let __s_b = self.gpu.stream();
6386 let mut b = __s_b.launch_builder(&f);
6387 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6388 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&one).arg(&one);
6389 unsafe { b.launch(cfg)?; }
6390 Ok(Some(((y0, s0), (y1, s1))))
6391 }
6392
6393 pub fn matmul_q8_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6401 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6402 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6403 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6409 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(),
6410 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6411 }
6412 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6413 Ok(Some(self.q8_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6414 }
6415
6416 #[allow(clippy::too_many_arguments)]
6417 fn q8_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6418 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6419 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6420 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6421 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6423 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6424 let f = self.func("qmatvec_q8_0_mmvq_fused2");
6425 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6426 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6427 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6428 shared_mem_bytes: 0 };
6429 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
6430 let __s_b = self.gpu.stream();
6431 let mut b = __s_b.launch_builder(&f);
6432 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6433 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl);
6434 unsafe { b.launch(cfg)?; }
6435 Ok((y0, y1))
6436 }
6437
6438 pub fn matmul_q8_fused2_x(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6444 x: &CudaSlice<f32>)
6445 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6446 if !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6447 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6448 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6449 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(),
6450 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6451 }
6452 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6453 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6454 Ok(Some(self.q8_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6455 }
6456
6457 #[allow(clippy::too_many_arguments)]
6460 pub fn qmatvec_q8_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
6461 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6462 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6463 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6464 self.q8_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes)
6465 }
6466
6467 pub fn matmul_q4_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6473 w2: &crate::model::GpuTensor,
6474 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6475 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6476 use crate::model::GpuTensor;
6477 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6478 match w {
6479 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6480 Some((*row_bytes, w.out_features())),
6481 _ => None,
6482 }
6483 };
6484 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6485 else { return Ok(None) };
6486 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6487 return Ok(None);
6488 }
6489 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6493 match w {
6494 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6495 Some(m) => (m, true),
6496 None => (bytes, *rp),
6497 },
6498 _ => unreachable!(),
6499 }
6500 }
6501 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6502 if rp0 != rp1 || rp1 != rp2 { return Ok(None); }
6503 let rp = rp0;
6504 let rpb: u32 = 4;
6505 let mr1 = rp && Self::q40_mr1_on();
6509 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6510 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6511 let grid = nb(o0) + nb(o1) + nb(o2);
6512 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6513 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6514 let mut y2 = self.alloc_uninit::<f32>(o2)?;
6515 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6516 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6517 else { "qmatvec_q4_0_mmvq_fused3" });
6518 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6519 let inf = w0.in_features() as i32;
6520 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6521 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6522 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6525 {
6526 use cudarc::driver::{DevicePtr, DevicePtrMut};
6527 let s = &self.gpu.stream();
6528 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6529 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6530 let (pad, _g4) = ad.device_ptr(s);
6531 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6532 let (py2, _g7) = y2.device_ptr_mut(s);
6533 let mut ps = [
6534 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6535 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6536 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6537 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6538 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6539 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6540 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6541 &r2 as *const _ as *mut _,
6542 ];
6543 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6544 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6545 }
6546 return Ok(Some((y0, y1, y2)));
6547 }
6548 let __s_b = self.gpu.stream();
6549 let mut b = __s_b.launch_builder(&f);
6550 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6551 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6552 unsafe { b.launch(cfg)?; }
6553 Ok(Some((y0, y1, y2)))
6554 }
6555
6556 #[allow(clippy::too_many_arguments)]
6559 pub fn matmul_q4_fused3_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6560 w2: &crate::model::GpuTensor,
6561 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6562 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>,
6563 y2: &mut CudaSlice<f32>)
6564 -> Result<bool, Box<dyn std::error::Error>> {
6565 use crate::model::GpuTensor;
6566 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6567 match w {
6568 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6569 Some((*row_bytes, w.out_features())),
6570 _ => None,
6571 }
6572 };
6573 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6574 else { return Ok(false) };
6575 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6576 return Ok(false);
6577 }
6578 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6579 match w {
6580 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6581 Some(m) => (m, true),
6582 None => (bytes, *rp),
6583 },
6584 _ => unreachable!(),
6585 }
6586 }
6587 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6588 if rp0 != rp1 || rp1 != rp2 { return Ok(false); }
6589 let rp = rp0;
6590 let rpb: u32 = 4;
6591 let mr1 = rp && Self::q40_mr1_on();
6592 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6593 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6594 let grid = nb(o0) + nb(o1) + nb(o2);
6595 debug_assert!(y0.len() >= o0 && y1.len() >= o1 && y2.len() >= o2);
6596 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6597 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6598 else { "qmatvec_q4_0_mmvq_fused3" });
6599 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6600 let inf = w0.in_features() as i32;
6601 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6602 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6603 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6605 use cudarc::driver::{DevicePtr, DevicePtrMut};
6606 let s = &self.gpu.stream();
6607 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6608 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6609 let (pad, _g4) = ad.device_ptr(s);
6610 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6611 let (py2, _g7) = y2.device_ptr_mut(s);
6612 let mut ps = [
6613 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6614 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6615 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6616 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6617 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6618 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6619 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6620 &r2 as *const _ as *mut _,
6621 ];
6622 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6623 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6624 return Ok(true);
6625 }
6626 let __s_b = self.gpu.stream();
6627 let mut b = __s_b.launch_builder(&f);
6628 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1).arg(&mut *y2)
6629 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6630 unsafe { b.launch(cfg)?; }
6631 Ok(true)
6632 }
6633
6634 pub fn matmul_q4_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6636 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6637 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6638 use crate::model::GpuTensor;
6639 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6640 match w {
6641 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6642 Some((*row_bytes, w.out_features())),
6643 _ => None,
6644 }
6645 };
6646 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6647 if w0.in_features() != w1.in_features() { return Ok(None); }
6648 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6650 match w {
6651 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6652 Some(m) => (m, true),
6653 None => (bytes, *rp),
6654 },
6655 _ => unreachable!(),
6656 }
6657 }
6658 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6659 if rp0 != rp1 { return Ok(None); }
6660 let rp = rp0;
6661 let rpb: u32 = 4;
6662 let mr1 = rp && Self::q40_mr1_on();
6664 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6665 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6666 let grid = nb(o0) + nb(o1);
6667 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6668 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6669 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6670 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6671 else { "qmatvec_q4_0_mmvq_fused2" });
6672 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6673 let inf = w0.in_features() as i32;
6674 let (oo0, oo1) = (o0 as i32, o1 as i32);
6675 let (r0, r1) = (rb0 as i64, rb1 as i64);
6676 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6678 {
6679 use cudarc::driver::{DevicePtr, DevicePtrMut};
6680 let s = &self.gpu.stream();
6681 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6682 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6683 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6684 let mut ps = [
6685 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6686 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6687 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6688 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6689 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6690 &r1 as *const _ as *mut _,
6691 ];
6692 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6693 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6694 }
6695 return Ok(Some((y0, y1)));
6696 }
6697 let __s_b = self.gpu.stream();
6698 let mut b = __s_b.launch_builder(&f);
6699 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6700 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6701 unsafe { b.launch(cfg)?; }
6702 Ok(Some((y0, y1)))
6703 }
6704
6705 pub fn matmul_q4_fused2_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6707 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6708 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>)
6709 -> Result<bool, Box<dyn std::error::Error>> {
6710 use crate::model::GpuTensor;
6711 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6712 match w {
6713 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6714 Some((*row_bytes, w.out_features())),
6715 _ => None,
6716 }
6717 };
6718 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(false) };
6719 if w0.in_features() != w1.in_features() { return Ok(false); }
6720 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6721 match w {
6722 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6723 Some(m) => (m, true),
6724 None => (bytes, *rp),
6725 },
6726 _ => unreachable!(),
6727 }
6728 }
6729 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6730 if rp0 != rp1 { return Ok(false); }
6731 let rp = rp0;
6732 let rpb: u32 = 4;
6733 let mr1 = rp && Self::q40_mr1_on();
6734 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6735 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6736 let grid = nb(o0) + nb(o1);
6737 debug_assert!(y0.len() >= o0 && y1.len() >= o1);
6738 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6739 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6740 else { "qmatvec_q4_0_mmvq_fused2" });
6741 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6742 let inf = w0.in_features() as i32;
6743 let (oo0, oo1) = (o0 as i32, o1 as i32);
6744 let (r0, r1) = (rb0 as i64, rb1 as i64);
6745 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6747 use cudarc::driver::{DevicePtr, DevicePtrMut};
6748 let s = &self.gpu.stream();
6749 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6750 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6751 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6752 let mut ps = [
6753 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6754 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6755 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6756 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6757 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6758 &r1 as *const _ as *mut _,
6759 ];
6760 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6761 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6762 return Ok(true);
6763 }
6764 let __s_b = self.gpu.stream();
6765 let mut b = __s_b.launch_builder(&f);
6766 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1)
6767 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6768 unsafe { b.launch(cfg)?; }
6769 Ok(true)
6770 }
6771
6772 pub fn matmul_q4_fused2_batched(&self, w0: &crate::model::GpuTensor,
6777 w1: &crate::model::GpuTensor,
6778 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6779 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6780 use crate::model::GpuTensor;
6781 if m < 2 || m > 8 { return Ok(None); }
6782 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6783 match w {
6784 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6785 Some((*row_bytes, w.out_features())),
6786 _ => None,
6787 }
6788 };
6789 let (Some((rb0, o0)), Some((_rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6790 if w0.in_features() != w1.in_features() { return Ok(None); }
6791 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6792 match w {
6793 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6794 Some(mr) => (mr, true),
6795 None => (bytes, *rp),
6796 },
6797 _ => unreachable!(),
6798 }
6799 }
6800 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6801 if !rp0 || !rp1 { return Ok(None); }
6802 let mcols = Self::batched_mcols(m);
6803 let rpb: u32 = 4;
6804 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6805 let grid = nb(o0) + nb(o1);
6806 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6807 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6808 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f2_rp",
6809 4 => "qmatvec_q4_0_mmvq_b4_f2_rp",
6810 _ => "qmatvec_q4_0_mmvq_b8_f2_rp" });
6811 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6812 shared_mem_bytes: 0 };
6813 let inf = w0.in_features() as i32;
6814 let (oo0, oo1, mi) = (o0 as i32, o1 as i32, m as i32);
6815 let rb = rb0 as i64;
6816 let __s_b = self.gpu.stream();
6817 let mut b = __s_b.launch_builder(&f);
6818 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6819 .arg(&inf).arg(&oo0).arg(&oo1).arg(&mi).arg(&rb);
6820 unsafe { b.launch(cfg)?; }
6821 Ok(Some((y0, y1)))
6822 }
6823
6824 #[allow(clippy::too_many_arguments)]
6827 pub fn matmul_q4_fused3_batched(&self, w0: &crate::model::GpuTensor,
6828 w1: &crate::model::GpuTensor, w2: &crate::model::GpuTensor,
6829 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6830 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6831 use crate::model::GpuTensor;
6832 if m < 2 || m > 8 { return Ok(None); }
6833 let q4 = |w: &GpuTensor| -> Option<usize> {
6834 match w {
6835 GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0 => Some(w.out_features()),
6836 _ => None,
6837 }
6838 };
6839 let (Some(o0), Some(o1), Some(o2)) = (q4(w0), q4(w1), q4(w2)) else { return Ok(None) };
6840 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6841 return Ok(None);
6842 }
6843 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6844 match w {
6845 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6846 Some(mr) => (mr, true),
6847 None => (bytes, *rp),
6848 },
6849 _ => unreachable!(),
6850 }
6851 }
6852 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6853 if !rp0 || !rp1 || !rp2 { return Ok(None); }
6854 let mcols = Self::batched_mcols(m);
6855 let rpb: u32 = 4;
6856 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6857 let grid = nb(o0) + nb(o1) + nb(o2);
6858 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6859 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6860 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
6861 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f3_rp",
6862 4 => "qmatvec_q4_0_mmvq_b4_f3_rp",
6863 _ => "qmatvec_q4_0_mmvq_b8_f3_rp" });
6864 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6865 shared_mem_bytes: 0 };
6866 let inf = w0.in_features() as i32;
6867 let (oo0, oo1, oo2, mi) = (o0 as i32, o1 as i32, o2 as i32, m as i32);
6868 let rb = 0i64;
6869 let __s_b = self.gpu.stream();
6870 let mut b = __s_b.launch_builder(&f);
6871 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6872 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&mi).arg(&rb);
6873 unsafe { b.launch(cfg)?; }
6874 Ok(Some((y0, y1, y2)))
6875 }
6876
6877 pub fn matmul_q8_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6878 w2: &crate::model::GpuTensor,
6879 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6880 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6881 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
6884 return Ok(Some(self.e4m3_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6885 p0.1, p1.1, p2.1, p0.2,
6886 p0.3, p1.3, p2.3)?));
6887 }
6888 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
6889 Ok(Some(self.q8_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6890 p0.1, p1.1, p2.1, p0.2)?))
6891 }
6892
6893 #[allow(clippy::too_many_arguments)]
6894 fn q8_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6895 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6896 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
6897 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6898 const ROWS_PER_BLOCK: u32 = 4;
6899 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6900 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6901 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6902 let f = self.func("qmatvec_q8_0_mmvq_fused3");
6903 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6904 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6905 let mut y2 = self.alloc_uninit::<f32>(out2)?;
6906 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6907 shared_mem_bytes: 0 };
6908 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32, row_bytes as i64);
6909 let __s_b = self.gpu.stream();
6910 let mut b = __s_b.launch_builder(&f);
6911 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6912 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl);
6913 unsafe { b.launch(cfg)?; }
6914 Ok((y0, y1, y2))
6915 }
6916
6917 #[allow(clippy::too_many_arguments)]
6919 pub fn qmatvec_q8_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6920 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
6921 out2: usize, row_bytes: usize)
6922 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6923 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6924 self.q8_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes)
6925 }
6926
6927 pub fn matmul_q8_fused2_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6938 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6939 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6940 if !(2..=8).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
6944 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6947 if m > 4 && !Self::b8_enabled() { return Ok(None); }
6948 return Ok(Some(self.e4m3_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(),
6949 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6950 }
6951 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6952 Ok(Some(self.q8_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(), p0.1, p1.1, p0.2)?))
6953 }
6954
6955 #[allow(clippy::too_many_arguments)]
6956 fn q8_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6957 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6958 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6959 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6960 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6962 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6963 let f = self.func(match Self::batched_mcols(m) {
6964 2 => "qmatvec_q8_0_mmvq_fused2_b2",
6965 4 => "qmatvec_q8_0_mmvq_fused2_b4",
6966 _ => "qmatvec_q8_0_mmvq_fused2_b8",
6968 });
6969 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6970 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6971 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6972 shared_mem_bytes: 0 };
6973 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32, row_bytes as i64);
6974 let __s_b = self.gpu.stream();
6975 let mut b = __s_b.launch_builder(&f);
6976 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6977 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
6978 unsafe { b.launch(cfg)?; }
6979 Ok((y0, y1))
6980 }
6981
6982 #[allow(clippy::too_many_arguments)]
6985 pub fn qmatvec_q8_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6986 x: &CudaSlice<f32>, m: usize,
6987 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6988 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6989 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6990 self.q8_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes)
6991 }
6992
6993 #[allow(clippy::too_many_arguments)]
6996 pub fn matmul_q8_fused3_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6997 w2: &crate::model::GpuTensor,
6998 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6999 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
7000 if !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
7001 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
7002 return Ok(Some(self.e4m3_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
7003 p0.1, p1.1, p2.1, p0.2,
7004 p0.3, p1.3, p2.3)?));
7005 }
7006 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
7007 Ok(Some(self.q8_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
7008 p0.1, p1.1, p2.1, p0.2)?))
7009 }
7010
7011 #[allow(clippy::too_many_arguments)]
7012 fn q8_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7013 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
7014 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
7015 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7016 const ROWS_PER_BLOCK: u32 = 4;
7017 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7018 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7019 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
7020 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_q8_0_mmvq_fused3_b2" }
7021 else { "qmatvec_q8_0_mmvq_fused3_b4" });
7022 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7023 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7024 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
7025 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7026 shared_mem_bytes: 0 };
7027 let (inf, o0, o1, o2, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
7028 m as i32, row_bytes as i64);
7029 let __s_b = self.gpu.stream();
7030 let mut b = __s_b.launch_builder(&f);
7031 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
7032 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&mi).arg(&rbl);
7033 unsafe { b.launch(cfg)?; }
7034 Ok((y0, y1, y2))
7035 }
7036
7037 #[allow(clippy::too_many_arguments)]
7039 pub fn qmatvec_q8_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7040 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
7041 out1: usize, out2: usize, row_bytes: usize)
7042 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7043 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7044 self.q8_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes)
7045 }
7046
7047 pub fn q8_ffn_fuse2_on(&self) -> bool {
7051 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7052 *ON.get_or_init(|| std::env::var("MEMRA_Q8_FFN_FUSE2").as_deref() != Ok("0"))
7053 }
7054
7055 #[allow(clippy::type_complexity)]
7061 fn q8_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
7062 -> Option<[(&'w CudaSlice<u8>, usize, usize); N]> {
7063 use crate::model::GpuTensor;
7064 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return None; }
7065 if std::env::var("MEMRA_Q8_DUAL").is_ok_and(|v| v == "0") { return None; }
7066 let in_f = ws[0].in_features();
7067 let mut out: [Option<(&CudaSlice<u8>, usize, usize)>; N] = [None; N];
7068 for (i, w) in ws.iter().enumerate() {
7069 match w {
7070 GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. }
7071 if *qtype == QT_Q8_0 && *scale == 1.0 && w.in_features() == in_f =>
7072 out[i] = Some((bytes, w.out_features(), *row_bytes)),
7073 _ => return None,
7074 }
7075 }
7076 Some(out.map(|o| o.unwrap()))
7077 }
7078
7079 pub fn e4m3_dual_on(&self) -> bool {
7082 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7083 *ON.get_or_init(|| std::env::var("MEMRA_E4M3_DUAL").as_deref() != Ok("0"))
7084 }
7085
7086 #[allow(clippy::type_complexity)]
7098 fn e4m3_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
7099 -> Option<[(&'w CudaSlice<u8>, usize, usize, f32); N]> {
7100 use crate::model::GpuTensor;
7101 if !self.e4m3_dual_on() { return None; }
7102 let in_f = ws[0].in_features();
7103 let mut out: [Option<(&CudaSlice<u8>, usize, usize, f32)>; N] = [None; N];
7104 for (i, w) in ws.iter().enumerate() {
7105 match w {
7106 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, rp4, .. }
7107 if *qtype == QT_F8_E4M3 && w.in_features() == in_f
7108 && *row_bytes == in_f && !*rp && rp4.is_none() =>
7109 out[i] = Some((bytes, w.out_features(), *row_bytes, *scale)),
7110 _ => return None,
7111 }
7112 }
7113 Some(out.map(|o| o.unwrap()))
7114 }
7115
7116 #[allow(clippy::too_many_arguments)]
7120 fn e4m3_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7121 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7122 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7123 ws0: f32, ws1: f32)
7124 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7125 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7127 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7128 let f = self.func("qmatvec_e4m3_mmvq_fused2");
7129 let mut y0 = self.alloc_uninit::<f32>(out0)?;
7130 let mut y1 = self.alloc_uninit::<f32>(out1)?;
7131 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7132 shared_mem_bytes: 0 };
7133 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
7134 let __s_b = self.gpu.stream();
7135 let mut b = __s_b.launch_builder(&f);
7136 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
7137 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl).arg(&ws0).arg(&ws1);
7138 unsafe { b.launch(cfg)?; }
7139 Ok((y0, y1))
7140 }
7141
7142 #[allow(clippy::too_many_arguments)]
7144 fn e4m3_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7145 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7146 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
7147 ws0: f32, ws1: f32, ws2: f32)
7148 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7149 const ROWS_PER_BLOCK: u32 = 4;
7150 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7151 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7152 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
7153 let f = self.func("qmatvec_e4m3_mmvq_fused3");
7154 let mut y0 = self.alloc_uninit::<f32>(out0)?;
7155 let mut y1 = self.alloc_uninit::<f32>(out1)?;
7156 let mut y2 = self.alloc_uninit::<f32>(out2)?;
7157 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7158 shared_mem_bytes: 0 };
7159 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
7160 row_bytes as i64);
7161 let __s_b = self.gpu.stream();
7162 let mut b = __s_b.launch_builder(&f);
7163 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
7164 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl).arg(&ws0).arg(&ws1).arg(&ws2);
7165 unsafe { b.launch(cfg)?; }
7166 Ok((y0, y1, y2))
7167 }
7168
7169 #[allow(clippy::too_many_arguments)]
7173 fn e4m3_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7174 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
7175 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7176 ws0: f32, ws1: f32)
7177 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7178 const ROWS_PER_BLOCK: u32 = 4;
7179 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7180 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7181 let f = self.func(match Self::batched_mcols(m) {
7182 2 => "qmatvec_e4m3_mmvq_fused2_b2",
7183 4 => "qmatvec_e4m3_mmvq_fused2_b4",
7184 _ => "qmatvec_e4m3_mmvq_fused2_b8",
7185 });
7186 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7187 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7188 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7189 shared_mem_bytes: 0 };
7190 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32,
7191 row_bytes as i64);
7192 let __s_b = self.gpu.stream();
7193 let mut b = __s_b.launch_builder(&f);
7194 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
7195 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
7196 unsafe { b.launch(cfg)?; }
7197 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
7198 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
7199 Ok((y0, y1))
7200 }
7201
7202 #[allow(clippy::too_many_arguments)]
7204 fn e4m3_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7205 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
7206 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
7207 ws0: f32, ws1: f32, ws2: f32)
7208 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7209 const ROWS_PER_BLOCK: u32 = 4;
7210 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7211 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7212 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
7213 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_e4m3_mmvq_fused3_b2" }
7214 else { "qmatvec_e4m3_mmvq_fused3_b4" });
7215 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7216 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7217 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
7218 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7219 shared_mem_bytes: 0 };
7220 let (inf, o0, o1, o2, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
7221 m as i32, row_bytes as i64);
7222 let __s_b = self.gpu.stream();
7223 let mut b = __s_b.launch_builder(&f);
7224 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
7225 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&mi).arg(&rbl);
7226 unsafe { b.launch(cfg)?; }
7227 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
7228 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
7229 if ws2 != 1.0 { self.scale_inplace(&mut y2, ws2, m * out2)?; }
7230 Ok((y0, y1, y2))
7231 }
7232
7233 pub fn qmatvec_e4m3_blk_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7243 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7244 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7245 scale_cols: usize)
7246 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7247 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_e4m3_blk_mmvq_into(bytes, aq, ad, scales, m, in_f, out_f, row_bytes,
7249 scale_cols, &mut y)?;
7250 Ok(y)
7251 }
7252
7253 #[allow(clippy::too_many_arguments)]
7255 pub fn qmatvec_e4m3_blk_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7256 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7257 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7258 scale_cols: usize, y: &mut CudaSlice<f32>)
7259 -> Result<(), Box<dyn std::error::Error>> {
7260 const ROWS_PER_BLOCK: u32 = 4; let f = self.func("qmatvec_e4m3_blk_mmvq");
7262 let cfg = LaunchConfig {
7263 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), m as u32, 1),
7264 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7267 let (inf, outf, mi, rb, sc) =
7268 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7269 let __s_b = self.gpu.stream();
7270 let mut b = __s_b.launch_builder(&f);
7271 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut *y)
7272 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7273 unsafe { b.launch(cfg)?; }
7274 Ok(())
7275 }
7276
7277 #[allow(clippy::too_many_arguments)]
7283 pub fn qmatvec_e4m3_blk_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7284 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7285 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7286 scale_cols: usize, mcols: usize)
7287 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7288 const ROWS_PER_BLOCK: u32 = 4; debug_assert!(mcols >= m, "blk batched: mcols {mcols} < m {m}");
7290 let name = match mcols {
7291 2 => "qmatvec_e4m3_blk_mmvq_b2",
7292 4 => "qmatvec_e4m3_blk_mmvq_b4",
7293 8 => "qmatvec_e4m3_blk_mmvq_b8",
7294 16 => "qmatvec_e4m3_blk_mmvq_b16",
7295 _ => return Err(format!("qmatvec_e4m3_blk_mmvq_batched: no kernel for mcols {mcols}").into()),
7296 };
7297 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7298 let f = self.func(name);
7299 let cfg = LaunchConfig {
7300 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
7301 block_dim: (32, ROWS_PER_BLOCK, 1),
7302 shared_mem_bytes: 0,
7303 };
7304 let (inf, outf, mi, rb, sc) =
7305 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7306 let __s_b = self.gpu.stream();
7307 let mut b = __s_b.launch_builder(&f);
7308 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut y)
7309 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7310 unsafe { b.launch(cfg)?; }
7311 Ok(y)
7312 }
7313
7314 #[allow(clippy::too_many_arguments)]
7317 pub fn qmatvec_e4m3_blk_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7318 scales: &CudaSlice<f32>, m: usize, in_f: usize,
7319 out_f: usize, row_bytes: usize, scale_cols: usize,
7320 mcols: usize)
7321 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7322 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7323 self.qmatvec_e4m3_blk_mmvq_batched(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes,
7324 scale_cols, mcols)
7325 }
7326
7327 #[allow(clippy::too_many_arguments)]
7330 pub fn qmatvec_e4m3_blk_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7331 scales: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
7332 row_bytes: usize, scale_cols: usize)
7333 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7334 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7335 self.qmatvec_e4m3_blk_mmvq(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols)
7336 }
7337
7338 #[allow(clippy::too_many_arguments)]
7341 pub fn qmatvec_e4m3_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
7342 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7343 ws0: f32, ws1: f32)
7344 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7345 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7346 self.e4m3_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes, ws0, ws1)
7347 }
7348
7349 #[allow(clippy::too_many_arguments)]
7350 pub fn qmatvec_e4m3_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7351 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
7352 out2: usize, row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7353 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7354 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7355 self.e4m3_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes,
7356 ws0, ws1, ws2)
7357 }
7358
7359 #[allow(clippy::too_many_arguments)]
7360 pub fn qmatvec_e4m3_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7361 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
7362 out1: usize, row_bytes: usize, ws0: f32, ws1: f32)
7363 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7364 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7365 self.e4m3_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes, ws0, ws1)
7366 }
7367
7368 #[allow(clippy::too_many_arguments)]
7369 pub fn qmatvec_e4m3_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7370 b2: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
7371 in_f: usize, out0: usize, out1: usize, out2: usize,
7372 row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7373 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7374 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7375 self.e4m3_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes,
7376 ws0, ws1, ws2)
7377 }
7378
7379 fn try_e4m3_blk_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
7390 ad: &CudaSlice<f32>, m: usize)
7391 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7392 use crate::model::GpuTensor;
7393 if let GpuTensor::Quant { bytes, qtype, row_bytes, blk: Some(g), .. } = w {
7394 if *qtype == QT_F8_E4M3_BLK {
7395 if (2..=16).contains(&m) && std::env::var("MEMRA_NO_BATCHED").is_err()
7401 && (m <= 4 || Self::b8_enabled()) {
7402 let mcols = Self::batched_mcols(m);
7403 return Ok(Some(self.qmatvec_e4m3_blk_mmvq_batched(
7404 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7405 *row_bytes, g.cols, mcols)?));
7406 }
7407 return Ok(Some(self.qmatvec_e4m3_blk_mmvq(
7408 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7409 *row_bytes, g.cols)?));
7410 }
7411 }
7412 Ok(None)
7413 }
7414
7415 fn try_e4m3_blk_prefill(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
7462 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7463 use crate::model::GpuTensor;
7464 let GpuTensor::Quant { bytes, qtype, blk: Some(g), .. } = w else { return Ok(None) };
7465 if *qtype != QT_F8_E4M3_BLK { return Ok(None) }
7466 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(Some(y)); }
7471 let (in_f, out_f) = (w.in_features(), w.out_features());
7472 let slab = self.fp8_blk_dequant_q8_0_dev(bytes, &g.scales, out_f, in_f)?;
7473 let tmp = GpuTensor::Quant {
7474 bytes: slab,
7475 qtype: QT_Q8_0,
7476 row_bytes: in_f / 32 * 34,
7477 ne: vec![in_f as u64, out_f as u64],
7478 scale: 1.0,
7479 rp: false,
7480 #[cfg(memra_cutlass)]
7481 cutlass: None,
7482 fp8: None, blk: None, f16: None, rp4: None,
7483 };
7484 Ok(Some(self.matmul(&tmp, x, m)?))
7486 }
7487
7488 pub fn matmul_pre_noscale(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7489 m: usize) -> Result<Option<(CudaSlice<f32>, f32)>, Box<dyn std::error::Error>> {
7490 use crate::model::GpuTensor;
7491 if m == 1 {
7495 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(Some((y, 1.0))); }
7496 }
7497 if m != 1 || !self.uses_q8_1_fast(w) { return Ok(None); }
7499 let in_f = w.in_features();
7500 let out_f = w.out_features();
7501 let (bytes, qtype, row_bytes, scale, rp) = match w {
7502 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
7503 _ => return Ok(None),
7504 };
7505 if self.mmvq_supports(qtype) {
7507 let (mbytes, mrp) = match w {
7509 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
7510 _ => (bytes, rp),
7511 };
7512 let y = self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, 1.0, mrp)?;
7513 return Ok(Some((y, scale)));
7514 }
7515 let name = match qtype {
7517 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
7518 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
7519 QT_Q3_K => "qmatvec_q3_K_dp4a",
7520 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
7521 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
7522 _ => return Ok(None),
7523 };
7524 let f = self.func(name);
7525 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7526 let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
7527 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7528 let __s_b = self.gpu.stream();
7529 let mut b = __s_b.launch_builder(&f);
7530 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7531 unsafe { b.launch(cfg)?; }
7532 Ok(Some((y, scale)))
7533 }
7534
7535 pub fn mmvq_supports(&self, qtype: i32) -> bool {
7538 if qtype == QT_F8_E4M3 { return true; }
7543 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return false; }
7544 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_Q4_0)
7545 }
7546
7547 pub fn qmatvec_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7552 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7553 rp: bool)
7554 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7555 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_mmvq_into(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp, &mut y)?;
7557 Ok(y)
7558 }
7559
7560 #[allow(clippy::too_many_arguments)]
7562 pub fn qmatvec_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7563 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7564 rp: bool, y: &mut CudaSlice<f32>)
7565 -> Result<(), Box<dyn std::error::Error>> {
7566 debug_assert!(y.len() >= m * out_f);
7567 const ROWS_PER_BLOCK: u32 = 4; if qtype == QT_Q8_0 && rp && m == 1 && out_f >= 64
7573 && (out_f as u32).div_ceil(ROWS_PER_BLOCK) < 4 * self.sm_count() as u32
7574 && {
7575 static G2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7576 *G2.get_or_init(|| std::env::var("MEMRA_Q80_G2").as_deref() != Ok("0"))
7577 }
7578 {
7579 let f = self.func("qmatvec_q8_0_mmvq_rp_g2");
7580 let cfg = LaunchConfig {
7581 grid_dim: ((out_f as u32).div_ceil(2), 1, 1),
7582 block_dim: (32, 2, 1),
7583 shared_mem_bytes: 0,
7584 };
7585 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, row_bytes as i64);
7586 let __s_b = self.gpu.stream();
7587 let mut b = __s_b.launch_builder(&f);
7588 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7589 unsafe { b.launch(cfg)?; }
7590 if scale != 1.0 { self.scale_inplace(y, scale, out_f)?; }
7591 return Ok(());
7592 }
7593 let mut mr: u32 = if m == 1 && (qtype == QT_NVFP4 || qtype == QT_Q5_K) { 2 } else { 1 };
7602 if m == 1 && qtype == QT_Q4_0 {
7607 static Q40MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7608 mr = *Q40MR.get_or_init(|| std::env::var("MEMRA_Q40_MR").ok()
7611 .and_then(|v| v.parse().ok()).unwrap_or(1));
7612 }
7613 let q5_mode = std::env::var("MEMRA_Q5K_ISSUE").ok();
7624 let q5_force = q5_mode.as_deref() == Some("2");
7625 let q5_il = qtype == QT_Q5_K && m == 1
7628 && (q5_force || q5_mode.as_deref().map(|v| v != "0").unwrap_or(true));
7629 if q5_il && !q5_force && out_f > 65536 { mr = 1; }
7630 if qtype == QT_Q4_0 && rp && mr != 1 { mr = 2; }
7633 if qtype == QT_Q8_0 && rp {
7637 static Q80MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7638 mr = *Q80MR.get_or_init(|| std::env::var("MEMRA_Q80_MR").ok()
7639 .and_then(|v| v.parse().ok()).unwrap_or(1));
7640 }
7641 let name = match (qtype, mr, rp) {
7642 (QT_NVFP4, 2, false) => "qmatvec_nvfp4_mmvq_mr2",
7643 (QT_NVFP4, 2, true) => "qmatvec_nvfp4_mmvq_mr2_rp",
7644 (QT_NVFP4, _, true) => "qmatvec_nvfp4_mmvq_rp",
7645 (QT_Q4_0, 1, true) => "qmatvec_q4_0_mmvq_rp",
7646 (QT_Q4_0, _, true) => "qmatvec_q4_0_mmvq_mr2_rp",
7647 (QT_Q5_K, 2, _) => if q5_il { "qmatvec_q5_K_mmvq_mr2_il" } else { "qmatvec_q5_K_mmvq_mr2" },
7648 (QT_Q8_0, 2, true) => "qmatvec_q8_0_mmvq_mr2_rp",
7649 (QT_Q8_0, _, true) if in_f % 1024 == 0 && {
7654 static CA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7655 *CA.get_or_init(|| std::env::var("MEMRA_Q80_CA").as_deref() == Ok("1"))
7656 } => "qmatvec_q8_0_mmvq_rpca",
7657 (QT_Q8_0, _, true) => "qmatvec_q8_0_mmvq_rp",
7658 (QT_Q8_0, _, _) => "qmatvec_q8_0_mmvq",
7659 (QT_Q4_K, _, true) => "qmatvec_q4_K_mmvq_rp",
7663 (QT_Q6_K, _, true) => "qmatvec_q6_K_mmvq_rp",
7664 (QT_Q4_K, _, _) => "qmatvec_q4_K_mmvq",
7665 (QT_Q4_0, 2, false) => "qmatvec_q4_0_mmvq_mr2",
7666 (QT_Q4_0, _, false) => "qmatvec_q4_0_mmvq",
7667 (QT_Q5_K, _, _) => if q5_il { "qmatvec_q5_K_mmvq_il" } else { "qmatvec_q5_K_mmvq" },
7668 (QT_Q6_K, _, _) => "qmatvec_q6_K_mmvq",
7669 (QT_NVFP4, _, false) => "qmatvec_nvfp4_mmvq",
7670 (QT_F8_E4M3, _, _) => "qmatvec_e4m3_mmvq",
7671 _ => panic!("qmatvec_mmvq: qtype {qtype} has no MMVQ kernel"),
7672 };
7673 let f = self.func(name);
7674 let rows_per_block = ROWS_PER_BLOCK * mr;
7676 let cfg = LaunchConfig {
7677 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, m as u32, 1),
7678 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7681 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7682 let __s_b = self.gpu.stream();
7683 let mut b = __s_b.launch_builder(&f);
7684 if qtype == QT_NVFP4 || qtype == QT_F8_E4M3 {
7689 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&scale);
7690 unsafe { b.launch(cfg)?; }
7691 } else if Self::pdl_on() && Self::pdl_mmvq_on()
7692 && matches!(name, "qmatvec_q4_0_mmvq_rp" | "qmatvec_q6_K_mmvq"
7693 | "qmatvec_q6_K_mmvq_rp") {
7694 {
7698 use cudarc::driver::{DevicePtr, DevicePtrMut};
7699 let s = &self.gpu.stream();
7700 let (pw, _g0) = bytes.device_ptr(s); let (paq, _g1) = aq.device_ptr(s);
7701 let (pad, _g2) = ad.device_ptr(s); let (py, _g3) = y.device_ptr_mut(s);
7702 let mut ps = [
7703 &pw as *const _ as *mut std::ffi::c_void, &paq as *const _ as *mut _,
7704 &pad as *const _ as *mut _, &py as *const _ as *mut _,
7705 &inf as *const _ as *mut _, &outf as *const _ as *mut _,
7706 &mi as *const _ as *mut _, &rb as *const _ as *mut _,
7707 ];
7708 unsafe { self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?; }
7709 }
7710 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7711 } else {
7712 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7713 unsafe { b.launch(cfg)?; }
7714 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7715 }
7716 Ok(())
7717 }
7718
7719 pub fn qmatvec_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
7723 out_f: usize, qtype: i32, row_bytes: usize, rp: bool)
7724 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7725 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7726 self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, 1.0, rp)
7727 }
7728
7729 pub fn batched_supports(&self, qtype: i32) -> bool {
7733 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_F8_E4M3 | QT_Q4_0)
7734 }
7735
7736 pub fn iq_fast_enabled() -> bool {
7744 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7745 *ON.get_or_init(|| std::env::var("MEMRA_IQ_FAST").map(|v| v != "0").unwrap_or(true))
7746 }
7747
7748 pub fn b8_enabled() -> bool {
7751 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7752 *ON.get_or_init(|| std::env::var("MEMRA_B8").map(|v| v != "0").unwrap_or(true))
7753 }
7754
7755 pub fn batched_mcols(m: usize) -> usize {
7757 if m == 2 { 2 } else if m <= 4 { 4 } else if m <= 8 { 8 } else { 16 }
7758 }
7759
7760 fn batched_kernel_name(qtype: i32, mcols: usize) -> Option<&'static str> {
7765 Some(match (qtype, mcols) {
7766 (QT_Q8_0, 2) => "qmatvec_q8_0_mmvq_b2", (QT_Q8_0, 4) => "qmatvec_q8_0_mmvq_b4",
7767 (QT_Q8_0, 8) => "qmatvec_q8_0_mmvq_b8",
7768 (QT_Q8_0, 16) => "qmatvec_q8_0_mmvq_b16",
7774 (QT_Q4_K, 2) => "qmatvec_q4_K_mmvq_b2", (QT_Q4_K, 4) => "qmatvec_q4_K_mmvq_b4",
7775 (QT_Q4_K, 8) => "qmatvec_q4_K_mmvq_b8",
7776 (QT_Q4_K, 16) => "qmatvec_q4_K_mmvq_b16",
7779 (QT_Q5_K, 2) => "qmatvec_q5_K_mmvq_b2", (QT_Q5_K, 4) => "qmatvec_q5_K_mmvq_b4",
7780 (QT_Q5_K, 8) => "qmatvec_q5_K_mmvq_b8",
7781 (QT_Q5_K, 16) => "qmatvec_q5_K_mmvq_b16",
7784 (QT_Q6_K, 2) => "qmatvec_q6_K_mmvq_b2", (QT_Q6_K, 4) => "qmatvec_q6_K_mmvq_b4",
7785 (QT_Q6_K, 8) => "qmatvec_q6_K_mmvq_b8", (QT_Q6_K, 16) => "qmatvec_q6_K_mmvq_b16",
7786 (QT_NVFP4, 2) => "qmatvec_nvfp4_mmvq_b2", (QT_NVFP4, 4) => "qmatvec_nvfp4_mmvq_b4",
7787 (QT_NVFP4, 8) => "qmatvec_nvfp4_mmvq_b8",
7788 (QT_NVFP4, 16) => "qmatvec_nvfp4_mmvq_b16",
7792 (QT_F8_E4M3, 2) => "qmatvec_e4m3_mmvq_b2", (QT_F8_E4M3, 4) => "qmatvec_e4m3_mmvq_b4",
7793 (QT_F8_E4M3, 8) => "qmatvec_e4m3_mmvq_b8",
7794 (QT_F8_E4M3, 16) => "qmatvec_e4m3_mmvq_b16",
7798 (QT_Q4_0, 2) => "qmatvec_q4_0_mmvq_b2", (QT_Q4_0, 4) => "qmatvec_q4_0_mmvq_b4",
7799 (QT_Q4_0, 8) => "qmatvec_q4_0_mmvq_b8", (QT_Q4_0, 16) => "qmatvec_q4_0_mmvq_b16",
7800 _ => return None,
7801 })
7802 }
7803
7804 pub fn sm_count(&self) -> i32 {
7839 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7840 *SMS.get_or_init(|| {
7841 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7842 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7843 })
7844 }
7845
7846 pub fn batched_variant(&self, _m: usize, in_f: usize, out_f: usize, qtype: i32,
7847 row_bytes: usize, mcols: usize, rp: bool) -> &'static str {
7848 if qtype == QT_Q8_0 {
7853 return if rp { "rp" } else { "base" };
7854 }
7855 static BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7856 let bv = *BV.get_or_init(|| match std::env::var("MEMRA_MMVQ_BV").as_deref() {
7857 Ok("base") => "base", Ok("pf") => "pf", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7858 Ok("pfr2") => "pfr2", Ok("ca") => "ca", Ok("car2") => "car2",
7859 Ok("rp") => "rp", Ok("rpr2") => "rpr2", Ok("rpr2w8") => "rpr2w8",
7862 Ok("rpca") => "rpca", Ok("rpcar2") => "rpcar2",
7865 Ok("rpsc") => "rpsc", Ok("rpms") => "rpms", Ok("rpmsc") => "rpmsc",
7872 Ok("rpks") => "rpks", Ok("rpksc") => "rpksc",
7873 _ => "auto",
7874 });
7875 let ca_ok = qtype == QT_NVFP4 && (row_bytes % 16 == 0) && (in_f % 1024 == 0);
7879 static KS_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7884 let ks_on = *KS_ON.get_or_init(|| std::env::var("MEMRA_KS").as_deref() != Ok("0"));
7885 let sc_ok = ks_on && qtype == QT_NVFP4 && (in_f % 256 == 0) && (in_f / 64 <= 272);
7886 let ks_ok = ks_on && qtype == QT_NVFP4 && (in_f % 512 == 0) && (in_f / 64 <= 272);
7887 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7888 let sms = *SMS.get_or_init(|| {
7889 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7890 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7891 });
7892 let kq_r2 = matches!(qtype, QT_Q4_K | QT_Q5_K | QT_Q6_K);
7912 static KQBV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7915 let kq_bv = *KQBV.get_or_init(|| match std::env::var("MEMRA_KQ_BV").as_deref() {
7916 Ok("base") => "base", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7917 _ => "auto",
7918 });
7919 let variant: &'static str = if qtype == QT_Q4_0 {
7920 static Q40BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7924 let q40 = *Q40BV.get_or_init(|| match std::env::var("MEMRA_Q40_BV").as_deref() {
7925 Ok("base") => "base", Ok("r2") => "r2", Ok("ms") => "ms", Ok("sm") => "sm",
7931 Ok("la") => "la", _ => "auto",
7932 });
7933 let v = if q40 != "auto" { q40 }
7934 else if (out_f as u32).div_ceil(8) >= 4 * sms as u32 { "r2" } else { "base" };
7935 if rp { match v { "ms" => "r2ms_rp", "sm" => "r2sm_rp", "la" => "r2la_rp",
7940 "r2" => "r2_rp", _ => "rp" } }
7941 else if matches!(v, "ms" | "sm" | "la") { "r2" } else { v }
7942 } else if qtype != QT_NVFP4 && !kq_r2 {
7943 "base"
7944 } else if kq_r2 && rp {
7945 "rp"
7949 } else if kq_r2 {
7950 if kq_bv != "auto" {
7953 if kq_bv == "r2w8" && mcols != 4 { "r2" } else { kq_bv }
7954 } else if bv != "auto" {
7955 match bv {
7956 "r2" | "pfr2" | "rpr2" | "car2" => "r2",
7957 "r2w8" | "rpr2w8" => if mcols != 4 { "r2" } else { "r2w8" },
7958 _ => "base", }
7960 } else {
7961 let blocks = (out_f + 7) / 8;
7962 let waves = blocks as f64 / (7 * sms as usize) as f64;
7963 let filled = blocks >= 4 * sms as usize;
7964 let use_r2 = if qtype == QT_Q4_K { filled } else { waves >= 2.0 };
7965 if use_r2 { "r2" } else { "base" }
7966 }
7967 } else if bv != "auto" {
7968 let v = if bv == "r2w8" && mcols == 2 { "r2" }
7973 else if bv == "ca" && (!ca_ok || mcols == 8) { "pf" }
7974 else if bv == "car2" && (!ca_ok || mcols == 8) { "r2" }
7975 else if bv == "pfr2" && mcols == 8 { "r2" }
7976 else if (bv == "rpr2w8" || bv == "rpr2") && mcols == 2 { "rpr2" }
7977 else if (bv == "rpca" || bv == "rpcar2") && (!ca_ok || mcols == 8) {
7979 if mcols == 8 { "rpr2w8" } else { "rpr2" }
7980 }
7981 else if bv == "rpcar2" && mcols == 2 { "rpca" }
7982 else if (bv == "rpsc" || bv == "rpmsc") && !sc_ok { "rpr2" }
7985 else if (bv == "rpks" || bv == "rpksc") && !ks_ok { "rpr2" }
7986 else { bv };
7987 if rp {
7988 match v {
7989 "base" | "pf" | "ca" | "rp" => "rp",
7990 "r2" | "pfr2" | "car2" | "rpr2" => "rpr2",
7991 "r2w8" | "rpr2w8" => if mcols == 2 { "rpr2" } else { "rpr2w8" },
7992 other => other, }
7994 } else { v }
7995 } else if mcols == 8 {
7996 if rp { if sc_ok { "rpsc" } else { "rpr2w8" } } else { "r2w8" }
8007 } else if mcols >= 4 {
8008 let blocks = (out_f + 7) / 8;
8012 let r7 = 7 * sms as usize;
8013 let r8 = 8 * sms as usize;
8014 let waves = blocks as f64 / r7 as f64;
8015 let filled = blocks >= 4 * sms as usize;
8016 if filled && blocks.div_ceil(r8) < blocks.div_ceil(r7) {
8020 if rp { "rpr2w8" } else { "r2w8" }
8024 } else if waves >= 2.0 || (waves <= 1.0 && filled) {
8025 if rp { "rpr2" } else { "r2" }
8028 } else {
8029 if rp { "rp" } else { "pf" }
8033 }
8034 } else if in_f >= 6144 {
8035 if rp { "rpr2" } else { "r2" }
8039 }
8040 else if rp {
8041 let waves = ((out_f + 7) / 8) as f64 / (7 * sms as usize) as f64;
8046 if sc_ok && waves >= 0.9 && waves <= 1.1 { "rpsc" } else { "rp" }
8047 } else { "base" };
8048 variant
8049 }
8050
8051 pub fn qmatvec_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
8052 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize,
8053 mcols: usize, scale: f32, rp: bool)
8054 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8055 const ROWS_PER_BLOCK: u32 = 4;
8056 let forced: Option<&'static str> = {
8061 static V: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
8062 V.get_or_init(|| std::env::var("MEMRA_BVAR").ok())
8063 .as_deref()
8064 .map(|s| Box::leak(s.to_string().into_boxed_str()) as &'static str)
8065 };
8066 let variant = match forced {
8067 Some(v) if !rp || v.contains("rp") => v,
8068 _ => self.batched_variant(m, in_f, out_f, qtype, row_bytes, mcols, rp),
8069 };
8070 let base_name = Self::batched_kernel_name(qtype, mcols)
8071 .ok_or_else(|| format!("qmatvec_mmvq_batched: no kernel for qtype {qtype} mcols {mcols}"))?;
8072 let variant = if mcols == 16 { if rp { "rp" } else { "base" } } else { variant };
8076 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8083 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
8084 if b567 && qtype == QT_NVFP4 && rp && mcols == 8 && (5..=7).contains(&m)
8085 && matches!(variant, "rpsc" | "rpr2w8") {
8086 let f = self.func(&format!("qmatvec_nvfp4_mmvq_b{m}_{variant}"));
8087 let rows_per_block = ROWS_PER_BLOCK * 2; let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8089 let cfg = LaunchConfig {
8090 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
8091 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0 };
8092 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8093 let __s_b = self.gpu.stream();
8094 let mut b = __s_b.launch_builder(&f);
8095 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8096 unsafe { b.launch(cfg)?; }
8097 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8098 return Ok(y);
8099 }
8100 let (name, rows_per_block): (std::borrow::Cow<'static, str>, u32) = match variant {
8101 "base" => (base_name.into(), ROWS_PER_BLOCK),
8102 "pf" => (format!("{base_name}_pf").into(), ROWS_PER_BLOCK),
8103 "ca" => (format!("{base_name}_ca").into(), ROWS_PER_BLOCK),
8104 "rp" => (format!("{base_name}_rp").into(), ROWS_PER_BLOCK),
8105 "rpca" => (format!("{base_name}_rpca").into(), ROWS_PER_BLOCK), "rpks" => (format!("{base_name}_rpks").into(), ROWS_PER_BLOCK),
8109 "rpksc" => (format!("{base_name}_rpksc").into(), ROWS_PER_BLOCK),
8110 "rpms" => (format!("{base_name}_rpms").into(), ROWS_PER_BLOCK),
8111 "rpmsc" => (format!("{base_name}_rpmsc").into(), ROWS_PER_BLOCK),
8112 "r2ms_rp" => (format!("{base_name}_r2ms_rp").into(), ROWS_PER_BLOCK),
8113 "r2sm_rp" => (format!("{base_name}_r2sm_rp").into(), ROWS_PER_BLOCK * 2),
8114 "r2la_rp" => (format!("{base_name}_r2la_rp").into(), ROWS_PER_BLOCK * 2),
8115 v => (format!("{base_name}_{v}").into(), ROWS_PER_BLOCK * 2), };
8117 debug_assert!(!rp || name.contains("_rp"), "rp weight dispatched to a GGUF-layout kernel");
8118 let f = self.func(&name);
8119 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8120 let smem = if name.contains("_r2sm_rp") { (mcols * 32 * 9 * 4 + mcols * 32 * 4) as u32 }
8122 else { 0 };
8123 let cfg = LaunchConfig {
8124 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
8125 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: smem };
8126 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8127 let __s_b = self.gpu.stream();
8128 let mut b = __s_b.launch_builder(&f);
8129 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8130 unsafe { b.launch(cfg)?; }
8131 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8132 Ok(y)
8133 }
8134
8135 pub fn qmatvec_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
8139 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, mcols: usize,
8140 rp: bool)
8141 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8142 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
8143 self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, 1.0, rp)
8144 }
8145
8146 pub fn qmatvec_nvfp4_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
8148 in_f: usize, out_f: usize, row_bytes: usize, mcols: usize,
8149 rp: bool)
8150 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8151 self.qmatvec_batched_raw(bytes, x, m, in_f, out_f, QT_NVFP4, row_bytes, mcols, rp)
8152 }
8153
8154 fn try_fp4_gemm(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize,
8158 in_f: usize, out_f: usize)
8159 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8160 use crate::model::GpuTensor;
8161 if cfg!(memra_portable_cuda) { return Ok(None); }
8162 if std::env::var("MEMRA_FP4").is_err() { return Ok(None); }
8163 #[cfg(memra_cutlass)]
8172 if m >= 128 && std::env::var("MEMRA_FP4_CUTLASS").is_ok() {
8173 if let GpuTensor::Quant { bytes, qtype, scale, row_bytes, cutlass, .. } = w {
8174 if *qtype == QT_NVFP4 && in_f % 64 == 0 {
8175 if let Some(cw) = cutlass {
8176 let y = self.cutlass_fp4_gemm(&cw.b_packed, &cw.sfb_swizzled, x, *scale,
8178 m, out_f, in_f)?;
8179 return Ok(Some(y));
8180 } else if std::env::var("MEMRA_FP4_CUTLASS_OTF").is_ok() {
8181 let (b_packed, sfb_sw) = self.build_cutlass_weight(bytes, out_f, in_f, *row_bytes)?;
8186 let y = self.cutlass_fp4_gemm(&b_packed, &sfb_sw, x, *scale, m, out_f, in_f)?;
8187 return Ok(Some(y));
8188 }
8189 }
8190 }
8191 }
8192 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } = w {
8193 if *qtype == QT_NVFP4 && in_f % 64 == 0 && !*rp {
8196 let y = self.qmatvec_gemm_nvfp4_fp4(bytes, x, m, in_f, out_f, *row_bytes, *scale)?;
8197 return Ok(Some(y));
8198 }
8199 }
8200 Ok(None)
8201 }
8202
8203 pub fn rms_norm_f16out(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>,
8207 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
8208 ncols: usize, nrows: usize, eps: f32)
8209 -> Result<(), Box<dyn std::error::Error>> {
8210 let f = self.func("rms_norm_f16out_f32");
8211 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
8212 let (nc, e) = (ncols as i32, eps);
8213 let __s_b = self.gpu.stream();
8214 let mut b = __s_b.launch_builder(&f);
8215 b.arg(x).arg(w).arg(dst).arg(dst16).arg(&nc).arg(&e);
8216 unsafe { b.launch(cfg)?; }
8217 Ok(())
8218 }
8219
8220 #[allow(clippy::too_many_arguments)]
8223 pub fn add_rms_norm_f16out(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
8224 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
8225 dst16: &mut CudaSlice<u8>, ncols: usize, nrows: usize, eps: f32)
8226 -> Result<(), Box<dyn std::error::Error>> {
8227 let f = self.func("add_rms_norm_f16out_f32");
8228 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
8229 let (nc, e) = (ncols as i32, eps);
8230 let __s_lb = self.gpu.stream();
8231 let mut lb = __s_lb.launch_builder(&f);
8232 lb.arg(a).arg(b).arg(w).arg(res).arg(dst).arg(dst16).arg(&nc).arg(&e);
8233 unsafe { lb.launch(cfg)?; }
8234 Ok(())
8235 }
8236
8237 pub fn matmul_group_xh(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>,
8240 xh: &CudaSlice<u8>, m: usize)
8241 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8242 let mut out = Vec::with_capacity(ws.len());
8243 let in_f = ws[0].in_features();
8244 for w in ws {
8245 if w.in_features() == in_f && m >= 16 && !self.verify_exact_on() {
8246 if let Some(y) = self.try_f16_gemm_pre(w, xh, m)? {
8247 out.push(y);
8248 continue;
8249 }
8250 }
8251 out.push(self.matmul(w, x, m)?);
8252 }
8253 Ok(out)
8254 }
8255
8256 pub fn gdn_pad_mask(&self, beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
8259 len_d: &CudaSlice<i32>, h: usize, t: usize)
8260 -> Result<(), Box<dyn std::error::Error>> {
8261 let f = self.func("gdn_pad_mask_f32");
8262 let cfg = LaunchConfig::for_num_elems((t * h) as u32);
8263 let (hi, ti) = (h as i32, t as i32);
8264 let __s_b = self.gpu.stream();
8265 let mut b = __s_b.launch_builder(&f);
8266 b.arg(beta).arg(g_log).arg(len_d).arg(&hi).arg(&ti);
8267 unsafe { b.launch(cfg)?; }
8268 Ok(())
8269 }
8270
8271 pub fn row_gather_dev(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
8274 len_d: &CudaSlice<i32>, ncols: usize)
8275 -> Result<(), Box<dyn std::error::Error>> {
8276 let f = self.func("row_gather_dev_f32");
8277 let cfg = LaunchConfig::for_num_elems(ncols as u32);
8278 let nc = ncols as i32;
8279 let __s_b = self.gpu.stream();
8280 let mut b = __s_b.launch_builder(&f);
8281 b.arg(src).arg(dst).arg(len_d).arg(&nc);
8282 unsafe { b.launch(cfg)?; }
8283 Ok(())
8284 }
8285
8286 pub fn matmul_group(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>, m: usize)
8293 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8294 use crate::model::GpuTensor;
8295 let mut out = Vec::with_capacity(ws.len());
8296 let any_mirror = ws.iter().any(|w| matches!(w, GpuTensor::Quant { f16: Some(_), .. }));
8297 if m >= 16 && any_mirror && !self.verify_exact_on() {
8298 let in_f = ws[0].in_features();
8299 let xh = self.f16_act(x, m * in_f, in_f)?;
8300 for w in ws {
8301 if w.in_features() == in_f {
8302 if let Some(y) = self.try_f16_gemm_pre(w, &xh, m)? {
8303 out.push(y);
8304 continue;
8305 }
8306 }
8307 out.push(self.matmul(w, x, m)?);
8308 }
8309 return Ok(out);
8310 }
8311 for w in ws {
8312 out.push(self.matmul(w, x, m)?);
8313 }
8314 Ok(out)
8315 }
8316
8317 pub fn matmul_group_multi(&self, ws: &[&crate::model::GpuTensor],
8324 xs: &[&CudaSlice<f32>], ms: &[usize])
8325 -> Result<Vec<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
8326 assert_eq!(xs.len(), ms.len());
8327 let in_f = ws[0].in_features();
8328 let total: usize = ms.iter().sum();
8329 let mut xcat = self.uninit(total * in_f)?;
8330 let mut off = 0usize;
8331 for (x, &m) in xs.iter().zip(ms) {
8332 self.copy_into(&mut xcat, off * in_f, x, m * in_f)?;
8333 off += m;
8334 }
8335 let ys = self.matmul_group(ws, &xcat, total)?;
8336 let mut out: Vec<Vec<CudaSlice<f32>>> = (0..xs.len()).map(|_| Vec::new()).collect();
8337 for (w, y) in ws.iter().zip(ys) {
8338 let out_f = w.out_features();
8339 let mut off = 0usize;
8340 for (s, &m) in ms.iter().enumerate() {
8341 let mut ys_s = self.uninit(m * out_f)?;
8342 let src = y.slice(off * out_f..(off + m) * out_f);
8343 self.gpu.stream().memcpy_dtod(&src, &mut ys_s)?;
8344 out[s].push(ys_s);
8345 off += m;
8346 }
8347 }
8348 Ok(out)
8349 }
8350
8351 pub fn gemm_supports(&self, w: &crate::model::GpuTensor) -> bool {
8361 use crate::model::GpuTensor;
8362 if !legacy_quant_gemm_allowed(
8363 cfg!(memra_portable_cuda),
8364 cfg!(memra_hopper_mma),
8365 std::env::var_os("MEMRA_NO_GEMM").is_some(),
8366 ) {
8367 return false;
8368 }
8369 match w {
8370 GpuTensor::Quant { qtype, .. } =>
8371 matches!(*qtype, QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q4_0)
8372 || (*qtype == QT_NVFP4 && w.in_features() % 64 == 0),
8373 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
8374 }
8375 }
8376
8377 pub fn qmatvec_gemm(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
8384 m: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8385 use crate::model::GpuTensor;
8386 let in_f = w.in_features();
8387 let out_f = w.out_features();
8388 let (bytes, qtype, row_bytes, scale, rp) = match w {
8389 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
8390 _ => unreachable!("gemm_supports guaranteed Quant"),
8391 };
8392 if cfg!(memra_hopper_mma) && qtype == QT_Q8_0 && out_f % 64 == 0 && wgmma_gemm_enabled() {
8398 if let GpuTensor::Quant { rp4: Some(m4), .. } = w {
8399 let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, aq, ad, m, in_f, out_f)?;
8400 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8401 return Ok(y);
8402 }
8403 }
8404 let name = match qtype {
8405 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8406 QT_Q4_0 => if rp { "qmatvec_gemm_q4_0_rp" } else { "qmatvec_gemm_q4_0" },
8407 QT_Q5_K => "qmatvec_gemm_q5_K",
8408 QT_Q6_K => "qmatvec_gemm_q6_K",
8409 QT_NVFP4 => if rp { "qmatvec_gemm_nvfp4_rp" } else { "qmatvec_gemm_nvfp4" },
8410 _ => unreachable!(),
8411 };
8412 let f = self.func(name);
8413 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let is_k1 = matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q4_0);
8418 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8420 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8421 let warps: u32 = if is_k1 { k1_tile.2 } else {
8422 match qtype { QT_NVFP4 => 8, _ => 4 }
8423 };
8424 let cfg = LaunchConfig {
8425 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8426 block_dim: (32, warps, 1),
8427 shared_mem_bytes: 0,
8428 };
8429 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8430 let __s_b = self.gpu.stream();
8431 let mut b = __s_b.launch_builder(&f);
8432 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8433 unsafe { b.launch(cfg)?; }
8434 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8435 Ok(y)
8436 }
8437
8438 pub fn qmatvec_gemm_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
8443 out_f: usize, qtype: i32, row_bytes: usize)
8444 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8445 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
8446 let name = match qtype {
8447 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8448 QT_Q4_0 => "qmatvec_gemm_q4_0",
8449 QT_Q5_K => "qmatvec_gemm_q5_K",
8450 QT_Q6_K => "qmatvec_gemm_q6_K", QT_NVFP4 => "qmatvec_gemm_nvfp4",
8451 QT_NVFP4_RP => "qmatvec_gemm_nvfp4_rp",
8452 _ => panic!("qmatvec_gemm_raw: qtype {qtype} has no GEMM kernel"),
8453 };
8454 let f = self.func(name);
8455 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let is_k1 = matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q4_0);
8459 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8461 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8462 let warps: u32 = if is_k1 { k1_tile.2 } else {
8463 match qtype { QT_NVFP4 | QT_NVFP4_RP => 8, _ => 4 }
8464 };
8465 let cfg = LaunchConfig {
8466 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8467 block_dim: (32, warps, 1), shared_mem_bytes: 0,
8468 };
8469 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8470 let __s_b = self.gpu.stream();
8471 let mut b = __s_b.launch_builder(&f);
8472 b.arg(bytes).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8473 unsafe { b.launch(cfg)?; }
8474 Ok(y)
8475 }
8476
8477 pub fn qmatvec_gemm_q8_0_wgmma_raw(&self, rp4: &CudaSlice<u8>, aq: &CudaSlice<i8>,
8484 ad: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize)
8485 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8486 assert!(out_f % 64 == 0 && in_f % 32 == 0, "wgmma GEMM needs out_f%64==0, in_f%32==0");
8487 let f = self.func("qmatvec_gemm_q8_0_wgmma");
8488 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
8490 grid_dim: ((out_f / 64) as u32, (m as u32).div_ceil(64), 1),
8491 block_dim: (128, 1, 1), shared_mem_bytes: 0,
8492 };
8493 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
8494 let __s_b = self.gpu.stream();
8495 let mut b = __s_b.launch_builder(&f);
8496 b.arg(rp4).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi);
8497 unsafe { b.launch(cfg)?; }
8498 Ok(y)
8499 }
8500
8501 pub fn scale_inplace(&self, y: &mut CudaSlice<f32>, s: f32, n: usize)
8503 -> Result<(), Box<dyn std::error::Error>> {
8504 let f = self.func("scale_f32");
8505 let cfg = LaunchConfig::for_num_elems(n as u32);
8506 let (sf, ni) = (s, n as i32);
8507 let __s_b = self.gpu.stream();
8508 let mut b = __s_b.launch_builder(&f);
8509 b.arg(y).arg(&sf).arg(&ni);
8510 unsafe { b.launch(cfg)?; }
8511 Ok(())
8512 }
8513
8514 pub fn bf16_to_f32(&self, data: &cudarc::driver::CudaView<'_, u8>, n: usize)
8519 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8520 let mut out = self.alloc_uninit::<f32>(n)?;
8521 let f = self.func("bf16_to_f32");
8522 let cfg = LaunchConfig::for_num_elems(n as u32);
8523 let ni = n as i32;
8524 let __s_b = self.gpu.stream();
8525 let mut b = __s_b.launch_builder(&f);
8526 b.arg(data).arg(&mut out).arg(&ni);
8527 unsafe { b.launch(cfg)?; }
8528 Ok(out)
8529 }
8530
8531 fn linear_bf16_chunked(&self, x: &CudaSlice<f32>, data: &CudaSlice<u8>, m: usize,
8538 in_f: usize, out_f: usize, exact: bool)
8539 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8540 const CHUNK_BYTES: usize = 256 << 20;
8541 let chunk_rows = (CHUNK_BYTES / (in_f * 4)).max(1).min(out_f);
8542 if chunk_rows >= out_f {
8543 let wf32 = self.bf16_to_f32(&data.slice(0..in_f * out_f * 2), in_f * out_f)?;
8544 return if exact { self.linear_decode_exact(x, &wf32, m, in_f, out_f) }
8545 else { self.linear(x, &wf32, m, in_f, out_f) };
8546 }
8547 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8548 let mut r0 = 0usize;
8549 while r0 < out_f {
8550 let rows = chunk_rows.min(out_f - r0);
8551 let wslice = data.slice(r0 * in_f * 2..(r0 + rows) * in_f * 2);
8552 let wf32 = self.bf16_to_f32(&wslice, in_f * rows)?;
8553 let yc = if exact { self.linear_decode_exact(x, &wf32, m, in_f, rows)? }
8554 else { self.linear(x, &wf32, m, in_f, rows)? };
8555 for mi in 0..m {
8557 let src = yc.slice(mi * rows..(mi + 1) * rows);
8558 let mut dst = y.slice_mut(mi * out_f + r0..mi * out_f + r0 + rows);
8559 self.gpu.stream().memcpy_dtod(&src, &mut dst)?;
8560 }
8561 r0 += rows;
8562 }
8563 Ok(y)
8564 }
8565
8566 pub fn linear_decode_exact(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize,
8573 in_f: usize, out_f: usize)
8574 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8575 if m_tokens == 1 { return self.linear(x, w, 1, in_f, out_f); }
8576 let xv = self.view(x, m_tokens * in_f);
8577 let mut y = self.alloc_uninit::<f32>(m_tokens * out_f)?;
8578 for t in 0..m_tokens {
8579 let row = xv.slice(t * in_f..(t + 1) * in_f);
8580 let mut xr = self.alloc_uninit::<f32>(in_f)?;
8581 self.copy_view_into(&mut xr, 0, &row, in_f)?;
8582 let yr = self.linear(&xr, w, 1, in_f, out_f)?;
8583 self.copy_into(&mut y, t * out_f, &yr, out_f)?;
8584 }
8585 Ok(y)
8586 }
8587
8588 pub fn linear(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize, in_f: usize, out_f: usize)
8589 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8590 use cudarc::cublaslt::{Matmul, MatmulConfig};
8591 let mut c = self.alloc_uninit::<f32>(m_tokens * out_f)?; let cfg = MatmulConfig {
8593 transa: true, transb: false, transc: false,
8594 m: out_f as u64, n: m_tokens as u64, k: in_f as u64,
8595 alpha: 1.0, lda: in_f as i64, ldb: in_f as i64, beta: 0.0, ldc: out_f as i64,
8596 stride_a: None, stride_b: None, stride_c: None, stride_bias: None, batch_size: None,
8597 };
8598 unsafe { self.gpu.blas.matmul(cfg, w, x, &mut c, None, None)?; }
8599 Ok(c)
8600 }
8601
8602 pub fn sdpa_naive(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8604 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8605 t: usize, t_kv: usize, scale: f32, causal: bool)
8606 -> Result<(), Box<dyn std::error::Error>> {
8607 let f = self.func("sdpa_naive_f32");
8608 let cfg = LaunchConfig {
8609 grid_dim: (n_head as u32, t as u32, 1),
8610 block_dim: (128, 1, 1),
8611 shared_mem_bytes: (t_kv * 4) as u32,
8612 };
8613 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8614 let __s_b = self.gpu.stream();
8615 let mut b = __s_b.launch_builder(&f);
8616 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8617 unsafe { b.launch(cfg)?; }
8618 Ok(())
8619 }
8620
8621 #[allow(clippy::too_many_arguments)]
8623 pub fn sdpa_naive_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8624 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8625 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8626 -> Result<(), Box<dyn std::error::Error>> {
8627 let f = self.func("sdpa_naive_w_f32");
8628 let cfg = LaunchConfig {
8629 grid_dim: (n_head as u32, t as u32, 1),
8630 block_dim: (128, 1, 1),
8631 shared_mem_bytes: (t_kv * 4) as u32,
8632 };
8633 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
8634 t as i32, t_kv as i32, causal as i32, window as i32);
8635 let __s_b = self.gpu.stream();
8636 let mut b = __s_b.launch_builder(&f);
8637 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8638 .arg(&scale).arg(&cz).arg(&wi);
8639 unsafe { b.launch(cfg)?; }
8640 Ok(())
8641 }
8642
8643 pub fn sdpa_naive_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<f32>,
8645 v: &cudarc::driver::CudaView<f32>, o: &mut CudaSlice<f32>,
8646 head_dim: usize, n_head: usize, n_head_kv: usize, t: usize, t_kv: usize,
8647 scale: f32, causal: bool) -> Result<(), Box<dyn std::error::Error>> {
8648 let f = self.func("sdpa_naive_f32");
8649 let cfg = LaunchConfig {
8650 grid_dim: (n_head as u32, t as u32, 1), block_dim: (128, 1, 1),
8651 shared_mem_bytes: (t_kv * 4) as u32,
8652 };
8653 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8654 let __s_b = self.gpu.stream();
8655 let mut b = __s_b.launch_builder(&f);
8656 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8657 unsafe { b.launch(cfg)?; }
8658 Ok(())
8659 }
8660
8661 #[allow(clippy::too_many_arguments)]
8669 pub fn fa_dequant_kv_view_f32(&self, k: &cudarc::driver::CudaView<u8>,
8670 v: &cudarc::driver::CudaView<u8>,
8671 kf: &mut CudaSlice<f32>, vf: &mut CudaSlice<f32>,
8672 kv_dim_k: usize, kv_dim_v: usize, t_kv: usize,
8673 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
8674 -> Result<(), Box<dyn std::error::Error>> {
8675 let f = if g { self.func_g("fa_dequant_kv_ws_f32") } else { self.func("fa_dequant_kv_ws_f32") };
8676 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
8677 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8678 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1),
8679 shared_mem_bytes: 0 };
8680 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
8681 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
8682 let __s_b = self.gpu.stream();
8683 let mut b = __s_b.launch_builder(&f);
8684 b.arg(k).arg(v).arg(&mut *kf).arg(&mut *vf).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
8685 unsafe { b.launch(cfg)?; }
8686 Ok(())
8687 }
8688
8689 #[allow(clippy::too_many_arguments)]
8690 pub fn sdpa_naive_quantized_view(
8691 &self,
8692 q: &CudaSlice<f32>,
8693 k: &cudarc::driver::CudaView<u8>,
8694 v: &cudarc::driver::CudaView<u8>,
8695 o: &mut CudaSlice<f32>,
8696 head_dim: usize,
8697 n_head: usize,
8698 n_head_kv: usize,
8699 t: usize,
8700 t_kv: usize,
8701 scale: f32,
8702 causal: bool,
8703 k_tok_bytes: usize,
8704 v_tok_bytes: usize,
8705 ) -> Result<(), Box<dyn std::error::Error>> {
8706 let kv_dim = n_head_kv * head_dim;
8707 let mut kf = self.uninit(t_kv * kv_dim)?;
8708 let mut vf = self.uninit(t_kv * kv_dim)?;
8709 let f = self.func("fa_dequant_kv_ws_f32");
8710 let total = (2 * t_kv * kv_dim) as u64;
8711 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8712 let cfg = LaunchConfig {
8713 grid_dim: (nblk.max(1), 1, 1),
8714 block_dim: (256, 1, 1),
8715 shared_mem_bytes: 0,
8716 };
8717 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8718 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8719 let __s_b = self.gpu.stream();
8720 let mut b = __s_b.launch_builder(&f);
8721 b.arg(k)
8722 .arg(v)
8723 .arg(&mut kf)
8724 .arg(&mut vf)
8725 .arg(&kv_dim_i)
8726 .arg(&kv_dim_i)
8727 .arg(&t_kv_i)
8728 .arg(&k_tok_bytes_i)
8729 .arg(&v_tok_bytes_i);
8730 unsafe { b.launch(cfg)? };
8731 self.sdpa_naive(
8732 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8733 )
8734 }
8735
8736 #[allow(clippy::too_many_arguments)]
8748 pub fn sdpa_naive_w_quantized_view(
8749 &self,
8750 q: &CudaSlice<f32>,
8751 k: &cudarc::driver::CudaView<u8>,
8752 v: &cudarc::driver::CudaView<u8>,
8753 o: &mut CudaSlice<f32>,
8754 head_dim: usize,
8755 n_head: usize,
8756 n_head_kv: usize,
8757 t: usize,
8758 t_kv: usize,
8759 scale: f32,
8760 causal: bool,
8761 window: usize,
8762 k_tok_bytes: usize,
8763 v_tok_bytes: usize,
8764 ) -> Result<(), Box<dyn std::error::Error>> {
8765 let kv_dim = n_head_kv * head_dim;
8766 let mut kf = self.uninit(t_kv * kv_dim)?;
8767 let mut vf = self.uninit(t_kv * kv_dim)?;
8768 let f = self.func("fa_dequant_kv_ws_f32");
8769 let total = (2 * t_kv * kv_dim) as u64;
8770 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8771 let cfg = LaunchConfig {
8772 grid_dim: (nblk.max(1), 1, 1),
8773 block_dim: (256, 1, 1),
8774 shared_mem_bytes: 0,
8775 };
8776 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8777 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8778 let __s_b = self.gpu.stream();
8779 let mut b = __s_b.launch_builder(&f);
8780 b.arg(k)
8781 .arg(v)
8782 .arg(&mut kf)
8783 .arg(&mut vf)
8784 .arg(&kv_dim_i)
8785 .arg(&kv_dim_i)
8786 .arg(&t_kv_i)
8787 .arg(&k_tok_bytes_i)
8788 .arg(&v_tok_bytes_i);
8789 unsafe { b.launch(cfg)? };
8790 self.sdpa_naive_w(
8791 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
8792 )
8793 }
8794
8795 pub fn fa_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8799 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8800 t: usize, t_kv: usize, scale: f32, causal: bool)
8801 -> Result<(), Box<dyn std::error::Error>> {
8802 if portable_mma_gated() {
8803 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
8804 t, t_kv, scale, causal);
8805 }
8806 let fa3_on = head_dim == 256 && causal && t == t_kv
8814 && match std::env::var("MEMRA_FA3").as_deref() {
8815 Ok("0") => false,
8816 Ok("1") => true,
8817 _ => cfg!(memra_hopper_mma),
8818 };
8819 if fa3_on {
8820 let n = t * n_head * head_dim;
8821 let nkv = t * n_head_kv * head_dim;
8822 let mut q16 = self.alloc_u8_uninit(n * 2)?;
8823 let mut k16 = self.alloc_u8_uninit(nkv * 2)?;
8824 let mut v16 = self.alloc_u8_uninit(nkv * 2)?;
8825 self.f32_to_bf16_into(q, &mut q16, n)?;
8826 self.f32_to_bf16_into(k, &mut k16, nkv)?;
8827 self.f32_to_bf16_into(v, &mut v16, nkv)?;
8828 let rc = {
8829 use cudarc::driver::{DevicePtr, DevicePtrMut};
8830 let stream = self.gpu.stream();
8831 let (qp, _g1) = q16.device_ptr(&stream);
8832 let (kp, _g2) = k16.device_ptr(&stream);
8833 let (vp, _g3) = v16.device_ptr(&stream);
8834 let (op, _g4) = o.device_ptr_mut(&stream);
8835 unsafe {
8836 memra_fa3_prefill(qp as *const core::ffi::c_void,
8837 kp as *const core::ffi::c_void,
8838 vp as *const core::ffi::c_void,
8839 op as *mut f32,
8840 t as i32, n_head as i32, n_head_kv as i32,
8841 head_dim as i32, scale,
8842 stream.cu_stream() as *mut core::ffi::c_void)
8843 }
8844 };
8845 if rc != 0 {
8846 return Err(format!("memra_fa3_prefill rc={rc}").into());
8847 }
8848 return Ok(());
8849 }
8850 static FA_P1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8855 let fa_p1 = *FA_P1.get_or_init(|| std::env::var("MEMRA_FA_P1").as_deref() == Ok("1"));
8856 if fa_p1 && head_dim == 256 && !std::env::var("MEMRA_FA_FLOOR").is_ok() {
8857 const BLOCK_Q: usize = 64; const BKX: usize = 32;
8858 let f = self.func("fa_prefill_bf16_p1");
8859 let shmem = (2 * (2 * BKX * head_dim + BLOCK_Q * BKX)
8860 + 4 * (BLOCK_Q * BKX + 2 * BLOCK_Q)) as u32;
8861 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8862 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8863 let cfg = LaunchConfig {
8864 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8865 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8866 };
8867 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32,
8868 n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8869 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8870 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8871 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8872 let __s_b = self.gpu.stream();
8873 let mut b = __s_b.launch_builder(&f);
8874 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti)
8875 .arg(&tkvi).arg(&scale).arg(&cz);
8876 unsafe { b.launch(cfg)?; }
8877 return Ok(());
8878 }
8879 const BK: usize = 32;
8885 let w2 = std::env::var("MEMRA_FA_PP_W2").as_deref() == Ok("1");
8888 let (block_q, warps, w2_sfx): (usize, u32, &str) =
8889 if w2 { (32, 2, "_w2") } else { (64, 4, "") };
8890 let hd_sfx = fa_hd_suffix(head_dim)?;
8894 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8895 let bf16kv = !floor && !w2
8900 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0");
8901 let (kb16, vb16) = if bf16kv {
8902 let n = t_kv * n_head_kv * head_dim;
8903 let mut kb = self.alloc_u8_uninit(n * 2)?;
8904 let mut vb = self.alloc_u8_uninit(n * 2)?;
8905 let fcv = self.func("f32_to_bf16_bulk");
8906 let ni = n as i64;
8907 let cfgc = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
8908 let __s_b = self.gpu.stream();
8909 let mut b = __s_b.launch_builder(&fcv);
8910 b.arg(k).arg(&mut kb).arg(&ni);
8911 unsafe { b.launch(cfgc)?; }
8912 let __s_b = self.gpu.stream();
8913 let mut b = __s_b.launch_builder(&fcv);
8914 b.arg(v).arg(&mut vb).arg(&ni);
8915 unsafe { b.launch(cfgc)?; }
8916 (Some(kb), Some(vb))
8917 } else {
8918 (None, None)
8919 };
8920 let f = self.func(&if bf16kv {
8921 format!("fa_prefill_bf16kv_pp{hd_sfx}")
8922 } else {
8923 format!("fa_prefill_f32{}{}{hd_sfx}",
8924 if floor { "" } else { "_pp" },
8925 if floor { "" } else { w2_sfx })
8926 });
8927 let kv_stages = if bf16kv { 2 } else { 1 };
8930 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
8931 + 4 * (block_q * BK + 2 * block_q)) as u32;
8932 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8933 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8934 let cfg = LaunchConfig {
8935 grid_dim: ((t as u32 + block_q as u32 - 1) / block_q as u32, n_head as u32, 1),
8936 block_dim: (32, warps, 1), shared_mem_bytes: shmem,
8937 };
8938 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8939 let __s_b = self.gpu.stream();
8940 let mut b = __s_b.launch_builder(&f);
8941 b.arg(q);
8942 match (&kb16, &vb16) {
8943 (Some(kb), Some(vb)) => { b.arg(kb).arg(vb); }
8944 _ => { b.arg(k).arg(v); }
8945 }
8946 b.arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8947 unsafe { b.launch(cfg)?; }
8948 Ok(())
8949 }
8950
8951 #[allow(clippy::too_many_arguments)]
8955 pub fn fa_prefill_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8956 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8957 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8958 -> Result<(), Box<dyn std::error::Error>> {
8959 if portable_mma_gated() {
8962 return self.sdpa_naive_w(q, k, v, o, head_dim, n_head, n_head_kv,
8963 t, t_kv, scale, causal, window);
8964 }
8965 static FAW_F32: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8969 let faw_f32 = *FAW_F32.get_or_init(|| {
8970 std::env::var("MEMRA_FAW_STAGE").as_deref() == Ok("f32")
8971 });
8972 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8973 self.fa_prefill_w_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8974 window, floor || faw_f32, floor)
8975 }
8976
8977 #[allow(clippy::too_many_arguments)]
8980 pub fn fa_prefill_w_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
8981 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8982 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
8983 window: usize, v_f16: bool)
8984 -> Result<(), Box<dyn std::error::Error>> {
8985 const BLOCK_Q: usize = 64; const BK: usize = 32;
8986 debug_assert_eq!(head_dim, 256);
8987 let hp = fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
8988 && (n_head / n_head_kv) % 2 == 0;
8989 debug_assert!(!v_f16 || hp, "f16 V emitted but the SWA hp arm is off");
8990 if hp {
8991 const BLOCK_QH: usize = 32;
8992 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
8995 let vh: &CudaSlice<u8> = if v_f16 { vb } else {
8996 let n = t_kv * n_head_kv * head_dim;
8997 if vguard.as_ref().map(|b| b.len() < n * 2).unwrap_or(true) {
8998 *vguard = Some(self.alloc_uninit::<u8>(n * 2)?);
8999 }
9000 self.bf16_to_f16_into(vb, n, vguard.as_mut().unwrap())?;
9001 vguard.as_ref().unwrap()
9002 };
9003 let f = self.func("fa_prefill_w_bf16_p1h2");
9004 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
9005 + 4 * (2 * BLOCK_QH)) as u32;
9006 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9007 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9008 let cfg = LaunchConfig {
9009 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
9010 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9011 };
9012 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9013 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9014 let __s_b = self.gpu.stream();
9015 let mut b = __s_b.launch_builder(&f);
9016 b.arg(qb).arg(kb).arg(vh).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9017 .arg(&scale).arg(&cz).arg(&wi);
9018 unsafe { b.launch(cfg)?; }
9019 return Ok(());
9020 }
9021 let f = self.func("fa_prefill_w_bf16_p1");
9022 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9023 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9024 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9025 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9026 let cfg = LaunchConfig {
9027 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9028 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9029 };
9030 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9031 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9032 let __s_b = self.gpu.stream();
9033 let mut b = __s_b.launch_builder(&f);
9034 b.arg(qb).arg(kb).arg(vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9035 .arg(&scale).arg(&cz).arg(&wi);
9036 unsafe { b.launch(cfg)?; }
9037 Ok(())
9038 }
9039
9040 #[allow(clippy::too_many_arguments)]
9042 pub fn fa_prefill_w_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9043 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9044 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9045 window: usize, f32_stage: bool, floor: bool)
9046 -> Result<(), Box<dyn std::error::Error>> {
9047 const BLOCK_Q: usize = 64; const BK: usize = 32;
9048 debug_assert_eq!(head_dim, 256, "fa_prefill_w is stamped hd256 only");
9049 static P1_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9053 let p1 = !floor && !f32_stage
9054 && *P1_ON.get_or_init(|| {
9055 std::env::var("MEMRA_FAW_P1").map(|v| v != "0").unwrap_or(true)
9056 });
9057 let hp = p1 && fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
9058 && (n_head / n_head_kv) % 2 == 0;
9059 if hp {
9060 const BLOCK_QH: usize = 32;
9061 let f = self.func("fa_prefill_w_bf16_p1h2");
9062 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
9063 + 4 * (2 * BLOCK_QH)) as u32;
9064 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9065 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9066 let cfg = LaunchConfig {
9067 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
9068 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9069 };
9070 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9071 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9072 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9073 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9074 let vh = self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?;
9075 let __s_b = self.gpu.stream();
9076 let mut b = __s_b.launch_builder(&f);
9077 b.arg(&qb).arg(&kb).arg(&vh).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9078 .arg(&scale).arg(&cz).arg(&wi);
9079 unsafe { b.launch(cfg)?; }
9080 return Ok(());
9081 }
9082 if p1 {
9083 let f = self.func("fa_prefill_w_bf16_p1");
9084 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9085 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9086 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9087 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9088 let cfg = LaunchConfig {
9089 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9090 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9091 };
9092 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9093 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9094 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9095 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9096 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9097 let __s_b = self.gpu.stream();
9098 let mut b = __s_b.launch_builder(&f);
9099 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9100 .arg(&scale).arg(&cz).arg(&wi);
9101 unsafe { b.launch(cfg)?; }
9102 return Ok(());
9103 }
9104 static G4_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9107 let g4 = !floor && !f32_stage && n_head_kv == 1 && n_head % 4 == 0
9108 && *G4_ON.get_or_init(|| {
9109 std::env::var("MEMRA_FAW_G4").map(|v| v != "0").unwrap_or(true)
9110 });
9111 if g4 {
9112 const SP_M: usize = 16;
9113 static O2_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9116 let o2 = *O2_ON.get_or_init(|| {
9117 std::env::var("MEMRA_FAW_O2").map(|v| v != "0").unwrap_or(true)
9118 });
9119 let f = self.func(if o2 { "fa_prefill_w_bf16_g4o2" } else { "fa_prefill_w_bf16_g4" });
9120 let shmem = if o2 {
9121 (2 * (4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M)) as u32
9122 } else {
9123 (2 * (2 * BK * head_dim + 4 * SP_M * head_dim + 4 * SP_M * BK)
9124 + 4 * (4 * SP_M)) as u32
9125 };
9126 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9127 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9128 let cfg = LaunchConfig {
9129 grid_dim: ((t as u32).div_ceil(SP_M as u32), (n_head / 4) as u32, 1),
9130 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9131 };
9132 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9133 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9134 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9135 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9136 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9137 let __s_b = self.gpu.stream();
9138 let mut b = __s_b.launch_builder(&f);
9139 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9140 .arg(&scale).arg(&cz).arg(&wi);
9141 unsafe { b.launch(cfg)?; }
9142 return Ok(());
9143 }
9144 let f = self.func(if floor { "fa_prefill_w_f32" }
9145 else if f32_stage { "fa_prefill_w_f32_pp" }
9146 else { "fa_prefill_w_bf16_pp" });
9147 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9148 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9149 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9150 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9151 let cfg = LaunchConfig {
9152 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9153 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9154 };
9155 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9156 t as i32, t_kv as i32, causal as i32, window as i32);
9157 if f32_stage {
9158 let __s_b = self.gpu.stream();
9159 let mut b = __s_b.launch_builder(&f);
9160 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9161 .arg(&scale).arg(&cz).arg(&wi);
9162 unsafe { b.launch(cfg)?; }
9163 } else {
9164 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9165 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9166 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9167 let __s_b = self.gpu.stream();
9168 let mut b = __s_b.launch_builder(&f);
9169 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9170 .arg(&scale).arg(&cz).arg(&wi);
9171 unsafe { b.launch(cfg)?; }
9172 }
9173 Ok(())
9174 }
9175
9176 #[allow(clippy::too_many_arguments)]
9180 pub fn fa_prefill_hd512(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9181 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9182 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool)
9183 -> Result<(), Box<dyn std::error::Error>> {
9184 if portable_mma_gated() {
9186 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
9187 t, t_kv, scale, causal);
9188 }
9189 static F32_STAGE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9195 let f32_stage = *F32_STAGE.get_or_init(|| {
9196 std::env::var("MEMRA_FA512_STAGE").as_deref() == Ok("f32")
9197 });
9198 static SP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9202 let sp = !f32_stage
9203 && *SP_ON.get_or_init(|| {
9204 std::env::var("MEMRA_FA512_SP").map(|v| v != "0").unwrap_or(true)
9205 });
9206 self.fa_prefill_hd512_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale,
9207 causal, f32_stage, sp, sp && fa_f16pv_on())
9208 }
9209
9210 #[allow(clippy::too_many_arguments)]
9212 pub fn fa_prefill_hd512_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
9213 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9214 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9215 v_f16: bool)
9216 -> Result<(), Box<dyn std::error::Error>> {
9217 debug_assert_eq!(head_dim, 512);
9218 const SP_M: usize = 16; const BKS: usize = 32;
9219 let f16pv = fa_f16pv_on();
9223 let nw = if f16pv { fa512_wide_warps() } else { 2 };
9224 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
9225 debug_assert!(!v_f16 || f16pv, "f16 V emitted without the door on");
9226 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
9227 let vref: &CudaSlice<u8> = if f16pv && !v_f16 {
9228 let n = t_kv * n_head_kv * head_dim;
9230 let need = n * 2;
9231 if vguard.as_ref().map(|b| b.len() < need).unwrap_or(true) {
9232 *vguard = Some(self.alloc_uninit::<u8>(need)?);
9233 }
9234 let dst = vguard.as_mut().unwrap();
9235 self.bf16_to_f16_into(vb, n, dst)?;
9236 vguard.as_ref().unwrap()
9237 } else { vb };
9238 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
9239 else { match (f16pv, nw) {
9240 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9241 (true, _) => "fa_prefill_bf16_hd512_sp16",
9242 _ => "fa_prefill_bf16_hd512_sp",
9243 } });
9244 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9245 let shmem = if hp {
9247 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9248 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9249 } else {
9250 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9251 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9252 };
9253 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9254 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9255 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9256 let cfg = LaunchConfig {
9257 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9258 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9259 };
9260 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9261 t as i32, t_kv as i32, causal as i32);
9262 let __s_b = self.gpu.stream();
9263 let mut b = __s_b.launch_builder(&f);
9264 b.arg(qb).arg(kb).arg(vref).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9265 .arg(&scale).arg(&cz);
9266 unsafe { b.launch(cfg)?; }
9267 Ok(())
9268 }
9269
9270 #[allow(clippy::too_many_arguments)]
9273 pub fn fa_prefill_hd512_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9274 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9275 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9276 f32_stage: bool, sp: bool, f16pv: bool)
9277 -> Result<(), Box<dyn std::error::Error>> {
9278 debug_assert_eq!(head_dim, 512, "fa_prefill_hd512 is hd512 only");
9279 if sp && !f32_stage {
9280 const SP_M: usize = 16; const BKS: usize = 32;
9284 let nw = if f16pv { fa512_wide_warps() } else { 2 };
9285 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
9286 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
9287 else { match (f16pv, nw) {
9288 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9289 (true, _) => "fa_prefill_bf16_hd512_sp16",
9290 _ => "fa_prefill_bf16_hd512_sp",
9291 } });
9292 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9293 let shmem = if hp {
9294 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9295 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9296 } else {
9297 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9298 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9299 };
9300 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9301 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9302 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9303 let cfg = LaunchConfig {
9304 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9305 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9306 };
9307 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9308 t as i32, t_kv as i32, causal as i32);
9309 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9310 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9311 let vb = if f16pv { self.f32_to_f16(v, t_kv * n_head_kv * head_dim)? }
9312 else { self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)? };
9313 let __s_b = self.gpu.stream();
9314 let mut b = __s_b.launch_builder(&f);
9315 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9316 .arg(&scale).arg(&cz);
9317 unsafe { b.launch(cfg)?; }
9318 return Ok(());
9319 }
9320 const BLOCK_Q: usize = 32; const BK: usize = 32; const HALF: usize = 256;
9321 let f = self.func(if f32_stage { "fa_prefill_f32_hd512" } else { "fa_prefill_bf16_hd512" });
9322 let shmem = (2 * (BLOCK_Q * head_dim + BK * head_dim + BK * HALF + BLOCK_Q * BK)
9324 + 4 * BLOCK_Q) as u32;
9325 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9326 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9327 let cfg = LaunchConfig {
9328 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 2),
9329 block_dim: (32, 2, 1), shared_mem_bytes: shmem,
9330 };
9331 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9332 t as i32, t_kv as i32, causal as i32);
9333 if f32_stage {
9334 let __s_b = self.gpu.stream();
9335 let mut b = __s_b.launch_builder(&f);
9336 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9337 .arg(&scale).arg(&cz);
9338 unsafe { b.launch(cfg)?; }
9339 } else {
9340 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9341 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9342 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9343 let __s_b = self.gpu.stream();
9344 let mut b = __s_b.launch_builder(&f);
9345 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9346 .arg(&scale).arg(&cz);
9347 unsafe { b.launch(cfg)?; }
9348 }
9349 Ok(())
9350 }
9351
9352 #[allow(clippy::too_many_arguments)]
9356 pub fn rope_neox2_bf16e(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
9357 qb: &mut CudaSlice<u8>, kb: &mut CudaSlice<u8>,
9358 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
9359 nh_q: usize, nh_k: usize, n_tokens: usize, base: f32,
9360 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
9361 -> Result<(), Box<dyn std::error::Error>> {
9362 let f = self.func("rope_neox2_bf16e_f32");
9363 let rows = ((nh_q + nh_k) * n_tokens) as u32;
9364 let cfg = LaunchConfig { grid_dim: (rows, 1, 1),
9365 block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
9366 let theta_scale = base.powf(-2.0 / n_dims as f32);
9367 let (hd, nd, nhq, nhk, nt) = (head_dim as i32, n_dims as i32, nh_q as i32,
9368 nh_k as i32, n_tokens as i32);
9369 let __s_b = self.gpu.stream();
9370 let mut b = __s_b.launch_builder(&f);
9371 match ff {
9372 Some(t) => { b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9373 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9374 .arg(&theta_scale).arg(&freq_scale).arg(t);
9375 unsafe { b.launch(cfg)?; } }
9376 None => { let null: u64 = 0;
9377 b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9378 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9379 .arg(&theta_scale).arg(&freq_scale).arg(&null);
9380 unsafe { b.launch(cfg)?; } }
9381 }
9382 Ok(())
9383 }
9384
9385 pub fn f32_to_bf16(&self, x: &CudaSlice<f32>, n: usize)
9388 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9389 assert!(n % 4 == 0, "f32_to_bf16 requires n % 4 == 0, got {n}");
9390 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9391 let f = self.func("f32_to_bf16_flat");
9392 let n_i = n as i64;
9393 let cfg = LaunchConfig {
9394 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9395 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9396 };
9397 let __s_b = self.gpu.stream();
9398 let mut b = __s_b.launch_builder(&f);
9399 b.arg(x).arg(&mut y).arg(&n_i);
9400 unsafe { b.launch(cfg)?; }
9401 Ok(y)
9402 }
9403
9404 pub fn f32_to_f16(&self, x: &CudaSlice<f32>, n: usize)
9405 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9406 assert!(n % 4 == 0, "f32_to_f16 requires n % 4 == 0, got {n}");
9407 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9408 let f = self.func("f32_to_f16_flat");
9409 let n_i = n as i64;
9410 let cfg = LaunchConfig {
9411 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9412 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9413 };
9414 let __s_b = self.gpu.stream();
9415 let mut b = __s_b.launch_builder(&f);
9416 b.arg(x).arg(&mut y).arg(&n_i);
9417 unsafe { b.launch(cfg)?; }
9418 Ok(y)
9419 }
9420
9421 pub fn bf16_to_f16(&self, xb: &CudaSlice<u8>, n: usize)
9423 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9424 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9425 self.bf16_to_f16_into(xb, n, &mut y)?;
9426 Ok(y)
9427 }
9428
9429 pub fn bf16_to_f16_into(&self, xb: &CudaSlice<u8>, n: usize, y: &mut CudaSlice<u8>)
9431 -> Result<(), Box<dyn std::error::Error>> {
9432 assert!(n % 2 == 0, "bf16_to_f16 requires n % 2 == 0, got {n}");
9433 assert!(y.len() >= n * 2);
9434 let f = self.func("bf16_to_f16_flat");
9435 let n2 = (n / 2) as i64;
9436 let cfg = LaunchConfig {
9437 grid_dim: (((n / 2) as u32).div_ceil(256), 1, 1),
9438 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9439 };
9440 let __s_b = self.gpu.stream();
9441 let mut b = __s_b.launch_builder(&f);
9442 b.arg(xb).arg(y).arg(&n2);
9443 unsafe { b.launch(cfg)?; }
9444 Ok(())
9445 }
9446
9447 #[allow(clippy::too_many_arguments)]
9452 pub fn fa_prefill_vl8(&self, seqs: &[FaSeqVl], head_dim: usize, n_head: usize,
9453 n_head_kv: usize, scale: f32)
9454 -> Result<(), Box<dyn std::error::Error>> {
9455 const BK: usize = 32;
9456 let b = seqs.len();
9457 assert!(b >= 1 && b <= 8);
9458 let mut packed = [FaSeqVl::default(); 8];
9459 packed[..b].copy_from_slice(seqs);
9460 let v = FaVl8(packed);
9461 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9462 let ept = (n_head_kv * head_dim) as i32;
9463 {
9464 let f = self.func("fa_mirror_vl");
9465 let max_n = (max_t as i64) * ept as i64;
9466 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
9467 for which in 0..2i32 {
9468 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9469 let __s_lb = self.gpu.stream();
9470 let mut lb = __s_lb.launch_builder(&f);
9471 lb.arg(&v).arg(&ept).arg(&which);
9472 unsafe { lb.launch(cfg)?; }
9473 }
9474 }
9475 let hd_sfx = fa_hd_suffix(head_dim)?;
9476 let f = self.func(&format!("fa_prefill_bf16kv_vl{hd_sfx}"));
9477 let block_q = 64usize;
9478 let kv_stages = 2usize;
9479 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
9480 + 4 * (block_q * BK + 2 * block_q)) as u32;
9481 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9482 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9483 let cfg = LaunchConfig {
9484 grid_dim: (max_t.div_ceil(block_q as u32), n_head as u32, b as u32),
9485 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9486 };
9487 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9488 let __s_lb = self.gpu.stream();
9489 let mut lb = __s_lb.launch_builder(&f);
9490 lb.arg(&v).arg(&hd).arg(&nh).arg(&nhkv).arg(&scale);
9491 unsafe { lb.launch(cfg)?; }
9492 Ok(())
9493 }
9494
9495 #[allow(clippy::too_many_arguments)]
9499 pub fn attn_pre_vl8(&self, seqs: &[AttnPreVl], wq: &CudaSlice<f32>, wk: &CudaSlice<f32>,
9500 head_dim: usize, rope_dims: usize, n_head: usize, n_head_kv: usize,
9501 eps: f32, freq_base: f32, freq_scale: f32,
9502 kv_dim_k: usize, kv_dim_v: usize,
9503 k_tok_bytes: usize, v_tok_bytes: usize)
9504 -> Result<(), Box<dyn std::error::Error>> {
9505 let b = seqs.len();
9506 assert!(b >= 1 && b <= 8);
9507 let mut packed = [AttnPreVl::default(); 8];
9508 packed[..b].copy_from_slice(seqs);
9509 let v = AttnPreVl8(packed);
9510 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9511 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9512 {
9513 let f = self.func("q_gate_split_vl");
9514 let n = max_t * (n_head * head_dim) as u32;
9515 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9516 let __s_lb = self.gpu.stream();
9517 let mut lb = __s_lb.launch_builder(&f);
9518 lb.arg(&v).arg(&hd).arg(&nh);
9519 unsafe { lb.launch(cfg)?; }
9520 }
9521 {
9522 let f = self.func("attn_rms_vl");
9523 let cfg = LaunchConfig { grid_dim: (max_t * n_head as u32, 2, b as u32), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
9524 let __s_lb = self.gpu.stream();
9525 let mut lb = __s_lb.launch_builder(&f);
9526 lb.arg(&v).arg(wq).arg(wk).arg(&hd).arg(&nh).arg(&nhkv).arg(&eps);
9527 unsafe { lb.launch(cfg)?; }
9528 }
9529 {
9530 let f = self.func("attn_rope_vl");
9531 let theta_scale = freq_base.powf(-2.0 / rope_dims as f32);
9532 let nd = rope_dims as i32;
9533 let cfg = LaunchConfig { grid_dim: (max_t * n_head as u32, 2, b as u32), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
9534 let __s_lb = self.gpu.stream();
9535 let mut lb = __s_lb.launch_builder(&f);
9536 lb.arg(&v).arg(&hd).arg(&nd).arg(&nh).arg(&nhkv).arg(&theta_scale).arg(&freq_scale);
9537 unsafe { lb.launch(cfg)?; }
9538 }
9539 {
9540 let f = self.func("append_kv_vl");
9541 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
9542 let cfg = LaunchConfig { grid_dim: (nblk, max_t, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
9543 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9544 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9545 let __s_lb = self.gpu.stream();
9546 let mut lb = __s_lb.launch_builder(&f);
9547 lb.arg(&v).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
9548 unsafe { lb.launch(cfg)?; }
9549 }
9550 Ok(())
9551 }
9552
9553 pub fn fa_prefill_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9558 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9559 head_dim: usize, n_head: usize, n_head_kv: usize,
9560 t: usize, t_kv: usize, scale: f32, causal: bool,
9561 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9562 -> Result<(), Box<dyn std::error::Error>> {
9563 if portable_mma_gated() {
9564 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9565 t, t_kv, scale, causal,
9566 k_tok_bytes, v_tok_bytes);
9567 }
9568 const BLOCK_Q: usize = 64; const BK: usize = 32;
9569 let name = format!("fa_prefill_q{}", fa_hd_suffix(head_dim)?);
9572 let f = if g { self.func_g(&name) } else { self.func(&name) };
9573 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9574 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9575 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9576 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9577 let cfg = LaunchConfig {
9578 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9579 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9580 };
9581 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
9582 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9583 let __s_b = self.gpu.stream();
9584 let mut b = __s_b.launch_builder(&f);
9585 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9586 .arg(&ktb).arg(&vtb);
9587 unsafe { b.launch(cfg)?; }
9588 Ok(())
9589 }
9590
9591 #[allow(clippy::too_many_arguments)]
9601 pub fn fa_prefill_view_ws(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9602 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9603 head_dim: usize, n_head: usize, n_head_kv: usize,
9604 t: usize, t_kv: usize, scale: f32, causal: bool,
9605 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9606 -> Result<(), Box<dyn std::error::Error>> {
9607 if portable_mma_gated() {
9608 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9609 t, t_kv, scale, causal,
9610 k_tok_bytes, v_tok_bytes);
9611 }
9612 const BLOCK_Q: usize = 64; const BK: usize = 32;
9613 let kv_dim_k = n_head_kv * head_dim;
9614 let kv_dim_v = n_head_kv * head_dim;
9615 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9617 let mut guard = self.prime_deqw_ws.lock().unwrap();
9619 let need_grow = match guard.as_ref() {
9620 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9621 None => true,
9622 };
9623 if need_grow {
9624 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9625 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9626 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9627 }
9628 let (kw, vw) = guard.as_mut().unwrap();
9629 {
9631 let f = if g { self.func_g("fa_dequant_kv_ws_bf16") } else { self.func("fa_dequant_kv_ws_bf16") };
9633 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9634 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9635 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9636 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9637 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9638 let __s_b = self.gpu.stream();
9639 let mut b = __s_b.launch_builder(&f);
9640 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9641 unsafe { b.launch(cfg)?; }
9642 }
9643 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9651 {
9652 let hd_sfx = fa_hd_suffix(head_dim)?;
9653 let f = self.func(&format!("fa_prefill_qw{}{hd_sfx}", if db { "_db" } else { "" }));
9654 let shmem = if db {
9655 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9657 } else {
9658 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9659 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9660 };
9661 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9662 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9663 let cfg = LaunchConfig {
9664 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9665 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9666 };
9667 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
9668 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9669 let __s_b = self.gpu.stream();
9670 let mut b = __s_b.launch_builder(&f);
9671 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9672 .arg(&kdk).arg(&kdv);
9673 unsafe { b.launch(cfg)?; }
9674 }
9675 Ok(())
9676 }
9677
9678 #[allow(clippy::too_many_arguments)]
9694 pub fn fa_prefill_view_ws_w_hd128(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9695 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9696 head_dim: usize, n_head: usize, n_head_kv: usize,
9697 t: usize, t_kv: usize, scale: f32, causal: bool,
9698 window: usize, k_tok_bytes: usize, v_tok_bytes: usize)
9699 -> Result<(), Box<dyn std::error::Error>> {
9700 assert_eq!(head_dim, 128, "fa_prefill_view_ws_w_hd128: only the hd128 twin is stamped");
9701 if portable_mma_gated() {
9702 return self.sdpa_naive_w_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9703 t, t_kv, scale, causal, window,
9704 k_tok_bytes, v_tok_bytes);
9705 }
9706 const BLOCK_Q: usize = 64; const BK: usize = 32;
9707 let kv_dim_k = n_head_kv * head_dim;
9708 let kv_dim_v = n_head_kv * head_dim;
9709 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9711 let mut guard = self.prime_deqw_ws.lock().unwrap();
9712 let need_grow = match guard.as_ref() {
9713 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9714 None => true,
9715 };
9716 if need_grow {
9717 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9718 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9719 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9720 }
9721 let (kw, vw) = guard.as_mut().unwrap();
9722 {
9725 let f = self.func("fa_dequant_kv_ws_bf16");
9726 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9727 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9728 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9729 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9730 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9731 let __s_b = self.gpu.stream();
9732 let mut b = __s_b.launch_builder(&f);
9733 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9734 unsafe { b.launch(cfg)?; }
9735 }
9736 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9738 {
9739 let f = self.func(if db { "fa_prefill_qw_db_w_hd128" } else { "fa_prefill_qw_w_hd128" });
9740 let shmem = if db {
9741 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9742 } else {
9743 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9744 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9745 };
9746 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9747 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9748 let cfg = LaunchConfig {
9749 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9750 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9751 };
9752 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
9753 let (kdk, kdv, wnd) = (kv_dim_k as i32, kv_dim_v as i32, window as i32);
9754 let __s_b = self.gpu.stream();
9755 let mut b = __s_b.launch_builder(&f);
9756 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9757 .arg(&kdk).arg(&kdv).arg(&wnd);
9758 unsafe { b.launch(cfg)?; }
9759 }
9760 Ok(())
9761 }
9762
9763 pub fn fa_decode(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9767 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9768 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9769 k_tok_bytes: usize, v_tok_bytes: usize)
9770 -> Result<(), Box<dyn std::error::Error>> {
9771 self.fa_decode_kvmod(q, k, v, o, head_dim, n_head, n_head_kv, t_kv, scale,
9772 k_tok_bytes, v_tok_bytes, false)
9773 }
9774
9775 #[allow(clippy::too_many_arguments)]
9779 #[allow(clippy::too_many_arguments)]
9783 #[allow(clippy::too_many_arguments)]
9784 fn fa_decode_scalar_unified(&self, q: &cudarc::driver::CudaView<f32>,
9785 k: &cudarc::driver::CudaView<u8>,
9786 v: &cudarc::driver::CudaView<u8>,
9787 o: &mut cudarc::driver::CudaViewMut<f32>,
9788 head_dim: usize, n_head: usize, n_head_kv: usize,
9789 t_kv_host: usize, t_kv_dev: Option<&CudaSlice<i32>>,
9790 scale: f32, n_splits: usize, split_keys: usize,
9791 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
9792 part_o: &mut CudaSlice<f32>, part_m: &mut CudaSlice<f32>,
9793 part_l: &mut CudaSlice<f32>,
9794 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
9795 -> Result<(), Box<dyn std::error::Error>> {
9796 let f = if g { self.func_g("fa_decode_f32") } else { self.fa_func("fa_decode_f32", head_dim) };
9797 let cfg = LaunchConfig { grid_dim: (n_head as u32, n_splits as u32, 1),
9798 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: (4 * (head_dim + 32)) as u32 };
9799 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
9800 let (ktb, vtb, tkvi, ski) = (k_tok_bytes as i64, v_tok_bytes as i64, t_kv_host as i32,
9801 split_keys as i32);
9802 let __s_b = self.gpu.stream();
9803 let mut b = __s_b.launch_builder(&f);
9804 match t_kv_dev {
9805 Some(d) => { b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9806 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(d).arg(&scale).arg(&nsp)
9807 .arg(&ski).arg(&ktb).arg(&vtb);
9808 unsafe { b.launch(cfg)?; } }
9809 None => { let null: u64 = 0;
9810 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9811 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&null).arg(&scale).arg(&nsp)
9812 .arg(&ski).arg(&ktb).arg(&vtb);
9813 unsafe { b.launch(cfg)?; } }
9814 }
9815 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1),
9816 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
9817 if let Some((oq, od)) = q8_out {
9818 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
9820 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
9821 let __s_b2 = self.gpu.stream();
9822 let mut b2 = __s_b2.launch_builder(&fc);
9823 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
9824 unsafe { b2.launch(cfg2)?; }
9825 return Ok(());
9826 }
9827 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
9828 let __s_b2 = self.gpu.stream();
9829 let mut b2 = __s_b2.launch_builder(&fc);
9830 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
9831 unsafe { b2.launch(cfg2)?; }
9832 Ok(())
9833 }
9834
9835 pub fn fa_decode_kvmod(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9836 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9837 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9838 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9839 -> Result<(), Box<dyn std::error::Error>> {
9840 let q_view = q.as_view();
9841 let mut o_view = o.as_view_mut();
9842 self.fa_decode_kvmod_view(&q_view, k, v, &mut o_view, head_dim, n_head, n_head_kv,
9843 t_kv, scale, k_tok_bytes, v_tok_bytes, g)
9844 }
9845
9846 #[allow(clippy::too_many_arguments)]
9851 pub fn fa_decode_kvmod_view(&self, q: &cudarc::driver::CudaView<f32>,
9852 k: &cudarc::driver::CudaView<u8>, v: &cudarc::driver::CudaView<u8>,
9853 o: &mut cudarc::driver::CudaViewMut<f32>,
9854 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9855 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9856 -> Result<(), Box<dyn std::error::Error>> {
9857 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
9878 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
9882 let sp = fa_split_keys(t_kv, n_head_kv);
9883 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
9884 let o_len = n_head * n_splits * head_dim;
9885 let ml_len = n_head * n_splits;
9886 let mut part_guard = self.fa_part_pool.lock().unwrap();
9887 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
9888 let old = part_guard.take();
9899 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
9900 if let Some(old) = old {
9901 self.fa_part_retired.lock().unwrap().push(old);
9902 }
9903 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
9904 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
9905 }
9906 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
9907 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
9908 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
9909 }
9910 let pg = part_guard.as_mut().unwrap();
9911 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
9912 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
9913 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
9914 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
9915 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
9916 let (hd, nh, nhkv, tkvi, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, t_kv as i32, n_splits as i32);
9917 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9918 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
9922 let fa512_min = fa512_min_tkv();
9927 let deep = fa_vec && head_dim == 256 && fa_v4_at(t_kv) && !g
9930 && fa_deep_at(t_kv) && !matches!(fa_v4_mode(), "noB3" | "stage");
9931 let (f, cfg) = if fa_vec && head_dim == 512 && t_kv >= fa512_min {
9932 let gqa = (n_head / n_head_kv).max(1) as u32;
9935 let fv = self.fa_func("fa_decode_vec_q_dpl16", head_dim);
9936 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9937 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
9938 } else if fa_vec && head_dim <= 256 {
9939 let gqa = (n_head / n_head_kv).max(1) as u32;
9940 static SMEM_TKV: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9951 let smem_tkv = *SMEM_TKV.get_or_init(|| {
9952 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
9953 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
9954 });
9955 if fa_v4_at(t_kv) && head_dim == 256 {
9956 let v4name = match fa_v4_mode() {
9960 "noB3" => "fa_decode_vec_q_v4_noB3", "stage" => "fa_decode_vec_q_v4_stage", _ if deep => "fa_decode_vec_q_v4_deep",
9963 _ => "fa_decode_vec_q_v4",
9964 };
9965 let fv = if g { self.func_g(v4name) } else { self.func(v4name) };
9966 let shmem = (if deep { 12160 } else { 11520 }
9969 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
9970 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9971 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9972 (fv,
9973 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9974 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9975 } else if fa_v3_active(head_dim) {
9976 let fv = if g { self.func_g("fa_decode_vec_q_v3") } else { self.func("fa_decode_vec_q_v3") };
9979 let shmem = (32 * head_dim * 2) as u32; (fv,
9981 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9982 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9983 } else if fa_v2_on() {
9984 let fv = if g { self.func_g("fa_decode_vec_q_v2") } else { self.func("fa_decode_vec_q_v2") };
9988 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
9990 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9991 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9992 } else if smem_tkv > 0 && t_kv >= smem_tkv && !g
9993 && !(head_dim == 512 && Self::gkv_on()) {
9994 let fv = if g { self.func_g("fa_decode_vec_q_smem") } else { self.func("fa_decode_vec_q_smem") };
9998 let shmem = (2 * 32 * head_dim * 2) as u32; use cudarc::driver::sys::CUfunction_attribute_enum as A;
10000 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
10001 (fv,
10002 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10003 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10004 } else {
10005 let fv = if g { self.func_g("fa_decode_vec_q") } else { self.func("fa_decode_vec_q") };
10008 (fv,
10009 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10010 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10011 }
10012 } else {
10013 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
10016 t_kv, None, scale, n_splits,
10017 if fa_vec { sp } else { 256 },
10018 k_tok_bytes, v_tok_bytes, g,
10019 part_o, part_m, part_l, None);
10020 };
10021 let __s_b = self.gpu.stream();
10022 let mut b = __s_b.launch_builder(&f);
10023 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10024 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&scale).arg(&nsp).arg(&ktb).arg(&vtb);
10025 unsafe { b.launch(cfg)?; }
10026 let (fc, cfg2) = (if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) },
10029 LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 });
10030 let __s_b2 = self.gpu.stream();
10031 let mut b2 = __s_b2.launch_builder(&fc);
10032 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
10033 unsafe { b2.launch(cfg2)?; }
10034 Ok(())
10035 }
10036
10037 #[allow(clippy::too_many_arguments)]
10048 pub fn fa_decode_batch_seqs_v4(&self, q: &CudaSlice<f32>,
10049 kv_ptrs: &cudarc::driver::CudaView<u64>,
10050 pos_seq: &CudaSlice<i32>, o: &mut CudaSlice<f32>,
10051 head_dim: usize, n_head: usize, n_head_kv: usize,
10052 b_n: usize, t_kv_max: usize, scale: f32,
10053 split_keys: usize, k_tok_bytes: usize, v_tok_bytes: usize)
10054 -> Result<(), Box<dyn std::error::Error>> {
10055 debug_assert!(head_dim == 256, "seqs twin is v4-stamped (hd256 only)");
10056 let n_splits_max = (t_kv_max + split_keys - 1) / split_keys;
10057 let o_len = b_n * n_head * n_splits_max * head_dim;
10058 let ml_len = b_n * n_head * n_splits_max;
10059 let mut part_guard = self.fa_part_pool.lock().unwrap();
10060 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10061 let old = part_guard.take();
10072 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10073 if let Some(old) = old {
10074 self.fa_part_retired.lock().unwrap().push(old);
10075 }
10076 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10077 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10078 }
10079 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10080 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10081 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10082 }
10083 let pg = part_guard.as_mut().unwrap();
10084 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10085 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10086 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10087 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10088 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10089 let (nspm, spk) = (n_splits_max as i32, split_keys as i32);
10090 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10091 let gqa = (n_head / n_head_kv).max(1) as u32;
10092 let f = self.func("fa_decode_vec_q_seqs_v4");
10093 let shmem = (11520 + 32 * head_dim * 2) as u32;
10095 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10096 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
10097 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, b_n as u32),
10098 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10099 {
10100 let __s_b = self.gpu.stream();
10101 let mut b = __s_b.launch_builder(&f);
10102 b.arg(q).arg(kv_ptrs).arg(pos_seq).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10103 .arg(&hd).arg(&nh).arg(&nhkv).arg(&scale).arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
10104 unsafe { b.launch(cfg)?; }
10105 }
10106 let fc = self.func("fa_decode_combine_seqs");
10107 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, b_n as u32, 1),
10108 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10109 let __s_b2 = self.gpu.stream();
10110 let mut b2 = __s_b2.launch_builder(&fc);
10111 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10112 .arg(pos_seq).arg(&nspm).arg(&spk);
10113 unsafe { b2.launch(cfg2)?; }
10114 Ok(())
10115 }
10116
10117 #[allow(clippy::too_many_arguments)]
10124 pub fn append_kv_quantized_seqs(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
10125 kv_ptrs: &cudarc::driver::CudaView<u64>,
10126 pos_seq: &CudaSlice<i32>, b_n: usize,
10127 kv_dim_k: usize, kv_dim_v: usize,
10128 k_tok_bytes: usize, v_tok_bytes: usize)
10129 -> Result<(), Box<dyn std::error::Error>> {
10130 let f = self.func("append_quantize_kv_q8_0_q5_1_seqs");
10131 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
10132 let cfg = LaunchConfig { grid_dim: (nblk, b_n as u32, 1),
10133 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
10134 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
10135 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10136 let __s_b = self.gpu.stream();
10137 let mut b = __s_b.launch_builder(&f);
10138 b.arg(k_rows).arg(v_rows).arg(kv_ptrs).arg(pos_seq)
10139 .arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
10140 unsafe { b.launch(cfg)?; }
10141 Ok(())
10142 }
10143
10144 pub fn fa_rows_eligible(&self, base_len: usize, head_dim: usize) -> bool {
10150 std::env::var("MEMRA_NO_FA_VEC").is_err()
10151 && std::env::var("MEMRA_FA_ROWS_OFF").is_err()
10152 && base_len + 1 >= fa_vec_min_tkv()
10153 && head_dim <= 256 && head_dim % 32 == 0
10154 }
10155
10156 #[allow(clippy::too_many_arguments)]
10165 pub fn fa_decode_rows(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10166 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10167 head_dim: usize, n_head: usize, n_head_kv: usize,
10168 base_len: usize, t: usize, scale: f32,
10169 k_tok_bytes: usize, v_tok_bytes: usize,
10170 base_dev: Option<(&CudaSlice<i32>, i32)>,
10174 kv_shared: bool,
10177 g: bool,
10181 mut q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10184 -> Result<(), Box<dyn std::error::Error>> {
10185 debug_assert!(base_len + 1 >= fa_vec_min_tkv() && head_dim <= 512 && head_dim % 32 == 0);
10186 let t_kv_max = base_len + t; let mut sp = fa_split_keys(t_kv_max, n_head_kv); if head_dim == 512 {
10193 static SP512: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10194 let v = *SP512.get_or_init(|| std::env::var("MEMRA_FA_SP512").ok()
10197 .and_then(|x| x.parse().ok()).unwrap_or(0));
10198 sp = if v >= 8 { v } else { FA_SP512_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) };
10199 }
10200 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10201 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10202 let gqa = (n_head / n_head_kv).max(1) as u32;
10203 let mut groups: Vec<(usize, usize, usize)> = Vec::new(); if head_dim == 512 || fa_split_keys(base_len + 1, n_head_kv) == sp {
10214 groups.push((0, t, sp));
10215 } else {
10216 let mut r0 = 0usize;
10217 while r0 < t {
10218 let sp_g = fa_split_keys(base_len + r0 + 1, n_head_kv);
10219 let mut r1 = r0 + 1;
10220 while r1 < t && fa_split_keys(base_len + r1 + 1, n_head_kv) == sp_g { r1 += 1; }
10221 groups.push((r0, r1 - r0, sp_g));
10222 r0 = r1;
10223 }
10224 }
10225 static SMEM_TKV_R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10229 let smem_tkv = *SMEM_TKV_R.get_or_init(|| {
10230 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
10231 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
10232 });
10233 let v4 = fa_v4_at(base_len + t) && head_dim == 256;
10234 let v3 = fa_v3_active(head_dim);
10235 let smem_rows = head_dim <= 256 && !v3 && !fa_v2_on() && smem_tkv > 0 && t_kv_max >= smem_tkv;
10236 let _ = kv_shared;
10241 let i2 = head_dim == 512 && std::env::var("MEMRA_FA_I2").as_deref() != Ok("0");
10244 static TB512: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10258 let tb512 = head_dim == 512 && sp <= 32 && n_head / n_head_kv.max(1) <= 16
10260 && *TB512.get_or_init(|| std::env::var("MEMRA_FA_TB512").as_deref() != Ok("0"));
10261 let fname = if tb512 { "fa_decode_vec_q_rows_v4_512_tb" }
10262 else if i2 { "fa_decode_vec_q_rows_dpl16_i2" }
10263 else if head_dim == 512 { "fa_decode_vec_q_rows_dpl16" } else if v4 { "fa_decode_vec_q_rows_v4" }
10265 else if v3 { "fa_decode_vec_q_rows_v3" }
10266 else if fa_v2_on() { "fa_decode_vec_q_rows_v2" }
10267 else if smem_rows { "fa_decode_vec_q_rows_smem" }
10268 else { "fa_decode_vec_q_rows" };
10269 let f = if head_dim == 512 { self.fa_func(fname, head_dim) }
10270 else if g {
10271 self.func_g(if smem_rows { "fa_decode_vec_q_rows" } else { fname })
10279 }
10280 else { self.func(fname) };
10281 let shmem = if tb512 {
10282 let gk = Self::gkv_on();
10284 let sh = (8192 + 1024 + 32 * 512 + 32 * 64
10285 + 32 * head_dim * if gk { 1 } else { 2 }) as u32;
10286 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10287 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10288 sh
10289 } else if v4 || v3 || smem_rows || fa_v2_on() {
10290 let sh = (if v4 { 11520 + 32 * head_dim * if g { 1 } else { 2 } }
10292 else if v3 { 32 * head_dim * 2 } else { 2 * 32 * head_dim * 2 }) as u32;
10293 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10294 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10295 sh
10296 } else { 0 };
10297 for &(r0, t_g, sp_g) in &groups {
10301 let n_splits_g = (base_len + r0 + t_g).div_ceil(sp_g);
10302 let (nspm, spk) = (n_splits_g as i32, sp_g as i32);
10303 let base_i = (base_len + r0) as i32;
10304 let o_len = t_g * n_head * n_splits_g * head_dim;
10305 let ml_len = t_g * n_head * n_splits_g;
10306 let mut part_guard = self.fa_part_pool.lock().unwrap();
10307 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10308 let old = part_guard.take();
10319 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10320 if let Some(old) = old {
10321 self.fa_part_retired.lock().unwrap().push(old);
10322 }
10323 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10324 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10325 }
10326 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10327 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10328 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10329 }
10330 let pg = part_guard.as_mut().unwrap();
10331 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10332 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10333 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10334 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10335 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
10336 let qv = self.view(q, t * n_head * head_dim);
10337 let q_g = qv.slice(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10338 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_g as u32, t_g as u32),
10339 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10340 {
10341 let __s_b = self.gpu.stream();
10342 let mut b = __s_b.launch_builder(&f);
10343 if tb512 {
10344 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10346 let plus_g = plus + r0 as i32;
10347 let nr = t_g as i32;
10348 if Self::pdl_on() && Self::pdl_wb_on() {
10349 use cudarc::driver::{DevicePtr, DevicePtrMut};
10351 let s = &self.gpu.stream();
10352 let (pq, _b0) = q_g.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10353 let (pv, _b2) = v.device_ptr(s);
10354 let (po, _b3) = part_o.device_ptr_mut(s);
10355 let (pm, _b4) = part_m.device_ptr_mut(s);
10356 let (pl, _b5) = part_l.device_ptr_mut(s);
10357 let (pb, _b6) = bd.device_ptr(s);
10358 let mut ps = [
10359 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10360 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10361 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10362 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10363 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10364 &plus_g as *const _ as *mut _, &scale as *const _ as *mut _,
10365 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10366 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10367 &nr as *const _ as *mut _,
10368 ];
10369 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10370 "fa_decode_vec_q_rows_v4_512_tb",
10371 (n_head_kv as u32, n_splits_g as u32, 1), (32, gqa, 1),
10372 shmem, &mut ps)?; }
10373 } else {
10374 let cfg_tb = LaunchConfig {
10375 grid_dim: (n_head_kv as u32, n_splits_g as u32, 1),
10376 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10377 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10378 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10379 .arg(&ktb).arg(&vtb).arg(&nr);
10380 unsafe { b.launch(cfg_tb)?; }
10381 }
10382 } else if head_dim == 512 {
10383 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10384 let plus_g = plus + r0 as i32;
10385 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10386 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10387 .arg(&ktb).arg(&vtb);
10388 unsafe { b.launch(cfg)?; }
10389 } else {
10390 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10391 .arg(&hd).arg(&nh).arg(&nhkv).arg(&base_i).arg(&scale).arg(&nspm).arg(&spk)
10392 .arg(&ktb).arg(&vtb);
10393 unsafe { b.launch(cfg)?; }
10394 }
10395 }
10396 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t_g as u32, 1),
10397 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10398 let mut o_g = o.slice_mut(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10399 if head_dim == 512 {
10400 let (bd, plus) = base_dev.unwrap();
10403 let plus_g = plus + r0 as i32;
10404 if let Some((oq, od)) = q8_out.as_mut() {
10405 debug_assert!(t == 1, "rows q8 emit is a t=1 decode arm");
10407 if Self::pdl_on() && Self::pdl_wb_on() {
10408 use cudarc::driver::{DevicePtr, DevicePtrMut};
10410 let s = &self.gpu.stream();
10411 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10412 let (pl, _g2) = part_l.device_ptr(s);
10413 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10414 let (pb, _g5) = bd.device_ptr(s);
10415 let mut ps = [
10416 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10417 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10418 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10419 &nh as *const _ as *mut _, &pb as *const _ as *mut _,
10420 &plus_g as *const _ as *mut _, &nspm as *const _ as *mut _,
10421 &spk as *const _ as *mut _,
10422 ];
10423 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10424 "fa_decode_combine_rows_dc_q8_1",
10425 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10426 continue;
10427 }
10428 let fc = self.fa_func("fa_decode_combine_rows_dc_q8_1", head_dim);
10429 let __s_b2 = self.gpu.stream();
10430 let mut b2 = __s_b2.launch_builder(&fc);
10431 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut **oq).arg(&mut **od)
10432 .arg(&hd).arg(&nh).arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10433 unsafe { b2.launch(cfg2)?; }
10434 continue;
10435 }
10436 let fc = self.fa_func("fa_decode_combine_rows_dc", head_dim);
10437 let __s_b2 = self.gpu.stream();
10438 let mut b2 = __s_b2.launch_builder(&fc);
10439 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10440 .arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10441 unsafe { b2.launch(cfg2)?; }
10442 } else {
10443 assert!(q8_out.is_none(), "rows q8 emit requires the hd512 dc combine");
10446 let fc = self.func("fa_decode_combine_rows");
10447 let __s_b2 = self.gpu.stream();
10448 let mut b2 = __s_b2.launch_builder(&fc);
10449 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10450 .arg(&base_i).arg(&nspm).arg(&spk);
10451 unsafe { b2.launch(cfg2)?; }
10452 }
10453 }
10454 Ok(())
10455 }
10456
10457 #[allow(clippy::too_many_arguments)]
10461 pub fn fa_decode_rows_w(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10462 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10463 head_dim: usize, n_head: usize, n_head_kv: usize,
10464 base_dev: &CudaSlice<i32>, base_plus: i32, t: usize, scale: f32,
10465 window: usize, k_tok_bytes: usize, v_tok_bytes: usize,
10466 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10467 -> Result<(), Box<dyn std::error::Error>> {
10468 debug_assert!(head_dim == 256);
10473 let sp = {
10481 static SPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10482 let v = *SPW.get_or_init(|| std::env::var("MEMRA_FA_SPW").ok()
10483 .and_then(|x| x.parse().ok()).unwrap_or(0));
10484 if v >= 8 { v } else { FA_SPW_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) }
10485 };
10486 let n_splits_max = (window + sp - 1) / sp;
10487 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10488 let (nspm, spk, wini) = (n_splits_max as i32, sp as i32, window as i32);
10489 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10490 let gqa = (n_head / n_head_kv).max(1) as u32;
10491 let o_len = t * n_head * n_splits_max * head_dim;
10492 let ml_len = t * n_head * n_splits_max;
10493 let mut part_guard = self.fa_part_pool.lock().unwrap();
10494 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10495 let old = part_guard.take();
10506 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10507 if let Some(old) = old {
10508 self.fa_part_retired.lock().unwrap().push(old);
10509 }
10510 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10511 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10512 }
10513 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10514 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10515 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10516 }
10517 let pg = part_guard.as_mut().unwrap();
10518 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10519 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10520 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10521 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10522 static SMEM_TKV_W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10528 let smem_tkv = *SMEM_TKV_W.get_or_init(|| {
10529 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
10530 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
10531 });
10532 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10538 let wg = Self::wkv_on();
10543 let sp2 = gqa <= 4 && fa_v4_at(window)
10546 && std::env::var("MEMRA_FA_SPW2").as_deref() != Ok("0");
10547 if sp2 {
10548 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10549 if Self::pdl_on() && Self::pdl_wb_on() {
10550 use cudarc::driver::{DevicePtr, DevicePtrMut};
10552 let s = &self.gpu.stream();
10553 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10554 let (pv, _b2) = v.device_ptr(s);
10555 let (po, _b3) = part_o.device_ptr_mut(s);
10556 let (pm, _b4) = part_m.device_ptr_mut(s);
10557 let (pl, _b5) = part_l.device_ptr_mut(s);
10558 let (pb, _b6) = base_dev.device_ptr(s);
10559 let mut ps = [
10560 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10561 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10562 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10563 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10564 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10565 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10566 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10567 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10568 &wini as *const _ as *mut _,
10569 ];
10570 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w_sp",
10571 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa + 1, 1),
10572 sh, &mut ps)?; }
10573 } else {
10574 let f = if wg { self.func_g("fa_decode_vec_q_rows_v4_w_sp") }
10575 else { self.func("fa_decode_vec_q_rows_v4_w_sp") };
10576 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10577 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10578 block_dim: (32, gqa + 1, 1), shared_mem_bytes: sh };
10579 let __s_b = self.gpu.stream();
10580 let mut b = __s_b.launch_builder(&f);
10581 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10582 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10583 .arg(&ktb).arg(&vtb).arg(&wini);
10584 unsafe { b.launch(cfg)?; }
10585 }
10586 } else {
10587 if fa_v4_at(window) && Self::pdl_on() && Self::pdl_wb_on() {
10588 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10590 use cudarc::driver::{DevicePtr, DevicePtrMut};
10591 let s = &self.gpu.stream();
10592 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10593 let (pv, _b2) = v.device_ptr(s);
10594 let (po, _b3) = part_o.device_ptr_mut(s);
10595 let (pm, _b4) = part_m.device_ptr_mut(s);
10596 let (pl, _b5) = part_l.device_ptr_mut(s);
10597 let (pb, _b6) = base_dev.device_ptr(s);
10598 let mut ps = [
10599 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10600 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10601 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10602 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10603 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10604 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10605 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10606 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10607 &wini as *const _ as *mut _,
10608 ];
10609 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w",
10610 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa, 1),
10611 sh, &mut ps)?; }
10612 } else {
10613 let pick = |name: &str| if wg { self.func_g(name) } else { self.func(name) };
10614 let (f, sh) = if fa_v4_at(window) {
10615 let f = pick("fa_decode_vec_q_rows_v4_w");
10616 (f, (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32)
10617 } else if smem_tkv > 0 && window >= smem_tkv {
10618 (pick("fa_decode_vec_q_rows_smem_w"), (2 * 32 * head_dim * 2) as u32)
10621 } else {
10622 (pick("fa_decode_vec_q_rows_reg_w"), 0u32)
10623 };
10624 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10625 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10626 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10627 let __s_b = self.gpu.stream();
10628 let mut b = __s_b.launch_builder(&f);
10629 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10630 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10631 .arg(&ktb).arg(&vtb).arg(&wini);
10632 unsafe { b.launch(cfg)?; }
10633 }
10634 }
10635 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10636 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10637 if let Some((oq, od)) = q8_out {
10638 if Self::pdl_on() && Self::pdl_wb_on() {
10641 use cudarc::driver::{DevicePtr, DevicePtrMut};
10643 let s = &self.gpu.stream();
10644 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10645 let (pl, _g2) = part_l.device_ptr(s);
10646 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10647 let mut ps = [
10648 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10649 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10650 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10651 &nh as *const _ as *mut _, &nspm as *const _ as *mut _,
10652 &spk as *const _ as *mut _, &wini as *const _ as *mut _,
10653 ];
10654 unsafe { self.launch_pdl_flash(wg, "fa_decode_combine_rows_w_q8_1",
10655 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10656 return Ok(());
10657 }
10658 let fc = if wg { self.func_g("fa_decode_combine_rows_w_q8_1") }
10659 else { self.func("fa_decode_combine_rows_w_q8_1") };
10660 let __s_b2 = self.gpu.stream();
10661 let mut b2 = __s_b2.launch_builder(&fc);
10662 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh)
10663 .arg(&nspm).arg(&spk).arg(&wini);
10664 unsafe { b2.launch(cfg2)?; }
10665 return Ok(());
10666 }
10667 let fc = if wg { self.func_g("fa_decode_combine_rows_w") }
10668 else { self.func("fa_decode_combine_rows_w") };
10669 let __s_b2 = self.gpu.stream();
10670 let mut b2 = __s_b2.launch_builder(&fc);
10671 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10672 .arg(&nspm).arg(&spk).arg(&wini);
10673 unsafe { b2.launch(cfg2)?; }
10674 Ok(())
10675 }
10676
10677 #[allow(clippy::too_many_arguments)]
10683 pub fn fa_decode_rows_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10684 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10685 head_dim: usize, n_head: usize, n_head_kv: usize,
10686 base_dev: &CudaSlice<i32>, t_kv_upper: usize, t: usize, scale: f32,
10687 k_tok_bytes: usize, v_tok_bytes: usize, base_plus: i32, g: bool)
10688 -> Result<(), Box<dyn std::error::Error>> {
10689 let v4 = head_dim == 256 && fa_v4_at(t_kv_upper);
10690 assert!(v4 || fa_v3_active(head_dim), "stream fa rows requires the v3 or v4 lane");
10691 assert!(v4 || base_plus == 0, "v3_dc kernel takes no plus arg");
10692 if v4 {
10693 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10694 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10695 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10696 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10697 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10698 let gqa = (n_head / n_head_kv).max(1) as u32;
10699 let o_len = t * n_head * n_splits_max * head_dim;
10700 let ml_len = t * n_head * n_splits_max;
10701 let mut part_guard = self.fa_part_pool.lock().unwrap();
10702 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10703 let old = part_guard.take();
10714 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10715 if let Some(old) = old {
10716 self.fa_part_retired.lock().unwrap().push(old);
10717 }
10718 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10719 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10720 }
10721 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10722 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10723 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10724 }
10725 let pg = part_guard.as_mut().unwrap();
10726 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10727 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10728 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10729 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10730 let f = if g { self.func_g("fa_decode_vec_q_rows_v4_dc") }
10731 else { self.func("fa_decode_vec_q_rows_v4_dc") };
10732 let sh = (11520 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10733 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10734 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10735 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10736 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10737 let __s_b = self.gpu.stream();
10738 let mut b = __s_b.launch_builder(&f);
10739 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10740 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale)
10741 .arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
10742 unsafe { b.launch(cfg)?; }
10743 let fc = self.func("fa_decode_combine_rows_dc");
10744 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10745 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10746 let __s_b2 = self.gpu.stream();
10747 let mut b2 = __s_b2.launch_builder(&fc);
10748 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10749 .arg(base_dev).arg(&base_plus).arg(&nspm).arg(&spk);
10750 unsafe { b2.launch(cfg2)?; }
10751 return Ok(());
10752 }
10753 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10754 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10755 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10756 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10757 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10758 let gqa = (n_head / n_head_kv).max(1) as u32;
10759 let o_len = t * n_head * n_splits_max * head_dim;
10760 let ml_len = t * n_head * n_splits_max;
10761 let mut part_guard = self.fa_part_pool.lock().unwrap();
10762 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10763 let old = part_guard.take();
10774 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10775 if let Some(old) = old {
10776 self.fa_part_retired.lock().unwrap().push(old);
10777 }
10778 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10779 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10780 }
10781 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10782 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10783 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10784 }
10785 let pg = part_guard.as_mut().unwrap();
10786 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10787 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10788 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10789 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10790 let f = self.func("fa_decode_vec_q_rows_v3_dc");
10791 let sh = (32 * head_dim * 2) as u32;
10792 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10793 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10794 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10795 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10796 let __s_b = self.gpu.stream();
10797 let mut b = __s_b.launch_builder(&f);
10798 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10799 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&scale).arg(&nspm).arg(&spk)
10800 .arg(&ktb).arg(&vtb);
10801 unsafe { b.launch(cfg)?; }
10802 let fc = self.func("fa_decode_combine_rows_dc");
10803 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10804 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10805 let plus0 = 0i32;
10806 let __s_b2 = self.gpu.stream();
10807 let mut b2 = __s_b2.launch_builder(&fc);
10808 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10809 .arg(base_dev).arg(&plus0).arg(&nspm).arg(&spk);
10810 unsafe { b2.launch(cfg2)?; }
10811 Ok(())
10812 }
10813
10814 pub fn fa_decode_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10825 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10826 head_dim: usize, n_head: usize, n_head_kv: usize,
10827 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10828 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
10829 -> Result<(), Box<dyn std::error::Error>> {
10830 self.fa_decode_dc_q8(q, k, v, o, head_dim, n_head, n_head_kv, t_kv_dev, bucket_max,
10831 scale, k_tok_bytes, v_tok_bytes, g, None)
10832 }
10833
10834 #[allow(clippy::too_many_arguments)]
10837 pub fn fa_decode_dc_q8(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10838 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10839 head_dim: usize, n_head: usize, n_head_kv: usize,
10840 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10841 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
10842 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10843 -> Result<(), Box<dyn std::error::Error>> {
10844 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
10852 if g && head_dim == 256 && !fa_v4_at(bucket_max) { fa_vec = false; } let sp = fa_split_keys(bucket_max, n_head_kv);
10854 let n_splits = if fa_vec { ((bucket_max + sp - 1) / sp).max(1) } else { ((bucket_max + 255) / 256).max(1) };
10855 let o_len = n_head * n_splits * head_dim;
10856 let ml_len = n_head * n_splits;
10857 let mut part_guard = self.fa_part_pool.lock().unwrap();
10858 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10859 let old = part_guard.take();
10870 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10871 if let Some(old) = old {
10872 self.fa_part_retired.lock().unwrap().push(old);
10873 }
10874 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10875 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10876 }
10877 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10878 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10879 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10880 }
10881 let pg = part_guard.as_mut().unwrap();
10882 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10883 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10884 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10885 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10886 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
10887 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10888 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
10889 let deep = fa_vec && head_dim == 256 && fa_v4_at(bucket_max) && !g
10892 && fa_deep_at(bucket_max) && !matches!(fa_v4_mode(), "noB3" | "stage");
10893 let (f, cfg) = if fa_vec && head_dim == 512 && bucket_max >= {
10894 static FA512_MIN_DC: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10895 *FA512_MIN_DC.get_or_init(|| std::env::var("MEMRA_FA512_MIN").ok()
10896 .and_then(|v| v.parse().ok()).unwrap_or(512))
10897 } {
10898 let gqa = (n_head / n_head_kv).max(1) as u32;
10900 (self.fa_func("fa_decode_vec_q_dpl16_dc", head_dim),
10901 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10902 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10903 } else if fa_vec && head_dim == 512 {
10904 let q_view = q.as_view();
10907 let mut o_view = o.as_view_mut();
10908 return self.fa_decode_scalar_unified(&q_view, k, v, &mut o_view,
10909 head_dim, n_head, n_head_kv,
10910 0, Some(t_kv_dev), scale, n_splits, sp,
10911 k_tok_bytes, v_tok_bytes, g,
10912 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10913 } else if fa_vec && head_dim == 256 && fa_v4_at(bucket_max) {
10914 let gqa = (n_head / n_head_kv).max(1) as u32;
10917 let fv = if g { self.func_g("fa_decode_vec_q_v4_dc") }
10918 else if deep { self.func("fa_decode_vec_q_v4_deep_dc") }
10919 else { self.func("fa_decode_vec_q_v4_dc") };
10920 let shmem = (if deep { 12160 } else { 11520 }
10921 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10922 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10923 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
10924 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10925 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10926 } else if fa_vec && fa_v3_active(head_dim) {
10927 let gqa = (n_head / n_head_kv).max(1) as u32;
10930 let fv = if g { self.func_g("fa_decode_vec_q_v3_dc") } else { self.func("fa_decode_vec_q_v3_dc") };
10931 let shmem = (32 * head_dim * 2) as u32; (fv,
10933 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10934 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10935 } else if fa_vec && fa_v2_on() {
10936 let gqa = (n_head / n_head_kv).max(1) as u32;
10940 let fv = if g { self.func_g("fa_decode_vec_q_v2_dc") } else { self.func("fa_decode_vec_q_v2_dc") };
10941 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
10943 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10944 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10945 } else if fa_vec {
10946 let gqa = (n_head / n_head_kv).max(1) as u32;
10947 let fv = if g { self.func_g("fa_decode_vec_q_dc") } else { self.func("fa_decode_vec_q_dc") };
10949 (fv,
10950 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10951 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10952 } else {
10953 let q_view = q.as_view();
10954 let mut o_view = o.as_view_mut();
10955 return self.fa_decode_scalar_unified(&q_view, k, v, &mut o_view,
10956 head_dim, n_head, n_head_kv,
10957 0, Some(t_kv_dev), scale, n_splits,
10958 if fa_vec { sp } else { 256 },
10959 k_tok_bytes, v_tok_bytes, g,
10960 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10961 };
10962 let ski = sp as i32; let __s_b = self.gpu.stream();
10964 let mut b = __s_b.launch_builder(&f);
10965 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10966 .arg(&hd).arg(&nh).arg(&nhkv).arg(t_kv_dev).arg(&scale).arg(&nsp).arg(&ski)
10967 .arg(&ktb).arg(&vtb);
10968 unsafe { b.launch(cfg)?; }
10969 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10970 if let Some((oq, od)) = q8_out {
10971 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
10972 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
10973 let __s_b2 = self.gpu.stream();
10974 let mut b2 = __s_b2.launch_builder(&fc);
10975 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
10976 unsafe { b2.launch(cfg2)?; }
10977 return Ok(());
10978 }
10979 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
10980 let __s_b2 = self.gpu.stream();
10981 let mut b2 = __s_b2.launch_builder(&fc);
10982 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
10983 unsafe { b2.launch(cfg2)?; }
10984 Ok(())
10985 }
10986
10987 pub fn fa_geom_eager(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
10993 let fa_ok = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
10997 let vec512 = fa_ok && head_dim == 512 && t_kv >= fa512_min_tkv();
11003 let mut fa_vec = vec512 || (fa_ok && head_dim <= 256 && head_dim % 32 == 0);
11004 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
11010 let sp = fa_split_keys(t_kv, n_head_kv);
11011 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
11012 (fa_vec, n_splits)
11013 }
11014
11015 pub fn fa_bucket_key(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
11021 self.fa_geom_eager(t_kv, head_dim, n_head_kv, g)
11022 }
11023
11024 pub fn capture_graph_retained<F>(&self, step: F)
11036 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
11037 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
11038 {
11039 use cudarc::driver::sys::CUgraphInstantiate_flags;
11040 self.capture_graph_retained_flags(
11041 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH, step)
11042 }
11043
11044 pub fn capture_graph_retained_flags<F>(&self,
11049 flags: cudarc::driver::sys::CUgraphInstantiate_flags, mut step: F)
11050 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
11051 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
11052 {
11053 use cudarc::driver::sys::CUstreamCaptureMode;
11054 self.capture_keep.lock().unwrap().clear();
11062 let was_tracking = self.gpu.ctx.is_event_tracking();
11063 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
11064 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
11065 self.capture_keep_on.store(true, std::sync::atomic::Ordering::Relaxed);
11066 let w = (|| { step(self)?; step(self) })();
11067 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
11068 w?;
11069 self.gpu.stream().synchronize()?;
11070 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
11071 let r = step(self);
11072 let g = self.gpu.stream().end_capture(flags);
11073 r?;
11074 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
11075 graph.upload()?;
11076 Ok(graph)
11077 };
11078 let result = run();
11079 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
11080 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
11081 let keeper = std::mem::take(&mut *self.capture_keep.lock().unwrap());
11082 Ok((result?, keeper))
11083 }
11084
11085 pub fn capture_graph<F>(&self, mut step: F) -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>>
11086 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
11087 {
11088 use cudarc::driver::sys::{CUstreamCaptureMode, CUgraphInstantiate_flags};
11089 let was_tracking = self.gpu.ctx.is_event_tracking();
11097 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
11098 let iflag = {
11105 static F: std::sync::OnceLock<CUgraphInstantiate_flags> = std::sync::OnceLock::new();
11106 *F.get_or_init(|| match std::env::var("MEMRA_GRAPH_IFLAG").as_deref() {
11107 Ok("upload") => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_UPLOAD,
11110 Ok("priority") =>
11111 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
11112 _ => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
11113 })
11114 };
11115 let ct = {
11122 static T: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11123 *T.get_or_init(|| std::env::var("MEMRA_GRAPH_CAPTIME").as_deref() == Ok("1"))
11124 };
11125 let warmups = {
11148 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
11149 *W.get_or_init(|| std::env::var("MEMRA_GRAPH_WARMUPS").ok()
11150 .and_then(|v| v.parse().ok()).filter(|n| *n >= 1).unwrap_or(1))
11151 };
11152 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
11153 let t_w = std::time::Instant::now();
11154 for _ in 0..warmups { step(self)?; }
11156 self.gpu.stream().synchronize()?;
11157 let ms_warm = t_w.elapsed().as_secs_f64() * 1e3;
11158 let t_c = std::time::Instant::now();
11160 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
11161 let r = step(self);
11164 let ms_body = t_c.elapsed().as_secs_f64() * 1e3;
11165 let t_i = std::time::Instant::now();
11166 let g = self.gpu.stream().end_capture(iflag);
11167 let ms_inst = t_i.elapsed().as_secs_f64() * 1e3;
11168 r?;
11169 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
11170 let t_u = std::time::Instant::now();
11171 graph.upload()?;
11172 if ct {
11173 println!("[graph-captime] warmup2x {ms_warm:.2} ms capture-body {ms_body:.2} ms \
11174 instantiate {ms_inst:.2} ms upload {:.2} ms",
11175 t_u.elapsed().as_secs_f64() * 1e3);
11176 }
11177 Ok(graph)
11178 };
11179 let result = run();
11180 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
11181 result
11182 }
11183
11184 pub fn gdn_scan_s128_view(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11186 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
11187 state_in: &cudarc::driver::CudaView<f32>,
11188 state_out: &mut cudarc::driver::CudaViewMut<f32>,
11189 o: &mut CudaSlice<f32>, n_head: usize, t: usize, scale: f32)
11190 -> Result<(), Box<dyn std::error::Error>> {
11191 let f = self.func("gdn_scan_s128");
11192 const S_V: u32 = 128; const WARP: u32 = 32; const COLS: u32 = 4;
11193 let cfg = LaunchConfig { grid_dim: (n_head as u32, 1, S_V / COLS), block_dim: (WARP, COLS, 1), shared_mem_bytes: 0 };
11194 let (h, ti) = (n_head as i32, t as i32);
11195 let __s_b = self.gpu.stream();
11196 let mut b = __s_b.launch_builder(&f);
11197 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in).arg(state_out).arg(o).arg(&h).arg(&ti).arg(&scale);
11198 unsafe { b.launch(cfg)?; }
11199 Ok(())
11200 }
11201
11202 pub fn ssm_conv1d_view(&self, x: &cudarc::driver::CudaView<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11204 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
11205 -> Result<(), Box<dyn std::error::Error>> {
11206 let f = self.func("ssm_conv1d_silu_f32");
11207 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
11209 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11210 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
11211 let __s_b = self.gpu.stream();
11212 let mut b = __s_b.launch_builder(&f);
11213 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
11214 unsafe { b.launch(cfg)?; }
11215 Ok(())
11216 }
11217
11218 pub fn ssm_conv1d_tm(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11225 conv_dim: usize, t: usize, d_conv: usize)
11226 -> Result<(), Box<dyn std::error::Error>> {
11227 let f = self.func("ssm_conv1d_tm_f32");
11228 let cfg = LaunchConfig {
11229 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11230 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11231 };
11232 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11233 let __s_b = self.gpu.stream();
11234 let mut b = __s_b.launch_builder(&f);
11235 b.arg(qkv_tm).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11236 unsafe { b.launch(cfg)?; }
11237 Ok(())
11238 }
11239
11240 pub fn ssm_conv1d_tm_state(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
11248 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11249 conv_dim: usize, t: usize, d_conv: usize)
11250 -> Result<(), Box<dyn std::error::Error>> {
11251 self.ssm_conv1d_tm_state_pad(qkv_tm, conv_state, w, y, conv_dim, t, d_conv, None)
11252 }
11253
11254 #[allow(clippy::too_many_arguments)]
11257 pub fn ssm_conv1d_tm_state_pad(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
11258 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11259 conv_dim: usize, t: usize, d_conv: usize,
11260 pad_len: Option<&CudaSlice<i32>>)
11261 -> Result<(), Box<dyn std::error::Error>> {
11262 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
11263 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11267 {
11268 let f = self.func("ssm_conv1d_tm_state_f32");
11269 let cfg = LaunchConfig {
11270 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11271 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11272 };
11273 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11274 let __s_b = self.gpu.stream();
11275 let mut b = __s_b.launch_builder(&f);
11276 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11277 unsafe { b.launch(cfg)?; }
11278 }
11279 match (ring_old, pad_len) {
11280 (None, Some(len_d)) => {
11281 let f = self.func("ssm_conv_ring_update_dev_f32");
11282 let n = conv_dim * (d_conv - 1);
11283 let cfg = LaunchConfig::for_num_elems(n as u32);
11284 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11285 let __s_b = self.gpu.stream();
11286 let mut b = __s_b.launch_builder(&f);
11287 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11288 unsafe { b.launch(cfg)?; }
11289 }
11290 (None, None) => {
11291 let f = self.func("ssm_conv_ring_update_f32");
11292 let n = conv_dim * (d_conv - 1);
11293 let cfg = LaunchConfig::for_num_elems(n as u32);
11294 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11295 let __s_b = self.gpu.stream();
11296 let mut b = __s_b.launch_builder(&f);
11297 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11298 unsafe { b.launch(cfg)?; }
11299 }
11300 (Some(old), _) => self.ssm_conv_ring_rebuild(qkv_tm, &old, conv_state, conv_dim, t, d_conv)?,
11301 }
11302 Ok(())
11303 }
11304
11305 pub fn ssm_conv1d_tm_state_pad_v(&self, qkv_tm: &cudarc::driver::CudaView<f32>, conv_state: &mut CudaSlice<f32>,
11307 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11308 conv_dim: usize, t: usize, d_conv: usize,
11309 pad_len: Option<&CudaSlice<i32>>)
11310 -> Result<(), Box<dyn std::error::Error>> {
11311 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
11312 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11316 {
11317 let f = self.func("ssm_conv1d_tm_state_f32");
11318 let cfg = LaunchConfig {
11319 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11320 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11321 };
11322 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11323 let __s_b = self.gpu.stream();
11324 let mut b = __s_b.launch_builder(&f);
11325 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11326 unsafe { b.launch(cfg)?; }
11327 }
11328 match (ring_old, pad_len) {
11329 (None, Some(len_d)) => {
11330 let f = self.func("ssm_conv_ring_update_dev_f32");
11331 let n = conv_dim * (d_conv - 1);
11332 let cfg = LaunchConfig::for_num_elems(n as u32);
11333 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11334 let __s_b = self.gpu.stream();
11335 let mut b = __s_b.launch_builder(&f);
11336 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11337 unsafe { b.launch(cfg)?; }
11338 }
11339 (None, None) => {
11340 let f = self.func("ssm_conv_ring_update_f32");
11341 let n = conv_dim * (d_conv - 1);
11342 let cfg = LaunchConfig::for_num_elems(n as u32);
11343 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11344 let __s_b = self.gpu.stream();
11345 let mut b = __s_b.launch_builder(&f);
11346 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11347 unsafe { b.launch(cfg)?; }
11348 }
11349 (Some(_), _) => unreachable!(
11350 "ssm_conv1d_tm_state_pad_v: T < d_conv-1 has no view path (PRIME_MIN_T gates it)"),
11351 }
11352 Ok(())
11353 }
11354
11355 pub fn ssm_conv_ring_rebuild(&self, qkv_tm: &CudaSlice<f32>, ring_old: &CudaSlice<f32>,
11360 conv_state: &mut CudaSlice<f32>,
11361 conv_dim: usize, tc: usize, d_conv: usize)
11362 -> Result<(), Box<dyn std::error::Error>> {
11363 let f = self.func("ssm_conv_ring_rebuild_f32");
11364 let n = conv_dim * (d_conv - 1);
11365 let cfg = LaunchConfig::for_num_elems(n as u32);
11366 let (cd, ti, dc) = (conv_dim as i32, tc as i32, d_conv as i32);
11367 let __s_b = self.gpu.stream();
11368 let mut b = __s_b.launch_builder(&f);
11369 b.arg(qkv_tm).arg(ring_old).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11370 unsafe { b.launch(cfg)?; }
11371 Ok(())
11372 }
11373
11374 #[allow(clippy::too_many_arguments)]
11379 pub fn gdn_prep_decode(&self, conv_out: &CudaSlice<f32>, beta_raw: &CudaSlice<f32>,
11380 alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11381 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11382 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11383 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32)
11384 -> Result<(), Box<dyn std::error::Error>> {
11385 let f = self.func("gdn_prep_decode_f32");
11386 let cfg = LaunchConfig { grid_dim: (num_v as u32, 1, 1), block_dim: (32, 4, 1), shared_mem_bytes: 0 };
11387 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11388 let __s_b = self.gpu.stream();
11389 let mut b = __s_b.launch_builder(&f);
11390 b.arg(conv_out).arg(beta_raw).arg(alpha).arg(dt_bias).arg(a)
11391 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11392 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps);
11393 unsafe { b.launch(cfg)?; }
11394 Ok(())
11395 }
11396
11397 #[allow(clippy::too_many_arguments)]
11401 pub fn ssm_conv1d_gdn(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>,
11402 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11403 conv_dim: usize, t: usize, d_conv: usize,
11404 d_state: usize, num_v: usize, num_k: usize, key_dim: usize)
11405 -> Result<(), Box<dyn std::error::Error>> {
11406 let f = self.func("ssm_conv1d_gdn_f32");
11407 let cfg = LaunchConfig {
11408 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11409 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11410 };
11411 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11412 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11413 let __s_b = self.gpu.stream();
11414 let mut b = __s_b.launch_builder(&f);
11415 b.arg(qkv_tm).arg(w).arg(q_g).arg(k_g).arg(v_g)
11416 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd);
11417 unsafe { b.launch(cfg)?; }
11418 Ok(())
11419 }
11420
11421 pub fn ssm_conv1d(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11422 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
11423 -> Result<(), Box<dyn std::error::Error>> {
11424 let f = self.func("ssm_conv1d_silu_f32");
11425 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
11426 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11427 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
11428 let __s_b = self.gpu.stream();
11429 let mut b = __s_b.launch_builder(&f);
11430 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
11431 unsafe { b.launch(cfg)?; }
11432 Ok(())
11433 }
11434
11435 pub fn gdn_scan_s128(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11438 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
11439 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11440 n_head: usize, t: usize, scale: f32)
11441 -> Result<(), Box<dyn std::error::Error>> {
11442 let f = self.func("gdn_scan_s128");
11443 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11444 let cfg = LaunchConfig {
11445 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
11446 block_dim: (WARP, COLS_PER_BLOCK, 1),
11447 shared_mem_bytes: 0,
11448 };
11449 let (h, ti) = (n_head as i32, t as i32);
11450 let __s_b = self.gpu.stream();
11451 let mut b = __s_b.launch_builder(&f);
11452 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in).arg(state_out).arg(o).arg(&h).arg(&ti).arg(&scale);
11453 unsafe { b.launch(cfg)?; }
11454 Ok(())
11455 }
11456
11457 #[allow(clippy::too_many_arguments)]
11462 pub fn ssm_conv1d_fused_decode_b(
11463 &self, qkv_cols: &CudaSlice<f32>, conv_state_ptrs: &cudarc::driver::CudaView<u64>,
11464 w: &CudaSlice<f32>, conv_outs: &mut CudaSlice<f32>, conv_dim: usize, d_conv: usize,
11465 b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11466 let f = self.func("ssm_conv1d_fused_decode_b_f32");
11467 let cfg = LaunchConfig {
11468 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
11469 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11470 };
11471 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11472 let __s_b = self.gpu.stream();
11473 let mut b = __s_b.launch_builder(&f);
11474 b.arg(qkv_cols).arg(conv_state_ptrs).arg(w).arg(conv_outs).arg(&cd).arg(&dc);
11475 unsafe { b.launch(cfg)?; }
11476 Ok(())
11477 }
11478
11479 #[allow(clippy::too_many_arguments)]
11480 pub fn gdn_prep_decode_b(
11481 &self, conv_outs: &CudaSlice<f32>, beta_raws: &CudaSlice<f32>, alphas: &CudaSlice<f32>,
11482 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11483 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11484 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11485 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32,
11486 conv_dim: usize, b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11487 let f = self.func("gdn_prep_decode_b_f32");
11488 let cfg = LaunchConfig {
11489 grid_dim: (num_v as u32, 1, b_n as u32),
11490 block_dim: (32, 4, 1), shared_mem_bytes: 0,
11491 };
11492 let (ds, nv, nk, kd, cd) =
11493 (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, conv_dim as i32);
11494 let __s_b = self.gpu.stream();
11495 let mut b = __s_b.launch_builder(&f);
11496 b.arg(conv_outs).arg(beta_raws).arg(alphas).arg(dt_bias).arg(a)
11497 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11498 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps).arg(&cd);
11499 unsafe { b.launch(cfg)?; }
11500 Ok(())
11501 }
11502
11503 #[allow(clippy::too_many_arguments)]
11504 pub fn gdn_scan_s128_batched(
11505 &self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11506 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
11507 state_in_ptrs: &cudarc::driver::CudaView<u64>,
11508 state_out_ptrs: &cudarc::driver::CudaView<u64>,
11509 o: &mut CudaSlice<f32>, n_head: usize, b_n: usize, scale: f32)
11510 -> Result<(), Box<dyn std::error::Error>> {
11511 let f = self.func("gdn_scan_s128_b");
11512 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11513 let cfg = LaunchConfig {
11514 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
11515 block_dim: (WARP, COLS_PER_BLOCK, 1), shared_mem_bytes: 0,
11516 };
11517 let h = n_head as i32;
11518 let __s_b = self.gpu.stream();
11519 let mut b = __s_b.launch_builder(&f);
11520 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in_ptrs).arg(state_out_ptrs)
11521 .arg(o).arg(&h).arg(&scale);
11522 unsafe { b.launch(cfg)?; }
11523 Ok(())
11524 }
11525
11526 pub fn gdn_chunked_enabled() -> bool {
11535 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11536 *E.get_or_init(|| std::env::var("MEMRA_GDN_CHUNKED").map(|v| v != "0").unwrap_or(true))
11537 }
11538
11539 pub fn gdn_chunk_size() -> usize {
11544 static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
11545 *C.get_or_init(|| {
11546 let c: usize = std::env::var("MEMRA_GDN_CHUNK").ok()
11547 .and_then(|v| v.parse().ok()).unwrap_or(32);
11548 c.clamp(32, 128) / 32 * 32
11549 })
11550 }
11551
11552 #[allow(clippy::too_many_arguments)]
11557 #[allow(clippy::too_many_arguments, clippy::type_complexity)]
11560 #[allow(clippy::too_many_arguments)]
11561 pub fn gdn_chunk_k123(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11562 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, wb16: Option<&mut CudaSlice<u8>>,
11563 n_head: usize, t: usize, c: usize, hk: usize,
11564 k2w: Option<(&CudaSlice<u8>, &CudaSlice<u8>, &mut CudaSlice<u8>)>)
11565 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11566 const D: usize = 128;
11567 let h = n_head;
11568 let nc = (t + c - 1) / c;
11569 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
11570 let mut gcum = self.uninit(t * h)?;
11571 let mut a = self.uninit(nc * h * c * c)?;
11572 let mut p = self.uninit(nc * h * c * c)?;
11573 let mut u = self.uninit(nc * h * c * D)?;
11574 let mut w = self.uninit(nc * h * c * D)?;
11575 { let f = self.func("gdn_chunk_cumgate_f32");
11577 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11578 let __s_b = self.gpu.stream();
11579 let mut b = __s_b.launch_builder(&f);
11580 b.arg(g).arg(&mut gcum).arg(&hi).arg(&ti).arg(&ci);
11581 unsafe { b.launch(cfg)?; }
11582 }
11583 if let Some((qb, kb, pb)) = k2w {
11584 assert!(c == 32, "gdn_k2_wgmma is a C==32 tile");
11587 let f = self.func("gdn_k2_wgmma");
11588 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11589 let hki = hk as i32;
11590 let __s_b = self.gpu.stream();
11591 let mut b = __s_b.launch_builder(&f);
11592 b.arg(qb).arg(kb).arg(&gcum).arg(beta).arg(&mut a).arg(&mut *pb).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11593 unsafe { b.launch(cfg)?; }
11594 } else if c <= 64 && !portable_mma_gated() { let f = self.func("gdn_chunk_attn_f32");
11596 let jt = ((c + 31) / 32) as u32;
11597 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11598 let hki = hk as i32;
11599 let __s_b = self.gpu.stream();
11600 let mut b = __s_b.launch_builder(&f);
11601 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11602 unsafe { b.launch(cfg)?; }
11603 } else { assert!(hk == h, "generic K2 is broadcast-only (de-broadcast rides C==32)");
11605 let f = self.func("gdn_chunk_attn_g_f32");
11606 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 8, 1), shared_mem_bytes: 0 };
11607 let __s_b = self.gpu.stream();
11608 let mut b = __s_b.launch_builder(&f);
11609 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci);
11610 unsafe { b.launch(cfg)?; }
11611 }
11612 { let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11614 match c {
11615 32 | 64 => {
11616 let f = self.func(if c == 32 { "gdn_chunk_solve32_f32" } else { "gdn_chunk_solve64_f32" });
11617 let wb: u64 = match wb16 { Some(d) => self.addr_u8(d), None => 0 };
11619 let hki = hk as i32;
11620 let __s_b = self.gpu.stream();
11621 let mut b = __s_b.launch_builder(&f);
11622 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&wb).arg(&hi).arg(&ti).arg(&hki);
11623 unsafe { b.launch(cfg)?; }
11624 }
11625 _ => {
11626 assert!(hk == h, "generic K3 is broadcast-only");
11627 let f = self.func("gdn_chunk_solve_f32");
11628 let __s_b = self.gpu.stream();
11629 let mut b = __s_b.launch_builder(&f);
11630 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&hi).arg(&ti).arg(&ci);
11631 unsafe { b.launch(cfg)?; }
11632 }
11633 }
11634 }
11635 Ok((gcum, p, u, w))
11636 }
11637
11638 pub fn gdn_db_on() -> bool {
11642 std::env::var("MEMRA_GDN_DB").as_deref() != Ok("0")
11643 }
11644
11645 pub fn gdn_mma_enabled(&self, c: usize) -> bool {
11648 !portable_mma_gated() && c == 32
11649 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11650 Ok("1") => true,
11651 Ok("0") => false,
11652 _ => cfg!(memra_hopper_mma),
11653 }
11654 }
11655
11656 pub fn gdn_wgmma_on(&self, c: usize) -> bool {
11659 self.gdn_mma_enabled(c)
11660 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
11661 Ok("0") => false,
11662 Ok("1") => true,
11663 _ => cfg!(memra_hopper_mma),
11664 }
11665 }
11666
11667 #[allow(clippy::too_many_arguments)]
11672 pub fn ssm_conv1d_gdn_state_pad(&self, qkv_tm: &cudarc::driver::CudaView<f32>,
11673 conv_state: &mut CudaSlice<f32>, w: &CudaSlice<f32>,
11674 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>,
11675 v_g: &mut CudaSlice<f32>,
11676 conv_dim: usize, t: usize, d_conv: usize,
11677 d_state: usize, num_v: usize, num_k: usize, key_dim: usize,
11678 hk: usize,
11679 pad_len: Option<&CudaSlice<i32>>)
11680 -> Result<(), Box<dyn std::error::Error>> {
11681 assert!(t >= d_conv - 1, "fused state conv requires T >= pad (PRIME_MIN_T gates)");
11682 {
11683 let f = self.func("ssm_conv1d_gdn_state_f32");
11684 let cfg = LaunchConfig {
11685 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11686 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11687 };
11688 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11689 let (ds, nv, nk, kd, hki) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, hk as i32);
11690 let __s_b = self.gpu.stream();
11691 let mut b = __s_b.launch_builder(&f);
11692 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(q_g).arg(k_g).arg(v_g)
11693 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&hki);
11694 unsafe { b.launch(cfg)?; }
11695 }
11696 match pad_len {
11697 Some(len_d) => {
11698 let f = self.func("ssm_conv_ring_update_dev_f32");
11699 let n = conv_dim * (d_conv - 1);
11700 let cfg = LaunchConfig::for_num_elems(n as u32);
11701 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11702 let __s_b = self.gpu.stream();
11703 let mut b = __s_b.launch_builder(&f);
11704 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11705 unsafe { b.launch(cfg)?; }
11706 }
11707 None => {
11708 let f = self.func("ssm_conv_ring_update_f32");
11709 let n = conv_dim * (d_conv - 1);
11710 let cfg = LaunchConfig::for_num_elems(n as u32);
11711 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11712 let __s_b = self.gpu.stream();
11713 let mut b = __s_b.launch_builder(&f);
11714 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11715 unsafe { b.launch(cfg)?; }
11716 }
11717 }
11718 Ok(())
11719 }
11720
11721 pub fn gdn_chunk_alloc(&self, n_head: usize, t: usize, c: usize, hk: usize)
11725 -> Result<GdnChunkBufs, Box<dyn std::error::Error>> {
11726 const D: usize = 128;
11727 assert!(c == 32, "gdn_chunk_alloc: varlen chain is the C==32 mma pair");
11728 let h = n_head;
11729 let nc = (t + c - 1) / c;
11730 Ok(GdnChunkBufs {
11731 gcum: self.uninit(t * h)?,
11732 a: self.uninit(nc * h * c * c)?,
11733 p: self.uninit(nc * h * c * c)?,
11734 u: self.uninit(nc * h * c * D)?,
11735 w: self.uninit(nc * h * c * D)?,
11736 kb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11737 wb16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11738 y16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11739 ssnap16: self.alloc_u8_uninit(nc * h * D * D * 2)?,
11740 qb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11741 pb16: self.alloc_u8_uninit(nc * h * c * c * 2)?,
11742 o: self.uninit(D * h * t)?,
11743 t, nc,
11744 })
11745 }
11746
11747 pub fn f32_to_bf16_v(&self, x: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<u8>, n: usize)
11749 -> Result<(), Box<dyn std::error::Error>> {
11750 let f = self.func("f32_to_bf16_bulk");
11751 let ni = n as i64;
11752 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11753 let __s_b = self.gpu.stream();
11754 let mut b = __s_b.launch_builder(&f);
11755 b.arg(x).arg(dst).arg(&ni);
11756 unsafe { b.launch(cfg)?; }
11757 Ok(())
11758 }
11759
11760 pub fn f32_to_bf16_into(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<u8>, n: usize)
11762 -> Result<(), Box<dyn std::error::Error>> {
11763 let f = self.func("f32_to_bf16_bulk");
11764 let ni = n as i64;
11765 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11766 let __s_b = self.gpu.stream();
11767 let mut b = __s_b.launch_builder(&f);
11768 b.arg(x).arg(dst).arg(&ni);
11769 unsafe { b.launch(cfg)?; }
11770 Ok(())
11771 }
11772
11773 pub fn gdn_chunk_k123_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, hk: usize,
11776 wq: Option<&GdnWVl8>)
11777 -> Result<(), Box<dyn std::error::Error>> {
11778 let b = seqs.len();
11779 assert!(b >= 1 && b <= 8, "gdn_chunk_k123_vl8: 1..=8 sequences");
11780 let mut packed = [GdnSeqVl::default(); 8];
11781 packed[..b].copy_from_slice(seqs);
11782 let v = GdnVl8(packed);
11783 let (hi, ci) = (n_head as i32, 32i32);
11784 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11785 {
11786 let f = self.func("gdn_chunk_cumgate_vl");
11787 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11788 let __s_lb = self.gpu.stream();
11789 let mut lb = __s_lb.launch_builder(&f);
11790 lb.arg(&v).arg(&hi).arg(&ci);
11791 unsafe { lb.launch(cfg)?; }
11792 }
11793 let hki = hk as i32;
11794 if let Some(w) = wq { let f = self.func("gdn_k2_wgmma_vl");
11796 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11797 let __s_lb = self.gpu.stream();
11798 let mut lb = __s_lb.launch_builder(&f);
11799 lb.arg(&v).arg(w).arg(&hi).arg(&ci).arg(&hki);
11800 unsafe { lb.launch(cfg)?; }
11801 } else {
11802 let f = self.func("gdn_chunk_attn_vl");
11803 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11804 let __s_lb = self.gpu.stream();
11805 let mut lb = __s_lb.launch_builder(&f);
11806 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11807 unsafe { lb.launch(cfg)?; }
11808 }
11809 {
11810 let f = self.func("gdn_chunk_solve32_vl");
11811 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11812 let __s_lb = self.gpu.stream();
11813 let mut lb = __s_lb.launch_builder(&f);
11814 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11815 unsafe { lb.launch(cfg)?; }
11816 }
11817 Ok(())
11818 }
11819
11820 #[allow(clippy::too_many_arguments)]
11824 pub fn gdn_prep_vl8(&self, seqs: &[GdnPrepVl], conv_w: &CudaSlice<f32>,
11825 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11826 conv_dim: usize, d_conv: usize, d_state: usize,
11827 num_v: usize, num_k: usize, key_dim: usize, hk: usize, eps: f32)
11828 -> Result<(), Box<dyn std::error::Error>> {
11829 let b = seqs.len();
11830 assert!(b >= 1 && b <= 8);
11831 let mut packed = [GdnPrepVl::default(); 8];
11832 packed[..b].copy_from_slice(seqs);
11833 let v = GdnPrepVl8(packed);
11834 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11835 let (cdi, dci) = (conv_dim as i32, d_conv as i32);
11836 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
11837 assert!(conv_fuse || hk == num_v, "de-broadcast requires the fused conv");
11838 if conv_fuse {
11839 let f = self.func("ssm_conv1d_gdn_state_vl");
11840 let cfg = LaunchConfig { grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11841 let (dsi, nvi, nki, kdi, hki) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, hk as i32);
11842 let __s_lb = self.gpu.stream();
11843 let mut lb = __s_lb.launch_builder(&f);
11844 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi).arg(&hki);
11845 unsafe { lb.launch(cfg)?; }
11846 } else {
11847 let f = self.func("ssm_conv1d_tm_state_vl");
11848 let cfg = LaunchConfig { grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11849 let __s_lb = self.gpu.stream();
11850 let mut lb = __s_lb.launch_builder(&f);
11851 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci);
11852 unsafe { lb.launch(cfg)?; }
11853 }
11854 {
11855 let f = self.func("ssm_conv_ring_update_vl");
11856 let n = (conv_dim * (d_conv - 1)) as u32;
11857 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11858 let __s_lb = self.gpu.stream();
11859 let mut lb = __s_lb.launch_builder(&f);
11860 lb.arg(&v).arg(&cdi).arg(&dci);
11861 unsafe { lb.launch(cfg)?; }
11862 }
11863 if !conv_fuse {
11864 let f = self.func("qkv_to_gdn_repack_vl");
11865 let n = max_t * (num_v * d_state) as u32;
11866 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11867 let (dsi, nvi, nki, kdi) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11868 let __s_lb = self.gpu.stream();
11869 let mut lb = __s_lb.launch_builder(&f);
11870 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi);
11871 unsafe { lb.launch(cfg)?; }
11872 }
11873 if Self::l2_v2_on(d_state) {
11874 let f = self.func("gdn_l2_v2_vl");
11875 let cfg = LaunchConfig { grid_dim: ((max_t * hk as u32).div_ceil(8), 2, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11876 let (dsi, nvi) = (d_state as i32, hk as i32);
11877 let __s_lb = self.gpu.stream();
11878 let mut lb = __s_lb.launch_builder(&f);
11879 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11880 unsafe { lb.launch(cfg)?; }
11881 } else {
11882 let f = self.func("gdn_l2_vl");
11883 let cfg = LaunchConfig { grid_dim: (max_t * hk as u32, 2, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11884 let (dsi, nvi) = (d_state as i32, hk as i32);
11885 let __s_lb = self.gpu.stream();
11886 let mut lb = __s_lb.launch_builder(&f);
11887 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11888 unsafe { lb.launch(cfg)?; }
11889 }
11890 {
11891 let f = self.func("gdn_gate_prep_vl");
11892 let n = max_t * num_v as u32;
11893 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11894 let nvi = num_v as i32;
11895 let __s_lb = self.gpu.stream();
11896 let mut lb = __s_lb.launch_builder(&f);
11897 lb.arg(&v).arg(dt_bias).arg(a).arg(&nvi);
11898 unsafe { lb.launch(cfg)?; }
11899 }
11900 Ok(())
11901 }
11902
11903 pub fn gdn_mirror_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, which: i32, hk: usize)
11905 -> Result<(), Box<dyn std::error::Error>> {
11906 let b = seqs.len();
11907 assert!(b >= 1 && b <= 8);
11908 let mut packed = [GdnSeqVl::default(); 8];
11909 packed[..b].copy_from_slice(seqs);
11910 let v = GdnVl8(packed);
11911 let ept = (if which == 0 { hk } else { n_head } * 128) as i32;
11912 let max_n = seqs.iter().map(|s| if which == 0 { s.t as i64 * ept as i64 }
11913 else { s.nc as i64 * ept as i64 * 32 }).max().unwrap();
11914 let f = self.func("gdn_mirror_vl");
11915 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
11916 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11917 let __s_lb = self.gpu.stream();
11918 let mut lb = __s_lb.launch_builder(&f);
11919 lb.arg(&v).arg(&ept).arg(&which);
11920 unsafe { lb.launch(cfg)?; }
11921 Ok(())
11922 }
11923
11924 pub fn gdn_tail_vl8(&self, seqs: &[GdnPrepVl], norm_w: &CudaSlice<f32>,
11926 d_state: usize, num_v: usize, eps: f32)
11927 -> Result<(), Box<dyn std::error::Error>> {
11928 let b = seqs.len();
11929 assert!(b >= 1 && b <= 8);
11930 let mut packed = [GdnPrepVl::default(); 8];
11931 packed[..b].copy_from_slice(seqs);
11932 let v = GdnPrepVl8(packed);
11933 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11934 let f = self.func("gated_rmsnorm_f16out_vl");
11935 let cfg = LaunchConfig { grid_dim: (max_t * num_v as u32, 1, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11937 let (dsi, nvi) = (d_state as i32, num_v as i32);
11938 let __s_lb = self.gpu.stream();
11939 let mut lb = __s_lb.launch_builder(&f);
11940 lb.arg(&v).arg(norm_w).arg(&dsi).arg(&nvi).arg(&eps);
11941 unsafe { lb.launch(cfg)?; }
11942 Ok(())
11943 }
11944
11945 pub fn addr_f32(&self, x: &CudaSlice<f32>) -> u64 {
11948 use cudarc::driver::DevicePtr;
11949 let s = self.gpu.stream();
11950 let (p, _g) = x.device_ptr(&s);
11951 p as u64
11952 }
11953 pub fn addr_f32_mut(&self, x: &mut CudaSlice<f32>) -> u64 {
11954 use cudarc::driver::DevicePtrMut;
11955 let s = self.gpu.stream();
11956 let (p, _g) = x.device_ptr_mut(&s);
11957 p as u64
11958 }
11959 pub fn addr_f32v(&self, x: &cudarc::driver::CudaView<f32>) -> u64 {
11960 use cudarc::driver::DevicePtr;
11961 let s = self.gpu.stream();
11962 let (p, _g) = x.device_ptr(&s);
11963 p as u64
11964 }
11965 pub fn addr_u8(&self, x: &CudaSlice<u8>) -> u64 {
11966 use cudarc::driver::DevicePtr;
11967 let s = self.gpu.stream();
11968 let (p, _g) = x.device_ptr(&s);
11969 p as u64
11970 }
11971
11972 pub fn gdn_chunk_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, scale: f32, hk: usize,
11976 wq: Option<&GdnWVl8>)
11977 -> Result<(), Box<dyn std::error::Error>> {
11978 const NSPLIT: u32 = 4;
11979 let b = seqs.len();
11980 assert!(b >= 1 && b <= 8, "gdn_chunk_vl8: 1..=8 sequences");
11981 let mut packed = [GdnSeqVl::default(); 8];
11982 packed[..b].copy_from_slice(seqs);
11983 let v = GdnVl8(packed);
11984 let (hi, ci) = (n_head as i32, 32i32);
11985 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11986 let hki = hk as i32;
11987 if let Some(w) = wq {
11988 let f = self.func("gdn_k45_wgmma_vl");
11990 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11991 let __s_lb = self.gpu.stream();
11992 let mut lb = __s_lb.launch_builder(&f);
11993 lb.arg(&v).arg(w).arg(&scale).arg(&hi).arg(&ci).arg(&hki);
11994 unsafe { lb.launch(cfg)?; }
11995 let _ = max_nc;
11996 return Ok(());
11997 }
11998 {
11999 let f = self.func("gdn_chunk_state_mma_vl");
12000 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12001 let __s_lb = self.gpu.stream();
12002 let mut lb = __s_lb.launch_builder(&f);
12003 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
12004 unsafe { lb.launch(cfg)?; }
12005 }
12006 {
12007 let f = self.func("gdn_chunk_output_mma_vl");
12008 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12009 let __s_lb = self.gpu.stream();
12010 let mut lb = __s_lb.launch_builder(&f);
12011 lb.arg(&v).arg(&hi).arg(&ci).arg(&scale).arg(&hki);
12012 unsafe { lb.launch(cfg)?; }
12013 }
12014 Ok(())
12015 }
12016 pub fn gdn_scan_chunked(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12017 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
12018 qb16_pre: Option<&CudaSlice<u8>>,
12019 state_in: &CudaSlice<f32>,
12020 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12021 n_head: usize, t: usize, scale: f32, c: usize, hk: usize)
12022 -> Result<(), Box<dyn std::error::Error>> {
12023 const D: usize = 128;
12024 const NSPLIT: u32 = 4;
12025 assert!(c >= 1 && c <= 128, "gdn_scan_chunked: C must be in 1..=128");
12026 let h = n_head;
12027 let nc = (t + c - 1) / c;
12028 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
12029 let gdn_mma_pre = !portable_mma_gated() && c == 32
12033 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
12034 Ok("1") => true,
12035 Ok("0") => false,
12036 _ => cfg!(memra_hopper_mma),
12037 };
12038 let mut wb16_pre: Option<CudaSlice<u8>> = if gdn_mma_pre {
12039 Some(self.alloc_u8_uninit(nc * h * c * D * 2)?)
12040 } else { None };
12041 let gdn_wgmma_pre = gdn_mma_pre
12045 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
12046 Ok("0") => false,
12047 Ok("1") => true,
12048 _ => cfg!(memra_hopper_mma),
12049 };
12050 let nk = t * hk * D;
12051 let mut kb16_local: Option<CudaSlice<u8>> = None;
12052 if gdn_mma_pre && kb16_pre.is_none() {
12053 let mut kb = self.alloc_u8_uninit(nk * 2)?;
12054 let f = self.func("f32_to_bf16_bulk");
12055 let n2 = nk as i64;
12056 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
12057 let __s_b = self.gpu.stream();
12058 let mut b = __s_b.launch_builder(&f);
12059 b.arg(k).arg(&mut kb).arg(&n2);
12060 unsafe { b.launch(cfg2)?; }
12061 kb16_local = Some(kb);
12062 }
12063 let kb16_ref0: Option<&CudaSlice<u8>> = kb16_local.as_ref().or(kb16_pre);
12064 if let Some(kb) = kb16_pre { assert!(kb.len() >= nk * 2, "kb16_pre too small"); }
12065 let mut qb16: Option<CudaSlice<u8>> = None;
12066 let mut pb16: Option<CudaSlice<u8>> = None;
12067 if gdn_wgmma_pre {
12068 if qb16_pre.is_none() {
12071 let mut qb = self.alloc_u8_uninit(nk * 2)?;
12072 let f = self.func("f32_to_bf16_bulk");
12073 let n2 = nk as i64;
12074 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
12075 let __s_b = self.gpu.stream();
12076 let mut b = __s_b.launch_builder(&f);
12077 b.arg(q).arg(&mut qb).arg(&n2);
12078 unsafe { b.launch(cfg2)?; }
12079 qb16 = Some(qb);
12080 } else if let Some(qb) = qb16_pre {
12081 assert!(qb.len() >= nk * 2, "qb16_pre too small");
12082 }
12083 pb16 = Some(self.alloc_u8_uninit(nc * h * c * c * 2)?);
12084 }
12085 let qb16_ref0: Option<&CudaSlice<u8>> = qb16.as_ref().or(qb16_pre);
12086 let k2w = if gdn_wgmma_pre {
12087 Some((*qb16_ref0.as_ref().unwrap(),
12088 *kb16_ref0.as_ref().unwrap(),
12089 pb16.as_mut().unwrap()))
12090 } else { None };
12091 let (gcum, p, u, w) = self.gdn_chunk_k123(q, k, v, g, beta, wb16_pre.as_mut(), n_head, t, c, hk, k2w)?;
12092 let _ = &w;
12093 let mut y = self.uninit(nc * h * c * D)?;
12094 let mut ssnap = self.uninit(nc * h * D * D)?; let gdn_mma = !portable_mma_gated() && c == 32
12106 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
12107 Ok("1") => true,
12108 Ok("0") => false,
12109 _ => cfg!(memra_hopper_mma),
12110 };
12111 if gdn_mma {
12112 let wb16 = wb16_pre.take().expect("mma path pre-allocates wb16 (K3 store fold)");
12113 let kb16_ref: &CudaSlice<u8> = kb16_ref0.expect("mma path pre-builds kb16 above K123");
12114 if gdn_wgmma_pre {
12126 let qb16 = qb16_ref0.unwrap();
12128 let pb16 = pb16.as_ref().unwrap();
12129 {
12130 let f = self.func("gdn_k45_wgmma");
12131 let cfg = LaunchConfig { grid_dim: (h as u32, 4, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12132 let hki = hk as i32;
12133 let __s_b = self.gpu.stream();
12134 let mut b = __s_b.launch_builder(&f);
12135 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(qb16).arg(pb16)
12136 .arg(o).arg(&scale).arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
12137 unsafe { b.launch(cfg)?; }
12138 }
12139 return Ok(());
12140 }
12141 let mut y16 = self.alloc_u8_uninit(nc * h * c * D * 2)?;
12145 let mut ssnap16 = self.alloc_u8_uninit(nc * h * D * D * 2)?;
12146 {
12147 let f = self.func("gdn_chunk_state_mma");
12148 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12149 let hki = hk as i32;
12150 let __s_b = self.gpu.stream();
12151 let mut b = __s_b.launch_builder(&f);
12152 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(&mut y16).arg(&mut ssnap16)
12153 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
12154 unsafe { b.launch(cfg)?; }
12155 }
12156 { let f = self.func("gdn_chunk_output_mma");
12158 let jt = ((c + 31) / 32) as u32;
12159 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12160 let hki = hk as i32;
12161 let __s_b = self.gpu.stream();
12162 let mut b = __s_b.launch_builder(&f);
12163 b.arg(q).arg(&gcum).arg(&p).arg(&y16).arg(&ssnap16).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale).arg(&hki);
12164 unsafe { b.launch(cfg)?; }
12165 }
12166 return Ok(());
12167 }
12168 { let f = self.func("gdn_chunk_state_f32");
12170 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12171 let __s_b = self.gpu.stream();
12172 let mut b = __s_b.launch_builder(&f);
12173 b.arg(k).arg(&gcum).arg(beta).arg(&u).arg(&w).arg(&mut y).arg(&mut ssnap)
12174 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci);
12175 unsafe { b.launch(cfg)?; }
12176 }
12177 { let f = self.func("gdn_chunk_output_f32");
12179 let jt = ((c + 31) / 32) as u32;
12180 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12181 let __s_b = self.gpu.stream();
12182 let mut b = __s_b.launch_builder(&f);
12183 b.arg(q).arg(&gcum).arg(&p).arg(&y).arg(&ssnap).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale);
12184 unsafe { b.launch(cfg)?; }
12185 }
12186 Ok(())
12187 }
12188
12189 #[allow(clippy::too_many_arguments)]
12198 #[allow(clippy::too_many_arguments)]
12199 pub fn gdn_scan_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12200 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
12201 qb16_pre: Option<&CudaSlice<u8>>,
12202 state_in: &CudaSlice<f32>,
12203 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12204 n_head: usize, t: usize, scale: f32, hk: usize)
12205 -> Result<(), Box<dyn std::error::Error>> {
12206 if std::env::var("MEMRA_GDN_DIFF").is_ok() && t >= 16 {
12207 assert!(hk == n_head, "GDN_DIFF oracle is broadcast-only");
12208 return self.gdn_scan_diff(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale);
12209 }
12210 if Self::gdn_chunked_enabled() && t >= 16 {
12211 self.gdn_scan_chunked(q, k, v, g, beta, kb16_pre, qb16_pre, state_in, state_out, o, n_head, t, scale,
12212 Self::gdn_chunk_size(), hk)
12213 } else {
12214 assert!(hk == n_head, "s128 scan is broadcast-only (prep guarantees by predicate)");
12215 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)
12216 }
12217 }
12218
12219 #[allow(clippy::too_many_arguments)]
12221 fn gdn_scan_diff(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12222 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
12223 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12224 n_head: usize, t: usize, scale: f32)
12225 -> Result<(), Box<dyn std::error::Error>> {
12226 static CALL: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
12227 let call = CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12228 let mut o_c = self.uninit(o.len())?;
12229 let mut st_c = self.uninit(state_out.len())?;
12230 self.gdn_scan_chunked(q, k, v, g, beta, None, None, state_in, &mut st_c, &mut o_c,
12231 n_head, t, scale, Self::gdn_chunk_size(), n_head)?;
12232 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)?;
12233 let (oh_s, oh_c) = (self.dtoh(o)?, self.dtoh(&o_c)?);
12234 let (sh_s, sh_c) = (self.dtoh(state_out)?, self.dtoh(&st_c)?);
12235 let stats = |a: &[f32], b: &[f32]| -> (f32, f32, f64) {
12236 let mut max_abs = 0f32; let mut max_rel = 0f32; let mut sum_rel = 0f64;
12237 for (x, y) in a.iter().zip(b) {
12238 let ad = (x - y).abs();
12239 let rel = ad / x.abs().max(y.abs()).max(1e-3);
12240 if ad > max_abs { max_abs = ad; }
12241 if rel > max_rel { max_rel = rel; }
12242 sum_rel += rel as f64;
12243 }
12244 (max_abs, max_rel, sum_rel / a.len() as f64)
12245 };
12246 let (o_ma, o_mr, o_mean) = stats(&oh_s, &oh_c);
12247 let (s_ma, s_mr, s_mean) = stats(&sh_s, &sh_c);
12248 println!("[gdn-diff call {call:3} T={t} C={}] out: max_abs={o_ma:.3e} max_rel={o_mr:.3e} mean_rel={o_mean:.3e} | \
12249 state: max_abs={s_ma:.3e} max_rel={s_mr:.3e} mean_rel={s_mean:.3e}",
12250 Self::gdn_chunk_size());
12251 Ok(())
12252 }
12253
12254 pub fn gdn_glog(&self, alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
12256 g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
12257 -> Result<(), Box<dyn std::error::Error>> {
12258 let f = self.func("gdn_glog_f32");
12259 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
12260 let (h, ti) = (n_head as i32, t as i32);
12261 let __s_b = self.gpu.stream();
12262 let mut b = __s_b.launch_builder(&f);
12263 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
12264 unsafe { b.launch(cfg)?; }
12265 Ok(())
12266 }
12267
12268 pub fn sigmoid_v(&self, x: &cudarc::driver::CudaView<f32>, y: &mut CudaSlice<f32>, n: usize)
12271 -> Result<(), Box<dyn std::error::Error>> {
12272 let f = self.func("sigmoid_f32");
12273 let cfg = LaunchConfig::for_num_elems(n as u32);
12274 let ni = n as i32;
12275 let __s_b = self.gpu.stream();
12276 let mut b = __s_b.launch_builder(&f);
12277 b.arg(x).arg(y).arg(&ni);
12278 unsafe { b.launch(cfg)?; }
12279 Ok(())
12280 }
12281
12282 pub fn gdn_glog_v(&self, alpha: &cudarc::driver::CudaView<f32>, dt_bias: &CudaSlice<f32>,
12283 a: &CudaSlice<f32>, g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
12284 -> Result<(), Box<dyn std::error::Error>> {
12285 let f = self.func("gdn_glog_f32");
12286 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
12287 let (h, ti) = (n_head as i32, t as i32);
12288 let __s_b = self.gpu.stream();
12289 let mut b = __s_b.launch_builder(&f);
12290 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
12291 unsafe { b.launch(cfg)?; }
12292 Ok(())
12293 }
12294
12295 pub fn sigmoid(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize)
12296 -> Result<(), Box<dyn std::error::Error>> {
12297 let f = self.func("sigmoid_f32");
12298 let cfg = LaunchConfig::for_num_elems(n as u32);
12299 let ni = n as i32;
12300 let __s_b = self.gpu.stream();
12301 let mut b = __s_b.launch_builder(&f);
12302 b.arg(x).arg(y).arg(&ni);
12303 unsafe { b.launch(cfg)?; }
12304 Ok(())
12305 }
12306
12307 pub fn sig_mul_f16out(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12310 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
12311 -> Result<(), Box<dyn std::error::Error>> {
12312 let f = self.func("sig_mul_f16out_f32");
12313 let cfg = LaunchConfig::for_num_elems(n as u32);
12314 let ni = n as i32;
12315 let __s_b = self.gpu.stream();
12316 let mut b = __s_b.launch_builder(&f);
12317 b.arg(a).arg(g).arg(dst).arg(dst16).arg(&ni);
12318 unsafe { b.launch(cfg)?; }
12319 Ok(())
12320 }
12321
12322 #[allow(clippy::too_many_arguments)]
12331 pub fn attn_head_gate(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12332 dst: &mut CudaSlice<f32>, dst16: Option<&mut CudaSlice<u8>>,
12333 head_dim: usize, n_head: usize, t: usize)
12334 -> Result<(), Box<dyn std::error::Error>> {
12335 let f = self.func("attn_head_gate_f32");
12336 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12337 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12338 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
12340 let __s_b = self.gpu.stream();
12341 let mut b = __s_b.launch_builder(&f);
12342 b.arg(a).arg(g).arg(dst).arg(&d16).arg(&hd).arg(&nh).arg(&ti);
12343 unsafe { b.launch(cfg)?; }
12344 Ok(())
12345 }
12346
12347 #[allow(clippy::too_many_arguments)]
12356 pub fn swiglu_clamped_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
12357 gs: f32, us: f32, limit: f32,
12358 dst: &mut CudaSlice<f32>, n: usize)
12359 -> Result<(), Box<dyn std::error::Error>> {
12360 debug_assert!(limit > 1e-6, "swiglu_clamped needs a live limit; use silu_mul_scaled");
12361 let f = self.func("swiglu_clamped_mul_scaled_f32");
12362 let cfg = LaunchConfig::for_num_elems(n as u32);
12363 let ni = n as i32;
12364 let __s_b = self.gpu.stream();
12365 let mut b = __s_b.launch_builder(&f);
12366 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&limit).arg(dst).arg(&ni);
12367 unsafe { b.launch(cfg)?; }
12368 Ok(())
12369 }
12370
12371 pub fn gated_rmsnorm(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12373 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12374 -> Result<(), Box<dyn std::error::Error>> {
12375 let f = self.func("gated_rmsnorm_f32");
12376 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12377 let (nc, e) = (ncols as i32, eps);
12378 let __s_b = self.gpu.stream();
12379 let mut b = __s_b.launch_builder(&f);
12380 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12381 unsafe { b.launch(cfg)?; }
12382 Ok(())
12383 }
12384
12385 pub fn gated_rmsnorm_f16out(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12388 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12389 ncols: usize, nrows: usize, eps: f32)
12390 -> Result<(), Box<dyn std::error::Error>> {
12391 let f = self.func("gated_rmsnorm_f16out_f32");
12392 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12394 let (nc, e) = (ncols as i32, eps);
12395 let __s_b = self.gpu.stream();
12396 let mut b = __s_b.launch_builder(&f);
12397 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12398 unsafe { b.launch(cfg)?; }
12399 Ok(())
12400 }
12401
12402 #[allow(clippy::too_many_arguments)]
12406 pub fn add_rms_norm_zq8(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
12407 res: &mut CudaSlice<f32>, z: &mut CudaSlice<f32>,
12408 ncols: usize, nrows: usize, eps: f32)
12409 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12410 assert!(ncols % 32 == 0);
12411 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
12412 let mut d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12413 let f = self.func("add_rms_norm_zq8");
12414 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
12415 let (nc, ep) = (ncols as i32, eps);
12416 let __s_b = self.gpu.stream();
12417 let mut b = __s_b.launch_builder(&f);
12418 b.arg(a).arg(b_in).arg(w).arg(res).arg(z).arg(&mut q).arg(&mut d).arg(&nc).arg(&ep);
12419 unsafe { b.launch(cfg)?; }
12420 Ok((q, d))
12421 }
12422
12423 pub fn gated_rmsnorm_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12428 z: &cudarc::driver::CudaView<f32>,
12429 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12430 -> Result<(), Box<dyn std::error::Error>> {
12431 let f = self.func("gated_rmsnorm_f32");
12432 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12433 let (nc, e) = (ncols as i32, eps);
12434 let __s_b = self.gpu.stream();
12435 let mut b = __s_b.launch_builder(&f);
12436 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12437 unsafe { b.launch(cfg)?; }
12438 Ok(())
12439 }
12440
12441 pub fn gated_rmsnorm_f16out_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12442 z: &cudarc::driver::CudaView<f32>,
12443 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12444 ncols: usize, nrows: usize, eps: f32)
12445 -> Result<(), Box<dyn std::error::Error>> {
12446 let f = self.func("gated_rmsnorm_f16out_f32");
12447 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12449 let (nc, e) = (ncols as i32, eps);
12450 let __s_b = self.gpu.stream();
12451 let mut b = __s_b.launch_builder(&f);
12452 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12453 unsafe { b.launch(cfg)?; }
12454 Ok(())
12455 }
12456
12457 pub fn gated_rmsnorm_q8_1(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12458 ncols: usize, nrows: usize, eps: f32)
12459 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12460 assert!(ncols % 32 == 0);
12461 let f = self.func("gated_rmsnorm_q8_1");
12462 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12463 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12464 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12465 let (nc, ep) = (ncols as i32, eps);
12466 let __s_b = self.gpu.stream();
12467 let mut b = __s_b.launch_builder(&f);
12468 b.arg(o).arg(w).arg(z).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
12469 unsafe { b.launch(cfg)?; }
12470 Ok((out_q, out_d))
12471 }
12472
12473 pub fn transpose(&self, inp: &CudaSlice<f32>, rows: usize, cols: usize)
12475 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12476 let f = self.func("transpose_f32");
12477 let mut out = self.zeros(rows * cols)?;
12478 let cfg = LaunchConfig::for_num_elems((rows * cols) as u32);
12479 let (r, c) = (rows as i32, cols as i32);
12480 let __s_b = self.gpu.stream();
12481 let mut b = __s_b.launch_builder(&f);
12482 b.arg(inp).arg(&mut out).arg(&r).arg(&c);
12483 unsafe { b.launch(cfg)?; }
12484 Ok(out)
12485 }
12486
12487 pub fn repeat_heads(&self, inp: &CudaSlice<f32>, out: &mut CudaSlice<f32>,
12489 head_dim: usize, n_in: usize, n_out: usize, t: usize)
12490 -> Result<(), Box<dyn std::error::Error>> {
12491 let f = self.func("repeat_heads_f32");
12492 let cfg = LaunchConfig::for_num_elems((head_dim * n_out * t) as u32);
12493 let (hd, ni, no, ti) = (head_dim as i32, n_in as i32, n_out as i32, t as i32);
12494 let __s_b = self.gpu.stream();
12495 let mut b = __s_b.launch_builder(&f);
12496 b.arg(inp).arg(out).arg(&hd).arg(&ni).arg(&no).arg(&ti);
12497 unsafe { b.launch(cfg)?; }
12498 Ok(())
12499 }
12500
12501 pub fn q_gate_split(&self, qf: &CudaSlice<f32>, q_out: &mut CudaSlice<f32>,
12504 gate_out: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, t: usize)
12505 -> Result<(), Box<dyn std::error::Error>> {
12506 let f = self.func("q_gate_split_f32");
12507 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12508 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12509 let __s_b = self.gpu.stream();
12510 let mut b = __s_b.launch_builder(&f);
12511 b.arg(qf).arg(q_out).arg(gate_out).arg(&hd).arg(&nh).arg(&ti);
12512 unsafe { b.launch(cfg)?; }
12513 Ok(())
12514 }
12515
12516 pub fn qkv_to_gdn_repack(&self, conv_out: &CudaSlice<f32>, q_g: &mut CudaSlice<f32>,
12520 k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
12521 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, t: usize)
12522 -> Result<(), Box<dyn std::error::Error>> {
12523 let f = self.func("qkv_to_gdn_repack_f32");
12524 let cfg = LaunchConfig::for_num_elems((d_state * num_v * t) as u32);
12525 let (ds, nv, nk, kd, ti) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, t as i32);
12526 let __s_b = self.gpu.stream();
12527 let mut b = __s_b.launch_builder(&f);
12528 b.arg(conv_out).arg(q_g).arg(k_g).arg(v_g).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&ti);
12529 unsafe { b.launch(cfg)?; }
12530 Ok(())
12531 }
12532
12533 pub fn conv_left_pad(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
12536 conv_dim: usize, t: usize, pad: usize)
12537 -> Result<(), Box<dyn std::error::Error>> {
12538 let f = self.func("conv_left_pad_f32");
12539 let cfg = LaunchConfig::for_num_elems((conv_dim * t) as u32);
12540 let (cd, ti, p) = (conv_dim as i32, t as i32, pad as i32);
12541 let __s_b = self.gpu.stream();
12542 let mut b = __s_b.launch_builder(&f);
12543 b.arg(src).arg(dst).arg(&cd).arg(&ti).arg(&p);
12544 unsafe { b.launch(cfg)?; }
12545 Ok(())
12546 }
12547
12548 pub fn conv_assemble_and_roll(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12552 conv_in: &mut CudaSlice<f32>, conv_dim: usize, pad: usize)
12553 -> Result<(), Box<dyn std::error::Error>> {
12554 let f = self.func("conv_assemble_and_roll_f32");
12555 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12556 let (cd, p) = (conv_dim as i32, pad as i32);
12557 let __s_b = self.gpu.stream();
12558 let mut b = __s_b.launch_builder(&f);
12559 b.arg(qkv_col).arg(conv_state).arg(conv_in).arg(&cd).arg(&p);
12560 unsafe { b.launch(cfg)?; }
12561 Ok(())
12562 }
12563
12564 pub fn ssm_conv1d_fused_decode(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12570 w: &CudaSlice<f32>, conv_out: &mut CudaSlice<f32>,
12571 conv_dim: usize, d_conv: usize)
12572 -> Result<(), Box<dyn std::error::Error>> {
12573 let f = self.func("ssm_conv1d_fused_decode_f32");
12574 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12575 let (cd, dc) = (conv_dim as i32, d_conv as i32);
12576 let __s_b = self.gpu.stream();
12577 let mut b = __s_b.launch_builder(&f);
12578 b.arg(qkv_col).arg(conv_state).arg(w).arg(conv_out).arg(&cd).arg(&dc);
12579 unsafe { b.launch(cfg)?; }
12580 Ok(())
12581 }
12582
12583 pub fn slice_range(&self, src: &CudaSlice<f32>, start: usize, len: usize)
12586 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12587 let host = self.gpu.stream().clone_dtoh(src)?;
12588 self.gpu.stream().synchronize()?;
12589 Ok(self.htod(&host[start..start + len])?)
12590 }
12591}
12592
12593#[cfg(test)]
12594mod target_dispatch_tests {
12595 use super::legacy_quant_gemm_allowed;
12596
12597 #[test]
12598 fn legacy_quant_gemm_arch_policy_honors_the_escape_hatch() {
12599 assert!(legacy_quant_gemm_allowed(false, false, false));
12601 assert!(!legacy_quant_gemm_allowed(false, false, true));
12602 assert!(!legacy_quant_gemm_allowed(true, false, false));
12604 assert!(!legacy_quant_gemm_allowed(true, false, true));
12605 assert!(legacy_quant_gemm_allowed(true, true, false));
12607 assert!(!legacy_quant_gemm_allowed(true, true, true));
12608 }
12609
12610 #[cfg(all(memra_portable_cuda, not(memra_hopper_mma)))]
12611 #[test]
12612 fn portable_build_disables_legacy_quant_gemm_without_an_env_override() {
12613 assert!(!legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12614 }
12615
12616 #[cfg(memra_hopper_mma)]
12617 #[test]
12618 fn hopper_mma_build_re_admits_legacy_quant_gemm() {
12619 assert!(legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12620 assert!(super::portable_mma_gated() == false);
12621 }
12622}
12623
12624impl memra_kv::KvDev for Engine {
12627 fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12628 Engine::zeros(self, n)
12629 }
12630 fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12631 Engine::uninit(self, n)
12632 }
12633 fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
12634 Engine::alloc_u8(self, n)
12635 }
12636 fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
12637 Engine::htod_i32(self, v)
12638 }
12639 fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12640 Engine::clone_dtod(self, src)
12641 }
12642 fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
12643 -> Result<(), Box<dyn std::error::Error>> {
12644 Engine::copy_into(self, dst, off, src, len)
12645 }
12646 fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
12647 Engine::set_i32_one(self, d, v)
12648 }
12649}