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_fork_valid(&self, acc: &CudaSlice<u32>, optimistic_pending: u32,
1924 valid: &mut CudaSlice<u32>)
1925 -> Result<(), Box<dyn std::error::Error>> {
1926 let f = self.func("spec_fork_valid");
1927 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1),
1928 shared_mem_bytes: 0 };
1929 let __s_bl = self.gpu.stream();
1930 let mut bl = __s_bl.launch_builder(&f);
1931 bl.arg(acc).arg(&optimistic_pending).arg(valid);
1932 unsafe { bl.launch(cfg)?; }
1933 Ok(())
1934 }
1935
1936 pub fn spec_fork_reconcile_kv(&self, len_ptrs: &CudaSlice<u64>, saved: &CudaSlice<i32>,
1938 valid: &CudaSlice<u32>, n_layer: usize)
1939 -> Result<(), Box<dyn std::error::Error>> {
1940 let f = self.func("spec_fork_reconcile_kv");
1941 let nl = n_layer as i32;
1942 let cfg = LaunchConfig { grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
1943 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1944 let __s_bl = self.gpu.stream();
1945 let mut bl = __s_bl.launch_builder(&f);
1946 bl.arg(len_ptrs).arg(saved).arg(valid).arg(&nl);
1947 unsafe { bl.launch(cfg)?; }
1948 Ok(())
1949 }
1950
1951 pub fn spec_fork_restore_f32(&self, snapshot: &CudaSlice<f32>, state: &mut CudaSlice<f32>,
1953 valid: &CudaSlice<u32>)
1954 -> Result<(), Box<dyn std::error::Error>> {
1955 assert_eq!(snapshot.len(), state.len(), "fork recurrent snapshot shape mismatch");
1956 let f = self.func("spec_fork_restore_f32");
1957 let n = state.len() as i32;
1958 let blocks = state.len().div_ceil(256).min(65535).max(1) as u32;
1959 let cfg = LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1),
1960 shared_mem_bytes: 0 };
1961 let __s_bl = self.gpu.stream();
1962 let mut bl = __s_bl.launch_builder(&f);
1963 bl.arg(snapshot).arg(state).arg(valid).arg(&n);
1964 unsafe { bl.launch(cfg)?; }
1965 Ok(())
1966 }
1967
1968 pub fn spec_seed_gather(&self, vx: &CudaSlice<f32>, fill_prev: &CudaSlice<f32>,
1971 acc: &CudaSlice<u32>, h_seed: &mut CudaSlice<f32>,
1972 base: usize, n_embd: usize)
1973 -> Result<(), Box<dyn std::error::Error>> {
1974 let f = self.func("spec_seed_gather");
1975 let (b, ne) = (base as i32, n_embd as i32);
1976 let cfg = LaunchConfig { grid_dim: (n_embd.div_ceil(256) as u32, 1, 1),
1977 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1978 let __s_bl = self.gpu.stream();
1979 let mut bl = __s_bl.launch_builder(&f);
1980 bl.arg(vx).arg(fill_prev).arg(acc).arg(h_seed).arg(&b).arg(&ne);
1981 unsafe { bl.launch(cfg)?; }
1982 Ok(())
1983 }
1984
1985
1986 pub fn spec_accept_greedy(&self, preds: &CudaSlice<u32>, draft: &CudaSlice<u32>,
1988 last_pred: u32, base: usize, k_round: usize,
1989 out: &mut CudaSlice<u32>)
1990 -> Result<(), Box<dyn std::error::Error>> {
1991 let f = self.func("spec_accept_greedy");
1992 let (b, k) = (base as i32, k_round as i32);
1993 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1994 let __s_bl = self.gpu.stream();
1995 let mut bl = __s_bl.launch_builder(&f);
1996 bl.arg(preds).arg(draft).arg(&last_pred).arg(&b).arg(&k).arg(out);
1997 unsafe { bl.launch(cfg)?; }
1998 Ok(())
1999 }
2000
2001 pub fn gumbel_perturb(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
2008 seed: u64, stream_pos: u32, temp: f32)
2009 -> Result<(), Box<dyn std::error::Error>> {
2010 let f = self.func("gumbel_perturb_f32");
2011 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2012 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2013 let __s_b = self.gpu.stream();
2014 let mut b = __s_b.launch_builder(&f);
2015 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
2016 unsafe { b.launch(cfg)?; }
2017 Ok(())
2018 }
2019
2020 pub fn mask_logits_col(&self, logits: &mut CudaSlice<f32>, mask: &CudaSlice<u32>,
2028 col: usize, n: usize, mask_words: usize)
2029 -> Result<(), Box<dyn std::error::Error>> {
2030 let f = self.func("mask_logits_f32");
2031 let (ci, ni, mw) = (col as i32, n as i32, mask_words as i32);
2032 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256).min(1024) as u32, 1, 1),
2033 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2034 let __s_b = self.gpu.stream();
2035 let mut b = __s_b.launch_builder(&f);
2036 b.arg(&mut *logits).arg(mask).arg(&ci).arg(&ni).arg(&mw);
2037 unsafe { b.launch(cfg)?; }
2038 Ok(())
2039 }
2040
2041 pub fn gumbel_perturb_col(&self, x: &CudaSlice<f32>, col: usize, y: &mut CudaSlice<f32>,
2048 n: usize, seed: u64, stream_pos: u32, temp: f32)
2049 -> Result<(), Box<dyn std::error::Error>> {
2050 let f = self.func("gumbel_perturb_f32");
2051 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2052 let col_view = x.slice(col * n..(col + 1) * n);
2053 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2054 let __s_b = self.gpu.stream();
2055 let mut b = __s_b.launch_builder(&f);
2056 b.arg(&col_view).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
2057 unsafe { b.launch(cfg)?; }
2058 Ok(())
2059 }
2060
2061 pub fn sctr_inc(&self, ctr: &mut CudaSlice<u32>) -> Result<(), Box<dyn std::error::Error>> {
2066 let f = self.func("memra_sctr_inc");
2067 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
2068 let __s_b = self.gpu.stream();
2069 let mut b = __s_b.launch_builder(&f);
2070 b.arg(&mut *ctr);
2071 unsafe { b.launch(cfg)?; }
2072 Ok(())
2073 }
2074
2075 pub fn gumbel_perturb_ctr(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
2080 seed: u64, ctr: &CudaSlice<u32>, temp: f32)
2081 -> Result<(), Box<dyn std::error::Error>> {
2082 let f = self.func("gumbel_perturb_ctr_f32");
2083 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2084 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2085 let __s_b = self.gpu.stream();
2086 let mut b = __s_b.launch_builder(&f);
2087 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(ctr).arg(&temp);
2088 unsafe { b.launch(cfg)?; }
2089 Ok(())
2090 }
2091
2092 pub fn softmax_gather(&self, x: &CudaSlice<f32>, row_stride: usize,
2096 ids: &CudaSlice<u32>, rows: &CudaSlice<i32>,
2097 out: &mut CudaSlice<f32>, n: usize, npair: usize, temp: f32)
2098 -> Result<(), Box<dyn std::error::Error>> {
2099 let f = self.func("softmax_gather_f32");
2100 let (ni, rs) = (n as i32, row_stride as i64);
2101 let np = npair as i32;
2102 let cfg = LaunchConfig { grid_dim: (npair as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2103 let __s_b = self.gpu.stream();
2104 let mut b = __s_b.launch_builder(&f);
2105 b.arg(x).arg(&rs).arg(ids).arg(rows).arg(&mut *out).arg(&ni).arg(&np).arg(&temp);
2106 unsafe { b.launch(cfg)?; }
2107 Ok(())
2108 }
2109
2110 pub fn residual_sample(&self, p: &CudaSlice<f32>, q: Option<&CudaSlice<f32>>, n: usize,
2114 temp: f32, seed: u64, stream_pos: u32,
2115 out_tok: &mut CudaSlice<u32>)
2116 -> Result<(), Box<dyn std::error::Error>> {
2117 let f = self.func("residual_sample_f32");
2118 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2119 let nth = 1024u32;
2120 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (nth, 1, 1), shared_mem_bytes: 0 };
2121 let has_q: i32 = q.is_some() as i32;
2122 let qbuf = q.unwrap_or(p); let __s_b = self.gpu.stream();
2124 let mut b = __s_b.launch_builder(&f);
2125 b.arg(p).arg(qbuf).arg(&has_q).arg(&ni).arg(&temp).arg(&slo).arg(&shi).arg(&stream_pos)
2126 .arg(&mut *out_tok);
2127 unsafe { b.launch(cfg)?; }
2128 Ok(())
2129 }
2130
2131 pub fn with_moe_cache<R>(&self, max_block_bytes: usize,
2136 f: impl FnOnce(&mut crate::moe_cache::MoeSlotCache, &Engine) -> Result<R, Box<dyn std::error::Error>>)
2137 -> Result<R, Box<dyn std::error::Error>> {
2138 let mut guard = self.moe_cache.lock().unwrap();
2139 if guard.is_none() {
2140 *guard = Some(crate::moe_cache::MoeSlotCache::new(self, max_block_bytes)?);
2141 }
2142 let cache = guard.as_mut().unwrap();
2143 f(cache, self)
2144 }
2145
2146 pub fn freeze_moe_cache(&self) {
2149 if let Some(cache) = self.moe_cache.lock().unwrap().as_mut() {
2150 cache.freeze();
2151 }
2152 }
2153
2154 pub fn export_moe_residency(&self) -> Option<Vec<(u16, u8, u16)>> {
2157 self.moe_cache
2158 .lock()
2159 .unwrap()
2160 .as_ref()
2161 .map(crate::moe_cache::MoeSlotCache::export_residency)
2162 }
2163
2164 pub(crate) fn moe_cache_frozen(&self) -> bool {
2165 self.moe_cache
2166 .lock()
2167 .unwrap()
2168 .as_ref()
2169 .is_some_and(crate::moe_cache::MoeSlotCache::is_frozen)
2170 }
2171
2172 pub fn frozen_cpu_experts_prefer_tokenwise_prime(&self) -> bool {
2179 crate::cpu_experts::configured()
2180 && self.moe_cache_frozen()
2181 && std::env::var("MEMRA_CPU_EXPERT_BATCHED_PRIME").as_deref() != Ok("1")
2182 }
2183
2184 pub(crate) fn configure_moe_cache_layout(&self, block_bytes: Vec<usize>) {
2186 assert!(
2187 self.moe_cache.lock().unwrap().is_none(),
2188 "MoE cache layout configured after cache construction"
2189 );
2190 *self.moe_cache_layout.lock().unwrap() = Some(block_bytes);
2191 }
2192
2193 pub(crate) fn moe_cache_layout(&self) -> Option<Vec<usize>> {
2194 self.moe_cache_layout.lock().unwrap().clone()
2195 }
2196
2197 pub fn moe_cache_enabled() -> bool {
2199 std::env::var("MEMRA_MOE_CACHE").as_deref() != Ok("0")
2200 }
2201
2202 pub fn moe_cache_stats(&self) -> Option<(u64, u64, u64, usize)> {
2205 let guard = self.moe_cache.lock().unwrap();
2206 guard.as_ref() .map(|c| (c.hits, c.misses, c.staged_bytes, c.n_slots()))
2207 }
2208
2209 pub fn cpu_expert_stats(
2213 &self,
2214 ) -> Option<(u64, u64, u64, u64, u64, u64, u64, u64, u64, u64, u64)> {
2215 crate::cpu_experts::configured().then(crate::cpu_experts::stats)
2216 }
2217
2218 pub fn cpu_expert_predictor_stats(&self) -> (u64, u64) {
2221 crate::cpu_experts::predictor_stats()
2222 }
2223
2224 pub fn cpu_expert_exposed_wait_ns(&self) -> Option<u64> {
2225 crate::cpu_experts::configured().then(crate::cpu_experts::exposed_wait_ns)
2226 }
2227
2228 pub fn cpu_expert_gpu_residency_stats(&self) -> Option<(u64, u64, u64)> {
2231 crate::cpu_experts::configured().then(crate::cpu_experts::incomplete_gpu_residency_stats)
2232 }
2233
2234 pub fn moe_pread_stats(&self) -> Option<(u64, u64, u64, u64, u64, u64, u64)> {
2237
2238 let guard = self.moe_cache.lock().unwrap();
2239 guard.as_ref().and_then(|cache| cache.pread_stats()).map(|stats| (
2240 stats.reads,
2241 stats.bytes,
2242 stats.read_errors,
2243 stats.short_reads,
2244 stats.fallbacks,
2245 stats.buffer_waits,
2246 stats.ring_full,
2247 ))
2248 }
2249
2250 pub fn moe_cache_reset_counters(&self) {
2252 if let Some(c) = self.moe_cache.lock().unwrap().as_mut() { c.reset_counters(); }
2253 }
2254
2255 pub fn htod_bytes(&self, v: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2256 Ok(self.gpu.stream().clone_htod(v)?)
2257 }
2258
2259 pub fn htod_bytes_padded(&self, v: &[u8], pad: usize)
2263 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2264 let mut d = self.alloc_u8_uninit(v.len() + pad)?;
2265 {
2266 let mut view = d.slice_mut(0..v.len());
2267 self.gpu.stream().memcpy_htod(v, &mut view)?;
2268 }
2269 Ok(d)
2270 }
2271
2272 pub fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
2274 -> Result<(), Box<dyn std::error::Error>> {
2275 let mut view = dst.slice_mut(off..off + len);
2276 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2277 Ok(())
2278 }
2279
2280 pub fn copy_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &CudaSlice<u8>, len: usize)
2283 -> Result<(), Box<dyn std::error::Error>> {
2284 let mut view = dst.slice_mut(off..off + len);
2285 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2286 Ok(())
2287 }
2288
2289 pub fn copy_u8_range_into(
2291 &self,
2292 dst: &mut CudaSlice<u8>,
2293 dst_off: usize,
2294 src: &CudaSlice<u8>,
2295 src_off: usize,
2296 len: usize,
2297 ) -> Result<(), Box<dyn std::error::Error>> {
2298 let mut dst_view = dst.slice_mut(dst_off..dst_off + len);
2299 self.gpu
2300 .stream()
2301 .memcpy_dtod(&src.slice(src_off..src_off + len), &mut dst_view)?;
2302 Ok(())
2303 }
2304
2305 pub fn prepare_kv_append(
2309 &self,
2310 kv: &mut crate::cache::KvLayer,
2311 retain_from: usize,
2312 append_rows: usize,
2313 ) -> Result<usize, Box<dyn std::error::Error>> {
2314 let Some(plan) = kv
2315 .ring
2316 .as_ref()
2317 .map(|ring| ring.append_plan(kv.len, retain_from, append_rows))
2318 .transpose()?
2319 else {
2320 return Ok(kv.len);
2321 };
2322 match plan {
2323 crate::cache::KvRingAppend::Contiguous { write_row } => Ok(write_row),
2324 crate::cache::KvRingAppend::Rebase {
2325 src_row,
2326 keep_rows,
2327 new_base,
2328 write_row,
2329 } => {
2330 if keep_rows > 0 {
2331 let k_len = keep_rows * kv.k_tok_bytes;
2332 let v_len = keep_rows * kv.v_tok_bytes;
2333 let mut k_tmp = self.alloc_u8_uninit(k_len)?;
2334 let mut v_tmp = self.alloc_u8_uninit(v_len)?;
2335 self.copy_u8_range_into(
2336 &mut k_tmp,
2337 0,
2338 &kv.k,
2339 src_row * kv.k_tok_bytes,
2340 k_len,
2341 )?;
2342 self.copy_u8_range_into(
2343 &mut v_tmp,
2344 0,
2345 &kv.v,
2346 src_row * kv.v_tok_bytes,
2347 v_len,
2348 )?;
2349 self.copy_u8_into(&mut kv.k, 0, &k_tmp, k_len)?;
2350 self.copy_u8_into(&mut kv.v, 0, &v_tmp, v_len)?;
2351 }
2352 kv.ring.as_mut().unwrap().apply_rebase(new_base);
2353 Ok(write_row)
2354 }
2355 }
2356 }
2357
2358 pub fn htod_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &[u8])
2361 -> Result<(), Box<dyn std::error::Error>> {
2362 let mut view = dst.slice_mut(off..off + src.len());
2363 self.gpu.stream().memcpy_htod(src, &mut view)?;
2364 Ok(())
2365 }
2366
2367 pub fn view<'a>(&self, b: &'a CudaSlice<f32>, len: usize) -> cudarc::driver::CudaView<'a, f32> {
2368 b.slice(0..len)
2369 }
2370
2371 pub fn view_u8_range<'a>(&self, b: &'a CudaSlice<u8>, start: usize, end: usize)
2374 -> cudarc::driver::CudaView<'a, u8> {
2375 b.slice(start..end)
2376 }
2377 pub fn view_u8<'a>(&self, b: &'a CudaSlice<u8>, len: usize) -> cudarc::driver::CudaView<'a, u8> {
2378 b.slice(0..len)
2379 }
2380
2381 pub fn append_kv_quantized(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2385 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2386 kv_dim_k: usize, kv_dim_v: usize,
2387 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2388 -> Result<(), Box<dyn std::error::Error>> {
2389 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") } else { self.func("append_quantize_kv_q8_0_q5_1") };
2390 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2391 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2392 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2393 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2394 let __s_b = self.gpu.stream();
2395 let mut b = __s_b.launch_builder(&f);
2396 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2397 unsafe { b.launch(cfg)?; }
2398 Ok(())
2399 }
2400
2401 pub fn append_kv_quantized_dc(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2405 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t_dev: &CudaSlice<i32>,
2406 kv_dim_k: usize, kv_dim_v: usize,
2407 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2408 -> Result<(), Box<dyn std::error::Error>> {
2409 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2410 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
2411 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2412 if Self::pdl_on() && Self::pdl_wb_on() {
2414 use cudarc::driver::{DevicePtr, DevicePtrMut};
2415 let s = &self.gpu.stream();
2416 let (pk, _g0) = k_row.device_ptr(s); let (pv, _g1) = v_row.device_ptr(s);
2417 let (pkc, _g2) = kc.device_ptr_mut(s); let (pvc, _g3) = vc.device_ptr_mut(s);
2418 let (pt, _g4) = t_dev.device_ptr(s);
2419 let mut ps = [
2420 &pk as *const _ as *mut std::ffi::c_void, &pv as *const _ as *mut _,
2421 &pkc as *const _ as *mut _, &pvc as *const _ as *mut _,
2422 &pt as *const _ as *mut _, &kdk as *const _ as *mut _,
2423 &kdv as *const _ as *mut _, &ktb as *const _ as *mut _,
2424 &vtb as *const _ as *mut _,
2425 ];
2426 unsafe { self.launch_pdl_flash(g, "append_quantize_kv_q8_0_q5_1_dc",
2427 (nblk, 1, 1), (32, 1, 1), 0, &mut ps)?; }
2428 return Ok(());
2429 }
2430 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") };
2431 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2432 let __s_b = self.gpu.stream();
2433 let mut b = __s_b.launch_builder(&f);
2434 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(t_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2435 unsafe { b.launch(cfg)?; }
2436 Ok(())
2437 }
2438
2439 #[allow(clippy::too_many_arguments)]
2446 pub fn append_kv_quantized_rows(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
2447 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
2448 t0: usize, t: usize, kv_dim_k: usize, kv_dim_v: usize,
2449 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2450 -> Result<(), Box<dyn std::error::Error>> {
2451 if std::env::var("MEMRA_PRIME_APPEND_LOOP").is_ok() {
2452 for i in 0..t {
2453 let k_row = k_rows.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
2454 let v_row = v_rows.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
2455 self.append_kv_quantized_view(&k_row, &v_row, kc, vc, t0 + i,
2456 kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes, g)?;
2457 }
2458 return Ok(());
2459 }
2460 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") };
2461 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2462 let cfg = LaunchConfig { grid_dim: (nblk, t as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2463 let (t0i, kdk, kdv) = (t0 as i32, kv_dim_k as i32, kv_dim_v as i32);
2464 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2465 let __s_b = self.gpu.stream();
2466 let mut b = __s_b.launch_builder(&f);
2467 b.arg(k_rows).arg(v_rows).arg(kc).arg(vc).arg(&t0i).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2468 unsafe { b.launch(cfg)?; }
2469 Ok(())
2470 }
2471
2472 pub fn inc_seqlen(&self, p: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
2476 let f = self.func("inc_i32");
2477 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
2478 let __s_b = self.gpu.stream();
2479 let mut b = __s_b.launch_builder(&f);
2480 b.arg(p);
2481 unsafe { b.launch(cfg)?; }
2482 Ok(())
2483 }
2484
2485 pub fn append_kv_quantized_view(&self, k_row: &cudarc::driver::CudaView<f32>,
2488 v_row: &cudarc::driver::CudaView<f32>,
2489 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2490 kv_dim_k: usize, kv_dim_v: usize,
2491 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2492 -> Result<(), Box<dyn std::error::Error>> {
2493 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") }
2494 else { self.func("append_quantize_kv_q8_0_q5_1") };
2495 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2496 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2497 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2498 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2499 let __s_b = self.gpu.stream();
2500 let mut b = __s_b.launch_builder(&f);
2501 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2502 unsafe { b.launch(cfg)?; }
2503 Ok(())
2504 }
2505
2506 pub fn copy_view_into(&self, dst: &mut CudaSlice<f32>, off: usize,
2509 src: &cudarc::driver::CudaView<f32>, len: usize)
2510 -> Result<(), Box<dyn std::error::Error>> {
2511 let mut view = dst.slice_mut(off..off + len);
2512 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2513 Ok(())
2514 }
2515
2516 pub fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2520 let mut dst = self.gpu.stream().alloc_zeros::<f32>(src.len())?;
2521 self.gpu.stream().memcpy_dtod(src, &mut dst)?;
2522 Ok(dst)
2523 }
2524
2525 pub fn dtod_copy_view(&self, src: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<f32>)
2528 -> Result<(), Box<dyn std::error::Error>> {
2529 self.gpu.stream().memcpy_dtod(src, dst)?;
2530 Ok(())
2531 }
2532
2533 pub fn dtod_copy_view_i8(&self, src: &cudarc::driver::CudaView<i8>, dst: &mut CudaSlice<i8>)
2535 -> Result<(), Box<dyn std::error::Error>> {
2536 self.gpu.stream().memcpy_dtod(src, dst)?;
2537 Ok(())
2538 }
2539
2540 pub fn dtod_copy_into(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, offset: usize)
2542 -> Result<(), Box<dyn std::error::Error>> {
2543 let n = src.len();
2544 let mut dv = dst.slice_mut(offset..offset + n);
2545 self.gpu.stream().memcpy_dtod(src, &mut dv)?;
2546 Ok(())
2547 }
2548
2549 pub fn uninit_i8(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
2551 self.alloc_uninit::<i8>(n)
2552 }
2553
2554 pub fn qmatvec(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
2556 qtype: i32, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2557 let f = self.func("qmatvec_f32");
2558 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 };
2560 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2561 let __s_b = self.gpu.stream();
2562 let mut b = __s_b.launch_builder(&f);
2563 b.arg(w).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2564 unsafe { b.launch(cfg)?; }
2565 Ok(y)
2566 }
2567
2568 pub fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2570 let s = self.gpu.stream().alloc_zeros::<u8>(n)?;
2571 self.keep_if_capturing(&s);
2572 Ok(s)
2573 }
2574
2575 pub fn alloc_u8_uninit(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2579 let s = unsafe { self.gpu.stream().alloc::<u8>(n)? };
2580 self.keep_if_capturing(&s);
2581 Ok(s)
2582 }
2583
2584 pub fn memset_zeros_view(&self, dst: &mut cudarc::driver::CudaViewMut<f32>)
2587 -> Result<(), Box<dyn std::error::Error>> {
2588 self.gpu.stream().memset_zeros(dst)?;
2589 Ok(())
2590 }
2591
2592 pub fn stage_expert(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2598 -> Result<(), Box<dyn std::error::Error>> {
2599 let mut dst = scratch.slice_mut(off..off + host_bytes.len()); self.gpu.stream().memcpy_htod(host_bytes, &mut dst)?; Ok(())
2602 }
2603
2604 pub fn moe_router_topk(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2610 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2611 let f = self.func("moe_router_topk_f32");
2612 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),
2615 shared_mem_bytes: 0 };
2616 let (ne, nu) = (n_expert as i32, n_used as i32);
2617 let __s_b = self.gpu.stream();
2618 let mut b = __s_b.launch_builder(&f);
2619 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2620 unsafe { b.launch(cfg)?; }
2621 Ok((sel_idx, sel_w))
2622 }
2623
2624 pub fn moe_router_topk_scaled(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize,
2627 n_used: usize, ex_scale: &CudaSlice<f32>)
2628 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2629 let f = self.func("moe_router_topk_scaled_f32");
2634 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
2635 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
2636 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2637 shared_mem_bytes: 0 };
2638 let (ne, nu) = (n_expert as i32, n_used as i32);
2639 let __s_b = self.gpu.stream();
2640 let mut b = __s_b.launch_builder(&f);
2641 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu).arg(ex_scale);
2642 unsafe { b.launch(cfg)?; }
2643 Ok((sel_idx, sel_w))
2644 }
2645
2646 pub fn moe_router_topk_host(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2654 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2655 let f = self.func("moe_router_topk_f32");
2656 let n = t * n_used;
2657 let mut sel_idx = self.alloc_uninit::<i32>(n)?;
2658 let mut sel_w = self.alloc_uninit::<f32>(n)?;
2659 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2660 shared_mem_bytes: 0 };
2661 let (ne, nu) = (n_expert as i32, n_used as i32);
2662 let __s_b = self.gpu.stream();
2663 let mut b = __s_b.launch_builder(&f);
2664 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2665 unsafe { b.launch(cfg)?; }
2666 let bytes = n * 8;
2668 let mut guard = self.router_stage.lock().unwrap();
2669 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
2670 *guard = Some(PinnedStage::new(bytes.max(4096))?);
2671 }
2672 let stage = guard.as_mut().unwrap();
2673 let (si, sw) = unsafe {
2674 (std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
2675 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n))
2676 };
2677 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()))
2681 }
2682
2683 pub fn stage_expert_async(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2687 -> Result<cudarc::driver::CudaEvent, Box<dyn std::error::Error>> {
2688 let mut dst = scratch.slice_mut(off..off + host_bytes.len());
2689 self.copy_stream.memcpy_htod(host_bytes, &mut dst)?;
2690 Ok(self.copy_stream.record_event(None)?)
2691 }
2692
2693 pub fn compute_wait(&self, ev: &cudarc::driver::CudaEvent) -> Result<(), Box<dyn std::error::Error>> {
2695 self.gpu.stream().wait(ev)?;
2696 Ok(())
2697 }
2698
2699 pub fn qmatvec_view(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2704 x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize, out_f: usize,
2705 qtype: i32, row_bytes: usize)
2706 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2707 let f = self.func("qmatvec_f32");
2708 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 };
2711 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2712 let __s_b = self.gpu.stream();
2713 let mut b = __s_b.launch_builder(&f);
2714 b.arg(&wv).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2715 unsafe { b.launch(cfg)?; }
2716 Ok(y)
2717 }
2718
2719 #[allow(clippy::too_many_arguments)]
2726 pub fn moe_gate_up_silu8_q8(&self, gp: WPtr8, up: WPtr8,
2730 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2731 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2732 rb_g: usize, rb_u: usize)
2733 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2734 let f = self.func("moe_gate_up_silu8_q8");
2735 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
2736 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2737 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2738 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2739 let __s_b = self.gpu.stream();
2740 let mut b = __s_b.launch_builder(&f);
2741 b.arg(&gp).arg(&up).arg(aq).arg(ad).arg(&mut act)
2742 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2743 unsafe { b.launch(cfg)?; }
2744 Ok(act)
2745 }
2746
2747 #[allow(clippy::too_many_arguments)]
2748 pub fn moe_down8_fma_q8(&self, dp: WPtr8, w: F32x8,
2749 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
2750 dst: &mut cudarc::driver::CudaViewMut<f32>,
2751 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2752 -> Result<(), Box<dyn std::error::Error>> {
2753 let f = self.func("moe_down8_fma_q8");
2754 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2755 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2756 let (inf, outf, nu, rbi) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2757 let __s_b = self.gpu.stream();
2758 let mut b = __s_b.launch_builder(&f);
2759 b.arg(&dp).arg(&w).arg(aq2).arg(ad2).arg(dst)
2760 .arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbi);
2761 unsafe { b.launch(cfg)?; }
2762 Ok(())
2763 }
2764
2765 pub fn qmatvec_expert_q8(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2767 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
2768 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize)
2769 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2770 let f = self.func("qmatvec_expert_q8");
2771 let wv = w.slice(range);
2772 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
2773 const ROWS: u32 = 4; let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, m as u32, 1),
2775 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2776 let (inf, outf, mi, rbi) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
2777 let __s_b = self.gpu.stream();
2778 let mut b = __s_b.launch_builder(&f);
2779 b.arg(&wv).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qtype).arg(&rbi);
2780 unsafe { b.launch(cfg)?; }
2781 Ok(y)
2782 }
2783
2784 pub fn moe_gate_up_silu8(&self, gp: WPtr8, up: WPtr8, x: &cudarc::driver::CudaView<f32>,
2785 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2786 rb_g: usize, rb_u: usize)
2787 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2788 let f = self.func("moe_gate_up_silu8_f32");
2789 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),
2791 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2792 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2793 let __s_b = self.gpu.stream();
2794 let mut b = __s_b.launch_builder(&f);
2795 b.arg(&gp).arg(&up).arg(x).arg(&mut act)
2796 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2797 unsafe { b.launch(cfg)?; }
2798 Ok(act)
2799 }
2800
2801 #[allow(clippy::too_many_arguments)]
2807 pub fn moe_down8_fma_into(&self, dp: WPtr8, w: F32x8, act: &CudaSlice<f32>,
2808 dst: &mut cudarc::driver::CudaViewMut<f32>,
2809 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2810 -> Result<(), Box<dyn std::error::Error>> {
2811 let f = self.func("moe_down8_fma_f32");
2812 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2813 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2814 let (inf, outf, nu, rbv) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2815 let __s_b = self.gpu.stream();
2816 let mut b = __s_b.launch_builder(&f);
2817 b.arg(&dp).arg(&w).arg(act).arg(dst).arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbv);
2818 unsafe { b.launch(cfg)?; }
2819 Ok(())
2820 }
2821
2822 #[allow(clippy::too_many_arguments)]
2827 #[allow(clippy::too_many_arguments)]
2842 #[allow(clippy::too_many_arguments)]
2844 pub fn moe_pairs_matvec_q8(&self, table: &CudaSlice<u64>, proj: i32,
2845 pair_tok: &CudaSlice<i32>, pair_ex: &CudaSlice<i32>,
2846 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2847 in_f: usize, out_f: usize, n_expert: usize, n_pairs: usize,
2848 qtype: i32, row_bytes: usize)
2849 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2850 let f = self.func("moe_pairs_matvec_q8");
2851 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2852 const ROWS: u32 = 4;
2853 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_pairs as u32, 1),
2854 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2855 let (inf, outf, ne, np, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2856 n_pairs as i32, row_bytes as i64);
2857 let __s_b = self.gpu.stream();
2858 let mut b = __s_b.launch_builder(&f);
2859 b.arg(table).arg(&proj).arg(pair_tok).arg(pair_ex).arg(aq).arg(ad).arg(&mut y)
2860 .arg(&inf).arg(&outf).arg(&ne).arg(&np).arg(&qtype).arg(&rbi);
2861 unsafe { b.launch(cfg)?; }
2862 Ok(y)
2863 }
2864
2865 #[allow(clippy::too_many_arguments)]
2867 pub fn moe_pairs_matvec_q8_em(&self, table: &CudaSlice<u64>, proj: i32,
2868 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2869 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2870 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2871 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2872 n_pairs: usize, qtype: i32, row_bytes: usize)
2873 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2874 let f = self.func("moe_pairs_matvec_q8_em");
2875 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2876 const ROWS: u32 = 4;
2877 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
2878 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2879 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2880 n_active as i32, row_bytes as i64);
2881 let __s_b = self.gpu.stream();
2882 let mut b = __s_b.launch_builder(&f);
2883 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
2884 .arg(aq).arg(ad).arg(&mut y)
2885 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
2886 unsafe { b.launch(cfg)?; }
2887 Ok(y)
2888 }
2889
2890 #[allow(clippy::too_many_arguments)]
2893 pub fn moe_pairs_matvec_q8_dec(&self, table: &CudaSlice<u64>, proj: i32,
2894 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2895 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2896 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2897 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2898 n_pairs: usize, qtype: i32, row_bytes: usize)
2899 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2900 let f = self.func("moe_pairs_matvec_q8_dec");
2901 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2902 const ROWS: u32 = 4;
2903 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
2904 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2905 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2906 n_active as i32, row_bytes as i64);
2907 let __s_b = self.gpu.stream();
2908 let mut b = __s_b.launch_builder(&f);
2909 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
2910 .arg(aq).arg(ad).arg(&mut y)
2911 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
2912 unsafe { b.launch(cfg)?; }
2913 Ok(y)
2914 }
2915
2916 pub fn moe_pairs_gelu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
2917 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2918 let f = self.func("moe_pairs_gelu_mul");
2919 let mut act = self.alloc_uninit::<f32>(n)?;
2920 let cfg = LaunchConfig::for_num_elems(n as u32);
2921 let nl = n as i64;
2922 let __s_b = self.gpu.stream();
2923 let mut b = __s_b.launch_builder(&f);
2924 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
2925 unsafe { b.launch(cfg)?; }
2926 Ok(act)
2927 }
2928
2929 pub fn moe_pairs_silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
2930 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2931 let f = self.func("moe_pairs_silu_mul");
2932 let mut act = self.alloc_uninit::<f32>(n)?;
2933 let cfg = LaunchConfig::for_num_elems(n as u32);
2934 let nl = n as i64;
2935 let __s_b = self.gpu.stream();
2936 let mut b = __s_b.launch_builder(&f);
2937 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
2938 unsafe { b.launch(cfg)?; }
2939 Ok(act)
2940 }
2941
2942 #[allow(clippy::too_many_arguments)]
2943 pub fn moe_pairs_scatter(&self, y_down: &CudaSlice<f32>, pair_w: &CudaSlice<f32>,
2944 tok_pair_off: &CudaSlice<i32>, tok_pair_ids: &CudaSlice<i32>,
2945 moe_out: &mut CudaSlice<f32>, t: usize, n_embd: usize)
2946 -> Result<(), Box<dyn std::error::Error>> {
2947 let f = self.func("moe_pairs_scatter");
2948 let cfg = LaunchConfig { grid_dim: (((n_embd + 255) / 256) as u32, t as u32, 1),
2949 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2950 let ne = n_embd as i32;
2951 let __s_b = self.gpu.stream();
2952 let mut b = __s_b.launch_builder(&f);
2953 b.arg(y_down).arg(pair_w).arg(tok_pair_off).arg(tok_pair_ids).arg(moe_out).arg(&ne);
2954 unsafe { b.launch(cfg)?; }
2955 Ok(())
2956 }
2957
2958 #[allow(clippy::too_many_arguments)]
2962 pub fn moe_gate_up_gelu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
2963 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2964 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
2965 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
2966 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2967 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
2968 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
2969 rb_g as i64, rb_u as i64);
2970 let f = self.func("moe_gate_up_gelu8_dev_q8");
2971 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2972 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2973 let __s_b = self.gpu.stream();
2974 let mut b = __s_b.launch_builder(&f);
2975 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
2976 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2977 unsafe { b.launch(cfg)?; }
2978 Ok(act)
2979 }
2980
2981 #[allow(clippy::too_many_arguments)]
2983 pub fn moe_gate_up_gelu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
2984 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
2985 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
2986 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
2987 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2988 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
2989 let (inf, nff, ne, rbg, rbu, nu) = (in_f as i32, n_ff as i32, n_expert as i32,
2990 rb_g as i64, rb_u as i64, n_used as i32);
2991 let f = self.func("moe_gate_up_gelu8_dev_q8_rows");
2992 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
2993 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2994 let __s_b = self.gpu.stream();
2995 let mut b = __s_b.launch_builder(&f);
2996 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
2997 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu);
2998 unsafe { b.launch(cfg)?; }
2999 Ok(act)
3000 }
3001
3002 #[allow(clippy::too_many_arguments)]
3004 pub fn moe_gate_up_gelu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3005 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, n_pairs: usize,
3006 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3007 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3008 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3009 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
3010 let (inf, nff, ne, rbg, rbu, nu, npi) = (in_f as i32, n_ff as i32, n_expert as i32,
3011 rb_g as i64, rb_u as i64, n_used as i32,
3012 n_pairs as i32);
3013 let f = self.func("moe_gate_up_gelu8_dev_q8_csr");
3014 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
3015 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3016 let __s_b = self.gpu.stream();
3017 let mut b = __s_b.launch_builder(&f);
3018 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3019 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
3020 unsafe { b.launch(cfg)?; }
3021 Ok(act)
3022 }
3023
3024 #[allow(clippy::too_many_arguments)]
3026 pub fn moe_down8_fma_dev_q8_rows_g(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3027 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3028 dst: &mut CudaSlice<f32>, t: usize,
3029 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3030 qt: i32, rb: usize)
3031 -> Result<(), Box<dyn std::error::Error>> {
3032 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3033 n_expert as i32, rb as i64);
3034 let f = self.func("moe_down8_fma_dev_q8_rows_g");
3035 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, t as u32),
3036 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3037 let __s_b = self.gpu.stream();
3038 let mut b = __s_b.launch_builder(&f);
3039 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3040 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3041 unsafe { b.launch(cfg)?; }
3042 Ok(())
3043 }
3044
3045 pub fn rp_probe_q4(&self, m: usize) -> Result<(f64, f64), Box<dyn std::error::Error>> {
3049 let (out_f, in_f) = (2048usize, 2816usize);
3050 let nblk = in_f / 32;
3051 let mut seed = 0x9E3779B97F4A7C15u64;
3052 let mut rng = move || { seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407); (seed >> 33) as u8 };
3053 let mut w = vec![0u8; out_f * nblk * 18];
3054 for b in w.iter_mut() { *b = rng(); }
3055 for r in 0..out_f {
3056 for g in 0..nblk {
3057 let off = (r * nblk + g) * 18;
3058 w[off] = 0x00; w[off + 1] = 0x2C; }
3060 }
3061 let qplane = out_f * nblk * 16;
3062 let mut wrp = vec![0u8; w.len()];
3063 for r in 0..out_f {
3064 for g in 0..nblk {
3065 let src = &w[(r * nblk + g) * 18..(r * nblk + g) * 18 + 18];
3066 wrp[qplane + (r * nblk + g) * 2..qplane + (r * nblk + g) * 2 + 2]
3067 .copy_from_slice(&src[0..2]);
3068 wrp[(r * nblk + g) * 16..(r * nblk + g) * 16 + 16].copy_from_slice(&src[2..18]);
3069 }
3070 }
3071 let w_d = self.htod_bytes(&w)?;
3072 let wrp_d = self.htod_bytes(&wrp)?;
3073 let mut aq = vec![0i8; m * in_f];
3074 for v in aq.iter_mut() { *v = rng() as i8; }
3075 let aq_d = self.htod_i8(&aq)?;
3076 let ad_d = self.htod(&vec![0.03125f32; m * nblk])?;
3077 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
3078 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
3079 const RPB: u32 = 4;
3080 let cfg = LaunchConfig { grid_dim: ((out_f as u32).div_ceil(RPB), 1, 1),
3081 block_dim: (32, RPB, 1), shared_mem_bytes: 0 };
3082 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
3083 let (rb, qp) = ((nblk * 18) as i64, qplane as i64);
3084 let fb = self.func("qmatvec_q4_0_mmvq_b4");
3085 let fr = self.func("qmatvec_q4_0_mmvq_b4_rp");
3086 {
3087 let __s_b = self.gpu.stream();
3088 let mut b = __s_b.launch_builder(&fb);
3089 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3090 unsafe { b.launch(cfg)?; }
3091 let __s_b = self.gpu.stream();
3092 let mut b = __s_b.launch_builder(&fr);
3093 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1).arg(&inf).arg(&outf).arg(&mi).arg(&qp);
3094 unsafe { b.launch(cfg)?; }
3095 }
3096 self.gpu.stream().synchronize()?;
3097 let (h0, h1) = (self.dtoh(&y0)?, self.dtoh(&y1)?);
3098 let nd = h0.iter().zip(&h1).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
3099 if nd != 0 { return Err(format!("rp twin not bitwise: {nd}/{} diffs", h0.len()).into()); }
3100 let mut time = |rp: bool| -> Result<f64, Box<dyn std::error::Error>> {
3101 self.gpu.stream().synchronize()?;
3102 let t0 = std::time::Instant::now();
3103 for _ in 0..500 {
3104 if rp {
3105 let __s_b = self.gpu.stream();
3106 let mut b = __s_b.launch_builder(&fr);
3107 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1)
3108 .arg(&inf).arg(&outf).arg(&mi).arg(&qp);
3109 unsafe { b.launch(cfg)?; }
3110 } else {
3111 let __s_b = self.gpu.stream();
3112 let mut b = __s_b.launch_builder(&fb);
3113 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0)
3114 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3115 unsafe { b.launch(cfg)?; }
3116 }
3117 }
3118 self.gpu.stream().synchronize()?;
3119 Ok(t0.elapsed().as_secs_f64() * 1e6 / 500.0)
3120 };
3121 let _ = time(false)?; let _ = time(true)?; Ok((time(false)?, time(true)?))
3123 }
3124
3125 pub fn build_q4_rp4(&self, t: &mut crate::model::GpuTensor)
3130 -> Result<(), Box<dyn std::error::Error>> {
3131 use crate::model::GpuTensor;
3132 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3133 if *qtype != QT_Q4_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3134 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3135 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 18 { return Ok(()); }
3136 let nblk = in_f / 32;
3137 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 18)?;
3138 let f = self.func("q4_0_split_rp_build");
3139 let n = (out_f * nblk) as i32;
3140 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3141 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3142 let (of, nb) = (out_f as i32, nblk as i32);
3143 let _ = n;
3144 let __s_b = self.gpu.stream();
3145 let mut b = __s_b.launch_builder(&f);
3146 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3147 unsafe { b.launch(cfg)?; }
3148 *rp4 = Some(dst);
3149 Ok(())
3150 }
3151
3152 pub fn build_q8_rp4(&self, t: &mut crate::model::GpuTensor)
3157 -> Result<(), Box<dyn std::error::Error>> {
3158 use crate::model::GpuTensor;
3159 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3160 if *qtype != QT_Q8_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3161 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3162 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 34 { return Ok(()); }
3163 *rp4 = Some(self.build_q8_rp4_raw(bytes, in_f, out_f)?);
3164 Ok(())
3165 }
3166
3167 pub fn build_q8_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize)
3170 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3171 assert!(in_f % 32 == 0);
3172 let nblk = in_f / 32;
3173 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 34)?;
3174 let f = self.func("q8_0_split_rp_build");
3175 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3176 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3177 let (of, nb) = (out_f as i32, nblk as i32);
3178 let __s_b = self.gpu.stream();
3179 let mut b = __s_b.launch_builder(&f);
3180 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3181 unsafe { b.launch(cfg)?; }
3182 Ok(dst)
3183 }
3184
3185 pub fn build_q4k_rp4(&self, t: &mut crate::model::GpuTensor)
3193 -> Result<(), Box<dyn std::error::Error>> {
3194 use crate::model::GpuTensor;
3195 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3196 if *qtype != QT_Q4_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3197 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3198 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 144 { return Ok(()); }
3199 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q4_K)?);
3200 Ok(())
3201 }
3202
3203 pub fn build_q6k_rp4(&self, t: &mut crate::model::GpuTensor)
3204 -> Result<(), Box<dyn std::error::Error>> {
3205 use crate::model::GpuTensor;
3206 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3207 if *qtype != QT_Q6_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3208 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3209 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 210 { return Ok(()); }
3210 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q6_K)?);
3211 Ok(())
3212 }
3213
3214 pub fn build_kq_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize, qtype: i32)
3216 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3217 assert!(in_f % 256 == 0);
3218 let nsbk = in_f / 256;
3219 let (sb_bytes, kname) = match qtype {
3220 QT_Q4_K => (144usize, "q4_K_split_rp_build"),
3221 QT_Q6_K => (210usize, "q6_K_split_rp_build"),
3222 _ => return Err(format!("build_kq_rp4_raw: qtype {qtype} has no rp mirror").into()),
3223 };
3224 let mut dst = self.alloc_uninit::<u8>(out_f * nsbk * sb_bytes)?;
3225 let f = self.func(kname);
3226 let cfg = LaunchConfig { grid_dim: (((out_f * nsbk) as u32).div_ceil(256), 1, 1),
3227 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3228 let (of, nb) = (out_f as i32, nsbk as i32);
3229 let __s_b = self.gpu.stream();
3230 let mut b = __s_b.launch_builder(&f);
3231 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3232 unsafe { b.launch(cfg)?; }
3233 Ok(dst)
3234 }
3235
3236 pub fn kqrp_enabled() -> bool {
3240 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3241 *ON.get_or_init(|| match std::env::var("MEMRA_KQRP").as_deref() {
3242 Ok("0") => false,
3243 Ok(_) => true,
3244 Err(_) => cfg!(memra_hopper_mma),
3245 })
3246 }
3247
3248 pub fn build_q4_rp_swap(&self, t: &mut crate::model::GpuTensor)
3254 -> Result<bool, Box<dyn std::error::Error>> {
3255 self.build_q4_rp4(t)?;
3256 self.gpu.stream().synchronize()?; use crate::model::GpuTensor;
3258 let GpuTensor::Quant { bytes, rp4, rp, .. } = t else { return Ok(false) };
3259 match rp4.take() {
3260 Some(split) => {
3261 *bytes = split; *rp = true;
3263 Ok(true)
3264 }
3265 None => Ok(false),
3266 }
3267 }
3268
3269 pub fn q4rp_enabled() -> bool {
3271 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3272 *ON.get_or_init(|| std::env::var("MEMRA_Q4RP").map(|v| v != "0").unwrap_or(true))
3273 }
3274
3275 pub fn copy_rows_strided(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
3278 row_elems: usize, n_rows: usize, src_stride: usize, src_off: usize)
3279 -> Result<(), Box<dyn std::error::Error>> {
3280 let f = self.func("copy_rows_strided_f32");
3281 let cfg = LaunchConfig { grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
3282 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3283 let (re, nr) = (row_elems as i32, n_rows as i32);
3284 let (st, off) = (src_stride as i64, src_off as i64);
3285 let __s_b = self.gpu.stream();
3286 let mut b = __s_b.launch_builder(&f);
3287 b.arg(src).arg(&mut *dst).arg(&re).arg(&nr).arg(&st).arg(&off);
3288 unsafe { b.launch(cfg)?; }
3289 Ok(())
3290 }
3291
3292 pub fn u32_set_k(&self, dst: &mut CudaSlice<u32>, v: u32, idx: usize)
3294 -> Result<(), Box<dyn std::error::Error>> {
3295 let f = self.func("u32_set_k");
3296 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3297 let ii = idx as i32;
3298 let __s_b = self.gpu.stream();
3299 let mut b = __s_b.launch_builder(&f);
3300 b.arg(dst).arg(&v).arg(&ii);
3301 unsafe { b.launch(cfg)?; }
3302 Ok(())
3303 }
3304
3305 pub fn i32_add_k(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
3307 let f = self.func("i32_add_k");
3308 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3309 let __s_b = self.gpu.stream();
3310 let mut b = __s_b.launch_builder(&f);
3311 b.arg(d).arg(&v);
3312 unsafe { b.launch(cfg)?; }
3313 Ok(())
3314 }
3315
3316 pub fn i32_iota_from(&self, ctr: &CudaSlice<i32>, dst: &mut CudaSlice<i32>, n: usize)
3318 -> Result<(), Box<dyn std::error::Error>> {
3319 let f = self.func("i32_iota_from");
3320 let cfg = LaunchConfig::for_num_elems(n as u32);
3321 let ni = n as i32;
3322 let __s_b = self.gpu.stream();
3323 let mut b = __s_b.launch_builder(&f);
3324 b.arg(ctr).arg(dst).arg(&ni);
3325 unsafe { b.launch(cfg)?; }
3326 Ok(())
3327 }
3328
3329 pub fn u32_map_k(&self, buf: &mut CudaSlice<u32>, map: &CudaSlice<u32>, idx: usize)
3331 -> Result<(), Box<dyn std::error::Error>> {
3332 let f = self.func("u32_map_k");
3333 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3334 let ii = idx as i32;
3335 let __s_b = self.gpu.stream();
3336 let mut b = __s_b.launch_builder(&f);
3337 b.arg(buf).arg(map).arg(&ii);
3338 unsafe { b.launch(cfg)?; }
3339 Ok(())
3340 }
3341
3342 #[allow(clippy::too_many_arguments)]
3344 pub fn u32_pack2(&self, a: &CudaSlice<u32>, off_a: usize, n1: usize,
3345 b_in: &CudaSlice<u32>, n2: usize, out: &mut CudaSlice<u32>)
3346 -> Result<(), Box<dyn std::error::Error>> {
3347 let f = self.func("u32_pack2");
3348 let cfg = LaunchConfig::for_num_elems((n1 + n2) as u32);
3349 let (oa, i1, i2) = (off_a as i32, n1 as i32, n2 as i32);
3350 let __s_b = self.gpu.stream();
3351 let mut b = __s_b.launch_builder(&f);
3352 b.arg(a).arg(&oa).arg(&i1).arg(b_in).arg(&i2).arg(out);
3353 unsafe { b.launch(cfg)?; }
3354 Ok(())
3355 }
3356
3357 pub fn moe_w_exscale(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3359 s: &CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
3360 let f = self.func("moe_w_exscale");
3361 let cfg = LaunchConfig::for_num_elems(n as u32);
3362 let ni = n as i32;
3363 let __s_b = self.gpu.stream();
3364 let mut b = __s_b.launch_builder(&f);
3365 b.arg(w).arg(sel).arg(s).arg(&ni);
3366 unsafe { b.launch(cfg)?; }
3367 Ok(())
3368 }
3369
3370 pub fn moe_w_scale_by_expert(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3373 macros: &CudaSlice<f32>, n_expert: usize, n: usize)
3374 -> Result<(), Box<dyn std::error::Error>> {
3375 let f = self.func("moe_w_scale_by_expert");
3376 let cfg = LaunchConfig { grid_dim: (n.div_ceil(64) as u32, 1, 1),
3377 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
3378 let (ne, nn) = (n_expert as i32, n as i32);
3379 let __s_b = self.gpu.stream();
3380 let mut b = __s_b.launch_builder(&f);
3381 b.arg(w).arg(sel).arg(macros).arg(&ne).arg(&nn);
3382 unsafe { b.launch(cfg)?; }
3383 Ok(())
3384 }
3385
3386 pub fn moe_gate_up_silu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3387 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3388 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3389 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3390 macros: &CudaSlice<f32>)
3391 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3392 static GU: std::sync::OnceLock<(String, u32)> = std::sync::OnceLock::new();
3393 let (mode, wpb) = GU.get_or_init(|| {
3394 let mode = std::env::var("MEMRA_MOE_DEVQ8_GU").unwrap_or_default();
3395 let wpb = std::env::var("MEMRA_MOE_DEVQ8_WPB").ok()
3396 .and_then(|v| v.parse().ok()).unwrap_or(4u32).clamp(1, 16);
3397 (mode, wpb)
3398 });
3399 let (mode, wpb) = (mode.as_str(), *wpb);
3400 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3401 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3402 rb_g as i64, rb_u as i64);
3403 let (f, cfg) = match mode {
3404 "1" | "2" | "4" => {
3405 let rpw: u32 = mode.parse().unwrap();
3406 let f = self.func(match rpw { 1 => "moe_gate_up_silu8_dev_q8_r1",
3407 2 => "moe_gate_up_silu8_dev_q8_r2",
3408 _ => "moe_gate_up_silu8_dev_q8_r4" });
3409 let rows_per_block = (rpw * wpb) as usize;
3410 let gx = n_ff.div_ceil(rows_per_block) as u32;
3411 (f, LaunchConfig { grid_dim: (gx, n_used as u32, 1),
3412 block_dim: (32, wpb, 1), shared_mem_bytes: 0 })
3413 }
3414 "j8" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8"),
3415 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3416 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3417 "vsm2" => {
3419 let f = self.func("moe_gate_up_silu8_dev_q8_vsm2");
3420 let sh = (rb_g + rb_u) as u32;
3421 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3422 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3423 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3424 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3425 }
3426 "vsm" => {
3427 let f = self.func("moe_gate_up_silu8_dev_q8_vsm");
3428 let sh = (rb_g + rb_u) as u32;
3429 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3430 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3431 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3432 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3433 }
3434 "sg" => (self.func("moe_gate_up_silu8_dev_q8_sg"),
3435 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3436 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3437 "j8sg" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8sg"),
3438 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3439 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3440 "u64" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_u64"),
3441 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3442 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3443 "gs4" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_gs4"),
3444 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3445 block_dim: (32, 4, 1), shared_mem_bytes: 0 }),
3446 "v" | "" => (self.func("moe_gate_up_silu8_dev_q8_v"),
3448 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3449 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3450 "s2" => (self.func("moe_gate_up_silu8_dev_q8_s2"),
3451 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3452 block_dim: (32, 2, 1), shared_mem_bytes: 0 }),
3453 "s2z" => {
3454 let rz = wpb.min(16); (self.func("moe_gate_up_silu8_dev_q8_s2z"),
3456 LaunchConfig { grid_dim: (n_ff.div_ceil(rz as usize) as u32, n_used as u32, 1),
3457 block_dim: (32, 2, rz), shared_mem_bytes: 0 })
3458 }
3459 _ => (self.func("moe_gate_up_silu8_dev_q8"),
3460 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3461 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3462 };
3463 let __s_b = self.gpu.stream();
3464 let mut b = __s_b.launch_builder(&f);
3465 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3466 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3467 unsafe { b.launch(cfg)?; }
3468 Ok(act)
3469 }
3470
3471 #[allow(clippy::too_many_arguments)]
3472 pub fn moe_down8_fma_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3473 w: &cudarc::driver::CudaView<f32>,
3474 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3475 dst: &mut cudarc::driver::CudaViewMut<f32>,
3476 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3477 qt: i32, rb: usize)
3478 -> Result<(), Box<dyn std::error::Error>> {
3479 static DOWN: std::sync::OnceLock<String> = std::sync::OnceLock::new();
3480 let mode = DOWN.get_or_init(|| std::env::var("MEMRA_MOE_DEVQ8_DOWN").unwrap_or_default());
3481 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3482 n_expert as i32, rb as i64);
3483 let (f, cfg) = match mode.as_str() {
3486 m @ ("1" | "2" | "4") if n_used <= 8 => {
3487 let rpw: usize = m.parse().unwrap();
3488 let f = self.func(match rpw { 1 => "moe_down8_fma_dev_q8_w8r1",
3489 2 => "moe_down8_fma_dev_q8_w8r2",
3490 _ => "moe_down8_fma_dev_q8_w8r4" });
3491 (f, LaunchConfig { grid_dim: (out_f.div_ceil(rpw) as u32, 1, 1),
3492 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3493 }
3494 "h2" if in_f == 512 => (self.func("moe_down8_fma_dev_q8_h2"),
3495 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3496 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3497 "" if in_f == 704 && n_used <= 8 =>
3500 (self.func("moe_down8_fma_dev_q8_w8r2"),
3501 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3502 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3503 "w8h2v" | "" if in_f == 512 && n_used <= 8 =>
3507 (self.func("moe_down8_fma_dev_q8_w8h2v"),
3508 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3509 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3510 "w8h2r2v" if in_f == 512 && n_used <= 8 =>
3511 (self.func("moe_down8_fma_dev_q8_w8h2r2v"),
3512 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3513 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3514 "w8h2r2" if in_f == 512 && n_used <= 8 =>
3515 (self.func("moe_down8_fma_dev_q8_w8h2r2"),
3516 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3517 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3518 "w8h2" if in_f == 512 && n_used <= 8 =>
3519 (self.func("moe_down8_fma_dev_q8_w8h2"),
3520 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3521 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
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)]
3541 pub fn moe_gate_up_silu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3542 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
3543 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3544 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3545 macros: &CudaSlice<f32>)
3546 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3547 let f = self.func("moe_gate_up_silu8_dev_q8_v_rows");
3548 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
3549 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
3550 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3551 let (inf, nff, ne, nu, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3552 n_used as i32, rb_g as i64, rb_u as i64);
3553 let __s_b = self.gpu.stream();
3554 let mut b = __s_b.launch_builder(&f);
3555 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3556 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(macros);
3557 unsafe { b.launch(cfg)?; }
3558 Ok(act)
3559 }
3560
3561 #[allow(clippy::too_many_arguments)]
3566 pub fn moe_down8_fma_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3567 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3568 dst: &mut CudaSlice<f32>, t: usize,
3569 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3570 qt: i32, rb: usize)
3571 -> Result<(), Box<dyn std::error::Error>> {
3572 assert!(in_f == 512 && n_used <= 8, "down rows twin is w8h2v shape-gated");
3573 let f = self.func("moe_down8_fma_dev_q8_w8h2v_rows");
3574 let cfg = LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, t as u32),
3575 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 };
3576 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3577 n_expert as i32, rb as i64);
3578 let __s_b = self.gpu.stream();
3579 let mut b = __s_b.launch_builder(&f);
3580 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3581 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3582 unsafe { b.launch(cfg)?; }
3583 Ok(())
3584 }
3585
3586 #[allow(clippy::too_many_arguments)]
3590 pub fn moe_gate_up_silu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3591 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3592 n_pairs: usize, in_f: usize, n_ff: usize, n_used: usize,
3593 n_expert: usize, qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3594 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3595 let f = self.func("moe_gate_up_silu8_dev_q8_csr_iq4");
3596 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
3597 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
3598 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3599 let (inf, nff, ne, nu, npi, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3600 n_used as i32, n_pairs as i32, rb_g as i64, rb_u as i64);
3601 let __s_b = self.gpu.stream();
3602 let mut b = __s_b.launch_builder(&f);
3603 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3604 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
3605 unsafe { b.launch(cfg)?; }
3606 Ok(act)
3607 }
3608
3609
3610 #[allow(clippy::too_many_arguments)]
3614 pub fn moe_down8_fma_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3615 sel: &cudarc::driver::CudaView<i32>,
3616 w: &cudarc::driver::CudaView<f32>,
3617 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3618 dst: &mut cudarc::driver::CudaViewMut<f32>,
3619 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3620 qt: i32, rb: usize)
3621 -> Result<(), Box<dyn std::error::Error>> {
3622 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3623 n_expert as i32, rb as i64);
3624 let (f, cfg) = match variant {
3625 "w8h2" | "w8h2v" => {
3626 (self.func(if variant == "w8h2" { "moe_down8_fma_dev_q8_w8h2" }
3627 else { "moe_down8_fma_dev_q8_w8h2v" }),
3628 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3629 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3630 }
3631 "w8h2r2" | "w8h2r2v" => {
3632 (self.func(if variant == "w8h2r2" { "moe_down8_fma_dev_q8_w8h2r2" }
3633 else { "moe_down8_fma_dev_q8_w8h2r2v" }),
3634 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3635 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3636 }
3637 _ => (self.func("moe_down8_fma_dev_q8"),
3638 LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3639 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3640 };
3641 let __s_b = self.gpu.stream();
3642 let mut b = __s_b.launch_builder(&f);
3643 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3644 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3645 unsafe { b.launch(cfg)?; }
3646 Ok(())
3647 }
3648
3649 #[allow(clippy::too_many_arguments)]
3651 pub fn moe_gate_up_silu8_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3652 sel: &cudarc::driver::CudaView<i32>,
3653 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3654 in_f: usize, n_ff: usize, n_used: usize,
3655 n_expert: usize, qt_g: i32, qt_u: i32,
3656 rb_g: usize, rb_u: usize)
3657 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3658 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3659 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3660 rb_g as i64, rb_u as i64);
3661 let f = self.func(if variant == "v" { "moe_gate_up_silu8_dev_q8_v" }
3662 else { "moe_gate_up_silu8_dev_q8" });
3663 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3664 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3665 let __s_b = self.gpu.stream();
3666 let mut b = __s_b.launch_builder(&f);
3667 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3668 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
3669 unsafe { b.launch(cfg)?; }
3670 Ok(act)
3671 }
3672
3673 pub fn moe_gate_up_silu8_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3674 x: &cudarc::driver::CudaView<f32>,
3675 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3676 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3677 macros: &CudaSlice<f32>)
3678 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3679 let f = self.func("moe_gate_up_silu8_dev");
3680 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),
3682 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3683 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3684 rb_g as i64, rb_u as i64);
3685 let __s_b = self.gpu.stream();
3686 let mut b = __s_b.launch_builder(&f);
3687 b.arg(table).arg(sel).arg(x).arg(&mut act)
3688 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3689 unsafe { b.launch(cfg)?; }
3690 Ok(act)
3691 }
3692
3693 #[allow(clippy::too_many_arguments)]
3696 pub fn moe_down8_fma_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3697 w: &cudarc::driver::CudaView<f32>, act: &CudaSlice<f32>,
3698 dst: &mut cudarc::driver::CudaViewMut<f32>,
3699 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3700 qt: i32, rb: usize)
3701 -> Result<(), Box<dyn std::error::Error>> {
3702 let f = self.func("moe_down8_fma_dev");
3703 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3704 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3705 let (inf, outf, nu, ne, rbv) = (in_f as i32, out_f as i32, n_used as i32,
3706 n_expert as i32, rb as i64);
3707 let __s_b = self.gpu.stream();
3708 let mut b = __s_b.launch_builder(&f);
3709 b.arg(table).arg(sel).arg(w).arg(act).arg(dst)
3710 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbv);
3711 unsafe { b.launch(cfg)?; }
3712 Ok(())
3713 }
3714
3715 pub fn axpy_into(&self, src: &CudaSlice<f32>, alpha: f32,
3717 dst: &mut cudarc::driver::CudaViewMut<f32>, n: usize)
3718 -> Result<(), Box<dyn std::error::Error>> {
3719 let f = self.func("axpy_f32");
3720 let cfg = LaunchConfig::for_num_elems(n as u32);
3721 let (a, ni) = (alpha, n as i32);
3722 let __s_b = self.gpu.stream();
3723 let mut b = __s_b.launch_builder(&f);
3724 b.arg(src).arg(dst).arg(&a).arg(&ni);
3725 unsafe { b.launch(cfg)?; }
3726 Ok(())
3727 }
3728
3729 pub fn add_scaled_rows(&self, src: &CudaSlice<f32>, scale: &CudaSlice<f32>,
3731 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
3732 -> Result<(), Box<dyn std::error::Error>> {
3733 let f = self.func("add_scaled_rows_f32");
3734 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
3735 let (nc, nr) = (ncols as i32, nrows as i32);
3736 let __s_b = self.gpu.stream();
3737 let mut b = __s_b.launch_builder(&f);
3738 b.arg(src).arg(scale).arg(dst).arg(&nc).arg(&nr);
3739 unsafe { b.launch(cfg)?; }
3740 Ok(())
3741 }
3742
3743 pub fn gather_rows(&self, src: &CudaSlice<f32>, idx: &CudaSlice<i32>,
3747 dst: &mut CudaSlice<f32>, ncols: usize, m_e: usize)
3748 -> Result<(), Box<dyn std::error::Error>> {
3749 let f = self.func("gather_rows_f32");
3750 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3751 let (nc, me) = (ncols as i32, m_e as i32);
3752 let __s_b = self.gpu.stream();
3753 let mut b = __s_b.launch_builder(&f);
3754 b.arg(src).arg(idx).arg(dst).arg(&nc).arg(&me);
3755 unsafe { b.launch(cfg)?; }
3756 Ok(())
3757 }
3758
3759 pub fn scatter_slot(&self, src: &CudaSlice<f32>, tok_idx: &CudaSlice<i32>,
3764 slot_idx: &CudaSlice<i32>, weight: &CudaSlice<f32>,
3765 dst: &mut CudaSlice<f32>, wbuf: &mut CudaSlice<f32>,
3766 ncols: usize, n_used: usize, m_e: usize)
3767 -> Result<(), Box<dyn std::error::Error>> {
3768 let f = self.func("scatter_add_slot_f32");
3769 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3770 let (nc, nu, me) = (ncols as i32, n_used as i32, m_e as i32);
3771 let __s_b = self.gpu.stream();
3772 let mut b = __s_b.launch_builder(&f);
3773 b.arg(src).arg(tok_idx).arg(slot_idx).arg(weight).arg(dst).arg(wbuf).arg(&nc).arg(&nu).arg(&me);
3774 unsafe { b.launch(cfg)?; }
3775 Ok(())
3776 }
3777
3778 pub fn reduce_slots(&self, slots: &CudaSlice<f32>, wbuf: &CudaSlice<f32>,
3782 dst: &mut CudaSlice<f32>, ncols: usize, n_used: usize, t: usize)
3783 -> Result<(), Box<dyn std::error::Error>> {
3784 let f = self.func("reduce_slots_f32");
3785 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
3786 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
3787 let __s_b = self.gpu.stream();
3788 let mut b = __s_b.launch_builder(&f);
3789 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
3790 unsafe { b.launch(cfg)?; }
3791 Ok(())
3792 }
3793
3794 pub fn quantize_q8_1_view(&self, x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize)
3801 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3802 let f = self.func("quantize_q8_1");
3803 let nblk = in_f / 32;
3804 let mut q = self.alloc_uninit::<i8>(m * in_f)?;
3805 let mut d = self.alloc_uninit::<f32>(m * nblk)?;
3806 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
3807 let (inf, mi) = (in_f as i32, m as i32);
3808 let __s_b = self.gpu.stream();
3809 let mut b = __s_b.launch_builder(&f);
3810 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3811 unsafe { b.launch(cfg)?; }
3812 Ok((q, d))
3813 }
3814
3815 pub fn quantize_q8_1(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3816 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3817 let nblk = in_f / 32;
3818 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);
3822 let (inf, mi) = (in_f as i32, m as i32);
3823 if Self::pdl_on() && Self::pdl_wb_on() {
3824 {
3825 use cudarc::driver::{DevicePtr, DevicePtrMut};
3826 let s = &self.gpu.stream();
3827 let (px, _g0) = x.device_ptr(s);
3828 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
3829 let mut ps = [
3830 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
3831 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
3832 &mi as *const _ as *mut _,
3833 ];
3834 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
3835 }
3836 return Ok((q, d));
3837 }
3838 let f = self.func("quantize_q8_1");
3839 let __s_b = self.gpu.stream();
3840 let mut b = __s_b.launch_builder(&f);
3841 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3842 unsafe { b.launch(cfg)?; }
3843 Ok((q, d))
3844 }
3845
3846 pub fn quantize_fp4_act(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3850 -> Result<(CudaSlice<u32>, CudaSlice<u8>), Box<dyn std::error::Error>> {
3851 let f = self.func("quantize_fp4_act");
3852 let nb16 = in_f / 16;
3853 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);
3856 let (inf, mi) = (in_f as i32, m as i32);
3857 let __s_b = self.gpu.stream();
3858 let mut b = __s_b.launch_builder(&f);
3859 b.arg(x).arg(&mut aq4).arg(&mut ad4).arg(&inf).arg(&mi);
3860 unsafe { b.launch(cfg)?; }
3861 Ok((aq4, ad4))
3862 }
3863
3864 pub fn qmatvec_gemm_nvfp4_fp4(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3869 in_f: usize, out_f: usize, row_bytes: usize, scale: f32)
3870 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3871 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
3872 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
3873 let mut y = self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)?;
3874 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
3875 Ok(y)
3876 }
3877
3878 fn fp4_gemm_launch(&self, bytes: &CudaSlice<u8>, aq4: &CudaSlice<u32>, ad4: &CudaSlice<u8>,
3881 m: usize, in_f: usize, out_f: usize, row_bytes: usize)
3882 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3883 let f = self.func("qmatvec_gemm_nvfp4_fp4");
3884 let mut y = self.alloc_uninit::<f32>(m * out_f)?; const BM: u32 = 64; const BN: u32 = 256;
3886 let cfg = LaunchConfig {
3887 grid_dim: ((out_f as u32 + BM - 1) / BM, (m as u32 + BN - 1) / BN, 1),
3888 block_dim: (32, 4, 1), shared_mem_bytes: 0,
3889 };
3890 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3891 let __s_b = self.gpu.stream();
3892 let mut b = __s_b.launch_builder(&f);
3893 b.arg(bytes).arg(aq4).arg(ad4).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3894 unsafe { b.launch(cfg)?; }
3895 Ok(y)
3896 }
3897
3898 pub fn qmatvec_gemm_nvfp4_fp4_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3900 in_f: usize, out_f: usize, row_bytes: usize)
3901 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3902 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
3903 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
3904 self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)
3905 }
3906
3907 pub fn qmatvec_q8_0_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3909 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3910 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3911 let f = self.func("qmatvec_q8_0_dp4a");
3912 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 };
3914 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3915 let __s_b = self.gpu.stream();
3916 let mut b = __s_b.launch_builder(&f);
3917 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3918 unsafe { b.launch(cfg)?; }
3919 Ok(y)
3920 }
3921
3922 #[allow(non_snake_case)] pub fn qmatvec_q4_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3925 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3926 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3927 let f = self.func("qmatvec_q4_K_dp4a");
3928 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 };
3930 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3931 let __s_b = self.gpu.stream();
3932 let mut b = __s_b.launch_builder(&f);
3933 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3934 unsafe { b.launch(cfg)?; }
3935 Ok(y)
3936 }
3937
3938 #[allow(non_snake_case)] pub fn qmatvec_q6_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3941 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3942 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3943 let f = self.func("qmatvec_q6_K_dp4a");
3944 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 };
3946 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3947 let __s_b = self.gpu.stream();
3948 let mut b = __s_b.launch_builder(&f);
3949 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3950 unsafe { b.launch(cfg)?; }
3951 Ok(y)
3952 }
3953
3954 #[allow(non_snake_case)] pub fn qmatvec_q5_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3957 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3958 self.qmatvec_dp4a_named("qmatvec_q5_K_dp4a", w, x, m, in_f, out_f, row_bytes)
3959 }
3960 #[allow(non_snake_case)] pub fn qmatvec_q3_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3963 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3964 self.qmatvec_dp4a_named("qmatvec_q3_K_dp4a", w, x, m, in_f, out_f, row_bytes)
3965 }
3966 pub fn qmatvec_nvfp4_fast_rp(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3968 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3969 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
3970 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a_rp", w, x, m, in_f, out_f, row_bytes)
3971 }
3972 pub fn qmatvec_nvfp4_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3974 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3975 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
3978 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a", w, x, m, in_f, out_f, row_bytes)
3979 }
3980 #[allow(non_snake_case)] pub fn qmatvec_iq4_XS_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3983 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3984 self.qmatvec_dp4a_named("qmatvec_iq4_XS_dp4a", w, x, m, in_f, out_f, row_bytes)
3985 }
3986
3987 fn qmatvec_dp4a_named(&self, name: &str, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3989 in_f: usize, out_f: usize, row_bytes: usize)
3990 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3991 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3992 let f = self.func(name);
3993 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 };
3995 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3996 let __s_b = self.gpu.stream();
3997 let mut b = __s_b.launch_builder(&f);
3998 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3999 unsafe { b.launch(cfg)?; }
4000 Ok(y)
4001 }
4002
4003 pub fn htod(&self, v: &[f32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4004 Ok(self.gpu.stream().clone_htod(v)?)
4005 }
4006 pub fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
4007 Ok(self.gpu.stream().clone_htod(v)?)
4008 }
4009 pub fn htod_i8(&self, v: &[i8]) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
4011 Ok(self.gpu.stream().clone_htod(v)?)
4012 }
4013 pub fn htod_u64(&self, v: &[u64]) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
4014 Ok(self.gpu.stream().clone_htod(v)?)
4015 }
4016 pub fn dtoh_view(&self, d: &cudarc::driver::CudaView<f32>)
4018 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
4019 let v = self.gpu.stream().clone_dtoh(d)?;
4020 self.gpu.stream().synchronize()?;
4021 Ok(v)
4022 }
4023 pub fn dtoh(&self, d: &CudaSlice<f32>) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
4024 let v = self.gpu.stream().clone_dtoh(d)?;
4025 self.gpu.stream().synchronize()?;
4026 Ok(v)
4027 }
4028 pub fn dtoh_pair(
4032 &self,
4033 a: &CudaSlice<f32>,
4034 b: &CudaSlice<f32>,
4035 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
4036 let av = self.gpu.stream().clone_dtoh(a)?;
4037 let bv = self.gpu.stream().clone_dtoh(b)?;
4038 self.gpu.stream().synchronize()?;
4039 Ok((av, bv))
4040 }
4041 pub fn dtoh_i32(&self, d: &CudaSlice<i32>) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
4043 let v = self.gpu.stream().clone_dtoh(d)?;
4044 self.gpu.stream().synchronize()?;
4045 Ok(v)
4046 }
4047 pub fn dtoh_u8(&self, d: &CudaSlice<u8>) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
4049 let v = self.gpu.stream().clone_dtoh(d)?;
4050 self.gpu.stream().synchronize()?;
4051 Ok(v)
4052 }
4053 pub fn dtoh_u8_view(&self, d: &cudarc::driver::CudaView<u8>)
4054 -> Result<Vec<u8>, Box<dyn std::error::Error>> {
4055 let v = self.gpu.stream().clone_dtoh(d)?;
4056 self.gpu.stream().synchronize()?;
4057 Ok(v)
4058 }
4059 pub fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4060 let s = self.gpu.stream().alloc_zeros::<f32>(n)?;
4061 self.keep_if_capturing(&s);
4062 Ok(s)
4063 }
4064
4065 pub fn prob_of_token_device(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>, n_vocab: usize)
4074 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4075 let nb = ARGMAX_NB;
4076 let mut part = self.alloc_uninit::<f32>(nb)?;
4077 let mut p = self.alloc_uninit::<f32>(1)?;
4078 let f1 = self.func("prob_of_token_partial_f32");
4079 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4080 let nv = n_vocab as i32;
4081 let __s_b1 = self.gpu.stream();
4082 let mut b1 = __s_b1.launch_builder(&f1);
4083 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
4084 unsafe { b1.launch(cfg1)?; }
4085 let f2 = self.func("prob_of_token_final_f32");
4086 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4087 let nbi = nb as i32;
4088 let __s_b2 = self.gpu.stream();
4089 let mut b2 = __s_b2.launch_builder(&f2);
4090 b2.arg(&part).arg(&mut p).arg(&nbi);
4091 unsafe { b2.launch(cfg2)?; }
4092 Ok(p)
4093 }
4094
4095 pub fn prob_of_token_device_col(&self, logits: &CudaSlice<f32>,
4102 tok_all: &CudaSlice<u32>, tok_idx: usize,
4103 p_out: &mut CudaSlice<f32>, p_idx: usize, n_vocab: usize)
4104 -> Result<(), Box<dyn std::error::Error>> {
4105 let tok_v = tok_all.slice(tok_idx..tok_idx + 1);
4106 let mut p_v = p_out.slice_mut(p_idx..p_idx + 1);
4107 let nb = ARGMAX_NB;
4108 let mut part = self.alloc_uninit::<f32>(nb)?;
4109 let f1 = self.func("prob_of_token_partial_f32");
4110 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4111 let nv = n_vocab as i32;
4112 let __s_b1 = self.gpu.stream();
4113 let mut b1 = __s_b1.launch_builder(&f1);
4114 b1.arg(logits).arg(&tok_v).arg(&mut part).arg(&nv);
4115 unsafe { b1.launch(cfg1)?; }
4116 let f2 = self.func("prob_of_token_final_f32");
4117 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4118 let nbi = nb as i32;
4119 let __s_b2 = self.gpu.stream();
4120 let mut b2 = __s_b2.launch_builder(&f2);
4121 b2.arg(&part).arg(&mut p_v).arg(&nbi);
4122 unsafe { b2.launch(cfg2)?; }
4123 Ok(())
4124 }
4125
4126 pub fn prob_of_token_device_into(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>,
4127 p_out: &mut CudaSlice<f32>, n_vocab: usize)
4128 -> Result<(), Box<dyn std::error::Error>> {
4129 let nb = ARGMAX_NB;
4130 let mut part = self.alloc_uninit::<f32>(nb)?;
4131 let f1 = self.func("prob_of_token_partial_f32");
4132 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4133 let nv = n_vocab as i32;
4134 let __s_b1 = self.gpu.stream();
4135 let mut b1 = __s_b1.launch_builder(&f1);
4136 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
4137 unsafe { b1.launch(cfg1)?; }
4138 let f2 = self.func("prob_of_token_final_f32");
4139 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4140 let nbi = nb as i32;
4141 let __s_b2 = self.gpu.stream();
4142 let mut b2 = __s_b2.launch_builder(&f2);
4143 b2.arg(&part).arg(p_out).arg(&nbi);
4144 unsafe { b2.launch(cfg2)?; }
4145 Ok(())
4146 }
4147
4148 pub fn argmax_token_device(&self, logits: &CudaSlice<f32>, n_vocab: usize)
4149 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4150 let mut tok = unsafe { self.gpu.stream().alloc::<u32>(1)? };
4151 self.argmax_token_device_into(logits, &mut tok, n_vocab)?;
4152 Ok(tok)
4153 }
4154 pub fn argmax_token_device_into(&self, logits: &CudaSlice<f32>, tok: &mut CudaSlice<u32>,
4161 n_vocab: usize) -> Result<(), Box<dyn std::error::Error>> {
4162 let nb = ARGMAX_NB;
4163 let f1 = self.func("argmax_partial_f32");
4164 let f2 = self.func("argmax_final_f32");
4165 let mut guard = self.argmax_partials.lock().unwrap();
4166 if guard.is_none() {
4167 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4170 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4171 *guard = Some((pv, pi));
4172 }
4173 let (part_v, part_i) = guard.as_mut().unwrap();
4174 let nv = n_vocab as i32;
4175 let nbi = nb as i32;
4176 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4178 let __s_b1 = self.gpu.stream();
4179 let mut b1 = __s_b1.launch_builder(&f1);
4180 b1.arg(logits).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4181 unsafe { b1.launch(cfg1)?; }
4182 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4184 let __s_b2 = self.gpu.stream();
4185 let mut b2 = __s_b2.launch_builder(&f2);
4186 b2.arg(&*part_v).arg(&*part_i).arg(tok).arg(&nbi);
4187 unsafe { b2.launch(cfg2)?; }
4188 Ok(())
4189 }
4190 pub fn argmax_token_device_col(&self, logits: &CudaSlice<f32>, col: usize, n_vocab: usize,
4196 toks: &mut CudaSlice<u32>, out_idx: usize)
4197 -> Result<(), Box<dyn std::error::Error>> {
4198 let nb = ARGMAX_NB;
4199 let f1 = self.func("argmax_partial_f32");
4200 let f2 = self.func("argmax_final_f32");
4201 let mut guard = self.argmax_partials.lock().unwrap();
4202 if guard.is_none() {
4203 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4204 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4205 *guard = Some((pv, pi));
4206 }
4207 let (part_v, part_i) = guard.as_mut().unwrap();
4208 let col_view = logits.slice(col * n_vocab..(col + 1) * n_vocab);
4209 let nv = n_vocab as i32;
4210 let nbi = nb as i32;
4211 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4212 let __s_b1 = self.gpu.stream();
4213 let mut b1 = __s_b1.launch_builder(&f1);
4214 b1.arg(&col_view).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4215 unsafe { b1.launch(cfg1)?; }
4216 let mut tok_view = toks.slice_mut(out_idx..out_idx + 1);
4217 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4218 let __s_b2 = self.gpu.stream();
4219 let mut b2 = __s_b2.launch_builder(&f2);
4220 b2.arg(&*part_v).arg(&*part_i).arg(&mut tok_view).arg(&nbi);
4221 unsafe { b2.launch(cfg2)?; }
4222 Ok(())
4223 }
4224 pub fn htod_u32_v(&self, v: &[u32]) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4226 Ok(self.gpu.stream().clone_htod(v)?)
4227 }
4228 pub fn dtoh_u32(&self, d: &CudaSlice<u32>) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
4229 let v = self.gpu.stream().clone_dtoh(d)?;
4230 self.gpu.stream().synchronize()?;
4231 Ok(v)
4232 }
4233 pub fn htod_u32_into(&self, dst: &mut CudaSlice<u32>, src: &[u32])
4237 -> Result<(), Box<dyn std::error::Error>> {
4238 let mut view = dst.slice_mut(0..src.len());
4239 self.gpu.stream().memcpy_htod(src, &mut view)?;
4240 Ok(())
4241 }
4242
4243 pub fn htod_i32_into(&self, dst: &mut CudaSlice<i32>, src: &[i32])
4246 -> Result<(), Box<dyn std::error::Error>> {
4247 let mut view = dst.slice_mut(0..src.len());
4248 self.gpu.stream().memcpy_htod(src, &mut view)?;
4249 Ok(())
4250 }
4251
4252 pub fn alloc_u32_zeroed(&self, n: usize) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4253 let s = self.gpu.stream().alloc_zeros::<u32>(n)?;
4254 self.keep_if_capturing(&s);
4255 Ok(s)
4256 }
4257 pub fn embed_gather_device_into(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4260 x_out: &mut CudaSlice<f32>, n_embd: usize, qtype: i32,
4261 row_bytes: usize) -> Result<(), Box<dyn std::error::Error>> {
4262 let f = self.func("embed_gather_u32");
4263 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
4264 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4265 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
4266 let __s_b = self.gpu.stream();
4267 let mut b = __s_b.launch_builder(&f);
4268 b.arg(embd).arg(token_d).arg(x_out).arg(&ne).arg(&qt).arg(&rb);
4269 unsafe { b.launch(cfg)?; }
4270 Ok(())
4271 }
4272 pub fn dtoh_i32_one(&self, d: &CudaSlice<i32>) -> Result<i32, Box<dyn std::error::Error>> {
4274 let v = self.gpu.stream().clone_dtoh(d)?;
4275 self.gpu.stream().synchronize()?;
4276 Ok(v[0])
4277 }
4278 pub fn i32_set_k(&self, dst: &mut CudaSlice<i32>, v: i32)
4285 -> Result<(), Box<dyn std::error::Error>> {
4286 let f = self.func("i32_set_k");
4287 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
4288 let idx = 0i32;
4289 let __s_b = self.gpu.stream();
4290 let mut b = __s_b.launch_builder(&f);
4291 b.arg(dst).arg(&v).arg(&idx);
4292 unsafe { b.launch(cfg)?; }
4293 Ok(())
4294 }
4295
4296 pub fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
4297 self.gpu.stream().memcpy_htod(&[v], d)?;
4298 Ok(())
4299 }
4300 pub fn set_u32_one(&self, d: &mut CudaSlice<u32>, v: u32) -> Result<(), Box<dyn std::error::Error>> {
4303 self.gpu.stream().memcpy_htod(&[v], d)?;
4304 Ok(())
4305 }
4306 pub fn dtoh_u32_one(&self, d: &CudaSlice<u32>) -> Result<u32, Box<dyn std::error::Error>> {
4308 let v = self.gpu.stream().clone_dtoh(d)?;
4309 self.gpu.stream().synchronize()?;
4310 Ok(v[0])
4311 }
4312 pub fn upload_u8(&self, bytes: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
4314 Ok(self.gpu.stream().clone_htod(bytes)?)
4315 }
4316 pub fn embed_gather_device(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4320 n_embd: usize, qtype: i32, row_bytes: usize)
4321 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4322 let f = self.func("embed_gather_u32");
4323 let mut x = self.alloc_uninit::<f32>(n_embd)?;
4324 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
4325 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4326 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
4327 let __s_b = self.gpu.stream();
4328 let mut b = __s_b.launch_builder(&f);
4329 b.arg(embd).arg(token_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb);
4330 unsafe { b.launch(cfg)?; }
4331 Ok(x)
4332 }
4333
4334
4335 pub fn embed_gather_device_t(&self, embd: &CudaSlice<u8>, tokens: &[u32],
4339 n_embd: usize, qtype: i32, row_bytes: usize)
4340 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4341 let t = tokens.len();
4342 let tok_d = self.gpu.stream().clone_htod(tokens)?;
4343 let f = self.func("embed_gather_u32_t");
4344 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4345 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4346 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4347 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4348 let __s_b = self.gpu.stream();
4349 let mut b = __s_b.launch_builder(&f);
4350 b.arg(embd).arg(&tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4351 unsafe { b.launch(cfg)?; }
4352 Ok(x)
4353 }
4354
4355 pub fn embed_gather_device_tv(&self, embd: &CudaSlice<u8>, tok_v: &cudarc::driver::CudaView<u32>,
4360 t: usize, n_embd: usize, qtype: i32, row_bytes: usize)
4361 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4362 let f = self.func("embed_gather_u32_t");
4363 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4364 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4365 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4366 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4367 let __s_b = self.gpu.stream();
4368 let mut b = __s_b.launch_builder(&f);
4369 b.arg(embd).arg(tok_v).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4370 unsafe { b.launch(cfg)?; }
4371 Ok(x)
4372 }
4373
4374 pub fn embed_gather_device_td(&self, embd: &CudaSlice<u8>, tok_d: &CudaSlice<u32>, t: usize,
4375 n_embd: usize, qtype: i32, row_bytes: usize)
4376 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4377 let f = self.func("embed_gather_u32_t");
4378 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4379 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4380 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4381 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4382 let __s_b = self.gpu.stream();
4383 let mut b = __s_b.launch_builder(&f);
4384 b.arg(embd).arg(tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4385 unsafe { b.launch(cfg)?; }
4386 Ok(x)
4387 }
4388
4389 #[inline]
4395 fn keep_if_capturing<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, s: &CudaSlice<T>) {
4397 if self.capture_keep_on.load(std::sync::atomic::Ordering::Relaxed) {
4398 self.capture_keep.lock().unwrap().push(Box::new(s.clone()));
4399 }
4400 }
4401
4402 fn alloc_uninit<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, n: usize)
4403 -> Result<CudaSlice<T>, Box<dyn std::error::Error>> {
4404 let mut s = unsafe { self.gpu.stream().alloc::<T>(n)? };
4405 {
4409 static Z: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4410 if *Z.get_or_init(|| std::env::var("MEMRA_DEBUG_ZERO_ALLOCS").as_deref() == Ok("1")) {
4411 use cudarc::driver::DevicePtrMut;
4413 let n_bytes = s.len() * std::mem::size_of::<T>();
4414 let stream = self.gpu.stream();
4415 let (p_, _g) = s.device_ptr_mut(&stream);
4416 unsafe {
4417 cudarc::driver::sys::cuMemsetD8Async(p_, 0, n_bytes, stream.cu_stream())
4418 .result()?;
4419 }
4420 }
4421 }
4422 self.keep_if_capturing(&s);
4423 Ok(s)
4424 }
4425
4426 pub fn uninit_q8_pair(&self, n: usize)
4431 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4432 Ok((self.alloc_uninit::<i8>(n)?, self.alloc_uninit::<f32>(n / 32)?))
4433 }
4434
4435 pub fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4436 self.alloc_uninit::<f32>(n)
4437 }
4438
4439 pub fn alloc_i8_uninit(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
4441 self.alloc_uninit::<i8>(n)
4442 }
4443
4444 #[allow(clippy::too_many_arguments)]
4448 pub fn rms_norm3(&self, x: &CudaSlice<f32>, w0: &CudaSlice<f32>, w1: &CudaSlice<f32>,
4449 w2: &CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4450 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4451 -> Result<(), Box<dyn std::error::Error>> {
4452 let f = self.func("rms_norm3_f32");
4453 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4454 let (nc, e) = (ncols as i32, eps);
4455 let __s_b = self.gpu.stream();
4456 let mut b = __s_b.launch_builder(&f);
4457 b.arg(x).arg(w0).arg(w1).arg(w2).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e);
4458 unsafe { b.launch(cfg)?; }
4459 Ok(())
4460 }
4461
4462 #[allow(clippy::too_many_arguments)]
4464 pub fn qkvnorm_w_on_prefill(rows: usize, ncols: usize) -> bool {
4467 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4468 *WARP_ON.get_or_init(|| {
4469 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4470 }) && ncols % 4 == 0 && rows >= 64
4471 }
4472
4473 #[allow(clippy::too_many_arguments)]
4476 pub fn rms_norm_qkv_w4b(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4477 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4478 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4479 dvb: &mut CudaSlice<u8>,
4480 ncols: usize, rq: usize, rk: usize, eps: f32, vf16: bool)
4481 -> Result<(), Box<dyn std::error::Error>> {
4482 assert!(ncols % 4 == 0 && rq + 2 * rk >= 64);
4483 let f = self.func("rms_norm_qkv_w4b_f32");
4484 let rows = (rq + 2 * rk) as u32;
4485 let cfg = LaunchConfig {
4486 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4487 };
4488 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4489 let vf = vf16 as i32;
4490 let __s_b = self.gpu.stream();
4491 let mut b = __s_b.launch_builder(&f);
4492 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv).arg(&mut *dvb)
4493 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e).arg(&vf);
4494 unsafe { b.launch(cfg)?; }
4495 Ok(())
4496 }
4497
4498 pub fn rms_norm_qkv(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4499 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4500 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4501 ncols: usize, rq: usize, rk: usize, eps: f32)
4502 -> Result<(), Box<dyn std::error::Error>> {
4503 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4507 let warp_on = *WARP_ON.get_or_init(|| {
4508 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4509 });
4510 if warp_on && ncols % 4 == 0 && rq + 2 * rk >= 64 {
4513 let f = self.func("rms_norm_qkv_w4_f32");
4514 let rows = (rq + 2 * rk) as u32;
4515 let cfg = LaunchConfig {
4516 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4517 };
4518 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4519 let __s_b = self.gpu.stream();
4520 let mut b = __s_b.launch_builder(&f);
4521 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4522 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e);
4523 unsafe { b.launch(cfg)?; }
4524 return Ok(());
4525 }
4526 let f = self.func("rms_norm_qkv_f32");
4527 let grid = (rq + 2 * rk) as u32;
4528 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4529 let (nc, rqi, rki, e) = (ncols as i32, rq as i32, rk as i32, eps);
4530 let __s_b = self.gpu.stream();
4531 let mut b = __s_b.launch_builder(&f);
4532 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4533 .arg(&nc).arg(&rqi).arg(&rki).arg(&e);
4534 unsafe { b.launch(cfg)?; }
4535 Ok(())
4536 }
4537
4538 #[allow(clippy::too_many_arguments)]
4540 pub fn rms_norm2x(&self, a: &CudaSlice<f32>, bb: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4541 wb: &CudaSlice<f32>, da: &mut CudaSlice<f32>, db: &mut CudaSlice<f32>,
4542 ncols: usize, nrows: usize, eps: f32)
4543 -> Result<(), Box<dyn std::error::Error>> {
4544 let f = self.func("rms_norm2x_f32");
4545 let cfg = LaunchConfig { grid_dim: (2 * nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4546 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
4547 let __s_b = self.gpu.stream();
4548 let mut b = __s_b.launch_builder(&f);
4549 b.arg(a).arg(bb).arg(wa).arg(wb).arg(da).arg(db).arg(&nc).arg(&nr).arg(&e);
4550 unsafe { b.launch(cfg)?; }
4551 Ok(())
4552 }
4553
4554 pub fn softcap(&self, y: &mut CudaSlice<f32>, cap: f32, n: usize)
4556 -> Result<(), Box<dyn std::error::Error>> {
4557 let f = self.func("softcap_f32");
4558 let cfg = LaunchConfig::for_num_elems(n as u32);
4559 let ni = n as i32;
4560 let __s_b = self.gpu.stream();
4561 let mut b = __s_b.launch_builder(&f);
4562 b.arg(y).arg(&cap).arg(&ni);
4563 unsafe { b.launch(cfg)?; }
4564 Ok(())
4565 }
4566
4567 pub fn mask_ids_rows(&self, y: &mut CudaSlice<f32>, ids: &CudaSlice<i32>, n_ids: usize,
4570 n_vocab: usize, t: usize)
4571 -> Result<(), Box<dyn std::error::Error>> {
4572 let f = self.func("mask_ids_rows_f32");
4573 let cfg = LaunchConfig::for_num_elems((n_ids * t) as u32);
4574 let (ni, nv, ti) = (n_ids as i32, n_vocab as i32, t as i32);
4575 let __s_b = self.gpu.stream();
4576 let mut b = __s_b.launch_builder(&f);
4577 b.arg(y).arg(ids).arg(&ni).arg(&nv).arg(&ti);
4578 unsafe { b.launch(cfg)?; }
4579 Ok(())
4580 }
4581
4582 #[allow(clippy::too_many_arguments)]
4584 pub fn add_scale_rms_norm(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4585 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4586 ncols: usize, nrows: usize, eps: f32)
4587 -> Result<(), Box<dyn std::error::Error>> {
4588 let f = self.func("add_scale_rms_norm_f32");
4589 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4590 let (nc, e2) = (ncols as i32, eps);
4591 let __s_b = self.gpu.stream();
4592 let mut b = __s_b.launch_builder(&f);
4593 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(dst).arg(&nc).arg(&e2);
4594 unsafe { b.launch(cfg)?; }
4595 Ok(())
4596 }
4597
4598 #[allow(clippy::too_many_arguments)]
4601 pub fn add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4602 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4603 ncols: usize, nrows: usize, eps: f32)
4604 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4605 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4606 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4607 let (nc, e2) = (ncols as i32, eps);
4608 if Self::pdl_on() && Self::pdl_wb_on() {
4609 {
4610 use cudarc::driver::{DevicePtr, DevicePtrMut};
4611 let s = &self.gpu.stream();
4612 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4613 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4614 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4615 let mut ps = [
4616 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4617 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4618 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4619 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4620 &e2 as *const _ as *mut _,
4621 ];
4622 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4623 (rms_block(), 1, 1), &mut ps)?; }
4624 }
4625 return Ok((out_q, out_d));
4626 }
4627 let f = self.func("add_scale_rms_norm_q8_1");
4628 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4629 let __s_b = self.gpu.stream();
4630 let mut b = __s_b.launch_builder(&f);
4631 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e2);
4632 unsafe { b.launch(cfg)?; }
4633 Ok((out_q, out_d))
4634 }
4635
4636 #[allow(clippy::too_many_arguments)]
4638 pub fn add_scale_rms_norm_q8_1_into(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4639 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4640 ncols: usize, nrows: usize, eps: f32,
4641 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4642 -> Result<(), Box<dyn std::error::Error>> {
4643 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4644 let (nc, e2) = (ncols as i32, eps);
4645 if Self::pdl_on() && Self::pdl_wb_on() {
4646 use cudarc::driver::{DevicePtr, DevicePtrMut};
4647 let s = &self.gpu.stream();
4648 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4649 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4650 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4651 let mut ps = [
4652 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4653 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4654 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4655 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4656 &e2 as *const _ as *mut _,
4657 ];
4658 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4659 (rms_block(), 1, 1), &mut ps)?; }
4660 return Ok(());
4661 }
4662 let f = self.func("add_scale_rms_norm_q8_1");
4663 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4664 let __s_b = self.gpu.stream();
4665 let mut b = __s_b.launch_builder(&f);
4666 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut *out_q).arg(&mut *out_d).arg(&nc).arg(&e2);
4667 unsafe { b.launch(cfg)?; }
4668 Ok(())
4669 }
4670
4671 #[allow(clippy::too_many_arguments)]
4674 pub fn rms_pre_add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4675 b_in: &CudaSlice<f32>, c: f32,
4676 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4677 ncols: usize, nrows: usize, eps: f32)
4678 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4679 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4680 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4681 let (nc, e2) = (ncols as i32, eps);
4682 if Self::pdl_on() {
4683 {
4684 use cudarc::driver::{DevicePtr, DevicePtrMut};
4685 let s = &self.gpu.stream();
4686 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
4687 let (pb, _g2) = b_in.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
4688 let (pr, _g4) = res.device_ptr_mut(s);
4689 let (pq, _g5) = out_q.device_ptr_mut(s); let (pd, _g6) = out_d.device_ptr_mut(s);
4690 let mut ps = [
4691 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
4692 &pb as *const _ as *mut _, &c as *const _ as *mut _,
4693 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
4694 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4695 &nc as *const _ as *mut _, &e2 as *const _ as *mut _,
4696 ];
4697 unsafe { self.launch_pdl("rms_pre_add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4698 (rms_block(), 1, 1), &mut ps)?; }
4699 }
4700 return Ok((out_q, out_d));
4701 }
4702 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
4703 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4704 let __s_b = self.gpu.stream();
4705 let mut b = __s_b.launch_builder(&f);
4706 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);
4707 unsafe { b.launch(cfg)?; }
4708 Ok((out_q, out_d))
4709 }
4710
4711 pub fn gelu_tanh_mul_q8_1(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4714 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
4715 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4716 debug_assert!(ncols % 128 == 0);
4717 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4718 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4719 let nc = ncols as i32;
4720 if Self::pdl_on() {
4721 {
4722 use cudarc::driver::{DevicePtr, DevicePtrMut};
4723 let s = &self.gpu.stream();
4724 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4725 let (pact, _g2) = act.device_ptr_mut(s);
4726 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4727 let mut ps = [
4728 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4729 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4730 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4731 ];
4732 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4733 (rms_block(), 1, 1), &mut ps)?; }
4734 }
4735 return Ok((out_q, out_d));
4736 }
4737 let f = self.func("gelu_tanh_mul_q8_1");
4738 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4739 let __s_b = self.gpu.stream();
4740 let mut b = __s_b.launch_builder(&f);
4741 b.arg(gate).arg(up).arg(act).arg(&mut out_q).arg(&mut out_d).arg(&nc);
4742 unsafe { b.launch(cfg)?; }
4743 Ok((out_q, out_d))
4744 }
4745
4746 #[allow(clippy::too_many_arguments)]
4748 pub fn gelu_tanh_mul_q8_1_into(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4749 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
4750 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4751 -> Result<(), Box<dyn std::error::Error>> {
4752 debug_assert!(ncols % 128 == 0);
4753 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4754 let nc = ncols as i32;
4755 if Self::pdl_on() {
4756 use cudarc::driver::{DevicePtr, DevicePtrMut};
4757 let s = &self.gpu.stream();
4758 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4759 let (pact, _g2) = act.device_ptr_mut(s);
4760 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4761 let mut ps = [
4762 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4763 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4764 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4765 ];
4766 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4767 (rms_block(), 1, 1), &mut ps)?; }
4768 return Ok(());
4769 }
4770 let f = self.func("gelu_tanh_mul_q8_1");
4771 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4772 let __s_b = self.gpu.stream();
4773 let mut b = __s_b.launch_builder(&f);
4774 b.arg(gate).arg(up).arg(&mut *act).arg(&mut *out_q).arg(&mut *out_d).arg(&nc);
4775 unsafe { b.launch(cfg)?; }
4776 Ok(())
4777 }
4778
4779 #[allow(clippy::too_many_arguments)]
4781 pub fn add_rms_norm3_q8z(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4782 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4783 res: &mut CudaSlice<f32>, out1: &mut CudaSlice<f32>,
4784 ncols: usize, nrows: usize, eps: f32)
4785 -> Result<((CudaSlice<i8>, CudaSlice<f32>), (CudaSlice<i8>, CudaSlice<f32>)), Box<dyn std::error::Error>> {
4786 let mut q0 = self.alloc_uninit::<i8>(nrows * ncols)?;
4787 let mut d0 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4788 let mut q2 = self.alloc_uninit::<i8>(nrows * ncols)?;
4789 let mut d2 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4790 let f = self.func("add_rms_norm3_q8z_f32");
4791 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4792 let (nc, e2) = (ncols as i32, eps);
4793 let __s_b = self.gpu.stream();
4794 let mut b = __s_b.launch_builder(&f);
4795 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res)
4796 .arg(&mut q0).arg(&mut d0).arg(out1).arg(&mut q2).arg(&mut d2).arg(&nc).arg(&e2);
4797 unsafe { b.launch(cfg)?; }
4798 Ok(((q0, d0), (q2, d2)))
4799 }
4800
4801 #[allow(clippy::too_many_arguments)]
4803 pub fn add_rms_norm3(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4804 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4805 res: &mut CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4806 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4807 -> Result<(), Box<dyn std::error::Error>> {
4808 let f = self.func("add_rms_norm3_f32");
4809 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4810 let (nc, e2) = (ncols as i32, eps);
4811 let __s_b = self.gpu.stream();
4812 let mut b = __s_b.launch_builder(&f);
4813 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e2);
4814 unsafe { b.launch(cfg)?; }
4815 Ok(())
4816 }
4817
4818 pub fn add_scale(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4820 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
4821 let f = self.func("add_scale_f32");
4822 let cfg = LaunchConfig::for_num_elems(n as u32);
4823 let ni = n as i32;
4824 let __s_b = self.gpu.stream();
4825 let mut b = __s_b.launch_builder(&f);
4826 b.arg(a).arg(b_in).arg(&c).arg(dst).arg(&ni);
4827 unsafe { b.launch(cfg)?; }
4828 Ok(())
4829 }
4830
4831 pub fn rms_norm(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4832 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4833 let (nc, e) = (ncols as i32, eps);
4834 if Self::pdl_on() && Self::pdl_wb_on() {
4835 use cudarc::driver::{DevicePtr, DevicePtrMut};
4836 let s = &self.gpu.stream();
4837 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4838 let (pd, _g2) = dst.device_ptr_mut(s);
4839 let mut ps = [
4840 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4841 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4842 &e as *const _ as *mut _,
4843 ];
4844 unsafe { self.launch_pdl("rms_norm_f32", (nrows as u32, 1, 1),
4845 (rms_block(), 1, 1), &mut ps)?; }
4846 return Ok(());
4847 }
4848 let f = self.func("rms_norm_f32");
4849 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4850 let __s_b = self.gpu.stream();
4851 let mut b = __s_b.launch_builder(&f);
4852 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4853 unsafe { b.launch(cfg)?; }
4854 Ok(())
4855 }
4856
4857 pub fn rms_norm_decode(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4865 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4866 let f = self.func("rms_norm_f32");
4867 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4868 let (nc, e) = (ncols as i32, eps);
4869 let __s_b = self.gpu.stream();
4870 let mut b = __s_b.launch_builder(&f);
4871 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4872 unsafe { b.launch(cfg)?; }
4873 Ok(())
4874 }
4875
4876 pub fn rms_norm_q8_1(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize, nrows: usize,
4880 eps: f32) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4881 let nblk = ncols / 32;
4882 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
4883 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
4884 let (nc, e) = (ncols as i32, eps);
4885 if Self::pdl_on() {
4886 {
4887 use cudarc::driver::{DevicePtr, DevicePtrMut};
4888 let s = &self.gpu.stream();
4889 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4890 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
4891 let mut ps = [
4892 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4893 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4894 &nc as *const _ as *mut _, &e as *const _ as *mut _,
4895 ];
4896 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
4897 &mut ps)?; }
4898 }
4899 return Ok((q, d));
4900 }
4901 let f = self.func("rms_norm_q8_1");
4902 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4905 let __s_b = self.gpu.stream();
4906 let mut b = __s_b.launch_builder(&f);
4907 b.arg(x).arg(w).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
4908 unsafe { b.launch(cfg)?; }
4909 Ok((q, d))
4910 }
4911
4912 pub fn rms_norm_q8_1_into(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize,
4915 nrows: usize, eps: f32,
4916 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
4917 -> Result<(), Box<dyn std::error::Error>> {
4918 let nblk = ncols / 32;
4919 debug_assert!(q.len() >= nrows * ncols && d.len() >= nrows * nblk);
4920 let (nc, e) = (ncols as i32, eps);
4921 if Self::pdl_on() {
4922 use cudarc::driver::{DevicePtr, DevicePtrMut};
4923 let s = &self.gpu.stream();
4924 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4925 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
4926 let mut ps = [
4927 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4928 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4929 &nc as *const _ as *mut _, &e as *const _ as *mut _,
4930 ];
4931 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
4932 &mut ps)?; }
4933 return Ok(());
4934 }
4935 let f = self.func("rms_norm_q8_1");
4936 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4937 let __s_b = self.gpu.stream();
4938 let mut b = __s_b.launch_builder(&f);
4939 b.arg(x).arg(w).arg(&mut *q).arg(&mut *d).arg(&nc).arg(&e);
4940 unsafe { b.launch(cfg)?; }
4941 Ok(())
4942 }
4943
4944 pub fn quantize_q8_1_into(&self, x: &CudaSlice<f32>, m: usize, in_f: usize,
4946 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
4947 -> Result<(), Box<dyn std::error::Error>> {
4948 let nblk = in_f / 32;
4949 debug_assert!(q.len() >= m * in_f && d.len() >= m * nblk);
4950 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
4951 let (inf, mi) = (in_f as i32, m as i32);
4952 if Self::pdl_on() && Self::pdl_wb_on() {
4953 use cudarc::driver::{DevicePtr, DevicePtrMut};
4954 let s = &self.gpu.stream();
4955 let (px, _g0) = x.device_ptr(s);
4956 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
4957 let mut ps = [
4958 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
4959 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
4960 &mi as *const _ as *mut _,
4961 ];
4962 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
4963 return Ok(());
4964 }
4965 let f = self.func("quantize_q8_1");
4966 let __s_b = self.gpu.stream();
4967 let mut b = __s_b.launch_builder(&f);
4968 b.arg(x).arg(&mut *q).arg(&mut *d).arg(&inf).arg(&mi);
4969 unsafe { b.launch(cfg)?; }
4970 Ok(())
4971 }
4972
4973 pub fn add_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
4977 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4978 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4979 let nblk = ncols / 32;
4980 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
4981 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
4982 let f = self.func("add_rms_norm_q8_1");
4983 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4985 let (nc, e) = (ncols as i32, eps);
4986 let __s_bld = self.gpu.stream();
4987 let mut bld = __s_bld.launch_builder(&f);
4988 bld.arg(a).arg(b_in).arg(w).arg(res).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
4989 unsafe { bld.launch(cfg)?; }
4990 Ok((q, d))
4991 }
4992
4993 pub fn add_rms_norm(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
4997 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
4998 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4999 let (nc, e) = (ncols as i32, eps);
5000 if Self::pdl_on() && Self::pdl_wb_on() {
5001 use cudarc::driver::{DevicePtr, DevicePtrMut};
5002 let s = &self.gpu.stream();
5003 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b.device_ptr(s);
5004 let (pw, _g2) = w.device_ptr(s);
5005 let (pr, _g3) = res.device_ptr_mut(s); let (pd, _g4) = dst.device_ptr_mut(s);
5006 let mut ps = [
5007 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
5008 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
5009 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
5010 &e as *const _ as *mut _,
5011 ];
5012 unsafe { self.launch_pdl("add_rms_norm_f32", (nrows as u32, 1, 1),
5013 (rms_block(), 1, 1), &mut ps)?; }
5014 return Ok(());
5015 }
5016 let f = self.func("add_rms_norm_f32");
5017 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5018 let __s_b2 = self.gpu.stream();
5019 let mut b2 = __s_b2.launch_builder(&f);
5020 b2.arg(a).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
5021 unsafe { b2.launch(cfg)?; }
5022 Ok(())
5023 }
5024
5025 #[allow(clippy::too_many_arguments)]
5028 pub fn rms_pre_add_rms_norm(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
5029 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
5030 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5031 ncols: usize, nrows: usize, eps: f32)
5032 -> Result<(), Box<dyn std::error::Error>> {
5033 let f = self.func("rms_pre_add_rms_norm_f32");
5034 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5035 let (nc, e) = (ncols as i32, eps);
5036 let __s_b2 = self.gpu.stream();
5037 let mut b2 = __s_b2.launch_builder(&f);
5038 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
5039 unsafe { b2.launch(cfg)?; }
5040 Ok(())
5041 }
5042
5043 #[allow(clippy::too_many_arguments)]
5045 pub fn rms_pre_add_rms_norm_q8z(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
5046 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
5047 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5048 ncols: usize, nrows: usize, eps: f32)
5049 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5050 debug_assert!(ncols % 128 == 0);
5051 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5052 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5053 let (nc, e) = (ncols as i32, eps);
5054 if Self::pdl_on() {
5055 {
5056 use cudarc::driver::{DevicePtr, DevicePtrMut};
5057 let s = &self.gpu.stream();
5058 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
5059 let (pb, _g2) = b.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
5060 let (pr, _g4) = res.device_ptr_mut(s); let (pdst, _g5) = dst.device_ptr_mut(s);
5061 let (pq, _g6) = out_q.device_ptr_mut(s); let (pd, _g7) = out_d.device_ptr_mut(s);
5062 let mut ps = [
5063 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
5064 &pb as *const _ as *mut _, &pw as *const _ as *mut _,
5065 &pr as *const _ as *mut _, &pdst as *const _ as *mut _,
5066 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
5067 &nc as *const _ as *mut _, &e as *const _ as *mut _,
5068 ];
5069 unsafe { self.launch_pdl("rms_pre_add_rms_norm_q8z_f32", (nrows as u32, 1, 1),
5070 (rms_block(), 1, 1), &mut ps)?; }
5071 }
5072 return Ok((out_q, out_d));
5073 }
5074 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
5075 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5076 let __s_b2 = self.gpu.stream();
5077 let mut b2 = __s_b2.launch_builder(&f);
5078 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst)
5079 .arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e);
5080 unsafe { b2.launch(cfg)?; }
5081 Ok((out_q, out_d))
5082 }
5083
5084 pub fn build_q4_out_concat3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
5088 w2: &crate::model::GpuTensor)
5089 -> Result<Option<crate::model::GpuTensor>, Box<dyn std::error::Error>> {
5090 use crate::model::GpuTensor;
5091 let part = |w: &GpuTensor| -> Option<(usize, usize)> {
5092 match w {
5093 GpuTensor::Quant { qtype, row_bytes, rp, .. }
5094 if *qtype == QT_Q4_0 && !*rp => Some((*row_bytes, w.out_features())),
5095 _ => None,
5096 }
5097 };
5098 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (part(w0), part(w1), part(w2))
5099 else { return Ok(None) };
5100 if rb0 != rb1 || rb0 != rb2
5101 || w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
5102 return Ok(None);
5103 }
5104 fn bytes_of(w: &crate::model::GpuTensor) -> &CudaSlice<u8> {
5105 match w { crate::model::GpuTensor::Quant { bytes, .. } => bytes, _ => unreachable!() }
5106 }
5107 let (b0, b1, b2) = (bytes_of(w0), bytes_of(w1), bytes_of(w2));
5108 let total = rb0 * (o0 + o1 + o2);
5109 let mut cat = self.alloc_u8(total)?;
5110 self.copy_u8_into(&mut cat, 0, b0, rb0 * o0)?;
5111 self.copy_u8_into(&mut cat, rb0 * o0, b1, rb1 * o1)?;
5112 self.copy_u8_into(&mut cat, rb0 * (o0 + o1), b2, rb2 * o2)?;
5113 Ok(Some(GpuTensor::Quant {
5114 bytes: cat, qtype: QT_Q4_0, row_bytes: rb0,
5115 ne: vec![w0.in_features() as u64, (o0 + o1 + o2) as u64], scale: 1.0, rp: false,
5116 #[cfg(memra_cutlass)]
5117 cutlass: None,
5118 fp8: None, blk: None, rp4: None, f16: None,
5119 }))
5120 }
5121
5122 #[allow(clippy::too_many_arguments)]
5124 pub fn rms_norm_qkv_rope_cat(&self, qkv: &CudaSlice<f32>,
5125 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5126 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5127 head_dim: usize, rq: usize, rk: usize,
5128 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5129 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32)
5130 -> Result<(), Box<dyn std::error::Error>> {
5131 let rows = rq + rk + rk;
5132 let theta_scale = base.powf(-2.0 / head_dim as f32);
5133 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5134 if Self::pdl_on() {
5135 use cudarc::driver::{DevicePtr, DevicePtrMut};
5136 let s = &self.gpu.stream();
5137 let (pqkv, _g0) = qkv.device_ptr(s);
5138 let (pwq, _g1) = wq.device_ptr(s); let (pwk, _g2) = wk.device_ptr(s);
5139 let (pwv, _g3) = wv.device_ptr(s);
5140 let (pq, _g4) = q.device_ptr_mut(s); let (pk, _g5) = k.device_ptr_mut(s);
5141 let (pv, _g6) = v.device_ptr_mut(s);
5142 let (ppos, _g7) = pos.device_ptr(s);
5143 let (pff, _g8) = match ff {
5144 Some(t) => { let (p, g) = t.device_ptr(s); (p, Some(g)) }
5145 None => (0, None),
5146 };
5147 let mut ps = [
5148 &pqkv as *const _ as *mut std::ffi::c_void,
5149 &pwq as *const _ as *mut _, &pwk as *const _ as *mut _,
5150 &pwv as *const _ as *mut _,
5151 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5152 &pv as *const _ as *mut _,
5153 &nc as *const _ as *mut _, &rqi as *const _ as *mut _,
5154 &rki as *const _ as *mut _, &ppos as *const _ as *mut _,
5155 &nhq as *const _ as *mut _, &nhk as *const _ as *mut _,
5156 &theta_scale as *const _ as *mut _, &freq_scale as *const _ as *mut _,
5157 &pff as *const _ as *mut _, &eps as *const _ as *mut _,
5158 ];
5159 unsafe { self.launch_pdl("rms_norm_qkv_rope_cat_f32", (rows as u32, 1, 1),
5160 (rms_block(), 1, 1), &mut ps)?; }
5161 return Ok(());
5162 }
5163 let f = self.func("rms_norm_qkv_rope_cat_f32");
5164 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5165 let __s_b = self.gpu.stream();
5166 let mut b = __s_b.launch_builder(&f);
5167 match ff {
5168 Some(t) => { b.arg(qkv).arg(wq).arg(wk).arg(wv)
5169 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5170 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5171 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5172 unsafe { b.launch(cfg)?; } }
5173 None => { let null: u64 = 0;
5174 b.arg(qkv).arg(wq).arg(wk).arg(wv)
5175 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5176 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5177 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5178 unsafe { b.launch(cfg)?; } }
5179 }
5180 Ok(())
5181 }
5182
5183 #[allow(clippy::too_many_arguments)]
5185 pub fn rms_norm_qkv_rope(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>, v0: &CudaSlice<f32>,
5186 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5187 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5188 head_dim: usize, rq: usize, rk: usize,
5189 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5190 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32)
5191 -> Result<(), Box<dyn std::error::Error>> {
5192 let f = self.func("rms_norm_qkv_rope_f32");
5193 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 };
5195 let theta_scale = base.powf(-2.0 / head_dim as f32);
5196 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5197 let __s_b = self.gpu.stream();
5198 let mut b = __s_b.launch_builder(&f);
5199 match ff {
5200 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5201 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5202 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5203 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5204 unsafe { b.launch(cfg)?; } }
5205 None => { let null: u64 = 0;
5206 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5207 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5208 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5209 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5210 unsafe { b.launch(cfg)?; } }
5211 }
5212 Ok(())
5213 }
5214
5215 #[allow(clippy::too_many_arguments)]
5219 pub fn rms_norm_qkv_rope_append_dc(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>,
5220 v0: &CudaSlice<f32>,
5221 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5222 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5223 head_dim: usize, rq: usize, rk: usize,
5224 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5225 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32,
5226 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
5227 t_dev: &CudaSlice<i32>, k_tok_bytes: usize, v_tok_bytes: usize,
5228 g: bool)
5229 -> Result<(), Box<dyn std::error::Error>> {
5230 let rows = rq + rk + rk;
5231 let theta_scale = base.powf(-2.0 / head_dim as f32);
5232 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5233 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
5234 if Self::pdl_on() && Self::pdl_wb_on() {
5235 use cudarc::driver::{DevicePtr, DevicePtrMut};
5236 let s = &self.gpu.stream();
5237 let (p0, _a0) = q0.device_ptr(s); let (p1, _a1) = k0.device_ptr(s);
5238 let (p2, _a2) = v0.device_ptr(s);
5239 let (pwq, _a3) = wq.device_ptr(s); let (pwk, _a4) = wk.device_ptr(s);
5240 let (pwv, _a5) = wv.device_ptr(s);
5241 let (pq, _a6) = q.device_ptr_mut(s); let (pk, _a7) = k.device_ptr_mut(s);
5242 let (pv, _a8) = v.device_ptr_mut(s);
5243 let (pp, _a9) = pos.device_ptr(s);
5244 let pff: u64 = match ff { Some(t) => { let (p, _gg) = t.device_ptr(s); p as u64 }
5245 None => 0 };
5246 let (pkc, _a10) = kc.device_ptr_mut(s); let (pvc, _a11) = vc.device_ptr_mut(s);
5247 let (pt, _a12) = t_dev.device_ptr(s);
5248 let mut ps = [
5249 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
5250 &p2 as *const _ as *mut _, &pwq as *const _ as *mut _,
5251 &pwk as *const _ as *mut _, &pwv as *const _ as *mut _,
5252 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5253 &pv as *const _ as *mut _, &nc as *const _ as *mut _,
5254 &rqi as *const _ as *mut _, &rki as *const _ as *mut _,
5255 &pp as *const _ as *mut _, &nhq as *const _ as *mut _,
5256 &nhk as *const _ as *mut _, &theta_scale as *const _ as *mut _,
5257 &freq_scale as *const _ as *mut _, &pff as *const _ as *mut _,
5258 &eps as *const _ as *mut _, &pkc as *const _ as *mut _,
5259 &pvc as *const _ as *mut _, &pt as *const _ as *mut _,
5260 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
5261 ];
5262 unsafe { self.launch_pdl_flash(g, "rms_norm_qkv_rope_append_dc_f32",
5263 (rows as u32, 1, 1), (rms_block(), 1, 1), 0, &mut ps)?; }
5264 return Ok(());
5265 }
5266 let f = if g { self.func_g("rms_norm_qkv_rope_append_dc_f32") }
5267 else { self.func("rms_norm_qkv_rope_append_dc_f32") };
5268 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5269 let __s_b = self.gpu.stream();
5270 let mut b = __s_b.launch_builder(&f);
5271 match ff {
5272 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5273 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5274 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5275 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps)
5276 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5277 unsafe { b.launch(cfg)?; } }
5278 None => { let null: u64 = 0;
5279 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5280 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5281 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5282 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps)
5283 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5284 unsafe { b.launch(cfg)?; } }
5285 }
5286 Ok(())
5287 }
5288
5289 pub fn add_q8_1(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
5291 ncols: usize, nrows: usize)
5292 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5293 debug_assert!(ncols % 128 == 0);
5294 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5295 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5296 let f = self.func("add_q8_1_f32");
5297 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5298 let nc = ncols as i32;
5299 let __s_b2 = self.gpu.stream();
5300 let mut b2 = __s_b2.launch_builder(&f);
5301 b2.arg(a).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc);
5302 unsafe { b2.launch(cfg)?; }
5303 Ok((out_q, out_d))
5304 }
5305
5306 pub fn rms_pre_add_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>, b: &CudaSlice<f32>,
5310 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
5311 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5312 debug_assert!(ncols % 128 == 0);
5313 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5314 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5315 let f = self.func("rms_pre_add_q8_1_f32");
5316 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1),
5317 shared_mem_bytes: 0 };
5318 let (nc, ep) = (ncols as i32, eps);
5319 let __s_b2 = self.gpu.stream();
5320 let mut b2 = __s_b2.launch_builder(&f);
5321 b2.arg(a).arg(wa).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
5322 unsafe { b2.launch(cfg)?; }
5323 Ok((out_q, out_d))
5324 }
5325
5326 pub fn l2_v2_on(ncols: usize) -> bool {
5330 ncols == 128 && std::env::var("MEMRA_L2_V2").as_deref() != Ok("0")
5331 }
5332
5333 pub fn l2_norm_pp(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5334 dst16: Option<&mut CudaSlice<u8>>, ncols: usize, nrows: usize,
5335 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5336 if Self::l2_v2_on(ncols) {
5337 let f = self.func("l2_norm_pp_v2_f32");
5338 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 };
5340 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
5341 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
5343 let __s_b = self.gpu.stream();
5344 let mut b = __s_b.launch_builder(&f);
5345 b.arg(x).arg(dst).arg(&d16).arg(&nc).arg(&nr).arg(&e);
5346 unsafe { b.launch(cfg)?; }
5347 return Ok(());
5348 }
5349 self.l2_norm(x, dst, ncols, nrows, eps)
5350 }
5351
5352 pub fn l2_norm(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
5353 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5354 let f = self.func("l2_norm_f32");
5355 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
5356 let (nc, e) = (ncols as i32, eps);
5357 let __s_b = self.gpu.stream();
5358 let mut b = __s_b.launch_builder(&f);
5359 b.arg(x).arg(dst).arg(&nc).arg(&e);
5360 unsafe { b.launch(cfg)?; }
5361 Ok(())
5362 }
5363
5364 pub fn l2_norm_decode(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize,
5370 nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5371 let f = self.func("l2_norm_f32");
5372 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
5373 let (nc, e) = (ncols as i32, eps);
5374 let __s_b = self.gpu.stream();
5375 let mut b = __s_b.launch_builder(&f);
5376 b.arg(x).arg(dst).arg(&nc).arg(&e);
5377 unsafe { b.launch(cfg)?; }
5378 Ok(())
5379 }
5380
5381 pub fn rope_neox(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5383 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32, freq_scale: f32)
5384 -> Result<(), Box<dyn std::error::Error>> {
5385 let f = self.func("rope_neox_f32");
5386 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5387 let grid = (n_heads * n_tokens) as u32;
5388 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5389 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5390 let __s_b = self.gpu.stream();
5391 let mut b = __s_b.launch_builder(&f);
5392 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale);
5393 unsafe { b.launch(cfg)?; }
5394 Ok(())
5395 }
5396
5397 pub fn rope_neox_ff(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5399 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32,
5400 freq_scale: f32, ff: &CudaSlice<f32>)
5401 -> Result<(), Box<dyn std::error::Error>> {
5402 let f = self.func("rope_neox_ff_f32");
5403 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5404 let grid = (n_heads * n_tokens) as u32;
5405 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5406 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5407 let __s_b = self.gpu.stream();
5408 let mut b = __s_b.launch_builder(&f);
5409 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale).arg(ff);
5410 unsafe { b.launch(cfg)?; }
5411 Ok(())
5412 }
5413
5414 #[allow(clippy::too_many_arguments)]
5416 pub fn rope_neox2(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
5417 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
5418 nh_q: usize, nh_k: usize, n_tokens: usize, freq_base: f32,
5419 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
5420 -> Result<(), Box<dyn std::error::Error>> {
5421 let f = self.func("rope_neox2_f32");
5422 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5423 let grid = ((nh_q + nh_k) * n_tokens) as u32;
5424 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5425 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);
5426 let __s_b = self.gpu.stream();
5427 let mut b = __s_b.launch_builder(&f);
5428 b.arg(q).arg(k).arg(pos).arg(&hd).arg(&nd).arg(&nq).arg(&nk).arg(&nt)
5429 .arg(&theta_scale).arg(&freq_scale);
5430 match ff {
5431 Some(ffv) => { b.arg(ffv); unsafe { b.launch(cfg)?; } }
5432 None => {
5433 let null: u64 = 0;
5434 b.arg(&null);
5435 unsafe { b.launch(cfg)?; }
5436 }
5437 }
5438 Ok(())
5439 }
5440
5441 pub fn gelu_tanh_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5443 -> Result<(), Box<dyn std::error::Error>> {
5444 let f = self.func("gelu_tanh_mul_f32");
5445 let cfg = LaunchConfig::for_num_elems(n as u32);
5446 let ni = n as i32;
5447 let __s_b = self.gpu.stream();
5448 let mut b = __s_b.launch_builder(&f);
5449 b.arg(gate).arg(up).arg(dst).arg(&ni);
5450 unsafe { b.launch(cfg)?; }
5451 Ok(())
5452 }
5453
5454 pub fn silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5455 -> Result<(), Box<dyn std::error::Error>> {
5456 let f = self.func("silu_mul_f32");
5457 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5459 let ni = n as i32;
5460 let __s_b = self.gpu.stream();
5461 let mut b = __s_b.launch_builder(&f);
5462 b.arg(gate).arg(up).arg(dst).arg(&ni);
5463 unsafe { b.launch(cfg)?; }
5464 Ok(())
5465 }
5466
5467 pub fn silu_mul_f16out(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
5470 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
5471 -> Result<(), Box<dyn std::error::Error>> {
5472 let f = self.func("silu_mul_f16out_f32");
5473 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5474 let ni = n as i32;
5475 let __s_b = self.gpu.stream();
5476 let mut b = __s_b.launch_builder(&f);
5477 b.arg(gate).arg(up).arg(dst).arg(dst16).arg(&ni);
5478 unsafe { b.launch(cfg)?; }
5479 Ok(())
5480 }
5481
5482 pub fn silu_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5489 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
5490 let f = self.func("silu_mul_scaled_f32");
5491 let cfg = LaunchConfig::for_num_elems(n as u32);
5492 let ni = n as i32;
5493 let (gsf, usf) = (gs, us);
5494 let __s_b = self.gpu.stream();
5495 let mut b = __s_b.launch_builder(&f);
5496 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(dst).arg(&ni);
5497 unsafe { b.launch(cfg)?; }
5498 Ok(())
5499 }
5500
5501 #[allow(clippy::too_many_arguments)]
5505 pub fn swigluoai_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5506 alpha: f32, limit: f32, dst: &mut CudaSlice<f32>, n: usize)
5507 -> Result<(), Box<dyn std::error::Error>> {
5508 let f = self.func("swigluoai_mul_scaled_f32");
5509 let cfg = LaunchConfig::for_num_elems(n as u32);
5510 let ni = n as i32;
5511 let __s_b = self.gpu.stream();
5512 let mut b = __s_b.launch_builder(&f);
5513 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&alpha).arg(&limit).arg(dst).arg(&ni);
5514 unsafe { b.launch(cfg)?; }
5515 Ok(())
5516 }
5517
5518 pub fn silu_mul_scaled_q8_1(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5526 n: usize)
5527 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5528 let f = self.func("silu_mul_scaled_q8_1");
5529 let nblk = n / 32;
5530 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);
5534 let (gsf, usf, ni) = (gs, us, n as i32);
5535 let __s_b = self.gpu.stream();
5536 let mut b = __s_b.launch_builder(&f);
5537 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(&mut aq).arg(&mut ad).arg(&ni);
5538 unsafe { b.launch(cfg)?; }
5539 Ok((aq, ad))
5540 }
5541
5542 pub fn add(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5543 -> Result<(), Box<dyn std::error::Error>> {
5544 let f = self.func("add_f32");
5545 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5547 let ni = n as i32;
5548 let __s_bld = self.gpu.stream();
5549 let mut bld = __s_bld.launch_builder(&f);
5550 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5551 unsafe { bld.launch(cfg)?; }
5552 Ok(())
5553 }
5554
5555 pub fn mul(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5556 -> Result<(), Box<dyn std::error::Error>> {
5557 let f = self.func("mul_f32");
5558 let cfg = LaunchConfig::for_num_elems(n as u32);
5559 let ni = n as i32;
5560 let __s_bld = self.gpu.stream();
5561 let mut bld = __s_bld.launch_builder(&f);
5562 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5563 unsafe { bld.launch(cfg)?; }
5564 Ok(())
5565 }
5566
5567 pub fn matmul(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
5570 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5571 use crate::model::GpuTensor;
5572 let in_f = w.in_features();
5573 let out_f = w.out_features();
5574 #[allow(non_snake_case)]
5582 let GEMM_M_THRESHOLD = if self.verify_exact_on() { usize::MAX } else { 16usize };
5585
5586 const GEMM_MIN_OUT_F: usize = 128; if m >= GEMM_M_THRESHOLD {
5611 if let Some(y) = self.try_fp8_gemm(w, x, m)? { return Ok(y); }
5612 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(y); }
5619 if let Some(y) = self.try_f16_gemm(w, x, m)? { return Ok(y); }
5622 }
5623 if let GpuTensor::Quant { qtype, .. } = w {
5638 if *qtype == QT_F8_E4M3_BLK {
5639 if m >= GEMM_M_THRESHOLD {
5640 if let Some(y) = self.try_e4m3_blk_prefill(w, x, m)? { return Ok(y); }
5641 }
5642 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5643 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
5644 }
5645 }
5646 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.mmq_supports(w) {
5647 return self.qmatvec_mmq(w, x, m);
5648 }
5649 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.gemm_supports(w) {
5650 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5651 return self.qmatvec_gemm(w, &aq, &ad, m);
5652 }
5653 if m >= GEMM_M_THRESHOLD {
5656 if let Some(y) = self.try_fp4_gemm(w, x, m, in_f, out_f)? { return Ok(y); }
5657 }
5658 let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
5662 if m == 1 && fast {
5667 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, scale, .. } = w {
5668 if self.mmvq_supports(*qtype) {
5669 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5673 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5674 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, rp);
5675 }
5676 }
5677 }
5678 if (2..=16).contains(&m) && fast && std::env::var("MEMRA_NO_BATCHED").is_err()
5694 && (m <= 4 || Self::b8_enabled()) {
5695 let m_ok = m <= 8 || matches!(w, GpuTensor::Quant { qtype, .. }
5705 if *qtype == QT_Q4_0 || *qtype == QT_Q6_K || *qtype == QT_F8_E4M3
5706 || *qtype == QT_NVFP4 || *qtype == QT_Q4_K || *qtype == QT_Q5_K || *qtype == QT_Q8_0);
5707 if m_ok {
5708 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, .. } = w {
5709 if self.batched_supports(*qtype) && self.mmvq_supports(*qtype) {
5710 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5711 let mcols = Self::batched_mcols(m);
5712 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5713 let mut y = self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, mcols, 1.0, rp)?;
5714 if let GpuTensor::Quant { scale, .. } = w {
5715 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5716 }
5717 return Ok(y);
5718 }
5719 }
5720 }
5721 }
5722 if fast {
5728 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. } = w {
5729 if *qtype == QT_F8_E4M3 {
5730 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5731 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes,
5732 *scale, false);
5733 }
5734 }
5735 }
5736 let mut y = match w {
5737 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q8_0 =>
5738 self.qmatvec_q8_0_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5739 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q4_K =>
5740 self.qmatvec_q4_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5741 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q6_K =>
5742 self.qmatvec_q6_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5743 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q5_K =>
5744 self.qmatvec_q5_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5745 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q3_K =>
5746 self.qmatvec_q3_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5747 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } if fast && *qtype == QT_NVFP4 =>
5748 self.qmatvec_dp4a_named(
5749 if *rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
5750 bytes, x, m, in_f, out_f, *row_bytes)?,
5751 GpuTensor::Quant { bytes, qtype, row_bytes, .. }
5755 if fast && *qtype == QT_IQ4_XS && Self::iq_fast_enabled() =>
5756 self.qmatvec_iq4_XS_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5757 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } =>
5762 self.qmatvec(bytes, x, m, in_f, out_f,
5765 if *rp && *qtype == QT_NVFP4 { QT_NVFP4_RP } else { *qtype },
5766 *row_bytes)?,
5767 GpuTensor::Float { data, .. } => self.linear(x, data, m, in_f, out_f)?,
5768 GpuTensor::FloatBf16 { data, .. } =>
5771 self.linear_bf16_chunked(x, data, m, in_f, out_f, false)?,
5772 };
5773 if let GpuTensor::Quant { scale, .. } = w {
5775 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5776 }
5777 Ok(y)
5778 }
5779
5780 pub fn uses_q8_1_fast(&self, w: &crate::model::GpuTensor) -> bool {
5783 use crate::model::GpuTensor;
5784 if std::env::var("MEMRA_FAST").as_deref() == Ok("0") { return false; }
5785 match w {
5786 GpuTensor::Quant { qtype, .. } => matches!(*qtype,
5793 QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q3_K | QT_NVFP4 | QT_F8_E4M3
5794 | QT_F8_E4M3_BLK | QT_Q4_0)
5795 || (*qtype == QT_IQ4_XS && Self::iq_fast_enabled()),
5796 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
5797 }
5798 }
5799
5800 pub fn matmul_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
5805 x_fallback: &CudaSlice<f32>, m: usize)
5806 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5807 use crate::model::GpuTensor;
5808 let x_raw_ok = x_fallback.len() >= m * w.in_features();
5814 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5817 if let Some(y) = self.try_fp8_gemm(w, x_fallback, m)? { return Ok(y); }
5818 if let Some(y) = self.try_fp8_blk_mmq(w, x_fallback, m)? { return Ok(y); }
5821 if let Some(y) = self.try_f16_gemm(w, x_fallback, m)? { return Ok(y); }
5823 }
5824 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5830 if let Some(y) = self.try_e4m3_blk_prefill(w, x_fallback, m)? { return Ok(y); }
5831 }
5832 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
5833 if m >= 16 && w.out_features() >= 128 && self.mmq_supports(w) && !self.verify_exact_on()
5838 && x_raw_ok {
5839 return self.qmatvec_mmq(w, x_fallback, m);
5840 }
5841 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5844 if let Some(y) = self.try_fp4_gemm(w, x_fallback, m, w.in_features(), w.out_features())? {
5845 return Ok(y);
5846 }
5847 }
5848 if m >= 16 && self.gemm_supports(w) && !self.verify_exact_on() {
5851 return self.qmatvec_gemm(w, aq, ad, m);
5852 }
5853 if !self.uses_q8_1_fast(w) { return self.matmul(w, x_fallback, m); }
5854 let in_f = w.in_features();
5855 let out_f = w.out_features();
5856 let (bytes, qtype, row_bytes, scale, rp) = match w {
5857 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
5858 _ => unreachable!("uses_q8_1_fast guaranteed Quant"),
5859 };
5860 let (mbytes, mrp) = match w {
5863 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5864 _ => (bytes, rp),
5865 };
5866 if m == 1 && self.mmvq_supports(qtype) {
5870 return self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, mrp);
5871 }
5872 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5885 && std::env::var("MEMRA_NO_BATCHED").is_err()
5886 && (m <= 4 || Self::b8_enabled())
5887 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_NVFP4
5891 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_F8_E4M3 || qtype == QT_Q8_0) {
5892 let mcols = Self::batched_mcols(m);
5893 return self.qmatvec_mmvq_batched(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, mrp);
5894 }
5895 if qtype == QT_F8_E4M3 || qtype == QT_Q4_0 {
5901 let (b2, r2) = if qtype == QT_Q4_0 { (mbytes, mrp) } else { (bytes, rp) };
5902 return self.qmatvec_mmvq(b2, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, r2);
5903 }
5904 let name = match qtype {
5905 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
5906 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
5907 QT_Q3_K => "qmatvec_q3_K_dp4a",
5908 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
5909 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
5910 _ => unreachable!(),
5911 };
5912 let f = self.func(name);
5913 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 };
5915 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
5916 let __s_b = self.gpu.stream();
5917 let mut b = __s_b.launch_builder(&f);
5918 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
5919 unsafe { b.launch(cfg)?; }
5920 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
5921 Ok(y)
5922 }
5923
5924 pub fn matmul_decode_exact(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
5932 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5933 use crate::model::GpuTensor;
5934 if let GpuTensor::Float { data, .. } = w {
5942 return self.linear_decode_exact(x, data, m, w.in_features(), w.out_features());
5943 }
5944 if let GpuTensor::FloatBf16 { data, .. } = w {
5947 let (in_f, out_f) = (w.in_features(), w.out_features());
5948 return self.linear_bf16_chunked(x, data, m, in_f, out_f, true);
5949 }
5950 if !self.uses_q8_1_fast(w) { return self.matmul(w, x, m); }
5951 let in_f = w.in_features();
5952 let out_f = w.out_features();
5953 let (bytes, qtype, row_bytes, scale, rp) = match w {
5954 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
5955 _ => return self.matmul(w, x, m),
5956 };
5957 let (bytes, rp) = match w {
5960 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5961 _ => (bytes, rp),
5962 };
5963 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5964 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
5968 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5977 && std::env::var("MEMRA_NO_BATCHED").is_err()
5978 && (m <= 4 || Self::b8_enabled())
5979 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
5982 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
5983 let mcols = Self::batched_mcols(m);
5984 return self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
5985 }
5986 if self.mmvq_supports(qtype) {
5987 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
5990 }
5991 self.matmul_pre(w, &aq, &ad, x, m)
5994 }
5995
5996 pub fn matmul_decode_exact_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
6006 ad: &CudaSlice<f32>, m: usize)
6007 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6008 use crate::model::GpuTensor;
6009 debug_assert!(self.uses_q8_1_fast(w),
6010 "matmul_decode_exact_pre: caller must guarantee q8_1-fast");
6011 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
6013 let in_f = w.in_features();
6014 let out_f = w.out_features();
6015 let (bytes, qtype, row_bytes, scale, rp) = match w {
6016 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
6017 (bytes, *qtype, *row_bytes, *scale, *rp),
6018 _ => return Err("matmul_decode_exact_pre: Quant tensor required (q8_1-fast contract)".into()),
6019 };
6020 let (bytes, rp) = match w {
6022 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
6023 _ => (bytes, rp),
6024 };
6025 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
6027 && std::env::var("MEMRA_NO_BATCHED").is_err()
6028 && (m <= 4 || Self::b8_enabled())
6029 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
6030 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
6031 let mcols = Self::batched_mcols(m);
6032 return self.qmatvec_mmvq_batched(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
6033 }
6034 if self.mmvq_supports(qtype) {
6035 return self.qmatvec_mmvq(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
6036 }
6037 let x0 = self.zeros(0)?;
6040 self.matmul_pre(w, aq, ad, &x0, m)
6041 }
6042
6043 pub fn matmul_decode_exact_dual_pre(&self, w0: &crate::model::GpuTensor,
6052 w1: &crate::model::GpuTensor,
6053 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6054 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
6055 use crate::model::GpuTensor;
6056 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6057 let on = *ON.get_or_init(|| {
6058 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
6059 });
6060 if !on || !(2..=7).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
6061 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
6062 return Ok(None);
6063 }
6064 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6069 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6070 if w1.in_features() != in_f || w1.out_features() != out_f {
6071 return Ok(None);
6072 }
6073 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
6074 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
6075 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
6076 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
6077 (b0, b1, *rb0, *s0, *s1, *rp0),
6078 _ => return Ok(None),
6079 };
6080 if m > 4 && !(rp && Self::b8_enabled()
6083 && std::env::var("MEMRA_B567").as_deref() != Ok("0")) {
6084 return Ok(None);
6085 }
6086 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, aq, ad, m, in_f, out_f, row_bytes, rp)?;
6087 Ok(Some(((y0, s0), (y1, s1))))
6088 }
6089
6090 pub fn matmul_decode_exact_dual(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6106 x: &CudaSlice<f32>, m: usize)
6107 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6108 use crate::model::GpuTensor;
6109 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6110 let on = *ON.get_or_init(|| {
6111 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
6112 });
6113 if !on || !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
6114 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
6115 return Ok(None);
6116 }
6117 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6122 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6123 if w1.in_features() != in_f || w1.out_features() != out_f {
6124 return Ok(None);
6125 }
6126 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
6127 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
6128 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
6129 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
6130 (b0, b1, *rb0, *s0, *s1, *rp0),
6131 _ => return Ok(None),
6132 };
6133 if std::env::var("MEMRA_DEBUG").is_ok() {
6136 static ONCE: std::sync::Once = std::sync::Once::new();
6137 ONCE.call_once(|| eprintln!("[memra] dual gate+up batched ENGAGED (m={m} rp={rp})"));
6138 }
6139 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6140 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, &aq, &ad, m, in_f, out_f, row_bytes, rp)?;
6141 let mut y0 = y0;
6142 let mut y1 = y1;
6143 if s0 != 1.0 { self.scale_inplace(&mut y0, s0, m * out_f)?; }
6144 if s1 != 1.0 { self.scale_inplace(&mut y1, s1, m * out_f)?; }
6145 Ok(Some((y0, y1)))
6146 }
6147
6148 #[allow(clippy::too_many_arguments)]
6153 pub fn qmatvec_batched_dual_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6154 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6155 m: usize, in_f: usize, out_f: usize, row_bytes: usize, rp: bool)
6156 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6157 const ROWS_PER_BLOCK: u32 = 4;
6158 let mcols = Self::batched_mcols(m);
6159 let (name, rows_per_block) = match (mcols, rp, m) {
6162 (2, false, _) => ("qmatvec_nvfp4_mmvq_dual_b2", ROWS_PER_BLOCK),
6163 (4, false, _) => ("qmatvec_nvfp4_mmvq_dual_b4_r2", ROWS_PER_BLOCK * 2),
6164 (2, true, _) => ("qmatvec_nvfp4_mmvq_dual_b2_rp", ROWS_PER_BLOCK),
6165 (4, true, _) => ("qmatvec_nvfp4_mmvq_dual_b4_rpr2", ROWS_PER_BLOCK * 2),
6166 (8, true, 5) => ("qmatvec_nvfp4_mmvq_dual_b5_rpr2", ROWS_PER_BLOCK * 2),
6167 (8, true, 6) => ("qmatvec_nvfp4_mmvq_dual_b6_rpr2", ROWS_PER_BLOCK * 2),
6168 (8, true, 7) => ("qmatvec_nvfp4_mmvq_dual_b7_rpr2", ROWS_PER_BLOCK * 2),
6169 _ => return Err(format!("qmatvec_batched_dual_raw: no dual kernel for m {m}").into()),
6170 };
6171 let f = self.func(name);
6172 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
6173 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
6174 let cfg = LaunchConfig {
6175 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6176 block_dim: (32, ROWS_PER_BLOCK, 1),
6177 shared_mem_bytes: 0,
6178 };
6179 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
6180 let __s_b = self.gpu.stream();
6181 let mut b = __s_b.launch_builder(&f);
6182 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6183 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
6184 unsafe { b.launch(cfg)?; }
6185 Ok((y0, y1))
6186 }
6187
6188 pub fn matmul_pre_dual_noscale(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6200 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6201 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
6202 use crate::model::GpuTensor;
6203 if m != 1 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6204 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6214 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6215 if w1.in_features() != in_f || w1.out_features() != out_f { return Ok(None); }
6216 let no_mirror = |w: &crate::model::GpuTensor| {
6229 !matches!(w, GpuTensor::Quant { rp4: Some(_), .. })
6230 };
6231 if self.q8_ffn_fuse2_on()
6232 && no_mirror(w0) && no_mirror(w1)
6233 && let Some([p0, p1]) = self.q8_fused_params(&[w0, w1])
6234 {
6235 let (y0, y1) = self.q8_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2)?;
6236 return Ok(Some(((y0, 1.0), (y1, 1.0))));
6237 }
6238 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6248 let (y0, y1) = self.e4m3_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2,
6249 1.0, 1.0)?;
6250 return Ok(Some(((y0, p0.3), (y1, p1.3))));
6251 }
6252 let (b0, q0, rb0, s0, rp0) = match w0 {
6253 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6254 _ => return Ok(None),
6255 };
6256 let (b1, q1, rb1, s1, rp1) = match w1 {
6257 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6258 _ => return Ok(None),
6259 };
6260 if q0 != QT_NVFP4 || q1 != QT_NVFP4 || rb0 != rb1 || rp0 != rp1 { return Ok(None); }
6261 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
6263 let rows_per_block = ROWS_PER_BLOCK * RPW;
6264 let f = self.func(if rp0 { "qmatvec_nvfp4_mmvq_dual_mr2_rp" } else { "qmatvec_nvfp4_mmvq_dual_mr2" });
6265 let mut y0 = self.alloc_uninit::<f32>(out_f)?;
6266 let mut y1 = self.alloc_uninit::<f32>(out_f)?;
6267 let cfg = LaunchConfig {
6268 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6269 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0,
6270 };
6271 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, rb0 as i64);
6272 let one = 1.0f32;
6275 let __s_b = self.gpu.stream();
6276 let mut b = __s_b.launch_builder(&f);
6277 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6278 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&one).arg(&one);
6279 unsafe { b.launch(cfg)?; }
6280 Ok(Some(((y0, s0), (y1, s1))))
6281 }
6282
6283 pub fn matmul_q8_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6291 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6292 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6293 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6299 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(),
6300 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6301 }
6302 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6303 Ok(Some(self.q8_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6304 }
6305
6306 #[allow(clippy::too_many_arguments)]
6307 fn q8_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6308 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6309 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6310 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6311 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6313 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6314 let f = self.func("qmatvec_q8_0_mmvq_fused2");
6315 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6316 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6317 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6318 shared_mem_bytes: 0 };
6319 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
6320 let __s_b = self.gpu.stream();
6321 let mut b = __s_b.launch_builder(&f);
6322 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6323 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl);
6324 unsafe { b.launch(cfg)?; }
6325 Ok((y0, y1))
6326 }
6327
6328 pub fn matmul_q8_fused2_x(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6334 x: &CudaSlice<f32>)
6335 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6336 if !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6337 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6338 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6339 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(),
6340 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6341 }
6342 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6343 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6344 Ok(Some(self.q8_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6345 }
6346
6347 #[allow(clippy::too_many_arguments)]
6350 pub fn qmatvec_q8_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
6351 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6352 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6353 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6354 self.q8_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes)
6355 }
6356
6357 pub fn matmul_q4_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6363 w2: &crate::model::GpuTensor,
6364 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6365 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6366 use crate::model::GpuTensor;
6367 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6368 match w {
6369 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6370 Some((*row_bytes, w.out_features())),
6371 _ => None,
6372 }
6373 };
6374 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6375 else { return Ok(None) };
6376 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6377 return Ok(None);
6378 }
6379 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6383 match w {
6384 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6385 Some(m) => (m, true),
6386 None => (bytes, *rp),
6387 },
6388 _ => unreachable!(),
6389 }
6390 }
6391 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6392 if rp0 != rp1 || rp1 != rp2 { return Ok(None); }
6393 let rp = rp0;
6394 let rpb: u32 = 4;
6395 let mr1 = rp && Self::q40_mr1_on();
6399 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6400 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6401 let grid = nb(o0) + nb(o1) + nb(o2);
6402 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6403 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6404 let mut y2 = self.alloc_uninit::<f32>(o2)?;
6405 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6406 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6407 else { "qmatvec_q4_0_mmvq_fused3" });
6408 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6409 let inf = w0.in_features() as i32;
6410 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6411 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6412 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6415 {
6416 use cudarc::driver::{DevicePtr, DevicePtrMut};
6417 let s = &self.gpu.stream();
6418 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6419 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6420 let (pad, _g4) = ad.device_ptr(s);
6421 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6422 let (py2, _g7) = y2.device_ptr_mut(s);
6423 let mut ps = [
6424 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6425 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6426 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6427 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6428 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6429 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6430 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6431 &r2 as *const _ as *mut _,
6432 ];
6433 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6434 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6435 }
6436 return Ok(Some((y0, y1, y2)));
6437 }
6438 let __s_b = self.gpu.stream();
6439 let mut b = __s_b.launch_builder(&f);
6440 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6441 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6442 unsafe { b.launch(cfg)?; }
6443 Ok(Some((y0, y1, y2)))
6444 }
6445
6446 #[allow(clippy::too_many_arguments)]
6449 pub fn matmul_q4_fused3_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6450 w2: &crate::model::GpuTensor,
6451 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6452 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>,
6453 y2: &mut CudaSlice<f32>)
6454 -> Result<bool, Box<dyn std::error::Error>> {
6455 use crate::model::GpuTensor;
6456 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6457 match w {
6458 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6459 Some((*row_bytes, w.out_features())),
6460 _ => None,
6461 }
6462 };
6463 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6464 else { return Ok(false) };
6465 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6466 return Ok(false);
6467 }
6468 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6469 match w {
6470 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6471 Some(m) => (m, true),
6472 None => (bytes, *rp),
6473 },
6474 _ => unreachable!(),
6475 }
6476 }
6477 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6478 if rp0 != rp1 || rp1 != rp2 { return Ok(false); }
6479 let rp = rp0;
6480 let rpb: u32 = 4;
6481 let mr1 = rp && Self::q40_mr1_on();
6482 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6483 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6484 let grid = nb(o0) + nb(o1) + nb(o2);
6485 debug_assert!(y0.len() >= o0 && y1.len() >= o1 && y2.len() >= o2);
6486 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6487 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6488 else { "qmatvec_q4_0_mmvq_fused3" });
6489 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6490 let inf = w0.in_features() as i32;
6491 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6492 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6493 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6495 use cudarc::driver::{DevicePtr, DevicePtrMut};
6496 let s = &self.gpu.stream();
6497 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6498 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6499 let (pad, _g4) = ad.device_ptr(s);
6500 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6501 let (py2, _g7) = y2.device_ptr_mut(s);
6502 let mut ps = [
6503 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6504 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6505 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6506 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6507 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6508 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6509 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6510 &r2 as *const _ as *mut _,
6511 ];
6512 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6513 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6514 return Ok(true);
6515 }
6516 let __s_b = self.gpu.stream();
6517 let mut b = __s_b.launch_builder(&f);
6518 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1).arg(&mut *y2)
6519 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6520 unsafe { b.launch(cfg)?; }
6521 Ok(true)
6522 }
6523
6524 pub fn matmul_q4_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6526 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6527 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6528 use crate::model::GpuTensor;
6529 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6530 match w {
6531 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6532 Some((*row_bytes, w.out_features())),
6533 _ => None,
6534 }
6535 };
6536 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6537 if w0.in_features() != w1.in_features() { return Ok(None); }
6538 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6540 match w {
6541 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6542 Some(m) => (m, true),
6543 None => (bytes, *rp),
6544 },
6545 _ => unreachable!(),
6546 }
6547 }
6548 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6549 if rp0 != rp1 { return Ok(None); }
6550 let rp = rp0;
6551 let rpb: u32 = 4;
6552 let mr1 = rp && Self::q40_mr1_on();
6554 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6555 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6556 let grid = nb(o0) + nb(o1);
6557 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6558 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6559 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6560 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6561 else { "qmatvec_q4_0_mmvq_fused2" });
6562 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6563 let inf = w0.in_features() as i32;
6564 let (oo0, oo1) = (o0 as i32, o1 as i32);
6565 let (r0, r1) = (rb0 as i64, rb1 as i64);
6566 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6568 {
6569 use cudarc::driver::{DevicePtr, DevicePtrMut};
6570 let s = &self.gpu.stream();
6571 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6572 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6573 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6574 let mut ps = [
6575 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6576 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6577 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6578 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6579 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6580 &r1 as *const _ as *mut _,
6581 ];
6582 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6583 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6584 }
6585 return Ok(Some((y0, y1)));
6586 }
6587 let __s_b = self.gpu.stream();
6588 let mut b = __s_b.launch_builder(&f);
6589 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6590 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6591 unsafe { b.launch(cfg)?; }
6592 Ok(Some((y0, y1)))
6593 }
6594
6595 pub fn matmul_q4_fused2_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6597 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6598 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>)
6599 -> Result<bool, Box<dyn std::error::Error>> {
6600 use crate::model::GpuTensor;
6601 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6602 match w {
6603 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6604 Some((*row_bytes, w.out_features())),
6605 _ => None,
6606 }
6607 };
6608 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(false) };
6609 if w0.in_features() != w1.in_features() { return Ok(false); }
6610 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6611 match w {
6612 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6613 Some(m) => (m, true),
6614 None => (bytes, *rp),
6615 },
6616 _ => unreachable!(),
6617 }
6618 }
6619 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6620 if rp0 != rp1 { return Ok(false); }
6621 let rp = rp0;
6622 let rpb: u32 = 4;
6623 let mr1 = rp && Self::q40_mr1_on();
6624 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6625 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6626 let grid = nb(o0) + nb(o1);
6627 debug_assert!(y0.len() >= o0 && y1.len() >= o1);
6628 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6629 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6630 else { "qmatvec_q4_0_mmvq_fused2" });
6631 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6632 let inf = w0.in_features() as i32;
6633 let (oo0, oo1) = (o0 as i32, o1 as i32);
6634 let (r0, r1) = (rb0 as i64, rb1 as i64);
6635 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6637 use cudarc::driver::{DevicePtr, DevicePtrMut};
6638 let s = &self.gpu.stream();
6639 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6640 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6641 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6642 let mut ps = [
6643 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6644 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6645 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6646 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6647 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6648 &r1 as *const _ as *mut _,
6649 ];
6650 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6651 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6652 return Ok(true);
6653 }
6654 let __s_b = self.gpu.stream();
6655 let mut b = __s_b.launch_builder(&f);
6656 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1)
6657 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6658 unsafe { b.launch(cfg)?; }
6659 Ok(true)
6660 }
6661
6662 pub fn matmul_q4_fused2_batched(&self, w0: &crate::model::GpuTensor,
6667 w1: &crate::model::GpuTensor,
6668 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6669 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6670 use crate::model::GpuTensor;
6671 if m < 2 || m > 8 { return Ok(None); }
6672 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6673 match w {
6674 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6675 Some((*row_bytes, w.out_features())),
6676 _ => None,
6677 }
6678 };
6679 let (Some((rb0, o0)), Some((_rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6680 if w0.in_features() != w1.in_features() { return Ok(None); }
6681 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6682 match w {
6683 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6684 Some(mr) => (mr, true),
6685 None => (bytes, *rp),
6686 },
6687 _ => unreachable!(),
6688 }
6689 }
6690 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6691 if !rp0 || !rp1 { return Ok(None); }
6692 let mcols = Self::batched_mcols(m);
6693 let rpb: u32 = 4;
6694 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6695 let grid = nb(o0) + nb(o1);
6696 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6697 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6698 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f2_rp",
6699 4 => "qmatvec_q4_0_mmvq_b4_f2_rp",
6700 _ => "qmatvec_q4_0_mmvq_b8_f2_rp" });
6701 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6702 shared_mem_bytes: 0 };
6703 let inf = w0.in_features() as i32;
6704 let (oo0, oo1, mi) = (o0 as i32, o1 as i32, m as i32);
6705 let rb = rb0 as i64;
6706 let __s_b = self.gpu.stream();
6707 let mut b = __s_b.launch_builder(&f);
6708 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6709 .arg(&inf).arg(&oo0).arg(&oo1).arg(&mi).arg(&rb);
6710 unsafe { b.launch(cfg)?; }
6711 Ok(Some((y0, y1)))
6712 }
6713
6714 #[allow(clippy::too_many_arguments)]
6717 pub fn matmul_q4_fused3_batched(&self, w0: &crate::model::GpuTensor,
6718 w1: &crate::model::GpuTensor, w2: &crate::model::GpuTensor,
6719 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6720 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6721 use crate::model::GpuTensor;
6722 if m < 2 || m > 8 { return Ok(None); }
6723 let q4 = |w: &GpuTensor| -> Option<usize> {
6724 match w {
6725 GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0 => Some(w.out_features()),
6726 _ => None,
6727 }
6728 };
6729 let (Some(o0), Some(o1), Some(o2)) = (q4(w0), q4(w1), q4(w2)) else { return Ok(None) };
6730 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6731 return Ok(None);
6732 }
6733 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6734 match w {
6735 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6736 Some(mr) => (mr, true),
6737 None => (bytes, *rp),
6738 },
6739 _ => unreachable!(),
6740 }
6741 }
6742 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6743 if !rp0 || !rp1 || !rp2 { return Ok(None); }
6744 let mcols = Self::batched_mcols(m);
6745 let rpb: u32 = 4;
6746 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6747 let grid = nb(o0) + nb(o1) + nb(o2);
6748 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6749 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6750 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
6751 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f3_rp",
6752 4 => "qmatvec_q4_0_mmvq_b4_f3_rp",
6753 _ => "qmatvec_q4_0_mmvq_b8_f3_rp" });
6754 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6755 shared_mem_bytes: 0 };
6756 let inf = w0.in_features() as i32;
6757 let (oo0, oo1, oo2, mi) = (o0 as i32, o1 as i32, o2 as i32, m as i32);
6758 let rb = 0i64;
6759 let __s_b = self.gpu.stream();
6760 let mut b = __s_b.launch_builder(&f);
6761 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6762 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&mi).arg(&rb);
6763 unsafe { b.launch(cfg)?; }
6764 Ok(Some((y0, y1, y2)))
6765 }
6766
6767 pub fn matmul_q8_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6768 w2: &crate::model::GpuTensor,
6769 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6770 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6771 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
6774 return Ok(Some(self.e4m3_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6775 p0.1, p1.1, p2.1, p0.2,
6776 p0.3, p1.3, p2.3)?));
6777 }
6778 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
6779 Ok(Some(self.q8_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6780 p0.1, p1.1, p2.1, p0.2)?))
6781 }
6782
6783 #[allow(clippy::too_many_arguments)]
6784 fn q8_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6785 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6786 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
6787 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6788 const ROWS_PER_BLOCK: u32 = 4;
6789 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6790 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6791 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6792 let f = self.func("qmatvec_q8_0_mmvq_fused3");
6793 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6794 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6795 let mut y2 = self.alloc_uninit::<f32>(out2)?;
6796 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6797 shared_mem_bytes: 0 };
6798 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32, row_bytes as i64);
6799 let __s_b = self.gpu.stream();
6800 let mut b = __s_b.launch_builder(&f);
6801 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6802 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl);
6803 unsafe { b.launch(cfg)?; }
6804 Ok((y0, y1, y2))
6805 }
6806
6807 #[allow(clippy::too_many_arguments)]
6809 pub fn qmatvec_q8_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6810 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
6811 out2: usize, row_bytes: usize)
6812 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6813 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6814 self.q8_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes)
6815 }
6816
6817 pub fn matmul_q8_fused2_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6828 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6829 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6830 if !(2..=8).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
6834 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6837 if m > 4 && !Self::b8_enabled() { return Ok(None); }
6838 return Ok(Some(self.e4m3_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(),
6839 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6840 }
6841 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6842 Ok(Some(self.q8_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(), p0.1, p1.1, p0.2)?))
6843 }
6844
6845 #[allow(clippy::too_many_arguments)]
6846 fn q8_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6847 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6848 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6849 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6850 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6852 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6853 let f = self.func(match Self::batched_mcols(m) {
6854 2 => "qmatvec_q8_0_mmvq_fused2_b2",
6855 4 => "qmatvec_q8_0_mmvq_fused2_b4",
6856 _ => "qmatvec_q8_0_mmvq_fused2_b8",
6858 });
6859 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6860 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6861 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6862 shared_mem_bytes: 0 };
6863 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32, row_bytes as i64);
6864 let __s_b = self.gpu.stream();
6865 let mut b = __s_b.launch_builder(&f);
6866 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6867 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
6868 unsafe { b.launch(cfg)?; }
6869 Ok((y0, y1))
6870 }
6871
6872 #[allow(clippy::too_many_arguments)]
6875 pub fn qmatvec_q8_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6876 x: &CudaSlice<f32>, m: usize,
6877 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6878 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6879 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6880 self.q8_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes)
6881 }
6882
6883 #[allow(clippy::too_many_arguments)]
6886 pub fn matmul_q8_fused3_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6887 w2: &crate::model::GpuTensor,
6888 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6889 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6890 if !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
6891 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
6892 return Ok(Some(self.e4m3_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
6893 p0.1, p1.1, p2.1, p0.2,
6894 p0.3, p1.3, p2.3)?));
6895 }
6896 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
6897 Ok(Some(self.q8_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
6898 p0.1, p1.1, p2.1, p0.2)?))
6899 }
6900
6901 #[allow(clippy::too_many_arguments)]
6902 fn q8_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6903 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6904 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
6905 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6906 const ROWS_PER_BLOCK: u32 = 4;
6907 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6908 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6909 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6910 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_q8_0_mmvq_fused3_b2" }
6911 else { "qmatvec_q8_0_mmvq_fused3_b4" });
6912 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6913 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6914 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
6915 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6916 shared_mem_bytes: 0 };
6917 let (inf, o0, o1, o2, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
6918 m as i32, row_bytes as i64);
6919 let __s_b = self.gpu.stream();
6920 let mut b = __s_b.launch_builder(&f);
6921 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6922 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&mi).arg(&rbl);
6923 unsafe { b.launch(cfg)?; }
6924 Ok((y0, y1, y2))
6925 }
6926
6927 #[allow(clippy::too_many_arguments)]
6929 pub fn qmatvec_q8_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6930 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
6931 out1: usize, out2: usize, row_bytes: usize)
6932 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6933 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6934 self.q8_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes)
6935 }
6936
6937 pub fn q8_ffn_fuse2_on(&self) -> bool {
6941 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6942 *ON.get_or_init(|| std::env::var("MEMRA_Q8_FFN_FUSE2").as_deref() != Ok("0"))
6943 }
6944
6945 #[allow(clippy::type_complexity)]
6951 fn q8_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
6952 -> Option<[(&'w CudaSlice<u8>, usize, usize); N]> {
6953 use crate::model::GpuTensor;
6954 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return None; }
6955 if std::env::var("MEMRA_Q8_DUAL").is_ok_and(|v| v == "0") { return None; }
6956 let in_f = ws[0].in_features();
6957 let mut out: [Option<(&CudaSlice<u8>, usize, usize)>; N] = [None; N];
6958 for (i, w) in ws.iter().enumerate() {
6959 match w {
6960 GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. }
6961 if *qtype == QT_Q8_0 && *scale == 1.0 && w.in_features() == in_f =>
6962 out[i] = Some((bytes, w.out_features(), *row_bytes)),
6963 _ => return None,
6964 }
6965 }
6966 Some(out.map(|o| o.unwrap()))
6967 }
6968
6969 pub fn e4m3_dual_on(&self) -> bool {
6972 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6973 *ON.get_or_init(|| std::env::var("MEMRA_E4M3_DUAL").as_deref() != Ok("0"))
6974 }
6975
6976 #[allow(clippy::type_complexity)]
6988 fn e4m3_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
6989 -> Option<[(&'w CudaSlice<u8>, usize, usize, f32); N]> {
6990 use crate::model::GpuTensor;
6991 if !self.e4m3_dual_on() { return None; }
6992 let in_f = ws[0].in_features();
6993 let mut out: [Option<(&CudaSlice<u8>, usize, usize, f32)>; N] = [None; N];
6994 for (i, w) in ws.iter().enumerate() {
6995 match w {
6996 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, rp4, .. }
6997 if *qtype == QT_F8_E4M3 && w.in_features() == in_f
6998 && *row_bytes == in_f && !*rp && rp4.is_none() =>
6999 out[i] = Some((bytes, w.out_features(), *row_bytes, *scale)),
7000 _ => return None,
7001 }
7002 }
7003 Some(out.map(|o| o.unwrap()))
7004 }
7005
7006 #[allow(clippy::too_many_arguments)]
7010 fn e4m3_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7011 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7012 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7013 ws0: f32, ws1: f32)
7014 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7015 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7017 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7018 let f = self.func("qmatvec_e4m3_mmvq_fused2");
7019 let mut y0 = self.alloc_uninit::<f32>(out0)?;
7020 let mut y1 = self.alloc_uninit::<f32>(out1)?;
7021 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7022 shared_mem_bytes: 0 };
7023 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
7024 let __s_b = self.gpu.stream();
7025 let mut b = __s_b.launch_builder(&f);
7026 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
7027 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl).arg(&ws0).arg(&ws1);
7028 unsafe { b.launch(cfg)?; }
7029 Ok((y0, y1))
7030 }
7031
7032 #[allow(clippy::too_many_arguments)]
7034 fn e4m3_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7035 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7036 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
7037 ws0: f32, ws1: f32, ws2: f32)
7038 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7039 const ROWS_PER_BLOCK: u32 = 4;
7040 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7041 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7042 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
7043 let f = self.func("qmatvec_e4m3_mmvq_fused3");
7044 let mut y0 = self.alloc_uninit::<f32>(out0)?;
7045 let mut y1 = self.alloc_uninit::<f32>(out1)?;
7046 let mut y2 = self.alloc_uninit::<f32>(out2)?;
7047 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7048 shared_mem_bytes: 0 };
7049 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
7050 row_bytes as i64);
7051 let __s_b = self.gpu.stream();
7052 let mut b = __s_b.launch_builder(&f);
7053 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
7054 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl).arg(&ws0).arg(&ws1).arg(&ws2);
7055 unsafe { b.launch(cfg)?; }
7056 Ok((y0, y1, y2))
7057 }
7058
7059 #[allow(clippy::too_many_arguments)]
7063 fn e4m3_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7064 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
7065 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7066 ws0: f32, ws1: f32)
7067 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7068 const ROWS_PER_BLOCK: u32 = 4;
7069 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7070 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7071 let f = self.func(match Self::batched_mcols(m) {
7072 2 => "qmatvec_e4m3_mmvq_fused2_b2",
7073 4 => "qmatvec_e4m3_mmvq_fused2_b4",
7074 _ => "qmatvec_e4m3_mmvq_fused2_b8",
7075 });
7076 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7077 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7078 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7079 shared_mem_bytes: 0 };
7080 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32,
7081 row_bytes as i64);
7082 let __s_b = self.gpu.stream();
7083 let mut b = __s_b.launch_builder(&f);
7084 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
7085 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
7086 unsafe { b.launch(cfg)?; }
7087 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
7088 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
7089 Ok((y0, y1))
7090 }
7091
7092 #[allow(clippy::too_many_arguments)]
7094 fn e4m3_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7095 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
7096 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
7097 ws0: f32, ws1: f32, ws2: f32)
7098 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7099 const ROWS_PER_BLOCK: u32 = 4;
7100 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7101 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7102 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
7103 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_e4m3_mmvq_fused3_b2" }
7104 else { "qmatvec_e4m3_mmvq_fused3_b4" });
7105 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7106 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7107 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
7108 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7109 shared_mem_bytes: 0 };
7110 let (inf, o0, o1, o2, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
7111 m as i32, row_bytes as i64);
7112 let __s_b = self.gpu.stream();
7113 let mut b = __s_b.launch_builder(&f);
7114 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
7115 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&mi).arg(&rbl);
7116 unsafe { b.launch(cfg)?; }
7117 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
7118 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
7119 if ws2 != 1.0 { self.scale_inplace(&mut y2, ws2, m * out2)?; }
7120 Ok((y0, y1, y2))
7121 }
7122
7123 pub fn qmatvec_e4m3_blk_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7133 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7134 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7135 scale_cols: usize)
7136 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7137 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,
7139 scale_cols, &mut y)?;
7140 Ok(y)
7141 }
7142
7143 #[allow(clippy::too_many_arguments)]
7145 pub fn qmatvec_e4m3_blk_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7146 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7147 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7148 scale_cols: usize, y: &mut CudaSlice<f32>)
7149 -> Result<(), Box<dyn std::error::Error>> {
7150 const ROWS_PER_BLOCK: u32 = 4; let f = self.func("qmatvec_e4m3_blk_mmvq");
7152 let cfg = LaunchConfig {
7153 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), m as u32, 1),
7154 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7157 let (inf, outf, mi, rb, sc) =
7158 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7159 let __s_b = self.gpu.stream();
7160 let mut b = __s_b.launch_builder(&f);
7161 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut *y)
7162 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7163 unsafe { b.launch(cfg)?; }
7164 Ok(())
7165 }
7166
7167 #[allow(clippy::too_many_arguments)]
7173 pub fn qmatvec_e4m3_blk_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7174 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7175 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7176 scale_cols: usize, mcols: usize)
7177 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7178 const ROWS_PER_BLOCK: u32 = 4; debug_assert!(mcols >= m, "blk batched: mcols {mcols} < m {m}");
7180 let name = match mcols {
7181 2 => "qmatvec_e4m3_blk_mmvq_b2",
7182 4 => "qmatvec_e4m3_blk_mmvq_b4",
7183 8 => "qmatvec_e4m3_blk_mmvq_b8",
7184 16 => "qmatvec_e4m3_blk_mmvq_b16",
7185 _ => return Err(format!("qmatvec_e4m3_blk_mmvq_batched: no kernel for mcols {mcols}").into()),
7186 };
7187 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7188 let f = self.func(name);
7189 let cfg = LaunchConfig {
7190 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
7191 block_dim: (32, ROWS_PER_BLOCK, 1),
7192 shared_mem_bytes: 0,
7193 };
7194 let (inf, outf, mi, rb, sc) =
7195 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7196 let __s_b = self.gpu.stream();
7197 let mut b = __s_b.launch_builder(&f);
7198 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut y)
7199 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7200 unsafe { b.launch(cfg)?; }
7201 Ok(y)
7202 }
7203
7204 #[allow(clippy::too_many_arguments)]
7207 pub fn qmatvec_e4m3_blk_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7208 scales: &CudaSlice<f32>, m: usize, in_f: usize,
7209 out_f: usize, row_bytes: usize, scale_cols: usize,
7210 mcols: usize)
7211 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7212 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7213 self.qmatvec_e4m3_blk_mmvq_batched(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes,
7214 scale_cols, mcols)
7215 }
7216
7217 #[allow(clippy::too_many_arguments)]
7220 pub fn qmatvec_e4m3_blk_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7221 scales: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
7222 row_bytes: usize, scale_cols: usize)
7223 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7224 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7225 self.qmatvec_e4m3_blk_mmvq(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols)
7226 }
7227
7228 #[allow(clippy::too_many_arguments)]
7231 pub fn qmatvec_e4m3_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
7232 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7233 ws0: f32, ws1: f32)
7234 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7235 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7236 self.e4m3_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes, ws0, ws1)
7237 }
7238
7239 #[allow(clippy::too_many_arguments)]
7240 pub fn qmatvec_e4m3_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7241 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
7242 out2: usize, row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7243 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7244 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7245 self.e4m3_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes,
7246 ws0, ws1, ws2)
7247 }
7248
7249 #[allow(clippy::too_many_arguments)]
7250 pub fn qmatvec_e4m3_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7251 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
7252 out1: usize, row_bytes: usize, ws0: f32, ws1: f32)
7253 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7254 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7255 self.e4m3_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes, ws0, ws1)
7256 }
7257
7258 #[allow(clippy::too_many_arguments)]
7259 pub fn qmatvec_e4m3_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7260 b2: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
7261 in_f: usize, out0: usize, out1: usize, out2: usize,
7262 row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7263 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7264 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7265 self.e4m3_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes,
7266 ws0, ws1, ws2)
7267 }
7268
7269 fn try_e4m3_blk_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
7280 ad: &CudaSlice<f32>, m: usize)
7281 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7282 use crate::model::GpuTensor;
7283 if let GpuTensor::Quant { bytes, qtype, row_bytes, blk: Some(g), .. } = w {
7284 if *qtype == QT_F8_E4M3_BLK {
7285 if (2..=16).contains(&m) && std::env::var("MEMRA_NO_BATCHED").is_err()
7291 && (m <= 4 || Self::b8_enabled()) {
7292 let mcols = Self::batched_mcols(m);
7293 return Ok(Some(self.qmatvec_e4m3_blk_mmvq_batched(
7294 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7295 *row_bytes, g.cols, mcols)?));
7296 }
7297 return Ok(Some(self.qmatvec_e4m3_blk_mmvq(
7298 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7299 *row_bytes, g.cols)?));
7300 }
7301 }
7302 Ok(None)
7303 }
7304
7305 fn try_e4m3_blk_prefill(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
7352 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7353 use crate::model::GpuTensor;
7354 let GpuTensor::Quant { bytes, qtype, blk: Some(g), .. } = w else { return Ok(None) };
7355 if *qtype != QT_F8_E4M3_BLK { return Ok(None) }
7356 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(Some(y)); }
7361 let (in_f, out_f) = (w.in_features(), w.out_features());
7362 let slab = self.fp8_blk_dequant_q8_0_dev(bytes, &g.scales, out_f, in_f)?;
7363 let tmp = GpuTensor::Quant {
7364 bytes: slab,
7365 qtype: QT_Q8_0,
7366 row_bytes: in_f / 32 * 34,
7367 ne: vec![in_f as u64, out_f as u64],
7368 scale: 1.0,
7369 rp: false,
7370 #[cfg(memra_cutlass)]
7371 cutlass: None,
7372 fp8: None, blk: None, f16: None, rp4: None,
7373 };
7374 Ok(Some(self.matmul(&tmp, x, m)?))
7376 }
7377
7378 pub fn matmul_pre_noscale(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7379 m: usize) -> Result<Option<(CudaSlice<f32>, f32)>, Box<dyn std::error::Error>> {
7380 use crate::model::GpuTensor;
7381 if m == 1 {
7385 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(Some((y, 1.0))); }
7386 }
7387 if m != 1 || !self.uses_q8_1_fast(w) { return Ok(None); }
7389 let in_f = w.in_features();
7390 let out_f = w.out_features();
7391 let (bytes, qtype, row_bytes, scale, rp) = match w {
7392 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
7393 _ => return Ok(None),
7394 };
7395 if self.mmvq_supports(qtype) {
7397 let (mbytes, mrp) = match w {
7399 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
7400 _ => (bytes, rp),
7401 };
7402 let y = self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, 1.0, mrp)?;
7403 return Ok(Some((y, scale)));
7404 }
7405 let name = match qtype {
7407 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
7408 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
7409 QT_Q3_K => "qmatvec_q3_K_dp4a",
7410 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
7411 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
7412 _ => return Ok(None),
7413 };
7414 let f = self.func(name);
7415 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7416 let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
7417 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7418 let __s_b = self.gpu.stream();
7419 let mut b = __s_b.launch_builder(&f);
7420 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7421 unsafe { b.launch(cfg)?; }
7422 Ok(Some((y, scale)))
7423 }
7424
7425 pub fn mmvq_supports(&self, qtype: i32) -> bool {
7428 if qtype == QT_F8_E4M3 { return true; }
7433 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return false; }
7434 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_Q4_0)
7435 }
7436
7437 pub fn qmatvec_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7442 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7443 rp: bool)
7444 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7445 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)?;
7447 Ok(y)
7448 }
7449
7450 #[allow(clippy::too_many_arguments)]
7452 pub fn qmatvec_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7453 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7454 rp: bool, y: &mut CudaSlice<f32>)
7455 -> Result<(), Box<dyn std::error::Error>> {
7456 debug_assert!(y.len() >= m * out_f);
7457 const ROWS_PER_BLOCK: u32 = 4; if qtype == QT_Q8_0 && rp && m == 1 && out_f >= 64
7463 && (out_f as u32).div_ceil(ROWS_PER_BLOCK) < 4 * self.sm_count() as u32
7464 && {
7465 static G2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7466 *G2.get_or_init(|| std::env::var("MEMRA_Q80_G2").as_deref() != Ok("0"))
7467 }
7468 {
7469 let f = self.func("qmatvec_q8_0_mmvq_rp_g2");
7470 let cfg = LaunchConfig {
7471 grid_dim: ((out_f as u32).div_ceil(2), 1, 1),
7472 block_dim: (32, 2, 1),
7473 shared_mem_bytes: 0,
7474 };
7475 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, row_bytes as i64);
7476 let __s_b = self.gpu.stream();
7477 let mut b = __s_b.launch_builder(&f);
7478 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7479 unsafe { b.launch(cfg)?; }
7480 if scale != 1.0 { self.scale_inplace(y, scale, out_f)?; }
7481 return Ok(());
7482 }
7483 let mut mr: u32 = if m == 1 && (qtype == QT_NVFP4 || qtype == QT_Q5_K) { 2 } else { 1 };
7492 if m == 1 && qtype == QT_Q4_0 {
7497 static Q40MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7498 mr = *Q40MR.get_or_init(|| std::env::var("MEMRA_Q40_MR").ok()
7501 .and_then(|v| v.parse().ok()).unwrap_or(1));
7502 }
7503 let q5_mode = std::env::var("MEMRA_Q5K_ISSUE").ok();
7514 let q5_force = q5_mode.as_deref() == Some("2");
7515 let q5_il = qtype == QT_Q5_K && m == 1
7518 && (q5_force || q5_mode.as_deref().map(|v| v != "0").unwrap_or(true));
7519 if q5_il && !q5_force && out_f > 65536 { mr = 1; }
7520 if qtype == QT_Q4_0 && rp && mr != 1 { mr = 2; }
7523 if qtype == QT_Q8_0 && rp {
7527 static Q80MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7528 mr = *Q80MR.get_or_init(|| std::env::var("MEMRA_Q80_MR").ok()
7529 .and_then(|v| v.parse().ok()).unwrap_or(1));
7530 }
7531 let name = match (qtype, mr, rp) {
7532 (QT_NVFP4, 2, false) => "qmatvec_nvfp4_mmvq_mr2",
7533 (QT_NVFP4, 2, true) => "qmatvec_nvfp4_mmvq_mr2_rp",
7534 (QT_NVFP4, _, true) => "qmatvec_nvfp4_mmvq_rp",
7535 (QT_Q4_0, 1, true) => "qmatvec_q4_0_mmvq_rp",
7536 (QT_Q4_0, _, true) => "qmatvec_q4_0_mmvq_mr2_rp",
7537 (QT_Q5_K, 2, _) => if q5_il { "qmatvec_q5_K_mmvq_mr2_il" } else { "qmatvec_q5_K_mmvq_mr2" },
7538 (QT_Q8_0, 2, true) => "qmatvec_q8_0_mmvq_mr2_rp",
7539 (QT_Q8_0, _, true) if in_f % 1024 == 0 && {
7544 static CA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7545 *CA.get_or_init(|| std::env::var("MEMRA_Q80_CA").as_deref() == Ok("1"))
7546 } => "qmatvec_q8_0_mmvq_rpca",
7547 (QT_Q8_0, _, true) => "qmatvec_q8_0_mmvq_rp",
7548 (QT_Q8_0, _, _) => "qmatvec_q8_0_mmvq",
7549 (QT_Q4_K, _, true) => "qmatvec_q4_K_mmvq_rp",
7553 (QT_Q6_K, _, true) => "qmatvec_q6_K_mmvq_rp",
7554 (QT_Q4_K, _, _) => "qmatvec_q4_K_mmvq",
7555 (QT_Q4_0, 2, false) => "qmatvec_q4_0_mmvq_mr2",
7556 (QT_Q4_0, _, false) => "qmatvec_q4_0_mmvq",
7557 (QT_Q5_K, _, _) => if q5_il { "qmatvec_q5_K_mmvq_il" } else { "qmatvec_q5_K_mmvq" },
7558 (QT_Q6_K, _, _) => "qmatvec_q6_K_mmvq",
7559 (QT_NVFP4, _, false) => "qmatvec_nvfp4_mmvq",
7560 (QT_F8_E4M3, _, _) => "qmatvec_e4m3_mmvq",
7561 _ => panic!("qmatvec_mmvq: qtype {qtype} has no MMVQ kernel"),
7562 };
7563 let f = self.func(name);
7564 let rows_per_block = ROWS_PER_BLOCK * mr;
7566 let cfg = LaunchConfig {
7567 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, m as u32, 1),
7568 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7571 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7572 let __s_b = self.gpu.stream();
7573 let mut b = __s_b.launch_builder(&f);
7574 if qtype == QT_NVFP4 || qtype == QT_F8_E4M3 {
7579 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&scale);
7580 unsafe { b.launch(cfg)?; }
7581 } else if Self::pdl_on() && Self::pdl_mmvq_on()
7582 && matches!(name, "qmatvec_q4_0_mmvq_rp" | "qmatvec_q6_K_mmvq"
7583 | "qmatvec_q6_K_mmvq_rp") {
7584 {
7588 use cudarc::driver::{DevicePtr, DevicePtrMut};
7589 let s = &self.gpu.stream();
7590 let (pw, _g0) = bytes.device_ptr(s); let (paq, _g1) = aq.device_ptr(s);
7591 let (pad, _g2) = ad.device_ptr(s); let (py, _g3) = y.device_ptr_mut(s);
7592 let mut ps = [
7593 &pw as *const _ as *mut std::ffi::c_void, &paq as *const _ as *mut _,
7594 &pad as *const _ as *mut _, &py as *const _ as *mut _,
7595 &inf as *const _ as *mut _, &outf as *const _ as *mut _,
7596 &mi as *const _ as *mut _, &rb as *const _ as *mut _,
7597 ];
7598 unsafe { self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?; }
7599 }
7600 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7601 } else {
7602 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7603 unsafe { b.launch(cfg)?; }
7604 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7605 }
7606 Ok(())
7607 }
7608
7609 pub fn qmatvec_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
7613 out_f: usize, qtype: i32, row_bytes: usize, rp: bool)
7614 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7615 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7616 self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, 1.0, rp)
7617 }
7618
7619 pub fn batched_supports(&self, qtype: i32) -> bool {
7623 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_F8_E4M3 | QT_Q4_0)
7624 }
7625
7626 pub fn iq_fast_enabled() -> bool {
7634 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7635 *ON.get_or_init(|| std::env::var("MEMRA_IQ_FAST").map(|v| v != "0").unwrap_or(true))
7636 }
7637
7638 pub fn b8_enabled() -> bool {
7641 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7642 *ON.get_or_init(|| std::env::var("MEMRA_B8").map(|v| v != "0").unwrap_or(true))
7643 }
7644
7645 pub fn batched_mcols(m: usize) -> usize {
7647 if m == 2 { 2 } else if m <= 4 { 4 } else if m <= 8 { 8 } else { 16 }
7648 }
7649
7650 fn batched_kernel_name(qtype: i32, mcols: usize) -> Option<&'static str> {
7655 Some(match (qtype, mcols) {
7656 (QT_Q8_0, 2) => "qmatvec_q8_0_mmvq_b2", (QT_Q8_0, 4) => "qmatvec_q8_0_mmvq_b4",
7657 (QT_Q8_0, 8) => "qmatvec_q8_0_mmvq_b8",
7658 (QT_Q8_0, 16) => "qmatvec_q8_0_mmvq_b16",
7664 (QT_Q4_K, 2) => "qmatvec_q4_K_mmvq_b2", (QT_Q4_K, 4) => "qmatvec_q4_K_mmvq_b4",
7665 (QT_Q4_K, 8) => "qmatvec_q4_K_mmvq_b8",
7666 (QT_Q4_K, 16) => "qmatvec_q4_K_mmvq_b16",
7669 (QT_Q5_K, 2) => "qmatvec_q5_K_mmvq_b2", (QT_Q5_K, 4) => "qmatvec_q5_K_mmvq_b4",
7670 (QT_Q5_K, 8) => "qmatvec_q5_K_mmvq_b8",
7671 (QT_Q5_K, 16) => "qmatvec_q5_K_mmvq_b16",
7674 (QT_Q6_K, 2) => "qmatvec_q6_K_mmvq_b2", (QT_Q6_K, 4) => "qmatvec_q6_K_mmvq_b4",
7675 (QT_Q6_K, 8) => "qmatvec_q6_K_mmvq_b8", (QT_Q6_K, 16) => "qmatvec_q6_K_mmvq_b16",
7676 (QT_NVFP4, 2) => "qmatvec_nvfp4_mmvq_b2", (QT_NVFP4, 4) => "qmatvec_nvfp4_mmvq_b4",
7677 (QT_NVFP4, 8) => "qmatvec_nvfp4_mmvq_b8",
7678 (QT_NVFP4, 16) => "qmatvec_nvfp4_mmvq_b16",
7682 (QT_F8_E4M3, 2) => "qmatvec_e4m3_mmvq_b2", (QT_F8_E4M3, 4) => "qmatvec_e4m3_mmvq_b4",
7683 (QT_F8_E4M3, 8) => "qmatvec_e4m3_mmvq_b8",
7684 (QT_F8_E4M3, 16) => "qmatvec_e4m3_mmvq_b16",
7688 (QT_Q4_0, 2) => "qmatvec_q4_0_mmvq_b2", (QT_Q4_0, 4) => "qmatvec_q4_0_mmvq_b4",
7689 (QT_Q4_0, 8) => "qmatvec_q4_0_mmvq_b8", (QT_Q4_0, 16) => "qmatvec_q4_0_mmvq_b16",
7690 _ => return None,
7691 })
7692 }
7693
7694 pub fn sm_count(&self) -> i32 {
7729 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7730 *SMS.get_or_init(|| {
7731 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7732 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7733 })
7734 }
7735
7736 pub fn batched_variant(&self, _m: usize, in_f: usize, out_f: usize, qtype: i32,
7737 row_bytes: usize, mcols: usize, rp: bool) -> &'static str {
7738 if qtype == QT_Q8_0 {
7743 return if rp { "rp" } else { "base" };
7744 }
7745 static BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7746 let bv = *BV.get_or_init(|| match std::env::var("MEMRA_MMVQ_BV").as_deref() {
7747 Ok("base") => "base", Ok("pf") => "pf", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7748 Ok("pfr2") => "pfr2", Ok("ca") => "ca", Ok("car2") => "car2",
7749 Ok("rp") => "rp", Ok("rpr2") => "rpr2", Ok("rpr2w8") => "rpr2w8",
7752 Ok("rpca") => "rpca", Ok("rpcar2") => "rpcar2",
7755 Ok("rpsc") => "rpsc", Ok("rpms") => "rpms", Ok("rpmsc") => "rpmsc",
7762 Ok("rpks") => "rpks", Ok("rpksc") => "rpksc",
7763 _ => "auto",
7764 });
7765 let ca_ok = qtype == QT_NVFP4 && (row_bytes % 16 == 0) && (in_f % 1024 == 0);
7769 static KS_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7774 let ks_on = *KS_ON.get_or_init(|| std::env::var("MEMRA_KS").as_deref() != Ok("0"));
7775 let sc_ok = ks_on && qtype == QT_NVFP4 && (in_f % 256 == 0) && (in_f / 64 <= 272);
7776 let ks_ok = ks_on && qtype == QT_NVFP4 && (in_f % 512 == 0) && (in_f / 64 <= 272);
7777 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7778 let sms = *SMS.get_or_init(|| {
7779 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7780 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7781 });
7782 let kq_r2 = matches!(qtype, QT_Q4_K | QT_Q5_K | QT_Q6_K);
7802 static KQBV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7805 let kq_bv = *KQBV.get_or_init(|| match std::env::var("MEMRA_KQ_BV").as_deref() {
7806 Ok("base") => "base", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7807 _ => "auto",
7808 });
7809 let variant: &'static str = if qtype == QT_Q4_0 {
7810 static Q40BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7814 let q40 = *Q40BV.get_or_init(|| match std::env::var("MEMRA_Q40_BV").as_deref() {
7815 Ok("base") => "base", Ok("r2") => "r2", Ok("ms") => "ms", Ok("sm") => "sm",
7821 Ok("la") => "la", _ => "auto",
7822 });
7823 let v = if q40 != "auto" { q40 }
7824 else if (out_f as u32).div_ceil(8) >= 4 * sms as u32 { "r2" } else { "base" };
7825 if rp { match v { "ms" => "r2ms_rp", "sm" => "r2sm_rp", "la" => "r2la_rp",
7830 "r2" => "r2_rp", _ => "rp" } }
7831 else if matches!(v, "ms" | "sm" | "la") { "r2" } else { v }
7832 } else if qtype != QT_NVFP4 && !kq_r2 {
7833 "base"
7834 } else if kq_r2 && rp {
7835 "rp"
7839 } else if kq_r2 {
7840 if kq_bv != "auto" {
7843 if kq_bv == "r2w8" && mcols != 4 { "r2" } else { kq_bv }
7844 } else if bv != "auto" {
7845 match bv {
7846 "r2" | "pfr2" | "rpr2" | "car2" => "r2",
7847 "r2w8" | "rpr2w8" => if mcols != 4 { "r2" } else { "r2w8" },
7848 _ => "base", }
7850 } else {
7851 let blocks = (out_f + 7) / 8;
7852 let waves = blocks as f64 / (7 * sms as usize) as f64;
7853 let filled = blocks >= 4 * sms as usize;
7854 let use_r2 = if qtype == QT_Q4_K { filled } else { waves >= 2.0 };
7855 if use_r2 { "r2" } else { "base" }
7856 }
7857 } else if bv != "auto" {
7858 let v = if bv == "r2w8" && mcols == 2 { "r2" }
7863 else if bv == "ca" && (!ca_ok || mcols == 8) { "pf" }
7864 else if bv == "car2" && (!ca_ok || mcols == 8) { "r2" }
7865 else if bv == "pfr2" && mcols == 8 { "r2" }
7866 else if (bv == "rpr2w8" || bv == "rpr2") && mcols == 2 { "rpr2" }
7867 else if (bv == "rpca" || bv == "rpcar2") && (!ca_ok || mcols == 8) {
7869 if mcols == 8 { "rpr2w8" } else { "rpr2" }
7870 }
7871 else if bv == "rpcar2" && mcols == 2 { "rpca" }
7872 else if (bv == "rpsc" || bv == "rpmsc") && !sc_ok { "rpr2" }
7875 else if (bv == "rpks" || bv == "rpksc") && !ks_ok { "rpr2" }
7876 else { bv };
7877 if rp {
7878 match v {
7879 "base" | "pf" | "ca" | "rp" => "rp",
7880 "r2" | "pfr2" | "car2" | "rpr2" => "rpr2",
7881 "r2w8" | "rpr2w8" => if mcols == 2 { "rpr2" } else { "rpr2w8" },
7882 other => other, }
7884 } else { v }
7885 } else if mcols == 8 {
7886 if rp { if sc_ok { "rpsc" } else { "rpr2w8" } } else { "r2w8" }
7897 } else if mcols >= 4 {
7898 let blocks = (out_f + 7) / 8;
7902 let r7 = 7 * sms as usize;
7903 let r8 = 8 * sms as usize;
7904 let waves = blocks as f64 / r7 as f64;
7905 let filled = blocks >= 4 * sms as usize;
7906 if filled && blocks.div_ceil(r8) < blocks.div_ceil(r7) {
7910 if rp { "rpr2w8" } else { "r2w8" }
7914 } else if waves >= 2.0 || (waves <= 1.0 && filled) {
7915 if rp { "rpr2" } else { "r2" }
7918 } else {
7919 if rp { "rp" } else { "pf" }
7923 }
7924 } else if in_f >= 6144 {
7925 if rp { "rpr2" } else { "r2" }
7929 }
7930 else if rp {
7931 let waves = ((out_f + 7) / 8) as f64 / (7 * sms as usize) as f64;
7936 if sc_ok && waves >= 0.9 && waves <= 1.1 { "rpsc" } else { "rp" }
7937 } else { "base" };
7938 variant
7939 }
7940
7941 pub fn qmatvec_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7942 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize,
7943 mcols: usize, scale: f32, rp: bool)
7944 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7945 const ROWS_PER_BLOCK: u32 = 4;
7946 let forced: Option<&'static str> = {
7951 static V: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
7952 V.get_or_init(|| std::env::var("MEMRA_BVAR").ok())
7953 .as_deref()
7954 .map(|s| Box::leak(s.to_string().into_boxed_str()) as &'static str)
7955 };
7956 let variant = match forced {
7957 Some(v) if !rp || v.contains("rp") => v,
7958 _ => self.batched_variant(m, in_f, out_f, qtype, row_bytes, mcols, rp),
7959 };
7960 let base_name = Self::batched_kernel_name(qtype, mcols)
7961 .ok_or_else(|| format!("qmatvec_mmvq_batched: no kernel for qtype {qtype} mcols {mcols}"))?;
7962 let variant = if mcols == 16 { if rp { "rp" } else { "base" } } else { variant };
7966 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7973 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
7974 if b567 && qtype == QT_NVFP4 && rp && mcols == 8 && (5..=7).contains(&m)
7975 && matches!(variant, "rpsc" | "rpr2w8") {
7976 let f = self.func(&format!("qmatvec_nvfp4_mmvq_b{m}_{variant}"));
7977 let rows_per_block = ROWS_PER_BLOCK * 2; let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7979 let cfg = LaunchConfig {
7980 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
7981 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0 };
7982 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7983 let __s_b = self.gpu.stream();
7984 let mut b = __s_b.launch_builder(&f);
7985 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7986 unsafe { b.launch(cfg)?; }
7987 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
7988 return Ok(y);
7989 }
7990 let (name, rows_per_block): (std::borrow::Cow<'static, str>, u32) = match variant {
7991 "base" => (base_name.into(), ROWS_PER_BLOCK),
7992 "pf" => (format!("{base_name}_pf").into(), ROWS_PER_BLOCK),
7993 "ca" => (format!("{base_name}_ca").into(), ROWS_PER_BLOCK),
7994 "rp" => (format!("{base_name}_rp").into(), ROWS_PER_BLOCK),
7995 "rpca" => (format!("{base_name}_rpca").into(), ROWS_PER_BLOCK), "rpks" => (format!("{base_name}_rpks").into(), ROWS_PER_BLOCK),
7999 "rpksc" => (format!("{base_name}_rpksc").into(), ROWS_PER_BLOCK),
8000 "rpms" => (format!("{base_name}_rpms").into(), ROWS_PER_BLOCK),
8001 "rpmsc" => (format!("{base_name}_rpmsc").into(), ROWS_PER_BLOCK),
8002 "r2ms_rp" => (format!("{base_name}_r2ms_rp").into(), ROWS_PER_BLOCK),
8003 "r2sm_rp" => (format!("{base_name}_r2sm_rp").into(), ROWS_PER_BLOCK * 2),
8004 "r2la_rp" => (format!("{base_name}_r2la_rp").into(), ROWS_PER_BLOCK * 2),
8005 v => (format!("{base_name}_{v}").into(), ROWS_PER_BLOCK * 2), };
8007 debug_assert!(!rp || name.contains("_rp"), "rp weight dispatched to a GGUF-layout kernel");
8008 let f = self.func(&name);
8009 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8010 let smem = if name.contains("_r2sm_rp") { (mcols * 32 * 9 * 4 + mcols * 32 * 4) as u32 }
8012 else { 0 };
8013 let cfg = LaunchConfig {
8014 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
8015 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: smem };
8016 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8017 let __s_b = self.gpu.stream();
8018 let mut b = __s_b.launch_builder(&f);
8019 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8020 unsafe { b.launch(cfg)?; }
8021 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8022 Ok(y)
8023 }
8024
8025 pub fn qmatvec_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
8029 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, mcols: usize,
8030 rp: bool)
8031 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8032 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
8033 self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, 1.0, rp)
8034 }
8035
8036 pub fn qmatvec_nvfp4_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
8038 in_f: usize, out_f: usize, row_bytes: usize, mcols: usize,
8039 rp: bool)
8040 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8041 self.qmatvec_batched_raw(bytes, x, m, in_f, out_f, QT_NVFP4, row_bytes, mcols, rp)
8042 }
8043
8044 fn try_fp4_gemm(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize,
8048 in_f: usize, out_f: usize)
8049 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8050 use crate::model::GpuTensor;
8051 if cfg!(memra_portable_cuda) { return Ok(None); }
8052 if std::env::var("MEMRA_FP4").is_err() { return Ok(None); }
8053 #[cfg(memra_cutlass)]
8062 if m >= 128 && std::env::var("MEMRA_FP4_CUTLASS").is_ok() {
8063 if let GpuTensor::Quant { bytes, qtype, scale, row_bytes, cutlass, .. } = w {
8064 if *qtype == QT_NVFP4 && in_f % 64 == 0 {
8065 if let Some(cw) = cutlass {
8066 let y = self.cutlass_fp4_gemm(&cw.b_packed, &cw.sfb_swizzled, x, *scale,
8068 m, out_f, in_f)?;
8069 return Ok(Some(y));
8070 } else if std::env::var("MEMRA_FP4_CUTLASS_OTF").is_ok() {
8071 let (b_packed, sfb_sw) = self.build_cutlass_weight(bytes, out_f, in_f, *row_bytes)?;
8076 let y = self.cutlass_fp4_gemm(&b_packed, &sfb_sw, x, *scale, m, out_f, in_f)?;
8077 return Ok(Some(y));
8078 }
8079 }
8080 }
8081 }
8082 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } = w {
8083 if *qtype == QT_NVFP4 && in_f % 64 == 0 && !*rp {
8086 let y = self.qmatvec_gemm_nvfp4_fp4(bytes, x, m, in_f, out_f, *row_bytes, *scale)?;
8087 return Ok(Some(y));
8088 }
8089 }
8090 Ok(None)
8091 }
8092
8093 pub fn rms_norm_f16out(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>,
8097 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
8098 ncols: usize, nrows: usize, eps: f32)
8099 -> Result<(), Box<dyn std::error::Error>> {
8100 let f = self.func("rms_norm_f16out_f32");
8101 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
8102 let (nc, e) = (ncols as i32, eps);
8103 let __s_b = self.gpu.stream();
8104 let mut b = __s_b.launch_builder(&f);
8105 b.arg(x).arg(w).arg(dst).arg(dst16).arg(&nc).arg(&e);
8106 unsafe { b.launch(cfg)?; }
8107 Ok(())
8108 }
8109
8110 #[allow(clippy::too_many_arguments)]
8113 pub fn add_rms_norm_f16out(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
8114 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
8115 dst16: &mut CudaSlice<u8>, ncols: usize, nrows: usize, eps: f32)
8116 -> Result<(), Box<dyn std::error::Error>> {
8117 let f = self.func("add_rms_norm_f16out_f32");
8118 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
8119 let (nc, e) = (ncols as i32, eps);
8120 let __s_lb = self.gpu.stream();
8121 let mut lb = __s_lb.launch_builder(&f);
8122 lb.arg(a).arg(b).arg(w).arg(res).arg(dst).arg(dst16).arg(&nc).arg(&e);
8123 unsafe { lb.launch(cfg)?; }
8124 Ok(())
8125 }
8126
8127 pub fn matmul_group_xh(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>,
8130 xh: &CudaSlice<u8>, m: usize)
8131 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8132 let mut out = Vec::with_capacity(ws.len());
8133 let in_f = ws[0].in_features();
8134 for w in ws {
8135 if w.in_features() == in_f && m >= 16 && !self.verify_exact_on() {
8136 if let Some(y) = self.try_f16_gemm_pre(w, xh, m)? {
8137 out.push(y);
8138 continue;
8139 }
8140 }
8141 out.push(self.matmul(w, x, m)?);
8142 }
8143 Ok(out)
8144 }
8145
8146 pub fn gdn_pad_mask(&self, beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
8149 len_d: &CudaSlice<i32>, h: usize, t: usize)
8150 -> Result<(), Box<dyn std::error::Error>> {
8151 let f = self.func("gdn_pad_mask_f32");
8152 let cfg = LaunchConfig::for_num_elems((t * h) as u32);
8153 let (hi, ti) = (h as i32, t as i32);
8154 let __s_b = self.gpu.stream();
8155 let mut b = __s_b.launch_builder(&f);
8156 b.arg(beta).arg(g_log).arg(len_d).arg(&hi).arg(&ti);
8157 unsafe { b.launch(cfg)?; }
8158 Ok(())
8159 }
8160
8161 pub fn row_gather_dev(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
8164 len_d: &CudaSlice<i32>, ncols: usize)
8165 -> Result<(), Box<dyn std::error::Error>> {
8166 let f = self.func("row_gather_dev_f32");
8167 let cfg = LaunchConfig::for_num_elems(ncols as u32);
8168 let nc = ncols as i32;
8169 let __s_b = self.gpu.stream();
8170 let mut b = __s_b.launch_builder(&f);
8171 b.arg(src).arg(dst).arg(len_d).arg(&nc);
8172 unsafe { b.launch(cfg)?; }
8173 Ok(())
8174 }
8175
8176 pub fn matmul_group(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>, m: usize)
8183 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8184 use crate::model::GpuTensor;
8185 let mut out = Vec::with_capacity(ws.len());
8186 let any_mirror = ws.iter().any(|w| matches!(w, GpuTensor::Quant { f16: Some(_), .. }));
8187 if m >= 16 && any_mirror && !self.verify_exact_on() {
8188 let in_f = ws[0].in_features();
8189 let xh = self.f16_act(x, m * in_f, in_f)?;
8190 for w in ws {
8191 if w.in_features() == in_f {
8192 if let Some(y) = self.try_f16_gemm_pre(w, &xh, m)? {
8193 out.push(y);
8194 continue;
8195 }
8196 }
8197 out.push(self.matmul(w, x, m)?);
8198 }
8199 return Ok(out);
8200 }
8201 for w in ws {
8202 out.push(self.matmul(w, x, m)?);
8203 }
8204 Ok(out)
8205 }
8206
8207 pub fn matmul_group_multi(&self, ws: &[&crate::model::GpuTensor],
8214 xs: &[&CudaSlice<f32>], ms: &[usize])
8215 -> Result<Vec<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
8216 assert_eq!(xs.len(), ms.len());
8217 let in_f = ws[0].in_features();
8218 let total: usize = ms.iter().sum();
8219 let mut xcat = self.uninit(total * in_f)?;
8220 let mut off = 0usize;
8221 for (x, &m) in xs.iter().zip(ms) {
8222 self.copy_into(&mut xcat, off * in_f, x, m * in_f)?;
8223 off += m;
8224 }
8225 let ys = self.matmul_group(ws, &xcat, total)?;
8226 let mut out: Vec<Vec<CudaSlice<f32>>> = (0..xs.len()).map(|_| Vec::new()).collect();
8227 for (w, y) in ws.iter().zip(ys) {
8228 let out_f = w.out_features();
8229 let mut off = 0usize;
8230 for (s, &m) in ms.iter().enumerate() {
8231 let mut ys_s = self.uninit(m * out_f)?;
8232 let src = y.slice(off * out_f..(off + m) * out_f);
8233 self.gpu.stream().memcpy_dtod(&src, &mut ys_s)?;
8234 out[s].push(ys_s);
8235 off += m;
8236 }
8237 }
8238 Ok(out)
8239 }
8240
8241 pub fn gemm_supports(&self, w: &crate::model::GpuTensor) -> bool {
8251 use crate::model::GpuTensor;
8252 if !legacy_quant_gemm_allowed(
8253 cfg!(memra_portable_cuda),
8254 cfg!(memra_hopper_mma),
8255 std::env::var_os("MEMRA_NO_GEMM").is_some(),
8256 ) {
8257 return false;
8258 }
8259 match w {
8260 GpuTensor::Quant { qtype, .. } =>
8261 matches!(*qtype, QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q4_0)
8262 || (*qtype == QT_NVFP4 && w.in_features() % 64 == 0),
8263 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
8264 }
8265 }
8266
8267 pub fn qmatvec_gemm(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
8274 m: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8275 use crate::model::GpuTensor;
8276 let in_f = w.in_features();
8277 let out_f = w.out_features();
8278 let (bytes, qtype, row_bytes, scale, rp) = match w {
8279 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
8280 _ => unreachable!("gemm_supports guaranteed Quant"),
8281 };
8282 if cfg!(memra_hopper_mma) && qtype == QT_Q8_0 && out_f % 64 == 0 && wgmma_gemm_enabled() {
8288 if let GpuTensor::Quant { rp4: Some(m4), .. } = w {
8289 let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, aq, ad, m, in_f, out_f)?;
8290 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8291 return Ok(y);
8292 }
8293 }
8294 let name = match qtype {
8295 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8296 QT_Q4_0 => if rp { "qmatvec_gemm_q4_0_rp" } else { "qmatvec_gemm_q4_0" },
8297 QT_Q5_K => "qmatvec_gemm_q5_K",
8298 QT_Q6_K => "qmatvec_gemm_q6_K",
8299 QT_NVFP4 => if rp { "qmatvec_gemm_nvfp4_rp" } else { "qmatvec_gemm_nvfp4" },
8300 _ => unreachable!(),
8301 };
8302 let f = self.func(name);
8303 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);
8308 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8310 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8311 let warps: u32 = if is_k1 { k1_tile.2 } else {
8312 match qtype { QT_NVFP4 => 8, _ => 4 }
8313 };
8314 let cfg = LaunchConfig {
8315 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8316 block_dim: (32, warps, 1),
8317 shared_mem_bytes: 0,
8318 };
8319 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8320 let __s_b = self.gpu.stream();
8321 let mut b = __s_b.launch_builder(&f);
8322 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8323 unsafe { b.launch(cfg)?; }
8324 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8325 Ok(y)
8326 }
8327
8328 pub fn qmatvec_gemm_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
8333 out_f: usize, qtype: i32, row_bytes: usize)
8334 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8335 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
8336 let name = match qtype {
8337 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8338 QT_Q4_0 => "qmatvec_gemm_q4_0",
8339 QT_Q5_K => "qmatvec_gemm_q5_K",
8340 QT_Q6_K => "qmatvec_gemm_q6_K", QT_NVFP4 => "qmatvec_gemm_nvfp4",
8341 QT_NVFP4_RP => "qmatvec_gemm_nvfp4_rp",
8342 _ => panic!("qmatvec_gemm_raw: qtype {qtype} has no GEMM kernel"),
8343 };
8344 let f = self.func(name);
8345 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);
8349 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8351 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8352 let warps: u32 = if is_k1 { k1_tile.2 } else {
8353 match qtype { QT_NVFP4 | QT_NVFP4_RP => 8, _ => 4 }
8354 };
8355 let cfg = LaunchConfig {
8356 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8357 block_dim: (32, warps, 1), shared_mem_bytes: 0,
8358 };
8359 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8360 let __s_b = self.gpu.stream();
8361 let mut b = __s_b.launch_builder(&f);
8362 b.arg(bytes).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8363 unsafe { b.launch(cfg)?; }
8364 Ok(y)
8365 }
8366
8367 pub fn qmatvec_gemm_q8_0_wgmma_raw(&self, rp4: &CudaSlice<u8>, aq: &CudaSlice<i8>,
8374 ad: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize)
8375 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8376 assert!(out_f % 64 == 0 && in_f % 32 == 0, "wgmma GEMM needs out_f%64==0, in_f%32==0");
8377 let f = self.func("qmatvec_gemm_q8_0_wgmma");
8378 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
8380 grid_dim: ((out_f / 64) as u32, (m as u32).div_ceil(64), 1),
8381 block_dim: (128, 1, 1), shared_mem_bytes: 0,
8382 };
8383 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
8384 let __s_b = self.gpu.stream();
8385 let mut b = __s_b.launch_builder(&f);
8386 b.arg(rp4).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi);
8387 unsafe { b.launch(cfg)?; }
8388 Ok(y)
8389 }
8390
8391 pub fn scale_inplace(&self, y: &mut CudaSlice<f32>, s: f32, n: usize)
8393 -> Result<(), Box<dyn std::error::Error>> {
8394 let f = self.func("scale_f32");
8395 let cfg = LaunchConfig::for_num_elems(n as u32);
8396 let (sf, ni) = (s, n as i32);
8397 let __s_b = self.gpu.stream();
8398 let mut b = __s_b.launch_builder(&f);
8399 b.arg(y).arg(&sf).arg(&ni);
8400 unsafe { b.launch(cfg)?; }
8401 Ok(())
8402 }
8403
8404 pub fn bf16_to_f32(&self, data: &cudarc::driver::CudaView<'_, u8>, n: usize)
8409 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8410 let mut out = self.alloc_uninit::<f32>(n)?;
8411 let f = self.func("bf16_to_f32");
8412 let cfg = LaunchConfig::for_num_elems(n as u32);
8413 let ni = n as i32;
8414 let __s_b = self.gpu.stream();
8415 let mut b = __s_b.launch_builder(&f);
8416 b.arg(data).arg(&mut out).arg(&ni);
8417 unsafe { b.launch(cfg)?; }
8418 Ok(out)
8419 }
8420
8421 fn linear_bf16_chunked(&self, x: &CudaSlice<f32>, data: &CudaSlice<u8>, m: usize,
8428 in_f: usize, out_f: usize, exact: bool)
8429 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8430 const CHUNK_BYTES: usize = 256 << 20;
8431 let chunk_rows = (CHUNK_BYTES / (in_f * 4)).max(1).min(out_f);
8432 if chunk_rows >= out_f {
8433 let wf32 = self.bf16_to_f32(&data.slice(0..in_f * out_f * 2), in_f * out_f)?;
8434 return if exact { self.linear_decode_exact(x, &wf32, m, in_f, out_f) }
8435 else { self.linear(x, &wf32, m, in_f, out_f) };
8436 }
8437 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8438 let mut r0 = 0usize;
8439 while r0 < out_f {
8440 let rows = chunk_rows.min(out_f - r0);
8441 let wslice = data.slice(r0 * in_f * 2..(r0 + rows) * in_f * 2);
8442 let wf32 = self.bf16_to_f32(&wslice, in_f * rows)?;
8443 let yc = if exact { self.linear_decode_exact(x, &wf32, m, in_f, rows)? }
8444 else { self.linear(x, &wf32, m, in_f, rows)? };
8445 for mi in 0..m {
8447 let src = yc.slice(mi * rows..(mi + 1) * rows);
8448 let mut dst = y.slice_mut(mi * out_f + r0..mi * out_f + r0 + rows);
8449 self.gpu.stream().memcpy_dtod(&src, &mut dst)?;
8450 }
8451 r0 += rows;
8452 }
8453 Ok(y)
8454 }
8455
8456 pub fn linear_decode_exact(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize,
8463 in_f: usize, out_f: usize)
8464 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8465 if m_tokens == 1 { return self.linear(x, w, 1, in_f, out_f); }
8466 let xv = self.view(x, m_tokens * in_f);
8467 let mut y = self.alloc_uninit::<f32>(m_tokens * out_f)?;
8468 for t in 0..m_tokens {
8469 let row = xv.slice(t * in_f..(t + 1) * in_f);
8470 let mut xr = self.alloc_uninit::<f32>(in_f)?;
8471 self.copy_view_into(&mut xr, 0, &row, in_f)?;
8472 let yr = self.linear(&xr, w, 1, in_f, out_f)?;
8473 self.copy_into(&mut y, t * out_f, &yr, out_f)?;
8474 }
8475 Ok(y)
8476 }
8477
8478 pub fn linear(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize, in_f: usize, out_f: usize)
8479 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8480 use cudarc::cublaslt::{Matmul, MatmulConfig};
8481 let mut c = self.alloc_uninit::<f32>(m_tokens * out_f)?; let cfg = MatmulConfig {
8483 transa: true, transb: false, transc: false,
8484 m: out_f as u64, n: m_tokens as u64, k: in_f as u64,
8485 alpha: 1.0, lda: in_f as i64, ldb: in_f as i64, beta: 0.0, ldc: out_f as i64,
8486 stride_a: None, stride_b: None, stride_c: None, stride_bias: None, batch_size: None,
8487 };
8488 unsafe { self.gpu.blas.matmul(cfg, w, x, &mut c, None, None)?; }
8489 Ok(c)
8490 }
8491
8492 pub fn sdpa_naive(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8494 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8495 t: usize, t_kv: usize, scale: f32, causal: bool)
8496 -> Result<(), Box<dyn std::error::Error>> {
8497 let f = self.func("sdpa_naive_f32");
8498 let cfg = LaunchConfig {
8499 grid_dim: (n_head as u32, t as u32, 1),
8500 block_dim: (128, 1, 1),
8501 shared_mem_bytes: (t_kv * 4) as u32,
8502 };
8503 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);
8504 let __s_b = self.gpu.stream();
8505 let mut b = __s_b.launch_builder(&f);
8506 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8507 unsafe { b.launch(cfg)?; }
8508 Ok(())
8509 }
8510
8511 #[allow(clippy::too_many_arguments)]
8513 pub fn sdpa_naive_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8514 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8515 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8516 -> Result<(), Box<dyn std::error::Error>> {
8517 let f = self.func("sdpa_naive_w_f32");
8518 let cfg = LaunchConfig {
8519 grid_dim: (n_head as u32, t as u32, 1),
8520 block_dim: (128, 1, 1),
8521 shared_mem_bytes: (t_kv * 4) as u32,
8522 };
8523 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
8524 t as i32, t_kv as i32, causal as i32, window as i32);
8525 let __s_b = self.gpu.stream();
8526 let mut b = __s_b.launch_builder(&f);
8527 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8528 .arg(&scale).arg(&cz).arg(&wi);
8529 unsafe { b.launch(cfg)?; }
8530 Ok(())
8531 }
8532
8533 pub fn sdpa_naive_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<f32>,
8535 v: &cudarc::driver::CudaView<f32>, o: &mut CudaSlice<f32>,
8536 head_dim: usize, n_head: usize, n_head_kv: usize, t: usize, t_kv: usize,
8537 scale: f32, causal: bool) -> Result<(), Box<dyn std::error::Error>> {
8538 let f = self.func("sdpa_naive_f32");
8539 let cfg = LaunchConfig {
8540 grid_dim: (n_head as u32, t as u32, 1), block_dim: (128, 1, 1),
8541 shared_mem_bytes: (t_kv * 4) as u32,
8542 };
8543 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);
8544 let __s_b = self.gpu.stream();
8545 let mut b = __s_b.launch_builder(&f);
8546 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8547 unsafe { b.launch(cfg)?; }
8548 Ok(())
8549 }
8550
8551 #[allow(clippy::too_many_arguments)]
8559 pub fn fa_dequant_kv_view_f32(&self, k: &cudarc::driver::CudaView<u8>,
8560 v: &cudarc::driver::CudaView<u8>,
8561 kf: &mut CudaSlice<f32>, vf: &mut CudaSlice<f32>,
8562 kv_dim_k: usize, kv_dim_v: usize, t_kv: usize,
8563 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
8564 -> Result<(), Box<dyn std::error::Error>> {
8565 let f = if g { self.func_g("fa_dequant_kv_ws_f32") } else { self.func("fa_dequant_kv_ws_f32") };
8566 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
8567 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8568 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1),
8569 shared_mem_bytes: 0 };
8570 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
8571 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
8572 let __s_b = self.gpu.stream();
8573 let mut b = __s_b.launch_builder(&f);
8574 b.arg(k).arg(v).arg(&mut *kf).arg(&mut *vf).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
8575 unsafe { b.launch(cfg)?; }
8576 Ok(())
8577 }
8578
8579 #[allow(clippy::too_many_arguments)]
8580 pub fn sdpa_naive_quantized_view(
8581 &self,
8582 q: &CudaSlice<f32>,
8583 k: &cudarc::driver::CudaView<u8>,
8584 v: &cudarc::driver::CudaView<u8>,
8585 o: &mut CudaSlice<f32>,
8586 head_dim: usize,
8587 n_head: usize,
8588 n_head_kv: usize,
8589 t: usize,
8590 t_kv: usize,
8591 scale: f32,
8592 causal: bool,
8593 k_tok_bytes: usize,
8594 v_tok_bytes: usize,
8595 ) -> Result<(), Box<dyn std::error::Error>> {
8596 let kv_dim = n_head_kv * head_dim;
8597 let mut kf = self.uninit(t_kv * kv_dim)?;
8598 let mut vf = self.uninit(t_kv * kv_dim)?;
8599 let f = self.func("fa_dequant_kv_ws_f32");
8600 let total = (2 * t_kv * kv_dim) as u64;
8601 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8602 let cfg = LaunchConfig {
8603 grid_dim: (nblk.max(1), 1, 1),
8604 block_dim: (256, 1, 1),
8605 shared_mem_bytes: 0,
8606 };
8607 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8608 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8609 let __s_b = self.gpu.stream();
8610 let mut b = __s_b.launch_builder(&f);
8611 b.arg(k)
8612 .arg(v)
8613 .arg(&mut kf)
8614 .arg(&mut vf)
8615 .arg(&kv_dim_i)
8616 .arg(&kv_dim_i)
8617 .arg(&t_kv_i)
8618 .arg(&k_tok_bytes_i)
8619 .arg(&v_tok_bytes_i);
8620 unsafe { b.launch(cfg)? };
8621 self.sdpa_naive(
8622 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8623 )
8624 }
8625
8626 #[allow(clippy::too_many_arguments)]
8638 pub fn sdpa_naive_w_quantized_view(
8639 &self,
8640 q: &CudaSlice<f32>,
8641 k: &cudarc::driver::CudaView<u8>,
8642 v: &cudarc::driver::CudaView<u8>,
8643 o: &mut CudaSlice<f32>,
8644 head_dim: usize,
8645 n_head: usize,
8646 n_head_kv: usize,
8647 t: usize,
8648 t_kv: usize,
8649 scale: f32,
8650 causal: bool,
8651 window: usize,
8652 k_tok_bytes: usize,
8653 v_tok_bytes: usize,
8654 ) -> Result<(), Box<dyn std::error::Error>> {
8655 let kv_dim = n_head_kv * head_dim;
8656 let mut kf = self.uninit(t_kv * kv_dim)?;
8657 let mut vf = self.uninit(t_kv * kv_dim)?;
8658 let f = self.func("fa_dequant_kv_ws_f32");
8659 let total = (2 * t_kv * kv_dim) as u64;
8660 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8661 let cfg = LaunchConfig {
8662 grid_dim: (nblk.max(1), 1, 1),
8663 block_dim: (256, 1, 1),
8664 shared_mem_bytes: 0,
8665 };
8666 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8667 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8668 let __s_b = self.gpu.stream();
8669 let mut b = __s_b.launch_builder(&f);
8670 b.arg(k)
8671 .arg(v)
8672 .arg(&mut kf)
8673 .arg(&mut vf)
8674 .arg(&kv_dim_i)
8675 .arg(&kv_dim_i)
8676 .arg(&t_kv_i)
8677 .arg(&k_tok_bytes_i)
8678 .arg(&v_tok_bytes_i);
8679 unsafe { b.launch(cfg)? };
8680 self.sdpa_naive_w(
8681 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
8682 )
8683 }
8684
8685 pub fn fa_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8689 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8690 t: usize, t_kv: usize, scale: f32, causal: bool)
8691 -> Result<(), Box<dyn std::error::Error>> {
8692 if portable_mma_gated() {
8693 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
8694 t, t_kv, scale, causal);
8695 }
8696 let fa3_on = head_dim == 256 && causal && t == t_kv
8704 && match std::env::var("MEMRA_FA3").as_deref() {
8705 Ok("0") => false,
8706 Ok("1") => true,
8707 _ => cfg!(memra_hopper_mma),
8708 };
8709 if fa3_on {
8710 let n = t * n_head * head_dim;
8711 let nkv = t * n_head_kv * head_dim;
8712 let mut q16 = self.alloc_u8_uninit(n * 2)?;
8713 let mut k16 = self.alloc_u8_uninit(nkv * 2)?;
8714 let mut v16 = self.alloc_u8_uninit(nkv * 2)?;
8715 self.f32_to_bf16_into(q, &mut q16, n)?;
8716 self.f32_to_bf16_into(k, &mut k16, nkv)?;
8717 self.f32_to_bf16_into(v, &mut v16, nkv)?;
8718 let rc = {
8719 use cudarc::driver::{DevicePtr, DevicePtrMut};
8720 let stream = self.gpu.stream();
8721 let (qp, _g1) = q16.device_ptr(&stream);
8722 let (kp, _g2) = k16.device_ptr(&stream);
8723 let (vp, _g3) = v16.device_ptr(&stream);
8724 let (op, _g4) = o.device_ptr_mut(&stream);
8725 unsafe {
8726 memra_fa3_prefill(qp as *const core::ffi::c_void,
8727 kp as *const core::ffi::c_void,
8728 vp as *const core::ffi::c_void,
8729 op as *mut f32,
8730 t as i32, n_head as i32, n_head_kv as i32,
8731 head_dim as i32, scale,
8732 stream.cu_stream() as *mut core::ffi::c_void)
8733 }
8734 };
8735 if rc != 0 {
8736 return Err(format!("memra_fa3_prefill rc={rc}").into());
8737 }
8738 return Ok(());
8739 }
8740 static FA_P1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8745 let fa_p1 = *FA_P1.get_or_init(|| std::env::var("MEMRA_FA_P1").as_deref() == Ok("1"));
8746 if fa_p1 && head_dim == 256 && !std::env::var("MEMRA_FA_FLOOR").is_ok() {
8747 const BLOCK_Q: usize = 64; const BKX: usize = 32;
8748 let f = self.func("fa_prefill_bf16_p1");
8749 let shmem = (2 * (2 * BKX * head_dim + BLOCK_Q * BKX)
8750 + 4 * (BLOCK_Q * BKX + 2 * BLOCK_Q)) as u32;
8751 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8752 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8753 let cfg = LaunchConfig {
8754 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8755 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8756 };
8757 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32,
8758 n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8759 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8760 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8761 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8762 let __s_b = self.gpu.stream();
8763 let mut b = __s_b.launch_builder(&f);
8764 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti)
8765 .arg(&tkvi).arg(&scale).arg(&cz);
8766 unsafe { b.launch(cfg)?; }
8767 return Ok(());
8768 }
8769 const BK: usize = 32;
8775 let w2 = std::env::var("MEMRA_FA_PP_W2").as_deref() == Ok("1");
8778 let (block_q, warps, w2_sfx): (usize, u32, &str) =
8779 if w2 { (32, 2, "_w2") } else { (64, 4, "") };
8780 let hd_sfx = fa_hd_suffix(head_dim)?;
8784 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8785 let bf16kv = !floor && !w2
8790 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0");
8791 let (kb16, vb16) = if bf16kv {
8792 let n = t_kv * n_head_kv * head_dim;
8793 let mut kb = self.alloc_u8_uninit(n * 2)?;
8794 let mut vb = self.alloc_u8_uninit(n * 2)?;
8795 let fcv = self.func("f32_to_bf16_bulk");
8796 let ni = n as i64;
8797 let cfgc = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
8798 let __s_b = self.gpu.stream();
8799 let mut b = __s_b.launch_builder(&fcv);
8800 b.arg(k).arg(&mut kb).arg(&ni);
8801 unsafe { b.launch(cfgc)?; }
8802 let __s_b = self.gpu.stream();
8803 let mut b = __s_b.launch_builder(&fcv);
8804 b.arg(v).arg(&mut vb).arg(&ni);
8805 unsafe { b.launch(cfgc)?; }
8806 (Some(kb), Some(vb))
8807 } else {
8808 (None, None)
8809 };
8810 let f = self.func(&if bf16kv {
8811 format!("fa_prefill_bf16kv_pp{hd_sfx}")
8812 } else {
8813 format!("fa_prefill_f32{}{}{hd_sfx}",
8814 if floor { "" } else { "_pp" },
8815 if floor { "" } else { w2_sfx })
8816 });
8817 let kv_stages = if bf16kv { 2 } else { 1 };
8820 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
8821 + 4 * (block_q * BK + 2 * block_q)) as u32;
8822 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8823 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8824 let cfg = LaunchConfig {
8825 grid_dim: ((t as u32 + block_q as u32 - 1) / block_q as u32, n_head as u32, 1),
8826 block_dim: (32, warps, 1), shared_mem_bytes: shmem,
8827 };
8828 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);
8829 let __s_b = self.gpu.stream();
8830 let mut b = __s_b.launch_builder(&f);
8831 b.arg(q);
8832 match (&kb16, &vb16) {
8833 (Some(kb), Some(vb)) => { b.arg(kb).arg(vb); }
8834 _ => { b.arg(k).arg(v); }
8835 }
8836 b.arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8837 unsafe { b.launch(cfg)?; }
8838 Ok(())
8839 }
8840
8841 #[allow(clippy::too_many_arguments)]
8845 pub fn fa_prefill_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8846 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8847 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8848 -> Result<(), Box<dyn std::error::Error>> {
8849 if portable_mma_gated() {
8852 return self.sdpa_naive_w(q, k, v, o, head_dim, n_head, n_head_kv,
8853 t, t_kv, scale, causal, window);
8854 }
8855 static FAW_F32: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8859 let faw_f32 = *FAW_F32.get_or_init(|| {
8860 std::env::var("MEMRA_FAW_STAGE").as_deref() == Ok("f32")
8861 });
8862 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8863 self.fa_prefill_w_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8864 window, floor || faw_f32, floor)
8865 }
8866
8867 #[allow(clippy::too_many_arguments)]
8870 pub fn fa_prefill_w_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
8871 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8872 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
8873 window: usize, v_f16: bool)
8874 -> Result<(), Box<dyn std::error::Error>> {
8875 const BLOCK_Q: usize = 64; const BK: usize = 32;
8876 debug_assert_eq!(head_dim, 256);
8877 let hp = fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
8878 && (n_head / n_head_kv) % 2 == 0;
8879 debug_assert!(!v_f16 || hp, "f16 V emitted but the SWA hp arm is off");
8880 if hp {
8881 const BLOCK_QH: usize = 32;
8882 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
8885 let vh: &CudaSlice<u8> = if v_f16 { vb } else {
8886 let n = t_kv * n_head_kv * head_dim;
8887 if vguard.as_ref().map(|b| b.len() < n * 2).unwrap_or(true) {
8888 *vguard = Some(self.alloc_uninit::<u8>(n * 2)?);
8889 }
8890 self.bf16_to_f16_into(vb, n, vguard.as_mut().unwrap())?;
8891 vguard.as_ref().unwrap()
8892 };
8893 let f = self.func("fa_prefill_w_bf16_p1h2");
8894 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
8895 + 4 * (2 * BLOCK_QH)) as u32;
8896 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8897 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8898 let cfg = LaunchConfig {
8899 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
8900 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8901 };
8902 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8903 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8904 let __s_b = self.gpu.stream();
8905 let mut b = __s_b.launch_builder(&f);
8906 b.arg(qb).arg(kb).arg(vh).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8907 .arg(&scale).arg(&cz).arg(&wi);
8908 unsafe { b.launch(cfg)?; }
8909 return Ok(());
8910 }
8911 let f = self.func("fa_prefill_w_bf16_p1");
8912 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
8913 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
8914 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8915 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8916 let cfg = LaunchConfig {
8917 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8918 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8919 };
8920 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8921 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8922 let __s_b = self.gpu.stream();
8923 let mut b = __s_b.launch_builder(&f);
8924 b.arg(qb).arg(kb).arg(vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8925 .arg(&scale).arg(&cz).arg(&wi);
8926 unsafe { b.launch(cfg)?; }
8927 Ok(())
8928 }
8929
8930 #[allow(clippy::too_many_arguments)]
8932 pub fn fa_prefill_w_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8933 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8934 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
8935 window: usize, f32_stage: bool, floor: bool)
8936 -> Result<(), Box<dyn std::error::Error>> {
8937 const BLOCK_Q: usize = 64; const BK: usize = 32;
8938 debug_assert_eq!(head_dim, 256, "fa_prefill_w is stamped hd256 only");
8939 static P1_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8943 let p1 = !floor && !f32_stage
8944 && *P1_ON.get_or_init(|| {
8945 std::env::var("MEMRA_FAW_P1").map(|v| v != "0").unwrap_or(true)
8946 });
8947 let hp = p1 && fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
8948 && (n_head / n_head_kv) % 2 == 0;
8949 if hp {
8950 const BLOCK_QH: usize = 32;
8951 let f = self.func("fa_prefill_w_bf16_p1h2");
8952 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
8953 + 4 * (2 * BLOCK_QH)) as u32;
8954 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8955 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8956 let cfg = LaunchConfig {
8957 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
8958 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8959 };
8960 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8961 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8962 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8963 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8964 let vh = self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?;
8965 let __s_b = self.gpu.stream();
8966 let mut b = __s_b.launch_builder(&f);
8967 b.arg(&qb).arg(&kb).arg(&vh).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8968 .arg(&scale).arg(&cz).arg(&wi);
8969 unsafe { b.launch(cfg)?; }
8970 return Ok(());
8971 }
8972 if p1 {
8973 let f = self.func("fa_prefill_w_bf16_p1");
8974 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
8975 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
8976 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8977 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8978 let cfg = LaunchConfig {
8979 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8980 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8981 };
8982 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8983 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8984 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8985 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8986 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8987 let __s_b = self.gpu.stream();
8988 let mut b = __s_b.launch_builder(&f);
8989 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8990 .arg(&scale).arg(&cz).arg(&wi);
8991 unsafe { b.launch(cfg)?; }
8992 return Ok(());
8993 }
8994 static G4_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8997 let g4 = !floor && !f32_stage && n_head_kv == 1 && n_head % 4 == 0
8998 && *G4_ON.get_or_init(|| {
8999 std::env::var("MEMRA_FAW_G4").map(|v| v != "0").unwrap_or(true)
9000 });
9001 if g4 {
9002 const SP_M: usize = 16;
9003 static O2_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9006 let o2 = *O2_ON.get_or_init(|| {
9007 std::env::var("MEMRA_FAW_O2").map(|v| v != "0").unwrap_or(true)
9008 });
9009 let f = self.func(if o2 { "fa_prefill_w_bf16_g4o2" } else { "fa_prefill_w_bf16_g4" });
9010 let shmem = if o2 {
9011 (2 * (4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M)) as u32
9012 } else {
9013 (2 * (2 * BK * head_dim + 4 * SP_M * head_dim + 4 * SP_M * BK)
9014 + 4 * (4 * SP_M)) as u32
9015 };
9016 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9017 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9018 let cfg = LaunchConfig {
9019 grid_dim: ((t as u32).div_ceil(SP_M as u32), (n_head / 4) as u32, 1),
9020 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9021 };
9022 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9023 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9024 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9025 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9026 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9027 let __s_b = self.gpu.stream();
9028 let mut b = __s_b.launch_builder(&f);
9029 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9030 .arg(&scale).arg(&cz).arg(&wi);
9031 unsafe { b.launch(cfg)?; }
9032 return Ok(());
9033 }
9034 let f = self.func(if floor { "fa_prefill_w_f32" }
9035 else if f32_stage { "fa_prefill_w_f32_pp" }
9036 else { "fa_prefill_w_bf16_pp" });
9037 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9038 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9039 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9040 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9041 let cfg = LaunchConfig {
9042 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9043 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9044 };
9045 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9046 t as i32, t_kv as i32, causal as i32, window as i32);
9047 if f32_stage {
9048 let __s_b = self.gpu.stream();
9049 let mut b = __s_b.launch_builder(&f);
9050 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9051 .arg(&scale).arg(&cz).arg(&wi);
9052 unsafe { b.launch(cfg)?; }
9053 } else {
9054 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9055 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9056 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9057 let __s_b = self.gpu.stream();
9058 let mut b = __s_b.launch_builder(&f);
9059 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9060 .arg(&scale).arg(&cz).arg(&wi);
9061 unsafe { b.launch(cfg)?; }
9062 }
9063 Ok(())
9064 }
9065
9066 #[allow(clippy::too_many_arguments)]
9070 pub fn fa_prefill_hd512(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9071 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9072 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool)
9073 -> Result<(), Box<dyn std::error::Error>> {
9074 if portable_mma_gated() {
9076 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
9077 t, t_kv, scale, causal);
9078 }
9079 static F32_STAGE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9085 let f32_stage = *F32_STAGE.get_or_init(|| {
9086 std::env::var("MEMRA_FA512_STAGE").as_deref() == Ok("f32")
9087 });
9088 static SP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9092 let sp = !f32_stage
9093 && *SP_ON.get_or_init(|| {
9094 std::env::var("MEMRA_FA512_SP").map(|v| v != "0").unwrap_or(true)
9095 });
9096 self.fa_prefill_hd512_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale,
9097 causal, f32_stage, sp, sp && fa_f16pv_on())
9098 }
9099
9100 #[allow(clippy::too_many_arguments)]
9102 pub fn fa_prefill_hd512_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
9103 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9104 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9105 v_f16: bool)
9106 -> Result<(), Box<dyn std::error::Error>> {
9107 debug_assert_eq!(head_dim, 512);
9108 const SP_M: usize = 16; const BKS: usize = 32;
9109 let f16pv = fa_f16pv_on();
9113 let nw = if f16pv { fa512_wide_warps() } else { 2 };
9114 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
9115 debug_assert!(!v_f16 || f16pv, "f16 V emitted without the door on");
9116 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
9117 let vref: &CudaSlice<u8> = if f16pv && !v_f16 {
9118 let n = t_kv * n_head_kv * head_dim;
9120 let need = n * 2;
9121 if vguard.as_ref().map(|b| b.len() < need).unwrap_or(true) {
9122 *vguard = Some(self.alloc_uninit::<u8>(need)?);
9123 }
9124 let dst = vguard.as_mut().unwrap();
9125 self.bf16_to_f16_into(vb, n, dst)?;
9126 vguard.as_ref().unwrap()
9127 } else { vb };
9128 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
9129 else { match (f16pv, nw) {
9130 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9131 (true, _) => "fa_prefill_bf16_hd512_sp16",
9132 _ => "fa_prefill_bf16_hd512_sp",
9133 } });
9134 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9135 let shmem = if hp {
9137 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9138 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9139 } else {
9140 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9141 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9142 };
9143 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9144 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9145 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9146 let cfg = LaunchConfig {
9147 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9148 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9149 };
9150 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9151 t as i32, t_kv as i32, causal as i32);
9152 let __s_b = self.gpu.stream();
9153 let mut b = __s_b.launch_builder(&f);
9154 b.arg(qb).arg(kb).arg(vref).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9155 .arg(&scale).arg(&cz);
9156 unsafe { b.launch(cfg)?; }
9157 Ok(())
9158 }
9159
9160 #[allow(clippy::too_many_arguments)]
9163 pub fn fa_prefill_hd512_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9164 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9165 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9166 f32_stage: bool, sp: bool, f16pv: bool)
9167 -> Result<(), Box<dyn std::error::Error>> {
9168 debug_assert_eq!(head_dim, 512, "fa_prefill_hd512 is hd512 only");
9169 if sp && !f32_stage {
9170 const SP_M: usize = 16; const BKS: usize = 32;
9174 let nw = if f16pv { fa512_wide_warps() } else { 2 };
9175 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
9176 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
9177 else { match (f16pv, nw) {
9178 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9179 (true, _) => "fa_prefill_bf16_hd512_sp16",
9180 _ => "fa_prefill_bf16_hd512_sp",
9181 } });
9182 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9183 let shmem = if hp {
9184 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9185 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9186 } else {
9187 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9188 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9189 };
9190 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9191 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9192 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9193 let cfg = LaunchConfig {
9194 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9195 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9196 };
9197 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9198 t as i32, t_kv as i32, causal as i32);
9199 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9200 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9201 let vb = if f16pv { self.f32_to_f16(v, t_kv * n_head_kv * head_dim)? }
9202 else { self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)? };
9203 let __s_b = self.gpu.stream();
9204 let mut b = __s_b.launch_builder(&f);
9205 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9206 .arg(&scale).arg(&cz);
9207 unsafe { b.launch(cfg)?; }
9208 return Ok(());
9209 }
9210 const BLOCK_Q: usize = 32; const BK: usize = 32; const HALF: usize = 256;
9211 let f = self.func(if f32_stage { "fa_prefill_f32_hd512" } else { "fa_prefill_bf16_hd512" });
9212 let shmem = (2 * (BLOCK_Q * head_dim + BK * head_dim + BK * HALF + BLOCK_Q * BK)
9214 + 4 * BLOCK_Q) as u32;
9215 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9216 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9217 let cfg = LaunchConfig {
9218 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 2),
9219 block_dim: (32, 2, 1), shared_mem_bytes: shmem,
9220 };
9221 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9222 t as i32, t_kv as i32, causal as i32);
9223 if f32_stage {
9224 let __s_b = self.gpu.stream();
9225 let mut b = __s_b.launch_builder(&f);
9226 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9227 .arg(&scale).arg(&cz);
9228 unsafe { b.launch(cfg)?; }
9229 } else {
9230 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9231 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9232 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9233 let __s_b = self.gpu.stream();
9234 let mut b = __s_b.launch_builder(&f);
9235 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9236 .arg(&scale).arg(&cz);
9237 unsafe { b.launch(cfg)?; }
9238 }
9239 Ok(())
9240 }
9241
9242 #[allow(clippy::too_many_arguments)]
9246 pub fn rope_neox2_bf16e(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
9247 qb: &mut CudaSlice<u8>, kb: &mut CudaSlice<u8>,
9248 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
9249 nh_q: usize, nh_k: usize, n_tokens: usize, base: f32,
9250 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
9251 -> Result<(), Box<dyn std::error::Error>> {
9252 let f = self.func("rope_neox2_bf16e_f32");
9253 let rows = ((nh_q + nh_k) * n_tokens) as u32;
9254 let cfg = LaunchConfig { grid_dim: (rows, 1, 1),
9255 block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
9256 let theta_scale = base.powf(-2.0 / n_dims as f32);
9257 let (hd, nd, nhq, nhk, nt) = (head_dim as i32, n_dims as i32, nh_q as i32,
9258 nh_k as i32, n_tokens as i32);
9259 let __s_b = self.gpu.stream();
9260 let mut b = __s_b.launch_builder(&f);
9261 match ff {
9262 Some(t) => { b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9263 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9264 .arg(&theta_scale).arg(&freq_scale).arg(t);
9265 unsafe { b.launch(cfg)?; } }
9266 None => { let null: u64 = 0;
9267 b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9268 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9269 .arg(&theta_scale).arg(&freq_scale).arg(&null);
9270 unsafe { b.launch(cfg)?; } }
9271 }
9272 Ok(())
9273 }
9274
9275 pub fn f32_to_bf16(&self, x: &CudaSlice<f32>, n: usize)
9278 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9279 assert!(n % 4 == 0, "f32_to_bf16 requires n % 4 == 0, got {n}");
9280 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9281 let f = self.func("f32_to_bf16_flat");
9282 let n_i = n as i64;
9283 let cfg = LaunchConfig {
9284 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9285 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9286 };
9287 let __s_b = self.gpu.stream();
9288 let mut b = __s_b.launch_builder(&f);
9289 b.arg(x).arg(&mut y).arg(&n_i);
9290 unsafe { b.launch(cfg)?; }
9291 Ok(y)
9292 }
9293
9294 pub fn f32_to_f16(&self, x: &CudaSlice<f32>, n: usize)
9295 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9296 assert!(n % 4 == 0, "f32_to_f16 requires n % 4 == 0, got {n}");
9297 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9298 let f = self.func("f32_to_f16_flat");
9299 let n_i = n as i64;
9300 let cfg = LaunchConfig {
9301 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9302 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9303 };
9304 let __s_b = self.gpu.stream();
9305 let mut b = __s_b.launch_builder(&f);
9306 b.arg(x).arg(&mut y).arg(&n_i);
9307 unsafe { b.launch(cfg)?; }
9308 Ok(y)
9309 }
9310
9311 pub fn bf16_to_f16(&self, xb: &CudaSlice<u8>, n: usize)
9313 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9314 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9315 self.bf16_to_f16_into(xb, n, &mut y)?;
9316 Ok(y)
9317 }
9318
9319 pub fn bf16_to_f16_into(&self, xb: &CudaSlice<u8>, n: usize, y: &mut CudaSlice<u8>)
9321 -> Result<(), Box<dyn std::error::Error>> {
9322 assert!(n % 2 == 0, "bf16_to_f16 requires n % 2 == 0, got {n}");
9323 assert!(y.len() >= n * 2);
9324 let f = self.func("bf16_to_f16_flat");
9325 let n2 = (n / 2) as i64;
9326 let cfg = LaunchConfig {
9327 grid_dim: (((n / 2) as u32).div_ceil(256), 1, 1),
9328 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9329 };
9330 let __s_b = self.gpu.stream();
9331 let mut b = __s_b.launch_builder(&f);
9332 b.arg(xb).arg(y).arg(&n2);
9333 unsafe { b.launch(cfg)?; }
9334 Ok(())
9335 }
9336
9337 #[allow(clippy::too_many_arguments)]
9342 pub fn fa_prefill_vl8(&self, seqs: &[FaSeqVl], head_dim: usize, n_head: usize,
9343 n_head_kv: usize, scale: f32)
9344 -> Result<(), Box<dyn std::error::Error>> {
9345 const BK: usize = 32;
9346 let b = seqs.len();
9347 assert!(b >= 1 && b <= 8);
9348 let mut packed = [FaSeqVl::default(); 8];
9349 packed[..b].copy_from_slice(seqs);
9350 let v = FaVl8(packed);
9351 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9352 let ept = (n_head_kv * head_dim) as i32;
9353 {
9354 let f = self.func("fa_mirror_vl");
9355 let max_n = (max_t as i64) * ept as i64;
9356 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
9357 for which in 0..2i32 {
9358 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9359 let __s_lb = self.gpu.stream();
9360 let mut lb = __s_lb.launch_builder(&f);
9361 lb.arg(&v).arg(&ept).arg(&which);
9362 unsafe { lb.launch(cfg)?; }
9363 }
9364 }
9365 let hd_sfx = fa_hd_suffix(head_dim)?;
9366 let f = self.func(&format!("fa_prefill_bf16kv_vl{hd_sfx}"));
9367 let block_q = 64usize;
9368 let kv_stages = 2usize;
9369 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
9370 + 4 * (block_q * BK + 2 * block_q)) as u32;
9371 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9372 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9373 let cfg = LaunchConfig {
9374 grid_dim: (max_t.div_ceil(block_q as u32), n_head as u32, b as u32),
9375 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9376 };
9377 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9378 let __s_lb = self.gpu.stream();
9379 let mut lb = __s_lb.launch_builder(&f);
9380 lb.arg(&v).arg(&hd).arg(&nh).arg(&nhkv).arg(&scale);
9381 unsafe { lb.launch(cfg)?; }
9382 Ok(())
9383 }
9384
9385 #[allow(clippy::too_many_arguments)]
9389 pub fn attn_pre_vl8(&self, seqs: &[AttnPreVl], wq: &CudaSlice<f32>, wk: &CudaSlice<f32>,
9390 head_dim: usize, rope_dims: usize, n_head: usize, n_head_kv: usize,
9391 eps: f32, freq_base: f32, freq_scale: f32,
9392 kv_dim_k: usize, kv_dim_v: usize,
9393 k_tok_bytes: usize, v_tok_bytes: usize)
9394 -> Result<(), Box<dyn std::error::Error>> {
9395 let b = seqs.len();
9396 assert!(b >= 1 && b <= 8);
9397 let mut packed = [AttnPreVl::default(); 8];
9398 packed[..b].copy_from_slice(seqs);
9399 let v = AttnPreVl8(packed);
9400 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9401 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9402 {
9403 let f = self.func("q_gate_split_vl");
9404 let n = max_t * (n_head * head_dim) as u32;
9405 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9406 let __s_lb = self.gpu.stream();
9407 let mut lb = __s_lb.launch_builder(&f);
9408 lb.arg(&v).arg(&hd).arg(&nh);
9409 unsafe { lb.launch(cfg)?; }
9410 }
9411 {
9412 let f = self.func("attn_rms_vl");
9413 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 };
9414 let __s_lb = self.gpu.stream();
9415 let mut lb = __s_lb.launch_builder(&f);
9416 lb.arg(&v).arg(wq).arg(wk).arg(&hd).arg(&nh).arg(&nhkv).arg(&eps);
9417 unsafe { lb.launch(cfg)?; }
9418 }
9419 {
9420 let f = self.func("attn_rope_vl");
9421 let theta_scale = freq_base.powf(-2.0 / rope_dims as f32);
9422 let nd = rope_dims as i32;
9423 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 };
9424 let __s_lb = self.gpu.stream();
9425 let mut lb = __s_lb.launch_builder(&f);
9426 lb.arg(&v).arg(&hd).arg(&nd).arg(&nh).arg(&nhkv).arg(&theta_scale).arg(&freq_scale);
9427 unsafe { lb.launch(cfg)?; }
9428 }
9429 {
9430 let f = self.func("append_kv_vl");
9431 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
9432 let cfg = LaunchConfig { grid_dim: (nblk, max_t, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
9433 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9434 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9435 let __s_lb = self.gpu.stream();
9436 let mut lb = __s_lb.launch_builder(&f);
9437 lb.arg(&v).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
9438 unsafe { lb.launch(cfg)?; }
9439 }
9440 Ok(())
9441 }
9442
9443 pub fn fa_prefill_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9448 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9449 head_dim: usize, n_head: usize, n_head_kv: usize,
9450 t: usize, t_kv: usize, scale: f32, causal: bool,
9451 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9452 -> Result<(), Box<dyn std::error::Error>> {
9453 if portable_mma_gated() {
9454 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9455 t, t_kv, scale, causal,
9456 k_tok_bytes, v_tok_bytes);
9457 }
9458 const BLOCK_Q: usize = 64; const BK: usize = 32;
9459 let name = format!("fa_prefill_q{}", fa_hd_suffix(head_dim)?);
9462 let f = if g { self.func_g(&name) } else { self.func(&name) };
9463 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9464 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9465 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9466 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9467 let cfg = LaunchConfig {
9468 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9469 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9470 };
9471 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);
9472 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9473 let __s_b = self.gpu.stream();
9474 let mut b = __s_b.launch_builder(&f);
9475 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9476 .arg(&ktb).arg(&vtb);
9477 unsafe { b.launch(cfg)?; }
9478 Ok(())
9479 }
9480
9481 #[allow(clippy::too_many_arguments)]
9491 pub fn fa_prefill_view_ws(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9492 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9493 head_dim: usize, n_head: usize, n_head_kv: usize,
9494 t: usize, t_kv: usize, scale: f32, causal: bool,
9495 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9496 -> Result<(), Box<dyn std::error::Error>> {
9497 if portable_mma_gated() {
9498 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9499 t, t_kv, scale, causal,
9500 k_tok_bytes, v_tok_bytes);
9501 }
9502 const BLOCK_Q: usize = 64; const BK: usize = 32;
9503 let kv_dim_k = n_head_kv * head_dim;
9504 let kv_dim_v = n_head_kv * head_dim;
9505 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9507 let mut guard = self.prime_deqw_ws.lock().unwrap();
9509 let need_grow = match guard.as_ref() {
9510 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9511 None => true,
9512 };
9513 if need_grow {
9514 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9515 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9516 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9517 }
9518 let (kw, vw) = guard.as_mut().unwrap();
9519 {
9521 let f = if g { self.func_g("fa_dequant_kv_ws_bf16") } else { self.func("fa_dequant_kv_ws_bf16") };
9523 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9524 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9525 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9526 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9527 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9528 let __s_b = self.gpu.stream();
9529 let mut b = __s_b.launch_builder(&f);
9530 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9531 unsafe { b.launch(cfg)?; }
9532 }
9533 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9541 {
9542 let hd_sfx = fa_hd_suffix(head_dim)?;
9543 let f = self.func(&format!("fa_prefill_qw{}{hd_sfx}", if db { "_db" } else { "" }));
9544 let shmem = if db {
9545 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9547 } else {
9548 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9549 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9550 };
9551 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9552 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9553 let cfg = LaunchConfig {
9554 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9555 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9556 };
9557 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);
9558 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9559 let __s_b = self.gpu.stream();
9560 let mut b = __s_b.launch_builder(&f);
9561 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9562 .arg(&kdk).arg(&kdv);
9563 unsafe { b.launch(cfg)?; }
9564 }
9565 Ok(())
9566 }
9567
9568 #[allow(clippy::too_many_arguments)]
9584 pub fn fa_prefill_view_ws_w_hd128(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9585 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9586 head_dim: usize, n_head: usize, n_head_kv: usize,
9587 t: usize, t_kv: usize, scale: f32, causal: bool,
9588 window: usize, k_tok_bytes: usize, v_tok_bytes: usize)
9589 -> Result<(), Box<dyn std::error::Error>> {
9590 assert_eq!(head_dim, 128, "fa_prefill_view_ws_w_hd128: only the hd128 twin is stamped");
9591 if portable_mma_gated() {
9592 return self.sdpa_naive_w_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9593 t, t_kv, scale, causal, window,
9594 k_tok_bytes, v_tok_bytes);
9595 }
9596 const BLOCK_Q: usize = 64; const BK: usize = 32;
9597 let kv_dim_k = n_head_kv * head_dim;
9598 let kv_dim_v = n_head_kv * head_dim;
9599 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9601 let mut guard = self.prime_deqw_ws.lock().unwrap();
9602 let need_grow = match guard.as_ref() {
9603 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9604 None => true,
9605 };
9606 if need_grow {
9607 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9608 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9609 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9610 }
9611 let (kw, vw) = guard.as_mut().unwrap();
9612 {
9615 let f = self.func("fa_dequant_kv_ws_bf16");
9616 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9617 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9618 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9619 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9620 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9621 let __s_b = self.gpu.stream();
9622 let mut b = __s_b.launch_builder(&f);
9623 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9624 unsafe { b.launch(cfg)?; }
9625 }
9626 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9628 {
9629 let f = self.func(if db { "fa_prefill_qw_db_w_hd128" } else { "fa_prefill_qw_w_hd128" });
9630 let shmem = if db {
9631 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9632 } else {
9633 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9634 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9635 };
9636 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9637 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9638 let cfg = LaunchConfig {
9639 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9640 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9641 };
9642 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);
9643 let (kdk, kdv, wnd) = (kv_dim_k as i32, kv_dim_v as i32, window as i32);
9644 let __s_b = self.gpu.stream();
9645 let mut b = __s_b.launch_builder(&f);
9646 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9647 .arg(&kdk).arg(&kdv).arg(&wnd);
9648 unsafe { b.launch(cfg)?; }
9649 }
9650 Ok(())
9651 }
9652
9653 pub fn fa_decode(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9657 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9658 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9659 k_tok_bytes: usize, v_tok_bytes: usize)
9660 -> Result<(), Box<dyn std::error::Error>> {
9661 self.fa_decode_kvmod(q, k, v, o, head_dim, n_head, n_head_kv, t_kv, scale,
9662 k_tok_bytes, v_tok_bytes, false)
9663 }
9664
9665 #[allow(clippy::too_many_arguments)]
9669 #[allow(clippy::too_many_arguments)]
9673 #[allow(clippy::too_many_arguments)]
9674 fn fa_decode_scalar_unified(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9675 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9676 head_dim: usize, n_head: usize, n_head_kv: usize,
9677 t_kv_host: usize, t_kv_dev: Option<&CudaSlice<i32>>,
9678 scale: f32, n_splits: usize, split_keys: usize,
9679 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
9680 part_o: &mut CudaSlice<f32>, part_m: &mut CudaSlice<f32>,
9681 part_l: &mut CudaSlice<f32>,
9682 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
9683 -> Result<(), Box<dyn std::error::Error>> {
9684 let f = if g { self.func_g("fa_decode_f32") } else { self.fa_func("fa_decode_f32", head_dim) };
9685 let cfg = LaunchConfig { grid_dim: (n_head as u32, n_splits as u32, 1),
9686 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: (4 * (head_dim + 32)) as u32 };
9687 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
9688 let (ktb, vtb, tkvi, ski) = (k_tok_bytes as i64, v_tok_bytes as i64, t_kv_host as i32,
9689 split_keys as i32);
9690 let __s_b = self.gpu.stream();
9691 let mut b = __s_b.launch_builder(&f);
9692 match t_kv_dev {
9693 Some(d) => { b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9694 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(d).arg(&scale).arg(&nsp)
9695 .arg(&ski).arg(&ktb).arg(&vtb);
9696 unsafe { b.launch(cfg)?; } }
9697 None => { let null: u64 = 0;
9698 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9699 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&null).arg(&scale).arg(&nsp)
9700 .arg(&ski).arg(&ktb).arg(&vtb);
9701 unsafe { b.launch(cfg)?; } }
9702 }
9703 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1),
9704 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
9705 if let Some((oq, od)) = q8_out {
9706 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
9708 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
9709 let __s_b2 = self.gpu.stream();
9710 let mut b2 = __s_b2.launch_builder(&fc);
9711 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
9712 unsafe { b2.launch(cfg2)?; }
9713 return Ok(());
9714 }
9715 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
9716 let __s_b2 = self.gpu.stream();
9717 let mut b2 = __s_b2.launch_builder(&fc);
9718 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
9719 unsafe { b2.launch(cfg2)?; }
9720 Ok(())
9721 }
9722
9723 pub fn fa_decode_kvmod(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9724 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9725 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9726 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9727 -> Result<(), Box<dyn std::error::Error>> {
9728 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
9749 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
9753 let sp = fa_split_keys(t_kv, n_head_kv);
9754 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
9755 let o_len = n_head * n_splits * head_dim;
9756 let ml_len = n_head * n_splits;
9757 let mut part_guard = self.fa_part_pool.lock().unwrap();
9758 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
9759 let old = part_guard.take();
9770 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
9771 if let Some(old) = old {
9772 self.fa_part_retired.lock().unwrap().push(old);
9773 }
9774 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
9775 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
9776 }
9777 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
9778 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
9779 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
9780 }
9781 let pg = part_guard.as_mut().unwrap();
9782 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
9783 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
9784 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
9785 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
9786 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
9787 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);
9788 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9789 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
9793 let fa512_min = fa512_min_tkv();
9798 let deep = fa_vec && head_dim == 256 && fa_v4_at(t_kv) && !g
9801 && fa_deep_at(t_kv) && !matches!(fa_v4_mode(), "noB3" | "stage");
9802 let (f, cfg) = if fa_vec && head_dim == 512 && t_kv >= fa512_min {
9803 let gqa = (n_head / n_head_kv).max(1) as u32;
9806 let fv = self.fa_func("fa_decode_vec_q_dpl16", head_dim);
9807 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9808 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
9809 } else if fa_vec && head_dim <= 256 {
9810 let gqa = (n_head / n_head_kv).max(1) as u32;
9811 static SMEM_TKV: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9822 let smem_tkv = *SMEM_TKV.get_or_init(|| {
9823 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
9824 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
9825 });
9826 if fa_v4_at(t_kv) && head_dim == 256 {
9827 let v4name = match fa_v4_mode() {
9831 "noB3" => "fa_decode_vec_q_v4_noB3", "stage" => "fa_decode_vec_q_v4_stage", _ if deep => "fa_decode_vec_q_v4_deep",
9834 _ => "fa_decode_vec_q_v4",
9835 };
9836 let fv = if g { self.func_g(v4name) } else { self.func(v4name) };
9837 let shmem = (if deep { 12160 } else { 11520 }
9840 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
9841 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9842 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9843 (fv,
9844 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9845 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9846 } else if fa_v3_active(head_dim) {
9847 let fv = if g { self.func_g("fa_decode_vec_q_v3") } else { self.func("fa_decode_vec_q_v3") };
9850 let shmem = (32 * head_dim * 2) as u32; (fv,
9852 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9853 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9854 } else if fa_v2_on() {
9855 let fv = if g { self.func_g("fa_decode_vec_q_v2") } else { self.func("fa_decode_vec_q_v2") };
9859 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
9861 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9862 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9863 } else if smem_tkv > 0 && t_kv >= smem_tkv && !g
9864 && !(head_dim == 512 && Self::gkv_on()) {
9865 let fv = if g { self.func_g("fa_decode_vec_q_smem") } else { self.func("fa_decode_vec_q_smem") };
9869 let shmem = (2 * 32 * head_dim * 2) as u32; use cudarc::driver::sys::CUfunction_attribute_enum as A;
9871 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9872 (fv,
9873 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9874 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9875 } else {
9876 let fv = if g { self.func_g("fa_decode_vec_q") } else { self.func("fa_decode_vec_q") };
9879 (fv,
9880 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9881 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
9882 }
9883 } else {
9884 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
9887 t_kv, None, scale, n_splits,
9888 if fa_vec { sp } else { 256 },
9889 k_tok_bytes, v_tok_bytes, g,
9890 part_o, part_m, part_l, None);
9891 };
9892 let __s_b = self.gpu.stream();
9893 let mut b = __s_b.launch_builder(&f);
9894 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9895 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&scale).arg(&nsp).arg(&ktb).arg(&vtb);
9896 unsafe { b.launch(cfg)?; }
9897 let (fc, cfg2) = (if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) },
9900 LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 });
9901 let __s_b2 = self.gpu.stream();
9902 let mut b2 = __s_b2.launch_builder(&fc);
9903 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
9904 unsafe { b2.launch(cfg2)?; }
9905 Ok(())
9906 }
9907
9908 #[allow(clippy::too_many_arguments)]
9919 pub fn fa_decode_batch_seqs_v4(&self, q: &CudaSlice<f32>,
9920 kv_ptrs: &cudarc::driver::CudaView<u64>,
9921 pos_seq: &CudaSlice<i32>, o: &mut CudaSlice<f32>,
9922 head_dim: usize, n_head: usize, n_head_kv: usize,
9923 b_n: usize, t_kv_max: usize, scale: f32,
9924 split_keys: usize, k_tok_bytes: usize, v_tok_bytes: usize)
9925 -> Result<(), Box<dyn std::error::Error>> {
9926 debug_assert!(head_dim == 256, "seqs twin is v4-stamped (hd256 only)");
9927 let n_splits_max = (t_kv_max + split_keys - 1) / split_keys;
9928 let o_len = b_n * n_head * n_splits_max * head_dim;
9929 let ml_len = b_n * n_head * n_splits_max;
9930 let mut part_guard = self.fa_part_pool.lock().unwrap();
9931 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
9932 let old = part_guard.take();
9943 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
9944 if let Some(old) = old {
9945 self.fa_part_retired.lock().unwrap().push(old);
9946 }
9947 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
9948 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
9949 }
9950 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
9951 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
9952 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
9953 }
9954 let pg = part_guard.as_mut().unwrap();
9955 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
9956 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
9957 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
9958 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
9959 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9960 let (nspm, spk) = (n_splits_max as i32, split_keys as i32);
9961 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9962 let gqa = (n_head / n_head_kv).max(1) as u32;
9963 let f = self.func("fa_decode_vec_q_seqs_v4");
9964 let shmem = (11520 + 32 * head_dim * 2) as u32;
9966 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9967 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9968 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, b_n as u32),
9969 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
9970 {
9971 let __s_b = self.gpu.stream();
9972 let mut b = __s_b.launch_builder(&f);
9973 b.arg(q).arg(kv_ptrs).arg(pos_seq).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9974 .arg(&hd).arg(&nh).arg(&nhkv).arg(&scale).arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
9975 unsafe { b.launch(cfg)?; }
9976 }
9977 let fc = self.func("fa_decode_combine_seqs");
9978 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, b_n as u32, 1),
9979 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
9980 let __s_b2 = self.gpu.stream();
9981 let mut b2 = __s_b2.launch_builder(&fc);
9982 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
9983 .arg(pos_seq).arg(&nspm).arg(&spk);
9984 unsafe { b2.launch(cfg2)?; }
9985 Ok(())
9986 }
9987
9988 #[allow(clippy::too_many_arguments)]
9995 pub fn append_kv_quantized_seqs(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
9996 kv_ptrs: &cudarc::driver::CudaView<u64>,
9997 pos_seq: &CudaSlice<i32>, b_n: usize,
9998 kv_dim_k: usize, kv_dim_v: usize,
9999 k_tok_bytes: usize, v_tok_bytes: usize)
10000 -> Result<(), Box<dyn std::error::Error>> {
10001 let f = self.func("append_quantize_kv_q8_0_q5_1_seqs");
10002 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
10003 let cfg = LaunchConfig { grid_dim: (nblk, b_n as u32, 1),
10004 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
10005 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
10006 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10007 let __s_b = self.gpu.stream();
10008 let mut b = __s_b.launch_builder(&f);
10009 b.arg(k_rows).arg(v_rows).arg(kv_ptrs).arg(pos_seq)
10010 .arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
10011 unsafe { b.launch(cfg)?; }
10012 Ok(())
10013 }
10014
10015 pub fn fa_rows_eligible(&self, base_len: usize, head_dim: usize) -> bool {
10021 std::env::var("MEMRA_NO_FA_VEC").is_err()
10022 && std::env::var("MEMRA_FA_ROWS_OFF").is_err()
10023 && base_len + 1 >= fa_vec_min_tkv()
10024 && head_dim <= 256 && head_dim % 32 == 0
10025 }
10026
10027 #[allow(clippy::too_many_arguments)]
10036 pub fn fa_decode_rows(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10037 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10038 head_dim: usize, n_head: usize, n_head_kv: usize,
10039 base_len: usize, t: usize, scale: f32,
10040 k_tok_bytes: usize, v_tok_bytes: usize,
10041 base_dev: Option<(&CudaSlice<i32>, i32)>,
10045 kv_shared: bool,
10048 g: bool,
10052 mut q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10055 -> Result<(), Box<dyn std::error::Error>> {
10056 debug_assert!(base_len + 1 >= fa_vec_min_tkv() && head_dim <= 512 && head_dim % 32 == 0);
10057 let t_kv_max = base_len + t; let mut sp = fa_split_keys(t_kv_max, n_head_kv); if head_dim == 512 {
10064 static SP512: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10065 let v = *SP512.get_or_init(|| std::env::var("MEMRA_FA_SP512").ok()
10068 .and_then(|x| x.parse().ok()).unwrap_or(0));
10069 sp = if v >= 8 { v } else { FA_SP512_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) };
10070 }
10071 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10072 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10073 let gqa = (n_head / n_head_kv).max(1) as u32;
10074 let mut groups: Vec<(usize, usize, usize)> = Vec::new(); if head_dim == 512 || fa_split_keys(base_len + 1, n_head_kv) == sp {
10085 groups.push((0, t, sp));
10086 } else {
10087 let mut r0 = 0usize;
10088 while r0 < t {
10089 let sp_g = fa_split_keys(base_len + r0 + 1, n_head_kv);
10090 let mut r1 = r0 + 1;
10091 while r1 < t && fa_split_keys(base_len + r1 + 1, n_head_kv) == sp_g { r1 += 1; }
10092 groups.push((r0, r1 - r0, sp_g));
10093 r0 = r1;
10094 }
10095 }
10096 static SMEM_TKV_R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10100 let smem_tkv = *SMEM_TKV_R.get_or_init(|| {
10101 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
10102 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
10103 });
10104 let v4 = fa_v4_at(base_len + t) && head_dim == 256;
10105 let v3 = fa_v3_active(head_dim);
10106 let smem_rows = head_dim <= 256 && !v3 && !fa_v2_on() && smem_tkv > 0 && t_kv_max >= smem_tkv;
10107 let _ = kv_shared;
10112 let i2 = head_dim == 512 && std::env::var("MEMRA_FA_I2").as_deref() != Ok("0");
10115 static TB512: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10129 let tb512 = head_dim == 512 && sp <= 32 && n_head / n_head_kv.max(1) <= 16
10131 && *TB512.get_or_init(|| std::env::var("MEMRA_FA_TB512").as_deref() != Ok("0"));
10132 let fname = if tb512 { "fa_decode_vec_q_rows_v4_512_tb" }
10133 else if i2 { "fa_decode_vec_q_rows_dpl16_i2" }
10134 else if head_dim == 512 { "fa_decode_vec_q_rows_dpl16" } else if v4 { "fa_decode_vec_q_rows_v4" }
10136 else if v3 { "fa_decode_vec_q_rows_v3" }
10137 else if fa_v2_on() { "fa_decode_vec_q_rows_v2" }
10138 else if smem_rows { "fa_decode_vec_q_rows_smem" }
10139 else { "fa_decode_vec_q_rows" };
10140 let f = if head_dim == 512 { self.fa_func(fname, head_dim) }
10141 else if g {
10142 self.func_g(if smem_rows { "fa_decode_vec_q_rows" } else { fname })
10150 }
10151 else { self.func(fname) };
10152 let shmem = if tb512 {
10153 let gk = Self::gkv_on();
10155 let sh = (8192 + 1024 + 32 * 512 + 32 * 64
10156 + 32 * head_dim * if gk { 1 } else { 2 }) as u32;
10157 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10158 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10159 sh
10160 } else if v4 || v3 || smem_rows || fa_v2_on() {
10161 let sh = (if v4 { 11520 + 32 * head_dim * if g { 1 } else { 2 } }
10163 else if v3 { 32 * head_dim * 2 } else { 2 * 32 * head_dim * 2 }) as u32;
10164 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10165 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10166 sh
10167 } else { 0 };
10168 for &(r0, t_g, sp_g) in &groups {
10172 let n_splits_g = (base_len + r0 + t_g).div_ceil(sp_g);
10173 let (nspm, spk) = (n_splits_g as i32, sp_g as i32);
10174 let base_i = (base_len + r0) as i32;
10175 let o_len = t_g * n_head * n_splits_g * head_dim;
10176 let ml_len = t_g * n_head * n_splits_g;
10177 let mut part_guard = self.fa_part_pool.lock().unwrap();
10178 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10179 let old = part_guard.take();
10190 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10191 if let Some(old) = old {
10192 self.fa_part_retired.lock().unwrap().push(old);
10193 }
10194 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10195 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10196 }
10197 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10198 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10199 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10200 }
10201 let pg = part_guard.as_mut().unwrap();
10202 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10203 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10204 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10205 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10206 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
10207 let qv = self.view(q, t * n_head * head_dim);
10208 let q_g = qv.slice(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10209 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_g as u32, t_g as u32),
10210 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10211 {
10212 let __s_b = self.gpu.stream();
10213 let mut b = __s_b.launch_builder(&f);
10214 if tb512 {
10215 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10217 let plus_g = plus + r0 as i32;
10218 let nr = t_g as i32;
10219 if Self::pdl_on() && Self::pdl_wb_on() {
10220 use cudarc::driver::{DevicePtr, DevicePtrMut};
10222 let s = &self.gpu.stream();
10223 let (pq, _b0) = q_g.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10224 let (pv, _b2) = v.device_ptr(s);
10225 let (po, _b3) = part_o.device_ptr_mut(s);
10226 let (pm, _b4) = part_m.device_ptr_mut(s);
10227 let (pl, _b5) = part_l.device_ptr_mut(s);
10228 let (pb, _b6) = bd.device_ptr(s);
10229 let mut ps = [
10230 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10231 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10232 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10233 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10234 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10235 &plus_g as *const _ as *mut _, &scale as *const _ as *mut _,
10236 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10237 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10238 &nr as *const _ as *mut _,
10239 ];
10240 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10241 "fa_decode_vec_q_rows_v4_512_tb",
10242 (n_head_kv as u32, n_splits_g as u32, 1), (32, gqa, 1),
10243 shmem, &mut ps)?; }
10244 } else {
10245 let cfg_tb = LaunchConfig {
10246 grid_dim: (n_head_kv as u32, n_splits_g as u32, 1),
10247 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10248 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10249 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10250 .arg(&ktb).arg(&vtb).arg(&nr);
10251 unsafe { b.launch(cfg_tb)?; }
10252 }
10253 } else if head_dim == 512 {
10254 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10255 let plus_g = plus + r0 as i32;
10256 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10257 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10258 .arg(&ktb).arg(&vtb);
10259 unsafe { b.launch(cfg)?; }
10260 } else {
10261 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10262 .arg(&hd).arg(&nh).arg(&nhkv).arg(&base_i).arg(&scale).arg(&nspm).arg(&spk)
10263 .arg(&ktb).arg(&vtb);
10264 unsafe { b.launch(cfg)?; }
10265 }
10266 }
10267 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t_g as u32, 1),
10268 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10269 let mut o_g = o.slice_mut(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10270 if head_dim == 512 {
10271 let (bd, plus) = base_dev.unwrap();
10274 let plus_g = plus + r0 as i32;
10275 if let Some((oq, od)) = q8_out.as_mut() {
10276 debug_assert!(t == 1, "rows q8 emit is a t=1 decode arm");
10278 if Self::pdl_on() && Self::pdl_wb_on() {
10279 use cudarc::driver::{DevicePtr, DevicePtrMut};
10281 let s = &self.gpu.stream();
10282 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10283 let (pl, _g2) = part_l.device_ptr(s);
10284 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10285 let (pb, _g5) = bd.device_ptr(s);
10286 let mut ps = [
10287 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10288 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10289 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10290 &nh as *const _ as *mut _, &pb as *const _ as *mut _,
10291 &plus_g as *const _ as *mut _, &nspm as *const _ as *mut _,
10292 &spk as *const _ as *mut _,
10293 ];
10294 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10295 "fa_decode_combine_rows_dc_q8_1",
10296 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10297 continue;
10298 }
10299 let fc = self.fa_func("fa_decode_combine_rows_dc_q8_1", head_dim);
10300 let __s_b2 = self.gpu.stream();
10301 let mut b2 = __s_b2.launch_builder(&fc);
10302 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut **oq).arg(&mut **od)
10303 .arg(&hd).arg(&nh).arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10304 unsafe { b2.launch(cfg2)?; }
10305 continue;
10306 }
10307 let fc = self.fa_func("fa_decode_combine_rows_dc", head_dim);
10308 let __s_b2 = self.gpu.stream();
10309 let mut b2 = __s_b2.launch_builder(&fc);
10310 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10311 .arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10312 unsafe { b2.launch(cfg2)?; }
10313 } else {
10314 assert!(q8_out.is_none(), "rows q8 emit requires the hd512 dc combine");
10317 let fc = self.func("fa_decode_combine_rows");
10318 let __s_b2 = self.gpu.stream();
10319 let mut b2 = __s_b2.launch_builder(&fc);
10320 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10321 .arg(&base_i).arg(&nspm).arg(&spk);
10322 unsafe { b2.launch(cfg2)?; }
10323 }
10324 }
10325 Ok(())
10326 }
10327
10328 #[allow(clippy::too_many_arguments)]
10332 pub fn fa_decode_rows_w(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10333 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10334 head_dim: usize, n_head: usize, n_head_kv: usize,
10335 base_dev: &CudaSlice<i32>, base_plus: i32, t: usize, scale: f32,
10336 window: usize, k_tok_bytes: usize, v_tok_bytes: usize,
10337 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10338 -> Result<(), Box<dyn std::error::Error>> {
10339 debug_assert!(head_dim == 256);
10344 let sp = {
10352 static SPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10353 let v = *SPW.get_or_init(|| std::env::var("MEMRA_FA_SPW").ok()
10354 .and_then(|x| x.parse().ok()).unwrap_or(0));
10355 if v >= 8 { v } else { FA_SPW_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) }
10356 };
10357 let n_splits_max = (window + sp - 1) / sp;
10358 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10359 let (nspm, spk, wini) = (n_splits_max as i32, sp as i32, window as i32);
10360 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10361 let gqa = (n_head / n_head_kv).max(1) as u32;
10362 let o_len = t * n_head * n_splits_max * head_dim;
10363 let ml_len = t * n_head * n_splits_max;
10364 let mut part_guard = self.fa_part_pool.lock().unwrap();
10365 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10366 let old = part_guard.take();
10377 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10378 if let Some(old) = old {
10379 self.fa_part_retired.lock().unwrap().push(old);
10380 }
10381 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10382 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10383 }
10384 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10385 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10386 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10387 }
10388 let pg = part_guard.as_mut().unwrap();
10389 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10390 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10391 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10392 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10393 static SMEM_TKV_W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10399 let smem_tkv = *SMEM_TKV_W.get_or_init(|| {
10400 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
10401 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
10402 });
10403 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10409 let wg = Self::wkv_on();
10414 let sp2 = gqa <= 4 && fa_v4_at(window)
10417 && std::env::var("MEMRA_FA_SPW2").as_deref() != Ok("0");
10418 if sp2 {
10419 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10420 if Self::pdl_on() && Self::pdl_wb_on() {
10421 use cudarc::driver::{DevicePtr, DevicePtrMut};
10423 let s = &self.gpu.stream();
10424 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10425 let (pv, _b2) = v.device_ptr(s);
10426 let (po, _b3) = part_o.device_ptr_mut(s);
10427 let (pm, _b4) = part_m.device_ptr_mut(s);
10428 let (pl, _b5) = part_l.device_ptr_mut(s);
10429 let (pb, _b6) = base_dev.device_ptr(s);
10430 let mut ps = [
10431 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10432 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10433 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10434 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10435 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10436 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10437 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10438 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10439 &wini as *const _ as *mut _,
10440 ];
10441 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w_sp",
10442 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa + 1, 1),
10443 sh, &mut ps)?; }
10444 } else {
10445 let f = if wg { self.func_g("fa_decode_vec_q_rows_v4_w_sp") }
10446 else { self.func("fa_decode_vec_q_rows_v4_w_sp") };
10447 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10448 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10449 block_dim: (32, gqa + 1, 1), shared_mem_bytes: sh };
10450 let __s_b = self.gpu.stream();
10451 let mut b = __s_b.launch_builder(&f);
10452 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10453 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10454 .arg(&ktb).arg(&vtb).arg(&wini);
10455 unsafe { b.launch(cfg)?; }
10456 }
10457 } else {
10458 if fa_v4_at(window) && Self::pdl_on() && Self::pdl_wb_on() {
10459 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10461 use cudarc::driver::{DevicePtr, DevicePtrMut};
10462 let s = &self.gpu.stream();
10463 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10464 let (pv, _b2) = v.device_ptr(s);
10465 let (po, _b3) = part_o.device_ptr_mut(s);
10466 let (pm, _b4) = part_m.device_ptr_mut(s);
10467 let (pl, _b5) = part_l.device_ptr_mut(s);
10468 let (pb, _b6) = base_dev.device_ptr(s);
10469 let mut ps = [
10470 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10471 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10472 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10473 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10474 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10475 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10476 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10477 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10478 &wini as *const _ as *mut _,
10479 ];
10480 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w",
10481 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa, 1),
10482 sh, &mut ps)?; }
10483 } else {
10484 let pick = |name: &str| if wg { self.func_g(name) } else { self.func(name) };
10485 let (f, sh) = if fa_v4_at(window) {
10486 let f = pick("fa_decode_vec_q_rows_v4_w");
10487 (f, (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32)
10488 } else if smem_tkv > 0 && window >= smem_tkv {
10489 (pick("fa_decode_vec_q_rows_smem_w"), (2 * 32 * head_dim * 2) as u32)
10492 } else {
10493 (pick("fa_decode_vec_q_rows_reg_w"), 0u32)
10494 };
10495 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10496 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10497 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10498 let __s_b = self.gpu.stream();
10499 let mut b = __s_b.launch_builder(&f);
10500 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10501 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10502 .arg(&ktb).arg(&vtb).arg(&wini);
10503 unsafe { b.launch(cfg)?; }
10504 }
10505 }
10506 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10507 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10508 if let Some((oq, od)) = q8_out {
10509 if Self::pdl_on() && Self::pdl_wb_on() {
10512 use cudarc::driver::{DevicePtr, DevicePtrMut};
10514 let s = &self.gpu.stream();
10515 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10516 let (pl, _g2) = part_l.device_ptr(s);
10517 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10518 let mut ps = [
10519 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10520 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10521 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10522 &nh as *const _ as *mut _, &nspm as *const _ as *mut _,
10523 &spk as *const _ as *mut _, &wini as *const _ as *mut _,
10524 ];
10525 unsafe { self.launch_pdl_flash(wg, "fa_decode_combine_rows_w_q8_1",
10526 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10527 return Ok(());
10528 }
10529 let fc = if wg { self.func_g("fa_decode_combine_rows_w_q8_1") }
10530 else { self.func("fa_decode_combine_rows_w_q8_1") };
10531 let __s_b2 = self.gpu.stream();
10532 let mut b2 = __s_b2.launch_builder(&fc);
10533 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh)
10534 .arg(&nspm).arg(&spk).arg(&wini);
10535 unsafe { b2.launch(cfg2)?; }
10536 return Ok(());
10537 }
10538 let fc = if wg { self.func_g("fa_decode_combine_rows_w") }
10539 else { self.func("fa_decode_combine_rows_w") };
10540 let __s_b2 = self.gpu.stream();
10541 let mut b2 = __s_b2.launch_builder(&fc);
10542 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10543 .arg(&nspm).arg(&spk).arg(&wini);
10544 unsafe { b2.launch(cfg2)?; }
10545 Ok(())
10546 }
10547
10548 #[allow(clippy::too_many_arguments)]
10554 pub fn fa_decode_rows_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10555 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10556 head_dim: usize, n_head: usize, n_head_kv: usize,
10557 base_dev: &CudaSlice<i32>, t_kv_upper: usize, t: usize, scale: f32,
10558 k_tok_bytes: usize, v_tok_bytes: usize, base_plus: i32, g: bool)
10559 -> Result<(), Box<dyn std::error::Error>> {
10560 let v4 = head_dim == 256 && fa_v4_at(t_kv_upper);
10561 assert!(v4 || fa_v3_active(head_dim), "stream fa rows requires the v3 or v4 lane");
10562 assert!(v4 || base_plus == 0, "v3_dc kernel takes no plus arg");
10563 if v4 {
10564 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10565 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10566 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10567 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10568 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10569 let gqa = (n_head / n_head_kv).max(1) as u32;
10570 let o_len = t * n_head * n_splits_max * head_dim;
10571 let ml_len = t * n_head * n_splits_max;
10572 let mut part_guard = self.fa_part_pool.lock().unwrap();
10573 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10574 let old = part_guard.take();
10585 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10586 if let Some(old) = old {
10587 self.fa_part_retired.lock().unwrap().push(old);
10588 }
10589 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10590 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10591 }
10592 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10593 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10594 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10595 }
10596 let pg = part_guard.as_mut().unwrap();
10597 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10598 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10599 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10600 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10601 let f = if g { self.func_g("fa_decode_vec_q_rows_v4_dc") }
10602 else { self.func("fa_decode_vec_q_rows_v4_dc") };
10603 let sh = (11520 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10604 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10605 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10606 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10607 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10608 let __s_b = self.gpu.stream();
10609 let mut b = __s_b.launch_builder(&f);
10610 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10611 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale)
10612 .arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
10613 unsafe { b.launch(cfg)?; }
10614 let fc = self.func("fa_decode_combine_rows_dc");
10615 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10616 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10617 let __s_b2 = self.gpu.stream();
10618 let mut b2 = __s_b2.launch_builder(&fc);
10619 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10620 .arg(base_dev).arg(&base_plus).arg(&nspm).arg(&spk);
10621 unsafe { b2.launch(cfg2)?; }
10622 return Ok(());
10623 }
10624 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10625 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10626 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10627 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10628 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10629 let gqa = (n_head / n_head_kv).max(1) as u32;
10630 let o_len = t * n_head * n_splits_max * head_dim;
10631 let ml_len = t * n_head * n_splits_max;
10632 let mut part_guard = self.fa_part_pool.lock().unwrap();
10633 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10634 let old = part_guard.take();
10645 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10646 if let Some(old) = old {
10647 self.fa_part_retired.lock().unwrap().push(old);
10648 }
10649 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10650 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10651 }
10652 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10653 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10654 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10655 }
10656 let pg = part_guard.as_mut().unwrap();
10657 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10658 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10659 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10660 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10661 let f = self.func("fa_decode_vec_q_rows_v3_dc");
10662 let sh = (32 * head_dim * 2) as u32;
10663 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10664 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10665 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10666 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10667 let __s_b = self.gpu.stream();
10668 let mut b = __s_b.launch_builder(&f);
10669 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10670 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&scale).arg(&nspm).arg(&spk)
10671 .arg(&ktb).arg(&vtb);
10672 unsafe { b.launch(cfg)?; }
10673 let fc = self.func("fa_decode_combine_rows_dc");
10674 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10675 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10676 let plus0 = 0i32;
10677 let __s_b2 = self.gpu.stream();
10678 let mut b2 = __s_b2.launch_builder(&fc);
10679 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10680 .arg(base_dev).arg(&plus0).arg(&nspm).arg(&spk);
10681 unsafe { b2.launch(cfg2)?; }
10682 Ok(())
10683 }
10684
10685 pub fn fa_decode_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10696 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10697 head_dim: usize, n_head: usize, n_head_kv: usize,
10698 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10699 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
10700 -> Result<(), Box<dyn std::error::Error>> {
10701 self.fa_decode_dc_q8(q, k, v, o, head_dim, n_head, n_head_kv, t_kv_dev, bucket_max,
10702 scale, k_tok_bytes, v_tok_bytes, g, None)
10703 }
10704
10705 #[allow(clippy::too_many_arguments)]
10708 pub fn fa_decode_dc_q8(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10709 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10710 head_dim: usize, n_head: usize, n_head_kv: usize,
10711 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10712 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
10713 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10714 -> Result<(), Box<dyn std::error::Error>> {
10715 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
10723 if g && head_dim == 256 && !fa_v4_at(bucket_max) { fa_vec = false; } let sp = fa_split_keys(bucket_max, n_head_kv);
10725 let n_splits = if fa_vec { ((bucket_max + sp - 1) / sp).max(1) } else { ((bucket_max + 255) / 256).max(1) };
10726 let o_len = n_head * n_splits * head_dim;
10727 let ml_len = n_head * n_splits;
10728 let mut part_guard = self.fa_part_pool.lock().unwrap();
10729 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10730 let old = part_guard.take();
10741 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10742 if let Some(old) = old {
10743 self.fa_part_retired.lock().unwrap().push(old);
10744 }
10745 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10746 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10747 }
10748 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10749 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10750 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10751 }
10752 let pg = part_guard.as_mut().unwrap();
10753 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10754 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10755 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10756 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10757 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
10758 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10759 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
10760 let deep = fa_vec && head_dim == 256 && fa_v4_at(bucket_max) && !g
10763 && fa_deep_at(bucket_max) && !matches!(fa_v4_mode(), "noB3" | "stage");
10764 let (f, cfg) = if fa_vec && head_dim == 512 && bucket_max >= {
10765 static FA512_MIN_DC: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10766 *FA512_MIN_DC.get_or_init(|| std::env::var("MEMRA_FA512_MIN").ok()
10767 .and_then(|v| v.parse().ok()).unwrap_or(512))
10768 } {
10769 let gqa = (n_head / n_head_kv).max(1) as u32;
10771 (self.fa_func("fa_decode_vec_q_dpl16_dc", head_dim),
10772 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10773 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10774 } else if fa_vec && head_dim == 512 {
10775 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
10778 0, Some(t_kv_dev), scale, n_splits, sp,
10779 k_tok_bytes, v_tok_bytes, g,
10780 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10781 } else if fa_vec && head_dim == 256 && fa_v4_at(bucket_max) {
10782 let gqa = (n_head / n_head_kv).max(1) as u32;
10785 let fv = if g { self.func_g("fa_decode_vec_q_v4_dc") }
10786 else if deep { self.func("fa_decode_vec_q_v4_deep_dc") }
10787 else { self.func("fa_decode_vec_q_v4_dc") };
10788 let shmem = (if deep { 12160 } else { 11520 }
10789 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10790 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10791 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
10792 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10793 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10794 } else if fa_vec && fa_v3_active(head_dim) {
10795 let gqa = (n_head / n_head_kv).max(1) as u32;
10798 let fv = if g { self.func_g("fa_decode_vec_q_v3_dc") } else { self.func("fa_decode_vec_q_v3_dc") };
10799 let shmem = (32 * head_dim * 2) as u32; (fv,
10801 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10802 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10803 } else if fa_vec && fa_v2_on() {
10804 let gqa = (n_head / n_head_kv).max(1) as u32;
10808 let fv = if g { self.func_g("fa_decode_vec_q_v2_dc") } else { self.func("fa_decode_vec_q_v2_dc") };
10809 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
10811 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10812 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10813 } else if fa_vec {
10814 let gqa = (n_head / n_head_kv).max(1) as u32;
10815 let fv = if g { self.func_g("fa_decode_vec_q_dc") } else { self.func("fa_decode_vec_q_dc") };
10817 (fv,
10818 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10819 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10820 } else {
10821 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
10822 0, Some(t_kv_dev), scale, n_splits,
10823 if fa_vec { sp } else { 256 },
10824 k_tok_bytes, v_tok_bytes, g,
10825 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10826 };
10827 let ski = sp as i32; let __s_b = self.gpu.stream();
10829 let mut b = __s_b.launch_builder(&f);
10830 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10831 .arg(&hd).arg(&nh).arg(&nhkv).arg(t_kv_dev).arg(&scale).arg(&nsp).arg(&ski)
10832 .arg(&ktb).arg(&vtb);
10833 unsafe { b.launch(cfg)?; }
10834 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10835 if let Some((oq, od)) = q8_out {
10836 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
10837 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
10838 let __s_b2 = self.gpu.stream();
10839 let mut b2 = __s_b2.launch_builder(&fc);
10840 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
10841 unsafe { b2.launch(cfg2)?; }
10842 return Ok(());
10843 }
10844 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
10845 let __s_b2 = self.gpu.stream();
10846 let mut b2 = __s_b2.launch_builder(&fc);
10847 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
10848 unsafe { b2.launch(cfg2)?; }
10849 Ok(())
10850 }
10851
10852 pub fn fa_geom_eager(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
10858 let fa_ok = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
10862 let vec512 = fa_ok && head_dim == 512 && t_kv >= fa512_min_tkv();
10868 let mut fa_vec = vec512 || (fa_ok && head_dim <= 256 && head_dim % 32 == 0);
10869 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
10875 let sp = fa_split_keys(t_kv, n_head_kv);
10876 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
10877 (fa_vec, n_splits)
10878 }
10879
10880 pub fn fa_bucket_key(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
10886 self.fa_geom_eager(t_kv, head_dim, n_head_kv, g)
10887 }
10888
10889 pub fn capture_graph_retained<F>(&self, step: F)
10901 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
10902 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
10903 {
10904 use cudarc::driver::sys::CUgraphInstantiate_flags;
10905 self.capture_graph_retained_flags(
10906 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH, step)
10907 }
10908
10909 pub fn capture_graph_retained_flags<F>(&self,
10914 flags: cudarc::driver::sys::CUgraphInstantiate_flags, mut step: F)
10915 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
10916 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
10917 {
10918 use cudarc::driver::sys::CUstreamCaptureMode;
10919 self.capture_keep.lock().unwrap().clear();
10927 let was_tracking = self.gpu.ctx.is_event_tracking();
10928 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
10929 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
10930 self.capture_keep_on.store(true, std::sync::atomic::Ordering::Relaxed);
10931 let w = (|| { step(self)?; step(self) })();
10932 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
10933 w?;
10934 self.gpu.stream().synchronize()?;
10935 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
10936 let r = step(self);
10937 let g = self.gpu.stream().end_capture(flags);
10938 r?;
10939 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
10940 graph.upload()?;
10941 Ok(graph)
10942 };
10943 let result = run();
10944 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
10945 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
10946 let keeper = std::mem::take(&mut *self.capture_keep.lock().unwrap());
10947 Ok((result?, keeper))
10948 }
10949
10950 pub fn capture_graph<F>(&self, mut step: F) -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>>
10951 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
10952 {
10953 use cudarc::driver::sys::{CUstreamCaptureMode, CUgraphInstantiate_flags};
10954 let was_tracking = self.gpu.ctx.is_event_tracking();
10962 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
10963 let iflag = {
10970 static F: std::sync::OnceLock<CUgraphInstantiate_flags> = std::sync::OnceLock::new();
10971 *F.get_or_init(|| match std::env::var("MEMRA_GRAPH_IFLAG").as_deref() {
10972 Ok("upload") => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_UPLOAD,
10975 Ok("priority") =>
10976 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
10977 _ => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
10978 })
10979 };
10980 let ct = {
10987 static T: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10988 *T.get_or_init(|| std::env::var("MEMRA_GRAPH_CAPTIME").as_deref() == Ok("1"))
10989 };
10990 let warmups = {
11013 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
11014 *W.get_or_init(|| std::env::var("MEMRA_GRAPH_WARMUPS").ok()
11015 .and_then(|v| v.parse().ok()).filter(|n| *n >= 1).unwrap_or(1))
11016 };
11017 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
11018 let t_w = std::time::Instant::now();
11019 for _ in 0..warmups { step(self)?; }
11021 self.gpu.stream().synchronize()?;
11022 let ms_warm = t_w.elapsed().as_secs_f64() * 1e3;
11023 let t_c = std::time::Instant::now();
11025 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
11026 let r = step(self);
11029 let ms_body = t_c.elapsed().as_secs_f64() * 1e3;
11030 let t_i = std::time::Instant::now();
11031 let g = self.gpu.stream().end_capture(iflag);
11032 let ms_inst = t_i.elapsed().as_secs_f64() * 1e3;
11033 r?;
11034 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
11035 let t_u = std::time::Instant::now();
11036 graph.upload()?;
11037 if ct {
11038 println!("[graph-captime] warmup2x {ms_warm:.2} ms capture-body {ms_body:.2} ms \
11039 instantiate {ms_inst:.2} ms upload {:.2} ms",
11040 t_u.elapsed().as_secs_f64() * 1e3);
11041 }
11042 Ok(graph)
11043 };
11044 let result = run();
11045 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
11046 result
11047 }
11048
11049 pub fn gdn_scan_s128_view(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11051 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
11052 state_in: &cudarc::driver::CudaView<f32>,
11053 state_out: &mut cudarc::driver::CudaViewMut<f32>,
11054 o: &mut CudaSlice<f32>, n_head: usize, t: usize, scale: f32)
11055 -> Result<(), Box<dyn std::error::Error>> {
11056 let f = self.func("gdn_scan_s128");
11057 const S_V: u32 = 128; const WARP: u32 = 32; const COLS: u32 = 4;
11058 let cfg = LaunchConfig { grid_dim: (n_head as u32, 1, S_V / COLS), block_dim: (WARP, COLS, 1), shared_mem_bytes: 0 };
11059 let (h, ti) = (n_head as i32, t as i32);
11060 let __s_b = self.gpu.stream();
11061 let mut b = __s_b.launch_builder(&f);
11062 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);
11063 unsafe { b.launch(cfg)?; }
11064 Ok(())
11065 }
11066
11067 pub fn ssm_conv1d_view(&self, x: &cudarc::driver::CudaView<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11069 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
11070 -> Result<(), Box<dyn std::error::Error>> {
11071 let f = self.func("ssm_conv1d_silu_f32");
11072 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
11074 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11075 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
11076 let __s_b = self.gpu.stream();
11077 let mut b = __s_b.launch_builder(&f);
11078 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
11079 unsafe { b.launch(cfg)?; }
11080 Ok(())
11081 }
11082
11083 pub fn ssm_conv1d_tm(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11090 conv_dim: usize, t: usize, d_conv: usize)
11091 -> Result<(), Box<dyn std::error::Error>> {
11092 let f = self.func("ssm_conv1d_tm_f32");
11093 let cfg = LaunchConfig {
11094 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11095 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11096 };
11097 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11098 let __s_b = self.gpu.stream();
11099 let mut b = __s_b.launch_builder(&f);
11100 b.arg(qkv_tm).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11101 unsafe { b.launch(cfg)?; }
11102 Ok(())
11103 }
11104
11105 pub fn ssm_conv1d_tm_state(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
11113 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11114 conv_dim: usize, t: usize, d_conv: usize)
11115 -> Result<(), Box<dyn std::error::Error>> {
11116 self.ssm_conv1d_tm_state_pad(qkv_tm, conv_state, w, y, conv_dim, t, d_conv, None)
11117 }
11118
11119 #[allow(clippy::too_many_arguments)]
11122 pub fn ssm_conv1d_tm_state_pad(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
11123 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11124 conv_dim: usize, t: usize, d_conv: usize,
11125 pad_len: Option<&CudaSlice<i32>>)
11126 -> Result<(), Box<dyn std::error::Error>> {
11127 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
11128 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11132 {
11133 let f = self.func("ssm_conv1d_tm_state_f32");
11134 let cfg = LaunchConfig {
11135 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11136 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11137 };
11138 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11139 let __s_b = self.gpu.stream();
11140 let mut b = __s_b.launch_builder(&f);
11141 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11142 unsafe { b.launch(cfg)?; }
11143 }
11144 match (ring_old, pad_len) {
11145 (None, Some(len_d)) => {
11146 let f = self.func("ssm_conv_ring_update_dev_f32");
11147 let n = conv_dim * (d_conv - 1);
11148 let cfg = LaunchConfig::for_num_elems(n as u32);
11149 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11150 let __s_b = self.gpu.stream();
11151 let mut b = __s_b.launch_builder(&f);
11152 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11153 unsafe { b.launch(cfg)?; }
11154 }
11155 (None, None) => {
11156 let f = self.func("ssm_conv_ring_update_f32");
11157 let n = conv_dim * (d_conv - 1);
11158 let cfg = LaunchConfig::for_num_elems(n as u32);
11159 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11160 let __s_b = self.gpu.stream();
11161 let mut b = __s_b.launch_builder(&f);
11162 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11163 unsafe { b.launch(cfg)?; }
11164 }
11165 (Some(old), _) => self.ssm_conv_ring_rebuild(qkv_tm, &old, conv_state, conv_dim, t, d_conv)?,
11166 }
11167 Ok(())
11168 }
11169
11170 pub fn ssm_conv1d_tm_state_pad_v(&self, qkv_tm: &cudarc::driver::CudaView<f32>, conv_state: &mut CudaSlice<f32>,
11172 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11173 conv_dim: usize, t: usize, d_conv: usize,
11174 pad_len: Option<&CudaSlice<i32>>)
11175 -> Result<(), Box<dyn std::error::Error>> {
11176 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
11177 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11181 {
11182 let f = self.func("ssm_conv1d_tm_state_f32");
11183 let cfg = LaunchConfig {
11184 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11185 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11186 };
11187 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11188 let __s_b = self.gpu.stream();
11189 let mut b = __s_b.launch_builder(&f);
11190 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11191 unsafe { b.launch(cfg)?; }
11192 }
11193 match (ring_old, pad_len) {
11194 (None, Some(len_d)) => {
11195 let f = self.func("ssm_conv_ring_update_dev_f32");
11196 let n = conv_dim * (d_conv - 1);
11197 let cfg = LaunchConfig::for_num_elems(n as u32);
11198 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11199 let __s_b = self.gpu.stream();
11200 let mut b = __s_b.launch_builder(&f);
11201 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11202 unsafe { b.launch(cfg)?; }
11203 }
11204 (None, None) => {
11205 let f = self.func("ssm_conv_ring_update_f32");
11206 let n = conv_dim * (d_conv - 1);
11207 let cfg = LaunchConfig::for_num_elems(n as u32);
11208 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11209 let __s_b = self.gpu.stream();
11210 let mut b = __s_b.launch_builder(&f);
11211 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11212 unsafe { b.launch(cfg)?; }
11213 }
11214 (Some(_), _) => unreachable!(
11215 "ssm_conv1d_tm_state_pad_v: T < d_conv-1 has no view path (PRIME_MIN_T gates it)"),
11216 }
11217 Ok(())
11218 }
11219
11220 pub fn ssm_conv_ring_rebuild(&self, qkv_tm: &CudaSlice<f32>, ring_old: &CudaSlice<f32>,
11225 conv_state: &mut CudaSlice<f32>,
11226 conv_dim: usize, tc: usize, d_conv: usize)
11227 -> Result<(), Box<dyn std::error::Error>> {
11228 let f = self.func("ssm_conv_ring_rebuild_f32");
11229 let n = conv_dim * (d_conv - 1);
11230 let cfg = LaunchConfig::for_num_elems(n as u32);
11231 let (cd, ti, dc) = (conv_dim as i32, tc as i32, d_conv as i32);
11232 let __s_b = self.gpu.stream();
11233 let mut b = __s_b.launch_builder(&f);
11234 b.arg(qkv_tm).arg(ring_old).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11235 unsafe { b.launch(cfg)?; }
11236 Ok(())
11237 }
11238
11239 #[allow(clippy::too_many_arguments)]
11244 pub fn gdn_prep_decode(&self, conv_out: &CudaSlice<f32>, beta_raw: &CudaSlice<f32>,
11245 alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11246 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11247 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11248 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32)
11249 -> Result<(), Box<dyn std::error::Error>> {
11250 let f = self.func("gdn_prep_decode_f32");
11251 let cfg = LaunchConfig { grid_dim: (num_v as u32, 1, 1), block_dim: (32, 4, 1), shared_mem_bytes: 0 };
11252 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11253 let __s_b = self.gpu.stream();
11254 let mut b = __s_b.launch_builder(&f);
11255 b.arg(conv_out).arg(beta_raw).arg(alpha).arg(dt_bias).arg(a)
11256 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11257 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps);
11258 unsafe { b.launch(cfg)?; }
11259 Ok(())
11260 }
11261
11262 #[allow(clippy::too_many_arguments)]
11266 pub fn ssm_conv1d_gdn(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>,
11267 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11268 conv_dim: usize, t: usize, d_conv: usize,
11269 d_state: usize, num_v: usize, num_k: usize, key_dim: usize)
11270 -> Result<(), Box<dyn std::error::Error>> {
11271 let f = self.func("ssm_conv1d_gdn_f32");
11272 let cfg = LaunchConfig {
11273 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11274 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11275 };
11276 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11277 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11278 let __s_b = self.gpu.stream();
11279 let mut b = __s_b.launch_builder(&f);
11280 b.arg(qkv_tm).arg(w).arg(q_g).arg(k_g).arg(v_g)
11281 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd);
11282 unsafe { b.launch(cfg)?; }
11283 Ok(())
11284 }
11285
11286 pub fn ssm_conv1d(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11287 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
11288 -> Result<(), Box<dyn std::error::Error>> {
11289 let f = self.func("ssm_conv1d_silu_f32");
11290 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
11291 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11292 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
11293 let __s_b = self.gpu.stream();
11294 let mut b = __s_b.launch_builder(&f);
11295 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
11296 unsafe { b.launch(cfg)?; }
11297 Ok(())
11298 }
11299
11300 pub fn gdn_scan_s128(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11303 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
11304 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11305 n_head: usize, t: usize, scale: f32)
11306 -> Result<(), Box<dyn std::error::Error>> {
11307 let f = self.func("gdn_scan_s128");
11308 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11309 let cfg = LaunchConfig {
11310 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
11311 block_dim: (WARP, COLS_PER_BLOCK, 1),
11312 shared_mem_bytes: 0,
11313 };
11314 let (h, ti) = (n_head as i32, t as i32);
11315 let __s_b = self.gpu.stream();
11316 let mut b = __s_b.launch_builder(&f);
11317 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);
11318 unsafe { b.launch(cfg)?; }
11319 Ok(())
11320 }
11321
11322 #[allow(clippy::too_many_arguments)]
11327 pub fn ssm_conv1d_fused_decode_b(
11328 &self, qkv_cols: &CudaSlice<f32>, conv_state_ptrs: &cudarc::driver::CudaView<u64>,
11329 w: &CudaSlice<f32>, conv_outs: &mut CudaSlice<f32>, conv_dim: usize, d_conv: usize,
11330 b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11331 let f = self.func("ssm_conv1d_fused_decode_b_f32");
11332 let cfg = LaunchConfig {
11333 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
11334 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11335 };
11336 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11337 let __s_b = self.gpu.stream();
11338 let mut b = __s_b.launch_builder(&f);
11339 b.arg(qkv_cols).arg(conv_state_ptrs).arg(w).arg(conv_outs).arg(&cd).arg(&dc);
11340 unsafe { b.launch(cfg)?; }
11341 Ok(())
11342 }
11343
11344 #[allow(clippy::too_many_arguments)]
11345 pub fn gdn_prep_decode_b(
11346 &self, conv_outs: &CudaSlice<f32>, beta_raws: &CudaSlice<f32>, alphas: &CudaSlice<f32>,
11347 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11348 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11349 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11350 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32,
11351 conv_dim: usize, b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11352 let f = self.func("gdn_prep_decode_b_f32");
11353 let cfg = LaunchConfig {
11354 grid_dim: (num_v as u32, 1, b_n as u32),
11355 block_dim: (32, 4, 1), shared_mem_bytes: 0,
11356 };
11357 let (ds, nv, nk, kd, cd) =
11358 (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, conv_dim as i32);
11359 let __s_b = self.gpu.stream();
11360 let mut b = __s_b.launch_builder(&f);
11361 b.arg(conv_outs).arg(beta_raws).arg(alphas).arg(dt_bias).arg(a)
11362 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11363 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps).arg(&cd);
11364 unsafe { b.launch(cfg)?; }
11365 Ok(())
11366 }
11367
11368 #[allow(clippy::too_many_arguments)]
11369 pub fn gdn_scan_s128_batched(
11370 &self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11371 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
11372 state_in_ptrs: &cudarc::driver::CudaView<u64>,
11373 state_out_ptrs: &cudarc::driver::CudaView<u64>,
11374 o: &mut CudaSlice<f32>, n_head: usize, b_n: usize, scale: f32)
11375 -> Result<(), Box<dyn std::error::Error>> {
11376 let f = self.func("gdn_scan_s128_b");
11377 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11378 let cfg = LaunchConfig {
11379 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
11380 block_dim: (WARP, COLS_PER_BLOCK, 1), shared_mem_bytes: 0,
11381 };
11382 let h = n_head as i32;
11383 let __s_b = self.gpu.stream();
11384 let mut b = __s_b.launch_builder(&f);
11385 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in_ptrs).arg(state_out_ptrs)
11386 .arg(o).arg(&h).arg(&scale);
11387 unsafe { b.launch(cfg)?; }
11388 Ok(())
11389 }
11390
11391 pub fn gdn_chunked_enabled() -> bool {
11400 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11401 *E.get_or_init(|| std::env::var("MEMRA_GDN_CHUNKED").map(|v| v != "0").unwrap_or(true))
11402 }
11403
11404 pub fn gdn_chunk_size() -> usize {
11409 static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
11410 *C.get_or_init(|| {
11411 let c: usize = std::env::var("MEMRA_GDN_CHUNK").ok()
11412 .and_then(|v| v.parse().ok()).unwrap_or(32);
11413 c.clamp(32, 128) / 32 * 32
11414 })
11415 }
11416
11417 #[allow(clippy::too_many_arguments)]
11422 #[allow(clippy::too_many_arguments, clippy::type_complexity)]
11425 #[allow(clippy::too_many_arguments)]
11426 pub fn gdn_chunk_k123(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11427 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, wb16: Option<&mut CudaSlice<u8>>,
11428 n_head: usize, t: usize, c: usize, hk: usize,
11429 k2w: Option<(&CudaSlice<u8>, &CudaSlice<u8>, &mut CudaSlice<u8>)>)
11430 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11431 const D: usize = 128;
11432 let h = n_head;
11433 let nc = (t + c - 1) / c;
11434 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
11435 let mut gcum = self.uninit(t * h)?;
11436 let mut a = self.uninit(nc * h * c * c)?;
11437 let mut p = self.uninit(nc * h * c * c)?;
11438 let mut u = self.uninit(nc * h * c * D)?;
11439 let mut w = self.uninit(nc * h * c * D)?;
11440 { let f = self.func("gdn_chunk_cumgate_f32");
11442 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11443 let __s_b = self.gpu.stream();
11444 let mut b = __s_b.launch_builder(&f);
11445 b.arg(g).arg(&mut gcum).arg(&hi).arg(&ti).arg(&ci);
11446 unsafe { b.launch(cfg)?; }
11447 }
11448 if let Some((qb, kb, pb)) = k2w {
11449 assert!(c == 32, "gdn_k2_wgmma is a C==32 tile");
11452 let f = self.func("gdn_k2_wgmma");
11453 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11454 let hki = hk as i32;
11455 let __s_b = self.gpu.stream();
11456 let mut b = __s_b.launch_builder(&f);
11457 b.arg(qb).arg(kb).arg(&gcum).arg(beta).arg(&mut a).arg(&mut *pb).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11458 unsafe { b.launch(cfg)?; }
11459 } else if c <= 64 && !portable_mma_gated() { let f = self.func("gdn_chunk_attn_f32");
11461 let jt = ((c + 31) / 32) as u32;
11462 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11463 let hki = hk as i32;
11464 let __s_b = self.gpu.stream();
11465 let mut b = __s_b.launch_builder(&f);
11466 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11467 unsafe { b.launch(cfg)?; }
11468 } else { assert!(hk == h, "generic K2 is broadcast-only (de-broadcast rides C==32)");
11470 let f = self.func("gdn_chunk_attn_g_f32");
11471 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 8, 1), shared_mem_bytes: 0 };
11472 let __s_b = self.gpu.stream();
11473 let mut b = __s_b.launch_builder(&f);
11474 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci);
11475 unsafe { b.launch(cfg)?; }
11476 }
11477 { let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11479 match c {
11480 32 | 64 => {
11481 let f = self.func(if c == 32 { "gdn_chunk_solve32_f32" } else { "gdn_chunk_solve64_f32" });
11482 let wb: u64 = match wb16 { Some(d) => self.addr_u8(d), None => 0 };
11484 let hki = hk as i32;
11485 let __s_b = self.gpu.stream();
11486 let mut b = __s_b.launch_builder(&f);
11487 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&wb).arg(&hi).arg(&ti).arg(&hki);
11488 unsafe { b.launch(cfg)?; }
11489 }
11490 _ => {
11491 assert!(hk == h, "generic K3 is broadcast-only");
11492 let f = self.func("gdn_chunk_solve_f32");
11493 let __s_b = self.gpu.stream();
11494 let mut b = __s_b.launch_builder(&f);
11495 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&hi).arg(&ti).arg(&ci);
11496 unsafe { b.launch(cfg)?; }
11497 }
11498 }
11499 }
11500 Ok((gcum, p, u, w))
11501 }
11502
11503 pub fn gdn_db_on() -> bool {
11507 std::env::var("MEMRA_GDN_DB").as_deref() != Ok("0")
11508 }
11509
11510 pub fn gdn_mma_enabled(&self, c: usize) -> bool {
11513 !portable_mma_gated() && c == 32
11514 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11515 Ok("1") => true,
11516 Ok("0") => false,
11517 _ => cfg!(memra_hopper_mma),
11518 }
11519 }
11520
11521 pub fn gdn_wgmma_on(&self, c: usize) -> bool {
11524 self.gdn_mma_enabled(c)
11525 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
11526 Ok("0") => false,
11527 Ok("1") => true,
11528 _ => cfg!(memra_hopper_mma),
11529 }
11530 }
11531
11532 #[allow(clippy::too_many_arguments)]
11537 pub fn ssm_conv1d_gdn_state_pad(&self, qkv_tm: &cudarc::driver::CudaView<f32>,
11538 conv_state: &mut CudaSlice<f32>, w: &CudaSlice<f32>,
11539 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>,
11540 v_g: &mut CudaSlice<f32>,
11541 conv_dim: usize, t: usize, d_conv: usize,
11542 d_state: usize, num_v: usize, num_k: usize, key_dim: usize,
11543 hk: usize,
11544 pad_len: Option<&CudaSlice<i32>>)
11545 -> Result<(), Box<dyn std::error::Error>> {
11546 assert!(t >= d_conv - 1, "fused state conv requires T >= pad (PRIME_MIN_T gates)");
11547 {
11548 let f = self.func("ssm_conv1d_gdn_state_f32");
11549 let cfg = LaunchConfig {
11550 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11551 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11552 };
11553 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11554 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);
11555 let __s_b = self.gpu.stream();
11556 let mut b = __s_b.launch_builder(&f);
11557 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(q_g).arg(k_g).arg(v_g)
11558 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&hki);
11559 unsafe { b.launch(cfg)?; }
11560 }
11561 match pad_len {
11562 Some(len_d) => {
11563 let f = self.func("ssm_conv_ring_update_dev_f32");
11564 let n = conv_dim * (d_conv - 1);
11565 let cfg = LaunchConfig::for_num_elems(n as u32);
11566 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11567 let __s_b = self.gpu.stream();
11568 let mut b = __s_b.launch_builder(&f);
11569 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11570 unsafe { b.launch(cfg)?; }
11571 }
11572 None => {
11573 let f = self.func("ssm_conv_ring_update_f32");
11574 let n = conv_dim * (d_conv - 1);
11575 let cfg = LaunchConfig::for_num_elems(n as u32);
11576 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11577 let __s_b = self.gpu.stream();
11578 let mut b = __s_b.launch_builder(&f);
11579 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11580 unsafe { b.launch(cfg)?; }
11581 }
11582 }
11583 Ok(())
11584 }
11585
11586 pub fn gdn_chunk_alloc(&self, n_head: usize, t: usize, c: usize, hk: usize)
11590 -> Result<GdnChunkBufs, Box<dyn std::error::Error>> {
11591 const D: usize = 128;
11592 assert!(c == 32, "gdn_chunk_alloc: varlen chain is the C==32 mma pair");
11593 let h = n_head;
11594 let nc = (t + c - 1) / c;
11595 Ok(GdnChunkBufs {
11596 gcum: self.uninit(t * h)?,
11597 a: self.uninit(nc * h * c * c)?,
11598 p: self.uninit(nc * h * c * c)?,
11599 u: self.uninit(nc * h * c * D)?,
11600 w: self.uninit(nc * h * c * D)?,
11601 kb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11602 wb16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11603 y16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11604 ssnap16: self.alloc_u8_uninit(nc * h * D * D * 2)?,
11605 qb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11606 pb16: self.alloc_u8_uninit(nc * h * c * c * 2)?,
11607 o: self.uninit(D * h * t)?,
11608 t, nc,
11609 })
11610 }
11611
11612 pub fn f32_to_bf16_v(&self, x: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<u8>, n: usize)
11614 -> Result<(), Box<dyn std::error::Error>> {
11615 let f = self.func("f32_to_bf16_bulk");
11616 let ni = n as i64;
11617 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11618 let __s_b = self.gpu.stream();
11619 let mut b = __s_b.launch_builder(&f);
11620 b.arg(x).arg(dst).arg(&ni);
11621 unsafe { b.launch(cfg)?; }
11622 Ok(())
11623 }
11624
11625 pub fn f32_to_bf16_into(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<u8>, n: usize)
11627 -> Result<(), Box<dyn std::error::Error>> {
11628 let f = self.func("f32_to_bf16_bulk");
11629 let ni = n as i64;
11630 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11631 let __s_b = self.gpu.stream();
11632 let mut b = __s_b.launch_builder(&f);
11633 b.arg(x).arg(dst).arg(&ni);
11634 unsafe { b.launch(cfg)?; }
11635 Ok(())
11636 }
11637
11638 pub fn gdn_chunk_k123_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, hk: usize,
11641 wq: Option<&GdnWVl8>)
11642 -> Result<(), Box<dyn std::error::Error>> {
11643 let b = seqs.len();
11644 assert!(b >= 1 && b <= 8, "gdn_chunk_k123_vl8: 1..=8 sequences");
11645 let mut packed = [GdnSeqVl::default(); 8];
11646 packed[..b].copy_from_slice(seqs);
11647 let v = GdnVl8(packed);
11648 let (hi, ci) = (n_head as i32, 32i32);
11649 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11650 {
11651 let f = self.func("gdn_chunk_cumgate_vl");
11652 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11653 let __s_lb = self.gpu.stream();
11654 let mut lb = __s_lb.launch_builder(&f);
11655 lb.arg(&v).arg(&hi).arg(&ci);
11656 unsafe { lb.launch(cfg)?; }
11657 }
11658 let hki = hk as i32;
11659 if let Some(w) = wq { let f = self.func("gdn_k2_wgmma_vl");
11661 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11662 let __s_lb = self.gpu.stream();
11663 let mut lb = __s_lb.launch_builder(&f);
11664 lb.arg(&v).arg(w).arg(&hi).arg(&ci).arg(&hki);
11665 unsafe { lb.launch(cfg)?; }
11666 } else {
11667 let f = self.func("gdn_chunk_attn_vl");
11668 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11669 let __s_lb = self.gpu.stream();
11670 let mut lb = __s_lb.launch_builder(&f);
11671 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11672 unsafe { lb.launch(cfg)?; }
11673 }
11674 {
11675 let f = self.func("gdn_chunk_solve32_vl");
11676 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11677 let __s_lb = self.gpu.stream();
11678 let mut lb = __s_lb.launch_builder(&f);
11679 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11680 unsafe { lb.launch(cfg)?; }
11681 }
11682 Ok(())
11683 }
11684
11685 #[allow(clippy::too_many_arguments)]
11689 pub fn gdn_prep_vl8(&self, seqs: &[GdnPrepVl], conv_w: &CudaSlice<f32>,
11690 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11691 conv_dim: usize, d_conv: usize, d_state: usize,
11692 num_v: usize, num_k: usize, key_dim: usize, hk: usize, eps: f32)
11693 -> Result<(), Box<dyn std::error::Error>> {
11694 let b = seqs.len();
11695 assert!(b >= 1 && b <= 8);
11696 let mut packed = [GdnPrepVl::default(); 8];
11697 packed[..b].copy_from_slice(seqs);
11698 let v = GdnPrepVl8(packed);
11699 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11700 let (cdi, dci) = (conv_dim as i32, d_conv as i32);
11701 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
11702 assert!(conv_fuse || hk == num_v, "de-broadcast requires the fused conv");
11703 if conv_fuse {
11704 let f = self.func("ssm_conv1d_gdn_state_vl");
11705 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 };
11706 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);
11707 let __s_lb = self.gpu.stream();
11708 let mut lb = __s_lb.launch_builder(&f);
11709 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi).arg(&hki);
11710 unsafe { lb.launch(cfg)?; }
11711 } else {
11712 let f = self.func("ssm_conv1d_tm_state_vl");
11713 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 };
11714 let __s_lb = self.gpu.stream();
11715 let mut lb = __s_lb.launch_builder(&f);
11716 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci);
11717 unsafe { lb.launch(cfg)?; }
11718 }
11719 {
11720 let f = self.func("ssm_conv_ring_update_vl");
11721 let n = (conv_dim * (d_conv - 1)) as u32;
11722 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11723 let __s_lb = self.gpu.stream();
11724 let mut lb = __s_lb.launch_builder(&f);
11725 lb.arg(&v).arg(&cdi).arg(&dci);
11726 unsafe { lb.launch(cfg)?; }
11727 }
11728 if !conv_fuse {
11729 let f = self.func("qkv_to_gdn_repack_vl");
11730 let n = max_t * (num_v * d_state) as u32;
11731 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11732 let (dsi, nvi, nki, kdi) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11733 let __s_lb = self.gpu.stream();
11734 let mut lb = __s_lb.launch_builder(&f);
11735 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi);
11736 unsafe { lb.launch(cfg)?; }
11737 }
11738 if Self::l2_v2_on(d_state) {
11739 let f = self.func("gdn_l2_v2_vl");
11740 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 };
11741 let (dsi, nvi) = (d_state as i32, hk as i32);
11742 let __s_lb = self.gpu.stream();
11743 let mut lb = __s_lb.launch_builder(&f);
11744 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11745 unsafe { lb.launch(cfg)?; }
11746 } else {
11747 let f = self.func("gdn_l2_vl");
11748 let cfg = LaunchConfig { grid_dim: (max_t * hk as u32, 2, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11749 let (dsi, nvi) = (d_state as i32, hk as i32);
11750 let __s_lb = self.gpu.stream();
11751 let mut lb = __s_lb.launch_builder(&f);
11752 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11753 unsafe { lb.launch(cfg)?; }
11754 }
11755 {
11756 let f = self.func("gdn_gate_prep_vl");
11757 let n = max_t * num_v as u32;
11758 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11759 let nvi = num_v as i32;
11760 let __s_lb = self.gpu.stream();
11761 let mut lb = __s_lb.launch_builder(&f);
11762 lb.arg(&v).arg(dt_bias).arg(a).arg(&nvi);
11763 unsafe { lb.launch(cfg)?; }
11764 }
11765 Ok(())
11766 }
11767
11768 pub fn gdn_mirror_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, which: i32, hk: usize)
11770 -> Result<(), Box<dyn std::error::Error>> {
11771 let b = seqs.len();
11772 assert!(b >= 1 && b <= 8);
11773 let mut packed = [GdnSeqVl::default(); 8];
11774 packed[..b].copy_from_slice(seqs);
11775 let v = GdnVl8(packed);
11776 let ept = (if which == 0 { hk } else { n_head } * 128) as i32;
11777 let max_n = seqs.iter().map(|s| if which == 0 { s.t as i64 * ept as i64 }
11778 else { s.nc as i64 * ept as i64 * 32 }).max().unwrap();
11779 let f = self.func("gdn_mirror_vl");
11780 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
11781 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11782 let __s_lb = self.gpu.stream();
11783 let mut lb = __s_lb.launch_builder(&f);
11784 lb.arg(&v).arg(&ept).arg(&which);
11785 unsafe { lb.launch(cfg)?; }
11786 Ok(())
11787 }
11788
11789 pub fn gdn_tail_vl8(&self, seqs: &[GdnPrepVl], norm_w: &CudaSlice<f32>,
11791 d_state: usize, num_v: usize, eps: f32)
11792 -> Result<(), Box<dyn std::error::Error>> {
11793 let b = seqs.len();
11794 assert!(b >= 1 && b <= 8);
11795 let mut packed = [GdnPrepVl::default(); 8];
11796 packed[..b].copy_from_slice(seqs);
11797 let v = GdnPrepVl8(packed);
11798 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11799 let f = self.func("gated_rmsnorm_f16out_vl");
11800 let cfg = LaunchConfig { grid_dim: (max_t * num_v as u32, 1, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11802 let (dsi, nvi) = (d_state as i32, num_v as i32);
11803 let __s_lb = self.gpu.stream();
11804 let mut lb = __s_lb.launch_builder(&f);
11805 lb.arg(&v).arg(norm_w).arg(&dsi).arg(&nvi).arg(&eps);
11806 unsafe { lb.launch(cfg)?; }
11807 Ok(())
11808 }
11809
11810 pub fn addr_f32(&self, x: &CudaSlice<f32>) -> u64 {
11813 use cudarc::driver::DevicePtr;
11814 let s = self.gpu.stream();
11815 let (p, _g) = x.device_ptr(&s);
11816 p as u64
11817 }
11818 pub fn addr_f32_mut(&self, x: &mut CudaSlice<f32>) -> u64 {
11819 use cudarc::driver::DevicePtrMut;
11820 let s = self.gpu.stream();
11821 let (p, _g) = x.device_ptr_mut(&s);
11822 p as u64
11823 }
11824 pub fn addr_f32v(&self, x: &cudarc::driver::CudaView<f32>) -> u64 {
11825 use cudarc::driver::DevicePtr;
11826 let s = self.gpu.stream();
11827 let (p, _g) = x.device_ptr(&s);
11828 p as u64
11829 }
11830 pub fn addr_u8(&self, x: &CudaSlice<u8>) -> u64 {
11831 use cudarc::driver::DevicePtr;
11832 let s = self.gpu.stream();
11833 let (p, _g) = x.device_ptr(&s);
11834 p as u64
11835 }
11836
11837 pub fn gdn_chunk_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, scale: f32, hk: usize,
11841 wq: Option<&GdnWVl8>)
11842 -> Result<(), Box<dyn std::error::Error>> {
11843 const NSPLIT: u32 = 4;
11844 let b = seqs.len();
11845 assert!(b >= 1 && b <= 8, "gdn_chunk_vl8: 1..=8 sequences");
11846 let mut packed = [GdnSeqVl::default(); 8];
11847 packed[..b].copy_from_slice(seqs);
11848 let v = GdnVl8(packed);
11849 let (hi, ci) = (n_head as i32, 32i32);
11850 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11851 let hki = hk as i32;
11852 if let Some(w) = wq {
11853 let f = self.func("gdn_k45_wgmma_vl");
11855 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11856 let __s_lb = self.gpu.stream();
11857 let mut lb = __s_lb.launch_builder(&f);
11858 lb.arg(&v).arg(w).arg(&scale).arg(&hi).arg(&ci).arg(&hki);
11859 unsafe { lb.launch(cfg)?; }
11860 let _ = max_nc;
11861 return Ok(());
11862 }
11863 {
11864 let f = self.func("gdn_chunk_state_mma_vl");
11865 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11866 let __s_lb = self.gpu.stream();
11867 let mut lb = __s_lb.launch_builder(&f);
11868 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11869 unsafe { lb.launch(cfg)?; }
11870 }
11871 {
11872 let f = self.func("gdn_chunk_output_mma_vl");
11873 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11874 let __s_lb = self.gpu.stream();
11875 let mut lb = __s_lb.launch_builder(&f);
11876 lb.arg(&v).arg(&hi).arg(&ci).arg(&scale).arg(&hki);
11877 unsafe { lb.launch(cfg)?; }
11878 }
11879 Ok(())
11880 }
11881 pub fn gdn_scan_chunked(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11882 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
11883 qb16_pre: Option<&CudaSlice<u8>>,
11884 state_in: &CudaSlice<f32>,
11885 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11886 n_head: usize, t: usize, scale: f32, c: usize, hk: usize)
11887 -> Result<(), Box<dyn std::error::Error>> {
11888 const D: usize = 128;
11889 const NSPLIT: u32 = 4;
11890 assert!(c >= 1 && c <= 128, "gdn_scan_chunked: C must be in 1..=128");
11891 let h = n_head;
11892 let nc = (t + c - 1) / c;
11893 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
11894 let gdn_mma_pre = !portable_mma_gated() && c == 32
11898 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11899 Ok("1") => true,
11900 Ok("0") => false,
11901 _ => cfg!(memra_hopper_mma),
11902 };
11903 let mut wb16_pre: Option<CudaSlice<u8>> = if gdn_mma_pre {
11904 Some(self.alloc_u8_uninit(nc * h * c * D * 2)?)
11905 } else { None };
11906 let gdn_wgmma_pre = gdn_mma_pre
11910 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
11911 Ok("0") => false,
11912 Ok("1") => true,
11913 _ => cfg!(memra_hopper_mma),
11914 };
11915 let nk = t * hk * D;
11916 let mut kb16_local: Option<CudaSlice<u8>> = None;
11917 if gdn_mma_pre && kb16_pre.is_none() {
11918 let mut kb = self.alloc_u8_uninit(nk * 2)?;
11919 let f = self.func("f32_to_bf16_bulk");
11920 let n2 = nk as i64;
11921 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
11922 let __s_b = self.gpu.stream();
11923 let mut b = __s_b.launch_builder(&f);
11924 b.arg(k).arg(&mut kb).arg(&n2);
11925 unsafe { b.launch(cfg2)?; }
11926 kb16_local = Some(kb);
11927 }
11928 let kb16_ref0: Option<&CudaSlice<u8>> = kb16_local.as_ref().or(kb16_pre);
11929 if let Some(kb) = kb16_pre { assert!(kb.len() >= nk * 2, "kb16_pre too small"); }
11930 let mut qb16: Option<CudaSlice<u8>> = None;
11931 let mut pb16: Option<CudaSlice<u8>> = None;
11932 if gdn_wgmma_pre {
11933 if qb16_pre.is_none() {
11936 let mut qb = self.alloc_u8_uninit(nk * 2)?;
11937 let f = self.func("f32_to_bf16_bulk");
11938 let n2 = nk as i64;
11939 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
11940 let __s_b = self.gpu.stream();
11941 let mut b = __s_b.launch_builder(&f);
11942 b.arg(q).arg(&mut qb).arg(&n2);
11943 unsafe { b.launch(cfg2)?; }
11944 qb16 = Some(qb);
11945 } else if let Some(qb) = qb16_pre {
11946 assert!(qb.len() >= nk * 2, "qb16_pre too small");
11947 }
11948 pb16 = Some(self.alloc_u8_uninit(nc * h * c * c * 2)?);
11949 }
11950 let qb16_ref0: Option<&CudaSlice<u8>> = qb16.as_ref().or(qb16_pre);
11951 let k2w = if gdn_wgmma_pre {
11952 Some((*qb16_ref0.as_ref().unwrap(),
11953 *kb16_ref0.as_ref().unwrap(),
11954 pb16.as_mut().unwrap()))
11955 } else { None };
11956 let (gcum, p, u, w) = self.gdn_chunk_k123(q, k, v, g, beta, wb16_pre.as_mut(), n_head, t, c, hk, k2w)?;
11957 let _ = &w;
11958 let mut y = self.uninit(nc * h * c * D)?;
11959 let mut ssnap = self.uninit(nc * h * D * D)?; let gdn_mma = !portable_mma_gated() && c == 32
11971 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11972 Ok("1") => true,
11973 Ok("0") => false,
11974 _ => cfg!(memra_hopper_mma),
11975 };
11976 if gdn_mma {
11977 let wb16 = wb16_pre.take().expect("mma path pre-allocates wb16 (K3 store fold)");
11978 let kb16_ref: &CudaSlice<u8> = kb16_ref0.expect("mma path pre-builds kb16 above K123");
11979 if gdn_wgmma_pre {
11991 let qb16 = qb16_ref0.unwrap();
11993 let pb16 = pb16.as_ref().unwrap();
11994 {
11995 let f = self.func("gdn_k45_wgmma");
11996 let cfg = LaunchConfig { grid_dim: (h as u32, 4, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11997 let hki = hk as i32;
11998 let __s_b = self.gpu.stream();
11999 let mut b = __s_b.launch_builder(&f);
12000 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(qb16).arg(pb16)
12001 .arg(o).arg(&scale).arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
12002 unsafe { b.launch(cfg)?; }
12003 }
12004 return Ok(());
12005 }
12006 let mut y16 = self.alloc_u8_uninit(nc * h * c * D * 2)?;
12010 let mut ssnap16 = self.alloc_u8_uninit(nc * h * D * D * 2)?;
12011 {
12012 let f = self.func("gdn_chunk_state_mma");
12013 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12014 let hki = hk as i32;
12015 let __s_b = self.gpu.stream();
12016 let mut b = __s_b.launch_builder(&f);
12017 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(&mut y16).arg(&mut ssnap16)
12018 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
12019 unsafe { b.launch(cfg)?; }
12020 }
12021 { let f = self.func("gdn_chunk_output_mma");
12023 let jt = ((c + 31) / 32) as u32;
12024 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12025 let hki = hk as i32;
12026 let __s_b = self.gpu.stream();
12027 let mut b = __s_b.launch_builder(&f);
12028 b.arg(q).arg(&gcum).arg(&p).arg(&y16).arg(&ssnap16).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale).arg(&hki);
12029 unsafe { b.launch(cfg)?; }
12030 }
12031 return Ok(());
12032 }
12033 { let f = self.func("gdn_chunk_state_f32");
12035 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12036 let __s_b = self.gpu.stream();
12037 let mut b = __s_b.launch_builder(&f);
12038 b.arg(k).arg(&gcum).arg(beta).arg(&u).arg(&w).arg(&mut y).arg(&mut ssnap)
12039 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci);
12040 unsafe { b.launch(cfg)?; }
12041 }
12042 { let f = self.func("gdn_chunk_output_f32");
12044 let jt = ((c + 31) / 32) as u32;
12045 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12046 let __s_b = self.gpu.stream();
12047 let mut b = __s_b.launch_builder(&f);
12048 b.arg(q).arg(&gcum).arg(&p).arg(&y).arg(&ssnap).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale);
12049 unsafe { b.launch(cfg)?; }
12050 }
12051 Ok(())
12052 }
12053
12054 #[allow(clippy::too_many_arguments)]
12063 #[allow(clippy::too_many_arguments)]
12064 pub fn gdn_scan_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12065 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
12066 qb16_pre: Option<&CudaSlice<u8>>,
12067 state_in: &CudaSlice<f32>,
12068 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12069 n_head: usize, t: usize, scale: f32, hk: usize)
12070 -> Result<(), Box<dyn std::error::Error>> {
12071 if std::env::var("MEMRA_GDN_DIFF").is_ok() && t >= 16 {
12072 assert!(hk == n_head, "GDN_DIFF oracle is broadcast-only");
12073 return self.gdn_scan_diff(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale);
12074 }
12075 if Self::gdn_chunked_enabled() && t >= 16 {
12076 self.gdn_scan_chunked(q, k, v, g, beta, kb16_pre, qb16_pre, state_in, state_out, o, n_head, t, scale,
12077 Self::gdn_chunk_size(), hk)
12078 } else {
12079 assert!(hk == n_head, "s128 scan is broadcast-only (prep guarantees by predicate)");
12080 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)
12081 }
12082 }
12083
12084 #[allow(clippy::too_many_arguments)]
12086 fn gdn_scan_diff(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12087 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
12088 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12089 n_head: usize, t: usize, scale: f32)
12090 -> Result<(), Box<dyn std::error::Error>> {
12091 static CALL: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
12092 let call = CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12093 let mut o_c = self.uninit(o.len())?;
12094 let mut st_c = self.uninit(state_out.len())?;
12095 self.gdn_scan_chunked(q, k, v, g, beta, None, None, state_in, &mut st_c, &mut o_c,
12096 n_head, t, scale, Self::gdn_chunk_size(), n_head)?;
12097 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)?;
12098 let (oh_s, oh_c) = (self.dtoh(o)?, self.dtoh(&o_c)?);
12099 let (sh_s, sh_c) = (self.dtoh(state_out)?, self.dtoh(&st_c)?);
12100 let stats = |a: &[f32], b: &[f32]| -> (f32, f32, f64) {
12101 let mut max_abs = 0f32; let mut max_rel = 0f32; let mut sum_rel = 0f64;
12102 for (x, y) in a.iter().zip(b) {
12103 let ad = (x - y).abs();
12104 let rel = ad / x.abs().max(y.abs()).max(1e-3);
12105 if ad > max_abs { max_abs = ad; }
12106 if rel > max_rel { max_rel = rel; }
12107 sum_rel += rel as f64;
12108 }
12109 (max_abs, max_rel, sum_rel / a.len() as f64)
12110 };
12111 let (o_ma, o_mr, o_mean) = stats(&oh_s, &oh_c);
12112 let (s_ma, s_mr, s_mean) = stats(&sh_s, &sh_c);
12113 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} | \
12114 state: max_abs={s_ma:.3e} max_rel={s_mr:.3e} mean_rel={s_mean:.3e}",
12115 Self::gdn_chunk_size());
12116 Ok(())
12117 }
12118
12119 pub fn gdn_glog(&self, alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
12121 g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
12122 -> Result<(), Box<dyn std::error::Error>> {
12123 let f = self.func("gdn_glog_f32");
12124 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
12125 let (h, ti) = (n_head as i32, t as i32);
12126 let __s_b = self.gpu.stream();
12127 let mut b = __s_b.launch_builder(&f);
12128 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
12129 unsafe { b.launch(cfg)?; }
12130 Ok(())
12131 }
12132
12133 pub fn sigmoid_v(&self, x: &cudarc::driver::CudaView<f32>, y: &mut CudaSlice<f32>, n: usize)
12136 -> Result<(), Box<dyn std::error::Error>> {
12137 let f = self.func("sigmoid_f32");
12138 let cfg = LaunchConfig::for_num_elems(n as u32);
12139 let ni = n as i32;
12140 let __s_b = self.gpu.stream();
12141 let mut b = __s_b.launch_builder(&f);
12142 b.arg(x).arg(y).arg(&ni);
12143 unsafe { b.launch(cfg)?; }
12144 Ok(())
12145 }
12146
12147 pub fn gdn_glog_v(&self, alpha: &cudarc::driver::CudaView<f32>, dt_bias: &CudaSlice<f32>,
12148 a: &CudaSlice<f32>, g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
12149 -> Result<(), Box<dyn std::error::Error>> {
12150 let f = self.func("gdn_glog_f32");
12151 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
12152 let (h, ti) = (n_head as i32, t as i32);
12153 let __s_b = self.gpu.stream();
12154 let mut b = __s_b.launch_builder(&f);
12155 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
12156 unsafe { b.launch(cfg)?; }
12157 Ok(())
12158 }
12159
12160 pub fn sigmoid(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize)
12161 -> Result<(), Box<dyn std::error::Error>> {
12162 let f = self.func("sigmoid_f32");
12163 let cfg = LaunchConfig::for_num_elems(n as u32);
12164 let ni = n as i32;
12165 let __s_b = self.gpu.stream();
12166 let mut b = __s_b.launch_builder(&f);
12167 b.arg(x).arg(y).arg(&ni);
12168 unsafe { b.launch(cfg)?; }
12169 Ok(())
12170 }
12171
12172 pub fn sig_mul_f16out(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12175 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
12176 -> Result<(), Box<dyn std::error::Error>> {
12177 let f = self.func("sig_mul_f16out_f32");
12178 let cfg = LaunchConfig::for_num_elems(n as u32);
12179 let ni = n as i32;
12180 let __s_b = self.gpu.stream();
12181 let mut b = __s_b.launch_builder(&f);
12182 b.arg(a).arg(g).arg(dst).arg(dst16).arg(&ni);
12183 unsafe { b.launch(cfg)?; }
12184 Ok(())
12185 }
12186
12187 #[allow(clippy::too_many_arguments)]
12196 pub fn attn_head_gate(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12197 dst: &mut CudaSlice<f32>, dst16: Option<&mut CudaSlice<u8>>,
12198 head_dim: usize, n_head: usize, t: usize)
12199 -> Result<(), Box<dyn std::error::Error>> {
12200 let f = self.func("attn_head_gate_f32");
12201 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12202 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12203 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
12205 let __s_b = self.gpu.stream();
12206 let mut b = __s_b.launch_builder(&f);
12207 b.arg(a).arg(g).arg(dst).arg(&d16).arg(&hd).arg(&nh).arg(&ti);
12208 unsafe { b.launch(cfg)?; }
12209 Ok(())
12210 }
12211
12212 #[allow(clippy::too_many_arguments)]
12221 pub fn swiglu_clamped_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
12222 gs: f32, us: f32, limit: f32,
12223 dst: &mut CudaSlice<f32>, n: usize)
12224 -> Result<(), Box<dyn std::error::Error>> {
12225 debug_assert!(limit > 1e-6, "swiglu_clamped needs a live limit; use silu_mul_scaled");
12226 let f = self.func("swiglu_clamped_mul_scaled_f32");
12227 let cfg = LaunchConfig::for_num_elems(n as u32);
12228 let ni = n as i32;
12229 let __s_b = self.gpu.stream();
12230 let mut b = __s_b.launch_builder(&f);
12231 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&limit).arg(dst).arg(&ni);
12232 unsafe { b.launch(cfg)?; }
12233 Ok(())
12234 }
12235
12236 pub fn gated_rmsnorm(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12238 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12239 -> Result<(), Box<dyn std::error::Error>> {
12240 let f = self.func("gated_rmsnorm_f32");
12241 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12242 let (nc, e) = (ncols as i32, eps);
12243 let __s_b = self.gpu.stream();
12244 let mut b = __s_b.launch_builder(&f);
12245 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12246 unsafe { b.launch(cfg)?; }
12247 Ok(())
12248 }
12249
12250 pub fn gated_rmsnorm_f16out(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12253 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12254 ncols: usize, nrows: usize, eps: f32)
12255 -> Result<(), Box<dyn std::error::Error>> {
12256 let f = self.func("gated_rmsnorm_f16out_f32");
12257 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12259 let (nc, e) = (ncols as i32, eps);
12260 let __s_b = self.gpu.stream();
12261 let mut b = __s_b.launch_builder(&f);
12262 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12263 unsafe { b.launch(cfg)?; }
12264 Ok(())
12265 }
12266
12267 #[allow(clippy::too_many_arguments)]
12271 pub fn add_rms_norm_zq8(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
12272 res: &mut CudaSlice<f32>, z: &mut CudaSlice<f32>,
12273 ncols: usize, nrows: usize, eps: f32)
12274 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12275 assert!(ncols % 32 == 0);
12276 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
12277 let mut d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12278 let f = self.func("add_rms_norm_zq8");
12279 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
12280 let (nc, ep) = (ncols as i32, eps);
12281 let __s_b = self.gpu.stream();
12282 let mut b = __s_b.launch_builder(&f);
12283 b.arg(a).arg(b_in).arg(w).arg(res).arg(z).arg(&mut q).arg(&mut d).arg(&nc).arg(&ep);
12284 unsafe { b.launch(cfg)?; }
12285 Ok((q, d))
12286 }
12287
12288 pub fn gated_rmsnorm_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12293 z: &cudarc::driver::CudaView<f32>,
12294 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12295 -> Result<(), Box<dyn std::error::Error>> {
12296 let f = self.func("gated_rmsnorm_f32");
12297 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12298 let (nc, e) = (ncols as i32, eps);
12299 let __s_b = self.gpu.stream();
12300 let mut b = __s_b.launch_builder(&f);
12301 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12302 unsafe { b.launch(cfg)?; }
12303 Ok(())
12304 }
12305
12306 pub fn gated_rmsnorm_f16out_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12307 z: &cudarc::driver::CudaView<f32>,
12308 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12309 ncols: usize, nrows: usize, eps: f32)
12310 -> Result<(), Box<dyn std::error::Error>> {
12311 let f = self.func("gated_rmsnorm_f16out_f32");
12312 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12314 let (nc, e) = (ncols as i32, eps);
12315 let __s_b = self.gpu.stream();
12316 let mut b = __s_b.launch_builder(&f);
12317 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12318 unsafe { b.launch(cfg)?; }
12319 Ok(())
12320 }
12321
12322 pub fn gated_rmsnorm_q8_1(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12323 ncols: usize, nrows: usize, eps: f32)
12324 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12325 assert!(ncols % 32 == 0);
12326 let f = self.func("gated_rmsnorm_q8_1");
12327 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12328 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12329 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12330 let (nc, ep) = (ncols as i32, eps);
12331 let __s_b = self.gpu.stream();
12332 let mut b = __s_b.launch_builder(&f);
12333 b.arg(o).arg(w).arg(z).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
12334 unsafe { b.launch(cfg)?; }
12335 Ok((out_q, out_d))
12336 }
12337
12338 pub fn transpose(&self, inp: &CudaSlice<f32>, rows: usize, cols: usize)
12340 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12341 let f = self.func("transpose_f32");
12342 let mut out = self.zeros(rows * cols)?;
12343 let cfg = LaunchConfig::for_num_elems((rows * cols) as u32);
12344 let (r, c) = (rows as i32, cols as i32);
12345 let __s_b = self.gpu.stream();
12346 let mut b = __s_b.launch_builder(&f);
12347 b.arg(inp).arg(&mut out).arg(&r).arg(&c);
12348 unsafe { b.launch(cfg)?; }
12349 Ok(out)
12350 }
12351
12352 pub fn repeat_heads(&self, inp: &CudaSlice<f32>, out: &mut CudaSlice<f32>,
12354 head_dim: usize, n_in: usize, n_out: usize, t: usize)
12355 -> Result<(), Box<dyn std::error::Error>> {
12356 let f = self.func("repeat_heads_f32");
12357 let cfg = LaunchConfig::for_num_elems((head_dim * n_out * t) as u32);
12358 let (hd, ni, no, ti) = (head_dim as i32, n_in as i32, n_out as i32, t as i32);
12359 let __s_b = self.gpu.stream();
12360 let mut b = __s_b.launch_builder(&f);
12361 b.arg(inp).arg(out).arg(&hd).arg(&ni).arg(&no).arg(&ti);
12362 unsafe { b.launch(cfg)?; }
12363 Ok(())
12364 }
12365
12366 pub fn q_gate_split(&self, qf: &CudaSlice<f32>, q_out: &mut CudaSlice<f32>,
12369 gate_out: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, t: usize)
12370 -> Result<(), Box<dyn std::error::Error>> {
12371 let f = self.func("q_gate_split_f32");
12372 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12373 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12374 let __s_b = self.gpu.stream();
12375 let mut b = __s_b.launch_builder(&f);
12376 b.arg(qf).arg(q_out).arg(gate_out).arg(&hd).arg(&nh).arg(&ti);
12377 unsafe { b.launch(cfg)?; }
12378 Ok(())
12379 }
12380
12381 pub fn qkv_to_gdn_repack(&self, conv_out: &CudaSlice<f32>, q_g: &mut CudaSlice<f32>,
12385 k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
12386 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, t: usize)
12387 -> Result<(), Box<dyn std::error::Error>> {
12388 let f = self.func("qkv_to_gdn_repack_f32");
12389 let cfg = LaunchConfig::for_num_elems((d_state * num_v * t) as u32);
12390 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);
12391 let __s_b = self.gpu.stream();
12392 let mut b = __s_b.launch_builder(&f);
12393 b.arg(conv_out).arg(q_g).arg(k_g).arg(v_g).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&ti);
12394 unsafe { b.launch(cfg)?; }
12395 Ok(())
12396 }
12397
12398 pub fn conv_left_pad(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
12401 conv_dim: usize, t: usize, pad: usize)
12402 -> Result<(), Box<dyn std::error::Error>> {
12403 let f = self.func("conv_left_pad_f32");
12404 let cfg = LaunchConfig::for_num_elems((conv_dim * t) as u32);
12405 let (cd, ti, p) = (conv_dim as i32, t as i32, pad as i32);
12406 let __s_b = self.gpu.stream();
12407 let mut b = __s_b.launch_builder(&f);
12408 b.arg(src).arg(dst).arg(&cd).arg(&ti).arg(&p);
12409 unsafe { b.launch(cfg)?; }
12410 Ok(())
12411 }
12412
12413 pub fn conv_assemble_and_roll(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12417 conv_in: &mut CudaSlice<f32>, conv_dim: usize, pad: usize)
12418 -> Result<(), Box<dyn std::error::Error>> {
12419 let f = self.func("conv_assemble_and_roll_f32");
12420 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12421 let (cd, p) = (conv_dim as i32, pad as i32);
12422 let __s_b = self.gpu.stream();
12423 let mut b = __s_b.launch_builder(&f);
12424 b.arg(qkv_col).arg(conv_state).arg(conv_in).arg(&cd).arg(&p);
12425 unsafe { b.launch(cfg)?; }
12426 Ok(())
12427 }
12428
12429 pub fn ssm_conv1d_fused_decode(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12435 w: &CudaSlice<f32>, conv_out: &mut CudaSlice<f32>,
12436 conv_dim: usize, d_conv: usize)
12437 -> Result<(), Box<dyn std::error::Error>> {
12438 let f = self.func("ssm_conv1d_fused_decode_f32");
12439 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12440 let (cd, dc) = (conv_dim as i32, d_conv as i32);
12441 let __s_b = self.gpu.stream();
12442 let mut b = __s_b.launch_builder(&f);
12443 b.arg(qkv_col).arg(conv_state).arg(w).arg(conv_out).arg(&cd).arg(&dc);
12444 unsafe { b.launch(cfg)?; }
12445 Ok(())
12446 }
12447
12448 pub fn slice_range(&self, src: &CudaSlice<f32>, start: usize, len: usize)
12451 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12452 let host = self.gpu.stream().clone_dtoh(src)?;
12453 self.gpu.stream().synchronize()?;
12454 Ok(self.htod(&host[start..start + len])?)
12455 }
12456}
12457
12458#[cfg(test)]
12459mod target_dispatch_tests {
12460 use super::legacy_quant_gemm_allowed;
12461
12462 #[test]
12463 fn legacy_quant_gemm_arch_policy_honors_the_escape_hatch() {
12464 assert!(legacy_quant_gemm_allowed(false, false, false));
12466 assert!(!legacy_quant_gemm_allowed(false, false, true));
12467 assert!(!legacy_quant_gemm_allowed(true, false, false));
12469 assert!(!legacy_quant_gemm_allowed(true, false, true));
12470 assert!(legacy_quant_gemm_allowed(true, true, false));
12472 assert!(!legacy_quant_gemm_allowed(true, true, true));
12473 }
12474
12475 #[cfg(all(memra_portable_cuda, not(memra_hopper_mma)))]
12476 #[test]
12477 fn portable_build_disables_legacy_quant_gemm_without_an_env_override() {
12478 assert!(!legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12479 }
12480
12481 #[cfg(memra_hopper_mma)]
12482 #[test]
12483 fn hopper_mma_build_re_admits_legacy_quant_gemm() {
12484 assert!(legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12485 assert!(super::portable_mma_gated() == false);
12486 }
12487}
12488
12489impl memra_kv::KvDev for Engine {
12492 fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12493 Engine::zeros(self, n)
12494 }
12495 fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12496 Engine::uninit(self, n)
12497 }
12498 fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
12499 Engine::alloc_u8(self, n)
12500 }
12501 fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
12502 Engine::htod_i32(self, v)
12503 }
12504 fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12505 Engine::clone_dtod(self, src)
12506 }
12507 fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
12508 -> Result<(), Box<dyn std::error::Error>> {
12509 Engine::copy_into(self, dst, off, src, len)
12510 }
12511 fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
12512 Engine::set_i32_one(self, d, v)
12513 }
12514}