1use std::sync::{Arc, Mutex};
4use cudarc::driver::{CudaContext, CudaStream, CudaModule, CudaFunction, CudaSlice, LaunchConfig, PushKernelArg};
5use cudarc::nvrtc::Ptx;
6
7pub use memra_gguf;
8pub use memra_runtime;
9
10pub mod model;
11pub mod forward;
12pub mod hybrid;
13pub mod hybrid_forward;
14pub mod cache {
17 pub use memra_kv::*;
18}
19pub mod decode;
20pub mod decode_batch;
21pub mod mla;
25pub mod pp;
26pub mod spec;
27pub mod gemma_spec;
28pub mod round_stream;
29pub mod graph_update;
30pub mod dflash;
31pub mod eagle;
32pub use memra_sampling as sampler;
33
34pub fn moe_f16g_mode() -> u8 {
78 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
79 *M.get_or_init(|| match std::env::var("MEMRA_MOE_F16G").as_deref() {
80 Ok("0") => 0,
81 Ok("2") => 2,
82 Ok("3") => 3,
83 Ok(_) => 1,
84 Err(_) => 2,
87 })
88}
89pub fn moe_f16g_sk_params() -> (i32, i32) {
103 static P: std::sync::OnceLock<(i32, i32)> = std::sync::OnceLock::new();
104 *P.get_or_init(|| match std::env::var("MEMRA_F16G_SK").as_deref() {
105 Ok("0") => (-1, 0),
106 Ok("32") => (0, i32::MAX),
107 Ok("128") => (0, 1),
108 _ => {
109 let cross = std::env::var("MEMRA_F16G_SK_CROSS").ok()
110 .and_then(|v| v.parse().ok()).unwrap_or(64);
111 (0, cross)
112 }
113 })
114}
115pub fn moe_f16g_direct_on(qtype: i32) -> bool {
126 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
127 let m = *M.get_or_init(|| match std::env::var("MEMRA_F16G_DIRECT").as_deref() {
128 Ok("0") => 0,
129 Ok("kq") => 1,
130 _ => 2,
131 });
132 match m {
133 0 => false,
134 1 => qtype == QT_Q4_K || qtype == QT_Q6_K,
135 _ => true,
136 }
137}
138pub fn moe_f16g_tail_on() -> bool {
147 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
148 *ON.get_or_init(|| std::env::var("MEMRA_F16G_TAIL").as_deref() != Ok("0"))
149}
150
151pub fn moe_f16g_gemma_on() -> bool {
158 static M: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
159 *M.get_or_init(|| !matches!(std::env::var("MEMRA_MOE_F16G").as_deref(), Ok("0") | Err(_)))
160}
161
162pub fn moe_fuse_actq_on() -> bool {
166 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
167 *ON.get_or_init(|| std::env::var("MEMRA_MOE_FUSE_ACTQ").as_deref() != Ok("0"))
168}
169
170pub fn router_prefill_exact_on() -> bool {
180 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
181 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_PREFILL_EXACT").as_deref() != Ok("0"))
182}
183
184pub fn router_kernel_on() -> bool {
185 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
186 *ON.get_or_init(|| {
187 let on = std::env::var("MEMRA_ROUTER_KERNEL").as_deref() != Ok("0");
188 if !on { eprintln!("[memra] router kernel OFF (rollback: per-column cuBLAS gemv)"); }
189 on
190 })
191}
192
193pub const ROUTER_BATCH_MIN_T: usize = 8;
208pub fn router_batch_on() -> bool {
209 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
210 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_BATCH").as_deref() != Ok("0"))
211}
212mod cpu_experts;
213pub mod moe_cache;
214pub mod spill;
215mod spill_pread;
216#[cfg(memra_cutlass)]
217pub mod cutlass_ffi;
218pub mod mmq_ffi;
219pub mod f16_ffi;
220pub mod prime_graph;
221pub mod fp8_ffi;
222
223const FATBIN: &[u8] = include_bytes!(env!("MEMRA_ENGINE_FATBIN"));
230const HYBRID_FATBIN: &[u8] = include_bytes!(env!("MEMRA_HYBRID_FATBIN"));
231const QMATVEC_FATBIN: &[u8] = include_bytes!(env!("MEMRA_QMATVEC_FATBIN"));
232const FLASH_FATBIN: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN"));
233const GEMM_FATBIN: &[u8] = include_bytes!(env!("MEMRA_GEMM_FATBIN"));
234const ROUTER_FATBIN: &[u8] = include_bytes!(env!("MEMRA_ROUTER_FATBIN"));
235const SAMPLE_FATBIN: &[u8] = include_bytes!(env!("MEMRA_SAMPLE_FATBIN"));
237
238fn gemm_fatbin_bytes() -> std::borrow::Cow<'static, [u8]> {
244 assert!(!(portable_mma_gated() && std::env::var_os("MEMRA_GEMM_FATBIN").is_some()),
245 "MEMRA_GEMM_FATBIN overrides are not allowed in the portable CUDA lane");
246 match std::env::var("MEMRA_GEMM_FATBIN") {
247 Ok(path) => std::borrow::Cow::Owned(
248 std::fs::read(&path).unwrap_or_else(|e| panic!("MEMRA_GEMM_FATBIN read {path}: {e}"))),
249 Err(_) => std::borrow::Cow::Borrowed(GEMM_FATBIN),
250 }
251}
252
253pub(crate) const fn portable_mma_gated() -> bool {
260 cfg!(memra_portable_cuda) && !cfg!(memra_hopper_mma)
261}
262
263const fn legacy_quant_gemm_allowed(portable_cuda: bool, hopper_mma: bool, no_gemm: bool) -> bool {
268 (!portable_cuda || hopper_mma) && !no_gemm
269}
270
271const FLASH_FATBIN_VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VQ4"));
279const FLASH_FATBIN_VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VF8"));
280const FLASH_FATBIN_KF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8"));
281const FLASH_FATBIN_KF8VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VQ4"));
282const FLASH_FATBIN_KF8VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VF8"));
283
284pub use memra_kv::{kv_blk_bytes, kv_cache_formats};
287
288fn flash_fatbin_bytes() -> &'static [u8] {
290 match kv_cache_formats() {
291 ("q8_0", "q5_1") => FLASH_FATBIN,
292 ("q8_0", "q4_0") => FLASH_FATBIN_VQ4,
293 ("q8_0", "fp8") => FLASH_FATBIN_VF8,
294 ("fp8", "q5_1") => FLASH_FATBIN_KF8,
295 ("fp8", "q4_0") => FLASH_FATBIN_KF8VQ4,
296 ("fp8", "fp8") => FLASH_FATBIN_KF8VF8,
297 other => unreachable!("kv_cache_formats returned {other:?}"),
298 }
299}
300
301fn k1_launch_override() -> Option<(u32, u32, u32)> {
308 static K1: std::sync::OnceLock<Option<(u32, u32, u32)>> = std::sync::OnceLock::new();
309 *K1.get_or_init(|| {
310 let v = std::env::var("MEMRA_GEMM_K1_LAUNCH").ok()?;
311 let p: Vec<u32> = v.split(',').filter_map(|s| s.trim().parse().ok()).collect();
312 match p.as_slice() { [bm, bn, w] => Some((*bm, *bn, *w)), _ => None }
313 })
314}
315
316pub(crate) fn wgmma_gemm_enabled() -> bool {
323 static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
324 *V.get_or_init(|| std::env::var("MEMRA_WGMMA").as_deref() == Ok("1"))
325}
326
327pub const FA_VEC_MIN_TKV: usize = 96;
342pub fn fa_vec_min_tkv() -> usize {
346 static V: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
347 *V.get_or_init(|| std::env::var("MEMRA_FA_VEC_MIN").ok()
348 .and_then(|v| v.parse().ok())
349 .unwrap_or_else(|| FA_VEC_MIN_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)))
350}
351
352pub fn fa_f16pv_on() -> bool {
363 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
364 *ON.get_or_init(|| std::env::var("MEMRA_FA_F16PV").map(|v| v != "0")
365 .unwrap_or_else(|_| std::env::var("MEMRA_DRAFT").is_err()))
366}
367
368pub fn fa512_hp_on() -> bool {
372 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
373 *ON.get_or_init(|| std::env::var("MEMRA_FA512_HP").as_deref() != Ok("0"))
374}
375
376pub fn faw_hp_on() -> bool {
380 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
381 *ON.get_or_init(|| std::env::var("MEMRA_FAW_HP").as_deref() != Ok("0"))
382}
383
384pub fn fa512_wide_warps() -> usize {
388 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
389 *N.get_or_init(|| match std::env::var("MEMRA_FA512_W4").as_deref() {
390 Ok("1") => 4, _ => 2,
391 })
392}
393
394pub fn fa512_min_tkv() -> usize {
397 static FA512_MIN: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
398 *FA512_MIN.get_or_init(|| std::env::var("MEMRA_FA512_MIN").ok()
399 .and_then(|v| v.parse().ok()).unwrap_or(512))
400}
401pub static FA_VEC_MIN_DEFAULT: std::sync::atomic::AtomicUsize =
405 std::sync::atomic::AtomicUsize::new(FA_VEC_MIN_TKV);
406pub static FA_SPW_DEFAULT: std::sync::atomic::AtomicUsize =
410 std::sync::atomic::AtomicUsize::new(32);
411pub static FUSED_MR1_DEFAULT: std::sync::atomic::AtomicBool =
417 std::sync::atomic::AtomicBool::new(false);
418pub static ROUTER_W8_DEFAULT: std::sync::atomic::AtomicBool =
425 std::sync::atomic::AtomicBool::new(true);
426pub static FA_SP512_DEFAULT: std::sync::atomic::AtomicUsize =
427 std::sync::atomic::AtomicUsize::new(16);
428pub static RMS_BLOCK_DEFAULT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(256);
433pub static FA_SP_GEMMA: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
435pub static MMQ_SK_FORCE: std::sync::atomic::AtomicI8 = std::sync::atomic::AtomicI8::new(-1);
440pub use memra_kv::KV_FP8_FORCE;
443pub(crate) fn rms_block() -> u32 {
444 static V: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
445 *V.get_or_init(|| std::env::var("MEMRA_RMS_BLOCK").ok()
446 .and_then(|v| v.parse().ok())
447 .unwrap_or_else(|| RMS_BLOCK_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)))
448}
449
450pub(crate) fn fa_split_keys(t_kv: usize, n_head_kv: usize) -> usize {
451 static S: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
452 if let Some(forced) = *S.get_or_init(|| {
453 std::env::var("MEMRA_FA_SPLIT").ok().and_then(|v| v.parse().ok())
454 .filter(|&s: &usize| s >= 8 && s % 8 == 0)
455 }) { return forced; }
456 if FA_SP_GEMMA.load(std::sync::atomic::Ordering::Relaxed)
474 && std::env::var("MEMRA_FA_SP16").as_deref() == Ok("1") {
475 return if t_kv <= 8192 { 16 } else if t_kv <= 16384 { 64 } else { 128 };
476 }
477 let big_rig = fa_sm_count() >= 128;
478 if big_rig {
479 let _ = n_head_kv;
480 if t_kv <= 2048 { 16 } else if t_kv <= 16384 { 64 } else { 128 }
481 } else if n_head_kv <= 4 {
482 if t_kv <= 512 { 8 } else if t_kv <= 16384 { 64 } else { 128 }
503 } else {
504 if t_kv <= 8192 { 32 } else if t_kv <= 16384 { 64 } else { 128 }
505 }
506}
507
508fn fa_sm_count() -> i32 {
511 static N: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
512 *N.get_or_init(|| {
513 cudarc::driver::result::init().ok();
514 cudarc::driver::result::device::get(0)
515 .and_then(|d| unsafe { cudarc::driver::result::device::get_attribute(
516 d, cudarc::driver::sys::CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT) })
517 .unwrap_or(82)
518 })
519}
520
521fn fa_hd_suffix(head_dim: usize) -> Result<&'static str, Box<dyn std::error::Error>> {
525 match head_dim {
526 256 => Ok(""),
527 128 => Ok("_hd128"),
528 d => Err(format!("fa_prefill: no kernel stamped for head_dim={d} (only 256/128); \
529 callers must gate to sdpa_naive").into()),
530 }
531}
532
533pub const QT_Q8_0: i32 = 0;
535pub const QT_Q4_K: i32 = 1;
536pub const QT_Q6_K: i32 = 2;
537pub const QT_Q5_K: i32 = 3;
538pub const QT_Q3_K: i32 = 4;
539pub const QT_IQ4_XS: i32 = 5;
540pub const QT_IQ3_S: i32 = 6;
541pub const QT_NVFP4: i32 = 7;
542pub const QT_F8_E4M3: i32 = 10;
548pub const QT_NVFP4_RP: i32 = 9;
551pub const QT_F32: i32 = 8;
553pub const QT_BF16: i32 = 11;
554pub const QT_Q4_0: i32 = 12; pub const QT_Q2_K: i32 = 13;
559pub const QT_F8_E4M3_BLK: i32 = 14;
575
576pub struct Engine {
578 pub gpu: memra_runtime::Gpu,
579 module: Arc<CudaModule>,
580 hybrid: Arc<CudaModule>,
581 qmatvec: Arc<CudaModule>,
582 flash: Arc<CudaModule>,
583 flash_g: std::sync::OnceLock<Arc<CudaModule>>,
587 gemm: Arc<CudaModule>,
588 router: Arc<CudaModule>,
589 sample: Arc<CudaModule>,
591 moe_cache: Mutex<Option<crate::moe_cache::MoeSlotCache>>,
595 moe_cache_layout: Mutex<Option<Vec<usize>>>,
599 capture_keep_on: std::sync::atomic::AtomicBool,
605 verify_exact: std::sync::atomic::AtomicBool,
610 capture_keep: Mutex<Vec<Box<dyn std::any::Any + Send>>>,
611 pub copy_stream: Arc<CudaStream>,
613 #[cfg(memra_cutlass)]
620 cutlass_scratch: Mutex<Option<crate::cutlass_ffi::CutlassScratch>>,
621 fp8_scratch: Mutex<Option<crate::fp8_ffi::Fp8Scratch>>,
625 fa_vf16_scratch: Mutex<Option<CudaSlice<u8>>>,
628 fa_part_pool: Mutex<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
632 fa_part_retired: Mutex<Vec<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
636 fn_cache: Mutex<std::collections::HashMap<String, CudaFunction>>,
638 f16_scratch: Mutex<Option<crate::f16_ffi::F16Scratch>>,
639 argmax_partials: Mutex<Option<(CudaSlice<f32>, CudaSlice<i32>)>>,
644 prime_deqw_ws: Mutex<Option<(CudaSlice<u8>, CudaSlice<u8>)>>,
649 router_stage: Mutex<Option<PinnedStage>>,
653}
654
655fn fa_v2_on() -> bool {
665 std::env::var("MEMRA_FA_V2").map(|v| v != "0").unwrap_or(true)
671}
672
673fn fa_v3_on() -> bool {
681 std::env::var("MEMRA_FA_V3").map(|v| v != "0").unwrap_or(true)
685}
686
687fn fa_v4_mode() -> &'static str {
692 static M: std::sync::OnceLock<String> = std::sync::OnceLock::new();
693 M.get_or_init(|| std::env::var("MEMRA_FA_V4").unwrap_or_default())
694}
695fn fa_v4_on() -> bool { fa_v4_mode() != "0" } pub static FA_SMEM_TKV_DEFAULT: std::sync::atomic::AtomicUsize =
704 std::sync::atomic::AtomicUsize::new(1024);
705pub static FA_V4_MAX_DEFAULT: std::sync::atomic::AtomicUsize =
706 std::sync::atomic::AtomicUsize::new(usize::MAX);
707pub fn fa_v4_at_pub(t_kv: usize) -> bool { fa_v4_at(t_kv) }
708fn fa_v4_at(t_kv: usize) -> bool {
709 static M: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
710 let mx = *M.get_or_init(|| std::env::var("MEMRA_FA_V4_MAX").ok()
711 .and_then(|v| v.parse().ok())
712 .unwrap_or_else(|| FA_V4_MAX_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)));
713 fa_v4_on() && t_kv < mx
714}
715pub const FA_DEEP_MIN_DEFAULT: usize = 0;
729fn fa_deep_at(t_kv: usize) -> bool {
730 if std::env::var("MEMRA_FA_DEEP").as_deref() == Ok("0") { return false; }
731 let min = std::env::var("MEMRA_FA_DEEP_MIN").ok().and_then(|v| v.parse().ok())
732 .unwrap_or(FA_DEEP_MIN_DEFAULT);
733 t_kv >= min
734}
735pub fn fa_deep_at_pub(t_kv: usize) -> bool { fa_deep_at(t_kv) }
737
738fn fa_v3_active(head_dim: usize) -> bool {
739 fa_v3_on() && head_dim % 128 == 0 && kv_cache_formats() == ("q8_0", "q5_1")
742 && !Engine::kv_fp8_on()
743}
744
745pub fn fa_seqs_eligible(t_kv: usize, head_dim: usize) -> bool {
753 std::env::var("MEMRA_NO_FA_VEC").is_err()
754 && t_kv >= fa_vec_min_tkv()
755 && head_dim == 256
756 && fa_v4_at(t_kv)
757 && !matches!(fa_v4_mode(), "noB3" | "stage")
758 && !Engine::kv_fp8_on()
759}
760pub fn fa_split_keys_pub(t_kv: usize, n_head_kv: usize) -> usize { fa_split_keys(t_kv, n_head_kv) }
762
763struct PinnedStage {
768 ptr: *mut u8,
769 cap: usize,
770}
771unsafe impl Send for PinnedStage {}
772impl PinnedStage {
773 fn new(cap: usize) -> Result<Self, Box<dyn std::error::Error>> {
774 let ptr = unsafe { cudarc::driver::result::malloc_host(cap, 0)? } as *mut u8;
775 Ok(PinnedStage { ptr, cap })
776 }
777}
778impl Drop for PinnedStage {
779 fn drop(&mut self) {
780 let _ = unsafe { cudarc::driver::result::free_host(self.ptr as _) };
781 }
782}
783
784pub const ARGMAX_NB: usize = 256;
787
788pub(crate) use memra_fa3_vl as fa3_vl_raw;
790
791unsafe extern "C" {
792 fn memra_fa3_prefill(q16: *const core::ffi::c_void, k16: *const core::ffi::c_void,
794 v16: *const core::ffi::c_void, o: *mut f32,
795 t: i32, h: i32, hkv: i32, d: i32, scale: f32,
796 stream: *mut core::ffi::c_void) -> i32;
797 pub(crate) fn memra_fa3_vl(q16s: *const *const core::ffi::c_void, k16s: *const *const core::ffi::c_void,
799 v16s: *const *const core::ffi::c_void, os: *const *mut f32,
800 ts: *const i32, b: i32, h: i32, hkv: i32, d: i32, scale: f32,
801 stream: *mut core::ffi::c_void) -> i32;
802}
803
804#[repr(C)]
809#[derive(Clone, Copy)]
810pub struct WPtr8(pub [u64; 8]);
811unsafe impl cudarc::driver::DeviceRepr for WPtr8 {}
812
813#[repr(C)]
818#[derive(Clone, Copy, Default)]
819pub struct GdnSeqVl {
820 pub kb16: u64, pub gcum: u64, pub beta: u64, pub u: u64, pub wb16: u64,
821 pub y: u64, pub ssnap: u64, pub state_in: u64, pub state_out: u64,
822 pub q: u64, pub p: u64, pub o: u64,
823 pub k: u64, pub v: u64, pub g: u64, pub a: u64, pub w: u64,
824 pub t: i32, pub nc: i32,
825}
826unsafe impl cudarc::driver::DeviceRepr for GdnSeqVl {}
827#[repr(C)]
828#[derive(Clone, Copy)]
829pub struct GdnVl8(pub [GdnSeqVl; 8]);
830unsafe impl cudarc::driver::DeviceRepr for GdnVl8 {}
831
832#[repr(C)]
835#[derive(Clone, Copy, Default)]
836pub struct GdnWVl { pub qb16: u64, pub pb16: u64 }
837unsafe impl cudarc::driver::DeviceRepr for GdnWVl {}
838#[repr(C)]
839#[derive(Clone, Copy)]
840pub struct GdnWVl8(pub [GdnWVl; 8]);
841unsafe impl cudarc::driver::DeviceRepr for GdnWVl8 {}
842
843#[repr(C)]
845#[derive(Clone, Copy, Default)]
846pub struct GdnPrepVl {
847 pub qkv: u64, pub conv_state: u64, pub conv_out: u64,
848 pub q_g: u64, pub k_g: u64, pub v_g: u64,
849 pub q_l2: u64, pub k_l2: u64,
850 pub beta_raw: u64, pub alpha: u64, pub beta: u64, pub g_log: u64,
851 pub o: u64, pub z: u64, pub gn: u64, pub gn16: u64,
852 pub kb16: u64,
853 pub qb16: u64,
854 pub t: i32, pub pad: i32,
855}
856unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl {}
857#[repr(C)]
858#[derive(Clone, Copy)]
859pub struct GdnPrepVl8(pub [GdnPrepVl; 8]);
860unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl8 {}
861
862#[repr(C)]
864#[derive(Clone, Copy, Default)]
865pub struct FaSeqVl {
866 pub q: u64, pub k16: u64, pub v16: u64, pub o: u64, pub kf: u64, pub vf: u64,
867 pub t: i32, pub pad: i32,
868}
869unsafe impl cudarc::driver::DeviceRepr for FaSeqVl {}
870#[repr(C)]
871#[derive(Clone, Copy)]
872pub struct FaVl8(pub [FaSeqVl; 8]);
873unsafe impl cudarc::driver::DeviceRepr for FaVl8 {}
874
875#[repr(C)]
877#[derive(Clone, Copy, Default)]
878pub struct AttnPreVl {
879 pub qf: u64, pub kf: u64, pub vf: u64,
880 pub q: u64, pub gate: u64, pub qn: u64, pub kn: u64,
881 pub kc: u64, pub vc: u64,
882 pub t: i32, pub pad: i32,
883}
884unsafe impl cudarc::driver::DeviceRepr for AttnPreVl {}
885#[repr(C)]
886#[derive(Clone, Copy)]
887pub struct AttnPreVl8(pub [AttnPreVl; 8]);
888unsafe impl cudarc::driver::DeviceRepr for AttnPreVl8 {}
889
890pub struct GdnChunkBufs {
893 pub gcum: CudaSlice<f32>,
894 pub a: CudaSlice<f32>,
895 pub p: CudaSlice<f32>,
896 pub u: CudaSlice<f32>,
897 pub w: CudaSlice<f32>,
898 pub kb16: CudaSlice<u8>,
899 pub wb16: CudaSlice<u8>,
900 pub y16: CudaSlice<u8>,
901 pub ssnap16: CudaSlice<u8>,
902 pub qb16: CudaSlice<u8>,
903 pub pb16: CudaSlice<u8>,
904 pub o: CudaSlice<f32>,
905 pub t: usize,
906 pub nc: usize,
907}
908
909#[repr(C)]
911#[derive(Clone, Copy)]
912pub struct F32x8(pub [f32; 8]);
913unsafe impl cudarc::driver::DeviceRepr for F32x8 {}
914
915pub static PRIME_NANOS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
919
920impl Engine {
921 pub fn new(ordinal: usize) -> Result<Self, Box<dyn std::error::Error>> {
922 let gpu = memra_runtime::Gpu::new(ordinal)?;
923 if std::env::var("MEMRA_ARCH_CHECK").as_deref() != Ok("0") {
927 use cudarc::driver::sys::CUdevice_attribute_enum as A;
928 let (maj, min) = cudarc::driver::result::device::get(ordinal as i32)
929 .and_then(|d| unsafe { Ok((
930 cudarc::driver::result::device::get_attribute(d, A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR)?,
931 cudarc::driver::result::device::get_attribute(d, A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR)?)) })
932 .unwrap_or((0, 0));
933 let built = env!("MEMRA_BUILT_CUDA_ARCH");
934 let ok = matches!((built, maj, min),
935 ("120a", 12, 0) | ("120a", 12, 1) | ("100a", 10, 0) | ("90a", 9, 0) | ("89", 8, 9));
936 if !ok {
937 return Err(format!(
938 "memra was built for sm_{built} but device {ordinal} reports compute \
939 capability {maj}.{min}. Rebuild on this machine (MEMRA_CUDA_ARCH \
940 auto-detects the GPU) or set MEMRA_ARCH_CHECK=0 to bypass.").into());
941 }
942 }
943 unsafe {
948 use cudarc::driver::sys;
949 let dev: sys::CUdevice = ordinal as sys::CUdevice;
950 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
951 if sys::cuDeviceGetDefaultMemPool(&mut pool, dev) == sys::CUresult::CUDA_SUCCESS {
952 let mut thresh: u64 = u64::MAX;
953 let _ = sys::cuMemPoolSetAttribute(
954 pool,
955 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
956 &mut thresh as *mut u64 as *mut core::ffi::c_void,
957 );
958 }
959 }
960 let module = gpu.ctx.load_module(Ptx::from_binary(FATBIN.to_vec()))?;
961 let hybrid = gpu.ctx.load_module(Ptx::from_binary(HYBRID_FATBIN.to_vec()))?;
962 let qmatvec = gpu.ctx.load_module(Ptx::from_binary(QMATVEC_FATBIN.to_vec()))?;
963 let flash = gpu.ctx.load_module(Ptx::from_binary(flash_fatbin_bytes().to_vec()))?;
964 let gemm = gpu.ctx.load_module(Ptx::from_binary(gemm_fatbin_bytes().into_owned()))?;
965 let router = gpu.ctx.load_module(Ptx::from_binary(ROUTER_FATBIN.to_vec()))?;
966 let sample = gpu.ctx.load_module(Ptx::from_binary(SAMPLE_FATBIN.to_vec()))?;
967 let copy_stream = gpu.ctx.new_stream()?;
968 if std::env::var("MEMRA_EVT").map(|v| v == "1").unwrap_or(false) {
984 } else {
986 unsafe { gpu.ctx.disable_event_tracking(); }
987 }
988 Ok(Self { gpu, module, hybrid, qmatvec, flash, flash_g: std::sync::OnceLock::new(), gemm, router, sample,
989 moe_cache: Mutex::new(None),
990 moe_cache_layout: Mutex::new(None),
991 copy_stream,
992 capture_keep_on: std::sync::atomic::AtomicBool::new(false),
993 verify_exact: std::sync::atomic::AtomicBool::new(false),
994 capture_keep: Mutex::new(Vec::new()),
995 argmax_partials: Mutex::new(None),
996 prime_deqw_ws: Mutex::new(None),
997 router_stage: Mutex::new(None),
998 fp8_scratch: Mutex::new(None),
999 fa_vf16_scratch: Mutex::new(None),
1000 fa_part_pool: Mutex::new(None),
1001 fa_part_retired: Mutex::new(Vec::new()),
1002 fn_cache: Mutex::new(Default::default()),
1003 f16_scratch: Mutex::new(None),
1004 #[cfg(memra_cutlass)]
1005 cutlass_scratch: Mutex::new(None) })
1006 }
1007
1008 pub fn ctx(&self) -> &Arc<CudaContext> { &self.gpu.ctx }
1009
1010 pub fn pool_cached_bytes(&self) -> usize {
1028 let (reserved, used) = self.pool_reserved_used();
1029 reserved.saturating_sub(used)
1030 }
1031
1032 pub fn pool_reserved_used(&self) -> (usize, usize) {
1039 use cudarc::driver::sys;
1040 unsafe {
1041 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
1042 if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
1043 != sys::CUresult::CUDA_SUCCESS
1044 {
1045 return (0, 0);
1046 }
1047 let (mut reserved, mut used) = (0u64, 0u64);
1048 if sys::cuMemPoolGetAttribute(
1049 pool,
1050 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RESERVED_MEM_CURRENT,
1051 &mut reserved as *mut u64 as *mut core::ffi::c_void,
1052 ) != sys::CUresult::CUDA_SUCCESS {
1053 return (0, 0);
1054 }
1055 if sys::cuMemPoolGetAttribute(
1056 pool,
1057 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_USED_MEM_CURRENT,
1058 &mut used as *mut u64 as *mut core::ffi::c_void,
1059 ) != sys::CUresult::CUDA_SUCCESS {
1060 return (0, 0);
1061 }
1062 (reserved as usize, used as usize)
1063 }
1064 }
1065
1066 pub fn stream(&self) -> Arc<CudaStream> { self.gpu.stream() }
1069 pub fn gkv_on() -> bool {
1072 memra_kv::gkv_on()
1073 }
1074
1075 pub fn wkv_on() -> bool {
1087 memra_kv::wkv_on()
1088 }
1089
1090 pub fn kv_fp8_on() -> bool {
1096 memra_kv::kv_fp8_on()
1097 }
1098
1099 fn fa_func(&self, name: &str, head_dim: usize) -> CudaFunction {
1102 if head_dim == 512 && Self::gkv_on() { self.func_g(name) } else { self.func(name) }
1103 }
1104
1105 fn func_g(&self, name: &str) -> CudaFunction {
1109 let m = self.flash_g.get_or_init(|| {
1110 self.gpu.ctx.load_module(cudarc::nvrtc::Ptx::from_binary(FLASH_FATBIN_KF8VF8.to_vec()))
1111 .expect("load kf8vf8 flash fatbin (fp8-globals arm)")
1112 });
1113 let key = format!("g:{name}");
1114 if let Some(f) = self.fn_cache.lock().unwrap().get(&key) { return f.clone(); }
1115 let f = match m.load_function(name) {
1116 Ok(f) => f,
1117 Err(_) => self.func(name),
1118 };
1119 self.fn_cache.lock().unwrap().insert(key, f.clone());
1120 f
1121 }
1122
1123 fn func(&self, name: &str) -> CudaFunction {
1124 if let Some(f) = self.fn_cache.lock().unwrap().get(name) { return f.clone(); }
1127 let f = self.module.load_function(name)
1128 .or_else(|_| self.hybrid.load_function(name))
1129 .or_else(|_| self.qmatvec.load_function(name))
1130 .or_else(|_| self.flash.load_function(name))
1131 .or_else(|_| self.gemm.load_function(name))
1132 .or_else(|_| self.router.load_function(name))
1133 .or_else(|_| self.sample.load_function(name))
1134 .unwrap_or_else(|_| panic!("kernel {name} not in any fatbin"));
1135 self.fn_cache.lock().unwrap().insert(name.to_string(), f.clone());
1136 f
1137 }
1138
1139 pub fn scatter_trim_logits(&self, src: &CudaSlice<f32>, d2t: &CudaSlice<u32>,
1142 dst: &mut CudaSlice<f32>, d_vocab: usize, n_vocab: usize)
1143 -> Result<(), Box<dyn std::error::Error>> {
1144 let f1 = self.func("scatter_trim_logits_f32");
1145 let f2 = self.func("scatter_trim_logits_pass2_f32");
1146 let (dv, nv) = (d_vocab as i32, n_vocab as i32);
1147 let cfg1 = LaunchConfig { grid_dim: (256, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1148 let __s_b1 = self.gpu.stream();
1149 let mut b1 = __s_b1.launch_builder(&f1);
1150 b1.arg(src).arg(d2t).arg(&mut *dst).arg(&dv).arg(&nv);
1151 unsafe { b1.launch(cfg1)?; }
1152 let cfg2 = LaunchConfig { grid_dim: (d_vocab.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1153 let __s_b2 = self.gpu.stream();
1154 let mut b2 = __s_b2.launch_builder(&f2);
1155 b2.arg(src).arg(d2t).arg(&mut *dst).arg(&dv);
1156 unsafe { b2.launch(cfg2)?; }
1157 Ok(())
1158 }
1159
1160 #[allow(clippy::too_many_arguments)]
1166 pub fn filter_stats(&self, x: &CudaSlice<f32>, row_stride: usize, rows: &CudaSlice<i32>,
1167 out_th: &mut CudaSlice<f32>, out_z: &mut CudaSlice<f32>,
1168 out_max: &mut CudaSlice<f32>, n: usize, nrow: usize,
1169 temp: f32, top_k: i32, top_p: f32, min_p: f32)
1170 -> Result<(), Box<dyn std::error::Error>> {
1171 let f = self.func("filter_stats_f32");
1172 let (ni, nr, rs) = (n as i32, nrow as i32, row_stride as i64);
1173 let cfg = LaunchConfig { grid_dim: (nrow as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
1174 let __s_b = self.gpu.stream();
1175 let mut b = __s_b.launch_builder(&f);
1176 b.arg(x).arg(&rs).arg(rows).arg(&mut *out_th).arg(&mut *out_z).arg(&mut *out_max)
1177 .arg(&ni).arg(&nr).arg(&temp).arg(&top_k).arg(&top_p).arg(&min_p);
1178 unsafe { b.launch(cfg)?; }
1179 Ok(())
1180 }
1181
1182 #[allow(clippy::too_many_arguments)]
1184 pub fn softmax_gather_filtered(&self, x: &CudaSlice<f32>, row_stride: usize,
1185 ids: &CudaSlice<u32>, rows: &CudaSlice<i32>,
1186 th: &CudaSlice<f32>, z: &CudaSlice<f32>,
1187 out: &mut CudaSlice<f32>, n: usize, npair: usize, temp: f32)
1188 -> Result<(), Box<dyn std::error::Error>> {
1189 let f = self.func("softmax_gather_filtered_f32");
1190 let (ni, np, rs) = (n as i32, npair as i32, row_stride as i64);
1191 let cfg = LaunchConfig { grid_dim: (npair as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1192 let __s_b = self.gpu.stream();
1193 let mut b = __s_b.launch_builder(&f);
1194 b.arg(x).arg(&rs).arg(ids).arg(rows).arg(th).arg(z).arg(&mut *out).arg(&ni).arg(&np).arg(&temp);
1195 unsafe { b.launch(cfg)?; }
1196 Ok(())
1197 }
1198
1199 #[allow(clippy::too_many_arguments)]
1201 pub fn residual_sample_filtered(&self, p: &CudaSlice<f32>, q: Option<&CudaSlice<f32>>, n: usize,
1202 temp: f32, seed: u64, stream_pos: u32,
1203 p_stats: (f32, f32, f32), q_stats: (f32, f32, f32),
1204 out_tok: &mut CudaSlice<u32>)
1205 -> Result<(), Box<dyn std::error::Error>> {
1206 let f = self.func("residual_sample_filtered_f32");
1207 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1208 let has_q: i32 = q.is_some() as i32;
1209 let qbuf = q.unwrap_or(p);
1210 let (pm, pth, pz) = p_stats; let (qm, qth, qz) = q_stats;
1211 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
1212 let __s_b = self.gpu.stream();
1213 let mut b = __s_b.launch_builder(&f);
1214 b.arg(p).arg(qbuf).arg(&has_q).arg(&ni).arg(&temp).arg(&slo).arg(&shi).arg(&stream_pos)
1215 .arg(&pm).arg(&pth).arg(&pz).arg(&qm).arg(&qth).arg(&qz).arg(&mut *out_tok);
1216 unsafe { b.launch(cfg)?; }
1217 Ok(())
1218 }
1219
1220 #[allow(clippy::too_many_arguments)]
1222 pub fn gumbel_perturb_filtered(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
1223 seed: u64, stream_pos: u32, temp: f32, row_max: f32, th: f32)
1224 -> Result<(), Box<dyn std::error::Error>> {
1225 let f = self.func("gumbel_perturb_filtered_f32");
1226 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1227 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1228 let __s_b = self.gpu.stream();
1229 let mut b = __s_b.launch_builder(&f);
1230 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp).arg(&row_max).arg(&th);
1231 unsafe { b.launch(cfg)?; }
1232 Ok(())
1233 }
1234
1235 #[allow(clippy::too_many_arguments)]
1239 pub fn penalize_logits(&self, x: &mut CudaSlice<f32>, hist: &CudaSlice<u32>, n_hist: usize,
1240 rep: f32, freq: f32, present: f32, n: usize)
1241 -> Result<(), Box<dyn std::error::Error>> {
1242 if n_hist == 0 { return Ok(()); }
1243 let f = self.func("penalize_logits_f32");
1244 let (nh, ni) = (n_hist as i32, n as i32);
1245 let cfg = LaunchConfig { grid_dim: (n_hist.div_ceil(128) as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
1246 let __s_b = self.gpu.stream();
1247 let mut b = __s_b.launch_builder(&f);
1248 b.arg(&mut *x).arg(hist).arg(&nh).arg(&rep).arg(&freq).arg(&present).arg(&ni);
1249 unsafe { b.launch(cfg)?; }
1250 Ok(())
1251 }
1252
1253 #[allow(clippy::too_many_arguments)]
1255 pub fn penalize_logits_rows(&self, x: &mut CudaSlice<f32>, hist: &CudaSlice<u32>, n_hist: usize,
1256 rep: f32, freq: f32, present: f32, n: usize, nrow: usize)
1257 -> Result<(), Box<dyn std::error::Error>> {
1258 if n_hist == 0 || nrow == 0 { return Ok(()); }
1259 let f = self.func("penalize_logits_rows_f32");
1260 let (nh, ni, nr) = (n_hist as i32, n as i32, nrow as i32);
1261 let cfg = LaunchConfig { grid_dim: (n_hist.div_ceil(128) as u32, nrow as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
1262 let __s_b = self.gpu.stream();
1263 let mut b = __s_b.launch_builder(&f);
1264 b.arg(&mut *x).arg(hist).arg(&nh).arg(&rep).arg(&freq).arg(&present).arg(&ni).arg(&nr);
1265 unsafe { b.launch(cfg)?; }
1266 Ok(())
1267 }
1268
1269 pub fn wpf_level() -> u32 {
1277 static ON: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
1278 *ON.get_or_init(|| std::env::var("MEMRA_WPF").ok()
1279 .and_then(|v| v.parse().ok()).unwrap_or(1))
1280 }
1281
1282 pub fn set_verify_exact(&self, on: bool) {
1294 self.verify_exact.store(on, std::sync::atomic::Ordering::Relaxed);
1295 }
1296 pub(crate) fn verify_exact_on(&self) -> bool {
1297 self.verify_exact.load(std::sync::atomic::Ordering::Relaxed)
1298 }
1299
1300 pub fn qkv_append_on() -> bool {
1303 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1304 *ON.get_or_init(|| std::env::var("MEMRA_QKV_APPEND").map(|v| v != "0").unwrap_or(true))
1305 }
1306
1307 pub fn pdl_wb_on() -> bool {
1310 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1311 *ON.get_or_init(|| std::env::var("MEMRA_PDL_WB").map(|v| v != "0").unwrap_or(true))
1312 }
1313
1314 pub fn pdl_mmvq_on() -> bool {
1318 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1319 *ON.get_or_init(|| std::env::var("MEMRA_PDL_MMVQ").map(|v| v != "0").unwrap_or(true))
1320 }
1321
1322 pub fn pdl_on() -> bool {
1323 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1324 *ON.get_or_init(|| std::env::var("MEMRA_PDL").map(|v| v != "0").unwrap_or(true))
1325 }
1326
1327 fn q40_mr1_on() -> bool {
1333 static Q40MR: std::sync::OnceLock<Option<u32>> = std::sync::OnceLock::new();
1334 match *Q40MR.get_or_init(|| std::env::var("MEMRA_Q40_MR").ok()
1335 .and_then(|v| v.parse().ok())) {
1336 Some(v) => v == 1,
1337 None => crate::FUSED_MR1_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
1338 }
1339 }
1340
1341 fn pdl_func_flash(&self, g: bool, name: &'static str)
1346 -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
1347 use cudarc::driver::sys as cu;
1348 static MODS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool), usize>>> =
1355 std::sync::Mutex::new(None);
1356 static FNS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool, &'static str), usize>>> =
1357 std::sync::Mutex::new(None);
1358 let ctx_key = self.ctx().cu_ctx() as usize;
1359 if let Some(&f) = FNS.lock().unwrap().get_or_insert_with(Default::default)
1360 .get(&(ctx_key, g, name)) { return Ok(f as cu::CUfunction); }
1361 let module = {
1362 let mut mods = MODS.lock().unwrap();
1363 let map = mods.get_or_insert_with(Default::default);
1364 match map.get(&(ctx_key, g)) {
1365 Some(&m) => m,
1366 None => {
1367 let m = self.pdl_load_module_in_ctx(
1368 if g { FLASH_FATBIN_KF8VF8 } else { FLASH_FATBIN })?;
1369 map.insert((ctx_key, g), m);
1370 m
1371 }
1372 }
1373 };
1374 let cname = std::ffi::CString::new(name)?;
1375 let mut f: cu::CUfunction = std::ptr::null_mut();
1376 let r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
1377 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("pdl_func_flash {name} (g={g}): {r:?}").into()); }
1378 FNS.lock().unwrap().get_or_insert_with(Default::default)
1379 .insert((ctx_key, g, name), f as usize);
1380 Ok(f)
1381 }
1382
1383 fn pdl_load_module_in_ctx(&self, bytes: &[u8]) -> Result<usize, Box<dyn std::error::Error>> {
1388 use cudarc::driver::sys as cu;
1389 let mut prev: cu::CUcontext = std::ptr::null_mut();
1390 unsafe { cu::cuCtxGetCurrent(&mut prev).result()?; }
1391 self.ctx().bind_to_thread()?;
1392 let mut m: cu::CUmodule = std::ptr::null_mut();
1393 let r = unsafe { cu::cuModuleLoadData(&mut m, bytes.as_ptr() as *const std::ffi::c_void) };
1394 let restore = if prev.is_null() { cu::CUresult::CUDA_SUCCESS }
1395 else { unsafe { cu::cuCtxSetCurrent(prev) } };
1396 if r != cu::CUresult::CUDA_SUCCESS {
1397 return Err(format!("pdl module load: {r:?}").into());
1398 }
1399 if restore != cu::CUresult::CUDA_SUCCESS {
1400 return Err(format!("pdl module load: ctx restore {restore:?}").into());
1401 }
1402 Ok(m as usize)
1403 }
1404
1405 fn pdl_func(&self, name: &'static str) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
1406 use cudarc::driver::sys as cu;
1407 static MODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
1410 std::sync::Mutex::new(None);
1411 static QMODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
1414 std::sync::Mutex::new(None);
1415 static FNS: std::sync::Mutex<Option<std::collections::HashMap<(usize, &'static str), usize>>> =
1416 std::sync::Mutex::new(None);
1417 let ctx_key = self.ctx().cu_ctx() as usize;
1418 if let Some(&f) = FNS.lock().unwrap().get_or_insert_with(Default::default)
1419 .get(&(ctx_key, name)) { return Ok(f as cu::CUfunction); }
1420 let module = {
1421 let mut mods = MODULES.lock().unwrap();
1422 let map = mods.get_or_insert_with(Default::default);
1423 match map.get(&ctx_key) {
1424 Some(&m) => m,
1425 None => {
1426 let m = self.pdl_load_module_in_ctx(FATBIN)?;
1427 map.insert(ctx_key, m);
1428 m
1429 }
1430 }
1431 };
1432 let cname = std::ffi::CString::new(name)?;
1433 let mut f: cu::CUfunction = std::ptr::null_mut();
1434 let mut r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
1435 if r == cu::CUresult::CUDA_ERROR_NOT_FOUND {
1436 let qmodule = {
1437 let mut mods = QMODULES.lock().unwrap();
1438 let map = mods.get_or_insert_with(Default::default);
1439 match map.get(&ctx_key) {
1440 Some(&m) => m,
1441 None => {
1442 let m = self.pdl_load_module_in_ctx(QMATVEC_FATBIN)?;
1443 map.insert(ctx_key, m);
1444 m
1445 }
1446 }
1447 };
1448 r = unsafe { cu::cuModuleGetFunction(&mut f, qmodule as cu::CUmodule, cname.as_ptr()) };
1449 }
1450 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("pdl_func {name}: {r:?}").into()); }
1451 FNS.lock().unwrap().get_or_insert_with(Default::default)
1452 .insert((ctx_key, name), f as usize);
1453 Ok(f)
1454 }
1455
1456 unsafe fn launch_pdl_flash(&self, g: bool, name: &'static str, grid: (u32, u32, u32),
1468 block: (u32, u32, u32), smem: u32,
1469 params: &mut [*mut std::ffi::c_void])
1470 -> Result<(), Box<dyn std::error::Error>> {
1471 use cudarc::driver::sys as cu;
1472 let f = self.pdl_func_flash(g, name)?;
1473 if smem > 0 {
1474 let r = unsafe { cu::cuFuncSetAttribute(f,
1476 cu::CUfunction_attribute_enum::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
1477 smem as i32) };
1478 if r != cu::CUresult::CUDA_SUCCESS {
1479 return Err(format!("pdl smem attr {name}: {r:?}").into());
1480 }
1481 }
1482 let mut attr = cu::CUlaunchAttribute {
1483 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
1484 pad: [0; 4],
1485 value: cu::CUlaunchAttributeValue { programmaticStreamSerializationAllowed: 1 },
1486 };
1487 let cfg = cu::CUlaunchConfig {
1488 gridDimX: grid.0, gridDimY: grid.1, gridDimZ: grid.2,
1489 blockDimX: block.0, blockDimY: block.1, blockDimZ: block.2,
1490 sharedMemBytes: smem, hStream: self.gpu.stream().cu_stream(),
1491 attrs: &mut attr, numAttrs: 1,
1492 };
1493 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
1494 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("launch_pdl_flash {name}: {r:?}").into()); }
1495 Ok(())
1496 }
1497
1498 unsafe fn launch_pdl(&self, name: &'static str, grid: (u32, u32, u32), block: (u32, u32, u32),
1499 params: &mut [*mut std::ffi::c_void])
1500 -> Result<(), Box<dyn std::error::Error>> {
1501 use cudarc::driver::sys as cu;
1502 let f = self.pdl_func(name)?;
1503 let mut attr = cu::CUlaunchAttribute {
1504 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
1505 pad: [0; 4],
1506 value: cu::CUlaunchAttributeValue { programmaticStreamSerializationAllowed: 1 },
1507 };
1508 let cfg = cu::CUlaunchConfig {
1509 gridDimX: grid.0, gridDimY: grid.1, gridDimZ: grid.2,
1510 blockDimX: block.0, blockDimY: block.1, blockDimZ: block.2,
1511 sharedMemBytes: 0, hStream: self.gpu.stream().cu_stream(),
1512 attrs: &mut attr, numAttrs: 1,
1513 };
1514 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
1515 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("launch_pdl {name}: {r:?}").into()); }
1516 Ok(())
1517 }
1518
1519 pub fn prefetch_weight_l2(&self, w: &crate::model::GpuTensor)
1522 -> Result<(), Box<dyn std::error::Error>> {
1523 if let crate::model::GpuTensor::Quant { bytes, rp4, .. } = w {
1524 let p = rp4.as_ref().unwrap_or(bytes);
1525 self.prefetch_l2(p, p.len())?;
1526 }
1527 Ok(())
1528 }
1529
1530 pub fn gather_row_bf16(&self, table: &CudaSlice<u8>, tok: &CudaSlice<u32>, idx: usize,
1533 dst: &mut CudaSlice<f32>, ncols: usize)
1534 -> Result<(), Box<dyn std::error::Error>> {
1535 let f = self.func("gather_row_bf16_f32");
1536 let cfg = LaunchConfig { grid_dim: (ncols.div_ceil(256) as u32, 1, 1),
1537 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1538 let (nc, ix) = (ncols as i32, idx as i32);
1539 let __s_b = self.gpu.stream();
1540 let mut b = __s_b.launch_builder(&f);
1541 b.arg(table).arg(tok).arg(&ix).arg(dst).arg(&nc);
1542 unsafe { b.launch(cfg)?; }
1543 Ok(())
1544 }
1545
1546 pub fn add_row_inplace(&self, logits: &mut CudaSlice<f32>, bias: &CudaSlice<f32>,
1548 n: usize, row_off: usize)
1549 -> Result<(), Box<dyn std::error::Error>> {
1550 let f = self.func("add_row_inplace_f32");
1551 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1),
1552 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1553 let (ni, off) = (n as i32, row_off as i64);
1554 let __s_b = self.gpu.stream();
1555 let mut b = __s_b.launch_builder(&f);
1556 b.arg(logits).arg(bias).arg(&ni).arg(&off);
1557 unsafe { b.launch(cfg)?; }
1558 Ok(())
1559 }
1560
1561 pub fn prefetch_l2(&self, p: &CudaSlice<u8>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
1563 let f = self.func("prefetch_l2_bytes");
1564 let lines = n.div_ceil(128);
1565 let ni = n as i64;
1566 let cfg = LaunchConfig { grid_dim: (lines.div_ceil(256) as u32, 1, 1),
1567 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1568 let __s_b = self.gpu.stream();
1569 let mut b = __s_b.launch_builder(&f);
1570 b.arg(p).arg(&ni);
1571 unsafe { b.launch(cfg)?; }
1572 Ok(())
1573 }
1574
1575 pub fn router_gemv(&self, w: &CudaSlice<f32>, x: &CudaSlice<f32>, n_embd: usize,
1578 n_experts: usize, t: usize)
1579 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1580 let w8 = match std::env::var("MEMRA_ROUTER_V2").as_deref() {
1586 Ok("0") => false,
1587 Ok(_) => true,
1588 Err(_) => ROUTER_W8_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
1589 };
1590 let batch = w8 && t >= ROUTER_BATCH_MIN_T && router_batch_on();
1600 self.router_gemv_form(w, x, n_embd, n_experts, t, w8, batch)
1601 }
1602
1603 pub fn router_gemv_form(&self, w: &CudaSlice<f32>, x: &CudaSlice<f32>, n_embd: usize,
1606 n_experts: usize, t: usize, w8: bool, batch: bool)
1607 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1608 debug_assert!(!batch || w8, "batch twin exists for the w8 form only");
1609 let mut y = self.alloc_uninit::<f32>(t * n_experts)?;
1610 let f = if batch { self.func("router_gemv_f32_w8_batch") }
1611 else if w8 { self.func("router_gemv_f32_w8") }
1612 else { self.func("router_gemv_f32") };
1613 let (ne, nx, ti) = (n_embd as i32, n_experts as i32, t as i32);
1614 let cfg = if batch {
1615 LaunchConfig { grid_dim: (n_experts.div_ceil(8) as u32, t.div_ceil(8) as u32, 1),
1616 block_dim: (32, 8, 1), shared_mem_bytes: 0 }
1617 } else {
1618 LaunchConfig { grid_dim: (n_experts as u32, t as u32, 1),
1619 block_dim: (32, if w8 { 8 } else { 1 }, 1), shared_mem_bytes: 0 }
1620 };
1621 let __s_b = self.gpu.stream();
1622 let mut b = __s_b.launch_builder(&f);
1623 b.arg(w).arg(x).arg(&mut y).arg(&ne).arg(&nx).arg(&ti);
1624 unsafe { b.launch(cfg)?; }
1625 Ok(y)
1626 }
1627
1628 pub fn rows_permute(&self, src: &CudaSlice<f32>, idx: &CudaSlice<i32>, nrows: usize,
1630 ncols: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1631 let mut dst = self.alloc_uninit::<f32>(nrows * ncols)?;
1632 let f = self.func("rows_permute_f32");
1633 let (nc, nr) = (ncols as i32, nrows as i32);
1634 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (256, 1, 1),
1635 shared_mem_bytes: 0 };
1636 let __s_b = self.gpu.stream();
1637 let mut b = __s_b.launch_builder(&f);
1638 b.arg(src).arg(idx).arg(&mut dst).arg(&nc).arg(&nr);
1639 unsafe { b.launch(cfg)?; }
1640 Ok(dst)
1641 }
1642
1643 pub fn sigmoid_dot_rows(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, n_embd: usize,
1648 t: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1649 static OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1652 if *OFF.get_or_init(|| std::env::var("MEMRA_SHEXP_DOT").as_deref() == Ok("0")) {
1653 let gs = self.linear(x, w, t, n_embd, 1)?;
1654 let mut g = self.uninit(t)?;
1655 self.sigmoid(&gs, &mut g, t)?;
1656 return Ok(g);
1657 }
1658 let mut g = self.alloc_uninit::<f32>(t)?;
1664 let f = self.func("sigmoid_dot_rows_f32");
1665 let (ne, ti) = (n_embd as i32, t as i32);
1666 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (32, 8, 1),
1667 shared_mem_bytes: 0 };
1668 let __s_b = self.gpu.stream();
1669 let mut b = __s_b.launch_builder(&f);
1670 b.arg(x).arg(w).arg(&mut g).arg(&ne).arg(&ti);
1671 unsafe { b.launch(cfg)?; }
1672 Ok(g)
1673 }
1674
1675 pub fn spec_rollback_stream(&self, len_ptrs: &CudaSlice<u64>, pos_start: &CudaSlice<i32>,
1677 acc: &CudaSlice<u32>, base: usize, n_rows: usize)
1678 -> Result<(), Box<dyn std::error::Error>> {
1679 let f = self.func("spec_rollback_stream");
1680 let (b, nr) = (base as i32, n_rows as i32);
1681 let cfg = LaunchConfig { grid_dim: (n_rows.div_ceil(64) as u32, 1, 1),
1682 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1683 let __s_bl = self.gpu.stream();
1684 let mut bl = __s_bl.launch_builder(&f);
1685 bl.arg(len_ptrs).arg(pos_start).arg(acc).arg(&b).arg(&nr);
1686 unsafe { bl.launch(cfg)?; }
1687 Ok(())
1688 }
1689
1690 pub fn plain_tok_ring(&self, vam: &CudaSlice<u32>, pos_start: &CudaSlice<i32>,
1692 base: usize, ring: &mut CudaSlice<u32>)
1693 -> Result<(), Box<dyn std::error::Error>> {
1694 let f = self.func("plain_tok_ring");
1695 let (b, cap) = (base as i32, ring.len() as i32);
1696 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1697 let __s_bl = self.gpu.stream();
1698 let mut bl = __s_bl.launch_builder(&f);
1699 bl.arg(vam).arg(pos_start).arg(&b).arg(&mut *ring).arg(&cap);
1700 unsafe { bl.launch(cfg)?; }
1701 Ok(())
1702 }
1703
1704 pub fn spec_ring_commit(&self, vtok: &CudaSlice<u32>, acc: &CudaSlice<u32>,
1706 brk: &CudaSlice<u32>, ring: &mut CudaSlice<u32>,
1707 pend: &mut CudaSlice<u32>)
1708 -> Result<(), Box<dyn std::error::Error>> {
1709 let f = self.func("spec_ring_commit");
1710 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1711 let __s_b = self.gpu.stream();
1712 let mut b = __s_b.launch_builder(&f);
1713 b.arg(vtok).arg(acc).arg(brk).arg(ring).arg(pend);
1714 unsafe { b.launch(cfg)?; }
1715 Ok(())
1716 }
1717 pub fn i32_copy_add(&self, src: &CudaSlice<i32>, dst: &mut CudaSlice<i32>, delta: i32)
1718 -> Result<(), Box<dyn std::error::Error>> {
1719 let f = self.func("i32_copy_add");
1720 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1721 let __s_b = self.gpu.stream();
1722 let mut b = __s_b.launch_builder(&f);
1723 b.arg(src).arg(dst).arg(&delta);
1724 unsafe { b.launch(cfg)?; }
1725 Ok(())
1726 }
1727 pub fn u32_copy(&self, src: &CudaSlice<u32>, dst: &mut CudaSlice<u32>)
1728 -> Result<(), Box<dyn std::error::Error>> {
1729 let f = self.func("u32_copy");
1730 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1731 let __s_b = self.gpu.stream();
1732 let mut b = __s_b.launch_builder(&f);
1733 b.arg(src).arg(dst);
1734 unsafe { b.launch(cfg)?; }
1735 Ok(())
1736 }
1737
1738 pub fn spec_adapt_k(&self, acc: &CudaSlice<u32>, brk: &mut CudaSlice<u32>,
1742 floor: usize, cap: usize)
1743 -> Result<(), Box<dyn std::error::Error>> {
1744 let f = self.func("spec_adapt_k");
1745 let (fl, cp) = (floor as i32, cap as i32);
1746 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1747 let __s_b = self.gpu.stream();
1748 let mut b = __s_b.launch_builder(&f);
1749 b.arg(acc).arg(brk).arg(&fl).arg(&cp);
1750 unsafe { b.launch(cfg)?; }
1751 Ok(())
1752 }
1753
1754 pub fn spec_accept_greedy_dc(&self, preds: &CudaSlice<u32>, vtok: &CudaSlice<u32>,
1756 last_pred: &CudaSlice<u32>, brk: &CudaSlice<u32>,
1757 out: &mut CudaSlice<u32>)
1758 -> Result<(), Box<dyn std::error::Error>> {
1759 let f = self.func("spec_accept_greedy_dc");
1760 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1761 let __s_b = self.gpu.stream();
1762 let mut b = __s_b.launch_builder(&f);
1763 b.arg(preds).arg(vtok).arg(last_pred).arg(brk).arg(out);
1764 unsafe { b.launch(cfg)?; }
1765 Ok(())
1766 }
1767
1768 pub fn pos_iota(&self, pos0: &CudaSlice<i32>, out: &mut CudaSlice<i32>, t: usize)
1770 -> Result<(), Box<dyn std::error::Error>> {
1771 let f = self.func("pos_iota_i32");
1772 let ti = t as i32;
1773 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (t.max(1) as u32, 1, 1),
1774 shared_mem_bytes: 0 };
1775 let __s_b = self.gpu.stream();
1776 let mut b = __s_b.launch_builder(&f);
1777 b.arg(pos0).arg(out).arg(&ti);
1778 unsafe { b.launch(cfg)?; }
1779 Ok(())
1780 }
1781 #[allow(clippy::too_many_arguments)]
1782 pub fn append_kv_quantized_rows_dc(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
1783 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
1784 t0_dev: &CudaSlice<i32>, t: usize,
1785 kv_dim_k: usize, kv_dim_v: usize,
1786 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
1787 -> Result<(), Box<dyn std::error::Error>> {
1788 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_rows_dc") }
1789 else { self.func("append_quantize_kv_q8_0_q5_1_rows_dc") };
1790 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
1791 let cfg = LaunchConfig { grid_dim: (nblk, t as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1792 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
1793 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
1794 let __s_b = self.gpu.stream();
1795 let mut b = __s_b.launch_builder(&f);
1796 b.arg(k_rows).arg(v_rows).arg(kc).arg(vc).arg(t0_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
1797 unsafe { b.launch(cfg)?; }
1798 Ok(())
1799 }
1800
1801 #[allow(clippy::too_many_arguments)]
1804 pub fn append_kv_quantized_row_dc_inc(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
1805 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
1806 t0_dev: &mut CudaSlice<i32>,
1807 kv_dim_k: usize, kv_dim_v: usize,
1808 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
1809 -> Result<(), Box<dyn std::error::Error>> {
1810 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_dc_inc") }
1811 else { self.func("append_quantize_kv_q8_0_q5_1_dc_inc") };
1812 let nthreads = ((kv_dim_k.max(kv_dim_v) / 32) * 32).min(1024) as u32;
1813 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (nthreads, 1, 1),
1814 shared_mem_bytes: 0 };
1815 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
1816 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
1817 let __s_b = self.gpu.stream();
1818 let mut b = __s_b.launch_builder(&f);
1819 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(t0_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
1820 unsafe { b.launch(cfg)?; }
1821 Ok(())
1822 }
1823
1824 pub fn pack_tok_p(&self, tok: &CudaSlice<u32>, p: &CudaSlice<f32>, out: &mut CudaSlice<u32>,
1826 slot: usize) -> Result<(), Box<dyn std::error::Error>> {
1827 let f = self.func("pack_tok_p");
1828 let sl = slot as i32;
1829 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1830 let __s_b = self.gpu.stream();
1831 let mut b = __s_b.launch_builder(&f);
1832 b.arg(tok).arg(p).arg(out).arg(&sl);
1833 unsafe { b.launch(cfg)?; }
1834 Ok(())
1835 }
1836 pub fn tok_map_u32(&self, tok: &mut CudaSlice<u32>, map: &CudaSlice<u32>)
1837 -> Result<(), Box<dyn std::error::Error>> {
1838 let f = self.func("tok_map_u32");
1839 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1840 let __s_b = self.gpu.stream();
1841 let mut b = __s_b.launch_builder(&f);
1842 b.arg(tok).arg(map);
1843 unsafe { b.launch(cfg)?; }
1844 Ok(())
1845 }
1846
1847 #[allow(clippy::too_many_arguments)]
1849 pub fn spec_assemble_verify(&self, tokp: &CudaSlice<u32>, pend: &CudaSlice<u32>,
1850 d2t: Option<&CudaSlice<u32>>, vtok: &mut CudaSlice<u32>,
1851 brk: &mut CudaSlice<u32>, p_min: f32, k: usize, pmin0: bool)
1852 -> Result<(), Box<dyn std::error::Error>> {
1853 let f = self.func("spec_assemble_verify");
1854 let (ki, pm) = (k as i32, if pmin0 { 1i32 } else { 0i32 });
1855 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1856 let __s_b = self.gpu.stream();
1857 let mut b = __s_b.launch_builder(&f);
1858 match d2t {
1859 Some(m) => { b.arg(tokp).arg(pend).arg(m).arg(vtok).arg(brk).arg(&p_min).arg(&ki).arg(&pm);
1860 unsafe { b.launch(cfg)?; } }
1861 None => { let null: u64 = 0;
1862 b.arg(tokp).arg(pend).arg(&null).arg(vtok).arg(brk).arg(&p_min).arg(&ki).arg(&pm);
1863 unsafe { b.launch(cfg)?; } }
1864 }
1865 Ok(())
1866 }
1867
1868 #[allow(clippy::too_many_arguments)]
1870 pub fn ssm_conv_ring_rebuild_dc(&self, qkv_tm: &CudaSlice<f32>, ring_old: &CudaSlice<f32>,
1871 conv_state: &mut CudaSlice<f32>, conv_dim: usize,
1872 acc: &CudaSlice<u32>, base: usize, t_v: usize, d_conv: usize)
1873 -> Result<(), Box<dyn std::error::Error>> {
1874 let f = self.func("ssm_conv_ring_rebuild_f32_dc");
1875 let n = conv_dim * (d_conv - 1);
1876 let cfg = LaunchConfig::for_num_elems(n as u32);
1877 let (cd, b0, tv, dc) = (conv_dim as i32, base as i32, t_v as i32, d_conv as i32);
1878 let __s_b = self.gpu.stream();
1879 let mut b = __s_b.launch_builder(&f);
1880 b.arg(qkv_tm).arg(ring_old).arg(conv_state).arg(&cd).arg(acc).arg(&b0).arg(&tv).arg(&dc);
1881 unsafe { b.launch(cfg)?; }
1882 Ok(())
1883 }
1884 #[allow(clippy::too_many_arguments)]
1885 pub fn gdn_scan_s128_dc(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
1886 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
1887 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
1888 n_head: usize, acc: &CudaSlice<u32>, base: usize, t_v: usize,
1889 scale: f32)
1890 -> Result<(), Box<dyn std::error::Error>> {
1891 let f = self.func("gdn_scan_s128_dc");
1892 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
1893 let cfg = LaunchConfig {
1894 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
1895 block_dim: (WARP, COLS_PER_BLOCK, 1),
1896 shared_mem_bytes: 0,
1897 };
1898 let (h, b0, tv) = (n_head as i32, base as i32, t_v as i32);
1899 let __s_b = self.gpu.stream();
1900 let mut b = __s_b.launch_builder(&f);
1901 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in).arg(state_out).arg(o)
1902 .arg(&h).arg(acc).arg(&b0).arg(&tv).arg(&scale);
1903 unsafe { b.launch(cfg)?; }
1904 Ok(())
1905 }
1906
1907 pub fn spec_rollback_kv(&self, len_ptrs: &CudaSlice<u64>, saved: &CudaSlice<i32>,
1909 acc: &CudaSlice<u32>, base: usize, n_layer: usize)
1910 -> Result<(), Box<dyn std::error::Error>> {
1911 let f = self.func("spec_rollback_kv");
1912 let (b, nl) = (base as i32, n_layer as i32);
1913 let cfg = LaunchConfig { grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
1914 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1915 let __s_bl = self.gpu.stream();
1916 let mut bl = __s_bl.launch_builder(&f);
1917 bl.arg(len_ptrs).arg(saved).arg(acc).arg(&b).arg(&nl);
1918 unsafe { bl.launch(cfg)?; }
1919 Ok(())
1920 }
1921
1922 pub fn spec_seed_gather(&self, vx: &CudaSlice<f32>, fill_prev: &CudaSlice<f32>,
1925 acc: &CudaSlice<u32>, h_seed: &mut CudaSlice<f32>,
1926 base: usize, n_embd: usize)
1927 -> Result<(), Box<dyn std::error::Error>> {
1928 let f = self.func("spec_seed_gather");
1929 let (b, ne) = (base as i32, n_embd as i32);
1930 let cfg = LaunchConfig { grid_dim: (n_embd.div_ceil(256) as u32, 1, 1),
1931 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1932 let __s_bl = self.gpu.stream();
1933 let mut bl = __s_bl.launch_builder(&f);
1934 bl.arg(vx).arg(fill_prev).arg(acc).arg(h_seed).arg(&b).arg(&ne);
1935 unsafe { bl.launch(cfg)?; }
1936 Ok(())
1937 }
1938
1939
1940 pub fn spec_accept_greedy(&self, preds: &CudaSlice<u32>, draft: &CudaSlice<u32>,
1942 last_pred: u32, base: usize, k_round: usize,
1943 out: &mut CudaSlice<u32>)
1944 -> Result<(), Box<dyn std::error::Error>> {
1945 let f = self.func("spec_accept_greedy");
1946 let (b, k) = (base as i32, k_round as i32);
1947 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1948 let __s_bl = self.gpu.stream();
1949 let mut bl = __s_bl.launch_builder(&f);
1950 bl.arg(preds).arg(draft).arg(&last_pred).arg(&b).arg(&k).arg(out);
1951 unsafe { bl.launch(cfg)?; }
1952 Ok(())
1953 }
1954
1955 pub fn gumbel_perturb(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
1962 seed: u64, stream_pos: u32, temp: f32)
1963 -> Result<(), Box<dyn std::error::Error>> {
1964 let f = self.func("gumbel_perturb_f32");
1965 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1966 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1967 let __s_b = self.gpu.stream();
1968 let mut b = __s_b.launch_builder(&f);
1969 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
1970 unsafe { b.launch(cfg)?; }
1971 Ok(())
1972 }
1973
1974 pub fn mask_logits_col(&self, logits: &mut CudaSlice<f32>, mask: &CudaSlice<u32>,
1982 col: usize, n: usize, mask_words: usize)
1983 -> Result<(), Box<dyn std::error::Error>> {
1984 let f = self.func("mask_logits_f32");
1985 let (ci, ni, mw) = (col as i32, n as i32, mask_words as i32);
1986 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256).min(1024) as u32, 1, 1),
1987 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1988 let __s_b = self.gpu.stream();
1989 let mut b = __s_b.launch_builder(&f);
1990 b.arg(&mut *logits).arg(mask).arg(&ci).arg(&ni).arg(&mw);
1991 unsafe { b.launch(cfg)?; }
1992 Ok(())
1993 }
1994
1995 pub fn gumbel_perturb_col(&self, x: &CudaSlice<f32>, col: usize, y: &mut CudaSlice<f32>,
2002 n: usize, seed: u64, stream_pos: u32, temp: f32)
2003 -> Result<(), Box<dyn std::error::Error>> {
2004 let f = self.func("gumbel_perturb_f32");
2005 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2006 let col_view = x.slice(col * n..(col + 1) * n);
2007 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2008 let __s_b = self.gpu.stream();
2009 let mut b = __s_b.launch_builder(&f);
2010 b.arg(&col_view).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
2011 unsafe { b.launch(cfg)?; }
2012 Ok(())
2013 }
2014
2015 pub fn sctr_inc(&self, ctr: &mut CudaSlice<u32>) -> Result<(), Box<dyn std::error::Error>> {
2020 let f = self.func("memra_sctr_inc");
2021 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
2022 let __s_b = self.gpu.stream();
2023 let mut b = __s_b.launch_builder(&f);
2024 b.arg(&mut *ctr);
2025 unsafe { b.launch(cfg)?; }
2026 Ok(())
2027 }
2028
2029 pub fn gumbel_perturb_ctr(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
2034 seed: u64, ctr: &CudaSlice<u32>, temp: f32)
2035 -> Result<(), Box<dyn std::error::Error>> {
2036 let f = self.func("gumbel_perturb_ctr_f32");
2037 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2038 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2039 let __s_b = self.gpu.stream();
2040 let mut b = __s_b.launch_builder(&f);
2041 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(ctr).arg(&temp);
2042 unsafe { b.launch(cfg)?; }
2043 Ok(())
2044 }
2045
2046 pub fn softmax_gather(&self, x: &CudaSlice<f32>, row_stride: usize,
2050 ids: &CudaSlice<u32>, rows: &CudaSlice<i32>,
2051 out: &mut CudaSlice<f32>, n: usize, npair: usize, temp: f32)
2052 -> Result<(), Box<dyn std::error::Error>> {
2053 let f = self.func("softmax_gather_f32");
2054 let (ni, rs) = (n as i32, row_stride as i64);
2055 let np = npair as i32;
2056 let cfg = LaunchConfig { grid_dim: (npair as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2057 let __s_b = self.gpu.stream();
2058 let mut b = __s_b.launch_builder(&f);
2059 b.arg(x).arg(&rs).arg(ids).arg(rows).arg(&mut *out).arg(&ni).arg(&np).arg(&temp);
2060 unsafe { b.launch(cfg)?; }
2061 Ok(())
2062 }
2063
2064 pub fn residual_sample(&self, p: &CudaSlice<f32>, q: Option<&CudaSlice<f32>>, n: usize,
2068 temp: f32, seed: u64, stream_pos: u32,
2069 out_tok: &mut CudaSlice<u32>)
2070 -> Result<(), Box<dyn std::error::Error>> {
2071 let f = self.func("residual_sample_f32");
2072 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2073 let nth = 1024u32;
2074 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (nth, 1, 1), shared_mem_bytes: 0 };
2075 let has_q: i32 = q.is_some() as i32;
2076 let qbuf = q.unwrap_or(p); let __s_b = self.gpu.stream();
2078 let mut b = __s_b.launch_builder(&f);
2079 b.arg(p).arg(qbuf).arg(&has_q).arg(&ni).arg(&temp).arg(&slo).arg(&shi).arg(&stream_pos)
2080 .arg(&mut *out_tok);
2081 unsafe { b.launch(cfg)?; }
2082 Ok(())
2083 }
2084
2085 pub fn with_moe_cache<R>(&self, max_block_bytes: usize,
2090 f: impl FnOnce(&mut crate::moe_cache::MoeSlotCache, &Engine) -> Result<R, Box<dyn std::error::Error>>)
2091 -> Result<R, Box<dyn std::error::Error>> {
2092 let mut guard = self.moe_cache.lock().unwrap();
2093 if guard.is_none() {
2094 *guard = Some(crate::moe_cache::MoeSlotCache::new(self, max_block_bytes)?);
2095 }
2096 let cache = guard.as_mut().unwrap();
2097 f(cache, self)
2098 }
2099
2100 pub fn freeze_moe_cache(&self) {
2103 if let Some(cache) = self.moe_cache.lock().unwrap().as_mut() {
2104 cache.freeze();
2105 }
2106 }
2107
2108 pub fn export_moe_residency(&self) -> Option<Vec<(u16, u8, u16)>> {
2111 self.moe_cache
2112 .lock()
2113 .unwrap()
2114 .as_ref()
2115 .map(crate::moe_cache::MoeSlotCache::export_residency)
2116 }
2117
2118 pub(crate) fn moe_cache_frozen(&self) -> bool {
2119 self.moe_cache
2120 .lock()
2121 .unwrap()
2122 .as_ref()
2123 .is_some_and(crate::moe_cache::MoeSlotCache::is_frozen)
2124 }
2125
2126 pub fn frozen_cpu_experts_prefer_tokenwise_prime(&self) -> bool {
2133 crate::cpu_experts::configured()
2134 && self.moe_cache_frozen()
2135 && std::env::var("MEMRA_CPU_EXPERT_BATCHED_PRIME").as_deref() != Ok("1")
2136 }
2137
2138 pub(crate) fn configure_moe_cache_layout(&self, block_bytes: Vec<usize>) {
2140 assert!(
2141 self.moe_cache.lock().unwrap().is_none(),
2142 "MoE cache layout configured after cache construction"
2143 );
2144 *self.moe_cache_layout.lock().unwrap() = Some(block_bytes);
2145 }
2146
2147 pub(crate) fn moe_cache_layout(&self) -> Option<Vec<usize>> {
2148 self.moe_cache_layout.lock().unwrap().clone()
2149 }
2150
2151 pub fn moe_cache_enabled() -> bool {
2153 std::env::var("MEMRA_MOE_CACHE").as_deref() != Ok("0")
2154 }
2155
2156 pub fn moe_cache_stats(&self) -> Option<(u64, u64, u64, usize)> {
2159 let guard = self.moe_cache.lock().unwrap();
2160 guard.as_ref() .map(|c| (c.hits, c.misses, c.staged_bytes, c.n_slots()))
2161 }
2162
2163 pub fn cpu_expert_stats(
2167 &self,
2168 ) -> Option<(u64, u64, u64, u64, u64, u64, u64, u64, u64, u64, u64)> {
2169 crate::cpu_experts::configured().then(crate::cpu_experts::stats)
2170 }
2171
2172 pub fn cpu_expert_predictor_stats(&self) -> (u64, u64) {
2175 crate::cpu_experts::predictor_stats()
2176 }
2177
2178 pub fn cpu_expert_exposed_wait_ns(&self) -> Option<u64> {
2179 crate::cpu_experts::configured().then(crate::cpu_experts::exposed_wait_ns)
2180 }
2181
2182 pub fn cpu_expert_gpu_residency_stats(&self) -> Option<(u64, u64, u64)> {
2185 crate::cpu_experts::configured().then(crate::cpu_experts::incomplete_gpu_residency_stats)
2186 }
2187
2188 pub fn moe_pread_stats(&self) -> Option<(u64, u64, u64, u64, u64, u64, u64)> {
2191
2192 let guard = self.moe_cache.lock().unwrap();
2193 guard.as_ref().and_then(|cache| cache.pread_stats()).map(|stats| (
2194 stats.reads,
2195 stats.bytes,
2196 stats.read_errors,
2197 stats.short_reads,
2198 stats.fallbacks,
2199 stats.buffer_waits,
2200 stats.ring_full,
2201 ))
2202 }
2203
2204 pub fn moe_cache_reset_counters(&self) {
2206 if let Some(c) = self.moe_cache.lock().unwrap().as_mut() { c.reset_counters(); }
2207 }
2208
2209 pub fn htod_bytes(&self, v: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2210 Ok(self.gpu.stream().clone_htod(v)?)
2211 }
2212
2213 pub fn htod_bytes_padded(&self, v: &[u8], pad: usize)
2217 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2218 let mut d = self.alloc_u8_uninit(v.len() + pad)?;
2219 {
2220 let mut view = d.slice_mut(0..v.len());
2221 self.gpu.stream().memcpy_htod(v, &mut view)?;
2222 }
2223 Ok(d)
2224 }
2225
2226 pub fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
2228 -> Result<(), Box<dyn std::error::Error>> {
2229 let mut view = dst.slice_mut(off..off + len);
2230 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2231 Ok(())
2232 }
2233
2234 pub fn copy_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &CudaSlice<u8>, len: usize)
2237 -> Result<(), Box<dyn std::error::Error>> {
2238 let mut view = dst.slice_mut(off..off + len);
2239 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2240 Ok(())
2241 }
2242
2243 pub fn copy_u8_range_into(
2245 &self,
2246 dst: &mut CudaSlice<u8>,
2247 dst_off: usize,
2248 src: &CudaSlice<u8>,
2249 src_off: usize,
2250 len: usize,
2251 ) -> Result<(), Box<dyn std::error::Error>> {
2252 let mut dst_view = dst.slice_mut(dst_off..dst_off + len);
2253 self.gpu
2254 .stream()
2255 .memcpy_dtod(&src.slice(src_off..src_off + len), &mut dst_view)?;
2256 Ok(())
2257 }
2258
2259 pub fn prepare_kv_append(
2263 &self,
2264 kv: &mut crate::cache::KvLayer,
2265 retain_from: usize,
2266 append_rows: usize,
2267 ) -> Result<usize, Box<dyn std::error::Error>> {
2268 let Some(plan) = kv
2269 .ring
2270 .as_ref()
2271 .map(|ring| ring.append_plan(kv.len, retain_from, append_rows))
2272 .transpose()?
2273 else {
2274 return Ok(kv.len);
2275 };
2276 match plan {
2277 crate::cache::KvRingAppend::Contiguous { write_row } => Ok(write_row),
2278 crate::cache::KvRingAppend::Rebase {
2279 src_row,
2280 keep_rows,
2281 new_base,
2282 write_row,
2283 } => {
2284 if keep_rows > 0 {
2285 let k_len = keep_rows * kv.k_tok_bytes;
2286 let v_len = keep_rows * kv.v_tok_bytes;
2287 let mut k_tmp = self.alloc_u8_uninit(k_len)?;
2288 let mut v_tmp = self.alloc_u8_uninit(v_len)?;
2289 self.copy_u8_range_into(
2290 &mut k_tmp,
2291 0,
2292 &kv.k,
2293 src_row * kv.k_tok_bytes,
2294 k_len,
2295 )?;
2296 self.copy_u8_range_into(
2297 &mut v_tmp,
2298 0,
2299 &kv.v,
2300 src_row * kv.v_tok_bytes,
2301 v_len,
2302 )?;
2303 self.copy_u8_into(&mut kv.k, 0, &k_tmp, k_len)?;
2304 self.copy_u8_into(&mut kv.v, 0, &v_tmp, v_len)?;
2305 }
2306 kv.ring.as_mut().unwrap().apply_rebase(new_base);
2307 Ok(write_row)
2308 }
2309 }
2310 }
2311
2312 pub fn htod_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &[u8])
2315 -> Result<(), Box<dyn std::error::Error>> {
2316 let mut view = dst.slice_mut(off..off + src.len());
2317 self.gpu.stream().memcpy_htod(src, &mut view)?;
2318 Ok(())
2319 }
2320
2321 pub fn view<'a>(&self, b: &'a CudaSlice<f32>, len: usize) -> cudarc::driver::CudaView<'a, f32> {
2322 b.slice(0..len)
2323 }
2324
2325 pub fn view_u8_range<'a>(&self, b: &'a CudaSlice<u8>, start: usize, end: usize)
2328 -> cudarc::driver::CudaView<'a, u8> {
2329 b.slice(start..end)
2330 }
2331 pub fn view_u8<'a>(&self, b: &'a CudaSlice<u8>, len: usize) -> cudarc::driver::CudaView<'a, u8> {
2332 b.slice(0..len)
2333 }
2334
2335 pub fn append_kv_quantized(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2339 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2340 kv_dim_k: usize, kv_dim_v: usize,
2341 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2342 -> Result<(), Box<dyn std::error::Error>> {
2343 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") } else { self.func("append_quantize_kv_q8_0_q5_1") };
2344 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2345 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2346 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2347 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2348 let __s_b = self.gpu.stream();
2349 let mut b = __s_b.launch_builder(&f);
2350 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2351 unsafe { b.launch(cfg)?; }
2352 Ok(())
2353 }
2354
2355 pub fn append_kv_quantized_dc(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2359 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t_dev: &CudaSlice<i32>,
2360 kv_dim_k: usize, kv_dim_v: usize,
2361 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2362 -> Result<(), Box<dyn std::error::Error>> {
2363 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2364 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
2365 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2366 if Self::pdl_on() && Self::pdl_wb_on() {
2368 use cudarc::driver::{DevicePtr, DevicePtrMut};
2369 let s = &self.gpu.stream();
2370 let (pk, _g0) = k_row.device_ptr(s); let (pv, _g1) = v_row.device_ptr(s);
2371 let (pkc, _g2) = kc.device_ptr_mut(s); let (pvc, _g3) = vc.device_ptr_mut(s);
2372 let (pt, _g4) = t_dev.device_ptr(s);
2373 let mut ps = [
2374 &pk as *const _ as *mut std::ffi::c_void, &pv as *const _ as *mut _,
2375 &pkc as *const _ as *mut _, &pvc as *const _ as *mut _,
2376 &pt as *const _ as *mut _, &kdk as *const _ as *mut _,
2377 &kdv as *const _ as *mut _, &ktb as *const _ as *mut _,
2378 &vtb as *const _ as *mut _,
2379 ];
2380 unsafe { self.launch_pdl_flash(g, "append_quantize_kv_q8_0_q5_1_dc",
2381 (nblk, 1, 1), (32, 1, 1), 0, &mut ps)?; }
2382 return Ok(());
2383 }
2384 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") };
2385 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2386 let __s_b = self.gpu.stream();
2387 let mut b = __s_b.launch_builder(&f);
2388 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(t_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2389 unsafe { b.launch(cfg)?; }
2390 Ok(())
2391 }
2392
2393 #[allow(clippy::too_many_arguments)]
2400 pub fn append_kv_quantized_rows(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
2401 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
2402 t0: usize, t: usize, kv_dim_k: usize, kv_dim_v: usize,
2403 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2404 -> Result<(), Box<dyn std::error::Error>> {
2405 if std::env::var("MEMRA_PRIME_APPEND_LOOP").is_ok() {
2406 for i in 0..t {
2407 let k_row = k_rows.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
2408 let v_row = v_rows.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
2409 self.append_kv_quantized_view(&k_row, &v_row, kc, vc, t0 + i,
2410 kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes, g)?;
2411 }
2412 return Ok(());
2413 }
2414 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") };
2415 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2416 let cfg = LaunchConfig { grid_dim: (nblk, t as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2417 let (t0i, kdk, kdv) = (t0 as i32, kv_dim_k as i32, kv_dim_v as i32);
2418 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2419 let __s_b = self.gpu.stream();
2420 let mut b = __s_b.launch_builder(&f);
2421 b.arg(k_rows).arg(v_rows).arg(kc).arg(vc).arg(&t0i).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2422 unsafe { b.launch(cfg)?; }
2423 Ok(())
2424 }
2425
2426 pub fn inc_seqlen(&self, p: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
2430 let f = self.func("inc_i32");
2431 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 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(p);
2435 unsafe { b.launch(cfg)?; }
2436 Ok(())
2437 }
2438
2439 pub fn append_kv_quantized_view(&self, k_row: &cudarc::driver::CudaView<f32>,
2442 v_row: &cudarc::driver::CudaView<f32>,
2443 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2444 kv_dim_k: usize, kv_dim_v: usize,
2445 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2446 -> Result<(), Box<dyn std::error::Error>> {
2447 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") }
2448 else { self.func("append_quantize_kv_q8_0_q5_1") };
2449 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2450 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2451 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2452 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2453 let __s_b = self.gpu.stream();
2454 let mut b = __s_b.launch_builder(&f);
2455 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2456 unsafe { b.launch(cfg)?; }
2457 Ok(())
2458 }
2459
2460 pub fn copy_view_into(&self, dst: &mut CudaSlice<f32>, off: usize,
2463 src: &cudarc::driver::CudaView<f32>, len: usize)
2464 -> Result<(), Box<dyn std::error::Error>> {
2465 let mut view = dst.slice_mut(off..off + len);
2466 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2467 Ok(())
2468 }
2469
2470 pub fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2474 let mut dst = self.gpu.stream().alloc_zeros::<f32>(src.len())?;
2475 self.gpu.stream().memcpy_dtod(src, &mut dst)?;
2476 Ok(dst)
2477 }
2478
2479 pub fn dtod_copy_view(&self, src: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<f32>)
2482 -> Result<(), Box<dyn std::error::Error>> {
2483 self.gpu.stream().memcpy_dtod(src, dst)?;
2484 Ok(())
2485 }
2486
2487 pub fn dtod_copy_view_i8(&self, src: &cudarc::driver::CudaView<i8>, dst: &mut CudaSlice<i8>)
2489 -> Result<(), Box<dyn std::error::Error>> {
2490 self.gpu.stream().memcpy_dtod(src, dst)?;
2491 Ok(())
2492 }
2493
2494 pub fn dtod_copy_into(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, offset: usize)
2496 -> Result<(), Box<dyn std::error::Error>> {
2497 let n = src.len();
2498 let mut dv = dst.slice_mut(offset..offset + n);
2499 self.gpu.stream().memcpy_dtod(src, &mut dv)?;
2500 Ok(())
2501 }
2502
2503 pub fn uninit_i8(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
2505 self.alloc_uninit::<i8>(n)
2506 }
2507
2508 pub fn qmatvec(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
2510 qtype: i32, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2511 let f = self.func("qmatvec_f32");
2512 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 };
2514 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2515 let __s_b = self.gpu.stream();
2516 let mut b = __s_b.launch_builder(&f);
2517 b.arg(w).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2518 unsafe { b.launch(cfg)?; }
2519 Ok(y)
2520 }
2521
2522 pub fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2524 let s = self.gpu.stream().alloc_zeros::<u8>(n)?;
2525 self.keep_if_capturing(&s);
2526 Ok(s)
2527 }
2528
2529 pub fn alloc_u8_uninit(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2533 let s = unsafe { self.gpu.stream().alloc::<u8>(n)? };
2534 self.keep_if_capturing(&s);
2535 Ok(s)
2536 }
2537
2538 pub fn memset_zeros_view(&self, dst: &mut cudarc::driver::CudaViewMut<f32>)
2541 -> Result<(), Box<dyn std::error::Error>> {
2542 self.gpu.stream().memset_zeros(dst)?;
2543 Ok(())
2544 }
2545
2546 pub fn stage_expert(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2552 -> Result<(), Box<dyn std::error::Error>> {
2553 let mut dst = scratch.slice_mut(off..off + host_bytes.len()); self.gpu.stream().memcpy_htod(host_bytes, &mut dst)?; Ok(())
2556 }
2557
2558 pub fn moe_router_topk(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2564 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2565 let f = self.func("moe_router_topk_f32");
2566 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),
2569 shared_mem_bytes: 0 };
2570 let (ne, nu) = (n_expert as i32, n_used as i32);
2571 let __s_b = self.gpu.stream();
2572 let mut b = __s_b.launch_builder(&f);
2573 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2574 unsafe { b.launch(cfg)?; }
2575 Ok((sel_idx, sel_w))
2576 }
2577
2578 pub fn moe_router_topk_scaled(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize,
2581 n_used: usize, ex_scale: &CudaSlice<f32>)
2582 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2583 let f = self.func("moe_router_topk_scaled_f32");
2588 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
2589 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
2590 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2591 shared_mem_bytes: 0 };
2592 let (ne, nu) = (n_expert as i32, n_used as i32);
2593 let __s_b = self.gpu.stream();
2594 let mut b = __s_b.launch_builder(&f);
2595 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu).arg(ex_scale);
2596 unsafe { b.launch(cfg)?; }
2597 Ok((sel_idx, sel_w))
2598 }
2599
2600 pub fn moe_router_topk_host(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2608 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2609 let f = self.func("moe_router_topk_f32");
2610 let n = t * n_used;
2611 let mut sel_idx = self.alloc_uninit::<i32>(n)?;
2612 let mut sel_w = self.alloc_uninit::<f32>(n)?;
2613 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2614 shared_mem_bytes: 0 };
2615 let (ne, nu) = (n_expert as i32, n_used as i32);
2616 let __s_b = self.gpu.stream();
2617 let mut b = __s_b.launch_builder(&f);
2618 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2619 unsafe { b.launch(cfg)?; }
2620 let bytes = n * 8;
2622 let mut guard = self.router_stage.lock().unwrap();
2623 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
2624 *guard = Some(PinnedStage::new(bytes.max(4096))?);
2625 }
2626 let stage = guard.as_mut().unwrap();
2627 let (si, sw) = unsafe {
2628 (std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
2629 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n))
2630 };
2631 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()))
2635 }
2636
2637 pub fn stage_expert_async(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2641 -> Result<cudarc::driver::CudaEvent, Box<dyn std::error::Error>> {
2642 let mut dst = scratch.slice_mut(off..off + host_bytes.len());
2643 self.copy_stream.memcpy_htod(host_bytes, &mut dst)?;
2644 Ok(self.copy_stream.record_event(None)?)
2645 }
2646
2647 pub fn compute_wait(&self, ev: &cudarc::driver::CudaEvent) -> Result<(), Box<dyn std::error::Error>> {
2649 self.gpu.stream().wait(ev)?;
2650 Ok(())
2651 }
2652
2653 pub fn qmatvec_view(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2658 x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize, out_f: usize,
2659 qtype: i32, row_bytes: usize)
2660 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2661 let f = self.func("qmatvec_f32");
2662 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 };
2665 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2666 let __s_b = self.gpu.stream();
2667 let mut b = __s_b.launch_builder(&f);
2668 b.arg(&wv).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2669 unsafe { b.launch(cfg)?; }
2670 Ok(y)
2671 }
2672
2673 #[allow(clippy::too_many_arguments)]
2680 pub fn moe_gate_up_silu8_q8(&self, gp: WPtr8, up: WPtr8,
2684 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2685 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2686 rb_g: usize, rb_u: usize)
2687 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2688 let f = self.func("moe_gate_up_silu8_q8");
2689 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
2690 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2691 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2692 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2693 let __s_b = self.gpu.stream();
2694 let mut b = __s_b.launch_builder(&f);
2695 b.arg(&gp).arg(&up).arg(aq).arg(ad).arg(&mut act)
2696 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2697 unsafe { b.launch(cfg)?; }
2698 Ok(act)
2699 }
2700
2701 #[allow(clippy::too_many_arguments)]
2702 pub fn moe_down8_fma_q8(&self, dp: WPtr8, w: F32x8,
2703 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
2704 dst: &mut cudarc::driver::CudaViewMut<f32>,
2705 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2706 -> Result<(), Box<dyn std::error::Error>> {
2707 let f = self.func("moe_down8_fma_q8");
2708 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2709 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2710 let (inf, outf, nu, rbi) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2711 let __s_b = self.gpu.stream();
2712 let mut b = __s_b.launch_builder(&f);
2713 b.arg(&dp).arg(&w).arg(aq2).arg(ad2).arg(dst)
2714 .arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbi);
2715 unsafe { b.launch(cfg)?; }
2716 Ok(())
2717 }
2718
2719 pub fn qmatvec_expert_q8(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2721 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
2722 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize)
2723 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2724 let f = self.func("qmatvec_expert_q8");
2725 let wv = w.slice(range);
2726 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
2727 const ROWS: u32 = 4; let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, m as u32, 1),
2729 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2730 let (inf, outf, mi, rbi) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
2731 let __s_b = self.gpu.stream();
2732 let mut b = __s_b.launch_builder(&f);
2733 b.arg(&wv).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qtype).arg(&rbi);
2734 unsafe { b.launch(cfg)?; }
2735 Ok(y)
2736 }
2737
2738 pub fn moe_gate_up_silu8(&self, gp: WPtr8, up: WPtr8, x: &cudarc::driver::CudaView<f32>,
2739 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2740 rb_g: usize, rb_u: usize)
2741 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2742 let f = self.func("moe_gate_up_silu8_f32");
2743 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),
2745 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2746 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2747 let __s_b = self.gpu.stream();
2748 let mut b = __s_b.launch_builder(&f);
2749 b.arg(&gp).arg(&up).arg(x).arg(&mut act)
2750 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2751 unsafe { b.launch(cfg)?; }
2752 Ok(act)
2753 }
2754
2755 #[allow(clippy::too_many_arguments)]
2761 pub fn moe_down8_fma_into(&self, dp: WPtr8, w: F32x8, act: &CudaSlice<f32>,
2762 dst: &mut cudarc::driver::CudaViewMut<f32>,
2763 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2764 -> Result<(), Box<dyn std::error::Error>> {
2765 let f = self.func("moe_down8_fma_f32");
2766 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2767 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2768 let (inf, outf, nu, rbv) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2769 let __s_b = self.gpu.stream();
2770 let mut b = __s_b.launch_builder(&f);
2771 b.arg(&dp).arg(&w).arg(act).arg(dst).arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbv);
2772 unsafe { b.launch(cfg)?; }
2773 Ok(())
2774 }
2775
2776 #[allow(clippy::too_many_arguments)]
2781 #[allow(clippy::too_many_arguments)]
2796 #[allow(clippy::too_many_arguments)]
2798 pub fn moe_pairs_matvec_q8(&self, table: &CudaSlice<u64>, proj: i32,
2799 pair_tok: &CudaSlice<i32>, pair_ex: &CudaSlice<i32>,
2800 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2801 in_f: usize, out_f: usize, n_expert: usize, n_pairs: usize,
2802 qtype: i32, row_bytes: usize)
2803 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2804 let f = self.func("moe_pairs_matvec_q8");
2805 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2806 const ROWS: u32 = 4;
2807 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_pairs as u32, 1),
2808 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2809 let (inf, outf, ne, np, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2810 n_pairs as i32, row_bytes as i64);
2811 let __s_b = self.gpu.stream();
2812 let mut b = __s_b.launch_builder(&f);
2813 b.arg(table).arg(&proj).arg(pair_tok).arg(pair_ex).arg(aq).arg(ad).arg(&mut y)
2814 .arg(&inf).arg(&outf).arg(&ne).arg(&np).arg(&qtype).arg(&rbi);
2815 unsafe { b.launch(cfg)?; }
2816 Ok(y)
2817 }
2818
2819 #[allow(clippy::too_many_arguments)]
2821 pub fn moe_pairs_matvec_q8_em(&self, table: &CudaSlice<u64>, proj: i32,
2822 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2823 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2824 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2825 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2826 n_pairs: usize, qtype: i32, row_bytes: usize)
2827 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2828 let f = self.func("moe_pairs_matvec_q8_em");
2829 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2830 const ROWS: u32 = 4;
2831 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
2832 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2833 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2834 n_active as i32, row_bytes as i64);
2835 let __s_b = self.gpu.stream();
2836 let mut b = __s_b.launch_builder(&f);
2837 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
2838 .arg(aq).arg(ad).arg(&mut y)
2839 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
2840 unsafe { b.launch(cfg)?; }
2841 Ok(y)
2842 }
2843
2844 #[allow(clippy::too_many_arguments)]
2847 pub fn moe_pairs_matvec_q8_dec(&self, table: &CudaSlice<u64>, proj: i32,
2848 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2849 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2850 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2851 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2852 n_pairs: usize, qtype: i32, row_bytes: usize)
2853 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2854 let f = self.func("moe_pairs_matvec_q8_dec");
2855 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2856 const ROWS: u32 = 4;
2857 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
2858 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2859 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2860 n_active as i32, row_bytes as i64);
2861 let __s_b = self.gpu.stream();
2862 let mut b = __s_b.launch_builder(&f);
2863 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
2864 .arg(aq).arg(ad).arg(&mut y)
2865 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
2866 unsafe { b.launch(cfg)?; }
2867 Ok(y)
2868 }
2869
2870 pub fn moe_pairs_gelu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
2871 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2872 let f = self.func("moe_pairs_gelu_mul");
2873 let mut act = self.alloc_uninit::<f32>(n)?;
2874 let cfg = LaunchConfig::for_num_elems(n as u32);
2875 let nl = n as i64;
2876 let __s_b = self.gpu.stream();
2877 let mut b = __s_b.launch_builder(&f);
2878 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
2879 unsafe { b.launch(cfg)?; }
2880 Ok(act)
2881 }
2882
2883 pub fn moe_pairs_silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
2884 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2885 let f = self.func("moe_pairs_silu_mul");
2886 let mut act = self.alloc_uninit::<f32>(n)?;
2887 let cfg = LaunchConfig::for_num_elems(n as u32);
2888 let nl = n as i64;
2889 let __s_b = self.gpu.stream();
2890 let mut b = __s_b.launch_builder(&f);
2891 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
2892 unsafe { b.launch(cfg)?; }
2893 Ok(act)
2894 }
2895
2896 #[allow(clippy::too_many_arguments)]
2897 pub fn moe_pairs_scatter(&self, y_down: &CudaSlice<f32>, pair_w: &CudaSlice<f32>,
2898 tok_pair_off: &CudaSlice<i32>, tok_pair_ids: &CudaSlice<i32>,
2899 moe_out: &mut CudaSlice<f32>, t: usize, n_embd: usize)
2900 -> Result<(), Box<dyn std::error::Error>> {
2901 let f = self.func("moe_pairs_scatter");
2902 let cfg = LaunchConfig { grid_dim: (((n_embd + 255) / 256) as u32, t as u32, 1),
2903 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2904 let ne = n_embd as i32;
2905 let __s_b = self.gpu.stream();
2906 let mut b = __s_b.launch_builder(&f);
2907 b.arg(y_down).arg(pair_w).arg(tok_pair_off).arg(tok_pair_ids).arg(moe_out).arg(&ne);
2908 unsafe { b.launch(cfg)?; }
2909 Ok(())
2910 }
2911
2912 #[allow(clippy::too_many_arguments)]
2916 pub fn moe_gate_up_gelu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
2917 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2918 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
2919 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
2920 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2921 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
2922 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
2923 rb_g as i64, rb_u as i64);
2924 let f = self.func("moe_gate_up_gelu8_dev_q8");
2925 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2926 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2927 let __s_b = self.gpu.stream();
2928 let mut b = __s_b.launch_builder(&f);
2929 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
2930 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2931 unsafe { b.launch(cfg)?; }
2932 Ok(act)
2933 }
2934
2935 #[allow(clippy::too_many_arguments)]
2937 pub fn moe_gate_up_gelu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
2938 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
2939 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
2940 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
2941 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2942 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
2943 let (inf, nff, ne, rbg, rbu, nu) = (in_f as i32, n_ff as i32, n_expert as i32,
2944 rb_g as i64, rb_u as i64, n_used as i32);
2945 let f = self.func("moe_gate_up_gelu8_dev_q8_rows");
2946 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
2947 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2948 let __s_b = self.gpu.stream();
2949 let mut b = __s_b.launch_builder(&f);
2950 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
2951 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu);
2952 unsafe { b.launch(cfg)?; }
2953 Ok(act)
2954 }
2955
2956 #[allow(clippy::too_many_arguments)]
2958 pub fn moe_gate_up_gelu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
2959 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, n_pairs: usize,
2960 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
2961 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
2962 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2963 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
2964 let (inf, nff, ne, rbg, rbu, nu, npi) = (in_f as i32, n_ff as i32, n_expert as i32,
2965 rb_g as i64, rb_u as i64, n_used as i32,
2966 n_pairs as i32);
2967 let f = self.func("moe_gate_up_gelu8_dev_q8_csr");
2968 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
2969 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2970 let __s_b = self.gpu.stream();
2971 let mut b = __s_b.launch_builder(&f);
2972 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
2973 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
2974 unsafe { b.launch(cfg)?; }
2975 Ok(act)
2976 }
2977
2978 #[allow(clippy::too_many_arguments)]
2980 pub fn moe_down8_fma_dev_q8_rows_g(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
2981 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
2982 dst: &mut CudaSlice<f32>, t: usize,
2983 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
2984 qt: i32, rb: usize)
2985 -> Result<(), Box<dyn std::error::Error>> {
2986 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
2987 n_expert as i32, rb as i64);
2988 let f = self.func("moe_down8_fma_dev_q8_rows_g");
2989 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, t as u32),
2990 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2991 let __s_b = self.gpu.stream();
2992 let mut b = __s_b.launch_builder(&f);
2993 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
2994 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
2995 unsafe { b.launch(cfg)?; }
2996 Ok(())
2997 }
2998
2999 pub fn rp_probe_q4(&self, m: usize) -> Result<(f64, f64), Box<dyn std::error::Error>> {
3003 let (out_f, in_f) = (2048usize, 2816usize);
3004 let nblk = in_f / 32;
3005 let mut seed = 0x9E3779B97F4A7C15u64;
3006 let mut rng = move || { seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407); (seed >> 33) as u8 };
3007 let mut w = vec![0u8; out_f * nblk * 18];
3008 for b in w.iter_mut() { *b = rng(); }
3009 for r in 0..out_f {
3010 for g in 0..nblk {
3011 let off = (r * nblk + g) * 18;
3012 w[off] = 0x00; w[off + 1] = 0x2C; }
3014 }
3015 let qplane = out_f * nblk * 16;
3016 let mut wrp = vec![0u8; w.len()];
3017 for r in 0..out_f {
3018 for g in 0..nblk {
3019 let src = &w[(r * nblk + g) * 18..(r * nblk + g) * 18 + 18];
3020 wrp[qplane + (r * nblk + g) * 2..qplane + (r * nblk + g) * 2 + 2]
3021 .copy_from_slice(&src[0..2]);
3022 wrp[(r * nblk + g) * 16..(r * nblk + g) * 16 + 16].copy_from_slice(&src[2..18]);
3023 }
3024 }
3025 let w_d = self.htod_bytes(&w)?;
3026 let wrp_d = self.htod_bytes(&wrp)?;
3027 let mut aq = vec![0i8; m * in_f];
3028 for v in aq.iter_mut() { *v = rng() as i8; }
3029 let aq_d = self.htod_i8(&aq)?;
3030 let ad_d = self.htod(&vec![0.03125f32; m * nblk])?;
3031 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
3032 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
3033 const RPB: u32 = 4;
3034 let cfg = LaunchConfig { grid_dim: ((out_f as u32).div_ceil(RPB), 1, 1),
3035 block_dim: (32, RPB, 1), shared_mem_bytes: 0 };
3036 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
3037 let (rb, qp) = ((nblk * 18) as i64, qplane as i64);
3038 let fb = self.func("qmatvec_q4_0_mmvq_b4");
3039 let fr = self.func("qmatvec_q4_0_mmvq_b4_rp");
3040 {
3041 let __s_b = self.gpu.stream();
3042 let mut b = __s_b.launch_builder(&fb);
3043 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3044 unsafe { b.launch(cfg)?; }
3045 let __s_b = self.gpu.stream();
3046 let mut b = __s_b.launch_builder(&fr);
3047 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1).arg(&inf).arg(&outf).arg(&mi).arg(&qp);
3048 unsafe { b.launch(cfg)?; }
3049 }
3050 self.gpu.stream().synchronize()?;
3051 let (h0, h1) = (self.dtoh(&y0)?, self.dtoh(&y1)?);
3052 let nd = h0.iter().zip(&h1).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
3053 if nd != 0 { return Err(format!("rp twin not bitwise: {nd}/{} diffs", h0.len()).into()); }
3054 let mut time = |rp: bool| -> Result<f64, Box<dyn std::error::Error>> {
3055 self.gpu.stream().synchronize()?;
3056 let t0 = std::time::Instant::now();
3057 for _ in 0..500 {
3058 if rp {
3059 let __s_b = self.gpu.stream();
3060 let mut b = __s_b.launch_builder(&fr);
3061 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1)
3062 .arg(&inf).arg(&outf).arg(&mi).arg(&qp);
3063 unsafe { b.launch(cfg)?; }
3064 } else {
3065 let __s_b = self.gpu.stream();
3066 let mut b = __s_b.launch_builder(&fb);
3067 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0)
3068 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3069 unsafe { b.launch(cfg)?; }
3070 }
3071 }
3072 self.gpu.stream().synchronize()?;
3073 Ok(t0.elapsed().as_secs_f64() * 1e6 / 500.0)
3074 };
3075 let _ = time(false)?; let _ = time(true)?; Ok((time(false)?, time(true)?))
3077 }
3078
3079 pub fn build_q4_rp4(&self, t: &mut crate::model::GpuTensor)
3084 -> Result<(), Box<dyn std::error::Error>> {
3085 use crate::model::GpuTensor;
3086 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3087 if *qtype != QT_Q4_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3088 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3089 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 18 { return Ok(()); }
3090 let nblk = in_f / 32;
3091 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 18)?;
3092 let f = self.func("q4_0_split_rp_build");
3093 let n = (out_f * nblk) as i32;
3094 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3095 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3096 let (of, nb) = (out_f as i32, nblk as i32);
3097 let _ = n;
3098 let __s_b = self.gpu.stream();
3099 let mut b = __s_b.launch_builder(&f);
3100 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3101 unsafe { b.launch(cfg)?; }
3102 *rp4 = Some(dst);
3103 Ok(())
3104 }
3105
3106 pub fn build_q8_rp4(&self, t: &mut crate::model::GpuTensor)
3111 -> Result<(), Box<dyn std::error::Error>> {
3112 use crate::model::GpuTensor;
3113 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3114 if *qtype != QT_Q8_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3115 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3116 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 34 { return Ok(()); }
3117 *rp4 = Some(self.build_q8_rp4_raw(bytes, in_f, out_f)?);
3118 Ok(())
3119 }
3120
3121 pub fn build_q8_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize)
3124 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3125 assert!(in_f % 32 == 0);
3126 let nblk = in_f / 32;
3127 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 34)?;
3128 let f = self.func("q8_0_split_rp_build");
3129 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3130 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3131 let (of, nb) = (out_f as i32, nblk as i32);
3132 let __s_b = self.gpu.stream();
3133 let mut b = __s_b.launch_builder(&f);
3134 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3135 unsafe { b.launch(cfg)?; }
3136 Ok(dst)
3137 }
3138
3139 pub fn build_q4k_rp4(&self, t: &mut crate::model::GpuTensor)
3147 -> Result<(), Box<dyn std::error::Error>> {
3148 use crate::model::GpuTensor;
3149 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3150 if *qtype != QT_Q4_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3151 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3152 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 144 { return Ok(()); }
3153 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q4_K)?);
3154 Ok(())
3155 }
3156
3157 pub fn build_q6k_rp4(&self, t: &mut crate::model::GpuTensor)
3158 -> Result<(), Box<dyn std::error::Error>> {
3159 use crate::model::GpuTensor;
3160 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3161 if *qtype != QT_Q6_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3162 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3163 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 210 { return Ok(()); }
3164 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q6_K)?);
3165 Ok(())
3166 }
3167
3168 pub fn build_kq_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize, qtype: i32)
3170 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3171 assert!(in_f % 256 == 0);
3172 let nsbk = in_f / 256;
3173 let (sb_bytes, kname) = match qtype {
3174 QT_Q4_K => (144usize, "q4_K_split_rp_build"),
3175 QT_Q6_K => (210usize, "q6_K_split_rp_build"),
3176 _ => return Err(format!("build_kq_rp4_raw: qtype {qtype} has no rp mirror").into()),
3177 };
3178 let mut dst = self.alloc_uninit::<u8>(out_f * nsbk * sb_bytes)?;
3179 let f = self.func(kname);
3180 let cfg = LaunchConfig { grid_dim: (((out_f * nsbk) as u32).div_ceil(256), 1, 1),
3181 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3182 let (of, nb) = (out_f as i32, nsbk as i32);
3183 let __s_b = self.gpu.stream();
3184 let mut b = __s_b.launch_builder(&f);
3185 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3186 unsafe { b.launch(cfg)?; }
3187 Ok(dst)
3188 }
3189
3190 pub fn kqrp_enabled() -> bool {
3194 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3195 *ON.get_or_init(|| match std::env::var("MEMRA_KQRP").as_deref() {
3196 Ok("0") => false,
3197 Ok(_) => true,
3198 Err(_) => cfg!(memra_hopper_mma),
3199 })
3200 }
3201
3202 pub fn build_q4_rp_swap(&self, t: &mut crate::model::GpuTensor)
3208 -> Result<bool, Box<dyn std::error::Error>> {
3209 self.build_q4_rp4(t)?;
3210 self.gpu.stream().synchronize()?; use crate::model::GpuTensor;
3212 let GpuTensor::Quant { bytes, rp4, rp, .. } = t else { return Ok(false) };
3213 match rp4.take() {
3214 Some(split) => {
3215 *bytes = split; *rp = true;
3217 Ok(true)
3218 }
3219 None => Ok(false),
3220 }
3221 }
3222
3223 pub fn q4rp_enabled() -> bool {
3225 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3226 *ON.get_or_init(|| std::env::var("MEMRA_Q4RP").map(|v| v != "0").unwrap_or(true))
3227 }
3228
3229 pub fn copy_rows_strided(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
3232 row_elems: usize, n_rows: usize, src_stride: usize, src_off: usize)
3233 -> Result<(), Box<dyn std::error::Error>> {
3234 let f = self.func("copy_rows_strided_f32");
3235 let cfg = LaunchConfig { grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
3236 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3237 let (re, nr) = (row_elems as i32, n_rows as i32);
3238 let (st, off) = (src_stride as i64, src_off as i64);
3239 let __s_b = self.gpu.stream();
3240 let mut b = __s_b.launch_builder(&f);
3241 b.arg(src).arg(&mut *dst).arg(&re).arg(&nr).arg(&st).arg(&off);
3242 unsafe { b.launch(cfg)?; }
3243 Ok(())
3244 }
3245
3246 pub fn u32_set_k(&self, dst: &mut CudaSlice<u32>, v: u32, idx: usize)
3248 -> Result<(), Box<dyn std::error::Error>> {
3249 let f = self.func("u32_set_k");
3250 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3251 let ii = idx as i32;
3252 let __s_b = self.gpu.stream();
3253 let mut b = __s_b.launch_builder(&f);
3254 b.arg(dst).arg(&v).arg(&ii);
3255 unsafe { b.launch(cfg)?; }
3256 Ok(())
3257 }
3258
3259 pub fn i32_add_k(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
3261 let f = self.func("i32_add_k");
3262 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3263 let __s_b = self.gpu.stream();
3264 let mut b = __s_b.launch_builder(&f);
3265 b.arg(d).arg(&v);
3266 unsafe { b.launch(cfg)?; }
3267 Ok(())
3268 }
3269
3270 pub fn i32_iota_from(&self, ctr: &CudaSlice<i32>, dst: &mut CudaSlice<i32>, n: usize)
3272 -> Result<(), Box<dyn std::error::Error>> {
3273 let f = self.func("i32_iota_from");
3274 let cfg = LaunchConfig::for_num_elems(n as u32);
3275 let ni = n as i32;
3276 let __s_b = self.gpu.stream();
3277 let mut b = __s_b.launch_builder(&f);
3278 b.arg(ctr).arg(dst).arg(&ni);
3279 unsafe { b.launch(cfg)?; }
3280 Ok(())
3281 }
3282
3283 pub fn u32_map_k(&self, buf: &mut CudaSlice<u32>, map: &CudaSlice<u32>, idx: usize)
3285 -> Result<(), Box<dyn std::error::Error>> {
3286 let f = self.func("u32_map_k");
3287 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3288 let ii = idx as i32;
3289 let __s_b = self.gpu.stream();
3290 let mut b = __s_b.launch_builder(&f);
3291 b.arg(buf).arg(map).arg(&ii);
3292 unsafe { b.launch(cfg)?; }
3293 Ok(())
3294 }
3295
3296 #[allow(clippy::too_many_arguments)]
3298 pub fn u32_pack2(&self, a: &CudaSlice<u32>, off_a: usize, n1: usize,
3299 b_in: &CudaSlice<u32>, n2: usize, out: &mut CudaSlice<u32>)
3300 -> Result<(), Box<dyn std::error::Error>> {
3301 let f = self.func("u32_pack2");
3302 let cfg = LaunchConfig::for_num_elems((n1 + n2) as u32);
3303 let (oa, i1, i2) = (off_a as i32, n1 as i32, n2 as i32);
3304 let __s_b = self.gpu.stream();
3305 let mut b = __s_b.launch_builder(&f);
3306 b.arg(a).arg(&oa).arg(&i1).arg(b_in).arg(&i2).arg(out);
3307 unsafe { b.launch(cfg)?; }
3308 Ok(())
3309 }
3310
3311 pub fn moe_w_exscale(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3313 s: &CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
3314 let f = self.func("moe_w_exscale");
3315 let cfg = LaunchConfig::for_num_elems(n as u32);
3316 let ni = n as i32;
3317 let __s_b = self.gpu.stream();
3318 let mut b = __s_b.launch_builder(&f);
3319 b.arg(w).arg(sel).arg(s).arg(&ni);
3320 unsafe { b.launch(cfg)?; }
3321 Ok(())
3322 }
3323
3324 pub fn moe_w_scale_by_expert(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3327 macros: &CudaSlice<f32>, n_expert: usize, n: usize)
3328 -> Result<(), Box<dyn std::error::Error>> {
3329 let f = self.func("moe_w_scale_by_expert");
3330 let cfg = LaunchConfig { grid_dim: (n.div_ceil(64) as u32, 1, 1),
3331 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
3332 let (ne, nn) = (n_expert as i32, n as i32);
3333 let __s_b = self.gpu.stream();
3334 let mut b = __s_b.launch_builder(&f);
3335 b.arg(w).arg(sel).arg(macros).arg(&ne).arg(&nn);
3336 unsafe { b.launch(cfg)?; }
3337 Ok(())
3338 }
3339
3340 pub fn moe_gate_up_silu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3341 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3342 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3343 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3344 macros: &CudaSlice<f32>)
3345 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3346 static GU: std::sync::OnceLock<(String, u32)> = std::sync::OnceLock::new();
3347 let (mode, wpb) = GU.get_or_init(|| {
3348 let mode = std::env::var("MEMRA_MOE_DEVQ8_GU").unwrap_or_default();
3349 let wpb = std::env::var("MEMRA_MOE_DEVQ8_WPB").ok()
3350 .and_then(|v| v.parse().ok()).unwrap_or(4u32).clamp(1, 16);
3351 (mode, wpb)
3352 });
3353 let (mode, wpb) = (mode.as_str(), *wpb);
3354 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3355 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3356 rb_g as i64, rb_u as i64);
3357 let (f, cfg) = match mode {
3358 "1" | "2" | "4" => {
3359 let rpw: u32 = mode.parse().unwrap();
3360 let f = self.func(match rpw { 1 => "moe_gate_up_silu8_dev_q8_r1",
3361 2 => "moe_gate_up_silu8_dev_q8_r2",
3362 _ => "moe_gate_up_silu8_dev_q8_r4" });
3363 let rows_per_block = (rpw * wpb) as usize;
3364 let gx = n_ff.div_ceil(rows_per_block) as u32;
3365 (f, LaunchConfig { grid_dim: (gx, n_used as u32, 1),
3366 block_dim: (32, wpb, 1), shared_mem_bytes: 0 })
3367 }
3368 "j8" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8"),
3369 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3370 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3371 "vsm2" => {
3373 let f = self.func("moe_gate_up_silu8_dev_q8_vsm2");
3374 let sh = (rb_g + rb_u) as u32;
3375 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3376 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3377 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3378 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3379 }
3380 "vsm" => {
3381 let f = self.func("moe_gate_up_silu8_dev_q8_vsm");
3382 let sh = (rb_g + rb_u) as u32;
3383 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3384 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3385 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3386 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3387 }
3388 "sg" => (self.func("moe_gate_up_silu8_dev_q8_sg"),
3389 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3390 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3391 "j8sg" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8sg"),
3392 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3393 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3394 "u64" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_u64"),
3395 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3396 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3397 "gs4" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_gs4"),
3398 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3399 block_dim: (32, 4, 1), shared_mem_bytes: 0 }),
3400 "v" | "" => (self.func("moe_gate_up_silu8_dev_q8_v"),
3402 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3403 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3404 "s2" => (self.func("moe_gate_up_silu8_dev_q8_s2"),
3405 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3406 block_dim: (32, 2, 1), shared_mem_bytes: 0 }),
3407 "s2z" => {
3408 let rz = wpb.min(16); (self.func("moe_gate_up_silu8_dev_q8_s2z"),
3410 LaunchConfig { grid_dim: (n_ff.div_ceil(rz as usize) as u32, n_used as u32, 1),
3411 block_dim: (32, 2, rz), shared_mem_bytes: 0 })
3412 }
3413 _ => (self.func("moe_gate_up_silu8_dev_q8"),
3414 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3415 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3416 };
3417 let __s_b = self.gpu.stream();
3418 let mut b = __s_b.launch_builder(&f);
3419 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3420 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3421 unsafe { b.launch(cfg)?; }
3422 Ok(act)
3423 }
3424
3425 #[allow(clippy::too_many_arguments)]
3426 pub fn moe_down8_fma_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3427 w: &cudarc::driver::CudaView<f32>,
3428 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3429 dst: &mut cudarc::driver::CudaViewMut<f32>,
3430 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3431 qt: i32, rb: usize)
3432 -> Result<(), Box<dyn std::error::Error>> {
3433 static DOWN: std::sync::OnceLock<String> = std::sync::OnceLock::new();
3434 let mode = DOWN.get_or_init(|| std::env::var("MEMRA_MOE_DEVQ8_DOWN").unwrap_or_default());
3435 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3436 n_expert as i32, rb as i64);
3437 let (f, cfg) = match mode.as_str() {
3440 m @ ("1" | "2" | "4") if n_used <= 8 => {
3441 let rpw: usize = m.parse().unwrap();
3442 let f = self.func(match rpw { 1 => "moe_down8_fma_dev_q8_w8r1",
3443 2 => "moe_down8_fma_dev_q8_w8r2",
3444 _ => "moe_down8_fma_dev_q8_w8r4" });
3445 (f, LaunchConfig { grid_dim: (out_f.div_ceil(rpw) as u32, 1, 1),
3446 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3447 }
3448 "h2" if in_f == 512 => (self.func("moe_down8_fma_dev_q8_h2"),
3449 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3450 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3451 "" if in_f == 704 && n_used <= 8 =>
3454 (self.func("moe_down8_fma_dev_q8_w8r2"),
3455 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3456 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3457 "w8h2v" | "" if in_f == 512 && n_used <= 8 =>
3461 (self.func("moe_down8_fma_dev_q8_w8h2v"),
3462 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3463 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3464 "w8h2r2v" if in_f == 512 && n_used <= 8 =>
3465 (self.func("moe_down8_fma_dev_q8_w8h2r2v"),
3466 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3467 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3468 "w8h2r2" if in_f == 512 && n_used <= 8 =>
3469 (self.func("moe_down8_fma_dev_q8_w8h2r2"),
3470 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3471 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3472 "w8h2" if in_f == 512 && n_used <= 8 =>
3473 (self.func("moe_down8_fma_dev_q8_w8h2"),
3474 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3475 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3476 _ => (self.func("moe_down8_fma_dev_q8"),
3477 LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3478 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3479 };
3480 let __s_b = self.gpu.stream();
3481 let mut b = __s_b.launch_builder(&f);
3482 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3483 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3484 unsafe { b.launch(cfg)?; }
3485 Ok(())
3486 }
3487
3488 #[allow(clippy::too_many_arguments)]
3495 pub fn moe_gate_up_silu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3496 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
3497 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3498 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3499 macros: &CudaSlice<f32>)
3500 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3501 let f = self.func("moe_gate_up_silu8_dev_q8_v_rows");
3502 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
3503 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
3504 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3505 let (inf, nff, ne, nu, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3506 n_used as i32, rb_g as i64, rb_u as i64);
3507 let __s_b = self.gpu.stream();
3508 let mut b = __s_b.launch_builder(&f);
3509 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3510 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(macros);
3511 unsafe { b.launch(cfg)?; }
3512 Ok(act)
3513 }
3514
3515 #[allow(clippy::too_many_arguments)]
3520 pub fn moe_down8_fma_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3521 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3522 dst: &mut CudaSlice<f32>, t: usize,
3523 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3524 qt: i32, rb: usize)
3525 -> Result<(), Box<dyn std::error::Error>> {
3526 assert!(in_f == 512 && n_used <= 8, "down rows twin is w8h2v shape-gated");
3527 let f = self.func("moe_down8_fma_dev_q8_w8h2v_rows");
3528 let cfg = LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, t as u32),
3529 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 };
3530 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3531 n_expert as i32, rb as i64);
3532 let __s_b = self.gpu.stream();
3533 let mut b = __s_b.launch_builder(&f);
3534 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3535 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3536 unsafe { b.launch(cfg)?; }
3537 Ok(())
3538 }
3539
3540 #[allow(clippy::too_many_arguments)]
3544 pub fn moe_gate_up_silu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3545 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3546 n_pairs: usize, in_f: usize, n_ff: usize, n_used: usize,
3547 n_expert: usize, qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3548 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3549 let f = self.func("moe_gate_up_silu8_dev_q8_csr_iq4");
3550 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
3551 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
3552 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3553 let (inf, nff, ne, nu, npi, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3554 n_used as i32, n_pairs as i32, rb_g as i64, rb_u as i64);
3555 let __s_b = self.gpu.stream();
3556 let mut b = __s_b.launch_builder(&f);
3557 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3558 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
3559 unsafe { b.launch(cfg)?; }
3560 Ok(act)
3561 }
3562
3563
3564 #[allow(clippy::too_many_arguments)]
3568 pub fn moe_down8_fma_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3569 sel: &cudarc::driver::CudaView<i32>,
3570 w: &cudarc::driver::CudaView<f32>,
3571 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3572 dst: &mut cudarc::driver::CudaViewMut<f32>,
3573 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3574 qt: i32, rb: usize)
3575 -> Result<(), Box<dyn std::error::Error>> {
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 (f, cfg) = match variant {
3579 "w8h2" | "w8h2v" => {
3580 (self.func(if variant == "w8h2" { "moe_down8_fma_dev_q8_w8h2" }
3581 else { "moe_down8_fma_dev_q8_w8h2v" }),
3582 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3583 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3584 }
3585 "w8h2r2" | "w8h2r2v" => {
3586 (self.func(if variant == "w8h2r2" { "moe_down8_fma_dev_q8_w8h2r2" }
3587 else { "moe_down8_fma_dev_q8_w8h2r2v" }),
3588 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3589 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3590 }
3591 _ => (self.func("moe_down8_fma_dev_q8"),
3592 LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3593 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3594 };
3595 let __s_b = self.gpu.stream();
3596 let mut b = __s_b.launch_builder(&f);
3597 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3598 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3599 unsafe { b.launch(cfg)?; }
3600 Ok(())
3601 }
3602
3603 #[allow(clippy::too_many_arguments)]
3605 pub fn moe_gate_up_silu8_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3606 sel: &cudarc::driver::CudaView<i32>,
3607 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3608 in_f: usize, n_ff: usize, n_used: usize,
3609 n_expert: usize, qt_g: i32, qt_u: i32,
3610 rb_g: usize, rb_u: usize)
3611 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3612 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3613 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3614 rb_g as i64, rb_u as i64);
3615 let f = self.func(if variant == "v" { "moe_gate_up_silu8_dev_q8_v" }
3616 else { "moe_gate_up_silu8_dev_q8" });
3617 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3618 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3619 let __s_b = self.gpu.stream();
3620 let mut b = __s_b.launch_builder(&f);
3621 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3622 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
3623 unsafe { b.launch(cfg)?; }
3624 Ok(act)
3625 }
3626
3627 pub fn moe_gate_up_silu8_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3628 x: &cudarc::driver::CudaView<f32>,
3629 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3630 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3631 macros: &CudaSlice<f32>)
3632 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3633 let f = self.func("moe_gate_up_silu8_dev");
3634 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),
3636 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3637 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3638 rb_g as i64, rb_u as i64);
3639 let __s_b = self.gpu.stream();
3640 let mut b = __s_b.launch_builder(&f);
3641 b.arg(table).arg(sel).arg(x).arg(&mut act)
3642 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3643 unsafe { b.launch(cfg)?; }
3644 Ok(act)
3645 }
3646
3647 #[allow(clippy::too_many_arguments)]
3650 pub fn moe_down8_fma_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3651 w: &cudarc::driver::CudaView<f32>, act: &CudaSlice<f32>,
3652 dst: &mut cudarc::driver::CudaViewMut<f32>,
3653 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3654 qt: i32, rb: usize)
3655 -> Result<(), Box<dyn std::error::Error>> {
3656 let f = self.func("moe_down8_fma_dev");
3657 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3658 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3659 let (inf, outf, nu, ne, rbv) = (in_f as i32, out_f as i32, n_used as i32,
3660 n_expert as i32, rb as i64);
3661 let __s_b = self.gpu.stream();
3662 let mut b = __s_b.launch_builder(&f);
3663 b.arg(table).arg(sel).arg(w).arg(act).arg(dst)
3664 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbv);
3665 unsafe { b.launch(cfg)?; }
3666 Ok(())
3667 }
3668
3669 pub fn axpy_into(&self, src: &CudaSlice<f32>, alpha: f32,
3671 dst: &mut cudarc::driver::CudaViewMut<f32>, n: usize)
3672 -> Result<(), Box<dyn std::error::Error>> {
3673 let f = self.func("axpy_f32");
3674 let cfg = LaunchConfig::for_num_elems(n as u32);
3675 let (a, ni) = (alpha, n as i32);
3676 let __s_b = self.gpu.stream();
3677 let mut b = __s_b.launch_builder(&f);
3678 b.arg(src).arg(dst).arg(&a).arg(&ni);
3679 unsafe { b.launch(cfg)?; }
3680 Ok(())
3681 }
3682
3683 pub fn add_scaled_rows(&self, src: &CudaSlice<f32>, scale: &CudaSlice<f32>,
3685 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
3686 -> Result<(), Box<dyn std::error::Error>> {
3687 let f = self.func("add_scaled_rows_f32");
3688 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
3689 let (nc, nr) = (ncols as i32, nrows as i32);
3690 let __s_b = self.gpu.stream();
3691 let mut b = __s_b.launch_builder(&f);
3692 b.arg(src).arg(scale).arg(dst).arg(&nc).arg(&nr);
3693 unsafe { b.launch(cfg)?; }
3694 Ok(())
3695 }
3696
3697 pub fn gather_rows(&self, src: &CudaSlice<f32>, idx: &CudaSlice<i32>,
3701 dst: &mut CudaSlice<f32>, ncols: usize, m_e: usize)
3702 -> Result<(), Box<dyn std::error::Error>> {
3703 let f = self.func("gather_rows_f32");
3704 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3705 let (nc, me) = (ncols as i32, m_e as i32);
3706 let __s_b = self.gpu.stream();
3707 let mut b = __s_b.launch_builder(&f);
3708 b.arg(src).arg(idx).arg(dst).arg(&nc).arg(&me);
3709 unsafe { b.launch(cfg)?; }
3710 Ok(())
3711 }
3712
3713 pub fn scatter_slot(&self, src: &CudaSlice<f32>, tok_idx: &CudaSlice<i32>,
3718 slot_idx: &CudaSlice<i32>, weight: &CudaSlice<f32>,
3719 dst: &mut CudaSlice<f32>, wbuf: &mut CudaSlice<f32>,
3720 ncols: usize, n_used: usize, m_e: usize)
3721 -> Result<(), Box<dyn std::error::Error>> {
3722 let f = self.func("scatter_add_slot_f32");
3723 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3724 let (nc, nu, me) = (ncols as i32, n_used as i32, m_e as i32);
3725 let __s_b = self.gpu.stream();
3726 let mut b = __s_b.launch_builder(&f);
3727 b.arg(src).arg(tok_idx).arg(slot_idx).arg(weight).arg(dst).arg(wbuf).arg(&nc).arg(&nu).arg(&me);
3728 unsafe { b.launch(cfg)?; }
3729 Ok(())
3730 }
3731
3732 pub fn reduce_slots(&self, slots: &CudaSlice<f32>, wbuf: &CudaSlice<f32>,
3736 dst: &mut CudaSlice<f32>, ncols: usize, n_used: usize, t: usize)
3737 -> Result<(), Box<dyn std::error::Error>> {
3738 let f = self.func("reduce_slots_f32");
3739 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
3740 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
3741 let __s_b = self.gpu.stream();
3742 let mut b = __s_b.launch_builder(&f);
3743 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
3744 unsafe { b.launch(cfg)?; }
3745 Ok(())
3746 }
3747
3748 pub fn quantize_q8_1_view(&self, x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize)
3755 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3756 let f = self.func("quantize_q8_1");
3757 let nblk = in_f / 32;
3758 let mut q = self.alloc_uninit::<i8>(m * in_f)?;
3759 let mut d = self.alloc_uninit::<f32>(m * nblk)?;
3760 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
3761 let (inf, mi) = (in_f as i32, m as i32);
3762 let __s_b = self.gpu.stream();
3763 let mut b = __s_b.launch_builder(&f);
3764 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3765 unsafe { b.launch(cfg)?; }
3766 Ok((q, d))
3767 }
3768
3769 pub fn quantize_q8_1(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3770 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3771 let nblk = in_f / 32;
3772 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);
3776 let (inf, mi) = (in_f as i32, m as i32);
3777 if Self::pdl_on() && Self::pdl_wb_on() {
3778 {
3779 use cudarc::driver::{DevicePtr, DevicePtrMut};
3780 let s = &self.gpu.stream();
3781 let (px, _g0) = x.device_ptr(s);
3782 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
3783 let mut ps = [
3784 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
3785 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
3786 &mi as *const _ as *mut _,
3787 ];
3788 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
3789 }
3790 return Ok((q, d));
3791 }
3792 let f = self.func("quantize_q8_1");
3793 let __s_b = self.gpu.stream();
3794 let mut b = __s_b.launch_builder(&f);
3795 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3796 unsafe { b.launch(cfg)?; }
3797 Ok((q, d))
3798 }
3799
3800 pub fn quantize_fp4_act(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3804 -> Result<(CudaSlice<u32>, CudaSlice<u8>), Box<dyn std::error::Error>> {
3805 let f = self.func("quantize_fp4_act");
3806 let nb16 = in_f / 16;
3807 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);
3810 let (inf, mi) = (in_f as i32, m as i32);
3811 let __s_b = self.gpu.stream();
3812 let mut b = __s_b.launch_builder(&f);
3813 b.arg(x).arg(&mut aq4).arg(&mut ad4).arg(&inf).arg(&mi);
3814 unsafe { b.launch(cfg)?; }
3815 Ok((aq4, ad4))
3816 }
3817
3818 pub fn qmatvec_gemm_nvfp4_fp4(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3823 in_f: usize, out_f: usize, row_bytes: usize, scale: f32)
3824 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3825 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
3826 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
3827 let mut y = self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)?;
3828 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
3829 Ok(y)
3830 }
3831
3832 fn fp4_gemm_launch(&self, bytes: &CudaSlice<u8>, aq4: &CudaSlice<u32>, ad4: &CudaSlice<u8>,
3835 m: usize, in_f: usize, out_f: usize, row_bytes: usize)
3836 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3837 let f = self.func("qmatvec_gemm_nvfp4_fp4");
3838 let mut y = self.alloc_uninit::<f32>(m * out_f)?; const BM: u32 = 64; const BN: u32 = 256;
3840 let cfg = LaunchConfig {
3841 grid_dim: ((out_f as u32 + BM - 1) / BM, (m as u32 + BN - 1) / BN, 1),
3842 block_dim: (32, 4, 1), shared_mem_bytes: 0,
3843 };
3844 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3845 let __s_b = self.gpu.stream();
3846 let mut b = __s_b.launch_builder(&f);
3847 b.arg(bytes).arg(aq4).arg(ad4).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3848 unsafe { b.launch(cfg)?; }
3849 Ok(y)
3850 }
3851
3852 pub fn qmatvec_gemm_nvfp4_fp4_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3854 in_f: usize, out_f: usize, row_bytes: usize)
3855 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3856 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
3857 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
3858 self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)
3859 }
3860
3861 pub fn qmatvec_q8_0_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3863 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3864 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3865 let f = self.func("qmatvec_q8_0_dp4a");
3866 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 };
3868 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3869 let __s_b = self.gpu.stream();
3870 let mut b = __s_b.launch_builder(&f);
3871 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3872 unsafe { b.launch(cfg)?; }
3873 Ok(y)
3874 }
3875
3876 #[allow(non_snake_case)] pub fn qmatvec_q4_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3879 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3880 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3881 let f = self.func("qmatvec_q4_K_dp4a");
3882 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 };
3884 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3885 let __s_b = self.gpu.stream();
3886 let mut b = __s_b.launch_builder(&f);
3887 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3888 unsafe { b.launch(cfg)?; }
3889 Ok(y)
3890 }
3891
3892 #[allow(non_snake_case)] pub fn qmatvec_q6_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3895 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3896 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3897 let f = self.func("qmatvec_q6_K_dp4a");
3898 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 };
3900 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3901 let __s_b = self.gpu.stream();
3902 let mut b = __s_b.launch_builder(&f);
3903 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3904 unsafe { b.launch(cfg)?; }
3905 Ok(y)
3906 }
3907
3908 #[allow(non_snake_case)] pub fn qmatvec_q5_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3911 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3912 self.qmatvec_dp4a_named("qmatvec_q5_K_dp4a", w, x, m, in_f, out_f, row_bytes)
3913 }
3914 #[allow(non_snake_case)] pub fn qmatvec_q3_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3917 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3918 self.qmatvec_dp4a_named("qmatvec_q3_K_dp4a", w, x, m, in_f, out_f, row_bytes)
3919 }
3920 pub fn qmatvec_nvfp4_fast_rp(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3922 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3923 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
3924 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a_rp", w, x, m, in_f, out_f, row_bytes)
3925 }
3926 pub fn qmatvec_nvfp4_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3928 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3929 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
3932 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a", w, x, m, in_f, out_f, row_bytes)
3933 }
3934 #[allow(non_snake_case)] pub fn qmatvec_iq4_XS_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3937 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3938 self.qmatvec_dp4a_named("qmatvec_iq4_XS_dp4a", w, x, m, in_f, out_f, row_bytes)
3939 }
3940
3941 fn qmatvec_dp4a_named(&self, name: &str, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3943 in_f: usize, out_f: usize, row_bytes: usize)
3944 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3945 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
3946 let f = self.func(name);
3947 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 };
3949 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3950 let __s_b = self.gpu.stream();
3951 let mut b = __s_b.launch_builder(&f);
3952 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3953 unsafe { b.launch(cfg)?; }
3954 Ok(y)
3955 }
3956
3957 pub fn htod(&self, v: &[f32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3958 Ok(self.gpu.stream().clone_htod(v)?)
3959 }
3960 pub fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
3961 Ok(self.gpu.stream().clone_htod(v)?)
3962 }
3963 pub fn htod_i8(&self, v: &[i8]) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
3965 Ok(self.gpu.stream().clone_htod(v)?)
3966 }
3967 pub fn htod_u64(&self, v: &[u64]) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
3968 Ok(self.gpu.stream().clone_htod(v)?)
3969 }
3970 pub fn dtoh_view(&self, d: &cudarc::driver::CudaView<f32>)
3972 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
3973 let v = self.gpu.stream().clone_dtoh(d)?;
3974 self.gpu.stream().synchronize()?;
3975 Ok(v)
3976 }
3977 pub fn dtoh(&self, d: &CudaSlice<f32>) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
3978 let v = self.gpu.stream().clone_dtoh(d)?;
3979 self.gpu.stream().synchronize()?;
3980 Ok(v)
3981 }
3982 pub fn dtoh_pair(
3986 &self,
3987 a: &CudaSlice<f32>,
3988 b: &CudaSlice<f32>,
3989 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
3990 let av = self.gpu.stream().clone_dtoh(a)?;
3991 let bv = self.gpu.stream().clone_dtoh(b)?;
3992 self.gpu.stream().synchronize()?;
3993 Ok((av, bv))
3994 }
3995 pub fn dtoh_i32(&self, d: &CudaSlice<i32>) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
3997 let v = self.gpu.stream().clone_dtoh(d)?;
3998 self.gpu.stream().synchronize()?;
3999 Ok(v)
4000 }
4001 pub fn dtoh_u8(&self, d: &CudaSlice<u8>) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
4003 let v = self.gpu.stream().clone_dtoh(d)?;
4004 self.gpu.stream().synchronize()?;
4005 Ok(v)
4006 }
4007 pub fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4008 let s = self.gpu.stream().alloc_zeros::<f32>(n)?;
4009 self.keep_if_capturing(&s);
4010 Ok(s)
4011 }
4012
4013 pub fn prob_of_token_device(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>, n_vocab: usize)
4022 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4023 let nb = ARGMAX_NB;
4024 let mut part = self.alloc_uninit::<f32>(nb)?;
4025 let mut p = self.alloc_uninit::<f32>(1)?;
4026 let f1 = self.func("prob_of_token_partial_f32");
4027 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4028 let nv = n_vocab as i32;
4029 let __s_b1 = self.gpu.stream();
4030 let mut b1 = __s_b1.launch_builder(&f1);
4031 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
4032 unsafe { b1.launch(cfg1)?; }
4033 let f2 = self.func("prob_of_token_final_f32");
4034 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4035 let nbi = nb as i32;
4036 let __s_b2 = self.gpu.stream();
4037 let mut b2 = __s_b2.launch_builder(&f2);
4038 b2.arg(&part).arg(&mut p).arg(&nbi);
4039 unsafe { b2.launch(cfg2)?; }
4040 Ok(p)
4041 }
4042
4043 pub fn prob_of_token_device_col(&self, logits: &CudaSlice<f32>,
4050 tok_all: &CudaSlice<u32>, tok_idx: usize,
4051 p_out: &mut CudaSlice<f32>, p_idx: usize, n_vocab: usize)
4052 -> Result<(), Box<dyn std::error::Error>> {
4053 let tok_v = tok_all.slice(tok_idx..tok_idx + 1);
4054 let mut p_v = p_out.slice_mut(p_idx..p_idx + 1);
4055 let nb = ARGMAX_NB;
4056 let mut part = self.alloc_uninit::<f32>(nb)?;
4057 let f1 = self.func("prob_of_token_partial_f32");
4058 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4059 let nv = n_vocab as i32;
4060 let __s_b1 = self.gpu.stream();
4061 let mut b1 = __s_b1.launch_builder(&f1);
4062 b1.arg(logits).arg(&tok_v).arg(&mut part).arg(&nv);
4063 unsafe { b1.launch(cfg1)?; }
4064 let f2 = self.func("prob_of_token_final_f32");
4065 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4066 let nbi = nb as i32;
4067 let __s_b2 = self.gpu.stream();
4068 let mut b2 = __s_b2.launch_builder(&f2);
4069 b2.arg(&part).arg(&mut p_v).arg(&nbi);
4070 unsafe { b2.launch(cfg2)?; }
4071 Ok(())
4072 }
4073
4074 pub fn prob_of_token_device_into(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>,
4075 p_out: &mut CudaSlice<f32>, n_vocab: usize)
4076 -> Result<(), Box<dyn std::error::Error>> {
4077 let nb = ARGMAX_NB;
4078 let mut part = self.alloc_uninit::<f32>(nb)?;
4079 let f1 = self.func("prob_of_token_partial_f32");
4080 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4081 let nv = n_vocab as i32;
4082 let __s_b1 = self.gpu.stream();
4083 let mut b1 = __s_b1.launch_builder(&f1);
4084 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
4085 unsafe { b1.launch(cfg1)?; }
4086 let f2 = self.func("prob_of_token_final_f32");
4087 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4088 let nbi = nb as i32;
4089 let __s_b2 = self.gpu.stream();
4090 let mut b2 = __s_b2.launch_builder(&f2);
4091 b2.arg(&part).arg(p_out).arg(&nbi);
4092 unsafe { b2.launch(cfg2)?; }
4093 Ok(())
4094 }
4095
4096 pub fn argmax_token_device(&self, logits: &CudaSlice<f32>, n_vocab: usize)
4097 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4098 let mut tok = unsafe { self.gpu.stream().alloc::<u32>(1)? };
4099 self.argmax_token_device_into(logits, &mut tok, n_vocab)?;
4100 Ok(tok)
4101 }
4102 pub fn argmax_token_device_into(&self, logits: &CudaSlice<f32>, tok: &mut CudaSlice<u32>,
4109 n_vocab: usize) -> Result<(), Box<dyn std::error::Error>> {
4110 let nb = ARGMAX_NB;
4111 let f1 = self.func("argmax_partial_f32");
4112 let f2 = self.func("argmax_final_f32");
4113 let mut guard = self.argmax_partials.lock().unwrap();
4114 if guard.is_none() {
4115 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4118 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4119 *guard = Some((pv, pi));
4120 }
4121 let (part_v, part_i) = guard.as_mut().unwrap();
4122 let nv = n_vocab as i32;
4123 let nbi = nb as i32;
4124 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4126 let __s_b1 = self.gpu.stream();
4127 let mut b1 = __s_b1.launch_builder(&f1);
4128 b1.arg(logits).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4129 unsafe { b1.launch(cfg1)?; }
4130 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4132 let __s_b2 = self.gpu.stream();
4133 let mut b2 = __s_b2.launch_builder(&f2);
4134 b2.arg(&*part_v).arg(&*part_i).arg(tok).arg(&nbi);
4135 unsafe { b2.launch(cfg2)?; }
4136 Ok(())
4137 }
4138 pub fn argmax_token_device_col(&self, logits: &CudaSlice<f32>, col: usize, n_vocab: usize,
4144 toks: &mut CudaSlice<u32>, out_idx: usize)
4145 -> Result<(), Box<dyn std::error::Error>> {
4146 let nb = ARGMAX_NB;
4147 let f1 = self.func("argmax_partial_f32");
4148 let f2 = self.func("argmax_final_f32");
4149 let mut guard = self.argmax_partials.lock().unwrap();
4150 if guard.is_none() {
4151 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4152 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4153 *guard = Some((pv, pi));
4154 }
4155 let (part_v, part_i) = guard.as_mut().unwrap();
4156 let col_view = logits.slice(col * n_vocab..(col + 1) * n_vocab);
4157 let nv = n_vocab as i32;
4158 let nbi = nb as i32;
4159 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4160 let __s_b1 = self.gpu.stream();
4161 let mut b1 = __s_b1.launch_builder(&f1);
4162 b1.arg(&col_view).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4163 unsafe { b1.launch(cfg1)?; }
4164 let mut tok_view = toks.slice_mut(out_idx..out_idx + 1);
4165 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4166 let __s_b2 = self.gpu.stream();
4167 let mut b2 = __s_b2.launch_builder(&f2);
4168 b2.arg(&*part_v).arg(&*part_i).arg(&mut tok_view).arg(&nbi);
4169 unsafe { b2.launch(cfg2)?; }
4170 Ok(())
4171 }
4172 pub fn htod_u32_v(&self, v: &[u32]) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4174 Ok(self.gpu.stream().clone_htod(v)?)
4175 }
4176 pub fn dtoh_u32(&self, d: &CudaSlice<u32>) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
4177 let v = self.gpu.stream().clone_dtoh(d)?;
4178 self.gpu.stream().synchronize()?;
4179 Ok(v)
4180 }
4181 pub fn htod_u32_into(&self, dst: &mut CudaSlice<u32>, src: &[u32])
4185 -> Result<(), Box<dyn std::error::Error>> {
4186 let mut view = dst.slice_mut(0..src.len());
4187 self.gpu.stream().memcpy_htod(src, &mut view)?;
4188 Ok(())
4189 }
4190
4191 pub fn alloc_u32_zeroed(&self, n: usize) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4192 let s = self.gpu.stream().alloc_zeros::<u32>(n)?;
4193 self.keep_if_capturing(&s);
4194 Ok(s)
4195 }
4196 pub fn embed_gather_device_into(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4199 x_out: &mut CudaSlice<f32>, n_embd: usize, qtype: i32,
4200 row_bytes: usize) -> Result<(), Box<dyn std::error::Error>> {
4201 let f = self.func("embed_gather_u32");
4202 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
4203 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4204 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
4205 let __s_b = self.gpu.stream();
4206 let mut b = __s_b.launch_builder(&f);
4207 b.arg(embd).arg(token_d).arg(x_out).arg(&ne).arg(&qt).arg(&rb);
4208 unsafe { b.launch(cfg)?; }
4209 Ok(())
4210 }
4211 pub fn dtoh_i32_one(&self, d: &CudaSlice<i32>) -> Result<i32, Box<dyn std::error::Error>> {
4213 let v = self.gpu.stream().clone_dtoh(d)?;
4214 self.gpu.stream().synchronize()?;
4215 Ok(v[0])
4216 }
4217 pub fn i32_set_k(&self, dst: &mut CudaSlice<i32>, v: i32)
4224 -> Result<(), Box<dyn std::error::Error>> {
4225 let f = self.func("i32_set_k");
4226 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
4227 let idx = 0i32;
4228 let __s_b = self.gpu.stream();
4229 let mut b = __s_b.launch_builder(&f);
4230 b.arg(dst).arg(&v).arg(&idx);
4231 unsafe { b.launch(cfg)?; }
4232 Ok(())
4233 }
4234
4235 pub fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
4236 self.gpu.stream().memcpy_htod(&[v], d)?;
4237 Ok(())
4238 }
4239 pub fn set_u32_one(&self, d: &mut CudaSlice<u32>, v: u32) -> Result<(), Box<dyn std::error::Error>> {
4242 self.gpu.stream().memcpy_htod(&[v], d)?;
4243 Ok(())
4244 }
4245 pub fn dtoh_u32_one(&self, d: &CudaSlice<u32>) -> Result<u32, Box<dyn std::error::Error>> {
4247 let v = self.gpu.stream().clone_dtoh(d)?;
4248 self.gpu.stream().synchronize()?;
4249 Ok(v[0])
4250 }
4251 pub fn upload_u8(&self, bytes: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
4253 Ok(self.gpu.stream().clone_htod(bytes)?)
4254 }
4255 pub fn embed_gather_device(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4259 n_embd: usize, qtype: i32, row_bytes: usize)
4260 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4261 let f = self.func("embed_gather_u32");
4262 let mut x = self.alloc_uninit::<f32>(n_embd)?;
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(&mut x).arg(&ne).arg(&qt).arg(&rb);
4269 unsafe { b.launch(cfg)?; }
4270 Ok(x)
4271 }
4272
4273
4274 pub fn embed_gather_device_t(&self, embd: &CudaSlice<u8>, tokens: &[u32],
4278 n_embd: usize, qtype: i32, row_bytes: usize)
4279 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4280 let t = tokens.len();
4281 let tok_d = self.gpu.stream().clone_htod(tokens)?;
4282 let f = self.func("embed_gather_u32_t");
4283 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4284 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4285 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4286 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4287 let __s_b = self.gpu.stream();
4288 let mut b = __s_b.launch_builder(&f);
4289 b.arg(embd).arg(&tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4290 unsafe { b.launch(cfg)?; }
4291 Ok(x)
4292 }
4293
4294 pub fn embed_gather_device_tv(&self, embd: &CudaSlice<u8>, tok_v: &cudarc::driver::CudaView<u32>,
4299 t: usize, n_embd: usize, qtype: i32, row_bytes: usize)
4300 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4301 let f = self.func("embed_gather_u32_t");
4302 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4303 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4304 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4305 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4306 let __s_b = self.gpu.stream();
4307 let mut b = __s_b.launch_builder(&f);
4308 b.arg(embd).arg(tok_v).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4309 unsafe { b.launch(cfg)?; }
4310 Ok(x)
4311 }
4312
4313 pub fn embed_gather_device_td(&self, embd: &CudaSlice<u8>, tok_d: &CudaSlice<u32>, t: usize,
4314 n_embd: usize, qtype: i32, row_bytes: usize)
4315 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4316 let f = self.func("embed_gather_u32_t");
4317 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4318 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4319 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4320 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4321 let __s_b = self.gpu.stream();
4322 let mut b = __s_b.launch_builder(&f);
4323 b.arg(embd).arg(tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4324 unsafe { b.launch(cfg)?; }
4325 Ok(x)
4326 }
4327
4328 #[inline]
4334 fn keep_if_capturing<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, s: &CudaSlice<T>) {
4336 if self.capture_keep_on.load(std::sync::atomic::Ordering::Relaxed) {
4337 self.capture_keep.lock().unwrap().push(Box::new(s.clone()));
4338 }
4339 }
4340
4341 fn alloc_uninit<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, n: usize)
4342 -> Result<CudaSlice<T>, Box<dyn std::error::Error>> {
4343 let mut s = unsafe { self.gpu.stream().alloc::<T>(n)? };
4344 {
4348 static Z: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4349 if *Z.get_or_init(|| std::env::var("MEMRA_DEBUG_ZERO_ALLOCS").as_deref() == Ok("1")) {
4350 use cudarc::driver::DevicePtrMut;
4352 let n_bytes = s.len() * std::mem::size_of::<T>();
4353 let stream = self.gpu.stream();
4354 let (p_, _g) = s.device_ptr_mut(&stream);
4355 unsafe {
4356 cudarc::driver::sys::cuMemsetD8Async(p_, 0, n_bytes, stream.cu_stream())
4357 .result()?;
4358 }
4359 }
4360 }
4361 self.keep_if_capturing(&s);
4362 Ok(s)
4363 }
4364
4365 pub fn uninit_q8_pair(&self, n: usize)
4370 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4371 Ok((self.alloc_uninit::<i8>(n)?, self.alloc_uninit::<f32>(n / 32)?))
4372 }
4373
4374 pub fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4375 self.alloc_uninit::<f32>(n)
4376 }
4377
4378 pub fn alloc_i8_uninit(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
4380 self.alloc_uninit::<i8>(n)
4381 }
4382
4383 #[allow(clippy::too_many_arguments)]
4387 pub fn rms_norm3(&self, x: &CudaSlice<f32>, w0: &CudaSlice<f32>, w1: &CudaSlice<f32>,
4388 w2: &CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4389 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4390 -> Result<(), Box<dyn std::error::Error>> {
4391 let f = self.func("rms_norm3_f32");
4392 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4393 let (nc, e) = (ncols as i32, eps);
4394 let __s_b = self.gpu.stream();
4395 let mut b = __s_b.launch_builder(&f);
4396 b.arg(x).arg(w0).arg(w1).arg(w2).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e);
4397 unsafe { b.launch(cfg)?; }
4398 Ok(())
4399 }
4400
4401 #[allow(clippy::too_many_arguments)]
4403 pub fn qkvnorm_w_on_prefill(rows: usize, ncols: usize) -> bool {
4406 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4407 *WARP_ON.get_or_init(|| {
4408 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4409 }) && ncols % 4 == 0 && rows >= 64
4410 }
4411
4412 #[allow(clippy::too_many_arguments)]
4415 pub fn rms_norm_qkv_w4b(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4416 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4417 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4418 dvb: &mut CudaSlice<u8>,
4419 ncols: usize, rq: usize, rk: usize, eps: f32, vf16: bool)
4420 -> Result<(), Box<dyn std::error::Error>> {
4421 assert!(ncols % 4 == 0 && rq + 2 * rk >= 64);
4422 let f = self.func("rms_norm_qkv_w4b_f32");
4423 let rows = (rq + 2 * rk) as u32;
4424 let cfg = LaunchConfig {
4425 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4426 };
4427 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4428 let vf = vf16 as i32;
4429 let __s_b = self.gpu.stream();
4430 let mut b = __s_b.launch_builder(&f);
4431 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv).arg(&mut *dvb)
4432 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e).arg(&vf);
4433 unsafe { b.launch(cfg)?; }
4434 Ok(())
4435 }
4436
4437 pub fn rms_norm_qkv(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4438 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4439 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4440 ncols: usize, rq: usize, rk: usize, eps: f32)
4441 -> Result<(), Box<dyn std::error::Error>> {
4442 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4446 let warp_on = *WARP_ON.get_or_init(|| {
4447 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4448 });
4449 if warp_on && ncols % 4 == 0 && rq + 2 * rk >= 64 {
4452 let f = self.func("rms_norm_qkv_w4_f32");
4453 let rows = (rq + 2 * rk) as u32;
4454 let cfg = LaunchConfig {
4455 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4456 };
4457 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4458 let __s_b = self.gpu.stream();
4459 let mut b = __s_b.launch_builder(&f);
4460 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4461 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e);
4462 unsafe { b.launch(cfg)?; }
4463 return Ok(());
4464 }
4465 let f = self.func("rms_norm_qkv_f32");
4466 let grid = (rq + 2 * rk) as u32;
4467 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4468 let (nc, rqi, rki, e) = (ncols as i32, rq as i32, rk as i32, eps);
4469 let __s_b = self.gpu.stream();
4470 let mut b = __s_b.launch_builder(&f);
4471 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4472 .arg(&nc).arg(&rqi).arg(&rki).arg(&e);
4473 unsafe { b.launch(cfg)?; }
4474 Ok(())
4475 }
4476
4477 #[allow(clippy::too_many_arguments)]
4479 pub fn rms_norm2x(&self, a: &CudaSlice<f32>, bb: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4480 wb: &CudaSlice<f32>, da: &mut CudaSlice<f32>, db: &mut CudaSlice<f32>,
4481 ncols: usize, nrows: usize, eps: f32)
4482 -> Result<(), Box<dyn std::error::Error>> {
4483 let f = self.func("rms_norm2x_f32");
4484 let cfg = LaunchConfig { grid_dim: (2 * nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4485 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
4486 let __s_b = self.gpu.stream();
4487 let mut b = __s_b.launch_builder(&f);
4488 b.arg(a).arg(bb).arg(wa).arg(wb).arg(da).arg(db).arg(&nc).arg(&nr).arg(&e);
4489 unsafe { b.launch(cfg)?; }
4490 Ok(())
4491 }
4492
4493 pub fn softcap(&self, y: &mut CudaSlice<f32>, cap: f32, n: usize)
4495 -> Result<(), Box<dyn std::error::Error>> {
4496 let f = self.func("softcap_f32");
4497 let cfg = LaunchConfig::for_num_elems(n as u32);
4498 let ni = n as i32;
4499 let __s_b = self.gpu.stream();
4500 let mut b = __s_b.launch_builder(&f);
4501 b.arg(y).arg(&cap).arg(&ni);
4502 unsafe { b.launch(cfg)?; }
4503 Ok(())
4504 }
4505
4506 pub fn mask_ids_rows(&self, y: &mut CudaSlice<f32>, ids: &CudaSlice<i32>, n_ids: usize,
4509 n_vocab: usize, t: usize)
4510 -> Result<(), Box<dyn std::error::Error>> {
4511 let f = self.func("mask_ids_rows_f32");
4512 let cfg = LaunchConfig::for_num_elems((n_ids * t) as u32);
4513 let (ni, nv, ti) = (n_ids as i32, n_vocab as i32, t as i32);
4514 let __s_b = self.gpu.stream();
4515 let mut b = __s_b.launch_builder(&f);
4516 b.arg(y).arg(ids).arg(&ni).arg(&nv).arg(&ti);
4517 unsafe { b.launch(cfg)?; }
4518 Ok(())
4519 }
4520
4521 #[allow(clippy::too_many_arguments)]
4523 pub fn add_scale_rms_norm(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4524 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4525 ncols: usize, nrows: usize, eps: f32)
4526 -> Result<(), Box<dyn std::error::Error>> {
4527 let f = self.func("add_scale_rms_norm_f32");
4528 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4529 let (nc, e2) = (ncols as i32, eps);
4530 let __s_b = self.gpu.stream();
4531 let mut b = __s_b.launch_builder(&f);
4532 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(dst).arg(&nc).arg(&e2);
4533 unsafe { b.launch(cfg)?; }
4534 Ok(())
4535 }
4536
4537 #[allow(clippy::too_many_arguments)]
4540 pub fn add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4541 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4542 ncols: usize, nrows: usize, eps: f32)
4543 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4544 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4545 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4546 let (nc, e2) = (ncols as i32, eps);
4547 if Self::pdl_on() && Self::pdl_wb_on() {
4548 {
4549 use cudarc::driver::{DevicePtr, DevicePtrMut};
4550 let s = &self.gpu.stream();
4551 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4552 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4553 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4554 let mut ps = [
4555 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4556 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4557 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4558 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4559 &e2 as *const _ as *mut _,
4560 ];
4561 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4562 (rms_block(), 1, 1), &mut ps)?; }
4563 }
4564 return Ok((out_q, out_d));
4565 }
4566 let f = self.func("add_scale_rms_norm_q8_1");
4567 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4568 let __s_b = self.gpu.stream();
4569 let mut b = __s_b.launch_builder(&f);
4570 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e2);
4571 unsafe { b.launch(cfg)?; }
4572 Ok((out_q, out_d))
4573 }
4574
4575 #[allow(clippy::too_many_arguments)]
4577 pub fn add_scale_rms_norm_q8_1_into(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4578 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4579 ncols: usize, nrows: usize, eps: f32,
4580 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4581 -> Result<(), Box<dyn std::error::Error>> {
4582 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4583 let (nc, e2) = (ncols as i32, eps);
4584 if Self::pdl_on() && Self::pdl_wb_on() {
4585 use cudarc::driver::{DevicePtr, DevicePtrMut};
4586 let s = &self.gpu.stream();
4587 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4588 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4589 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4590 let mut ps = [
4591 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4592 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4593 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4594 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4595 &e2 as *const _ as *mut _,
4596 ];
4597 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4598 (rms_block(), 1, 1), &mut ps)?; }
4599 return Ok(());
4600 }
4601 let f = self.func("add_scale_rms_norm_q8_1");
4602 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4603 let __s_b = self.gpu.stream();
4604 let mut b = __s_b.launch_builder(&f);
4605 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut *out_q).arg(&mut *out_d).arg(&nc).arg(&e2);
4606 unsafe { b.launch(cfg)?; }
4607 Ok(())
4608 }
4609
4610 #[allow(clippy::too_many_arguments)]
4613 pub fn rms_pre_add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4614 b_in: &CudaSlice<f32>, c: f32,
4615 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4616 ncols: usize, nrows: usize, eps: f32)
4617 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4618 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4619 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4620 let (nc, e2) = (ncols as i32, eps);
4621 if Self::pdl_on() {
4622 {
4623 use cudarc::driver::{DevicePtr, DevicePtrMut};
4624 let s = &self.gpu.stream();
4625 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
4626 let (pb, _g2) = b_in.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
4627 let (pr, _g4) = res.device_ptr_mut(s);
4628 let (pq, _g5) = out_q.device_ptr_mut(s); let (pd, _g6) = out_d.device_ptr_mut(s);
4629 let mut ps = [
4630 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
4631 &pb as *const _ as *mut _, &c as *const _ as *mut _,
4632 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
4633 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4634 &nc as *const _ as *mut _, &e2 as *const _ as *mut _,
4635 ];
4636 unsafe { self.launch_pdl("rms_pre_add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4637 (rms_block(), 1, 1), &mut ps)?; }
4638 }
4639 return Ok((out_q, out_d));
4640 }
4641 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
4642 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4643 let __s_b = self.gpu.stream();
4644 let mut b = __s_b.launch_builder(&f);
4645 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);
4646 unsafe { b.launch(cfg)?; }
4647 Ok((out_q, out_d))
4648 }
4649
4650 pub fn gelu_tanh_mul_q8_1(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4653 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
4654 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4655 debug_assert!(ncols % 128 == 0);
4656 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4657 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4658 let nc = ncols as i32;
4659 if Self::pdl_on() {
4660 {
4661 use cudarc::driver::{DevicePtr, DevicePtrMut};
4662 let s = &self.gpu.stream();
4663 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4664 let (pact, _g2) = act.device_ptr_mut(s);
4665 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4666 let mut ps = [
4667 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4668 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4669 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4670 ];
4671 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4672 (rms_block(), 1, 1), &mut ps)?; }
4673 }
4674 return Ok((out_q, out_d));
4675 }
4676 let f = self.func("gelu_tanh_mul_q8_1");
4677 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4678 let __s_b = self.gpu.stream();
4679 let mut b = __s_b.launch_builder(&f);
4680 b.arg(gate).arg(up).arg(act).arg(&mut out_q).arg(&mut out_d).arg(&nc);
4681 unsafe { b.launch(cfg)?; }
4682 Ok((out_q, out_d))
4683 }
4684
4685 #[allow(clippy::too_many_arguments)]
4687 pub fn gelu_tanh_mul_q8_1_into(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4688 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
4689 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4690 -> Result<(), Box<dyn std::error::Error>> {
4691 debug_assert!(ncols % 128 == 0);
4692 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4693 let nc = ncols as i32;
4694 if Self::pdl_on() {
4695 use cudarc::driver::{DevicePtr, DevicePtrMut};
4696 let s = &self.gpu.stream();
4697 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4698 let (pact, _g2) = act.device_ptr_mut(s);
4699 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4700 let mut ps = [
4701 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4702 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4703 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4704 ];
4705 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4706 (rms_block(), 1, 1), &mut ps)?; }
4707 return Ok(());
4708 }
4709 let f = self.func("gelu_tanh_mul_q8_1");
4710 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4711 let __s_b = self.gpu.stream();
4712 let mut b = __s_b.launch_builder(&f);
4713 b.arg(gate).arg(up).arg(&mut *act).arg(&mut *out_q).arg(&mut *out_d).arg(&nc);
4714 unsafe { b.launch(cfg)?; }
4715 Ok(())
4716 }
4717
4718 #[allow(clippy::too_many_arguments)]
4720 pub fn add_rms_norm3_q8z(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4721 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4722 res: &mut CudaSlice<f32>, out1: &mut CudaSlice<f32>,
4723 ncols: usize, nrows: usize, eps: f32)
4724 -> Result<((CudaSlice<i8>, CudaSlice<f32>), (CudaSlice<i8>, CudaSlice<f32>)), Box<dyn std::error::Error>> {
4725 let mut q0 = self.alloc_uninit::<i8>(nrows * ncols)?;
4726 let mut d0 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4727 let mut q2 = self.alloc_uninit::<i8>(nrows * ncols)?;
4728 let mut d2 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4729 let f = self.func("add_rms_norm3_q8z_f32");
4730 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4731 let (nc, e2) = (ncols as i32, eps);
4732 let __s_b = self.gpu.stream();
4733 let mut b = __s_b.launch_builder(&f);
4734 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res)
4735 .arg(&mut q0).arg(&mut d0).arg(out1).arg(&mut q2).arg(&mut d2).arg(&nc).arg(&e2);
4736 unsafe { b.launch(cfg)?; }
4737 Ok(((q0, d0), (q2, d2)))
4738 }
4739
4740 #[allow(clippy::too_many_arguments)]
4742 pub fn add_rms_norm3(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4743 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4744 res: &mut CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4745 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4746 -> Result<(), Box<dyn std::error::Error>> {
4747 let f = self.func("add_rms_norm3_f32");
4748 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4749 let (nc, e2) = (ncols as i32, eps);
4750 let __s_b = self.gpu.stream();
4751 let mut b = __s_b.launch_builder(&f);
4752 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e2);
4753 unsafe { b.launch(cfg)?; }
4754 Ok(())
4755 }
4756
4757 pub fn add_scale(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4759 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
4760 let f = self.func("add_scale_f32");
4761 let cfg = LaunchConfig::for_num_elems(n as u32);
4762 let ni = n as i32;
4763 let __s_b = self.gpu.stream();
4764 let mut b = __s_b.launch_builder(&f);
4765 b.arg(a).arg(b_in).arg(&c).arg(dst).arg(&ni);
4766 unsafe { b.launch(cfg)?; }
4767 Ok(())
4768 }
4769
4770 pub fn rms_norm(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4771 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4772 let (nc, e) = (ncols as i32, eps);
4773 if Self::pdl_on() && Self::pdl_wb_on() {
4774 use cudarc::driver::{DevicePtr, DevicePtrMut};
4775 let s = &self.gpu.stream();
4776 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4777 let (pd, _g2) = dst.device_ptr_mut(s);
4778 let mut ps = [
4779 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4780 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4781 &e as *const _ as *mut _,
4782 ];
4783 unsafe { self.launch_pdl("rms_norm_f32", (nrows as u32, 1, 1),
4784 (rms_block(), 1, 1), &mut ps)?; }
4785 return Ok(());
4786 }
4787 let f = self.func("rms_norm_f32");
4788 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4789 let __s_b = self.gpu.stream();
4790 let mut b = __s_b.launch_builder(&f);
4791 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4792 unsafe { b.launch(cfg)?; }
4793 Ok(())
4794 }
4795
4796 pub fn rms_norm_decode(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4804 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4805 let f = self.func("rms_norm_f32");
4806 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4807 let (nc, e) = (ncols as i32, eps);
4808 let __s_b = self.gpu.stream();
4809 let mut b = __s_b.launch_builder(&f);
4810 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4811 unsafe { b.launch(cfg)?; }
4812 Ok(())
4813 }
4814
4815 pub fn rms_norm_q8_1(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize, nrows: usize,
4819 eps: f32) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4820 let nblk = ncols / 32;
4821 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
4822 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
4823 let (nc, e) = (ncols as i32, eps);
4824 if Self::pdl_on() {
4825 {
4826 use cudarc::driver::{DevicePtr, DevicePtrMut};
4827 let s = &self.gpu.stream();
4828 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4829 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
4830 let mut ps = [
4831 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4832 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4833 &nc as *const _ as *mut _, &e as *const _ as *mut _,
4834 ];
4835 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
4836 &mut ps)?; }
4837 }
4838 return Ok((q, d));
4839 }
4840 let f = self.func("rms_norm_q8_1");
4841 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4844 let __s_b = self.gpu.stream();
4845 let mut b = __s_b.launch_builder(&f);
4846 b.arg(x).arg(w).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
4847 unsafe { b.launch(cfg)?; }
4848 Ok((q, d))
4849 }
4850
4851 pub fn rms_norm_q8_1_into(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize,
4854 nrows: usize, eps: f32,
4855 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
4856 -> Result<(), Box<dyn std::error::Error>> {
4857 let nblk = ncols / 32;
4858 debug_assert!(q.len() >= nrows * ncols && d.len() >= nrows * nblk);
4859 let (nc, e) = (ncols as i32, eps);
4860 if Self::pdl_on() {
4861 use cudarc::driver::{DevicePtr, DevicePtrMut};
4862 let s = &self.gpu.stream();
4863 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4864 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
4865 let mut ps = [
4866 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4867 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4868 &nc as *const _ as *mut _, &e as *const _ as *mut _,
4869 ];
4870 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
4871 &mut ps)?; }
4872 return Ok(());
4873 }
4874 let f = self.func("rms_norm_q8_1");
4875 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4876 let __s_b = self.gpu.stream();
4877 let mut b = __s_b.launch_builder(&f);
4878 b.arg(x).arg(w).arg(&mut *q).arg(&mut *d).arg(&nc).arg(&e);
4879 unsafe { b.launch(cfg)?; }
4880 Ok(())
4881 }
4882
4883 pub fn quantize_q8_1_into(&self, x: &CudaSlice<f32>, m: usize, in_f: usize,
4885 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
4886 -> Result<(), Box<dyn std::error::Error>> {
4887 let nblk = in_f / 32;
4888 debug_assert!(q.len() >= m * in_f && d.len() >= m * nblk);
4889 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
4890 let (inf, mi) = (in_f as i32, m as i32);
4891 if Self::pdl_on() && Self::pdl_wb_on() {
4892 use cudarc::driver::{DevicePtr, DevicePtrMut};
4893 let s = &self.gpu.stream();
4894 let (px, _g0) = x.device_ptr(s);
4895 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
4896 let mut ps = [
4897 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
4898 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
4899 &mi as *const _ as *mut _,
4900 ];
4901 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
4902 return Ok(());
4903 }
4904 let f = self.func("quantize_q8_1");
4905 let __s_b = self.gpu.stream();
4906 let mut b = __s_b.launch_builder(&f);
4907 b.arg(x).arg(&mut *q).arg(&mut *d).arg(&inf).arg(&mi);
4908 unsafe { b.launch(cfg)?; }
4909 Ok(())
4910 }
4911
4912 pub fn add_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
4916 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4917 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4918 let nblk = ncols / 32;
4919 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
4920 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
4921 let f = self.func("add_rms_norm_q8_1");
4922 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4924 let (nc, e) = (ncols as i32, eps);
4925 let __s_bld = self.gpu.stream();
4926 let mut bld = __s_bld.launch_builder(&f);
4927 bld.arg(a).arg(b_in).arg(w).arg(res).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
4928 unsafe { bld.launch(cfg)?; }
4929 Ok((q, d))
4930 }
4931
4932 pub fn add_rms_norm(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
4936 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
4937 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4938 let (nc, e) = (ncols as i32, eps);
4939 if Self::pdl_on() && Self::pdl_wb_on() {
4940 use cudarc::driver::{DevicePtr, DevicePtrMut};
4941 let s = &self.gpu.stream();
4942 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b.device_ptr(s);
4943 let (pw, _g2) = w.device_ptr(s);
4944 let (pr, _g3) = res.device_ptr_mut(s); let (pd, _g4) = dst.device_ptr_mut(s);
4945 let mut ps = [
4946 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4947 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
4948 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4949 &e as *const _ as *mut _,
4950 ];
4951 unsafe { self.launch_pdl("add_rms_norm_f32", (nrows as u32, 1, 1),
4952 (rms_block(), 1, 1), &mut ps)?; }
4953 return Ok(());
4954 }
4955 let f = self.func("add_rms_norm_f32");
4956 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4957 let __s_b2 = self.gpu.stream();
4958 let mut b2 = __s_b2.launch_builder(&f);
4959 b2.arg(a).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
4960 unsafe { b2.launch(cfg)?; }
4961 Ok(())
4962 }
4963
4964 #[allow(clippy::too_many_arguments)]
4967 pub fn rms_pre_add_rms_norm(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4968 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
4969 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4970 ncols: usize, nrows: usize, eps: f32)
4971 -> Result<(), Box<dyn std::error::Error>> {
4972 let f = self.func("rms_pre_add_rms_norm_f32");
4973 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4974 let (nc, e) = (ncols as i32, eps);
4975 let __s_b2 = self.gpu.stream();
4976 let mut b2 = __s_b2.launch_builder(&f);
4977 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
4978 unsafe { b2.launch(cfg)?; }
4979 Ok(())
4980 }
4981
4982 #[allow(clippy::too_many_arguments)]
4984 pub fn rms_pre_add_rms_norm_q8z(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4985 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
4986 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4987 ncols: usize, nrows: usize, eps: f32)
4988 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4989 debug_assert!(ncols % 128 == 0);
4990 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4991 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4992 let (nc, e) = (ncols as i32, eps);
4993 if Self::pdl_on() {
4994 {
4995 use cudarc::driver::{DevicePtr, DevicePtrMut};
4996 let s = &self.gpu.stream();
4997 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
4998 let (pb, _g2) = b.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
4999 let (pr, _g4) = res.device_ptr_mut(s); let (pdst, _g5) = dst.device_ptr_mut(s);
5000 let (pq, _g6) = out_q.device_ptr_mut(s); let (pd, _g7) = out_d.device_ptr_mut(s);
5001 let mut ps = [
5002 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
5003 &pb as *const _ as *mut _, &pw as *const _ as *mut _,
5004 &pr as *const _ as *mut _, &pdst as *const _ as *mut _,
5005 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
5006 &nc as *const _ as *mut _, &e as *const _ as *mut _,
5007 ];
5008 unsafe { self.launch_pdl("rms_pre_add_rms_norm_q8z_f32", (nrows as u32, 1, 1),
5009 (rms_block(), 1, 1), &mut ps)?; }
5010 }
5011 return Ok((out_q, out_d));
5012 }
5013 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
5014 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5015 let __s_b2 = self.gpu.stream();
5016 let mut b2 = __s_b2.launch_builder(&f);
5017 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst)
5018 .arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e);
5019 unsafe { b2.launch(cfg)?; }
5020 Ok((out_q, out_d))
5021 }
5022
5023 pub fn build_q4_out_concat3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
5027 w2: &crate::model::GpuTensor)
5028 -> Result<Option<crate::model::GpuTensor>, Box<dyn std::error::Error>> {
5029 use crate::model::GpuTensor;
5030 let part = |w: &GpuTensor| -> Option<(usize, usize)> {
5031 match w {
5032 GpuTensor::Quant { qtype, row_bytes, rp, .. }
5033 if *qtype == QT_Q4_0 && !*rp => Some((*row_bytes, w.out_features())),
5034 _ => None,
5035 }
5036 };
5037 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (part(w0), part(w1), part(w2))
5038 else { return Ok(None) };
5039 if rb0 != rb1 || rb0 != rb2
5040 || w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
5041 return Ok(None);
5042 }
5043 fn bytes_of(w: &crate::model::GpuTensor) -> &CudaSlice<u8> {
5044 match w { crate::model::GpuTensor::Quant { bytes, .. } => bytes, _ => unreachable!() }
5045 }
5046 let (b0, b1, b2) = (bytes_of(w0), bytes_of(w1), bytes_of(w2));
5047 let total = rb0 * (o0 + o1 + o2);
5048 let mut cat = self.alloc_u8(total)?;
5049 self.copy_u8_into(&mut cat, 0, b0, rb0 * o0)?;
5050 self.copy_u8_into(&mut cat, rb0 * o0, b1, rb1 * o1)?;
5051 self.copy_u8_into(&mut cat, rb0 * (o0 + o1), b2, rb2 * o2)?;
5052 Ok(Some(GpuTensor::Quant {
5053 bytes: cat, qtype: QT_Q4_0, row_bytes: rb0,
5054 ne: vec![w0.in_features() as u64, (o0 + o1 + o2) as u64], scale: 1.0, rp: false,
5055 #[cfg(memra_cutlass)]
5056 cutlass: None,
5057 fp8: None, blk: None, rp4: None, f16: None,
5058 }))
5059 }
5060
5061 #[allow(clippy::too_many_arguments)]
5063 pub fn rms_norm_qkv_rope_cat(&self, qkv: &CudaSlice<f32>,
5064 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5065 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5066 head_dim: usize, rq: usize, rk: usize,
5067 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5068 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32)
5069 -> Result<(), Box<dyn std::error::Error>> {
5070 let rows = rq + rk + rk;
5071 let theta_scale = base.powf(-2.0 / head_dim as f32);
5072 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5073 if Self::pdl_on() {
5074 use cudarc::driver::{DevicePtr, DevicePtrMut};
5075 let s = &self.gpu.stream();
5076 let (pqkv, _g0) = qkv.device_ptr(s);
5077 let (pwq, _g1) = wq.device_ptr(s); let (pwk, _g2) = wk.device_ptr(s);
5078 let (pwv, _g3) = wv.device_ptr(s);
5079 let (pq, _g4) = q.device_ptr_mut(s); let (pk, _g5) = k.device_ptr_mut(s);
5080 let (pv, _g6) = v.device_ptr_mut(s);
5081 let (ppos, _g7) = pos.device_ptr(s);
5082 let (pff, _g8) = match ff {
5083 Some(t) => { let (p, g) = t.device_ptr(s); (p, Some(g)) }
5084 None => (0, None),
5085 };
5086 let mut ps = [
5087 &pqkv as *const _ as *mut std::ffi::c_void,
5088 &pwq as *const _ as *mut _, &pwk as *const _ as *mut _,
5089 &pwv as *const _ as *mut _,
5090 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5091 &pv as *const _ as *mut _,
5092 &nc as *const _ as *mut _, &rqi as *const _ as *mut _,
5093 &rki as *const _ as *mut _, &ppos as *const _ as *mut _,
5094 &nhq as *const _ as *mut _, &nhk as *const _ as *mut _,
5095 &theta_scale as *const _ as *mut _, &freq_scale as *const _ as *mut _,
5096 &pff as *const _ as *mut _, &eps as *const _ as *mut _,
5097 ];
5098 unsafe { self.launch_pdl("rms_norm_qkv_rope_cat_f32", (rows as u32, 1, 1),
5099 (rms_block(), 1, 1), &mut ps)?; }
5100 return Ok(());
5101 }
5102 let f = self.func("rms_norm_qkv_rope_cat_f32");
5103 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5104 let __s_b = self.gpu.stream();
5105 let mut b = __s_b.launch_builder(&f);
5106 match ff {
5107 Some(t) => { b.arg(qkv).arg(wq).arg(wk).arg(wv)
5108 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5109 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5110 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5111 unsafe { b.launch(cfg)?; } }
5112 None => { let null: u64 = 0;
5113 b.arg(qkv).arg(wq).arg(wk).arg(wv)
5114 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5115 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5116 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5117 unsafe { b.launch(cfg)?; } }
5118 }
5119 Ok(())
5120 }
5121
5122 #[allow(clippy::too_many_arguments)]
5124 pub fn rms_norm_qkv_rope(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>, v0: &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 f = self.func("rms_norm_qkv_rope_f32");
5132 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 };
5134 let theta_scale = base.powf(-2.0 / head_dim as f32);
5135 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5136 let __s_b = self.gpu.stream();
5137 let mut b = __s_b.launch_builder(&f);
5138 match ff {
5139 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5140 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5141 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5142 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5143 unsafe { b.launch(cfg)?; } }
5144 None => { let null: u64 = 0;
5145 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5146 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5147 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5148 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5149 unsafe { b.launch(cfg)?; } }
5150 }
5151 Ok(())
5152 }
5153
5154 #[allow(clippy::too_many_arguments)]
5158 pub fn rms_norm_qkv_rope_append_dc(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>,
5159 v0: &CudaSlice<f32>,
5160 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5161 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5162 head_dim: usize, rq: usize, rk: usize,
5163 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5164 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32,
5165 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
5166 t_dev: &CudaSlice<i32>, k_tok_bytes: usize, v_tok_bytes: usize,
5167 g: bool)
5168 -> Result<(), Box<dyn std::error::Error>> {
5169 let rows = rq + rk + rk;
5170 let theta_scale = base.powf(-2.0 / head_dim as f32);
5171 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5172 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
5173 if Self::pdl_on() && Self::pdl_wb_on() {
5174 use cudarc::driver::{DevicePtr, DevicePtrMut};
5175 let s = &self.gpu.stream();
5176 let (p0, _a0) = q0.device_ptr(s); let (p1, _a1) = k0.device_ptr(s);
5177 let (p2, _a2) = v0.device_ptr(s);
5178 let (pwq, _a3) = wq.device_ptr(s); let (pwk, _a4) = wk.device_ptr(s);
5179 let (pwv, _a5) = wv.device_ptr(s);
5180 let (pq, _a6) = q.device_ptr_mut(s); let (pk, _a7) = k.device_ptr_mut(s);
5181 let (pv, _a8) = v.device_ptr_mut(s);
5182 let (pp, _a9) = pos.device_ptr(s);
5183 let pff: u64 = match ff { Some(t) => { let (p, _gg) = t.device_ptr(s); p as u64 }
5184 None => 0 };
5185 let (pkc, _a10) = kc.device_ptr_mut(s); let (pvc, _a11) = vc.device_ptr_mut(s);
5186 let (pt, _a12) = t_dev.device_ptr(s);
5187 let mut ps = [
5188 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
5189 &p2 as *const _ as *mut _, &pwq as *const _ as *mut _,
5190 &pwk as *const _ as *mut _, &pwv as *const _ as *mut _,
5191 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5192 &pv as *const _ as *mut _, &nc as *const _ as *mut _,
5193 &rqi as *const _ as *mut _, &rki as *const _ as *mut _,
5194 &pp as *const _ as *mut _, &nhq as *const _ as *mut _,
5195 &nhk as *const _ as *mut _, &theta_scale as *const _ as *mut _,
5196 &freq_scale as *const _ as *mut _, &pff as *const _ as *mut _,
5197 &eps as *const _ as *mut _, &pkc as *const _ as *mut _,
5198 &pvc as *const _ as *mut _, &pt as *const _ as *mut _,
5199 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
5200 ];
5201 unsafe { self.launch_pdl_flash(g, "rms_norm_qkv_rope_append_dc_f32",
5202 (rows as u32, 1, 1), (rms_block(), 1, 1), 0, &mut ps)?; }
5203 return Ok(());
5204 }
5205 let f = if g { self.func_g("rms_norm_qkv_rope_append_dc_f32") }
5206 else { self.func("rms_norm_qkv_rope_append_dc_f32") };
5207 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5208 let __s_b = self.gpu.stream();
5209 let mut b = __s_b.launch_builder(&f);
5210 match ff {
5211 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5212 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5213 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5214 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps)
5215 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5216 unsafe { b.launch(cfg)?; } }
5217 None => { let null: u64 = 0;
5218 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5219 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5220 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5221 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps)
5222 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5223 unsafe { b.launch(cfg)?; } }
5224 }
5225 Ok(())
5226 }
5227
5228 pub fn add_q8_1(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
5230 ncols: usize, nrows: usize)
5231 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5232 debug_assert!(ncols % 128 == 0);
5233 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5234 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5235 let f = self.func("add_q8_1_f32");
5236 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5237 let nc = ncols as i32;
5238 let __s_b2 = self.gpu.stream();
5239 let mut b2 = __s_b2.launch_builder(&f);
5240 b2.arg(a).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc);
5241 unsafe { b2.launch(cfg)?; }
5242 Ok((out_q, out_d))
5243 }
5244
5245 pub fn rms_pre_add_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>, b: &CudaSlice<f32>,
5249 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
5250 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5251 debug_assert!(ncols % 128 == 0);
5252 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5253 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5254 let f = self.func("rms_pre_add_q8_1_f32");
5255 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1),
5256 shared_mem_bytes: 0 };
5257 let (nc, ep) = (ncols as i32, eps);
5258 let __s_b2 = self.gpu.stream();
5259 let mut b2 = __s_b2.launch_builder(&f);
5260 b2.arg(a).arg(wa).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
5261 unsafe { b2.launch(cfg)?; }
5262 Ok((out_q, out_d))
5263 }
5264
5265 pub fn l2_v2_on(ncols: usize) -> bool {
5269 ncols == 128 && std::env::var("MEMRA_L2_V2").as_deref() != Ok("0")
5270 }
5271
5272 pub fn l2_norm_pp(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5273 dst16: Option<&mut CudaSlice<u8>>, ncols: usize, nrows: usize,
5274 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5275 if Self::l2_v2_on(ncols) {
5276 let f = self.func("l2_norm_pp_v2_f32");
5277 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 };
5279 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
5280 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
5282 let __s_b = self.gpu.stream();
5283 let mut b = __s_b.launch_builder(&f);
5284 b.arg(x).arg(dst).arg(&d16).arg(&nc).arg(&nr).arg(&e);
5285 unsafe { b.launch(cfg)?; }
5286 return Ok(());
5287 }
5288 self.l2_norm(x, dst, ncols, nrows, eps)
5289 }
5290
5291 pub fn l2_norm(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
5292 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5293 let f = self.func("l2_norm_f32");
5294 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
5295 let (nc, e) = (ncols as i32, eps);
5296 let __s_b = self.gpu.stream();
5297 let mut b = __s_b.launch_builder(&f);
5298 b.arg(x).arg(dst).arg(&nc).arg(&e);
5299 unsafe { b.launch(cfg)?; }
5300 Ok(())
5301 }
5302
5303 pub fn l2_norm_decode(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize,
5309 nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5310 let f = self.func("l2_norm_f32");
5311 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
5312 let (nc, e) = (ncols as i32, eps);
5313 let __s_b = self.gpu.stream();
5314 let mut b = __s_b.launch_builder(&f);
5315 b.arg(x).arg(dst).arg(&nc).arg(&e);
5316 unsafe { b.launch(cfg)?; }
5317 Ok(())
5318 }
5319
5320 pub fn rope_neox(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5322 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32, freq_scale: f32)
5323 -> Result<(), Box<dyn std::error::Error>> {
5324 let f = self.func("rope_neox_f32");
5325 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5326 let grid = (n_heads * n_tokens) as u32;
5327 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5328 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5329 let __s_b = self.gpu.stream();
5330 let mut b = __s_b.launch_builder(&f);
5331 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale);
5332 unsafe { b.launch(cfg)?; }
5333 Ok(())
5334 }
5335
5336 pub fn rope_neox_ff(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5338 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32,
5339 freq_scale: f32, ff: &CudaSlice<f32>)
5340 -> Result<(), Box<dyn std::error::Error>> {
5341 let f = self.func("rope_neox_ff_f32");
5342 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5343 let grid = (n_heads * n_tokens) as u32;
5344 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5345 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5346 let __s_b = self.gpu.stream();
5347 let mut b = __s_b.launch_builder(&f);
5348 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale).arg(ff);
5349 unsafe { b.launch(cfg)?; }
5350 Ok(())
5351 }
5352
5353 #[allow(clippy::too_many_arguments)]
5355 pub fn rope_neox2(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
5356 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
5357 nh_q: usize, nh_k: usize, n_tokens: usize, freq_base: f32,
5358 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
5359 -> Result<(), Box<dyn std::error::Error>> {
5360 let f = self.func("rope_neox2_f32");
5361 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5362 let grid = ((nh_q + nh_k) * n_tokens) as u32;
5363 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5364 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);
5365 let __s_b = self.gpu.stream();
5366 let mut b = __s_b.launch_builder(&f);
5367 b.arg(q).arg(k).arg(pos).arg(&hd).arg(&nd).arg(&nq).arg(&nk).arg(&nt)
5368 .arg(&theta_scale).arg(&freq_scale);
5369 match ff {
5370 Some(ffv) => { b.arg(ffv); unsafe { b.launch(cfg)?; } }
5371 None => {
5372 let null: u64 = 0;
5373 b.arg(&null);
5374 unsafe { b.launch(cfg)?; }
5375 }
5376 }
5377 Ok(())
5378 }
5379
5380 pub fn gelu_tanh_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5382 -> Result<(), Box<dyn std::error::Error>> {
5383 let f = self.func("gelu_tanh_mul_f32");
5384 let cfg = LaunchConfig::for_num_elems(n as u32);
5385 let ni = n as i32;
5386 let __s_b = self.gpu.stream();
5387 let mut b = __s_b.launch_builder(&f);
5388 b.arg(gate).arg(up).arg(dst).arg(&ni);
5389 unsafe { b.launch(cfg)?; }
5390 Ok(())
5391 }
5392
5393 pub fn silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5394 -> Result<(), Box<dyn std::error::Error>> {
5395 let f = self.func("silu_mul_f32");
5396 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5398 let ni = n as i32;
5399 let __s_b = self.gpu.stream();
5400 let mut b = __s_b.launch_builder(&f);
5401 b.arg(gate).arg(up).arg(dst).arg(&ni);
5402 unsafe { b.launch(cfg)?; }
5403 Ok(())
5404 }
5405
5406 pub fn silu_mul_f16out(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
5409 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
5410 -> Result<(), Box<dyn std::error::Error>> {
5411 let f = self.func("silu_mul_f16out_f32");
5412 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5413 let ni = n as i32;
5414 let __s_b = self.gpu.stream();
5415 let mut b = __s_b.launch_builder(&f);
5416 b.arg(gate).arg(up).arg(dst).arg(dst16).arg(&ni);
5417 unsafe { b.launch(cfg)?; }
5418 Ok(())
5419 }
5420
5421 pub fn silu_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5428 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
5429 let f = self.func("silu_mul_scaled_f32");
5430 let cfg = LaunchConfig::for_num_elems(n as u32);
5431 let ni = n as i32;
5432 let (gsf, usf) = (gs, us);
5433 let __s_b = self.gpu.stream();
5434 let mut b = __s_b.launch_builder(&f);
5435 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(dst).arg(&ni);
5436 unsafe { b.launch(cfg)?; }
5437 Ok(())
5438 }
5439
5440 #[allow(clippy::too_many_arguments)]
5444 pub fn swigluoai_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5445 alpha: f32, limit: f32, dst: &mut CudaSlice<f32>, n: usize)
5446 -> Result<(), Box<dyn std::error::Error>> {
5447 let f = self.func("swigluoai_mul_scaled_f32");
5448 let cfg = LaunchConfig::for_num_elems(n as u32);
5449 let ni = n as i32;
5450 let __s_b = self.gpu.stream();
5451 let mut b = __s_b.launch_builder(&f);
5452 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&alpha).arg(&limit).arg(dst).arg(&ni);
5453 unsafe { b.launch(cfg)?; }
5454 Ok(())
5455 }
5456
5457 pub fn silu_mul_scaled_q8_1(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5465 n: usize)
5466 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5467 let f = self.func("silu_mul_scaled_q8_1");
5468 let nblk = n / 32;
5469 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);
5473 let (gsf, usf, ni) = (gs, us, n as i32);
5474 let __s_b = self.gpu.stream();
5475 let mut b = __s_b.launch_builder(&f);
5476 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(&mut aq).arg(&mut ad).arg(&ni);
5477 unsafe { b.launch(cfg)?; }
5478 Ok((aq, ad))
5479 }
5480
5481 pub fn add(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5482 -> Result<(), Box<dyn std::error::Error>> {
5483 let f = self.func("add_f32");
5484 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5486 let ni = n as i32;
5487 let __s_bld = self.gpu.stream();
5488 let mut bld = __s_bld.launch_builder(&f);
5489 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5490 unsafe { bld.launch(cfg)?; }
5491 Ok(())
5492 }
5493
5494 pub fn mul(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5495 -> Result<(), Box<dyn std::error::Error>> {
5496 let f = self.func("mul_f32");
5497 let cfg = LaunchConfig::for_num_elems(n as u32);
5498 let ni = n as i32;
5499 let __s_bld = self.gpu.stream();
5500 let mut bld = __s_bld.launch_builder(&f);
5501 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5502 unsafe { bld.launch(cfg)?; }
5503 Ok(())
5504 }
5505
5506 pub fn matmul(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
5509 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5510 use crate::model::GpuTensor;
5511 let in_f = w.in_features();
5512 let out_f = w.out_features();
5513 #[allow(non_snake_case)]
5521 let GEMM_M_THRESHOLD = if self.verify_exact_on() { usize::MAX } else { 16usize };
5524
5525 const GEMM_MIN_OUT_F: usize = 128; if m >= GEMM_M_THRESHOLD {
5550 if let Some(y) = self.try_fp8_gemm(w, x, m)? { return Ok(y); }
5551 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(y); }
5558 if let Some(y) = self.try_f16_gemm(w, x, m)? { return Ok(y); }
5561 }
5562 if let GpuTensor::Quant { qtype, .. } = w {
5577 if *qtype == QT_F8_E4M3_BLK {
5578 if m >= GEMM_M_THRESHOLD {
5579 if let Some(y) = self.try_e4m3_blk_prefill(w, x, m)? { return Ok(y); }
5580 }
5581 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5582 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
5583 }
5584 }
5585 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.mmq_supports(w) {
5586 return self.qmatvec_mmq(w, x, m);
5587 }
5588 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.gemm_supports(w) {
5589 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5590 return self.qmatvec_gemm(w, &aq, &ad, m);
5591 }
5592 if m >= GEMM_M_THRESHOLD {
5595 if let Some(y) = self.try_fp4_gemm(w, x, m, in_f, out_f)? { return Ok(y); }
5596 }
5597 let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
5601 if m == 1 && fast {
5606 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, scale, .. } = w {
5607 if self.mmvq_supports(*qtype) {
5608 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5612 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5613 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, rp);
5614 }
5615 }
5616 }
5617 if (2..=16).contains(&m) && fast && std::env::var("MEMRA_NO_BATCHED").is_err()
5633 && (m <= 4 || Self::b8_enabled()) {
5634 let m_ok = m <= 8 || matches!(w, GpuTensor::Quant { qtype, .. }
5644 if *qtype == QT_Q4_0 || *qtype == QT_Q6_K || *qtype == QT_F8_E4M3
5645 || *qtype == QT_NVFP4 || *qtype == QT_Q4_K || *qtype == QT_Q5_K || *qtype == QT_Q8_0);
5646 if m_ok {
5647 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, .. } = w {
5648 if self.batched_supports(*qtype) && self.mmvq_supports(*qtype) {
5649 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5650 let mcols = Self::batched_mcols(m);
5651 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5652 let mut y = self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, mcols, 1.0, rp)?;
5653 if let GpuTensor::Quant { scale, .. } = w {
5654 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5655 }
5656 return Ok(y);
5657 }
5658 }
5659 }
5660 }
5661 if fast {
5667 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. } = w {
5668 if *qtype == QT_F8_E4M3 {
5669 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5670 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes,
5671 *scale, false);
5672 }
5673 }
5674 }
5675 let mut y = match w {
5676 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q8_0 =>
5677 self.qmatvec_q8_0_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5678 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q4_K =>
5679 self.qmatvec_q4_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5680 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q6_K =>
5681 self.qmatvec_q6_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5682 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q5_K =>
5683 self.qmatvec_q5_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5684 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q3_K =>
5685 self.qmatvec_q3_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5686 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } if fast && *qtype == QT_NVFP4 =>
5687 self.qmatvec_dp4a_named(
5688 if *rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
5689 bytes, x, m, in_f, out_f, *row_bytes)?,
5690 GpuTensor::Quant { bytes, qtype, row_bytes, .. }
5694 if fast && *qtype == QT_IQ4_XS && Self::iq_fast_enabled() =>
5695 self.qmatvec_iq4_XS_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5696 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } =>
5701 self.qmatvec(bytes, x, m, in_f, out_f,
5704 if *rp && *qtype == QT_NVFP4 { QT_NVFP4_RP } else { *qtype },
5705 *row_bytes)?,
5706 GpuTensor::Float { data, .. } => self.linear(x, data, m, in_f, out_f)?,
5707 GpuTensor::FloatBf16 { data, .. } =>
5710 self.linear_bf16_chunked(x, data, m, in_f, out_f, false)?,
5711 };
5712 if let GpuTensor::Quant { scale, .. } = w {
5714 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5715 }
5716 Ok(y)
5717 }
5718
5719 pub fn uses_q8_1_fast(&self, w: &crate::model::GpuTensor) -> bool {
5722 use crate::model::GpuTensor;
5723 if std::env::var("MEMRA_FAST").as_deref() == Ok("0") { return false; }
5724 match w {
5725 GpuTensor::Quant { qtype, .. } => matches!(*qtype,
5732 QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q3_K | QT_NVFP4 | QT_F8_E4M3
5733 | QT_F8_E4M3_BLK | QT_Q4_0)
5734 || (*qtype == QT_IQ4_XS && Self::iq_fast_enabled()),
5735 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
5736 }
5737 }
5738
5739 pub fn matmul_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
5744 x_fallback: &CudaSlice<f32>, m: usize)
5745 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5746 use crate::model::GpuTensor;
5747 let x_raw_ok = x_fallback.len() >= m * w.in_features();
5753 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5756 if let Some(y) = self.try_fp8_gemm(w, x_fallback, m)? { return Ok(y); }
5757 if let Some(y) = self.try_fp8_blk_mmq(w, x_fallback, m)? { return Ok(y); }
5760 if let Some(y) = self.try_f16_gemm(w, x_fallback, m)? { return Ok(y); }
5762 }
5763 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5769 if let Some(y) = self.try_e4m3_blk_prefill(w, x_fallback, m)? { return Ok(y); }
5770 }
5771 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
5772 if m >= 16 && w.out_features() >= 128 && self.mmq_supports(w) && !self.verify_exact_on()
5777 && x_raw_ok {
5778 return self.qmatvec_mmq(w, x_fallback, m);
5779 }
5780 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5783 if let Some(y) = self.try_fp4_gemm(w, x_fallback, m, w.in_features(), w.out_features())? {
5784 return Ok(y);
5785 }
5786 }
5787 if m >= 16 && self.gemm_supports(w) && !self.verify_exact_on() {
5790 return self.qmatvec_gemm(w, aq, ad, m);
5791 }
5792 if !self.uses_q8_1_fast(w) { return self.matmul(w, x_fallback, m); }
5793 let in_f = w.in_features();
5794 let out_f = w.out_features();
5795 let (bytes, qtype, row_bytes, scale, rp) = match w {
5796 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
5797 _ => unreachable!("uses_q8_1_fast guaranteed Quant"),
5798 };
5799 let (mbytes, mrp) = match w {
5802 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5803 _ => (bytes, rp),
5804 };
5805 if m == 1 && self.mmvq_supports(qtype) {
5809 return self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, mrp);
5810 }
5811 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5824 && std::env::var("MEMRA_NO_BATCHED").is_err()
5825 && (m <= 4 || Self::b8_enabled())
5826 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_NVFP4
5830 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_F8_E4M3 || qtype == QT_Q8_0) {
5831 let mcols = Self::batched_mcols(m);
5832 return self.qmatvec_mmvq_batched(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, mrp);
5833 }
5834 if qtype == QT_F8_E4M3 || qtype == QT_Q4_0 {
5840 let (b2, r2) = if qtype == QT_Q4_0 { (mbytes, mrp) } else { (bytes, rp) };
5841 return self.qmatvec_mmvq(b2, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, r2);
5842 }
5843 let name = match qtype {
5844 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
5845 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
5846 QT_Q3_K => "qmatvec_q3_K_dp4a",
5847 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
5848 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
5849 _ => unreachable!(),
5850 };
5851 let f = self.func(name);
5852 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 };
5854 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
5855 let __s_b = self.gpu.stream();
5856 let mut b = __s_b.launch_builder(&f);
5857 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
5858 unsafe { b.launch(cfg)?; }
5859 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
5860 Ok(y)
5861 }
5862
5863 pub fn matmul_decode_exact(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
5871 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5872 use crate::model::GpuTensor;
5873 if let GpuTensor::Float { data, .. } = w {
5881 return self.linear_decode_exact(x, data, m, w.in_features(), w.out_features());
5882 }
5883 if let GpuTensor::FloatBf16 { data, .. } = w {
5886 let (in_f, out_f) = (w.in_features(), w.out_features());
5887 return self.linear_bf16_chunked(x, data, m, in_f, out_f, true);
5888 }
5889 if !self.uses_q8_1_fast(w) { return self.matmul(w, x, m); }
5890 let in_f = w.in_features();
5891 let out_f = w.out_features();
5892 let (bytes, qtype, row_bytes, scale, rp) = match w {
5893 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
5894 _ => return self.matmul(w, x, m),
5895 };
5896 let (bytes, rp) = match w {
5899 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5900 _ => (bytes, rp),
5901 };
5902 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5903 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
5907 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5916 && std::env::var("MEMRA_NO_BATCHED").is_err()
5917 && (m <= 4 || Self::b8_enabled())
5918 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
5921 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
5922 let mcols = Self::batched_mcols(m);
5923 return self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
5924 }
5925 if self.mmvq_supports(qtype) {
5926 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
5929 }
5930 self.matmul_pre(w, &aq, &ad, x, m)
5933 }
5934
5935 pub fn matmul_decode_exact_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
5945 ad: &CudaSlice<f32>, m: usize)
5946 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5947 use crate::model::GpuTensor;
5948 debug_assert!(self.uses_q8_1_fast(w),
5949 "matmul_decode_exact_pre: caller must guarantee q8_1-fast");
5950 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
5952 let in_f = w.in_features();
5953 let out_f = w.out_features();
5954 let (bytes, qtype, row_bytes, scale, rp) = match w {
5955 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
5956 (bytes, *qtype, *row_bytes, *scale, *rp),
5957 _ => return Err("matmul_decode_exact_pre: Quant tensor required (q8_1-fast contract)".into()),
5958 };
5959 let (bytes, rp) = match w {
5961 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5962 _ => (bytes, rp),
5963 };
5964 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5966 && std::env::var("MEMRA_NO_BATCHED").is_err()
5967 && (m <= 4 || Self::b8_enabled())
5968 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
5969 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
5970 let mcols = Self::batched_mcols(m);
5971 return self.qmatvec_mmvq_batched(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
5972 }
5973 if self.mmvq_supports(qtype) {
5974 return self.qmatvec_mmvq(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
5975 }
5976 let x0 = self.zeros(0)?;
5979 self.matmul_pre(w, aq, ad, &x0, m)
5980 }
5981
5982 pub fn matmul_decode_exact_dual_pre(&self, w0: &crate::model::GpuTensor,
5991 w1: &crate::model::GpuTensor,
5992 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
5993 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
5994 use crate::model::GpuTensor;
5995 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5996 let on = *ON.get_or_init(|| {
5997 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
5998 });
5999 if !on || !(2..=7).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
6000 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
6001 return Ok(None);
6002 }
6003 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6008 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6009 if w1.in_features() != in_f || w1.out_features() != out_f {
6010 return Ok(None);
6011 }
6012 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
6013 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
6014 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
6015 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
6016 (b0, b1, *rb0, *s0, *s1, *rp0),
6017 _ => return Ok(None),
6018 };
6019 if m > 4 && !(rp && Self::b8_enabled()
6022 && std::env::var("MEMRA_B567").as_deref() != Ok("0")) {
6023 return Ok(None);
6024 }
6025 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, aq, ad, m, in_f, out_f, row_bytes, rp)?;
6026 Ok(Some(((y0, s0), (y1, s1))))
6027 }
6028
6029 pub fn matmul_decode_exact_dual(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6045 x: &CudaSlice<f32>, m: usize)
6046 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6047 use crate::model::GpuTensor;
6048 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6049 let on = *ON.get_or_init(|| {
6050 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
6051 });
6052 if !on || !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
6053 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
6054 return Ok(None);
6055 }
6056 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6061 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6062 if w1.in_features() != in_f || w1.out_features() != out_f {
6063 return Ok(None);
6064 }
6065 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
6066 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
6067 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
6068 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
6069 (b0, b1, *rb0, *s0, *s1, *rp0),
6070 _ => return Ok(None),
6071 };
6072 if std::env::var("MEMRA_DEBUG").is_ok() {
6075 static ONCE: std::sync::Once = std::sync::Once::new();
6076 ONCE.call_once(|| eprintln!("[memra] dual gate+up batched ENGAGED (m={m} rp={rp})"));
6077 }
6078 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6079 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, &aq, &ad, m, in_f, out_f, row_bytes, rp)?;
6080 let mut y0 = y0;
6081 let mut y1 = y1;
6082 if s0 != 1.0 { self.scale_inplace(&mut y0, s0, m * out_f)?; }
6083 if s1 != 1.0 { self.scale_inplace(&mut y1, s1, m * out_f)?; }
6084 Ok(Some((y0, y1)))
6085 }
6086
6087 #[allow(clippy::too_many_arguments)]
6092 pub fn qmatvec_batched_dual_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6093 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6094 m: usize, in_f: usize, out_f: usize, row_bytes: usize, rp: bool)
6095 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6096 const ROWS_PER_BLOCK: u32 = 4;
6097 let mcols = Self::batched_mcols(m);
6098 let (name, rows_per_block) = match (mcols, rp, m) {
6101 (2, false, _) => ("qmatvec_nvfp4_mmvq_dual_b2", ROWS_PER_BLOCK),
6102 (4, false, _) => ("qmatvec_nvfp4_mmvq_dual_b4_r2", ROWS_PER_BLOCK * 2),
6103 (2, true, _) => ("qmatvec_nvfp4_mmvq_dual_b2_rp", ROWS_PER_BLOCK),
6104 (4, true, _) => ("qmatvec_nvfp4_mmvq_dual_b4_rpr2", ROWS_PER_BLOCK * 2),
6105 (8, true, 5) => ("qmatvec_nvfp4_mmvq_dual_b5_rpr2", ROWS_PER_BLOCK * 2),
6106 (8, true, 6) => ("qmatvec_nvfp4_mmvq_dual_b6_rpr2", ROWS_PER_BLOCK * 2),
6107 (8, true, 7) => ("qmatvec_nvfp4_mmvq_dual_b7_rpr2", ROWS_PER_BLOCK * 2),
6108 _ => return Err(format!("qmatvec_batched_dual_raw: no dual kernel for m {m}").into()),
6109 };
6110 let f = self.func(name);
6111 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
6112 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
6113 let cfg = LaunchConfig {
6114 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6115 block_dim: (32, ROWS_PER_BLOCK, 1),
6116 shared_mem_bytes: 0,
6117 };
6118 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
6119 let __s_b = self.gpu.stream();
6120 let mut b = __s_b.launch_builder(&f);
6121 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6122 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
6123 unsafe { b.launch(cfg)?; }
6124 Ok((y0, y1))
6125 }
6126
6127 pub fn matmul_pre_dual_noscale(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6139 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6140 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
6141 use crate::model::GpuTensor;
6142 if m != 1 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6143 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6153 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6154 if w1.in_features() != in_f || w1.out_features() != out_f { return Ok(None); }
6155 let no_mirror = |w: &crate::model::GpuTensor| {
6168 !matches!(w, GpuTensor::Quant { rp4: Some(_), .. })
6169 };
6170 if self.q8_ffn_fuse2_on()
6171 && no_mirror(w0) && no_mirror(w1)
6172 && let Some([p0, p1]) = self.q8_fused_params(&[w0, w1])
6173 {
6174 let (y0, y1) = self.q8_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2)?;
6175 return Ok(Some(((y0, 1.0), (y1, 1.0))));
6176 }
6177 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6187 let (y0, y1) = self.e4m3_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2,
6188 1.0, 1.0)?;
6189 return Ok(Some(((y0, p0.3), (y1, p1.3))));
6190 }
6191 let (b0, q0, rb0, s0, rp0) = match w0 {
6192 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6193 _ => return Ok(None),
6194 };
6195 let (b1, q1, rb1, s1, rp1) = match w1 {
6196 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6197 _ => return Ok(None),
6198 };
6199 if q0 != QT_NVFP4 || q1 != QT_NVFP4 || rb0 != rb1 || rp0 != rp1 { return Ok(None); }
6200 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
6202 let rows_per_block = ROWS_PER_BLOCK * RPW;
6203 let f = self.func(if rp0 { "qmatvec_nvfp4_mmvq_dual_mr2_rp" } else { "qmatvec_nvfp4_mmvq_dual_mr2" });
6204 let mut y0 = self.alloc_uninit::<f32>(out_f)?;
6205 let mut y1 = self.alloc_uninit::<f32>(out_f)?;
6206 let cfg = LaunchConfig {
6207 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6208 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0,
6209 };
6210 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, rb0 as i64);
6211 let one = 1.0f32;
6214 let __s_b = self.gpu.stream();
6215 let mut b = __s_b.launch_builder(&f);
6216 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6217 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&one).arg(&one);
6218 unsafe { b.launch(cfg)?; }
6219 Ok(Some(((y0, s0), (y1, s1))))
6220 }
6221
6222 pub fn matmul_q8_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6230 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6231 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6232 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6238 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(),
6239 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6240 }
6241 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6242 Ok(Some(self.q8_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6243 }
6244
6245 #[allow(clippy::too_many_arguments)]
6246 fn q8_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6247 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6248 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6249 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6250 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6252 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6253 let f = self.func("qmatvec_q8_0_mmvq_fused2");
6254 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6255 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6256 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6257 shared_mem_bytes: 0 };
6258 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
6259 let __s_b = self.gpu.stream();
6260 let mut b = __s_b.launch_builder(&f);
6261 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6262 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl);
6263 unsafe { b.launch(cfg)?; }
6264 Ok((y0, y1))
6265 }
6266
6267 pub fn matmul_q8_fused2_x(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6273 x: &CudaSlice<f32>)
6274 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6275 if !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6276 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6277 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6278 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(),
6279 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6280 }
6281 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6282 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6283 Ok(Some(self.q8_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6284 }
6285
6286 #[allow(clippy::too_many_arguments)]
6289 pub fn qmatvec_q8_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
6290 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6291 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6292 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6293 self.q8_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes)
6294 }
6295
6296 pub fn matmul_q4_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6302 w2: &crate::model::GpuTensor,
6303 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6304 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6305 use crate::model::GpuTensor;
6306 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6307 match w {
6308 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6309 Some((*row_bytes, w.out_features())),
6310 _ => None,
6311 }
6312 };
6313 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6314 else { return Ok(None) };
6315 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6316 return Ok(None);
6317 }
6318 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6322 match w {
6323 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6324 Some(m) => (m, true),
6325 None => (bytes, *rp),
6326 },
6327 _ => unreachable!(),
6328 }
6329 }
6330 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6331 if rp0 != rp1 || rp1 != rp2 { return Ok(None); }
6332 let rp = rp0;
6333 let rpb: u32 = 4;
6334 let mr1 = rp && Self::q40_mr1_on();
6338 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6339 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6340 let grid = nb(o0) + nb(o1) + nb(o2);
6341 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6342 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6343 let mut y2 = self.alloc_uninit::<f32>(o2)?;
6344 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6345 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6346 else { "qmatvec_q4_0_mmvq_fused3" });
6347 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6348 let inf = w0.in_features() as i32;
6349 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6350 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6351 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6354 {
6355 use cudarc::driver::{DevicePtr, DevicePtrMut};
6356 let s = &self.gpu.stream();
6357 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6358 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6359 let (pad, _g4) = ad.device_ptr(s);
6360 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6361 let (py2, _g7) = y2.device_ptr_mut(s);
6362 let mut ps = [
6363 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6364 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6365 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6366 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6367 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6368 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6369 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6370 &r2 as *const _ as *mut _,
6371 ];
6372 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6373 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6374 }
6375 return Ok(Some((y0, y1, y2)));
6376 }
6377 let __s_b = self.gpu.stream();
6378 let mut b = __s_b.launch_builder(&f);
6379 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6380 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6381 unsafe { b.launch(cfg)?; }
6382 Ok(Some((y0, y1, y2)))
6383 }
6384
6385 #[allow(clippy::too_many_arguments)]
6388 pub fn matmul_q4_fused3_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6389 w2: &crate::model::GpuTensor,
6390 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6391 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>,
6392 y2: &mut CudaSlice<f32>)
6393 -> Result<bool, Box<dyn std::error::Error>> {
6394 use crate::model::GpuTensor;
6395 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6396 match w {
6397 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6398 Some((*row_bytes, w.out_features())),
6399 _ => None,
6400 }
6401 };
6402 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6403 else { return Ok(false) };
6404 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6405 return Ok(false);
6406 }
6407 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6408 match w {
6409 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6410 Some(m) => (m, true),
6411 None => (bytes, *rp),
6412 },
6413 _ => unreachable!(),
6414 }
6415 }
6416 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6417 if rp0 != rp1 || rp1 != rp2 { return Ok(false); }
6418 let rp = rp0;
6419 let rpb: u32 = 4;
6420 let mr1 = rp && Self::q40_mr1_on();
6421 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6422 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6423 let grid = nb(o0) + nb(o1) + nb(o2);
6424 debug_assert!(y0.len() >= o0 && y1.len() >= o1 && y2.len() >= o2);
6425 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6426 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6427 else { "qmatvec_q4_0_mmvq_fused3" });
6428 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6429 let inf = w0.in_features() as i32;
6430 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6431 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6432 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6434 use cudarc::driver::{DevicePtr, DevicePtrMut};
6435 let s = &self.gpu.stream();
6436 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6437 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6438 let (pad, _g4) = ad.device_ptr(s);
6439 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6440 let (py2, _g7) = y2.device_ptr_mut(s);
6441 let mut ps = [
6442 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6443 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6444 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6445 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6446 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6447 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6448 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6449 &r2 as *const _ as *mut _,
6450 ];
6451 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6452 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6453 return Ok(true);
6454 }
6455 let __s_b = self.gpu.stream();
6456 let mut b = __s_b.launch_builder(&f);
6457 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1).arg(&mut *y2)
6458 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6459 unsafe { b.launch(cfg)?; }
6460 Ok(true)
6461 }
6462
6463 pub fn matmul_q4_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6465 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6466 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6467 use crate::model::GpuTensor;
6468 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6469 match w {
6470 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6471 Some((*row_bytes, w.out_features())),
6472 _ => None,
6473 }
6474 };
6475 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6476 if w0.in_features() != w1.in_features() { return Ok(None); }
6477 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6479 match w {
6480 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6481 Some(m) => (m, true),
6482 None => (bytes, *rp),
6483 },
6484 _ => unreachable!(),
6485 }
6486 }
6487 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6488 if rp0 != rp1 { return Ok(None); }
6489 let rp = rp0;
6490 let rpb: u32 = 4;
6491 let mr1 = rp && Self::q40_mr1_on();
6493 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6494 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6495 let grid = nb(o0) + nb(o1);
6496 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6497 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6498 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6499 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6500 else { "qmatvec_q4_0_mmvq_fused2" });
6501 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6502 let inf = w0.in_features() as i32;
6503 let (oo0, oo1) = (o0 as i32, o1 as i32);
6504 let (r0, r1) = (rb0 as i64, rb1 as i64);
6505 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6507 {
6508 use cudarc::driver::{DevicePtr, DevicePtrMut};
6509 let s = &self.gpu.stream();
6510 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6511 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6512 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6513 let mut ps = [
6514 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6515 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6516 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6517 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6518 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6519 &r1 as *const _ as *mut _,
6520 ];
6521 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6522 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6523 }
6524 return Ok(Some((y0, y1)));
6525 }
6526 let __s_b = self.gpu.stream();
6527 let mut b = __s_b.launch_builder(&f);
6528 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6529 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6530 unsafe { b.launch(cfg)?; }
6531 Ok(Some((y0, y1)))
6532 }
6533
6534 pub fn matmul_q4_fused2_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6536 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6537 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>)
6538 -> Result<bool, Box<dyn std::error::Error>> {
6539 use crate::model::GpuTensor;
6540 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6541 match w {
6542 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6543 Some((*row_bytes, w.out_features())),
6544 _ => None,
6545 }
6546 };
6547 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(false) };
6548 if w0.in_features() != w1.in_features() { return Ok(false); }
6549 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6550 match w {
6551 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6552 Some(m) => (m, true),
6553 None => (bytes, *rp),
6554 },
6555 _ => unreachable!(),
6556 }
6557 }
6558 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6559 if rp0 != rp1 { return Ok(false); }
6560 let rp = rp0;
6561 let rpb: u32 = 4;
6562 let mr1 = rp && Self::q40_mr1_on();
6563 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6564 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6565 let grid = nb(o0) + nb(o1);
6566 debug_assert!(y0.len() >= o0 && y1.len() >= o1);
6567 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6568 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6569 else { "qmatvec_q4_0_mmvq_fused2" });
6570 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6571 let inf = w0.in_features() as i32;
6572 let (oo0, oo1) = (o0 as i32, o1 as i32);
6573 let (r0, r1) = (rb0 as i64, rb1 as i64);
6574 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6576 use cudarc::driver::{DevicePtr, DevicePtrMut};
6577 let s = &self.gpu.stream();
6578 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6579 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6580 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6581 let mut ps = [
6582 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6583 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6584 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6585 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6586 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6587 &r1 as *const _ as *mut _,
6588 ];
6589 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6590 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6591 return Ok(true);
6592 }
6593 let __s_b = self.gpu.stream();
6594 let mut b = __s_b.launch_builder(&f);
6595 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1)
6596 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6597 unsafe { b.launch(cfg)?; }
6598 Ok(true)
6599 }
6600
6601 pub fn matmul_q4_fused2_batched(&self, w0: &crate::model::GpuTensor,
6606 w1: &crate::model::GpuTensor,
6607 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6608 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6609 use crate::model::GpuTensor;
6610 if m < 2 || m > 8 { return Ok(None); }
6611 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6612 match w {
6613 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6614 Some((*row_bytes, w.out_features())),
6615 _ => None,
6616 }
6617 };
6618 let (Some((rb0, o0)), Some((_rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6619 if w0.in_features() != w1.in_features() { return Ok(None); }
6620 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6621 match w {
6622 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6623 Some(mr) => (mr, true),
6624 None => (bytes, *rp),
6625 },
6626 _ => unreachable!(),
6627 }
6628 }
6629 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6630 if !rp0 || !rp1 { return Ok(None); }
6631 let mcols = Self::batched_mcols(m);
6632 let rpb: u32 = 4;
6633 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6634 let grid = nb(o0) + nb(o1);
6635 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6636 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6637 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f2_rp",
6638 4 => "qmatvec_q4_0_mmvq_b4_f2_rp",
6639 _ => "qmatvec_q4_0_mmvq_b8_f2_rp" });
6640 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6641 shared_mem_bytes: 0 };
6642 let inf = w0.in_features() as i32;
6643 let (oo0, oo1, mi) = (o0 as i32, o1 as i32, m as i32);
6644 let rb = rb0 as i64;
6645 let __s_b = self.gpu.stream();
6646 let mut b = __s_b.launch_builder(&f);
6647 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6648 .arg(&inf).arg(&oo0).arg(&oo1).arg(&mi).arg(&rb);
6649 unsafe { b.launch(cfg)?; }
6650 Ok(Some((y0, y1)))
6651 }
6652
6653 #[allow(clippy::too_many_arguments)]
6656 pub fn matmul_q4_fused3_batched(&self, w0: &crate::model::GpuTensor,
6657 w1: &crate::model::GpuTensor, w2: &crate::model::GpuTensor,
6658 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6659 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6660 use crate::model::GpuTensor;
6661 if m < 2 || m > 8 { return Ok(None); }
6662 let q4 = |w: &GpuTensor| -> Option<usize> {
6663 match w {
6664 GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0 => Some(w.out_features()),
6665 _ => None,
6666 }
6667 };
6668 let (Some(o0), Some(o1), Some(o2)) = (q4(w0), q4(w1), q4(w2)) else { return Ok(None) };
6669 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6670 return Ok(None);
6671 }
6672 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6673 match w {
6674 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6675 Some(mr) => (mr, true),
6676 None => (bytes, *rp),
6677 },
6678 _ => unreachable!(),
6679 }
6680 }
6681 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6682 if !rp0 || !rp1 || !rp2 { return Ok(None); }
6683 let mcols = Self::batched_mcols(m);
6684 let rpb: u32 = 4;
6685 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6686 let grid = nb(o0) + nb(o1) + nb(o2);
6687 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6688 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6689 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
6690 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f3_rp",
6691 4 => "qmatvec_q4_0_mmvq_b4_f3_rp",
6692 _ => "qmatvec_q4_0_mmvq_b8_f3_rp" });
6693 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6694 shared_mem_bytes: 0 };
6695 let inf = w0.in_features() as i32;
6696 let (oo0, oo1, oo2, mi) = (o0 as i32, o1 as i32, o2 as i32, m as i32);
6697 let rb = 0i64;
6698 let __s_b = self.gpu.stream();
6699 let mut b = __s_b.launch_builder(&f);
6700 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6701 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&mi).arg(&rb);
6702 unsafe { b.launch(cfg)?; }
6703 Ok(Some((y0, y1, y2)))
6704 }
6705
6706 pub fn matmul_q8_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6707 w2: &crate::model::GpuTensor,
6708 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6709 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6710 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
6713 return Ok(Some(self.e4m3_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6714 p0.1, p1.1, p2.1, p0.2,
6715 p0.3, p1.3, p2.3)?));
6716 }
6717 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
6718 Ok(Some(self.q8_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6719 p0.1, p1.1, p2.1, p0.2)?))
6720 }
6721
6722 #[allow(clippy::too_many_arguments)]
6723 fn q8_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6724 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6725 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
6726 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6727 const ROWS_PER_BLOCK: u32 = 4;
6728 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6729 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6730 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6731 let f = self.func("qmatvec_q8_0_mmvq_fused3");
6732 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6733 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6734 let mut y2 = self.alloc_uninit::<f32>(out2)?;
6735 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6736 shared_mem_bytes: 0 };
6737 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32, row_bytes as i64);
6738 let __s_b = self.gpu.stream();
6739 let mut b = __s_b.launch_builder(&f);
6740 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6741 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl);
6742 unsafe { b.launch(cfg)?; }
6743 Ok((y0, y1, y2))
6744 }
6745
6746 #[allow(clippy::too_many_arguments)]
6748 pub fn qmatvec_q8_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6749 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
6750 out2: usize, row_bytes: usize)
6751 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6752 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6753 self.q8_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes)
6754 }
6755
6756 pub fn matmul_q8_fused2_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6767 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6768 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6769 if !(2..=8).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
6773 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6776 if m > 4 && !Self::b8_enabled() { return Ok(None); }
6777 return Ok(Some(self.e4m3_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(),
6778 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6779 }
6780 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6781 Ok(Some(self.q8_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(), p0.1, p1.1, p0.2)?))
6782 }
6783
6784 #[allow(clippy::too_many_arguments)]
6785 fn q8_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6786 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6787 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6788 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6789 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6791 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6792 let f = self.func(match Self::batched_mcols(m) {
6793 2 => "qmatvec_q8_0_mmvq_fused2_b2",
6794 4 => "qmatvec_q8_0_mmvq_fused2_b4",
6795 _ => "qmatvec_q8_0_mmvq_fused2_b8",
6797 });
6798 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6799 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6800 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6801 shared_mem_bytes: 0 };
6802 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32, row_bytes as i64);
6803 let __s_b = self.gpu.stream();
6804 let mut b = __s_b.launch_builder(&f);
6805 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6806 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
6807 unsafe { b.launch(cfg)?; }
6808 Ok((y0, y1))
6809 }
6810
6811 #[allow(clippy::too_many_arguments)]
6814 pub fn qmatvec_q8_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6815 x: &CudaSlice<f32>, m: usize,
6816 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6817 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6818 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6819 self.q8_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes)
6820 }
6821
6822 #[allow(clippy::too_many_arguments)]
6825 pub fn matmul_q8_fused3_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6826 w2: &crate::model::GpuTensor,
6827 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6828 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6829 if !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
6830 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
6831 return Ok(Some(self.e4m3_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
6832 p0.1, p1.1, p2.1, p0.2,
6833 p0.3, p1.3, p2.3)?));
6834 }
6835 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
6836 Ok(Some(self.q8_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
6837 p0.1, p1.1, p2.1, p0.2)?))
6838 }
6839
6840 #[allow(clippy::too_many_arguments)]
6841 fn q8_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6842 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6843 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
6844 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6845 const ROWS_PER_BLOCK: u32 = 4;
6846 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6847 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6848 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6849 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_q8_0_mmvq_fused3_b2" }
6850 else { "qmatvec_q8_0_mmvq_fused3_b4" });
6851 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6852 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6853 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
6854 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6855 shared_mem_bytes: 0 };
6856 let (inf, o0, o1, o2, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
6857 m as i32, row_bytes as i64);
6858 let __s_b = self.gpu.stream();
6859 let mut b = __s_b.launch_builder(&f);
6860 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6861 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&mi).arg(&rbl);
6862 unsafe { b.launch(cfg)?; }
6863 Ok((y0, y1, y2))
6864 }
6865
6866 #[allow(clippy::too_many_arguments)]
6868 pub fn qmatvec_q8_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6869 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
6870 out1: usize, out2: usize, row_bytes: usize)
6871 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6872 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6873 self.q8_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes)
6874 }
6875
6876 pub fn q8_ffn_fuse2_on(&self) -> bool {
6880 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6881 *ON.get_or_init(|| std::env::var("MEMRA_Q8_FFN_FUSE2").as_deref() != Ok("0"))
6882 }
6883
6884 #[allow(clippy::type_complexity)]
6890 fn q8_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
6891 -> Option<[(&'w CudaSlice<u8>, usize, usize); N]> {
6892 use crate::model::GpuTensor;
6893 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return None; }
6894 if std::env::var("MEMRA_Q8_DUAL").is_ok_and(|v| v == "0") { return None; }
6895 let in_f = ws[0].in_features();
6896 let mut out: [Option<(&CudaSlice<u8>, usize, usize)>; N] = [None; N];
6897 for (i, w) in ws.iter().enumerate() {
6898 match w {
6899 GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. }
6900 if *qtype == QT_Q8_0 && *scale == 1.0 && w.in_features() == in_f =>
6901 out[i] = Some((bytes, w.out_features(), *row_bytes)),
6902 _ => return None,
6903 }
6904 }
6905 Some(out.map(|o| o.unwrap()))
6906 }
6907
6908 pub fn e4m3_dual_on(&self) -> bool {
6911 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6912 *ON.get_or_init(|| std::env::var("MEMRA_E4M3_DUAL").as_deref() != Ok("0"))
6913 }
6914
6915 #[allow(clippy::type_complexity)]
6927 fn e4m3_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
6928 -> Option<[(&'w CudaSlice<u8>, usize, usize, f32); N]> {
6929 use crate::model::GpuTensor;
6930 if !self.e4m3_dual_on() { return None; }
6931 let in_f = ws[0].in_features();
6932 let mut out: [Option<(&CudaSlice<u8>, usize, usize, f32)>; N] = [None; N];
6933 for (i, w) in ws.iter().enumerate() {
6934 match w {
6935 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, rp4, .. }
6936 if *qtype == QT_F8_E4M3 && w.in_features() == in_f
6937 && *row_bytes == in_f && !*rp && rp4.is_none() =>
6938 out[i] = Some((bytes, w.out_features(), *row_bytes, *scale)),
6939 _ => return None,
6940 }
6941 }
6942 Some(out.map(|o| o.unwrap()))
6943 }
6944
6945 #[allow(clippy::too_many_arguments)]
6949 fn e4m3_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6950 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6951 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
6952 ws0: f32, ws1: f32)
6953 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6954 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6956 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6957 let f = self.func("qmatvec_e4m3_mmvq_fused2");
6958 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6959 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6960 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6961 shared_mem_bytes: 0 };
6962 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
6963 let __s_b = self.gpu.stream();
6964 let mut b = __s_b.launch_builder(&f);
6965 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6966 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl).arg(&ws0).arg(&ws1);
6967 unsafe { b.launch(cfg)?; }
6968 Ok((y0, y1))
6969 }
6970
6971 #[allow(clippy::too_many_arguments)]
6973 fn e4m3_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6974 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6975 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
6976 ws0: f32, ws1: f32, ws2: f32)
6977 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6978 const ROWS_PER_BLOCK: u32 = 4;
6979 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6980 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6981 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6982 let f = self.func("qmatvec_e4m3_mmvq_fused3");
6983 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6984 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6985 let mut y2 = self.alloc_uninit::<f32>(out2)?;
6986 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6987 shared_mem_bytes: 0 };
6988 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
6989 row_bytes as i64);
6990 let __s_b = self.gpu.stream();
6991 let mut b = __s_b.launch_builder(&f);
6992 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6993 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl).arg(&ws0).arg(&ws1).arg(&ws2);
6994 unsafe { b.launch(cfg)?; }
6995 Ok((y0, y1, y2))
6996 }
6997
6998 #[allow(clippy::too_many_arguments)]
7002 fn e4m3_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7003 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
7004 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7005 ws0: f32, ws1: f32)
7006 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7007 const ROWS_PER_BLOCK: u32 = 4;
7008 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7009 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7010 let f = self.func(match Self::batched_mcols(m) {
7011 2 => "qmatvec_e4m3_mmvq_fused2_b2",
7012 4 => "qmatvec_e4m3_mmvq_fused2_b4",
7013 _ => "qmatvec_e4m3_mmvq_fused2_b8",
7014 });
7015 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7016 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7017 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7018 shared_mem_bytes: 0 };
7019 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32,
7020 row_bytes as i64);
7021 let __s_b = self.gpu.stream();
7022 let mut b = __s_b.launch_builder(&f);
7023 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
7024 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
7025 unsafe { b.launch(cfg)?; }
7026 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
7027 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
7028 Ok((y0, y1))
7029 }
7030
7031 #[allow(clippy::too_many_arguments)]
7033 fn e4m3_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7034 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
7035 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
7036 ws0: f32, ws1: f32, ws2: f32)
7037 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7038 const ROWS_PER_BLOCK: u32 = 4;
7039 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7040 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7041 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
7042 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_e4m3_mmvq_fused3_b2" }
7043 else { "qmatvec_e4m3_mmvq_fused3_b4" });
7044 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7045 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7046 let mut y2 = self.alloc_uninit::<f32>(m * 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, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
7050 m as i32, 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(&mi).arg(&rbl);
7055 unsafe { b.launch(cfg)?; }
7056 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
7057 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
7058 if ws2 != 1.0 { self.scale_inplace(&mut y2, ws2, m * out2)?; }
7059 Ok((y0, y1, y2))
7060 }
7061
7062 pub fn qmatvec_e4m3_blk_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7072 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7073 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7074 scale_cols: usize)
7075 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7076 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,
7078 scale_cols, &mut y)?;
7079 Ok(y)
7080 }
7081
7082 #[allow(clippy::too_many_arguments)]
7084 pub fn qmatvec_e4m3_blk_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7085 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7086 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7087 scale_cols: usize, y: &mut CudaSlice<f32>)
7088 -> Result<(), Box<dyn std::error::Error>> {
7089 const ROWS_PER_BLOCK: u32 = 4; let f = self.func("qmatvec_e4m3_blk_mmvq");
7091 let cfg = LaunchConfig {
7092 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), m as u32, 1),
7093 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7096 let (inf, outf, mi, rb, sc) =
7097 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7098 let __s_b = self.gpu.stream();
7099 let mut b = __s_b.launch_builder(&f);
7100 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut *y)
7101 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7102 unsafe { b.launch(cfg)?; }
7103 Ok(())
7104 }
7105
7106 #[allow(clippy::too_many_arguments)]
7112 pub fn qmatvec_e4m3_blk_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7113 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7114 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7115 scale_cols: usize, mcols: usize)
7116 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7117 const ROWS_PER_BLOCK: u32 = 4; debug_assert!(mcols >= m, "blk batched: mcols {mcols} < m {m}");
7119 let name = match mcols {
7120 2 => "qmatvec_e4m3_blk_mmvq_b2",
7121 4 => "qmatvec_e4m3_blk_mmvq_b4",
7122 8 => "qmatvec_e4m3_blk_mmvq_b8",
7123 16 => "qmatvec_e4m3_blk_mmvq_b16",
7124 _ => return Err(format!("qmatvec_e4m3_blk_mmvq_batched: no kernel for mcols {mcols}").into()),
7125 };
7126 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7127 let f = self.func(name);
7128 let cfg = LaunchConfig {
7129 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
7130 block_dim: (32, ROWS_PER_BLOCK, 1),
7131 shared_mem_bytes: 0,
7132 };
7133 let (inf, outf, mi, rb, sc) =
7134 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7135 let __s_b = self.gpu.stream();
7136 let mut b = __s_b.launch_builder(&f);
7137 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut y)
7138 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7139 unsafe { b.launch(cfg)?; }
7140 Ok(y)
7141 }
7142
7143 #[allow(clippy::too_many_arguments)]
7146 pub fn qmatvec_e4m3_blk_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7147 scales: &CudaSlice<f32>, m: usize, in_f: usize,
7148 out_f: usize, row_bytes: usize, scale_cols: usize,
7149 mcols: usize)
7150 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7151 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7152 self.qmatvec_e4m3_blk_mmvq_batched(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes,
7153 scale_cols, mcols)
7154 }
7155
7156 #[allow(clippy::too_many_arguments)]
7159 pub fn qmatvec_e4m3_blk_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7160 scales: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
7161 row_bytes: usize, scale_cols: usize)
7162 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7163 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7164 self.qmatvec_e4m3_blk_mmvq(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols)
7165 }
7166
7167 #[allow(clippy::too_many_arguments)]
7170 pub fn qmatvec_e4m3_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
7171 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7172 ws0: f32, ws1: f32)
7173 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7174 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7175 self.e4m3_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes, ws0, ws1)
7176 }
7177
7178 #[allow(clippy::too_many_arguments)]
7179 pub fn qmatvec_e4m3_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7180 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
7181 out2: usize, row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7182 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7183 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7184 self.e4m3_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes,
7185 ws0, ws1, ws2)
7186 }
7187
7188 #[allow(clippy::too_many_arguments)]
7189 pub fn qmatvec_e4m3_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7190 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
7191 out1: usize, row_bytes: usize, ws0: f32, ws1: f32)
7192 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7193 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7194 self.e4m3_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes, ws0, ws1)
7195 }
7196
7197 #[allow(clippy::too_many_arguments)]
7198 pub fn qmatvec_e4m3_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7199 b2: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
7200 in_f: usize, out0: usize, out1: usize, out2: usize,
7201 row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7202 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7203 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7204 self.e4m3_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes,
7205 ws0, ws1, ws2)
7206 }
7207
7208 fn try_e4m3_blk_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
7219 ad: &CudaSlice<f32>, m: usize)
7220 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7221 use crate::model::GpuTensor;
7222 if let GpuTensor::Quant { bytes, qtype, row_bytes, blk: Some(g), .. } = w {
7223 if *qtype == QT_F8_E4M3_BLK {
7224 if (2..=16).contains(&m) && std::env::var("MEMRA_NO_BATCHED").is_err()
7230 && (m <= 4 || Self::b8_enabled()) {
7231 let mcols = Self::batched_mcols(m);
7232 return Ok(Some(self.qmatvec_e4m3_blk_mmvq_batched(
7233 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7234 *row_bytes, g.cols, mcols)?));
7235 }
7236 return Ok(Some(self.qmatvec_e4m3_blk_mmvq(
7237 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7238 *row_bytes, g.cols)?));
7239 }
7240 }
7241 Ok(None)
7242 }
7243
7244 fn try_e4m3_blk_prefill(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
7291 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7292 use crate::model::GpuTensor;
7293 let GpuTensor::Quant { bytes, qtype, blk: Some(g), .. } = w else { return Ok(None) };
7294 if *qtype != QT_F8_E4M3_BLK { return Ok(None) }
7295 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(Some(y)); }
7300 let (in_f, out_f) = (w.in_features(), w.out_features());
7301 let slab = self.fp8_blk_dequant_q8_0_dev(bytes, &g.scales, out_f, in_f)?;
7302 let tmp = GpuTensor::Quant {
7303 bytes: slab,
7304 qtype: QT_Q8_0,
7305 row_bytes: in_f / 32 * 34,
7306 ne: vec![in_f as u64, out_f as u64],
7307 scale: 1.0,
7308 rp: false,
7309 #[cfg(memra_cutlass)]
7310 cutlass: None,
7311 fp8: None, blk: None, f16: None, rp4: None,
7312 };
7313 Ok(Some(self.matmul(&tmp, x, m)?))
7315 }
7316
7317 pub fn matmul_pre_noscale(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7318 m: usize) -> Result<Option<(CudaSlice<f32>, f32)>, Box<dyn std::error::Error>> {
7319 use crate::model::GpuTensor;
7320 if m == 1 {
7324 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(Some((y, 1.0))); }
7325 }
7326 if m != 1 || !self.uses_q8_1_fast(w) { return Ok(None); }
7328 let in_f = w.in_features();
7329 let out_f = w.out_features();
7330 let (bytes, qtype, row_bytes, scale, rp) = match w {
7331 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
7332 _ => return Ok(None),
7333 };
7334 if self.mmvq_supports(qtype) {
7336 let (mbytes, mrp) = match w {
7338 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
7339 _ => (bytes, rp),
7340 };
7341 let y = self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, 1.0, mrp)?;
7342 return Ok(Some((y, scale)));
7343 }
7344 let name = match qtype {
7346 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
7347 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
7348 QT_Q3_K => "qmatvec_q3_K_dp4a",
7349 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
7350 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
7351 _ => return Ok(None),
7352 };
7353 let f = self.func(name);
7354 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7355 let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
7356 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7357 let __s_b = self.gpu.stream();
7358 let mut b = __s_b.launch_builder(&f);
7359 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7360 unsafe { b.launch(cfg)?; }
7361 Ok(Some((y, scale)))
7362 }
7363
7364 pub fn mmvq_supports(&self, qtype: i32) -> bool {
7367 if qtype == QT_F8_E4M3 { return true; }
7372 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return false; }
7373 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_Q4_0)
7374 }
7375
7376 pub fn qmatvec_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7381 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7382 rp: bool)
7383 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7384 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)?;
7386 Ok(y)
7387 }
7388
7389 #[allow(clippy::too_many_arguments)]
7391 pub fn qmatvec_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7392 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7393 rp: bool, y: &mut CudaSlice<f32>)
7394 -> Result<(), Box<dyn std::error::Error>> {
7395 debug_assert!(y.len() >= m * out_f);
7396 const ROWS_PER_BLOCK: u32 = 4; if qtype == QT_Q8_0 && rp && m == 1 && out_f >= 64
7402 && (out_f as u32).div_ceil(ROWS_PER_BLOCK) < 4 * self.sm_count() as u32
7403 && {
7404 static G2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7405 *G2.get_or_init(|| std::env::var("MEMRA_Q80_G2").as_deref() != Ok("0"))
7406 }
7407 {
7408 let f = self.func("qmatvec_q8_0_mmvq_rp_g2");
7409 let cfg = LaunchConfig {
7410 grid_dim: ((out_f as u32).div_ceil(2), 1, 1),
7411 block_dim: (32, 2, 1),
7412 shared_mem_bytes: 0,
7413 };
7414 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, row_bytes as i64);
7415 let __s_b = self.gpu.stream();
7416 let mut b = __s_b.launch_builder(&f);
7417 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7418 unsafe { b.launch(cfg)?; }
7419 if scale != 1.0 { self.scale_inplace(y, scale, out_f)?; }
7420 return Ok(());
7421 }
7422 let mut mr: u32 = if m == 1 && (qtype == QT_NVFP4 || qtype == QT_Q5_K) { 2 } else { 1 };
7431 if m == 1 && qtype == QT_Q4_0 {
7436 static Q40MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7437 mr = *Q40MR.get_or_init(|| std::env::var("MEMRA_Q40_MR").ok()
7440 .and_then(|v| v.parse().ok()).unwrap_or(1));
7441 }
7442 let q5_mode = std::env::var("MEMRA_Q5K_ISSUE").ok();
7453 let q5_force = q5_mode.as_deref() == Some("2");
7454 let q5_il = qtype == QT_Q5_K && m == 1
7457 && (q5_force || q5_mode.as_deref().map(|v| v != "0").unwrap_or(true));
7458 if q5_il && !q5_force && out_f > 65536 { mr = 1; }
7459 if qtype == QT_Q4_0 && rp && mr != 1 { mr = 2; }
7462 if qtype == QT_Q8_0 && rp {
7466 static Q80MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7467 mr = *Q80MR.get_or_init(|| std::env::var("MEMRA_Q80_MR").ok()
7468 .and_then(|v| v.parse().ok()).unwrap_or(1));
7469 }
7470 let name = match (qtype, mr, rp) {
7471 (QT_NVFP4, 2, false) => "qmatvec_nvfp4_mmvq_mr2",
7472 (QT_NVFP4, 2, true) => "qmatvec_nvfp4_mmvq_mr2_rp",
7473 (QT_NVFP4, _, true) => "qmatvec_nvfp4_mmvq_rp",
7474 (QT_Q4_0, 1, true) => "qmatvec_q4_0_mmvq_rp",
7475 (QT_Q4_0, _, true) => "qmatvec_q4_0_mmvq_mr2_rp",
7476 (QT_Q5_K, 2, _) => if q5_il { "qmatvec_q5_K_mmvq_mr2_il" } else { "qmatvec_q5_K_mmvq_mr2" },
7477 (QT_Q8_0, 2, true) => "qmatvec_q8_0_mmvq_mr2_rp",
7478 (QT_Q8_0, _, true) if in_f % 1024 == 0 && {
7483 static CA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7484 *CA.get_or_init(|| std::env::var("MEMRA_Q80_CA").as_deref() == Ok("1"))
7485 } => "qmatvec_q8_0_mmvq_rpca",
7486 (QT_Q8_0, _, true) => "qmatvec_q8_0_mmvq_rp",
7487 (QT_Q8_0, _, _) => "qmatvec_q8_0_mmvq",
7488 (QT_Q4_K, _, true) => "qmatvec_q4_K_mmvq_rp",
7492 (QT_Q6_K, _, true) => "qmatvec_q6_K_mmvq_rp",
7493 (QT_Q4_K, _, _) => "qmatvec_q4_K_mmvq",
7494 (QT_Q4_0, 2, false) => "qmatvec_q4_0_mmvq_mr2",
7495 (QT_Q4_0, _, false) => "qmatvec_q4_0_mmvq",
7496 (QT_Q5_K, _, _) => if q5_il { "qmatvec_q5_K_mmvq_il" } else { "qmatvec_q5_K_mmvq" },
7497 (QT_Q6_K, _, _) => "qmatvec_q6_K_mmvq",
7498 (QT_NVFP4, _, false) => "qmatvec_nvfp4_mmvq",
7499 (QT_F8_E4M3, _, _) => "qmatvec_e4m3_mmvq",
7500 _ => panic!("qmatvec_mmvq: qtype {qtype} has no MMVQ kernel"),
7501 };
7502 let f = self.func(name);
7503 let rows_per_block = ROWS_PER_BLOCK * mr;
7505 let cfg = LaunchConfig {
7506 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, m as u32, 1),
7507 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7510 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7511 let __s_b = self.gpu.stream();
7512 let mut b = __s_b.launch_builder(&f);
7513 if qtype == QT_NVFP4 || qtype == QT_F8_E4M3 {
7518 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&scale);
7519 unsafe { b.launch(cfg)?; }
7520 } else if Self::pdl_on() && Self::pdl_mmvq_on()
7521 && matches!(name, "qmatvec_q4_0_mmvq_rp" | "qmatvec_q6_K_mmvq"
7522 | "qmatvec_q6_K_mmvq_rp") {
7523 {
7527 use cudarc::driver::{DevicePtr, DevicePtrMut};
7528 let s = &self.gpu.stream();
7529 let (pw, _g0) = bytes.device_ptr(s); let (paq, _g1) = aq.device_ptr(s);
7530 let (pad, _g2) = ad.device_ptr(s); let (py, _g3) = y.device_ptr_mut(s);
7531 let mut ps = [
7532 &pw as *const _ as *mut std::ffi::c_void, &paq as *const _ as *mut _,
7533 &pad as *const _ as *mut _, &py as *const _ as *mut _,
7534 &inf as *const _ as *mut _, &outf as *const _ as *mut _,
7535 &mi as *const _ as *mut _, &rb as *const _ as *mut _,
7536 ];
7537 unsafe { self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?; }
7538 }
7539 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7540 } else {
7541 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7542 unsafe { b.launch(cfg)?; }
7543 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7544 }
7545 Ok(())
7546 }
7547
7548 pub fn qmatvec_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
7552 out_f: usize, qtype: i32, row_bytes: usize, rp: bool)
7553 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7554 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7555 self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, 1.0, rp)
7556 }
7557
7558 pub fn batched_supports(&self, qtype: i32) -> bool {
7562 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_F8_E4M3 | QT_Q4_0)
7563 }
7564
7565 pub fn iq_fast_enabled() -> bool {
7573 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7574 *ON.get_or_init(|| std::env::var("MEMRA_IQ_FAST").map(|v| v != "0").unwrap_or(true))
7575 }
7576
7577 pub fn b8_enabled() -> bool {
7580 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7581 *ON.get_or_init(|| std::env::var("MEMRA_B8").map(|v| v != "0").unwrap_or(true))
7582 }
7583
7584 pub fn batched_mcols(m: usize) -> usize {
7586 if m == 2 { 2 } else if m <= 4 { 4 } else if m <= 8 { 8 } else { 16 }
7587 }
7588
7589 fn batched_kernel_name(qtype: i32, mcols: usize) -> Option<&'static str> {
7594 Some(match (qtype, mcols) {
7595 (QT_Q8_0, 2) => "qmatvec_q8_0_mmvq_b2", (QT_Q8_0, 4) => "qmatvec_q8_0_mmvq_b4",
7596 (QT_Q8_0, 8) => "qmatvec_q8_0_mmvq_b8",
7597 (QT_Q8_0, 16) => "qmatvec_q8_0_mmvq_b16",
7603 (QT_Q4_K, 2) => "qmatvec_q4_K_mmvq_b2", (QT_Q4_K, 4) => "qmatvec_q4_K_mmvq_b4",
7604 (QT_Q4_K, 8) => "qmatvec_q4_K_mmvq_b8",
7605 (QT_Q4_K, 16) => "qmatvec_q4_K_mmvq_b16",
7608 (QT_Q5_K, 2) => "qmatvec_q5_K_mmvq_b2", (QT_Q5_K, 4) => "qmatvec_q5_K_mmvq_b4",
7609 (QT_Q5_K, 8) => "qmatvec_q5_K_mmvq_b8",
7610 (QT_Q5_K, 16) => "qmatvec_q5_K_mmvq_b16",
7613 (QT_Q6_K, 2) => "qmatvec_q6_K_mmvq_b2", (QT_Q6_K, 4) => "qmatvec_q6_K_mmvq_b4",
7614 (QT_Q6_K, 8) => "qmatvec_q6_K_mmvq_b8", (QT_Q6_K, 16) => "qmatvec_q6_K_mmvq_b16",
7615 (QT_NVFP4, 2) => "qmatvec_nvfp4_mmvq_b2", (QT_NVFP4, 4) => "qmatvec_nvfp4_mmvq_b4",
7616 (QT_NVFP4, 8) => "qmatvec_nvfp4_mmvq_b8",
7617 (QT_NVFP4, 16) => "qmatvec_nvfp4_mmvq_b16",
7621 (QT_F8_E4M3, 2) => "qmatvec_e4m3_mmvq_b2", (QT_F8_E4M3, 4) => "qmatvec_e4m3_mmvq_b4",
7622 (QT_F8_E4M3, 8) => "qmatvec_e4m3_mmvq_b8",
7623 (QT_F8_E4M3, 16) => "qmatvec_e4m3_mmvq_b16",
7627 (QT_Q4_0, 2) => "qmatvec_q4_0_mmvq_b2", (QT_Q4_0, 4) => "qmatvec_q4_0_mmvq_b4",
7628 (QT_Q4_0, 8) => "qmatvec_q4_0_mmvq_b8", (QT_Q4_0, 16) => "qmatvec_q4_0_mmvq_b16",
7629 _ => return None,
7630 })
7631 }
7632
7633 pub fn sm_count(&self) -> i32 {
7668 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7669 *SMS.get_or_init(|| {
7670 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7671 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7672 })
7673 }
7674
7675 pub fn batched_variant(&self, _m: usize, in_f: usize, out_f: usize, qtype: i32,
7676 row_bytes: usize, mcols: usize, rp: bool) -> &'static str {
7677 if qtype == QT_Q8_0 {
7682 return if rp { "rp" } else { "base" };
7683 }
7684 static BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7685 let bv = *BV.get_or_init(|| match std::env::var("MEMRA_MMVQ_BV").as_deref() {
7686 Ok("base") => "base", Ok("pf") => "pf", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7687 Ok("pfr2") => "pfr2", Ok("ca") => "ca", Ok("car2") => "car2",
7688 Ok("rp") => "rp", Ok("rpr2") => "rpr2", Ok("rpr2w8") => "rpr2w8",
7691 Ok("rpca") => "rpca", Ok("rpcar2") => "rpcar2",
7694 Ok("rpsc") => "rpsc", Ok("rpms") => "rpms", Ok("rpmsc") => "rpmsc",
7701 Ok("rpks") => "rpks", Ok("rpksc") => "rpksc",
7702 _ => "auto",
7703 });
7704 let ca_ok = qtype == QT_NVFP4 && (row_bytes % 16 == 0) && (in_f % 1024 == 0);
7708 static KS_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7713 let ks_on = *KS_ON.get_or_init(|| std::env::var("MEMRA_KS").as_deref() != Ok("0"));
7714 let sc_ok = ks_on && qtype == QT_NVFP4 && (in_f % 256 == 0) && (in_f / 64 <= 272);
7715 let ks_ok = ks_on && qtype == QT_NVFP4 && (in_f % 512 == 0) && (in_f / 64 <= 272);
7716 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7717 let sms = *SMS.get_or_init(|| {
7718 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7719 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7720 });
7721 let kq_r2 = matches!(qtype, QT_Q4_K | QT_Q5_K | QT_Q6_K);
7741 static KQBV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7744 let kq_bv = *KQBV.get_or_init(|| match std::env::var("MEMRA_KQ_BV").as_deref() {
7745 Ok("base") => "base", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7746 _ => "auto",
7747 });
7748 let variant: &'static str = if qtype == QT_Q4_0 {
7749 static Q40BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7753 let q40 = *Q40BV.get_or_init(|| match std::env::var("MEMRA_Q40_BV").as_deref() {
7754 Ok("base") => "base", Ok("r2") => "r2", Ok("ms") => "ms", Ok("sm") => "sm",
7760 Ok("la") => "la", _ => "auto",
7761 });
7762 let v = if q40 != "auto" { q40 }
7763 else if (out_f as u32).div_ceil(8) >= 4 * sms as u32 { "r2" } else { "base" };
7764 if rp { match v { "ms" => "r2ms_rp", "sm" => "r2sm_rp", "la" => "r2la_rp",
7769 "r2" => "r2_rp", _ => "rp" } }
7770 else if matches!(v, "ms" | "sm" | "la") { "r2" } else { v }
7771 } else if qtype != QT_NVFP4 && !kq_r2 {
7772 "base"
7773 } else if kq_r2 && rp {
7774 "rp"
7778 } else if kq_r2 {
7779 if kq_bv != "auto" {
7782 if kq_bv == "r2w8" && mcols != 4 { "r2" } else { kq_bv }
7783 } else if bv != "auto" {
7784 match bv {
7785 "r2" | "pfr2" | "rpr2" | "car2" => "r2",
7786 "r2w8" | "rpr2w8" => if mcols != 4 { "r2" } else { "r2w8" },
7787 _ => "base", }
7789 } else {
7790 let blocks = (out_f + 7) / 8;
7791 let waves = blocks as f64 / (7 * sms as usize) as f64;
7792 let filled = blocks >= 4 * sms as usize;
7793 let use_r2 = if qtype == QT_Q4_K { filled } else { waves >= 2.0 };
7794 if use_r2 { "r2" } else { "base" }
7795 }
7796 } else if bv != "auto" {
7797 let v = if bv == "r2w8" && mcols == 2 { "r2" }
7802 else if bv == "ca" && (!ca_ok || mcols == 8) { "pf" }
7803 else if bv == "car2" && (!ca_ok || mcols == 8) { "r2" }
7804 else if bv == "pfr2" && mcols == 8 { "r2" }
7805 else if (bv == "rpr2w8" || bv == "rpr2") && mcols == 2 { "rpr2" }
7806 else if (bv == "rpca" || bv == "rpcar2") && (!ca_ok || mcols == 8) {
7808 if mcols == 8 { "rpr2w8" } else { "rpr2" }
7809 }
7810 else if bv == "rpcar2" && mcols == 2 { "rpca" }
7811 else if (bv == "rpsc" || bv == "rpmsc") && !sc_ok { "rpr2" }
7814 else if (bv == "rpks" || bv == "rpksc") && !ks_ok { "rpr2" }
7815 else { bv };
7816 if rp {
7817 match v {
7818 "base" | "pf" | "ca" | "rp" => "rp",
7819 "r2" | "pfr2" | "car2" | "rpr2" => "rpr2",
7820 "r2w8" | "rpr2w8" => if mcols == 2 { "rpr2" } else { "rpr2w8" },
7821 other => other, }
7823 } else { v }
7824 } else if mcols == 8 {
7825 if rp { if sc_ok { "rpsc" } else { "rpr2w8" } } else { "r2w8" }
7836 } else if mcols >= 4 {
7837 let blocks = (out_f + 7) / 8;
7841 let r7 = 7 * sms as usize;
7842 let r8 = 8 * sms as usize;
7843 let waves = blocks as f64 / r7 as f64;
7844 let filled = blocks >= 4 * sms as usize;
7845 if filled && blocks.div_ceil(r8) < blocks.div_ceil(r7) {
7849 if rp { "rpr2w8" } else { "r2w8" }
7853 } else if waves >= 2.0 || (waves <= 1.0 && filled) {
7854 if rp { "rpr2" } else { "r2" }
7857 } else {
7858 if rp { "rp" } else { "pf" }
7862 }
7863 } else if in_f >= 6144 {
7864 if rp { "rpr2" } else { "r2" }
7868 }
7869 else if rp {
7870 let waves = ((out_f + 7) / 8) as f64 / (7 * sms as usize) as f64;
7875 if sc_ok && waves >= 0.9 && waves <= 1.1 { "rpsc" } else { "rp" }
7876 } else { "base" };
7877 variant
7878 }
7879
7880 pub fn qmatvec_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7881 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize,
7882 mcols: usize, scale: f32, rp: bool)
7883 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7884 const ROWS_PER_BLOCK: u32 = 4;
7885 let forced: Option<&'static str> = {
7890 static V: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
7891 V.get_or_init(|| std::env::var("MEMRA_BVAR").ok())
7892 .as_deref()
7893 .map(|s| Box::leak(s.to_string().into_boxed_str()) as &'static str)
7894 };
7895 let variant = match forced {
7896 Some(v) if !rp || v.contains("rp") => v,
7897 _ => self.batched_variant(m, in_f, out_f, qtype, row_bytes, mcols, rp),
7898 };
7899 let base_name = Self::batched_kernel_name(qtype, mcols)
7900 .ok_or_else(|| format!("qmatvec_mmvq_batched: no kernel for qtype {qtype} mcols {mcols}"))?;
7901 let variant = if mcols == 16 { if rp { "rp" } else { "base" } } else { variant };
7905 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7912 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
7913 if b567 && qtype == QT_NVFP4 && rp && mcols == 8 && (5..=7).contains(&m)
7914 && matches!(variant, "rpsc" | "rpr2w8") {
7915 let f = self.func(&format!("qmatvec_nvfp4_mmvq_b{m}_{variant}"));
7916 let rows_per_block = ROWS_PER_BLOCK * 2; let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7918 let cfg = LaunchConfig {
7919 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
7920 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0 };
7921 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7922 let __s_b = self.gpu.stream();
7923 let mut b = __s_b.launch_builder(&f);
7924 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7925 unsafe { b.launch(cfg)?; }
7926 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
7927 return Ok(y);
7928 }
7929 let (name, rows_per_block): (std::borrow::Cow<'static, str>, u32) = match variant {
7930 "base" => (base_name.into(), ROWS_PER_BLOCK),
7931 "pf" => (format!("{base_name}_pf").into(), ROWS_PER_BLOCK),
7932 "ca" => (format!("{base_name}_ca").into(), ROWS_PER_BLOCK),
7933 "rp" => (format!("{base_name}_rp").into(), ROWS_PER_BLOCK),
7934 "rpca" => (format!("{base_name}_rpca").into(), ROWS_PER_BLOCK), "rpks" => (format!("{base_name}_rpks").into(), ROWS_PER_BLOCK),
7938 "rpksc" => (format!("{base_name}_rpksc").into(), ROWS_PER_BLOCK),
7939 "rpms" => (format!("{base_name}_rpms").into(), ROWS_PER_BLOCK),
7940 "rpmsc" => (format!("{base_name}_rpmsc").into(), ROWS_PER_BLOCK),
7941 "r2ms_rp" => (format!("{base_name}_r2ms_rp").into(), ROWS_PER_BLOCK),
7942 "r2sm_rp" => (format!("{base_name}_r2sm_rp").into(), ROWS_PER_BLOCK * 2),
7943 "r2la_rp" => (format!("{base_name}_r2la_rp").into(), ROWS_PER_BLOCK * 2),
7944 v => (format!("{base_name}_{v}").into(), ROWS_PER_BLOCK * 2), };
7946 debug_assert!(!rp || name.contains("_rp"), "rp weight dispatched to a GGUF-layout kernel");
7947 let f = self.func(&name);
7948 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7949 let smem = if name.contains("_r2sm_rp") { (mcols * 32 * 9 * 4 + mcols * 32 * 4) as u32 }
7951 else { 0 };
7952 let cfg = LaunchConfig {
7953 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
7954 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: smem };
7955 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7956 let __s_b = self.gpu.stream();
7957 let mut b = __s_b.launch_builder(&f);
7958 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7959 unsafe { b.launch(cfg)?; }
7960 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
7961 Ok(y)
7962 }
7963
7964 pub fn qmatvec_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
7968 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, mcols: usize,
7969 rp: bool)
7970 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7971 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7972 self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, 1.0, rp)
7973 }
7974
7975 pub fn qmatvec_nvfp4_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
7977 in_f: usize, out_f: usize, row_bytes: usize, mcols: usize,
7978 rp: bool)
7979 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7980 self.qmatvec_batched_raw(bytes, x, m, in_f, out_f, QT_NVFP4, row_bytes, mcols, rp)
7981 }
7982
7983 fn try_fp4_gemm(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize,
7987 in_f: usize, out_f: usize)
7988 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7989 use crate::model::GpuTensor;
7990 if cfg!(memra_portable_cuda) { return Ok(None); }
7991 if std::env::var("MEMRA_FP4").is_err() { return Ok(None); }
7992 #[cfg(memra_cutlass)]
8001 if m >= 128 && std::env::var("MEMRA_FP4_CUTLASS").is_ok() {
8002 if let GpuTensor::Quant { bytes, qtype, scale, row_bytes, cutlass, .. } = w {
8003 if *qtype == QT_NVFP4 && in_f % 64 == 0 {
8004 if let Some(cw) = cutlass {
8005 let y = self.cutlass_fp4_gemm(&cw.b_packed, &cw.sfb_swizzled, x, *scale,
8007 m, out_f, in_f)?;
8008 return Ok(Some(y));
8009 } else if std::env::var("MEMRA_FP4_CUTLASS_OTF").is_ok() {
8010 let (b_packed, sfb_sw) = self.build_cutlass_weight(bytes, out_f, in_f, *row_bytes)?;
8015 let y = self.cutlass_fp4_gemm(&b_packed, &sfb_sw, x, *scale, m, out_f, in_f)?;
8016 return Ok(Some(y));
8017 }
8018 }
8019 }
8020 }
8021 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } = w {
8022 if *qtype == QT_NVFP4 && in_f % 64 == 0 && !*rp {
8025 let y = self.qmatvec_gemm_nvfp4_fp4(bytes, x, m, in_f, out_f, *row_bytes, *scale)?;
8026 return Ok(Some(y));
8027 }
8028 }
8029 Ok(None)
8030 }
8031
8032 pub fn rms_norm_f16out(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>,
8036 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
8037 ncols: usize, nrows: usize, eps: f32)
8038 -> Result<(), Box<dyn std::error::Error>> {
8039 let f = self.func("rms_norm_f16out_f32");
8040 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
8041 let (nc, e) = (ncols as i32, eps);
8042 let __s_b = self.gpu.stream();
8043 let mut b = __s_b.launch_builder(&f);
8044 b.arg(x).arg(w).arg(dst).arg(dst16).arg(&nc).arg(&e);
8045 unsafe { b.launch(cfg)?; }
8046 Ok(())
8047 }
8048
8049 #[allow(clippy::too_many_arguments)]
8052 pub fn add_rms_norm_f16out(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
8053 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
8054 dst16: &mut CudaSlice<u8>, ncols: usize, nrows: usize, eps: f32)
8055 -> Result<(), Box<dyn std::error::Error>> {
8056 let f = self.func("add_rms_norm_f16out_f32");
8057 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
8058 let (nc, e) = (ncols as i32, eps);
8059 let __s_lb = self.gpu.stream();
8060 let mut lb = __s_lb.launch_builder(&f);
8061 lb.arg(a).arg(b).arg(w).arg(res).arg(dst).arg(dst16).arg(&nc).arg(&e);
8062 unsafe { lb.launch(cfg)?; }
8063 Ok(())
8064 }
8065
8066 pub fn matmul_group_xh(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>,
8069 xh: &CudaSlice<u8>, m: usize)
8070 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8071 let mut out = Vec::with_capacity(ws.len());
8072 let in_f = ws[0].in_features();
8073 for w in ws {
8074 if w.in_features() == in_f && m >= 16 && !self.verify_exact_on() {
8075 if let Some(y) = self.try_f16_gemm_pre(w, xh, m)? {
8076 out.push(y);
8077 continue;
8078 }
8079 }
8080 out.push(self.matmul(w, x, m)?);
8081 }
8082 Ok(out)
8083 }
8084
8085 pub fn gdn_pad_mask(&self, beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
8088 len_d: &CudaSlice<i32>, h: usize, t: usize)
8089 -> Result<(), Box<dyn std::error::Error>> {
8090 let f = self.func("gdn_pad_mask_f32");
8091 let cfg = LaunchConfig::for_num_elems((t * h) as u32);
8092 let (hi, ti) = (h as i32, t as i32);
8093 let __s_b = self.gpu.stream();
8094 let mut b = __s_b.launch_builder(&f);
8095 b.arg(beta).arg(g_log).arg(len_d).arg(&hi).arg(&ti);
8096 unsafe { b.launch(cfg)?; }
8097 Ok(())
8098 }
8099
8100 pub fn row_gather_dev(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
8103 len_d: &CudaSlice<i32>, ncols: usize)
8104 -> Result<(), Box<dyn std::error::Error>> {
8105 let f = self.func("row_gather_dev_f32");
8106 let cfg = LaunchConfig::for_num_elems(ncols as u32);
8107 let nc = ncols as i32;
8108 let __s_b = self.gpu.stream();
8109 let mut b = __s_b.launch_builder(&f);
8110 b.arg(src).arg(dst).arg(len_d).arg(&nc);
8111 unsafe { b.launch(cfg)?; }
8112 Ok(())
8113 }
8114
8115 pub fn matmul_group(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>, m: usize)
8122 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8123 use crate::model::GpuTensor;
8124 let mut out = Vec::with_capacity(ws.len());
8125 let any_mirror = ws.iter().any(|w| matches!(w, GpuTensor::Quant { f16: Some(_), .. }));
8126 if m >= 16 && any_mirror && !self.verify_exact_on() {
8127 let in_f = ws[0].in_features();
8128 let xh = self.f16_act(x, m * in_f, in_f)?;
8129 for w in ws {
8130 if w.in_features() == in_f {
8131 if let Some(y) = self.try_f16_gemm_pre(w, &xh, m)? {
8132 out.push(y);
8133 continue;
8134 }
8135 }
8136 out.push(self.matmul(w, x, m)?);
8137 }
8138 return Ok(out);
8139 }
8140 for w in ws {
8141 out.push(self.matmul(w, x, m)?);
8142 }
8143 Ok(out)
8144 }
8145
8146 pub fn matmul_group_multi(&self, ws: &[&crate::model::GpuTensor],
8153 xs: &[&CudaSlice<f32>], ms: &[usize])
8154 -> Result<Vec<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
8155 assert_eq!(xs.len(), ms.len());
8156 let in_f = ws[0].in_features();
8157 let total: usize = ms.iter().sum();
8158 let mut xcat = self.uninit(total * in_f)?;
8159 let mut off = 0usize;
8160 for (x, &m) in xs.iter().zip(ms) {
8161 self.copy_into(&mut xcat, off * in_f, x, m * in_f)?;
8162 off += m;
8163 }
8164 let ys = self.matmul_group(ws, &xcat, total)?;
8165 let mut out: Vec<Vec<CudaSlice<f32>>> = (0..xs.len()).map(|_| Vec::new()).collect();
8166 for (w, y) in ws.iter().zip(ys) {
8167 let out_f = w.out_features();
8168 let mut off = 0usize;
8169 for (s, &m) in ms.iter().enumerate() {
8170 let mut ys_s = self.uninit(m * out_f)?;
8171 let src = y.slice(off * out_f..(off + m) * out_f);
8172 self.gpu.stream().memcpy_dtod(&src, &mut ys_s)?;
8173 out[s].push(ys_s);
8174 off += m;
8175 }
8176 }
8177 Ok(out)
8178 }
8179
8180 pub fn gemm_supports(&self, w: &crate::model::GpuTensor) -> bool {
8190 use crate::model::GpuTensor;
8191 if !legacy_quant_gemm_allowed(
8192 cfg!(memra_portable_cuda),
8193 cfg!(memra_hopper_mma),
8194 std::env::var_os("MEMRA_NO_GEMM").is_some(),
8195 ) {
8196 return false;
8197 }
8198 match w {
8199 GpuTensor::Quant { qtype, .. } =>
8200 matches!(*qtype, QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q4_0)
8201 || (*qtype == QT_NVFP4 && w.in_features() % 64 == 0),
8202 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
8203 }
8204 }
8205
8206 pub fn qmatvec_gemm(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
8213 m: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8214 use crate::model::GpuTensor;
8215 let in_f = w.in_features();
8216 let out_f = w.out_features();
8217 let (bytes, qtype, row_bytes, scale, rp) = match w {
8218 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
8219 _ => unreachable!("gemm_supports guaranteed Quant"),
8220 };
8221 if cfg!(memra_hopper_mma) && qtype == QT_Q8_0 && out_f % 64 == 0 && wgmma_gemm_enabled() {
8227 if let GpuTensor::Quant { rp4: Some(m4), .. } = w {
8228 let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, aq, ad, m, in_f, out_f)?;
8229 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8230 return Ok(y);
8231 }
8232 }
8233 let name = match qtype {
8234 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8235 QT_Q4_0 => if rp { "qmatvec_gemm_q4_0_rp" } else { "qmatvec_gemm_q4_0" },
8236 QT_Q5_K => "qmatvec_gemm_q5_K",
8237 QT_Q6_K => "qmatvec_gemm_q6_K",
8238 QT_NVFP4 => if rp { "qmatvec_gemm_nvfp4_rp" } else { "qmatvec_gemm_nvfp4" },
8239 _ => unreachable!(),
8240 };
8241 let f = self.func(name);
8242 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);
8247 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8249 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8250 let warps: u32 = if is_k1 { k1_tile.2 } else {
8251 match qtype { QT_NVFP4 => 8, _ => 4 }
8252 };
8253 let cfg = LaunchConfig {
8254 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8255 block_dim: (32, warps, 1),
8256 shared_mem_bytes: 0,
8257 };
8258 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8259 let __s_b = self.gpu.stream();
8260 let mut b = __s_b.launch_builder(&f);
8261 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8262 unsafe { b.launch(cfg)?; }
8263 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8264 Ok(y)
8265 }
8266
8267 pub fn qmatvec_gemm_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
8272 out_f: usize, qtype: i32, row_bytes: usize)
8273 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8274 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
8275 let name = match qtype {
8276 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8277 QT_Q4_0 => "qmatvec_gemm_q4_0",
8278 QT_Q5_K => "qmatvec_gemm_q5_K",
8279 QT_Q6_K => "qmatvec_gemm_q6_K", QT_NVFP4 => "qmatvec_gemm_nvfp4",
8280 QT_NVFP4_RP => "qmatvec_gemm_nvfp4_rp",
8281 _ => panic!("qmatvec_gemm_raw: qtype {qtype} has no GEMM kernel"),
8282 };
8283 let f = self.func(name);
8284 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);
8288 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8290 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8291 let warps: u32 = if is_k1 { k1_tile.2 } else {
8292 match qtype { QT_NVFP4 | QT_NVFP4_RP => 8, _ => 4 }
8293 };
8294 let cfg = LaunchConfig {
8295 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8296 block_dim: (32, warps, 1), shared_mem_bytes: 0,
8297 };
8298 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8299 let __s_b = self.gpu.stream();
8300 let mut b = __s_b.launch_builder(&f);
8301 b.arg(bytes).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8302 unsafe { b.launch(cfg)?; }
8303 Ok(y)
8304 }
8305
8306 pub fn qmatvec_gemm_q8_0_wgmma_raw(&self, rp4: &CudaSlice<u8>, aq: &CudaSlice<i8>,
8313 ad: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize)
8314 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8315 assert!(out_f % 64 == 0 && in_f % 32 == 0, "wgmma GEMM needs out_f%64==0, in_f%32==0");
8316 let f = self.func("qmatvec_gemm_q8_0_wgmma");
8317 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
8319 grid_dim: ((out_f / 64) as u32, (m as u32).div_ceil(64), 1),
8320 block_dim: (128, 1, 1), shared_mem_bytes: 0,
8321 };
8322 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
8323 let __s_b = self.gpu.stream();
8324 let mut b = __s_b.launch_builder(&f);
8325 b.arg(rp4).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi);
8326 unsafe { b.launch(cfg)?; }
8327 Ok(y)
8328 }
8329
8330 pub fn scale_inplace(&self, y: &mut CudaSlice<f32>, s: f32, n: usize)
8332 -> Result<(), Box<dyn std::error::Error>> {
8333 let f = self.func("scale_f32");
8334 let cfg = LaunchConfig::for_num_elems(n as u32);
8335 let (sf, ni) = (s, n as i32);
8336 let __s_b = self.gpu.stream();
8337 let mut b = __s_b.launch_builder(&f);
8338 b.arg(y).arg(&sf).arg(&ni);
8339 unsafe { b.launch(cfg)?; }
8340 Ok(())
8341 }
8342
8343 pub fn bf16_to_f32(&self, data: &cudarc::driver::CudaView<'_, u8>, n: usize)
8348 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8349 let mut out = self.alloc_uninit::<f32>(n)?;
8350 let f = self.func("bf16_to_f32");
8351 let cfg = LaunchConfig::for_num_elems(n as u32);
8352 let ni = n as i32;
8353 let __s_b = self.gpu.stream();
8354 let mut b = __s_b.launch_builder(&f);
8355 b.arg(data).arg(&mut out).arg(&ni);
8356 unsafe { b.launch(cfg)?; }
8357 Ok(out)
8358 }
8359
8360 fn linear_bf16_chunked(&self, x: &CudaSlice<f32>, data: &CudaSlice<u8>, m: usize,
8367 in_f: usize, out_f: usize, exact: bool)
8368 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8369 const CHUNK_BYTES: usize = 256 << 20;
8370 let chunk_rows = (CHUNK_BYTES / (in_f * 4)).max(1).min(out_f);
8371 if chunk_rows >= out_f {
8372 let wf32 = self.bf16_to_f32(&data.slice(0..in_f * out_f * 2), in_f * out_f)?;
8373 return if exact { self.linear_decode_exact(x, &wf32, m, in_f, out_f) }
8374 else { self.linear(x, &wf32, m, in_f, out_f) };
8375 }
8376 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8377 let mut r0 = 0usize;
8378 while r0 < out_f {
8379 let rows = chunk_rows.min(out_f - r0);
8380 let wslice = data.slice(r0 * in_f * 2..(r0 + rows) * in_f * 2);
8381 let wf32 = self.bf16_to_f32(&wslice, in_f * rows)?;
8382 let yc = if exact { self.linear_decode_exact(x, &wf32, m, in_f, rows)? }
8383 else { self.linear(x, &wf32, m, in_f, rows)? };
8384 for mi in 0..m {
8386 let src = yc.slice(mi * rows..(mi + 1) * rows);
8387 let mut dst = y.slice_mut(mi * out_f + r0..mi * out_f + r0 + rows);
8388 self.gpu.stream().memcpy_dtod(&src, &mut dst)?;
8389 }
8390 r0 += rows;
8391 }
8392 Ok(y)
8393 }
8394
8395 pub fn linear_decode_exact(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize,
8402 in_f: usize, out_f: usize)
8403 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8404 if m_tokens == 1 { return self.linear(x, w, 1, in_f, out_f); }
8405 let xv = self.view(x, m_tokens * in_f);
8406 let mut y = self.alloc_uninit::<f32>(m_tokens * out_f)?;
8407 for t in 0..m_tokens {
8408 let row = xv.slice(t * in_f..(t + 1) * in_f);
8409 let mut xr = self.alloc_uninit::<f32>(in_f)?;
8410 self.copy_view_into(&mut xr, 0, &row, in_f)?;
8411 let yr = self.linear(&xr, w, 1, in_f, out_f)?;
8412 self.copy_into(&mut y, t * out_f, &yr, out_f)?;
8413 }
8414 Ok(y)
8415 }
8416
8417 pub fn linear(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize, in_f: usize, out_f: usize)
8418 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8419 use cudarc::cublaslt::{Matmul, MatmulConfig};
8420 let mut c = self.alloc_uninit::<f32>(m_tokens * out_f)?; let cfg = MatmulConfig {
8422 transa: true, transb: false, transc: false,
8423 m: out_f as u64, n: m_tokens as u64, k: in_f as u64,
8424 alpha: 1.0, lda: in_f as i64, ldb: in_f as i64, beta: 0.0, ldc: out_f as i64,
8425 stride_a: None, stride_b: None, stride_c: None, stride_bias: None, batch_size: None,
8426 };
8427 unsafe { self.gpu.blas.matmul(cfg, w, x, &mut c, None, None)?; }
8428 Ok(c)
8429 }
8430
8431 pub fn sdpa_naive(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8433 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8434 t: usize, t_kv: usize, scale: f32, causal: bool)
8435 -> Result<(), Box<dyn std::error::Error>> {
8436 let f = self.func("sdpa_naive_f32");
8437 let cfg = LaunchConfig {
8438 grid_dim: (n_head as u32, t as u32, 1),
8439 block_dim: (128, 1, 1),
8440 shared_mem_bytes: (t_kv * 4) as u32,
8441 };
8442 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);
8443 let __s_b = self.gpu.stream();
8444 let mut b = __s_b.launch_builder(&f);
8445 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8446 unsafe { b.launch(cfg)?; }
8447 Ok(())
8448 }
8449
8450 #[allow(clippy::too_many_arguments)]
8452 pub fn sdpa_naive_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8453 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8454 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8455 -> Result<(), Box<dyn std::error::Error>> {
8456 let f = self.func("sdpa_naive_w_f32");
8457 let cfg = LaunchConfig {
8458 grid_dim: (n_head as u32, t as u32, 1),
8459 block_dim: (128, 1, 1),
8460 shared_mem_bytes: (t_kv * 4) as u32,
8461 };
8462 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
8463 t as i32, t_kv as i32, causal as i32, window as i32);
8464 let __s_b = self.gpu.stream();
8465 let mut b = __s_b.launch_builder(&f);
8466 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8467 .arg(&scale).arg(&cz).arg(&wi);
8468 unsafe { b.launch(cfg)?; }
8469 Ok(())
8470 }
8471
8472 pub fn sdpa_naive_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<f32>,
8474 v: &cudarc::driver::CudaView<f32>, o: &mut CudaSlice<f32>,
8475 head_dim: usize, n_head: usize, n_head_kv: usize, t: usize, t_kv: usize,
8476 scale: f32, causal: bool) -> Result<(), Box<dyn std::error::Error>> {
8477 let f = self.func("sdpa_naive_f32");
8478 let cfg = LaunchConfig {
8479 grid_dim: (n_head as u32, t as u32, 1), block_dim: (128, 1, 1),
8480 shared_mem_bytes: (t_kv * 4) as u32,
8481 };
8482 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);
8483 let __s_b = self.gpu.stream();
8484 let mut b = __s_b.launch_builder(&f);
8485 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8486 unsafe { b.launch(cfg)?; }
8487 Ok(())
8488 }
8489
8490 #[allow(clippy::too_many_arguments)]
8498 pub fn fa_dequant_kv_view_f32(&self, k: &cudarc::driver::CudaView<u8>,
8499 v: &cudarc::driver::CudaView<u8>,
8500 kf: &mut CudaSlice<f32>, vf: &mut CudaSlice<f32>,
8501 kv_dim_k: usize, kv_dim_v: usize, t_kv: usize,
8502 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
8503 -> Result<(), Box<dyn std::error::Error>> {
8504 let f = if g { self.func_g("fa_dequant_kv_ws_f32") } else { self.func("fa_dequant_kv_ws_f32") };
8505 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
8506 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8507 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1),
8508 shared_mem_bytes: 0 };
8509 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
8510 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
8511 let __s_b = self.gpu.stream();
8512 let mut b = __s_b.launch_builder(&f);
8513 b.arg(k).arg(v).arg(&mut *kf).arg(&mut *vf).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
8514 unsafe { b.launch(cfg)?; }
8515 Ok(())
8516 }
8517
8518 #[allow(clippy::too_many_arguments)]
8519 pub fn sdpa_naive_quantized_view(
8520 &self,
8521 q: &CudaSlice<f32>,
8522 k: &cudarc::driver::CudaView<u8>,
8523 v: &cudarc::driver::CudaView<u8>,
8524 o: &mut CudaSlice<f32>,
8525 head_dim: usize,
8526 n_head: usize,
8527 n_head_kv: usize,
8528 t: usize,
8529 t_kv: usize,
8530 scale: f32,
8531 causal: bool,
8532 k_tok_bytes: usize,
8533 v_tok_bytes: usize,
8534 ) -> Result<(), Box<dyn std::error::Error>> {
8535 let kv_dim = n_head_kv * head_dim;
8536 let mut kf = self.uninit(t_kv * kv_dim)?;
8537 let mut vf = self.uninit(t_kv * kv_dim)?;
8538 let f = self.func("fa_dequant_kv_ws_f32");
8539 let total = (2 * t_kv * kv_dim) as u64;
8540 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8541 let cfg = LaunchConfig {
8542 grid_dim: (nblk.max(1), 1, 1),
8543 block_dim: (256, 1, 1),
8544 shared_mem_bytes: 0,
8545 };
8546 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8547 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8548 let __s_b = self.gpu.stream();
8549 let mut b = __s_b.launch_builder(&f);
8550 b.arg(k)
8551 .arg(v)
8552 .arg(&mut kf)
8553 .arg(&mut vf)
8554 .arg(&kv_dim_i)
8555 .arg(&kv_dim_i)
8556 .arg(&t_kv_i)
8557 .arg(&k_tok_bytes_i)
8558 .arg(&v_tok_bytes_i);
8559 unsafe { b.launch(cfg)? };
8560 self.sdpa_naive(
8561 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8562 )
8563 }
8564
8565 #[allow(clippy::too_many_arguments)]
8577 pub fn sdpa_naive_w_quantized_view(
8578 &self,
8579 q: &CudaSlice<f32>,
8580 k: &cudarc::driver::CudaView<u8>,
8581 v: &cudarc::driver::CudaView<u8>,
8582 o: &mut CudaSlice<f32>,
8583 head_dim: usize,
8584 n_head: usize,
8585 n_head_kv: usize,
8586 t: usize,
8587 t_kv: usize,
8588 scale: f32,
8589 causal: bool,
8590 window: usize,
8591 k_tok_bytes: usize,
8592 v_tok_bytes: usize,
8593 ) -> Result<(), Box<dyn std::error::Error>> {
8594 let kv_dim = n_head_kv * head_dim;
8595 let mut kf = self.uninit(t_kv * kv_dim)?;
8596 let mut vf = self.uninit(t_kv * kv_dim)?;
8597 let f = self.func("fa_dequant_kv_ws_f32");
8598 let total = (2 * t_kv * kv_dim) as u64;
8599 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8600 let cfg = LaunchConfig {
8601 grid_dim: (nblk.max(1), 1, 1),
8602 block_dim: (256, 1, 1),
8603 shared_mem_bytes: 0,
8604 };
8605 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8606 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8607 let __s_b = self.gpu.stream();
8608 let mut b = __s_b.launch_builder(&f);
8609 b.arg(k)
8610 .arg(v)
8611 .arg(&mut kf)
8612 .arg(&mut vf)
8613 .arg(&kv_dim_i)
8614 .arg(&kv_dim_i)
8615 .arg(&t_kv_i)
8616 .arg(&k_tok_bytes_i)
8617 .arg(&v_tok_bytes_i);
8618 unsafe { b.launch(cfg)? };
8619 self.sdpa_naive_w(
8620 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
8621 )
8622 }
8623
8624 pub fn fa_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8628 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8629 t: usize, t_kv: usize, scale: f32, causal: bool)
8630 -> Result<(), Box<dyn std::error::Error>> {
8631 if portable_mma_gated() {
8632 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
8633 t, t_kv, scale, causal);
8634 }
8635 let fa3_on = head_dim == 256 && causal && t == t_kv
8643 && match std::env::var("MEMRA_FA3").as_deref() {
8644 Ok("0") => false,
8645 Ok("1") => true,
8646 _ => cfg!(memra_hopper_mma),
8647 };
8648 if fa3_on {
8649 let n = t * n_head * head_dim;
8650 let nkv = t * n_head_kv * head_dim;
8651 let mut q16 = self.alloc_u8_uninit(n * 2)?;
8652 let mut k16 = self.alloc_u8_uninit(nkv * 2)?;
8653 let mut v16 = self.alloc_u8_uninit(nkv * 2)?;
8654 self.f32_to_bf16_into(q, &mut q16, n)?;
8655 self.f32_to_bf16_into(k, &mut k16, nkv)?;
8656 self.f32_to_bf16_into(v, &mut v16, nkv)?;
8657 let rc = {
8658 use cudarc::driver::{DevicePtr, DevicePtrMut};
8659 let stream = self.gpu.stream();
8660 let (qp, _g1) = q16.device_ptr(&stream);
8661 let (kp, _g2) = k16.device_ptr(&stream);
8662 let (vp, _g3) = v16.device_ptr(&stream);
8663 let (op, _g4) = o.device_ptr_mut(&stream);
8664 unsafe {
8665 memra_fa3_prefill(qp as *const core::ffi::c_void,
8666 kp as *const core::ffi::c_void,
8667 vp as *const core::ffi::c_void,
8668 op as *mut f32,
8669 t as i32, n_head as i32, n_head_kv as i32,
8670 head_dim as i32, scale,
8671 stream.cu_stream() as *mut core::ffi::c_void)
8672 }
8673 };
8674 if rc != 0 {
8675 return Err(format!("memra_fa3_prefill rc={rc}").into());
8676 }
8677 return Ok(());
8678 }
8679 static FA_P1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8684 let fa_p1 = *FA_P1.get_or_init(|| std::env::var("MEMRA_FA_P1").as_deref() == Ok("1"));
8685 if fa_p1 && head_dim == 256 && !std::env::var("MEMRA_FA_FLOOR").is_ok() {
8686 const BLOCK_Q: usize = 64; const BKX: usize = 32;
8687 let f = self.func("fa_prefill_bf16_p1");
8688 let shmem = (2 * (2 * BKX * head_dim + BLOCK_Q * BKX)
8689 + 4 * (BLOCK_Q * BKX + 2 * BLOCK_Q)) as u32;
8690 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8691 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8692 let cfg = LaunchConfig {
8693 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8694 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8695 };
8696 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32,
8697 n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8698 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8699 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8700 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8701 let __s_b = self.gpu.stream();
8702 let mut b = __s_b.launch_builder(&f);
8703 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti)
8704 .arg(&tkvi).arg(&scale).arg(&cz);
8705 unsafe { b.launch(cfg)?; }
8706 return Ok(());
8707 }
8708 const BK: usize = 32;
8714 let w2 = std::env::var("MEMRA_FA_PP_W2").as_deref() == Ok("1");
8717 let (block_q, warps, w2_sfx): (usize, u32, &str) =
8718 if w2 { (32, 2, "_w2") } else { (64, 4, "") };
8719 let hd_sfx = fa_hd_suffix(head_dim)?;
8723 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8724 let bf16kv = !floor && !w2
8729 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0");
8730 let (kb16, vb16) = if bf16kv {
8731 let n = t_kv * n_head_kv * head_dim;
8732 let mut kb = self.alloc_u8_uninit(n * 2)?;
8733 let mut vb = self.alloc_u8_uninit(n * 2)?;
8734 let fcv = self.func("f32_to_bf16_bulk");
8735 let ni = n as i64;
8736 let cfgc = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
8737 let __s_b = self.gpu.stream();
8738 let mut b = __s_b.launch_builder(&fcv);
8739 b.arg(k).arg(&mut kb).arg(&ni);
8740 unsafe { b.launch(cfgc)?; }
8741 let __s_b = self.gpu.stream();
8742 let mut b = __s_b.launch_builder(&fcv);
8743 b.arg(v).arg(&mut vb).arg(&ni);
8744 unsafe { b.launch(cfgc)?; }
8745 (Some(kb), Some(vb))
8746 } else {
8747 (None, None)
8748 };
8749 let f = self.func(&if bf16kv {
8750 format!("fa_prefill_bf16kv_pp{hd_sfx}")
8751 } else {
8752 format!("fa_prefill_f32{}{}{hd_sfx}",
8753 if floor { "" } else { "_pp" },
8754 if floor { "" } else { w2_sfx })
8755 });
8756 let kv_stages = if bf16kv { 2 } else { 1 };
8759 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
8760 + 4 * (block_q * BK + 2 * block_q)) as u32;
8761 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8762 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8763 let cfg = LaunchConfig {
8764 grid_dim: ((t as u32 + block_q as u32 - 1) / block_q as u32, n_head as u32, 1),
8765 block_dim: (32, warps, 1), shared_mem_bytes: shmem,
8766 };
8767 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);
8768 let __s_b = self.gpu.stream();
8769 let mut b = __s_b.launch_builder(&f);
8770 b.arg(q);
8771 match (&kb16, &vb16) {
8772 (Some(kb), Some(vb)) => { b.arg(kb).arg(vb); }
8773 _ => { b.arg(k).arg(v); }
8774 }
8775 b.arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8776 unsafe { b.launch(cfg)?; }
8777 Ok(())
8778 }
8779
8780 #[allow(clippy::too_many_arguments)]
8784 pub fn fa_prefill_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8785 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8786 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8787 -> Result<(), Box<dyn std::error::Error>> {
8788 if portable_mma_gated() {
8791 return self.sdpa_naive_w(q, k, v, o, head_dim, n_head, n_head_kv,
8792 t, t_kv, scale, causal, window);
8793 }
8794 static FAW_F32: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8798 let faw_f32 = *FAW_F32.get_or_init(|| {
8799 std::env::var("MEMRA_FAW_STAGE").as_deref() == Ok("f32")
8800 });
8801 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8802 self.fa_prefill_w_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8803 window, floor || faw_f32, floor)
8804 }
8805
8806 #[allow(clippy::too_many_arguments)]
8809 pub fn fa_prefill_w_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
8810 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8811 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
8812 window: usize, v_f16: bool)
8813 -> Result<(), Box<dyn std::error::Error>> {
8814 const BLOCK_Q: usize = 64; const BK: usize = 32;
8815 debug_assert_eq!(head_dim, 256);
8816 let hp = fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
8817 && (n_head / n_head_kv) % 2 == 0;
8818 debug_assert!(!v_f16 || hp, "f16 V emitted but the SWA hp arm is off");
8819 if hp {
8820 const BLOCK_QH: usize = 32;
8821 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
8824 let vh: &CudaSlice<u8> = if v_f16 { vb } else {
8825 let n = t_kv * n_head_kv * head_dim;
8826 if vguard.as_ref().map(|b| b.len() < n * 2).unwrap_or(true) {
8827 *vguard = Some(self.alloc_uninit::<u8>(n * 2)?);
8828 }
8829 self.bf16_to_f16_into(vb, n, vguard.as_mut().unwrap())?;
8830 vguard.as_ref().unwrap()
8831 };
8832 let f = self.func("fa_prefill_w_bf16_p1h2");
8833 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
8834 + 4 * (2 * BLOCK_QH)) as u32;
8835 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8836 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8837 let cfg = LaunchConfig {
8838 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
8839 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8840 };
8841 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8842 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8843 let __s_b = self.gpu.stream();
8844 let mut b = __s_b.launch_builder(&f);
8845 b.arg(qb).arg(kb).arg(vh).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8846 .arg(&scale).arg(&cz).arg(&wi);
8847 unsafe { b.launch(cfg)?; }
8848 return Ok(());
8849 }
8850 let f = self.func("fa_prefill_w_bf16_p1");
8851 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
8852 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
8853 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8854 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8855 let cfg = LaunchConfig {
8856 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8857 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8858 };
8859 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8860 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8861 let __s_b = self.gpu.stream();
8862 let mut b = __s_b.launch_builder(&f);
8863 b.arg(qb).arg(kb).arg(vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8864 .arg(&scale).arg(&cz).arg(&wi);
8865 unsafe { b.launch(cfg)?; }
8866 Ok(())
8867 }
8868
8869 #[allow(clippy::too_many_arguments)]
8871 pub fn fa_prefill_w_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8872 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8873 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
8874 window: usize, f32_stage: bool, floor: bool)
8875 -> Result<(), Box<dyn std::error::Error>> {
8876 const BLOCK_Q: usize = 64; const BK: usize = 32;
8877 debug_assert_eq!(head_dim, 256, "fa_prefill_w is stamped hd256 only");
8878 static P1_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8882 let p1 = !floor && !f32_stage
8883 && *P1_ON.get_or_init(|| {
8884 std::env::var("MEMRA_FAW_P1").map(|v| v != "0").unwrap_or(true)
8885 });
8886 let hp = p1 && fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
8887 && (n_head / n_head_kv) % 2 == 0;
8888 if hp {
8889 const BLOCK_QH: usize = 32;
8890 let f = self.func("fa_prefill_w_bf16_p1h2");
8891 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
8892 + 4 * (2 * BLOCK_QH)) as u32;
8893 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8894 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8895 let cfg = LaunchConfig {
8896 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
8897 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8898 };
8899 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8900 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8901 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8902 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8903 let vh = self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?;
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 if p1 {
8912 let f = self.func("fa_prefill_w_bf16_p1");
8913 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
8914 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
8915 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8916 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8917 let cfg = LaunchConfig {
8918 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8919 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8920 };
8921 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8922 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8923 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8924 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8925 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8926 let __s_b = self.gpu.stream();
8927 let mut b = __s_b.launch_builder(&f);
8928 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8929 .arg(&scale).arg(&cz).arg(&wi);
8930 unsafe { b.launch(cfg)?; }
8931 return Ok(());
8932 }
8933 static G4_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8936 let g4 = !floor && !f32_stage && n_head_kv == 1 && n_head % 4 == 0
8937 && *G4_ON.get_or_init(|| {
8938 std::env::var("MEMRA_FAW_G4").map(|v| v != "0").unwrap_or(true)
8939 });
8940 if g4 {
8941 const SP_M: usize = 16;
8942 static O2_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8945 let o2 = *O2_ON.get_or_init(|| {
8946 std::env::var("MEMRA_FAW_O2").map(|v| v != "0").unwrap_or(true)
8947 });
8948 let f = self.func(if o2 { "fa_prefill_w_bf16_g4o2" } else { "fa_prefill_w_bf16_g4" });
8949 let shmem = if o2 {
8950 (2 * (4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M)) as u32
8951 } else {
8952 (2 * (2 * BK * head_dim + 4 * SP_M * head_dim + 4 * SP_M * BK)
8953 + 4 * (4 * SP_M)) as u32
8954 };
8955 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8956 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8957 let cfg = LaunchConfig {
8958 grid_dim: ((t as u32).div_ceil(SP_M as u32), (n_head / 4) as u32, 1),
8959 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8960 };
8961 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8962 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8963 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8964 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8965 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8966 let __s_b = self.gpu.stream();
8967 let mut b = __s_b.launch_builder(&f);
8968 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8969 .arg(&scale).arg(&cz).arg(&wi);
8970 unsafe { b.launch(cfg)?; }
8971 return Ok(());
8972 }
8973 let f = self.func(if floor { "fa_prefill_w_f32" }
8974 else if f32_stage { "fa_prefill_w_f32_pp" }
8975 else { "fa_prefill_w_bf16_pp" });
8976 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
8977 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
8978 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8979 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8980 let cfg = LaunchConfig {
8981 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8982 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8983 };
8984 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
8985 t as i32, t_kv as i32, causal as i32, window as i32);
8986 if f32_stage {
8987 let __s_b = self.gpu.stream();
8988 let mut b = __s_b.launch_builder(&f);
8989 b.arg(q).arg(k).arg(v).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 } else {
8993 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8994 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8995 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8996 let __s_b = self.gpu.stream();
8997 let mut b = __s_b.launch_builder(&f);
8998 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8999 .arg(&scale).arg(&cz).arg(&wi);
9000 unsafe { b.launch(cfg)?; }
9001 }
9002 Ok(())
9003 }
9004
9005 #[allow(clippy::too_many_arguments)]
9009 pub fn fa_prefill_hd512(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9010 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9011 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool)
9012 -> Result<(), Box<dyn std::error::Error>> {
9013 if portable_mma_gated() {
9015 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
9016 t, t_kv, scale, causal);
9017 }
9018 static F32_STAGE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9024 let f32_stage = *F32_STAGE.get_or_init(|| {
9025 std::env::var("MEMRA_FA512_STAGE").as_deref() == Ok("f32")
9026 });
9027 static SP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9031 let sp = !f32_stage
9032 && *SP_ON.get_or_init(|| {
9033 std::env::var("MEMRA_FA512_SP").map(|v| v != "0").unwrap_or(true)
9034 });
9035 self.fa_prefill_hd512_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale,
9036 causal, f32_stage, sp, sp && fa_f16pv_on())
9037 }
9038
9039 #[allow(clippy::too_many_arguments)]
9041 pub fn fa_prefill_hd512_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
9042 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9043 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9044 v_f16: bool)
9045 -> Result<(), Box<dyn std::error::Error>> {
9046 debug_assert_eq!(head_dim, 512);
9047 const SP_M: usize = 16; const BKS: usize = 32;
9048 let f16pv = fa_f16pv_on();
9052 let nw = if f16pv { fa512_wide_warps() } else { 2 };
9053 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
9054 debug_assert!(!v_f16 || f16pv, "f16 V emitted without the door on");
9055 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
9056 let vref: &CudaSlice<u8> = if f16pv && !v_f16 {
9057 let n = t_kv * n_head_kv * head_dim;
9059 let need = n * 2;
9060 if vguard.as_ref().map(|b| b.len() < need).unwrap_or(true) {
9061 *vguard = Some(self.alloc_uninit::<u8>(need)?);
9062 }
9063 let dst = vguard.as_mut().unwrap();
9064 self.bf16_to_f16_into(vb, n, dst)?;
9065 vguard.as_ref().unwrap()
9066 } else { vb };
9067 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
9068 else { match (f16pv, nw) {
9069 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9070 (true, _) => "fa_prefill_bf16_hd512_sp16",
9071 _ => "fa_prefill_bf16_hd512_sp",
9072 } });
9073 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9074 let shmem = if hp {
9076 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9077 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9078 } else {
9079 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9080 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9081 };
9082 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9083 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9084 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9085 let cfg = LaunchConfig {
9086 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9087 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9088 };
9089 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9090 t as i32, t_kv as i32, causal as i32);
9091 let __s_b = self.gpu.stream();
9092 let mut b = __s_b.launch_builder(&f);
9093 b.arg(qb).arg(kb).arg(vref).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9094 .arg(&scale).arg(&cz);
9095 unsafe { b.launch(cfg)?; }
9096 Ok(())
9097 }
9098
9099 #[allow(clippy::too_many_arguments)]
9102 pub fn fa_prefill_hd512_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
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 f32_stage: bool, sp: bool, f16pv: bool)
9106 -> Result<(), Box<dyn std::error::Error>> {
9107 debug_assert_eq!(head_dim, 512, "fa_prefill_hd512 is hd512 only");
9108 if sp && !f32_stage {
9109 const SP_M: usize = 16; const BKS: usize = 32;
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 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
9116 else { match (f16pv, nw) {
9117 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9118 (true, _) => "fa_prefill_bf16_hd512_sp16",
9119 _ => "fa_prefill_bf16_hd512_sp",
9120 } });
9121 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9122 let shmem = if hp {
9123 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9124 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9125 } else {
9126 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9127 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9128 };
9129 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9130 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9131 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9132 let cfg = LaunchConfig {
9133 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9134 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9135 };
9136 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9137 t as i32, t_kv as i32, causal as i32);
9138 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9139 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9140 let vb = if f16pv { self.f32_to_f16(v, t_kv * n_head_kv * head_dim)? }
9141 else { self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)? };
9142 let __s_b = self.gpu.stream();
9143 let mut b = __s_b.launch_builder(&f);
9144 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9145 .arg(&scale).arg(&cz);
9146 unsafe { b.launch(cfg)?; }
9147 return Ok(());
9148 }
9149 const BLOCK_Q: usize = 32; const BK: usize = 32; const HALF: usize = 256;
9150 let f = self.func(if f32_stage { "fa_prefill_f32_hd512" } else { "fa_prefill_bf16_hd512" });
9151 let shmem = (2 * (BLOCK_Q * head_dim + BK * head_dim + BK * HALF + BLOCK_Q * BK)
9153 + 4 * BLOCK_Q) as u32;
9154 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9155 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9156 let cfg = LaunchConfig {
9157 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 2),
9158 block_dim: (32, 2, 1), shared_mem_bytes: shmem,
9159 };
9160 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9161 t as i32, t_kv as i32, causal as i32);
9162 if f32_stage {
9163 let __s_b = self.gpu.stream();
9164 let mut b = __s_b.launch_builder(&f);
9165 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9166 .arg(&scale).arg(&cz);
9167 unsafe { b.launch(cfg)?; }
9168 } else {
9169 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9170 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9171 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9172 let __s_b = self.gpu.stream();
9173 let mut b = __s_b.launch_builder(&f);
9174 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9175 .arg(&scale).arg(&cz);
9176 unsafe { b.launch(cfg)?; }
9177 }
9178 Ok(())
9179 }
9180
9181 #[allow(clippy::too_many_arguments)]
9185 pub fn rope_neox2_bf16e(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
9186 qb: &mut CudaSlice<u8>, kb: &mut CudaSlice<u8>,
9187 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
9188 nh_q: usize, nh_k: usize, n_tokens: usize, base: f32,
9189 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
9190 -> Result<(), Box<dyn std::error::Error>> {
9191 let f = self.func("rope_neox2_bf16e_f32");
9192 let rows = ((nh_q + nh_k) * n_tokens) as u32;
9193 let cfg = LaunchConfig { grid_dim: (rows, 1, 1),
9194 block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
9195 let theta_scale = base.powf(-2.0 / n_dims as f32);
9196 let (hd, nd, nhq, nhk, nt) = (head_dim as i32, n_dims as i32, nh_q as i32,
9197 nh_k as i32, n_tokens as i32);
9198 let __s_b = self.gpu.stream();
9199 let mut b = __s_b.launch_builder(&f);
9200 match ff {
9201 Some(t) => { b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9202 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9203 .arg(&theta_scale).arg(&freq_scale).arg(t);
9204 unsafe { b.launch(cfg)?; } }
9205 None => { let null: u64 = 0;
9206 b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9207 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9208 .arg(&theta_scale).arg(&freq_scale).arg(&null);
9209 unsafe { b.launch(cfg)?; } }
9210 }
9211 Ok(())
9212 }
9213
9214 pub fn f32_to_bf16(&self, x: &CudaSlice<f32>, n: usize)
9217 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9218 assert!(n % 4 == 0, "f32_to_bf16 requires n % 4 == 0, got {n}");
9219 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9220 let f = self.func("f32_to_bf16_flat");
9221 let n_i = n as i64;
9222 let cfg = LaunchConfig {
9223 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9224 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9225 };
9226 let __s_b = self.gpu.stream();
9227 let mut b = __s_b.launch_builder(&f);
9228 b.arg(x).arg(&mut y).arg(&n_i);
9229 unsafe { b.launch(cfg)?; }
9230 Ok(y)
9231 }
9232
9233 pub fn f32_to_f16(&self, x: &CudaSlice<f32>, n: usize)
9234 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9235 assert!(n % 4 == 0, "f32_to_f16 requires n % 4 == 0, got {n}");
9236 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9237 let f = self.func("f32_to_f16_flat");
9238 let n_i = n as i64;
9239 let cfg = LaunchConfig {
9240 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9241 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9242 };
9243 let __s_b = self.gpu.stream();
9244 let mut b = __s_b.launch_builder(&f);
9245 b.arg(x).arg(&mut y).arg(&n_i);
9246 unsafe { b.launch(cfg)?; }
9247 Ok(y)
9248 }
9249
9250 pub fn bf16_to_f16(&self, xb: &CudaSlice<u8>, n: usize)
9252 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9253 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9254 self.bf16_to_f16_into(xb, n, &mut y)?;
9255 Ok(y)
9256 }
9257
9258 pub fn bf16_to_f16_into(&self, xb: &CudaSlice<u8>, n: usize, y: &mut CudaSlice<u8>)
9260 -> Result<(), Box<dyn std::error::Error>> {
9261 assert!(n % 2 == 0, "bf16_to_f16 requires n % 2 == 0, got {n}");
9262 assert!(y.len() >= n * 2);
9263 let f = self.func("bf16_to_f16_flat");
9264 let n2 = (n / 2) as i64;
9265 let cfg = LaunchConfig {
9266 grid_dim: (((n / 2) as u32).div_ceil(256), 1, 1),
9267 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9268 };
9269 let __s_b = self.gpu.stream();
9270 let mut b = __s_b.launch_builder(&f);
9271 b.arg(xb).arg(y).arg(&n2);
9272 unsafe { b.launch(cfg)?; }
9273 Ok(())
9274 }
9275
9276 #[allow(clippy::too_many_arguments)]
9281 pub fn fa_prefill_vl8(&self, seqs: &[FaSeqVl], head_dim: usize, n_head: usize,
9282 n_head_kv: usize, scale: f32)
9283 -> Result<(), Box<dyn std::error::Error>> {
9284 const BK: usize = 32;
9285 let b = seqs.len();
9286 assert!(b >= 1 && b <= 8);
9287 let mut packed = [FaSeqVl::default(); 8];
9288 packed[..b].copy_from_slice(seqs);
9289 let v = FaVl8(packed);
9290 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9291 let ept = (n_head_kv * head_dim) as i32;
9292 {
9293 let f = self.func("fa_mirror_vl");
9294 let max_n = (max_t as i64) * ept as i64;
9295 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
9296 for which in 0..2i32 {
9297 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9298 let __s_lb = self.gpu.stream();
9299 let mut lb = __s_lb.launch_builder(&f);
9300 lb.arg(&v).arg(&ept).arg(&which);
9301 unsafe { lb.launch(cfg)?; }
9302 }
9303 }
9304 let hd_sfx = fa_hd_suffix(head_dim)?;
9305 let f = self.func(&format!("fa_prefill_bf16kv_vl{hd_sfx}"));
9306 let block_q = 64usize;
9307 let kv_stages = 2usize;
9308 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
9309 + 4 * (block_q * BK + 2 * block_q)) as u32;
9310 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9311 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9312 let cfg = LaunchConfig {
9313 grid_dim: (max_t.div_ceil(block_q as u32), n_head as u32, b as u32),
9314 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9315 };
9316 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9317 let __s_lb = self.gpu.stream();
9318 let mut lb = __s_lb.launch_builder(&f);
9319 lb.arg(&v).arg(&hd).arg(&nh).arg(&nhkv).arg(&scale);
9320 unsafe { lb.launch(cfg)?; }
9321 Ok(())
9322 }
9323
9324 #[allow(clippy::too_many_arguments)]
9328 pub fn attn_pre_vl8(&self, seqs: &[AttnPreVl], wq: &CudaSlice<f32>, wk: &CudaSlice<f32>,
9329 head_dim: usize, rope_dims: usize, n_head: usize, n_head_kv: usize,
9330 eps: f32, freq_base: f32, freq_scale: f32,
9331 kv_dim_k: usize, kv_dim_v: usize,
9332 k_tok_bytes: usize, v_tok_bytes: usize)
9333 -> Result<(), Box<dyn std::error::Error>> {
9334 let b = seqs.len();
9335 assert!(b >= 1 && b <= 8);
9336 let mut packed = [AttnPreVl::default(); 8];
9337 packed[..b].copy_from_slice(seqs);
9338 let v = AttnPreVl8(packed);
9339 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9340 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9341 {
9342 let f = self.func("q_gate_split_vl");
9343 let n = max_t * (n_head * head_dim) as u32;
9344 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9345 let __s_lb = self.gpu.stream();
9346 let mut lb = __s_lb.launch_builder(&f);
9347 lb.arg(&v).arg(&hd).arg(&nh);
9348 unsafe { lb.launch(cfg)?; }
9349 }
9350 {
9351 let f = self.func("attn_rms_vl");
9352 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 };
9353 let __s_lb = self.gpu.stream();
9354 let mut lb = __s_lb.launch_builder(&f);
9355 lb.arg(&v).arg(wq).arg(wk).arg(&hd).arg(&nh).arg(&nhkv).arg(&eps);
9356 unsafe { lb.launch(cfg)?; }
9357 }
9358 {
9359 let f = self.func("attn_rope_vl");
9360 let theta_scale = freq_base.powf(-2.0 / rope_dims as f32);
9361 let nd = rope_dims as i32;
9362 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 };
9363 let __s_lb = self.gpu.stream();
9364 let mut lb = __s_lb.launch_builder(&f);
9365 lb.arg(&v).arg(&hd).arg(&nd).arg(&nh).arg(&nhkv).arg(&theta_scale).arg(&freq_scale);
9366 unsafe { lb.launch(cfg)?; }
9367 }
9368 {
9369 let f = self.func("append_kv_vl");
9370 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
9371 let cfg = LaunchConfig { grid_dim: (nblk, max_t, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
9372 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9373 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9374 let __s_lb = self.gpu.stream();
9375 let mut lb = __s_lb.launch_builder(&f);
9376 lb.arg(&v).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
9377 unsafe { lb.launch(cfg)?; }
9378 }
9379 Ok(())
9380 }
9381
9382 pub fn fa_prefill_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9387 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9388 head_dim: usize, n_head: usize, n_head_kv: usize,
9389 t: usize, t_kv: usize, scale: f32, causal: bool,
9390 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9391 -> Result<(), Box<dyn std::error::Error>> {
9392 if portable_mma_gated() {
9393 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9394 t, t_kv, scale, causal,
9395 k_tok_bytes, v_tok_bytes);
9396 }
9397 const BLOCK_Q: usize = 64; const BK: usize = 32;
9398 let name = format!("fa_prefill_q{}", fa_hd_suffix(head_dim)?);
9401 let f = if g { self.func_g(&name) } else { self.func(&name) };
9402 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9403 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9404 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9405 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9406 let cfg = LaunchConfig {
9407 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9408 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9409 };
9410 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);
9411 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9412 let __s_b = self.gpu.stream();
9413 let mut b = __s_b.launch_builder(&f);
9414 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9415 .arg(&ktb).arg(&vtb);
9416 unsafe { b.launch(cfg)?; }
9417 Ok(())
9418 }
9419
9420 #[allow(clippy::too_many_arguments)]
9430 pub fn fa_prefill_view_ws(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9431 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9432 head_dim: usize, n_head: usize, n_head_kv: usize,
9433 t: usize, t_kv: usize, scale: f32, causal: bool,
9434 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9435 -> Result<(), Box<dyn std::error::Error>> {
9436 if portable_mma_gated() {
9437 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9438 t, t_kv, scale, causal,
9439 k_tok_bytes, v_tok_bytes);
9440 }
9441 const BLOCK_Q: usize = 64; const BK: usize = 32;
9442 let kv_dim_k = n_head_kv * head_dim;
9443 let kv_dim_v = n_head_kv * head_dim;
9444 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9446 let mut guard = self.prime_deqw_ws.lock().unwrap();
9448 let need_grow = match guard.as_ref() {
9449 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9450 None => true,
9451 };
9452 if need_grow {
9453 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9454 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9455 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9456 }
9457 let (kw, vw) = guard.as_mut().unwrap();
9458 {
9460 let f = if g { self.func_g("fa_dequant_kv_ws_bf16") } else { self.func("fa_dequant_kv_ws_bf16") };
9462 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9463 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9464 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9465 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9466 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9467 let __s_b = self.gpu.stream();
9468 let mut b = __s_b.launch_builder(&f);
9469 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9470 unsafe { b.launch(cfg)?; }
9471 }
9472 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9480 {
9481 let hd_sfx = fa_hd_suffix(head_dim)?;
9482 let f = self.func(&format!("fa_prefill_qw{}{hd_sfx}", if db { "_db" } else { "" }));
9483 let shmem = if db {
9484 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9486 } else {
9487 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9488 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9489 };
9490 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9491 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9492 let cfg = LaunchConfig {
9493 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9494 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9495 };
9496 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);
9497 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9498 let __s_b = self.gpu.stream();
9499 let mut b = __s_b.launch_builder(&f);
9500 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9501 .arg(&kdk).arg(&kdv);
9502 unsafe { b.launch(cfg)?; }
9503 }
9504 Ok(())
9505 }
9506
9507 #[allow(clippy::too_many_arguments)]
9523 pub fn fa_prefill_view_ws_w_hd128(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9524 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9525 head_dim: usize, n_head: usize, n_head_kv: usize,
9526 t: usize, t_kv: usize, scale: f32, causal: bool,
9527 window: usize, k_tok_bytes: usize, v_tok_bytes: usize)
9528 -> Result<(), Box<dyn std::error::Error>> {
9529 assert_eq!(head_dim, 128, "fa_prefill_view_ws_w_hd128: only the hd128 twin is stamped");
9530 if portable_mma_gated() {
9531 return self.sdpa_naive_w_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9532 t, t_kv, scale, causal, window,
9533 k_tok_bytes, v_tok_bytes);
9534 }
9535 const BLOCK_Q: usize = 64; const BK: usize = 32;
9536 let kv_dim_k = n_head_kv * head_dim;
9537 let kv_dim_v = n_head_kv * head_dim;
9538 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9540 let mut guard = self.prime_deqw_ws.lock().unwrap();
9541 let need_grow = match guard.as_ref() {
9542 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9543 None => true,
9544 };
9545 if need_grow {
9546 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9547 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9548 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9549 }
9550 let (kw, vw) = guard.as_mut().unwrap();
9551 {
9554 let f = self.func("fa_dequant_kv_ws_bf16");
9555 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9556 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9557 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9558 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9559 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9560 let __s_b = self.gpu.stream();
9561 let mut b = __s_b.launch_builder(&f);
9562 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9563 unsafe { b.launch(cfg)?; }
9564 }
9565 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9567 {
9568 let f = self.func(if db { "fa_prefill_qw_db_w_hd128" } else { "fa_prefill_qw_w_hd128" });
9569 let shmem = if db {
9570 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9571 } else {
9572 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9573 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9574 };
9575 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9576 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9577 let cfg = LaunchConfig {
9578 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9579 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9580 };
9581 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
9582 let (kdk, kdv, wnd) = (kv_dim_k as i32, kv_dim_v as i32, window as i32);
9583 let __s_b = self.gpu.stream();
9584 let mut b = __s_b.launch_builder(&f);
9585 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9586 .arg(&kdk).arg(&kdv).arg(&wnd);
9587 unsafe { b.launch(cfg)?; }
9588 }
9589 Ok(())
9590 }
9591
9592 pub fn fa_decode(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9596 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9597 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9598 k_tok_bytes: usize, v_tok_bytes: usize)
9599 -> Result<(), Box<dyn std::error::Error>> {
9600 self.fa_decode_kvmod(q, k, v, o, head_dim, n_head, n_head_kv, t_kv, scale,
9601 k_tok_bytes, v_tok_bytes, false)
9602 }
9603
9604 #[allow(clippy::too_many_arguments)]
9608 #[allow(clippy::too_many_arguments)]
9612 #[allow(clippy::too_many_arguments)]
9613 fn fa_decode_scalar_unified(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9614 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9615 head_dim: usize, n_head: usize, n_head_kv: usize,
9616 t_kv_host: usize, t_kv_dev: Option<&CudaSlice<i32>>,
9617 scale: f32, n_splits: usize, split_keys: usize,
9618 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
9619 part_o: &mut CudaSlice<f32>, part_m: &mut CudaSlice<f32>,
9620 part_l: &mut CudaSlice<f32>,
9621 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
9622 -> Result<(), Box<dyn std::error::Error>> {
9623 let f = if g { self.func_g("fa_decode_f32") } else { self.fa_func("fa_decode_f32", head_dim) };
9624 let cfg = LaunchConfig { grid_dim: (n_head as u32, n_splits as u32, 1),
9625 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: (4 * (head_dim + 32)) as u32 };
9626 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
9627 let (ktb, vtb, tkvi, ski) = (k_tok_bytes as i64, v_tok_bytes as i64, t_kv_host as i32,
9628 split_keys as i32);
9629 let __s_b = self.gpu.stream();
9630 let mut b = __s_b.launch_builder(&f);
9631 match t_kv_dev {
9632 Some(d) => { b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9633 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(d).arg(&scale).arg(&nsp)
9634 .arg(&ski).arg(&ktb).arg(&vtb);
9635 unsafe { b.launch(cfg)?; } }
9636 None => { let null: u64 = 0;
9637 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9638 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&null).arg(&scale).arg(&nsp)
9639 .arg(&ski).arg(&ktb).arg(&vtb);
9640 unsafe { b.launch(cfg)?; } }
9641 }
9642 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1),
9643 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
9644 if let Some((oq, od)) = q8_out {
9645 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
9647 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
9648 let __s_b2 = self.gpu.stream();
9649 let mut b2 = __s_b2.launch_builder(&fc);
9650 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
9651 unsafe { b2.launch(cfg2)?; }
9652 return Ok(());
9653 }
9654 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
9655 let __s_b2 = self.gpu.stream();
9656 let mut b2 = __s_b2.launch_builder(&fc);
9657 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
9658 unsafe { b2.launch(cfg2)?; }
9659 Ok(())
9660 }
9661
9662 pub fn fa_decode_kvmod(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9663 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9664 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9665 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9666 -> Result<(), Box<dyn std::error::Error>> {
9667 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
9688 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
9692 let sp = fa_split_keys(t_kv, n_head_kv);
9693 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
9694 let o_len = n_head * n_splits * head_dim;
9695 let ml_len = n_head * n_splits;
9696 let mut part_guard = self.fa_part_pool.lock().unwrap();
9697 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
9698 let old = part_guard.take();
9709 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
9710 if let Some(old) = old {
9711 self.fa_part_retired.lock().unwrap().push(old);
9712 }
9713 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
9714 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
9715 }
9716 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
9717 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
9718 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
9719 }
9720 let pg = part_guard.as_mut().unwrap();
9721 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
9722 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
9723 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
9724 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
9725 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
9726 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);
9727 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9728 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
9732 let fa512_min = fa512_min_tkv();
9737 let deep = fa_vec && head_dim == 256 && fa_v4_at(t_kv) && !g
9740 && fa_deep_at(t_kv) && !matches!(fa_v4_mode(), "noB3" | "stage");
9741 let (f, cfg) = if fa_vec && head_dim == 512 && t_kv >= fa512_min {
9742 let gqa = (n_head / n_head_kv).max(1) as u32;
9745 let fv = self.fa_func("fa_decode_vec_q_dpl16", head_dim);
9746 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9747 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
9748 } else if fa_vec && head_dim <= 256 {
9749 let gqa = (n_head / n_head_kv).max(1) as u32;
9750 static SMEM_TKV: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9761 let smem_tkv = *SMEM_TKV.get_or_init(|| {
9762 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
9763 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
9764 });
9765 if fa_v4_at(t_kv) && head_dim == 256 {
9766 let v4name = match fa_v4_mode() {
9770 "noB3" => "fa_decode_vec_q_v4_noB3", "stage" => "fa_decode_vec_q_v4_stage", _ if deep => "fa_decode_vec_q_v4_deep",
9773 _ => "fa_decode_vec_q_v4",
9774 };
9775 let fv = if g { self.func_g(v4name) } else { self.func(v4name) };
9776 let shmem = (if deep { 12160 } else { 11520 }
9779 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
9780 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9781 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9782 (fv,
9783 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9784 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9785 } else if fa_v3_active(head_dim) {
9786 let fv = if g { self.func_g("fa_decode_vec_q_v3") } else { self.func("fa_decode_vec_q_v3") };
9789 let shmem = (32 * head_dim * 2) as u32; (fv,
9791 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9792 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9793 } else if fa_v2_on() {
9794 let fv = if g { self.func_g("fa_decode_vec_q_v2") } else { self.func("fa_decode_vec_q_v2") };
9798 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
9800 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9801 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9802 } else if smem_tkv > 0 && t_kv >= smem_tkv && !g
9803 && !(head_dim == 512 && Self::gkv_on()) {
9804 let fv = if g { self.func_g("fa_decode_vec_q_smem") } else { self.func("fa_decode_vec_q_smem") };
9808 let shmem = (2 * 32 * head_dim * 2) as u32; use cudarc::driver::sys::CUfunction_attribute_enum as A;
9810 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9811 (fv,
9812 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9813 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9814 } else {
9815 let fv = if g { self.func_g("fa_decode_vec_q") } else { self.func("fa_decode_vec_q") };
9818 (fv,
9819 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9820 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
9821 }
9822 } else {
9823 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
9826 t_kv, None, scale, n_splits,
9827 if fa_vec { sp } else { 256 },
9828 k_tok_bytes, v_tok_bytes, g,
9829 part_o, part_m, part_l, None);
9830 };
9831 let __s_b = self.gpu.stream();
9832 let mut b = __s_b.launch_builder(&f);
9833 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9834 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&scale).arg(&nsp).arg(&ktb).arg(&vtb);
9835 unsafe { b.launch(cfg)?; }
9836 let (fc, cfg2) = (if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) },
9839 LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 });
9840 let __s_b2 = self.gpu.stream();
9841 let mut b2 = __s_b2.launch_builder(&fc);
9842 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
9843 unsafe { b2.launch(cfg2)?; }
9844 Ok(())
9845 }
9846
9847 #[allow(clippy::too_many_arguments)]
9858 pub fn fa_decode_batch_seqs_v4(&self, q: &CudaSlice<f32>,
9859 kv_ptrs: &cudarc::driver::CudaView<u64>,
9860 pos_seq: &CudaSlice<i32>, o: &mut CudaSlice<f32>,
9861 head_dim: usize, n_head: usize, n_head_kv: usize,
9862 b_n: usize, t_kv_max: usize, scale: f32,
9863 split_keys: usize, k_tok_bytes: usize, v_tok_bytes: usize)
9864 -> Result<(), Box<dyn std::error::Error>> {
9865 debug_assert!(head_dim == 256, "seqs twin is v4-stamped (hd256 only)");
9866 let n_splits_max = (t_kv_max + split_keys - 1) / split_keys;
9867 let o_len = b_n * n_head * n_splits_max * head_dim;
9868 let ml_len = b_n * n_head * n_splits_max;
9869 let mut part_guard = self.fa_part_pool.lock().unwrap();
9870 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
9871 let old = part_guard.take();
9882 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
9883 if let Some(old) = old {
9884 self.fa_part_retired.lock().unwrap().push(old);
9885 }
9886 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
9887 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
9888 }
9889 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
9890 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
9891 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
9892 }
9893 let pg = part_guard.as_mut().unwrap();
9894 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
9895 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
9896 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
9897 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
9898 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9899 let (nspm, spk) = (n_splits_max as i32, split_keys as i32);
9900 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9901 let gqa = (n_head / n_head_kv).max(1) as u32;
9902 let f = self.func("fa_decode_vec_q_seqs_v4");
9903 let shmem = (11520 + 32 * head_dim * 2) as u32;
9905 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9906 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9907 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, b_n as u32),
9908 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
9909 {
9910 let __s_b = self.gpu.stream();
9911 let mut b = __s_b.launch_builder(&f);
9912 b.arg(q).arg(kv_ptrs).arg(pos_seq).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9913 .arg(&hd).arg(&nh).arg(&nhkv).arg(&scale).arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
9914 unsafe { b.launch(cfg)?; }
9915 }
9916 let fc = self.func("fa_decode_combine_seqs");
9917 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, b_n as u32, 1),
9918 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
9919 let __s_b2 = self.gpu.stream();
9920 let mut b2 = __s_b2.launch_builder(&fc);
9921 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
9922 .arg(pos_seq).arg(&nspm).arg(&spk);
9923 unsafe { b2.launch(cfg2)?; }
9924 Ok(())
9925 }
9926
9927 #[allow(clippy::too_many_arguments)]
9934 pub fn append_kv_quantized_seqs(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
9935 kv_ptrs: &cudarc::driver::CudaView<u64>,
9936 pos_seq: &CudaSlice<i32>, b_n: usize,
9937 kv_dim_k: usize, kv_dim_v: usize,
9938 k_tok_bytes: usize, v_tok_bytes: usize)
9939 -> Result<(), Box<dyn std::error::Error>> {
9940 let f = self.func("append_quantize_kv_q8_0_q5_1_seqs");
9941 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
9942 let cfg = LaunchConfig { grid_dim: (nblk, b_n as u32, 1),
9943 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
9944 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9945 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9946 let __s_b = self.gpu.stream();
9947 let mut b = __s_b.launch_builder(&f);
9948 b.arg(k_rows).arg(v_rows).arg(kv_ptrs).arg(pos_seq)
9949 .arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
9950 unsafe { b.launch(cfg)?; }
9951 Ok(())
9952 }
9953
9954 pub fn fa_rows_eligible(&self, base_len: usize, head_dim: usize) -> bool {
9960 std::env::var("MEMRA_NO_FA_VEC").is_err()
9961 && std::env::var("MEMRA_FA_ROWS_OFF").is_err()
9962 && base_len + 1 >= fa_vec_min_tkv()
9963 && head_dim <= 256 && head_dim % 32 == 0
9964 }
9965
9966 #[allow(clippy::too_many_arguments)]
9975 pub fn fa_decode_rows(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9976 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9977 head_dim: usize, n_head: usize, n_head_kv: usize,
9978 base_len: usize, t: usize, scale: f32,
9979 k_tok_bytes: usize, v_tok_bytes: usize,
9980 base_dev: Option<(&CudaSlice<i32>, i32)>,
9984 kv_shared: bool,
9987 g: bool,
9991 mut q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
9994 -> Result<(), Box<dyn std::error::Error>> {
9995 debug_assert!(base_len + 1 >= fa_vec_min_tkv() && head_dim <= 512 && head_dim % 32 == 0);
9996 let t_kv_max = base_len + t; let mut sp = fa_split_keys(t_kv_max, n_head_kv); if head_dim == 512 {
10003 static SP512: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10004 let v = *SP512.get_or_init(|| std::env::var("MEMRA_FA_SP512").ok()
10007 .and_then(|x| x.parse().ok()).unwrap_or(0));
10008 sp = if v >= 8 { v } else { FA_SP512_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) };
10009 }
10010 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10011 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10012 let gqa = (n_head / n_head_kv).max(1) as u32;
10013 let mut groups: Vec<(usize, usize, usize)> = Vec::new(); if head_dim == 512 || fa_split_keys(base_len + 1, n_head_kv) == sp {
10024 groups.push((0, t, sp));
10025 } else {
10026 let mut r0 = 0usize;
10027 while r0 < t {
10028 let sp_g = fa_split_keys(base_len + r0 + 1, n_head_kv);
10029 let mut r1 = r0 + 1;
10030 while r1 < t && fa_split_keys(base_len + r1 + 1, n_head_kv) == sp_g { r1 += 1; }
10031 groups.push((r0, r1 - r0, sp_g));
10032 r0 = r1;
10033 }
10034 }
10035 static SMEM_TKV_R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10039 let smem_tkv = *SMEM_TKV_R.get_or_init(|| {
10040 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
10041 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
10042 });
10043 let v4 = fa_v4_at(base_len + t) && head_dim == 256;
10044 let v3 = fa_v3_active(head_dim);
10045 let smem_rows = head_dim <= 256 && !v3 && !fa_v2_on() && smem_tkv > 0 && t_kv_max >= smem_tkv;
10046 let _ = kv_shared;
10051 let i2 = head_dim == 512 && std::env::var("MEMRA_FA_I2").as_deref() != Ok("0");
10054 static TB512: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10068 let tb512 = head_dim == 512 && sp <= 32 && n_head / n_head_kv.max(1) <= 16
10070 && *TB512.get_or_init(|| std::env::var("MEMRA_FA_TB512").as_deref() != Ok("0"));
10071 let fname = if tb512 { "fa_decode_vec_q_rows_v4_512_tb" }
10072 else if i2 { "fa_decode_vec_q_rows_dpl16_i2" }
10073 else if head_dim == 512 { "fa_decode_vec_q_rows_dpl16" } else if v4 { "fa_decode_vec_q_rows_v4" }
10075 else if v3 { "fa_decode_vec_q_rows_v3" }
10076 else if fa_v2_on() { "fa_decode_vec_q_rows_v2" }
10077 else if smem_rows { "fa_decode_vec_q_rows_smem" }
10078 else { "fa_decode_vec_q_rows" };
10079 let f = if head_dim == 512 { self.fa_func(fname, head_dim) }
10080 else if g {
10081 self.func_g(if smem_rows { "fa_decode_vec_q_rows" } else { fname })
10089 }
10090 else { self.func(fname) };
10091 let shmem = if tb512 {
10092 let gk = Self::gkv_on();
10094 let sh = (8192 + 1024 + 32 * 512 + 32 * 64
10095 + 32 * head_dim * if gk { 1 } else { 2 }) as u32;
10096 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10097 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10098 sh
10099 } else if v4 || v3 || smem_rows || fa_v2_on() {
10100 let sh = (if v4 { 11520 + 32 * head_dim * if g { 1 } else { 2 } }
10102 else if v3 { 32 * head_dim * 2 } else { 2 * 32 * head_dim * 2 }) as u32;
10103 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10104 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10105 sh
10106 } else { 0 };
10107 for &(r0, t_g, sp_g) in &groups {
10111 let n_splits_g = (base_len + r0 + t_g).div_ceil(sp_g);
10112 let (nspm, spk) = (n_splits_g as i32, sp_g as i32);
10113 let base_i = (base_len + r0) as i32;
10114 let o_len = t_g * n_head * n_splits_g * head_dim;
10115 let ml_len = t_g * n_head * n_splits_g;
10116 let mut part_guard = self.fa_part_pool.lock().unwrap();
10117 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10118 let old = part_guard.take();
10129 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10130 if let Some(old) = old {
10131 self.fa_part_retired.lock().unwrap().push(old);
10132 }
10133 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10134 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10135 }
10136 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10137 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10138 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10139 }
10140 let pg = part_guard.as_mut().unwrap();
10141 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10142 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10143 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10144 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10145 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
10146 let qv = self.view(q, t * n_head * head_dim);
10147 let q_g = qv.slice(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10148 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_g as u32, t_g as u32),
10149 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10150 {
10151 let __s_b = self.gpu.stream();
10152 let mut b = __s_b.launch_builder(&f);
10153 if tb512 {
10154 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10156 let plus_g = plus + r0 as i32;
10157 let nr = t_g as i32;
10158 if Self::pdl_on() && Self::pdl_wb_on() {
10159 use cudarc::driver::{DevicePtr, DevicePtrMut};
10161 let s = &self.gpu.stream();
10162 let (pq, _b0) = q_g.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10163 let (pv, _b2) = v.device_ptr(s);
10164 let (po, _b3) = part_o.device_ptr_mut(s);
10165 let (pm, _b4) = part_m.device_ptr_mut(s);
10166 let (pl, _b5) = part_l.device_ptr_mut(s);
10167 let (pb, _b6) = bd.device_ptr(s);
10168 let mut ps = [
10169 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10170 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10171 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10172 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10173 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10174 &plus_g as *const _ as *mut _, &scale as *const _ as *mut _,
10175 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10176 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10177 &nr as *const _ as *mut _,
10178 ];
10179 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10180 "fa_decode_vec_q_rows_v4_512_tb",
10181 (n_head_kv as u32, n_splits_g as u32, 1), (32, gqa, 1),
10182 shmem, &mut ps)?; }
10183 } else {
10184 let cfg_tb = LaunchConfig {
10185 grid_dim: (n_head_kv as u32, n_splits_g as u32, 1),
10186 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10187 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10188 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10189 .arg(&ktb).arg(&vtb).arg(&nr);
10190 unsafe { b.launch(cfg_tb)?; }
10191 }
10192 } else if head_dim == 512 {
10193 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10194 let plus_g = plus + r0 as i32;
10195 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10196 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10197 .arg(&ktb).arg(&vtb);
10198 unsafe { b.launch(cfg)?; }
10199 } else {
10200 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10201 .arg(&hd).arg(&nh).arg(&nhkv).arg(&base_i).arg(&scale).arg(&nspm).arg(&spk)
10202 .arg(&ktb).arg(&vtb);
10203 unsafe { b.launch(cfg)?; }
10204 }
10205 }
10206 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t_g as u32, 1),
10207 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10208 let mut o_g = o.slice_mut(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10209 if head_dim == 512 {
10210 let (bd, plus) = base_dev.unwrap();
10213 let plus_g = plus + r0 as i32;
10214 if let Some((oq, od)) = q8_out.as_mut() {
10215 debug_assert!(t == 1, "rows q8 emit is a t=1 decode arm");
10217 if Self::pdl_on() && Self::pdl_wb_on() {
10218 use cudarc::driver::{DevicePtr, DevicePtrMut};
10220 let s = &self.gpu.stream();
10221 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10222 let (pl, _g2) = part_l.device_ptr(s);
10223 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10224 let (pb, _g5) = bd.device_ptr(s);
10225 let mut ps = [
10226 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10227 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10228 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10229 &nh as *const _ as *mut _, &pb as *const _ as *mut _,
10230 &plus_g as *const _ as *mut _, &nspm as *const _ as *mut _,
10231 &spk as *const _ as *mut _,
10232 ];
10233 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10234 "fa_decode_combine_rows_dc_q8_1",
10235 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10236 continue;
10237 }
10238 let fc = self.fa_func("fa_decode_combine_rows_dc_q8_1", head_dim);
10239 let __s_b2 = self.gpu.stream();
10240 let mut b2 = __s_b2.launch_builder(&fc);
10241 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut **oq).arg(&mut **od)
10242 .arg(&hd).arg(&nh).arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10243 unsafe { b2.launch(cfg2)?; }
10244 continue;
10245 }
10246 let fc = self.fa_func("fa_decode_combine_rows_dc", head_dim);
10247 let __s_b2 = self.gpu.stream();
10248 let mut b2 = __s_b2.launch_builder(&fc);
10249 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10250 .arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10251 unsafe { b2.launch(cfg2)?; }
10252 } else {
10253 assert!(q8_out.is_none(), "rows q8 emit requires the hd512 dc combine");
10256 let fc = self.func("fa_decode_combine_rows");
10257 let __s_b2 = self.gpu.stream();
10258 let mut b2 = __s_b2.launch_builder(&fc);
10259 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10260 .arg(&base_i).arg(&nspm).arg(&spk);
10261 unsafe { b2.launch(cfg2)?; }
10262 }
10263 }
10264 Ok(())
10265 }
10266
10267 #[allow(clippy::too_many_arguments)]
10271 pub fn fa_decode_rows_w(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10272 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10273 head_dim: usize, n_head: usize, n_head_kv: usize,
10274 base_dev: &CudaSlice<i32>, base_plus: i32, t: usize, scale: f32,
10275 window: usize, k_tok_bytes: usize, v_tok_bytes: usize,
10276 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10277 -> Result<(), Box<dyn std::error::Error>> {
10278 debug_assert!(head_dim == 256);
10283 let sp = {
10291 static SPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10292 let v = *SPW.get_or_init(|| std::env::var("MEMRA_FA_SPW").ok()
10293 .and_then(|x| x.parse().ok()).unwrap_or(0));
10294 if v >= 8 { v } else { FA_SPW_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) }
10295 };
10296 let n_splits_max = (window + sp - 1) / sp;
10297 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10298 let (nspm, spk, wini) = (n_splits_max as i32, sp as i32, window as i32);
10299 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10300 let gqa = (n_head / n_head_kv).max(1) as u32;
10301 let o_len = t * n_head * n_splits_max * head_dim;
10302 let ml_len = t * n_head * n_splits_max;
10303 let mut part_guard = self.fa_part_pool.lock().unwrap();
10304 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10305 let old = part_guard.take();
10316 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10317 if let Some(old) = old {
10318 self.fa_part_retired.lock().unwrap().push(old);
10319 }
10320 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10321 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10322 }
10323 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10324 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10325 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10326 }
10327 let pg = part_guard.as_mut().unwrap();
10328 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10329 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10330 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10331 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10332 static SMEM_TKV_W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10338 let smem_tkv = *SMEM_TKV_W.get_or_init(|| {
10339 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
10340 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
10341 });
10342 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10348 let wg = Self::wkv_on();
10353 let sp2 = gqa <= 4 && fa_v4_at(window)
10356 && std::env::var("MEMRA_FA_SPW2").as_deref() != Ok("0");
10357 if sp2 {
10358 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10359 if Self::pdl_on() && Self::pdl_wb_on() {
10360 use cudarc::driver::{DevicePtr, DevicePtrMut};
10362 let s = &self.gpu.stream();
10363 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10364 let (pv, _b2) = v.device_ptr(s);
10365 let (po, _b3) = part_o.device_ptr_mut(s);
10366 let (pm, _b4) = part_m.device_ptr_mut(s);
10367 let (pl, _b5) = part_l.device_ptr_mut(s);
10368 let (pb, _b6) = base_dev.device_ptr(s);
10369 let mut ps = [
10370 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10371 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10372 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10373 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10374 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10375 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10376 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10377 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10378 &wini as *const _ as *mut _,
10379 ];
10380 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w_sp",
10381 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa + 1, 1),
10382 sh, &mut ps)?; }
10383 } else {
10384 let f = if wg { self.func_g("fa_decode_vec_q_rows_v4_w_sp") }
10385 else { self.func("fa_decode_vec_q_rows_v4_w_sp") };
10386 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10387 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10388 block_dim: (32, gqa + 1, 1), shared_mem_bytes: sh };
10389 let __s_b = self.gpu.stream();
10390 let mut b = __s_b.launch_builder(&f);
10391 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10392 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10393 .arg(&ktb).arg(&vtb).arg(&wini);
10394 unsafe { b.launch(cfg)?; }
10395 }
10396 } else {
10397 if fa_v4_at(window) && Self::pdl_on() && Self::pdl_wb_on() {
10398 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10400 use cudarc::driver::{DevicePtr, DevicePtrMut};
10401 let s = &self.gpu.stream();
10402 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10403 let (pv, _b2) = v.device_ptr(s);
10404 let (po, _b3) = part_o.device_ptr_mut(s);
10405 let (pm, _b4) = part_m.device_ptr_mut(s);
10406 let (pl, _b5) = part_l.device_ptr_mut(s);
10407 let (pb, _b6) = base_dev.device_ptr(s);
10408 let mut ps = [
10409 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10410 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10411 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10412 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10413 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10414 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10415 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10416 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10417 &wini as *const _ as *mut _,
10418 ];
10419 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w",
10420 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa, 1),
10421 sh, &mut ps)?; }
10422 } else {
10423 let pick = |name: &str| if wg { self.func_g(name) } else { self.func(name) };
10424 let (f, sh) = if fa_v4_at(window) {
10425 let f = pick("fa_decode_vec_q_rows_v4_w");
10426 (f, (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32)
10427 } else if smem_tkv > 0 && window >= smem_tkv {
10428 (pick("fa_decode_vec_q_rows_smem_w"), (2 * 32 * head_dim * 2) as u32)
10431 } else {
10432 (pick("fa_decode_vec_q_rows_reg_w"), 0u32)
10433 };
10434 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10435 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10436 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10437 let __s_b = self.gpu.stream();
10438 let mut b = __s_b.launch_builder(&f);
10439 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10440 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10441 .arg(&ktb).arg(&vtb).arg(&wini);
10442 unsafe { b.launch(cfg)?; }
10443 }
10444 }
10445 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10446 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10447 if let Some((oq, od)) = q8_out {
10448 if Self::pdl_on() && Self::pdl_wb_on() {
10451 use cudarc::driver::{DevicePtr, DevicePtrMut};
10453 let s = &self.gpu.stream();
10454 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10455 let (pl, _g2) = part_l.device_ptr(s);
10456 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10457 let mut ps = [
10458 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10459 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10460 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10461 &nh as *const _ as *mut _, &nspm as *const _ as *mut _,
10462 &spk as *const _ as *mut _, &wini as *const _ as *mut _,
10463 ];
10464 unsafe { self.launch_pdl_flash(wg, "fa_decode_combine_rows_w_q8_1",
10465 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10466 return Ok(());
10467 }
10468 let fc = if wg { self.func_g("fa_decode_combine_rows_w_q8_1") }
10469 else { self.func("fa_decode_combine_rows_w_q8_1") };
10470 let __s_b2 = self.gpu.stream();
10471 let mut b2 = __s_b2.launch_builder(&fc);
10472 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh)
10473 .arg(&nspm).arg(&spk).arg(&wini);
10474 unsafe { b2.launch(cfg2)?; }
10475 return Ok(());
10476 }
10477 let fc = if wg { self.func_g("fa_decode_combine_rows_w") }
10478 else { self.func("fa_decode_combine_rows_w") };
10479 let __s_b2 = self.gpu.stream();
10480 let mut b2 = __s_b2.launch_builder(&fc);
10481 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10482 .arg(&nspm).arg(&spk).arg(&wini);
10483 unsafe { b2.launch(cfg2)?; }
10484 Ok(())
10485 }
10486
10487 #[allow(clippy::too_many_arguments)]
10493 pub fn fa_decode_rows_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10494 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10495 head_dim: usize, n_head: usize, n_head_kv: usize,
10496 base_dev: &CudaSlice<i32>, t_kv_upper: usize, t: usize, scale: f32,
10497 k_tok_bytes: usize, v_tok_bytes: usize, base_plus: i32, g: bool)
10498 -> Result<(), Box<dyn std::error::Error>> {
10499 let v4 = head_dim == 256 && fa_v4_at(t_kv_upper);
10500 assert!(v4 || fa_v3_active(head_dim), "stream fa rows requires the v3 or v4 lane");
10501 assert!(v4 || base_plus == 0, "v3_dc kernel takes no plus arg");
10502 if v4 {
10503 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10504 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10505 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10506 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10507 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10508 let gqa = (n_head / n_head_kv).max(1) as u32;
10509 let o_len = t * n_head * n_splits_max * head_dim;
10510 let ml_len = t * n_head * n_splits_max;
10511 let mut part_guard = self.fa_part_pool.lock().unwrap();
10512 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10513 let old = part_guard.take();
10524 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10525 if let Some(old) = old {
10526 self.fa_part_retired.lock().unwrap().push(old);
10527 }
10528 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10529 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10530 }
10531 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10532 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10533 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10534 }
10535 let pg = part_guard.as_mut().unwrap();
10536 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10537 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10538 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10539 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10540 let f = if g { self.func_g("fa_decode_vec_q_rows_v4_dc") }
10541 else { self.func("fa_decode_vec_q_rows_v4_dc") };
10542 let sh = (11520 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10543 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10544 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10545 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10546 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10547 let __s_b = self.gpu.stream();
10548 let mut b = __s_b.launch_builder(&f);
10549 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10550 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale)
10551 .arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
10552 unsafe { b.launch(cfg)?; }
10553 let fc = self.func("fa_decode_combine_rows_dc");
10554 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10555 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10556 let __s_b2 = self.gpu.stream();
10557 let mut b2 = __s_b2.launch_builder(&fc);
10558 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10559 .arg(base_dev).arg(&base_plus).arg(&nspm).arg(&spk);
10560 unsafe { b2.launch(cfg2)?; }
10561 return Ok(());
10562 }
10563 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10564 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10565 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10566 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10567 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10568 let gqa = (n_head / n_head_kv).max(1) as u32;
10569 let o_len = t * n_head * n_splits_max * head_dim;
10570 let ml_len = t * n_head * n_splits_max;
10571 let mut part_guard = self.fa_part_pool.lock().unwrap();
10572 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10573 let old = part_guard.take();
10584 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10585 if let Some(old) = old {
10586 self.fa_part_retired.lock().unwrap().push(old);
10587 }
10588 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10589 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10590 }
10591 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10592 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10593 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10594 }
10595 let pg = part_guard.as_mut().unwrap();
10596 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10597 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10598 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10599 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10600 let f = self.func("fa_decode_vec_q_rows_v3_dc");
10601 let sh = (32 * head_dim * 2) as u32;
10602 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10603 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10604 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10605 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10606 let __s_b = self.gpu.stream();
10607 let mut b = __s_b.launch_builder(&f);
10608 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10609 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&scale).arg(&nspm).arg(&spk)
10610 .arg(&ktb).arg(&vtb);
10611 unsafe { b.launch(cfg)?; }
10612 let fc = self.func("fa_decode_combine_rows_dc");
10613 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10614 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10615 let plus0 = 0i32;
10616 let __s_b2 = self.gpu.stream();
10617 let mut b2 = __s_b2.launch_builder(&fc);
10618 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10619 .arg(base_dev).arg(&plus0).arg(&nspm).arg(&spk);
10620 unsafe { b2.launch(cfg2)?; }
10621 Ok(())
10622 }
10623
10624 pub fn fa_decode_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10635 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10636 head_dim: usize, n_head: usize, n_head_kv: usize,
10637 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10638 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
10639 -> Result<(), Box<dyn std::error::Error>> {
10640 self.fa_decode_dc_q8(q, k, v, o, head_dim, n_head, n_head_kv, t_kv_dev, bucket_max,
10641 scale, k_tok_bytes, v_tok_bytes, g, None)
10642 }
10643
10644 #[allow(clippy::too_many_arguments)]
10647 pub fn fa_decode_dc_q8(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10648 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10649 head_dim: usize, n_head: usize, n_head_kv: usize,
10650 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10651 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
10652 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10653 -> Result<(), Box<dyn std::error::Error>> {
10654 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
10662 if g && head_dim == 256 && !fa_v4_at(bucket_max) { fa_vec = false; } let sp = fa_split_keys(bucket_max, n_head_kv);
10664 let n_splits = if fa_vec { ((bucket_max + sp - 1) / sp).max(1) } else { ((bucket_max + 255) / 256).max(1) };
10665 let o_len = n_head * n_splits * head_dim;
10666 let ml_len = n_head * n_splits;
10667 let mut part_guard = self.fa_part_pool.lock().unwrap();
10668 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10669 let old = part_guard.take();
10680 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10681 if let Some(old) = old {
10682 self.fa_part_retired.lock().unwrap().push(old);
10683 }
10684 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10685 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10686 }
10687 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10688 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10689 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10690 }
10691 let pg = part_guard.as_mut().unwrap();
10692 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10693 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10694 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10695 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10696 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
10697 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10698 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
10699 let deep = fa_vec && head_dim == 256 && fa_v4_at(bucket_max) && !g
10702 && fa_deep_at(bucket_max) && !matches!(fa_v4_mode(), "noB3" | "stage");
10703 let (f, cfg) = if fa_vec && head_dim == 512 && bucket_max >= {
10704 static FA512_MIN_DC: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10705 *FA512_MIN_DC.get_or_init(|| std::env::var("MEMRA_FA512_MIN").ok()
10706 .and_then(|v| v.parse().ok()).unwrap_or(512))
10707 } {
10708 let gqa = (n_head / n_head_kv).max(1) as u32;
10710 (self.fa_func("fa_decode_vec_q_dpl16_dc", head_dim),
10711 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10712 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10713 } else if fa_vec && head_dim == 512 {
10714 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
10717 0, Some(t_kv_dev), scale, n_splits, sp,
10718 k_tok_bytes, v_tok_bytes, g,
10719 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10720 } else if fa_vec && head_dim == 256 && fa_v4_at(bucket_max) {
10721 let gqa = (n_head / n_head_kv).max(1) as u32;
10724 let fv = if g { self.func_g("fa_decode_vec_q_v4_dc") }
10725 else if deep { self.func("fa_decode_vec_q_v4_deep_dc") }
10726 else { self.func("fa_decode_vec_q_v4_dc") };
10727 let shmem = (if deep { 12160 } else { 11520 }
10728 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10729 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10730 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
10731 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10732 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10733 } else if fa_vec && fa_v3_active(head_dim) {
10734 let gqa = (n_head / n_head_kv).max(1) as u32;
10737 let fv = if g { self.func_g("fa_decode_vec_q_v3_dc") } else { self.func("fa_decode_vec_q_v3_dc") };
10738 let shmem = (32 * head_dim * 2) as u32; (fv,
10740 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10741 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10742 } else if fa_vec && fa_v2_on() {
10743 let gqa = (n_head / n_head_kv).max(1) as u32;
10747 let fv = if g { self.func_g("fa_decode_vec_q_v2_dc") } else { self.func("fa_decode_vec_q_v2_dc") };
10748 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
10750 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10751 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10752 } else if fa_vec {
10753 let gqa = (n_head / n_head_kv).max(1) as u32;
10754 let fv = if g { self.func_g("fa_decode_vec_q_dc") } else { self.func("fa_decode_vec_q_dc") };
10756 (fv,
10757 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10758 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10759 } else {
10760 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
10761 0, Some(t_kv_dev), scale, n_splits,
10762 if fa_vec { sp } else { 256 },
10763 k_tok_bytes, v_tok_bytes, g,
10764 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10765 };
10766 let ski = sp as i32; let __s_b = self.gpu.stream();
10768 let mut b = __s_b.launch_builder(&f);
10769 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10770 .arg(&hd).arg(&nh).arg(&nhkv).arg(t_kv_dev).arg(&scale).arg(&nsp).arg(&ski)
10771 .arg(&ktb).arg(&vtb);
10772 unsafe { b.launch(cfg)?; }
10773 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10774 if let Some((oq, od)) = q8_out {
10775 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
10776 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
10777 let __s_b2 = self.gpu.stream();
10778 let mut b2 = __s_b2.launch_builder(&fc);
10779 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
10780 unsafe { b2.launch(cfg2)?; }
10781 return Ok(());
10782 }
10783 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
10784 let __s_b2 = self.gpu.stream();
10785 let mut b2 = __s_b2.launch_builder(&fc);
10786 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
10787 unsafe { b2.launch(cfg2)?; }
10788 Ok(())
10789 }
10790
10791 pub fn fa_geom_eager(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
10797 let fa_ok = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
10801 let vec512 = fa_ok && head_dim == 512 && t_kv >= fa512_min_tkv();
10807 let mut fa_vec = vec512 || (fa_ok && head_dim <= 256 && head_dim % 32 == 0);
10808 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
10814 let sp = fa_split_keys(t_kv, n_head_kv);
10815 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
10816 (fa_vec, n_splits)
10817 }
10818
10819 pub fn fa_bucket_key(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
10825 self.fa_geom_eager(t_kv, head_dim, n_head_kv, g)
10826 }
10827
10828 pub fn capture_graph_retained<F>(&self, step: F)
10840 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
10841 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
10842 {
10843 use cudarc::driver::sys::CUgraphInstantiate_flags;
10844 self.capture_graph_retained_flags(
10845 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH, step)
10846 }
10847
10848 pub fn capture_graph_retained_flags<F>(&self,
10853 flags: cudarc::driver::sys::CUgraphInstantiate_flags, mut step: F)
10854 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
10855 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
10856 {
10857 use cudarc::driver::sys::CUstreamCaptureMode;
10858 self.capture_keep.lock().unwrap().clear();
10866 let was_tracking = self.gpu.ctx.is_event_tracking();
10867 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
10868 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
10869 self.capture_keep_on.store(true, std::sync::atomic::Ordering::Relaxed);
10870 let w = (|| { step(self)?; step(self) })();
10871 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
10872 w?;
10873 self.gpu.stream().synchronize()?;
10874 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
10875 let r = step(self);
10876 let g = self.gpu.stream().end_capture(flags);
10877 r?;
10878 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
10879 graph.upload()?;
10880 Ok(graph)
10881 };
10882 let result = run();
10883 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
10884 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
10885 let keeper = std::mem::take(&mut *self.capture_keep.lock().unwrap());
10886 Ok((result?, keeper))
10887 }
10888
10889 pub fn capture_graph<F>(&self, mut step: F) -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>>
10890 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
10891 {
10892 use cudarc::driver::sys::{CUstreamCaptureMode, CUgraphInstantiate_flags};
10893 let was_tracking = self.gpu.ctx.is_event_tracking();
10901 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
10902 let iflag = {
10909 static F: std::sync::OnceLock<CUgraphInstantiate_flags> = std::sync::OnceLock::new();
10910 *F.get_or_init(|| match std::env::var("MEMRA_GRAPH_IFLAG").as_deref() {
10911 Ok("upload") => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_UPLOAD,
10914 Ok("priority") =>
10915 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
10916 _ => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
10917 })
10918 };
10919 let ct = {
10926 static T: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10927 *T.get_or_init(|| std::env::var("MEMRA_GRAPH_CAPTIME").as_deref() == Ok("1"))
10928 };
10929 let warmups = {
10952 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10953 *W.get_or_init(|| std::env::var("MEMRA_GRAPH_WARMUPS").ok()
10954 .and_then(|v| v.parse().ok()).filter(|n| *n >= 1).unwrap_or(1))
10955 };
10956 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
10957 let t_w = std::time::Instant::now();
10958 for _ in 0..warmups { step(self)?; }
10960 self.gpu.stream().synchronize()?;
10961 let ms_warm = t_w.elapsed().as_secs_f64() * 1e3;
10962 let t_c = std::time::Instant::now();
10964 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
10965 let r = step(self);
10968 let ms_body = t_c.elapsed().as_secs_f64() * 1e3;
10969 let t_i = std::time::Instant::now();
10970 let g = self.gpu.stream().end_capture(iflag);
10971 let ms_inst = t_i.elapsed().as_secs_f64() * 1e3;
10972 r?;
10973 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
10974 let t_u = std::time::Instant::now();
10975 graph.upload()?;
10976 if ct {
10977 println!("[graph-captime] warmup2x {ms_warm:.2} ms capture-body {ms_body:.2} ms \
10978 instantiate {ms_inst:.2} ms upload {:.2} ms",
10979 t_u.elapsed().as_secs_f64() * 1e3);
10980 }
10981 Ok(graph)
10982 };
10983 let result = run();
10984 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
10985 result
10986 }
10987
10988 pub fn gdn_scan_s128_view(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
10990 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
10991 state_in: &cudarc::driver::CudaView<f32>,
10992 state_out: &mut cudarc::driver::CudaViewMut<f32>,
10993 o: &mut CudaSlice<f32>, n_head: usize, t: usize, scale: f32)
10994 -> Result<(), Box<dyn std::error::Error>> {
10995 let f = self.func("gdn_scan_s128");
10996 const S_V: u32 = 128; const WARP: u32 = 32; const COLS: u32 = 4;
10997 let cfg = LaunchConfig { grid_dim: (n_head as u32, 1, S_V / COLS), block_dim: (WARP, COLS, 1), shared_mem_bytes: 0 };
10998 let (h, ti) = (n_head as i32, t as i32);
10999 let __s_b = self.gpu.stream();
11000 let mut b = __s_b.launch_builder(&f);
11001 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);
11002 unsafe { b.launch(cfg)?; }
11003 Ok(())
11004 }
11005
11006 pub fn ssm_conv1d_view(&self, x: &cudarc::driver::CudaView<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11008 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
11009 -> Result<(), Box<dyn std::error::Error>> {
11010 let f = self.func("ssm_conv1d_silu_f32");
11011 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
11013 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11014 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
11015 let __s_b = self.gpu.stream();
11016 let mut b = __s_b.launch_builder(&f);
11017 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
11018 unsafe { b.launch(cfg)?; }
11019 Ok(())
11020 }
11021
11022 pub fn ssm_conv1d_tm(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11029 conv_dim: usize, t: usize, d_conv: usize)
11030 -> Result<(), Box<dyn std::error::Error>> {
11031 let f = self.func("ssm_conv1d_tm_f32");
11032 let cfg = LaunchConfig {
11033 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11034 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11035 };
11036 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11037 let __s_b = self.gpu.stream();
11038 let mut b = __s_b.launch_builder(&f);
11039 b.arg(qkv_tm).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11040 unsafe { b.launch(cfg)?; }
11041 Ok(())
11042 }
11043
11044 pub fn ssm_conv1d_tm_state(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
11052 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11053 conv_dim: usize, t: usize, d_conv: usize)
11054 -> Result<(), Box<dyn std::error::Error>> {
11055 self.ssm_conv1d_tm_state_pad(qkv_tm, conv_state, w, y, conv_dim, t, d_conv, None)
11056 }
11057
11058 #[allow(clippy::too_many_arguments)]
11061 pub fn ssm_conv1d_tm_state_pad(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
11062 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11063 conv_dim: usize, t: usize, d_conv: usize,
11064 pad_len: Option<&CudaSlice<i32>>)
11065 -> Result<(), Box<dyn std::error::Error>> {
11066 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
11067 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11071 {
11072 let f = self.func("ssm_conv1d_tm_state_f32");
11073 let cfg = LaunchConfig {
11074 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11075 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11076 };
11077 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11078 let __s_b = self.gpu.stream();
11079 let mut b = __s_b.launch_builder(&f);
11080 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11081 unsafe { b.launch(cfg)?; }
11082 }
11083 match (ring_old, pad_len) {
11084 (None, Some(len_d)) => {
11085 let f = self.func("ssm_conv_ring_update_dev_f32");
11086 let n = conv_dim * (d_conv - 1);
11087 let cfg = LaunchConfig::for_num_elems(n as u32);
11088 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11089 let __s_b = self.gpu.stream();
11090 let mut b = __s_b.launch_builder(&f);
11091 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11092 unsafe { b.launch(cfg)?; }
11093 }
11094 (None, None) => {
11095 let f = self.func("ssm_conv_ring_update_f32");
11096 let n = conv_dim * (d_conv - 1);
11097 let cfg = LaunchConfig::for_num_elems(n as u32);
11098 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11099 let __s_b = self.gpu.stream();
11100 let mut b = __s_b.launch_builder(&f);
11101 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11102 unsafe { b.launch(cfg)?; }
11103 }
11104 (Some(old), _) => self.ssm_conv_ring_rebuild(qkv_tm, &old, conv_state, conv_dim, t, d_conv)?,
11105 }
11106 Ok(())
11107 }
11108
11109 pub fn ssm_conv1d_tm_state_pad_v(&self, qkv_tm: &cudarc::driver::CudaView<f32>, conv_state: &mut CudaSlice<f32>,
11111 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11112 conv_dim: usize, t: usize, d_conv: usize,
11113 pad_len: Option<&CudaSlice<i32>>)
11114 -> Result<(), Box<dyn std::error::Error>> {
11115 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
11116 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11120 {
11121 let f = self.func("ssm_conv1d_tm_state_f32");
11122 let cfg = LaunchConfig {
11123 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11124 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11125 };
11126 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11127 let __s_b = self.gpu.stream();
11128 let mut b = __s_b.launch_builder(&f);
11129 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11130 unsafe { b.launch(cfg)?; }
11131 }
11132 match (ring_old, pad_len) {
11133 (None, Some(len_d)) => {
11134 let f = self.func("ssm_conv_ring_update_dev_f32");
11135 let n = conv_dim * (d_conv - 1);
11136 let cfg = LaunchConfig::for_num_elems(n as u32);
11137 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11138 let __s_b = self.gpu.stream();
11139 let mut b = __s_b.launch_builder(&f);
11140 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11141 unsafe { b.launch(cfg)?; }
11142 }
11143 (None, None) => {
11144 let f = self.func("ssm_conv_ring_update_f32");
11145 let n = conv_dim * (d_conv - 1);
11146 let cfg = LaunchConfig::for_num_elems(n as u32);
11147 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11148 let __s_b = self.gpu.stream();
11149 let mut b = __s_b.launch_builder(&f);
11150 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11151 unsafe { b.launch(cfg)?; }
11152 }
11153 (Some(_), _) => unreachable!(
11154 "ssm_conv1d_tm_state_pad_v: T < d_conv-1 has no view path (PRIME_MIN_T gates it)"),
11155 }
11156 Ok(())
11157 }
11158
11159 pub fn ssm_conv_ring_rebuild(&self, qkv_tm: &CudaSlice<f32>, ring_old: &CudaSlice<f32>,
11164 conv_state: &mut CudaSlice<f32>,
11165 conv_dim: usize, tc: usize, d_conv: usize)
11166 -> Result<(), Box<dyn std::error::Error>> {
11167 let f = self.func("ssm_conv_ring_rebuild_f32");
11168 let n = conv_dim * (d_conv - 1);
11169 let cfg = LaunchConfig::for_num_elems(n as u32);
11170 let (cd, ti, dc) = (conv_dim as i32, tc as i32, d_conv as i32);
11171 let __s_b = self.gpu.stream();
11172 let mut b = __s_b.launch_builder(&f);
11173 b.arg(qkv_tm).arg(ring_old).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11174 unsafe { b.launch(cfg)?; }
11175 Ok(())
11176 }
11177
11178 #[allow(clippy::too_many_arguments)]
11183 pub fn gdn_prep_decode(&self, conv_out: &CudaSlice<f32>, beta_raw: &CudaSlice<f32>,
11184 alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11185 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11186 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11187 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32)
11188 -> Result<(), Box<dyn std::error::Error>> {
11189 let f = self.func("gdn_prep_decode_f32");
11190 let cfg = LaunchConfig { grid_dim: (num_v as u32, 1, 1), block_dim: (32, 4, 1), shared_mem_bytes: 0 };
11191 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11192 let __s_b = self.gpu.stream();
11193 let mut b = __s_b.launch_builder(&f);
11194 b.arg(conv_out).arg(beta_raw).arg(alpha).arg(dt_bias).arg(a)
11195 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11196 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps);
11197 unsafe { b.launch(cfg)?; }
11198 Ok(())
11199 }
11200
11201 #[allow(clippy::too_many_arguments)]
11205 pub fn ssm_conv1d_gdn(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>,
11206 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11207 conv_dim: usize, t: usize, d_conv: usize,
11208 d_state: usize, num_v: usize, num_k: usize, key_dim: usize)
11209 -> Result<(), Box<dyn std::error::Error>> {
11210 let f = self.func("ssm_conv1d_gdn_f32");
11211 let cfg = LaunchConfig {
11212 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11213 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11214 };
11215 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11216 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11217 let __s_b = self.gpu.stream();
11218 let mut b = __s_b.launch_builder(&f);
11219 b.arg(qkv_tm).arg(w).arg(q_g).arg(k_g).arg(v_g)
11220 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd);
11221 unsafe { b.launch(cfg)?; }
11222 Ok(())
11223 }
11224
11225 pub fn ssm_conv1d(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11226 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
11227 -> Result<(), Box<dyn std::error::Error>> {
11228 let f = self.func("ssm_conv1d_silu_f32");
11229 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
11230 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11231 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
11232 let __s_b = self.gpu.stream();
11233 let mut b = __s_b.launch_builder(&f);
11234 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
11235 unsafe { b.launch(cfg)?; }
11236 Ok(())
11237 }
11238
11239 pub fn gdn_scan_s128(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11242 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
11243 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11244 n_head: usize, t: usize, scale: f32)
11245 -> Result<(), Box<dyn std::error::Error>> {
11246 let f = self.func("gdn_scan_s128");
11247 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11248 let cfg = LaunchConfig {
11249 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
11250 block_dim: (WARP, COLS_PER_BLOCK, 1),
11251 shared_mem_bytes: 0,
11252 };
11253 let (h, ti) = (n_head as i32, t as i32);
11254 let __s_b = self.gpu.stream();
11255 let mut b = __s_b.launch_builder(&f);
11256 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);
11257 unsafe { b.launch(cfg)?; }
11258 Ok(())
11259 }
11260
11261 #[allow(clippy::too_many_arguments)]
11266 pub fn ssm_conv1d_fused_decode_b(
11267 &self, qkv_cols: &CudaSlice<f32>, conv_state_ptrs: &cudarc::driver::CudaView<u64>,
11268 w: &CudaSlice<f32>, conv_outs: &mut CudaSlice<f32>, conv_dim: usize, d_conv: usize,
11269 b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11270 let f = self.func("ssm_conv1d_fused_decode_b_f32");
11271 let cfg = LaunchConfig {
11272 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
11273 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11274 };
11275 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11276 let __s_b = self.gpu.stream();
11277 let mut b = __s_b.launch_builder(&f);
11278 b.arg(qkv_cols).arg(conv_state_ptrs).arg(w).arg(conv_outs).arg(&cd).arg(&dc);
11279 unsafe { b.launch(cfg)?; }
11280 Ok(())
11281 }
11282
11283 #[allow(clippy::too_many_arguments)]
11284 pub fn gdn_prep_decode_b(
11285 &self, conv_outs: &CudaSlice<f32>, beta_raws: &CudaSlice<f32>, alphas: &CudaSlice<f32>,
11286 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11287 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11288 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11289 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32,
11290 conv_dim: usize, b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11291 let f = self.func("gdn_prep_decode_b_f32");
11292 let cfg = LaunchConfig {
11293 grid_dim: (num_v as u32, 1, b_n as u32),
11294 block_dim: (32, 4, 1), shared_mem_bytes: 0,
11295 };
11296 let (ds, nv, nk, kd, cd) =
11297 (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, conv_dim as i32);
11298 let __s_b = self.gpu.stream();
11299 let mut b = __s_b.launch_builder(&f);
11300 b.arg(conv_outs).arg(beta_raws).arg(alphas).arg(dt_bias).arg(a)
11301 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11302 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps).arg(&cd);
11303 unsafe { b.launch(cfg)?; }
11304 Ok(())
11305 }
11306
11307 #[allow(clippy::too_many_arguments)]
11308 pub fn gdn_scan_s128_batched(
11309 &self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11310 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
11311 state_in_ptrs: &cudarc::driver::CudaView<u64>,
11312 state_out_ptrs: &cudarc::driver::CudaView<u64>,
11313 o: &mut CudaSlice<f32>, n_head: usize, b_n: usize, scale: f32)
11314 -> Result<(), Box<dyn std::error::Error>> {
11315 let f = self.func("gdn_scan_s128_b");
11316 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11317 let cfg = LaunchConfig {
11318 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
11319 block_dim: (WARP, COLS_PER_BLOCK, 1), shared_mem_bytes: 0,
11320 };
11321 let h = n_head as i32;
11322 let __s_b = self.gpu.stream();
11323 let mut b = __s_b.launch_builder(&f);
11324 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in_ptrs).arg(state_out_ptrs)
11325 .arg(o).arg(&h).arg(&scale);
11326 unsafe { b.launch(cfg)?; }
11327 Ok(())
11328 }
11329
11330 pub fn gdn_chunked_enabled() -> bool {
11339 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11340 *E.get_or_init(|| std::env::var("MEMRA_GDN_CHUNKED").map(|v| v != "0").unwrap_or(true))
11341 }
11342
11343 pub fn gdn_chunk_size() -> usize {
11348 static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
11349 *C.get_or_init(|| {
11350 let c: usize = std::env::var("MEMRA_GDN_CHUNK").ok()
11351 .and_then(|v| v.parse().ok()).unwrap_or(32);
11352 c.clamp(32, 128) / 32 * 32
11353 })
11354 }
11355
11356 #[allow(clippy::too_many_arguments)]
11361 #[allow(clippy::too_many_arguments, clippy::type_complexity)]
11364 #[allow(clippy::too_many_arguments)]
11365 pub fn gdn_chunk_k123(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11366 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, wb16: Option<&mut CudaSlice<u8>>,
11367 n_head: usize, t: usize, c: usize, hk: usize,
11368 k2w: Option<(&CudaSlice<u8>, &CudaSlice<u8>, &mut CudaSlice<u8>)>)
11369 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11370 const D: usize = 128;
11371 let h = n_head;
11372 let nc = (t + c - 1) / c;
11373 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
11374 let mut gcum = self.uninit(t * h)?;
11375 let mut a = self.uninit(nc * h * c * c)?;
11376 let mut p = self.uninit(nc * h * c * c)?;
11377 let mut u = self.uninit(nc * h * c * D)?;
11378 let mut w = self.uninit(nc * h * c * D)?;
11379 { let f = self.func("gdn_chunk_cumgate_f32");
11381 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11382 let __s_b = self.gpu.stream();
11383 let mut b = __s_b.launch_builder(&f);
11384 b.arg(g).arg(&mut gcum).arg(&hi).arg(&ti).arg(&ci);
11385 unsafe { b.launch(cfg)?; }
11386 }
11387 if let Some((qb, kb, pb)) = k2w {
11388 assert!(c == 32, "gdn_k2_wgmma is a C==32 tile");
11391 let f = self.func("gdn_k2_wgmma");
11392 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11393 let hki = hk as i32;
11394 let __s_b = self.gpu.stream();
11395 let mut b = __s_b.launch_builder(&f);
11396 b.arg(qb).arg(kb).arg(&gcum).arg(beta).arg(&mut a).arg(&mut *pb).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11397 unsafe { b.launch(cfg)?; }
11398 } else if c <= 64 && !portable_mma_gated() { let f = self.func("gdn_chunk_attn_f32");
11400 let jt = ((c + 31) / 32) as u32;
11401 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11402 let hki = hk as i32;
11403 let __s_b = self.gpu.stream();
11404 let mut b = __s_b.launch_builder(&f);
11405 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11406 unsafe { b.launch(cfg)?; }
11407 } else { assert!(hk == h, "generic K2 is broadcast-only (de-broadcast rides C==32)");
11409 let f = self.func("gdn_chunk_attn_g_f32");
11410 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 8, 1), shared_mem_bytes: 0 };
11411 let __s_b = self.gpu.stream();
11412 let mut b = __s_b.launch_builder(&f);
11413 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci);
11414 unsafe { b.launch(cfg)?; }
11415 }
11416 { let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11418 match c {
11419 32 | 64 => {
11420 let f = self.func(if c == 32 { "gdn_chunk_solve32_f32" } else { "gdn_chunk_solve64_f32" });
11421 let wb: u64 = match wb16 { Some(d) => self.addr_u8(d), None => 0 };
11423 let hki = hk as i32;
11424 let __s_b = self.gpu.stream();
11425 let mut b = __s_b.launch_builder(&f);
11426 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&wb).arg(&hi).arg(&ti).arg(&hki);
11427 unsafe { b.launch(cfg)?; }
11428 }
11429 _ => {
11430 assert!(hk == h, "generic K3 is broadcast-only");
11431 let f = self.func("gdn_chunk_solve_f32");
11432 let __s_b = self.gpu.stream();
11433 let mut b = __s_b.launch_builder(&f);
11434 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&hi).arg(&ti).arg(&ci);
11435 unsafe { b.launch(cfg)?; }
11436 }
11437 }
11438 }
11439 Ok((gcum, p, u, w))
11440 }
11441
11442 pub fn gdn_db_on() -> bool {
11446 std::env::var("MEMRA_GDN_DB").as_deref() != Ok("0")
11447 }
11448
11449 pub fn gdn_mma_enabled(&self, c: usize) -> bool {
11452 !portable_mma_gated() && c == 32
11453 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11454 Ok("1") => true,
11455 Ok("0") => false,
11456 _ => cfg!(memra_hopper_mma),
11457 }
11458 }
11459
11460 pub fn gdn_wgmma_on(&self, c: usize) -> bool {
11463 self.gdn_mma_enabled(c)
11464 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
11465 Ok("0") => false,
11466 Ok("1") => true,
11467 _ => cfg!(memra_hopper_mma),
11468 }
11469 }
11470
11471 #[allow(clippy::too_many_arguments)]
11476 pub fn ssm_conv1d_gdn_state_pad(&self, qkv_tm: &cudarc::driver::CudaView<f32>,
11477 conv_state: &mut CudaSlice<f32>, w: &CudaSlice<f32>,
11478 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>,
11479 v_g: &mut CudaSlice<f32>,
11480 conv_dim: usize, t: usize, d_conv: usize,
11481 d_state: usize, num_v: usize, num_k: usize, key_dim: usize,
11482 hk: usize,
11483 pad_len: Option<&CudaSlice<i32>>)
11484 -> Result<(), Box<dyn std::error::Error>> {
11485 assert!(t >= d_conv - 1, "fused state conv requires T >= pad (PRIME_MIN_T gates)");
11486 {
11487 let f = self.func("ssm_conv1d_gdn_state_f32");
11488 let cfg = LaunchConfig {
11489 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11490 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11491 };
11492 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11493 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);
11494 let __s_b = self.gpu.stream();
11495 let mut b = __s_b.launch_builder(&f);
11496 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(q_g).arg(k_g).arg(v_g)
11497 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&hki);
11498 unsafe { b.launch(cfg)?; }
11499 }
11500 match pad_len {
11501 Some(len_d) => {
11502 let f = self.func("ssm_conv_ring_update_dev_f32");
11503 let n = conv_dim * (d_conv - 1);
11504 let cfg = LaunchConfig::for_num_elems(n as u32);
11505 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11506 let __s_b = self.gpu.stream();
11507 let mut b = __s_b.launch_builder(&f);
11508 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11509 unsafe { b.launch(cfg)?; }
11510 }
11511 None => {
11512 let f = self.func("ssm_conv_ring_update_f32");
11513 let n = conv_dim * (d_conv - 1);
11514 let cfg = LaunchConfig::for_num_elems(n as u32);
11515 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11516 let __s_b = self.gpu.stream();
11517 let mut b = __s_b.launch_builder(&f);
11518 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11519 unsafe { b.launch(cfg)?; }
11520 }
11521 }
11522 Ok(())
11523 }
11524
11525 pub fn gdn_chunk_alloc(&self, n_head: usize, t: usize, c: usize, hk: usize)
11529 -> Result<GdnChunkBufs, Box<dyn std::error::Error>> {
11530 const D: usize = 128;
11531 assert!(c == 32, "gdn_chunk_alloc: varlen chain is the C==32 mma pair");
11532 let h = n_head;
11533 let nc = (t + c - 1) / c;
11534 Ok(GdnChunkBufs {
11535 gcum: self.uninit(t * h)?,
11536 a: self.uninit(nc * h * c * c)?,
11537 p: self.uninit(nc * h * c * c)?,
11538 u: self.uninit(nc * h * c * D)?,
11539 w: self.uninit(nc * h * c * D)?,
11540 kb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11541 wb16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11542 y16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11543 ssnap16: self.alloc_u8_uninit(nc * h * D * D * 2)?,
11544 qb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11545 pb16: self.alloc_u8_uninit(nc * h * c * c * 2)?,
11546 o: self.uninit(D * h * t)?,
11547 t, nc,
11548 })
11549 }
11550
11551 pub fn f32_to_bf16_v(&self, x: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<u8>, n: usize)
11553 -> Result<(), Box<dyn std::error::Error>> {
11554 let f = self.func("f32_to_bf16_bulk");
11555 let ni = n as i64;
11556 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11557 let __s_b = self.gpu.stream();
11558 let mut b = __s_b.launch_builder(&f);
11559 b.arg(x).arg(dst).arg(&ni);
11560 unsafe { b.launch(cfg)?; }
11561 Ok(())
11562 }
11563
11564 pub fn f32_to_bf16_into(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<u8>, n: usize)
11566 -> Result<(), Box<dyn std::error::Error>> {
11567 let f = self.func("f32_to_bf16_bulk");
11568 let ni = n as i64;
11569 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11570 let __s_b = self.gpu.stream();
11571 let mut b = __s_b.launch_builder(&f);
11572 b.arg(x).arg(dst).arg(&ni);
11573 unsafe { b.launch(cfg)?; }
11574 Ok(())
11575 }
11576
11577 pub fn gdn_chunk_k123_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, hk: usize,
11580 wq: Option<&GdnWVl8>)
11581 -> Result<(), Box<dyn std::error::Error>> {
11582 let b = seqs.len();
11583 assert!(b >= 1 && b <= 8, "gdn_chunk_k123_vl8: 1..=8 sequences");
11584 let mut packed = [GdnSeqVl::default(); 8];
11585 packed[..b].copy_from_slice(seqs);
11586 let v = GdnVl8(packed);
11587 let (hi, ci) = (n_head as i32, 32i32);
11588 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11589 {
11590 let f = self.func("gdn_chunk_cumgate_vl");
11591 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11592 let __s_lb = self.gpu.stream();
11593 let mut lb = __s_lb.launch_builder(&f);
11594 lb.arg(&v).arg(&hi).arg(&ci);
11595 unsafe { lb.launch(cfg)?; }
11596 }
11597 let hki = hk as i32;
11598 if let Some(w) = wq { let f = self.func("gdn_k2_wgmma_vl");
11600 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11601 let __s_lb = self.gpu.stream();
11602 let mut lb = __s_lb.launch_builder(&f);
11603 lb.arg(&v).arg(w).arg(&hi).arg(&ci).arg(&hki);
11604 unsafe { lb.launch(cfg)?; }
11605 } else {
11606 let f = self.func("gdn_chunk_attn_vl");
11607 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11608 let __s_lb = self.gpu.stream();
11609 let mut lb = __s_lb.launch_builder(&f);
11610 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11611 unsafe { lb.launch(cfg)?; }
11612 }
11613 {
11614 let f = self.func("gdn_chunk_solve32_vl");
11615 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11616 let __s_lb = self.gpu.stream();
11617 let mut lb = __s_lb.launch_builder(&f);
11618 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11619 unsafe { lb.launch(cfg)?; }
11620 }
11621 Ok(())
11622 }
11623
11624 #[allow(clippy::too_many_arguments)]
11628 pub fn gdn_prep_vl8(&self, seqs: &[GdnPrepVl], conv_w: &CudaSlice<f32>,
11629 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11630 conv_dim: usize, d_conv: usize, d_state: usize,
11631 num_v: usize, num_k: usize, key_dim: usize, hk: usize, eps: f32)
11632 -> Result<(), Box<dyn std::error::Error>> {
11633 let b = seqs.len();
11634 assert!(b >= 1 && b <= 8);
11635 let mut packed = [GdnPrepVl::default(); 8];
11636 packed[..b].copy_from_slice(seqs);
11637 let v = GdnPrepVl8(packed);
11638 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11639 let (cdi, dci) = (conv_dim as i32, d_conv as i32);
11640 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
11641 assert!(conv_fuse || hk == num_v, "de-broadcast requires the fused conv");
11642 if conv_fuse {
11643 let f = self.func("ssm_conv1d_gdn_state_vl");
11644 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 };
11645 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);
11646 let __s_lb = self.gpu.stream();
11647 let mut lb = __s_lb.launch_builder(&f);
11648 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi).arg(&hki);
11649 unsafe { lb.launch(cfg)?; }
11650 } else {
11651 let f = self.func("ssm_conv1d_tm_state_vl");
11652 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 };
11653 let __s_lb = self.gpu.stream();
11654 let mut lb = __s_lb.launch_builder(&f);
11655 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci);
11656 unsafe { lb.launch(cfg)?; }
11657 }
11658 {
11659 let f = self.func("ssm_conv_ring_update_vl");
11660 let n = (conv_dim * (d_conv - 1)) as u32;
11661 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 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(&cdi).arg(&dci);
11665 unsafe { lb.launch(cfg)?; }
11666 }
11667 if !conv_fuse {
11668 let f = self.func("qkv_to_gdn_repack_vl");
11669 let n = max_t * (num_v * d_state) as u32;
11670 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11671 let (dsi, nvi, nki, kdi) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11672 let __s_lb = self.gpu.stream();
11673 let mut lb = __s_lb.launch_builder(&f);
11674 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi);
11675 unsafe { lb.launch(cfg)?; }
11676 }
11677 if Self::l2_v2_on(d_state) {
11678 let f = self.func("gdn_l2_v2_vl");
11679 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 };
11680 let (dsi, nvi) = (d_state as i32, hk as i32);
11681 let __s_lb = self.gpu.stream();
11682 let mut lb = __s_lb.launch_builder(&f);
11683 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11684 unsafe { lb.launch(cfg)?; }
11685 } else {
11686 let f = self.func("gdn_l2_vl");
11687 let cfg = LaunchConfig { grid_dim: (max_t * hk as u32, 2, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11688 let (dsi, nvi) = (d_state as i32, hk as i32);
11689 let __s_lb = self.gpu.stream();
11690 let mut lb = __s_lb.launch_builder(&f);
11691 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11692 unsafe { lb.launch(cfg)?; }
11693 }
11694 {
11695 let f = self.func("gdn_gate_prep_vl");
11696 let n = max_t * num_v as u32;
11697 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11698 let nvi = num_v as i32;
11699 let __s_lb = self.gpu.stream();
11700 let mut lb = __s_lb.launch_builder(&f);
11701 lb.arg(&v).arg(dt_bias).arg(a).arg(&nvi);
11702 unsafe { lb.launch(cfg)?; }
11703 }
11704 Ok(())
11705 }
11706
11707 pub fn gdn_mirror_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, which: i32, hk: usize)
11709 -> Result<(), Box<dyn std::error::Error>> {
11710 let b = seqs.len();
11711 assert!(b >= 1 && b <= 8);
11712 let mut packed = [GdnSeqVl::default(); 8];
11713 packed[..b].copy_from_slice(seqs);
11714 let v = GdnVl8(packed);
11715 let ept = (if which == 0 { hk } else { n_head } * 128) as i32;
11716 let max_n = seqs.iter().map(|s| if which == 0 { s.t as i64 * ept as i64 }
11717 else { s.nc as i64 * ept as i64 * 32 }).max().unwrap();
11718 let f = self.func("gdn_mirror_vl");
11719 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
11720 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11721 let __s_lb = self.gpu.stream();
11722 let mut lb = __s_lb.launch_builder(&f);
11723 lb.arg(&v).arg(&ept).arg(&which);
11724 unsafe { lb.launch(cfg)?; }
11725 Ok(())
11726 }
11727
11728 pub fn gdn_tail_vl8(&self, seqs: &[GdnPrepVl], norm_w: &CudaSlice<f32>,
11730 d_state: usize, num_v: usize, eps: f32)
11731 -> Result<(), Box<dyn std::error::Error>> {
11732 let b = seqs.len();
11733 assert!(b >= 1 && b <= 8);
11734 let mut packed = [GdnPrepVl::default(); 8];
11735 packed[..b].copy_from_slice(seqs);
11736 let v = GdnPrepVl8(packed);
11737 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11738 let f = self.func("gated_rmsnorm_f16out_vl");
11739 let cfg = LaunchConfig { grid_dim: (max_t * num_v as u32, 1, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11741 let (dsi, nvi) = (d_state as i32, num_v as i32);
11742 let __s_lb = self.gpu.stream();
11743 let mut lb = __s_lb.launch_builder(&f);
11744 lb.arg(&v).arg(norm_w).arg(&dsi).arg(&nvi).arg(&eps);
11745 unsafe { lb.launch(cfg)?; }
11746 Ok(())
11747 }
11748
11749 pub fn addr_f32(&self, x: &CudaSlice<f32>) -> u64 {
11752 use cudarc::driver::DevicePtr;
11753 let s = self.gpu.stream();
11754 let (p, _g) = x.device_ptr(&s);
11755 p as u64
11756 }
11757 pub fn addr_f32_mut(&self, x: &mut CudaSlice<f32>) -> u64 {
11758 use cudarc::driver::DevicePtrMut;
11759 let s = self.gpu.stream();
11760 let (p, _g) = x.device_ptr_mut(&s);
11761 p as u64
11762 }
11763 pub fn addr_f32v(&self, x: &cudarc::driver::CudaView<f32>) -> u64 {
11764 use cudarc::driver::DevicePtr;
11765 let s = self.gpu.stream();
11766 let (p, _g) = x.device_ptr(&s);
11767 p as u64
11768 }
11769 pub fn addr_u8(&self, x: &CudaSlice<u8>) -> u64 {
11770 use cudarc::driver::DevicePtr;
11771 let s = self.gpu.stream();
11772 let (p, _g) = x.device_ptr(&s);
11773 p as u64
11774 }
11775
11776 pub fn gdn_chunk_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, scale: f32, hk: usize,
11780 wq: Option<&GdnWVl8>)
11781 -> Result<(), Box<dyn std::error::Error>> {
11782 const NSPLIT: u32 = 4;
11783 let b = seqs.len();
11784 assert!(b >= 1 && b <= 8, "gdn_chunk_vl8: 1..=8 sequences");
11785 let mut packed = [GdnSeqVl::default(); 8];
11786 packed[..b].copy_from_slice(seqs);
11787 let v = GdnVl8(packed);
11788 let (hi, ci) = (n_head as i32, 32i32);
11789 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11790 let hki = hk as i32;
11791 if let Some(w) = wq {
11792 let f = self.func("gdn_k45_wgmma_vl");
11794 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11795 let __s_lb = self.gpu.stream();
11796 let mut lb = __s_lb.launch_builder(&f);
11797 lb.arg(&v).arg(w).arg(&scale).arg(&hi).arg(&ci).arg(&hki);
11798 unsafe { lb.launch(cfg)?; }
11799 let _ = max_nc;
11800 return Ok(());
11801 }
11802 {
11803 let f = self.func("gdn_chunk_state_mma_vl");
11804 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11805 let __s_lb = self.gpu.stream();
11806 let mut lb = __s_lb.launch_builder(&f);
11807 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11808 unsafe { lb.launch(cfg)?; }
11809 }
11810 {
11811 let f = self.func("gdn_chunk_output_mma_vl");
11812 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11813 let __s_lb = self.gpu.stream();
11814 let mut lb = __s_lb.launch_builder(&f);
11815 lb.arg(&v).arg(&hi).arg(&ci).arg(&scale).arg(&hki);
11816 unsafe { lb.launch(cfg)?; }
11817 }
11818 Ok(())
11819 }
11820 pub fn gdn_scan_chunked(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11821 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
11822 qb16_pre: Option<&CudaSlice<u8>>,
11823 state_in: &CudaSlice<f32>,
11824 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11825 n_head: usize, t: usize, scale: f32, c: usize, hk: usize)
11826 -> Result<(), Box<dyn std::error::Error>> {
11827 const D: usize = 128;
11828 const NSPLIT: u32 = 4;
11829 assert!(c >= 1 && c <= 128, "gdn_scan_chunked: C must be in 1..=128");
11830 let h = n_head;
11831 let nc = (t + c - 1) / c;
11832 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
11833 let gdn_mma_pre = !portable_mma_gated() && c == 32
11837 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11838 Ok("1") => true,
11839 Ok("0") => false,
11840 _ => cfg!(memra_hopper_mma),
11841 };
11842 let mut wb16_pre: Option<CudaSlice<u8>> = if gdn_mma_pre {
11843 Some(self.alloc_u8_uninit(nc * h * c * D * 2)?)
11844 } else { None };
11845 let gdn_wgmma_pre = gdn_mma_pre
11849 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
11850 Ok("0") => false,
11851 Ok("1") => true,
11852 _ => cfg!(memra_hopper_mma),
11853 };
11854 let nk = t * hk * D;
11855 let mut kb16_local: Option<CudaSlice<u8>> = None;
11856 if gdn_mma_pre && kb16_pre.is_none() {
11857 let mut kb = self.alloc_u8_uninit(nk * 2)?;
11858 let f = self.func("f32_to_bf16_bulk");
11859 let n2 = nk as i64;
11860 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
11861 let __s_b = self.gpu.stream();
11862 let mut b = __s_b.launch_builder(&f);
11863 b.arg(k).arg(&mut kb).arg(&n2);
11864 unsafe { b.launch(cfg2)?; }
11865 kb16_local = Some(kb);
11866 }
11867 let kb16_ref0: Option<&CudaSlice<u8>> = kb16_local.as_ref().or(kb16_pre);
11868 if let Some(kb) = kb16_pre { assert!(kb.len() >= nk * 2, "kb16_pre too small"); }
11869 let mut qb16: Option<CudaSlice<u8>> = None;
11870 let mut pb16: Option<CudaSlice<u8>> = None;
11871 if gdn_wgmma_pre {
11872 if qb16_pre.is_none() {
11875 let mut qb = self.alloc_u8_uninit(nk * 2)?;
11876 let f = self.func("f32_to_bf16_bulk");
11877 let n2 = nk as i64;
11878 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
11879 let __s_b = self.gpu.stream();
11880 let mut b = __s_b.launch_builder(&f);
11881 b.arg(q).arg(&mut qb).arg(&n2);
11882 unsafe { b.launch(cfg2)?; }
11883 qb16 = Some(qb);
11884 } else if let Some(qb) = qb16_pre {
11885 assert!(qb.len() >= nk * 2, "qb16_pre too small");
11886 }
11887 pb16 = Some(self.alloc_u8_uninit(nc * h * c * c * 2)?);
11888 }
11889 let qb16_ref0: Option<&CudaSlice<u8>> = qb16.as_ref().or(qb16_pre);
11890 let k2w = if gdn_wgmma_pre {
11891 Some((*qb16_ref0.as_ref().unwrap(),
11892 *kb16_ref0.as_ref().unwrap(),
11893 pb16.as_mut().unwrap()))
11894 } else { None };
11895 let (gcum, p, u, w) = self.gdn_chunk_k123(q, k, v, g, beta, wb16_pre.as_mut(), n_head, t, c, hk, k2w)?;
11896 let _ = &w;
11897 let mut y = self.uninit(nc * h * c * D)?;
11898 let mut ssnap = self.uninit(nc * h * D * D)?; let gdn_mma = !portable_mma_gated() && c == 32
11910 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11911 Ok("1") => true,
11912 Ok("0") => false,
11913 _ => cfg!(memra_hopper_mma),
11914 };
11915 if gdn_mma {
11916 let wb16 = wb16_pre.take().expect("mma path pre-allocates wb16 (K3 store fold)");
11917 let kb16_ref: &CudaSlice<u8> = kb16_ref0.expect("mma path pre-builds kb16 above K123");
11918 if gdn_wgmma_pre {
11930 let qb16 = qb16_ref0.unwrap();
11932 let pb16 = pb16.as_ref().unwrap();
11933 {
11934 let f = self.func("gdn_k45_wgmma");
11935 let cfg = LaunchConfig { grid_dim: (h as u32, 4, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11936 let hki = hk as i32;
11937 let __s_b = self.gpu.stream();
11938 let mut b = __s_b.launch_builder(&f);
11939 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(qb16).arg(pb16)
11940 .arg(o).arg(&scale).arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11941 unsafe { b.launch(cfg)?; }
11942 }
11943 return Ok(());
11944 }
11945 let mut y16 = self.alloc_u8_uninit(nc * h * c * D * 2)?;
11949 let mut ssnap16 = self.alloc_u8_uninit(nc * h * D * D * 2)?;
11950 {
11951 let f = self.func("gdn_chunk_state_mma");
11952 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11953 let hki = hk as i32;
11954 let __s_b = self.gpu.stream();
11955 let mut b = __s_b.launch_builder(&f);
11956 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(&mut y16).arg(&mut ssnap16)
11957 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11958 unsafe { b.launch(cfg)?; }
11959 }
11960 { let f = self.func("gdn_chunk_output_mma");
11962 let jt = ((c + 31) / 32) as u32;
11963 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11964 let hki = hk as i32;
11965 let __s_b = self.gpu.stream();
11966 let mut b = __s_b.launch_builder(&f);
11967 b.arg(q).arg(&gcum).arg(&p).arg(&y16).arg(&ssnap16).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale).arg(&hki);
11968 unsafe { b.launch(cfg)?; }
11969 }
11970 return Ok(());
11971 }
11972 { let f = self.func("gdn_chunk_state_f32");
11974 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11975 let __s_b = self.gpu.stream();
11976 let mut b = __s_b.launch_builder(&f);
11977 b.arg(k).arg(&gcum).arg(beta).arg(&u).arg(&w).arg(&mut y).arg(&mut ssnap)
11978 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci);
11979 unsafe { b.launch(cfg)?; }
11980 }
11981 { let f = self.func("gdn_chunk_output_f32");
11983 let jt = ((c + 31) / 32) as u32;
11984 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11985 let __s_b = self.gpu.stream();
11986 let mut b = __s_b.launch_builder(&f);
11987 b.arg(q).arg(&gcum).arg(&p).arg(&y).arg(&ssnap).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale);
11988 unsafe { b.launch(cfg)?; }
11989 }
11990 Ok(())
11991 }
11992
11993 #[allow(clippy::too_many_arguments)]
12002 #[allow(clippy::too_many_arguments)]
12003 pub fn gdn_scan_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12004 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
12005 qb16_pre: Option<&CudaSlice<u8>>,
12006 state_in: &CudaSlice<f32>,
12007 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12008 n_head: usize, t: usize, scale: f32, hk: usize)
12009 -> Result<(), Box<dyn std::error::Error>> {
12010 if std::env::var("MEMRA_GDN_DIFF").is_ok() && t >= 16 {
12011 assert!(hk == n_head, "GDN_DIFF oracle is broadcast-only");
12012 return self.gdn_scan_diff(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale);
12013 }
12014 if Self::gdn_chunked_enabled() && t >= 16 {
12015 self.gdn_scan_chunked(q, k, v, g, beta, kb16_pre, qb16_pre, state_in, state_out, o, n_head, t, scale,
12016 Self::gdn_chunk_size(), hk)
12017 } else {
12018 assert!(hk == n_head, "s128 scan is broadcast-only (prep guarantees by predicate)");
12019 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)
12020 }
12021 }
12022
12023 #[allow(clippy::too_many_arguments)]
12025 fn gdn_scan_diff(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12026 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
12027 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12028 n_head: usize, t: usize, scale: f32)
12029 -> Result<(), Box<dyn std::error::Error>> {
12030 static CALL: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
12031 let call = CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12032 let mut o_c = self.uninit(o.len())?;
12033 let mut st_c = self.uninit(state_out.len())?;
12034 self.gdn_scan_chunked(q, k, v, g, beta, None, None, state_in, &mut st_c, &mut o_c,
12035 n_head, t, scale, Self::gdn_chunk_size(), n_head)?;
12036 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)?;
12037 let (oh_s, oh_c) = (self.dtoh(o)?, self.dtoh(&o_c)?);
12038 let (sh_s, sh_c) = (self.dtoh(state_out)?, self.dtoh(&st_c)?);
12039 let stats = |a: &[f32], b: &[f32]| -> (f32, f32, f64) {
12040 let mut max_abs = 0f32; let mut max_rel = 0f32; let mut sum_rel = 0f64;
12041 for (x, y) in a.iter().zip(b) {
12042 let ad = (x - y).abs();
12043 let rel = ad / x.abs().max(y.abs()).max(1e-3);
12044 if ad > max_abs { max_abs = ad; }
12045 if rel > max_rel { max_rel = rel; }
12046 sum_rel += rel as f64;
12047 }
12048 (max_abs, max_rel, sum_rel / a.len() as f64)
12049 };
12050 let (o_ma, o_mr, o_mean) = stats(&oh_s, &oh_c);
12051 let (s_ma, s_mr, s_mean) = stats(&sh_s, &sh_c);
12052 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} | \
12053 state: max_abs={s_ma:.3e} max_rel={s_mr:.3e} mean_rel={s_mean:.3e}",
12054 Self::gdn_chunk_size());
12055 Ok(())
12056 }
12057
12058 pub fn gdn_glog(&self, alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
12060 g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
12061 -> Result<(), Box<dyn std::error::Error>> {
12062 let f = self.func("gdn_glog_f32");
12063 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
12064 let (h, ti) = (n_head as i32, t as i32);
12065 let __s_b = self.gpu.stream();
12066 let mut b = __s_b.launch_builder(&f);
12067 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
12068 unsafe { b.launch(cfg)?; }
12069 Ok(())
12070 }
12071
12072 pub fn sigmoid_v(&self, x: &cudarc::driver::CudaView<f32>, y: &mut CudaSlice<f32>, n: usize)
12075 -> Result<(), Box<dyn std::error::Error>> {
12076 let f = self.func("sigmoid_f32");
12077 let cfg = LaunchConfig::for_num_elems(n as u32);
12078 let ni = n as i32;
12079 let __s_b = self.gpu.stream();
12080 let mut b = __s_b.launch_builder(&f);
12081 b.arg(x).arg(y).arg(&ni);
12082 unsafe { b.launch(cfg)?; }
12083 Ok(())
12084 }
12085
12086 pub fn gdn_glog_v(&self, alpha: &cudarc::driver::CudaView<f32>, dt_bias: &CudaSlice<f32>,
12087 a: &CudaSlice<f32>, g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
12088 -> Result<(), Box<dyn std::error::Error>> {
12089 let f = self.func("gdn_glog_f32");
12090 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
12091 let (h, ti) = (n_head as i32, t as i32);
12092 let __s_b = self.gpu.stream();
12093 let mut b = __s_b.launch_builder(&f);
12094 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
12095 unsafe { b.launch(cfg)?; }
12096 Ok(())
12097 }
12098
12099 pub fn sigmoid(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize)
12100 -> Result<(), Box<dyn std::error::Error>> {
12101 let f = self.func("sigmoid_f32");
12102 let cfg = LaunchConfig::for_num_elems(n as u32);
12103 let ni = n as i32;
12104 let __s_b = self.gpu.stream();
12105 let mut b = __s_b.launch_builder(&f);
12106 b.arg(x).arg(y).arg(&ni);
12107 unsafe { b.launch(cfg)?; }
12108 Ok(())
12109 }
12110
12111 pub fn sig_mul_f16out(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12114 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
12115 -> Result<(), Box<dyn std::error::Error>> {
12116 let f = self.func("sig_mul_f16out_f32");
12117 let cfg = LaunchConfig::for_num_elems(n as u32);
12118 let ni = n as i32;
12119 let __s_b = self.gpu.stream();
12120 let mut b = __s_b.launch_builder(&f);
12121 b.arg(a).arg(g).arg(dst).arg(dst16).arg(&ni);
12122 unsafe { b.launch(cfg)?; }
12123 Ok(())
12124 }
12125
12126 #[allow(clippy::too_many_arguments)]
12135 pub fn attn_head_gate(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12136 dst: &mut CudaSlice<f32>, dst16: Option<&mut CudaSlice<u8>>,
12137 head_dim: usize, n_head: usize, t: usize)
12138 -> Result<(), Box<dyn std::error::Error>> {
12139 let f = self.func("attn_head_gate_f32");
12140 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12141 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12142 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
12144 let __s_b = self.gpu.stream();
12145 let mut b = __s_b.launch_builder(&f);
12146 b.arg(a).arg(g).arg(dst).arg(&d16).arg(&hd).arg(&nh).arg(&ti);
12147 unsafe { b.launch(cfg)?; }
12148 Ok(())
12149 }
12150
12151 #[allow(clippy::too_many_arguments)]
12160 pub fn swiglu_clamped_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
12161 gs: f32, us: f32, limit: f32,
12162 dst: &mut CudaSlice<f32>, n: usize)
12163 -> Result<(), Box<dyn std::error::Error>> {
12164 debug_assert!(limit > 1e-6, "swiglu_clamped needs a live limit; use silu_mul_scaled");
12165 let f = self.func("swiglu_clamped_mul_scaled_f32");
12166 let cfg = LaunchConfig::for_num_elems(n as u32);
12167 let ni = n as i32;
12168 let __s_b = self.gpu.stream();
12169 let mut b = __s_b.launch_builder(&f);
12170 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&limit).arg(dst).arg(&ni);
12171 unsafe { b.launch(cfg)?; }
12172 Ok(())
12173 }
12174
12175 pub fn gated_rmsnorm(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12177 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12178 -> Result<(), Box<dyn std::error::Error>> {
12179 let f = self.func("gated_rmsnorm_f32");
12180 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12181 let (nc, e) = (ncols as i32, eps);
12182 let __s_b = self.gpu.stream();
12183 let mut b = __s_b.launch_builder(&f);
12184 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12185 unsafe { b.launch(cfg)?; }
12186 Ok(())
12187 }
12188
12189 pub fn gated_rmsnorm_f16out(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12192 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12193 ncols: usize, nrows: usize, eps: f32)
12194 -> Result<(), Box<dyn std::error::Error>> {
12195 let f = self.func("gated_rmsnorm_f16out_f32");
12196 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12198 let (nc, e) = (ncols as i32, eps);
12199 let __s_b = self.gpu.stream();
12200 let mut b = __s_b.launch_builder(&f);
12201 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12202 unsafe { b.launch(cfg)?; }
12203 Ok(())
12204 }
12205
12206 #[allow(clippy::too_many_arguments)]
12210 pub fn add_rms_norm_zq8(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
12211 res: &mut CudaSlice<f32>, z: &mut CudaSlice<f32>,
12212 ncols: usize, nrows: usize, eps: f32)
12213 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12214 assert!(ncols % 32 == 0);
12215 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
12216 let mut d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12217 let f = self.func("add_rms_norm_zq8");
12218 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
12219 let (nc, ep) = (ncols as i32, eps);
12220 let __s_b = self.gpu.stream();
12221 let mut b = __s_b.launch_builder(&f);
12222 b.arg(a).arg(b_in).arg(w).arg(res).arg(z).arg(&mut q).arg(&mut d).arg(&nc).arg(&ep);
12223 unsafe { b.launch(cfg)?; }
12224 Ok((q, d))
12225 }
12226
12227 pub fn gated_rmsnorm_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12232 z: &cudarc::driver::CudaView<f32>,
12233 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12234 -> Result<(), Box<dyn std::error::Error>> {
12235 let f = self.func("gated_rmsnorm_f32");
12236 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12237 let (nc, e) = (ncols as i32, eps);
12238 let __s_b = self.gpu.stream();
12239 let mut b = __s_b.launch_builder(&f);
12240 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12241 unsafe { b.launch(cfg)?; }
12242 Ok(())
12243 }
12244
12245 pub fn gated_rmsnorm_f16out_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12246 z: &cudarc::driver::CudaView<f32>,
12247 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12248 ncols: usize, nrows: usize, eps: f32)
12249 -> Result<(), Box<dyn std::error::Error>> {
12250 let f = self.func("gated_rmsnorm_f16out_f32");
12251 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12253 let (nc, e) = (ncols as i32, eps);
12254 let __s_b = self.gpu.stream();
12255 let mut b = __s_b.launch_builder(&f);
12256 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12257 unsafe { b.launch(cfg)?; }
12258 Ok(())
12259 }
12260
12261 pub fn gated_rmsnorm_q8_1(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12262 ncols: usize, nrows: usize, eps: f32)
12263 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12264 assert!(ncols % 32 == 0);
12265 let f = self.func("gated_rmsnorm_q8_1");
12266 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12267 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12268 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12269 let (nc, ep) = (ncols as i32, eps);
12270 let __s_b = self.gpu.stream();
12271 let mut b = __s_b.launch_builder(&f);
12272 b.arg(o).arg(w).arg(z).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
12273 unsafe { b.launch(cfg)?; }
12274 Ok((out_q, out_d))
12275 }
12276
12277 pub fn transpose(&self, inp: &CudaSlice<f32>, rows: usize, cols: usize)
12279 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12280 let f = self.func("transpose_f32");
12281 let mut out = self.zeros(rows * cols)?;
12282 let cfg = LaunchConfig::for_num_elems((rows * cols) as u32);
12283 let (r, c) = (rows as i32, cols as i32);
12284 let __s_b = self.gpu.stream();
12285 let mut b = __s_b.launch_builder(&f);
12286 b.arg(inp).arg(&mut out).arg(&r).arg(&c);
12287 unsafe { b.launch(cfg)?; }
12288 Ok(out)
12289 }
12290
12291 pub fn repeat_heads(&self, inp: &CudaSlice<f32>, out: &mut CudaSlice<f32>,
12293 head_dim: usize, n_in: usize, n_out: usize, t: usize)
12294 -> Result<(), Box<dyn std::error::Error>> {
12295 let f = self.func("repeat_heads_f32");
12296 let cfg = LaunchConfig::for_num_elems((head_dim * n_out * t) as u32);
12297 let (hd, ni, no, ti) = (head_dim as i32, n_in as i32, n_out as i32, t as i32);
12298 let __s_b = self.gpu.stream();
12299 let mut b = __s_b.launch_builder(&f);
12300 b.arg(inp).arg(out).arg(&hd).arg(&ni).arg(&no).arg(&ti);
12301 unsafe { b.launch(cfg)?; }
12302 Ok(())
12303 }
12304
12305 pub fn q_gate_split(&self, qf: &CudaSlice<f32>, q_out: &mut CudaSlice<f32>,
12308 gate_out: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, t: usize)
12309 -> Result<(), Box<dyn std::error::Error>> {
12310 let f = self.func("q_gate_split_f32");
12311 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12312 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12313 let __s_b = self.gpu.stream();
12314 let mut b = __s_b.launch_builder(&f);
12315 b.arg(qf).arg(q_out).arg(gate_out).arg(&hd).arg(&nh).arg(&ti);
12316 unsafe { b.launch(cfg)?; }
12317 Ok(())
12318 }
12319
12320 pub fn qkv_to_gdn_repack(&self, conv_out: &CudaSlice<f32>, q_g: &mut CudaSlice<f32>,
12324 k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
12325 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, t: usize)
12326 -> Result<(), Box<dyn std::error::Error>> {
12327 let f = self.func("qkv_to_gdn_repack_f32");
12328 let cfg = LaunchConfig::for_num_elems((d_state * num_v * t) as u32);
12329 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);
12330 let __s_b = self.gpu.stream();
12331 let mut b = __s_b.launch_builder(&f);
12332 b.arg(conv_out).arg(q_g).arg(k_g).arg(v_g).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&ti);
12333 unsafe { b.launch(cfg)?; }
12334 Ok(())
12335 }
12336
12337 pub fn conv_left_pad(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
12340 conv_dim: usize, t: usize, pad: usize)
12341 -> Result<(), Box<dyn std::error::Error>> {
12342 let f = self.func("conv_left_pad_f32");
12343 let cfg = LaunchConfig::for_num_elems((conv_dim * t) as u32);
12344 let (cd, ti, p) = (conv_dim as i32, t as i32, pad as i32);
12345 let __s_b = self.gpu.stream();
12346 let mut b = __s_b.launch_builder(&f);
12347 b.arg(src).arg(dst).arg(&cd).arg(&ti).arg(&p);
12348 unsafe { b.launch(cfg)?; }
12349 Ok(())
12350 }
12351
12352 pub fn conv_assemble_and_roll(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12356 conv_in: &mut CudaSlice<f32>, conv_dim: usize, pad: usize)
12357 -> Result<(), Box<dyn std::error::Error>> {
12358 let f = self.func("conv_assemble_and_roll_f32");
12359 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12360 let (cd, p) = (conv_dim as i32, pad as i32);
12361 let __s_b = self.gpu.stream();
12362 let mut b = __s_b.launch_builder(&f);
12363 b.arg(qkv_col).arg(conv_state).arg(conv_in).arg(&cd).arg(&p);
12364 unsafe { b.launch(cfg)?; }
12365 Ok(())
12366 }
12367
12368 pub fn ssm_conv1d_fused_decode(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12374 w: &CudaSlice<f32>, conv_out: &mut CudaSlice<f32>,
12375 conv_dim: usize, d_conv: usize)
12376 -> Result<(), Box<dyn std::error::Error>> {
12377 let f = self.func("ssm_conv1d_fused_decode_f32");
12378 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12379 let (cd, dc) = (conv_dim as i32, d_conv as i32);
12380 let __s_b = self.gpu.stream();
12381 let mut b = __s_b.launch_builder(&f);
12382 b.arg(qkv_col).arg(conv_state).arg(w).arg(conv_out).arg(&cd).arg(&dc);
12383 unsafe { b.launch(cfg)?; }
12384 Ok(())
12385 }
12386
12387 pub fn slice_range(&self, src: &CudaSlice<f32>, start: usize, len: usize)
12390 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12391 let host = self.gpu.stream().clone_dtoh(src)?;
12392 self.gpu.stream().synchronize()?;
12393 Ok(self.htod(&host[start..start + len])?)
12394 }
12395}
12396
12397#[cfg(test)]
12398mod target_dispatch_tests {
12399 use super::legacy_quant_gemm_allowed;
12400
12401 #[test]
12402 fn legacy_quant_gemm_arch_policy_honors_the_escape_hatch() {
12403 assert!(legacy_quant_gemm_allowed(false, false, false));
12405 assert!(!legacy_quant_gemm_allowed(false, false, true));
12406 assert!(!legacy_quant_gemm_allowed(true, false, false));
12408 assert!(!legacy_quant_gemm_allowed(true, false, true));
12409 assert!(legacy_quant_gemm_allowed(true, true, false));
12411 assert!(!legacy_quant_gemm_allowed(true, true, true));
12412 }
12413
12414 #[cfg(all(memra_portable_cuda, not(memra_hopper_mma)))]
12415 #[test]
12416 fn portable_build_disables_legacy_quant_gemm_without_an_env_override() {
12417 assert!(!legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12418 }
12419
12420 #[cfg(memra_hopper_mma)]
12421 #[test]
12422 fn hopper_mma_build_re_admits_legacy_quant_gemm() {
12423 assert!(legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12424 assert!(super::portable_mma_gated() == false);
12425 }
12426}
12427
12428impl memra_kv::KvDev for Engine {
12431 fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12432 Engine::zeros(self, n)
12433 }
12434 fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12435 Engine::uninit(self, n)
12436 }
12437 fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
12438 Engine::alloc_u8(self, n)
12439 }
12440 fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
12441 Engine::htod_i32(self, v)
12442 }
12443 fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12444 Engine::clone_dtod(self, src)
12445 }
12446 fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
12447 -> Result<(), Box<dyn std::error::Error>> {
12448 Engine::copy_into(self, dst, off, src, len)
12449 }
12450 fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
12451 Engine::set_i32_one(self, d, v)
12452 }
12453}