1use std::sync::{Arc, Mutex};
4use cudarc::driver::{CudaContext, CudaStream, CudaModule, CudaFunction, CudaSlice, LaunchConfig, PushKernelArg};
5use cudarc::nvrtc::Ptx;
6
7pub use memra_gguf;
8pub use memra_runtime;
9
10pub mod model;
11pub mod forward;
12pub mod hybrid;
13pub mod hybrid_forward;
14pub mod cache {
17 pub use memra_kv::*;
18}
19pub mod decode;
20pub mod decode_batch;
21pub mod mla;
25pub mod pp;
26pub mod spec;
27pub mod gemma_spec;
28pub mod round_stream;
29pub mod graph_update;
30pub mod dflash;
31pub mod eagle;
32pub use memra_sampling as sampler;
33
34pub fn moe_f16g_mode() -> u8 {
78 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
79 *M.get_or_init(|| match std::env::var("MEMRA_MOE_F16G").as_deref() {
80 Ok("0") => 0,
81 Ok("2") => 2,
82 Ok("3") => 3,
83 Ok(_) => 1,
84 Err(_) => 2,
87 })
88}
89pub fn moe_f16g_sk_params() -> (i32, i32) {
103 static P: std::sync::OnceLock<(i32, i32)> = std::sync::OnceLock::new();
104 *P.get_or_init(|| match std::env::var("MEMRA_F16G_SK").as_deref() {
105 Ok("0") => (-1, 0),
106 Ok("32") => (0, i32::MAX),
107 Ok("128") => (0, 1),
108 _ => {
109 let cross = std::env::var("MEMRA_F16G_SK_CROSS").ok()
110 .and_then(|v| v.parse().ok()).unwrap_or(64);
111 (0, cross)
112 }
113 })
114}
115pub fn moe_f16g_direct_on(qtype: i32) -> bool {
126 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
127 let m = *M.get_or_init(|| match std::env::var("MEMRA_F16G_DIRECT").as_deref() {
128 Ok("0") => 0,
129 Ok("kq") => 1,
130 _ => 2,
131 });
132 match m {
133 0 => false,
134 1 => qtype == QT_Q4_K || qtype == QT_Q6_K,
135 _ => true,
136 }
137}
138pub fn moe_f16g_tail_on() -> bool {
147 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
148 *ON.get_or_init(|| std::env::var("MEMRA_F16G_TAIL").as_deref() != Ok("0"))
149}
150
151pub fn moe_f16g_gemma_on() -> bool {
158 static M: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
159 *M.get_or_init(|| !matches!(std::env::var("MEMRA_MOE_F16G").as_deref(), Ok("0") | Err(_)))
160}
161
162pub fn moe_fuse_actq_on() -> bool {
166 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
167 *ON.get_or_init(|| std::env::var("MEMRA_MOE_FUSE_ACTQ").as_deref() != Ok("0"))
168}
169
170pub fn router_prefill_exact_on() -> bool {
180 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
181 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_PREFILL_EXACT").as_deref() != Ok("0"))
182}
183
184pub fn router_kernel_on() -> bool {
185 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
186 *ON.get_or_init(|| {
187 let on = std::env::var("MEMRA_ROUTER_KERNEL").as_deref() != Ok("0");
188 if !on { eprintln!("[memra] router kernel OFF (rollback: per-column cuBLAS gemv)"); }
189 on
190 })
191}
192
193pub const ROUTER_BATCH_MIN_T: usize = 8;
208pub fn router_batch_on() -> bool {
209 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
210 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_BATCH").as_deref() != Ok("0"))
211}
212mod cpu_experts;
213pub mod moe_cache;
214pub mod spill;
215mod spill_pread;
216#[cfg(memra_cutlass)]
217pub mod cutlass_ffi;
218pub mod mmq_ffi;
219pub mod f16_ffi;
220pub mod prime_graph;
221pub mod fp8_ffi;
222
223const FATBIN: &[u8] = include_bytes!(env!("MEMRA_ENGINE_FATBIN"));
230const HYBRID_FATBIN: &[u8] = include_bytes!(env!("MEMRA_HYBRID_FATBIN"));
231const QMATVEC_FATBIN: &[u8] = include_bytes!(env!("MEMRA_QMATVEC_FATBIN"));
232const FLASH_FATBIN: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN"));
233const GEMM_FATBIN: &[u8] = include_bytes!(env!("MEMRA_GEMM_FATBIN"));
234const ROUTER_FATBIN: &[u8] = include_bytes!(env!("MEMRA_ROUTER_FATBIN"));
235const SAMPLE_FATBIN: &[u8] = include_bytes!(env!("MEMRA_SAMPLE_FATBIN"));
237
238fn gemm_fatbin_bytes() -> std::borrow::Cow<'static, [u8]> {
244 assert!(!(portable_mma_gated() && std::env::var_os("MEMRA_GEMM_FATBIN").is_some()),
245 "MEMRA_GEMM_FATBIN overrides are not allowed in the portable CUDA lane");
246 match std::env::var("MEMRA_GEMM_FATBIN") {
247 Ok(path) => std::borrow::Cow::Owned(
248 std::fs::read(&path).unwrap_or_else(|e| panic!("MEMRA_GEMM_FATBIN read {path}: {e}"))),
249 Err(_) => std::borrow::Cow::Borrowed(GEMM_FATBIN),
250 }
251}
252
253pub(crate) const fn portable_mma_gated() -> bool {
260 cfg!(memra_portable_cuda) && !cfg!(memra_hopper_mma)
261}
262
263const fn legacy_quant_gemm_allowed(portable_cuda: bool, hopper_mma: bool, no_gemm: bool) -> bool {
268 (!portable_cuda || hopper_mma) && !no_gemm
269}
270
271const FLASH_FATBIN_VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VQ4"));
279const FLASH_FATBIN_VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VF8"));
280const FLASH_FATBIN_KF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8"));
281const FLASH_FATBIN_KF8VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VQ4"));
282const FLASH_FATBIN_KF8VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VF8"));
283
284pub use memra_kv::{kv_blk_bytes, kv_cache_formats};
287
288fn flash_fatbin_bytes() -> &'static [u8] {
290 match kv_cache_formats() {
291 ("q8_0", "q5_1") => FLASH_FATBIN,
292 ("q8_0", "q4_0") => FLASH_FATBIN_VQ4,
293 ("q8_0", "fp8") => FLASH_FATBIN_VF8,
294 ("fp8", "q5_1") => FLASH_FATBIN_KF8,
295 ("fp8", "q4_0") => FLASH_FATBIN_KF8VQ4,
296 ("fp8", "fp8") => FLASH_FATBIN_KF8VF8,
297 other => unreachable!("kv_cache_formats returned {other:?}"),
298 }
299}
300
301fn k1_launch_override() -> Option<(u32, u32, u32)> {
308 static K1: std::sync::OnceLock<Option<(u32, u32, u32)>> = std::sync::OnceLock::new();
309 *K1.get_or_init(|| {
310 let v = std::env::var("MEMRA_GEMM_K1_LAUNCH").ok()?;
311 let p: Vec<u32> = v.split(',').filter_map(|s| s.trim().parse().ok()).collect();
312 match p.as_slice() { [bm, bn, w] => Some((*bm, *bn, *w)), _ => None }
313 })
314}
315
316pub(crate) fn wgmma_gemm_enabled() -> bool {
323 static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
324 *V.get_or_init(|| std::env::var("MEMRA_WGMMA").as_deref() == Ok("1"))
325}
326
327pub const FA_VEC_MIN_TKV: usize = 96;
342pub fn fa_vec_min_tkv() -> usize {
346 static V: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
347 *V.get_or_init(|| std::env::var("MEMRA_FA_VEC_MIN").ok()
348 .and_then(|v| v.parse().ok())
349 .unwrap_or_else(|| FA_VEC_MIN_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)))
350}
351
352pub fn fa_f16pv_on() -> bool {
363 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
364 *ON.get_or_init(|| std::env::var("MEMRA_FA_F16PV").map(|v| v != "0")
365 .unwrap_or_else(|_| std::env::var("MEMRA_DRAFT").is_err()))
366}
367
368pub fn fa512_hp_on() -> bool {
372 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
373 *ON.get_or_init(|| std::env::var("MEMRA_FA512_HP").as_deref() != Ok("0"))
374}
375
376pub fn faw_hp_on() -> bool {
380 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
381 *ON.get_or_init(|| std::env::var("MEMRA_FAW_HP").as_deref() != Ok("0"))
382}
383
384pub fn fa512_wide_warps() -> usize {
388 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
389 *N.get_or_init(|| match std::env::var("MEMRA_FA512_W4").as_deref() {
390 Ok("1") => 4, _ => 2,
391 })
392}
393
394pub fn fa512_min_tkv() -> usize {
397 static FA512_MIN: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
398 *FA512_MIN.get_or_init(|| std::env::var("MEMRA_FA512_MIN").ok()
399 .and_then(|v| v.parse().ok()).unwrap_or(512))
400}
401pub static FA_VEC_MIN_DEFAULT: std::sync::atomic::AtomicUsize =
405 std::sync::atomic::AtomicUsize::new(FA_VEC_MIN_TKV);
406pub static FA_SPW_DEFAULT: std::sync::atomic::AtomicUsize =
410 std::sync::atomic::AtomicUsize::new(32);
411pub static FUSED_MR1_DEFAULT: std::sync::atomic::AtomicBool =
417 std::sync::atomic::AtomicBool::new(false);
418pub static ROUTER_W8_DEFAULT: std::sync::atomic::AtomicBool =
425 std::sync::atomic::AtomicBool::new(true);
426pub static FA_SP512_DEFAULT: std::sync::atomic::AtomicUsize =
427 std::sync::atomic::AtomicUsize::new(16);
428pub static RMS_BLOCK_DEFAULT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(256);
433pub static FA_SP_GEMMA: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
435pub static MMQ_SK_FORCE: std::sync::atomic::AtomicI8 = std::sync::atomic::AtomicI8::new(-1);
440pub use memra_kv::KV_FP8_FORCE;
443pub(crate) fn rms_block() -> u32 {
444 static V: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
445 *V.get_or_init(|| std::env::var("MEMRA_RMS_BLOCK").ok()
446 .and_then(|v| v.parse().ok())
447 .unwrap_or_else(|| RMS_BLOCK_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)))
448}
449
450pub(crate) fn fa_split_keys(t_kv: usize, n_head_kv: usize) -> usize {
451 static S: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
452 if let Some(forced) = *S.get_or_init(|| {
453 std::env::var("MEMRA_FA_SPLIT").ok().and_then(|v| v.parse().ok())
454 .filter(|&s: &usize| s >= 8 && s % 8 == 0)
455 }) { return forced; }
456 if FA_SP_GEMMA.load(std::sync::atomic::Ordering::Relaxed)
474 && std::env::var("MEMRA_FA_SP16").as_deref() == Ok("1") {
475 return if t_kv <= 8192 { 16 } else if t_kv <= 16384 { 64 } else { 128 };
476 }
477 let big_rig = fa_sm_count() >= 128;
478 if big_rig {
479 let _ = n_head_kv;
480 if t_kv <= 2048 { 16 } else if t_kv <= 16384 { 64 } else { 128 }
481 } else if n_head_kv <= 4 {
482 if t_kv <= 512 { 8 } else if t_kv <= 16384 { 64 } else { 128 }
503 } else {
504 if t_kv <= 8192 { 32 } else if t_kv <= 16384 { 64 } else { 128 }
505 }
506}
507
508fn fa_sm_count() -> i32 {
511 static N: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
512 *N.get_or_init(|| {
513 cudarc::driver::result::init().ok();
514 cudarc::driver::result::device::get(0)
515 .and_then(|d| unsafe { cudarc::driver::result::device::get_attribute(
516 d, cudarc::driver::sys::CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT) })
517 .unwrap_or(82)
518 })
519}
520
521fn fa_hd_suffix(head_dim: usize) -> Result<&'static str, Box<dyn std::error::Error>> {
525 match head_dim {
526 256 => Ok(""),
527 128 => Ok("_hd128"),
528 d => Err(format!("fa_prefill: no kernel stamped for head_dim={d} (only 256/128); \
529 callers must gate to sdpa_naive").into()),
530 }
531}
532
533pub const QT_Q8_0: i32 = 0;
535pub const QT_Q4_K: i32 = 1;
536pub const QT_Q6_K: i32 = 2;
537pub const QT_Q5_K: i32 = 3;
538pub const QT_Q3_K: i32 = 4;
539pub const QT_IQ4_XS: i32 = 5;
540pub const QT_IQ3_S: i32 = 6;
541pub const QT_NVFP4: i32 = 7;
542pub const QT_F8_E4M3: i32 = 10;
548pub const QT_NVFP4_RP: i32 = 9;
551pub const QT_F32: i32 = 8;
553pub const QT_BF16: i32 = 11;
554pub const QT_Q4_0: i32 = 12; pub const QT_Q2_K: i32 = 13;
559pub const QT_F8_E4M3_BLK: i32 = 14;
575
576pub struct Engine {
578 pub gpu: memra_runtime::Gpu,
579 module: Arc<CudaModule>,
580 hybrid: Arc<CudaModule>,
581 qmatvec: Arc<CudaModule>,
582 flash: Arc<CudaModule>,
583 flash_g: std::sync::OnceLock<Arc<CudaModule>>,
587 gemm: Arc<CudaModule>,
588 router: Arc<CudaModule>,
589 sample: Arc<CudaModule>,
591 moe_cache: Mutex<Option<crate::moe_cache::MoeSlotCache>>,
595 moe_cache_layout: Mutex<Option<Vec<usize>>>,
599 capture_keep_on: std::sync::atomic::AtomicBool,
605 verify_exact: std::sync::atomic::AtomicBool,
610 capture_keep: Mutex<Vec<Box<dyn std::any::Any + Send>>>,
611 pub copy_stream: Arc<CudaStream>,
613 #[cfg(memra_cutlass)]
620 cutlass_scratch: Mutex<Option<crate::cutlass_ffi::CutlassScratch>>,
621 fp8_scratch: Mutex<Option<crate::fp8_ffi::Fp8Scratch>>,
625 fa_vf16_scratch: Mutex<Option<CudaSlice<u8>>>,
628 fa_part_pool: Mutex<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
632 fa_part_retired: Mutex<Vec<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
636 fn_cache: Mutex<std::collections::HashMap<String, CudaFunction>>,
638 f16_scratch: Mutex<Option<crate::f16_ffi::F16Scratch>>,
639 argmax_partials: Mutex<Option<(CudaSlice<f32>, CudaSlice<i32>)>>,
644 prime_deqw_ws: Mutex<Option<(CudaSlice<u8>, CudaSlice<u8>)>>,
649 router_stage: Mutex<Option<PinnedStage>>,
653}
654
655fn fa_v2_on() -> bool {
665 std::env::var("MEMRA_FA_V2").map(|v| v != "0").unwrap_or(true)
671}
672
673fn fa_v3_on() -> bool {
681 std::env::var("MEMRA_FA_V3").map(|v| v != "0").unwrap_or(true)
685}
686
687fn fa_v4_mode() -> &'static str {
692 static M: std::sync::OnceLock<String> = std::sync::OnceLock::new();
693 M.get_or_init(|| std::env::var("MEMRA_FA_V4").unwrap_or_default())
694}
695fn fa_v4_on() -> bool { fa_v4_mode() != "0" } pub static FA_SMEM_TKV_DEFAULT: std::sync::atomic::AtomicUsize =
704 std::sync::atomic::AtomicUsize::new(1024);
705pub static FA_V4_MAX_DEFAULT: std::sync::atomic::AtomicUsize =
706 std::sync::atomic::AtomicUsize::new(usize::MAX);
707pub fn fa_v4_at_pub(t_kv: usize) -> bool { fa_v4_at(t_kv) }
708fn fa_v4_at(t_kv: usize) -> bool {
709 static M: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
710 let mx = *M.get_or_init(|| std::env::var("MEMRA_FA_V4_MAX").ok()
711 .and_then(|v| v.parse().ok())
712 .unwrap_or_else(|| FA_V4_MAX_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)));
713 fa_v4_on() && t_kv < mx
714}
715pub const FA_DEEP_MIN_DEFAULT: usize = 0;
729fn fa_deep_at(t_kv: usize) -> bool {
730 if std::env::var("MEMRA_FA_DEEP").as_deref() == Ok("0") { return false; }
731 let min = std::env::var("MEMRA_FA_DEEP_MIN").ok().and_then(|v| v.parse().ok())
732 .unwrap_or(FA_DEEP_MIN_DEFAULT);
733 t_kv >= min
734}
735pub fn fa_deep_at_pub(t_kv: usize) -> bool { fa_deep_at(t_kv) }
737
738fn fa_v3_active(head_dim: usize) -> bool {
739 fa_v3_on() && head_dim % 128 == 0 && kv_cache_formats() == ("q8_0", "q5_1")
742 && !Engine::kv_fp8_on()
743}
744
745pub fn fa_seqs_eligible(t_kv: usize, head_dim: usize) -> bool {
753 std::env::var("MEMRA_NO_FA_VEC").is_err()
754 && t_kv >= fa_vec_min_tkv()
755 && head_dim == 256
756 && fa_v4_at(t_kv)
757 && !matches!(fa_v4_mode(), "noB3" | "stage")
758 && !Engine::kv_fp8_on()
759}
760pub fn fa_split_keys_pub(t_kv: usize, n_head_kv: usize) -> usize { fa_split_keys(t_kv, n_head_kv) }
762
763struct PinnedStage {
768 ptr: *mut u8,
769 cap: usize,
770}
771unsafe impl Send for PinnedStage {}
772impl PinnedStage {
773 fn new(cap: usize) -> Result<Self, Box<dyn std::error::Error>> {
774 let ptr = unsafe { cudarc::driver::result::malloc_host(cap, 0)? } as *mut u8;
775 Ok(PinnedStage { ptr, cap })
776 }
777}
778impl Drop for PinnedStage {
779 fn drop(&mut self) {
780 let _ = unsafe { cudarc::driver::result::free_host(self.ptr as _) };
781 }
782}
783
784pub const ARGMAX_NB: usize = 256;
787
788pub(crate) use memra_fa3_vl as fa3_vl_raw;
790
791unsafe extern "C" {
792 fn memra_fa3_prefill(q16: *const core::ffi::c_void, k16: *const core::ffi::c_void,
794 v16: *const core::ffi::c_void, o: *mut f32,
795 t: i32, h: i32, hkv: i32, d: i32, scale: f32,
796 stream: *mut core::ffi::c_void) -> i32;
797 pub(crate) fn memra_fa3_vl(q16s: *const *const core::ffi::c_void, k16s: *const *const core::ffi::c_void,
799 v16s: *const *const core::ffi::c_void, os: *const *mut f32,
800 ts: *const i32, b: i32, h: i32, hkv: i32, d: i32, scale: f32,
801 stream: *mut core::ffi::c_void) -> i32;
802}
803
804#[repr(C)]
809#[derive(Clone, Copy)]
810pub struct WPtr8(pub [u64; 8]);
811unsafe impl cudarc::driver::DeviceRepr for WPtr8 {}
812
813#[repr(C)]
818#[derive(Clone, Copy, Default)]
819pub struct GdnSeqVl {
820 pub kb16: u64, pub gcum: u64, pub beta: u64, pub u: u64, pub wb16: u64,
821 pub y: u64, pub ssnap: u64, pub state_in: u64, pub state_out: u64,
822 pub q: u64, pub p: u64, pub o: u64,
823 pub k: u64, pub v: u64, pub g: u64, pub a: u64, pub w: u64,
824 pub t: i32, pub nc: i32,
825}
826unsafe impl cudarc::driver::DeviceRepr for GdnSeqVl {}
827#[repr(C)]
828#[derive(Clone, Copy)]
829pub struct GdnVl8(pub [GdnSeqVl; 8]);
830unsafe impl cudarc::driver::DeviceRepr for GdnVl8 {}
831
832#[repr(C)]
835#[derive(Clone, Copy, Default)]
836pub struct GdnWVl { pub qb16: u64, pub pb16: u64 }
837unsafe impl cudarc::driver::DeviceRepr for GdnWVl {}
838#[repr(C)]
839#[derive(Clone, Copy)]
840pub struct GdnWVl8(pub [GdnWVl; 8]);
841unsafe impl cudarc::driver::DeviceRepr for GdnWVl8 {}
842
843#[repr(C)]
845#[derive(Clone, Copy, Default)]
846pub struct GdnPrepVl {
847 pub qkv: u64, pub conv_state: u64, pub conv_out: u64,
848 pub q_g: u64, pub k_g: u64, pub v_g: u64,
849 pub q_l2: u64, pub k_l2: u64,
850 pub beta_raw: u64, pub alpha: u64, pub beta: u64, pub g_log: u64,
851 pub o: u64, pub z: u64, pub gn: u64, pub gn16: u64,
852 pub kb16: u64,
853 pub qb16: u64,
854 pub t: i32, pub pad: i32,
855}
856unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl {}
857#[repr(C)]
858#[derive(Clone, Copy)]
859pub struct GdnPrepVl8(pub [GdnPrepVl; 8]);
860unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl8 {}
861
862#[repr(C)]
864#[derive(Clone, Copy, Default)]
865pub struct FaSeqVl {
866 pub q: u64, pub k16: u64, pub v16: u64, pub o: u64, pub kf: u64, pub vf: u64,
867 pub t: i32, pub pad: i32,
868}
869unsafe impl cudarc::driver::DeviceRepr for FaSeqVl {}
870#[repr(C)]
871#[derive(Clone, Copy)]
872pub struct FaVl8(pub [FaSeqVl; 8]);
873unsafe impl cudarc::driver::DeviceRepr for FaVl8 {}
874
875#[repr(C)]
877#[derive(Clone, Copy, Default)]
878pub struct AttnPreVl {
879 pub qf: u64, pub kf: u64, pub vf: u64,
880 pub q: u64, pub gate: u64, pub qn: u64, pub kn: u64,
881 pub kc: u64, pub vc: u64,
882 pub t: i32, pub pad: i32,
883}
884unsafe impl cudarc::driver::DeviceRepr for AttnPreVl {}
885#[repr(C)]
886#[derive(Clone, Copy)]
887pub struct AttnPreVl8(pub [AttnPreVl; 8]);
888unsafe impl cudarc::driver::DeviceRepr for AttnPreVl8 {}
889
890pub struct GdnChunkBufs {
893 pub gcum: CudaSlice<f32>,
894 pub a: CudaSlice<f32>,
895 pub p: CudaSlice<f32>,
896 pub u: CudaSlice<f32>,
897 pub w: CudaSlice<f32>,
898 pub kb16: CudaSlice<u8>,
899 pub wb16: CudaSlice<u8>,
900 pub y16: CudaSlice<u8>,
901 pub ssnap16: CudaSlice<u8>,
902 pub qb16: CudaSlice<u8>,
903 pub pb16: CudaSlice<u8>,
904 pub o: CudaSlice<f32>,
905 pub t: usize,
906 pub nc: usize,
907}
908
909#[repr(C)]
911#[derive(Clone, Copy)]
912pub struct F32x8(pub [f32; 8]);
913unsafe impl cudarc::driver::DeviceRepr for F32x8 {}
914
915pub static PRIME_NANOS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
919
920impl Engine {
921 pub fn new(ordinal: usize) -> Result<Self, Box<dyn std::error::Error>> {
922 let gpu = memra_runtime::Gpu::new(ordinal)?;
923 if std::env::var("MEMRA_ARCH_CHECK").as_deref() != Ok("0") {
927 use cudarc::driver::sys::CUdevice_attribute_enum as A;
928 let (maj, min) = cudarc::driver::result::device::get(ordinal as i32)
929 .and_then(|d| unsafe { Ok((
930 cudarc::driver::result::device::get_attribute(d, A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR)?,
931 cudarc::driver::result::device::get_attribute(d, A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR)?)) })
932 .unwrap_or((0, 0));
933 let built = env!("MEMRA_BUILT_CUDA_ARCH");
934 let ok = matches!((built, maj, min),
935 ("120a", 12, 0) | ("120a", 12, 1) | ("100a", 10, 0) | ("90a", 9, 0) | ("89", 8, 9));
936 if !ok {
937 return Err(format!(
938 "memra was built for sm_{built} but device {ordinal} reports compute \
939 capability {maj}.{min}. Rebuild on this machine (MEMRA_CUDA_ARCH \
940 auto-detects the GPU) or set MEMRA_ARCH_CHECK=0 to bypass.").into());
941 }
942 }
943 unsafe {
948 use cudarc::driver::sys;
949 let dev: sys::CUdevice = ordinal as sys::CUdevice;
950 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
951 if sys::cuDeviceGetDefaultMemPool(&mut pool, dev) == sys::CUresult::CUDA_SUCCESS {
952 let mut thresh: u64 = u64::MAX;
953 let _ = sys::cuMemPoolSetAttribute(
954 pool,
955 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
956 &mut thresh as *mut u64 as *mut core::ffi::c_void,
957 );
958 }
959 }
960 let module = gpu.ctx.load_module(Ptx::from_binary(FATBIN.to_vec()))?;
961 let hybrid = gpu.ctx.load_module(Ptx::from_binary(HYBRID_FATBIN.to_vec()))?;
962 let qmatvec = gpu.ctx.load_module(Ptx::from_binary(QMATVEC_FATBIN.to_vec()))?;
963 let flash = gpu.ctx.load_module(Ptx::from_binary(flash_fatbin_bytes().to_vec()))?;
964 let gemm = gpu.ctx.load_module(Ptx::from_binary(gemm_fatbin_bytes().into_owned()))?;
965 let router = gpu.ctx.load_module(Ptx::from_binary(ROUTER_FATBIN.to_vec()))?;
966 let sample = gpu.ctx.load_module(Ptx::from_binary(SAMPLE_FATBIN.to_vec()))?;
967 let copy_stream = gpu.ctx.new_stream()?;
968 if std::env::var("MEMRA_EVT").map(|v| v == "1").unwrap_or(false) {
984 } else {
986 unsafe { gpu.ctx.disable_event_tracking(); }
987 }
988 Ok(Self { gpu, module, hybrid, qmatvec, flash, flash_g: std::sync::OnceLock::new(), gemm, router, sample,
989 moe_cache: Mutex::new(None),
990 moe_cache_layout: Mutex::new(None),
991 copy_stream,
992 capture_keep_on: std::sync::atomic::AtomicBool::new(false),
993 verify_exact: std::sync::atomic::AtomicBool::new(false),
994 capture_keep: Mutex::new(Vec::new()),
995 argmax_partials: Mutex::new(None),
996 prime_deqw_ws: Mutex::new(None),
997 router_stage: Mutex::new(None),
998 fp8_scratch: Mutex::new(None),
999 fa_vf16_scratch: Mutex::new(None),
1000 fa_part_pool: Mutex::new(None),
1001 fa_part_retired: Mutex::new(Vec::new()),
1002 fn_cache: Mutex::new(Default::default()),
1003 f16_scratch: Mutex::new(None),
1004 #[cfg(memra_cutlass)]
1005 cutlass_scratch: Mutex::new(None) })
1006 }
1007
1008 pub fn ctx(&self) -> &Arc<CudaContext> { &self.gpu.ctx }
1009
1010 pub fn pool_cached_bytes(&self) -> usize {
1028 let (reserved, used) = self.pool_reserved_used();
1029 reserved.saturating_sub(used)
1030 }
1031
1032 pub fn pool_reserved_used(&self) -> (usize, usize) {
1039 use cudarc::driver::sys;
1040 unsafe {
1041 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
1042 if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
1043 != sys::CUresult::CUDA_SUCCESS
1044 {
1045 return (0, 0);
1046 }
1047 let (mut reserved, mut used) = (0u64, 0u64);
1048 if sys::cuMemPoolGetAttribute(
1049 pool,
1050 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RESERVED_MEM_CURRENT,
1051 &mut reserved as *mut u64 as *mut core::ffi::c_void,
1052 ) != sys::CUresult::CUDA_SUCCESS {
1053 return (0, 0);
1054 }
1055 if sys::cuMemPoolGetAttribute(
1056 pool,
1057 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_USED_MEM_CURRENT,
1058 &mut used as *mut u64 as *mut core::ffi::c_void,
1059 ) != sys::CUresult::CUDA_SUCCESS {
1060 return (0, 0);
1061 }
1062 (reserved as usize, used as usize)
1063 }
1064 }
1065
1066 pub fn stream(&self) -> Arc<CudaStream> { self.gpu.stream() }
1069 pub fn gkv_on() -> bool {
1072 memra_kv::gkv_on()
1073 }
1074
1075 pub fn wkv_on() -> bool {
1087 memra_kv::wkv_on()
1088 }
1089
1090 pub fn kv_fp8_on() -> bool {
1096 memra_kv::kv_fp8_on()
1097 }
1098
1099 fn fa_func(&self, name: &str, head_dim: usize) -> CudaFunction {
1102 if head_dim == 512 && Self::gkv_on() { self.func_g(name) } else { self.func(name) }
1103 }
1104
1105 fn func_g(&self, name: &str) -> CudaFunction {
1109 let m = self.flash_g.get_or_init(|| {
1110 self.gpu.ctx.load_module(cudarc::nvrtc::Ptx::from_binary(FLASH_FATBIN_KF8VF8.to_vec()))
1111 .expect("load kf8vf8 flash fatbin (fp8-globals arm)")
1112 });
1113 let key = format!("g:{name}");
1114 if let Some(f) = self.fn_cache.lock().unwrap().get(&key) { return f.clone(); }
1115 let f = match m.load_function(name) {
1116 Ok(f) => f,
1117 Err(_) => self.func(name),
1118 };
1119 self.fn_cache.lock().unwrap().insert(key, f.clone());
1120 f
1121 }
1122
1123 fn func(&self, name: &str) -> CudaFunction {
1124 if let Some(f) = self.fn_cache.lock().unwrap().get(name) { return f.clone(); }
1127 let f = self.module.load_function(name)
1128 .or_else(|_| self.hybrid.load_function(name))
1129 .or_else(|_| self.qmatvec.load_function(name))
1130 .or_else(|_| self.flash.load_function(name))
1131 .or_else(|_| self.gemm.load_function(name))
1132 .or_else(|_| self.router.load_function(name))
1133 .or_else(|_| self.sample.load_function(name))
1134 .unwrap_or_else(|_| panic!("kernel {name} not in any fatbin"));
1135 self.fn_cache.lock().unwrap().insert(name.to_string(), f.clone());
1136 f
1137 }
1138
1139 pub fn scatter_trim_logits(&self, src: &CudaSlice<f32>, d2t: &CudaSlice<u32>,
1142 dst: &mut CudaSlice<f32>, d_vocab: usize, n_vocab: usize)
1143 -> Result<(), Box<dyn std::error::Error>> {
1144 let f1 = self.func("scatter_trim_logits_f32");
1145 let f2 = self.func("scatter_trim_logits_pass2_f32");
1146 let (dv, nv) = (d_vocab as i32, n_vocab as i32);
1147 let cfg1 = LaunchConfig { grid_dim: (256, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1148 let __s_b1 = self.gpu.stream();
1149 let mut b1 = __s_b1.launch_builder(&f1);
1150 b1.arg(src).arg(d2t).arg(&mut *dst).arg(&dv).arg(&nv);
1151 unsafe { b1.launch(cfg1)?; }
1152 let cfg2 = LaunchConfig { grid_dim: (d_vocab.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1153 let __s_b2 = self.gpu.stream();
1154 let mut b2 = __s_b2.launch_builder(&f2);
1155 b2.arg(src).arg(d2t).arg(&mut *dst).arg(&dv);
1156 unsafe { b2.launch(cfg2)?; }
1157 Ok(())
1158 }
1159
1160 #[allow(clippy::too_many_arguments)]
1166 pub fn filter_stats(&self, x: &CudaSlice<f32>, row_stride: usize, rows: &CudaSlice<i32>,
1167 out_th: &mut CudaSlice<f32>, out_z: &mut CudaSlice<f32>,
1168 out_max: &mut CudaSlice<f32>, n: usize, nrow: usize,
1169 temp: f32, top_k: i32, top_p: f32, min_p: f32)
1170 -> Result<(), Box<dyn std::error::Error>> {
1171 let f = self.func("filter_stats_f32");
1172 let (ni, nr, rs) = (n as i32, nrow as i32, row_stride as i64);
1173 let cfg = LaunchConfig { grid_dim: (nrow as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
1174 let __s_b = self.gpu.stream();
1175 let mut b = __s_b.launch_builder(&f);
1176 b.arg(x).arg(&rs).arg(rows).arg(&mut *out_th).arg(&mut *out_z).arg(&mut *out_max)
1177 .arg(&ni).arg(&nr).arg(&temp).arg(&top_k).arg(&top_p).arg(&min_p);
1178 unsafe { b.launch(cfg)?; }
1179 Ok(())
1180 }
1181
1182 #[allow(clippy::too_many_arguments)]
1184 pub fn softmax_gather_filtered(&self, x: &CudaSlice<f32>, row_stride: usize,
1185 ids: &CudaSlice<u32>, rows: &CudaSlice<i32>,
1186 th: &CudaSlice<f32>, z: &CudaSlice<f32>,
1187 out: &mut CudaSlice<f32>, n: usize, npair: usize, temp: f32)
1188 -> Result<(), Box<dyn std::error::Error>> {
1189 let f = self.func("softmax_gather_filtered_f32");
1190 let (ni, np, rs) = (n as i32, npair as i32, row_stride as i64);
1191 let cfg = LaunchConfig { grid_dim: (npair as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1192 let __s_b = self.gpu.stream();
1193 let mut b = __s_b.launch_builder(&f);
1194 b.arg(x).arg(&rs).arg(ids).arg(rows).arg(th).arg(z).arg(&mut *out).arg(&ni).arg(&np).arg(&temp);
1195 unsafe { b.launch(cfg)?; }
1196 Ok(())
1197 }
1198
1199 #[allow(clippy::too_many_arguments)]
1201 pub fn residual_sample_filtered(&self, p: &CudaSlice<f32>, q: Option<&CudaSlice<f32>>, n: usize,
1202 temp: f32, seed: u64, stream_pos: u32,
1203 p_stats: (f32, f32, f32), q_stats: (f32, f32, f32),
1204 out_tok: &mut CudaSlice<u32>)
1205 -> Result<(), Box<dyn std::error::Error>> {
1206 let f = self.func("residual_sample_filtered_f32");
1207 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1208 let has_q: i32 = q.is_some() as i32;
1209 let qbuf = q.unwrap_or(p);
1210 let (pm, pth, pz) = p_stats; let (qm, qth, qz) = q_stats;
1211 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
1212 let __s_b = self.gpu.stream();
1213 let mut b = __s_b.launch_builder(&f);
1214 b.arg(p).arg(qbuf).arg(&has_q).arg(&ni).arg(&temp).arg(&slo).arg(&shi).arg(&stream_pos)
1215 .arg(&pm).arg(&pth).arg(&pz).arg(&qm).arg(&qth).arg(&qz).arg(&mut *out_tok);
1216 unsafe { b.launch(cfg)?; }
1217 Ok(())
1218 }
1219
1220 #[allow(clippy::too_many_arguments)]
1222 pub fn gumbel_perturb_filtered(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
1223 seed: u64, stream_pos: u32, temp: f32, row_max: f32, th: f32)
1224 -> Result<(), Box<dyn std::error::Error>> {
1225 let f = self.func("gumbel_perturb_filtered_f32");
1226 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1227 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 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(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp).arg(&row_max).arg(&th);
1231 unsafe { b.launch(cfg)?; }
1232 Ok(())
1233 }
1234
1235 #[allow(clippy::too_many_arguments)]
1239 pub fn penalize_logits(&self, x: &mut CudaSlice<f32>, hist: &CudaSlice<u32>, n_hist: usize,
1240 rep: f32, freq: f32, present: f32, n: usize)
1241 -> Result<(), Box<dyn std::error::Error>> {
1242 if n_hist == 0 { return Ok(()); }
1243 let f = self.func("penalize_logits_f32");
1244 let (nh, ni) = (n_hist as i32, n as i32);
1245 let cfg = LaunchConfig { grid_dim: (n_hist.div_ceil(128) as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
1246 let __s_b = self.gpu.stream();
1247 let mut b = __s_b.launch_builder(&f);
1248 b.arg(&mut *x).arg(hist).arg(&nh).arg(&rep).arg(&freq).arg(&present).arg(&ni);
1249 unsafe { b.launch(cfg)?; }
1250 Ok(())
1251 }
1252
1253 #[allow(clippy::too_many_arguments)]
1255 pub fn penalize_logits_rows(&self, x: &mut CudaSlice<f32>, hist: &CudaSlice<u32>, n_hist: usize,
1256 rep: f32, freq: f32, present: f32, n: usize, nrow: usize)
1257 -> Result<(), Box<dyn std::error::Error>> {
1258 if n_hist == 0 || nrow == 0 { return Ok(()); }
1259 let f = self.func("penalize_logits_rows_f32");
1260 let (nh, ni, nr) = (n_hist as i32, n as i32, nrow as i32);
1261 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 };
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).arg(&nr);
1265 unsafe { b.launch(cfg)?; }
1266 Ok(())
1267 }
1268
1269 pub fn wpf_level() -> u32 {
1277 static ON: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
1278 *ON.get_or_init(|| std::env::var("MEMRA_WPF").ok()
1279 .and_then(|v| v.parse().ok()).unwrap_or(1))
1280 }
1281
1282 pub fn set_verify_exact(&self, on: bool) {
1294 self.verify_exact.store(on, std::sync::atomic::Ordering::Relaxed);
1295 }
1296 pub(crate) fn verify_exact_on(&self) -> bool {
1297 self.verify_exact.load(std::sync::atomic::Ordering::Relaxed)
1298 }
1299
1300 pub fn qkv_append_on() -> bool {
1303 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1304 *ON.get_or_init(|| std::env::var("MEMRA_QKV_APPEND").map(|v| v != "0").unwrap_or(true))
1305 }
1306
1307 pub fn pdl_wb_on() -> bool {
1310 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1311 *ON.get_or_init(|| std::env::var("MEMRA_PDL_WB").map(|v| v != "0").unwrap_or(true))
1312 }
1313
1314 pub fn pdl_mmvq_on() -> bool {
1318 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1319 *ON.get_or_init(|| std::env::var("MEMRA_PDL_MMVQ").map(|v| v != "0").unwrap_or(true))
1320 }
1321
1322 pub fn pdl_on() -> bool {
1323 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1324 *ON.get_or_init(|| std::env::var("MEMRA_PDL").map(|v| v != "0").unwrap_or(true))
1325 }
1326
1327 fn q40_mr1_on() -> bool {
1333 static Q40MR: std::sync::OnceLock<Option<u32>> = std::sync::OnceLock::new();
1334 match *Q40MR.get_or_init(|| std::env::var("MEMRA_Q40_MR").ok()
1335 .and_then(|v| v.parse().ok())) {
1336 Some(v) => v == 1,
1337 None => crate::FUSED_MR1_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
1338 }
1339 }
1340
1341 fn pdl_func_flash(&self, g: bool, name: &'static str)
1346 -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
1347 use cudarc::driver::sys as cu;
1348 static MODS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool), usize>>> =
1355 std::sync::Mutex::new(None);
1356 static FNS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool, &'static str), usize>>> =
1357 std::sync::Mutex::new(None);
1358 let ctx_key = self.ctx().cu_ctx() as usize;
1359 if let Some(&f) = FNS.lock().unwrap().get_or_insert_with(Default::default)
1360 .get(&(ctx_key, g, name)) { return Ok(f as cu::CUfunction); }
1361 let module = {
1362 let mut mods = MODS.lock().unwrap();
1363 let map = mods.get_or_insert_with(Default::default);
1364 match map.get(&(ctx_key, g)) {
1365 Some(&m) => m,
1366 None => {
1367 let m = self.pdl_load_module_in_ctx(
1368 if g { FLASH_FATBIN_KF8VF8 } else { FLASH_FATBIN })?;
1369 map.insert((ctx_key, g), m);
1370 m
1371 }
1372 }
1373 };
1374 let cname = std::ffi::CString::new(name)?;
1375 let mut f: cu::CUfunction = std::ptr::null_mut();
1376 let r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
1377 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("pdl_func_flash {name} (g={g}): {r:?}").into()); }
1378 FNS.lock().unwrap().get_or_insert_with(Default::default)
1379 .insert((ctx_key, g, name), f as usize);
1380 Ok(f)
1381 }
1382
1383 fn pdl_load_module_in_ctx(&self, bytes: &[u8]) -> Result<usize, Box<dyn std::error::Error>> {
1388 use cudarc::driver::sys as cu;
1389 let mut prev: cu::CUcontext = std::ptr::null_mut();
1390 unsafe { cu::cuCtxGetCurrent(&mut prev).result()?; }
1391 self.ctx().bind_to_thread()?;
1392 let mut m: cu::CUmodule = std::ptr::null_mut();
1393 let r = unsafe { cu::cuModuleLoadData(&mut m, bytes.as_ptr() as *const std::ffi::c_void) };
1394 let restore = if prev.is_null() { cu::CUresult::CUDA_SUCCESS }
1395 else { unsafe { cu::cuCtxSetCurrent(prev) } };
1396 if r != cu::CUresult::CUDA_SUCCESS {
1397 return Err(format!("pdl module load: {r:?}").into());
1398 }
1399 if restore != cu::CUresult::CUDA_SUCCESS {
1400 return Err(format!("pdl module load: ctx restore {restore:?}").into());
1401 }
1402 Ok(m as usize)
1403 }
1404
1405 fn pdl_func(&self, name: &'static str) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
1406 use cudarc::driver::sys as cu;
1407 static MODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
1410 std::sync::Mutex::new(None);
1411 static QMODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
1414 std::sync::Mutex::new(None);
1415 static FNS: std::sync::Mutex<Option<std::collections::HashMap<(usize, &'static str), usize>>> =
1416 std::sync::Mutex::new(None);
1417 let ctx_key = self.ctx().cu_ctx() as usize;
1418 if let Some(&f) = FNS.lock().unwrap().get_or_insert_with(Default::default)
1419 .get(&(ctx_key, name)) { return Ok(f as cu::CUfunction); }
1420 let module = {
1421 let mut mods = MODULES.lock().unwrap();
1422 let map = mods.get_or_insert_with(Default::default);
1423 match map.get(&ctx_key) {
1424 Some(&m) => m,
1425 None => {
1426 let m = self.pdl_load_module_in_ctx(FATBIN)?;
1427 map.insert(ctx_key, m);
1428 m
1429 }
1430 }
1431 };
1432 let cname = std::ffi::CString::new(name)?;
1433 let mut f: cu::CUfunction = std::ptr::null_mut();
1434 let mut r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
1435 if r == cu::CUresult::CUDA_ERROR_NOT_FOUND {
1436 let qmodule = {
1437 let mut mods = QMODULES.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(QMATVEC_FATBIN)?;
1443 map.insert(ctx_key, m);
1444 m
1445 }
1446 }
1447 };
1448 r = unsafe { cu::cuModuleGetFunction(&mut f, qmodule as cu::CUmodule, cname.as_ptr()) };
1449 }
1450 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("pdl_func {name}: {r:?}").into()); }
1451 FNS.lock().unwrap().get_or_insert_with(Default::default)
1452 .insert((ctx_key, name), f as usize);
1453 Ok(f)
1454 }
1455
1456 unsafe fn launch_pdl_flash(&self, g: bool, name: &'static str, grid: (u32, u32, u32),
1468 block: (u32, u32, u32), smem: u32,
1469 params: &mut [*mut std::ffi::c_void])
1470 -> Result<(), Box<dyn std::error::Error>> {
1471 use cudarc::driver::sys as cu;
1472 let f = self.pdl_func_flash(g, name)?;
1473 if smem > 0 {
1474 let r = unsafe { cu::cuFuncSetAttribute(f,
1476 cu::CUfunction_attribute_enum::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
1477 smem as i32) };
1478 if r != cu::CUresult::CUDA_SUCCESS {
1479 return Err(format!("pdl smem attr {name}: {r:?}").into());
1480 }
1481 }
1482 let mut attr = cu::CUlaunchAttribute {
1483 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
1484 pad: [0; 4],
1485 value: cu::CUlaunchAttributeValue { programmaticStreamSerializationAllowed: 1 },
1486 };
1487 let cfg = cu::CUlaunchConfig {
1488 gridDimX: grid.0, gridDimY: grid.1, gridDimZ: grid.2,
1489 blockDimX: block.0, blockDimY: block.1, blockDimZ: block.2,
1490 sharedMemBytes: smem, hStream: self.gpu.stream().cu_stream(),
1491 attrs: &mut attr, numAttrs: 1,
1492 };
1493 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
1494 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("launch_pdl_flash {name}: {r:?}").into()); }
1495 Ok(())
1496 }
1497
1498 unsafe fn launch_pdl(&self, name: &'static str, grid: (u32, u32, u32), block: (u32, u32, u32),
1499 params: &mut [*mut std::ffi::c_void])
1500 -> Result<(), Box<dyn std::error::Error>> {
1501 use cudarc::driver::sys as cu;
1502 let f = self.pdl_func(name)?;
1503 let mut attr = cu::CUlaunchAttribute {
1504 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
1505 pad: [0; 4],
1506 value: cu::CUlaunchAttributeValue { programmaticStreamSerializationAllowed: 1 },
1507 };
1508 let cfg = cu::CUlaunchConfig {
1509 gridDimX: grid.0, gridDimY: grid.1, gridDimZ: grid.2,
1510 blockDimX: block.0, blockDimY: block.1, blockDimZ: block.2,
1511 sharedMemBytes: 0, hStream: self.gpu.stream().cu_stream(),
1512 attrs: &mut attr, numAttrs: 1,
1513 };
1514 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
1515 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("launch_pdl {name}: {r:?}").into()); }
1516 Ok(())
1517 }
1518
1519 pub fn prefetch_weight_l2(&self, w: &crate::model::GpuTensor)
1522 -> Result<(), Box<dyn std::error::Error>> {
1523 if let crate::model::GpuTensor::Quant { bytes, rp4, .. } = w {
1524 let p = rp4.as_ref().unwrap_or(bytes);
1525 self.prefetch_l2(p, p.len())?;
1526 }
1527 Ok(())
1528 }
1529
1530 pub fn gather_row_bf16(&self, table: &CudaSlice<u8>, tok: &CudaSlice<u32>, idx: usize,
1533 dst: &mut CudaSlice<f32>, ncols: usize)
1534 -> Result<(), Box<dyn std::error::Error>> {
1535 let f = self.func("gather_row_bf16_f32");
1536 let cfg = LaunchConfig { grid_dim: (ncols.div_ceil(256) as u32, 1, 1),
1537 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1538 let (nc, ix) = (ncols as i32, idx as i32);
1539 let __s_b = self.gpu.stream();
1540 let mut b = __s_b.launch_builder(&f);
1541 b.arg(table).arg(tok).arg(&ix).arg(dst).arg(&nc);
1542 unsafe { b.launch(cfg)?; }
1543 Ok(())
1544 }
1545
1546 pub fn add_row_inplace(&self, logits: &mut CudaSlice<f32>, bias: &CudaSlice<f32>,
1548 n: usize, row_off: usize)
1549 -> Result<(), Box<dyn std::error::Error>> {
1550 let f = self.func("add_row_inplace_f32");
1551 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1),
1552 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1553 let (ni, off) = (n as i32, row_off as i64);
1554 let __s_b = self.gpu.stream();
1555 let mut b = __s_b.launch_builder(&f);
1556 b.arg(logits).arg(bias).arg(&ni).arg(&off);
1557 unsafe { b.launch(cfg)?; }
1558 Ok(())
1559 }
1560
1561 pub fn prefetch_l2(&self, p: &CudaSlice<u8>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
1563 let f = self.func("prefetch_l2_bytes");
1564 let lines = n.div_ceil(128);
1565 let ni = n as i64;
1566 let cfg = LaunchConfig { grid_dim: (lines.div_ceil(256) as u32, 1, 1),
1567 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1568 let __s_b = self.gpu.stream();
1569 let mut b = __s_b.launch_builder(&f);
1570 b.arg(p).arg(&ni);
1571 unsafe { b.launch(cfg)?; }
1572 Ok(())
1573 }
1574
1575 pub fn router_gemv(&self, w: &CudaSlice<f32>, x: &CudaSlice<f32>, n_embd: usize,
1578 n_experts: usize, t: usize)
1579 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1580 let w8 = match std::env::var("MEMRA_ROUTER_V2").as_deref() {
1586 Ok("0") => false,
1587 Ok(_) => true,
1588 Err(_) => ROUTER_W8_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
1589 };
1590 let batch = w8 && t >= ROUTER_BATCH_MIN_T && router_batch_on();
1600 self.router_gemv_form(w, x, n_embd, n_experts, t, w8, batch)
1601 }
1602
1603 pub fn router_gemv_form(&self, w: &CudaSlice<f32>, x: &CudaSlice<f32>, n_embd: usize,
1606 n_experts: usize, t: usize, w8: bool, batch: bool)
1607 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1608 debug_assert!(!batch || w8, "batch twin exists for the w8 form only");
1609 let mut y = self.alloc_uninit::<f32>(t * n_experts)?;
1610 let f = if batch { self.func("router_gemv_f32_w8_batch") }
1611 else if w8 { self.func("router_gemv_f32_w8") }
1612 else { self.func("router_gemv_f32") };
1613 let (ne, nx, ti) = (n_embd as i32, n_experts as i32, t as i32);
1614 let cfg = if batch {
1615 LaunchConfig { grid_dim: (n_experts.div_ceil(8) as u32, t.div_ceil(8) as u32, 1),
1616 block_dim: (32, 8, 1), shared_mem_bytes: 0 }
1617 } else {
1618 LaunchConfig { grid_dim: (n_experts as u32, t as u32, 1),
1619 block_dim: (32, if w8 { 8 } else { 1 }, 1), shared_mem_bytes: 0 }
1620 };
1621 let __s_b = self.gpu.stream();
1622 let mut b = __s_b.launch_builder(&f);
1623 b.arg(w).arg(x).arg(&mut y).arg(&ne).arg(&nx).arg(&ti);
1624 unsafe { b.launch(cfg)?; }
1625 Ok(y)
1626 }
1627
1628 pub fn rows_permute(&self, src: &CudaSlice<f32>, idx: &CudaSlice<i32>, nrows: usize,
1630 ncols: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1631 let mut dst = self.alloc_uninit::<f32>(nrows * ncols)?;
1632 let f = self.func("rows_permute_f32");
1633 let (nc, nr) = (ncols as i32, nrows as i32);
1634 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (256, 1, 1),
1635 shared_mem_bytes: 0 };
1636 let __s_b = self.gpu.stream();
1637 let mut b = __s_b.launch_builder(&f);
1638 b.arg(src).arg(idx).arg(&mut dst).arg(&nc).arg(&nr);
1639 unsafe { b.launch(cfg)?; }
1640 Ok(dst)
1641 }
1642
1643 pub fn sigmoid_dot_rows(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, n_embd: usize,
1648 t: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1649 static OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1652 if *OFF.get_or_init(|| std::env::var("MEMRA_SHEXP_DOT").as_deref() == Ok("0")) {
1653 let gs = self.linear(x, w, t, n_embd, 1)?;
1654 let mut g = self.uninit(t)?;
1655 self.sigmoid(&gs, &mut g, t)?;
1656 return Ok(g);
1657 }
1658 let mut g = self.alloc_uninit::<f32>(t)?;
1664 let f = self.func("sigmoid_dot_rows_f32");
1665 let (ne, ti) = (n_embd as i32, t as i32);
1666 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (32, 8, 1),
1667 shared_mem_bytes: 0 };
1668 let __s_b = self.gpu.stream();
1669 let mut b = __s_b.launch_builder(&f);
1670 b.arg(x).arg(w).arg(&mut g).arg(&ne).arg(&ti);
1671 unsafe { b.launch(cfg)?; }
1672 Ok(g)
1673 }
1674
1675 pub fn spec_rollback_stream(&self, len_ptrs: &CudaSlice<u64>, pos_start: &CudaSlice<i32>,
1677 acc: &CudaSlice<u32>, base: usize, n_rows: usize)
1678 -> Result<(), Box<dyn std::error::Error>> {
1679 let f = self.func("spec_rollback_stream");
1680 let (b, nr) = (base as i32, n_rows as i32);
1681 let cfg = LaunchConfig { grid_dim: (n_rows.div_ceil(64) as u32, 1, 1),
1682 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1683 let __s_bl = self.gpu.stream();
1684 let mut bl = __s_bl.launch_builder(&f);
1685 bl.arg(len_ptrs).arg(pos_start).arg(acc).arg(&b).arg(&nr);
1686 unsafe { bl.launch(cfg)?; }
1687 Ok(())
1688 }
1689
1690 pub fn plain_tok_ring(&self, vam: &CudaSlice<u32>, pos_start: &CudaSlice<i32>,
1692 base: usize, ring: &mut CudaSlice<u32>)
1693 -> Result<(), Box<dyn std::error::Error>> {
1694 let f = self.func("plain_tok_ring");
1695 let (b, cap) = (base as i32, ring.len() as i32);
1696 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1697 let __s_bl = self.gpu.stream();
1698 let mut bl = __s_bl.launch_builder(&f);
1699 bl.arg(vam).arg(pos_start).arg(&b).arg(&mut *ring).arg(&cap);
1700 unsafe { bl.launch(cfg)?; }
1701 Ok(())
1702 }
1703
1704 pub fn spec_ring_commit(&self, vtok: &CudaSlice<u32>, acc: &CudaSlice<u32>,
1706 brk: &CudaSlice<u32>, ring: &mut CudaSlice<u32>,
1707 pend: &mut CudaSlice<u32>)
1708 -> Result<(), Box<dyn std::error::Error>> {
1709 let f = self.func("spec_ring_commit");
1710 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1711 let __s_b = self.gpu.stream();
1712 let mut b = __s_b.launch_builder(&f);
1713 b.arg(vtok).arg(acc).arg(brk).arg(ring).arg(pend);
1714 unsafe { b.launch(cfg)?; }
1715 Ok(())
1716 }
1717 pub fn i32_copy_add(&self, src: &CudaSlice<i32>, dst: &mut CudaSlice<i32>, delta: i32)
1718 -> Result<(), Box<dyn std::error::Error>> {
1719 let f = self.func("i32_copy_add");
1720 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1721 let __s_b = self.gpu.stream();
1722 let mut b = __s_b.launch_builder(&f);
1723 b.arg(src).arg(dst).arg(&delta);
1724 unsafe { b.launch(cfg)?; }
1725 Ok(())
1726 }
1727 pub fn u32_copy(&self, src: &CudaSlice<u32>, dst: &mut CudaSlice<u32>)
1728 -> Result<(), Box<dyn std::error::Error>> {
1729 let f = self.func("u32_copy");
1730 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1731 let __s_b = self.gpu.stream();
1732 let mut b = __s_b.launch_builder(&f);
1733 b.arg(src).arg(dst);
1734 unsafe { b.launch(cfg)?; }
1735 Ok(())
1736 }
1737
1738 pub fn spec_adapt_k(&self, acc: &CudaSlice<u32>, brk: &mut CudaSlice<u32>,
1742 floor: usize, cap: usize)
1743 -> Result<(), Box<dyn std::error::Error>> {
1744 let f = self.func("spec_adapt_k");
1745 let (fl, cp) = (floor as i32, cap as i32);
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(acc).arg(brk).arg(&fl).arg(&cp);
1750 unsafe { b.launch(cfg)?; }
1751 Ok(())
1752 }
1753
1754 pub fn spec_accept_greedy_dc(&self, preds: &CudaSlice<u32>, vtok: &CudaSlice<u32>,
1756 last_pred: &CudaSlice<u32>, brk: &CudaSlice<u32>,
1757 out: &mut CudaSlice<u32>)
1758 -> Result<(), Box<dyn std::error::Error>> {
1759 let f = self.func("spec_accept_greedy_dc");
1760 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1761 let __s_b = self.gpu.stream();
1762 let mut b = __s_b.launch_builder(&f);
1763 b.arg(preds).arg(vtok).arg(last_pred).arg(brk).arg(out);
1764 unsafe { b.launch(cfg)?; }
1765 Ok(())
1766 }
1767
1768 pub fn pos_iota(&self, pos0: &CudaSlice<i32>, out: &mut CudaSlice<i32>, t: usize)
1770 -> Result<(), Box<dyn std::error::Error>> {
1771 let f = self.func("pos_iota_i32");
1772 let ti = t as i32;
1773 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (t.max(1) as u32, 1, 1),
1774 shared_mem_bytes: 0 };
1775 let __s_b = self.gpu.stream();
1776 let mut b = __s_b.launch_builder(&f);
1777 b.arg(pos0).arg(out).arg(&ti);
1778 unsafe { b.launch(cfg)?; }
1779 Ok(())
1780 }
1781 #[allow(clippy::too_many_arguments)]
1782 pub fn append_kv_quantized_rows_dc(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
1783 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
1784 t0_dev: &CudaSlice<i32>, t: usize,
1785 kv_dim_k: usize, kv_dim_v: usize,
1786 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
1787 -> Result<(), Box<dyn std::error::Error>> {
1788 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_rows_dc") }
1789 else { self.func("append_quantize_kv_q8_0_q5_1_rows_dc") };
1790 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
1791 let cfg = LaunchConfig { grid_dim: (nblk, t as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1792 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
1793 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
1794 let __s_b = self.gpu.stream();
1795 let mut b = __s_b.launch_builder(&f);
1796 b.arg(k_rows).arg(v_rows).arg(kc).arg(vc).arg(t0_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
1797 unsafe { b.launch(cfg)?; }
1798 Ok(())
1799 }
1800
1801 #[allow(clippy::too_many_arguments)]
1804 pub fn append_kv_quantized_row_dc_inc(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
1805 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
1806 t0_dev: &mut CudaSlice<i32>,
1807 kv_dim_k: usize, kv_dim_v: usize,
1808 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
1809 -> Result<(), Box<dyn std::error::Error>> {
1810 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_dc_inc") }
1811 else { self.func("append_quantize_kv_q8_0_q5_1_dc_inc") };
1812 let nthreads = ((kv_dim_k.max(kv_dim_v) / 32) * 32).min(1024) as u32;
1813 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (nthreads, 1, 1),
1814 shared_mem_bytes: 0 };
1815 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
1816 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
1817 let __s_b = self.gpu.stream();
1818 let mut b = __s_b.launch_builder(&f);
1819 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(t0_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
1820 unsafe { b.launch(cfg)?; }
1821 Ok(())
1822 }
1823
1824 pub fn pack_tok_p(&self, tok: &CudaSlice<u32>, p: &CudaSlice<f32>, out: &mut CudaSlice<u32>,
1826 slot: usize) -> Result<(), Box<dyn std::error::Error>> {
1827 let f = self.func("pack_tok_p");
1828 let sl = slot as i32;
1829 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1830 let __s_b = self.gpu.stream();
1831 let mut b = __s_b.launch_builder(&f);
1832 b.arg(tok).arg(p).arg(out).arg(&sl);
1833 unsafe { b.launch(cfg)?; }
1834 Ok(())
1835 }
1836 pub fn tok_map_u32(&self, tok: &mut CudaSlice<u32>, map: &CudaSlice<u32>)
1837 -> Result<(), Box<dyn std::error::Error>> {
1838 let f = self.func("tok_map_u32");
1839 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1840 let __s_b = self.gpu.stream();
1841 let mut b = __s_b.launch_builder(&f);
1842 b.arg(tok).arg(map);
1843 unsafe { b.launch(cfg)?; }
1844 Ok(())
1845 }
1846
1847 #[allow(clippy::too_many_arguments)]
1849 pub fn spec_assemble_verify(&self, tokp: &CudaSlice<u32>, pend: &CudaSlice<u32>,
1850 d2t: Option<&CudaSlice<u32>>, vtok: &mut CudaSlice<u32>,
1851 brk: &mut CudaSlice<u32>, p_min: f32, k: usize, pmin0: bool)
1852 -> Result<(), Box<dyn std::error::Error>> {
1853 let f = self.func("spec_assemble_verify");
1854 let (ki, pm) = (k as i32, if pmin0 { 1i32 } else { 0i32 });
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 match d2t {
1859 Some(m) => { b.arg(tokp).arg(pend).arg(m).arg(vtok).arg(brk).arg(&p_min).arg(&ki).arg(&pm);
1860 unsafe { b.launch(cfg)?; } }
1861 None => { let null: u64 = 0;
1862 b.arg(tokp).arg(pend).arg(&null).arg(vtok).arg(brk).arg(&p_min).arg(&ki).arg(&pm);
1863 unsafe { b.launch(cfg)?; } }
1864 }
1865 Ok(())
1866 }
1867
1868 #[allow(clippy::too_many_arguments)]
1870 pub fn ssm_conv_ring_rebuild_dc(&self, qkv_tm: &CudaSlice<f32>, ring_old: &CudaSlice<f32>,
1871 conv_state: &mut CudaSlice<f32>, conv_dim: usize,
1872 acc: &CudaSlice<u32>, base: usize, t_v: usize, d_conv: usize)
1873 -> Result<(), Box<dyn std::error::Error>> {
1874 let f = self.func("ssm_conv_ring_rebuild_f32_dc");
1875 let n = conv_dim * (d_conv - 1);
1876 let cfg = LaunchConfig::for_num_elems(n as u32);
1877 let (cd, b0, tv, dc) = (conv_dim as i32, base as i32, t_v as i32, d_conv as i32);
1878 let __s_b = self.gpu.stream();
1879 let mut b = __s_b.launch_builder(&f);
1880 b.arg(qkv_tm).arg(ring_old).arg(conv_state).arg(&cd).arg(acc).arg(&b0).arg(&tv).arg(&dc);
1881 unsafe { b.launch(cfg)?; }
1882 Ok(())
1883 }
1884 #[allow(clippy::too_many_arguments)]
1885 pub fn gdn_scan_s128_dc(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
1886 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
1887 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
1888 n_head: usize, acc: &CudaSlice<u32>, base: usize, t_v: usize,
1889 scale: f32)
1890 -> Result<(), Box<dyn std::error::Error>> {
1891 let f = self.func("gdn_scan_s128_dc");
1892 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
1893 let cfg = LaunchConfig {
1894 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
1895 block_dim: (WARP, COLS_PER_BLOCK, 1),
1896 shared_mem_bytes: 0,
1897 };
1898 let (h, b0, tv) = (n_head as i32, base as i32, t_v as i32);
1899 let __s_b = self.gpu.stream();
1900 let mut b = __s_b.launch_builder(&f);
1901 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in).arg(state_out).arg(o)
1902 .arg(&h).arg(acc).arg(&b0).arg(&tv).arg(&scale);
1903 unsafe { b.launch(cfg)?; }
1904 Ok(())
1905 }
1906
1907 pub fn spec_rollback_kv(&self, len_ptrs: &CudaSlice<u64>, saved: &CudaSlice<i32>,
1909 acc: &CudaSlice<u32>, base: usize, n_layer: usize)
1910 -> Result<(), Box<dyn std::error::Error>> {
1911 let f = self.func("spec_rollback_kv");
1912 let (b, nl) = (base as i32, n_layer as i32);
1913 let cfg = LaunchConfig { grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
1914 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1915 let __s_bl = self.gpu.stream();
1916 let mut bl = __s_bl.launch_builder(&f);
1917 bl.arg(len_ptrs).arg(saved).arg(acc).arg(&b).arg(&nl);
1918 unsafe { bl.launch(cfg)?; }
1919 Ok(())
1920 }
1921
1922 pub fn spec_seed_gather(&self, vx: &CudaSlice<f32>, fill_prev: &CudaSlice<f32>,
1925 acc: &CudaSlice<u32>, h_seed: &mut CudaSlice<f32>,
1926 base: usize, n_embd: usize)
1927 -> Result<(), Box<dyn std::error::Error>> {
1928 let f = self.func("spec_seed_gather");
1929 let (b, ne) = (base as i32, n_embd as i32);
1930 let cfg = LaunchConfig { grid_dim: (n_embd.div_ceil(256) as u32, 1, 1),
1931 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1932 let __s_bl = self.gpu.stream();
1933 let mut bl = __s_bl.launch_builder(&f);
1934 bl.arg(vx).arg(fill_prev).arg(acc).arg(h_seed).arg(&b).arg(&ne);
1935 unsafe { bl.launch(cfg)?; }
1936 Ok(())
1937 }
1938
1939
1940 pub fn spec_accept_greedy(&self, preds: &CudaSlice<u32>, draft: &CudaSlice<u32>,
1942 last_pred: u32, base: usize, k_round: usize,
1943 out: &mut CudaSlice<u32>)
1944 -> Result<(), Box<dyn std::error::Error>> {
1945 let f = self.func("spec_accept_greedy");
1946 let (b, k) = (base as i32, k_round as i32);
1947 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1948 let __s_bl = self.gpu.stream();
1949 let mut bl = __s_bl.launch_builder(&f);
1950 bl.arg(preds).arg(draft).arg(&last_pred).arg(&b).arg(&k).arg(out);
1951 unsafe { bl.launch(cfg)?; }
1952 Ok(())
1953 }
1954
1955 pub fn gumbel_perturb(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
1962 seed: u64, stream_pos: u32, temp: f32)
1963 -> Result<(), Box<dyn std::error::Error>> {
1964 let f = self.func("gumbel_perturb_f32");
1965 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1966 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1967 let __s_b = self.gpu.stream();
1968 let mut b = __s_b.launch_builder(&f);
1969 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
1970 unsafe { b.launch(cfg)?; }
1971 Ok(())
1972 }
1973
1974 pub fn mask_logits_col(&self, logits: &mut CudaSlice<f32>, mask: &CudaSlice<u32>,
1982 col: usize, n: usize, mask_words: usize)
1983 -> Result<(), Box<dyn std::error::Error>> {
1984 let f = self.func("mask_logits_f32");
1985 let (ci, ni, mw) = (col as i32, n as i32, mask_words as i32);
1986 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256).min(1024) as u32, 1, 1),
1987 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1988 let __s_b = self.gpu.stream();
1989 let mut b = __s_b.launch_builder(&f);
1990 b.arg(&mut *logits).arg(mask).arg(&ci).arg(&ni).arg(&mw);
1991 unsafe { b.launch(cfg)?; }
1992 Ok(())
1993 }
1994
1995 pub fn gumbel_perturb_col(&self, x: &CudaSlice<f32>, col: usize, y: &mut CudaSlice<f32>,
2002 n: usize, seed: u64, stream_pos: u32, temp: f32)
2003 -> Result<(), Box<dyn std::error::Error>> {
2004 let f = self.func("gumbel_perturb_f32");
2005 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2006 let col_view = x.slice(col * n..(col + 1) * n);
2007 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2008 let __s_b = self.gpu.stream();
2009 let mut b = __s_b.launch_builder(&f);
2010 b.arg(&col_view).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
2011 unsafe { b.launch(cfg)?; }
2012 Ok(())
2013 }
2014
2015 pub fn sctr_inc(&self, ctr: &mut CudaSlice<u32>) -> Result<(), Box<dyn std::error::Error>> {
2020 let f = self.func("memra_sctr_inc");
2021 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
2022 let __s_b = self.gpu.stream();
2023 let mut b = __s_b.launch_builder(&f);
2024 b.arg(&mut *ctr);
2025 unsafe { b.launch(cfg)?; }
2026 Ok(())
2027 }
2028
2029 pub fn gumbel_perturb_ctr(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
2034 seed: u64, ctr: &CudaSlice<u32>, temp: f32)
2035 -> Result<(), Box<dyn std::error::Error>> {
2036 let f = self.func("gumbel_perturb_ctr_f32");
2037 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2038 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2039 let __s_b = self.gpu.stream();
2040 let mut b = __s_b.launch_builder(&f);
2041 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(ctr).arg(&temp);
2042 unsafe { b.launch(cfg)?; }
2043 Ok(())
2044 }
2045
2046 pub fn softmax_gather(&self, x: &CudaSlice<f32>, row_stride: usize,
2050 ids: &CudaSlice<u32>, rows: &CudaSlice<i32>,
2051 out: &mut CudaSlice<f32>, n: usize, npair: usize, temp: f32)
2052 -> Result<(), Box<dyn std::error::Error>> {
2053 let f = self.func("softmax_gather_f32");
2054 let (ni, rs) = (n as i32, row_stride as i64);
2055 let np = npair as i32;
2056 let cfg = LaunchConfig { grid_dim: (npair as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2057 let __s_b = self.gpu.stream();
2058 let mut b = __s_b.launch_builder(&f);
2059 b.arg(x).arg(&rs).arg(ids).arg(rows).arg(&mut *out).arg(&ni).arg(&np).arg(&temp);
2060 unsafe { b.launch(cfg)?; }
2061 Ok(())
2062 }
2063
2064 pub fn residual_sample(&self, p: &CudaSlice<f32>, q: Option<&CudaSlice<f32>>, n: usize,
2068 temp: f32, seed: u64, stream_pos: u32,
2069 out_tok: &mut CudaSlice<u32>)
2070 -> Result<(), Box<dyn std::error::Error>> {
2071 let f = self.func("residual_sample_f32");
2072 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2073 let nth = 1024u32;
2074 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (nth, 1, 1), shared_mem_bytes: 0 };
2075 let has_q: i32 = q.is_some() as i32;
2076 let qbuf = q.unwrap_or(p); let __s_b = self.gpu.stream();
2078 let mut b = __s_b.launch_builder(&f);
2079 b.arg(p).arg(qbuf).arg(&has_q).arg(&ni).arg(&temp).arg(&slo).arg(&shi).arg(&stream_pos)
2080 .arg(&mut *out_tok);
2081 unsafe { b.launch(cfg)?; }
2082 Ok(())
2083 }
2084
2085 pub fn with_moe_cache<R>(&self, max_block_bytes: usize,
2090 f: impl FnOnce(&mut crate::moe_cache::MoeSlotCache, &Engine) -> Result<R, Box<dyn std::error::Error>>)
2091 -> Result<R, Box<dyn std::error::Error>> {
2092 let mut guard = self.moe_cache.lock().unwrap();
2093 if guard.is_none() {
2094 *guard = Some(crate::moe_cache::MoeSlotCache::new(self, max_block_bytes)?);
2095 }
2096 let cache = guard.as_mut().unwrap();
2097 f(cache, self)
2098 }
2099
2100 pub fn freeze_moe_cache(&self) {
2103 if let Some(cache) = self.moe_cache.lock().unwrap().as_mut() {
2104 cache.freeze();
2105 }
2106 }
2107
2108 pub fn export_moe_residency(&self) -> Option<Vec<(u16, u8, u16)>> {
2111 self.moe_cache
2112 .lock()
2113 .unwrap()
2114 .as_ref()
2115 .map(crate::moe_cache::MoeSlotCache::export_residency)
2116 }
2117
2118 pub(crate) fn moe_cache_frozen(&self) -> bool {
2119 self.moe_cache
2120 .lock()
2121 .unwrap()
2122 .as_ref()
2123 .is_some_and(crate::moe_cache::MoeSlotCache::is_frozen)
2124 }
2125
2126 pub fn frozen_cpu_experts_prefer_tokenwise_prime(&self) -> bool {
2133 crate::cpu_experts::configured()
2134 && self.moe_cache_frozen()
2135 && std::env::var("MEMRA_CPU_EXPERT_BATCHED_PRIME").as_deref() != Ok("1")
2136 }
2137
2138 pub(crate) fn configure_moe_cache_layout(&self, block_bytes: Vec<usize>) {
2140 assert!(
2141 self.moe_cache.lock().unwrap().is_none(),
2142 "MoE cache layout configured after cache construction"
2143 );
2144 *self.moe_cache_layout.lock().unwrap() = Some(block_bytes);
2145 }
2146
2147 pub(crate) fn moe_cache_layout(&self) -> Option<Vec<usize>> {
2148 self.moe_cache_layout.lock().unwrap().clone()
2149 }
2150
2151 pub fn moe_cache_enabled() -> bool {
2153 std::env::var("MEMRA_MOE_CACHE").as_deref() != Ok("0")
2154 }
2155
2156 pub fn moe_cache_stats(&self) -> Option<(u64, u64, u64, usize)> {
2159 let guard = self.moe_cache.lock().unwrap();
2160 guard.as_ref() .map(|c| (c.hits, c.misses, c.staged_bytes, c.n_slots()))
2161 }
2162
2163 pub fn cpu_expert_stats(
2167 &self,
2168 ) -> Option<(u64, u64, u64, u64, u64, u64, u64, u64, u64, u64, u64)> {
2169 crate::cpu_experts::configured().then(crate::cpu_experts::stats)
2170 }
2171
2172 pub fn cpu_expert_predictor_stats(&self) -> (u64, u64) {
2175 crate::cpu_experts::predictor_stats()
2176 }
2177
2178 pub fn cpu_expert_exposed_wait_ns(&self) -> Option<u64> {
2179 crate::cpu_experts::configured().then(crate::cpu_experts::exposed_wait_ns)
2180 }
2181
2182 pub fn cpu_expert_gpu_residency_stats(&self) -> Option<(u64, u64, u64)> {
2185 crate::cpu_experts::configured().then(crate::cpu_experts::incomplete_gpu_residency_stats)
2186 }
2187
2188 pub fn moe_pread_stats(&self) -> Option<(u64, u64, u64, u64, u64, u64, u64)> {
2191
2192 let guard = self.moe_cache.lock().unwrap();
2193 guard.as_ref().and_then(|cache| cache.pread_stats()).map(|stats| (
2194 stats.reads,
2195 stats.bytes,
2196 stats.read_errors,
2197 stats.short_reads,
2198 stats.fallbacks,
2199 stats.buffer_waits,
2200 stats.ring_full,
2201 ))
2202 }
2203
2204 pub fn moe_cache_reset_counters(&self) {
2206 if let Some(c) = self.moe_cache.lock().unwrap().as_mut() { c.reset_counters(); }
2207 }
2208
2209 pub fn htod_bytes(&self, v: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2210 Ok(self.gpu.stream().clone_htod(v)?)
2211 }
2212
2213 pub fn htod_bytes_padded(&self, v: &[u8], pad: usize)
2217 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2218 let mut d = self.alloc_u8_uninit(v.len() + pad)?;
2219 {
2220 let mut view = d.slice_mut(0..v.len());
2221 self.gpu.stream().memcpy_htod(v, &mut view)?;
2222 }
2223 Ok(d)
2224 }
2225
2226 pub fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
2228 -> Result<(), Box<dyn std::error::Error>> {
2229 let mut view = dst.slice_mut(off..off + len);
2230 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2231 Ok(())
2232 }
2233
2234 pub fn copy_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &CudaSlice<u8>, len: usize)
2237 -> Result<(), Box<dyn std::error::Error>> {
2238 let mut view = dst.slice_mut(off..off + len);
2239 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2240 Ok(())
2241 }
2242
2243 pub fn htod_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &[u8])
2246 -> Result<(), Box<dyn std::error::Error>> {
2247 let mut view = dst.slice_mut(off..off + src.len());
2248 self.gpu.stream().memcpy_htod(src, &mut view)?;
2249 Ok(())
2250 }
2251
2252 pub fn view<'a>(&self, b: &'a CudaSlice<f32>, len: usize) -> cudarc::driver::CudaView<'a, f32> {
2253 b.slice(0..len)
2254 }
2255
2256 pub fn view_u8_range<'a>(&self, b: &'a CudaSlice<u8>, start: usize, end: usize)
2259 -> cudarc::driver::CudaView<'a, u8> {
2260 b.slice(start..end)
2261 }
2262 pub fn view_u8<'a>(&self, b: &'a CudaSlice<u8>, len: usize) -> cudarc::driver::CudaView<'a, u8> {
2263 b.slice(0..len)
2264 }
2265
2266 pub fn append_kv_quantized(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2270 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2271 kv_dim_k: usize, kv_dim_v: usize,
2272 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2273 -> Result<(), Box<dyn std::error::Error>> {
2274 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") } else { self.func("append_quantize_kv_q8_0_q5_1") };
2275 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2276 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2277 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2278 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2279 let __s_b = self.gpu.stream();
2280 let mut b = __s_b.launch_builder(&f);
2281 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2282 unsafe { b.launch(cfg)?; }
2283 Ok(())
2284 }
2285
2286 pub fn append_kv_quantized_dc(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2290 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t_dev: &CudaSlice<i32>,
2291 kv_dim_k: usize, kv_dim_v: usize,
2292 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2293 -> Result<(), Box<dyn std::error::Error>> {
2294 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2295 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
2296 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2297 if Self::pdl_on() && Self::pdl_wb_on() {
2299 use cudarc::driver::{DevicePtr, DevicePtrMut};
2300 let s = &self.gpu.stream();
2301 let (pk, _g0) = k_row.device_ptr(s); let (pv, _g1) = v_row.device_ptr(s);
2302 let (pkc, _g2) = kc.device_ptr_mut(s); let (pvc, _g3) = vc.device_ptr_mut(s);
2303 let (pt, _g4) = t_dev.device_ptr(s);
2304 let mut ps = [
2305 &pk as *const _ as *mut std::ffi::c_void, &pv as *const _ as *mut _,
2306 &pkc as *const _ as *mut _, &pvc as *const _ as *mut _,
2307 &pt as *const _ as *mut _, &kdk as *const _ as *mut _,
2308 &kdv as *const _ as *mut _, &ktb as *const _ as *mut _,
2309 &vtb as *const _ as *mut _,
2310 ];
2311 unsafe { self.launch_pdl_flash(g, "append_quantize_kv_q8_0_q5_1_dc",
2312 (nblk, 1, 1), (32, 1, 1), 0, &mut ps)?; }
2313 return Ok(());
2314 }
2315 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") };
2316 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2317 let __s_b = self.gpu.stream();
2318 let mut b = __s_b.launch_builder(&f);
2319 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(t_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2320 unsafe { b.launch(cfg)?; }
2321 Ok(())
2322 }
2323
2324 #[allow(clippy::too_many_arguments)]
2331 pub fn append_kv_quantized_rows(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
2332 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
2333 t0: usize, t: usize, kv_dim_k: usize, kv_dim_v: usize,
2334 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2335 -> Result<(), Box<dyn std::error::Error>> {
2336 if std::env::var("MEMRA_PRIME_APPEND_LOOP").is_ok() {
2337 for i in 0..t {
2338 let k_row = k_rows.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
2339 let v_row = v_rows.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
2340 self.append_kv_quantized_view(&k_row, &v_row, kc, vc, t0 + i,
2341 kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes, g)?;
2342 }
2343 return Ok(());
2344 }
2345 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") };
2346 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2347 let cfg = LaunchConfig { grid_dim: (nblk, t as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2348 let (t0i, kdk, kdv) = (t0 as i32, kv_dim_k as i32, kv_dim_v as i32);
2349 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2350 let __s_b = self.gpu.stream();
2351 let mut b = __s_b.launch_builder(&f);
2352 b.arg(k_rows).arg(v_rows).arg(kc).arg(vc).arg(&t0i).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2353 unsafe { b.launch(cfg)?; }
2354 Ok(())
2355 }
2356
2357 pub fn inc_seqlen(&self, p: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
2361 let f = self.func("inc_i32");
2362 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
2363 let __s_b = self.gpu.stream();
2364 let mut b = __s_b.launch_builder(&f);
2365 b.arg(p);
2366 unsafe { b.launch(cfg)?; }
2367 Ok(())
2368 }
2369
2370 pub fn append_kv_quantized_view(&self, k_row: &cudarc::driver::CudaView<f32>,
2373 v_row: &cudarc::driver::CudaView<f32>,
2374 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2375 kv_dim_k: usize, kv_dim_v: usize,
2376 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2377 -> Result<(), Box<dyn std::error::Error>> {
2378 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") }
2379 else { self.func("append_quantize_kv_q8_0_q5_1") };
2380 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2381 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2382 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2383 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2384 let __s_b = self.gpu.stream();
2385 let mut b = __s_b.launch_builder(&f);
2386 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2387 unsafe { b.launch(cfg)?; }
2388 Ok(())
2389 }
2390
2391 pub fn copy_view_into(&self, dst: &mut CudaSlice<f32>, off: usize,
2394 src: &cudarc::driver::CudaView<f32>, len: usize)
2395 -> Result<(), Box<dyn std::error::Error>> {
2396 let mut view = dst.slice_mut(off..off + len);
2397 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2398 Ok(())
2399 }
2400
2401 pub fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2405 let mut dst = self.gpu.stream().alloc_zeros::<f32>(src.len())?;
2406 self.gpu.stream().memcpy_dtod(src, &mut dst)?;
2407 Ok(dst)
2408 }
2409
2410 pub fn dtod_copy_view(&self, src: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<f32>)
2413 -> Result<(), Box<dyn std::error::Error>> {
2414 self.gpu.stream().memcpy_dtod(src, dst)?;
2415 Ok(())
2416 }
2417
2418 pub fn dtod_copy_view_i8(&self, src: &cudarc::driver::CudaView<i8>, dst: &mut CudaSlice<i8>)
2420 -> Result<(), Box<dyn std::error::Error>> {
2421 self.gpu.stream().memcpy_dtod(src, dst)?;
2422 Ok(())
2423 }
2424
2425 pub fn dtod_copy_into(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, offset: usize)
2427 -> Result<(), Box<dyn std::error::Error>> {
2428 let n = src.len();
2429 let mut dv = dst.slice_mut(offset..offset + n);
2430 self.gpu.stream().memcpy_dtod(src, &mut dv)?;
2431 Ok(())
2432 }
2433
2434 pub fn uninit_i8(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
2436 self.alloc_uninit::<i8>(n)
2437 }
2438
2439 pub fn qmatvec(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
2441 qtype: i32, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2442 let f = self.func("qmatvec_f32");
2443 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 };
2445 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2446 let __s_b = self.gpu.stream();
2447 let mut b = __s_b.launch_builder(&f);
2448 b.arg(w).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2449 unsafe { b.launch(cfg)?; }
2450 Ok(y)
2451 }
2452
2453 pub fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2455 let s = self.gpu.stream().alloc_zeros::<u8>(n)?;
2456 self.keep_if_capturing(&s);
2457 Ok(s)
2458 }
2459
2460 pub fn alloc_u8_uninit(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2464 let s = unsafe { self.gpu.stream().alloc::<u8>(n)? };
2465 self.keep_if_capturing(&s);
2466 Ok(s)
2467 }
2468
2469 pub fn memset_zeros_view(&self, dst: &mut cudarc::driver::CudaViewMut<f32>)
2472 -> Result<(), Box<dyn std::error::Error>> {
2473 self.gpu.stream().memset_zeros(dst)?;
2474 Ok(())
2475 }
2476
2477 pub fn stage_expert(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2483 -> Result<(), Box<dyn std::error::Error>> {
2484 let mut dst = scratch.slice_mut(off..off + host_bytes.len()); self.gpu.stream().memcpy_htod(host_bytes, &mut dst)?; Ok(())
2487 }
2488
2489 pub fn moe_router_topk(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2495 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2496 let f = self.func("moe_router_topk_f32");
2497 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),
2500 shared_mem_bytes: 0 };
2501 let (ne, nu) = (n_expert as i32, n_used as i32);
2502 let __s_b = self.gpu.stream();
2503 let mut b = __s_b.launch_builder(&f);
2504 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2505 unsafe { b.launch(cfg)?; }
2506 Ok((sel_idx, sel_w))
2507 }
2508
2509 pub fn moe_router_topk_scaled(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize,
2512 n_used: usize, ex_scale: &CudaSlice<f32>)
2513 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2514 let f = self.func("moe_router_topk_scaled_f32");
2519 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
2520 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
2521 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2522 shared_mem_bytes: 0 };
2523 let (ne, nu) = (n_expert as i32, n_used as i32);
2524 let __s_b = self.gpu.stream();
2525 let mut b = __s_b.launch_builder(&f);
2526 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu).arg(ex_scale);
2527 unsafe { b.launch(cfg)?; }
2528 Ok((sel_idx, sel_w))
2529 }
2530
2531 pub fn moe_router_topk_host(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2539 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2540 let f = self.func("moe_router_topk_f32");
2541 let n = t * n_used;
2542 let mut sel_idx = self.alloc_uninit::<i32>(n)?;
2543 let mut sel_w = self.alloc_uninit::<f32>(n)?;
2544 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2545 shared_mem_bytes: 0 };
2546 let (ne, nu) = (n_expert as i32, n_used as i32);
2547 let __s_b = self.gpu.stream();
2548 let mut b = __s_b.launch_builder(&f);
2549 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2550 unsafe { b.launch(cfg)?; }
2551 let bytes = n * 8;
2553 let mut guard = self.router_stage.lock().unwrap();
2554 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
2555 *guard = Some(PinnedStage::new(bytes.max(4096))?);
2556 }
2557 let stage = guard.as_mut().unwrap();
2558 let (si, sw) = unsafe {
2559 (std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
2560 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n))
2561 };
2562 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()))
2566 }
2567
2568 pub fn stage_expert_async(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2572 -> Result<cudarc::driver::CudaEvent, Box<dyn std::error::Error>> {
2573 let mut dst = scratch.slice_mut(off..off + host_bytes.len());
2574 self.copy_stream.memcpy_htod(host_bytes, &mut dst)?;
2575 Ok(self.copy_stream.record_event(None)?)
2576 }
2577
2578 pub fn compute_wait(&self, ev: &cudarc::driver::CudaEvent) -> Result<(), Box<dyn std::error::Error>> {
2580 self.gpu.stream().wait(ev)?;
2581 Ok(())
2582 }
2583
2584 pub fn qmatvec_view(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2589 x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize, out_f: usize,
2590 qtype: i32, row_bytes: usize)
2591 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2592 let f = self.func("qmatvec_f32");
2593 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 };
2596 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2597 let __s_b = self.gpu.stream();
2598 let mut b = __s_b.launch_builder(&f);
2599 b.arg(&wv).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2600 unsafe { b.launch(cfg)?; }
2601 Ok(y)
2602 }
2603
2604 #[allow(clippy::too_many_arguments)]
2611 pub fn moe_gate_up_silu8_q8(&self, gp: WPtr8, up: WPtr8,
2615 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2616 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2617 rb_g: usize, rb_u: usize)
2618 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2619 let f = self.func("moe_gate_up_silu8_q8");
2620 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
2621 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2622 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2623 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2624 let __s_b = self.gpu.stream();
2625 let mut b = __s_b.launch_builder(&f);
2626 b.arg(&gp).arg(&up).arg(aq).arg(ad).arg(&mut act)
2627 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2628 unsafe { b.launch(cfg)?; }
2629 Ok(act)
2630 }
2631
2632 #[allow(clippy::too_many_arguments)]
2633 pub fn moe_down8_fma_q8(&self, dp: WPtr8, w: F32x8,
2634 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
2635 dst: &mut cudarc::driver::CudaViewMut<f32>,
2636 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2637 -> Result<(), Box<dyn std::error::Error>> {
2638 let f = self.func("moe_down8_fma_q8");
2639 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2640 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2641 let (inf, outf, nu, rbi) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2642 let __s_b = self.gpu.stream();
2643 let mut b = __s_b.launch_builder(&f);
2644 b.arg(&dp).arg(&w).arg(aq2).arg(ad2).arg(dst)
2645 .arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbi);
2646 unsafe { b.launch(cfg)?; }
2647 Ok(())
2648 }
2649
2650 pub fn qmatvec_expert_q8(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2652 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
2653 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize)
2654 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2655 let f = self.func("qmatvec_expert_q8");
2656 let wv = w.slice(range);
2657 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
2658 const ROWS: u32 = 4; let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, m as u32, 1),
2660 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2661 let (inf, outf, mi, rbi) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
2662 let __s_b = self.gpu.stream();
2663 let mut b = __s_b.launch_builder(&f);
2664 b.arg(&wv).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qtype).arg(&rbi);
2665 unsafe { b.launch(cfg)?; }
2666 Ok(y)
2667 }
2668
2669 pub fn moe_gate_up_silu8(&self, gp: WPtr8, up: WPtr8, x: &cudarc::driver::CudaView<f32>,
2670 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2671 rb_g: usize, rb_u: usize)
2672 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2673 let f = self.func("moe_gate_up_silu8_f32");
2674 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),
2676 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2677 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2678 let __s_b = self.gpu.stream();
2679 let mut b = __s_b.launch_builder(&f);
2680 b.arg(&gp).arg(&up).arg(x).arg(&mut act)
2681 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2682 unsafe { b.launch(cfg)?; }
2683 Ok(act)
2684 }
2685
2686 #[allow(clippy::too_many_arguments)]
2692 pub fn moe_down8_fma_into(&self, dp: WPtr8, w: F32x8, act: &CudaSlice<f32>,
2693 dst: &mut cudarc::driver::CudaViewMut<f32>,
2694 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2695 -> Result<(), Box<dyn std::error::Error>> {
2696 let f = self.func("moe_down8_fma_f32");
2697 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2698 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2699 let (inf, outf, nu, rbv) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2700 let __s_b = self.gpu.stream();
2701 let mut b = __s_b.launch_builder(&f);
2702 b.arg(&dp).arg(&w).arg(act).arg(dst).arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbv);
2703 unsafe { b.launch(cfg)?; }
2704 Ok(())
2705 }
2706
2707 #[allow(clippy::too_many_arguments)]
2712 #[allow(clippy::too_many_arguments)]
2727 #[allow(clippy::too_many_arguments)]
2729 pub fn moe_pairs_matvec_q8(&self, table: &CudaSlice<u64>, proj: i32,
2730 pair_tok: &CudaSlice<i32>, pair_ex: &CudaSlice<i32>,
2731 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2732 in_f: usize, out_f: usize, n_expert: usize, n_pairs: usize,
2733 qtype: i32, row_bytes: usize)
2734 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2735 let f = self.func("moe_pairs_matvec_q8");
2736 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2737 const ROWS: u32 = 4;
2738 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_pairs as u32, 1),
2739 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2740 let (inf, outf, ne, np, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2741 n_pairs as i32, row_bytes as i64);
2742 let __s_b = self.gpu.stream();
2743 let mut b = __s_b.launch_builder(&f);
2744 b.arg(table).arg(&proj).arg(pair_tok).arg(pair_ex).arg(aq).arg(ad).arg(&mut y)
2745 .arg(&inf).arg(&outf).arg(&ne).arg(&np).arg(&qtype).arg(&rbi);
2746 unsafe { b.launch(cfg)?; }
2747 Ok(y)
2748 }
2749
2750 #[allow(clippy::too_many_arguments)]
2752 pub fn moe_pairs_matvec_q8_em(&self, table: &CudaSlice<u64>, proj: i32,
2753 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2754 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2755 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2756 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2757 n_pairs: usize, qtype: i32, row_bytes: usize)
2758 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2759 let f = self.func("moe_pairs_matvec_q8_em");
2760 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2761 const ROWS: u32 = 4;
2762 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
2763 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2764 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2765 n_active as i32, row_bytes as i64);
2766 let __s_b = self.gpu.stream();
2767 let mut b = __s_b.launch_builder(&f);
2768 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
2769 .arg(aq).arg(ad).arg(&mut y)
2770 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
2771 unsafe { b.launch(cfg)?; }
2772 Ok(y)
2773 }
2774
2775 #[allow(clippy::too_many_arguments)]
2778 pub fn moe_pairs_matvec_q8_dec(&self, table: &CudaSlice<u64>, proj: i32,
2779 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2780 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2781 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2782 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2783 n_pairs: usize, qtype: i32, row_bytes: usize)
2784 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2785 let f = self.func("moe_pairs_matvec_q8_dec");
2786 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2787 const ROWS: u32 = 4;
2788 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
2789 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2790 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2791 n_active as i32, row_bytes as i64);
2792 let __s_b = self.gpu.stream();
2793 let mut b = __s_b.launch_builder(&f);
2794 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
2795 .arg(aq).arg(ad).arg(&mut y)
2796 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
2797 unsafe { b.launch(cfg)?; }
2798 Ok(y)
2799 }
2800
2801 pub fn moe_pairs_gelu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
2802 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2803 let f = self.func("moe_pairs_gelu_mul");
2804 let mut act = self.alloc_uninit::<f32>(n)?;
2805 let cfg = LaunchConfig::for_num_elems(n as u32);
2806 let nl = n as i64;
2807 let __s_b = self.gpu.stream();
2808 let mut b = __s_b.launch_builder(&f);
2809 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
2810 unsafe { b.launch(cfg)?; }
2811 Ok(act)
2812 }
2813
2814 pub fn moe_pairs_silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
2815 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2816 let f = self.func("moe_pairs_silu_mul");
2817 let mut act = self.alloc_uninit::<f32>(n)?;
2818 let cfg = LaunchConfig::for_num_elems(n as u32);
2819 let nl = n as i64;
2820 let __s_b = self.gpu.stream();
2821 let mut b = __s_b.launch_builder(&f);
2822 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
2823 unsafe { b.launch(cfg)?; }
2824 Ok(act)
2825 }
2826
2827 #[allow(clippy::too_many_arguments)]
2828 pub fn moe_pairs_scatter(&self, y_down: &CudaSlice<f32>, pair_w: &CudaSlice<f32>,
2829 tok_pair_off: &CudaSlice<i32>, tok_pair_ids: &CudaSlice<i32>,
2830 moe_out: &mut CudaSlice<f32>, t: usize, n_embd: usize)
2831 -> Result<(), Box<dyn std::error::Error>> {
2832 let f = self.func("moe_pairs_scatter");
2833 let cfg = LaunchConfig { grid_dim: (((n_embd + 255) / 256) as u32, t as u32, 1),
2834 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2835 let ne = n_embd as i32;
2836 let __s_b = self.gpu.stream();
2837 let mut b = __s_b.launch_builder(&f);
2838 b.arg(y_down).arg(pair_w).arg(tok_pair_off).arg(tok_pair_ids).arg(moe_out).arg(&ne);
2839 unsafe { b.launch(cfg)?; }
2840 Ok(())
2841 }
2842
2843 #[allow(clippy::too_many_arguments)]
2847 pub fn moe_gate_up_gelu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
2848 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2849 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
2850 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
2851 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2852 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
2853 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
2854 rb_g as i64, rb_u as i64);
2855 let f = self.func("moe_gate_up_gelu8_dev_q8");
2856 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2857 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2858 let __s_b = self.gpu.stream();
2859 let mut b = __s_b.launch_builder(&f);
2860 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
2861 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2862 unsafe { b.launch(cfg)?; }
2863 Ok(act)
2864 }
2865
2866 #[allow(clippy::too_many_arguments)]
2868 pub fn moe_gate_up_gelu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
2869 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
2870 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
2871 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
2872 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2873 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
2874 let (inf, nff, ne, rbg, rbu, nu) = (in_f as i32, n_ff as i32, n_expert as i32,
2875 rb_g as i64, rb_u as i64, n_used as i32);
2876 let f = self.func("moe_gate_up_gelu8_dev_q8_rows");
2877 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
2878 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2879 let __s_b = self.gpu.stream();
2880 let mut b = __s_b.launch_builder(&f);
2881 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
2882 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu);
2883 unsafe { b.launch(cfg)?; }
2884 Ok(act)
2885 }
2886
2887 #[allow(clippy::too_many_arguments)]
2889 pub fn moe_gate_up_gelu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
2890 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, n_pairs: usize,
2891 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
2892 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
2893 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2894 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
2895 let (inf, nff, ne, rbg, rbu, nu, npi) = (in_f as i32, n_ff as i32, n_expert as i32,
2896 rb_g as i64, rb_u as i64, n_used as i32,
2897 n_pairs as i32);
2898 let f = self.func("moe_gate_up_gelu8_dev_q8_csr");
2899 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
2900 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2901 let __s_b = self.gpu.stream();
2902 let mut b = __s_b.launch_builder(&f);
2903 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
2904 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
2905 unsafe { b.launch(cfg)?; }
2906 Ok(act)
2907 }
2908
2909 #[allow(clippy::too_many_arguments)]
2911 pub fn moe_down8_fma_dev_q8_rows_g(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
2912 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
2913 dst: &mut CudaSlice<f32>, t: usize,
2914 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
2915 qt: i32, rb: usize)
2916 -> Result<(), Box<dyn std::error::Error>> {
2917 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
2918 n_expert as i32, rb as i64);
2919 let f = self.func("moe_down8_fma_dev_q8_rows_g");
2920 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, t as u32),
2921 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2922 let __s_b = self.gpu.stream();
2923 let mut b = __s_b.launch_builder(&f);
2924 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
2925 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
2926 unsafe { b.launch(cfg)?; }
2927 Ok(())
2928 }
2929
2930 pub fn rp_probe_q4(&self, m: usize) -> Result<(f64, f64), Box<dyn std::error::Error>> {
2934 let (out_f, in_f) = (2048usize, 2816usize);
2935 let nblk = in_f / 32;
2936 let mut seed = 0x9E3779B97F4A7C15u64;
2937 let mut rng = move || { seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407); (seed >> 33) as u8 };
2938 let mut w = vec![0u8; out_f * nblk * 18];
2939 for b in w.iter_mut() { *b = rng(); }
2940 for r in 0..out_f {
2941 for g in 0..nblk {
2942 let off = (r * nblk + g) * 18;
2943 w[off] = 0x00; w[off + 1] = 0x2C; }
2945 }
2946 let qplane = out_f * nblk * 16;
2947 let mut wrp = vec![0u8; w.len()];
2948 for r in 0..out_f {
2949 for g in 0..nblk {
2950 let src = &w[(r * nblk + g) * 18..(r * nblk + g) * 18 + 18];
2951 wrp[qplane + (r * nblk + g) * 2..qplane + (r * nblk + g) * 2 + 2]
2952 .copy_from_slice(&src[0..2]);
2953 wrp[(r * nblk + g) * 16..(r * nblk + g) * 16 + 16].copy_from_slice(&src[2..18]);
2954 }
2955 }
2956 let w_d = self.htod_bytes(&w)?;
2957 let wrp_d = self.htod_bytes(&wrp)?;
2958 let mut aq = vec![0i8; m * in_f];
2959 for v in aq.iter_mut() { *v = rng() as i8; }
2960 let aq_d = self.htod_i8(&aq)?;
2961 let ad_d = self.htod(&vec![0.03125f32; m * nblk])?;
2962 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
2963 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
2964 const RPB: u32 = 4;
2965 let cfg = LaunchConfig { grid_dim: ((out_f as u32).div_ceil(RPB), 1, 1),
2966 block_dim: (32, RPB, 1), shared_mem_bytes: 0 };
2967 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
2968 let (rb, qp) = ((nblk * 18) as i64, qplane as i64);
2969 let fb = self.func("qmatvec_q4_0_mmvq_b4");
2970 let fr = self.func("qmatvec_q4_0_mmvq_b4_rp");
2971 {
2972 let __s_b = self.gpu.stream();
2973 let mut b = __s_b.launch_builder(&fb);
2974 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
2975 unsafe { b.launch(cfg)?; }
2976 let __s_b = self.gpu.stream();
2977 let mut b = __s_b.launch_builder(&fr);
2978 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1).arg(&inf).arg(&outf).arg(&mi).arg(&qp);
2979 unsafe { b.launch(cfg)?; }
2980 }
2981 self.gpu.stream().synchronize()?;
2982 let (h0, h1) = (self.dtoh(&y0)?, self.dtoh(&y1)?);
2983 let nd = h0.iter().zip(&h1).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
2984 if nd != 0 { return Err(format!("rp twin not bitwise: {nd}/{} diffs", h0.len()).into()); }
2985 let mut time = |rp: bool| -> Result<f64, Box<dyn std::error::Error>> {
2986 self.gpu.stream().synchronize()?;
2987 let t0 = std::time::Instant::now();
2988 for _ in 0..500 {
2989 if rp {
2990 let __s_b = self.gpu.stream();
2991 let mut b = __s_b.launch_builder(&fr);
2992 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1)
2993 .arg(&inf).arg(&outf).arg(&mi).arg(&qp);
2994 unsafe { b.launch(cfg)?; }
2995 } else {
2996 let __s_b = self.gpu.stream();
2997 let mut b = __s_b.launch_builder(&fb);
2998 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0)
2999 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3000 unsafe { b.launch(cfg)?; }
3001 }
3002 }
3003 self.gpu.stream().synchronize()?;
3004 Ok(t0.elapsed().as_secs_f64() * 1e6 / 500.0)
3005 };
3006 let _ = time(false)?; let _ = time(true)?; Ok((time(false)?, time(true)?))
3008 }
3009
3010 pub fn build_q4_rp4(&self, t: &mut crate::model::GpuTensor)
3015 -> Result<(), Box<dyn std::error::Error>> {
3016 use crate::model::GpuTensor;
3017 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3018 if *qtype != QT_Q4_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3019 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3020 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 18 { return Ok(()); }
3021 let nblk = in_f / 32;
3022 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 18)?;
3023 let f = self.func("q4_0_split_rp_build");
3024 let n = (out_f * nblk) as i32;
3025 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3026 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3027 let (of, nb) = (out_f as i32, nblk as i32);
3028 let _ = n;
3029 let __s_b = self.gpu.stream();
3030 let mut b = __s_b.launch_builder(&f);
3031 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3032 unsafe { b.launch(cfg)?; }
3033 *rp4 = Some(dst);
3034 Ok(())
3035 }
3036
3037 pub fn build_q8_rp4(&self, t: &mut crate::model::GpuTensor)
3042 -> Result<(), Box<dyn std::error::Error>> {
3043 use crate::model::GpuTensor;
3044 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3045 if *qtype != QT_Q8_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3046 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3047 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 34 { return Ok(()); }
3048 *rp4 = Some(self.build_q8_rp4_raw(bytes, in_f, out_f)?);
3049 Ok(())
3050 }
3051
3052 pub fn build_q8_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize)
3055 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3056 assert!(in_f % 32 == 0);
3057 let nblk = in_f / 32;
3058 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 34)?;
3059 let f = self.func("q8_0_split_rp_build");
3060 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3061 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3062 let (of, nb) = (out_f as i32, nblk as i32);
3063 let __s_b = self.gpu.stream();
3064 let mut b = __s_b.launch_builder(&f);
3065 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3066 unsafe { b.launch(cfg)?; }
3067 Ok(dst)
3068 }
3069
3070 pub fn build_q4k_rp4(&self, t: &mut crate::model::GpuTensor)
3078 -> Result<(), Box<dyn std::error::Error>> {
3079 use crate::model::GpuTensor;
3080 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3081 if *qtype != QT_Q4_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3082 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3083 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 144 { return Ok(()); }
3084 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q4_K)?);
3085 Ok(())
3086 }
3087
3088 pub fn build_q6k_rp4(&self, t: &mut crate::model::GpuTensor)
3089 -> Result<(), Box<dyn std::error::Error>> {
3090 use crate::model::GpuTensor;
3091 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3092 if *qtype != QT_Q6_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3093 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3094 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 210 { return Ok(()); }
3095 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q6_K)?);
3096 Ok(())
3097 }
3098
3099 pub fn build_kq_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize, qtype: i32)
3101 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3102 assert!(in_f % 256 == 0);
3103 let nsbk = in_f / 256;
3104 let (sb_bytes, kname) = match qtype {
3105 QT_Q4_K => (144usize, "q4_K_split_rp_build"),
3106 QT_Q6_K => (210usize, "q6_K_split_rp_build"),
3107 _ => return Err(format!("build_kq_rp4_raw: qtype {qtype} has no rp mirror").into()),
3108 };
3109 let mut dst = self.alloc_uninit::<u8>(out_f * nsbk * sb_bytes)?;
3110 let f = self.func(kname);
3111 let cfg = LaunchConfig { grid_dim: (((out_f * nsbk) as u32).div_ceil(256), 1, 1),
3112 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3113 let (of, nb) = (out_f as i32, nsbk as i32);
3114 let __s_b = self.gpu.stream();
3115 let mut b = __s_b.launch_builder(&f);
3116 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3117 unsafe { b.launch(cfg)?; }
3118 Ok(dst)
3119 }
3120
3121 pub fn kqrp_enabled() -> bool {
3125 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3126 *ON.get_or_init(|| match std::env::var("MEMRA_KQRP").as_deref() {
3127 Ok("0") => false,
3128 Ok(_) => true,
3129 Err(_) => cfg!(memra_hopper_mma),
3130 })
3131 }
3132
3133 pub fn build_q4_rp_swap(&self, t: &mut crate::model::GpuTensor)
3139 -> Result<bool, Box<dyn std::error::Error>> {
3140 self.build_q4_rp4(t)?;
3141 self.gpu.stream().synchronize()?; use crate::model::GpuTensor;
3143 let GpuTensor::Quant { bytes, rp4, rp, .. } = t else { return Ok(false) };
3144 match rp4.take() {
3145 Some(split) => {
3146 *bytes = split; *rp = true;
3148 Ok(true)
3149 }
3150 None => Ok(false),
3151 }
3152 }
3153
3154 pub fn q4rp_enabled() -> bool {
3156 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3157 *ON.get_or_init(|| std::env::var("MEMRA_Q4RP").map(|v| v != "0").unwrap_or(true))
3158 }
3159
3160 pub fn copy_rows_strided(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
3163 row_elems: usize, n_rows: usize, src_stride: usize, src_off: usize)
3164 -> Result<(), Box<dyn std::error::Error>> {
3165 let f = self.func("copy_rows_strided_f32");
3166 let cfg = LaunchConfig { grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
3167 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3168 let (re, nr) = (row_elems as i32, n_rows as i32);
3169 let (st, off) = (src_stride as i64, src_off as i64);
3170 let __s_b = self.gpu.stream();
3171 let mut b = __s_b.launch_builder(&f);
3172 b.arg(src).arg(&mut *dst).arg(&re).arg(&nr).arg(&st).arg(&off);
3173 unsafe { b.launch(cfg)?; }
3174 Ok(())
3175 }
3176
3177 pub fn u32_set_k(&self, dst: &mut CudaSlice<u32>, v: u32, idx: usize)
3179 -> Result<(), Box<dyn std::error::Error>> {
3180 let f = self.func("u32_set_k");
3181 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3182 let ii = idx as i32;
3183 let __s_b = self.gpu.stream();
3184 let mut b = __s_b.launch_builder(&f);
3185 b.arg(dst).arg(&v).arg(&ii);
3186 unsafe { b.launch(cfg)?; }
3187 Ok(())
3188 }
3189
3190 pub fn i32_add_k(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
3192 let f = self.func("i32_add_k");
3193 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3194 let __s_b = self.gpu.stream();
3195 let mut b = __s_b.launch_builder(&f);
3196 b.arg(d).arg(&v);
3197 unsafe { b.launch(cfg)?; }
3198 Ok(())
3199 }
3200
3201 pub fn i32_iota_from(&self, ctr: &CudaSlice<i32>, dst: &mut CudaSlice<i32>, n: usize)
3203 -> Result<(), Box<dyn std::error::Error>> {
3204 let f = self.func("i32_iota_from");
3205 let cfg = LaunchConfig::for_num_elems(n as u32);
3206 let ni = n as i32;
3207 let __s_b = self.gpu.stream();
3208 let mut b = __s_b.launch_builder(&f);
3209 b.arg(ctr).arg(dst).arg(&ni);
3210 unsafe { b.launch(cfg)?; }
3211 Ok(())
3212 }
3213
3214 pub fn u32_map_k(&self, buf: &mut CudaSlice<u32>, map: &CudaSlice<u32>, idx: usize)
3216 -> Result<(), Box<dyn std::error::Error>> {
3217 let f = self.func("u32_map_k");
3218 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3219 let ii = idx as i32;
3220 let __s_b = self.gpu.stream();
3221 let mut b = __s_b.launch_builder(&f);
3222 b.arg(buf).arg(map).arg(&ii);
3223 unsafe { b.launch(cfg)?; }
3224 Ok(())
3225 }
3226
3227 #[allow(clippy::too_many_arguments)]
3229 pub fn u32_pack2(&self, a: &CudaSlice<u32>, off_a: usize, n1: usize,
3230 b_in: &CudaSlice<u32>, n2: usize, out: &mut CudaSlice<u32>)
3231 -> Result<(), Box<dyn std::error::Error>> {
3232 let f = self.func("u32_pack2");
3233 let cfg = LaunchConfig::for_num_elems((n1 + n2) as u32);
3234 let (oa, i1, i2) = (off_a as i32, n1 as i32, n2 as i32);
3235 let __s_b = self.gpu.stream();
3236 let mut b = __s_b.launch_builder(&f);
3237 b.arg(a).arg(&oa).arg(&i1).arg(b_in).arg(&i2).arg(out);
3238 unsafe { b.launch(cfg)?; }
3239 Ok(())
3240 }
3241
3242 pub fn moe_w_exscale(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3244 s: &CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
3245 let f = self.func("moe_w_exscale");
3246 let cfg = LaunchConfig::for_num_elems(n as u32);
3247 let ni = n as i32;
3248 let __s_b = self.gpu.stream();
3249 let mut b = __s_b.launch_builder(&f);
3250 b.arg(w).arg(sel).arg(s).arg(&ni);
3251 unsafe { b.launch(cfg)?; }
3252 Ok(())
3253 }
3254
3255 pub fn moe_w_scale_by_expert(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3258 macros: &CudaSlice<f32>, n_expert: usize, n: usize)
3259 -> Result<(), Box<dyn std::error::Error>> {
3260 let f = self.func("moe_w_scale_by_expert");
3261 let cfg = LaunchConfig { grid_dim: (n.div_ceil(64) as u32, 1, 1),
3262 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
3263 let (ne, nn) = (n_expert as i32, n as i32);
3264 let __s_b = self.gpu.stream();
3265 let mut b = __s_b.launch_builder(&f);
3266 b.arg(w).arg(sel).arg(macros).arg(&ne).arg(&nn);
3267 unsafe { b.launch(cfg)?; }
3268 Ok(())
3269 }
3270
3271 pub fn moe_gate_up_silu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3272 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3273 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3274 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3275 macros: &CudaSlice<f32>)
3276 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3277 static GU: std::sync::OnceLock<(String, u32)> = std::sync::OnceLock::new();
3278 let (mode, wpb) = GU.get_or_init(|| {
3279 let mode = std::env::var("MEMRA_MOE_DEVQ8_GU").unwrap_or_default();
3280 let wpb = std::env::var("MEMRA_MOE_DEVQ8_WPB").ok()
3281 .and_then(|v| v.parse().ok()).unwrap_or(4u32).clamp(1, 16);
3282 (mode, wpb)
3283 });
3284 let (mode, wpb) = (mode.as_str(), *wpb);
3285 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3286 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3287 rb_g as i64, rb_u as i64);
3288 let (f, cfg) = match mode {
3289 "1" | "2" | "4" => {
3290 let rpw: u32 = mode.parse().unwrap();
3291 let f = self.func(match rpw { 1 => "moe_gate_up_silu8_dev_q8_r1",
3292 2 => "moe_gate_up_silu8_dev_q8_r2",
3293 _ => "moe_gate_up_silu8_dev_q8_r4" });
3294 let rows_per_block = (rpw * wpb) as usize;
3295 let gx = n_ff.div_ceil(rows_per_block) as u32;
3296 (f, LaunchConfig { grid_dim: (gx, n_used as u32, 1),
3297 block_dim: (32, wpb, 1), shared_mem_bytes: 0 })
3298 }
3299 "j8" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8"),
3300 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3301 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3302 "vsm2" => {
3304 let f = self.func("moe_gate_up_silu8_dev_q8_vsm2");
3305 let sh = (rb_g + rb_u) as u32;
3306 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3307 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3308 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3309 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3310 }
3311 "vsm" => {
3312 let f = self.func("moe_gate_up_silu8_dev_q8_vsm");
3313 let sh = (rb_g + rb_u) as u32;
3314 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3315 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3316 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3317 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3318 }
3319 "sg" => (self.func("moe_gate_up_silu8_dev_q8_sg"),
3320 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3321 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3322 "j8sg" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8sg"),
3323 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3324 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3325 "u64" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_u64"),
3326 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3327 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3328 "gs4" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_gs4"),
3329 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3330 block_dim: (32, 4, 1), shared_mem_bytes: 0 }),
3331 "v" | "" => (self.func("moe_gate_up_silu8_dev_q8_v"),
3333 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3334 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3335 "s2" => (self.func("moe_gate_up_silu8_dev_q8_s2"),
3336 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3337 block_dim: (32, 2, 1), shared_mem_bytes: 0 }),
3338 "s2z" => {
3339 let rz = wpb.min(16); (self.func("moe_gate_up_silu8_dev_q8_s2z"),
3341 LaunchConfig { grid_dim: (n_ff.div_ceil(rz as usize) as u32, n_used as u32, 1),
3342 block_dim: (32, 2, rz), shared_mem_bytes: 0 })
3343 }
3344 _ => (self.func("moe_gate_up_silu8_dev_q8"),
3345 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3346 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3347 };
3348 let __s_b = self.gpu.stream();
3349 let mut b = __s_b.launch_builder(&f);
3350 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3351 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3352 unsafe { b.launch(cfg)?; }
3353 Ok(act)
3354 }
3355
3356 #[allow(clippy::too_many_arguments)]
3357 pub fn moe_down8_fma_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3358 w: &cudarc::driver::CudaView<f32>,
3359 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3360 dst: &mut cudarc::driver::CudaViewMut<f32>,
3361 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3362 qt: i32, rb: usize)
3363 -> Result<(), Box<dyn std::error::Error>> {
3364 static DOWN: std::sync::OnceLock<String> = std::sync::OnceLock::new();
3365 let mode = DOWN.get_or_init(|| std::env::var("MEMRA_MOE_DEVQ8_DOWN").unwrap_or_default());
3366 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3367 n_expert as i32, rb as i64);
3368 let (f, cfg) = match mode.as_str() {
3371 m @ ("1" | "2" | "4") if n_used <= 8 => {
3372 let rpw: usize = m.parse().unwrap();
3373 let f = self.func(match rpw { 1 => "moe_down8_fma_dev_q8_w8r1",
3374 2 => "moe_down8_fma_dev_q8_w8r2",
3375 _ => "moe_down8_fma_dev_q8_w8r4" });
3376 (f, LaunchConfig { grid_dim: (out_f.div_ceil(rpw) as u32, 1, 1),
3377 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3378 }
3379 "h2" if in_f == 512 => (self.func("moe_down8_fma_dev_q8_h2"),
3380 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3381 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3382 "" if in_f == 704 && n_used <= 8 =>
3385 (self.func("moe_down8_fma_dev_q8_w8r2"),
3386 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3387 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3388 "w8h2v" | "" if in_f == 512 && n_used <= 8 =>
3392 (self.func("moe_down8_fma_dev_q8_w8h2v"),
3393 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3394 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3395 "w8h2r2v" if in_f == 512 && n_used <= 8 =>
3396 (self.func("moe_down8_fma_dev_q8_w8h2r2v"),
3397 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3398 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3399 "w8h2r2" if in_f == 512 && n_used <= 8 =>
3400 (self.func("moe_down8_fma_dev_q8_w8h2r2"),
3401 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3402 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3403 "w8h2" if in_f == 512 && n_used <= 8 =>
3404 (self.func("moe_down8_fma_dev_q8_w8h2"),
3405 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3406 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3407 _ => (self.func("moe_down8_fma_dev_q8"),
3408 LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3409 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3410 };
3411 let __s_b = self.gpu.stream();
3412 let mut b = __s_b.launch_builder(&f);
3413 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3414 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3415 unsafe { b.launch(cfg)?; }
3416 Ok(())
3417 }
3418
3419 #[allow(clippy::too_many_arguments)]
3426 pub fn moe_gate_up_silu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3427 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
3428 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3429 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3430 macros: &CudaSlice<f32>)
3431 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3432 let f = self.func("moe_gate_up_silu8_dev_q8_v_rows");
3433 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
3434 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
3435 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3436 let (inf, nff, ne, nu, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3437 n_used as i32, rb_g as i64, rb_u as i64);
3438 let __s_b = self.gpu.stream();
3439 let mut b = __s_b.launch_builder(&f);
3440 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3441 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(macros);
3442 unsafe { b.launch(cfg)?; }
3443 Ok(act)
3444 }
3445
3446 #[allow(clippy::too_many_arguments)]
3451 pub fn moe_down8_fma_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3452 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3453 dst: &mut CudaSlice<f32>, t: usize,
3454 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3455 qt: i32, rb: usize)
3456 -> Result<(), Box<dyn std::error::Error>> {
3457 assert!(in_f == 512 && n_used <= 8, "down rows twin is w8h2v shape-gated");
3458 let f = self.func("moe_down8_fma_dev_q8_w8h2v_rows");
3459 let cfg = LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, t as u32),
3460 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 };
3461 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3462 n_expert as i32, rb as i64);
3463 let __s_b = self.gpu.stream();
3464 let mut b = __s_b.launch_builder(&f);
3465 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3466 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3467 unsafe { b.launch(cfg)?; }
3468 Ok(())
3469 }
3470
3471 #[allow(clippy::too_many_arguments)]
3475 pub fn moe_gate_up_silu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3476 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3477 n_pairs: usize, in_f: usize, n_ff: usize, n_used: usize,
3478 n_expert: usize, qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3479 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3480 let f = self.func("moe_gate_up_silu8_dev_q8_csr_iq4");
3481 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
3482 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
3483 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3484 let (inf, nff, ne, nu, npi, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3485 n_used as i32, n_pairs as i32, rb_g as i64, rb_u as i64);
3486 let __s_b = self.gpu.stream();
3487 let mut b = __s_b.launch_builder(&f);
3488 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3489 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
3490 unsafe { b.launch(cfg)?; }
3491 Ok(act)
3492 }
3493
3494
3495 #[allow(clippy::too_many_arguments)]
3499 pub fn moe_down8_fma_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3500 sel: &cudarc::driver::CudaView<i32>,
3501 w: &cudarc::driver::CudaView<f32>,
3502 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3503 dst: &mut cudarc::driver::CudaViewMut<f32>,
3504 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3505 qt: i32, rb: usize)
3506 -> Result<(), Box<dyn std::error::Error>> {
3507 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3508 n_expert as i32, rb as i64);
3509 let (f, cfg) = match variant {
3510 "w8h2" | "w8h2v" => {
3511 (self.func(if variant == "w8h2" { "moe_down8_fma_dev_q8_w8h2" }
3512 else { "moe_down8_fma_dev_q8_w8h2v" }),
3513 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3514 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3515 }
3516 "w8h2r2" | "w8h2r2v" => {
3517 (self.func(if variant == "w8h2r2" { "moe_down8_fma_dev_q8_w8h2r2" }
3518 else { "moe_down8_fma_dev_q8_w8h2r2v" }),
3519 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3520 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3521 }
3522 _ => (self.func("moe_down8_fma_dev_q8"),
3523 LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3524 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3525 };
3526 let __s_b = self.gpu.stream();
3527 let mut b = __s_b.launch_builder(&f);
3528 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3529 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3530 unsafe { b.launch(cfg)?; }
3531 Ok(())
3532 }
3533
3534 #[allow(clippy::too_many_arguments)]
3536 pub fn moe_gate_up_silu8_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3537 sel: &cudarc::driver::CudaView<i32>,
3538 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3539 in_f: usize, n_ff: usize, n_used: usize,
3540 n_expert: usize, qt_g: i32, qt_u: i32,
3541 rb_g: usize, rb_u: usize)
3542 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3543 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3544 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3545 rb_g as i64, rb_u as i64);
3546 let f = self.func(if variant == "v" { "moe_gate_up_silu8_dev_q8_v" }
3547 else { "moe_gate_up_silu8_dev_q8" });
3548 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3549 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3550 let __s_b = self.gpu.stream();
3551 let mut b = __s_b.launch_builder(&f);
3552 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3553 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
3554 unsafe { b.launch(cfg)?; }
3555 Ok(act)
3556 }
3557
3558 pub fn moe_gate_up_silu8_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3559 x: &cudarc::driver::CudaView<f32>,
3560 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3561 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3562 macros: &CudaSlice<f32>)
3563 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3564 let f = self.func("moe_gate_up_silu8_dev");
3565 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),
3567 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3568 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3569 rb_g as i64, rb_u as i64);
3570 let __s_b = self.gpu.stream();
3571 let mut b = __s_b.launch_builder(&f);
3572 b.arg(table).arg(sel).arg(x).arg(&mut act)
3573 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3574 unsafe { b.launch(cfg)?; }
3575 Ok(act)
3576 }
3577
3578 #[allow(clippy::too_many_arguments)]
3581 pub fn moe_down8_fma_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3582 w: &cudarc::driver::CudaView<f32>, act: &CudaSlice<f32>,
3583 dst: &mut cudarc::driver::CudaViewMut<f32>,
3584 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3585 qt: i32, rb: usize)
3586 -> Result<(), Box<dyn std::error::Error>> {
3587 let f = self.func("moe_down8_fma_dev");
3588 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3589 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3590 let (inf, outf, nu, ne, rbv) = (in_f as i32, out_f as i32, n_used as i32,
3591 n_expert as i32, rb as i64);
3592 let __s_b = self.gpu.stream();
3593 let mut b = __s_b.launch_builder(&f);
3594 b.arg(table).arg(sel).arg(w).arg(act).arg(dst)
3595 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbv);
3596 unsafe { b.launch(cfg)?; }
3597 Ok(())
3598 }
3599
3600 pub fn axpy_into(&self, src: &CudaSlice<f32>, alpha: f32,
3602 dst: &mut cudarc::driver::CudaViewMut<f32>, n: usize)
3603 -> Result<(), Box<dyn std::error::Error>> {
3604 let f = self.func("axpy_f32");
3605 let cfg = LaunchConfig::for_num_elems(n as u32);
3606 let (a, ni) = (alpha, n as i32);
3607 let __s_b = self.gpu.stream();
3608 let mut b = __s_b.launch_builder(&f);
3609 b.arg(src).arg(dst).arg(&a).arg(&ni);
3610 unsafe { b.launch(cfg)?; }
3611 Ok(())
3612 }
3613
3614 pub fn add_scaled_rows(&self, src: &CudaSlice<f32>, scale: &CudaSlice<f32>,
3616 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
3617 -> Result<(), Box<dyn std::error::Error>> {
3618 let f = self.func("add_scaled_rows_f32");
3619 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
3620 let (nc, nr) = (ncols as i32, nrows as i32);
3621 let __s_b = self.gpu.stream();
3622 let mut b = __s_b.launch_builder(&f);
3623 b.arg(src).arg(scale).arg(dst).arg(&nc).arg(&nr);
3624 unsafe { b.launch(cfg)?; }
3625 Ok(())
3626 }
3627
3628 pub fn gather_rows(&self, src: &CudaSlice<f32>, idx: &CudaSlice<i32>,
3632 dst: &mut CudaSlice<f32>, ncols: usize, m_e: usize)
3633 -> Result<(), Box<dyn std::error::Error>> {
3634 let f = self.func("gather_rows_f32");
3635 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3636 let (nc, me) = (ncols as i32, m_e as i32);
3637 let __s_b = self.gpu.stream();
3638 let mut b = __s_b.launch_builder(&f);
3639 b.arg(src).arg(idx).arg(dst).arg(&nc).arg(&me);
3640 unsafe { b.launch(cfg)?; }
3641 Ok(())
3642 }
3643
3644 pub fn scatter_slot(&self, src: &CudaSlice<f32>, tok_idx: &CudaSlice<i32>,
3649 slot_idx: &CudaSlice<i32>, weight: &CudaSlice<f32>,
3650 dst: &mut CudaSlice<f32>, wbuf: &mut CudaSlice<f32>,
3651 ncols: usize, n_used: usize, m_e: usize)
3652 -> Result<(), Box<dyn std::error::Error>> {
3653 let f = self.func("scatter_add_slot_f32");
3654 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3655 let (nc, nu, me) = (ncols as i32, n_used as i32, m_e as i32);
3656 let __s_b = self.gpu.stream();
3657 let mut b = __s_b.launch_builder(&f);
3658 b.arg(src).arg(tok_idx).arg(slot_idx).arg(weight).arg(dst).arg(wbuf).arg(&nc).arg(&nu).arg(&me);
3659 unsafe { b.launch(cfg)?; }
3660 Ok(())
3661 }
3662
3663 pub fn reduce_slots(&self, slots: &CudaSlice<f32>, wbuf: &CudaSlice<f32>,
3667 dst: &mut CudaSlice<f32>, ncols: usize, n_used: usize, t: usize)
3668 -> Result<(), Box<dyn std::error::Error>> {
3669 let f = self.func("reduce_slots_f32");
3670 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
3671 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
3672 let __s_b = self.gpu.stream();
3673 let mut b = __s_b.launch_builder(&f);
3674 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
3675 unsafe { b.launch(cfg)?; }
3676 Ok(())
3677 }
3678
3679 pub fn quantize_q8_1_view(&self, x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize)
3686 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3687 let f = self.func("quantize_q8_1");
3688 let nblk = in_f / 32;
3689 let mut q = self.alloc_uninit::<i8>(m * in_f)?;
3690 let mut d = self.alloc_uninit::<f32>(m * nblk)?;
3691 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
3692 let (inf, mi) = (in_f as i32, m as i32);
3693 let __s_b = self.gpu.stream();
3694 let mut b = __s_b.launch_builder(&f);
3695 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3696 unsafe { b.launch(cfg)?; }
3697 Ok((q, d))
3698 }
3699
3700 pub fn quantize_q8_1(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3701 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3702 let nblk = in_f / 32;
3703 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);
3707 let (inf, mi) = (in_f as i32, m as i32);
3708 if Self::pdl_on() && Self::pdl_wb_on() {
3709 {
3710 use cudarc::driver::{DevicePtr, DevicePtrMut};
3711 let s = &self.gpu.stream();
3712 let (px, _g0) = x.device_ptr(s);
3713 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
3714 let mut ps = [
3715 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
3716 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
3717 &mi as *const _ as *mut _,
3718 ];
3719 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
3720 }
3721 return Ok((q, d));
3722 }
3723 let f = self.func("quantize_q8_1");
3724 let __s_b = self.gpu.stream();
3725 let mut b = __s_b.launch_builder(&f);
3726 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3727 unsafe { b.launch(cfg)?; }
3728 Ok((q, d))
3729 }
3730
3731 pub fn quantize_fp4_act(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3735 -> Result<(CudaSlice<u32>, CudaSlice<u8>), Box<dyn std::error::Error>> {
3736 let f = self.func("quantize_fp4_act");
3737 let nb16 = in_f / 16;
3738 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);
3741 let (inf, mi) = (in_f as i32, m as i32);
3742 let __s_b = self.gpu.stream();
3743 let mut b = __s_b.launch_builder(&f);
3744 b.arg(x).arg(&mut aq4).arg(&mut ad4).arg(&inf).arg(&mi);
3745 unsafe { b.launch(cfg)?; }
3746 Ok((aq4, ad4))
3747 }
3748
3749 pub fn qmatvec_gemm_nvfp4_fp4(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3754 in_f: usize, out_f: usize, row_bytes: usize, scale: f32)
3755 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3756 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
3757 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
3758 let mut y = self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)?;
3759 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
3760 Ok(y)
3761 }
3762
3763 fn fp4_gemm_launch(&self, bytes: &CudaSlice<u8>, aq4: &CudaSlice<u32>, ad4: &CudaSlice<u8>,
3766 m: usize, in_f: usize, out_f: usize, row_bytes: usize)
3767 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3768 let f = self.func("qmatvec_gemm_nvfp4_fp4");
3769 let mut y = self.alloc_uninit::<f32>(m * out_f)?; const BM: u32 = 64; const BN: u32 = 256;
3771 let cfg = LaunchConfig {
3772 grid_dim: ((out_f as u32 + BM - 1) / BM, (m as u32 + BN - 1) / BN, 1),
3773 block_dim: (32, 4, 1), shared_mem_bytes: 0,
3774 };
3775 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3776 let __s_b = self.gpu.stream();
3777 let mut b = __s_b.launch_builder(&f);
3778 b.arg(bytes).arg(aq4).arg(ad4).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3779 unsafe { b.launch(cfg)?; }
3780 Ok(y)
3781 }
3782
3783 pub fn qmatvec_gemm_nvfp4_fp4_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3785 in_f: usize, out_f: usize, row_bytes: usize)
3786 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3787 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
3788 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
3789 self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)
3790 }
3791
3792 pub fn qmatvec_q8_0_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3794 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3795 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3796 let f = self.func("qmatvec_q8_0_dp4a");
3797 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 };
3799 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3800 let __s_b = self.gpu.stream();
3801 let mut b = __s_b.launch_builder(&f);
3802 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3803 unsafe { b.launch(cfg)?; }
3804 Ok(y)
3805 }
3806
3807 #[allow(non_snake_case)] pub fn qmatvec_q4_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3810 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3811 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3812 let f = self.func("qmatvec_q4_K_dp4a");
3813 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 };
3815 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3816 let __s_b = self.gpu.stream();
3817 let mut b = __s_b.launch_builder(&f);
3818 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3819 unsafe { b.launch(cfg)?; }
3820 Ok(y)
3821 }
3822
3823 #[allow(non_snake_case)] pub fn qmatvec_q6_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3826 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3827 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3828 let f = self.func("qmatvec_q6_K_dp4a");
3829 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 };
3831 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3832 let __s_b = self.gpu.stream();
3833 let mut b = __s_b.launch_builder(&f);
3834 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3835 unsafe { b.launch(cfg)?; }
3836 Ok(y)
3837 }
3838
3839 #[allow(non_snake_case)] pub fn qmatvec_q5_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3842 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3843 self.qmatvec_dp4a_named("qmatvec_q5_K_dp4a", w, x, m, in_f, out_f, row_bytes)
3844 }
3845 #[allow(non_snake_case)] pub fn qmatvec_q3_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3848 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3849 self.qmatvec_dp4a_named("qmatvec_q3_K_dp4a", w, x, m, in_f, out_f, row_bytes)
3850 }
3851 pub fn qmatvec_nvfp4_fast_rp(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3853 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3854 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
3855 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a_rp", w, x, m, in_f, out_f, row_bytes)
3856 }
3857 pub fn qmatvec_nvfp4_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3859 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3860 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
3863 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a", w, x, m, in_f, out_f, row_bytes)
3864 }
3865 #[allow(non_snake_case)] pub fn qmatvec_iq4_XS_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3868 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3869 self.qmatvec_dp4a_named("qmatvec_iq4_XS_dp4a", w, x, m, in_f, out_f, row_bytes)
3870 }
3871
3872 fn qmatvec_dp4a_named(&self, name: &str, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3874 in_f: usize, out_f: usize, row_bytes: usize)
3875 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3876 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3877 let f = self.func(name);
3878 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 };
3880 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3881 let __s_b = self.gpu.stream();
3882 let mut b = __s_b.launch_builder(&f);
3883 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3884 unsafe { b.launch(cfg)?; }
3885 Ok(y)
3886 }
3887
3888 pub fn htod(&self, v: &[f32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3889 Ok(self.gpu.stream().clone_htod(v)?)
3890 }
3891 pub fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
3892 Ok(self.gpu.stream().clone_htod(v)?)
3893 }
3894 pub fn htod_i8(&self, v: &[i8]) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
3896 Ok(self.gpu.stream().clone_htod(v)?)
3897 }
3898 pub fn htod_u64(&self, v: &[u64]) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
3899 Ok(self.gpu.stream().clone_htod(v)?)
3900 }
3901 pub fn dtoh_view(&self, d: &cudarc::driver::CudaView<f32>)
3903 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
3904 let v = self.gpu.stream().clone_dtoh(d)?;
3905 self.gpu.stream().synchronize()?;
3906 Ok(v)
3907 }
3908 pub fn dtoh(&self, d: &CudaSlice<f32>) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
3909 let v = self.gpu.stream().clone_dtoh(d)?;
3910 self.gpu.stream().synchronize()?;
3911 Ok(v)
3912 }
3913 pub fn dtoh_pair(
3917 &self,
3918 a: &CudaSlice<f32>,
3919 b: &CudaSlice<f32>,
3920 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
3921 let av = self.gpu.stream().clone_dtoh(a)?;
3922 let bv = self.gpu.stream().clone_dtoh(b)?;
3923 self.gpu.stream().synchronize()?;
3924 Ok((av, bv))
3925 }
3926 pub fn dtoh_i32(&self, d: &CudaSlice<i32>) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
3928 let v = self.gpu.stream().clone_dtoh(d)?;
3929 self.gpu.stream().synchronize()?;
3930 Ok(v)
3931 }
3932 pub fn dtoh_u8(&self, d: &CudaSlice<u8>) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
3934 let v = self.gpu.stream().clone_dtoh(d)?;
3935 self.gpu.stream().synchronize()?;
3936 Ok(v)
3937 }
3938 pub fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3939 let s = self.gpu.stream().alloc_zeros::<f32>(n)?;
3940 self.keep_if_capturing(&s);
3941 Ok(s)
3942 }
3943
3944 pub fn prob_of_token_device(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>, n_vocab: usize)
3953 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3954 let nb = ARGMAX_NB;
3955 let mut part = self.alloc_uninit::<f32>(nb)?;
3956 let mut p = self.alloc_uninit::<f32>(1)?;
3957 let f1 = self.func("prob_of_token_partial_f32");
3958 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3959 let nv = n_vocab as i32;
3960 let __s_b1 = self.gpu.stream();
3961 let mut b1 = __s_b1.launch_builder(&f1);
3962 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
3963 unsafe { b1.launch(cfg1)?; }
3964 let f2 = self.func("prob_of_token_final_f32");
3965 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3966 let nbi = nb as i32;
3967 let __s_b2 = self.gpu.stream();
3968 let mut b2 = __s_b2.launch_builder(&f2);
3969 b2.arg(&part).arg(&mut p).arg(&nbi);
3970 unsafe { b2.launch(cfg2)?; }
3971 Ok(p)
3972 }
3973
3974 pub fn prob_of_token_device_col(&self, logits: &CudaSlice<f32>,
3981 tok_all: &CudaSlice<u32>, tok_idx: usize,
3982 p_out: &mut CudaSlice<f32>, p_idx: usize, n_vocab: usize)
3983 -> Result<(), Box<dyn std::error::Error>> {
3984 let tok_v = tok_all.slice(tok_idx..tok_idx + 1);
3985 let mut p_v = p_out.slice_mut(p_idx..p_idx + 1);
3986 let nb = ARGMAX_NB;
3987 let mut part = self.alloc_uninit::<f32>(nb)?;
3988 let f1 = self.func("prob_of_token_partial_f32");
3989 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3990 let nv = n_vocab as i32;
3991 let __s_b1 = self.gpu.stream();
3992 let mut b1 = __s_b1.launch_builder(&f1);
3993 b1.arg(logits).arg(&tok_v).arg(&mut part).arg(&nv);
3994 unsafe { b1.launch(cfg1)?; }
3995 let f2 = self.func("prob_of_token_final_f32");
3996 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3997 let nbi = nb as i32;
3998 let __s_b2 = self.gpu.stream();
3999 let mut b2 = __s_b2.launch_builder(&f2);
4000 b2.arg(&part).arg(&mut p_v).arg(&nbi);
4001 unsafe { b2.launch(cfg2)?; }
4002 Ok(())
4003 }
4004
4005 pub fn prob_of_token_device_into(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>,
4006 p_out: &mut CudaSlice<f32>, n_vocab: usize)
4007 -> Result<(), Box<dyn std::error::Error>> {
4008 let nb = ARGMAX_NB;
4009 let mut part = self.alloc_uninit::<f32>(nb)?;
4010 let f1 = self.func("prob_of_token_partial_f32");
4011 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4012 let nv = n_vocab as i32;
4013 let __s_b1 = self.gpu.stream();
4014 let mut b1 = __s_b1.launch_builder(&f1);
4015 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
4016 unsafe { b1.launch(cfg1)?; }
4017 let f2 = self.func("prob_of_token_final_f32");
4018 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4019 let nbi = nb as i32;
4020 let __s_b2 = self.gpu.stream();
4021 let mut b2 = __s_b2.launch_builder(&f2);
4022 b2.arg(&part).arg(p_out).arg(&nbi);
4023 unsafe { b2.launch(cfg2)?; }
4024 Ok(())
4025 }
4026
4027 pub fn argmax_token_device(&self, logits: &CudaSlice<f32>, n_vocab: usize)
4028 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4029 let mut tok = unsafe { self.gpu.stream().alloc::<u32>(1)? };
4030 self.argmax_token_device_into(logits, &mut tok, n_vocab)?;
4031 Ok(tok)
4032 }
4033 pub fn argmax_token_device_into(&self, logits: &CudaSlice<f32>, tok: &mut CudaSlice<u32>,
4040 n_vocab: usize) -> Result<(), Box<dyn std::error::Error>> {
4041 let nb = ARGMAX_NB;
4042 let f1 = self.func("argmax_partial_f32");
4043 let f2 = self.func("argmax_final_f32");
4044 let mut guard = self.argmax_partials.lock().unwrap();
4045 if guard.is_none() {
4046 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4049 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4050 *guard = Some((pv, pi));
4051 }
4052 let (part_v, part_i) = guard.as_mut().unwrap();
4053 let nv = n_vocab as i32;
4054 let nbi = nb as i32;
4055 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4057 let __s_b1 = self.gpu.stream();
4058 let mut b1 = __s_b1.launch_builder(&f1);
4059 b1.arg(logits).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4060 unsafe { b1.launch(cfg1)?; }
4061 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4063 let __s_b2 = self.gpu.stream();
4064 let mut b2 = __s_b2.launch_builder(&f2);
4065 b2.arg(&*part_v).arg(&*part_i).arg(tok).arg(&nbi);
4066 unsafe { b2.launch(cfg2)?; }
4067 Ok(())
4068 }
4069 pub fn argmax_token_device_col(&self, logits: &CudaSlice<f32>, col: usize, n_vocab: usize,
4075 toks: &mut CudaSlice<u32>, out_idx: usize)
4076 -> Result<(), Box<dyn std::error::Error>> {
4077 let nb = ARGMAX_NB;
4078 let f1 = self.func("argmax_partial_f32");
4079 let f2 = self.func("argmax_final_f32");
4080 let mut guard = self.argmax_partials.lock().unwrap();
4081 if guard.is_none() {
4082 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4083 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4084 *guard = Some((pv, pi));
4085 }
4086 let (part_v, part_i) = guard.as_mut().unwrap();
4087 let col_view = logits.slice(col * n_vocab..(col + 1) * n_vocab);
4088 let nv = n_vocab as i32;
4089 let nbi = nb as i32;
4090 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4091 let __s_b1 = self.gpu.stream();
4092 let mut b1 = __s_b1.launch_builder(&f1);
4093 b1.arg(&col_view).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4094 unsafe { b1.launch(cfg1)?; }
4095 let mut tok_view = toks.slice_mut(out_idx..out_idx + 1);
4096 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4097 let __s_b2 = self.gpu.stream();
4098 let mut b2 = __s_b2.launch_builder(&f2);
4099 b2.arg(&*part_v).arg(&*part_i).arg(&mut tok_view).arg(&nbi);
4100 unsafe { b2.launch(cfg2)?; }
4101 Ok(())
4102 }
4103 pub fn htod_u32_v(&self, v: &[u32]) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4105 Ok(self.gpu.stream().clone_htod(v)?)
4106 }
4107 pub fn dtoh_u32(&self, d: &CudaSlice<u32>) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
4108 let v = self.gpu.stream().clone_dtoh(d)?;
4109 self.gpu.stream().synchronize()?;
4110 Ok(v)
4111 }
4112 pub fn htod_u32_into(&self, dst: &mut CudaSlice<u32>, src: &[u32])
4116 -> Result<(), Box<dyn std::error::Error>> {
4117 let mut view = dst.slice_mut(0..src.len());
4118 self.gpu.stream().memcpy_htod(src, &mut view)?;
4119 Ok(())
4120 }
4121
4122 pub fn alloc_u32_zeroed(&self, n: usize) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4123 let s = self.gpu.stream().alloc_zeros::<u32>(n)?;
4124 self.keep_if_capturing(&s);
4125 Ok(s)
4126 }
4127 pub fn embed_gather_device_into(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4130 x_out: &mut CudaSlice<f32>, n_embd: usize, qtype: i32,
4131 row_bytes: usize) -> Result<(), Box<dyn std::error::Error>> {
4132 let f = self.func("embed_gather_u32");
4133 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
4134 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4135 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
4136 let __s_b = self.gpu.stream();
4137 let mut b = __s_b.launch_builder(&f);
4138 b.arg(embd).arg(token_d).arg(x_out).arg(&ne).arg(&qt).arg(&rb);
4139 unsafe { b.launch(cfg)?; }
4140 Ok(())
4141 }
4142 pub fn dtoh_i32_one(&self, d: &CudaSlice<i32>) -> Result<i32, Box<dyn std::error::Error>> {
4144 let v = self.gpu.stream().clone_dtoh(d)?;
4145 self.gpu.stream().synchronize()?;
4146 Ok(v[0])
4147 }
4148 pub fn i32_set_k(&self, dst: &mut CudaSlice<i32>, v: i32)
4155 -> Result<(), Box<dyn std::error::Error>> {
4156 let f = self.func("i32_set_k");
4157 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
4158 let idx = 0i32;
4159 let __s_b = self.gpu.stream();
4160 let mut b = __s_b.launch_builder(&f);
4161 b.arg(dst).arg(&v).arg(&idx);
4162 unsafe { b.launch(cfg)?; }
4163 Ok(())
4164 }
4165
4166 pub fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
4167 self.gpu.stream().memcpy_htod(&[v], d)?;
4168 Ok(())
4169 }
4170 pub fn set_u32_one(&self, d: &mut CudaSlice<u32>, v: u32) -> Result<(), Box<dyn std::error::Error>> {
4173 self.gpu.stream().memcpy_htod(&[v], d)?;
4174 Ok(())
4175 }
4176 pub fn dtoh_u32_one(&self, d: &CudaSlice<u32>) -> Result<u32, Box<dyn std::error::Error>> {
4178 let v = self.gpu.stream().clone_dtoh(d)?;
4179 self.gpu.stream().synchronize()?;
4180 Ok(v[0])
4181 }
4182 pub fn upload_u8(&self, bytes: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
4184 Ok(self.gpu.stream().clone_htod(bytes)?)
4185 }
4186 pub fn embed_gather_device(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4190 n_embd: usize, qtype: i32, row_bytes: usize)
4191 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4192 let f = self.func("embed_gather_u32");
4193 let mut x = self.alloc_uninit::<f32>(n_embd)?;
4194 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
4195 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4196 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
4197 let __s_b = self.gpu.stream();
4198 let mut b = __s_b.launch_builder(&f);
4199 b.arg(embd).arg(token_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb);
4200 unsafe { b.launch(cfg)?; }
4201 Ok(x)
4202 }
4203
4204
4205 pub fn embed_gather_device_t(&self, embd: &CudaSlice<u8>, tokens: &[u32],
4209 n_embd: usize, qtype: i32, row_bytes: usize)
4210 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4211 let t = tokens.len();
4212 let tok_d = self.gpu.stream().clone_htod(tokens)?;
4213 let f = self.func("embed_gather_u32_t");
4214 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4215 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4216 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4217 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4218 let __s_b = self.gpu.stream();
4219 let mut b = __s_b.launch_builder(&f);
4220 b.arg(embd).arg(&tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4221 unsafe { b.launch(cfg)?; }
4222 Ok(x)
4223 }
4224
4225 pub fn embed_gather_device_tv(&self, embd: &CudaSlice<u8>, tok_v: &cudarc::driver::CudaView<u32>,
4230 t: usize, n_embd: usize, qtype: i32, row_bytes: usize)
4231 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4232 let f = self.func("embed_gather_u32_t");
4233 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4234 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4235 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4236 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4237 let __s_b = self.gpu.stream();
4238 let mut b = __s_b.launch_builder(&f);
4239 b.arg(embd).arg(tok_v).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4240 unsafe { b.launch(cfg)?; }
4241 Ok(x)
4242 }
4243
4244 pub fn embed_gather_device_td(&self, embd: &CudaSlice<u8>, tok_d: &CudaSlice<u32>, t: usize,
4245 n_embd: usize, qtype: i32, row_bytes: usize)
4246 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4247 let f = self.func("embed_gather_u32_t");
4248 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4249 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4250 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4251 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4252 let __s_b = self.gpu.stream();
4253 let mut b = __s_b.launch_builder(&f);
4254 b.arg(embd).arg(tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4255 unsafe { b.launch(cfg)?; }
4256 Ok(x)
4257 }
4258
4259 #[inline]
4265 fn keep_if_capturing<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, s: &CudaSlice<T>) {
4267 if self.capture_keep_on.load(std::sync::atomic::Ordering::Relaxed) {
4268 self.capture_keep.lock().unwrap().push(Box::new(s.clone()));
4269 }
4270 }
4271
4272 fn alloc_uninit<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, n: usize)
4273 -> Result<CudaSlice<T>, Box<dyn std::error::Error>> {
4274 let mut s = unsafe { self.gpu.stream().alloc::<T>(n)? };
4275 {
4279 static Z: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4280 if *Z.get_or_init(|| std::env::var("MEMRA_DEBUG_ZERO_ALLOCS").as_deref() == Ok("1")) {
4281 use cudarc::driver::DevicePtrMut;
4283 let n_bytes = s.len() * std::mem::size_of::<T>();
4284 let stream = self.gpu.stream();
4285 let (p_, _g) = s.device_ptr_mut(&stream);
4286 unsafe {
4287 cudarc::driver::sys::cuMemsetD8Async(p_, 0, n_bytes, stream.cu_stream())
4288 .result()?;
4289 }
4290 }
4291 }
4292 self.keep_if_capturing(&s);
4293 Ok(s)
4294 }
4295
4296 pub fn uninit_q8_pair(&self, n: usize)
4301 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4302 Ok((self.alloc_uninit::<i8>(n)?, self.alloc_uninit::<f32>(n / 32)?))
4303 }
4304
4305 pub fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4306 self.alloc_uninit::<f32>(n)
4307 }
4308
4309 pub fn alloc_i8_uninit(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
4311 self.alloc_uninit::<i8>(n)
4312 }
4313
4314 #[allow(clippy::too_many_arguments)]
4318 pub fn rms_norm3(&self, x: &CudaSlice<f32>, w0: &CudaSlice<f32>, w1: &CudaSlice<f32>,
4319 w2: &CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4320 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4321 -> Result<(), Box<dyn std::error::Error>> {
4322 let f = self.func("rms_norm3_f32");
4323 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4324 let (nc, e) = (ncols as i32, eps);
4325 let __s_b = self.gpu.stream();
4326 let mut b = __s_b.launch_builder(&f);
4327 b.arg(x).arg(w0).arg(w1).arg(w2).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e);
4328 unsafe { b.launch(cfg)?; }
4329 Ok(())
4330 }
4331
4332 #[allow(clippy::too_many_arguments)]
4334 pub fn qkvnorm_w_on_prefill(rows: usize, ncols: usize) -> bool {
4337 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4338 *WARP_ON.get_or_init(|| {
4339 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4340 }) && ncols % 4 == 0 && rows >= 64
4341 }
4342
4343 #[allow(clippy::too_many_arguments)]
4346 pub fn rms_norm_qkv_w4b(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4347 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4348 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4349 dvb: &mut CudaSlice<u8>,
4350 ncols: usize, rq: usize, rk: usize, eps: f32, vf16: bool)
4351 -> Result<(), Box<dyn std::error::Error>> {
4352 assert!(ncols % 4 == 0 && rq + 2 * rk >= 64);
4353 let f = self.func("rms_norm_qkv_w4b_f32");
4354 let rows = (rq + 2 * rk) as u32;
4355 let cfg = LaunchConfig {
4356 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4357 };
4358 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4359 let vf = vf16 as i32;
4360 let __s_b = self.gpu.stream();
4361 let mut b = __s_b.launch_builder(&f);
4362 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv).arg(&mut *dvb)
4363 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e).arg(&vf);
4364 unsafe { b.launch(cfg)?; }
4365 Ok(())
4366 }
4367
4368 pub fn rms_norm_qkv(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4369 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4370 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4371 ncols: usize, rq: usize, rk: usize, eps: f32)
4372 -> Result<(), Box<dyn std::error::Error>> {
4373 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4377 let warp_on = *WARP_ON.get_or_init(|| {
4378 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4379 });
4380 if warp_on && ncols % 4 == 0 && rq + 2 * rk >= 64 {
4383 let f = self.func("rms_norm_qkv_w4_f32");
4384 let rows = (rq + 2 * rk) as u32;
4385 let cfg = LaunchConfig {
4386 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4387 };
4388 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4389 let __s_b = self.gpu.stream();
4390 let mut b = __s_b.launch_builder(&f);
4391 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4392 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e);
4393 unsafe { b.launch(cfg)?; }
4394 return Ok(());
4395 }
4396 let f = self.func("rms_norm_qkv_f32");
4397 let grid = (rq + 2 * rk) as u32;
4398 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4399 let (nc, rqi, rki, e) = (ncols as i32, rq as i32, rk as i32, eps);
4400 let __s_b = self.gpu.stream();
4401 let mut b = __s_b.launch_builder(&f);
4402 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4403 .arg(&nc).arg(&rqi).arg(&rki).arg(&e);
4404 unsafe { b.launch(cfg)?; }
4405 Ok(())
4406 }
4407
4408 #[allow(clippy::too_many_arguments)]
4410 pub fn rms_norm2x(&self, a: &CudaSlice<f32>, bb: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4411 wb: &CudaSlice<f32>, da: &mut CudaSlice<f32>, db: &mut CudaSlice<f32>,
4412 ncols: usize, nrows: usize, eps: f32)
4413 -> Result<(), Box<dyn std::error::Error>> {
4414 let f = self.func("rms_norm2x_f32");
4415 let cfg = LaunchConfig { grid_dim: (2 * nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4416 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
4417 let __s_b = self.gpu.stream();
4418 let mut b = __s_b.launch_builder(&f);
4419 b.arg(a).arg(bb).arg(wa).arg(wb).arg(da).arg(db).arg(&nc).arg(&nr).arg(&e);
4420 unsafe { b.launch(cfg)?; }
4421 Ok(())
4422 }
4423
4424 pub fn softcap(&self, y: &mut CudaSlice<f32>, cap: f32, n: usize)
4426 -> Result<(), Box<dyn std::error::Error>> {
4427 let f = self.func("softcap_f32");
4428 let cfg = LaunchConfig::for_num_elems(n as u32);
4429 let ni = n as i32;
4430 let __s_b = self.gpu.stream();
4431 let mut b = __s_b.launch_builder(&f);
4432 b.arg(y).arg(&cap).arg(&ni);
4433 unsafe { b.launch(cfg)?; }
4434 Ok(())
4435 }
4436
4437 pub fn mask_ids_rows(&self, y: &mut CudaSlice<f32>, ids: &CudaSlice<i32>, n_ids: usize,
4440 n_vocab: usize, t: usize)
4441 -> Result<(), Box<dyn std::error::Error>> {
4442 let f = self.func("mask_ids_rows_f32");
4443 let cfg = LaunchConfig::for_num_elems((n_ids * t) as u32);
4444 let (ni, nv, ti) = (n_ids as i32, n_vocab as i32, t as i32);
4445 let __s_b = self.gpu.stream();
4446 let mut b = __s_b.launch_builder(&f);
4447 b.arg(y).arg(ids).arg(&ni).arg(&nv).arg(&ti);
4448 unsafe { b.launch(cfg)?; }
4449 Ok(())
4450 }
4451
4452 #[allow(clippy::too_many_arguments)]
4454 pub fn add_scale_rms_norm(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4455 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4456 ncols: usize, nrows: usize, eps: f32)
4457 -> Result<(), Box<dyn std::error::Error>> {
4458 let f = self.func("add_scale_rms_norm_f32");
4459 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4460 let (nc, e2) = (ncols as i32, eps);
4461 let __s_b = self.gpu.stream();
4462 let mut b = __s_b.launch_builder(&f);
4463 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(dst).arg(&nc).arg(&e2);
4464 unsafe { b.launch(cfg)?; }
4465 Ok(())
4466 }
4467
4468 #[allow(clippy::too_many_arguments)]
4471 pub fn add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4472 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4473 ncols: usize, nrows: usize, eps: f32)
4474 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4475 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4476 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4477 let (nc, e2) = (ncols as i32, eps);
4478 if Self::pdl_on() && Self::pdl_wb_on() {
4479 {
4480 use cudarc::driver::{DevicePtr, DevicePtrMut};
4481 let s = &self.gpu.stream();
4482 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4483 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4484 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4485 let mut ps = [
4486 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4487 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4488 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4489 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4490 &e2 as *const _ as *mut _,
4491 ];
4492 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4493 (rms_block(), 1, 1), &mut ps)?; }
4494 }
4495 return Ok((out_q, out_d));
4496 }
4497 let f = self.func("add_scale_rms_norm_q8_1");
4498 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4499 let __s_b = self.gpu.stream();
4500 let mut b = __s_b.launch_builder(&f);
4501 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e2);
4502 unsafe { b.launch(cfg)?; }
4503 Ok((out_q, out_d))
4504 }
4505
4506 #[allow(clippy::too_many_arguments)]
4508 pub fn add_scale_rms_norm_q8_1_into(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4509 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4510 ncols: usize, nrows: usize, eps: f32,
4511 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4512 -> Result<(), Box<dyn std::error::Error>> {
4513 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4514 let (nc, e2) = (ncols as i32, eps);
4515 if Self::pdl_on() && Self::pdl_wb_on() {
4516 use cudarc::driver::{DevicePtr, DevicePtrMut};
4517 let s = &self.gpu.stream();
4518 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4519 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4520 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4521 let mut ps = [
4522 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4523 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4524 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4525 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4526 &e2 as *const _ as *mut _,
4527 ];
4528 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4529 (rms_block(), 1, 1), &mut ps)?; }
4530 return Ok(());
4531 }
4532 let f = self.func("add_scale_rms_norm_q8_1");
4533 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4534 let __s_b = self.gpu.stream();
4535 let mut b = __s_b.launch_builder(&f);
4536 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut *out_q).arg(&mut *out_d).arg(&nc).arg(&e2);
4537 unsafe { b.launch(cfg)?; }
4538 Ok(())
4539 }
4540
4541 #[allow(clippy::too_many_arguments)]
4544 pub fn rms_pre_add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4545 b_in: &CudaSlice<f32>, c: f32,
4546 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4547 ncols: usize, nrows: usize, eps: f32)
4548 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4549 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4550 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4551 let (nc, e2) = (ncols as i32, eps);
4552 if Self::pdl_on() {
4553 {
4554 use cudarc::driver::{DevicePtr, DevicePtrMut};
4555 let s = &self.gpu.stream();
4556 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
4557 let (pb, _g2) = b_in.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
4558 let (pr, _g4) = res.device_ptr_mut(s);
4559 let (pq, _g5) = out_q.device_ptr_mut(s); let (pd, _g6) = out_d.device_ptr_mut(s);
4560 let mut ps = [
4561 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
4562 &pb as *const _ as *mut _, &c as *const _ as *mut _,
4563 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
4564 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4565 &nc as *const _ as *mut _, &e2 as *const _ as *mut _,
4566 ];
4567 unsafe { self.launch_pdl("rms_pre_add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4568 (rms_block(), 1, 1), &mut ps)?; }
4569 }
4570 return Ok((out_q, out_d));
4571 }
4572 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
4573 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4574 let __s_b = self.gpu.stream();
4575 let mut b = __s_b.launch_builder(&f);
4576 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);
4577 unsafe { b.launch(cfg)?; }
4578 Ok((out_q, out_d))
4579 }
4580
4581 pub fn gelu_tanh_mul_q8_1(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4584 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
4585 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4586 debug_assert!(ncols % 128 == 0);
4587 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4588 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4589 let nc = ncols as i32;
4590 if Self::pdl_on() {
4591 {
4592 use cudarc::driver::{DevicePtr, DevicePtrMut};
4593 let s = &self.gpu.stream();
4594 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4595 let (pact, _g2) = act.device_ptr_mut(s);
4596 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4597 let mut ps = [
4598 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4599 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4600 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4601 ];
4602 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4603 (rms_block(), 1, 1), &mut ps)?; }
4604 }
4605 return Ok((out_q, out_d));
4606 }
4607 let f = self.func("gelu_tanh_mul_q8_1");
4608 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4609 let __s_b = self.gpu.stream();
4610 let mut b = __s_b.launch_builder(&f);
4611 b.arg(gate).arg(up).arg(act).arg(&mut out_q).arg(&mut out_d).arg(&nc);
4612 unsafe { b.launch(cfg)?; }
4613 Ok((out_q, out_d))
4614 }
4615
4616 #[allow(clippy::too_many_arguments)]
4618 pub fn gelu_tanh_mul_q8_1_into(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4619 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
4620 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4621 -> Result<(), Box<dyn std::error::Error>> {
4622 debug_assert!(ncols % 128 == 0);
4623 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4624 let nc = ncols as i32;
4625 if Self::pdl_on() {
4626 use cudarc::driver::{DevicePtr, DevicePtrMut};
4627 let s = &self.gpu.stream();
4628 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4629 let (pact, _g2) = act.device_ptr_mut(s);
4630 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4631 let mut ps = [
4632 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4633 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4634 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4635 ];
4636 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4637 (rms_block(), 1, 1), &mut ps)?; }
4638 return Ok(());
4639 }
4640 let f = self.func("gelu_tanh_mul_q8_1");
4641 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4642 let __s_b = self.gpu.stream();
4643 let mut b = __s_b.launch_builder(&f);
4644 b.arg(gate).arg(up).arg(&mut *act).arg(&mut *out_q).arg(&mut *out_d).arg(&nc);
4645 unsafe { b.launch(cfg)?; }
4646 Ok(())
4647 }
4648
4649 #[allow(clippy::too_many_arguments)]
4651 pub fn add_rms_norm3_q8z(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4652 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4653 res: &mut CudaSlice<f32>, out1: &mut CudaSlice<f32>,
4654 ncols: usize, nrows: usize, eps: f32)
4655 -> Result<((CudaSlice<i8>, CudaSlice<f32>), (CudaSlice<i8>, CudaSlice<f32>)), Box<dyn std::error::Error>> {
4656 let mut q0 = self.alloc_uninit::<i8>(nrows * ncols)?;
4657 let mut d0 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4658 let mut q2 = self.alloc_uninit::<i8>(nrows * ncols)?;
4659 let mut d2 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4660 let f = self.func("add_rms_norm3_q8z_f32");
4661 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4662 let (nc, e2) = (ncols as i32, eps);
4663 let __s_b = self.gpu.stream();
4664 let mut b = __s_b.launch_builder(&f);
4665 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res)
4666 .arg(&mut q0).arg(&mut d0).arg(out1).arg(&mut q2).arg(&mut d2).arg(&nc).arg(&e2);
4667 unsafe { b.launch(cfg)?; }
4668 Ok(((q0, d0), (q2, d2)))
4669 }
4670
4671 #[allow(clippy::too_many_arguments)]
4673 pub fn add_rms_norm3(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4674 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4675 res: &mut CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4676 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4677 -> Result<(), Box<dyn std::error::Error>> {
4678 let f = self.func("add_rms_norm3_f32");
4679 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4680 let (nc, e2) = (ncols as i32, eps);
4681 let __s_b = self.gpu.stream();
4682 let mut b = __s_b.launch_builder(&f);
4683 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e2);
4684 unsafe { b.launch(cfg)?; }
4685 Ok(())
4686 }
4687
4688 pub fn add_scale(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4690 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
4691 let f = self.func("add_scale_f32");
4692 let cfg = LaunchConfig::for_num_elems(n as u32);
4693 let ni = n as i32;
4694 let __s_b = self.gpu.stream();
4695 let mut b = __s_b.launch_builder(&f);
4696 b.arg(a).arg(b_in).arg(&c).arg(dst).arg(&ni);
4697 unsafe { b.launch(cfg)?; }
4698 Ok(())
4699 }
4700
4701 pub fn rms_norm(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4702 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4703 let (nc, e) = (ncols as i32, eps);
4704 if Self::pdl_on() && Self::pdl_wb_on() {
4705 use cudarc::driver::{DevicePtr, DevicePtrMut};
4706 let s = &self.gpu.stream();
4707 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4708 let (pd, _g2) = dst.device_ptr_mut(s);
4709 let mut ps = [
4710 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4711 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4712 &e as *const _ as *mut _,
4713 ];
4714 unsafe { self.launch_pdl("rms_norm_f32", (nrows as u32, 1, 1),
4715 (rms_block(), 1, 1), &mut ps)?; }
4716 return Ok(());
4717 }
4718 let f = self.func("rms_norm_f32");
4719 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4720 let __s_b = self.gpu.stream();
4721 let mut b = __s_b.launch_builder(&f);
4722 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4723 unsafe { b.launch(cfg)?; }
4724 Ok(())
4725 }
4726
4727 pub fn rms_norm_decode(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4735 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4736 let f = self.func("rms_norm_f32");
4737 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4738 let (nc, e) = (ncols as i32, eps);
4739 let __s_b = self.gpu.stream();
4740 let mut b = __s_b.launch_builder(&f);
4741 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4742 unsafe { b.launch(cfg)?; }
4743 Ok(())
4744 }
4745
4746 pub fn rms_norm_q8_1(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize, nrows: usize,
4750 eps: f32) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4751 let nblk = ncols / 32;
4752 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
4753 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
4754 let (nc, e) = (ncols as i32, eps);
4755 if Self::pdl_on() {
4756 {
4757 use cudarc::driver::{DevicePtr, DevicePtrMut};
4758 let s = &self.gpu.stream();
4759 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4760 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
4761 let mut ps = [
4762 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4763 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4764 &nc as *const _ as *mut _, &e as *const _ as *mut _,
4765 ];
4766 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
4767 &mut ps)?; }
4768 }
4769 return Ok((q, d));
4770 }
4771 let f = self.func("rms_norm_q8_1");
4772 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4775 let __s_b = self.gpu.stream();
4776 let mut b = __s_b.launch_builder(&f);
4777 b.arg(x).arg(w).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
4778 unsafe { b.launch(cfg)?; }
4779 Ok((q, d))
4780 }
4781
4782 pub fn rms_norm_q8_1_into(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize,
4785 nrows: usize, eps: f32,
4786 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
4787 -> Result<(), Box<dyn std::error::Error>> {
4788 let nblk = ncols / 32;
4789 debug_assert!(q.len() >= nrows * ncols && d.len() >= nrows * nblk);
4790 let (nc, e) = (ncols as i32, eps);
4791 if Self::pdl_on() {
4792 use cudarc::driver::{DevicePtr, DevicePtrMut};
4793 let s = &self.gpu.stream();
4794 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4795 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
4796 let mut ps = [
4797 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4798 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4799 &nc as *const _ as *mut _, &e as *const _ as *mut _,
4800 ];
4801 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
4802 &mut ps)?; }
4803 return Ok(());
4804 }
4805 let f = self.func("rms_norm_q8_1");
4806 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4807 let __s_b = self.gpu.stream();
4808 let mut b = __s_b.launch_builder(&f);
4809 b.arg(x).arg(w).arg(&mut *q).arg(&mut *d).arg(&nc).arg(&e);
4810 unsafe { b.launch(cfg)?; }
4811 Ok(())
4812 }
4813
4814 pub fn quantize_q8_1_into(&self, x: &CudaSlice<f32>, m: usize, in_f: usize,
4816 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
4817 -> Result<(), Box<dyn std::error::Error>> {
4818 let nblk = in_f / 32;
4819 debug_assert!(q.len() >= m * in_f && d.len() >= m * nblk);
4820 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
4821 let (inf, mi) = (in_f as i32, m as i32);
4822 if Self::pdl_on() && Self::pdl_wb_on() {
4823 use cudarc::driver::{DevicePtr, DevicePtrMut};
4824 let s = &self.gpu.stream();
4825 let (px, _g0) = x.device_ptr(s);
4826 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
4827 let mut ps = [
4828 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
4829 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
4830 &mi as *const _ as *mut _,
4831 ];
4832 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
4833 return Ok(());
4834 }
4835 let f = self.func("quantize_q8_1");
4836 let __s_b = self.gpu.stream();
4837 let mut b = __s_b.launch_builder(&f);
4838 b.arg(x).arg(&mut *q).arg(&mut *d).arg(&inf).arg(&mi);
4839 unsafe { b.launch(cfg)?; }
4840 Ok(())
4841 }
4842
4843 pub fn add_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
4847 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4848 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4849 let nblk = ncols / 32;
4850 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
4851 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
4852 let f = self.func("add_rms_norm_q8_1");
4853 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4855 let (nc, e) = (ncols as i32, eps);
4856 let __s_bld = self.gpu.stream();
4857 let mut bld = __s_bld.launch_builder(&f);
4858 bld.arg(a).arg(b_in).arg(w).arg(res).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
4859 unsafe { bld.launch(cfg)?; }
4860 Ok((q, d))
4861 }
4862
4863 pub fn add_rms_norm(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
4867 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
4868 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4869 let (nc, e) = (ncols as i32, eps);
4870 if Self::pdl_on() && Self::pdl_wb_on() {
4871 use cudarc::driver::{DevicePtr, DevicePtrMut};
4872 let s = &self.gpu.stream();
4873 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b.device_ptr(s);
4874 let (pw, _g2) = w.device_ptr(s);
4875 let (pr, _g3) = res.device_ptr_mut(s); let (pd, _g4) = dst.device_ptr_mut(s);
4876 let mut ps = [
4877 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4878 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
4879 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4880 &e as *const _ as *mut _,
4881 ];
4882 unsafe { self.launch_pdl("add_rms_norm_f32", (nrows as u32, 1, 1),
4883 (rms_block(), 1, 1), &mut ps)?; }
4884 return Ok(());
4885 }
4886 let f = self.func("add_rms_norm_f32");
4887 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4888 let __s_b2 = self.gpu.stream();
4889 let mut b2 = __s_b2.launch_builder(&f);
4890 b2.arg(a).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
4891 unsafe { b2.launch(cfg)?; }
4892 Ok(())
4893 }
4894
4895 #[allow(clippy::too_many_arguments)]
4898 pub fn rms_pre_add_rms_norm(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4899 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
4900 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4901 ncols: usize, nrows: usize, eps: f32)
4902 -> Result<(), Box<dyn std::error::Error>> {
4903 let f = self.func("rms_pre_add_rms_norm_f32");
4904 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4905 let (nc, e) = (ncols as i32, eps);
4906 let __s_b2 = self.gpu.stream();
4907 let mut b2 = __s_b2.launch_builder(&f);
4908 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
4909 unsafe { b2.launch(cfg)?; }
4910 Ok(())
4911 }
4912
4913 #[allow(clippy::too_many_arguments)]
4915 pub fn rms_pre_add_rms_norm_q8z(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4916 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
4917 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4918 ncols: usize, nrows: usize, eps: f32)
4919 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4920 debug_assert!(ncols % 128 == 0);
4921 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4922 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4923 let (nc, e) = (ncols as i32, eps);
4924 if Self::pdl_on() {
4925 {
4926 use cudarc::driver::{DevicePtr, DevicePtrMut};
4927 let s = &self.gpu.stream();
4928 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
4929 let (pb, _g2) = b.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
4930 let (pr, _g4) = res.device_ptr_mut(s); let (pdst, _g5) = dst.device_ptr_mut(s);
4931 let (pq, _g6) = out_q.device_ptr_mut(s); let (pd, _g7) = out_d.device_ptr_mut(s);
4932 let mut ps = [
4933 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
4934 &pb as *const _ as *mut _, &pw as *const _ as *mut _,
4935 &pr as *const _ as *mut _, &pdst as *const _ as *mut _,
4936 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4937 &nc as *const _ as *mut _, &e as *const _ as *mut _,
4938 ];
4939 unsafe { self.launch_pdl("rms_pre_add_rms_norm_q8z_f32", (nrows as u32, 1, 1),
4940 (rms_block(), 1, 1), &mut ps)?; }
4941 }
4942 return Ok((out_q, out_d));
4943 }
4944 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
4945 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4946 let __s_b2 = self.gpu.stream();
4947 let mut b2 = __s_b2.launch_builder(&f);
4948 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst)
4949 .arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e);
4950 unsafe { b2.launch(cfg)?; }
4951 Ok((out_q, out_d))
4952 }
4953
4954 pub fn build_q4_out_concat3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
4958 w2: &crate::model::GpuTensor)
4959 -> Result<Option<crate::model::GpuTensor>, Box<dyn std::error::Error>> {
4960 use crate::model::GpuTensor;
4961 let part = |w: &GpuTensor| -> Option<(usize, usize)> {
4962 match w {
4963 GpuTensor::Quant { qtype, row_bytes, rp, .. }
4964 if *qtype == QT_Q4_0 && !*rp => Some((*row_bytes, w.out_features())),
4965 _ => None,
4966 }
4967 };
4968 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (part(w0), part(w1), part(w2))
4969 else { return Ok(None) };
4970 if rb0 != rb1 || rb0 != rb2
4971 || w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
4972 return Ok(None);
4973 }
4974 fn bytes_of(w: &crate::model::GpuTensor) -> &CudaSlice<u8> {
4975 match w { crate::model::GpuTensor::Quant { bytes, .. } => bytes, _ => unreachable!() }
4976 }
4977 let (b0, b1, b2) = (bytes_of(w0), bytes_of(w1), bytes_of(w2));
4978 let total = rb0 * (o0 + o1 + o2);
4979 let mut cat = self.alloc_u8(total)?;
4980 self.copy_u8_into(&mut cat, 0, b0, rb0 * o0)?;
4981 self.copy_u8_into(&mut cat, rb0 * o0, b1, rb1 * o1)?;
4982 self.copy_u8_into(&mut cat, rb0 * (o0 + o1), b2, rb2 * o2)?;
4983 Ok(Some(GpuTensor::Quant {
4984 bytes: cat, qtype: QT_Q4_0, row_bytes: rb0,
4985 ne: vec![w0.in_features() as u64, (o0 + o1 + o2) as u64], scale: 1.0, rp: false,
4986 #[cfg(memra_cutlass)]
4987 cutlass: None,
4988 fp8: None, blk: None, rp4: None, f16: None,
4989 }))
4990 }
4991
4992 #[allow(clippy::too_many_arguments)]
4994 pub fn rms_norm_qkv_rope_cat(&self, qkv: &CudaSlice<f32>,
4995 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4996 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
4997 head_dim: usize, rq: usize, rk: usize,
4998 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
4999 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32)
5000 -> Result<(), Box<dyn std::error::Error>> {
5001 let rows = rq + rk + rk;
5002 let theta_scale = base.powf(-2.0 / head_dim as f32);
5003 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5004 if Self::pdl_on() {
5005 use cudarc::driver::{DevicePtr, DevicePtrMut};
5006 let s = &self.gpu.stream();
5007 let (pqkv, _g0) = qkv.device_ptr(s);
5008 let (pwq, _g1) = wq.device_ptr(s); let (pwk, _g2) = wk.device_ptr(s);
5009 let (pwv, _g3) = wv.device_ptr(s);
5010 let (pq, _g4) = q.device_ptr_mut(s); let (pk, _g5) = k.device_ptr_mut(s);
5011 let (pv, _g6) = v.device_ptr_mut(s);
5012 let (ppos, _g7) = pos.device_ptr(s);
5013 let (pff, _g8) = match ff {
5014 Some(t) => { let (p, g) = t.device_ptr(s); (p, Some(g)) }
5015 None => (0, None),
5016 };
5017 let mut ps = [
5018 &pqkv as *const _ as *mut std::ffi::c_void,
5019 &pwq as *const _ as *mut _, &pwk as *const _ as *mut _,
5020 &pwv as *const _ as *mut _,
5021 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5022 &pv as *const _ as *mut _,
5023 &nc as *const _ as *mut _, &rqi as *const _ as *mut _,
5024 &rki as *const _ as *mut _, &ppos as *const _ as *mut _,
5025 &nhq as *const _ as *mut _, &nhk as *const _ as *mut _,
5026 &theta_scale as *const _ as *mut _, &freq_scale as *const _ as *mut _,
5027 &pff as *const _ as *mut _, &eps as *const _ as *mut _,
5028 ];
5029 unsafe { self.launch_pdl("rms_norm_qkv_rope_cat_f32", (rows as u32, 1, 1),
5030 (rms_block(), 1, 1), &mut ps)?; }
5031 return Ok(());
5032 }
5033 let f = self.func("rms_norm_qkv_rope_cat_f32");
5034 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5035 let __s_b = self.gpu.stream();
5036 let mut b = __s_b.launch_builder(&f);
5037 match ff {
5038 Some(t) => { b.arg(qkv).arg(wq).arg(wk).arg(wv)
5039 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5040 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5041 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5042 unsafe { b.launch(cfg)?; } }
5043 None => { let null: u64 = 0;
5044 b.arg(qkv).arg(wq).arg(wk).arg(wv)
5045 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5046 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5047 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5048 unsafe { b.launch(cfg)?; } }
5049 }
5050 Ok(())
5051 }
5052
5053 #[allow(clippy::too_many_arguments)]
5055 pub fn rms_norm_qkv_rope(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>, v0: &CudaSlice<f32>,
5056 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5057 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5058 head_dim: usize, rq: usize, rk: usize,
5059 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5060 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32)
5061 -> Result<(), Box<dyn std::error::Error>> {
5062 let f = self.func("rms_norm_qkv_rope_f32");
5063 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 };
5065 let theta_scale = base.powf(-2.0 / head_dim as f32);
5066 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5067 let __s_b = self.gpu.stream();
5068 let mut b = __s_b.launch_builder(&f);
5069 match ff {
5070 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5071 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5072 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5073 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5074 unsafe { b.launch(cfg)?; } }
5075 None => { let null: u64 = 0;
5076 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5077 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5078 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5079 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5080 unsafe { b.launch(cfg)?; } }
5081 }
5082 Ok(())
5083 }
5084
5085 #[allow(clippy::too_many_arguments)]
5089 pub fn rms_norm_qkv_rope_append_dc(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>,
5090 v0: &CudaSlice<f32>,
5091 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5092 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5093 head_dim: usize, rq: usize, rk: usize,
5094 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5095 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32,
5096 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
5097 t_dev: &CudaSlice<i32>, k_tok_bytes: usize, v_tok_bytes: usize,
5098 g: bool)
5099 -> Result<(), Box<dyn std::error::Error>> {
5100 let rows = rq + rk + rk;
5101 let theta_scale = base.powf(-2.0 / head_dim as f32);
5102 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5103 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
5104 if Self::pdl_on() && Self::pdl_wb_on() {
5105 use cudarc::driver::{DevicePtr, DevicePtrMut};
5106 let s = &self.gpu.stream();
5107 let (p0, _a0) = q0.device_ptr(s); let (p1, _a1) = k0.device_ptr(s);
5108 let (p2, _a2) = v0.device_ptr(s);
5109 let (pwq, _a3) = wq.device_ptr(s); let (pwk, _a4) = wk.device_ptr(s);
5110 let (pwv, _a5) = wv.device_ptr(s);
5111 let (pq, _a6) = q.device_ptr_mut(s); let (pk, _a7) = k.device_ptr_mut(s);
5112 let (pv, _a8) = v.device_ptr_mut(s);
5113 let (pp, _a9) = pos.device_ptr(s);
5114 let pff: u64 = match ff { Some(t) => { let (p, _gg) = t.device_ptr(s); p as u64 }
5115 None => 0 };
5116 let (pkc, _a10) = kc.device_ptr_mut(s); let (pvc, _a11) = vc.device_ptr_mut(s);
5117 let (pt, _a12) = t_dev.device_ptr(s);
5118 let mut ps = [
5119 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
5120 &p2 as *const _ as *mut _, &pwq as *const _ as *mut _,
5121 &pwk as *const _ as *mut _, &pwv as *const _ as *mut _,
5122 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5123 &pv as *const _ as *mut _, &nc as *const _ as *mut _,
5124 &rqi as *const _ as *mut _, &rki as *const _ as *mut _,
5125 &pp as *const _ as *mut _, &nhq as *const _ as *mut _,
5126 &nhk as *const _ as *mut _, &theta_scale as *const _ as *mut _,
5127 &freq_scale as *const _ as *mut _, &pff as *const _ as *mut _,
5128 &eps as *const _ as *mut _, &pkc as *const _ as *mut _,
5129 &pvc as *const _ as *mut _, &pt as *const _ as *mut _,
5130 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
5131 ];
5132 unsafe { self.launch_pdl_flash(g, "rms_norm_qkv_rope_append_dc_f32",
5133 (rows as u32, 1, 1), (rms_block(), 1, 1), 0, &mut ps)?; }
5134 return Ok(());
5135 }
5136 let f = if g { self.func_g("rms_norm_qkv_rope_append_dc_f32") }
5137 else { self.func("rms_norm_qkv_rope_append_dc_f32") };
5138 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5139 let __s_b = self.gpu.stream();
5140 let mut b = __s_b.launch_builder(&f);
5141 match ff {
5142 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5143 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5144 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5145 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps)
5146 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5147 unsafe { b.launch(cfg)?; } }
5148 None => { let null: u64 = 0;
5149 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5150 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5151 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5152 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps)
5153 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5154 unsafe { b.launch(cfg)?; } }
5155 }
5156 Ok(())
5157 }
5158
5159 pub fn add_q8_1(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
5161 ncols: usize, nrows: usize)
5162 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5163 debug_assert!(ncols % 128 == 0);
5164 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5165 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5166 let f = self.func("add_q8_1_f32");
5167 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5168 let nc = ncols as i32;
5169 let __s_b2 = self.gpu.stream();
5170 let mut b2 = __s_b2.launch_builder(&f);
5171 b2.arg(a).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc);
5172 unsafe { b2.launch(cfg)?; }
5173 Ok((out_q, out_d))
5174 }
5175
5176 pub fn rms_pre_add_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>, b: &CudaSlice<f32>,
5180 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
5181 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5182 debug_assert!(ncols % 128 == 0);
5183 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5184 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5185 let f = self.func("rms_pre_add_q8_1_f32");
5186 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1),
5187 shared_mem_bytes: 0 };
5188 let (nc, ep) = (ncols as i32, eps);
5189 let __s_b2 = self.gpu.stream();
5190 let mut b2 = __s_b2.launch_builder(&f);
5191 b2.arg(a).arg(wa).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
5192 unsafe { b2.launch(cfg)?; }
5193 Ok((out_q, out_d))
5194 }
5195
5196 pub fn l2_v2_on(ncols: usize) -> bool {
5200 ncols == 128 && std::env::var("MEMRA_L2_V2").as_deref() != Ok("0")
5201 }
5202
5203 pub fn l2_norm_pp(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5204 dst16: Option<&mut CudaSlice<u8>>, ncols: usize, nrows: usize,
5205 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5206 if Self::l2_v2_on(ncols) {
5207 let f = self.func("l2_norm_pp_v2_f32");
5208 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 };
5210 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
5211 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
5213 let __s_b = self.gpu.stream();
5214 let mut b = __s_b.launch_builder(&f);
5215 b.arg(x).arg(dst).arg(&d16).arg(&nc).arg(&nr).arg(&e);
5216 unsafe { b.launch(cfg)?; }
5217 return Ok(());
5218 }
5219 self.l2_norm(x, dst, ncols, nrows, eps)
5220 }
5221
5222 pub fn l2_norm(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
5223 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5224 let f = self.func("l2_norm_f32");
5225 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
5226 let (nc, e) = (ncols as i32, eps);
5227 let __s_b = self.gpu.stream();
5228 let mut b = __s_b.launch_builder(&f);
5229 b.arg(x).arg(dst).arg(&nc).arg(&e);
5230 unsafe { b.launch(cfg)?; }
5231 Ok(())
5232 }
5233
5234 pub fn l2_norm_decode(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize,
5240 nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5241 let f = self.func("l2_norm_f32");
5242 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
5243 let (nc, e) = (ncols as i32, eps);
5244 let __s_b = self.gpu.stream();
5245 let mut b = __s_b.launch_builder(&f);
5246 b.arg(x).arg(dst).arg(&nc).arg(&e);
5247 unsafe { b.launch(cfg)?; }
5248 Ok(())
5249 }
5250
5251 pub fn rope_neox(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5253 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32, freq_scale: f32)
5254 -> Result<(), Box<dyn std::error::Error>> {
5255 let f = self.func("rope_neox_f32");
5256 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5257 let grid = (n_heads * n_tokens) as u32;
5258 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5259 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5260 let __s_b = self.gpu.stream();
5261 let mut b = __s_b.launch_builder(&f);
5262 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale);
5263 unsafe { b.launch(cfg)?; }
5264 Ok(())
5265 }
5266
5267 pub fn rope_neox_ff(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5269 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32,
5270 freq_scale: f32, ff: &CudaSlice<f32>)
5271 -> Result<(), Box<dyn std::error::Error>> {
5272 let f = self.func("rope_neox_ff_f32");
5273 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5274 let grid = (n_heads * n_tokens) as u32;
5275 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5276 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5277 let __s_b = self.gpu.stream();
5278 let mut b = __s_b.launch_builder(&f);
5279 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale).arg(ff);
5280 unsafe { b.launch(cfg)?; }
5281 Ok(())
5282 }
5283
5284 #[allow(clippy::too_many_arguments)]
5286 pub fn rope_neox2(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
5287 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
5288 nh_q: usize, nh_k: usize, n_tokens: usize, freq_base: f32,
5289 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
5290 -> Result<(), Box<dyn std::error::Error>> {
5291 let f = self.func("rope_neox2_f32");
5292 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5293 let grid = ((nh_q + nh_k) * n_tokens) as u32;
5294 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5295 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);
5296 let __s_b = self.gpu.stream();
5297 let mut b = __s_b.launch_builder(&f);
5298 b.arg(q).arg(k).arg(pos).arg(&hd).arg(&nd).arg(&nq).arg(&nk).arg(&nt)
5299 .arg(&theta_scale).arg(&freq_scale);
5300 match ff {
5301 Some(ffv) => { b.arg(ffv); unsafe { b.launch(cfg)?; } }
5302 None => {
5303 let null: u64 = 0;
5304 b.arg(&null);
5305 unsafe { b.launch(cfg)?; }
5306 }
5307 }
5308 Ok(())
5309 }
5310
5311 pub fn gelu_tanh_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5313 -> Result<(), Box<dyn std::error::Error>> {
5314 let f = self.func("gelu_tanh_mul_f32");
5315 let cfg = LaunchConfig::for_num_elems(n as u32);
5316 let ni = n as i32;
5317 let __s_b = self.gpu.stream();
5318 let mut b = __s_b.launch_builder(&f);
5319 b.arg(gate).arg(up).arg(dst).arg(&ni);
5320 unsafe { b.launch(cfg)?; }
5321 Ok(())
5322 }
5323
5324 pub fn silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5325 -> Result<(), Box<dyn std::error::Error>> {
5326 let f = self.func("silu_mul_f32");
5327 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5329 let ni = n as i32;
5330 let __s_b = self.gpu.stream();
5331 let mut b = __s_b.launch_builder(&f);
5332 b.arg(gate).arg(up).arg(dst).arg(&ni);
5333 unsafe { b.launch(cfg)?; }
5334 Ok(())
5335 }
5336
5337 pub fn silu_mul_f16out(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
5340 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
5341 -> Result<(), Box<dyn std::error::Error>> {
5342 let f = self.func("silu_mul_f16out_f32");
5343 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5344 let ni = n as i32;
5345 let __s_b = self.gpu.stream();
5346 let mut b = __s_b.launch_builder(&f);
5347 b.arg(gate).arg(up).arg(dst).arg(dst16).arg(&ni);
5348 unsafe { b.launch(cfg)?; }
5349 Ok(())
5350 }
5351
5352 pub fn silu_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5359 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
5360 let f = self.func("silu_mul_scaled_f32");
5361 let cfg = LaunchConfig::for_num_elems(n as u32);
5362 let ni = n as i32;
5363 let (gsf, usf) = (gs, us);
5364 let __s_b = self.gpu.stream();
5365 let mut b = __s_b.launch_builder(&f);
5366 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(dst).arg(&ni);
5367 unsafe { b.launch(cfg)?; }
5368 Ok(())
5369 }
5370
5371 #[allow(clippy::too_many_arguments)]
5375 pub fn swigluoai_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5376 alpha: f32, limit: f32, dst: &mut CudaSlice<f32>, n: usize)
5377 -> Result<(), Box<dyn std::error::Error>> {
5378 let f = self.func("swigluoai_mul_scaled_f32");
5379 let cfg = LaunchConfig::for_num_elems(n as u32);
5380 let ni = n as i32;
5381 let __s_b = self.gpu.stream();
5382 let mut b = __s_b.launch_builder(&f);
5383 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&alpha).arg(&limit).arg(dst).arg(&ni);
5384 unsafe { b.launch(cfg)?; }
5385 Ok(())
5386 }
5387
5388 pub fn silu_mul_scaled_q8_1(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5396 n: usize)
5397 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5398 let f = self.func("silu_mul_scaled_q8_1");
5399 let nblk = n / 32;
5400 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);
5404 let (gsf, usf, ni) = (gs, us, n as i32);
5405 let __s_b = self.gpu.stream();
5406 let mut b = __s_b.launch_builder(&f);
5407 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(&mut aq).arg(&mut ad).arg(&ni);
5408 unsafe { b.launch(cfg)?; }
5409 Ok((aq, ad))
5410 }
5411
5412 pub fn add(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5413 -> Result<(), Box<dyn std::error::Error>> {
5414 let f = self.func("add_f32");
5415 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5417 let ni = n as i32;
5418 let __s_bld = self.gpu.stream();
5419 let mut bld = __s_bld.launch_builder(&f);
5420 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5421 unsafe { bld.launch(cfg)?; }
5422 Ok(())
5423 }
5424
5425 pub fn mul(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5426 -> Result<(), Box<dyn std::error::Error>> {
5427 let f = self.func("mul_f32");
5428 let cfg = LaunchConfig::for_num_elems(n as u32);
5429 let ni = n as i32;
5430 let __s_bld = self.gpu.stream();
5431 let mut bld = __s_bld.launch_builder(&f);
5432 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5433 unsafe { bld.launch(cfg)?; }
5434 Ok(())
5435 }
5436
5437 pub fn matmul(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
5440 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5441 use crate::model::GpuTensor;
5442 let in_f = w.in_features();
5443 let out_f = w.out_features();
5444 #[allow(non_snake_case)]
5452 let GEMM_M_THRESHOLD = if self.verify_exact_on() { usize::MAX } else { 16usize };
5455
5456 const GEMM_MIN_OUT_F: usize = 128; if m >= GEMM_M_THRESHOLD {
5481 if let Some(y) = self.try_fp8_gemm(w, x, m)? { return Ok(y); }
5482 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(y); }
5489 if let Some(y) = self.try_f16_gemm(w, x, m)? { return Ok(y); }
5492 }
5493 if let GpuTensor::Quant { qtype, .. } = w {
5508 if *qtype == QT_F8_E4M3_BLK {
5509 if m >= GEMM_M_THRESHOLD {
5510 if let Some(y) = self.try_e4m3_blk_prefill(w, x, m)? { return Ok(y); }
5511 }
5512 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5513 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
5514 }
5515 }
5516 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.mmq_supports(w) {
5517 return self.qmatvec_mmq(w, x, m);
5518 }
5519 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.gemm_supports(w) {
5520 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5521 return self.qmatvec_gemm(w, &aq, &ad, m);
5522 }
5523 if m >= GEMM_M_THRESHOLD {
5526 if let Some(y) = self.try_fp4_gemm(w, x, m, in_f, out_f)? { return Ok(y); }
5527 }
5528 let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
5532 if m == 1 && fast {
5537 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, scale, .. } = w {
5538 if self.mmvq_supports(*qtype) {
5539 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5543 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5544 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, rp);
5545 }
5546 }
5547 }
5548 if (2..=16).contains(&m) && fast && std::env::var("MEMRA_NO_BATCHED").is_err()
5564 && (m <= 4 || Self::b8_enabled()) {
5565 let m_ok = m <= 8 || matches!(w, GpuTensor::Quant { qtype, .. }
5575 if *qtype == QT_Q4_0 || *qtype == QT_Q6_K || *qtype == QT_F8_E4M3
5576 || *qtype == QT_NVFP4 || *qtype == QT_Q4_K || *qtype == QT_Q5_K || *qtype == QT_Q8_0);
5577 if m_ok {
5578 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, .. } = w {
5579 if self.batched_supports(*qtype) && self.mmvq_supports(*qtype) {
5580 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5581 let mcols = Self::batched_mcols(m);
5582 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5583 let mut y = self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, mcols, 1.0, rp)?;
5584 if let GpuTensor::Quant { scale, .. } = w {
5585 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5586 }
5587 return Ok(y);
5588 }
5589 }
5590 }
5591 }
5592 if fast {
5598 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. } = w {
5599 if *qtype == QT_F8_E4M3 {
5600 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5601 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes,
5602 *scale, false);
5603 }
5604 }
5605 }
5606 let mut y = match w {
5607 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q8_0 =>
5608 self.qmatvec_q8_0_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5609 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q4_K =>
5610 self.qmatvec_q4_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5611 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q6_K =>
5612 self.qmatvec_q6_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5613 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q5_K =>
5614 self.qmatvec_q5_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5615 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q3_K =>
5616 self.qmatvec_q3_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5617 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } if fast && *qtype == QT_NVFP4 =>
5618 self.qmatvec_dp4a_named(
5619 if *rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
5620 bytes, x, m, in_f, out_f, *row_bytes)?,
5621 GpuTensor::Quant { bytes, qtype, row_bytes, .. }
5625 if fast && *qtype == QT_IQ4_XS && Self::iq_fast_enabled() =>
5626 self.qmatvec_iq4_XS_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5627 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } =>
5632 self.qmatvec(bytes, x, m, in_f, out_f,
5635 if *rp && *qtype == QT_NVFP4 { QT_NVFP4_RP } else { *qtype },
5636 *row_bytes)?,
5637 GpuTensor::Float { data, .. } => self.linear(x, data, m, in_f, out_f)?,
5638 GpuTensor::FloatBf16 { data, .. } =>
5641 self.linear_bf16_chunked(x, data, m, in_f, out_f, false)?,
5642 };
5643 if let GpuTensor::Quant { scale, .. } = w {
5645 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5646 }
5647 Ok(y)
5648 }
5649
5650 pub fn uses_q8_1_fast(&self, w: &crate::model::GpuTensor) -> bool {
5653 use crate::model::GpuTensor;
5654 if std::env::var("MEMRA_FAST").as_deref() == Ok("0") { return false; }
5655 match w {
5656 GpuTensor::Quant { qtype, .. } => matches!(*qtype,
5663 QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q3_K | QT_NVFP4 | QT_F8_E4M3
5664 | QT_F8_E4M3_BLK | QT_Q4_0)
5665 || (*qtype == QT_IQ4_XS && Self::iq_fast_enabled()),
5666 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
5667 }
5668 }
5669
5670 pub fn matmul_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
5675 x_fallback: &CudaSlice<f32>, m: usize)
5676 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5677 use crate::model::GpuTensor;
5678 let x_raw_ok = x_fallback.len() >= m * w.in_features();
5684 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5687 if let Some(y) = self.try_fp8_gemm(w, x_fallback, m)? { return Ok(y); }
5688 if let Some(y) = self.try_fp8_blk_mmq(w, x_fallback, m)? { return Ok(y); }
5691 if let Some(y) = self.try_f16_gemm(w, x_fallback, m)? { return Ok(y); }
5693 }
5694 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5700 if let Some(y) = self.try_e4m3_blk_prefill(w, x_fallback, m)? { return Ok(y); }
5701 }
5702 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
5703 if m >= 16 && w.out_features() >= 128 && self.mmq_supports(w) && !self.verify_exact_on()
5708 && x_raw_ok {
5709 return self.qmatvec_mmq(w, x_fallback, m);
5710 }
5711 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5714 if let Some(y) = self.try_fp4_gemm(w, x_fallback, m, w.in_features(), w.out_features())? {
5715 return Ok(y);
5716 }
5717 }
5718 if m >= 16 && self.gemm_supports(w) && !self.verify_exact_on() {
5721 return self.qmatvec_gemm(w, aq, ad, m);
5722 }
5723 if !self.uses_q8_1_fast(w) { return self.matmul(w, x_fallback, m); }
5724 let in_f = w.in_features();
5725 let out_f = w.out_features();
5726 let (bytes, qtype, row_bytes, scale, rp) = match w {
5727 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
5728 _ => unreachable!("uses_q8_1_fast guaranteed Quant"),
5729 };
5730 let (mbytes, mrp) = match w {
5733 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5734 _ => (bytes, rp),
5735 };
5736 if m == 1 && self.mmvq_supports(qtype) {
5740 return self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, mrp);
5741 }
5742 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5755 && std::env::var("MEMRA_NO_BATCHED").is_err()
5756 && (m <= 4 || Self::b8_enabled())
5757 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_NVFP4
5761 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_F8_E4M3 || qtype == QT_Q8_0) {
5762 let mcols = Self::batched_mcols(m);
5763 return self.qmatvec_mmvq_batched(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, mrp);
5764 }
5765 if qtype == QT_F8_E4M3 || qtype == QT_Q4_0 {
5771 let (b2, r2) = if qtype == QT_Q4_0 { (mbytes, mrp) } else { (bytes, rp) };
5772 return self.qmatvec_mmvq(b2, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, r2);
5773 }
5774 let name = match qtype {
5775 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
5776 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
5777 QT_Q3_K => "qmatvec_q3_K_dp4a",
5778 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
5779 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
5780 _ => unreachable!(),
5781 };
5782 let f = self.func(name);
5783 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 };
5785 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
5786 let __s_b = self.gpu.stream();
5787 let mut b = __s_b.launch_builder(&f);
5788 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
5789 unsafe { b.launch(cfg)?; }
5790 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
5791 Ok(y)
5792 }
5793
5794 pub fn matmul_decode_exact(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
5802 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5803 use crate::model::GpuTensor;
5804 if let GpuTensor::Float { data, .. } = w {
5812 return self.linear_decode_exact(x, data, m, w.in_features(), w.out_features());
5813 }
5814 if let GpuTensor::FloatBf16 { data, .. } = w {
5817 let (in_f, out_f) = (w.in_features(), w.out_features());
5818 return self.linear_bf16_chunked(x, data, m, in_f, out_f, true);
5819 }
5820 if !self.uses_q8_1_fast(w) { return self.matmul(w, x, m); }
5821 let in_f = w.in_features();
5822 let out_f = w.out_features();
5823 let (bytes, qtype, row_bytes, scale, rp) = match w {
5824 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
5825 _ => return self.matmul(w, x, m),
5826 };
5827 let (bytes, rp) = match w {
5830 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5831 _ => (bytes, rp),
5832 };
5833 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5834 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
5838 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5847 && std::env::var("MEMRA_NO_BATCHED").is_err()
5848 && (m <= 4 || Self::b8_enabled())
5849 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
5852 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
5853 let mcols = Self::batched_mcols(m);
5854 return self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
5855 }
5856 if self.mmvq_supports(qtype) {
5857 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
5860 }
5861 self.matmul_pre(w, &aq, &ad, x, m)
5864 }
5865
5866 pub fn matmul_decode_exact_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
5876 ad: &CudaSlice<f32>, m: usize)
5877 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5878 use crate::model::GpuTensor;
5879 debug_assert!(self.uses_q8_1_fast(w),
5880 "matmul_decode_exact_pre: caller must guarantee q8_1-fast");
5881 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
5883 let in_f = w.in_features();
5884 let out_f = w.out_features();
5885 let (bytes, qtype, row_bytes, scale, rp) = match w {
5886 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
5887 (bytes, *qtype, *row_bytes, *scale, *rp),
5888 _ => return Err("matmul_decode_exact_pre: Quant tensor required (q8_1-fast contract)".into()),
5889 };
5890 let (bytes, rp) = match w {
5892 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5893 _ => (bytes, rp),
5894 };
5895 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5897 && std::env::var("MEMRA_NO_BATCHED").is_err()
5898 && (m <= 4 || Self::b8_enabled())
5899 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
5900 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
5901 let mcols = Self::batched_mcols(m);
5902 return self.qmatvec_mmvq_batched(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
5903 }
5904 if self.mmvq_supports(qtype) {
5905 return self.qmatvec_mmvq(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
5906 }
5907 let x0 = self.zeros(0)?;
5910 self.matmul_pre(w, aq, ad, &x0, m)
5911 }
5912
5913 pub fn matmul_decode_exact_dual_pre(&self, w0: &crate::model::GpuTensor,
5922 w1: &crate::model::GpuTensor,
5923 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
5924 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
5925 use crate::model::GpuTensor;
5926 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5927 let on = *ON.get_or_init(|| {
5928 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
5929 });
5930 if !on || !(2..=7).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
5931 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
5932 return Ok(None);
5933 }
5934 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
5939 let (in_f, out_f) = (w0.in_features(), w0.out_features());
5940 if w1.in_features() != in_f || w1.out_features() != out_f {
5941 return Ok(None);
5942 }
5943 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
5944 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
5945 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
5946 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
5947 (b0, b1, *rb0, *s0, *s1, *rp0),
5948 _ => return Ok(None),
5949 };
5950 if m > 4 && !(rp && Self::b8_enabled()
5953 && std::env::var("MEMRA_B567").as_deref() != Ok("0")) {
5954 return Ok(None);
5955 }
5956 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, aq, ad, m, in_f, out_f, row_bytes, rp)?;
5957 Ok(Some(((y0, s0), (y1, s1))))
5958 }
5959
5960 pub fn matmul_decode_exact_dual(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
5976 x: &CudaSlice<f32>, m: usize)
5977 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
5978 use crate::model::GpuTensor;
5979 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5980 let on = *ON.get_or_init(|| {
5981 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
5982 });
5983 if !on || !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
5984 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
5985 return Ok(None);
5986 }
5987 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
5992 let (in_f, out_f) = (w0.in_features(), w0.out_features());
5993 if w1.in_features() != in_f || w1.out_features() != out_f {
5994 return Ok(None);
5995 }
5996 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
5997 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
5998 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
5999 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
6000 (b0, b1, *rb0, *s0, *s1, *rp0),
6001 _ => return Ok(None),
6002 };
6003 if std::env::var("MEMRA_DEBUG").is_ok() {
6006 static ONCE: std::sync::Once = std::sync::Once::new();
6007 ONCE.call_once(|| eprintln!("[memra] dual gate+up batched ENGAGED (m={m} rp={rp})"));
6008 }
6009 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6010 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, &aq, &ad, m, in_f, out_f, row_bytes, rp)?;
6011 let mut y0 = y0;
6012 let mut y1 = y1;
6013 if s0 != 1.0 { self.scale_inplace(&mut y0, s0, m * out_f)?; }
6014 if s1 != 1.0 { self.scale_inplace(&mut y1, s1, m * out_f)?; }
6015 Ok(Some((y0, y1)))
6016 }
6017
6018 #[allow(clippy::too_many_arguments)]
6023 pub fn qmatvec_batched_dual_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6024 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6025 m: usize, in_f: usize, out_f: usize, row_bytes: usize, rp: bool)
6026 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6027 const ROWS_PER_BLOCK: u32 = 4;
6028 let mcols = Self::batched_mcols(m);
6029 let (name, rows_per_block) = match (mcols, rp, m) {
6032 (2, false, _) => ("qmatvec_nvfp4_mmvq_dual_b2", ROWS_PER_BLOCK),
6033 (4, false, _) => ("qmatvec_nvfp4_mmvq_dual_b4_r2", ROWS_PER_BLOCK * 2),
6034 (2, true, _) => ("qmatvec_nvfp4_mmvq_dual_b2_rp", ROWS_PER_BLOCK),
6035 (4, true, _) => ("qmatvec_nvfp4_mmvq_dual_b4_rpr2", ROWS_PER_BLOCK * 2),
6036 (8, true, 5) => ("qmatvec_nvfp4_mmvq_dual_b5_rpr2", ROWS_PER_BLOCK * 2),
6037 (8, true, 6) => ("qmatvec_nvfp4_mmvq_dual_b6_rpr2", ROWS_PER_BLOCK * 2),
6038 (8, true, 7) => ("qmatvec_nvfp4_mmvq_dual_b7_rpr2", ROWS_PER_BLOCK * 2),
6039 _ => return Err(format!("qmatvec_batched_dual_raw: no dual kernel for m {m}").into()),
6040 };
6041 let f = self.func(name);
6042 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
6043 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
6044 let cfg = LaunchConfig {
6045 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6046 block_dim: (32, ROWS_PER_BLOCK, 1),
6047 shared_mem_bytes: 0,
6048 };
6049 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
6050 let __s_b = self.gpu.stream();
6051 let mut b = __s_b.launch_builder(&f);
6052 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6053 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
6054 unsafe { b.launch(cfg)?; }
6055 Ok((y0, y1))
6056 }
6057
6058 pub fn matmul_pre_dual_noscale(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6070 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6071 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
6072 use crate::model::GpuTensor;
6073 if m != 1 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6074 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6084 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6085 if w1.in_features() != in_f || w1.out_features() != out_f { return Ok(None); }
6086 let no_mirror = |w: &crate::model::GpuTensor| {
6099 !matches!(w, GpuTensor::Quant { rp4: Some(_), .. })
6100 };
6101 if self.q8_ffn_fuse2_on()
6102 && no_mirror(w0) && no_mirror(w1)
6103 && let Some([p0, p1]) = self.q8_fused_params(&[w0, w1])
6104 {
6105 let (y0, y1) = self.q8_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2)?;
6106 return Ok(Some(((y0, 1.0), (y1, 1.0))));
6107 }
6108 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6118 let (y0, y1) = self.e4m3_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2,
6119 1.0, 1.0)?;
6120 return Ok(Some(((y0, p0.3), (y1, p1.3))));
6121 }
6122 let (b0, q0, rb0, s0, rp0) = match w0 {
6123 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6124 _ => return Ok(None),
6125 };
6126 let (b1, q1, rb1, s1, rp1) = match w1 {
6127 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6128 _ => return Ok(None),
6129 };
6130 if q0 != QT_NVFP4 || q1 != QT_NVFP4 || rb0 != rb1 || rp0 != rp1 { return Ok(None); }
6131 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
6133 let rows_per_block = ROWS_PER_BLOCK * RPW;
6134 let f = self.func(if rp0 { "qmatvec_nvfp4_mmvq_dual_mr2_rp" } else { "qmatvec_nvfp4_mmvq_dual_mr2" });
6135 let mut y0 = self.alloc_uninit::<f32>(out_f)?;
6136 let mut y1 = self.alloc_uninit::<f32>(out_f)?;
6137 let cfg = LaunchConfig {
6138 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6139 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0,
6140 };
6141 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, rb0 as i64);
6142 let one = 1.0f32;
6145 let __s_b = self.gpu.stream();
6146 let mut b = __s_b.launch_builder(&f);
6147 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6148 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&one).arg(&one);
6149 unsafe { b.launch(cfg)?; }
6150 Ok(Some(((y0, s0), (y1, s1))))
6151 }
6152
6153 pub fn matmul_q8_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6161 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6162 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6163 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6169 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(),
6170 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6171 }
6172 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6173 Ok(Some(self.q8_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6174 }
6175
6176 #[allow(clippy::too_many_arguments)]
6177 fn q8_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6178 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6179 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6180 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6181 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6183 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6184 let f = self.func("qmatvec_q8_0_mmvq_fused2");
6185 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6186 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6187 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6188 shared_mem_bytes: 0 };
6189 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
6190 let __s_b = self.gpu.stream();
6191 let mut b = __s_b.launch_builder(&f);
6192 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6193 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl);
6194 unsafe { b.launch(cfg)?; }
6195 Ok((y0, y1))
6196 }
6197
6198 pub fn matmul_q8_fused2_x(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6204 x: &CudaSlice<f32>)
6205 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6206 if !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6207 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6208 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6209 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(),
6210 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6211 }
6212 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6213 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6214 Ok(Some(self.q8_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6215 }
6216
6217 #[allow(clippy::too_many_arguments)]
6220 pub fn qmatvec_q8_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
6221 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6222 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6223 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6224 self.q8_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes)
6225 }
6226
6227 pub fn matmul_q4_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6233 w2: &crate::model::GpuTensor,
6234 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6235 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6236 use crate::model::GpuTensor;
6237 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6238 match w {
6239 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6240 Some((*row_bytes, w.out_features())),
6241 _ => None,
6242 }
6243 };
6244 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6245 else { return Ok(None) };
6246 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6247 return Ok(None);
6248 }
6249 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6253 match w {
6254 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6255 Some(m) => (m, true),
6256 None => (bytes, *rp),
6257 },
6258 _ => unreachable!(),
6259 }
6260 }
6261 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6262 if rp0 != rp1 || rp1 != rp2 { return Ok(None); }
6263 let rp = rp0;
6264 let rpb: u32 = 4;
6265 let mr1 = rp && Self::q40_mr1_on();
6269 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6270 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6271 let grid = nb(o0) + nb(o1) + nb(o2);
6272 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6273 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6274 let mut y2 = self.alloc_uninit::<f32>(o2)?;
6275 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6276 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6277 else { "qmatvec_q4_0_mmvq_fused3" });
6278 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6279 let inf = w0.in_features() as i32;
6280 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6281 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6282 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6285 {
6286 use cudarc::driver::{DevicePtr, DevicePtrMut};
6287 let s = &self.gpu.stream();
6288 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6289 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6290 let (pad, _g4) = ad.device_ptr(s);
6291 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6292 let (py2, _g7) = y2.device_ptr_mut(s);
6293 let mut ps = [
6294 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6295 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6296 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6297 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6298 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6299 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6300 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6301 &r2 as *const _ as *mut _,
6302 ];
6303 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6304 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6305 }
6306 return Ok(Some((y0, y1, y2)));
6307 }
6308 let __s_b = self.gpu.stream();
6309 let mut b = __s_b.launch_builder(&f);
6310 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6311 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6312 unsafe { b.launch(cfg)?; }
6313 Ok(Some((y0, y1, y2)))
6314 }
6315
6316 #[allow(clippy::too_many_arguments)]
6319 pub fn matmul_q4_fused3_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6320 w2: &crate::model::GpuTensor,
6321 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6322 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>,
6323 y2: &mut CudaSlice<f32>)
6324 -> Result<bool, Box<dyn std::error::Error>> {
6325 use crate::model::GpuTensor;
6326 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6327 match w {
6328 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6329 Some((*row_bytes, w.out_features())),
6330 _ => None,
6331 }
6332 };
6333 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6334 else { return Ok(false) };
6335 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6336 return Ok(false);
6337 }
6338 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6339 match w {
6340 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6341 Some(m) => (m, true),
6342 None => (bytes, *rp),
6343 },
6344 _ => unreachable!(),
6345 }
6346 }
6347 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6348 if rp0 != rp1 || rp1 != rp2 { return Ok(false); }
6349 let rp = rp0;
6350 let rpb: u32 = 4;
6351 let mr1 = rp && Self::q40_mr1_on();
6352 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6353 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6354 let grid = nb(o0) + nb(o1) + nb(o2);
6355 debug_assert!(y0.len() >= o0 && y1.len() >= o1 && y2.len() >= o2);
6356 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6357 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6358 else { "qmatvec_q4_0_mmvq_fused3" });
6359 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6360 let inf = w0.in_features() as i32;
6361 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6362 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6363 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6365 use cudarc::driver::{DevicePtr, DevicePtrMut};
6366 let s = &self.gpu.stream();
6367 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6368 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6369 let (pad, _g4) = ad.device_ptr(s);
6370 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6371 let (py2, _g7) = y2.device_ptr_mut(s);
6372 let mut ps = [
6373 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6374 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6375 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6376 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6377 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6378 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6379 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6380 &r2 as *const _ as *mut _,
6381 ];
6382 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6383 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6384 return Ok(true);
6385 }
6386 let __s_b = self.gpu.stream();
6387 let mut b = __s_b.launch_builder(&f);
6388 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1).arg(&mut *y2)
6389 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6390 unsafe { b.launch(cfg)?; }
6391 Ok(true)
6392 }
6393
6394 pub fn matmul_q4_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6396 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6397 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6398 use crate::model::GpuTensor;
6399 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6400 match w {
6401 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6402 Some((*row_bytes, w.out_features())),
6403 _ => None,
6404 }
6405 };
6406 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6407 if w0.in_features() != w1.in_features() { return Ok(None); }
6408 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6410 match w {
6411 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6412 Some(m) => (m, true),
6413 None => (bytes, *rp),
6414 },
6415 _ => unreachable!(),
6416 }
6417 }
6418 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6419 if rp0 != rp1 { return Ok(None); }
6420 let rp = rp0;
6421 let rpb: u32 = 4;
6422 let mr1 = rp && Self::q40_mr1_on();
6424 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6425 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6426 let grid = nb(o0) + nb(o1);
6427 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6428 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6429 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6430 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6431 else { "qmatvec_q4_0_mmvq_fused2" });
6432 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6433 let inf = w0.in_features() as i32;
6434 let (oo0, oo1) = (o0 as i32, o1 as i32);
6435 let (r0, r1) = (rb0 as i64, rb1 as i64);
6436 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6438 {
6439 use cudarc::driver::{DevicePtr, DevicePtrMut};
6440 let s = &self.gpu.stream();
6441 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6442 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6443 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6444 let mut ps = [
6445 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6446 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6447 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6448 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6449 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6450 &r1 as *const _ as *mut _,
6451 ];
6452 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6453 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6454 }
6455 return Ok(Some((y0, y1)));
6456 }
6457 let __s_b = self.gpu.stream();
6458 let mut b = __s_b.launch_builder(&f);
6459 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6460 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6461 unsafe { b.launch(cfg)?; }
6462 Ok(Some((y0, y1)))
6463 }
6464
6465 pub fn matmul_q4_fused2_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6467 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6468 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>)
6469 -> Result<bool, Box<dyn std::error::Error>> {
6470 use crate::model::GpuTensor;
6471 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6472 match w {
6473 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6474 Some((*row_bytes, w.out_features())),
6475 _ => None,
6476 }
6477 };
6478 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(false) };
6479 if w0.in_features() != w1.in_features() { return Ok(false); }
6480 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6481 match w {
6482 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6483 Some(m) => (m, true),
6484 None => (bytes, *rp),
6485 },
6486 _ => unreachable!(),
6487 }
6488 }
6489 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6490 if rp0 != rp1 { return Ok(false); }
6491 let rp = rp0;
6492 let rpb: u32 = 4;
6493 let mr1 = rp && Self::q40_mr1_on();
6494 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6495 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6496 let grid = nb(o0) + nb(o1);
6497 debug_assert!(y0.len() >= o0 && y1.len() >= o1);
6498 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6499 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6500 else { "qmatvec_q4_0_mmvq_fused2" });
6501 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6502 let inf = w0.in_features() as i32;
6503 let (oo0, oo1) = (o0 as i32, o1 as i32);
6504 let (r0, r1) = (rb0 as i64, rb1 as i64);
6505 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6507 use cudarc::driver::{DevicePtr, DevicePtrMut};
6508 let s = &self.gpu.stream();
6509 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6510 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6511 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6512 let mut ps = [
6513 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6514 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6515 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6516 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6517 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6518 &r1 as *const _ as *mut _,
6519 ];
6520 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6521 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6522 return Ok(true);
6523 }
6524 let __s_b = self.gpu.stream();
6525 let mut b = __s_b.launch_builder(&f);
6526 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1)
6527 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6528 unsafe { b.launch(cfg)?; }
6529 Ok(true)
6530 }
6531
6532 pub fn matmul_q4_fused2_batched(&self, w0: &crate::model::GpuTensor,
6537 w1: &crate::model::GpuTensor,
6538 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6539 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6540 use crate::model::GpuTensor;
6541 if m < 2 || m > 8 { return Ok(None); }
6542 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6543 match w {
6544 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6545 Some((*row_bytes, w.out_features())),
6546 _ => None,
6547 }
6548 };
6549 let (Some((rb0, o0)), Some((_rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6550 if w0.in_features() != w1.in_features() { return Ok(None); }
6551 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6552 match w {
6553 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6554 Some(mr) => (mr, true),
6555 None => (bytes, *rp),
6556 },
6557 _ => unreachable!(),
6558 }
6559 }
6560 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6561 if !rp0 || !rp1 { return Ok(None); }
6562 let mcols = Self::batched_mcols(m);
6563 let rpb: u32 = 4;
6564 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6565 let grid = nb(o0) + nb(o1);
6566 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6567 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6568 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f2_rp",
6569 4 => "qmatvec_q4_0_mmvq_b4_f2_rp",
6570 _ => "qmatvec_q4_0_mmvq_b8_f2_rp" });
6571 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6572 shared_mem_bytes: 0 };
6573 let inf = w0.in_features() as i32;
6574 let (oo0, oo1, mi) = (o0 as i32, o1 as i32, m as i32);
6575 let rb = rb0 as i64;
6576 let __s_b = self.gpu.stream();
6577 let mut b = __s_b.launch_builder(&f);
6578 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6579 .arg(&inf).arg(&oo0).arg(&oo1).arg(&mi).arg(&rb);
6580 unsafe { b.launch(cfg)?; }
6581 Ok(Some((y0, y1)))
6582 }
6583
6584 #[allow(clippy::too_many_arguments)]
6587 pub fn matmul_q4_fused3_batched(&self, w0: &crate::model::GpuTensor,
6588 w1: &crate::model::GpuTensor, w2: &crate::model::GpuTensor,
6589 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6590 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6591 use crate::model::GpuTensor;
6592 if m < 2 || m > 8 { return Ok(None); }
6593 let q4 = |w: &GpuTensor| -> Option<usize> {
6594 match w {
6595 GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0 => Some(w.out_features()),
6596 _ => None,
6597 }
6598 };
6599 let (Some(o0), Some(o1), Some(o2)) = (q4(w0), q4(w1), q4(w2)) else { return Ok(None) };
6600 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6601 return Ok(None);
6602 }
6603 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6604 match w {
6605 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6606 Some(mr) => (mr, true),
6607 None => (bytes, *rp),
6608 },
6609 _ => unreachable!(),
6610 }
6611 }
6612 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6613 if !rp0 || !rp1 || !rp2 { return Ok(None); }
6614 let mcols = Self::batched_mcols(m);
6615 let rpb: u32 = 4;
6616 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6617 let grid = nb(o0) + nb(o1) + nb(o2);
6618 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6619 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6620 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
6621 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f3_rp",
6622 4 => "qmatvec_q4_0_mmvq_b4_f3_rp",
6623 _ => "qmatvec_q4_0_mmvq_b8_f3_rp" });
6624 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6625 shared_mem_bytes: 0 };
6626 let inf = w0.in_features() as i32;
6627 let (oo0, oo1, oo2, mi) = (o0 as i32, o1 as i32, o2 as i32, m as i32);
6628 let rb = 0i64;
6629 let __s_b = self.gpu.stream();
6630 let mut b = __s_b.launch_builder(&f);
6631 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6632 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&mi).arg(&rb);
6633 unsafe { b.launch(cfg)?; }
6634 Ok(Some((y0, y1, y2)))
6635 }
6636
6637 pub fn matmul_q8_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6638 w2: &crate::model::GpuTensor,
6639 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6640 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6641 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
6644 return Ok(Some(self.e4m3_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6645 p0.1, p1.1, p2.1, p0.2,
6646 p0.3, p1.3, p2.3)?));
6647 }
6648 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
6649 Ok(Some(self.q8_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6650 p0.1, p1.1, p2.1, p0.2)?))
6651 }
6652
6653 #[allow(clippy::too_many_arguments)]
6654 fn q8_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6655 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6656 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
6657 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6658 const ROWS_PER_BLOCK: u32 = 4;
6659 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6660 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6661 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6662 let f = self.func("qmatvec_q8_0_mmvq_fused3");
6663 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6664 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6665 let mut y2 = self.alloc_uninit::<f32>(out2)?;
6666 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6667 shared_mem_bytes: 0 };
6668 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32, row_bytes as i64);
6669 let __s_b = self.gpu.stream();
6670 let mut b = __s_b.launch_builder(&f);
6671 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6672 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl);
6673 unsafe { b.launch(cfg)?; }
6674 Ok((y0, y1, y2))
6675 }
6676
6677 #[allow(clippy::too_many_arguments)]
6679 pub fn qmatvec_q8_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6680 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
6681 out2: usize, row_bytes: usize)
6682 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6683 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6684 self.q8_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes)
6685 }
6686
6687 pub fn matmul_q8_fused2_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6698 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6699 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6700 if !(2..=8).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
6704 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6707 if m > 4 && !Self::b8_enabled() { return Ok(None); }
6708 return Ok(Some(self.e4m3_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(),
6709 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6710 }
6711 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6712 Ok(Some(self.q8_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(), p0.1, p1.1, p0.2)?))
6713 }
6714
6715 #[allow(clippy::too_many_arguments)]
6716 fn q8_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6717 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6718 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6719 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6720 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6722 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6723 let f = self.func(match Self::batched_mcols(m) {
6724 2 => "qmatvec_q8_0_mmvq_fused2_b2",
6725 4 => "qmatvec_q8_0_mmvq_fused2_b4",
6726 _ => "qmatvec_q8_0_mmvq_fused2_b8",
6728 });
6729 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6730 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6731 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6732 shared_mem_bytes: 0 };
6733 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32, row_bytes as i64);
6734 let __s_b = self.gpu.stream();
6735 let mut b = __s_b.launch_builder(&f);
6736 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6737 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
6738 unsafe { b.launch(cfg)?; }
6739 Ok((y0, y1))
6740 }
6741
6742 #[allow(clippy::too_many_arguments)]
6745 pub fn qmatvec_q8_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6746 x: &CudaSlice<f32>, m: usize,
6747 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6748 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6749 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6750 self.q8_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes)
6751 }
6752
6753 #[allow(clippy::too_many_arguments)]
6756 pub fn matmul_q8_fused3_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6757 w2: &crate::model::GpuTensor,
6758 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6759 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6760 if !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
6761 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
6762 return Ok(Some(self.e4m3_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
6763 p0.1, p1.1, p2.1, p0.2,
6764 p0.3, p1.3, p2.3)?));
6765 }
6766 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
6767 Ok(Some(self.q8_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
6768 p0.1, p1.1, p2.1, p0.2)?))
6769 }
6770
6771 #[allow(clippy::too_many_arguments)]
6772 fn q8_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6773 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6774 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
6775 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6776 const ROWS_PER_BLOCK: u32 = 4;
6777 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6778 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6779 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6780 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_q8_0_mmvq_fused3_b2" }
6781 else { "qmatvec_q8_0_mmvq_fused3_b4" });
6782 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6783 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6784 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
6785 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6786 shared_mem_bytes: 0 };
6787 let (inf, o0, o1, o2, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
6788 m as i32, row_bytes as i64);
6789 let __s_b = self.gpu.stream();
6790 let mut b = __s_b.launch_builder(&f);
6791 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6792 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&mi).arg(&rbl);
6793 unsafe { b.launch(cfg)?; }
6794 Ok((y0, y1, y2))
6795 }
6796
6797 #[allow(clippy::too_many_arguments)]
6799 pub fn qmatvec_q8_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6800 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
6801 out1: usize, out2: usize, row_bytes: usize)
6802 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6803 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6804 self.q8_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes)
6805 }
6806
6807 pub fn q8_ffn_fuse2_on(&self) -> bool {
6811 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6812 *ON.get_or_init(|| std::env::var("MEMRA_Q8_FFN_FUSE2").as_deref() != Ok("0"))
6813 }
6814
6815 #[allow(clippy::type_complexity)]
6821 fn q8_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
6822 -> Option<[(&'w CudaSlice<u8>, usize, usize); N]> {
6823 use crate::model::GpuTensor;
6824 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return None; }
6825 if std::env::var("MEMRA_Q8_DUAL").is_ok_and(|v| v == "0") { return None; }
6826 let in_f = ws[0].in_features();
6827 let mut out: [Option<(&CudaSlice<u8>, usize, usize)>; N] = [None; N];
6828 for (i, w) in ws.iter().enumerate() {
6829 match w {
6830 GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. }
6831 if *qtype == QT_Q8_0 && *scale == 1.0 && w.in_features() == in_f =>
6832 out[i] = Some((bytes, w.out_features(), *row_bytes)),
6833 _ => return None,
6834 }
6835 }
6836 Some(out.map(|o| o.unwrap()))
6837 }
6838
6839 pub fn e4m3_dual_on(&self) -> bool {
6842 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6843 *ON.get_or_init(|| std::env::var("MEMRA_E4M3_DUAL").as_deref() != Ok("0"))
6844 }
6845
6846 #[allow(clippy::type_complexity)]
6858 fn e4m3_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
6859 -> Option<[(&'w CudaSlice<u8>, usize, usize, f32); N]> {
6860 use crate::model::GpuTensor;
6861 if !self.e4m3_dual_on() { return None; }
6862 let in_f = ws[0].in_features();
6863 let mut out: [Option<(&CudaSlice<u8>, usize, usize, f32)>; N] = [None; N];
6864 for (i, w) in ws.iter().enumerate() {
6865 match w {
6866 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, rp4, .. }
6867 if *qtype == QT_F8_E4M3 && w.in_features() == in_f
6868 && *row_bytes == in_f && !*rp && rp4.is_none() =>
6869 out[i] = Some((bytes, w.out_features(), *row_bytes, *scale)),
6870 _ => return None,
6871 }
6872 }
6873 Some(out.map(|o| o.unwrap()))
6874 }
6875
6876 #[allow(clippy::too_many_arguments)]
6880 fn e4m3_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6881 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6882 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
6883 ws0: f32, ws1: f32)
6884 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6885 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6887 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6888 let f = self.func("qmatvec_e4m3_mmvq_fused2");
6889 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6890 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6891 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6892 shared_mem_bytes: 0 };
6893 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
6894 let __s_b = self.gpu.stream();
6895 let mut b = __s_b.launch_builder(&f);
6896 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6897 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl).arg(&ws0).arg(&ws1);
6898 unsafe { b.launch(cfg)?; }
6899 Ok((y0, y1))
6900 }
6901
6902 #[allow(clippy::too_many_arguments)]
6904 fn e4m3_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6905 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6906 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
6907 ws0: f32, ws1: f32, ws2: f32)
6908 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6909 const ROWS_PER_BLOCK: u32 = 4;
6910 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6911 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6912 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6913 let f = self.func("qmatvec_e4m3_mmvq_fused3");
6914 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6915 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6916 let mut y2 = self.alloc_uninit::<f32>(out2)?;
6917 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6918 shared_mem_bytes: 0 };
6919 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
6920 row_bytes as i64);
6921 let __s_b = self.gpu.stream();
6922 let mut b = __s_b.launch_builder(&f);
6923 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6924 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl).arg(&ws0).arg(&ws1).arg(&ws2);
6925 unsafe { b.launch(cfg)?; }
6926 Ok((y0, y1, y2))
6927 }
6928
6929 #[allow(clippy::too_many_arguments)]
6933 fn e4m3_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6934 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6935 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
6936 ws0: f32, ws1: f32)
6937 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6938 const ROWS_PER_BLOCK: u32 = 4;
6939 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6940 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6941 let f = self.func(match Self::batched_mcols(m) {
6942 2 => "qmatvec_e4m3_mmvq_fused2_b2",
6943 4 => "qmatvec_e4m3_mmvq_fused2_b4",
6944 _ => "qmatvec_e4m3_mmvq_fused2_b8",
6945 });
6946 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6947 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6948 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6949 shared_mem_bytes: 0 };
6950 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32,
6951 row_bytes as i64);
6952 let __s_b = self.gpu.stream();
6953 let mut b = __s_b.launch_builder(&f);
6954 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6955 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
6956 unsafe { b.launch(cfg)?; }
6957 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
6958 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
6959 Ok((y0, y1))
6960 }
6961
6962 #[allow(clippy::too_many_arguments)]
6964 fn e4m3_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6965 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6966 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
6967 ws0: f32, ws1: f32, ws2: f32)
6968 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6969 const ROWS_PER_BLOCK: u32 = 4;
6970 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6971 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6972 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6973 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_e4m3_mmvq_fused3_b2" }
6974 else { "qmatvec_e4m3_mmvq_fused3_b4" });
6975 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6976 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6977 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
6978 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6979 shared_mem_bytes: 0 };
6980 let (inf, o0, o1, o2, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
6981 m as i32, row_bytes as i64);
6982 let __s_b = self.gpu.stream();
6983 let mut b = __s_b.launch_builder(&f);
6984 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6985 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&mi).arg(&rbl);
6986 unsafe { b.launch(cfg)?; }
6987 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
6988 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
6989 if ws2 != 1.0 { self.scale_inplace(&mut y2, ws2, m * out2)?; }
6990 Ok((y0, y1, y2))
6991 }
6992
6993 pub fn qmatvec_e4m3_blk_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7003 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7004 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7005 scale_cols: usize)
7006 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7007 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,
7009 scale_cols, &mut y)?;
7010 Ok(y)
7011 }
7012
7013 #[allow(clippy::too_many_arguments)]
7015 pub fn qmatvec_e4m3_blk_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7016 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7017 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7018 scale_cols: usize, y: &mut CudaSlice<f32>)
7019 -> Result<(), Box<dyn std::error::Error>> {
7020 const ROWS_PER_BLOCK: u32 = 4; let f = self.func("qmatvec_e4m3_blk_mmvq");
7022 let cfg = LaunchConfig {
7023 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), m as u32, 1),
7024 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7027 let (inf, outf, mi, rb, sc) =
7028 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7029 let __s_b = self.gpu.stream();
7030 let mut b = __s_b.launch_builder(&f);
7031 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut *y)
7032 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7033 unsafe { b.launch(cfg)?; }
7034 Ok(())
7035 }
7036
7037 #[allow(clippy::too_many_arguments)]
7043 pub fn qmatvec_e4m3_blk_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7044 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7045 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7046 scale_cols: usize, mcols: usize)
7047 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7048 const ROWS_PER_BLOCK: u32 = 4; debug_assert!(mcols >= m, "blk batched: mcols {mcols} < m {m}");
7050 let name = match mcols {
7051 2 => "qmatvec_e4m3_blk_mmvq_b2",
7052 4 => "qmatvec_e4m3_blk_mmvq_b4",
7053 8 => "qmatvec_e4m3_blk_mmvq_b8",
7054 16 => "qmatvec_e4m3_blk_mmvq_b16",
7055 _ => return Err(format!("qmatvec_e4m3_blk_mmvq_batched: no kernel for mcols {mcols}").into()),
7056 };
7057 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7058 let f = self.func(name);
7059 let cfg = LaunchConfig {
7060 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
7061 block_dim: (32, ROWS_PER_BLOCK, 1),
7062 shared_mem_bytes: 0,
7063 };
7064 let (inf, outf, mi, rb, sc) =
7065 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7066 let __s_b = self.gpu.stream();
7067 let mut b = __s_b.launch_builder(&f);
7068 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut y)
7069 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7070 unsafe { b.launch(cfg)?; }
7071 Ok(y)
7072 }
7073
7074 #[allow(clippy::too_many_arguments)]
7077 pub fn qmatvec_e4m3_blk_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7078 scales: &CudaSlice<f32>, m: usize, in_f: usize,
7079 out_f: usize, row_bytes: usize, scale_cols: usize,
7080 mcols: usize)
7081 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7082 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7083 self.qmatvec_e4m3_blk_mmvq_batched(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes,
7084 scale_cols, mcols)
7085 }
7086
7087 #[allow(clippy::too_many_arguments)]
7090 pub fn qmatvec_e4m3_blk_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7091 scales: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
7092 row_bytes: usize, scale_cols: usize)
7093 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7094 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7095 self.qmatvec_e4m3_blk_mmvq(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols)
7096 }
7097
7098 #[allow(clippy::too_many_arguments)]
7101 pub fn qmatvec_e4m3_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
7102 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7103 ws0: f32, ws1: f32)
7104 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7105 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7106 self.e4m3_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes, ws0, ws1)
7107 }
7108
7109 #[allow(clippy::too_many_arguments)]
7110 pub fn qmatvec_e4m3_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7111 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
7112 out2: usize, row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7113 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7114 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7115 self.e4m3_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes,
7116 ws0, ws1, ws2)
7117 }
7118
7119 #[allow(clippy::too_many_arguments)]
7120 pub fn qmatvec_e4m3_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7121 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
7122 out1: usize, row_bytes: usize, ws0: f32, ws1: f32)
7123 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7124 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7125 self.e4m3_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes, ws0, ws1)
7126 }
7127
7128 #[allow(clippy::too_many_arguments)]
7129 pub fn qmatvec_e4m3_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7130 b2: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
7131 in_f: usize, out0: usize, out1: usize, out2: usize,
7132 row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7133 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7134 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7135 self.e4m3_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes,
7136 ws0, ws1, ws2)
7137 }
7138
7139 fn try_e4m3_blk_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
7150 ad: &CudaSlice<f32>, m: usize)
7151 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7152 use crate::model::GpuTensor;
7153 if let GpuTensor::Quant { bytes, qtype, row_bytes, blk: Some(g), .. } = w {
7154 if *qtype == QT_F8_E4M3_BLK {
7155 if (2..=16).contains(&m) && std::env::var("MEMRA_NO_BATCHED").is_err()
7161 && (m <= 4 || Self::b8_enabled()) {
7162 let mcols = Self::batched_mcols(m);
7163 return Ok(Some(self.qmatvec_e4m3_blk_mmvq_batched(
7164 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7165 *row_bytes, g.cols, mcols)?));
7166 }
7167 return Ok(Some(self.qmatvec_e4m3_blk_mmvq(
7168 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7169 *row_bytes, g.cols)?));
7170 }
7171 }
7172 Ok(None)
7173 }
7174
7175 fn try_e4m3_blk_prefill(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
7222 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7223 use crate::model::GpuTensor;
7224 let GpuTensor::Quant { bytes, qtype, blk: Some(g), .. } = w else { return Ok(None) };
7225 if *qtype != QT_F8_E4M3_BLK { return Ok(None) }
7226 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(Some(y)); }
7231 let (in_f, out_f) = (w.in_features(), w.out_features());
7232 let slab = self.fp8_blk_dequant_q8_0_dev(bytes, &g.scales, out_f, in_f)?;
7233 let tmp = GpuTensor::Quant {
7234 bytes: slab,
7235 qtype: QT_Q8_0,
7236 row_bytes: in_f / 32 * 34,
7237 ne: vec![in_f as u64, out_f as u64],
7238 scale: 1.0,
7239 rp: false,
7240 #[cfg(memra_cutlass)]
7241 cutlass: None,
7242 fp8: None, blk: None, f16: None, rp4: None,
7243 };
7244 Ok(Some(self.matmul(&tmp, x, m)?))
7246 }
7247
7248 pub fn matmul_pre_noscale(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7249 m: usize) -> Result<Option<(CudaSlice<f32>, f32)>, Box<dyn std::error::Error>> {
7250 use crate::model::GpuTensor;
7251 if m == 1 {
7255 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(Some((y, 1.0))); }
7256 }
7257 if m != 1 || !self.uses_q8_1_fast(w) { return Ok(None); }
7259 let in_f = w.in_features();
7260 let out_f = w.out_features();
7261 let (bytes, qtype, row_bytes, scale, rp) = match w {
7262 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
7263 _ => return Ok(None),
7264 };
7265 if self.mmvq_supports(qtype) {
7267 let (mbytes, mrp) = match w {
7269 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
7270 _ => (bytes, rp),
7271 };
7272 let y = self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, 1.0, mrp)?;
7273 return Ok(Some((y, scale)));
7274 }
7275 let name = match qtype {
7277 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
7278 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
7279 QT_Q3_K => "qmatvec_q3_K_dp4a",
7280 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
7281 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
7282 _ => return Ok(None),
7283 };
7284 let f = self.func(name);
7285 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7286 let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
7287 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7288 let __s_b = self.gpu.stream();
7289 let mut b = __s_b.launch_builder(&f);
7290 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7291 unsafe { b.launch(cfg)?; }
7292 Ok(Some((y, scale)))
7293 }
7294
7295 pub fn mmvq_supports(&self, qtype: i32) -> bool {
7298 if qtype == QT_F8_E4M3 { return true; }
7303 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return false; }
7304 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_Q4_0)
7305 }
7306
7307 pub fn qmatvec_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7312 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7313 rp: bool)
7314 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7315 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)?;
7317 Ok(y)
7318 }
7319
7320 #[allow(clippy::too_many_arguments)]
7322 pub fn qmatvec_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7323 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7324 rp: bool, y: &mut CudaSlice<f32>)
7325 -> Result<(), Box<dyn std::error::Error>> {
7326 debug_assert!(y.len() >= m * out_f);
7327 const ROWS_PER_BLOCK: u32 = 4; if qtype == QT_Q8_0 && rp && m == 1 && out_f >= 64
7333 && (out_f as u32).div_ceil(ROWS_PER_BLOCK) < 4 * self.sm_count() as u32
7334 && {
7335 static G2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7336 *G2.get_or_init(|| std::env::var("MEMRA_Q80_G2").as_deref() != Ok("0"))
7337 }
7338 {
7339 let f = self.func("qmatvec_q8_0_mmvq_rp_g2");
7340 let cfg = LaunchConfig {
7341 grid_dim: ((out_f as u32).div_ceil(2), 1, 1),
7342 block_dim: (32, 2, 1),
7343 shared_mem_bytes: 0,
7344 };
7345 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, row_bytes as i64);
7346 let __s_b = self.gpu.stream();
7347 let mut b = __s_b.launch_builder(&f);
7348 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7349 unsafe { b.launch(cfg)?; }
7350 if scale != 1.0 { self.scale_inplace(y, scale, out_f)?; }
7351 return Ok(());
7352 }
7353 let mut mr: u32 = if m == 1 && (qtype == QT_NVFP4 || qtype == QT_Q5_K) { 2 } else { 1 };
7362 if m == 1 && qtype == QT_Q4_0 {
7367 static Q40MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7368 mr = *Q40MR.get_or_init(|| std::env::var("MEMRA_Q40_MR").ok()
7371 .and_then(|v| v.parse().ok()).unwrap_or(1));
7372 }
7373 let q5_mode = std::env::var("MEMRA_Q5K_ISSUE").ok();
7384 let q5_force = q5_mode.as_deref() == Some("2");
7385 let q5_il = qtype == QT_Q5_K && m == 1
7388 && (q5_force || q5_mode.as_deref().map(|v| v != "0").unwrap_or(true));
7389 if q5_il && !q5_force && out_f > 65536 { mr = 1; }
7390 if qtype == QT_Q4_0 && rp && mr != 1 { mr = 2; }
7393 if qtype == QT_Q8_0 && rp {
7397 static Q80MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7398 mr = *Q80MR.get_or_init(|| std::env::var("MEMRA_Q80_MR").ok()
7399 .and_then(|v| v.parse().ok()).unwrap_or(1));
7400 }
7401 let name = match (qtype, mr, rp) {
7402 (QT_NVFP4, 2, false) => "qmatvec_nvfp4_mmvq_mr2",
7403 (QT_NVFP4, 2, true) => "qmatvec_nvfp4_mmvq_mr2_rp",
7404 (QT_NVFP4, _, true) => "qmatvec_nvfp4_mmvq_rp",
7405 (QT_Q4_0, 1, true) => "qmatvec_q4_0_mmvq_rp",
7406 (QT_Q4_0, _, true) => "qmatvec_q4_0_mmvq_mr2_rp",
7407 (QT_Q5_K, 2, _) => if q5_il { "qmatvec_q5_K_mmvq_mr2_il" } else { "qmatvec_q5_K_mmvq_mr2" },
7408 (QT_Q8_0, 2, true) => "qmatvec_q8_0_mmvq_mr2_rp",
7409 (QT_Q8_0, _, true) if in_f % 1024 == 0 && {
7414 static CA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7415 *CA.get_or_init(|| std::env::var("MEMRA_Q80_CA").as_deref() == Ok("1"))
7416 } => "qmatvec_q8_0_mmvq_rpca",
7417 (QT_Q8_0, _, true) => "qmatvec_q8_0_mmvq_rp",
7418 (QT_Q8_0, _, _) => "qmatvec_q8_0_mmvq",
7419 (QT_Q4_K, _, true) => "qmatvec_q4_K_mmvq_rp",
7423 (QT_Q6_K, _, true) => "qmatvec_q6_K_mmvq_rp",
7424 (QT_Q4_K, _, _) => "qmatvec_q4_K_mmvq",
7425 (QT_Q4_0, 2, false) => "qmatvec_q4_0_mmvq_mr2",
7426 (QT_Q4_0, _, false) => "qmatvec_q4_0_mmvq",
7427 (QT_Q5_K, _, _) => if q5_il { "qmatvec_q5_K_mmvq_il" } else { "qmatvec_q5_K_mmvq" },
7428 (QT_Q6_K, _, _) => "qmatvec_q6_K_mmvq",
7429 (QT_NVFP4, _, false) => "qmatvec_nvfp4_mmvq",
7430 (QT_F8_E4M3, _, _) => "qmatvec_e4m3_mmvq",
7431 _ => panic!("qmatvec_mmvq: qtype {qtype} has no MMVQ kernel"),
7432 };
7433 let f = self.func(name);
7434 let rows_per_block = ROWS_PER_BLOCK * mr;
7436 let cfg = LaunchConfig {
7437 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, m as u32, 1),
7438 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7441 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7442 let __s_b = self.gpu.stream();
7443 let mut b = __s_b.launch_builder(&f);
7444 if qtype == QT_NVFP4 || qtype == QT_F8_E4M3 {
7449 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&scale);
7450 unsafe { b.launch(cfg)?; }
7451 } else if Self::pdl_on() && Self::pdl_mmvq_on()
7452 && matches!(name, "qmatvec_q4_0_mmvq_rp" | "qmatvec_q6_K_mmvq"
7453 | "qmatvec_q6_K_mmvq_rp") {
7454 {
7458 use cudarc::driver::{DevicePtr, DevicePtrMut};
7459 let s = &self.gpu.stream();
7460 let (pw, _g0) = bytes.device_ptr(s); let (paq, _g1) = aq.device_ptr(s);
7461 let (pad, _g2) = ad.device_ptr(s); let (py, _g3) = y.device_ptr_mut(s);
7462 let mut ps = [
7463 &pw as *const _ as *mut std::ffi::c_void, &paq as *const _ as *mut _,
7464 &pad as *const _ as *mut _, &py as *const _ as *mut _,
7465 &inf as *const _ as *mut _, &outf as *const _ as *mut _,
7466 &mi as *const _ as *mut _, &rb as *const _ as *mut _,
7467 ];
7468 unsafe { self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?; }
7469 }
7470 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7471 } else {
7472 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7473 unsafe { b.launch(cfg)?; }
7474 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7475 }
7476 Ok(())
7477 }
7478
7479 pub fn qmatvec_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
7483 out_f: usize, qtype: i32, row_bytes: usize, rp: bool)
7484 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7485 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7486 self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, 1.0, rp)
7487 }
7488
7489 pub fn batched_supports(&self, qtype: i32) -> bool {
7493 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_F8_E4M3 | QT_Q4_0)
7494 }
7495
7496 pub fn iq_fast_enabled() -> bool {
7504 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7505 *ON.get_or_init(|| std::env::var("MEMRA_IQ_FAST").map(|v| v != "0").unwrap_or(true))
7506 }
7507
7508 pub fn b8_enabled() -> bool {
7511 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7512 *ON.get_or_init(|| std::env::var("MEMRA_B8").map(|v| v != "0").unwrap_or(true))
7513 }
7514
7515 pub fn batched_mcols(m: usize) -> usize {
7517 if m == 2 { 2 } else if m <= 4 { 4 } else if m <= 8 { 8 } else { 16 }
7518 }
7519
7520 fn batched_kernel_name(qtype: i32, mcols: usize) -> Option<&'static str> {
7525 Some(match (qtype, mcols) {
7526 (QT_Q8_0, 2) => "qmatvec_q8_0_mmvq_b2", (QT_Q8_0, 4) => "qmatvec_q8_0_mmvq_b4",
7527 (QT_Q8_0, 8) => "qmatvec_q8_0_mmvq_b8",
7528 (QT_Q8_0, 16) => "qmatvec_q8_0_mmvq_b16",
7534 (QT_Q4_K, 2) => "qmatvec_q4_K_mmvq_b2", (QT_Q4_K, 4) => "qmatvec_q4_K_mmvq_b4",
7535 (QT_Q4_K, 8) => "qmatvec_q4_K_mmvq_b8",
7536 (QT_Q4_K, 16) => "qmatvec_q4_K_mmvq_b16",
7539 (QT_Q5_K, 2) => "qmatvec_q5_K_mmvq_b2", (QT_Q5_K, 4) => "qmatvec_q5_K_mmvq_b4",
7540 (QT_Q5_K, 8) => "qmatvec_q5_K_mmvq_b8",
7541 (QT_Q5_K, 16) => "qmatvec_q5_K_mmvq_b16",
7544 (QT_Q6_K, 2) => "qmatvec_q6_K_mmvq_b2", (QT_Q6_K, 4) => "qmatvec_q6_K_mmvq_b4",
7545 (QT_Q6_K, 8) => "qmatvec_q6_K_mmvq_b8", (QT_Q6_K, 16) => "qmatvec_q6_K_mmvq_b16",
7546 (QT_NVFP4, 2) => "qmatvec_nvfp4_mmvq_b2", (QT_NVFP4, 4) => "qmatvec_nvfp4_mmvq_b4",
7547 (QT_NVFP4, 8) => "qmatvec_nvfp4_mmvq_b8",
7548 (QT_NVFP4, 16) => "qmatvec_nvfp4_mmvq_b16",
7552 (QT_F8_E4M3, 2) => "qmatvec_e4m3_mmvq_b2", (QT_F8_E4M3, 4) => "qmatvec_e4m3_mmvq_b4",
7553 (QT_F8_E4M3, 8) => "qmatvec_e4m3_mmvq_b8",
7554 (QT_F8_E4M3, 16) => "qmatvec_e4m3_mmvq_b16",
7558 (QT_Q4_0, 2) => "qmatvec_q4_0_mmvq_b2", (QT_Q4_0, 4) => "qmatvec_q4_0_mmvq_b4",
7559 (QT_Q4_0, 8) => "qmatvec_q4_0_mmvq_b8", (QT_Q4_0, 16) => "qmatvec_q4_0_mmvq_b16",
7560 _ => return None,
7561 })
7562 }
7563
7564 pub fn sm_count(&self) -> i32 {
7599 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7600 *SMS.get_or_init(|| {
7601 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7602 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7603 })
7604 }
7605
7606 pub fn batched_variant(&self, _m: usize, in_f: usize, out_f: usize, qtype: i32,
7607 row_bytes: usize, mcols: usize, rp: bool) -> &'static str {
7608 if qtype == QT_Q8_0 {
7613 return if rp { "rp" } else { "base" };
7614 }
7615 static BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7616 let bv = *BV.get_or_init(|| match std::env::var("MEMRA_MMVQ_BV").as_deref() {
7617 Ok("base") => "base", Ok("pf") => "pf", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7618 Ok("pfr2") => "pfr2", Ok("ca") => "ca", Ok("car2") => "car2",
7619 Ok("rp") => "rp", Ok("rpr2") => "rpr2", Ok("rpr2w8") => "rpr2w8",
7622 Ok("rpca") => "rpca", Ok("rpcar2") => "rpcar2",
7625 Ok("rpsc") => "rpsc", Ok("rpms") => "rpms", Ok("rpmsc") => "rpmsc",
7632 Ok("rpks") => "rpks", Ok("rpksc") => "rpksc",
7633 _ => "auto",
7634 });
7635 let ca_ok = qtype == QT_NVFP4 && (row_bytes % 16 == 0) && (in_f % 1024 == 0);
7639 static KS_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7644 let ks_on = *KS_ON.get_or_init(|| std::env::var("MEMRA_KS").as_deref() != Ok("0"));
7645 let sc_ok = ks_on && qtype == QT_NVFP4 && (in_f % 256 == 0) && (in_f / 64 <= 272);
7646 let ks_ok = ks_on && qtype == QT_NVFP4 && (in_f % 512 == 0) && (in_f / 64 <= 272);
7647 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7648 let sms = *SMS.get_or_init(|| {
7649 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7650 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7651 });
7652 let kq_r2 = matches!(qtype, QT_Q4_K | QT_Q5_K | QT_Q6_K);
7672 static KQBV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7675 let kq_bv = *KQBV.get_or_init(|| match std::env::var("MEMRA_KQ_BV").as_deref() {
7676 Ok("base") => "base", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7677 _ => "auto",
7678 });
7679 let variant: &'static str = if qtype == QT_Q4_0 {
7680 static Q40BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7684 let q40 = *Q40BV.get_or_init(|| match std::env::var("MEMRA_Q40_BV").as_deref() {
7685 Ok("base") => "base", Ok("r2") => "r2", Ok("ms") => "ms", Ok("sm") => "sm",
7691 Ok("la") => "la", _ => "auto",
7692 });
7693 let v = if q40 != "auto" { q40 }
7694 else if (out_f as u32).div_ceil(8) >= 4 * sms as u32 { "r2" } else { "base" };
7695 if rp { match v { "ms" => "r2ms_rp", "sm" => "r2sm_rp", "la" => "r2la_rp",
7700 "r2" => "r2_rp", _ => "rp" } }
7701 else if matches!(v, "ms" | "sm" | "la") { "r2" } else { v }
7702 } else if qtype != QT_NVFP4 && !kq_r2 {
7703 "base"
7704 } else if kq_r2 && rp {
7705 "rp"
7709 } else if kq_r2 {
7710 if kq_bv != "auto" {
7713 if kq_bv == "r2w8" && mcols != 4 { "r2" } else { kq_bv }
7714 } else if bv != "auto" {
7715 match bv {
7716 "r2" | "pfr2" | "rpr2" | "car2" => "r2",
7717 "r2w8" | "rpr2w8" => if mcols != 4 { "r2" } else { "r2w8" },
7718 _ => "base", }
7720 } else {
7721 let blocks = (out_f + 7) / 8;
7722 let waves = blocks as f64 / (7 * sms as usize) as f64;
7723 let filled = blocks >= 4 * sms as usize;
7724 let use_r2 = if qtype == QT_Q4_K { filled } else { waves >= 2.0 };
7725 if use_r2 { "r2" } else { "base" }
7726 }
7727 } else if bv != "auto" {
7728 let v = if bv == "r2w8" && mcols == 2 { "r2" }
7733 else if bv == "ca" && (!ca_ok || mcols == 8) { "pf" }
7734 else if bv == "car2" && (!ca_ok || mcols == 8) { "r2" }
7735 else if bv == "pfr2" && mcols == 8 { "r2" }
7736 else if (bv == "rpr2w8" || bv == "rpr2") && mcols == 2 { "rpr2" }
7737 else if (bv == "rpca" || bv == "rpcar2") && (!ca_ok || mcols == 8) {
7739 if mcols == 8 { "rpr2w8" } else { "rpr2" }
7740 }
7741 else if bv == "rpcar2" && mcols == 2 { "rpca" }
7742 else if (bv == "rpsc" || bv == "rpmsc") && !sc_ok { "rpr2" }
7745 else if (bv == "rpks" || bv == "rpksc") && !ks_ok { "rpr2" }
7746 else { bv };
7747 if rp {
7748 match v {
7749 "base" | "pf" | "ca" | "rp" => "rp",
7750 "r2" | "pfr2" | "car2" | "rpr2" => "rpr2",
7751 "r2w8" | "rpr2w8" => if mcols == 2 { "rpr2" } else { "rpr2w8" },
7752 other => other, }
7754 } else { v }
7755 } else if mcols == 8 {
7756 if rp { if sc_ok { "rpsc" } else { "rpr2w8" } } else { "r2w8" }
7767 } else if mcols >= 4 {
7768 let blocks = (out_f + 7) / 8;
7772 let r7 = 7 * sms as usize;
7773 let r8 = 8 * sms as usize;
7774 let waves = blocks as f64 / r7 as f64;
7775 let filled = blocks >= 4 * sms as usize;
7776 if filled && blocks.div_ceil(r8) < blocks.div_ceil(r7) {
7780 if rp { "rpr2w8" } else { "r2w8" }
7784 } else if waves >= 2.0 || (waves <= 1.0 && filled) {
7785 if rp { "rpr2" } else { "r2" }
7788 } else {
7789 if rp { "rp" } else { "pf" }
7793 }
7794 } else if in_f >= 6144 {
7795 if rp { "rpr2" } else { "r2" }
7799 }
7800 else if rp {
7801 let waves = ((out_f + 7) / 8) as f64 / (7 * sms as usize) as f64;
7806 if sc_ok && waves >= 0.9 && waves <= 1.1 { "rpsc" } else { "rp" }
7807 } else { "base" };
7808 variant
7809 }
7810
7811 pub fn qmatvec_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7812 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize,
7813 mcols: usize, scale: f32, rp: bool)
7814 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7815 const ROWS_PER_BLOCK: u32 = 4;
7816 let forced: Option<&'static str> = {
7821 static V: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
7822 V.get_or_init(|| std::env::var("MEMRA_BVAR").ok())
7823 .as_deref()
7824 .map(|s| Box::leak(s.to_string().into_boxed_str()) as &'static str)
7825 };
7826 let variant = match forced {
7827 Some(v) if !rp || v.contains("rp") => v,
7828 _ => self.batched_variant(m, in_f, out_f, qtype, row_bytes, mcols, rp),
7829 };
7830 let base_name = Self::batched_kernel_name(qtype, mcols)
7831 .ok_or_else(|| format!("qmatvec_mmvq_batched: no kernel for qtype {qtype} mcols {mcols}"))?;
7832 let variant = if mcols == 16 { if rp { "rp" } else { "base" } } else { variant };
7836 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7843 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
7844 if b567 && qtype == QT_NVFP4 && rp && mcols == 8 && (5..=7).contains(&m)
7845 && matches!(variant, "rpsc" | "rpr2w8") {
7846 let f = self.func(&format!("qmatvec_nvfp4_mmvq_b{m}_{variant}"));
7847 let rows_per_block = ROWS_PER_BLOCK * 2; let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7849 let cfg = LaunchConfig {
7850 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
7851 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0 };
7852 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7853 let __s_b = self.gpu.stream();
7854 let mut b = __s_b.launch_builder(&f);
7855 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7856 unsafe { b.launch(cfg)?; }
7857 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
7858 return Ok(y);
7859 }
7860 let (name, rows_per_block): (std::borrow::Cow<'static, str>, u32) = match variant {
7861 "base" => (base_name.into(), ROWS_PER_BLOCK),
7862 "pf" => (format!("{base_name}_pf").into(), ROWS_PER_BLOCK),
7863 "ca" => (format!("{base_name}_ca").into(), ROWS_PER_BLOCK),
7864 "rp" => (format!("{base_name}_rp").into(), ROWS_PER_BLOCK),
7865 "rpca" => (format!("{base_name}_rpca").into(), ROWS_PER_BLOCK), "rpks" => (format!("{base_name}_rpks").into(), ROWS_PER_BLOCK),
7869 "rpksc" => (format!("{base_name}_rpksc").into(), ROWS_PER_BLOCK),
7870 "rpms" => (format!("{base_name}_rpms").into(), ROWS_PER_BLOCK),
7871 "rpmsc" => (format!("{base_name}_rpmsc").into(), ROWS_PER_BLOCK),
7872 "r2ms_rp" => (format!("{base_name}_r2ms_rp").into(), ROWS_PER_BLOCK),
7873 "r2sm_rp" => (format!("{base_name}_r2sm_rp").into(), ROWS_PER_BLOCK * 2),
7874 "r2la_rp" => (format!("{base_name}_r2la_rp").into(), ROWS_PER_BLOCK * 2),
7875 v => (format!("{base_name}_{v}").into(), ROWS_PER_BLOCK * 2), };
7877 debug_assert!(!rp || name.contains("_rp"), "rp weight dispatched to a GGUF-layout kernel");
7878 let f = self.func(&name);
7879 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7880 let smem = if name.contains("_r2sm_rp") { (mcols * 32 * 9 * 4 + mcols * 32 * 4) as u32 }
7882 else { 0 };
7883 let cfg = LaunchConfig {
7884 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
7885 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: smem };
7886 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7887 let __s_b = self.gpu.stream();
7888 let mut b = __s_b.launch_builder(&f);
7889 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7890 unsafe { b.launch(cfg)?; }
7891 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
7892 Ok(y)
7893 }
7894
7895 pub fn qmatvec_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
7899 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, mcols: usize,
7900 rp: bool)
7901 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7902 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7903 self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, 1.0, rp)
7904 }
7905
7906 pub fn qmatvec_nvfp4_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
7908 in_f: usize, out_f: usize, row_bytes: usize, mcols: usize,
7909 rp: bool)
7910 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7911 self.qmatvec_batched_raw(bytes, x, m, in_f, out_f, QT_NVFP4, row_bytes, mcols, rp)
7912 }
7913
7914 fn try_fp4_gemm(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize,
7918 in_f: usize, out_f: usize)
7919 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7920 use crate::model::GpuTensor;
7921 if cfg!(memra_portable_cuda) { return Ok(None); }
7922 if std::env::var("MEMRA_FP4").is_err() { return Ok(None); }
7923 #[cfg(memra_cutlass)]
7932 if m >= 128 && std::env::var("MEMRA_FP4_CUTLASS").is_ok() {
7933 if let GpuTensor::Quant { bytes, qtype, scale, row_bytes, cutlass, .. } = w {
7934 if *qtype == QT_NVFP4 && in_f % 64 == 0 {
7935 if let Some(cw) = cutlass {
7936 let y = self.cutlass_fp4_gemm(&cw.b_packed, &cw.sfb_swizzled, x, *scale,
7938 m, out_f, in_f)?;
7939 return Ok(Some(y));
7940 } else if std::env::var("MEMRA_FP4_CUTLASS_OTF").is_ok() {
7941 let (b_packed, sfb_sw) = self.build_cutlass_weight(bytes, out_f, in_f, *row_bytes)?;
7946 let y = self.cutlass_fp4_gemm(&b_packed, &sfb_sw, x, *scale, m, out_f, in_f)?;
7947 return Ok(Some(y));
7948 }
7949 }
7950 }
7951 }
7952 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } = w {
7953 if *qtype == QT_NVFP4 && in_f % 64 == 0 && !*rp {
7956 let y = self.qmatvec_gemm_nvfp4_fp4(bytes, x, m, in_f, out_f, *row_bytes, *scale)?;
7957 return Ok(Some(y));
7958 }
7959 }
7960 Ok(None)
7961 }
7962
7963 pub fn rms_norm_f16out(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>,
7967 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
7968 ncols: usize, nrows: usize, eps: f32)
7969 -> Result<(), Box<dyn std::error::Error>> {
7970 let f = self.func("rms_norm_f16out_f32");
7971 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
7972 let (nc, e) = (ncols as i32, eps);
7973 let __s_b = self.gpu.stream();
7974 let mut b = __s_b.launch_builder(&f);
7975 b.arg(x).arg(w).arg(dst).arg(dst16).arg(&nc).arg(&e);
7976 unsafe { b.launch(cfg)?; }
7977 Ok(())
7978 }
7979
7980 #[allow(clippy::too_many_arguments)]
7983 pub fn add_rms_norm_f16out(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
7984 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
7985 dst16: &mut CudaSlice<u8>, ncols: usize, nrows: usize, eps: f32)
7986 -> Result<(), Box<dyn std::error::Error>> {
7987 let f = self.func("add_rms_norm_f16out_f32");
7988 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
7989 let (nc, e) = (ncols as i32, eps);
7990 let __s_lb = self.gpu.stream();
7991 let mut lb = __s_lb.launch_builder(&f);
7992 lb.arg(a).arg(b).arg(w).arg(res).arg(dst).arg(dst16).arg(&nc).arg(&e);
7993 unsafe { lb.launch(cfg)?; }
7994 Ok(())
7995 }
7996
7997 pub fn matmul_group_xh(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>,
8000 xh: &CudaSlice<u8>, m: usize)
8001 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8002 let mut out = Vec::with_capacity(ws.len());
8003 let in_f = ws[0].in_features();
8004 for w in ws {
8005 if w.in_features() == in_f && m >= 16 && !self.verify_exact_on() {
8006 if let Some(y) = self.try_f16_gemm_pre(w, xh, m)? {
8007 out.push(y);
8008 continue;
8009 }
8010 }
8011 out.push(self.matmul(w, x, m)?);
8012 }
8013 Ok(out)
8014 }
8015
8016 pub fn gdn_pad_mask(&self, beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
8019 len_d: &CudaSlice<i32>, h: usize, t: usize)
8020 -> Result<(), Box<dyn std::error::Error>> {
8021 let f = self.func("gdn_pad_mask_f32");
8022 let cfg = LaunchConfig::for_num_elems((t * h) as u32);
8023 let (hi, ti) = (h as i32, t as i32);
8024 let __s_b = self.gpu.stream();
8025 let mut b = __s_b.launch_builder(&f);
8026 b.arg(beta).arg(g_log).arg(len_d).arg(&hi).arg(&ti);
8027 unsafe { b.launch(cfg)?; }
8028 Ok(())
8029 }
8030
8031 pub fn row_gather_dev(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
8034 len_d: &CudaSlice<i32>, ncols: usize)
8035 -> Result<(), Box<dyn std::error::Error>> {
8036 let f = self.func("row_gather_dev_f32");
8037 let cfg = LaunchConfig::for_num_elems(ncols as u32);
8038 let nc = ncols as i32;
8039 let __s_b = self.gpu.stream();
8040 let mut b = __s_b.launch_builder(&f);
8041 b.arg(src).arg(dst).arg(len_d).arg(&nc);
8042 unsafe { b.launch(cfg)?; }
8043 Ok(())
8044 }
8045
8046 pub fn matmul_group(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>, m: usize)
8053 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8054 use crate::model::GpuTensor;
8055 let mut out = Vec::with_capacity(ws.len());
8056 let any_mirror = ws.iter().any(|w| matches!(w, GpuTensor::Quant { f16: Some(_), .. }));
8057 if m >= 16 && any_mirror && !self.verify_exact_on() {
8058 let in_f = ws[0].in_features();
8059 let xh = self.f16_act(x, m * in_f, in_f)?;
8060 for w in ws {
8061 if w.in_features() == in_f {
8062 if let Some(y) = self.try_f16_gemm_pre(w, &xh, m)? {
8063 out.push(y);
8064 continue;
8065 }
8066 }
8067 out.push(self.matmul(w, x, m)?);
8068 }
8069 return Ok(out);
8070 }
8071 for w in ws {
8072 out.push(self.matmul(w, x, m)?);
8073 }
8074 Ok(out)
8075 }
8076
8077 pub fn matmul_group_multi(&self, ws: &[&crate::model::GpuTensor],
8084 xs: &[&CudaSlice<f32>], ms: &[usize])
8085 -> Result<Vec<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
8086 assert_eq!(xs.len(), ms.len());
8087 let in_f = ws[0].in_features();
8088 let total: usize = ms.iter().sum();
8089 let mut xcat = self.uninit(total * in_f)?;
8090 let mut off = 0usize;
8091 for (x, &m) in xs.iter().zip(ms) {
8092 self.copy_into(&mut xcat, off * in_f, x, m * in_f)?;
8093 off += m;
8094 }
8095 let ys = self.matmul_group(ws, &xcat, total)?;
8096 let mut out: Vec<Vec<CudaSlice<f32>>> = (0..xs.len()).map(|_| Vec::new()).collect();
8097 for (w, y) in ws.iter().zip(ys) {
8098 let out_f = w.out_features();
8099 let mut off = 0usize;
8100 for (s, &m) in ms.iter().enumerate() {
8101 let mut ys_s = self.uninit(m * out_f)?;
8102 let src = y.slice(off * out_f..(off + m) * out_f);
8103 self.gpu.stream().memcpy_dtod(&src, &mut ys_s)?;
8104 out[s].push(ys_s);
8105 off += m;
8106 }
8107 }
8108 Ok(out)
8109 }
8110
8111 pub fn gemm_supports(&self, w: &crate::model::GpuTensor) -> bool {
8121 use crate::model::GpuTensor;
8122 if !legacy_quant_gemm_allowed(
8123 cfg!(memra_portable_cuda),
8124 cfg!(memra_hopper_mma),
8125 std::env::var_os("MEMRA_NO_GEMM").is_some(),
8126 ) {
8127 return false;
8128 }
8129 match w {
8130 GpuTensor::Quant { qtype, .. } =>
8131 matches!(*qtype, QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q4_0)
8132 || (*qtype == QT_NVFP4 && w.in_features() % 64 == 0),
8133 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
8134 }
8135 }
8136
8137 pub fn qmatvec_gemm(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
8144 m: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8145 use crate::model::GpuTensor;
8146 let in_f = w.in_features();
8147 let out_f = w.out_features();
8148 let (bytes, qtype, row_bytes, scale, rp) = match w {
8149 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
8150 _ => unreachable!("gemm_supports guaranteed Quant"),
8151 };
8152 if cfg!(memra_hopper_mma) && qtype == QT_Q8_0 && out_f % 64 == 0 && wgmma_gemm_enabled() {
8158 if let GpuTensor::Quant { rp4: Some(m4), .. } = w {
8159 let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, aq, ad, m, in_f, out_f)?;
8160 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8161 return Ok(y);
8162 }
8163 }
8164 let name = match qtype {
8165 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8166 QT_Q4_0 => if rp { "qmatvec_gemm_q4_0_rp" } else { "qmatvec_gemm_q4_0" },
8167 QT_Q5_K => "qmatvec_gemm_q5_K",
8168 QT_Q6_K => "qmatvec_gemm_q6_K",
8169 QT_NVFP4 => if rp { "qmatvec_gemm_nvfp4_rp" } else { "qmatvec_gemm_nvfp4" },
8170 _ => unreachable!(),
8171 };
8172 let f = self.func(name);
8173 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);
8178 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8180 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8181 let warps: u32 = if is_k1 { k1_tile.2 } else {
8182 match qtype { QT_NVFP4 => 8, _ => 4 }
8183 };
8184 let cfg = LaunchConfig {
8185 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8186 block_dim: (32, warps, 1),
8187 shared_mem_bytes: 0,
8188 };
8189 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8190 let __s_b = self.gpu.stream();
8191 let mut b = __s_b.launch_builder(&f);
8192 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8193 unsafe { b.launch(cfg)?; }
8194 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8195 Ok(y)
8196 }
8197
8198 pub fn qmatvec_gemm_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
8203 out_f: usize, qtype: i32, row_bytes: usize)
8204 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8205 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
8206 let name = match qtype {
8207 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8208 QT_Q4_0 => "qmatvec_gemm_q4_0",
8209 QT_Q5_K => "qmatvec_gemm_q5_K",
8210 QT_Q6_K => "qmatvec_gemm_q6_K", QT_NVFP4 => "qmatvec_gemm_nvfp4",
8211 QT_NVFP4_RP => "qmatvec_gemm_nvfp4_rp",
8212 _ => panic!("qmatvec_gemm_raw: qtype {qtype} has no GEMM kernel"),
8213 };
8214 let f = self.func(name);
8215 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);
8219 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8221 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8222 let warps: u32 = if is_k1 { k1_tile.2 } else {
8223 match qtype { QT_NVFP4 | QT_NVFP4_RP => 8, _ => 4 }
8224 };
8225 let cfg = LaunchConfig {
8226 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8227 block_dim: (32, warps, 1), shared_mem_bytes: 0,
8228 };
8229 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8230 let __s_b = self.gpu.stream();
8231 let mut b = __s_b.launch_builder(&f);
8232 b.arg(bytes).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8233 unsafe { b.launch(cfg)?; }
8234 Ok(y)
8235 }
8236
8237 pub fn qmatvec_gemm_q8_0_wgmma_raw(&self, rp4: &CudaSlice<u8>, aq: &CudaSlice<i8>,
8244 ad: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize)
8245 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8246 assert!(out_f % 64 == 0 && in_f % 32 == 0, "wgmma GEMM needs out_f%64==0, in_f%32==0");
8247 let f = self.func("qmatvec_gemm_q8_0_wgmma");
8248 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
8250 grid_dim: ((out_f / 64) as u32, (m as u32).div_ceil(64), 1),
8251 block_dim: (128, 1, 1), shared_mem_bytes: 0,
8252 };
8253 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
8254 let __s_b = self.gpu.stream();
8255 let mut b = __s_b.launch_builder(&f);
8256 b.arg(rp4).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi);
8257 unsafe { b.launch(cfg)?; }
8258 Ok(y)
8259 }
8260
8261 pub fn scale_inplace(&self, y: &mut CudaSlice<f32>, s: f32, n: usize)
8263 -> Result<(), Box<dyn std::error::Error>> {
8264 let f = self.func("scale_f32");
8265 let cfg = LaunchConfig::for_num_elems(n as u32);
8266 let (sf, ni) = (s, n as i32);
8267 let __s_b = self.gpu.stream();
8268 let mut b = __s_b.launch_builder(&f);
8269 b.arg(y).arg(&sf).arg(&ni);
8270 unsafe { b.launch(cfg)?; }
8271 Ok(())
8272 }
8273
8274 pub fn bf16_to_f32(&self, data: &cudarc::driver::CudaView<'_, u8>, n: usize)
8279 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8280 let mut out = self.alloc_uninit::<f32>(n)?;
8281 let f = self.func("bf16_to_f32");
8282 let cfg = LaunchConfig::for_num_elems(n as u32);
8283 let ni = n as i32;
8284 let __s_b = self.gpu.stream();
8285 let mut b = __s_b.launch_builder(&f);
8286 b.arg(data).arg(&mut out).arg(&ni);
8287 unsafe { b.launch(cfg)?; }
8288 Ok(out)
8289 }
8290
8291 fn linear_bf16_chunked(&self, x: &CudaSlice<f32>, data: &CudaSlice<u8>, m: usize,
8298 in_f: usize, out_f: usize, exact: bool)
8299 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8300 const CHUNK_BYTES: usize = 256 << 20;
8301 let chunk_rows = (CHUNK_BYTES / (in_f * 4)).max(1).min(out_f);
8302 if chunk_rows >= out_f {
8303 let wf32 = self.bf16_to_f32(&data.slice(0..in_f * out_f * 2), in_f * out_f)?;
8304 return if exact { self.linear_decode_exact(x, &wf32, m, in_f, out_f) }
8305 else { self.linear(x, &wf32, m, in_f, out_f) };
8306 }
8307 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8308 let mut r0 = 0usize;
8309 while r0 < out_f {
8310 let rows = chunk_rows.min(out_f - r0);
8311 let wslice = data.slice(r0 * in_f * 2..(r0 + rows) * in_f * 2);
8312 let wf32 = self.bf16_to_f32(&wslice, in_f * rows)?;
8313 let yc = if exact { self.linear_decode_exact(x, &wf32, m, in_f, rows)? }
8314 else { self.linear(x, &wf32, m, in_f, rows)? };
8315 for mi in 0..m {
8317 let src = yc.slice(mi * rows..(mi + 1) * rows);
8318 let mut dst = y.slice_mut(mi * out_f + r0..mi * out_f + r0 + rows);
8319 self.gpu.stream().memcpy_dtod(&src, &mut dst)?;
8320 }
8321 r0 += rows;
8322 }
8323 Ok(y)
8324 }
8325
8326 pub fn linear_decode_exact(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize,
8333 in_f: usize, out_f: usize)
8334 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8335 if m_tokens == 1 { return self.linear(x, w, 1, in_f, out_f); }
8336 let xv = self.view(x, m_tokens * in_f);
8337 let mut y = self.alloc_uninit::<f32>(m_tokens * out_f)?;
8338 for t in 0..m_tokens {
8339 let row = xv.slice(t * in_f..(t + 1) * in_f);
8340 let mut xr = self.alloc_uninit::<f32>(in_f)?;
8341 self.copy_view_into(&mut xr, 0, &row, in_f)?;
8342 let yr = self.linear(&xr, w, 1, in_f, out_f)?;
8343 self.copy_into(&mut y, t * out_f, &yr, out_f)?;
8344 }
8345 Ok(y)
8346 }
8347
8348 pub fn linear(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize, in_f: usize, out_f: usize)
8349 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8350 use cudarc::cublaslt::{Matmul, MatmulConfig};
8351 let mut c = self.alloc_uninit::<f32>(m_tokens * out_f)?; let cfg = MatmulConfig {
8353 transa: true, transb: false, transc: false,
8354 m: out_f as u64, n: m_tokens as u64, k: in_f as u64,
8355 alpha: 1.0, lda: in_f as i64, ldb: in_f as i64, beta: 0.0, ldc: out_f as i64,
8356 stride_a: None, stride_b: None, stride_c: None, stride_bias: None, batch_size: None,
8357 };
8358 unsafe { self.gpu.blas.matmul(cfg, w, x, &mut c, None, None)?; }
8359 Ok(c)
8360 }
8361
8362 pub fn sdpa_naive(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8364 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8365 t: usize, t_kv: usize, scale: f32, causal: bool)
8366 -> Result<(), Box<dyn std::error::Error>> {
8367 let f = self.func("sdpa_naive_f32");
8368 let cfg = LaunchConfig {
8369 grid_dim: (n_head as u32, t as u32, 1),
8370 block_dim: (128, 1, 1),
8371 shared_mem_bytes: (t_kv * 4) as u32,
8372 };
8373 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);
8374 let __s_b = self.gpu.stream();
8375 let mut b = __s_b.launch_builder(&f);
8376 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8377 unsafe { b.launch(cfg)?; }
8378 Ok(())
8379 }
8380
8381 #[allow(clippy::too_many_arguments)]
8383 pub fn sdpa_naive_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8384 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8385 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8386 -> Result<(), Box<dyn std::error::Error>> {
8387 let f = self.func("sdpa_naive_w_f32");
8388 let cfg = LaunchConfig {
8389 grid_dim: (n_head as u32, t as u32, 1),
8390 block_dim: (128, 1, 1),
8391 shared_mem_bytes: (t_kv * 4) as u32,
8392 };
8393 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
8394 t as i32, t_kv as i32, causal as i32, window as i32);
8395 let __s_b = self.gpu.stream();
8396 let mut b = __s_b.launch_builder(&f);
8397 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8398 .arg(&scale).arg(&cz).arg(&wi);
8399 unsafe { b.launch(cfg)?; }
8400 Ok(())
8401 }
8402
8403 pub fn sdpa_naive_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<f32>,
8405 v: &cudarc::driver::CudaView<f32>, o: &mut CudaSlice<f32>,
8406 head_dim: usize, n_head: usize, n_head_kv: usize, t: usize, t_kv: usize,
8407 scale: f32, causal: bool) -> Result<(), Box<dyn std::error::Error>> {
8408 let f = self.func("sdpa_naive_f32");
8409 let cfg = LaunchConfig {
8410 grid_dim: (n_head as u32, t as u32, 1), block_dim: (128, 1, 1),
8411 shared_mem_bytes: (t_kv * 4) as u32,
8412 };
8413 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);
8414 let __s_b = self.gpu.stream();
8415 let mut b = __s_b.launch_builder(&f);
8416 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8417 unsafe { b.launch(cfg)?; }
8418 Ok(())
8419 }
8420
8421 #[allow(clippy::too_many_arguments)]
8429 pub fn fa_dequant_kv_view_f32(&self, k: &cudarc::driver::CudaView<u8>,
8430 v: &cudarc::driver::CudaView<u8>,
8431 kf: &mut CudaSlice<f32>, vf: &mut CudaSlice<f32>,
8432 kv_dim_k: usize, kv_dim_v: usize, t_kv: usize,
8433 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
8434 -> Result<(), Box<dyn std::error::Error>> {
8435 let f = if g { self.func_g("fa_dequant_kv_ws_f32") } else { self.func("fa_dequant_kv_ws_f32") };
8436 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
8437 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8438 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1),
8439 shared_mem_bytes: 0 };
8440 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
8441 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
8442 let __s_b = self.gpu.stream();
8443 let mut b = __s_b.launch_builder(&f);
8444 b.arg(k).arg(v).arg(&mut *kf).arg(&mut *vf).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
8445 unsafe { b.launch(cfg)?; }
8446 Ok(())
8447 }
8448
8449 #[allow(clippy::too_many_arguments)]
8450 pub fn sdpa_naive_quantized_view(
8451 &self,
8452 q: &CudaSlice<f32>,
8453 k: &cudarc::driver::CudaView<u8>,
8454 v: &cudarc::driver::CudaView<u8>,
8455 o: &mut CudaSlice<f32>,
8456 head_dim: usize,
8457 n_head: usize,
8458 n_head_kv: usize,
8459 t: usize,
8460 t_kv: usize,
8461 scale: f32,
8462 causal: bool,
8463 k_tok_bytes: usize,
8464 v_tok_bytes: usize,
8465 ) -> Result<(), Box<dyn std::error::Error>> {
8466 let kv_dim = n_head_kv * head_dim;
8467 let mut kf = self.uninit(t_kv * kv_dim)?;
8468 let mut vf = self.uninit(t_kv * kv_dim)?;
8469 let f = self.func("fa_dequant_kv_ws_f32");
8470 let total = (2 * t_kv * kv_dim) as u64;
8471 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8472 let cfg = LaunchConfig {
8473 grid_dim: (nblk.max(1), 1, 1),
8474 block_dim: (256, 1, 1),
8475 shared_mem_bytes: 0,
8476 };
8477 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8478 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8479 let __s_b = self.gpu.stream();
8480 let mut b = __s_b.launch_builder(&f);
8481 b.arg(k)
8482 .arg(v)
8483 .arg(&mut kf)
8484 .arg(&mut vf)
8485 .arg(&kv_dim_i)
8486 .arg(&kv_dim_i)
8487 .arg(&t_kv_i)
8488 .arg(&k_tok_bytes_i)
8489 .arg(&v_tok_bytes_i);
8490 unsafe { b.launch(cfg)? };
8491 self.sdpa_naive(
8492 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8493 )
8494 }
8495
8496 #[allow(clippy::too_many_arguments)]
8508 pub fn sdpa_naive_w_quantized_view(
8509 &self,
8510 q: &CudaSlice<f32>,
8511 k: &cudarc::driver::CudaView<u8>,
8512 v: &cudarc::driver::CudaView<u8>,
8513 o: &mut CudaSlice<f32>,
8514 head_dim: usize,
8515 n_head: usize,
8516 n_head_kv: usize,
8517 t: usize,
8518 t_kv: usize,
8519 scale: f32,
8520 causal: bool,
8521 window: usize,
8522 k_tok_bytes: usize,
8523 v_tok_bytes: usize,
8524 ) -> Result<(), Box<dyn std::error::Error>> {
8525 let kv_dim = n_head_kv * head_dim;
8526 let mut kf = self.uninit(t_kv * kv_dim)?;
8527 let mut vf = self.uninit(t_kv * kv_dim)?;
8528 let f = self.func("fa_dequant_kv_ws_f32");
8529 let total = (2 * t_kv * kv_dim) as u64;
8530 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8531 let cfg = LaunchConfig {
8532 grid_dim: (nblk.max(1), 1, 1),
8533 block_dim: (256, 1, 1),
8534 shared_mem_bytes: 0,
8535 };
8536 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8537 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8538 let __s_b = self.gpu.stream();
8539 let mut b = __s_b.launch_builder(&f);
8540 b.arg(k)
8541 .arg(v)
8542 .arg(&mut kf)
8543 .arg(&mut vf)
8544 .arg(&kv_dim_i)
8545 .arg(&kv_dim_i)
8546 .arg(&t_kv_i)
8547 .arg(&k_tok_bytes_i)
8548 .arg(&v_tok_bytes_i);
8549 unsafe { b.launch(cfg)? };
8550 self.sdpa_naive_w(
8551 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
8552 )
8553 }
8554
8555 pub fn fa_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8559 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8560 t: usize, t_kv: usize, scale: f32, causal: bool)
8561 -> Result<(), Box<dyn std::error::Error>> {
8562 if portable_mma_gated() {
8563 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
8564 t, t_kv, scale, causal);
8565 }
8566 let fa3_on = head_dim == 256 && causal && t == t_kv
8574 && match std::env::var("MEMRA_FA3").as_deref() {
8575 Ok("0") => false,
8576 Ok("1") => true,
8577 _ => cfg!(memra_hopper_mma),
8578 };
8579 if fa3_on {
8580 let n = t * n_head * head_dim;
8581 let nkv = t * n_head_kv * head_dim;
8582 let mut q16 = self.alloc_u8_uninit(n * 2)?;
8583 let mut k16 = self.alloc_u8_uninit(nkv * 2)?;
8584 let mut v16 = self.alloc_u8_uninit(nkv * 2)?;
8585 self.f32_to_bf16_into(q, &mut q16, n)?;
8586 self.f32_to_bf16_into(k, &mut k16, nkv)?;
8587 self.f32_to_bf16_into(v, &mut v16, nkv)?;
8588 let rc = {
8589 use cudarc::driver::{DevicePtr, DevicePtrMut};
8590 let stream = self.gpu.stream();
8591 let (qp, _g1) = q16.device_ptr(&stream);
8592 let (kp, _g2) = k16.device_ptr(&stream);
8593 let (vp, _g3) = v16.device_ptr(&stream);
8594 let (op, _g4) = o.device_ptr_mut(&stream);
8595 unsafe {
8596 memra_fa3_prefill(qp as *const core::ffi::c_void,
8597 kp as *const core::ffi::c_void,
8598 vp as *const core::ffi::c_void,
8599 op as *mut f32,
8600 t as i32, n_head as i32, n_head_kv as i32,
8601 head_dim as i32, scale,
8602 stream.cu_stream() as *mut core::ffi::c_void)
8603 }
8604 };
8605 if rc != 0 {
8606 return Err(format!("memra_fa3_prefill rc={rc}").into());
8607 }
8608 return Ok(());
8609 }
8610 static FA_P1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8615 let fa_p1 = *FA_P1.get_or_init(|| std::env::var("MEMRA_FA_P1").as_deref() == Ok("1"));
8616 if fa_p1 && head_dim == 256 && !std::env::var("MEMRA_FA_FLOOR").is_ok() {
8617 const BLOCK_Q: usize = 64; const BKX: usize = 32;
8618 let f = self.func("fa_prefill_bf16_p1");
8619 let shmem = (2 * (2 * BKX * head_dim + BLOCK_Q * BKX)
8620 + 4 * (BLOCK_Q * BKX + 2 * BLOCK_Q)) as u32;
8621 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8622 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8623 let cfg = LaunchConfig {
8624 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8625 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8626 };
8627 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32,
8628 n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8629 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8630 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8631 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8632 let __s_b = self.gpu.stream();
8633 let mut b = __s_b.launch_builder(&f);
8634 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti)
8635 .arg(&tkvi).arg(&scale).arg(&cz);
8636 unsafe { b.launch(cfg)?; }
8637 return Ok(());
8638 }
8639 const BK: usize = 32;
8645 let w2 = std::env::var("MEMRA_FA_PP_W2").as_deref() == Ok("1");
8648 let (block_q, warps, w2_sfx): (usize, u32, &str) =
8649 if w2 { (32, 2, "_w2") } else { (64, 4, "") };
8650 let hd_sfx = fa_hd_suffix(head_dim)?;
8654 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8655 let bf16kv = !floor && !w2
8660 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0");
8661 let (kb16, vb16) = if bf16kv {
8662 let n = t_kv * n_head_kv * head_dim;
8663 let mut kb = self.alloc_u8_uninit(n * 2)?;
8664 let mut vb = self.alloc_u8_uninit(n * 2)?;
8665 let fcv = self.func("f32_to_bf16_bulk");
8666 let ni = n as i64;
8667 let cfgc = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
8668 let __s_b = self.gpu.stream();
8669 let mut b = __s_b.launch_builder(&fcv);
8670 b.arg(k).arg(&mut kb).arg(&ni);
8671 unsafe { b.launch(cfgc)?; }
8672 let __s_b = self.gpu.stream();
8673 let mut b = __s_b.launch_builder(&fcv);
8674 b.arg(v).arg(&mut vb).arg(&ni);
8675 unsafe { b.launch(cfgc)?; }
8676 (Some(kb), Some(vb))
8677 } else {
8678 (None, None)
8679 };
8680 let f = self.func(&if bf16kv {
8681 format!("fa_prefill_bf16kv_pp{hd_sfx}")
8682 } else {
8683 format!("fa_prefill_f32{}{}{hd_sfx}",
8684 if floor { "" } else { "_pp" },
8685 if floor { "" } else { w2_sfx })
8686 });
8687 let kv_stages = if bf16kv { 2 } else { 1 };
8690 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
8691 + 4 * (block_q * BK + 2 * block_q)) as u32;
8692 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8693 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8694 let cfg = LaunchConfig {
8695 grid_dim: ((t as u32 + block_q as u32 - 1) / block_q as u32, n_head as u32, 1),
8696 block_dim: (32, warps, 1), shared_mem_bytes: shmem,
8697 };
8698 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);
8699 let __s_b = self.gpu.stream();
8700 let mut b = __s_b.launch_builder(&f);
8701 b.arg(q);
8702 match (&kb16, &vb16) {
8703 (Some(kb), Some(vb)) => { b.arg(kb).arg(vb); }
8704 _ => { b.arg(k).arg(v); }
8705 }
8706 b.arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8707 unsafe { b.launch(cfg)?; }
8708 Ok(())
8709 }
8710
8711 #[allow(clippy::too_many_arguments)]
8715 pub fn fa_prefill_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8716 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8717 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8718 -> Result<(), Box<dyn std::error::Error>> {
8719 if portable_mma_gated() {
8722 return self.sdpa_naive_w(q, k, v, o, head_dim, n_head, n_head_kv,
8723 t, t_kv, scale, causal, window);
8724 }
8725 static FAW_F32: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8729 let faw_f32 = *FAW_F32.get_or_init(|| {
8730 std::env::var("MEMRA_FAW_STAGE").as_deref() == Ok("f32")
8731 });
8732 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8733 self.fa_prefill_w_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8734 window, floor || faw_f32, floor)
8735 }
8736
8737 #[allow(clippy::too_many_arguments)]
8740 pub fn fa_prefill_w_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
8741 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8742 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
8743 window: usize, v_f16: bool)
8744 -> Result<(), Box<dyn std::error::Error>> {
8745 const BLOCK_Q: usize = 64; const BK: usize = 32;
8746 debug_assert_eq!(head_dim, 256);
8747 let hp = fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
8748 && (n_head / n_head_kv) % 2 == 0;
8749 debug_assert!(!v_f16 || hp, "f16 V emitted but the SWA hp arm is off");
8750 if hp {
8751 const BLOCK_QH: usize = 32;
8752 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
8755 let vh: &CudaSlice<u8> = if v_f16 { vb } else {
8756 let n = t_kv * n_head_kv * head_dim;
8757 if vguard.as_ref().map(|b| b.len() < n * 2).unwrap_or(true) {
8758 *vguard = Some(self.alloc_uninit::<u8>(n * 2)?);
8759 }
8760 self.bf16_to_f16_into(vb, n, vguard.as_mut().unwrap())?;
8761 vguard.as_ref().unwrap()
8762 };
8763 let f = self.func("fa_prefill_w_bf16_p1h2");
8764 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
8765 + 4 * (2 * BLOCK_QH)) as u32;
8766 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8767 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8768 let cfg = LaunchConfig {
8769 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
8770 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8771 };
8772 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8773 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8774 let __s_b = self.gpu.stream();
8775 let mut b = __s_b.launch_builder(&f);
8776 b.arg(qb).arg(kb).arg(vh).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8777 .arg(&scale).arg(&cz).arg(&wi);
8778 unsafe { b.launch(cfg)?; }
8779 return Ok(());
8780 }
8781 let f = self.func("fa_prefill_w_bf16_p1");
8782 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
8783 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
8784 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8785 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8786 let cfg = LaunchConfig {
8787 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8788 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8789 };
8790 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8791 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8792 let __s_b = self.gpu.stream();
8793 let mut b = __s_b.launch_builder(&f);
8794 b.arg(qb).arg(kb).arg(vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8795 .arg(&scale).arg(&cz).arg(&wi);
8796 unsafe { b.launch(cfg)?; }
8797 Ok(())
8798 }
8799
8800 #[allow(clippy::too_many_arguments)]
8802 pub fn fa_prefill_w_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8803 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8804 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
8805 window: usize, f32_stage: bool, floor: bool)
8806 -> Result<(), Box<dyn std::error::Error>> {
8807 const BLOCK_Q: usize = 64; const BK: usize = 32;
8808 debug_assert_eq!(head_dim, 256, "fa_prefill_w is stamped hd256 only");
8809 static P1_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8813 let p1 = !floor && !f32_stage
8814 && *P1_ON.get_or_init(|| {
8815 std::env::var("MEMRA_FAW_P1").map(|v| v != "0").unwrap_or(true)
8816 });
8817 let hp = p1 && fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
8818 && (n_head / n_head_kv) % 2 == 0;
8819 if hp {
8820 const BLOCK_QH: usize = 32;
8821 let f = self.func("fa_prefill_w_bf16_p1h2");
8822 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
8823 + 4 * (2 * BLOCK_QH)) as u32;
8824 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8825 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8826 let cfg = LaunchConfig {
8827 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
8828 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8829 };
8830 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8831 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8832 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8833 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8834 let vh = self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?;
8835 let __s_b = self.gpu.stream();
8836 let mut b = __s_b.launch_builder(&f);
8837 b.arg(&qb).arg(&kb).arg(&vh).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8838 .arg(&scale).arg(&cz).arg(&wi);
8839 unsafe { b.launch(cfg)?; }
8840 return Ok(());
8841 }
8842 if p1 {
8843 let f = self.func("fa_prefill_w_bf16_p1");
8844 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
8845 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
8846 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8847 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8848 let cfg = LaunchConfig {
8849 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8850 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8851 };
8852 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8853 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8854 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8855 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8856 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8857 let __s_b = self.gpu.stream();
8858 let mut b = __s_b.launch_builder(&f);
8859 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8860 .arg(&scale).arg(&cz).arg(&wi);
8861 unsafe { b.launch(cfg)?; }
8862 return Ok(());
8863 }
8864 static G4_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8867 let g4 = !floor && !f32_stage && n_head_kv == 1 && n_head % 4 == 0
8868 && *G4_ON.get_or_init(|| {
8869 std::env::var("MEMRA_FAW_G4").map(|v| v != "0").unwrap_or(true)
8870 });
8871 if g4 {
8872 const SP_M: usize = 16;
8873 static O2_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8876 let o2 = *O2_ON.get_or_init(|| {
8877 std::env::var("MEMRA_FAW_O2").map(|v| v != "0").unwrap_or(true)
8878 });
8879 let f = self.func(if o2 { "fa_prefill_w_bf16_g4o2" } else { "fa_prefill_w_bf16_g4" });
8880 let shmem = if o2 {
8881 (2 * (4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M)) as u32
8882 } else {
8883 (2 * (2 * BK * head_dim + 4 * SP_M * head_dim + 4 * SP_M * BK)
8884 + 4 * (4 * SP_M)) as u32
8885 };
8886 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8887 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8888 let cfg = LaunchConfig {
8889 grid_dim: ((t as u32).div_ceil(SP_M as u32), (n_head / 4) as u32, 1),
8890 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8891 };
8892 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8893 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8894 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8895 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8896 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8897 let __s_b = self.gpu.stream();
8898 let mut b = __s_b.launch_builder(&f);
8899 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8900 .arg(&scale).arg(&cz).arg(&wi);
8901 unsafe { b.launch(cfg)?; }
8902 return Ok(());
8903 }
8904 let f = self.func(if floor { "fa_prefill_w_f32" }
8905 else if f32_stage { "fa_prefill_w_f32_pp" }
8906 else { "fa_prefill_w_bf16_pp" });
8907 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
8908 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
8909 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8910 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8911 let cfg = LaunchConfig {
8912 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8913 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8914 };
8915 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
8916 t as i32, t_kv as i32, causal as i32, window as i32);
8917 if f32_stage {
8918 let __s_b = self.gpu.stream();
8919 let mut b = __s_b.launch_builder(&f);
8920 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8921 .arg(&scale).arg(&cz).arg(&wi);
8922 unsafe { b.launch(cfg)?; }
8923 } else {
8924 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8925 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8926 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8927 let __s_b = self.gpu.stream();
8928 let mut b = __s_b.launch_builder(&f);
8929 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8930 .arg(&scale).arg(&cz).arg(&wi);
8931 unsafe { b.launch(cfg)?; }
8932 }
8933 Ok(())
8934 }
8935
8936 #[allow(clippy::too_many_arguments)]
8940 pub fn fa_prefill_hd512(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8941 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8942 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool)
8943 -> Result<(), Box<dyn std::error::Error>> {
8944 if portable_mma_gated() {
8946 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
8947 t, t_kv, scale, causal);
8948 }
8949 static F32_STAGE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8955 let f32_stage = *F32_STAGE.get_or_init(|| {
8956 std::env::var("MEMRA_FA512_STAGE").as_deref() == Ok("f32")
8957 });
8958 static SP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8962 let sp = !f32_stage
8963 && *SP_ON.get_or_init(|| {
8964 std::env::var("MEMRA_FA512_SP").map(|v| v != "0").unwrap_or(true)
8965 });
8966 self.fa_prefill_hd512_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale,
8967 causal, f32_stage, sp, sp && fa_f16pv_on())
8968 }
8969
8970 #[allow(clippy::too_many_arguments)]
8972 pub fn fa_prefill_hd512_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
8973 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8974 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
8975 v_f16: bool)
8976 -> Result<(), Box<dyn std::error::Error>> {
8977 debug_assert_eq!(head_dim, 512);
8978 const SP_M: usize = 16; const BKS: usize = 32;
8979 let f16pv = fa_f16pv_on();
8983 let nw = if f16pv { fa512_wide_warps() } else { 2 };
8984 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
8985 debug_assert!(!v_f16 || f16pv, "f16 V emitted without the door on");
8986 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
8987 let vref: &CudaSlice<u8> = if f16pv && !v_f16 {
8988 let n = t_kv * n_head_kv * head_dim;
8990 let need = n * 2;
8991 if vguard.as_ref().map(|b| b.len() < need).unwrap_or(true) {
8992 *vguard = Some(self.alloc_uninit::<u8>(need)?);
8993 }
8994 let dst = vguard.as_mut().unwrap();
8995 self.bf16_to_f16_into(vb, n, dst)?;
8996 vguard.as_ref().unwrap()
8997 } else { vb };
8998 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
8999 else { match (f16pv, nw) {
9000 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9001 (true, _) => "fa_prefill_bf16_hd512_sp16",
9002 _ => "fa_prefill_bf16_hd512_sp",
9003 } });
9004 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9005 let shmem = if hp {
9007 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9008 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9009 } else {
9010 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9011 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9012 };
9013 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9014 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9015 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9016 let cfg = LaunchConfig {
9017 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9018 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9019 };
9020 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9021 t as i32, t_kv as i32, causal as i32);
9022 let __s_b = self.gpu.stream();
9023 let mut b = __s_b.launch_builder(&f);
9024 b.arg(qb).arg(kb).arg(vref).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9025 .arg(&scale).arg(&cz);
9026 unsafe { b.launch(cfg)?; }
9027 Ok(())
9028 }
9029
9030 #[allow(clippy::too_many_arguments)]
9033 pub fn fa_prefill_hd512_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9034 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9035 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9036 f32_stage: bool, sp: bool, f16pv: bool)
9037 -> Result<(), Box<dyn std::error::Error>> {
9038 debug_assert_eq!(head_dim, 512, "fa_prefill_hd512 is hd512 only");
9039 if sp && !f32_stage {
9040 const SP_M: usize = 16; const BKS: usize = 32;
9044 let nw = if f16pv { fa512_wide_warps() } else { 2 };
9045 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
9046 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
9047 else { match (f16pv, nw) {
9048 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9049 (true, _) => "fa_prefill_bf16_hd512_sp16",
9050 _ => "fa_prefill_bf16_hd512_sp",
9051 } });
9052 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9053 let shmem = if hp {
9054 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9055 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9056 } else {
9057 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9058 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9059 };
9060 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9061 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9062 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9063 let cfg = LaunchConfig {
9064 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9065 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9066 };
9067 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9068 t as i32, t_kv as i32, causal as i32);
9069 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9070 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9071 let vb = if f16pv { self.f32_to_f16(v, t_kv * n_head_kv * head_dim)? }
9072 else { self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)? };
9073 let __s_b = self.gpu.stream();
9074 let mut b = __s_b.launch_builder(&f);
9075 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9076 .arg(&scale).arg(&cz);
9077 unsafe { b.launch(cfg)?; }
9078 return Ok(());
9079 }
9080 const BLOCK_Q: usize = 32; const BK: usize = 32; const HALF: usize = 256;
9081 let f = self.func(if f32_stage { "fa_prefill_f32_hd512" } else { "fa_prefill_bf16_hd512" });
9082 let shmem = (2 * (BLOCK_Q * head_dim + BK * head_dim + BK * HALF + BLOCK_Q * BK)
9084 + 4 * BLOCK_Q) as u32;
9085 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9086 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9087 let cfg = LaunchConfig {
9088 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 2),
9089 block_dim: (32, 2, 1), shared_mem_bytes: shmem,
9090 };
9091 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9092 t as i32, t_kv as i32, causal as i32);
9093 if f32_stage {
9094 let __s_b = self.gpu.stream();
9095 let mut b = __s_b.launch_builder(&f);
9096 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9097 .arg(&scale).arg(&cz);
9098 unsafe { b.launch(cfg)?; }
9099 } else {
9100 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9101 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9102 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9103 let __s_b = self.gpu.stream();
9104 let mut b = __s_b.launch_builder(&f);
9105 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9106 .arg(&scale).arg(&cz);
9107 unsafe { b.launch(cfg)?; }
9108 }
9109 Ok(())
9110 }
9111
9112 #[allow(clippy::too_many_arguments)]
9116 pub fn rope_neox2_bf16e(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
9117 qb: &mut CudaSlice<u8>, kb: &mut CudaSlice<u8>,
9118 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
9119 nh_q: usize, nh_k: usize, n_tokens: usize, base: f32,
9120 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
9121 -> Result<(), Box<dyn std::error::Error>> {
9122 let f = self.func("rope_neox2_bf16e_f32");
9123 let rows = ((nh_q + nh_k) * n_tokens) as u32;
9124 let cfg = LaunchConfig { grid_dim: (rows, 1, 1),
9125 block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
9126 let theta_scale = base.powf(-2.0 / n_dims as f32);
9127 let (hd, nd, nhq, nhk, nt) = (head_dim as i32, n_dims as i32, nh_q as i32,
9128 nh_k as i32, n_tokens as i32);
9129 let __s_b = self.gpu.stream();
9130 let mut b = __s_b.launch_builder(&f);
9131 match ff {
9132 Some(t) => { b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9133 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9134 .arg(&theta_scale).arg(&freq_scale).arg(t);
9135 unsafe { b.launch(cfg)?; } }
9136 None => { let null: u64 = 0;
9137 b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9138 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9139 .arg(&theta_scale).arg(&freq_scale).arg(&null);
9140 unsafe { b.launch(cfg)?; } }
9141 }
9142 Ok(())
9143 }
9144
9145 pub fn f32_to_bf16(&self, x: &CudaSlice<f32>, n: usize)
9148 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9149 assert!(n % 4 == 0, "f32_to_bf16 requires n % 4 == 0, got {n}");
9150 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9151 let f = self.func("f32_to_bf16_flat");
9152 let n_i = n as i64;
9153 let cfg = LaunchConfig {
9154 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9155 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9156 };
9157 let __s_b = self.gpu.stream();
9158 let mut b = __s_b.launch_builder(&f);
9159 b.arg(x).arg(&mut y).arg(&n_i);
9160 unsafe { b.launch(cfg)?; }
9161 Ok(y)
9162 }
9163
9164 pub fn f32_to_f16(&self, x: &CudaSlice<f32>, n: usize)
9165 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9166 assert!(n % 4 == 0, "f32_to_f16 requires n % 4 == 0, got {n}");
9167 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9168 let f = self.func("f32_to_f16_flat");
9169 let n_i = n as i64;
9170 let cfg = LaunchConfig {
9171 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9172 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9173 };
9174 let __s_b = self.gpu.stream();
9175 let mut b = __s_b.launch_builder(&f);
9176 b.arg(x).arg(&mut y).arg(&n_i);
9177 unsafe { b.launch(cfg)?; }
9178 Ok(y)
9179 }
9180
9181 pub fn bf16_to_f16(&self, xb: &CudaSlice<u8>, n: usize)
9183 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9184 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9185 self.bf16_to_f16_into(xb, n, &mut y)?;
9186 Ok(y)
9187 }
9188
9189 pub fn bf16_to_f16_into(&self, xb: &CudaSlice<u8>, n: usize, y: &mut CudaSlice<u8>)
9191 -> Result<(), Box<dyn std::error::Error>> {
9192 assert!(n % 2 == 0, "bf16_to_f16 requires n % 2 == 0, got {n}");
9193 assert!(y.len() >= n * 2);
9194 let f = self.func("bf16_to_f16_flat");
9195 let n2 = (n / 2) as i64;
9196 let cfg = LaunchConfig {
9197 grid_dim: (((n / 2) as u32).div_ceil(256), 1, 1),
9198 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9199 };
9200 let __s_b = self.gpu.stream();
9201 let mut b = __s_b.launch_builder(&f);
9202 b.arg(xb).arg(y).arg(&n2);
9203 unsafe { b.launch(cfg)?; }
9204 Ok(())
9205 }
9206
9207 #[allow(clippy::too_many_arguments)]
9212 pub fn fa_prefill_vl8(&self, seqs: &[FaSeqVl], head_dim: usize, n_head: usize,
9213 n_head_kv: usize, scale: f32)
9214 -> Result<(), Box<dyn std::error::Error>> {
9215 const BK: usize = 32;
9216 let b = seqs.len();
9217 assert!(b >= 1 && b <= 8);
9218 let mut packed = [FaSeqVl::default(); 8];
9219 packed[..b].copy_from_slice(seqs);
9220 let v = FaVl8(packed);
9221 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9222 let ept = (n_head_kv * head_dim) as i32;
9223 {
9224 let f = self.func("fa_mirror_vl");
9225 let max_n = (max_t as i64) * ept as i64;
9226 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
9227 for which in 0..2i32 {
9228 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9229 let __s_lb = self.gpu.stream();
9230 let mut lb = __s_lb.launch_builder(&f);
9231 lb.arg(&v).arg(&ept).arg(&which);
9232 unsafe { lb.launch(cfg)?; }
9233 }
9234 }
9235 let hd_sfx = fa_hd_suffix(head_dim)?;
9236 let f = self.func(&format!("fa_prefill_bf16kv_vl{hd_sfx}"));
9237 let block_q = 64usize;
9238 let kv_stages = 2usize;
9239 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
9240 + 4 * (block_q * BK + 2 * block_q)) as u32;
9241 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9242 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9243 let cfg = LaunchConfig {
9244 grid_dim: (max_t.div_ceil(block_q as u32), n_head as u32, b as u32),
9245 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9246 };
9247 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9248 let __s_lb = self.gpu.stream();
9249 let mut lb = __s_lb.launch_builder(&f);
9250 lb.arg(&v).arg(&hd).arg(&nh).arg(&nhkv).arg(&scale);
9251 unsafe { lb.launch(cfg)?; }
9252 Ok(())
9253 }
9254
9255 #[allow(clippy::too_many_arguments)]
9259 pub fn attn_pre_vl8(&self, seqs: &[AttnPreVl], wq: &CudaSlice<f32>, wk: &CudaSlice<f32>,
9260 head_dim: usize, rope_dims: usize, n_head: usize, n_head_kv: usize,
9261 eps: f32, freq_base: f32, freq_scale: f32,
9262 kv_dim_k: usize, kv_dim_v: usize,
9263 k_tok_bytes: usize, v_tok_bytes: usize)
9264 -> Result<(), Box<dyn std::error::Error>> {
9265 let b = seqs.len();
9266 assert!(b >= 1 && b <= 8);
9267 let mut packed = [AttnPreVl::default(); 8];
9268 packed[..b].copy_from_slice(seqs);
9269 let v = AttnPreVl8(packed);
9270 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9271 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9272 {
9273 let f = self.func("q_gate_split_vl");
9274 let n = max_t * (n_head * head_dim) as u32;
9275 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9276 let __s_lb = self.gpu.stream();
9277 let mut lb = __s_lb.launch_builder(&f);
9278 lb.arg(&v).arg(&hd).arg(&nh);
9279 unsafe { lb.launch(cfg)?; }
9280 }
9281 {
9282 let f = self.func("attn_rms_vl");
9283 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 };
9284 let __s_lb = self.gpu.stream();
9285 let mut lb = __s_lb.launch_builder(&f);
9286 lb.arg(&v).arg(wq).arg(wk).arg(&hd).arg(&nh).arg(&nhkv).arg(&eps);
9287 unsafe { lb.launch(cfg)?; }
9288 }
9289 {
9290 let f = self.func("attn_rope_vl");
9291 let theta_scale = freq_base.powf(-2.0 / rope_dims as f32);
9292 let nd = rope_dims as i32;
9293 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 };
9294 let __s_lb = self.gpu.stream();
9295 let mut lb = __s_lb.launch_builder(&f);
9296 lb.arg(&v).arg(&hd).arg(&nd).arg(&nh).arg(&nhkv).arg(&theta_scale).arg(&freq_scale);
9297 unsafe { lb.launch(cfg)?; }
9298 }
9299 {
9300 let f = self.func("append_kv_vl");
9301 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
9302 let cfg = LaunchConfig { grid_dim: (nblk, max_t, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
9303 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9304 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9305 let __s_lb = self.gpu.stream();
9306 let mut lb = __s_lb.launch_builder(&f);
9307 lb.arg(&v).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
9308 unsafe { lb.launch(cfg)?; }
9309 }
9310 Ok(())
9311 }
9312
9313 pub fn fa_prefill_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9318 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9319 head_dim: usize, n_head: usize, n_head_kv: usize,
9320 t: usize, t_kv: usize, scale: f32, causal: bool,
9321 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9322 -> Result<(), Box<dyn std::error::Error>> {
9323 if portable_mma_gated() {
9324 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9325 t, t_kv, scale, causal,
9326 k_tok_bytes, v_tok_bytes);
9327 }
9328 const BLOCK_Q: usize = 64; const BK: usize = 32;
9329 let name = format!("fa_prefill_q{}", fa_hd_suffix(head_dim)?);
9332 let f = if g { self.func_g(&name) } else { self.func(&name) };
9333 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9334 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9335 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9336 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9337 let cfg = LaunchConfig {
9338 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9339 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9340 };
9341 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);
9342 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9343 let __s_b = self.gpu.stream();
9344 let mut b = __s_b.launch_builder(&f);
9345 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9346 .arg(&ktb).arg(&vtb);
9347 unsafe { b.launch(cfg)?; }
9348 Ok(())
9349 }
9350
9351 #[allow(clippy::too_many_arguments)]
9361 pub fn fa_prefill_view_ws(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9362 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9363 head_dim: usize, n_head: usize, n_head_kv: usize,
9364 t: usize, t_kv: usize, scale: f32, causal: bool,
9365 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9366 -> Result<(), Box<dyn std::error::Error>> {
9367 if portable_mma_gated() {
9368 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9369 t, t_kv, scale, causal,
9370 k_tok_bytes, v_tok_bytes);
9371 }
9372 const BLOCK_Q: usize = 64; const BK: usize = 32;
9373 let kv_dim_k = n_head_kv * head_dim;
9374 let kv_dim_v = n_head_kv * head_dim;
9375 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9377 let mut guard = self.prime_deqw_ws.lock().unwrap();
9379 let need_grow = match guard.as_ref() {
9380 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9381 None => true,
9382 };
9383 if need_grow {
9384 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9385 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9386 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9387 }
9388 let (kw, vw) = guard.as_mut().unwrap();
9389 {
9391 let f = if g { self.func_g("fa_dequant_kv_ws_bf16") } else { self.func("fa_dequant_kv_ws_bf16") };
9393 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9394 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9395 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9396 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9397 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9398 let __s_b = self.gpu.stream();
9399 let mut b = __s_b.launch_builder(&f);
9400 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9401 unsafe { b.launch(cfg)?; }
9402 }
9403 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9411 {
9412 let hd_sfx = fa_hd_suffix(head_dim)?;
9413 let f = self.func(&format!("fa_prefill_qw{}{hd_sfx}", if db { "_db" } else { "" }));
9414 let shmem = if db {
9415 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9417 } else {
9418 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9419 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9420 };
9421 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9422 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9423 let cfg = LaunchConfig {
9424 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9425 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9426 };
9427 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);
9428 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9429 let __s_b = self.gpu.stream();
9430 let mut b = __s_b.launch_builder(&f);
9431 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9432 .arg(&kdk).arg(&kdv);
9433 unsafe { b.launch(cfg)?; }
9434 }
9435 Ok(())
9436 }
9437
9438 #[allow(clippy::too_many_arguments)]
9454 pub fn fa_prefill_view_ws_w_hd128(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9455 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9456 head_dim: usize, n_head: usize, n_head_kv: usize,
9457 t: usize, t_kv: usize, scale: f32, causal: bool,
9458 window: usize, k_tok_bytes: usize, v_tok_bytes: usize)
9459 -> Result<(), Box<dyn std::error::Error>> {
9460 assert_eq!(head_dim, 128, "fa_prefill_view_ws_w_hd128: only the hd128 twin is stamped");
9461 if portable_mma_gated() {
9462 return self.sdpa_naive_w_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9463 t, t_kv, scale, causal, window,
9464 k_tok_bytes, v_tok_bytes);
9465 }
9466 const BLOCK_Q: usize = 64; const BK: usize = 32;
9467 let kv_dim_k = n_head_kv * head_dim;
9468 let kv_dim_v = n_head_kv * head_dim;
9469 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9471 let mut guard = self.prime_deqw_ws.lock().unwrap();
9472 let need_grow = match guard.as_ref() {
9473 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9474 None => true,
9475 };
9476 if need_grow {
9477 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9478 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9479 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9480 }
9481 let (kw, vw) = guard.as_mut().unwrap();
9482 {
9485 let f = self.func("fa_dequant_kv_ws_bf16");
9486 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9487 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9488 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9489 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9490 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9491 let __s_b = self.gpu.stream();
9492 let mut b = __s_b.launch_builder(&f);
9493 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9494 unsafe { b.launch(cfg)?; }
9495 }
9496 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9498 {
9499 let f = self.func(if db { "fa_prefill_qw_db_w_hd128" } else { "fa_prefill_qw_w_hd128" });
9500 let shmem = if db {
9501 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9502 } else {
9503 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9504 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9505 };
9506 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9507 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9508 let cfg = LaunchConfig {
9509 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9510 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9511 };
9512 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);
9513 let (kdk, kdv, wnd) = (kv_dim_k as i32, kv_dim_v as i32, window as i32);
9514 let __s_b = self.gpu.stream();
9515 let mut b = __s_b.launch_builder(&f);
9516 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9517 .arg(&kdk).arg(&kdv).arg(&wnd);
9518 unsafe { b.launch(cfg)?; }
9519 }
9520 Ok(())
9521 }
9522
9523 pub fn fa_decode(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9527 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9528 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9529 k_tok_bytes: usize, v_tok_bytes: usize)
9530 -> Result<(), Box<dyn std::error::Error>> {
9531 self.fa_decode_kvmod(q, k, v, o, head_dim, n_head, n_head_kv, t_kv, scale,
9532 k_tok_bytes, v_tok_bytes, false)
9533 }
9534
9535 #[allow(clippy::too_many_arguments)]
9539 #[allow(clippy::too_many_arguments)]
9543 #[allow(clippy::too_many_arguments)]
9544 fn fa_decode_scalar_unified(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9545 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9546 head_dim: usize, n_head: usize, n_head_kv: usize,
9547 t_kv_host: usize, t_kv_dev: Option<&CudaSlice<i32>>,
9548 scale: f32, n_splits: usize, split_keys: usize,
9549 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
9550 part_o: &mut CudaSlice<f32>, part_m: &mut CudaSlice<f32>,
9551 part_l: &mut CudaSlice<f32>,
9552 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
9553 -> Result<(), Box<dyn std::error::Error>> {
9554 let f = if g { self.func_g("fa_decode_f32") } else { self.fa_func("fa_decode_f32", head_dim) };
9555 let cfg = LaunchConfig { grid_dim: (n_head as u32, n_splits as u32, 1),
9556 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: (4 * (head_dim + 32)) as u32 };
9557 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
9558 let (ktb, vtb, tkvi, ski) = (k_tok_bytes as i64, v_tok_bytes as i64, t_kv_host as i32,
9559 split_keys as i32);
9560 let __s_b = self.gpu.stream();
9561 let mut b = __s_b.launch_builder(&f);
9562 match t_kv_dev {
9563 Some(d) => { b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9564 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(d).arg(&scale).arg(&nsp)
9565 .arg(&ski).arg(&ktb).arg(&vtb);
9566 unsafe { b.launch(cfg)?; } }
9567 None => { let null: u64 = 0;
9568 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9569 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&null).arg(&scale).arg(&nsp)
9570 .arg(&ski).arg(&ktb).arg(&vtb);
9571 unsafe { b.launch(cfg)?; } }
9572 }
9573 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1),
9574 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
9575 if let Some((oq, od)) = q8_out {
9576 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
9578 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
9579 let __s_b2 = self.gpu.stream();
9580 let mut b2 = __s_b2.launch_builder(&fc);
9581 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
9582 unsafe { b2.launch(cfg2)?; }
9583 return Ok(());
9584 }
9585 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
9586 let __s_b2 = self.gpu.stream();
9587 let mut b2 = __s_b2.launch_builder(&fc);
9588 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
9589 unsafe { b2.launch(cfg2)?; }
9590 Ok(())
9591 }
9592
9593 pub fn fa_decode_kvmod(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9594 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9595 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9596 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9597 -> Result<(), Box<dyn std::error::Error>> {
9598 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
9619 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
9623 let sp = fa_split_keys(t_kv, n_head_kv);
9624 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
9625 let o_len = n_head * n_splits * head_dim;
9626 let ml_len = n_head * n_splits;
9627 let mut part_guard = self.fa_part_pool.lock().unwrap();
9628 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
9629 let old = part_guard.take();
9640 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
9641 if let Some(old) = old {
9642 self.fa_part_retired.lock().unwrap().push(old);
9643 }
9644 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
9645 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
9646 }
9647 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
9648 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
9649 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
9650 }
9651 let pg = part_guard.as_mut().unwrap();
9652 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
9653 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
9654 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
9655 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
9656 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
9657 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);
9658 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9659 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
9663 let fa512_min = fa512_min_tkv();
9668 let deep = fa_vec && head_dim == 256 && fa_v4_at(t_kv) && !g
9671 && fa_deep_at(t_kv) && !matches!(fa_v4_mode(), "noB3" | "stage");
9672 let (f, cfg) = if fa_vec && head_dim == 512 && t_kv >= fa512_min {
9673 let gqa = (n_head / n_head_kv).max(1) as u32;
9676 let fv = self.fa_func("fa_decode_vec_q_dpl16", head_dim);
9677 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9678 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
9679 } else if fa_vec && head_dim <= 256 {
9680 let gqa = (n_head / n_head_kv).max(1) as u32;
9681 static SMEM_TKV: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9692 let smem_tkv = *SMEM_TKV.get_or_init(|| {
9693 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
9694 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
9695 });
9696 if fa_v4_at(t_kv) && head_dim == 256 {
9697 let v4name = match fa_v4_mode() {
9701 "noB3" => "fa_decode_vec_q_v4_noB3", "stage" => "fa_decode_vec_q_v4_stage", _ if deep => "fa_decode_vec_q_v4_deep",
9704 _ => "fa_decode_vec_q_v4",
9705 };
9706 let fv = if g { self.func_g(v4name) } else { self.func(v4name) };
9707 let shmem = (if deep { 12160 } else { 11520 }
9710 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
9711 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9712 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9713 (fv,
9714 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9715 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9716 } else if fa_v3_active(head_dim) {
9717 let fv = if g { self.func_g("fa_decode_vec_q_v3") } else { self.func("fa_decode_vec_q_v3") };
9720 let shmem = (32 * head_dim * 2) as u32; (fv,
9722 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9723 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9724 } else if fa_v2_on() {
9725 let fv = if g { self.func_g("fa_decode_vec_q_v2") } else { self.func("fa_decode_vec_q_v2") };
9729 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
9731 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9732 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9733 } else if smem_tkv > 0 && t_kv >= smem_tkv && !g
9734 && !(head_dim == 512 && Self::gkv_on()) {
9735 let fv = if g { self.func_g("fa_decode_vec_q_smem") } else { self.func("fa_decode_vec_q_smem") };
9739 let shmem = (2 * 32 * head_dim * 2) as u32; use cudarc::driver::sys::CUfunction_attribute_enum as A;
9741 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9742 (fv,
9743 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9744 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9745 } else {
9746 let fv = if g { self.func_g("fa_decode_vec_q") } else { self.func("fa_decode_vec_q") };
9749 (fv,
9750 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9751 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
9752 }
9753 } else {
9754 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
9757 t_kv, None, scale, n_splits,
9758 if fa_vec { sp } else { 256 },
9759 k_tok_bytes, v_tok_bytes, g,
9760 part_o, part_m, part_l, None);
9761 };
9762 let __s_b = self.gpu.stream();
9763 let mut b = __s_b.launch_builder(&f);
9764 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9765 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&scale).arg(&nsp).arg(&ktb).arg(&vtb);
9766 unsafe { b.launch(cfg)?; }
9767 let (fc, cfg2) = (if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) },
9770 LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 });
9771 let __s_b2 = self.gpu.stream();
9772 let mut b2 = __s_b2.launch_builder(&fc);
9773 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
9774 unsafe { b2.launch(cfg2)?; }
9775 Ok(())
9776 }
9777
9778 #[allow(clippy::too_many_arguments)]
9789 pub fn fa_decode_batch_seqs_v4(&self, q: &CudaSlice<f32>,
9790 kv_ptrs: &cudarc::driver::CudaView<u64>,
9791 pos_seq: &CudaSlice<i32>, o: &mut CudaSlice<f32>,
9792 head_dim: usize, n_head: usize, n_head_kv: usize,
9793 b_n: usize, t_kv_max: usize, scale: f32,
9794 split_keys: usize, k_tok_bytes: usize, v_tok_bytes: usize)
9795 -> Result<(), Box<dyn std::error::Error>> {
9796 debug_assert!(head_dim == 256, "seqs twin is v4-stamped (hd256 only)");
9797 let n_splits_max = (t_kv_max + split_keys - 1) / split_keys;
9798 let o_len = b_n * n_head * n_splits_max * head_dim;
9799 let ml_len = b_n * n_head * n_splits_max;
9800 let mut part_guard = self.fa_part_pool.lock().unwrap();
9801 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
9802 let old = part_guard.take();
9813 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
9814 if let Some(old) = old {
9815 self.fa_part_retired.lock().unwrap().push(old);
9816 }
9817 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
9818 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
9819 }
9820 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
9821 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
9822 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
9823 }
9824 let pg = part_guard.as_mut().unwrap();
9825 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
9826 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
9827 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
9828 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
9829 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9830 let (nspm, spk) = (n_splits_max as i32, split_keys as i32);
9831 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9832 let gqa = (n_head / n_head_kv).max(1) as u32;
9833 let f = self.func("fa_decode_vec_q_seqs_v4");
9834 let shmem = (11520 + 32 * head_dim * 2) as u32;
9836 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9837 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9838 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, b_n as u32),
9839 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
9840 {
9841 let __s_b = self.gpu.stream();
9842 let mut b = __s_b.launch_builder(&f);
9843 b.arg(q).arg(kv_ptrs).arg(pos_seq).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9844 .arg(&hd).arg(&nh).arg(&nhkv).arg(&scale).arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
9845 unsafe { b.launch(cfg)?; }
9846 }
9847 let fc = self.func("fa_decode_combine_seqs");
9848 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, b_n as u32, 1),
9849 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
9850 let __s_b2 = self.gpu.stream();
9851 let mut b2 = __s_b2.launch_builder(&fc);
9852 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
9853 .arg(pos_seq).arg(&nspm).arg(&spk);
9854 unsafe { b2.launch(cfg2)?; }
9855 Ok(())
9856 }
9857
9858 #[allow(clippy::too_many_arguments)]
9865 pub fn append_kv_quantized_seqs(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
9866 kv_ptrs: &cudarc::driver::CudaView<u64>,
9867 pos_seq: &CudaSlice<i32>, b_n: usize,
9868 kv_dim_k: usize, kv_dim_v: usize,
9869 k_tok_bytes: usize, v_tok_bytes: usize)
9870 -> Result<(), Box<dyn std::error::Error>> {
9871 let f = self.func("append_quantize_kv_q8_0_q5_1_seqs");
9872 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
9873 let cfg = LaunchConfig { grid_dim: (nblk, b_n as u32, 1),
9874 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
9875 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9876 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9877 let __s_b = self.gpu.stream();
9878 let mut b = __s_b.launch_builder(&f);
9879 b.arg(k_rows).arg(v_rows).arg(kv_ptrs).arg(pos_seq)
9880 .arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
9881 unsafe { b.launch(cfg)?; }
9882 Ok(())
9883 }
9884
9885 pub fn fa_rows_eligible(&self, base_len: usize, head_dim: usize) -> bool {
9891 std::env::var("MEMRA_NO_FA_VEC").is_err()
9892 && std::env::var("MEMRA_FA_ROWS_OFF").is_err()
9893 && base_len + 1 >= fa_vec_min_tkv()
9894 && head_dim <= 256 && head_dim % 32 == 0
9895 }
9896
9897 #[allow(clippy::too_many_arguments)]
9906 pub fn fa_decode_rows(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9907 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9908 head_dim: usize, n_head: usize, n_head_kv: usize,
9909 base_len: usize, t: usize, scale: f32,
9910 k_tok_bytes: usize, v_tok_bytes: usize,
9911 base_dev: Option<(&CudaSlice<i32>, i32)>,
9915 kv_shared: bool,
9918 g: bool,
9922 mut q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
9925 -> Result<(), Box<dyn std::error::Error>> {
9926 debug_assert!(base_len + 1 >= fa_vec_min_tkv() && head_dim <= 512 && head_dim % 32 == 0);
9927 let t_kv_max = base_len + t; let mut sp = fa_split_keys(t_kv_max, n_head_kv); if head_dim == 512 {
9934 static SP512: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9935 let v = *SP512.get_or_init(|| std::env::var("MEMRA_FA_SP512").ok()
9938 .and_then(|x| x.parse().ok()).unwrap_or(0));
9939 sp = if v >= 8 { v } else { FA_SP512_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) };
9940 }
9941 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9942 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9943 let gqa = (n_head / n_head_kv).max(1) as u32;
9944 let mut groups: Vec<(usize, usize, usize)> = Vec::new(); if head_dim == 512 || fa_split_keys(base_len + 1, n_head_kv) == sp {
9955 groups.push((0, t, sp));
9956 } else {
9957 let mut r0 = 0usize;
9958 while r0 < t {
9959 let sp_g = fa_split_keys(base_len + r0 + 1, n_head_kv);
9960 let mut r1 = r0 + 1;
9961 while r1 < t && fa_split_keys(base_len + r1 + 1, n_head_kv) == sp_g { r1 += 1; }
9962 groups.push((r0, r1 - r0, sp_g));
9963 r0 = r1;
9964 }
9965 }
9966 static SMEM_TKV_R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9970 let smem_tkv = *SMEM_TKV_R.get_or_init(|| {
9971 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
9972 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
9973 });
9974 let v4 = fa_v4_at(base_len + t) && head_dim == 256;
9975 let v3 = fa_v3_active(head_dim);
9976 let smem_rows = head_dim <= 256 && !v3 && !fa_v2_on() && smem_tkv > 0 && t_kv_max >= smem_tkv;
9977 let _ = kv_shared;
9982 let i2 = head_dim == 512 && std::env::var("MEMRA_FA_I2").as_deref() != Ok("0");
9985 static TB512: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9999 let tb512 = head_dim == 512 && sp <= 32 && n_head / n_head_kv.max(1) <= 16
10001 && *TB512.get_or_init(|| std::env::var("MEMRA_FA_TB512").as_deref() != Ok("0"));
10002 let fname = if tb512 { "fa_decode_vec_q_rows_v4_512_tb" }
10003 else if i2 { "fa_decode_vec_q_rows_dpl16_i2" }
10004 else if head_dim == 512 { "fa_decode_vec_q_rows_dpl16" } else if v4 { "fa_decode_vec_q_rows_v4" }
10006 else if v3 { "fa_decode_vec_q_rows_v3" }
10007 else if fa_v2_on() { "fa_decode_vec_q_rows_v2" }
10008 else if smem_rows { "fa_decode_vec_q_rows_smem" }
10009 else { "fa_decode_vec_q_rows" };
10010 let f = if head_dim == 512 { self.fa_func(fname, head_dim) }
10011 else if g {
10012 self.func_g(if smem_rows { "fa_decode_vec_q_rows" } else { fname })
10020 }
10021 else { self.func(fname) };
10022 let shmem = if tb512 {
10023 let gk = Self::gkv_on();
10025 let sh = (8192 + 1024 + 32 * 512 + 32 * 64
10026 + 32 * head_dim * if gk { 1 } else { 2 }) as u32;
10027 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10028 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10029 sh
10030 } else if v4 || v3 || smem_rows || fa_v2_on() {
10031 let sh = (if v4 { 11520 + 32 * head_dim * if g { 1 } else { 2 } }
10033 else if v3 { 32 * head_dim * 2 } else { 2 * 32 * head_dim * 2 }) as u32;
10034 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10035 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10036 sh
10037 } else { 0 };
10038 for &(r0, t_g, sp_g) in &groups {
10042 let n_splits_g = (base_len + r0 + t_g).div_ceil(sp_g);
10043 let (nspm, spk) = (n_splits_g as i32, sp_g as i32);
10044 let base_i = (base_len + r0) as i32;
10045 let o_len = t_g * n_head * n_splits_g * head_dim;
10046 let ml_len = t_g * n_head * n_splits_g;
10047 let mut part_guard = self.fa_part_pool.lock().unwrap();
10048 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10049 let old = part_guard.take();
10060 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10061 if let Some(old) = old {
10062 self.fa_part_retired.lock().unwrap().push(old);
10063 }
10064 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10065 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10066 }
10067 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10068 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10069 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10070 }
10071 let pg = part_guard.as_mut().unwrap();
10072 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10073 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10074 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10075 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10076 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
10077 let qv = self.view(q, t * n_head * head_dim);
10078 let q_g = qv.slice(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10079 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_g as u32, t_g as u32),
10080 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10081 {
10082 let __s_b = self.gpu.stream();
10083 let mut b = __s_b.launch_builder(&f);
10084 if tb512 {
10085 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10087 let plus_g = plus + r0 as i32;
10088 let nr = t_g as i32;
10089 if Self::pdl_on() && Self::pdl_wb_on() {
10090 use cudarc::driver::{DevicePtr, DevicePtrMut};
10092 let s = &self.gpu.stream();
10093 let (pq, _b0) = q_g.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10094 let (pv, _b2) = v.device_ptr(s);
10095 let (po, _b3) = part_o.device_ptr_mut(s);
10096 let (pm, _b4) = part_m.device_ptr_mut(s);
10097 let (pl, _b5) = part_l.device_ptr_mut(s);
10098 let (pb, _b6) = bd.device_ptr(s);
10099 let mut ps = [
10100 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10101 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10102 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10103 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10104 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10105 &plus_g as *const _ as *mut _, &scale as *const _ as *mut _,
10106 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10107 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10108 &nr as *const _ as *mut _,
10109 ];
10110 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10111 "fa_decode_vec_q_rows_v4_512_tb",
10112 (n_head_kv as u32, n_splits_g as u32, 1), (32, gqa, 1),
10113 shmem, &mut ps)?; }
10114 } else {
10115 let cfg_tb = LaunchConfig {
10116 grid_dim: (n_head_kv as u32, n_splits_g as u32, 1),
10117 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10118 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10119 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10120 .arg(&ktb).arg(&vtb).arg(&nr);
10121 unsafe { b.launch(cfg_tb)?; }
10122 }
10123 } else if head_dim == 512 {
10124 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10125 let plus_g = plus + r0 as i32;
10126 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10127 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10128 .arg(&ktb).arg(&vtb);
10129 unsafe { b.launch(cfg)?; }
10130 } else {
10131 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10132 .arg(&hd).arg(&nh).arg(&nhkv).arg(&base_i).arg(&scale).arg(&nspm).arg(&spk)
10133 .arg(&ktb).arg(&vtb);
10134 unsafe { b.launch(cfg)?; }
10135 }
10136 }
10137 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t_g as u32, 1),
10138 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10139 let mut o_g = o.slice_mut(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10140 if head_dim == 512 {
10141 let (bd, plus) = base_dev.unwrap();
10144 let plus_g = plus + r0 as i32;
10145 if let Some((oq, od)) = q8_out.as_mut() {
10146 debug_assert!(t == 1, "rows q8 emit is a t=1 decode arm");
10148 if Self::pdl_on() && Self::pdl_wb_on() {
10149 use cudarc::driver::{DevicePtr, DevicePtrMut};
10151 let s = &self.gpu.stream();
10152 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10153 let (pl, _g2) = part_l.device_ptr(s);
10154 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10155 let (pb, _g5) = bd.device_ptr(s);
10156 let mut ps = [
10157 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10158 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10159 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10160 &nh as *const _ as *mut _, &pb as *const _ as *mut _,
10161 &plus_g as *const _ as *mut _, &nspm as *const _ as *mut _,
10162 &spk as *const _ as *mut _,
10163 ];
10164 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10165 "fa_decode_combine_rows_dc_q8_1",
10166 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10167 continue;
10168 }
10169 let fc = self.fa_func("fa_decode_combine_rows_dc_q8_1", head_dim);
10170 let __s_b2 = self.gpu.stream();
10171 let mut b2 = __s_b2.launch_builder(&fc);
10172 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut **oq).arg(&mut **od)
10173 .arg(&hd).arg(&nh).arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10174 unsafe { b2.launch(cfg2)?; }
10175 continue;
10176 }
10177 let fc = self.fa_func("fa_decode_combine_rows_dc", head_dim);
10178 let __s_b2 = self.gpu.stream();
10179 let mut b2 = __s_b2.launch_builder(&fc);
10180 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10181 .arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10182 unsafe { b2.launch(cfg2)?; }
10183 } else {
10184 assert!(q8_out.is_none(), "rows q8 emit requires the hd512 dc combine");
10187 let fc = self.func("fa_decode_combine_rows");
10188 let __s_b2 = self.gpu.stream();
10189 let mut b2 = __s_b2.launch_builder(&fc);
10190 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10191 .arg(&base_i).arg(&nspm).arg(&spk);
10192 unsafe { b2.launch(cfg2)?; }
10193 }
10194 }
10195 Ok(())
10196 }
10197
10198 #[allow(clippy::too_many_arguments)]
10202 pub fn fa_decode_rows_w(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10203 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10204 head_dim: usize, n_head: usize, n_head_kv: usize,
10205 base_dev: &CudaSlice<i32>, base_plus: i32, t: usize, scale: f32,
10206 window: usize, k_tok_bytes: usize, v_tok_bytes: usize,
10207 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10208 -> Result<(), Box<dyn std::error::Error>> {
10209 debug_assert!(head_dim == 256);
10214 let sp = {
10222 static SPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10223 let v = *SPW.get_or_init(|| std::env::var("MEMRA_FA_SPW").ok()
10224 .and_then(|x| x.parse().ok()).unwrap_or(0));
10225 if v >= 8 { v } else { FA_SPW_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) }
10226 };
10227 let n_splits_max = (window + sp - 1) / sp;
10228 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10229 let (nspm, spk, wini) = (n_splits_max as i32, sp as i32, window as i32);
10230 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10231 let gqa = (n_head / n_head_kv).max(1) as u32;
10232 let o_len = t * n_head * n_splits_max * head_dim;
10233 let ml_len = t * n_head * n_splits_max;
10234 let mut part_guard = self.fa_part_pool.lock().unwrap();
10235 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10236 let old = part_guard.take();
10247 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10248 if let Some(old) = old {
10249 self.fa_part_retired.lock().unwrap().push(old);
10250 }
10251 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10252 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10253 }
10254 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10255 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10256 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10257 }
10258 let pg = part_guard.as_mut().unwrap();
10259 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10260 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10261 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10262 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10263 static SMEM_TKV_W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10269 let smem_tkv = *SMEM_TKV_W.get_or_init(|| {
10270 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
10271 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
10272 });
10273 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10279 let wg = Self::wkv_on();
10284 let sp2 = gqa <= 4 && fa_v4_at(window)
10287 && std::env::var("MEMRA_FA_SPW2").as_deref() != Ok("0");
10288 if sp2 {
10289 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10290 if Self::pdl_on() && Self::pdl_wb_on() {
10291 use cudarc::driver::{DevicePtr, DevicePtrMut};
10293 let s = &self.gpu.stream();
10294 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10295 let (pv, _b2) = v.device_ptr(s);
10296 let (po, _b3) = part_o.device_ptr_mut(s);
10297 let (pm, _b4) = part_m.device_ptr_mut(s);
10298 let (pl, _b5) = part_l.device_ptr_mut(s);
10299 let (pb, _b6) = base_dev.device_ptr(s);
10300 let mut ps = [
10301 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10302 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10303 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10304 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10305 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10306 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10307 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10308 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10309 &wini as *const _ as *mut _,
10310 ];
10311 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w_sp",
10312 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa + 1, 1),
10313 sh, &mut ps)?; }
10314 } else {
10315 let f = if wg { self.func_g("fa_decode_vec_q_rows_v4_w_sp") }
10316 else { self.func("fa_decode_vec_q_rows_v4_w_sp") };
10317 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10318 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10319 block_dim: (32, gqa + 1, 1), shared_mem_bytes: sh };
10320 let __s_b = self.gpu.stream();
10321 let mut b = __s_b.launch_builder(&f);
10322 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10323 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10324 .arg(&ktb).arg(&vtb).arg(&wini);
10325 unsafe { b.launch(cfg)?; }
10326 }
10327 } else {
10328 if fa_v4_at(window) && Self::pdl_on() && Self::pdl_wb_on() {
10329 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10331 use cudarc::driver::{DevicePtr, DevicePtrMut};
10332 let s = &self.gpu.stream();
10333 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10334 let (pv, _b2) = v.device_ptr(s);
10335 let (po, _b3) = part_o.device_ptr_mut(s);
10336 let (pm, _b4) = part_m.device_ptr_mut(s);
10337 let (pl, _b5) = part_l.device_ptr_mut(s);
10338 let (pb, _b6) = base_dev.device_ptr(s);
10339 let mut ps = [
10340 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10341 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10342 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10343 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10344 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10345 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10346 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10347 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10348 &wini as *const _ as *mut _,
10349 ];
10350 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w",
10351 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa, 1),
10352 sh, &mut ps)?; }
10353 } else {
10354 let pick = |name: &str| if wg { self.func_g(name) } else { self.func(name) };
10355 let (f, sh) = if fa_v4_at(window) {
10356 let f = pick("fa_decode_vec_q_rows_v4_w");
10357 (f, (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32)
10358 } else if smem_tkv > 0 && window >= smem_tkv {
10359 (pick("fa_decode_vec_q_rows_smem_w"), (2 * 32 * head_dim * 2) as u32)
10362 } else {
10363 (pick("fa_decode_vec_q_rows_reg_w"), 0u32)
10364 };
10365 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10366 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10367 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10368 let __s_b = self.gpu.stream();
10369 let mut b = __s_b.launch_builder(&f);
10370 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10371 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10372 .arg(&ktb).arg(&vtb).arg(&wini);
10373 unsafe { b.launch(cfg)?; }
10374 }
10375 }
10376 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10377 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10378 if let Some((oq, od)) = q8_out {
10379 if Self::pdl_on() && Self::pdl_wb_on() {
10382 use cudarc::driver::{DevicePtr, DevicePtrMut};
10384 let s = &self.gpu.stream();
10385 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10386 let (pl, _g2) = part_l.device_ptr(s);
10387 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10388 let mut ps = [
10389 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10390 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10391 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10392 &nh as *const _ as *mut _, &nspm as *const _ as *mut _,
10393 &spk as *const _ as *mut _, &wini as *const _ as *mut _,
10394 ];
10395 unsafe { self.launch_pdl_flash(wg, "fa_decode_combine_rows_w_q8_1",
10396 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10397 return Ok(());
10398 }
10399 let fc = if wg { self.func_g("fa_decode_combine_rows_w_q8_1") }
10400 else { self.func("fa_decode_combine_rows_w_q8_1") };
10401 let __s_b2 = self.gpu.stream();
10402 let mut b2 = __s_b2.launch_builder(&fc);
10403 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh)
10404 .arg(&nspm).arg(&spk).arg(&wini);
10405 unsafe { b2.launch(cfg2)?; }
10406 return Ok(());
10407 }
10408 let fc = if wg { self.func_g("fa_decode_combine_rows_w") }
10409 else { self.func("fa_decode_combine_rows_w") };
10410 let __s_b2 = self.gpu.stream();
10411 let mut b2 = __s_b2.launch_builder(&fc);
10412 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10413 .arg(&nspm).arg(&spk).arg(&wini);
10414 unsafe { b2.launch(cfg2)?; }
10415 Ok(())
10416 }
10417
10418 #[allow(clippy::too_many_arguments)]
10424 pub fn fa_decode_rows_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10425 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10426 head_dim: usize, n_head: usize, n_head_kv: usize,
10427 base_dev: &CudaSlice<i32>, t_kv_upper: usize, t: usize, scale: f32,
10428 k_tok_bytes: usize, v_tok_bytes: usize, base_plus: i32, g: bool)
10429 -> Result<(), Box<dyn std::error::Error>> {
10430 let v4 = head_dim == 256 && fa_v4_at(t_kv_upper);
10431 assert!(v4 || fa_v3_active(head_dim), "stream fa rows requires the v3 or v4 lane");
10432 assert!(v4 || base_plus == 0, "v3_dc kernel takes no plus arg");
10433 if v4 {
10434 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10435 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10436 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10437 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10438 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10439 let gqa = (n_head / n_head_kv).max(1) as u32;
10440 let o_len = t * n_head * n_splits_max * head_dim;
10441 let ml_len = t * n_head * n_splits_max;
10442 let mut part_guard = self.fa_part_pool.lock().unwrap();
10443 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10444 let old = part_guard.take();
10455 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10456 if let Some(old) = old {
10457 self.fa_part_retired.lock().unwrap().push(old);
10458 }
10459 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10460 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10461 }
10462 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10463 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10464 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10465 }
10466 let pg = part_guard.as_mut().unwrap();
10467 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10468 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10469 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10470 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10471 let f = if g { self.func_g("fa_decode_vec_q_rows_v4_dc") }
10472 else { self.func("fa_decode_vec_q_rows_v4_dc") };
10473 let sh = (11520 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10474 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10475 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10476 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10477 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10478 let __s_b = self.gpu.stream();
10479 let mut b = __s_b.launch_builder(&f);
10480 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10481 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale)
10482 .arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
10483 unsafe { b.launch(cfg)?; }
10484 let fc = self.func("fa_decode_combine_rows_dc");
10485 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10486 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10487 let __s_b2 = self.gpu.stream();
10488 let mut b2 = __s_b2.launch_builder(&fc);
10489 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10490 .arg(base_dev).arg(&base_plus).arg(&nspm).arg(&spk);
10491 unsafe { b2.launch(cfg2)?; }
10492 return Ok(());
10493 }
10494 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10495 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10496 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10497 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10498 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10499 let gqa = (n_head / n_head_kv).max(1) as u32;
10500 let o_len = t * n_head * n_splits_max * head_dim;
10501 let ml_len = t * n_head * n_splits_max;
10502 let mut part_guard = self.fa_part_pool.lock().unwrap();
10503 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10504 let old = part_guard.take();
10515 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10516 if let Some(old) = old {
10517 self.fa_part_retired.lock().unwrap().push(old);
10518 }
10519 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10520 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10521 }
10522 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10523 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10524 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10525 }
10526 let pg = part_guard.as_mut().unwrap();
10527 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10528 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10529 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10530 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10531 let f = self.func("fa_decode_vec_q_rows_v3_dc");
10532 let sh = (32 * head_dim * 2) as u32;
10533 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10534 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10535 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10536 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10537 let __s_b = self.gpu.stream();
10538 let mut b = __s_b.launch_builder(&f);
10539 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10540 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&scale).arg(&nspm).arg(&spk)
10541 .arg(&ktb).arg(&vtb);
10542 unsafe { b.launch(cfg)?; }
10543 let fc = self.func("fa_decode_combine_rows_dc");
10544 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10545 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10546 let plus0 = 0i32;
10547 let __s_b2 = self.gpu.stream();
10548 let mut b2 = __s_b2.launch_builder(&fc);
10549 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10550 .arg(base_dev).arg(&plus0).arg(&nspm).arg(&spk);
10551 unsafe { b2.launch(cfg2)?; }
10552 Ok(())
10553 }
10554
10555 pub fn fa_decode_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10566 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10567 head_dim: usize, n_head: usize, n_head_kv: usize,
10568 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10569 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
10570 -> Result<(), Box<dyn std::error::Error>> {
10571 self.fa_decode_dc_q8(q, k, v, o, head_dim, n_head, n_head_kv, t_kv_dev, bucket_max,
10572 scale, k_tok_bytes, v_tok_bytes, g, None)
10573 }
10574
10575 #[allow(clippy::too_many_arguments)]
10578 pub fn fa_decode_dc_q8(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10579 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10580 head_dim: usize, n_head: usize, n_head_kv: usize,
10581 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10582 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
10583 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10584 -> Result<(), Box<dyn std::error::Error>> {
10585 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
10593 if g && head_dim == 256 && !fa_v4_at(bucket_max) { fa_vec = false; } let sp = fa_split_keys(bucket_max, n_head_kv);
10595 let n_splits = if fa_vec { ((bucket_max + sp - 1) / sp).max(1) } else { ((bucket_max + 255) / 256).max(1) };
10596 let o_len = n_head * n_splits * head_dim;
10597 let ml_len = n_head * n_splits;
10598 let mut part_guard = self.fa_part_pool.lock().unwrap();
10599 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10600 let old = part_guard.take();
10611 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10612 if let Some(old) = old {
10613 self.fa_part_retired.lock().unwrap().push(old);
10614 }
10615 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10616 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10617 }
10618 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10619 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10620 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10621 }
10622 let pg = part_guard.as_mut().unwrap();
10623 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10624 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10625 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10626 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10627 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
10628 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10629 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
10630 let deep = fa_vec && head_dim == 256 && fa_v4_at(bucket_max) && !g
10633 && fa_deep_at(bucket_max) && !matches!(fa_v4_mode(), "noB3" | "stage");
10634 let (f, cfg) = if fa_vec && head_dim == 512 && bucket_max >= {
10635 static FA512_MIN_DC: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10636 *FA512_MIN_DC.get_or_init(|| std::env::var("MEMRA_FA512_MIN").ok()
10637 .and_then(|v| v.parse().ok()).unwrap_or(512))
10638 } {
10639 let gqa = (n_head / n_head_kv).max(1) as u32;
10641 (self.fa_func("fa_decode_vec_q_dpl16_dc", head_dim),
10642 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10643 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10644 } else if fa_vec && head_dim == 512 {
10645 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
10648 0, Some(t_kv_dev), scale, n_splits, sp,
10649 k_tok_bytes, v_tok_bytes, g,
10650 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10651 } else if fa_vec && head_dim == 256 && fa_v4_at(bucket_max) {
10652 let gqa = (n_head / n_head_kv).max(1) as u32;
10655 let fv = if g { self.func_g("fa_decode_vec_q_v4_dc") }
10656 else if deep { self.func("fa_decode_vec_q_v4_deep_dc") }
10657 else { self.func("fa_decode_vec_q_v4_dc") };
10658 let shmem = (if deep { 12160 } else { 11520 }
10659 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10660 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10661 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
10662 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10663 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10664 } else if fa_vec && fa_v3_active(head_dim) {
10665 let gqa = (n_head / n_head_kv).max(1) as u32;
10668 let fv = if g { self.func_g("fa_decode_vec_q_v3_dc") } else { self.func("fa_decode_vec_q_v3_dc") };
10669 let shmem = (32 * head_dim * 2) as u32; (fv,
10671 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10672 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10673 } else if fa_vec && fa_v2_on() {
10674 let gqa = (n_head / n_head_kv).max(1) as u32;
10678 let fv = if g { self.func_g("fa_decode_vec_q_v2_dc") } else { self.func("fa_decode_vec_q_v2_dc") };
10679 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
10681 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10682 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10683 } else if fa_vec {
10684 let gqa = (n_head / n_head_kv).max(1) as u32;
10685 let fv = if g { self.func_g("fa_decode_vec_q_dc") } else { self.func("fa_decode_vec_q_dc") };
10687 (fv,
10688 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10689 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10690 } else {
10691 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
10692 0, Some(t_kv_dev), scale, n_splits,
10693 if fa_vec { sp } else { 256 },
10694 k_tok_bytes, v_tok_bytes, g,
10695 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10696 };
10697 let ski = sp as i32; let __s_b = self.gpu.stream();
10699 let mut b = __s_b.launch_builder(&f);
10700 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10701 .arg(&hd).arg(&nh).arg(&nhkv).arg(t_kv_dev).arg(&scale).arg(&nsp).arg(&ski)
10702 .arg(&ktb).arg(&vtb);
10703 unsafe { b.launch(cfg)?; }
10704 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10705 if let Some((oq, od)) = q8_out {
10706 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
10707 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
10708 let __s_b2 = self.gpu.stream();
10709 let mut b2 = __s_b2.launch_builder(&fc);
10710 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
10711 unsafe { b2.launch(cfg2)?; }
10712 return Ok(());
10713 }
10714 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
10715 let __s_b2 = self.gpu.stream();
10716 let mut b2 = __s_b2.launch_builder(&fc);
10717 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
10718 unsafe { b2.launch(cfg2)?; }
10719 Ok(())
10720 }
10721
10722 pub fn fa_geom_eager(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
10728 let fa_ok = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
10732 let vec512 = fa_ok && head_dim == 512 && t_kv >= fa512_min_tkv();
10738 let mut fa_vec = vec512 || (fa_ok && head_dim <= 256 && head_dim % 32 == 0);
10739 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
10745 let sp = fa_split_keys(t_kv, n_head_kv);
10746 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
10747 (fa_vec, n_splits)
10748 }
10749
10750 pub fn fa_bucket_key(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
10756 self.fa_geom_eager(t_kv, head_dim, n_head_kv, g)
10757 }
10758
10759 pub fn capture_graph_retained<F>(&self, step: F)
10771 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
10772 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
10773 {
10774 use cudarc::driver::sys::CUgraphInstantiate_flags;
10775 self.capture_graph_retained_flags(
10776 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH, step)
10777 }
10778
10779 pub fn capture_graph_retained_flags<F>(&self,
10784 flags: cudarc::driver::sys::CUgraphInstantiate_flags, mut step: F)
10785 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
10786 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
10787 {
10788 use cudarc::driver::sys::CUstreamCaptureMode;
10789 self.capture_keep.lock().unwrap().clear();
10797 let was_tracking = self.gpu.ctx.is_event_tracking();
10798 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
10799 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
10800 self.capture_keep_on.store(true, std::sync::atomic::Ordering::Relaxed);
10801 let w = (|| { step(self)?; step(self) })();
10802 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
10803 w?;
10804 self.gpu.stream().synchronize()?;
10805 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
10806 let r = step(self);
10807 let g = self.gpu.stream().end_capture(flags);
10808 r?;
10809 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
10810 graph.upload()?;
10811 Ok(graph)
10812 };
10813 let result = run();
10814 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
10815 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
10816 let keeper = std::mem::take(&mut *self.capture_keep.lock().unwrap());
10817 Ok((result?, keeper))
10818 }
10819
10820 pub fn capture_graph<F>(&self, mut step: F) -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>>
10821 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
10822 {
10823 use cudarc::driver::sys::{CUstreamCaptureMode, CUgraphInstantiate_flags};
10824 let was_tracking = self.gpu.ctx.is_event_tracking();
10832 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
10833 let iflag = {
10840 static F: std::sync::OnceLock<CUgraphInstantiate_flags> = std::sync::OnceLock::new();
10841 *F.get_or_init(|| match std::env::var("MEMRA_GRAPH_IFLAG").as_deref() {
10842 Ok("upload") => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_UPLOAD,
10845 Ok("priority") =>
10846 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
10847 _ => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
10848 })
10849 };
10850 let ct = {
10857 static T: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10858 *T.get_or_init(|| std::env::var("MEMRA_GRAPH_CAPTIME").as_deref() == Ok("1"))
10859 };
10860 let warmups = {
10883 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10884 *W.get_or_init(|| std::env::var("MEMRA_GRAPH_WARMUPS").ok()
10885 .and_then(|v| v.parse().ok()).filter(|n| *n >= 1).unwrap_or(1))
10886 };
10887 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
10888 let t_w = std::time::Instant::now();
10889 for _ in 0..warmups { step(self)?; }
10891 self.gpu.stream().synchronize()?;
10892 let ms_warm = t_w.elapsed().as_secs_f64() * 1e3;
10893 let t_c = std::time::Instant::now();
10895 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
10896 let r = step(self);
10899 let ms_body = t_c.elapsed().as_secs_f64() * 1e3;
10900 let t_i = std::time::Instant::now();
10901 let g = self.gpu.stream().end_capture(iflag);
10902 let ms_inst = t_i.elapsed().as_secs_f64() * 1e3;
10903 r?;
10904 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
10905 let t_u = std::time::Instant::now();
10906 graph.upload()?;
10907 if ct {
10908 println!("[graph-captime] warmup2x {ms_warm:.2} ms capture-body {ms_body:.2} ms \
10909 instantiate {ms_inst:.2} ms upload {:.2} ms",
10910 t_u.elapsed().as_secs_f64() * 1e3);
10911 }
10912 Ok(graph)
10913 };
10914 let result = run();
10915 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
10916 result
10917 }
10918
10919 pub fn gdn_scan_s128_view(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
10921 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
10922 state_in: &cudarc::driver::CudaView<f32>,
10923 state_out: &mut cudarc::driver::CudaViewMut<f32>,
10924 o: &mut CudaSlice<f32>, n_head: usize, t: usize, scale: f32)
10925 -> Result<(), Box<dyn std::error::Error>> {
10926 let f = self.func("gdn_scan_s128");
10927 const S_V: u32 = 128; const WARP: u32 = 32; const COLS: u32 = 4;
10928 let cfg = LaunchConfig { grid_dim: (n_head as u32, 1, S_V / COLS), block_dim: (WARP, COLS, 1), shared_mem_bytes: 0 };
10929 let (h, ti) = (n_head as i32, t as i32);
10930 let __s_b = self.gpu.stream();
10931 let mut b = __s_b.launch_builder(&f);
10932 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);
10933 unsafe { b.launch(cfg)?; }
10934 Ok(())
10935 }
10936
10937 pub fn ssm_conv1d_view(&self, x: &cudarc::driver::CudaView<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
10939 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
10940 -> Result<(), Box<dyn std::error::Error>> {
10941 let f = self.func("ssm_conv1d_silu_f32");
10942 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
10944 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
10945 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
10946 let __s_b = self.gpu.stream();
10947 let mut b = __s_b.launch_builder(&f);
10948 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
10949 unsafe { b.launch(cfg)?; }
10950 Ok(())
10951 }
10952
10953 pub fn ssm_conv1d_tm(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
10960 conv_dim: usize, t: usize, d_conv: usize)
10961 -> Result<(), Box<dyn std::error::Error>> {
10962 let f = self.func("ssm_conv1d_tm_f32");
10963 let cfg = LaunchConfig {
10964 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
10965 block_dim: (256, 1, 1), shared_mem_bytes: 0,
10966 };
10967 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
10968 let __s_b = self.gpu.stream();
10969 let mut b = __s_b.launch_builder(&f);
10970 b.arg(qkv_tm).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
10971 unsafe { b.launch(cfg)?; }
10972 Ok(())
10973 }
10974
10975 pub fn ssm_conv1d_tm_state(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
10983 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
10984 conv_dim: usize, t: usize, d_conv: usize)
10985 -> Result<(), Box<dyn std::error::Error>> {
10986 self.ssm_conv1d_tm_state_pad(qkv_tm, conv_state, w, y, conv_dim, t, d_conv, None)
10987 }
10988
10989 #[allow(clippy::too_many_arguments)]
10992 pub fn ssm_conv1d_tm_state_pad(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
10993 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
10994 conv_dim: usize, t: usize, d_conv: usize,
10995 pad_len: Option<&CudaSlice<i32>>)
10996 -> Result<(), Box<dyn std::error::Error>> {
10997 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
10998 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11002 {
11003 let f = self.func("ssm_conv1d_tm_state_f32");
11004 let cfg = LaunchConfig {
11005 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11006 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11007 };
11008 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11009 let __s_b = self.gpu.stream();
11010 let mut b = __s_b.launch_builder(&f);
11011 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11012 unsafe { b.launch(cfg)?; }
11013 }
11014 match (ring_old, pad_len) {
11015 (None, Some(len_d)) => {
11016 let f = self.func("ssm_conv_ring_update_dev_f32");
11017 let n = conv_dim * (d_conv - 1);
11018 let cfg = LaunchConfig::for_num_elems(n as u32);
11019 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11020 let __s_b = self.gpu.stream();
11021 let mut b = __s_b.launch_builder(&f);
11022 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11023 unsafe { b.launch(cfg)?; }
11024 }
11025 (None, None) => {
11026 let f = self.func("ssm_conv_ring_update_f32");
11027 let n = conv_dim * (d_conv - 1);
11028 let cfg = LaunchConfig::for_num_elems(n as u32);
11029 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11030 let __s_b = self.gpu.stream();
11031 let mut b = __s_b.launch_builder(&f);
11032 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11033 unsafe { b.launch(cfg)?; }
11034 }
11035 (Some(old), _) => self.ssm_conv_ring_rebuild(qkv_tm, &old, conv_state, conv_dim, t, d_conv)?,
11036 }
11037 Ok(())
11038 }
11039
11040 pub fn ssm_conv1d_tm_state_pad_v(&self, qkv_tm: &cudarc::driver::CudaView<f32>, conv_state: &mut CudaSlice<f32>,
11042 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11043 conv_dim: usize, t: usize, d_conv: usize,
11044 pad_len: Option<&CudaSlice<i32>>)
11045 -> Result<(), Box<dyn std::error::Error>> {
11046 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
11047 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11051 {
11052 let f = self.func("ssm_conv1d_tm_state_f32");
11053 let cfg = LaunchConfig {
11054 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11055 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11056 };
11057 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11058 let __s_b = self.gpu.stream();
11059 let mut b = __s_b.launch_builder(&f);
11060 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11061 unsafe { b.launch(cfg)?; }
11062 }
11063 match (ring_old, pad_len) {
11064 (None, Some(len_d)) => {
11065 let f = self.func("ssm_conv_ring_update_dev_f32");
11066 let n = conv_dim * (d_conv - 1);
11067 let cfg = LaunchConfig::for_num_elems(n as u32);
11068 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11069 let __s_b = self.gpu.stream();
11070 let mut b = __s_b.launch_builder(&f);
11071 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11072 unsafe { b.launch(cfg)?; }
11073 }
11074 (None, None) => {
11075 let f = self.func("ssm_conv_ring_update_f32");
11076 let n = conv_dim * (d_conv - 1);
11077 let cfg = LaunchConfig::for_num_elems(n as u32);
11078 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11079 let __s_b = self.gpu.stream();
11080 let mut b = __s_b.launch_builder(&f);
11081 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11082 unsafe { b.launch(cfg)?; }
11083 }
11084 (Some(_), _) => unreachable!(
11085 "ssm_conv1d_tm_state_pad_v: T < d_conv-1 has no view path (PRIME_MIN_T gates it)"),
11086 }
11087 Ok(())
11088 }
11089
11090 pub fn ssm_conv_ring_rebuild(&self, qkv_tm: &CudaSlice<f32>, ring_old: &CudaSlice<f32>,
11095 conv_state: &mut CudaSlice<f32>,
11096 conv_dim: usize, tc: usize, d_conv: usize)
11097 -> Result<(), Box<dyn std::error::Error>> {
11098 let f = self.func("ssm_conv_ring_rebuild_f32");
11099 let n = conv_dim * (d_conv - 1);
11100 let cfg = LaunchConfig::for_num_elems(n as u32);
11101 let (cd, ti, dc) = (conv_dim as i32, tc as i32, d_conv as i32);
11102 let __s_b = self.gpu.stream();
11103 let mut b = __s_b.launch_builder(&f);
11104 b.arg(qkv_tm).arg(ring_old).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11105 unsafe { b.launch(cfg)?; }
11106 Ok(())
11107 }
11108
11109 #[allow(clippy::too_many_arguments)]
11114 pub fn gdn_prep_decode(&self, conv_out: &CudaSlice<f32>, beta_raw: &CudaSlice<f32>,
11115 alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11116 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11117 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11118 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32)
11119 -> Result<(), Box<dyn std::error::Error>> {
11120 let f = self.func("gdn_prep_decode_f32");
11121 let cfg = LaunchConfig { grid_dim: (num_v as u32, 1, 1), block_dim: (32, 4, 1), shared_mem_bytes: 0 };
11122 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11123 let __s_b = self.gpu.stream();
11124 let mut b = __s_b.launch_builder(&f);
11125 b.arg(conv_out).arg(beta_raw).arg(alpha).arg(dt_bias).arg(a)
11126 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11127 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps);
11128 unsafe { b.launch(cfg)?; }
11129 Ok(())
11130 }
11131
11132 #[allow(clippy::too_many_arguments)]
11136 pub fn ssm_conv1d_gdn(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>,
11137 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11138 conv_dim: usize, t: usize, d_conv: usize,
11139 d_state: usize, num_v: usize, num_k: usize, key_dim: usize)
11140 -> Result<(), Box<dyn std::error::Error>> {
11141 let f = self.func("ssm_conv1d_gdn_f32");
11142 let cfg = LaunchConfig {
11143 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11144 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11145 };
11146 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11147 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11148 let __s_b = self.gpu.stream();
11149 let mut b = __s_b.launch_builder(&f);
11150 b.arg(qkv_tm).arg(w).arg(q_g).arg(k_g).arg(v_g)
11151 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd);
11152 unsafe { b.launch(cfg)?; }
11153 Ok(())
11154 }
11155
11156 pub fn ssm_conv1d(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11157 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
11158 -> Result<(), Box<dyn std::error::Error>> {
11159 let f = self.func("ssm_conv1d_silu_f32");
11160 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
11161 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11162 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
11163 let __s_b = self.gpu.stream();
11164 let mut b = __s_b.launch_builder(&f);
11165 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
11166 unsafe { b.launch(cfg)?; }
11167 Ok(())
11168 }
11169
11170 pub fn gdn_scan_s128(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11173 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
11174 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11175 n_head: usize, t: usize, scale: f32)
11176 -> Result<(), Box<dyn std::error::Error>> {
11177 let f = self.func("gdn_scan_s128");
11178 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11179 let cfg = LaunchConfig {
11180 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
11181 block_dim: (WARP, COLS_PER_BLOCK, 1),
11182 shared_mem_bytes: 0,
11183 };
11184 let (h, ti) = (n_head as i32, t as i32);
11185 let __s_b = self.gpu.stream();
11186 let mut b = __s_b.launch_builder(&f);
11187 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);
11188 unsafe { b.launch(cfg)?; }
11189 Ok(())
11190 }
11191
11192 #[allow(clippy::too_many_arguments)]
11197 pub fn ssm_conv1d_fused_decode_b(
11198 &self, qkv_cols: &CudaSlice<f32>, conv_state_ptrs: &cudarc::driver::CudaView<u64>,
11199 w: &CudaSlice<f32>, conv_outs: &mut CudaSlice<f32>, conv_dim: usize, d_conv: usize,
11200 b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11201 let f = self.func("ssm_conv1d_fused_decode_b_f32");
11202 let cfg = LaunchConfig {
11203 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
11204 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11205 };
11206 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11207 let __s_b = self.gpu.stream();
11208 let mut b = __s_b.launch_builder(&f);
11209 b.arg(qkv_cols).arg(conv_state_ptrs).arg(w).arg(conv_outs).arg(&cd).arg(&dc);
11210 unsafe { b.launch(cfg)?; }
11211 Ok(())
11212 }
11213
11214 #[allow(clippy::too_many_arguments)]
11215 pub fn gdn_prep_decode_b(
11216 &self, conv_outs: &CudaSlice<f32>, beta_raws: &CudaSlice<f32>, alphas: &CudaSlice<f32>,
11217 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11218 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11219 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11220 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32,
11221 conv_dim: usize, b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11222 let f = self.func("gdn_prep_decode_b_f32");
11223 let cfg = LaunchConfig {
11224 grid_dim: (num_v as u32, 1, b_n as u32),
11225 block_dim: (32, 4, 1), shared_mem_bytes: 0,
11226 };
11227 let (ds, nv, nk, kd, cd) =
11228 (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, conv_dim as i32);
11229 let __s_b = self.gpu.stream();
11230 let mut b = __s_b.launch_builder(&f);
11231 b.arg(conv_outs).arg(beta_raws).arg(alphas).arg(dt_bias).arg(a)
11232 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11233 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps).arg(&cd);
11234 unsafe { b.launch(cfg)?; }
11235 Ok(())
11236 }
11237
11238 #[allow(clippy::too_many_arguments)]
11239 pub fn gdn_scan_s128_batched(
11240 &self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11241 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
11242 state_in_ptrs: &cudarc::driver::CudaView<u64>,
11243 state_out_ptrs: &cudarc::driver::CudaView<u64>,
11244 o: &mut CudaSlice<f32>, n_head: usize, b_n: usize, scale: f32)
11245 -> Result<(), Box<dyn std::error::Error>> {
11246 let f = self.func("gdn_scan_s128_b");
11247 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11248 let cfg = LaunchConfig {
11249 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
11250 block_dim: (WARP, COLS_PER_BLOCK, 1), shared_mem_bytes: 0,
11251 };
11252 let h = n_head as i32;
11253 let __s_b = self.gpu.stream();
11254 let mut b = __s_b.launch_builder(&f);
11255 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in_ptrs).arg(state_out_ptrs)
11256 .arg(o).arg(&h).arg(&scale);
11257 unsafe { b.launch(cfg)?; }
11258 Ok(())
11259 }
11260
11261 pub fn gdn_chunked_enabled() -> bool {
11270 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11271 *E.get_or_init(|| std::env::var("MEMRA_GDN_CHUNKED").map(|v| v != "0").unwrap_or(true))
11272 }
11273
11274 pub fn gdn_chunk_size() -> usize {
11279 static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
11280 *C.get_or_init(|| {
11281 let c: usize = std::env::var("MEMRA_GDN_CHUNK").ok()
11282 .and_then(|v| v.parse().ok()).unwrap_or(32);
11283 c.clamp(32, 128) / 32 * 32
11284 })
11285 }
11286
11287 #[allow(clippy::too_many_arguments)]
11292 #[allow(clippy::too_many_arguments, clippy::type_complexity)]
11295 #[allow(clippy::too_many_arguments)]
11296 pub fn gdn_chunk_k123(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11297 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, wb16: Option<&mut CudaSlice<u8>>,
11298 n_head: usize, t: usize, c: usize, hk: usize,
11299 k2w: Option<(&CudaSlice<u8>, &CudaSlice<u8>, &mut CudaSlice<u8>)>)
11300 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11301 const D: usize = 128;
11302 let h = n_head;
11303 let nc = (t + c - 1) / c;
11304 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
11305 let mut gcum = self.uninit(t * h)?;
11306 let mut a = self.uninit(nc * h * c * c)?;
11307 let mut p = self.uninit(nc * h * c * c)?;
11308 let mut u = self.uninit(nc * h * c * D)?;
11309 let mut w = self.uninit(nc * h * c * D)?;
11310 { let f = self.func("gdn_chunk_cumgate_f32");
11312 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11313 let __s_b = self.gpu.stream();
11314 let mut b = __s_b.launch_builder(&f);
11315 b.arg(g).arg(&mut gcum).arg(&hi).arg(&ti).arg(&ci);
11316 unsafe { b.launch(cfg)?; }
11317 }
11318 if let Some((qb, kb, pb)) = k2w {
11319 assert!(c == 32, "gdn_k2_wgmma is a C==32 tile");
11322 let f = self.func("gdn_k2_wgmma");
11323 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11324 let hki = hk as i32;
11325 let __s_b = self.gpu.stream();
11326 let mut b = __s_b.launch_builder(&f);
11327 b.arg(qb).arg(kb).arg(&gcum).arg(beta).arg(&mut a).arg(&mut *pb).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11328 unsafe { b.launch(cfg)?; }
11329 } else if c <= 64 && !portable_mma_gated() { let f = self.func("gdn_chunk_attn_f32");
11331 let jt = ((c + 31) / 32) as u32;
11332 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11333 let hki = hk as i32;
11334 let __s_b = self.gpu.stream();
11335 let mut b = __s_b.launch_builder(&f);
11336 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11337 unsafe { b.launch(cfg)?; }
11338 } else { assert!(hk == h, "generic K2 is broadcast-only (de-broadcast rides C==32)");
11340 let f = self.func("gdn_chunk_attn_g_f32");
11341 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 8, 1), shared_mem_bytes: 0 };
11342 let __s_b = self.gpu.stream();
11343 let mut b = __s_b.launch_builder(&f);
11344 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci);
11345 unsafe { b.launch(cfg)?; }
11346 }
11347 { let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11349 match c {
11350 32 | 64 => {
11351 let f = self.func(if c == 32 { "gdn_chunk_solve32_f32" } else { "gdn_chunk_solve64_f32" });
11352 let wb: u64 = match wb16 { Some(d) => self.addr_u8(d), None => 0 };
11354 let hki = hk as i32;
11355 let __s_b = self.gpu.stream();
11356 let mut b = __s_b.launch_builder(&f);
11357 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&wb).arg(&hi).arg(&ti).arg(&hki);
11358 unsafe { b.launch(cfg)?; }
11359 }
11360 _ => {
11361 assert!(hk == h, "generic K3 is broadcast-only");
11362 let f = self.func("gdn_chunk_solve_f32");
11363 let __s_b = self.gpu.stream();
11364 let mut b = __s_b.launch_builder(&f);
11365 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&hi).arg(&ti).arg(&ci);
11366 unsafe { b.launch(cfg)?; }
11367 }
11368 }
11369 }
11370 Ok((gcum, p, u, w))
11371 }
11372
11373 pub fn gdn_db_on() -> bool {
11377 std::env::var("MEMRA_GDN_DB").as_deref() != Ok("0")
11378 }
11379
11380 pub fn gdn_mma_enabled(&self, c: usize) -> bool {
11383 !portable_mma_gated() && c == 32
11384 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11385 Ok("1") => true,
11386 Ok("0") => false,
11387 _ => cfg!(memra_hopper_mma),
11388 }
11389 }
11390
11391 pub fn gdn_wgmma_on(&self, c: usize) -> bool {
11394 self.gdn_mma_enabled(c)
11395 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
11396 Ok("0") => false,
11397 Ok("1") => true,
11398 _ => cfg!(memra_hopper_mma),
11399 }
11400 }
11401
11402 #[allow(clippy::too_many_arguments)]
11407 pub fn ssm_conv1d_gdn_state_pad(&self, qkv_tm: &cudarc::driver::CudaView<f32>,
11408 conv_state: &mut CudaSlice<f32>, w: &CudaSlice<f32>,
11409 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>,
11410 v_g: &mut CudaSlice<f32>,
11411 conv_dim: usize, t: usize, d_conv: usize,
11412 d_state: usize, num_v: usize, num_k: usize, key_dim: usize,
11413 hk: usize,
11414 pad_len: Option<&CudaSlice<i32>>)
11415 -> Result<(), Box<dyn std::error::Error>> {
11416 assert!(t >= d_conv - 1, "fused state conv requires T >= pad (PRIME_MIN_T gates)");
11417 {
11418 let f = self.func("ssm_conv1d_gdn_state_f32");
11419 let cfg = LaunchConfig {
11420 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11421 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11422 };
11423 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11424 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);
11425 let __s_b = self.gpu.stream();
11426 let mut b = __s_b.launch_builder(&f);
11427 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(q_g).arg(k_g).arg(v_g)
11428 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&hki);
11429 unsafe { b.launch(cfg)?; }
11430 }
11431 match pad_len {
11432 Some(len_d) => {
11433 let f = self.func("ssm_conv_ring_update_dev_f32");
11434 let n = conv_dim * (d_conv - 1);
11435 let cfg = LaunchConfig::for_num_elems(n as u32);
11436 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11437 let __s_b = self.gpu.stream();
11438 let mut b = __s_b.launch_builder(&f);
11439 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11440 unsafe { b.launch(cfg)?; }
11441 }
11442 None => {
11443 let f = self.func("ssm_conv_ring_update_f32");
11444 let n = conv_dim * (d_conv - 1);
11445 let cfg = LaunchConfig::for_num_elems(n as u32);
11446 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11447 let __s_b = self.gpu.stream();
11448 let mut b = __s_b.launch_builder(&f);
11449 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11450 unsafe { b.launch(cfg)?; }
11451 }
11452 }
11453 Ok(())
11454 }
11455
11456 pub fn gdn_chunk_alloc(&self, n_head: usize, t: usize, c: usize, hk: usize)
11460 -> Result<GdnChunkBufs, Box<dyn std::error::Error>> {
11461 const D: usize = 128;
11462 assert!(c == 32, "gdn_chunk_alloc: varlen chain is the C==32 mma pair");
11463 let h = n_head;
11464 let nc = (t + c - 1) / c;
11465 Ok(GdnChunkBufs {
11466 gcum: self.uninit(t * h)?,
11467 a: self.uninit(nc * h * c * c)?,
11468 p: self.uninit(nc * h * c * c)?,
11469 u: self.uninit(nc * h * c * D)?,
11470 w: self.uninit(nc * h * c * D)?,
11471 kb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11472 wb16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11473 y16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11474 ssnap16: self.alloc_u8_uninit(nc * h * D * D * 2)?,
11475 qb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11476 pb16: self.alloc_u8_uninit(nc * h * c * c * 2)?,
11477 o: self.uninit(D * h * t)?,
11478 t, nc,
11479 })
11480 }
11481
11482 pub fn f32_to_bf16_v(&self, x: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<u8>, n: usize)
11484 -> Result<(), Box<dyn std::error::Error>> {
11485 let f = self.func("f32_to_bf16_bulk");
11486 let ni = n as i64;
11487 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11488 let __s_b = self.gpu.stream();
11489 let mut b = __s_b.launch_builder(&f);
11490 b.arg(x).arg(dst).arg(&ni);
11491 unsafe { b.launch(cfg)?; }
11492 Ok(())
11493 }
11494
11495 pub fn f32_to_bf16_into(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<u8>, n: usize)
11497 -> Result<(), Box<dyn std::error::Error>> {
11498 let f = self.func("f32_to_bf16_bulk");
11499 let ni = n as i64;
11500 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11501 let __s_b = self.gpu.stream();
11502 let mut b = __s_b.launch_builder(&f);
11503 b.arg(x).arg(dst).arg(&ni);
11504 unsafe { b.launch(cfg)?; }
11505 Ok(())
11506 }
11507
11508 pub fn gdn_chunk_k123_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, hk: usize,
11511 wq: Option<&GdnWVl8>)
11512 -> Result<(), Box<dyn std::error::Error>> {
11513 let b = seqs.len();
11514 assert!(b >= 1 && b <= 8, "gdn_chunk_k123_vl8: 1..=8 sequences");
11515 let mut packed = [GdnSeqVl::default(); 8];
11516 packed[..b].copy_from_slice(seqs);
11517 let v = GdnVl8(packed);
11518 let (hi, ci) = (n_head as i32, 32i32);
11519 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11520 {
11521 let f = self.func("gdn_chunk_cumgate_vl");
11522 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11523 let __s_lb = self.gpu.stream();
11524 let mut lb = __s_lb.launch_builder(&f);
11525 lb.arg(&v).arg(&hi).arg(&ci);
11526 unsafe { lb.launch(cfg)?; }
11527 }
11528 let hki = hk as i32;
11529 if let Some(w) = wq { let f = self.func("gdn_k2_wgmma_vl");
11531 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11532 let __s_lb = self.gpu.stream();
11533 let mut lb = __s_lb.launch_builder(&f);
11534 lb.arg(&v).arg(w).arg(&hi).arg(&ci).arg(&hki);
11535 unsafe { lb.launch(cfg)?; }
11536 } else {
11537 let f = self.func("gdn_chunk_attn_vl");
11538 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11539 let __s_lb = self.gpu.stream();
11540 let mut lb = __s_lb.launch_builder(&f);
11541 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11542 unsafe { lb.launch(cfg)?; }
11543 }
11544 {
11545 let f = self.func("gdn_chunk_solve32_vl");
11546 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11547 let __s_lb = self.gpu.stream();
11548 let mut lb = __s_lb.launch_builder(&f);
11549 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11550 unsafe { lb.launch(cfg)?; }
11551 }
11552 Ok(())
11553 }
11554
11555 #[allow(clippy::too_many_arguments)]
11559 pub fn gdn_prep_vl8(&self, seqs: &[GdnPrepVl], conv_w: &CudaSlice<f32>,
11560 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11561 conv_dim: usize, d_conv: usize, d_state: usize,
11562 num_v: usize, num_k: usize, key_dim: usize, hk: usize, eps: f32)
11563 -> Result<(), Box<dyn std::error::Error>> {
11564 let b = seqs.len();
11565 assert!(b >= 1 && b <= 8);
11566 let mut packed = [GdnPrepVl::default(); 8];
11567 packed[..b].copy_from_slice(seqs);
11568 let v = GdnPrepVl8(packed);
11569 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11570 let (cdi, dci) = (conv_dim as i32, d_conv as i32);
11571 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
11572 assert!(conv_fuse || hk == num_v, "de-broadcast requires the fused conv");
11573 if conv_fuse {
11574 let f = self.func("ssm_conv1d_gdn_state_vl");
11575 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 };
11576 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);
11577 let __s_lb = self.gpu.stream();
11578 let mut lb = __s_lb.launch_builder(&f);
11579 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi).arg(&hki);
11580 unsafe { lb.launch(cfg)?; }
11581 } else {
11582 let f = self.func("ssm_conv1d_tm_state_vl");
11583 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 };
11584 let __s_lb = self.gpu.stream();
11585 let mut lb = __s_lb.launch_builder(&f);
11586 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci);
11587 unsafe { lb.launch(cfg)?; }
11588 }
11589 {
11590 let f = self.func("ssm_conv_ring_update_vl");
11591 let n = (conv_dim * (d_conv - 1)) as u32;
11592 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11593 let __s_lb = self.gpu.stream();
11594 let mut lb = __s_lb.launch_builder(&f);
11595 lb.arg(&v).arg(&cdi).arg(&dci);
11596 unsafe { lb.launch(cfg)?; }
11597 }
11598 if !conv_fuse {
11599 let f = self.func("qkv_to_gdn_repack_vl");
11600 let n = max_t * (num_v * d_state) as u32;
11601 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11602 let (dsi, nvi, nki, kdi) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11603 let __s_lb = self.gpu.stream();
11604 let mut lb = __s_lb.launch_builder(&f);
11605 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi);
11606 unsafe { lb.launch(cfg)?; }
11607 }
11608 if Self::l2_v2_on(d_state) {
11609 let f = self.func("gdn_l2_v2_vl");
11610 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 };
11611 let (dsi, nvi) = (d_state as i32, hk as i32);
11612 let __s_lb = self.gpu.stream();
11613 let mut lb = __s_lb.launch_builder(&f);
11614 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11615 unsafe { lb.launch(cfg)?; }
11616 } else {
11617 let f = self.func("gdn_l2_vl");
11618 let cfg = LaunchConfig { grid_dim: (max_t * hk as u32, 2, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11619 let (dsi, nvi) = (d_state as i32, hk as i32);
11620 let __s_lb = self.gpu.stream();
11621 let mut lb = __s_lb.launch_builder(&f);
11622 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11623 unsafe { lb.launch(cfg)?; }
11624 }
11625 {
11626 let f = self.func("gdn_gate_prep_vl");
11627 let n = max_t * num_v as u32;
11628 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11629 let nvi = num_v as i32;
11630 let __s_lb = self.gpu.stream();
11631 let mut lb = __s_lb.launch_builder(&f);
11632 lb.arg(&v).arg(dt_bias).arg(a).arg(&nvi);
11633 unsafe { lb.launch(cfg)?; }
11634 }
11635 Ok(())
11636 }
11637
11638 pub fn gdn_mirror_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, which: i32, hk: usize)
11640 -> Result<(), Box<dyn std::error::Error>> {
11641 let b = seqs.len();
11642 assert!(b >= 1 && b <= 8);
11643 let mut packed = [GdnSeqVl::default(); 8];
11644 packed[..b].copy_from_slice(seqs);
11645 let v = GdnVl8(packed);
11646 let ept = (if which == 0 { hk } else { n_head } * 128) as i32;
11647 let max_n = seqs.iter().map(|s| if which == 0 { s.t as i64 * ept as i64 }
11648 else { s.nc as i64 * ept as i64 * 32 }).max().unwrap();
11649 let f = self.func("gdn_mirror_vl");
11650 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
11651 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11652 let __s_lb = self.gpu.stream();
11653 let mut lb = __s_lb.launch_builder(&f);
11654 lb.arg(&v).arg(&ept).arg(&which);
11655 unsafe { lb.launch(cfg)?; }
11656 Ok(())
11657 }
11658
11659 pub fn gdn_tail_vl8(&self, seqs: &[GdnPrepVl], norm_w: &CudaSlice<f32>,
11661 d_state: usize, num_v: usize, eps: f32)
11662 -> Result<(), Box<dyn std::error::Error>> {
11663 let b = seqs.len();
11664 assert!(b >= 1 && b <= 8);
11665 let mut packed = [GdnPrepVl::default(); 8];
11666 packed[..b].copy_from_slice(seqs);
11667 let v = GdnPrepVl8(packed);
11668 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11669 let f = self.func("gated_rmsnorm_f16out_vl");
11670 let cfg = LaunchConfig { grid_dim: (max_t * num_v as u32, 1, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11672 let (dsi, nvi) = (d_state as i32, num_v as i32);
11673 let __s_lb = self.gpu.stream();
11674 let mut lb = __s_lb.launch_builder(&f);
11675 lb.arg(&v).arg(norm_w).arg(&dsi).arg(&nvi).arg(&eps);
11676 unsafe { lb.launch(cfg)?; }
11677 Ok(())
11678 }
11679
11680 pub fn addr_f32(&self, x: &CudaSlice<f32>) -> u64 {
11683 use cudarc::driver::DevicePtr;
11684 let s = self.gpu.stream();
11685 let (p, _g) = x.device_ptr(&s);
11686 p as u64
11687 }
11688 pub fn addr_f32_mut(&self, x: &mut CudaSlice<f32>) -> u64 {
11689 use cudarc::driver::DevicePtrMut;
11690 let s = self.gpu.stream();
11691 let (p, _g) = x.device_ptr_mut(&s);
11692 p as u64
11693 }
11694 pub fn addr_f32v(&self, x: &cudarc::driver::CudaView<f32>) -> u64 {
11695 use cudarc::driver::DevicePtr;
11696 let s = self.gpu.stream();
11697 let (p, _g) = x.device_ptr(&s);
11698 p as u64
11699 }
11700 pub fn addr_u8(&self, x: &CudaSlice<u8>) -> u64 {
11701 use cudarc::driver::DevicePtr;
11702 let s = self.gpu.stream();
11703 let (p, _g) = x.device_ptr(&s);
11704 p as u64
11705 }
11706
11707 pub fn gdn_chunk_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, scale: f32, hk: usize,
11711 wq: Option<&GdnWVl8>)
11712 -> Result<(), Box<dyn std::error::Error>> {
11713 const NSPLIT: u32 = 4;
11714 let b = seqs.len();
11715 assert!(b >= 1 && b <= 8, "gdn_chunk_vl8: 1..=8 sequences");
11716 let mut packed = [GdnSeqVl::default(); 8];
11717 packed[..b].copy_from_slice(seqs);
11718 let v = GdnVl8(packed);
11719 let (hi, ci) = (n_head as i32, 32i32);
11720 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11721 let hki = hk as i32;
11722 if let Some(w) = wq {
11723 let f = self.func("gdn_k45_wgmma_vl");
11725 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11726 let __s_lb = self.gpu.stream();
11727 let mut lb = __s_lb.launch_builder(&f);
11728 lb.arg(&v).arg(w).arg(&scale).arg(&hi).arg(&ci).arg(&hki);
11729 unsafe { lb.launch(cfg)?; }
11730 let _ = max_nc;
11731 return Ok(());
11732 }
11733 {
11734 let f = self.func("gdn_chunk_state_mma_vl");
11735 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11736 let __s_lb = self.gpu.stream();
11737 let mut lb = __s_lb.launch_builder(&f);
11738 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11739 unsafe { lb.launch(cfg)?; }
11740 }
11741 {
11742 let f = self.func("gdn_chunk_output_mma_vl");
11743 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11744 let __s_lb = self.gpu.stream();
11745 let mut lb = __s_lb.launch_builder(&f);
11746 lb.arg(&v).arg(&hi).arg(&ci).arg(&scale).arg(&hki);
11747 unsafe { lb.launch(cfg)?; }
11748 }
11749 Ok(())
11750 }
11751 pub fn gdn_scan_chunked(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11752 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
11753 qb16_pre: Option<&CudaSlice<u8>>,
11754 state_in: &CudaSlice<f32>,
11755 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11756 n_head: usize, t: usize, scale: f32, c: usize, hk: usize)
11757 -> Result<(), Box<dyn std::error::Error>> {
11758 const D: usize = 128;
11759 const NSPLIT: u32 = 4;
11760 assert!(c >= 1 && c <= 128, "gdn_scan_chunked: C must be in 1..=128");
11761 let h = n_head;
11762 let nc = (t + c - 1) / c;
11763 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
11764 let gdn_mma_pre = !portable_mma_gated() && c == 32
11768 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11769 Ok("1") => true,
11770 Ok("0") => false,
11771 _ => cfg!(memra_hopper_mma),
11772 };
11773 let mut wb16_pre: Option<CudaSlice<u8>> = if gdn_mma_pre {
11774 Some(self.alloc_u8_uninit(nc * h * c * D * 2)?)
11775 } else { None };
11776 let gdn_wgmma_pre = gdn_mma_pre
11780 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
11781 Ok("0") => false,
11782 Ok("1") => true,
11783 _ => cfg!(memra_hopper_mma),
11784 };
11785 let nk = t * hk * D;
11786 let mut kb16_local: Option<CudaSlice<u8>> = None;
11787 if gdn_mma_pre && kb16_pre.is_none() {
11788 let mut kb = self.alloc_u8_uninit(nk * 2)?;
11789 let f = self.func("f32_to_bf16_bulk");
11790 let n2 = nk as i64;
11791 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
11792 let __s_b = self.gpu.stream();
11793 let mut b = __s_b.launch_builder(&f);
11794 b.arg(k).arg(&mut kb).arg(&n2);
11795 unsafe { b.launch(cfg2)?; }
11796 kb16_local = Some(kb);
11797 }
11798 let kb16_ref0: Option<&CudaSlice<u8>> = kb16_local.as_ref().or(kb16_pre);
11799 if let Some(kb) = kb16_pre { assert!(kb.len() >= nk * 2, "kb16_pre too small"); }
11800 let mut qb16: Option<CudaSlice<u8>> = None;
11801 let mut pb16: Option<CudaSlice<u8>> = None;
11802 if gdn_wgmma_pre {
11803 if qb16_pre.is_none() {
11806 let mut qb = self.alloc_u8_uninit(nk * 2)?;
11807 let f = self.func("f32_to_bf16_bulk");
11808 let n2 = nk as i64;
11809 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
11810 let __s_b = self.gpu.stream();
11811 let mut b = __s_b.launch_builder(&f);
11812 b.arg(q).arg(&mut qb).arg(&n2);
11813 unsafe { b.launch(cfg2)?; }
11814 qb16 = Some(qb);
11815 } else if let Some(qb) = qb16_pre {
11816 assert!(qb.len() >= nk * 2, "qb16_pre too small");
11817 }
11818 pb16 = Some(self.alloc_u8_uninit(nc * h * c * c * 2)?);
11819 }
11820 let qb16_ref0: Option<&CudaSlice<u8>> = qb16.as_ref().or(qb16_pre);
11821 let k2w = if gdn_wgmma_pre {
11822 Some((*qb16_ref0.as_ref().unwrap(),
11823 *kb16_ref0.as_ref().unwrap(),
11824 pb16.as_mut().unwrap()))
11825 } else { None };
11826 let (gcum, p, u, w) = self.gdn_chunk_k123(q, k, v, g, beta, wb16_pre.as_mut(), n_head, t, c, hk, k2w)?;
11827 let _ = &w;
11828 let mut y = self.uninit(nc * h * c * D)?;
11829 let mut ssnap = self.uninit(nc * h * D * D)?; let gdn_mma = !portable_mma_gated() && c == 32
11841 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11842 Ok("1") => true,
11843 Ok("0") => false,
11844 _ => cfg!(memra_hopper_mma),
11845 };
11846 if gdn_mma {
11847 let wb16 = wb16_pre.take().expect("mma path pre-allocates wb16 (K3 store fold)");
11848 let kb16_ref: &CudaSlice<u8> = kb16_ref0.expect("mma path pre-builds kb16 above K123");
11849 if gdn_wgmma_pre {
11861 let qb16 = qb16_ref0.unwrap();
11863 let pb16 = pb16.as_ref().unwrap();
11864 {
11865 let f = self.func("gdn_k45_wgmma");
11866 let cfg = LaunchConfig { grid_dim: (h as u32, 4, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11867 let hki = hk as i32;
11868 let __s_b = self.gpu.stream();
11869 let mut b = __s_b.launch_builder(&f);
11870 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(qb16).arg(pb16)
11871 .arg(o).arg(&scale).arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11872 unsafe { b.launch(cfg)?; }
11873 }
11874 return Ok(());
11875 }
11876 let mut y16 = self.alloc_u8_uninit(nc * h * c * D * 2)?;
11880 let mut ssnap16 = self.alloc_u8_uninit(nc * h * D * D * 2)?;
11881 {
11882 let f = self.func("gdn_chunk_state_mma");
11883 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11884 let hki = hk as i32;
11885 let __s_b = self.gpu.stream();
11886 let mut b = __s_b.launch_builder(&f);
11887 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(&mut y16).arg(&mut ssnap16)
11888 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11889 unsafe { b.launch(cfg)?; }
11890 }
11891 { let f = self.func("gdn_chunk_output_mma");
11893 let jt = ((c + 31) / 32) as u32;
11894 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11895 let hki = hk as i32;
11896 let __s_b = self.gpu.stream();
11897 let mut b = __s_b.launch_builder(&f);
11898 b.arg(q).arg(&gcum).arg(&p).arg(&y16).arg(&ssnap16).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale).arg(&hki);
11899 unsafe { b.launch(cfg)?; }
11900 }
11901 return Ok(());
11902 }
11903 { let f = self.func("gdn_chunk_state_f32");
11905 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11906 let __s_b = self.gpu.stream();
11907 let mut b = __s_b.launch_builder(&f);
11908 b.arg(k).arg(&gcum).arg(beta).arg(&u).arg(&w).arg(&mut y).arg(&mut ssnap)
11909 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci);
11910 unsafe { b.launch(cfg)?; }
11911 }
11912 { let f = self.func("gdn_chunk_output_f32");
11914 let jt = ((c + 31) / 32) as u32;
11915 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11916 let __s_b = self.gpu.stream();
11917 let mut b = __s_b.launch_builder(&f);
11918 b.arg(q).arg(&gcum).arg(&p).arg(&y).arg(&ssnap).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale);
11919 unsafe { b.launch(cfg)?; }
11920 }
11921 Ok(())
11922 }
11923
11924 #[allow(clippy::too_many_arguments)]
11933 #[allow(clippy::too_many_arguments)]
11934 pub fn gdn_scan_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11935 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
11936 qb16_pre: Option<&CudaSlice<u8>>,
11937 state_in: &CudaSlice<f32>,
11938 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11939 n_head: usize, t: usize, scale: f32, hk: usize)
11940 -> Result<(), Box<dyn std::error::Error>> {
11941 if std::env::var("MEMRA_GDN_DIFF").is_ok() && t >= 16 {
11942 assert!(hk == n_head, "GDN_DIFF oracle is broadcast-only");
11943 return self.gdn_scan_diff(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale);
11944 }
11945 if Self::gdn_chunked_enabled() && t >= 16 {
11946 self.gdn_scan_chunked(q, k, v, g, beta, kb16_pre, qb16_pre, state_in, state_out, o, n_head, t, scale,
11947 Self::gdn_chunk_size(), hk)
11948 } else {
11949 assert!(hk == n_head, "s128 scan is broadcast-only (prep guarantees by predicate)");
11950 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)
11951 }
11952 }
11953
11954 #[allow(clippy::too_many_arguments)]
11956 fn gdn_scan_diff(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11957 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
11958 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11959 n_head: usize, t: usize, scale: f32)
11960 -> Result<(), Box<dyn std::error::Error>> {
11961 static CALL: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
11962 let call = CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11963 let mut o_c = self.uninit(o.len())?;
11964 let mut st_c = self.uninit(state_out.len())?;
11965 self.gdn_scan_chunked(q, k, v, g, beta, None, None, state_in, &mut st_c, &mut o_c,
11966 n_head, t, scale, Self::gdn_chunk_size(), n_head)?;
11967 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)?;
11968 let (oh_s, oh_c) = (self.dtoh(o)?, self.dtoh(&o_c)?);
11969 let (sh_s, sh_c) = (self.dtoh(state_out)?, self.dtoh(&st_c)?);
11970 let stats = |a: &[f32], b: &[f32]| -> (f32, f32, f64) {
11971 let mut max_abs = 0f32; let mut max_rel = 0f32; let mut sum_rel = 0f64;
11972 for (x, y) in a.iter().zip(b) {
11973 let ad = (x - y).abs();
11974 let rel = ad / x.abs().max(y.abs()).max(1e-3);
11975 if ad > max_abs { max_abs = ad; }
11976 if rel > max_rel { max_rel = rel; }
11977 sum_rel += rel as f64;
11978 }
11979 (max_abs, max_rel, sum_rel / a.len() as f64)
11980 };
11981 let (o_ma, o_mr, o_mean) = stats(&oh_s, &oh_c);
11982 let (s_ma, s_mr, s_mean) = stats(&sh_s, &sh_c);
11983 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} | \
11984 state: max_abs={s_ma:.3e} max_rel={s_mr:.3e} mean_rel={s_mean:.3e}",
11985 Self::gdn_chunk_size());
11986 Ok(())
11987 }
11988
11989 pub fn gdn_glog(&self, alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11991 g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
11992 -> Result<(), Box<dyn std::error::Error>> {
11993 let f = self.func("gdn_glog_f32");
11994 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
11995 let (h, ti) = (n_head as i32, t as i32);
11996 let __s_b = self.gpu.stream();
11997 let mut b = __s_b.launch_builder(&f);
11998 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
11999 unsafe { b.launch(cfg)?; }
12000 Ok(())
12001 }
12002
12003 pub fn sigmoid_v(&self, x: &cudarc::driver::CudaView<f32>, y: &mut CudaSlice<f32>, n: usize)
12006 -> Result<(), Box<dyn std::error::Error>> {
12007 let f = self.func("sigmoid_f32");
12008 let cfg = LaunchConfig::for_num_elems(n as u32);
12009 let ni = n as i32;
12010 let __s_b = self.gpu.stream();
12011 let mut b = __s_b.launch_builder(&f);
12012 b.arg(x).arg(y).arg(&ni);
12013 unsafe { b.launch(cfg)?; }
12014 Ok(())
12015 }
12016
12017 pub fn gdn_glog_v(&self, alpha: &cudarc::driver::CudaView<f32>, dt_bias: &CudaSlice<f32>,
12018 a: &CudaSlice<f32>, g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
12019 -> Result<(), Box<dyn std::error::Error>> {
12020 let f = self.func("gdn_glog_f32");
12021 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
12022 let (h, ti) = (n_head as i32, t as i32);
12023 let __s_b = self.gpu.stream();
12024 let mut b = __s_b.launch_builder(&f);
12025 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
12026 unsafe { b.launch(cfg)?; }
12027 Ok(())
12028 }
12029
12030 pub fn sigmoid(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize)
12031 -> Result<(), Box<dyn std::error::Error>> {
12032 let f = self.func("sigmoid_f32");
12033 let cfg = LaunchConfig::for_num_elems(n as u32);
12034 let ni = n as i32;
12035 let __s_b = self.gpu.stream();
12036 let mut b = __s_b.launch_builder(&f);
12037 b.arg(x).arg(y).arg(&ni);
12038 unsafe { b.launch(cfg)?; }
12039 Ok(())
12040 }
12041
12042 pub fn sig_mul_f16out(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12045 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
12046 -> Result<(), Box<dyn std::error::Error>> {
12047 let f = self.func("sig_mul_f16out_f32");
12048 let cfg = LaunchConfig::for_num_elems(n as u32);
12049 let ni = n as i32;
12050 let __s_b = self.gpu.stream();
12051 let mut b = __s_b.launch_builder(&f);
12052 b.arg(a).arg(g).arg(dst).arg(dst16).arg(&ni);
12053 unsafe { b.launch(cfg)?; }
12054 Ok(())
12055 }
12056
12057 #[allow(clippy::too_many_arguments)]
12066 pub fn attn_head_gate(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12067 dst: &mut CudaSlice<f32>, dst16: Option<&mut CudaSlice<u8>>,
12068 head_dim: usize, n_head: usize, t: usize)
12069 -> Result<(), Box<dyn std::error::Error>> {
12070 let f = self.func("attn_head_gate_f32");
12071 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12072 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12073 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
12075 let __s_b = self.gpu.stream();
12076 let mut b = __s_b.launch_builder(&f);
12077 b.arg(a).arg(g).arg(dst).arg(&d16).arg(&hd).arg(&nh).arg(&ti);
12078 unsafe { b.launch(cfg)?; }
12079 Ok(())
12080 }
12081
12082 #[allow(clippy::too_many_arguments)]
12091 pub fn swiglu_clamped_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
12092 gs: f32, us: f32, limit: f32,
12093 dst: &mut CudaSlice<f32>, n: usize)
12094 -> Result<(), Box<dyn std::error::Error>> {
12095 debug_assert!(limit > 1e-6, "swiglu_clamped needs a live limit; use silu_mul_scaled");
12096 let f = self.func("swiglu_clamped_mul_scaled_f32");
12097 let cfg = LaunchConfig::for_num_elems(n as u32);
12098 let ni = n as i32;
12099 let __s_b = self.gpu.stream();
12100 let mut b = __s_b.launch_builder(&f);
12101 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&limit).arg(dst).arg(&ni);
12102 unsafe { b.launch(cfg)?; }
12103 Ok(())
12104 }
12105
12106 pub fn gated_rmsnorm(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12108 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12109 -> Result<(), Box<dyn std::error::Error>> {
12110 let f = self.func("gated_rmsnorm_f32");
12111 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12112 let (nc, e) = (ncols as i32, eps);
12113 let __s_b = self.gpu.stream();
12114 let mut b = __s_b.launch_builder(&f);
12115 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12116 unsafe { b.launch(cfg)?; }
12117 Ok(())
12118 }
12119
12120 pub fn gated_rmsnorm_f16out(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12123 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12124 ncols: usize, nrows: usize, eps: f32)
12125 -> Result<(), Box<dyn std::error::Error>> {
12126 let f = self.func("gated_rmsnorm_f16out_f32");
12127 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12129 let (nc, e) = (ncols as i32, eps);
12130 let __s_b = self.gpu.stream();
12131 let mut b = __s_b.launch_builder(&f);
12132 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12133 unsafe { b.launch(cfg)?; }
12134 Ok(())
12135 }
12136
12137 #[allow(clippy::too_many_arguments)]
12141 pub fn add_rms_norm_zq8(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
12142 res: &mut CudaSlice<f32>, z: &mut CudaSlice<f32>,
12143 ncols: usize, nrows: usize, eps: f32)
12144 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12145 assert!(ncols % 32 == 0);
12146 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
12147 let mut d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12148 let f = self.func("add_rms_norm_zq8");
12149 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
12150 let (nc, ep) = (ncols as i32, eps);
12151 let __s_b = self.gpu.stream();
12152 let mut b = __s_b.launch_builder(&f);
12153 b.arg(a).arg(b_in).arg(w).arg(res).arg(z).arg(&mut q).arg(&mut d).arg(&nc).arg(&ep);
12154 unsafe { b.launch(cfg)?; }
12155 Ok((q, d))
12156 }
12157
12158 pub fn gated_rmsnorm_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12163 z: &cudarc::driver::CudaView<f32>,
12164 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12165 -> Result<(), Box<dyn std::error::Error>> {
12166 let f = self.func("gated_rmsnorm_f32");
12167 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12168 let (nc, e) = (ncols as i32, eps);
12169 let __s_b = self.gpu.stream();
12170 let mut b = __s_b.launch_builder(&f);
12171 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12172 unsafe { b.launch(cfg)?; }
12173 Ok(())
12174 }
12175
12176 pub fn gated_rmsnorm_f16out_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12177 z: &cudarc::driver::CudaView<f32>,
12178 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12179 ncols: usize, nrows: usize, eps: f32)
12180 -> Result<(), Box<dyn std::error::Error>> {
12181 let f = self.func("gated_rmsnorm_f16out_f32");
12182 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12184 let (nc, e) = (ncols as i32, eps);
12185 let __s_b = self.gpu.stream();
12186 let mut b = __s_b.launch_builder(&f);
12187 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12188 unsafe { b.launch(cfg)?; }
12189 Ok(())
12190 }
12191
12192 pub fn gated_rmsnorm_q8_1(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12193 ncols: usize, nrows: usize, eps: f32)
12194 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12195 assert!(ncols % 32 == 0);
12196 let f = self.func("gated_rmsnorm_q8_1");
12197 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12198 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12199 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12200 let (nc, ep) = (ncols as i32, eps);
12201 let __s_b = self.gpu.stream();
12202 let mut b = __s_b.launch_builder(&f);
12203 b.arg(o).arg(w).arg(z).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
12204 unsafe { b.launch(cfg)?; }
12205 Ok((out_q, out_d))
12206 }
12207
12208 pub fn transpose(&self, inp: &CudaSlice<f32>, rows: usize, cols: usize)
12210 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12211 let f = self.func("transpose_f32");
12212 let mut out = self.zeros(rows * cols)?;
12213 let cfg = LaunchConfig::for_num_elems((rows * cols) as u32);
12214 let (r, c) = (rows as i32, cols as i32);
12215 let __s_b = self.gpu.stream();
12216 let mut b = __s_b.launch_builder(&f);
12217 b.arg(inp).arg(&mut out).arg(&r).arg(&c);
12218 unsafe { b.launch(cfg)?; }
12219 Ok(out)
12220 }
12221
12222 pub fn repeat_heads(&self, inp: &CudaSlice<f32>, out: &mut CudaSlice<f32>,
12224 head_dim: usize, n_in: usize, n_out: usize, t: usize)
12225 -> Result<(), Box<dyn std::error::Error>> {
12226 let f = self.func("repeat_heads_f32");
12227 let cfg = LaunchConfig::for_num_elems((head_dim * n_out * t) as u32);
12228 let (hd, ni, no, ti) = (head_dim as i32, n_in as i32, n_out as i32, t as i32);
12229 let __s_b = self.gpu.stream();
12230 let mut b = __s_b.launch_builder(&f);
12231 b.arg(inp).arg(out).arg(&hd).arg(&ni).arg(&no).arg(&ti);
12232 unsafe { b.launch(cfg)?; }
12233 Ok(())
12234 }
12235
12236 pub fn q_gate_split(&self, qf: &CudaSlice<f32>, q_out: &mut CudaSlice<f32>,
12239 gate_out: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, t: usize)
12240 -> Result<(), Box<dyn std::error::Error>> {
12241 let f = self.func("q_gate_split_f32");
12242 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12243 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12244 let __s_b = self.gpu.stream();
12245 let mut b = __s_b.launch_builder(&f);
12246 b.arg(qf).arg(q_out).arg(gate_out).arg(&hd).arg(&nh).arg(&ti);
12247 unsafe { b.launch(cfg)?; }
12248 Ok(())
12249 }
12250
12251 pub fn qkv_to_gdn_repack(&self, conv_out: &CudaSlice<f32>, q_g: &mut CudaSlice<f32>,
12255 k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
12256 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, t: usize)
12257 -> Result<(), Box<dyn std::error::Error>> {
12258 let f = self.func("qkv_to_gdn_repack_f32");
12259 let cfg = LaunchConfig::for_num_elems((d_state * num_v * t) as u32);
12260 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);
12261 let __s_b = self.gpu.stream();
12262 let mut b = __s_b.launch_builder(&f);
12263 b.arg(conv_out).arg(q_g).arg(k_g).arg(v_g).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&ti);
12264 unsafe { b.launch(cfg)?; }
12265 Ok(())
12266 }
12267
12268 pub fn conv_left_pad(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
12271 conv_dim: usize, t: usize, pad: usize)
12272 -> Result<(), Box<dyn std::error::Error>> {
12273 let f = self.func("conv_left_pad_f32");
12274 let cfg = LaunchConfig::for_num_elems((conv_dim * t) as u32);
12275 let (cd, ti, p) = (conv_dim as i32, t as i32, pad as i32);
12276 let __s_b = self.gpu.stream();
12277 let mut b = __s_b.launch_builder(&f);
12278 b.arg(src).arg(dst).arg(&cd).arg(&ti).arg(&p);
12279 unsafe { b.launch(cfg)?; }
12280 Ok(())
12281 }
12282
12283 pub fn conv_assemble_and_roll(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12287 conv_in: &mut CudaSlice<f32>, conv_dim: usize, pad: usize)
12288 -> Result<(), Box<dyn std::error::Error>> {
12289 let f = self.func("conv_assemble_and_roll_f32");
12290 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12291 let (cd, p) = (conv_dim as i32, pad as i32);
12292 let __s_b = self.gpu.stream();
12293 let mut b = __s_b.launch_builder(&f);
12294 b.arg(qkv_col).arg(conv_state).arg(conv_in).arg(&cd).arg(&p);
12295 unsafe { b.launch(cfg)?; }
12296 Ok(())
12297 }
12298
12299 pub fn ssm_conv1d_fused_decode(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12305 w: &CudaSlice<f32>, conv_out: &mut CudaSlice<f32>,
12306 conv_dim: usize, d_conv: usize)
12307 -> Result<(), Box<dyn std::error::Error>> {
12308 let f = self.func("ssm_conv1d_fused_decode_f32");
12309 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12310 let (cd, dc) = (conv_dim as i32, d_conv as i32);
12311 let __s_b = self.gpu.stream();
12312 let mut b = __s_b.launch_builder(&f);
12313 b.arg(qkv_col).arg(conv_state).arg(w).arg(conv_out).arg(&cd).arg(&dc);
12314 unsafe { b.launch(cfg)?; }
12315 Ok(())
12316 }
12317
12318 pub fn slice_range(&self, src: &CudaSlice<f32>, start: usize, len: usize)
12321 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12322 let host = self.gpu.stream().clone_dtoh(src)?;
12323 self.gpu.stream().synchronize()?;
12324 Ok(self.htod(&host[start..start + len])?)
12325 }
12326}
12327
12328#[cfg(test)]
12329mod target_dispatch_tests {
12330 use super::legacy_quant_gemm_allowed;
12331
12332 #[test]
12333 fn legacy_quant_gemm_arch_policy_honors_the_escape_hatch() {
12334 assert!(legacy_quant_gemm_allowed(false, false, false));
12336 assert!(!legacy_quant_gemm_allowed(false, false, true));
12337 assert!(!legacy_quant_gemm_allowed(true, false, false));
12339 assert!(!legacy_quant_gemm_allowed(true, false, true));
12340 assert!(legacy_quant_gemm_allowed(true, true, false));
12342 assert!(!legacy_quant_gemm_allowed(true, true, true));
12343 }
12344
12345 #[cfg(all(memra_portable_cuda, not(memra_hopper_mma)))]
12346 #[test]
12347 fn portable_build_disables_legacy_quant_gemm_without_an_env_override() {
12348 assert!(!legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12349 }
12350
12351 #[cfg(memra_hopper_mma)]
12352 #[test]
12353 fn hopper_mma_build_re_admits_legacy_quant_gemm() {
12354 assert!(legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12355 assert!(super::portable_mma_gated() == false);
12356 }
12357}
12358
12359impl memra_kv::KvDev for Engine {
12362 fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12363 Engine::zeros(self, n)
12364 }
12365 fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12366 Engine::uninit(self, n)
12367 }
12368 fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
12369 Engine::alloc_u8(self, n)
12370 }
12371 fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
12372 Engine::htod_i32(self, v)
12373 }
12374 fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12375 Engine::clone_dtod(self, src)
12376 }
12377 fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
12378 -> Result<(), Box<dyn std::error::Error>> {
12379 Engine::copy_into(self, dst, off, src, len)
12380 }
12381 fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
12382 Engine::set_i32_one(self, d, v)
12383 }
12384}