1use cudarc::driver::sys::CUfunction_attribute_enum::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES;
4use cudarc::driver::{
5 CudaContext, CudaFunction, CudaModule, CudaSlice, CudaStream, DeviceSlice, LaunchConfig,
6 PushKernelArg,
7};
8use cudarc::nvrtc::Ptx;
9use std::sync::{Arc, Mutex};
10
11const GDN_K2_DYNAMIC_SHARED_BYTES: u32 = 67_072;
12
13const SDPA_NAIVE_SMEM_MAX: usize = 48 * 1024;
18
19const SDPA_NAIVE_GMEM_WS_MAX: usize = 1 << 30;
23
24#[cfg(debug_assertions)]
25pub(crate) fn debug_assert_tensor_stream_device<T>(
26 tensor: &CudaSlice<T>,
27 stream: &CudaStream,
28 site: &str,
29) {
30 let tensor_dev = tensor.ordinal();
31 let stream_dev = stream.context().ordinal();
32 assert_eq!(
33 tensor_dev, stream_dev,
34 "PP cross-device tensor read at {site}: tensor on dev{tensor_dev}, stream on dev{stream_dev}"
35 );
36}
37
38fn ensure_tensor_stream_device<T>(
39 tensor: &impl DeviceSlice<T>,
40 stream: &CudaStream,
41 site: &str,
42) -> Result<(), Box<dyn std::error::Error>> {
43 let tensor_dev = tensor.stream().context().ordinal();
44 let stream_dev = stream.context().ordinal();
45 if tensor_dev != stream_dev {
46 return Err(format!(
47 "PP cross-device tensor access at {site}: tensor on dev{tensor_dev}, \
48 stream on dev{stream_dev}"
49 )
50 .into());
51 }
52 Ok(())
53}
54
55pub use memra_gguf;
56pub use memra_runtime;
57
58pub mod forward;
59pub mod hybrid;
60pub mod hybrid_forward;
61pub mod hyper;
62pub mod model;
63pub mod sigrouter_contract;
64pub mod vision;
65pub mod vision_gemma;
66pub mod vision_glm5;
67pub mod vision_pre;
68pub mod vision_step;
69pub mod cache {
72 pub use memra_kv::*;
73}
74pub mod decode;
75pub mod decode_batch;
76pub mod dflash;
77pub mod eagle;
78pub mod ep_map;
83pub mod gemma_spec;
84pub mod glm5_tp;
85pub mod glm_spec;
89pub mod graph_update;
90pub mod kda;
91pub mod mla;
95pub mod mla_ffi;
96pub mod moesd;
97pub mod parallel;
98pub mod plan_backend;
99pub mod pp;
100pub mod qwen4exp_gpu;
103pub mod round_stream;
104pub mod spec;
105pub mod spec_phase;
110pub mod tp;
111pub mod tp_transport;
112pub use memra_sampling as sampler;
113
114pub fn moe_f16g_mode() -> u8 {
158 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
159 *M.get_or_init(|| match std::env::var("MEMRA_MOE_F16G").as_deref() {
160 Ok("0") => 0,
161 Ok("2") => 2,
162 Ok("3") => 3,
163 Ok(_) => 1,
164 Err(_) => 2,
167 })
168}
169pub fn moe_f16g_sk_params() -> (i32, i32) {
183 static P: std::sync::OnceLock<(i32, i32)> = std::sync::OnceLock::new();
184 *P.get_or_init(|| match std::env::var("MEMRA_F16G_SK").as_deref() {
185 Ok("0") => (-1, 0),
186 Ok("32") => (0, i32::MAX),
187 Ok("128") => (0, 1),
188 _ => {
189 let cross = std::env::var("MEMRA_F16G_SK_CROSS")
190 .ok()
191 .and_then(|v| v.parse().ok())
192 .unwrap_or(64);
193 (0, cross)
194 }
195 })
196}
197pub fn moe_f16g_direct_on(qtype: i32) -> bool {
208 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
209 let m = *M.get_or_init(|| match std::env::var("MEMRA_F16G_DIRECT").as_deref() {
210 Ok("0") => 0,
211 Ok("kq") => 1,
212 _ => 2,
213 });
214 match m {
215 0 => false,
216 1 => qtype == QT_Q4_K || qtype == QT_Q6_K,
217 _ => true,
218 }
219}
220pub fn moe_f16g_tail_on() -> bool {
229 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
230 *ON.get_or_init(|| std::env::var("MEMRA_F16G_TAIL").as_deref() != Ok("0"))
231}
232
233pub fn moe_f16g_gemma_on() -> bool {
240 static M: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
241 *M.get_or_init(|| !matches!(std::env::var("MEMRA_MOE_F16G").as_deref(), Ok("0") | Err(_)))
242}
243
244pub fn moe_fuse_actq_on() -> bool {
248 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
249 *ON.get_or_init(|| std::env::var("MEMRA_MOE_FUSE_ACTQ").as_deref() != Ok("0"))
250}
251
252pub fn router_prefill_exact_on() -> bool {
262 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
263 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_PREFILL_EXACT").as_deref() != Ok("0"))
264}
265
266pub fn router_kernel_on() -> bool {
267 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
268 *ON.get_or_init(|| {
269 let on = std::env::var("MEMRA_ROUTER_KERNEL").as_deref() != Ok("0");
270 if !on {
271 eprintln!("[memra] router kernel OFF (rollback: per-column cuBLAS gemv)");
272 }
273 on
274 })
275}
276
277pub const ROUTER_BATCH_MIN_T: usize = 8;
292pub fn router_batch_on() -> bool {
293 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
294 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_BATCH").as_deref() != Ok("0"))
295}
296mod cpu_experts;
297#[cfg(memra_cutlass)]
298pub mod cutlass_ffi;
299pub mod dsv4_ffi;
300pub mod dsv4_gpu;
301pub mod f16_ffi;
302pub mod fp8_ffi;
303pub mod mmq_ffi;
304pub mod moe_cache;
305pub mod prime_graph;
306pub mod spill;
307mod spill_pread;
308
309const FATBIN: &[u8] = include_bytes!(env!("MEMRA_ENGINE_FATBIN"));
316const HYBRID_FATBIN: &[u8] = include_bytes!(env!("MEMRA_HYBRID_FATBIN"));
317const KDA_FATBIN: &[u8] = include_bytes!(env!("MEMRA_KDA_FATBIN"));
319const QMATVEC_FATBIN: &[u8] = include_bytes!(env!("MEMRA_QMATVEC_FATBIN"));
320const FLASH_FATBIN: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN"));
321const GEMM_FATBIN: &[u8] = include_bytes!(env!("MEMRA_GEMM_FATBIN"));
322const ROUTER_FATBIN: &[u8] = include_bytes!(env!("MEMRA_ROUTER_FATBIN"));
323const SAMPLE_FATBIN: &[u8] = include_bytes!(env!("MEMRA_SAMPLE_FATBIN"));
325
326fn gemm_fatbin_bytes() -> std::borrow::Cow<'static, [u8]> {
332 assert!(
333 !(portable_mma_gated() && std::env::var_os("MEMRA_GEMM_FATBIN").is_some()),
334 "MEMRA_GEMM_FATBIN overrides are not allowed in the portable CUDA lane"
335 );
336 match std::env::var("MEMRA_GEMM_FATBIN") {
337 Ok(path) => std::borrow::Cow::Owned(
338 std::fs::read(&path).unwrap_or_else(|e| panic!("MEMRA_GEMM_FATBIN read {path}: {e}")),
339 ),
340 Err(_) => std::borrow::Cow::Borrowed(GEMM_FATBIN),
341 }
342}
343
344pub(crate) const fn portable_mma_gated() -> bool {
351 cfg!(memra_portable_cuda) && !cfg!(memra_hopper_mma)
352}
353
354#[track_caller]
371pub(crate) fn refuse_portable_force(var: &str, needs: &str) {
372 assert!(
373 !portable_mma_gated(),
374 "{var} forces a kernel path this build does not contain: it needs {needs}, and this is a \
375 portable-CUDA build (sm_89). Unset {var} — the default path serves this arch."
376 );
377}
378
379pub(crate) const fn gdn_mma_default_on() -> bool {
388 cfg!(memra_hopper_mma) || konst_eq(env!("MEMRA_BUILT_CUDA_ARCH"), "120a")
389}
390
391const fn konst_eq(a: &str, b: &str) -> bool {
393 let (a, b) = (a.as_bytes(), b.as_bytes());
394 if a.len() != b.len() {
395 return false;
396 }
397 let mut i = 0;
398 while i < a.len() {
399 if a[i] != b[i] {
400 return false;
401 }
402 i += 1;
403 }
404 true
405}
406
407const fn legacy_quant_gemm_allowed(portable_cuda: bool, hopper_mma: bool, no_gemm: bool) -> bool {
412 (!portable_cuda || hopper_mma) && !no_gemm
413}
414
415const FLASH_FATBIN_VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VQ4"));
423const FLASH_FATBIN_VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VF8"));
424const FLASH_FATBIN_KF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8"));
425const FLASH_FATBIN_KF8VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VQ4"));
426const FLASH_FATBIN_KF8VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VF8"));
427
428pub use memra_kv::{kv_blk_bytes, kv_cache_formats};
431
432fn flash_fatbin_bytes() -> &'static [u8] {
434 match kv_cache_formats() {
435 ("q8_0", "q5_1") => FLASH_FATBIN,
436 ("q8_0", "q4_0") => FLASH_FATBIN_VQ4,
437 ("q8_0", "fp8") => FLASH_FATBIN_VF8,
438 ("fp8", "q5_1") => FLASH_FATBIN_KF8,
439 ("fp8", "q4_0") => FLASH_FATBIN_KF8VQ4,
440 ("fp8", "fp8") => FLASH_FATBIN_KF8VF8,
441 other => unreachable!("kv_cache_formats returned {other:?}"),
442 }
443}
444
445fn k1_launch_override() -> Option<(u32, u32, u32)> {
452 static K1: std::sync::OnceLock<Option<(u32, u32, u32)>> = std::sync::OnceLock::new();
453 *K1.get_or_init(|| {
454 let v = std::env::var("MEMRA_GEMM_K1_LAUNCH").ok()?;
455 let p: Vec<u32> = v.split(',').filter_map(|s| s.trim().parse().ok()).collect();
456 match p.as_slice() {
457 [bm, bn, w] => Some((*bm, *bn, *w)),
458 _ => None,
459 }
460 })
461}
462
463pub(crate) fn wgmma_gemm_enabled() -> bool {
470 static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
471 *V.get_or_init(|| std::env::var("MEMRA_WGMMA").as_deref() == Ok("1"))
472}
473
474pub const FA_VEC_MIN_TKV: usize = 96;
489pub fn fa_vec_min_tkv() -> usize {
493 static V: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
494 *V.get_or_init(|| {
495 std::env::var("MEMRA_FA_VEC_MIN")
496 .ok()
497 .and_then(|v| v.parse().ok())
498 .unwrap_or_else(|| FA_VEC_MIN_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
499 })
500}
501
502pub fn fa_f16pv_on() -> bool {
513 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
514 *ON.get_or_init(|| {
515 std::env::var("MEMRA_FA_F16PV")
516 .map(|v| v != "0")
517 .unwrap_or_else(|_| std::env::var("MEMRA_DRAFT").is_err())
518 })
519}
520
521pub fn fa512_hp_on() -> bool {
525 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
526 *ON.get_or_init(|| std::env::var("MEMRA_FA512_HP").as_deref() != Ok("0"))
527}
528
529pub fn faw_hp_on() -> bool {
533 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
534 *ON.get_or_init(|| std::env::var("MEMRA_FAW_HP").as_deref() != Ok("0"))
535}
536
537pub fn fa512_wide_warps() -> usize {
541 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
542 *N.get_or_init(|| match std::env::var("MEMRA_FA512_W4").as_deref() {
543 Ok("1") => 4,
544 _ => 2,
545 })
546}
547
548pub fn fa512_min_tkv() -> usize {
551 static FA512_MIN: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
552 *FA512_MIN.get_or_init(|| {
553 std::env::var("MEMRA_FA512_MIN")
554 .ok()
555 .and_then(|v| v.parse().ok())
556 .unwrap_or(512)
557 })
558}
559pub static FA_VEC_MIN_DEFAULT: std::sync::atomic::AtomicUsize =
563 std::sync::atomic::AtomicUsize::new(FA_VEC_MIN_TKV);
564pub static FA_SPW_DEFAULT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(32);
568pub static FUSED_MR1_DEFAULT: std::sync::atomic::AtomicBool =
574 std::sync::atomic::AtomicBool::new(false);
575pub static ROUTER_W8_DEFAULT: std::sync::atomic::AtomicBool =
582 std::sync::atomic::AtomicBool::new(true);
583pub static FA_SP512_DEFAULT: std::sync::atomic::AtomicUsize =
584 std::sync::atomic::AtomicUsize::new(16);
585pub static RMS_BLOCK_DEFAULT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(256);
590pub static FA_SP_GEMMA: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
592pub static MMQ_SK_FORCE: std::sync::atomic::AtomicI8 = std::sync::atomic::AtomicI8::new(-1);
597pub use memra_kv::KV_FP8_FORCE;
600pub(crate) fn mmv_block() -> u32 {
605 static V: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
606 *V.get_or_init(|| {
607 std::env::var("MEMRA_MMV_BLOCK")
608 .ok()
609 .and_then(|v| v.parse().ok())
610 .filter(|&b: &u32| (64..=256).contains(&b) && b % 32 == 0)
611 .unwrap_or(128)
612 })
613}
614
615static STEP37_SERVING_DEFAULTS: std::sync::atomic::AtomicBool =
641 std::sync::atomic::AtomicBool::new(false);
642
643pub fn arm_step37_serving_defaults() {
644 STEP37_SERVING_DEFAULTS.store(true, std::sync::atomic::Ordering::Relaxed);
645 crate::cache::set_swa_ring_default(true);
646 eprintln!(
647 "[step37-defaults] serving doors armed ON for the SlidingGatedMoe program \
648 (per-flag =0 kills, =1 forces; owner flip 2026-08-27)"
649 );
650}
651
652pub(crate) fn step37_defaults_armed() -> bool {
653 STEP37_SERVING_DEFAULTS.load(std::sync::atomic::Ordering::Relaxed)
654}
655
656pub(crate) fn step37_door(cell: &'static std::sync::OnceLock<Option<bool>>, name: &str) -> bool {
660 match *cell.get_or_init(|| match std::env::var(name).ok().as_deref() {
661 Some("1") => Some(true),
662 Some("0") => Some(false),
663 _ => None,
664 }) {
665 Some(forced) => forced,
666 None => step37_defaults_armed(),
667 }
668}
669
670pub(crate) fn w8_hybrid_on() -> bool {
671 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
672 step37_door(&ENV, "MEMRA_W8_HYBRID")
673}
674
675pub(crate) fn step_tp_w8_on() -> bool {
676 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
677 step37_door(&ENV, "MEMRA_STEP_TP_W8")
678}
679
680pub(crate) fn w8_view_on() -> bool {
685 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
686 *ON.get_or_init(|| std::env::var("MEMRA_W8_VIEW").as_deref() == Ok("1"))
687}
688
689pub(crate) fn step_gemm_prime_on() -> bool {
703 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
704 step37_door(&ENV, "MEMRA_STEP_GEMM_PRIME")
705}
706
707pub(crate) fn q8t_wonce_on() -> bool {
708 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
709 step37_door(&ENV, "MEMRA_Q8T_WONCE")
710}
711
712pub(crate) fn sig_expf_dev_on() -> bool {
717 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
718 *ON.get_or_init(|| std::env::var("MEMRA_SIG_EXPF_DEV").as_deref() == Ok("1"))
719}
720
721pub(crate) fn topk_fast_on() -> bool {
722 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
723 *ON.get_or_init(|| std::env::var("MEMRA_TOPK_FAST").as_deref() == Ok("1"))
724}
725
726fn sigmoid_topk_kernel(sig_expf: bool, fast: bool, n_used: usize) -> &'static str {
731 match (sig_expf, fast && n_used <= 8) {
732 (true, true) => "moe_router_sigmoid_topk_f32_dexp_fast",
733 (true, false) => "moe_router_sigmoid_topk_f32_dexp",
734 (false, true) => "moe_router_sigmoid_topk_f32_fast",
735 (false, false) => "moe_router_sigmoid_topk_f32",
736 }
737}
738
739#[cfg(test)]
740mod sigmoid_topk_dispatch_tests {
741 #[test]
742 fn fast_kernel_refuses_wide_topk_and_composes_with_dexp() {
743 use super::sigmoid_topk_kernel;
744
745 assert_eq!(
746 sigmoid_topk_kernel(false, true, 8),
747 "moe_router_sigmoid_topk_f32_fast"
748 );
749 assert_eq!(
750 sigmoid_topk_kernel(true, true, 8),
751 "moe_router_sigmoid_topk_f32_dexp_fast"
752 );
753 assert_eq!(
754 sigmoid_topk_kernel(false, true, 9),
755 "moe_router_sigmoid_topk_f32"
756 );
757 assert_eq!(
758 sigmoid_topk_kernel(true, true, 9),
759 "moe_router_sigmoid_topk_f32_dexp"
760 );
761 }
762}
763
764pub(crate) fn rms_block() -> u32 {
765 static V: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
766 *V.get_or_init(|| {
767 std::env::var("MEMRA_RMS_BLOCK")
768 .ok()
769 .and_then(|v| v.parse().ok())
770 .unwrap_or_else(|| RMS_BLOCK_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
771 })
772}
773
774pub(crate) fn fa_split_keys(t_kv: usize, n_head_kv: usize) -> usize {
775 static S: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
776 if let Some(forced) = *S.get_or_init(|| {
777 std::env::var("MEMRA_FA_SPLIT")
778 .ok()
779 .and_then(|v| v.parse().ok())
780 .filter(|&s: &usize| s >= 8 && s % 8 == 0)
781 }) {
782 return forced;
783 }
784 if FA_SP_GEMMA.load(std::sync::atomic::Ordering::Relaxed)
802 && std::env::var("MEMRA_FA_SP16").as_deref() == Ok("1")
803 {
804 return if t_kv <= 8192 {
805 16
806 } else if t_kv <= 16384 {
807 64
808 } else {
809 128
810 };
811 }
812 let big_rig = fa_sm_count() >= 128;
813 if big_rig {
814 let _ = n_head_kv;
815 if t_kv <= 2048 {
816 static SHORT: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
825 if let Some(sp) = *SHORT.get_or_init(|| {
826 std::env::var("MEMRA_FA_SP_SHORT")
827 .ok()
828 .and_then(|v| v.parse().ok())
829 .filter(|&s: &usize| s >= 8 && s % 8 == 0)
830 }) {
831 return sp;
832 }
833 16
834 } else if t_kv <= 16384 {
835 64
836 } else {
837 128
838 }
839 } else if n_head_kv <= 4 {
840 if t_kv <= 512 {
861 8
862 } else if t_kv <= 16384 {
863 64
864 } else {
865 128
866 }
867 } else {
868 if t_kv <= 8192 {
869 32
870 } else if t_kv <= 16384 {
871 64
872 } else {
873 128
874 }
875 }
876}
877
878pub(crate) fn fa_sm_count() -> i32 {
881 static N: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
882 *N.get_or_init(|| {
883 cudarc::driver::result::init().ok();
884 cudarc::driver::result::device::get(0)
885 .and_then(|d| unsafe { cudarc::driver::result::device::get_attribute(
886 d, cudarc::driver::sys::CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT) })
887 .unwrap_or(82)
888 })
889}
890
891#[allow(clippy::type_complexity)] fn fa_hd_suffix(head_dim: usize) -> Result<&'static str, Box<dyn std::error::Error>> {
896 match head_dim {
897 256 => Ok(""),
898 128 => Ok("_hd128"),
899 d => Err(format!(
900 "fa_prefill: no kernel stamped for head_dim={d} (only 256/128); \
901 callers must gate to sdpa_naive"
902 )
903 .into()),
904 }
905}
906
907pub const QT_Q8_0: i32 = 0;
909pub const QT_Q4_K: i32 = 1;
910pub const QT_Q6_K: i32 = 2;
911pub const QT_Q5_K: i32 = 3;
912pub const QT_Q3_K: i32 = 4;
913pub const QT_IQ4_XS: i32 = 5;
914pub const QT_IQ3_S: i32 = 6;
915pub const QT_NVFP4: i32 = 7;
916pub const QT_NVFP4_V2: i32 = 107;
919pub const QT_F8_E4M3: i32 = 10;
925pub const QT_NVFP4_RP: i32 = 9;
928pub const QT_F32: i32 = 8;
930pub const QT_BF16: i32 = 11;
931pub const QT_Q4_0: i32 = 12; pub const QT_Q2_K: i32 = 13;
936pub const QT_F8_E4M3_BLK: i32 = 14;
952
953pub struct Engine {
955 pub gpu: memra_runtime::Gpu,
956 module: Arc<CudaModule>,
957 hybrid: Arc<CudaModule>,
958 kda: Arc<CudaModule>,
960 qmatvec: Arc<CudaModule>,
961 flash: Arc<CudaModule>,
962 flash_g: std::sync::OnceLock<Arc<CudaModule>>,
966 gemm: Arc<CudaModule>,
967 router: Arc<CudaModule>,
968 sample: Arc<CudaModule>,
970 moe_cache: Mutex<Option<crate::moe_cache::MoeSlotCache>>,
974 w8_mirrors: Mutex<std::collections::HashMap<(u64, u32, u32), CudaSlice<u8>>>,
983 #[allow(clippy::type_complexity)]
986 w8_act: Mutex<std::collections::HashMap<usize, (CudaSlice<i8>, CudaSlice<f32>)>>,
988 moe_cache_layout: Mutex<Option<Vec<usize>>>,
992 capture_keep_on: std::sync::atomic::AtomicBool,
998 verify_exact: std::sync::atomic::AtomicBool,
1003 capture_keep: Mutex<Vec<Box<dyn std::any::Any + Send>>>,
1004 pub copy_stream: Arc<CudaStream>,
1006 #[cfg(memra_cutlass)]
1013 cutlass_scratch: Mutex<Option<crate::cutlass_ffi::CutlassScratch>>,
1014 fp8_scratch: Mutex<Option<crate::fp8_ffi::Fp8Scratch>>,
1018 fa_vf16_scratch: Mutex<Option<CudaSlice<u8>>>,
1021 #[allow(clippy::type_complexity)]
1025 fa_part_pool: Mutex<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
1027 #[allow(clippy::type_complexity)]
1031 fa_part_retired: Mutex<Vec<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
1033 fn_cache: Mutex<std::collections::HashMap<String, CudaFunction>>,
1035 f16_scratch: Mutex<Option<crate::f16_ffi::F16Scratch>>,
1036 argmax_partials: Mutex<Option<(CudaSlice<f32>, CudaSlice<i32>)>>,
1041 prime_deqw_ws: Mutex<Option<(CudaSlice<u8>, CudaSlice<u8>)>>,
1046 router_stage: Mutex<Option<PinnedStage>>,
1050 hyper_decode_ws: Mutex<Option<crate::hyper::HyperDecodeWs>>,
1056 verify_ws: Mutex<VerifyWs>,
1068 vrows_macro_dev: Mutex<std::collections::HashMap<(u16, u8), CudaSlice<f32>>>,
1076 shexp_ones: Mutex<Option<CudaSlice<f32>>>,
1084}
1085
1086#[derive(Default)]
1091pub struct VerifyWs {
1092 f32_pool: std::collections::HashMap<usize, Vec<CudaSlice<f32>>>,
1093 i8_pool: std::collections::HashMap<usize, Vec<CudaSlice<i8>>>,
1094 u64_pool: std::collections::HashMap<usize, Vec<CudaSlice<u64>>>,
1095 held_bytes: usize,
1096}
1097
1098const VWS_PER_CLASS_CAP: usize = 16;
1101const VWS_HELD_BYTES_CAP: usize = 256 << 20;
1104
1105impl VerifyWs {
1106 fn take<T>(
1107 pool: &mut std::collections::HashMap<usize, Vec<CudaSlice<T>>>,
1108 held: &mut usize,
1109 n: usize,
1110 ) -> Option<CudaSlice<T>> {
1111 let s = pool.get_mut(&n)?.pop()?;
1112 *held -= n * std::mem::size_of::<T>();
1113 Some(s)
1114 }
1115 fn put<T>(
1116 pool: &mut std::collections::HashMap<usize, Vec<CudaSlice<T>>>,
1117 held: &mut usize,
1118 s: CudaSlice<T>,
1119 ) {
1120 let n = s.len();
1121 let bytes = n * std::mem::size_of::<T>();
1122 if *held + bytes > VWS_HELD_BYTES_CAP {
1123 return; }
1125 let v = pool.entry(n).or_default();
1126 if v.len() >= VWS_PER_CLASS_CAP {
1127 return;
1128 }
1129 v.push(s);
1130 *held += bytes;
1131 }
1132}
1133
1134pub static SCRATCH_ALLOC_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
1140
1141fn fa_v2_on() -> bool {
1151 std::env::var("MEMRA_FA_V2")
1157 .map(|v| v != "0")
1158 .unwrap_or(true)
1159}
1160
1161pub(crate) fn fa_part_zero_on() -> bool {
1171 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1172 *ON.get_or_init(|| std::env::var("MEMRA_FA_PART_ZERO").as_deref() == Ok("1"))
1173}
1174
1175pub(crate) fn fa_v3_on() -> bool {
1176 std::env::var("MEMRA_FA_V3")
1180 .map(|v| v != "0")
1181 .unwrap_or(true)
1182}
1183
1184fn fa_v4_mode() -> &'static str {
1189 static M: std::sync::OnceLock<String> = std::sync::OnceLock::new();
1190 M.get_or_init(|| std::env::var("MEMRA_FA_V4").unwrap_or_default())
1191}
1192fn fa_v4_on() -> bool {
1193 fa_v4_mode() != "0"
1194} pub static FA_SMEM_TKV_DEFAULT: std::sync::atomic::AtomicUsize =
1203 std::sync::atomic::AtomicUsize::new(1024);
1204pub static FA_V4_MAX_DEFAULT: std::sync::atomic::AtomicUsize =
1205 std::sync::atomic::AtomicUsize::new(usize::MAX);
1206pub fn fa_v4_at_pub(t_kv: usize) -> bool {
1207 fa_v4_at(t_kv)
1208}
1209fn fa_v4_at(t_kv: usize) -> bool {
1210 static M: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
1211 let mx = *M.get_or_init(|| {
1212 std::env::var("MEMRA_FA_V4_MAX")
1213 .ok()
1214 .and_then(|v| v.parse().ok())
1215 .unwrap_or_else(|| FA_V4_MAX_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
1216 });
1217 fa_v4_on() && t_kv < mx
1218}
1219pub const FA_DEEP_MIN_DEFAULT: usize = 0;
1233fn fa_deep_at(t_kv: usize) -> bool {
1234 if std::env::var("MEMRA_FA_DEEP").as_deref() == Ok("0") {
1235 return false;
1236 }
1237 let min = std::env::var("MEMRA_FA_DEEP_MIN")
1238 .ok()
1239 .and_then(|v| v.parse().ok())
1240 .unwrap_or(FA_DEEP_MIN_DEFAULT);
1241 t_kv >= min
1242}
1243pub fn fa_deep_at_pub(t_kv: usize) -> bool {
1245 fa_deep_at(t_kv)
1246}
1247
1248fn fa_v3_active(head_dim: usize) -> bool {
1249 fa_v3_on()
1252 && head_dim.is_multiple_of(128)
1253 && kv_cache_formats() == ("q8_0", "q5_1")
1254 && !Engine::kv_fp8_on()
1255}
1256
1257pub fn fa_seqs_eligible(t_kv: usize, head_dim: usize) -> bool {
1265 std::env::var("MEMRA_NO_FA_VEC").is_err()
1266 && t_kv >= fa_vec_min_tkv()
1267 && head_dim == 256
1268 && fa_v4_at(t_kv)
1269 && !matches!(fa_v4_mode(), "noB3" | "stage")
1270 && !Engine::kv_fp8_on()
1271}
1272pub fn fa_split_keys_pub(t_kv: usize, n_head_kv: usize) -> usize {
1274 fa_split_keys(t_kv, n_head_kv)
1275}
1276
1277struct PinnedStage {
1282 ptr: *mut u8,
1283 cap: usize,
1284}
1285unsafe impl Send for PinnedStage {}
1286impl PinnedStage {
1287 fn new(cap: usize) -> Result<Self, Box<dyn std::error::Error>> {
1288 let ptr = unsafe { cudarc::driver::result::malloc_host(cap, 0)? } as *mut u8;
1289 Ok(PinnedStage { ptr, cap })
1290 }
1291}
1292impl Drop for PinnedStage {
1293 fn drop(&mut self) {
1294 let _ = unsafe { cudarc::driver::result::free_host(self.ptr as _) };
1295 }
1296}
1297
1298pub struct PinnedHostBuf {
1305 ptr: *mut u8,
1306 len: usize,
1307}
1308unsafe impl Send for PinnedHostBuf {}
1311impl PinnedHostBuf {
1312 pub fn new(len: usize) -> Result<Self, Box<dyn std::error::Error>> {
1315 let ptr = unsafe { cudarc::driver::result::malloc_host(len.max(1), 0)? } as *mut u8;
1316 Ok(PinnedHostBuf { ptr, len })
1317 }
1318 pub fn len(&self) -> usize {
1319 self.len
1320 }
1321 pub fn is_empty(&self) -> bool {
1322 self.len == 0
1323 }
1324 pub fn as_slice(&self) -> &[u8] {
1325 unsafe { std::slice::from_raw_parts(self.ptr, self.len) }
1326 }
1327 pub fn as_mut_slice(&mut self) -> &mut [u8] {
1328 unsafe { std::slice::from_raw_parts_mut(self.ptr, self.len) }
1329 }
1330}
1331impl Drop for PinnedHostBuf {
1332 fn drop(&mut self) {
1333 let _ = unsafe { cudarc::driver::result::free_host(self.ptr as _) };
1334 }
1335}
1336
1337pub const ARGMAX_NB: usize = 256;
1340
1341pub(crate) use memra_fa3_vl as fa3_vl_raw;
1343
1344unsafe extern "C" {
1345 fn memra_fa3_prefill(
1347 q16: *const core::ffi::c_void,
1348 k16: *const core::ffi::c_void,
1349 v16: *const core::ffi::c_void,
1350 o: *mut f32,
1351 t: i32,
1352 h: i32,
1353 hkv: i32,
1354 d: i32,
1355 scale: f32,
1356 stream: *mut core::ffi::c_void,
1357 ) -> i32;
1358 pub(crate) fn memra_fa3_vl(
1360 q16s: *const *const core::ffi::c_void,
1361 k16s: *const *const core::ffi::c_void,
1362 v16s: *const *const core::ffi::c_void,
1363 os: *const *mut f32,
1364 ts: *const i32,
1365 b: i32,
1366 h: i32,
1367 hkv: i32,
1368 d: i32,
1369 scale: f32,
1370 stream: *mut core::ffi::c_void,
1371 ) -> i32;
1372}
1373
1374#[repr(C)]
1379#[derive(Clone, Copy)]
1380pub struct WPtr8(pub [u64; 8]);
1381unsafe impl cudarc::driver::DeviceRepr for WPtr8 {}
1382
1383#[repr(C)]
1388#[derive(Clone, Copy, Default)]
1389pub struct GdnSeqVl {
1390 pub kb16: u64,
1391 pub gcum: u64,
1392 pub beta: u64,
1393 pub u: u64,
1394 pub wb16: u64,
1395 pub y: u64,
1396 pub ssnap: u64,
1397 pub state_in: u64,
1398 pub state_out: u64,
1399 pub q: u64,
1400 pub p: u64,
1401 pub o: u64,
1402 pub k: u64,
1403 pub v: u64,
1404 pub g: u64,
1405 pub a: u64,
1406 pub w: u64,
1407 pub t: i32,
1408 pub nc: i32,
1409}
1410unsafe impl cudarc::driver::DeviceRepr for GdnSeqVl {}
1411#[repr(C)]
1412#[derive(Clone, Copy)]
1413pub struct GdnVl8(pub [GdnSeqVl; 8]);
1414unsafe impl cudarc::driver::DeviceRepr for GdnVl8 {}
1415
1416#[repr(C)]
1419#[derive(Clone, Copy, Default)]
1420pub struct GdnWVl {
1421 pub qb16: u64,
1422 pub pb16: u64,
1423}
1424unsafe impl cudarc::driver::DeviceRepr for GdnWVl {}
1425#[repr(C)]
1426#[derive(Clone, Copy)]
1427pub struct GdnWVl8(pub [GdnWVl; 8]);
1428unsafe impl cudarc::driver::DeviceRepr for GdnWVl8 {}
1429
1430#[repr(C)]
1432#[derive(Clone, Copy, Default)]
1433pub struct GdnPrepVl {
1434 pub qkv: u64,
1435 pub conv_state: u64,
1436 pub conv_out: u64,
1437 pub q_g: u64,
1438 pub k_g: u64,
1439 pub v_g: u64,
1440 pub q_l2: u64,
1441 pub k_l2: u64,
1442 pub beta_raw: u64,
1443 pub alpha: u64,
1444 pub beta: u64,
1445 pub g_log: u64,
1446 pub o: u64,
1447 pub z: u64,
1448 pub gn: u64,
1449 pub gn16: u64,
1450 pub kb16: u64,
1451 pub qb16: u64,
1452 pub t: i32,
1453 pub pad: i32,
1454}
1455unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl {}
1456#[repr(C)]
1457#[derive(Clone, Copy)]
1458pub struct GdnPrepVl8(pub [GdnPrepVl; 8]);
1459unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl8 {}
1460
1461#[repr(C)]
1463#[derive(Clone, Copy, Default)]
1464pub struct FaSeqVl {
1465 pub q: u64,
1466 pub k16: u64,
1467 pub v16: u64,
1468 pub o: u64,
1469 pub kf: u64,
1470 pub vf: u64,
1471 pub t: i32,
1472 pub pad: i32,
1473}
1474unsafe impl cudarc::driver::DeviceRepr for FaSeqVl {}
1475#[repr(C)]
1476#[derive(Clone, Copy)]
1477pub struct FaVl8(pub [FaSeqVl; 8]);
1478unsafe impl cudarc::driver::DeviceRepr for FaVl8 {}
1479
1480#[repr(C)]
1482#[derive(Clone, Copy, Default)]
1483pub struct AttnPreVl {
1484 pub qf: u64,
1485 pub kf: u64,
1486 pub vf: u64,
1487 pub q: u64,
1488 pub gate: u64,
1489 pub qn: u64,
1490 pub kn: u64,
1491 pub kc: u64,
1492 pub vc: u64,
1493 pub t: i32,
1494 pub pad: i32,
1495}
1496unsafe impl cudarc::driver::DeviceRepr for AttnPreVl {}
1497#[repr(C)]
1498#[derive(Clone, Copy)]
1499pub struct AttnPreVl8(pub [AttnPreVl; 8]);
1500unsafe impl cudarc::driver::DeviceRepr for AttnPreVl8 {}
1501
1502pub struct GdnChunkBufs {
1505 pub gcum: CudaSlice<f32>,
1506 pub a: CudaSlice<f32>,
1507 pub p: CudaSlice<f32>,
1508 pub u: CudaSlice<f32>,
1509 pub w: CudaSlice<f32>,
1510 pub kb16: CudaSlice<u8>,
1511 pub wb16: CudaSlice<u8>,
1512 pub y16: CudaSlice<u8>,
1513 pub ssnap16: CudaSlice<u8>,
1514 pub qb16: CudaSlice<u8>,
1515 pub pb16: CudaSlice<u8>,
1516 pub o: CudaSlice<f32>,
1517 pub t: usize,
1518 pub nc: usize,
1519}
1520
1521#[repr(C)]
1523#[derive(Clone, Copy)]
1524pub struct F32x8(pub [f32; 8]);
1525unsafe impl cudarc::driver::DeviceRepr for F32x8 {}
1526
1527pub static PRIME_NANOS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
1531
1532pub static MOE_FUSED_EPI_DISPATCHES: std::sync::atomic::AtomicU64 =
1541 std::sync::atomic::AtomicU64::new(0);
1542
1543pub fn moe_fused_epilogue_dispatches() -> u64 {
1547 MOE_FUSED_EPI_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
1548}
1549
1550pub static MOE_VROWS_DISPATCHES: std::sync::atomic::AtomicU64 =
1557 std::sync::atomic::AtomicU64::new(0);
1558
1559pub fn moe_vrows_dispatches() -> u64 {
1561 MOE_VROWS_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
1562}
1563
1564fn bf16_tcols_wide_on() -> bool {
1573 std::env::var("MEMRA_BF16_TCOLS_WIDE").as_deref() != Ok("0")
1574}
1575
1576pub static BF16_TCOLS_WIDE_DISPATCHES: std::sync::atomic::AtomicU64 =
1580 std::sync::atomic::AtomicU64::new(0);
1581
1582pub fn bf16_tcols_wide_dispatches() -> u64 {
1584 BF16_TCOLS_WIDE_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
1585}
1586
1587fn bf16_tcols_x1_on() -> bool {
1595 std::env::var("MEMRA_BF16_TCOLS_X1").as_deref() != Ok("0")
1596}
1597
1598pub static BF16_TCOLS_X1_DISPATCHES: std::sync::atomic::AtomicU64 =
1600 std::sync::atomic::AtomicU64::new(0);
1601
1602pub fn bf16_tcols_x1_dispatches() -> u64 {
1604 BF16_TCOLS_X1_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
1605}
1606
1607fn bf16_tcols_red_fused_on() -> bool {
1620 std::env::var("MEMRA_BF16_TCOLS_RED_FUSED").as_deref() == Ok("1")
1621}
1622
1623pub static BF16_TCOLS_RED_FUSED_DISPATCHES: std::sync::atomic::AtomicU64 =
1626 std::sync::atomic::AtomicU64::new(0);
1627
1628pub fn bf16_tcols_red_fused_dispatches() -> u64 {
1630 BF16_TCOLS_RED_FUSED_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
1631}
1632
1633fn moe_vrows_pack_on() -> bool {
1640 std::env::var("MEMRA_MOE_VROWS_PACK").as_deref() == Ok("1")
1641}
1642
1643pub static MOE_VROWS_PACK_DISPATCHES: std::sync::atomic::AtomicU64 =
1645 std::sync::atomic::AtomicU64::new(0);
1646
1647pub fn moe_vrows_pack_dispatches() -> u64 {
1649 MOE_VROWS_PACK_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
1650}
1651
1652fn moe_vrows_dev_tables_on() -> bool {
1662 std::env::var("MEMRA_MOE_VROWS_DEV_TABLES").as_deref() == Ok("1")
1663}
1664
1665pub static MOE_VROWS_DEV_TABLES_DISPATCHES: std::sync::atomic::AtomicU64 =
1667 std::sync::atomic::AtomicU64::new(0);
1668
1669pub fn moe_vrows_dev_tables_dispatches() -> u64 {
1671 MOE_VROWS_DEV_TABLES_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
1672}
1673
1674pub static MOE_VROWS_ROUTER_SYNCS_AVOIDED: std::sync::atomic::AtomicU64 =
1678 std::sync::atomic::AtomicU64::new(0);
1679
1680pub fn moe_vrows_router_syncs_avoided() -> u64 {
1682 MOE_VROWS_ROUTER_SYNCS_AVOIDED.load(std::sync::atomic::Ordering::Relaxed)
1683}
1684
1685fn moe_vrows_dedup_stat_on() -> bool {
1695 std::env::var("MEMRA_MOE_VROWS_DEDUP_STAT").as_deref() == Ok("1")
1696}
1697
1698pub static MOE_VROWS_PAIR_VISITS: std::sync::atomic::AtomicU64 =
1700 std::sync::atomic::AtomicU64::new(0);
1701
1702pub static MOE_VROWS_PAIR_DISTINCT: std::sync::atomic::AtomicU64 =
1706 std::sync::atomic::AtomicU64::new(0);
1707
1708pub(crate) fn vrows_overlap_counts(sel_all: &[u32]) -> (u64, u64) {
1715 let mut seen = std::collections::HashSet::with_capacity(sel_all.len());
1716 for &ex in sel_all {
1717 seen.insert(ex);
1718 }
1719 (sel_all.len() as u64, seen.len() as u64)
1720}
1721
1722static MOE_VROWS_DEDUP_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
1724
1725fn moe_vrows_dedup_report() {
1732 let n = MOE_VROWS_DEDUP_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1733 if n != 0 && !n.is_multiple_of(42) {
1734 return;
1735 }
1736 let (visits, distinct) = moe_vrows_pair_overlap();
1737 if visits == 0 {
1738 return;
1739 }
1740 let repeat = 100.0 * (1.0 - distinct as f64 / visits as f64);
1741 eprintln!(
1742 "[moe-vrows-dedup] layer-calls={} visits={visits} distinct={distinct} \
1743 repeat={repeat:.2}% = the cross-row expert-slab dedup ceiling on the vrows pair \
1744 (MEMRA_MOE_VROWS_DEDUP_STAT=1)",
1745 n + 1
1746 );
1747}
1748
1749pub fn vrows_overlap_counts_for_test(sel_all: &[u32]) -> (u64, u64) {
1752 vrows_overlap_counts(sel_all)
1753}
1754
1755pub fn moe_vrows_pair_overlap() -> (u64, u64) {
1757 (
1758 MOE_VROWS_PAIR_VISITS.load(std::sync::atomic::Ordering::Relaxed),
1759 MOE_VROWS_PAIR_DISTINCT.load(std::sync::atomic::Ordering::Relaxed),
1760 )
1761}
1762
1763fn moe_vrows_dedup_order_on() -> bool {
1780 std::env::var("MEMRA_MOE_VROWS_DEDUP_ORDER").as_deref() == Ok("1")
1781}
1782
1783fn moe_vrows_down_tmaj_on() -> bool {
1792 std::env::var("MEMRA_MOE_VROWS_DOWN_TMAJ").as_deref() == Ok("1")
1793}
1794
1795pub static MOE_VROWS_DEDUP_ORDER_DISPATCHES: std::sync::atomic::AtomicU64 =
1797 std::sync::atomic::AtomicU64::new(0);
1798
1799pub fn moe_vrows_dedup_order_dispatches() -> u64 {
1801 MOE_VROWS_DEDUP_ORDER_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
1802}
1803
1804pub static MOE_VROWS_DOWN_TMAJ_DISPATCHES: std::sync::atomic::AtomicU64 =
1806 std::sync::atomic::AtomicU64::new(0);
1807
1808pub fn moe_vrows_down_tmaj_dispatches() -> u64 {
1810 MOE_VROWS_DOWN_TMAJ_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
1811}
1812
1813pub static MOE_VROWS_SLAB_READS_AVOIDED: std::sync::atomic::AtomicU64 =
1824 std::sync::atomic::AtomicU64::new(0);
1825
1826pub fn moe_vrows_slab_reads_avoided() -> u64 {
1828 MOE_VROWS_SLAB_READS_AVOIDED.load(std::sync::atomic::Ordering::Relaxed)
1829}
1830
1831pub(crate) fn vrows_expert_major_order(sel_all: &[u32]) -> Vec<u64> {
1836 let mut ord: Vec<u64> = (0..sel_all.len() as u64).collect();
1837 ord.sort_by_key(|&p| sel_all[p as usize]);
1840 ord
1841}
1842
1843pub fn vrows_expert_major_order_for_test(sel_all: &[u32]) -> Vec<u64> {
1846 vrows_expert_major_order(sel_all)
1847}
1848
1849pub(crate) fn alias_door_from(
1870 general: (&'static str, Option<&str>),
1871 alias: (&'static str, Option<&str>),
1872) -> Result<(bool, &'static str), String> {
1873 match (general.1, alias.1) {
1874 (Some(g), Some(a)) if g != a => Err(format!(
1875 "{}={g:?} and {}={a:?} disagree — the alias and the general flag name ONE door \
1876 (unset one); refused rather than silently picking a precedence winner, and the \
1877 door falls closed to the shipped program",
1878 general.0, alias.0
1879 )),
1880 (Some(g), _) => Ok((g == "1", general.0)),
1881 (None, Some(a)) => Ok((a == "1", alias.0)),
1882 (None, None) => Ok((false, general.0)),
1883 }
1884}
1885
1886fn alias_door(
1903 general: &'static str,
1904 alias: &'static str,
1905 latch: &'static std::sync::atomic::AtomicBool,
1906) -> (bool, &'static str) {
1907 let g = std::env::var(general).ok();
1908 let a = std::env::var(alias).ok();
1909 match alias_door_from((general, g.as_deref()), (alias, a.as_deref())) {
1910 Ok(resolved) => resolved,
1911 Err(msg) => {
1912 if !latch.swap(true, std::sync::atomic::Ordering::Relaxed) {
1913 eprintln!("[flag-alias] {msg}");
1914 }
1915 (false, general)
1916 }
1917 }
1918}
1919
1920pub fn htod_diet_on() -> bool {
1940 htod_diet_armed().0
1941}
1942
1943static HTOD_DIET_ALIAS_WARNED: std::sync::atomic::AtomicBool =
1945 std::sync::atomic::AtomicBool::new(false);
1946
1947pub(crate) fn htod_diet_armed() -> (bool, &'static str) {
1949 alias_door(
1950 "MEMRA_HTOD_DIET",
1951 "MEMRA_GLM5_HTOD_DIET",
1952 &HTOD_DIET_ALIAS_WARNED,
1953 )
1954}
1955
1956pub static HTOD_DIET_AVOIDED: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
1959
1960pub fn htod_diet_avoided() -> u64 {
1962 HTOD_DIET_AVOIDED.load(std::sync::atomic::Ordering::Relaxed)
1963}
1964
1965pub fn ep_diet_on() -> bool {
1986 ep_diet_armed().0
1987}
1988
1989static EP_DIET_ALIAS_WARNED: std::sync::atomic::AtomicBool =
1991 std::sync::atomic::AtomicBool::new(false);
1992
1993pub(crate) fn ep_diet_armed() -> (bool, &'static str) {
1996 alias_door("MEMRA_EP_DIET", "MEMRA_GLM5_EP_DIET", &EP_DIET_ALIAS_WARNED)
1997}
1998
1999pub fn ep_grouped_prime_on() -> bool {
2017 ep_grouped_prime_armed().0
2018}
2019
2020static EP_GROUPED_PRIME_ALIAS_WARNED: std::sync::atomic::AtomicBool =
2022 std::sync::atomic::AtomicBool::new(false);
2023
2024pub(crate) fn ep_grouped_prime_armed() -> (bool, &'static str) {
2026 alias_door(
2027 "MEMRA_EP_GROUPED_PRIME",
2028 "MEMRA_GLM5_EP_GROUPED_PRIME",
2029 &EP_GROUPED_PRIME_ALIAS_WARNED,
2030 )
2031}
2032
2033fn topk_shards_on() -> bool {
2042 std::env::var("MEMRA_TOPK_SHARDS").as_deref() != Ok("0")
2043}
2044
2045pub static TOPK_SHARDS_DISPATCHES: std::sync::atomic::AtomicU64 =
2047 std::sync::atomic::AtomicU64::new(0);
2048
2049pub fn topk_shards_dispatches() -> u64 {
2051 TOPK_SHARDS_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
2052}
2053
2054fn verify_ws_on() -> bool {
2066 verify_ws_on_from(
2067 std::env::var("MEMRA_VERIFY_WS").ok().as_deref(),
2068 std::env::var("MEMRA_GLM5_VERIFY_WS").ok().as_deref(),
2069 )
2070}
2071
2072fn verify_ws_on_from(general: Option<&str>, glm5_alias: Option<&str>) -> bool {
2075 general != Some("0") && glm5_alias != Some("0")
2076}
2077
2078pub static VERIFY_WS_HITS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
2082
2083pub fn verify_ws_hits() -> u64 {
2085 VERIFY_WS_HITS.load(std::sync::atomic::Ordering::Relaxed)
2086}
2087
2088pub static MLA_TC_PREFILL_DISPATCHES: std::sync::atomic::AtomicU64 =
2095 std::sync::atomic::AtomicU64::new(0);
2096
2097pub fn mla_tc_prefill_dispatches() -> u64 {
2101 MLA_TC_PREFILL_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
2102}
2103
2104pub static MOE_GROUPED_PREFILL_DISPATCHES: std::sync::atomic::AtomicU64 =
2110 std::sync::atomic::AtomicU64::new(0);
2111
2112pub fn moe_grouped_prefill_dispatches() -> u64 {
2115 MOE_GROUPED_PREFILL_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
2116}
2117
2118#[must_use = "dropping immediately ends the exact scope"]
2123pub struct ExactScope<'a> {
2124 flag: &'a std::sync::atomic::AtomicBool,
2125 prev: bool,
2126}
2127
2128impl<'a> ExactScope<'a> {
2129 pub(crate) fn set(flag: &'a std::sync::atomic::AtomicBool, on: bool) -> Self {
2130 let prev = flag.load(std::sync::atomic::Ordering::Relaxed);
2131 flag.store(on, std::sync::atomic::Ordering::Relaxed);
2132 ExactScope { flag, prev }
2133 }
2134}
2135
2136impl Drop for ExactScope<'_> {
2137 fn drop(&mut self) {
2138 self.flag
2139 .store(self.prev, std::sync::atomic::Ordering::Relaxed);
2140 }
2141}
2142
2143#[cfg(test)]
2144mod verify_ws_flag_tests {
2145 use super::verify_ws_on_from;
2146
2147 #[test]
2148 fn off_wins_across_general_and_alias() {
2149 assert!(verify_ws_on_from(None, None));
2151 assert!(!verify_ws_on_from(Some("0"), None));
2154 assert!(!verify_ws_on_from(None, Some("0")));
2155 assert!(!verify_ws_on_from(Some("1"), Some("0")));
2156 assert!(!verify_ws_on_from(Some("0"), Some("1")));
2157 assert!(verify_ws_on_from(Some("1"), None));
2159 assert!(verify_ws_on_from(None, Some("1")));
2160 }
2161}
2162
2163#[cfg(test)]
2164mod alias_door_tests {
2165 use super::alias_door_from;
2166
2167 const G: &str = "MEMRA_EP_DIET";
2168 const A: &str = "MEMRA_GLM5_EP_DIET";
2169
2170 fn r(g: Option<&str>, a: Option<&str>) -> Result<(bool, &'static str), String> {
2171 alias_door_from((G, g), (A, a))
2172 }
2173
2174 #[test]
2175 fn default_off_and_either_name_arms() {
2176 assert_eq!(r(None, None).unwrap(), (false, G));
2178 assert_eq!(r(Some("1"), None).unwrap(), (true, G));
2180 assert_eq!(r(None, Some("1")).unwrap(), (true, A));
2181 assert_eq!(r(Some("0"), None).unwrap(), (false, G));
2183 assert_eq!(r(None, Some("0")).unwrap(), (false, A));
2184 assert_eq!(r(None, Some("on")).unwrap(), (false, A));
2186 assert_eq!(r(Some(""), None).unwrap(), (false, G));
2187 }
2188
2189 #[test]
2190 fn agreeing_pair_resolves_to_the_general_name() {
2191 assert_eq!(r(Some("1"), Some("1")).unwrap(), (true, G));
2192 assert_eq!(r(Some("0"), Some("0")).unwrap(), (false, G));
2193 }
2194
2195 #[test]
2196 fn disagreeing_pair_refuses_and_names_both() {
2197 for (g, a) in [("1", "0"), ("0", "1")] {
2198 let err = r(Some(g), Some(a)).expect_err("a disagreeing pair must refuse");
2199 assert!(
2200 err.contains(G),
2201 "the refusal must name the general flag: {err}"
2202 );
2203 assert!(err.contains(A), "the refusal must name the alias: {err}");
2204 assert!(err.contains("falls closed"), "{err}");
2207 }
2208 }
2209}
2210
2211#[cfg(test)]
2212mod exact_scope_tests {
2213 use std::sync::atomic::{AtomicBool, Ordering};
2214
2215 #[test]
2216 fn error_path_restores_verify_exact() {
2217 let flag = AtomicBool::new(false);
2222 let failing = |flag: &AtomicBool| -> Result<(), &'static str> {
2223 let _scope = super::ExactScope::set(flag, true);
2224 assert!(flag.load(Ordering::Relaxed), "scope arms the flag");
2225 Err("draft forward failed")? };
2227 assert!(failing(&flag).is_err());
2228 assert!(
2229 !flag.load(Ordering::Relaxed),
2230 "error propagation must restore the pre-scope value"
2231 );
2232 let flag = AtomicBool::new(true);
2234 {
2235 let _scope = super::ExactScope::set(&flag, true);
2236 }
2237 assert!(flag.load(Ordering::Relaxed));
2238 let flag = AtomicBool::new(false);
2240 let scope = super::ExactScope::set(&flag, true);
2241 drop(scope);
2242 assert!(!flag.load(Ordering::Relaxed));
2243 }
2244}
2245
2246impl Engine {
2247 pub fn new(ordinal: usize) -> Result<Self, Box<dyn std::error::Error>> {
2248 let gpu = memra_runtime::Gpu::new(ordinal)?;
2249 if std::env::var("MEMRA_ARCH_CHECK").as_deref() != Ok("0") {
2253 use cudarc::driver::sys::CUdevice_attribute_enum as A;
2254 let (maj, min) = cudarc::driver::result::device::get(ordinal as i32)
2255 .and_then(|d| unsafe {
2256 Ok((
2257 cudarc::driver::result::device::get_attribute(
2258 d,
2259 A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR,
2260 )?,
2261 cudarc::driver::result::device::get_attribute(
2262 d,
2263 A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR,
2264 )?,
2265 ))
2266 })
2267 .unwrap_or((0, 0));
2268 let built = env!("MEMRA_BUILT_CUDA_ARCH");
2269 let ok = matches!(
2270 (built, maj, min),
2271 ("120a", 12, 0) | ("120a", 12, 1) | ("100a", 10, 0) | ("90a", 9, 0) | ("89", 8, 9)
2272 );
2273 if !ok {
2274 return Err(format!(
2275 "memra was built for sm_{built} but device {ordinal} reports compute \
2276 capability {maj}.{min}. Rebuild on this machine (MEMRA_CUDA_ARCH \
2277 auto-detects the GPU) or set MEMRA_ARCH_CHECK=0 to bypass."
2278 )
2279 .into());
2280 }
2281 }
2282 unsafe {
2287 use cudarc::driver::sys;
2288 let dev: sys::CUdevice = ordinal as sys::CUdevice;
2289 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
2290 if sys::cuDeviceGetDefaultMemPool(&mut pool, dev) == sys::CUresult::CUDA_SUCCESS {
2291 let mut thresh: u64 = u64::MAX;
2292 let _ = sys::cuMemPoolSetAttribute(
2293 pool,
2294 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
2295 &mut thresh as *mut u64 as *mut core::ffi::c_void,
2296 );
2297 }
2298 }
2299 let module = gpu.ctx.load_module(Ptx::from_binary(FATBIN.to_vec()))?;
2300 let hybrid = gpu
2301 .ctx
2302 .load_module(Ptx::from_binary(HYBRID_FATBIN.to_vec()))?;
2303 let kda = gpu.ctx.load_module(Ptx::from_binary(KDA_FATBIN.to_vec()))?;
2304 let qmatvec = gpu
2305 .ctx
2306 .load_module(Ptx::from_binary(QMATVEC_FATBIN.to_vec()))?;
2307 let flash = gpu
2308 .ctx
2309 .load_module(Ptx::from_binary(flash_fatbin_bytes().to_vec()))?;
2310 let gemm = gpu
2311 .ctx
2312 .load_module(Ptx::from_binary(gemm_fatbin_bytes().into_owned()))?;
2313 let router = gpu
2314 .ctx
2315 .load_module(Ptx::from_binary(ROUTER_FATBIN.to_vec()))?;
2316 let sample = gpu
2317 .ctx
2318 .load_module(Ptx::from_binary(SAMPLE_FATBIN.to_vec()))?;
2319 let copy_stream = gpu.ctx.new_stream()?;
2320 if std::env::var("MEMRA_EVT")
2336 .map(|v| v == "1")
2337 .unwrap_or(false)
2338 {
2339 } else {
2341 unsafe {
2342 gpu.ctx.disable_event_tracking();
2343 }
2344 }
2345 Ok(Self {
2346 gpu,
2347 module,
2348 hybrid,
2349 kda,
2350 qmatvec,
2351 flash,
2352 flash_g: std::sync::OnceLock::new(),
2353 gemm,
2354 router,
2355 sample,
2356 moe_cache: Mutex::new(None),
2357 w8_mirrors: Mutex::new(std::collections::HashMap::new()),
2358 w8_act: Mutex::new(std::collections::HashMap::new()),
2359 moe_cache_layout: Mutex::new(None),
2360 copy_stream,
2361 capture_keep_on: std::sync::atomic::AtomicBool::new(false),
2362 verify_exact: std::sync::atomic::AtomicBool::new(false),
2363 capture_keep: Mutex::new(Vec::new()),
2364 argmax_partials: Mutex::new(None),
2365 prime_deqw_ws: Mutex::new(None),
2366 router_stage: Mutex::new(None),
2367 hyper_decode_ws: Mutex::new(None),
2368 verify_ws: Mutex::new(VerifyWs::default()),
2369 vrows_macro_dev: Mutex::new(std::collections::HashMap::new()),
2370 shexp_ones: Mutex::new(None),
2371 fp8_scratch: Mutex::new(None),
2372 fa_vf16_scratch: Mutex::new(None),
2373 fa_part_pool: Mutex::new(None),
2374 fa_part_retired: Mutex::new(Vec::new()),
2375 fn_cache: Mutex::new(Default::default()),
2376 f16_scratch: Mutex::new(None),
2377 #[cfg(memra_cutlass)]
2378 cutlass_scratch: Mutex::new(None),
2379 })
2380 }
2381
2382 pub fn ctx(&self) -> &Arc<CudaContext> {
2383 &self.gpu.ctx
2384 }
2385
2386 pub fn pool_cached_bytes(&self) -> usize {
2404 let (reserved, used) = self.pool_reserved_used();
2405 reserved.saturating_sub(used)
2406 }
2407
2408 pub fn device_graph_mem_reserved(&self) -> usize {
2418 use cudarc::driver::sys as cus;
2419 let Ok(dev) = cudarc::driver::result::device::get(self.gpu.ctx.ordinal() as i32) else {
2420 return 0;
2421 };
2422 let mut bytes: u64 = 0;
2423 let rc = unsafe {
2424 cus::cuDeviceGetGraphMemAttribute(
2425 dev,
2426 cus::CUgraphMem_attribute::CU_GRAPH_MEM_ATTR_RESERVED_MEM_CURRENT,
2427 &mut bytes as *mut u64 as *mut std::ffi::c_void,
2428 )
2429 };
2430 if rc == cus::cudaError_enum::CUDA_SUCCESS {
2431 bytes as usize
2432 } else {
2433 0
2434 }
2435 }
2436
2437 pub fn pool_trim_to_zero(&self) -> usize {
2451 use cudarc::driver::sys;
2452 let (before, _) = self.pool_reserved_used();
2453 unsafe {
2454 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
2455 if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
2456 != sys::CUresult::CUDA_SUCCESS
2457 {
2458 return 0;
2459 }
2460 let _ = sys::cuMemPoolTrimTo(pool, 0);
2461 }
2462 let (after, _) = self.pool_reserved_used();
2463 before.saturating_sub(after)
2464 }
2465
2466 pub fn pool_reserved_used(&self) -> (usize, usize) {
2467 use cudarc::driver::sys;
2468 unsafe {
2469 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
2470 if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
2471 != sys::CUresult::CUDA_SUCCESS
2472 {
2473 return (0, 0);
2474 }
2475 let (mut reserved, mut used) = (0u64, 0u64);
2476 if sys::cuMemPoolGetAttribute(
2477 pool,
2478 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RESERVED_MEM_CURRENT,
2479 &mut reserved as *mut u64 as *mut core::ffi::c_void,
2480 ) != sys::CUresult::CUDA_SUCCESS
2481 {
2482 return (0, 0);
2483 }
2484 if sys::cuMemPoolGetAttribute(
2485 pool,
2486 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_USED_MEM_CURRENT,
2487 &mut used as *mut u64 as *mut core::ffi::c_void,
2488 ) != sys::CUresult::CUDA_SUCCESS
2489 {
2490 return (0, 0);
2491 }
2492 (reserved as usize, used as usize)
2493 }
2494 }
2495
2496 pub fn pool_high_water_reset(&self) -> (usize, usize) {
2505 use cudarc::driver::sys;
2506 unsafe {
2507 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
2508 if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
2509 != sys::CUresult::CUDA_SUCCESS
2510 {
2511 return (0, 0);
2512 }
2513 let (mut reserved, mut used) = (0u64, 0u64);
2514 if sys::cuMemPoolGetAttribute(
2515 pool,
2516 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RESERVED_MEM_HIGH,
2517 &mut reserved as *mut u64 as *mut core::ffi::c_void,
2518 ) != sys::CUresult::CUDA_SUCCESS
2519 {
2520 return (0, 0);
2521 }
2522 if sys::cuMemPoolGetAttribute(
2523 pool,
2524 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_USED_MEM_HIGH,
2525 &mut used as *mut u64 as *mut core::ffi::c_void,
2526 ) != sys::CUresult::CUDA_SUCCESS
2527 {
2528 return (0, 0);
2529 }
2530 let mut zero: u64 = 0;
2533 let _ = sys::cuMemPoolSetAttribute(
2534 pool,
2535 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RESERVED_MEM_HIGH,
2536 &mut zero as *mut u64 as *mut core::ffi::c_void,
2537 );
2538 let mut zero2: u64 = 0;
2539 let _ = sys::cuMemPoolSetAttribute(
2540 pool,
2541 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_USED_MEM_HIGH,
2542 &mut zero2 as *mut u64 as *mut core::ffi::c_void,
2543 );
2544 (reserved as usize, used as usize)
2545 }
2546 }
2547
2548 pub fn stream(&self) -> Arc<CudaStream> {
2551 self.gpu.stream()
2552 }
2553 pub fn gkv_on() -> bool {
2556 memra_kv::gkv_on()
2557 }
2558
2559 pub fn wkv_on() -> bool {
2571 memra_kv::wkv_on()
2572 }
2573
2574 pub fn kv_fp8_on() -> bool {
2580 memra_kv::kv_fp8_on()
2581 }
2582
2583 fn fa_func(&self, name: &str, head_dim: usize) -> CudaFunction {
2586 if head_dim == 512 && Self::gkv_on() {
2587 self.func_g(name)
2588 } else {
2589 self.func(name)
2590 }
2591 }
2592
2593 fn func_g(&self, name: &str) -> CudaFunction {
2597 let m = self.flash_g.get_or_init(|| {
2598 self.gpu
2599 .ctx
2600 .load_module(cudarc::nvrtc::Ptx::from_binary(
2601 FLASH_FATBIN_KF8VF8.to_vec(),
2602 ))
2603 .expect("load kf8vf8 flash fatbin (fp8-globals arm)")
2604 });
2605 let key = format!("g:{name}");
2606 if let Some(f) = self.fn_cache.lock().unwrap().get(&key) {
2607 return f.clone();
2608 }
2609 let f = match m.load_function(name) {
2610 Ok(f) => f,
2611 Err(_) => self.func(name),
2612 };
2613 self.fn_cache.lock().unwrap().insert(key, f.clone());
2614 f
2615 }
2616
2617 fn func(&self, name: &str) -> CudaFunction {
2618 if let Some(f) = self.fn_cache.lock().unwrap().get(name) {
2621 return f.clone();
2622 }
2623 let f = self
2624 .module
2625 .load_function(name)
2626 .or_else(|_| self.hybrid.load_function(name))
2627 .or_else(|_| self.kda.load_function(name))
2628 .or_else(|_| self.qmatvec.load_function(name))
2629 .or_else(|_| self.flash.load_function(name))
2630 .or_else(|_| self.gemm.load_function(name))
2631 .or_else(|_| self.router.load_function(name))
2632 .or_else(|_| self.sample.load_function(name))
2633 .unwrap_or_else(|_| panic!("kernel {name} not in any fatbin"));
2634 self.fn_cache
2635 .lock()
2636 .unwrap()
2637 .insert(name.to_string(), f.clone());
2638 f
2639 }
2640
2641 pub fn scatter_trim_logits(
2644 &self,
2645 src: &CudaSlice<f32>,
2646 d2t: &CudaSlice<u32>,
2647 dst: &mut CudaSlice<f32>,
2648 d_vocab: usize,
2649 n_vocab: usize,
2650 ) -> Result<(), Box<dyn std::error::Error>> {
2651 let f1 = self.func("scatter_trim_logits_f32");
2652 let f2 = self.func("scatter_trim_logits_pass2_f32");
2653 let (dv, nv) = (d_vocab as i32, n_vocab as i32);
2654 let cfg1 = LaunchConfig {
2655 grid_dim: (256, 1, 1),
2656 block_dim: (256, 1, 1),
2657 shared_mem_bytes: 0,
2658 };
2659 let __s_b1 = self.gpu.stream();
2660 let mut b1 = __s_b1.launch_builder(&f1);
2661 b1.arg(src).arg(d2t).arg(&mut *dst).arg(&dv).arg(&nv);
2662 unsafe {
2663 b1.launch(cfg1)?;
2664 }
2665 let cfg2 = LaunchConfig {
2666 grid_dim: (d_vocab.div_ceil(256) as u32, 1, 1),
2667 block_dim: (256, 1, 1),
2668 shared_mem_bytes: 0,
2669 };
2670 let __s_b2 = self.gpu.stream();
2671 let mut b2 = __s_b2.launch_builder(&f2);
2672 b2.arg(src).arg(d2t).arg(&mut *dst).arg(&dv);
2673 unsafe {
2674 b2.launch(cfg2)?;
2675 }
2676 Ok(())
2677 }
2678
2679 #[allow(clippy::too_many_arguments)]
2685 pub fn filter_stats(
2686 &self,
2687 x: &CudaSlice<f32>,
2688 row_stride: usize,
2689 rows: &CudaSlice<i32>,
2690 out_th: &mut CudaSlice<f32>,
2691 out_z: &mut CudaSlice<f32>,
2692 out_max: &mut CudaSlice<f32>,
2693 n: usize,
2694 nrow: usize,
2695 temp: f32,
2696 top_k: i32,
2697 top_p: f32,
2698 min_p: f32,
2699 ) -> Result<(), Box<dyn std::error::Error>> {
2700 static COOP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2727 let coop_on =
2728 *COOP_ON.get_or_init(|| std::env::var("MEMRA_FILTER_COOP").as_deref() != Ok("0"));
2729 if coop_on && self.sm_count() >= 16 {
2730 let cap = self.sm_count() as usize / 16;
2731 let mut done = 0usize;
2732 while done < nrow {
2733 let chunk = cap.min(nrow - done);
2734 self.filter_stats_coop_chunk(
2735 x, row_stride, rows, done, out_th, out_z, out_max, n, chunk, temp, top_k,
2736 top_p, min_p,
2737 )?;
2738 done += chunk;
2739 }
2740 return Ok(());
2741 }
2742 self.filter_stats_plain_program(
2743 x, row_stride, rows, out_th, out_z, out_max, n, nrow, temp, top_k, top_p, min_p,
2744 )
2745 }
2746
2747 #[allow(clippy::too_many_arguments)]
2752 pub fn filter_stats_coop_chunk(
2753 &self,
2754 x: &CudaSlice<f32>,
2755 row_stride: usize,
2756 rows: &CudaSlice<i32>,
2757 row0: usize,
2758 out_th: &mut CudaSlice<f32>,
2759 out_z: &mut CudaSlice<f32>,
2760 out_max: &mut CudaSlice<f32>,
2761 n: usize,
2762 chunk: usize,
2763 temp: f32,
2764 top_k: i32,
2765 top_p: f32,
2766 min_p: f32,
2767 ) -> Result<(), Box<dyn std::error::Error>> {
2768 let (ni, nr, rs) = (n as i32, chunk as i32, row_stride as i64);
2769 let f = self.func("filter_stats_coop_f32");
2770 let mut ws = self.alloc_uninit::<f32>(chunk * (2 * 16 + 2))?;
2771 let cfg = LaunchConfig {
2772 grid_dim: (16, chunk as u32, 1),
2773 block_dim: (512, 1, 1),
2774 shared_mem_bytes: 0,
2775 };
2776 let rows_v = rows.slice(row0..row0 + chunk);
2777 let mut th_v = out_th.slice_mut(row0..row0 + chunk);
2778 let mut z_v = out_z.slice_mut(row0..row0 + chunk);
2779 let mut mx_v = out_max.slice_mut(row0..row0 + chunk);
2780 let __s_b = self.gpu.stream();
2781 let mut b = __s_b.launch_builder(&f);
2782 b.arg(x)
2783 .arg(&rs)
2784 .arg(&rows_v)
2785 .arg(&mut th_v)
2786 .arg(&mut z_v)
2787 .arg(&mut mx_v)
2788 .arg(&mut ws)
2789 .arg(&ni)
2790 .arg(&nr)
2791 .arg(&temp)
2792 .arg(&top_k)
2793 .arg(&top_p)
2794 .arg(&min_p);
2795 unsafe {
2796 b.launch_cooperative(cfg)?;
2797 }
2798 Ok(())
2799 }
2800
2801 #[allow(clippy::too_many_arguments)]
2805 pub fn filter_stats_plain_program(
2806 &self,
2807 x: &CudaSlice<f32>,
2808 row_stride: usize,
2809 rows: &CudaSlice<i32>,
2810 out_th: &mut CudaSlice<f32>,
2811 out_z: &mut CudaSlice<f32>,
2812 out_max: &mut CudaSlice<f32>,
2813 n: usize,
2814 nrow: usize,
2815 temp: f32,
2816 top_k: i32,
2817 top_p: f32,
2818 min_p: f32,
2819 ) -> Result<(), Box<dyn std::error::Error>> {
2820 let (ni, nr, rs) = (n as i32, nrow as i32, row_stride as i64);
2821 let f = self.func("filter_stats_f32");
2822 let cfg = LaunchConfig {
2823 grid_dim: (nrow as u32, 1, 1),
2824 block_dim: (1024, 1, 1),
2825 shared_mem_bytes: 0,
2826 };
2827 let __s_b = self.gpu.stream();
2828 let mut b = __s_b.launch_builder(&f);
2829 b.arg(x)
2830 .arg(&rs)
2831 .arg(rows)
2832 .arg(&mut *out_th)
2833 .arg(&mut *out_z)
2834 .arg(&mut *out_max)
2835 .arg(&ni)
2836 .arg(&nr)
2837 .arg(&temp)
2838 .arg(&top_k)
2839 .arg(&top_p)
2840 .arg(&min_p);
2841 unsafe {
2842 b.launch(cfg)?;
2843 }
2844 Ok(())
2845 }
2846
2847 #[allow(clippy::too_many_arguments)]
2849 pub fn softmax_gather_filtered(
2850 &self,
2851 x: &CudaSlice<f32>,
2852 row_stride: usize,
2853 ids: &CudaSlice<u32>,
2854 rows: &CudaSlice<i32>,
2855 th: &CudaSlice<f32>,
2856 z: &CudaSlice<f32>,
2857 out: &mut CudaSlice<f32>,
2858 n: usize,
2859 npair: usize,
2860 temp: f32,
2861 ) -> Result<(), Box<dyn std::error::Error>> {
2862 let f = self.func("softmax_gather_filtered_f32");
2863 let (ni, np, rs) = (n as i32, npair as i32, row_stride as i64);
2864 let cfg = LaunchConfig {
2865 grid_dim: (npair as u32, 1, 1),
2866 block_dim: (256, 1, 1),
2867 shared_mem_bytes: 0,
2868 };
2869 let __s_b = self.gpu.stream();
2870 let mut b = __s_b.launch_builder(&f);
2871 b.arg(x)
2872 .arg(&rs)
2873 .arg(ids)
2874 .arg(rows)
2875 .arg(th)
2876 .arg(z)
2877 .arg(&mut *out)
2878 .arg(&ni)
2879 .arg(&np)
2880 .arg(&temp);
2881 unsafe {
2882 b.launch(cfg)?;
2883 }
2884 Ok(())
2885 }
2886
2887 #[allow(clippy::too_many_arguments)]
2889 pub fn residual_sample_filtered(
2890 &self,
2891 p: &CudaSlice<f32>,
2892 q: Option<&CudaSlice<f32>>,
2893 n: usize,
2894 temp: f32,
2895 seed: u64,
2896 stream_pos: u32,
2897 p_stats: (f32, f32, f32),
2898 q_stats: (f32, f32, f32),
2899 out_tok: &mut CudaSlice<u32>,
2900 ) -> Result<(), Box<dyn std::error::Error>> {
2901 let f = self.func("residual_sample_filtered_f32");
2902 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2903 let has_q: i32 = q.is_some() as i32;
2904 let qbuf = q.unwrap_or(p);
2905 let (pm, pth, pz) = p_stats;
2906 let (qm, qth, qz) = q_stats;
2907 let cfg = LaunchConfig {
2908 grid_dim: (1, 1, 1),
2909 block_dim: (1024, 1, 1),
2910 shared_mem_bytes: 0,
2911 };
2912 let __s_b = self.gpu.stream();
2913 let mut b = __s_b.launch_builder(&f);
2914 b.arg(p)
2915 .arg(qbuf)
2916 .arg(&has_q)
2917 .arg(&ni)
2918 .arg(&temp)
2919 .arg(&slo)
2920 .arg(&shi)
2921 .arg(&stream_pos)
2922 .arg(&pm)
2923 .arg(&pth)
2924 .arg(&pz)
2925 .arg(&qm)
2926 .arg(&qth)
2927 .arg(&qz)
2928 .arg(&mut *out_tok);
2929 unsafe {
2930 b.launch(cfg)?;
2931 }
2932 Ok(())
2933 }
2934
2935 #[allow(clippy::too_many_arguments)]
2941 pub fn residual_sample_sparse_q(
2942 &self,
2943 p: &CudaSlice<f32>,
2944 cand_ids: &CudaSlice<u32>,
2945 q_probs: &CudaSlice<f32>,
2946 n_cand: usize,
2947 n: usize,
2948 temp: f32,
2949 seed: u64,
2950 stream_pos: u32,
2951 p_stats: (f32, f32, f32),
2952 out_tok: &mut CudaSlice<u32>,
2953 ) -> Result<(), Box<dyn std::error::Error>> {
2954 assert!(
2955 (1..=32).contains(&n_cand),
2956 "residual_sample_sparse_q supports 1..=32 candidates, got {n_cand}"
2957 );
2958 let f = self.func("residual_sample_sparse_q_f32");
2959 let (ni, nc) = (n as i32, n_cand as i32);
2960 let (slo, shi) = ((seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2961 let (pm, pth, pz) = p_stats;
2962 let cfg = LaunchConfig {
2963 grid_dim: (1, 1, 1),
2964 block_dim: (1024, 1, 1),
2965 shared_mem_bytes: 0,
2966 };
2967 let __s_b = self.gpu.stream();
2968 let mut b = __s_b.launch_builder(&f);
2969 b.arg(p)
2970 .arg(cand_ids)
2971 .arg(q_probs)
2972 .arg(&nc)
2973 .arg(&ni)
2974 .arg(&temp)
2975 .arg(&slo)
2976 .arg(&shi)
2977 .arg(&stream_pos)
2978 .arg(&pm)
2979 .arg(&pth)
2980 .arg(&pz)
2981 .arg(&mut *out_tok);
2982 unsafe {
2983 b.launch(cfg)?;
2984 }
2985 Ok(())
2986 }
2987
2988 #[allow(clippy::too_many_arguments)]
2990 pub fn gumbel_perturb_filtered(
2991 &self,
2992 x: &CudaSlice<f32>,
2993 y: &mut CudaSlice<f32>,
2994 n: usize,
2995 seed: u64,
2996 stream_pos: u32,
2997 temp: f32,
2998 row_max: f32,
2999 th: f32,
3000 ) -> Result<(), Box<dyn std::error::Error>> {
3001 let f = self.func("gumbel_perturb_filtered_f32");
3002 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
3003 let cfg = LaunchConfig {
3004 grid_dim: (n.div_ceil(256) as u32, 1, 1),
3005 block_dim: (256, 1, 1),
3006 shared_mem_bytes: 0,
3007 };
3008 let __s_b = self.gpu.stream();
3009 let mut b = __s_b.launch_builder(&f);
3010 b.arg(x)
3011 .arg(&mut *y)
3012 .arg(&ni)
3013 .arg(&slo)
3014 .arg(&shi)
3015 .arg(&stream_pos)
3016 .arg(&temp)
3017 .arg(&row_max)
3018 .arg(&th);
3019 unsafe {
3020 b.launch(cfg)?;
3021 }
3022 Ok(())
3023 }
3024
3025 #[allow(clippy::too_many_arguments)]
3029 pub fn penalize_logits(
3030 &self,
3031 x: &mut CudaSlice<f32>,
3032 hist: &CudaSlice<u32>,
3033 n_hist: usize,
3034 rep: f32,
3035 freq: f32,
3036 present: f32,
3037 n: usize,
3038 ) -> Result<(), Box<dyn std::error::Error>> {
3039 if n_hist == 0 {
3040 return Ok(());
3041 }
3042 let f = self.func("penalize_logits_f32");
3043 let (nh, ni) = (n_hist as i32, n as i32);
3044 let cfg = LaunchConfig {
3045 grid_dim: (n_hist.div_ceil(128) as u32, 1, 1),
3046 block_dim: (128, 1, 1),
3047 shared_mem_bytes: 0,
3048 };
3049 let __s_b = self.gpu.stream();
3050 let mut b = __s_b.launch_builder(&f);
3051 b.arg(&mut *x)
3052 .arg(hist)
3053 .arg(&nh)
3054 .arg(&rep)
3055 .arg(&freq)
3056 .arg(&present)
3057 .arg(&ni);
3058 unsafe {
3059 b.launch(cfg)?;
3060 }
3061 Ok(())
3062 }
3063
3064 #[allow(clippy::too_many_arguments)]
3066 pub fn penalize_logits_rows(
3067 &self,
3068 x: &mut CudaSlice<f32>,
3069 hist: &CudaSlice<u32>,
3070 n_hist: usize,
3071 rep: f32,
3072 freq: f32,
3073 present: f32,
3074 n: usize,
3075 nrow: usize,
3076 ) -> Result<(), Box<dyn std::error::Error>> {
3077 if n_hist == 0 || nrow == 0 {
3078 return Ok(());
3079 }
3080 let f = self.func("penalize_logits_rows_f32");
3081 let (nh, ni, nr) = (n_hist as i32, n as i32, nrow as i32);
3082 let cfg = LaunchConfig {
3083 grid_dim: (n_hist.div_ceil(128) as u32, nrow as u32, 1),
3084 block_dim: (128, 1, 1),
3085 shared_mem_bytes: 0,
3086 };
3087 let __s_b = self.gpu.stream();
3088 let mut b = __s_b.launch_builder(&f);
3089 b.arg(&mut *x)
3090 .arg(hist)
3091 .arg(&nh)
3092 .arg(&rep)
3093 .arg(&freq)
3094 .arg(&present)
3095 .arg(&ni)
3096 .arg(&nr);
3097 unsafe {
3098 b.launch(cfg)?;
3099 }
3100 Ok(())
3101 }
3102
3103 #[allow(clippy::too_many_arguments)]
3109 pub fn penalize_logits_sparse_rows(
3110 &self,
3111 x: &mut CudaSlice<f32>,
3112 ids: &[u32],
3113 counts: &[u32],
3114 offsets: &[i32],
3115 rows: &[i32],
3116 reps: &[f32],
3117 freqs: &[f32],
3118 presents: &[f32],
3119 n: usize,
3120 ) -> Result<(), Box<dyn std::error::Error>> {
3121 let nrow = rows.len();
3122 if nrow == 0 {
3123 return Ok(());
3124 }
3125 let _ni = i32::try_from(n).map_err(|_| "sparse penalty logits width must fit CUDA i32")?;
3126 let _nr = i32::try_from(nrow).map_err(|_| "sparse penalty row count must fit CUDA i32")?;
3127 let entry_count =
3128 i32::try_from(ids.len()).map_err(|_| "sparse penalty entry count must fit CUDA i32")?;
3129 if ids.len() != counts.len()
3130 || offsets.len() != nrow + 1
3131 || reps.len() != nrow
3132 || freqs.len() != nrow
3133 || presents.len() != nrow
3134 || offsets.first().copied() != Some(0)
3135 || offsets.last().copied() != Some(entry_count)
3136 {
3137 return Err("sparse penalty row metadata shape mismatch".into());
3138 }
3139 if counts.contains(&0) {
3140 return Err("sparse penalty counts must be positive".into());
3141 }
3142 let mut max_len = 0usize;
3143 for pair in offsets.windows(2) {
3144 if pair[0] < 0 || pair[1] < pair[0] {
3145 return Err("sparse penalty offsets must be monotonic".into());
3146 }
3147 max_len = max_len.max((pair[1] - pair[0]) as usize);
3148 }
3149 if max_len == 0 {
3150 return Ok(());
3151 }
3152
3153 let mut seen = std::collections::HashSet::with_capacity(ids.len());
3154 for (r, &row) in rows.iter().enumerate() {
3155 if row < 0 || (row as usize + 1).saturating_mul(n) > x.len() {
3156 return Err("sparse penalty row index exceeds logits shape".into());
3157 }
3158 let begin = offsets[r] as usize;
3159 let end = offsets[r + 1] as usize;
3160 for &id in &ids[begin..end] {
3161 if id as usize >= n {
3162 return Err("sparse penalty token id exceeds logits row".into());
3163 }
3164 if !seen.insert((row, id)) {
3165 return Err("sparse penalty entries must be unique per logits row".into());
3166 }
3167 }
3168 }
3169
3170 unsafe {
3172 self.penalize_logits_sparse_rows_unchecked(
3173 x, ids, counts, offsets, rows, reps, freqs, presents, n,
3174 )
3175 }
3176 }
3177
3178 #[allow(clippy::too_many_arguments)]
3187 pub(crate) unsafe fn penalize_logits_sparse_rows_unchecked(
3188 &self,
3189 x: &mut CudaSlice<f32>,
3190 ids: &[u32],
3191 counts: &[u32],
3192 offsets: &[i32],
3193 rows: &[i32],
3194 reps: &[f32],
3195 freqs: &[f32],
3196 presents: &[f32],
3197 n: usize,
3198 ) -> Result<(), Box<dyn std::error::Error>> {
3199 let nrow = rows.len();
3200 if nrow == 0 {
3201 return Ok(());
3202 }
3203 let max_len = offsets
3204 .windows(2)
3205 .map(|pair| (pair[1] - pair[0]) as usize)
3206 .max()
3207 .unwrap_or(0);
3208 if max_len == 0 {
3209 return Ok(());
3210 }
3211 let ids_d = self.htod_u32_v(ids)?;
3212 let counts_d = self.htod_u32_v(counts)?;
3213 let offsets_d = self.htod_i32(offsets)?;
3214 let rows_d = self.htod_i32(rows)?;
3215 let reps_d = self.htod(reps)?;
3216 let freqs_d = self.htod(freqs)?;
3217 let presents_d = self.htod(presents)?;
3218 let f = self.func("penalize_logits_sparse_rows_f32");
3219 let ni = i32::try_from(n).map_err(|_| "sparse penalty logits width must fit CUDA i32")?;
3220 let nr = i32::try_from(nrow).map_err(|_| "sparse penalty row count must fit CUDA i32")?;
3221 let cfg = LaunchConfig {
3222 grid_dim: (max_len.div_ceil(128) as u32, nrow as u32, 1),
3223 block_dim: (128, 1, 1),
3224 shared_mem_bytes: 0,
3225 };
3226 let __s_b = self.gpu.stream();
3227 let mut b = __s_b.launch_builder(&f);
3228 b.arg(&mut *x)
3229 .arg(&ids_d)
3230 .arg(&counts_d)
3231 .arg(&offsets_d)
3232 .arg(&rows_d)
3233 .arg(&reps_d)
3234 .arg(&freqs_d)
3235 .arg(&presents_d)
3236 .arg(&ni)
3237 .arg(&nr);
3238 unsafe {
3239 b.launch(cfg)?;
3240 }
3241 Ok(())
3242 }
3243
3244 #[allow(clippy::too_many_arguments)]
3252 pub fn penalize_logits_rows_inc(
3253 &self,
3254 x: &mut CudaSlice<f32>,
3255 hist: &CudaSlice<u32>,
3256 n_hist0: usize,
3257 rep: f32,
3258 freq: f32,
3259 present: f32,
3260 n: usize,
3261 nrow: usize,
3262 win: usize,
3263 ) -> Result<(), Box<dyn std::error::Error>> {
3264 if nrow == 0 || win == 0 || (n_hist0 == 0 && nrow == 1) {
3265 return Ok(());
3266 }
3267 debug_assert!(
3268 hist.len() >= n_hist0 + nrow - 1,
3269 "rows-inc hist must carry n_hist0 + nrow - 1 ids"
3270 );
3271 let f = self.func("penalize_logits_rows_inc_f32");
3272 let max_len = win.min(n_hist0 + nrow - 1).max(1);
3273 let (nh, ni, nr, wi) = (n_hist0 as i32, n as i32, nrow as i32, win as i32);
3274 let cfg = LaunchConfig {
3275 grid_dim: (max_len.div_ceil(128) as u32, nrow as u32, 1),
3276 block_dim: (128, 1, 1),
3277 shared_mem_bytes: 0,
3278 };
3279 let __s_b = self.gpu.stream();
3280 let mut b = __s_b.launch_builder(&f);
3281 b.arg(&mut *x)
3282 .arg(hist)
3283 .arg(&nh)
3284 .arg(&rep)
3285 .arg(&freq)
3286 .arg(&present)
3287 .arg(&ni)
3288 .arg(&nr)
3289 .arg(&wi);
3290 unsafe {
3291 b.launch(cfg)?;
3292 }
3293 Ok(())
3294 }
3295
3296 pub fn wpf_level() -> u32 {
3304 static ON: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
3305 *ON.get_or_init(|| {
3306 std::env::var("MEMRA_WPF")
3307 .ok()
3308 .and_then(|v| v.parse().ok())
3309 .unwrap_or(1)
3310 })
3311 }
3312
3313 pub fn set_verify_exact(&self, on: bool) {
3329 self.verify_exact
3330 .store(on, std::sync::atomic::Ordering::Relaxed);
3331 }
3332 pub(crate) fn verify_exact_on(&self) -> bool {
3333 self.verify_exact.load(std::sync::atomic::Ordering::Relaxed)
3334 }
3335
3336 pub fn exact_scope(&self, on: bool) -> ExactScope<'_> {
3342 ExactScope::set(&self.verify_exact, on)
3343 }
3344
3345 pub fn qkv_append_on() -> bool {
3348 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3349 *ON.get_or_init(|| {
3350 std::env::var("MEMRA_QKV_APPEND")
3351 .map(|v| v != "0")
3352 .unwrap_or(true)
3353 })
3354 }
3355
3356 pub fn pdl_wb_on() -> bool {
3359 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3360 *ON.get_or_init(|| {
3361 std::env::var("MEMRA_PDL_WB")
3362 .map(|v| v != "0")
3363 .unwrap_or(true)
3364 })
3365 }
3366
3367 pub fn norm_ilp_on() -> bool {
3375 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3376 *ON.get_or_init(|| {
3377 std::env::var("MEMRA_NORM_ILP")
3378 .map(|v| v != "0")
3379 .unwrap_or(true)
3380 })
3381 }
3382
3383 pub fn tk_ffn_dual_on() -> bool {
3391 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3392 *ON.get_or_init(|| {
3393 std::env::var("MEMRA_TK_FFN_DUAL")
3394 .map(|v| v != "0")
3395 .unwrap_or(true)
3396 })
3397 }
3398
3399 pub fn pdl_mmvq_on() -> bool {
3403 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3404 *ON.get_or_init(|| {
3405 std::env::var("MEMRA_PDL_MMVQ")
3406 .map(|v| v != "0")
3407 .unwrap_or(true)
3408 })
3409 }
3410
3411 pub fn pdl_on() -> bool {
3412 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3413 *ON.get_or_init(|| std::env::var("MEMRA_PDL").map(|v| v != "0").unwrap_or(true))
3414 }
3415
3416 pub fn pdl_nvfp4q8_on() -> bool {
3422 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3423 *ON.get_or_init(|| {
3424 std::env::var("MEMRA_PDL_NVFP4")
3425 .map(|v| v != "0")
3426 .unwrap_or(true)
3427 })
3428 }
3429
3430 fn q40_mr1_on() -> bool {
3436 static Q40MR: std::sync::OnceLock<Option<u32>> = std::sync::OnceLock::new();
3437 match *Q40MR.get_or_init(|| {
3438 std::env::var("MEMRA_Q40_MR")
3439 .ok()
3440 .and_then(|v| v.parse().ok())
3441 }) {
3442 Some(v) => v == 1,
3443 None => crate::FUSED_MR1_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
3444 }
3445 }
3446
3447 fn pdl_func_flash(
3452 &self,
3453 g: bool,
3454 name: &'static str,
3455 ) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
3456 use cudarc::driver::sys as cu;
3457 static MODS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool), usize>>> =
3464 std::sync::Mutex::new(None);
3465 #[allow(clippy::type_complexity)] static FNS: std::sync::Mutex<
3467 Option<std::collections::HashMap<(usize, bool, &'static str), usize>>,
3468 > = std::sync::Mutex::new(None);
3469 let ctx_key = self.ctx().cu_ctx() as usize;
3470 if let Some(&f) = FNS
3471 .lock()
3472 .unwrap()
3473 .get_or_insert_with(Default::default)
3474 .get(&(ctx_key, g, name))
3475 {
3476 return Ok(f as cu::CUfunction);
3477 }
3478 let module = {
3479 let mut mods = MODS.lock().unwrap();
3480 let map = mods.get_or_insert_with(Default::default);
3481 match map.get(&(ctx_key, g)) {
3482 Some(&m) => m,
3483 None => {
3484 let m = self.pdl_load_module_in_ctx(if g {
3485 FLASH_FATBIN_KF8VF8
3486 } else {
3487 FLASH_FATBIN
3488 })?;
3489 map.insert((ctx_key, g), m);
3490 m
3491 }
3492 }
3493 };
3494 let cname = std::ffi::CString::new(name)?;
3495 let mut f: cu::CUfunction = std::ptr::null_mut();
3496 let r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
3497 if r != cu::CUresult::CUDA_SUCCESS {
3498 return Err(format!("pdl_func_flash {name} (g={g}): {r:?}").into());
3499 }
3500 FNS.lock()
3501 .unwrap()
3502 .get_or_insert_with(Default::default)
3503 .insert((ctx_key, g, name), f as usize);
3504 Ok(f)
3505 }
3506
3507 fn pdl_load_module_in_ctx(&self, bytes: &[u8]) -> Result<usize, Box<dyn std::error::Error>> {
3512 use cudarc::driver::sys as cu;
3513 let mut prev: cu::CUcontext = std::ptr::null_mut();
3514 unsafe {
3515 cu::cuCtxGetCurrent(&mut prev).result()?;
3516 }
3517 self.ctx().bind_to_thread()?;
3518 let mut m: cu::CUmodule = std::ptr::null_mut();
3519 let r = unsafe { cu::cuModuleLoadData(&mut m, bytes.as_ptr() as *const std::ffi::c_void) };
3520 let restore = if prev.is_null() {
3521 cu::CUresult::CUDA_SUCCESS
3522 } else {
3523 unsafe { cu::cuCtxSetCurrent(prev) }
3524 };
3525 if r != cu::CUresult::CUDA_SUCCESS {
3526 return Err(format!("pdl module load: {r:?}").into());
3527 }
3528 if restore != cu::CUresult::CUDA_SUCCESS {
3529 return Err(format!("pdl module load: ctx restore {restore:?}").into());
3530 }
3531 Ok(m as usize)
3532 }
3533
3534 pub fn raw_kernel_function(
3537 &self,
3538 name: &'static str,
3539 ) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
3540 self.pdl_func(name)
3541 }
3542
3543 fn pdl_func(
3544 &self,
3545 name: &'static str,
3546 ) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
3547 use cudarc::driver::sys as cu;
3548 static MODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
3551 std::sync::Mutex::new(None);
3552 static QMODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
3555 std::sync::Mutex::new(None);
3556 static FNS: std::sync::Mutex<
3557 Option<std::collections::HashMap<(usize, &'static str), usize>>,
3558 > = std::sync::Mutex::new(None);
3559 let ctx_key = self.ctx().cu_ctx() as usize;
3560 if let Some(&f) = FNS
3561 .lock()
3562 .unwrap()
3563 .get_or_insert_with(Default::default)
3564 .get(&(ctx_key, name))
3565 {
3566 return Ok(f as cu::CUfunction);
3567 }
3568 let module = {
3569 let mut mods = MODULES.lock().unwrap();
3570 let map = mods.get_or_insert_with(Default::default);
3571 match map.get(&ctx_key) {
3572 Some(&m) => m,
3573 None => {
3574 let m = self.pdl_load_module_in_ctx(FATBIN)?;
3575 map.insert(ctx_key, m);
3576 m
3577 }
3578 }
3579 };
3580 let cname = std::ffi::CString::new(name)?;
3581 let mut f: cu::CUfunction = std::ptr::null_mut();
3582 let mut r =
3583 unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
3584 if r == cu::CUresult::CUDA_ERROR_NOT_FOUND {
3585 let qmodule = {
3586 let mut mods = QMODULES.lock().unwrap();
3587 let map = mods.get_or_insert_with(Default::default);
3588 match map.get(&ctx_key) {
3589 Some(&m) => m,
3590 None => {
3591 let m = self.pdl_load_module_in_ctx(QMATVEC_FATBIN)?;
3592 map.insert(ctx_key, m);
3593 m
3594 }
3595 }
3596 };
3597 r = unsafe { cu::cuModuleGetFunction(&mut f, qmodule as cu::CUmodule, cname.as_ptr()) };
3598 }
3599 if r != cu::CUresult::CUDA_SUCCESS {
3600 return Err(format!("pdl_func {name}: {r:?}").into());
3601 }
3602 FNS.lock()
3603 .unwrap()
3604 .get_or_insert_with(Default::default)
3605 .insert((ctx_key, name), f as usize);
3606 Ok(f)
3607 }
3608
3609 unsafe fn launch_pdl_flash(
3621 &self,
3622 g: bool,
3623 name: &'static str,
3624 grid: (u32, u32, u32),
3625 block: (u32, u32, u32),
3626 smem: u32,
3627 params: &mut [*mut std::ffi::c_void],
3628 ) -> Result<(), Box<dyn std::error::Error>> {
3629 use cudarc::driver::sys as cu;
3630 let f = self.pdl_func_flash(g, name)?;
3631 if smem > 0 {
3632 let r =
3634 unsafe {
3635 cu::cuFuncSetAttribute(f,
3636 cu::CUfunction_attribute_enum::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
3637 smem as i32)
3638 };
3639 if r != cu::CUresult::CUDA_SUCCESS {
3640 return Err(format!("pdl smem attr {name}: {r:?}").into());
3641 }
3642 }
3643 let mut attr = cu::CUlaunchAttribute {
3644 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
3645 pad: [0; 4],
3646 value: cu::CUlaunchAttributeValue {
3647 programmaticStreamSerializationAllowed: 1,
3648 },
3649 };
3650 let cfg = cu::CUlaunchConfig {
3651 gridDimX: grid.0,
3652 gridDimY: grid.1,
3653 gridDimZ: grid.2,
3654 blockDimX: block.0,
3655 blockDimY: block.1,
3656 blockDimZ: block.2,
3657 sharedMemBytes: smem,
3658 hStream: self.gpu.stream().cu_stream(),
3659 attrs: &mut attr,
3660 numAttrs: 1,
3661 };
3662 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
3663 if r != cu::CUresult::CUDA_SUCCESS {
3664 return Err(format!("launch_pdl_flash {name}: {r:?}").into());
3665 }
3666 Ok(())
3667 }
3668
3669 unsafe fn launch_pdl(
3670 &self,
3671 name: &'static str,
3672 grid: (u32, u32, u32),
3673 block: (u32, u32, u32),
3674 params: &mut [*mut std::ffi::c_void],
3675 ) -> Result<(), Box<dyn std::error::Error>> {
3676 use cudarc::driver::sys as cu;
3677 let f = self.pdl_func(name)?;
3678 let mut attr = cu::CUlaunchAttribute {
3679 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
3680 pad: [0; 4],
3681 value: cu::CUlaunchAttributeValue {
3682 programmaticStreamSerializationAllowed: 1,
3683 },
3684 };
3685 let cfg = cu::CUlaunchConfig {
3686 gridDimX: grid.0,
3687 gridDimY: grid.1,
3688 gridDimZ: grid.2,
3689 blockDimX: block.0,
3690 blockDimY: block.1,
3691 blockDimZ: block.2,
3692 sharedMemBytes: 0,
3693 hStream: self.gpu.stream().cu_stream(),
3694 attrs: &mut attr,
3695 numAttrs: 1,
3696 };
3697 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
3698 if r != cu::CUresult::CUDA_SUCCESS {
3699 return Err(format!("launch_pdl {name}: {r:?}").into());
3700 }
3701 Ok(())
3702 }
3703
3704 pub fn prefetch_weight_l2(
3707 &self,
3708 w: &crate::model::GpuTensor,
3709 ) -> Result<(), Box<dyn std::error::Error>> {
3710 if let crate::model::GpuTensor::Quant { bytes, rp4, .. } = w {
3711 let p = rp4.as_ref().unwrap_or(bytes);
3712 self.prefetch_l2(p, p.len())?;
3713 }
3714 Ok(())
3715 }
3716
3717 pub fn gather_row_bf16(
3720 &self,
3721 table: &CudaSlice<u8>,
3722 tok: &CudaSlice<u32>,
3723 idx: usize,
3724 dst: &mut CudaSlice<f32>,
3725 ncols: usize,
3726 ) -> Result<(), Box<dyn std::error::Error>> {
3727 let f = self.func("gather_row_bf16_f32");
3728 let cfg = LaunchConfig {
3729 grid_dim: (ncols.div_ceil(256) as u32, 1, 1),
3730 block_dim: (256, 1, 1),
3731 shared_mem_bytes: 0,
3732 };
3733 let (nc, ix) = (ncols as i32, idx as i32);
3734 let __s_b = self.gpu.stream();
3735 let mut b = __s_b.launch_builder(&f);
3736 b.arg(table).arg(tok).arg(&ix).arg(dst).arg(&nc);
3737 unsafe {
3738 b.launch(cfg)?;
3739 }
3740 Ok(())
3741 }
3742
3743 #[allow(clippy::too_many_arguments)]
3749 pub fn dflash2_dynconv(
3750 &self,
3751 x: &CudaSlice<f32>,
3752 dyn_: &CudaSlice<f32>,
3753 base: &CudaSlice<f32>,
3754 out: &mut CudaSlice<f32>,
3755 rows: usize,
3756 hidden: usize,
3757 group_size: usize,
3758 ksize: usize,
3759 half: usize,
3760 ) -> Result<(), Box<dyn std::error::Error>> {
3761 assert_eq!(hidden % group_size, 0, "hidden % group_size != 0");
3762 let f = self.func("dflash2_dynconv_f32");
3763 let n = rows * hidden;
3764 let cfg = LaunchConfig {
3765 grid_dim: (n.div_ceil(256) as u32, 1, 1),
3766 block_dim: (256, 1, 1),
3767 shared_mem_bytes: 0,
3768 };
3769 let (ri, hi, gi, ki, hf) = (
3770 rows as i32,
3771 hidden as i32,
3772 group_size as i32,
3773 ksize as i32,
3774 half as i32,
3775 );
3776 let __s_b = self.gpu.stream();
3777 let mut b = __s_b.launch_builder(&f);
3778 b.arg(x)
3779 .arg(dyn_)
3780 .arg(base)
3781 .arg(out)
3782 .arg(&ri)
3783 .arg(&hi)
3784 .arg(&gi)
3785 .arg(&ki)
3786 .arg(&hf);
3787 unsafe {
3788 b.launch(cfg)?;
3789 }
3790 Ok(())
3791 }
3792
3793 pub fn topk_rows(
3797 &self,
3798 logits: &CudaSlice<f32>,
3799 n_rows: usize,
3800 n_cols: usize,
3801 k: usize,
3802 ) -> Result<(CudaSlice<f32>, CudaSlice<u32>), Box<dyn std::error::Error>> {
3803 assert!((1..=32).contains(&k), "topk_rows supports 1..=32, got {k}");
3804 assert!(k <= n_cols, "topk_rows: k {k} > n_cols {n_cols}");
3805 if topk_shards_on() && n_cols >= 16 * 1024 && k <= n_cols / 16 {
3812 if TOPK_SHARDS_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
3813 eprintln!(
3814 "[topk-shards] engaged: rows={n_rows} cols={n_cols} k={k} shards=16 \
3815 (MEMRA_TOPK_SHARDS=1)"
3816 );
3817 }
3818 return self.topk_rows_sharded(logits, n_rows, n_cols, k, 16);
3819 }
3820 let f = self.func("topk_rows_f32");
3821 let nth = 256usize;
3822 let mut vals = self.uninit(n_rows * k)?;
3823 let mut idxs = self.gpu.stream().alloc_zeros::<u32>(n_rows * k)?;
3824 let cfg = LaunchConfig {
3825 grid_dim: (n_rows as u32, 1, 1),
3826 block_dim: (nth as u32, 1, 1),
3827 shared_mem_bytes: (nth * k * 8) as u32,
3828 };
3829 let (nr, nc, ki) = (n_rows as i32, n_cols as i32, k as i32);
3830 let __s_b = self.gpu.stream();
3831 let mut b = __s_b.launch_builder(&f);
3832 b.arg(logits)
3833 .arg(&nr)
3834 .arg(&nc)
3835 .arg(&ki)
3836 .arg(&mut vals)
3837 .arg(&mut idxs);
3838 unsafe {
3839 b.launch(cfg)?;
3840 }
3841 Ok((vals, idxs))
3842 }
3843
3844 fn topk_rows_sharded(
3850 &self,
3851 logits: &CudaSlice<f32>,
3852 n_rows: usize,
3853 n_cols: usize,
3854 k: usize,
3855 n_shards: usize,
3856 ) -> Result<(CudaSlice<f32>, CudaSlice<u32>), Box<dyn std::error::Error>> {
3857 assert!((1..=64).contains(&n_shards), "shard merge head cap is 64");
3858 let nth = 256usize;
3859 let mut pvals = self.uninit(n_rows * n_shards * k)?;
3860 let mut pidxs = self.alloc_uninit::<u32>(n_rows * n_shards * k)?;
3861 let f1 = self.func("topk_rows_shard_f32");
3862 let cfg1 = LaunchConfig {
3863 grid_dim: (n_rows as u32, n_shards as u32, 1),
3864 block_dim: (nth as u32, 1, 1),
3865 shared_mem_bytes: (nth * k * 8) as u32,
3866 };
3867 let (nr, nc, ki, ns) = (n_rows as i32, n_cols as i32, k as i32, n_shards as i32);
3868 {
3869 let __s_b = self.gpu.stream();
3870 let mut b = __s_b.launch_builder(&f1);
3871 b.arg(logits)
3872 .arg(&nr)
3873 .arg(&nc)
3874 .arg(&ki)
3875 .arg(&ns)
3876 .arg(&mut pvals)
3877 .arg(&mut pidxs);
3878 unsafe {
3879 b.launch(cfg1)?;
3880 }
3881 }
3882 let mut vals = self.uninit(n_rows * k)?;
3883 let mut idxs = self.alloc_uninit::<u32>(n_rows * k)?;
3884 let f2 = self.func("topk_rows_shard_merge_f32");
3885 let cfg2 = LaunchConfig {
3886 grid_dim: (n_rows as u32, 1, 1),
3887 block_dim: (32, 1, 1),
3888 shared_mem_bytes: 0,
3889 };
3890 let __s_b = self.gpu.stream();
3891 let mut b = __s_b.launch_builder(&f2);
3892 b.arg(&pvals)
3893 .arg(&pidxs)
3894 .arg(&nr)
3895 .arg(&ns)
3896 .arg(&ki)
3897 .arg(&mut vals)
3898 .arg(&mut idxs);
3899 unsafe {
3900 b.launch(cfg2)?;
3901 }
3902 Ok((vals, idxs))
3903 }
3904
3905 pub fn add_row_inplace(
3907 &self,
3908 logits: &mut CudaSlice<f32>,
3909 bias: &CudaSlice<f32>,
3910 n: usize,
3911 row_off: usize,
3912 ) -> Result<(), Box<dyn std::error::Error>> {
3913 let f = self.func("add_row_inplace_f32");
3914 let cfg = LaunchConfig {
3915 grid_dim: (n.div_ceil(256) as u32, 1, 1),
3916 block_dim: (256, 1, 1),
3917 shared_mem_bytes: 0,
3918 };
3919 let (ni, off) = (n as i32, row_off as i64);
3920 let __s_b = self.gpu.stream();
3921 let mut b = __s_b.launch_builder(&f);
3922 b.arg(logits).arg(bias).arg(&ni).arg(&off);
3923 unsafe {
3924 b.launch(cfg)?;
3925 }
3926 Ok(())
3927 }
3928
3929 pub fn prefetch_l2(
3931 &self,
3932 p: &CudaSlice<u8>,
3933 n: usize,
3934 ) -> Result<(), Box<dyn std::error::Error>> {
3935 let f = self.func("prefetch_l2_bytes");
3936 let lines = n.div_ceil(128);
3937 let ni = n as i64;
3938 let cfg = LaunchConfig {
3939 grid_dim: (lines.div_ceil(256) as u32, 1, 1),
3940 block_dim: (256, 1, 1),
3941 shared_mem_bytes: 0,
3942 };
3943 let __s_b = self.gpu.stream();
3944 let mut b = __s_b.launch_builder(&f);
3945 b.arg(p).arg(&ni);
3946 unsafe {
3947 b.launch(cfg)?;
3948 }
3949 Ok(())
3950 }
3951
3952 pub fn router_gemv(
3955 &self,
3956 w: &CudaSlice<f32>,
3957 x: &CudaSlice<f32>,
3958 n_embd: usize,
3959 n_experts: usize,
3960 t: usize,
3961 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3962 let w8 = match std::env::var("MEMRA_ROUTER_V2").as_deref() {
3968 Ok("0") => false,
3969 Ok(_) => true,
3970 Err(_) => ROUTER_W8_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
3971 };
3972 let batch = w8 && t >= ROUTER_BATCH_MIN_T && router_batch_on();
3982 self.router_gemv_form(w, x, n_embd, n_experts, t, w8, batch)
3983 }
3984
3985 #[allow(clippy::too_many_arguments)] pub fn router_gemv_form(
3989 &self,
3990 w: &CudaSlice<f32>,
3991 x: &CudaSlice<f32>,
3992 n_embd: usize,
3993 n_experts: usize,
3994 t: usize,
3995 w8: bool,
3996 batch: bool,
3997 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3998 debug_assert!(!batch || w8, "batch twin exists for the w8 form only");
3999 let mut y = self.alloc_uninit::<f32>(t * n_experts)?;
4000 let f = if batch {
4001 self.func("router_gemv_f32_w8_batch")
4002 } else if w8 {
4003 self.func("router_gemv_f32_w8")
4004 } else {
4005 self.func("router_gemv_f32")
4006 };
4007 let (ne, nx, ti) = (n_embd as i32, n_experts as i32, t as i32);
4008 let cfg = if batch {
4009 LaunchConfig {
4010 grid_dim: (n_experts.div_ceil(8) as u32, t.div_ceil(8) as u32, 1),
4011 block_dim: (32, 8, 1),
4012 shared_mem_bytes: 0,
4013 }
4014 } else {
4015 LaunchConfig {
4016 grid_dim: (n_experts as u32, t as u32, 1),
4017 block_dim: (32, if w8 { 8 } else { 1 }, 1),
4018 shared_mem_bytes: 0,
4019 }
4020 };
4021 let __s_b = self.gpu.stream();
4022 let mut b = __s_b.launch_builder(&f);
4023 b.arg(w).arg(x).arg(&mut y).arg(&ne).arg(&nx).arg(&ti);
4024 unsafe {
4025 b.launch(cfg)?;
4026 }
4027 Ok(y)
4028 }
4029
4030 pub fn router_gemv_into(
4033 &self,
4034 w: &CudaSlice<f32>,
4035 x: &CudaSlice<f32>,
4036 y: &mut CudaSlice<f32>,
4037 n_embd: usize,
4038 n_experts: usize,
4039 t: usize,
4040 ) -> Result<(), Box<dyn std::error::Error>> {
4041 if y.len() < t * n_experts {
4042 return Err("router_gemv_into output too small".into());
4043 }
4044 let w8 = match std::env::var("MEMRA_ROUTER_V2").as_deref() {
4045 Ok("0") => false,
4046 Ok(_) => true,
4047 Err(_) => ROUTER_W8_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
4048 };
4049 let f = if w8 {
4050 self.func("router_gemv_f32_w8")
4051 } else {
4052 self.func("router_gemv_f32")
4053 };
4054 let (ne, nx, ti) = (n_embd as i32, n_experts as i32, t as i32);
4055 let cfg = LaunchConfig {
4056 grid_dim: (n_experts as u32, t as u32, 1),
4057 block_dim: (32, if w8 { 8 } else { 1 }, 1),
4058 shared_mem_bytes: 0,
4059 };
4060 let __s_b = self.gpu.stream();
4061 let mut b = __s_b.launch_builder(&f);
4062 b.arg(w).arg(x).arg(&mut *y).arg(&ne).arg(&nx).arg(&ti);
4063 unsafe {
4064 b.launch(cfg)?;
4065 }
4066 Ok(())
4067 }
4068
4069 pub fn rows_permute(
4071 &self,
4072 src: &CudaSlice<f32>,
4073 idx: &CudaSlice<i32>,
4074 nrows: usize,
4075 ncols: usize,
4076 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4077 let mut dst = self.alloc_uninit::<f32>(nrows * ncols)?;
4078 let f = self.func("rows_permute_f32");
4079 let (nc, nr) = (ncols as i32, nrows as i32);
4080 let cfg = LaunchConfig {
4081 grid_dim: (nrows as u32, 1, 1),
4082 block_dim: (256, 1, 1),
4083 shared_mem_bytes: 0,
4084 };
4085 let __s_b = self.gpu.stream();
4086 let mut b = __s_b.launch_builder(&f);
4087 b.arg(src).arg(idx).arg(&mut dst).arg(&nc).arg(&nr);
4088 unsafe {
4089 b.launch(cfg)?;
4090 }
4091 Ok(dst)
4092 }
4093
4094 pub fn sigmoid_dot_rows(
4099 &self,
4100 x: &CudaSlice<f32>,
4101 w: &CudaSlice<f32>,
4102 n_embd: usize,
4103 t: usize,
4104 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4105 static OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4108 if *OFF.get_or_init(|| std::env::var("MEMRA_SHEXP_DOT").as_deref() == Ok("0")) {
4109 let gs = self.linear(x, w, t, n_embd, 1)?;
4110 let mut g = self.uninit(t)?;
4111 self.sigmoid(&gs, &mut g, t)?;
4112 return Ok(g);
4113 }
4114 let mut g = self.alloc_uninit::<f32>(t)?;
4120 let f = self.func("sigmoid_dot_rows_f32");
4121 let (ne, ti) = (n_embd as i32, t as i32);
4122 let cfg = LaunchConfig {
4123 grid_dim: (t as u32, 1, 1),
4124 block_dim: (32, 8, 1),
4125 shared_mem_bytes: 0,
4126 };
4127 let __s_b = self.gpu.stream();
4128 let mut b = __s_b.launch_builder(&f);
4129 b.arg(x).arg(w).arg(&mut g).arg(&ne).arg(&ti);
4130 unsafe {
4131 b.launch(cfg)?;
4132 }
4133 Ok(g)
4134 }
4135
4136 pub fn sigmoid_dot_rows_into(
4138 &self,
4139 x: &CudaSlice<f32>,
4140 w: &CudaSlice<f32>,
4141 g: &mut CudaSlice<f32>,
4142 n_embd: usize,
4143 t: usize,
4144 ) -> Result<(), Box<dyn std::error::Error>> {
4145 if g.len() < t {
4146 return Err("sigmoid_dot_rows_into output too small".into());
4147 }
4148 let f = self.func("sigmoid_dot_rows_f32");
4149 let (ne, ti) = (n_embd as i32, t as i32);
4150 let cfg = LaunchConfig {
4151 grid_dim: (t as u32, 1, 1),
4152 block_dim: (32, 8, 1),
4153 shared_mem_bytes: 0,
4154 };
4155 let __s_b = self.gpu.stream();
4156 let mut b = __s_b.launch_builder(&f);
4157 b.arg(x).arg(w).arg(&mut *g).arg(&ne).arg(&ti);
4158 unsafe {
4159 b.launch(cfg)?;
4160 }
4161 Ok(())
4162 }
4163
4164 pub fn spec_rollback_stream(
4166 &self,
4167 len_ptrs: &CudaSlice<u64>,
4168 pos_start: &CudaSlice<i32>,
4169 acc: &CudaSlice<u32>,
4170 base: usize,
4171 n_rows: usize,
4172 ) -> Result<(), Box<dyn std::error::Error>> {
4173 let f = self.func("spec_rollback_stream");
4174 let (b, nr) = (base as i32, n_rows as i32);
4175 let cfg = LaunchConfig {
4176 grid_dim: (n_rows.div_ceil(64) as u32, 1, 1),
4177 block_dim: (64, 1, 1),
4178 shared_mem_bytes: 0,
4179 };
4180 let __s_bl = self.gpu.stream();
4181 let mut bl = __s_bl.launch_builder(&f);
4182 bl.arg(len_ptrs).arg(pos_start).arg(acc).arg(&b).arg(&nr);
4183 unsafe {
4184 bl.launch(cfg)?;
4185 }
4186 Ok(())
4187 }
4188
4189 pub fn plain_tok_ring(
4191 &self,
4192 vam: &CudaSlice<u32>,
4193 pos_start: &CudaSlice<i32>,
4194 base: usize,
4195 ring: &mut CudaSlice<u32>,
4196 ) -> Result<(), Box<dyn std::error::Error>> {
4197 let f = self.func("plain_tok_ring");
4198 let (b, cap) = (base as i32, ring.len() as i32);
4199 let cfg = LaunchConfig {
4200 grid_dim: (1, 1, 1),
4201 block_dim: (32, 1, 1),
4202 shared_mem_bytes: 0,
4203 };
4204 let __s_bl = self.gpu.stream();
4205 let mut bl = __s_bl.launch_builder(&f);
4206 bl.arg(vam).arg(pos_start).arg(&b).arg(&mut *ring).arg(&cap);
4207 unsafe {
4208 bl.launch(cfg)?;
4209 }
4210 Ok(())
4211 }
4212
4213 pub fn spec_ring_commit(
4215 &self,
4216 vtok: &CudaSlice<u32>,
4217 acc: &CudaSlice<u32>,
4218 brk: &CudaSlice<u32>,
4219 ring: &mut CudaSlice<u32>,
4220 pend: &mut CudaSlice<u32>,
4221 ) -> Result<(), Box<dyn std::error::Error>> {
4222 let f = self.func("spec_ring_commit");
4223 let cfg = LaunchConfig {
4224 grid_dim: (1, 1, 1),
4225 block_dim: (32, 1, 1),
4226 shared_mem_bytes: 0,
4227 };
4228 let __s_b = self.gpu.stream();
4229 let mut b = __s_b.launch_builder(&f);
4230 b.arg(vtok).arg(acc).arg(brk).arg(ring).arg(pend);
4231 unsafe {
4232 b.launch(cfg)?;
4233 }
4234 Ok(())
4235 }
4236 pub fn i32_copy_add(
4237 &self,
4238 src: &CudaSlice<i32>,
4239 dst: &mut CudaSlice<i32>,
4240 delta: i32,
4241 ) -> Result<(), Box<dyn std::error::Error>> {
4242 let f = self.func("i32_copy_add");
4243 let cfg = LaunchConfig {
4244 grid_dim: (1, 1, 1),
4245 block_dim: (32, 1, 1),
4246 shared_mem_bytes: 0,
4247 };
4248 let __s_b = self.gpu.stream();
4249 let mut b = __s_b.launch_builder(&f);
4250 b.arg(src).arg(dst).arg(&delta);
4251 unsafe {
4252 b.launch(cfg)?;
4253 }
4254 Ok(())
4255 }
4256 pub fn u32_copy(
4257 &self,
4258 src: &CudaSlice<u32>,
4259 dst: &mut CudaSlice<u32>,
4260 ) -> Result<(), Box<dyn std::error::Error>> {
4261 let f = self.func("u32_copy");
4262 let cfg = LaunchConfig {
4263 grid_dim: (1, 1, 1),
4264 block_dim: (32, 1, 1),
4265 shared_mem_bytes: 0,
4266 };
4267 let __s_b = self.gpu.stream();
4268 let mut b = __s_b.launch_builder(&f);
4269 b.arg(src).arg(dst);
4270 unsafe {
4271 b.launch(cfg)?;
4272 }
4273 Ok(())
4274 }
4275
4276 pub fn spec_adapt_k(
4280 &self,
4281 acc: &CudaSlice<u32>,
4282 brk: &mut CudaSlice<u32>,
4283 floor: usize,
4284 cap: usize,
4285 ) -> Result<(), Box<dyn std::error::Error>> {
4286 let f = self.func("spec_adapt_k");
4287 let (fl, cp) = (floor as i32, cap as i32);
4288 let cfg = LaunchConfig {
4289 grid_dim: (1, 1, 1),
4290 block_dim: (32, 1, 1),
4291 shared_mem_bytes: 0,
4292 };
4293 let __s_b = self.gpu.stream();
4294 let mut b = __s_b.launch_builder(&f);
4295 b.arg(acc).arg(brk).arg(&fl).arg(&cp);
4296 unsafe {
4297 b.launch(cfg)?;
4298 }
4299 Ok(())
4300 }
4301
4302 pub fn spec_accept_greedy_dc(
4304 &self,
4305 preds: &CudaSlice<u32>,
4306 vtok: &CudaSlice<u32>,
4307 last_pred: &CudaSlice<u32>,
4308 brk: &CudaSlice<u32>,
4309 out: &mut CudaSlice<u32>,
4310 ) -> Result<(), Box<dyn std::error::Error>> {
4311 let f = self.func("spec_accept_greedy_dc");
4312 let cfg = LaunchConfig {
4313 grid_dim: (1, 1, 1),
4314 block_dim: (32, 1, 1),
4315 shared_mem_bytes: 0,
4316 };
4317 let __s_b = self.gpu.stream();
4318 let mut b = __s_b.launch_builder(&f);
4319 b.arg(preds).arg(vtok).arg(last_pred).arg(brk).arg(out);
4320 unsafe {
4321 b.launch(cfg)?;
4322 }
4323 Ok(())
4324 }
4325
4326 pub fn pos_iota(
4328 &self,
4329 pos0: &CudaSlice<i32>,
4330 out: &mut CudaSlice<i32>,
4331 t: usize,
4332 ) -> Result<(), Box<dyn std::error::Error>> {
4333 let f = self.func("pos_iota_i32");
4334 let ti = t as i32;
4335 let cfg = LaunchConfig {
4336 grid_dim: (1, 1, 1),
4337 block_dim: (t.max(1) as u32, 1, 1),
4338 shared_mem_bytes: 0,
4339 };
4340 let __s_b = self.gpu.stream();
4341 let mut b = __s_b.launch_builder(&f);
4342 b.arg(pos0).arg(out).arg(&ti);
4343 unsafe {
4344 b.launch(cfg)?;
4345 }
4346 Ok(())
4347 }
4348 #[allow(clippy::too_many_arguments)]
4349 pub fn append_kv_quantized_rows_dc(
4350 &self,
4351 k_rows: &CudaSlice<f32>,
4352 v_rows: &CudaSlice<f32>,
4353 kc: &mut CudaSlice<u8>,
4354 vc: &mut CudaSlice<u8>,
4355 t0_dev: &CudaSlice<i32>,
4356 t: usize,
4357 kv_dim_k: usize,
4358 kv_dim_v: usize,
4359 k_tok_bytes: usize,
4360 v_tok_bytes: usize,
4361 g: bool,
4362 ) -> Result<(), Box<dyn std::error::Error>> {
4363 let f = if g {
4364 self.func_g("append_quantize_kv_q8_0_q5_1_rows_dc")
4365 } else {
4366 self.func("append_quantize_kv_q8_0_q5_1_rows_dc")
4367 };
4368 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
4369 let cfg = LaunchConfig {
4370 grid_dim: (nblk, t as u32, 1),
4371 block_dim: (32, 1, 1),
4372 shared_mem_bytes: 0,
4373 };
4374 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
4375 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
4376 let __s_b = self.gpu.stream();
4377 let mut b = __s_b.launch_builder(&f);
4378 b.arg(k_rows)
4379 .arg(v_rows)
4380 .arg(kc)
4381 .arg(vc)
4382 .arg(t0_dev)
4383 .arg(&kdk)
4384 .arg(&kdv)
4385 .arg(&ktb)
4386 .arg(&vtb);
4387 unsafe {
4388 b.launch(cfg)?;
4389 }
4390 Ok(())
4391 }
4392
4393 #[allow(clippy::too_many_arguments)]
4396 pub fn append_kv_quantized_row_dc_inc(
4397 &self,
4398 k_row: &CudaSlice<f32>,
4399 v_row: &CudaSlice<f32>,
4400 kc: &mut CudaSlice<u8>,
4401 vc: &mut CudaSlice<u8>,
4402 t0_dev: &mut CudaSlice<i32>,
4403 kv_dim_k: usize,
4404 kv_dim_v: usize,
4405 k_tok_bytes: usize,
4406 v_tok_bytes: usize,
4407 g: bool,
4408 ) -> Result<(), Box<dyn std::error::Error>> {
4409 let f = if g {
4410 self.func_g("append_quantize_kv_q8_0_q5_1_dc_inc")
4411 } else {
4412 self.func("append_quantize_kv_q8_0_q5_1_dc_inc")
4413 };
4414 let nthreads = ((kv_dim_k.max(kv_dim_v) / 32) * 32).min(1024) as u32;
4415 let cfg = LaunchConfig {
4416 grid_dim: (1, 1, 1),
4417 block_dim: (nthreads, 1, 1),
4418 shared_mem_bytes: 0,
4419 };
4420 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
4421 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
4422 let __s_b = self.gpu.stream();
4423 let mut b = __s_b.launch_builder(&f);
4424 b.arg(k_row)
4425 .arg(v_row)
4426 .arg(kc)
4427 .arg(vc)
4428 .arg(t0_dev)
4429 .arg(&kdk)
4430 .arg(&kdv)
4431 .arg(&ktb)
4432 .arg(&vtb);
4433 unsafe {
4434 b.launch(cfg)?;
4435 }
4436 Ok(())
4437 }
4438
4439 pub fn pack_tok_p(
4441 &self,
4442 tok: &CudaSlice<u32>,
4443 p: &CudaSlice<f32>,
4444 out: &mut CudaSlice<u32>,
4445 slot: usize,
4446 ) -> Result<(), Box<dyn std::error::Error>> {
4447 let f = self.func("pack_tok_p");
4448 let sl = slot as i32;
4449 let cfg = LaunchConfig {
4450 grid_dim: (1, 1, 1),
4451 block_dim: (32, 1, 1),
4452 shared_mem_bytes: 0,
4453 };
4454 let __s_b = self.gpu.stream();
4455 let mut b = __s_b.launch_builder(&f);
4456 b.arg(tok).arg(p).arg(out).arg(&sl);
4457 unsafe {
4458 b.launch(cfg)?;
4459 }
4460 Ok(())
4461 }
4462 pub fn tok_map_u32(
4463 &self,
4464 tok: &mut CudaSlice<u32>,
4465 map: &CudaSlice<u32>,
4466 ) -> Result<(), Box<dyn std::error::Error>> {
4467 let f = self.func("tok_map_u32");
4468 let cfg = LaunchConfig {
4469 grid_dim: (1, 1, 1),
4470 block_dim: (32, 1, 1),
4471 shared_mem_bytes: 0,
4472 };
4473 let __s_b = self.gpu.stream();
4474 let mut b = __s_b.launch_builder(&f);
4475 b.arg(tok).arg(map);
4476 unsafe {
4477 b.launch(cfg)?;
4478 }
4479 Ok(())
4480 }
4481
4482 #[allow(clippy::too_many_arguments)]
4484 pub fn spec_assemble_verify(
4485 &self,
4486 tokp: &CudaSlice<u32>,
4487 pend: &CudaSlice<u32>,
4488 d2t: Option<&CudaSlice<u32>>,
4489 vtok: &mut CudaSlice<u32>,
4490 brk: &mut CudaSlice<u32>,
4491 p_min: f32,
4492 k: usize,
4493 pmin0: bool,
4494 ) -> Result<(), Box<dyn std::error::Error>> {
4495 let f = self.func("spec_assemble_verify");
4496 let (ki, pm) = (k as i32, if pmin0 { 1i32 } else { 0i32 });
4497 let cfg = LaunchConfig {
4498 grid_dim: (1, 1, 1),
4499 block_dim: (32, 1, 1),
4500 shared_mem_bytes: 0,
4501 };
4502 let __s_b = self.gpu.stream();
4503 let mut b = __s_b.launch_builder(&f);
4504 match d2t {
4505 Some(m) => {
4506 b.arg(tokp)
4507 .arg(pend)
4508 .arg(m)
4509 .arg(vtok)
4510 .arg(brk)
4511 .arg(&p_min)
4512 .arg(&ki)
4513 .arg(&pm);
4514 unsafe {
4515 b.launch(cfg)?;
4516 }
4517 }
4518 None => {
4519 let null: u64 = 0;
4520 b.arg(tokp)
4521 .arg(pend)
4522 .arg(&null)
4523 .arg(vtok)
4524 .arg(brk)
4525 .arg(&p_min)
4526 .arg(&ki)
4527 .arg(&pm);
4528 unsafe {
4529 b.launch(cfg)?;
4530 }
4531 }
4532 }
4533 Ok(())
4534 }
4535
4536 #[allow(clippy::too_many_arguments)]
4538 pub fn ssm_conv_ring_rebuild_dc(
4539 &self,
4540 qkv_tm: &CudaSlice<f32>,
4541 ring_old: &CudaSlice<f32>,
4542 conv_state: &mut CudaSlice<f32>,
4543 conv_dim: usize,
4544 acc: &CudaSlice<u32>,
4545 base: usize,
4546 t_v: usize,
4547 d_conv: usize,
4548 ) -> Result<(), Box<dyn std::error::Error>> {
4549 let f = self.func("ssm_conv_ring_rebuild_f32_dc");
4550 let n = conv_dim * (d_conv - 1);
4551 let cfg = LaunchConfig::for_num_elems(n as u32);
4552 let (cd, b0, tv, dc) = (conv_dim as i32, base as i32, t_v as i32, d_conv as i32);
4553 let __s_b = self.gpu.stream();
4554 let mut b = __s_b.launch_builder(&f);
4555 b.arg(qkv_tm)
4556 .arg(ring_old)
4557 .arg(conv_state)
4558 .arg(&cd)
4559 .arg(acc)
4560 .arg(&b0)
4561 .arg(&tv)
4562 .arg(&dc);
4563 unsafe {
4564 b.launch(cfg)?;
4565 }
4566 Ok(())
4567 }
4568 #[allow(clippy::too_many_arguments)]
4569 pub fn gdn_scan_s128_dc(
4570 &self,
4571 q: &CudaSlice<f32>,
4572 k: &CudaSlice<f32>,
4573 v: &CudaSlice<f32>,
4574 g: &CudaSlice<f32>,
4575 beta: &CudaSlice<f32>,
4576 state_in: &CudaSlice<f32>,
4577 state_out: &mut CudaSlice<f32>,
4578 o: &mut CudaSlice<f32>,
4579 n_head: usize,
4580 acc: &CudaSlice<u32>,
4581 base: usize,
4582 t_v: usize,
4583 scale: f32,
4584 ) -> Result<(), Box<dyn std::error::Error>> {
4585 let f = self.func("gdn_scan_s128_dc");
4586 const S_V: u32 = 128;
4587 const WARP: u32 = 32;
4588 const COLS_PER_BLOCK: u32 = 4;
4589 let cfg = LaunchConfig {
4590 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
4591 block_dim: (WARP, COLS_PER_BLOCK, 1),
4592 shared_mem_bytes: 0,
4593 };
4594 let (h, b0, tv) = (n_head as i32, base as i32, t_v as i32);
4595 let __s_b = self.gpu.stream();
4596 let mut b = __s_b.launch_builder(&f);
4597 b.arg(q)
4598 .arg(k)
4599 .arg(v)
4600 .arg(g)
4601 .arg(beta)
4602 .arg(state_in)
4603 .arg(state_out)
4604 .arg(o)
4605 .arg(&h)
4606 .arg(acc)
4607 .arg(&b0)
4608 .arg(&tv)
4609 .arg(&scale);
4610 unsafe {
4611 b.launch(cfg)?;
4612 }
4613 Ok(())
4614 }
4615
4616 pub fn spec_rollback_kv(
4618 &self,
4619 len_ptrs: &CudaSlice<u64>,
4620 saved: &CudaSlice<i32>,
4621 acc: &CudaSlice<u32>,
4622 base: usize,
4623 n_layer: usize,
4624 ) -> Result<(), Box<dyn std::error::Error>> {
4625 let f = self.func("spec_rollback_kv");
4626 let (b, nl) = (base as i32, n_layer as i32);
4627 let cfg = LaunchConfig {
4628 grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
4629 block_dim: (64, 1, 1),
4630 shared_mem_bytes: 0,
4631 };
4632 let __s_bl = self.gpu.stream();
4633 let mut bl = __s_bl.launch_builder(&f);
4634 bl.arg(len_ptrs).arg(saved).arg(acc).arg(&b).arg(&nl);
4635 unsafe {
4636 bl.launch(cfg)?;
4637 }
4638 Ok(())
4639 }
4640
4641 pub fn spec_fork_valid(
4643 &self,
4644 acc: &CudaSlice<u32>,
4645 optimistic_pending: u32,
4646 valid: &mut CudaSlice<u32>,
4647 ) -> Result<(), Box<dyn std::error::Error>> {
4648 let f = self.func("spec_fork_valid");
4649 let cfg = LaunchConfig {
4650 grid_dim: (1, 1, 1),
4651 block_dim: (1, 1, 1),
4652 shared_mem_bytes: 0,
4653 };
4654 let __s_bl = self.gpu.stream();
4655 let mut bl = __s_bl.launch_builder(&f);
4656 bl.arg(acc).arg(&optimistic_pending).arg(valid);
4657 unsafe {
4658 bl.launch(cfg)?;
4659 }
4660 Ok(())
4661 }
4662
4663 pub fn spec_fork_reconcile_kv(
4665 &self,
4666 len_ptrs: &CudaSlice<u64>,
4667 saved: &CudaSlice<i32>,
4668 acc: &CudaSlice<u32>,
4669 valid: &CudaSlice<u32>,
4670 base: usize,
4671 n_layer: usize,
4672 ) -> Result<(), Box<dyn std::error::Error>> {
4673 let f = self.func("spec_fork_reconcile_kv");
4674 let (b, nl) = (base as i32, n_layer as i32);
4675 let cfg = LaunchConfig {
4676 grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
4677 block_dim: (64, 1, 1),
4678 shared_mem_bytes: 0,
4679 };
4680 let __s_bl = self.gpu.stream();
4681 let mut bl = __s_bl.launch_builder(&f);
4682 bl.arg(len_ptrs)
4683 .arg(saved)
4684 .arg(acc)
4685 .arg(valid)
4686 .arg(&b)
4687 .arg(&nl);
4688 unsafe {
4689 bl.launch(cfg)?;
4690 }
4691 Ok(())
4692 }
4693
4694 pub fn spec_fork_restore_f32(
4696 &self,
4697 snapshot: &CudaSlice<f32>,
4698 state: &mut CudaSlice<f32>,
4699 valid: &CudaSlice<u32>,
4700 ) -> Result<(), Box<dyn std::error::Error>> {
4701 assert_eq!(
4702 snapshot.len(),
4703 state.len(),
4704 "fork recurrent snapshot shape mismatch"
4705 );
4706 let f = self.func("spec_fork_restore_f32");
4707 let n = state.len() as i32;
4708 #[allow(clippy::manual_clamp)]
4709 let blocks = state.len().div_ceil(256).min(65535).max(1) as u32;
4711 let cfg = LaunchConfig {
4712 grid_dim: (blocks, 1, 1),
4713 block_dim: (256, 1, 1),
4714 shared_mem_bytes: 0,
4715 };
4716 let __s_bl = self.gpu.stream();
4717 let mut bl = __s_bl.launch_builder(&f);
4718 bl.arg(snapshot).arg(state).arg(valid).arg(&n);
4719 unsafe {
4720 bl.launch(cfg)?;
4721 }
4722 Ok(())
4723 }
4724
4725 pub fn spec_seed_gather(
4728 &self,
4729 vx: &CudaSlice<f32>,
4730 fill_prev: &CudaSlice<f32>,
4731 acc: &CudaSlice<u32>,
4732 h_seed: &mut CudaSlice<f32>,
4733 base: usize,
4734 n_embd: usize,
4735 ) -> Result<(), Box<dyn std::error::Error>> {
4736 let f = self.func("spec_seed_gather");
4737 let (b, ne) = (base as i32, n_embd as i32);
4738 let cfg = LaunchConfig {
4739 grid_dim: (n_embd.div_ceil(256) as u32, 1, 1),
4740 block_dim: (256, 1, 1),
4741 shared_mem_bytes: 0,
4742 };
4743 let __s_bl = self.gpu.stream();
4744 let mut bl = __s_bl.launch_builder(&f);
4745 bl.arg(vx)
4746 .arg(fill_prev)
4747 .arg(acc)
4748 .arg(h_seed)
4749 .arg(&b)
4750 .arg(&ne);
4751 unsafe {
4752 bl.launch(cfg)?;
4753 }
4754 Ok(())
4755 }
4756
4757 pub fn spec_accept_greedy(
4759 &self,
4760 preds: &CudaSlice<u32>,
4761 draft: &CudaSlice<u32>,
4762 last_pred: u32,
4763 base: usize,
4764 k_round: usize,
4765 out: &mut CudaSlice<u32>,
4766 ) -> Result<(), Box<dyn std::error::Error>> {
4767 let f = self.func("spec_accept_greedy");
4768 let (b, k) = (base as i32, k_round as i32);
4769 let cfg = LaunchConfig {
4770 grid_dim: (1, 1, 1),
4771 block_dim: (32, 1, 1),
4772 shared_mem_bytes: 0,
4773 };
4774 let __s_bl = self.gpu.stream();
4775 let mut bl = __s_bl.launch_builder(&f);
4776 bl.arg(preds)
4777 .arg(draft)
4778 .arg(&last_pred)
4779 .arg(&b)
4780 .arg(&k)
4781 .arg(out);
4782 unsafe {
4783 bl.launch(cfg)?;
4784 }
4785 Ok(())
4786 }
4787
4788 pub fn gumbel_perturb(
4795 &self,
4796 x: &CudaSlice<f32>,
4797 y: &mut CudaSlice<f32>,
4798 n: usize,
4799 seed: u64,
4800 stream_pos: u32,
4801 temp: f32,
4802 ) -> Result<(), Box<dyn std::error::Error>> {
4803 let f = self.func("gumbel_perturb_f32");
4804 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
4805 let cfg = LaunchConfig {
4806 grid_dim: (n.div_ceil(256) as u32, 1, 1),
4807 block_dim: (256, 1, 1),
4808 shared_mem_bytes: 0,
4809 };
4810 let __s_b = self.gpu.stream();
4811 let mut b = __s_b.launch_builder(&f);
4812 b.arg(x)
4813 .arg(&mut *y)
4814 .arg(&ni)
4815 .arg(&slo)
4816 .arg(&shi)
4817 .arg(&stream_pos)
4818 .arg(&temp);
4819 unsafe {
4820 b.launch(cfg)?;
4821 }
4822 Ok(())
4823 }
4824
4825 pub fn mask_logits_col(
4833 &self,
4834 logits: &mut CudaSlice<f32>,
4835 mask: &CudaSlice<u32>,
4836 col: usize,
4837 n: usize,
4838 mask_words: usize,
4839 ) -> Result<(), Box<dyn std::error::Error>> {
4840 let f = self.func("mask_logits_f32");
4841 let (ci, ni, mw) = (col as i32, n as i32, mask_words as i32);
4842 let cfg = LaunchConfig {
4843 grid_dim: (n.div_ceil(256).min(1024) as u32, 1, 1),
4844 block_dim: (256, 1, 1),
4845 shared_mem_bytes: 0,
4846 };
4847 let __s_b = self.gpu.stream();
4848 let mut b = __s_b.launch_builder(&f);
4849 b.arg(&mut *logits).arg(mask).arg(&ci).arg(&ni).arg(&mw);
4850 unsafe {
4851 b.launch(cfg)?;
4852 }
4853 Ok(())
4854 }
4855
4856 #[allow(clippy::too_many_arguments)] pub fn gumbel_perturb_col(
4864 &self,
4865 x: &CudaSlice<f32>,
4866 col: usize,
4867 y: &mut CudaSlice<f32>,
4868 n: usize,
4869 seed: u64,
4870 stream_pos: u32,
4871 temp: f32,
4872 ) -> Result<(), Box<dyn std::error::Error>> {
4873 let f = self.func("gumbel_perturb_f32");
4874 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
4875 let col_view = x.slice(col * n..(col + 1) * n);
4876 let cfg = LaunchConfig {
4877 grid_dim: (n.div_ceil(256) as u32, 1, 1),
4878 block_dim: (256, 1, 1),
4879 shared_mem_bytes: 0,
4880 };
4881 let __s_b = self.gpu.stream();
4882 let mut b = __s_b.launch_builder(&f);
4883 b.arg(&col_view)
4884 .arg(&mut *y)
4885 .arg(&ni)
4886 .arg(&slo)
4887 .arg(&shi)
4888 .arg(&stream_pos)
4889 .arg(&temp);
4890 unsafe {
4891 b.launch(cfg)?;
4892 }
4893 Ok(())
4894 }
4895
4896 #[allow(clippy::too_many_arguments)]
4902 pub fn gumbel_perturb_filtered_col(
4903 &self,
4904 x: &CudaSlice<f32>,
4905 col: usize,
4906 y: &mut CudaSlice<f32>,
4907 n: usize,
4908 seed: u64,
4909 stream_pos: u32,
4910 temp: f32,
4911 stat_max: &CudaSlice<f32>,
4912 stat_th: &CudaSlice<f32>,
4913 stat_idx: usize,
4914 ) -> Result<(), Box<dyn std::error::Error>> {
4915 let f = self.func("gumbel_perturb_filtered_col_f32");
4916 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
4917 let (ci, si) = (col as i32, stat_idx as i32);
4918 let cfg = LaunchConfig {
4919 grid_dim: (n.div_ceil(256) as u32, 1, 1),
4920 block_dim: (256, 1, 1),
4921 shared_mem_bytes: 0,
4922 };
4923 let __s_b = self.gpu.stream();
4924 let mut b = __s_b.launch_builder(&f);
4925 b.arg(x)
4926 .arg(&ci)
4927 .arg(&mut *y)
4928 .arg(&ni)
4929 .arg(&slo)
4930 .arg(&shi)
4931 .arg(&stream_pos)
4932 .arg(&temp)
4933 .arg(stat_max)
4934 .arg(stat_th)
4935 .arg(&si);
4936 unsafe {
4937 b.launch(cfg)?;
4938 }
4939 Ok(())
4940 }
4941
4942 pub fn sctr_inc(&self, ctr: &mut CudaSlice<u32>) -> Result<(), Box<dyn std::error::Error>> {
4947 let f = self.func("memra_sctr_inc");
4948 let cfg = LaunchConfig {
4949 grid_dim: (1, 1, 1),
4950 block_dim: (1, 1, 1),
4951 shared_mem_bytes: 0,
4952 };
4953 let __s_b = self.gpu.stream();
4954 let mut b = __s_b.launch_builder(&f);
4955 b.arg(&mut *ctr);
4956 unsafe {
4957 b.launch(cfg)?;
4958 }
4959 Ok(())
4960 }
4961
4962 pub fn gumbel_perturb_ctr(
4967 &self,
4968 x: &CudaSlice<f32>,
4969 y: &mut CudaSlice<f32>,
4970 n: usize,
4971 seed: u64,
4972 ctr: &CudaSlice<u32>,
4973 temp: f32,
4974 ) -> Result<(), Box<dyn std::error::Error>> {
4975 let f = self.func("gumbel_perturb_ctr_f32");
4976 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
4977 let cfg = LaunchConfig {
4978 grid_dim: (n.div_ceil(256) as u32, 1, 1),
4979 block_dim: (256, 1, 1),
4980 shared_mem_bytes: 0,
4981 };
4982 let __s_b = self.gpu.stream();
4983 let mut b = __s_b.launch_builder(&f);
4984 b.arg(x)
4985 .arg(&mut *y)
4986 .arg(&ni)
4987 .arg(&slo)
4988 .arg(&shi)
4989 .arg(ctr)
4990 .arg(&temp);
4991 unsafe {
4992 b.launch(cfg)?;
4993 }
4994 Ok(())
4995 }
4996
4997 #[allow(clippy::too_many_arguments)]
5005 pub fn gumbel_perturb_filtered_ctr(
5006 &self,
5007 x: &CudaSlice<f32>,
5008 y: &mut CudaSlice<f32>,
5009 n: usize,
5010 seed: u64,
5011 ctr: &CudaSlice<u32>,
5012 temp: f32,
5013 stat_max: &CudaSlice<f32>,
5014 stat_th: &CudaSlice<f32>,
5015 ) -> Result<(), Box<dyn std::error::Error>> {
5016 let f = self.func("gumbel_perturb_filtered_ctr_f32");
5017 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
5018 let cfg = LaunchConfig {
5019 grid_dim: (n.div_ceil(256) as u32, 1, 1),
5020 block_dim: (256, 1, 1),
5021 shared_mem_bytes: 0,
5022 };
5023 let __s_b = self.gpu.stream();
5024 let mut b = __s_b.launch_builder(&f);
5025 b.arg(x)
5026 .arg(&mut *y)
5027 .arg(&ni)
5028 .arg(&slo)
5029 .arg(&shi)
5030 .arg(ctr)
5031 .arg(&temp)
5032 .arg(stat_max)
5033 .arg(stat_th);
5034 unsafe {
5035 b.launch(cfg)?;
5036 }
5037 Ok(())
5038 }
5039
5040 #[allow(clippy::too_many_arguments)] pub fn softmax_gather(
5045 &self,
5046 x: &CudaSlice<f32>,
5047 row_stride: usize,
5048 ids: &CudaSlice<u32>,
5049 rows: &CudaSlice<i32>,
5050 out: &mut CudaSlice<f32>,
5051 n: usize,
5052 npair: usize,
5053 temp: f32,
5054 ) -> Result<(), Box<dyn std::error::Error>> {
5055 let f = self.func("softmax_gather_f32");
5056 let (ni, rs) = (n as i32, row_stride as i64);
5057 let np = npair as i32;
5058 let cfg = LaunchConfig {
5059 grid_dim: (npair as u32, 1, 1),
5060 block_dim: (256, 1, 1),
5061 shared_mem_bytes: 0,
5062 };
5063 let __s_b = self.gpu.stream();
5064 let mut b = __s_b.launch_builder(&f);
5065 b.arg(x)
5066 .arg(&rs)
5067 .arg(ids)
5068 .arg(rows)
5069 .arg(&mut *out)
5070 .arg(&ni)
5071 .arg(&np)
5072 .arg(&temp);
5073 unsafe {
5074 b.launch(cfg)?;
5075 }
5076 Ok(())
5077 }
5078
5079 #[allow(clippy::too_many_arguments)] pub fn residual_sample(
5084 &self,
5085 p: &CudaSlice<f32>,
5086 q: Option<&CudaSlice<f32>>,
5087 n: usize,
5088 temp: f32,
5089 seed: u64,
5090 stream_pos: u32,
5091 out_tok: &mut CudaSlice<u32>,
5092 ) -> Result<(), Box<dyn std::error::Error>> {
5093 let f = self.func("residual_sample_f32");
5094 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
5095 let nth = 1024u32;
5096 let cfg = LaunchConfig {
5097 grid_dim: (1, 1, 1),
5098 block_dim: (nth, 1, 1),
5099 shared_mem_bytes: 0,
5100 };
5101 let has_q: i32 = q.is_some() as i32;
5102 let qbuf = q.unwrap_or(p); let __s_b = self.gpu.stream();
5104 let mut b = __s_b.launch_builder(&f);
5105 b.arg(p)
5106 .arg(qbuf)
5107 .arg(&has_q)
5108 .arg(&ni)
5109 .arg(&temp)
5110 .arg(&slo)
5111 .arg(&shi)
5112 .arg(&stream_pos)
5113 .arg(&mut *out_tok);
5114 unsafe {
5115 b.launch(cfg)?;
5116 }
5117 Ok(())
5118 }
5119
5120 pub fn with_moe_cache<R>(
5125 &self,
5126 max_block_bytes: usize,
5127 f: impl FnOnce(
5128 &mut crate::moe_cache::MoeSlotCache,
5129 &Engine,
5130 ) -> Result<R, Box<dyn std::error::Error>>,
5131 ) -> Result<R, Box<dyn std::error::Error>> {
5132 let mut guard = self.moe_cache.lock().unwrap();
5133 if guard.is_none() {
5134 *guard = Some(crate::moe_cache::MoeSlotCache::new(self, max_block_bytes)?);
5135 }
5136 let cache = guard.as_mut().unwrap();
5137 f(cache, self)
5138 }
5139
5140 pub fn freeze_moe_cache(&self) {
5143 if let Some(cache) = self.moe_cache.lock().unwrap().as_mut() {
5144 cache.freeze();
5145 }
5146 }
5147
5148 pub fn export_moe_residency(&self) -> Option<Vec<(u16, u8, u16)>> {
5151 self.moe_cache
5152 .lock()
5153 .unwrap()
5154 .as_ref()
5155 .map(crate::moe_cache::MoeSlotCache::export_residency)
5156 }
5157
5158 pub(crate) fn moe_cache_frozen(&self) -> bool {
5159 self.moe_cache
5160 .lock()
5161 .unwrap()
5162 .as_ref()
5163 .is_some_and(crate::moe_cache::MoeSlotCache::is_frozen)
5164 }
5165
5166 pub fn frozen_cpu_experts_prefer_tokenwise_prime(&self) -> bool {
5173 crate::cpu_experts::configured()
5174 && self.moe_cache_frozen()
5175 && std::env::var("MEMRA_CPU_EXPERT_BATCHED_PRIME").as_deref() != Ok("1")
5176 }
5177
5178 pub(crate) fn configure_moe_cache_layout(&self, block_bytes: Vec<usize>) {
5180 assert!(
5181 self.moe_cache.lock().unwrap().is_none(),
5182 "MoE cache layout configured after cache construction"
5183 );
5184 *self.moe_cache_layout.lock().unwrap() = Some(block_bytes);
5185 }
5186
5187 pub(crate) fn moe_cache_layout(&self) -> Option<Vec<usize>> {
5188 self.moe_cache_layout.lock().unwrap().clone()
5189 }
5190
5191 pub fn moe_cache_enabled() -> bool {
5193 std::env::var("MEMRA_MOE_CACHE").as_deref() != Ok("0")
5194 }
5195
5196 pub fn moe_cache_stats(&self) -> Option<(u64, u64, u64, usize)> {
5199 let guard = self.moe_cache.lock().unwrap();
5200 guard
5201 .as_ref()
5202 .map(|c| (c.hits, c.misses, c.staged_bytes, c.n_slots()))
5203 }
5204
5205 #[allow(clippy::type_complexity)] pub fn cpu_expert_stats(
5210 &self,
5211 ) -> Option<(u64, u64, u64, u64, u64, u64, u64, u64, u64, u64, u64)> {
5212 crate::cpu_experts::configured().then(crate::cpu_experts::stats)
5213 }
5214
5215 pub fn cpu_expert_predictor_stats(&self) -> (u64, u64) {
5218 crate::cpu_experts::predictor_stats()
5219 }
5220
5221 pub fn cpu_expert_exposed_wait_ns(&self) -> Option<u64> {
5222 crate::cpu_experts::configured().then(crate::cpu_experts::exposed_wait_ns)
5223 }
5224
5225 pub fn cpu_expert_gpu_residency_stats(&self) -> Option<(u64, u64, u64)> {
5228 crate::cpu_experts::configured().then(crate::cpu_experts::incomplete_gpu_residency_stats)
5229 }
5230
5231 pub fn moe_pread_stats(&self) -> Option<(u64, u64, u64, u64, u64, u64, u64)> {
5234 let guard = self.moe_cache.lock().unwrap();
5235 guard
5236 .as_ref()
5237 .and_then(|cache| cache.pread_stats())
5238 .map(|stats| {
5239 (
5240 stats.reads,
5241 stats.bytes,
5242 stats.read_errors,
5243 stats.short_reads,
5244 stats.fallbacks,
5245 stats.buffer_waits,
5246 stats.ring_full,
5247 )
5248 })
5249 }
5250
5251 pub fn spill_config_fallbacks(&self) -> u64 {
5253 crate::spill_pread::config_fallbacks()
5254 }
5255
5256 pub fn moe_cache_reset_counters(&self) {
5258 if let Some(c) = self.moe_cache.lock().unwrap().as_mut() {
5259 c.reset_counters();
5260 }
5261 }
5262
5263 pub fn htod_bytes(&self, v: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
5264 Ok(self.gpu.stream().clone_htod(v)?)
5265 }
5266
5267 pub fn htod_bytes_padded(
5271 &self,
5272 v: &[u8],
5273 pad: usize,
5274 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
5275 let mut d = self.alloc_u8_uninit(v.len() + pad)?;
5276 {
5277 let mut view = d.slice_mut(0..v.len());
5278 self.gpu.stream().memcpy_htod(v, &mut view)?;
5279 }
5280 Ok(d)
5281 }
5282
5283 pub fn copy_into(
5285 &self,
5286 dst: &mut CudaSlice<f32>,
5287 off: usize,
5288 src: &CudaSlice<f32>,
5289 len: usize,
5290 ) -> Result<(), Box<dyn std::error::Error>> {
5291 let mut view = dst.slice_mut(off..off + len);
5292 self.gpu
5293 .stream()
5294 .memcpy_dtod(&src.slice(0..len), &mut view)?;
5295 Ok(())
5296 }
5297
5298 pub fn copy_range_into(
5302 &self,
5303 dst: &mut CudaSlice<f32>,
5304 dst_off: usize,
5305 src: &CudaSlice<f32>,
5306 src_off: usize,
5307 len: usize,
5308 ) -> Result<(), Box<dyn std::error::Error>> {
5309 let mut view = dst.slice_mut(dst_off..dst_off + len);
5310 self.gpu
5311 .stream()
5312 .memcpy_dtod(&src.slice(src_off..src_off + len), &mut view)?;
5313 Ok(())
5314 }
5315
5316 pub fn copy_u8_into(
5319 &self,
5320 dst: &mut CudaSlice<u8>,
5321 off: usize,
5322 src: &CudaSlice<u8>,
5323 len: usize,
5324 ) -> Result<(), Box<dyn std::error::Error>> {
5325 let cap = dst.len();
5329 let mut view = dst.try_slice_mut(off..off + len).ok_or_else(|| {
5330 format!(
5331 "copy_u8_into dst range [{off},{}) exceeds capacity {cap}",
5332 off + len,
5333 )
5334 })?;
5335 self.gpu
5336 .stream()
5337 .memcpy_dtod(&src.slice(0..len), &mut view)?;
5338 Ok(())
5339 }
5340
5341 pub fn copy_u8_range_into(
5343 &self,
5344 dst: &mut CudaSlice<u8>,
5345 dst_off: usize,
5346 src: &CudaSlice<u8>,
5347 src_off: usize,
5348 len: usize,
5349 ) -> Result<(), Box<dyn std::error::Error>> {
5350 let cap = dst.len();
5353 let mut dst_view = dst.try_slice_mut(dst_off..dst_off + len).ok_or_else(|| {
5354 format!(
5355 "copy_u8_range_into dst range [{dst_off},{}) exceeds capacity {cap}",
5356 dst_off + len,
5357 )
5358 })?;
5359 self.gpu
5360 .stream()
5361 .memcpy_dtod(&src.slice(src_off..src_off + len), &mut dst_view)?;
5362 Ok(())
5363 }
5364
5365 #[track_caller]
5375 pub fn prepare_kv_append(
5376 &self,
5377 kv: &mut crate::cache::KvLayer,
5378 retain_from: usize,
5379 append_rows: usize,
5380 ) -> Result<usize, Box<dyn std::error::Error>> {
5381 let caller = std::panic::Location::caller();
5382 let base_before = kv.ring.as_ref().map(|r| r.base());
5383 let Some(plan) = kv
5384 .ring
5385 .as_ref()
5386 .map(|ring| ring.append_plan(kv.len, retain_from, append_rows))
5387 .transpose()
5388 .map_err(|err| -> Box<dyn std::error::Error> {
5389 format!(
5390 "{err} [append len={} retain_from={retain_from} append_rows={append_rows} base={base_before:?} called from {caller}]",
5391 kv.len
5392 )
5393 .into()
5394 })?
5395 else {
5396 return Ok(kv.len);
5397 };
5398 match plan {
5399 crate::cache::KvRingAppend::Contiguous { write_row } => Ok(write_row),
5400 crate::cache::KvRingAppend::Rebase {
5401 src_row,
5402 keep_rows,
5403 new_base,
5404 write_row,
5405 } => {
5406 if keep_rows > 0 {
5407 let k_len = keep_rows * kv.k_tok_bytes;
5408 let v_len = keep_rows * kv.v_tok_bytes;
5409 let mut k_tmp = self.alloc_u8_uninit(k_len)?;
5410 let mut v_tmp = self.alloc_u8_uninit(v_len)?;
5411 self.copy_u8_range_into(&mut k_tmp, 0, &kv.k, src_row * kv.k_tok_bytes, k_len)?;
5412 self.copy_u8_range_into(&mut v_tmp, 0, &kv.v, src_row * kv.v_tok_bytes, v_len)?;
5413 self.copy_u8_into(&mut kv.k, 0, &k_tmp, k_len)?;
5414 self.copy_u8_into(&mut kv.v, 0, &v_tmp, v_len)?;
5415 }
5416 if std::env::var("MEMRA_KV_REBASE_TRACE").as_deref() == Ok("1") {
5419 eprintln!(
5420 "[kv-rebase] new_base={new_base} keep_rows={keep_rows} len={} \
5421 retain_from={retain_from} called from {caller}",
5422 kv.len
5423 );
5424 }
5425 kv.ring.as_mut().unwrap().apply_rebase(new_base);
5426 if let Some(base_d) = kv.base_d.as_mut() {
5430 self.set_i32_one(base_d, new_base as i32)?;
5431 }
5432 Ok(write_row)
5433 }
5434 }
5435 }
5436
5437 pub fn htod_u8_into(
5440 &self,
5441 dst: &mut CudaSlice<u8>,
5442 off: usize,
5443 src: &[u8],
5444 ) -> Result<(), Box<dyn std::error::Error>> {
5445 let mut view = dst.slice_mut(off..off + src.len());
5446 self.gpu.stream().memcpy_htod(src, &mut view)?;
5447 Ok(())
5448 }
5449
5450 pub fn view<'a>(&self, b: &'a CudaSlice<f32>, len: usize) -> cudarc::driver::CudaView<'a, f32> {
5451 b.slice(0..len)
5452 }
5453
5454 pub fn view_u8_range<'a>(
5457 &self,
5458 b: &'a CudaSlice<u8>,
5459 start: usize,
5460 end: usize,
5461 ) -> cudarc::driver::CudaView<'a, u8> {
5462 b.slice(start..end)
5463 }
5464 pub fn view_u8<'a>(
5465 &self,
5466 b: &'a CudaSlice<u8>,
5467 len: usize,
5468 ) -> cudarc::driver::CudaView<'a, u8> {
5469 b.slice(0..len)
5470 }
5471
5472 #[allow(clippy::too_many_arguments)] pub fn append_kv_quantized(
5477 &self,
5478 k_row: &CudaSlice<f32>,
5479 v_row: &CudaSlice<f32>,
5480 kc: &mut CudaSlice<u8>,
5481 vc: &mut CudaSlice<u8>,
5482 t: usize,
5483 kv_dim_k: usize,
5484 kv_dim_v: usize,
5485 k_tok_bytes: usize,
5486 v_tok_bytes: usize,
5487 g: bool,
5488 ) -> Result<(), Box<dyn std::error::Error>> {
5489 let f = if g {
5490 self.func_g("append_quantize_kv_q8_0_q5_1")
5491 } else {
5492 self.func("append_quantize_kv_q8_0_q5_1")
5493 };
5494 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
5495 let cfg = LaunchConfig {
5496 grid_dim: (nblk, 1, 1),
5497 block_dim: (32, 1, 1),
5498 shared_mem_bytes: 0,
5499 };
5500 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
5501 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
5502 let __s_b = self.gpu.stream();
5503 let mut b = __s_b.launch_builder(&f);
5504 b.arg(k_row)
5505 .arg(v_row)
5506 .arg(kc)
5507 .arg(vc)
5508 .arg(&ti)
5509 .arg(&kdk)
5510 .arg(&kdv)
5511 .arg(&ktb)
5512 .arg(&vtb);
5513 unsafe {
5514 b.launch(cfg)?;
5515 }
5516 Ok(())
5517 }
5518
5519 #[allow(clippy::too_many_arguments)] pub fn append_kv_quantized_dc(
5524 &self,
5525 k_row: &CudaSlice<f32>,
5526 v_row: &CudaSlice<f32>,
5527 kc: &mut CudaSlice<u8>,
5528 vc: &mut CudaSlice<u8>,
5529 t_dev: &CudaSlice<i32>,
5530 kv_dim_k: usize,
5531 kv_dim_v: usize,
5532 k_tok_bytes: usize,
5533 v_tok_bytes: usize,
5534 g: bool,
5535 ) -> Result<(), Box<dyn std::error::Error>> {
5536 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
5537 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
5538 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
5539 if Self::pdl_on() && Self::pdl_wb_on() {
5541 use cudarc::driver::{DevicePtr, DevicePtrMut};
5542 let s = &self.gpu.stream();
5543 let (pk, _g0) = k_row.device_ptr(s);
5544 let (pv, _g1) = v_row.device_ptr(s);
5545 let (pkc, _g2) = kc.device_ptr_mut(s);
5546 let (pvc, _g3) = vc.device_ptr_mut(s);
5547 let (pt, _g4) = t_dev.device_ptr(s);
5548 let mut ps = [
5549 &pk as *const _ as *mut std::ffi::c_void,
5550 &pv as *const _ as *mut _,
5551 &pkc as *const _ as *mut _,
5552 &pvc as *const _ as *mut _,
5553 &pt as *const _ as *mut _,
5554 &kdk as *const _ as *mut _,
5555 &kdv as *const _ as *mut _,
5556 &ktb as *const _ as *mut _,
5557 &vtb as *const _ as *mut _,
5558 ];
5559 unsafe {
5560 self.launch_pdl_flash(
5561 g,
5562 "append_quantize_kv_q8_0_q5_1_dc",
5563 (nblk, 1, 1),
5564 (32, 1, 1),
5565 0,
5566 &mut ps,
5567 )?;
5568 }
5569 return Ok(());
5570 }
5571 let f = if g {
5572 self.func_g("append_quantize_kv_q8_0_q5_1_dc")
5573 } else {
5574 self.func("append_quantize_kv_q8_0_q5_1_dc")
5575 };
5576 let cfg = LaunchConfig {
5577 grid_dim: (nblk, 1, 1),
5578 block_dim: (32, 1, 1),
5579 shared_mem_bytes: 0,
5580 };
5581 let __s_b = self.gpu.stream();
5582 let mut b = __s_b.launch_builder(&f);
5583 b.arg(k_row)
5584 .arg(v_row)
5585 .arg(kc)
5586 .arg(vc)
5587 .arg(t_dev)
5588 .arg(&kdk)
5589 .arg(&kdv)
5590 .arg(&ktb)
5591 .arg(&vtb);
5592 unsafe {
5593 b.launch(cfg)?;
5594 }
5595 Ok(())
5596 }
5597
5598 #[allow(clippy::too_many_arguments)]
5605 pub fn append_kv_quantized_rows(
5606 &self,
5607 k_rows: &CudaSlice<f32>,
5608 v_rows: &CudaSlice<f32>,
5609 kc: &mut CudaSlice<u8>,
5610 vc: &mut CudaSlice<u8>,
5611 t0: usize,
5612 t: usize,
5613 kv_dim_k: usize,
5614 kv_dim_v: usize,
5615 k_tok_bytes: usize,
5616 v_tok_bytes: usize,
5617 g: bool,
5618 ) -> Result<(), Box<dyn std::error::Error>> {
5619 if std::env::var("MEMRA_PRIME_APPEND_LOOP").is_ok() {
5620 for i in 0..t {
5621 let k_row = k_rows.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
5622 let v_row = v_rows.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
5623 self.append_kv_quantized_view(
5624 &k_row,
5625 &v_row,
5626 kc,
5627 vc,
5628 t0 + i,
5629 kv_dim_k,
5630 kv_dim_v,
5631 k_tok_bytes,
5632 v_tok_bytes,
5633 g,
5634 )?;
5635 }
5636 return Ok(());
5637 }
5638 let f = if g {
5639 self.func_g("append_quantize_kv_q8_0_q5_1_rows")
5640 } else {
5641 self.func("append_quantize_kv_q8_0_q5_1_rows")
5642 };
5643 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
5644 let cfg = LaunchConfig {
5645 grid_dim: (nblk, t as u32, 1),
5646 block_dim: (32, 1, 1),
5647 shared_mem_bytes: 0,
5648 };
5649 let (t0i, kdk, kdv) = (t0 as i32, kv_dim_k as i32, kv_dim_v as i32);
5650 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
5651 let __s_b = self.gpu.stream();
5652 let mut b = __s_b.launch_builder(&f);
5653 b.arg(k_rows)
5654 .arg(v_rows)
5655 .arg(kc)
5656 .arg(vc)
5657 .arg(&t0i)
5658 .arg(&kdk)
5659 .arg(&kdv)
5660 .arg(&ktb)
5661 .arg(&vtb);
5662 unsafe {
5663 b.launch(cfg)?;
5664 }
5665 Ok(())
5666 }
5667
5668 pub fn inc_seqlen(&self, p: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
5672 let f = self.func("inc_i32");
5673 let cfg = LaunchConfig {
5674 grid_dim: (1, 1, 1),
5675 block_dim: (1, 1, 1),
5676 shared_mem_bytes: 0,
5677 };
5678 let __s_b = self.gpu.stream();
5679 let mut b = __s_b.launch_builder(&f);
5680 b.arg(p);
5681 unsafe {
5682 b.launch(cfg)?;
5683 }
5684 Ok(())
5685 }
5686
5687 #[allow(clippy::too_many_arguments)] pub fn append_kv_quantized_view(
5691 &self,
5692 k_row: &cudarc::driver::CudaView<f32>,
5693 v_row: &cudarc::driver::CudaView<f32>,
5694 kc: &mut CudaSlice<u8>,
5695 vc: &mut CudaSlice<u8>,
5696 t: usize,
5697 kv_dim_k: usize,
5698 kv_dim_v: usize,
5699 k_tok_bytes: usize,
5700 v_tok_bytes: usize,
5701 g: bool,
5702 ) -> Result<(), Box<dyn std::error::Error>> {
5703 let stream = self.gpu.stream();
5704 ensure_tensor_stream_device(k_row, &stream, "append_kv_quantized_view.k_row")?;
5705 ensure_tensor_stream_device(v_row, &stream, "append_kv_quantized_view.v_row")?;
5706 ensure_tensor_stream_device(kc, &stream, "append_kv_quantized_view.k_cache")?;
5707 ensure_tensor_stream_device(vc, &stream, "append_kv_quantized_view.v_cache")?;
5708 let f = if g {
5709 self.func_g("append_quantize_kv_q8_0_q5_1")
5710 } else {
5711 self.func("append_quantize_kv_q8_0_q5_1")
5712 };
5713 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
5714 let cfg = LaunchConfig {
5715 grid_dim: (nblk, 1, 1),
5716 block_dim: (32, 1, 1),
5717 shared_mem_bytes: 0,
5718 };
5719 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
5720 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
5721 let mut b = stream.launch_builder(&f);
5722 b.arg(k_row)
5723 .arg(v_row)
5724 .arg(kc)
5725 .arg(vc)
5726 .arg(&ti)
5727 .arg(&kdk)
5728 .arg(&kdv)
5729 .arg(&ktb)
5730 .arg(&vtb);
5731 unsafe {
5732 b.launch(cfg)?;
5733 }
5734 Ok(())
5735 }
5736
5737 pub fn copy_view_into(
5740 &self,
5741 dst: &mut CudaSlice<f32>,
5742 off: usize,
5743 src: &cudarc::driver::CudaView<f32>,
5744 len: usize,
5745 ) -> Result<(), Box<dyn std::error::Error>> {
5746 let mut view = dst.slice_mut(off..off + len);
5747 self.gpu
5748 .stream()
5749 .memcpy_dtod(&src.slice(0..len), &mut view)?;
5750 Ok(())
5751 }
5752
5753 pub fn clone_dtod(
5774 &self,
5775 src: &CudaSlice<f32>,
5776 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5777 let mut dst = self.gpu.stream().alloc_zeros::<f32>(src.len())?;
5778 self.gpu.stream().memcpy_dtod(src, &mut dst)?;
5779 Ok(dst)
5780 }
5781
5782 pub fn dtod_copy_view(
5785 &self,
5786 src: &cudarc::driver::CudaView<f32>,
5787 dst: &mut CudaSlice<f32>,
5788 ) -> Result<(), Box<dyn std::error::Error>> {
5789 self.gpu.stream().memcpy_dtod(src, dst)?;
5790 Ok(())
5791 }
5792
5793 pub fn dtod_copy_view_i8(
5795 &self,
5796 src: &cudarc::driver::CudaView<i8>,
5797 dst: &mut CudaSlice<i8>,
5798 ) -> Result<(), Box<dyn std::error::Error>> {
5799 self.gpu.stream().memcpy_dtod(src, dst)?;
5800 Ok(())
5801 }
5802
5803 pub fn dtod_copy_into(
5805 &self,
5806 src: &CudaSlice<f32>,
5807 dst: &mut CudaSlice<f32>,
5808 offset: usize,
5809 ) -> Result<(), Box<dyn std::error::Error>> {
5810 let n = src.len();
5811 let mut dv = dst.slice_mut(offset..offset + n);
5812 self.gpu.stream().memcpy_dtod(src, &mut dv)?;
5813 Ok(())
5814 }
5815
5816 pub fn copy_batch_uniform_f32(
5822 &self,
5823 table: &CudaSlice<u64>,
5824 n: usize,
5825 words: usize,
5826 ) -> Result<(), Box<dyn std::error::Error>> {
5827 if n == 0 || words == 0 {
5828 return Ok(());
5829 }
5830 debug_assert!(
5831 table.len() >= 2 * n,
5832 "pointer table must hold n srcs + n dsts"
5833 );
5834 let f = self.func("copy_batch_uniform_f32");
5835 let chunks = (words / 4).max(1).div_ceil(256).min(48) as u32;
5838 let (ni, wi) = (n as i32, words as i32);
5839 let cfg = LaunchConfig {
5840 grid_dim: (chunks, n as u32, 1),
5841 block_dim: (256, 1, 1),
5842 shared_mem_bytes: 0,
5843 };
5844 let __s = self.gpu.stream();
5845 let mut b = __s.launch_builder(&f);
5846 b.arg(table).arg(&ni).arg(&wi);
5847 unsafe {
5848 b.launch(cfg)?;
5849 }
5850 Ok(())
5851 }
5852
5853 #[allow(clippy::too_many_arguments)] pub fn copy_batch_uniform_kv_u8_set_len(
5858 &self,
5859 table: &CudaSlice<u64>,
5860 n: usize,
5861 rows: usize,
5862 k_row_bytes: usize,
5863 v_row_bytes: usize,
5864 k_src_stride: usize,
5865 v_src_stride: usize,
5866 logical_len: usize,
5867 ) -> Result<(), Box<dyn std::error::Error>> {
5868 if n == 0 || rows == 0 || (k_row_bytes == 0 && v_row_bytes == 0) {
5869 return Ok(());
5870 }
5871 if table.len() < 5 * n {
5872 return Err(format!(
5873 "TP KV repair table has {} words, expected at least {}",
5874 table.len(),
5875 5 * n
5876 )
5877 .into());
5878 }
5879 let ni = i32::try_from(n).map_err(|_| "TP KV repair layer count exceeds i32")?;
5880 let rows = i32::try_from(rows).map_err(|_| "TP KV repair rows exceed i32")?;
5881 let kb = i32::try_from(k_row_bytes).map_err(|_| "TP KV repair K bytes exceed i32")?;
5882 let vb = i32::try_from(v_row_bytes).map_err(|_| "TP KV repair V bytes exceed i32")?;
5883 let ks = i32::try_from(k_src_stride).map_err(|_| "TP KV repair K stride exceeds i32")?;
5884 let vs = i32::try_from(v_src_stride).map_err(|_| "TP KV repair V stride exceeds i32")?;
5885 let len = i32::try_from(logical_len).map_err(|_| "TP KV repair length exceeds i32")?;
5886 let f = self.func("copy_batch_uniform_kv_u8_set_len");
5887 let cfg = LaunchConfig {
5888 grid_dim: (n as u32, 1, 1),
5889 block_dim: (256, 1, 1),
5890 shared_mem_bytes: 0,
5891 };
5892 let stream = self.gpu.stream();
5893 let mut builder = stream.launch_builder(&f);
5894 builder
5895 .arg(table)
5896 .arg(&ni)
5897 .arg(&rows)
5898 .arg(&kb)
5899 .arg(&vb)
5900 .arg(&ks)
5901 .arg(&vs)
5902 .arg(&len);
5903 unsafe {
5904 builder.launch(cfg)?;
5905 }
5906 Ok(())
5907 }
5908
5909 pub fn htod_u64_into(
5912 &self,
5913 v: &[u64],
5914 dst: &mut CudaSlice<u64>,
5915 ) -> Result<(), Box<dyn std::error::Error>> {
5916 let mut view = dst.slice_mut(0..v.len());
5917 self.gpu.stream().memcpy_htod(v, &mut view)?;
5918 Ok(())
5919 }
5920
5921 pub fn htod_f32_into(
5924 &self,
5925 v: &[f32],
5926 dst: &mut CudaSlice<f32>,
5927 ) -> Result<(), Box<dyn std::error::Error>> {
5928 let mut view = dst.slice_mut(0..v.len());
5929 self.gpu.stream().memcpy_htod(v, &mut view)?;
5930 Ok(())
5931 }
5932
5933 pub fn htod_f32_into_at(
5937 &self,
5938 v: &[f32],
5939 dst: &mut CudaSlice<f32>,
5940 off: usize,
5941 ) -> Result<(), Box<dyn std::error::Error>> {
5942 if off + v.len() > dst.len() {
5943 return Err(format!(
5944 "htod_f32_into_at range {}..{} exceeds dst {}",
5945 off,
5946 off + v.len(),
5947 dst.len()
5948 )
5949 .into());
5950 }
5951 let mut view = dst.slice_mut(off..off + v.len());
5952 self.gpu.stream().memcpy_htod(v, &mut view)?;
5953 Ok(())
5954 }
5955
5956 pub fn copy_indirect_src_f32(
5961 &self,
5962 src_entry: &cudarc::driver::CudaView<u64>,
5963 dst: &mut CudaSlice<f32>,
5964 dst_off: usize,
5965 words: usize,
5966 ) -> Result<(), Box<dyn std::error::Error>> {
5967 let f = self.func("copy_indirect_src_f32");
5968 let chunks = (words / 4).max(1).div_ceil(256).min(48) as u32;
5969 let wi = words as i32;
5970 let cfg = LaunchConfig {
5971 grid_dim: (chunks, 1, 1),
5972 block_dim: (256, 1, 1),
5973 shared_mem_bytes: 0,
5974 };
5975 let mut dv = dst.slice_mut(dst_off..dst_off + words);
5976 let __s = self.gpu.stream();
5977 let mut b = __s.launch_builder(&f);
5978 b.arg(src_entry).arg(&mut dv).arg(&wi);
5979 unsafe {
5980 b.launch(cfg)?;
5981 }
5982 Ok(())
5983 }
5984
5985 pub fn uninit_i8(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
5987 self.alloc_uninit::<i8>(n)
5988 }
5989
5990 #[allow(clippy::too_many_arguments)] pub fn qmatvec(
5993 &self,
5994 w: &CudaSlice<u8>,
5995 x: &CudaSlice<f32>,
5996 m: usize,
5997 in_f: usize,
5998 out_f: usize,
5999 qtype: i32,
6000 row_bytes: usize,
6001 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6002 let f = self.func("qmatvec_f32");
6003 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
6005 grid_dim: (out_f as u32, m as u32, 1),
6006 block_dim: (256, 1, 1),
6007 shared_mem_bytes: 0,
6008 };
6009 let (inf, outf, mi, qt, rb) =
6010 (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
6011 let __s_b = self.gpu.stream();
6012 let mut b = __s_b.launch_builder(&f);
6013 b.arg(w)
6014 .arg(x)
6015 .arg(&mut y)
6016 .arg(&inf)
6017 .arg(&outf)
6018 .arg(&mi)
6019 .arg(&qt)
6020 .arg(&rb);
6021 unsafe {
6022 b.launch(cfg)?;
6023 }
6024 Ok(y)
6025 }
6026
6027 pub fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
6029 let s = self.gpu.stream().alloc_zeros::<u8>(n)?;
6030 self.keep_if_capturing(&s);
6031 Ok(s)
6032 }
6033
6034 pub fn alloc_u8_uninit(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
6038 let s = unsafe { self.gpu.stream().alloc::<u8>(n)? };
6039 self.keep_if_capturing(&s);
6040 Ok(s)
6041 }
6042
6043 pub fn memset_zeros_view(
6046 &self,
6047 dst: &mut cudarc::driver::CudaViewMut<f32>,
6048 ) -> Result<(), Box<dyn std::error::Error>> {
6049 self.gpu.stream().memset_zeros(dst)?;
6050 Ok(())
6051 }
6052
6053 pub fn stage_expert(
6059 &self,
6060 host_bytes: &[u8],
6061 scratch: &mut CudaSlice<u8>,
6062 off: usize,
6063 ) -> Result<(), Box<dyn std::error::Error>> {
6064 let mut dst = scratch.slice_mut(off..off + host_bytes.len()); self.gpu.stream().memcpy_htod(host_bytes, &mut dst)?; Ok(())
6067 }
6068
6069 pub fn moe_router_topk(
6075 &self,
6076 logits: &CudaSlice<f32>,
6077 t: usize,
6078 n_expert: usize,
6079 n_used: usize,
6080 ) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6081 let f = self.func("moe_router_topk_f32");
6082 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 {
6085 grid_dim: (t as u32, 1, 1),
6086 block_dim: (n_expert as u32, 1, 1),
6087 shared_mem_bytes: 0,
6088 };
6089 let (ne, nu) = (n_expert as i32, n_used as i32);
6090 let __s_b = self.gpu.stream();
6091 let mut b = __s_b.launch_builder(&f);
6092 b.arg(logits)
6093 .arg(&mut sel_idx)
6094 .arg(&mut sel_w)
6095 .arg(&ne)
6096 .arg(&nu);
6097 unsafe {
6098 b.launch(cfg)?;
6099 }
6100 Ok((sel_idx, sel_w))
6101 }
6102
6103 pub fn moe_router_topk_scaled(
6106 &self,
6107 logits: &CudaSlice<f32>,
6108 t: usize,
6109 n_expert: usize,
6110 n_used: usize,
6111 ex_scale: &CudaSlice<f32>,
6112 ) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6113 let f = self.func("moe_router_topk_scaled_f32");
6118 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
6119 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
6120 let cfg = LaunchConfig {
6121 grid_dim: (t as u32, 1, 1),
6122 block_dim: (n_expert as u32, 1, 1),
6123 shared_mem_bytes: 0,
6124 };
6125 let (ne, nu) = (n_expert as i32, n_used as i32);
6126 let __s_b = self.gpu.stream();
6127 let mut b = __s_b.launch_builder(&f);
6128 b.arg(logits)
6129 .arg(&mut sel_idx)
6130 .arg(&mut sel_w)
6131 .arg(&ne)
6132 .arg(&nu)
6133 .arg(ex_scale);
6134 unsafe {
6135 b.launch(cfg)?;
6136 }
6137 Ok((sel_idx, sel_w))
6138 }
6139
6140 pub fn moe_router_topk_host(
6148 &self,
6149 logits: &CudaSlice<f32>,
6150 t: usize,
6151 n_expert: usize,
6152 n_used: usize,
6153 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
6154 let f = self.func("moe_router_topk_f32");
6155 let n = t * n_used;
6156 let mut sel_idx = self.alloc_uninit::<i32>(n)?;
6157 let mut sel_w = self.alloc_uninit::<f32>(n)?;
6158 let cfg = LaunchConfig {
6159 grid_dim: (t as u32, 1, 1),
6160 block_dim: (n_expert as u32, 1, 1),
6161 shared_mem_bytes: 0,
6162 };
6163 let (ne, nu) = (n_expert as i32, n_used as i32);
6164 let __s_b = self.gpu.stream();
6165 let mut b = __s_b.launch_builder(&f);
6166 b.arg(logits)
6167 .arg(&mut sel_idx)
6168 .arg(&mut sel_w)
6169 .arg(&ne)
6170 .arg(&nu);
6171 unsafe {
6172 b.launch(cfg)?;
6173 }
6174 let bytes = n * 8;
6176 let mut guard = self.router_stage.lock().unwrap();
6177 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
6178 *guard = Some(PinnedStage::new(bytes.max(4096))?);
6179 }
6180 let stage = guard.as_mut().unwrap();
6181 let (si, sw) = unsafe {
6182 (
6183 std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
6184 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n),
6185 )
6186 };
6187 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()))
6191 }
6192
6193 #[allow(clippy::too_many_arguments)]
6197 pub fn moe_router_sigmoid_topk(
6198 &self,
6199 logits: &CudaSlice<f32>,
6200 t: usize,
6201 n_expert: usize,
6202 n_used: usize,
6203 active_count: usize,
6204 correction_bias: &CudaSlice<f32>,
6205 active: &CudaSlice<u8>,
6206 scaling_factor: f32,
6207 route_norm: bool,
6208 ) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6209 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
6210 if n_expert == 0 || n_expert > 1024 || n_used == 0 || n_used > n_expert {
6211 return Err(format!(
6212 "sigmoid router shape unsupported: n_expert={n_expert}, n_used={n_used}",
6213 )
6214 .into());
6215 }
6216 if logits.len() < t * n_expert
6217 || correction_bias.len() != n_expert
6218 || active.len() != n_expert
6219 {
6220 return Err(format!(
6221 "sigmoid router buffer mismatch: logits={} bias={} active={} expected logits>={} row={}",
6222 logits.len(), correction_bias.len(), active.len(), t * n_expert, n_expert,
6223 ).into());
6224 }
6225 let f = self.func(crate::sigmoid_topk_kernel(
6226 crate::sig_expf_dev_on(),
6227 crate::topk_fast_on(),
6228 n_used,
6229 ));
6230 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
6231 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
6232 let threads = n_expert.div_ceil(32) * 32;
6233 let cfg = LaunchConfig {
6234 grid_dim: (t as u32, 1, 1),
6235 block_dim: (threads as u32, 1, 1),
6236 shared_mem_bytes: 0,
6237 };
6238 let (ne, nu, rn) = (n_expert as i32, n_used as i32, i32::from(route_norm));
6239 let __s_b = self.gpu.stream();
6240 let mut b = __s_b.launch_builder(&f);
6241 b.arg(logits)
6242 .arg(correction_bias)
6243 .arg(active)
6244 .arg(&mut sel_idx)
6245 .arg(&mut sel_w)
6246 .arg(&ne)
6247 .arg(&nu)
6248 .arg(&scaling_factor)
6249 .arg(&rn);
6250 unsafe {
6251 b.launch(cfg)?;
6252 }
6253 Ok((sel_idx, sel_w))
6254 }
6255
6256 #[allow(clippy::too_many_arguments)]
6259 pub fn ring_flag_raw(&self, ptr: u64, value: u32) -> Result<(), Box<dyn std::error::Error>> {
6263 if ptr == 0 {
6264 return Err("ring_flag_raw: unarmed flag".into());
6265 }
6266 let f = self.func("memra_ring_flag");
6267 let cfg = LaunchConfig {
6268 grid_dim: (1, 1, 1),
6269 block_dim: (32, 1, 1),
6270 shared_mem_bytes: 0,
6271 };
6272 let __s_b = self.gpu.stream();
6273 let mut b = __s_b.launch_builder(&f);
6274 b.arg(&ptr).arg(&value);
6275 unsafe {
6276 b.launch(cfg)?;
6277 }
6278 Ok(())
6279 }
6280
6281 pub fn moe_sel_w_mirror(
6284 &self,
6285 sel_src: &CudaSlice<i32>,
6286 w_src: &CudaSlice<f32>,
6287 sel_dst: &mut CudaSlice<i32>,
6288 w_dst: &mut CudaSlice<f32>,
6289 n: usize,
6290 ) -> Result<(), Box<dyn std::error::Error>> {
6291 if n == 0
6292 || n > i32::MAX as usize
6293 || sel_src.len() < n
6294 || w_src.len() < n
6295 || sel_dst.len() < n
6296 || w_dst.len() < n
6297 {
6298 return Err(format!("moe_sel_w_mirror geometry n={n}").into());
6299 }
6300 let f = self.func("moe_sel_w_mirror");
6301 let threads = if n <= 32 { 32 } else { 128 };
6302 let cfg = LaunchConfig {
6303 grid_dim: ((n as u32).div_ceil(threads), 1, 1),
6304 block_dim: (threads, 1, 1),
6305 shared_mem_bytes: 0,
6306 };
6307 let ni = n as i32;
6308 let __s_b = self.gpu.stream();
6309 let mut b = __s_b.launch_builder(&f);
6310 b.arg(sel_src).arg(w_src).arg(sel_dst).arg(w_dst).arg(&ni);
6311 unsafe {
6312 b.launch(cfg)?;
6313 }
6314 Ok(())
6315 }
6316
6317 #[allow(clippy::too_many_arguments)]
6321 pub fn nvfp4_ep_stage_inputs(
6322 &self,
6323 input_src: &CudaSlice<f32>,
6324 sel_src: &CudaSlice<i32>,
6325 w_src: &CudaSlice<f32>,
6326 input_bf16_dst: &mut CudaSlice<u8>,
6327 sel_dst: &mut CudaSlice<i32>,
6328 w_dst: &mut CudaSlice<f32>,
6329 input_values: usize,
6330 pairs: usize,
6331 copy_weights: bool,
6332 ) -> Result<(), Box<dyn std::error::Error>> {
6333 if input_values == 0
6334 || pairs == 0
6335 || input_src.len() < input_values
6336 || sel_src.len() < pairs
6337 || w_src.len() < pairs
6338 || input_bf16_dst.len() < 2 * input_values
6339 || sel_dst.len() < pairs
6340 || w_dst.len() < pairs
6341 {
6342 return Err(format!(
6343 "W4A16 EP stage geometry input={} sel={} weights={} input_bf16={} \
6344 sel_dst={} weights_dst={} active={input_values} pairs={pairs}",
6345 input_src.len(),
6346 sel_src.len(),
6347 w_src.len(),
6348 input_bf16_dst.len(),
6349 sel_dst.len(),
6350 w_dst.len(),
6351 )
6352 .into());
6353 }
6354 let f = self.func("nvfp4_ep_stage_inputs");
6355 let n = input_values.max(pairs);
6356 let cfg = LaunchConfig::for_num_elems(n as u32);
6357 let (input_values, pairs, copy_weights) =
6358 (input_values as i32, pairs as i32, i32::from(copy_weights));
6359 let __s_b = self.gpu.stream();
6360 let mut b = __s_b.launch_builder(&f);
6361 b.arg(input_src)
6362 .arg(sel_src)
6363 .arg(w_src)
6364 .arg(input_bf16_dst)
6365 .arg(sel_dst)
6366 .arg(w_dst)
6367 .arg(&input_values)
6368 .arg(&pairs)
6369 .arg(©_weights);
6370 unsafe {
6371 b.launch(cfg)?;
6372 }
6373 Ok(())
6374 }
6375
6376 #[allow(clippy::too_many_arguments)]
6379 pub fn nvfp4_ep_stage_inputs_raw(
6380 &self,
6381 input_src: u64,
6382 sel_src: u64,
6383 w_src: u64,
6384 input_bf16_dst: &mut CudaSlice<u8>,
6385 sel_dst: &mut CudaSlice<i32>,
6386 w_dst: &mut CudaSlice<f32>,
6387 input_values: usize,
6388 pairs: usize,
6389 copy_weights: bool,
6390 ) -> Result<(), Box<dyn std::error::Error>> {
6391 if input_src == 0
6392 || sel_src == 0
6393 || w_src == 0
6394 || input_values == 0
6395 || pairs == 0
6396 || input_bf16_dst.len() < 2 * input_values
6397 || sel_dst.len() < pairs
6398 || w_dst.len() < pairs
6399 {
6400 return Err(format!(
6401 "W4A16 EP raw stage geometry input={input_src:#x} sel={sel_src:#x} \
6402 weights={w_src:#x} input_bf16={} sel_dst={} weights_dst={} \
6403 active={input_values} pairs={pairs}",
6404 input_bf16_dst.len(),
6405 sel_dst.len(),
6406 w_dst.len(),
6407 )
6408 .into());
6409 }
6410 let f = self.func("nvfp4_ep_stage_inputs");
6411 let n = input_values.max(pairs);
6412 let cfg = LaunchConfig::for_num_elems(n as u32);
6413 let (input_values, pairs, copy_weights) =
6414 (input_values as i32, pairs as i32, i32::from(copy_weights));
6415 let __s_b = self.gpu.stream();
6416 let mut b = __s_b.launch_builder(&f);
6417 b.arg(&input_src)
6418 .arg(&sel_src)
6419 .arg(&w_src)
6420 .arg(input_bf16_dst)
6421 .arg(sel_dst)
6422 .arg(w_dst)
6423 .arg(&input_values)
6424 .arg(&pairs)
6425 .arg(©_weights);
6426 unsafe {
6427 b.launch(cfg)?;
6428 }
6429 Ok(())
6430 }
6431
6432 #[allow(clippy::too_many_arguments)] pub fn moe_router_sigmoid_topk_into(
6434 &self,
6435 logits: &CudaSlice<f32>,
6436 t: usize,
6437 n_expert: usize,
6438 n_used: usize,
6439 active_count: usize,
6440 correction_bias: &CudaSlice<f32>,
6441 active: &CudaSlice<u8>,
6442 scaling_factor: f32,
6443 route_norm: bool,
6444 sel_idx: &mut CudaSlice<i32>,
6445 sel_w: &mut CudaSlice<f32>,
6446 ) -> Result<(), Box<dyn std::error::Error>> {
6447 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
6448 if n_expert == 0
6449 || n_expert > 1024
6450 || n_used == 0
6451 || n_used > 32 || n_used > n_expert
6453 || logits.len() < t * n_expert
6454 || correction_bias.len() != n_expert
6455 || active.len() != n_expert
6456 || sel_idx.len() < t * n_used
6457 || sel_w.len() < t * n_used
6458 {
6459 return Err("sigmoid router _into geometry mismatch".into());
6460 }
6461 let f = self.func(crate::sigmoid_topk_kernel(
6462 crate::sig_expf_dev_on(),
6463 crate::topk_fast_on(),
6464 n_used,
6465 ));
6466 let threads = n_expert.div_ceil(32) * 32;
6467 let cfg = LaunchConfig {
6468 grid_dim: (t as u32, 1, 1),
6469 block_dim: (threads as u32, 1, 1),
6470 shared_mem_bytes: 0,
6471 };
6472 let (ne, nu, rn) = (n_expert as i32, n_used as i32, i32::from(route_norm));
6473 let __s_b = self.gpu.stream();
6474 let mut b = __s_b.launch_builder(&f);
6475 b.arg(logits)
6476 .arg(correction_bias)
6477 .arg(active)
6478 .arg(&mut *sel_idx)
6479 .arg(&mut *sel_w)
6480 .arg(&ne)
6481 .arg(&nu)
6482 .arg(&scaling_factor)
6483 .arg(&rn);
6484 unsafe {
6485 b.launch(cfg)?;
6486 }
6487 Ok(())
6488 }
6489
6490 #[allow(clippy::too_many_arguments)]
6493 pub fn moe_router_sigmoid_topk_host(
6494 &self,
6495 logits: &CudaSlice<f32>,
6496 t: usize,
6497 n_expert: usize,
6498 n_used: usize,
6499 active_count: usize,
6500 correction_bias: &CudaSlice<f32>,
6501 active: &CudaSlice<u8>,
6502 scaling_factor: f32,
6503 route_norm: bool,
6504 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
6505 let (sel_idx, sel_w) = self.moe_router_sigmoid_topk(
6506 logits,
6507 t,
6508 n_expert,
6509 n_used,
6510 active_count,
6511 correction_bias,
6512 active,
6513 scaling_factor,
6514 route_norm,
6515 )?;
6516 let n = t * n_used;
6517 let bytes = n * 8;
6518 let mut guard = self.router_stage.lock().unwrap();
6519 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
6520 *guard = Some(PinnedStage::new(bytes.max(4096))?);
6521 }
6522 let stage = guard.as_mut().unwrap();
6523 let (si, sw) = unsafe {
6524 (
6525 std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
6526 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n),
6527 )
6528 };
6529 self.gpu.stream().memcpy_dtoh(&sel_idx, si)?;
6530 self.gpu.stream().memcpy_dtoh(&sel_w, sw)?;
6531 self.gpu.stream().synchronize()?;
6532 Ok((si.iter().map(|&i| i as u32).collect(), sw.to_vec()))
6533 }
6534
6535 pub fn stage_expert_async(
6539 &self,
6540 host_bytes: &[u8],
6541 scratch: &mut CudaSlice<u8>,
6542 off: usize,
6543 ) -> Result<cudarc::driver::CudaEvent, Box<dyn std::error::Error>> {
6544 let mut dst = scratch.slice_mut(off..off + host_bytes.len());
6545 self.copy_stream.memcpy_htod(host_bytes, &mut dst)?;
6546 Ok(self.copy_stream.record_event(None)?)
6547 }
6548
6549 pub fn compute_wait(
6551 &self,
6552 ev: &cudarc::driver::CudaEvent,
6553 ) -> Result<(), Box<dyn std::error::Error>> {
6554 self.gpu.stream().wait(ev)?;
6555 Ok(())
6556 }
6557
6558 #[allow(clippy::too_many_arguments)] pub fn qmatvec_view(
6564 &self,
6565 w: &CudaSlice<u8>,
6566 range: std::ops::Range<usize>,
6567 x: &cudarc::driver::CudaView<f32>,
6568 m: usize,
6569 in_f: usize,
6570 out_f: usize,
6571 qtype: i32,
6572 row_bytes: usize,
6573 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6574 self.qmatvec_view_inner(w, range, x, m, in_f, out_f, qtype, row_bytes)
6575 }
6576
6577 #[allow(clippy::too_many_arguments)]
6581 pub fn qmatvec_view_bf16_activation(
6582 &self,
6583 w: &CudaSlice<u8>,
6584 range: std::ops::Range<usize>,
6585 x: &cudarc::driver::CudaView<f32>,
6586 m: usize,
6587 in_f: usize,
6588 out_f: usize,
6589 qtype: i32,
6590 row_bytes: usize,
6591 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6592 let n = m * in_f;
6593 if x.len() != n {
6594 return Err(format!(
6595 "W4A16 BF16 activation input length {} != {m}x{in_f}",
6596 x.len()
6597 )
6598 .into());
6599 }
6600 let mut x_bf16 = self.alloc_u8_uninit(n * 2)?;
6601 self.f32_to_bf16_v(x, &mut x_bf16, n)?;
6602 let x_f32 = self.bf16_to_f32(&x_bf16.slice(0..n * 2), n)?;
6603 self.qmatvec_view_inner(
6604 w,
6605 range,
6606 &x_f32.slice(0..n),
6607 m,
6608 in_f,
6609 out_f,
6610 qtype,
6611 row_bytes,
6612 )
6613 }
6614
6615 #[allow(clippy::too_many_arguments)]
6616 fn qmatvec_view_inner(
6617 &self,
6618 w: &CudaSlice<u8>,
6619 range: std::ops::Range<usize>,
6620 x: &cudarc::driver::CudaView<f32>,
6621 m: usize,
6622 in_f: usize,
6623 out_f: usize,
6624 qtype: i32,
6625 row_bytes: usize,
6626 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6627 let f = self.func("qmatvec_f32");
6628 let wv = w.slice(range); let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
6631 grid_dim: (out_f as u32, m as u32, 1),
6632 block_dim: (256, 1, 1),
6633 shared_mem_bytes: 0,
6634 };
6635 let (inf, outf, mi, qt, rb) =
6636 (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
6637 let __s_b = self.gpu.stream();
6638 let mut b = __s_b.launch_builder(&f);
6639 b.arg(&wv)
6640 .arg(x)
6641 .arg(&mut y)
6642 .arg(&inf)
6643 .arg(&outf)
6644 .arg(&mi)
6645 .arg(&qt)
6646 .arg(&rb);
6647 unsafe {
6648 b.launch(cfg)?;
6649 }
6650 Ok(y)
6651 }
6652
6653 #[allow(clippy::too_many_arguments)]
6660 pub fn moe_gate_up_silu8_q8(
6664 &self,
6665 gp: WPtr8,
6666 up: WPtr8,
6667 aq: &CudaSlice<i8>,
6668 ad: &CudaSlice<f32>,
6669 in_f: usize,
6670 n_ff: usize,
6671 n_used: usize,
6672 qt_g: i32,
6673 qt_u: i32,
6674 rb_g: usize,
6675 rb_u: usize,
6676 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6677 let f = self.func("moe_gate_up_silu8_q8");
6678 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
6679 let cfg = LaunchConfig {
6680 grid_dim: (n_ff as u32, n_used as u32, 1),
6681 block_dim: (32, 1, 1),
6682 shared_mem_bytes: 0,
6683 };
6684 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
6685 let __s_b = self.gpu.stream();
6686 let mut b = __s_b.launch_builder(&f);
6687 b.arg(&gp)
6688 .arg(&up)
6689 .arg(aq)
6690 .arg(ad)
6691 .arg(&mut act)
6692 .arg(&inf)
6693 .arg(&nff)
6694 .arg(&qt_g)
6695 .arg(&qt_u)
6696 .arg(&rbg)
6697 .arg(&rbu);
6698 unsafe {
6699 b.launch(cfg)?;
6700 }
6701 Ok(act)
6702 }
6703
6704 #[allow(clippy::too_many_arguments)]
6717 pub fn moe_gate_up_preclamp8_q8(
6718 &self,
6719 gp: WPtr8,
6720 up: WPtr8,
6721 aq: &CudaSlice<i8>,
6722 ad: &CudaSlice<f32>,
6723 gs: F32x8,
6724 us: F32x8,
6725 limit: f32,
6726 in_f: usize,
6727 n_ff: usize,
6728 n_used: usize,
6729 qt_g: i32,
6730 qt_u: i32,
6731 rb_g: usize,
6732 rb_u: usize,
6733 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6734 debug_assert!(
6735 limit > 1e-6,
6736 "moe_gate_up_preclamp8_q8 needs a live limit; use moe_gate_up_silu8_q8"
6737 );
6738 let f = self.func("moe_gate_up_preclamp8_q8");
6739 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
6740 let cfg = LaunchConfig {
6741 grid_dim: (n_ff as u32, n_used as u32, 1),
6742 block_dim: (32, 1, 1),
6743 shared_mem_bytes: 0,
6744 };
6745 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
6746 let __s_b = self.gpu.stream();
6747 let mut b = __s_b.launch_builder(&f);
6748 b.arg(&gp)
6749 .arg(&up)
6750 .arg(aq)
6751 .arg(ad)
6752 .arg(&gs)
6753 .arg(&us)
6754 .arg(&limit)
6755 .arg(&mut act)
6756 .arg(&inf)
6757 .arg(&nff)
6758 .arg(&qt_g)
6759 .arg(&qt_u)
6760 .arg(&rbg)
6761 .arg(&rbu);
6762 unsafe {
6763 b.launch(cfg)?;
6764 }
6765 Ok(act)
6766 }
6767
6768 #[allow(clippy::too_many_arguments)]
6769 pub fn moe_down8_fma_q8(
6770 &self,
6771 dp: WPtr8,
6772 w: F32x8,
6773 aq2: &CudaSlice<i8>,
6774 ad2: &CudaSlice<f32>,
6775 dst: &mut cudarc::driver::CudaViewMut<f32>,
6776 in_f: usize,
6777 out_f: usize,
6778 n_used: usize,
6779 qt: i32,
6780 rb: usize,
6781 ) -> Result<(), Box<dyn std::error::Error>> {
6782 let f = self.func("moe_down8_fma_q8");
6783 let cfg = LaunchConfig {
6784 grid_dim: (out_f as u32, 1, 1),
6785 block_dim: (32, 1, 1),
6786 shared_mem_bytes: 0,
6787 };
6788 let (inf, outf, nu, rbi) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
6789 let __s_b = self.gpu.stream();
6790 let mut b = __s_b.launch_builder(&f);
6791 b.arg(&dp)
6792 .arg(&w)
6793 .arg(aq2)
6794 .arg(ad2)
6795 .arg(dst)
6796 .arg(&inf)
6797 .arg(&outf)
6798 .arg(&nu)
6799 .arg(&qt)
6800 .arg(&rbi);
6801 unsafe {
6802 b.launch(cfg)?;
6803 }
6804 Ok(())
6805 }
6806
6807 #[allow(clippy::too_many_arguments)]
6819 pub fn moe_vrows_tables_from_sel(
6821 &self,
6822 sel: &CudaSlice<i32>,
6823 selw: &CudaSlice<f32>,
6824 il: u16,
6825 macros: Option<(&[f32], &[f32], &[f32])>,
6826 (pg, pu, pd): (u64, u64, u64),
6827 (sg, su, sd): (usize, usize, usize),
6828 n_pairs: usize,
6829 ptrs: &mut CudaSlice<u64>,
6830 scl: &mut CudaSlice<f32>,
6831 ) -> Result<(), Box<dyn std::error::Error>> {
6832 debug_assert!(sel.len() >= n_pairs && selw.len() >= n_pairs);
6833 debug_assert!(ptrs.len() >= 3 * n_pairs);
6835 debug_assert_eq!(scl.len(), 3 * n_pairs);
6836 let mut mac = self
6839 .vrows_macro_dev
6840 .lock()
6841 .map_err(|_| "vrows macro mirror map is poisoned")?;
6842 if let Some((hg, hu, hd)) = macros {
6843 for (plane, host) in [(0u8, hg), (1u8, hu), (2u8, hd)] {
6844 if let std::collections::hash_map::Entry::Vacant(slot) = mac.entry((il, plane)) {
6847 slot.insert(self.htod(host)?);
6848 }
6849 }
6850 }
6851 let (mg, mu, md, have) = match macros {
6854 Some(_) => (
6855 mac.get(&(il, 0)).expect("gate macro mirror built above"),
6856 mac.get(&(il, 1)).expect("up macro mirror built above"),
6857 mac.get(&(il, 2)).expect("down macro mirror built above"),
6858 1i32,
6859 ),
6860 None => (selw, selw, selw, 0i32),
6861 };
6862 let f = self.func("moe_vrows_tables_from_sel");
6863 let threads = 128u32;
6864 let cfg = LaunchConfig {
6865 grid_dim: ((n_pairs as u32).div_ceil(threads), 1, 1),
6866 block_dim: (threads, 1, 1),
6867 shared_mem_bytes: 0,
6868 };
6869 let (sgi, sui, sdi) = (sg as i64, su as i64, sd as i64);
6870 let (npi, havei) = (n_pairs as i32, have);
6871 let __s_b = self.gpu.stream();
6872 let mut b = __s_b.launch_builder(&f);
6873 b.arg(sel)
6874 .arg(selw)
6875 .arg(mg)
6876 .arg(mu)
6877 .arg(md)
6878 .arg(&mut *ptrs)
6879 .arg(&mut *scl)
6880 .arg(&pg)
6881 .arg(&pu)
6882 .arg(&pd)
6883 .arg(&sgi)
6884 .arg(&sui)
6885 .arg(&sdi)
6886 .arg(&npi)
6887 .arg(&havei);
6888 unsafe {
6889 b.launch(cfg)?;
6890 }
6891 Ok(())
6892 }
6893
6894 pub fn moe_vrows_order_from_sel(
6906 &self,
6907 sel: &CudaSlice<i32>,
6908 n_pairs: usize,
6909 ptrs: &mut CudaSlice<u64>,
6910 ) -> Result<(), Box<dyn std::error::Error>> {
6911 debug_assert!(sel.len() >= n_pairs);
6912 debug_assert!(
6913 ptrs.len() >= 4 * n_pairs,
6914 "the order plane lives at ptrs[3*n_pairs .. 4*n_pairs)"
6915 );
6916 let f = self.func("moe_vrows_order_from_sel");
6917 let threads = 128u32;
6918 let cfg = LaunchConfig {
6919 grid_dim: ((n_pairs as u32).div_ceil(threads), 1, 1),
6920 block_dim: (threads, 1, 1),
6921 shared_mem_bytes: 0,
6922 };
6923 let np = n_pairs as i32;
6924 let __s_b = self.gpu.stream();
6925 let mut b = __s_b.launch_builder(&f);
6926 b.arg(sel).arg(&mut *ptrs).arg(&np);
6927 unsafe {
6928 b.launch(cfg)?;
6929 }
6930 Ok(())
6931 }
6932
6933 #[allow(clippy::too_many_arguments)]
6939 pub fn moe_gate_up_preclamp8_q8_rows(
6941 &self,
6942 ptrs: &CudaSlice<u64>,
6943 scl: &CudaSlice<f32>,
6944 aq: &CudaSlice<i8>,
6945 ad: &CudaSlice<f32>,
6946 limit: f32,
6947 in_f: usize,
6948 n_ff: usize,
6949 n_used: usize,
6950 n_pairs: usize,
6951 qt_g: i32,
6952 qt_u: i32,
6953 rb_g: usize,
6954 rb_u: usize,
6955 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6956 debug_assert!(
6957 limit > 1e-6,
6958 "moe_gate_up_preclamp8_q8_rows needs a live limit; the kernel collapses every gate \
6959 to silu(0) at limit 0"
6960 );
6961 debug_assert!(ptrs.len() >= 3 * n_pairs);
6962 debug_assert_eq!(scl.len(), 3 * n_pairs);
6963 let packed = moe_vrows_pack_on();
6971 let ordered =
6972 !packed && moe_vrows_dedup_order_on() && ptrs.len() >= 4 * n_pairs && n_ff <= 65535;
6973 let (f, cfg) = if packed {
6974 if MOE_VROWS_PACK_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
6975 eprintln!(
6976 "[moe-vrows-pack] engaged: 4-warp blocks on the verify-rows MoE pair \
6977 (MEMRA_MOE_VROWS_PACK=1)"
6978 );
6979 }
6980 (
6981 self.func("moe_gate_up_preclamp8_q8_rows_w4"),
6982 LaunchConfig {
6983 grid_dim: ((n_ff as u32).div_ceil(4), n_pairs as u32, 1),
6984 block_dim: (32, 4, 1),
6985 shared_mem_bytes: 0,
6986 },
6987 )
6988 } else if ordered {
6989 if MOE_VROWS_DEDUP_ORDER_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
6990 == 0
6991 {
6992 eprintln!(
6993 "[moe-vrows-dedup-order] engaged: verify-rows gate/up walks the pair union \
6994 EXPERT-MAJOR with the pair index as the fastest grid dimension, so the \
6995 21.96%-measured repeat visits read a shared expert slab's rows in adjacent \
6996 blocks (MEMRA_MOE_VROWS_DEDUP_ORDER=1)"
6997 );
6998 }
6999 (
7000 self.func("moe_gate_up_preclamp8_q8_rows_ord"),
7001 LaunchConfig {
7002 grid_dim: (n_pairs as u32, n_ff as u32, 1),
7003 block_dim: (32, 1, 1),
7004 shared_mem_bytes: 0,
7005 },
7006 )
7007 } else {
7008 (
7009 self.func("moe_gate_up_preclamp8_q8_rows"),
7010 LaunchConfig {
7011 grid_dim: (n_ff as u32, n_pairs as u32, 1),
7012 block_dim: (32, 1, 1),
7013 shared_mem_bytes: 0,
7014 },
7015 )
7016 };
7017 let mut act = self.vws_uninit(n_pairs * n_ff)?;
7019 let (inf, nff, nu, np) = (in_f as i32, n_ff as i32, n_used as i32, n_pairs as i32);
7020 let (rbg, rbu) = (rb_g as i64, rb_u as i64);
7021 let __s_b = self.gpu.stream();
7022 let mut b = __s_b.launch_builder(&f);
7023 b.arg(ptrs)
7024 .arg(scl)
7025 .arg(aq)
7026 .arg(ad)
7027 .arg(&limit)
7028 .arg(&mut act)
7029 .arg(&inf)
7030 .arg(&nff)
7031 .arg(&nu)
7032 .arg(&np)
7033 .arg(&qt_g)
7034 .arg(&qt_u)
7035 .arg(&rbg)
7036 .arg(&rbu);
7037 unsafe {
7038 b.launch(cfg)?;
7039 }
7040 Ok(act)
7041 }
7042
7043 #[allow(clippy::too_many_arguments)]
7047 pub fn moe_down8_fma_q8_rows(
7049 &self,
7050 ptrs: &CudaSlice<u64>,
7051 scl: &CudaSlice<f32>,
7052 aq2: &CudaSlice<i8>,
7053 ad2: &CudaSlice<f32>,
7054 dst: &mut CudaSlice<f32>,
7055 in_f: usize,
7056 out_f: usize,
7057 n_used: usize,
7058 n_pairs: usize,
7059 qt: i32,
7060 rb: usize,
7061 ) -> Result<(), Box<dyn std::error::Error>> {
7062 debug_assert!(ptrs.len() >= 3 * n_pairs);
7063 debug_assert_eq!(scl.len(), 3 * n_pairs);
7064 debug_assert_eq!(n_pairs % n_used, 0, "pairs are dense slot-major");
7065 let t = n_pairs / n_used;
7066 debug_assert!(dst.len() >= t * out_f);
7067 let packed = moe_vrows_pack_on();
7069 let tmaj = !packed && moe_vrows_down_tmaj_on() && out_f <= 65535;
7075 let (f, cfg) = if packed {
7076 (
7077 self.func("moe_down8_fma_q8_rows_w4"),
7078 LaunchConfig {
7079 grid_dim: ((out_f as u32).div_ceil(4), t as u32, 1),
7080 block_dim: (32, 4, 1),
7081 shared_mem_bytes: 0,
7082 },
7083 )
7084 } else if tmaj {
7085 if MOE_VROWS_DOWN_TMAJ_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
7086 == 0
7087 {
7088 eprintln!(
7089 "[moe-vrows-down-tmaj] engaged: verify-rows down/FMA grid transposed to \
7090 (t, out_f) so the verify rows at one output row are adjacent blocks; the \
7091 slot-ordered FMA chain is unchanged (MEMRA_MOE_VROWS_DOWN_TMAJ=1)"
7092 );
7093 }
7094 (
7095 self.func("moe_down8_fma_q8_rows_tmaj"),
7096 LaunchConfig {
7097 grid_dim: (t as u32, out_f as u32, 1),
7098 block_dim: (32, 1, 1),
7099 shared_mem_bytes: 0,
7100 },
7101 )
7102 } else {
7103 (
7104 self.func("moe_down8_fma_q8_rows"),
7105 LaunchConfig {
7106 grid_dim: (out_f as u32, t as u32, 1),
7107 block_dim: (32, 1, 1),
7108 shared_mem_bytes: 0,
7109 },
7110 )
7111 };
7112 let (inf, outf, nu, np, rbi) = (
7113 in_f as i32,
7114 out_f as i32,
7115 n_used as i32,
7116 n_pairs as i32,
7117 rb as i64,
7118 );
7119 let __s_b = self.gpu.stream();
7120 let mut b = __s_b.launch_builder(&f);
7121 b.arg(ptrs)
7122 .arg(scl)
7123 .arg(aq2)
7124 .arg(ad2)
7125 .arg(dst)
7126 .arg(&inf)
7127 .arg(&outf)
7128 .arg(&nu)
7129 .arg(&np)
7130 .arg(&qt)
7131 .arg(&rbi);
7132 unsafe {
7133 b.launch(cfg)?;
7134 }
7135 Ok(())
7136 }
7137
7138 #[allow(clippy::too_many_arguments)]
7140 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_expert_q8(
7143 &self,
7144 w: &CudaSlice<u8>,
7145 range: std::ops::Range<usize>,
7146 aq: &CudaSlice<i8>,
7147 ad: &CudaSlice<f32>,
7148 m: usize,
7149 in_f: usize,
7150 out_f: usize,
7151 qtype: i32,
7152 row_bytes: usize,
7153 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7154 let f = self.func("qmatvec_expert_q8");
7155 let wv = w.slice(range);
7156 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7157 const ROWS: u32 = 4; let cfg = LaunchConfig {
7159 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, m as u32, 1),
7160 block_dim: (32, ROWS, 1),
7161 shared_mem_bytes: 0,
7162 };
7163 let (inf, outf, mi, rbi) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7164 let __s_b = self.gpu.stream();
7165 let mut b = __s_b.launch_builder(&f);
7166 b.arg(&wv)
7167 .arg(aq)
7168 .arg(ad)
7169 .arg(&mut y)
7170 .arg(&inf)
7171 .arg(&outf)
7172 .arg(&mi)
7173 .arg(&qtype)
7174 .arg(&rbi);
7175 unsafe {
7176 b.launch(cfg)?;
7177 }
7178 Ok(y)
7179 }
7180
7181 #[allow(clippy::too_many_arguments)] pub fn moe_gate_up_silu8(
7183 &self,
7184 gp: WPtr8,
7185 up: WPtr8,
7186 x: &cudarc::driver::CudaView<f32>,
7187 in_f: usize,
7188 n_ff: usize,
7189 n_used: usize,
7190 qt_g: i32,
7191 qt_u: i32,
7192 rb_g: usize,
7193 rb_u: usize,
7194 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7195 let f = self.func("moe_gate_up_silu8_f32");
7196 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig {
7198 grid_dim: (n_ff as u32, n_used as u32, 1),
7199 block_dim: (256, 1, 1),
7200 shared_mem_bytes: 0,
7201 };
7202 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
7203 let __s_b = self.gpu.stream();
7204 let mut b = __s_b.launch_builder(&f);
7205 b.arg(&gp)
7206 .arg(&up)
7207 .arg(x)
7208 .arg(&mut act)
7209 .arg(&inf)
7210 .arg(&nff)
7211 .arg(&qt_g)
7212 .arg(&qt_u)
7213 .arg(&rbg)
7214 .arg(&rbu);
7215 unsafe {
7216 b.launch(cfg)?;
7217 }
7218 Ok(act)
7219 }
7220
7221 #[allow(clippy::too_many_arguments)]
7227 pub fn moe_down8_fma_into(
7228 &self,
7229 dp: WPtr8,
7230 w: F32x8,
7231 act: &CudaSlice<f32>,
7232 dst: &mut cudarc::driver::CudaViewMut<f32>,
7233 in_f: usize,
7234 out_f: usize,
7235 n_used: usize,
7236 qt: i32,
7237 rb: usize,
7238 ) -> Result<(), Box<dyn std::error::Error>> {
7239 let f = self.func("moe_down8_fma_f32");
7240 let cfg = LaunchConfig {
7241 grid_dim: (out_f as u32, 1, 1),
7242 block_dim: (256, 1, 1),
7243 shared_mem_bytes: 0,
7244 };
7245 let (inf, outf, nu, rbv) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
7246 let __s_b = self.gpu.stream();
7247 let mut b = __s_b.launch_builder(&f);
7248 b.arg(&dp)
7249 .arg(&w)
7250 .arg(act)
7251 .arg(dst)
7252 .arg(&inf)
7253 .arg(&outf)
7254 .arg(&nu)
7255 .arg(&qt)
7256 .arg(&rbv);
7257 unsafe {
7258 b.launch(cfg)?;
7259 }
7260 Ok(())
7261 }
7262
7263 #[allow(clippy::too_many_arguments)]
7268 #[allow(clippy::too_many_arguments)]
7283 #[allow(clippy::too_many_arguments)]
7285 #[allow(clippy::manual_div_ceil)] pub fn moe_pairs_matvec_q8(
7287 &self,
7288 table: &CudaSlice<u64>,
7289 proj: i32,
7290 pair_tok: &CudaSlice<i32>,
7291 pair_ex: &CudaSlice<i32>,
7292 aq: &CudaSlice<i8>,
7293 ad: &CudaSlice<f32>,
7294 in_f: usize,
7295 out_f: usize,
7296 n_expert: usize,
7297 n_pairs: usize,
7298 qtype: i32,
7299 row_bytes: usize,
7300 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7301 let f = self.func("moe_pairs_matvec_q8");
7302 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
7303 const ROWS: u32 = 4;
7304 let cfg = LaunchConfig {
7305 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_pairs as u32, 1),
7306 block_dim: (32, ROWS, 1),
7307 shared_mem_bytes: 0,
7308 };
7309 let (inf, outf, ne, np, rbi) = (
7310 in_f as i32,
7311 out_f as i32,
7312 n_expert as i32,
7313 n_pairs as i32,
7314 row_bytes as i64,
7315 );
7316 let __s_b = self.gpu.stream();
7317 let mut b = __s_b.launch_builder(&f);
7318 b.arg(table)
7319 .arg(&proj)
7320 .arg(pair_tok)
7321 .arg(pair_ex)
7322 .arg(aq)
7323 .arg(ad)
7324 .arg(&mut y)
7325 .arg(&inf)
7326 .arg(&outf)
7327 .arg(&ne)
7328 .arg(&np)
7329 .arg(&qtype)
7330 .arg(&rbi);
7331 unsafe {
7332 b.launch(cfg)?;
7333 }
7334 Ok(y)
7335 }
7336
7337 #[allow(clippy::too_many_arguments)]
7339 #[allow(clippy::manual_div_ceil)] pub fn moe_pairs_matvec_q8_em(
7341 &self,
7342 table: &CudaSlice<u64>,
7343 proj: i32,
7344 ex_ids: &CudaSlice<i32>,
7345 ex_off: &CudaSlice<i32>,
7346 ex_pairs: &CudaSlice<i32>,
7347 pair_tok: &CudaSlice<i32>,
7348 aq: &CudaSlice<i8>,
7349 ad: &CudaSlice<f32>,
7350 in_f: usize,
7351 out_f: usize,
7352 n_expert: usize,
7353 n_active: usize,
7354 n_pairs: usize,
7355 qtype: i32,
7356 row_bytes: usize,
7357 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7358 let f = self.func("moe_pairs_matvec_q8_em");
7359 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
7360 const ROWS: u32 = 4;
7361 let cfg = LaunchConfig {
7362 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
7363 block_dim: (32, ROWS, 1),
7364 shared_mem_bytes: 0,
7365 };
7366 let (inf, outf, ne, na, rbi) = (
7367 in_f as i32,
7368 out_f as i32,
7369 n_expert as i32,
7370 n_active as i32,
7371 row_bytes as i64,
7372 );
7373 let __s_b = self.gpu.stream();
7374 let mut b = __s_b.launch_builder(&f);
7375 b.arg(table)
7376 .arg(&proj)
7377 .arg(ex_ids)
7378 .arg(ex_off)
7379 .arg(ex_pairs)
7380 .arg(pair_tok)
7381 .arg(aq)
7382 .arg(ad)
7383 .arg(&mut y)
7384 .arg(&inf)
7385 .arg(&outf)
7386 .arg(&ne)
7387 .arg(&na)
7388 .arg(&qtype)
7389 .arg(&rbi);
7390 unsafe {
7391 b.launch(cfg)?;
7392 }
7393 Ok(y)
7394 }
7395
7396 #[allow(clippy::too_many_arguments)]
7399 #[allow(clippy::manual_div_ceil)] pub fn moe_pairs_matvec_q8_dec(
7401 &self,
7402 table: &CudaSlice<u64>,
7403 proj: i32,
7404 ex_ids: &CudaSlice<i32>,
7405 ex_off: &CudaSlice<i32>,
7406 ex_pairs: &CudaSlice<i32>,
7407 pair_tok: &CudaSlice<i32>,
7408 aq: &CudaSlice<i8>,
7409 ad: &CudaSlice<f32>,
7410 in_f: usize,
7411 out_f: usize,
7412 n_expert: usize,
7413 n_active: usize,
7414 n_pairs: usize,
7415 qtype: i32,
7416 row_bytes: usize,
7417 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7418 let f = self.func("moe_pairs_matvec_q8_dec");
7419 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
7420 const ROWS: u32 = 4;
7421 let cfg = LaunchConfig {
7422 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
7423 block_dim: (32, ROWS, 1),
7424 shared_mem_bytes: 0,
7425 };
7426 let (inf, outf, ne, na, rbi) = (
7427 in_f as i32,
7428 out_f as i32,
7429 n_expert as i32,
7430 n_active as i32,
7431 row_bytes as i64,
7432 );
7433 let __s_b = self.gpu.stream();
7434 let mut b = __s_b.launch_builder(&f);
7435 b.arg(table)
7436 .arg(&proj)
7437 .arg(ex_ids)
7438 .arg(ex_off)
7439 .arg(ex_pairs)
7440 .arg(pair_tok)
7441 .arg(aq)
7442 .arg(ad)
7443 .arg(&mut y)
7444 .arg(&inf)
7445 .arg(&outf)
7446 .arg(&ne)
7447 .arg(&na)
7448 .arg(&qtype)
7449 .arg(&rbi);
7450 unsafe {
7451 b.launch(cfg)?;
7452 }
7453 Ok(y)
7454 }
7455
7456 pub fn moe_pairs_gelu_mul(
7457 &self,
7458 gate: &CudaSlice<f32>,
7459 up: &CudaSlice<f32>,
7460 n: usize,
7461 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7462 let f = self.func("moe_pairs_gelu_mul");
7463 let mut act = self.alloc_uninit::<f32>(n)?;
7464 let cfg = LaunchConfig::for_num_elems(n as u32);
7465 let nl = n as i64;
7466 let __s_b = self.gpu.stream();
7467 let mut b = __s_b.launch_builder(&f);
7468 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
7469 unsafe {
7470 b.launch(cfg)?;
7471 }
7472 Ok(act)
7473 }
7474
7475 pub fn moe_pairs_silu_mul(
7476 &self,
7477 gate: &CudaSlice<f32>,
7478 up: &CudaSlice<f32>,
7479 n: usize,
7480 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7481 let f = self.func("moe_pairs_silu_mul");
7482 let mut act = self.alloc_uninit::<f32>(n)?;
7483 let cfg = LaunchConfig::for_num_elems(n as u32);
7484 let nl = n as i64;
7485 let __s_b = self.gpu.stream();
7486 let mut b = __s_b.launch_builder(&f);
7487 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
7488 unsafe {
7489 b.launch(cfg)?;
7490 }
7491 Ok(act)
7492 }
7493
7494 #[allow(clippy::too_many_arguments)]
7495 #[allow(clippy::manual_div_ceil)] pub fn moe_pairs_scatter(
7497 &self,
7498 y_down: &CudaSlice<f32>,
7499 pair_w: &CudaSlice<f32>,
7500 tok_pair_off: &CudaSlice<i32>,
7501 tok_pair_ids: &CudaSlice<i32>,
7502 moe_out: &mut CudaSlice<f32>,
7503 t: usize,
7504 n_embd: usize,
7505 ) -> Result<(), Box<dyn std::error::Error>> {
7506 let f = self.func("moe_pairs_scatter");
7507 let cfg = LaunchConfig {
7508 grid_dim: (((n_embd + 255) / 256) as u32, t as u32, 1),
7509 block_dim: (256, 1, 1),
7510 shared_mem_bytes: 0,
7511 };
7512 let ne = n_embd as i32;
7513 let __s_b = self.gpu.stream();
7514 let mut b = __s_b.launch_builder(&f);
7515 b.arg(y_down)
7516 .arg(pair_w)
7517 .arg(tok_pair_off)
7518 .arg(tok_pair_ids)
7519 .arg(moe_out)
7520 .arg(&ne);
7521 unsafe {
7522 b.launch(cfg)?;
7523 }
7524 Ok(())
7525 }
7526
7527 #[allow(clippy::too_many_arguments)]
7531 pub fn moe_gate_up_gelu8_dev_q8(
7532 &self,
7533 table: &CudaSlice<u64>,
7534 sel: &cudarc::driver::CudaView<i32>,
7535 aq: &CudaSlice<i8>,
7536 ad: &CudaSlice<f32>,
7537 in_f: usize,
7538 n_ff: usize,
7539 n_used: usize,
7540 n_expert: usize,
7541 qt_g: i32,
7542 qt_u: i32,
7543 rb_g: usize,
7544 rb_u: usize,
7545 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7546 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
7547 let (inf, nff, ne, rbg, rbu) = (
7548 in_f as i32,
7549 n_ff as i32,
7550 n_expert as i32,
7551 rb_g as i64,
7552 rb_u as i64,
7553 );
7554 let f = self.func("moe_gate_up_gelu8_dev_q8");
7555 let cfg = LaunchConfig {
7556 grid_dim: (n_ff as u32, n_used as u32, 1),
7557 block_dim: (32, 1, 1),
7558 shared_mem_bytes: 0,
7559 };
7560 let __s_b = self.gpu.stream();
7561 let mut b = __s_b.launch_builder(&f);
7562 b.arg(table)
7563 .arg(sel)
7564 .arg(aq)
7565 .arg(ad)
7566 .arg(&mut act)
7567 .arg(&inf)
7568 .arg(&nff)
7569 .arg(&ne)
7570 .arg(&qt_g)
7571 .arg(&qt_u)
7572 .arg(&rbg)
7573 .arg(&rbu);
7574 unsafe {
7575 b.launch(cfg)?;
7576 }
7577 Ok(act)
7578 }
7579
7580 #[allow(clippy::too_many_arguments)]
7582 pub fn moe_gate_up_gelu8_dev_q8_rows(
7583 &self,
7584 table: &CudaSlice<u64>,
7585 sel: &CudaSlice<i32>,
7586 aq: &CudaSlice<i8>,
7587 ad: &CudaSlice<f32>,
7588 t: usize,
7589 in_f: usize,
7590 n_ff: usize,
7591 n_used: usize,
7592 n_expert: usize,
7593 qt_g: i32,
7594 qt_u: i32,
7595 rb_g: usize,
7596 rb_u: usize,
7597 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7598 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
7599 let (inf, nff, ne, rbg, rbu, nu) = (
7600 in_f as i32,
7601 n_ff as i32,
7602 n_expert as i32,
7603 rb_g as i64,
7604 rb_u as i64,
7605 n_used as i32,
7606 );
7607 let f = self.func("moe_gate_up_gelu8_dev_q8_rows");
7608 let cfg = LaunchConfig {
7609 grid_dim: (n_ff as u32, n_used as u32, t as u32),
7610 block_dim: (32, 1, 1),
7611 shared_mem_bytes: 0,
7612 };
7613 let __s_b = self.gpu.stream();
7614 let mut b = __s_b.launch_builder(&f);
7615 b.arg(table)
7616 .arg(sel)
7617 .arg(aq)
7618 .arg(ad)
7619 .arg(&mut act)
7620 .arg(&inf)
7621 .arg(&nff)
7622 .arg(&ne)
7623 .arg(&qt_g)
7624 .arg(&qt_u)
7625 .arg(&rbg)
7626 .arg(&rbu)
7627 .arg(&nu);
7628 unsafe {
7629 b.launch(cfg)?;
7630 }
7631 Ok(act)
7632 }
7633
7634 #[allow(clippy::too_many_arguments)]
7636 pub fn moe_gate_up_gelu8_dev_q8_csr(
7637 &self,
7638 table: &CudaSlice<u64>,
7639 sel: &CudaSlice<i32>,
7640 aq: &CudaSlice<i8>,
7641 ad: &CudaSlice<f32>,
7642 n_pairs: usize,
7643 in_f: usize,
7644 n_ff: usize,
7645 n_used: usize,
7646 n_expert: usize,
7647 qt_g: i32,
7648 qt_u: i32,
7649 rb_g: usize,
7650 rb_u: usize,
7651 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7652 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
7653 let (inf, nff, ne, rbg, rbu, nu, npi) = (
7654 in_f as i32,
7655 n_ff as i32,
7656 n_expert as i32,
7657 rb_g as i64,
7658 rb_u as i64,
7659 n_used as i32,
7660 n_pairs as i32,
7661 );
7662 let f = self.func("moe_gate_up_gelu8_dev_q8_csr");
7663 let cfg = LaunchConfig {
7664 grid_dim: (n_ff as u32, n_pairs as u32, 1),
7665 block_dim: (32, 1, 1),
7666 shared_mem_bytes: 0,
7667 };
7668 let __s_b = self.gpu.stream();
7669 let mut b = __s_b.launch_builder(&f);
7670 b.arg(table)
7671 .arg(sel)
7672 .arg(aq)
7673 .arg(ad)
7674 .arg(&mut act)
7675 .arg(&inf)
7676 .arg(&nff)
7677 .arg(&ne)
7678 .arg(&qt_g)
7679 .arg(&qt_u)
7680 .arg(&rbg)
7681 .arg(&rbu)
7682 .arg(&nu)
7683 .arg(&npi);
7684 unsafe {
7685 b.launch(cfg)?;
7686 }
7687 Ok(act)
7688 }
7689
7690 #[allow(clippy::too_many_arguments)]
7692 pub fn moe_down8_fma_dev_q8_rows_g(
7693 &self,
7694 table: &CudaSlice<u64>,
7695 sel: &CudaSlice<i32>,
7696 w: &CudaSlice<f32>,
7697 aq2: &CudaSlice<i8>,
7698 ad2: &CudaSlice<f32>,
7699 dst: &mut CudaSlice<f32>,
7700 t: usize,
7701 in_f: usize,
7702 out_f: usize,
7703 n_used: usize,
7704 n_expert: usize,
7705 qt: i32,
7706 rb: usize,
7707 ) -> Result<(), Box<dyn std::error::Error>> {
7708 let (inf, outf, nu, ne, rbi) = (
7709 in_f as i32,
7710 out_f as i32,
7711 n_used as i32,
7712 n_expert as i32,
7713 rb as i64,
7714 );
7715 let step_b1_w8 = t == 1 && in_f == 1280 && out_f == 4096 && n_used == 8 && qt == QT_IQ4_XS;
7719 let f = self.func(if step_b1_w8 {
7720 "moe_down8_fma_dev_q8_rows_w8"
7721 } else {
7722 "moe_down8_fma_dev_q8_rows_g"
7723 });
7724 let cfg = LaunchConfig {
7725 grid_dim: (out_f as u32, 1, t as u32),
7726 block_dim: (32, if step_b1_w8 { 8 } else { 1 }, 1),
7727 shared_mem_bytes: 0,
7728 };
7729 let __s_b = self.gpu.stream();
7730 let mut b = __s_b.launch_builder(&f);
7731 b.arg(table)
7732 .arg(sel)
7733 .arg(w)
7734 .arg(aq2)
7735 .arg(ad2)
7736 .arg(dst)
7737 .arg(&inf)
7738 .arg(&outf)
7739 .arg(&nu)
7740 .arg(&ne)
7741 .arg(&qt)
7742 .arg(&rbi);
7743 unsafe {
7744 b.launch(cfg)?;
7745 }
7746 Ok(())
7747 }
7748
7749 pub fn rp_probe_q4(&self, m: usize) -> Result<(f64, f64), Box<dyn std::error::Error>> {
7753 let (out_f, in_f) = (2048usize, 2816usize);
7754 let nblk = in_f / 32;
7755 let mut seed = 0x9E3779B97F4A7C15u64;
7756 let mut rng = move || {
7757 seed = seed
7758 .wrapping_mul(6364136223846793005)
7759 .wrapping_add(1442695040888963407);
7760 (seed >> 33) as u8
7761 };
7762 let mut w = vec![0u8; out_f * nblk * 18];
7763 for b in w.iter_mut() {
7764 *b = rng();
7765 }
7766 for r in 0..out_f {
7767 for g in 0..nblk {
7768 let off = (r * nblk + g) * 18;
7769 w[off] = 0x00;
7770 w[off + 1] = 0x2C; }
7772 }
7773 let qplane = out_f * nblk * 16;
7774 let mut wrp = vec![0u8; w.len()];
7775 for r in 0..out_f {
7776 for g in 0..nblk {
7777 let src = &w[(r * nblk + g) * 18..(r * nblk + g) * 18 + 18];
7778 wrp[qplane + (r * nblk + g) * 2..qplane + (r * nblk + g) * 2 + 2]
7779 .copy_from_slice(&src[0..2]);
7780 wrp[(r * nblk + g) * 16..(r * nblk + g) * 16 + 16].copy_from_slice(&src[2..18]);
7781 }
7782 }
7783 let w_d = self.htod_bytes(&w)?;
7784 let wrp_d = self.htod_bytes(&wrp)?;
7785 let mut aq = vec![0i8; m * in_f];
7786 for v in aq.iter_mut() {
7787 *v = rng() as i8;
7788 }
7789 let aq_d = self.htod_i8(&aq)?;
7790 let ad_d = self.htod(&vec![0.03125f32; m * nblk])?;
7791 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
7792 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
7793 const RPB: u32 = 4;
7794 let cfg = LaunchConfig {
7795 grid_dim: ((out_f as u32).div_ceil(RPB), 1, 1),
7796 block_dim: (32, RPB, 1),
7797 shared_mem_bytes: 0,
7798 };
7799 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
7800 let (rb, qp) = ((nblk * 18) as i64, qplane as i64);
7801 let fb = self.func("qmatvec_q4_0_mmvq_b4");
7802 let fr = self.func("qmatvec_q4_0_mmvq_b4_rp");
7803 {
7804 let __s_b = self.gpu.stream();
7805 let mut b = __s_b.launch_builder(&fb);
7806 b.arg(&w_d)
7807 .arg(&aq_d)
7808 .arg(&ad_d)
7809 .arg(&mut y0)
7810 .arg(&inf)
7811 .arg(&outf)
7812 .arg(&mi)
7813 .arg(&rb);
7814 unsafe {
7815 b.launch(cfg)?;
7816 }
7817 let __s_b = self.gpu.stream();
7818 let mut b = __s_b.launch_builder(&fr);
7819 b.arg(&wrp_d)
7820 .arg(&aq_d)
7821 .arg(&ad_d)
7822 .arg(&mut y1)
7823 .arg(&inf)
7824 .arg(&outf)
7825 .arg(&mi)
7826 .arg(&qp);
7827 unsafe {
7828 b.launch(cfg)?;
7829 }
7830 }
7831 self.gpu.stream().synchronize()?;
7832 let (h0, h1) = (self.dtoh(&y0)?, self.dtoh(&y1)?);
7833 let nd = h0
7834 .iter()
7835 .zip(&h1)
7836 .filter(|(a, b)| a.to_bits() != b.to_bits())
7837 .count();
7838 if nd != 0 {
7839 return Err(format!("rp twin not bitwise: {nd}/{} diffs", h0.len()).into());
7840 }
7841 let mut time = |rp: bool| -> Result<f64, Box<dyn std::error::Error>> {
7842 self.gpu.stream().synchronize()?;
7843 let t0 = std::time::Instant::now();
7844 for _ in 0..500 {
7845 if rp {
7846 let __s_b = self.gpu.stream();
7847 let mut b = __s_b.launch_builder(&fr);
7848 b.arg(&wrp_d)
7849 .arg(&aq_d)
7850 .arg(&ad_d)
7851 .arg(&mut y1)
7852 .arg(&inf)
7853 .arg(&outf)
7854 .arg(&mi)
7855 .arg(&qp);
7856 unsafe {
7857 b.launch(cfg)?;
7858 }
7859 } else {
7860 let __s_b = self.gpu.stream();
7861 let mut b = __s_b.launch_builder(&fb);
7862 b.arg(&w_d)
7863 .arg(&aq_d)
7864 .arg(&ad_d)
7865 .arg(&mut y0)
7866 .arg(&inf)
7867 .arg(&outf)
7868 .arg(&mi)
7869 .arg(&rb);
7870 unsafe {
7871 b.launch(cfg)?;
7872 }
7873 }
7874 }
7875 self.gpu.stream().synchronize()?;
7876 Ok(t0.elapsed().as_secs_f64() * 1e6 / 500.0)
7877 };
7878 let _ = time(false)?;
7879 let _ = time(true)?; Ok((time(false)?, time(true)?))
7881 }
7882
7883 pub fn build_q4_rp4(
7888 &self,
7889 t: &mut crate::model::GpuTensor,
7890 ) -> Result<(), Box<dyn std::error::Error>> {
7891 use crate::model::GpuTensor;
7892 let GpuTensor::Quant {
7893 bytes,
7894 qtype,
7895 row_bytes,
7896 ne,
7897 rp4,
7898 ..
7899 } = t
7900 else {
7901 return Ok(());
7902 };
7903 if *qtype != QT_Q4_0 || rp4.is_some() || ne.len() != 2 {
7904 return Ok(());
7905 }
7906 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
7907 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 18 {
7908 return Ok(());
7909 }
7910 let nblk = in_f / 32;
7911 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 18)?;
7912 let f = self.func("q4_0_split_rp_build");
7913 let n = (out_f * nblk) as i32;
7914 let cfg = LaunchConfig {
7915 grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
7916 block_dim: (256, 1, 1),
7917 shared_mem_bytes: 0,
7918 };
7919 let (of, nb) = (out_f as i32, nblk as i32);
7920 let _ = n;
7921 let __s_b = self.gpu.stream();
7922 let mut b = __s_b.launch_builder(&f);
7923 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
7924 unsafe {
7925 b.launch(cfg)?;
7926 }
7927 *rp4 = Some(dst);
7928 Ok(())
7929 }
7930
7931 pub fn build_q8_rp4(
7936 &self,
7937 t: &mut crate::model::GpuTensor,
7938 ) -> Result<(), Box<dyn std::error::Error>> {
7939 use crate::model::GpuTensor;
7940 let GpuTensor::Quant {
7941 bytes,
7942 qtype,
7943 row_bytes,
7944 ne,
7945 rp4,
7946 ..
7947 } = t
7948 else {
7949 return Ok(());
7950 };
7951 if *qtype != QT_Q8_0 || rp4.is_some() || ne.len() != 2 {
7952 return Ok(());
7953 }
7954 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
7955 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 34 {
7956 return Ok(());
7957 }
7958 *rp4 = Some(self.build_q8_rp4_raw(bytes, in_f, out_f)?);
7959 Ok(())
7960 }
7961
7962 pub fn build_q8_rp4_raw(
7965 &self,
7966 bytes: &CudaSlice<u8>,
7967 in_f: usize,
7968 out_f: usize,
7969 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
7970 assert!(in_f.is_multiple_of(32));
7971 let nblk = in_f / 32;
7972 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 34)?;
7973 let f = self.func("q8_0_split_rp_build");
7974 let cfg = LaunchConfig {
7975 grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
7976 block_dim: (256, 1, 1),
7977 shared_mem_bytes: 0,
7978 };
7979 let (of, nb) = (out_f as i32, nblk as i32);
7980 let __s_b = self.gpu.stream();
7981 let mut b = __s_b.launch_builder(&f);
7982 b.arg(bytes).arg(&mut dst).arg(&of).arg(&nb);
7983 unsafe {
7984 b.launch(cfg)?;
7985 }
7986 Ok(dst)
7987 }
7988
7989 pub fn build_q4k_rp4(
7997 &self,
7998 t: &mut crate::model::GpuTensor,
7999 ) -> Result<(), Box<dyn std::error::Error>> {
8000 use crate::model::GpuTensor;
8001 let GpuTensor::Quant {
8002 bytes,
8003 qtype,
8004 row_bytes,
8005 ne,
8006 rp4,
8007 ..
8008 } = t
8009 else {
8010 return Ok(());
8011 };
8012 if *qtype != QT_Q4_K || rp4.is_some() || ne.len() != 2 {
8013 return Ok(());
8014 }
8015 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
8016 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 144 {
8017 return Ok(());
8018 }
8019 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q4_K)?);
8020 Ok(())
8021 }
8022
8023 pub fn build_q6k_rp4(
8024 &self,
8025 t: &mut crate::model::GpuTensor,
8026 ) -> Result<(), Box<dyn std::error::Error>> {
8027 use crate::model::GpuTensor;
8028 let GpuTensor::Quant {
8029 bytes,
8030 qtype,
8031 row_bytes,
8032 ne,
8033 rp4,
8034 ..
8035 } = t
8036 else {
8037 return Ok(());
8038 };
8039 if *qtype != QT_Q6_K || rp4.is_some() || ne.len() != 2 {
8040 return Ok(());
8041 }
8042 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
8043 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 210 {
8044 return Ok(());
8045 }
8046 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q6_K)?);
8047 Ok(())
8048 }
8049
8050 pub fn build_kq_rp4_raw(
8052 &self,
8053 bytes: &CudaSlice<u8>,
8054 in_f: usize,
8055 out_f: usize,
8056 qtype: i32,
8057 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
8058 assert!(in_f.is_multiple_of(256));
8059 let nsbk = in_f / 256;
8060 let (sb_bytes, kname) = match qtype {
8061 QT_Q4_K => (144usize, "q4_K_split_rp_build"),
8062 QT_Q6_K => (210usize, "q6_K_split_rp_build"),
8063 _ => return Err(format!("build_kq_rp4_raw: qtype {qtype} has no rp mirror").into()),
8064 };
8065 let mut dst = self.alloc_uninit::<u8>(out_f * nsbk * sb_bytes)?;
8066 let f = self.func(kname);
8067 let cfg = LaunchConfig {
8068 grid_dim: (((out_f * nsbk) as u32).div_ceil(256), 1, 1),
8069 block_dim: (256, 1, 1),
8070 shared_mem_bytes: 0,
8071 };
8072 let (of, nb) = (out_f as i32, nsbk as i32);
8073 let __s_b = self.gpu.stream();
8074 let mut b = __s_b.launch_builder(&f);
8075 b.arg(bytes).arg(&mut dst).arg(&of).arg(&nb);
8076 unsafe {
8077 b.launch(cfg)?;
8078 }
8079 Ok(dst)
8080 }
8081
8082 pub fn kqrp_enabled() -> bool {
8086 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8087 *ON.get_or_init(|| match std::env::var("MEMRA_KQRP").as_deref() {
8088 Ok("0") => false,
8089 Ok(_) => true,
8090 Err(_) => cfg!(memra_hopper_mma),
8091 })
8092 }
8093
8094 pub fn build_q4_rp_swap(
8100 &self,
8101 t: &mut crate::model::GpuTensor,
8102 ) -> Result<bool, Box<dyn std::error::Error>> {
8103 use crate::model::GpuTensor;
8104 if !matches!(t, GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0) {
8114 return Ok(false);
8115 }
8116 self.build_q4_rp4(t)?;
8117 self.gpu.stream().synchronize()?; let GpuTensor::Quant { bytes, rp4, rp, .. } = t else {
8119 return Ok(false);
8120 };
8121 match rp4.take() {
8122 Some(split) => {
8123 *bytes = split; *rp = true;
8125 Ok(true)
8126 }
8127 None => Ok(false),
8128 }
8129 }
8130
8131 pub fn q4rp_enabled() -> bool {
8133 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8134 *ON.get_or_init(|| {
8135 std::env::var("MEMRA_Q4RP")
8136 .map(|v| v != "0")
8137 .unwrap_or(true)
8138 })
8139 }
8140
8141 #[allow(clippy::manual_div_ceil)] pub fn copy_rows_strided(
8145 &self,
8146 src: &CudaSlice<f32>,
8147 dst: &mut CudaSlice<f32>,
8148 row_elems: usize,
8149 n_rows: usize,
8150 src_stride: usize,
8151 src_off: usize,
8152 ) -> Result<(), Box<dyn std::error::Error>> {
8153 let f = self.func("copy_rows_strided_f32");
8154 let cfg = LaunchConfig {
8155 grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
8156 block_dim: (256, 1, 1),
8157 shared_mem_bytes: 0,
8158 };
8159 let (re, nr) = (row_elems as i32, n_rows as i32);
8160 let (st, off) = (src_stride as i64, src_off as i64);
8161 let __s_b = self.gpu.stream();
8162 let mut b = __s_b.launch_builder(&f);
8163 b.arg(src)
8164 .arg(&mut *dst)
8165 .arg(&re)
8166 .arg(&nr)
8167 .arg(&st)
8168 .arg(&off);
8169 unsafe {
8170 b.launch(cfg)?;
8171 }
8172 Ok(())
8173 }
8174
8175 #[allow(clippy::manual_div_ceil)] pub fn place_rows_strided(
8182 &self,
8183 src: &CudaSlice<f32>,
8184 dst: &mut CudaSlice<f32>,
8185 row_elems: usize,
8186 n_rows: usize,
8187 dst_stride: usize,
8188 dst_off: usize,
8189 ) -> Result<(), Box<dyn std::error::Error>> {
8190 if row_elems == 0 || n_rows == 0 {
8191 return Err("strided row placement requires nonzero rows and row width".into());
8192 }
8193 let src_len = n_rows
8194 .checked_mul(row_elems)
8195 .ok_or("strided row placement source size overflow")?;
8196 let dst_len = n_rows
8197 .checked_sub(1)
8198 .and_then(|rows| rows.checked_mul(dst_stride))
8199 .and_then(|base| base.checked_add(dst_off))
8200 .and_then(|base| base.checked_add(row_elems))
8201 .ok_or("strided row placement destination size overflow")?;
8202 let row_end = dst_off
8203 .checked_add(row_elems)
8204 .ok_or("strided row placement row size overflow")?;
8205 if src.len() < src_len || dst.len() < dst_len || row_end > dst_stride {
8206 return Err(format!(
8207 "strided row placement geometry mismatch: src={} need_src={src_len} \
8208 dst={} need_dst={dst_len} row_elems={row_elems} rows={n_rows} \
8209 dst_stride={dst_stride} dst_off={dst_off}",
8210 src.len(),
8211 dst.len(),
8212 )
8213 .into());
8214 }
8215 if row_elems > i32::MAX as usize || n_rows > i32::MAX as usize {
8216 return Err("strided row placement exceeds CUDA kernel geometry".into());
8217 }
8218 let f = self.func("place_rows_strided_f32");
8219 let cfg = LaunchConfig {
8220 grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
8221 block_dim: (256, 1, 1),
8222 shared_mem_bytes: 0,
8223 };
8224 let (re, nr) = (row_elems as i32, n_rows as i32);
8225 let (st, off) = (dst_stride as i64, dst_off as i64);
8226 let __s_b = self.gpu.stream();
8227 let mut b = __s_b.launch_builder(&f);
8228 b.arg(src)
8229 .arg(&mut *dst)
8230 .arg(&re)
8231 .arg(&nr)
8232 .arg(&st)
8233 .arg(&off);
8234 unsafe {
8235 b.launch(cfg)?;
8236 }
8237 Ok(())
8238 }
8239
8240 pub fn u32_set_k(
8242 &self,
8243 dst: &mut CudaSlice<u32>,
8244 v: u32,
8245 idx: usize,
8246 ) -> Result<(), Box<dyn std::error::Error>> {
8247 let f = self.func("u32_set_k");
8248 let cfg = LaunchConfig {
8249 grid_dim: (1, 1, 1),
8250 block_dim: (1, 1, 1),
8251 shared_mem_bytes: 0,
8252 };
8253 let ii = idx as i32;
8254 let __s_b = self.gpu.stream();
8255 let mut b = __s_b.launch_builder(&f);
8256 b.arg(dst).arg(&v).arg(&ii);
8257 unsafe {
8258 b.launch(cfg)?;
8259 }
8260 Ok(())
8261 }
8262
8263 pub fn i32_add_k(
8265 &self,
8266 d: &mut CudaSlice<i32>,
8267 v: i32,
8268 ) -> Result<(), Box<dyn std::error::Error>> {
8269 let f = self.func("i32_add_k");
8270 let cfg = LaunchConfig {
8271 grid_dim: (1, 1, 1),
8272 block_dim: (32, 1, 1),
8273 shared_mem_bytes: 0,
8274 };
8275 let __s_b = self.gpu.stream();
8276 let mut b = __s_b.launch_builder(&f);
8277 b.arg(d).arg(&v);
8278 unsafe {
8279 b.launch(cfg)?;
8280 }
8281 Ok(())
8282 }
8283
8284 pub fn i32_iota_from(
8286 &self,
8287 ctr: &CudaSlice<i32>,
8288 dst: &mut CudaSlice<i32>,
8289 n: usize,
8290 ) -> Result<(), Box<dyn std::error::Error>> {
8291 let f = self.func("i32_iota_from");
8292 let cfg = LaunchConfig::for_num_elems(n as u32);
8293 let ni = n as i32;
8294 let __s_b = self.gpu.stream();
8295 let mut b = __s_b.launch_builder(&f);
8296 b.arg(ctr).arg(dst).arg(&ni);
8297 unsafe {
8298 b.launch(cfg)?;
8299 }
8300 Ok(())
8301 }
8302
8303 pub fn u32_map_k(
8305 &self,
8306 buf: &mut CudaSlice<u32>,
8307 map: &CudaSlice<u32>,
8308 idx: usize,
8309 ) -> Result<(), Box<dyn std::error::Error>> {
8310 let f = self.func("u32_map_k");
8311 let cfg = LaunchConfig {
8312 grid_dim: (1, 1, 1),
8313 block_dim: (1, 1, 1),
8314 shared_mem_bytes: 0,
8315 };
8316 let ii = idx as i32;
8317 let __s_b = self.gpu.stream();
8318 let mut b = __s_b.launch_builder(&f);
8319 b.arg(buf).arg(map).arg(&ii);
8320 unsafe {
8321 b.launch(cfg)?;
8322 }
8323 Ok(())
8324 }
8325
8326 #[allow(clippy::too_many_arguments)]
8328 pub fn u32_pack2(
8329 &self,
8330 a: &CudaSlice<u32>,
8331 off_a: usize,
8332 n1: usize,
8333 b_in: &CudaSlice<u32>,
8334 n2: usize,
8335 out: &mut CudaSlice<u32>,
8336 ) -> Result<(), Box<dyn std::error::Error>> {
8337 let f = self.func("u32_pack2");
8338 let cfg = LaunchConfig::for_num_elems((n1 + n2) as u32);
8339 let (oa, i1, i2) = (off_a as i32, n1 as i32, n2 as i32);
8340 let __s_b = self.gpu.stream();
8341 let mut b = __s_b.launch_builder(&f);
8342 b.arg(a).arg(&oa).arg(&i1).arg(b_in).arg(&i2).arg(out);
8343 unsafe {
8344 b.launch(cfg)?;
8345 }
8346 Ok(())
8347 }
8348
8349 pub fn moe_w_exscale(
8351 &self,
8352 w: &mut CudaSlice<f32>,
8353 sel: &CudaSlice<i32>,
8354 s: &CudaSlice<f32>,
8355 n: usize,
8356 ) -> Result<(), Box<dyn std::error::Error>> {
8357 let f = self.func("moe_w_exscale");
8358 let cfg = LaunchConfig::for_num_elems(n as u32);
8359 let ni = n as i32;
8360 let __s_b = self.gpu.stream();
8361 let mut b = __s_b.launch_builder(&f);
8362 b.arg(w).arg(sel).arg(s).arg(&ni);
8363 unsafe {
8364 b.launch(cfg)?;
8365 }
8366 Ok(())
8367 }
8368
8369 pub fn moe_w_scale_by_expert(
8372 &self,
8373 w: &mut CudaSlice<f32>,
8374 sel: &CudaSlice<i32>,
8375 macros: &CudaSlice<f32>,
8376 n_expert: usize,
8377 n: usize,
8378 ) -> Result<(), Box<dyn std::error::Error>> {
8379 let f = self.func("moe_w_scale_by_expert");
8380 let cfg = LaunchConfig {
8381 grid_dim: (n.div_ceil(64) as u32, 1, 1),
8382 block_dim: (64, 1, 1),
8383 shared_mem_bytes: 0,
8384 };
8385 let (ne, nn) = (n_expert as i32, n as i32);
8386 let __s_b = self.gpu.stream();
8387 let mut b = __s_b.launch_builder(&f);
8388 b.arg(w).arg(sel).arg(macros).arg(&ne).arg(&nn);
8389 unsafe {
8390 b.launch(cfg)?;
8391 }
8392 Ok(())
8393 }
8394
8395 #[allow(clippy::too_many_arguments)] pub fn moe_gate_up_silu8_dev_q8(
8397 &self,
8398 table: &CudaSlice<u64>,
8399 sel: &cudarc::driver::CudaView<i32>,
8400 aq: &CudaSlice<i8>,
8401 ad: &CudaSlice<f32>,
8402 in_f: usize,
8403 n_ff: usize,
8404 n_used: usize,
8405 n_expert: usize,
8406 qt_g: i32,
8407 qt_u: i32,
8408 rb_g: usize,
8409 rb_u: usize,
8410 macros: &CudaSlice<f32>,
8411 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8412 static GU: std::sync::OnceLock<(String, u32)> = std::sync::OnceLock::new();
8413 let (mode, wpb) = GU.get_or_init(|| {
8414 let mode = std::env::var("MEMRA_MOE_DEVQ8_GU").unwrap_or_default();
8415 let wpb = std::env::var("MEMRA_MOE_DEVQ8_WPB")
8416 .ok()
8417 .and_then(|v| v.parse().ok())
8418 .unwrap_or(4u32)
8419 .clamp(1, 16);
8420 (mode, wpb)
8421 });
8422 let (mode, wpb) = (mode.as_str(), *wpb);
8423 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
8424 let (inf, nff, ne, rbg, rbu) = (
8425 in_f as i32,
8426 n_ff as i32,
8427 n_expert as i32,
8428 rb_g as i64,
8429 rb_u as i64,
8430 );
8431 let (f, cfg) = match mode {
8432 "1" | "2" | "4" => {
8433 let rpw: u32 = mode.parse().unwrap();
8434 let f = self.func(match rpw {
8435 1 => "moe_gate_up_silu8_dev_q8_r1",
8436 2 => "moe_gate_up_silu8_dev_q8_r2",
8437 _ => "moe_gate_up_silu8_dev_q8_r4",
8438 });
8439 let rows_per_block = (rpw * wpb) as usize;
8440 let gx = n_ff.div_ceil(rows_per_block) as u32;
8441 (
8442 f,
8443 LaunchConfig {
8444 grid_dim: (gx, n_used as u32, 1),
8445 block_dim: (32, wpb, 1),
8446 shared_mem_bytes: 0,
8447 },
8448 )
8449 }
8450 "j8" if n_used <= 32 => (
8451 self.func("moe_gate_up_silu8_dev_q8_j8"),
8452 LaunchConfig {
8453 grid_dim: (n_ff as u32, 1, 1),
8454 block_dim: (32, n_used as u32, 1),
8455 shared_mem_bytes: 0,
8456 },
8457 ),
8458 "vsm2" => {
8460 let f = self.func("moe_gate_up_silu8_dev_q8_vsm2");
8461 let sh = (rb_g + rb_u) as u32;
8462 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8463 f.set_attribute(
8464 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
8465 sh as i32,
8466 )?;
8467 (
8468 f,
8469 LaunchConfig {
8470 grid_dim: (n_ff as u32, n_used as u32, 1),
8471 block_dim: (32, 1, 1),
8472 shared_mem_bytes: sh,
8473 },
8474 )
8475 }
8476 "vsm" => {
8477 let f = self.func("moe_gate_up_silu8_dev_q8_vsm");
8478 let sh = (rb_g + rb_u) as u32;
8479 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8480 f.set_attribute(
8481 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
8482 sh as i32,
8483 )?;
8484 (
8485 f,
8486 LaunchConfig {
8487 grid_dim: (n_ff as u32, n_used as u32, 1),
8488 block_dim: (32, 1, 1),
8489 shared_mem_bytes: sh,
8490 },
8491 )
8492 }
8493 "sg" => (
8494 self.func("moe_gate_up_silu8_dev_q8_sg"),
8495 LaunchConfig {
8496 grid_dim: (n_ff as u32, n_used as u32, 1),
8497 block_dim: (32, 1, 1),
8498 shared_mem_bytes: 0,
8499 },
8500 ),
8501 "j8sg" if n_used <= 32 => (
8502 self.func("moe_gate_up_silu8_dev_q8_j8sg"),
8503 LaunchConfig {
8504 grid_dim: (n_ff as u32, 1, 1),
8505 block_dim: (32, n_used as u32, 1),
8506 shared_mem_bytes: 0,
8507 },
8508 ),
8509 "u64" if in_f == 2048 => (
8510 self.func("moe_gate_up_silu8_dev_q8_u64"),
8511 LaunchConfig {
8512 grid_dim: (n_ff as u32, n_used as u32, 1),
8513 block_dim: (32, 1, 1),
8514 shared_mem_bytes: 0,
8515 },
8516 ),
8517 "gs4" if in_f == 2048 => (
8518 self.func("moe_gate_up_silu8_dev_q8_gs4"),
8519 LaunchConfig {
8520 grid_dim: (n_ff as u32, n_used as u32, 1),
8521 block_dim: (32, 4, 1),
8522 shared_mem_bytes: 0,
8523 },
8524 ),
8525 "v" | "" => (
8527 self.func("moe_gate_up_silu8_dev_q8_v"),
8528 LaunchConfig {
8529 grid_dim: (n_ff as u32, n_used as u32, 1),
8530 block_dim: (32, 1, 1),
8531 shared_mem_bytes: 0,
8532 },
8533 ),
8534 "s2" => (
8535 self.func("moe_gate_up_silu8_dev_q8_s2"),
8536 LaunchConfig {
8537 grid_dim: (n_ff as u32, n_used as u32, 1),
8538 block_dim: (32, 2, 1),
8539 shared_mem_bytes: 0,
8540 },
8541 ),
8542 "s2z" => {
8543 let rz = wpb.min(16); (
8545 self.func("moe_gate_up_silu8_dev_q8_s2z"),
8546 LaunchConfig {
8547 grid_dim: (n_ff.div_ceil(rz as usize) as u32, n_used as u32, 1),
8548 block_dim: (32, 2, rz),
8549 shared_mem_bytes: 0,
8550 },
8551 )
8552 }
8553 _ => (
8554 self.func("moe_gate_up_silu8_dev_q8"),
8555 LaunchConfig {
8556 grid_dim: (n_ff as u32, n_used as u32, 1),
8557 block_dim: (32, 1, 1),
8558 shared_mem_bytes: 0,
8559 },
8560 ),
8561 };
8562 let __s_b = self.gpu.stream();
8563 let mut b = __s_b.launch_builder(&f);
8564 b.arg(table)
8565 .arg(sel)
8566 .arg(aq)
8567 .arg(ad)
8568 .arg(&mut act)
8569 .arg(&inf)
8570 .arg(&nff)
8571 .arg(&ne)
8572 .arg(&qt_g)
8573 .arg(&qt_u)
8574 .arg(&rbg)
8575 .arg(&rbu)
8576 .arg(macros);
8577 unsafe {
8578 b.launch(cfg)?;
8579 }
8580 Ok(act)
8581 }
8582
8583 #[allow(clippy::too_many_arguments)]
8584 pub fn moe_down8_fma_dev_q8(
8585 &self,
8586 table: &CudaSlice<u64>,
8587 sel: &cudarc::driver::CudaView<i32>,
8588 w: &cudarc::driver::CudaView<f32>,
8589 aq2: &CudaSlice<i8>,
8590 ad2: &CudaSlice<f32>,
8591 dst: &mut cudarc::driver::CudaViewMut<f32>,
8592 in_f: usize,
8593 out_f: usize,
8594 n_used: usize,
8595 n_expert: usize,
8596 qt: i32,
8597 rb: usize,
8598 ) -> Result<(), Box<dyn std::error::Error>> {
8599 static DOWN: std::sync::OnceLock<String> = std::sync::OnceLock::new();
8600 let mode = DOWN.get_or_init(|| std::env::var("MEMRA_MOE_DEVQ8_DOWN").unwrap_or_default());
8601 let (inf, outf, nu, ne, rbi) = (
8602 in_f as i32,
8603 out_f as i32,
8604 n_used as i32,
8605 n_expert as i32,
8606 rb as i64,
8607 );
8608 let (f, cfg) = match mode.as_str() {
8611 m @ ("1" | "2" | "4") if n_used <= 8 => {
8612 let rpw: usize = m.parse().unwrap();
8613 let f = self.func(match rpw {
8614 1 => "moe_down8_fma_dev_q8_w8r1",
8615 2 => "moe_down8_fma_dev_q8_w8r2",
8616 _ => "moe_down8_fma_dev_q8_w8r4",
8617 });
8618 (
8619 f,
8620 LaunchConfig {
8621 grid_dim: (out_f.div_ceil(rpw) as u32, 1, 1),
8622 block_dim: (32, n_used as u32, 1),
8623 shared_mem_bytes: 0,
8624 },
8625 )
8626 }
8627 "h2" if in_f == 512 => (
8628 self.func("moe_down8_fma_dev_q8_h2"),
8629 LaunchConfig {
8630 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8631 block_dim: (32, 1, 1),
8632 shared_mem_bytes: 0,
8633 },
8634 ),
8635 "" if in_f == 704 && n_used <= 8 => (
8638 self.func("moe_down8_fma_dev_q8_w8r2"),
8639 LaunchConfig {
8640 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8641 block_dim: (32, n_used as u32, 1),
8642 shared_mem_bytes: 0,
8643 },
8644 ),
8645 "w8h2v" | "" if in_f == 512 && n_used <= 8 => (
8649 self.func("moe_down8_fma_dev_q8_w8h2v"),
8650 LaunchConfig {
8651 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8652 block_dim: (32, n_used as u32, 1),
8653 shared_mem_bytes: 0,
8654 },
8655 ),
8656 "w8h2r2v" if in_f == 512 && n_used <= 8 => (
8657 self.func("moe_down8_fma_dev_q8_w8h2r2v"),
8658 LaunchConfig {
8659 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
8660 block_dim: (32, n_used as u32, 1),
8661 shared_mem_bytes: 0,
8662 },
8663 ),
8664 "w8h2r2" if in_f == 512 && n_used <= 8 => (
8665 self.func("moe_down8_fma_dev_q8_w8h2r2"),
8666 LaunchConfig {
8667 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
8668 block_dim: (32, n_used as u32, 1),
8669 shared_mem_bytes: 0,
8670 },
8671 ),
8672 "w8h2" if in_f == 512 && n_used <= 8 => (
8673 self.func("moe_down8_fma_dev_q8_w8h2"),
8674 LaunchConfig {
8675 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8676 block_dim: (32, n_used as u32, 1),
8677 shared_mem_bytes: 0,
8678 },
8679 ),
8680 _ => (
8681 self.func("moe_down8_fma_dev_q8"),
8682 LaunchConfig {
8683 grid_dim: (out_f as u32, 1, 1),
8684 block_dim: (32, 1, 1),
8685 shared_mem_bytes: 0,
8686 },
8687 ),
8688 };
8689 let __s_b = self.gpu.stream();
8690 let mut b = __s_b.launch_builder(&f);
8691 b.arg(table)
8692 .arg(sel)
8693 .arg(w)
8694 .arg(aq2)
8695 .arg(ad2)
8696 .arg(dst)
8697 .arg(&inf)
8698 .arg(&outf)
8699 .arg(&nu)
8700 .arg(&ne)
8701 .arg(&qt)
8702 .arg(&rbi);
8703 unsafe {
8704 b.launch(cfg)?;
8705 }
8706 Ok(())
8707 }
8708
8709 #[allow(clippy::too_many_arguments)]
8716 pub fn moe_gate_up_silu8_dev_q8_rows(
8717 &self,
8718 table: &CudaSlice<u64>,
8719 sel: &CudaSlice<i32>,
8720 aq: &CudaSlice<i8>,
8721 ad: &CudaSlice<f32>,
8722 t: usize,
8723 in_f: usize,
8724 n_ff: usize,
8725 n_used: usize,
8726 n_expert: usize,
8727 qt_g: i32,
8728 qt_u: i32,
8729 rb_g: usize,
8730 rb_u: usize,
8731 macros: &CudaSlice<f32>,
8732 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8733 let f = self.func("moe_gate_up_silu8_dev_q8_v_rows");
8734 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
8735 let cfg = LaunchConfig {
8736 grid_dim: (n_ff as u32, n_used as u32, t as u32),
8737 block_dim: (32, 1, 1),
8738 shared_mem_bytes: 0,
8739 };
8740 let (inf, nff, ne, nu, rbg, rbu) = (
8741 in_f as i32,
8742 n_ff as i32,
8743 n_expert as i32,
8744 n_used as i32,
8745 rb_g as i64,
8746 rb_u as i64,
8747 );
8748 let __s_b = self.gpu.stream();
8749 let mut b = __s_b.launch_builder(&f);
8750 b.arg(table)
8751 .arg(sel)
8752 .arg(aq)
8753 .arg(ad)
8754 .arg(&mut act)
8755 .arg(&inf)
8756 .arg(&nff)
8757 .arg(&ne)
8758 .arg(&qt_g)
8759 .arg(&qt_u)
8760 .arg(&rbg)
8761 .arg(&rbu)
8762 .arg(&nu)
8763 .arg(macros);
8764 unsafe {
8765 b.launch(cfg)?;
8766 }
8767 Ok(act)
8768 }
8769
8770 #[allow(clippy::too_many_arguments)]
8775 pub fn moe_down8_fma_dev_q8_rows(
8776 &self,
8777 table: &CudaSlice<u64>,
8778 sel: &CudaSlice<i32>,
8779 w: &CudaSlice<f32>,
8780 aq2: &CudaSlice<i8>,
8781 ad2: &CudaSlice<f32>,
8782 dst: &mut CudaSlice<f32>,
8783 t: usize,
8784 in_f: usize,
8785 out_f: usize,
8786 n_used: usize,
8787 n_expert: usize,
8788 qt: i32,
8789 rb: usize,
8790 ) -> Result<(), Box<dyn std::error::Error>> {
8791 assert!(
8792 in_f == 512 && n_used <= 8,
8793 "down rows twin is w8h2v shape-gated"
8794 );
8795 let f = self.func("moe_down8_fma_dev_q8_w8h2v_rows");
8796 let cfg = LaunchConfig {
8797 grid_dim: (out_f.div_ceil(2) as u32, 1, t as u32),
8798 block_dim: (32, n_used as u32, 1),
8799 shared_mem_bytes: 0,
8800 };
8801 let (inf, outf, nu, ne, rbi) = (
8802 in_f as i32,
8803 out_f as i32,
8804 n_used as i32,
8805 n_expert as i32,
8806 rb as i64,
8807 );
8808 let __s_b = self.gpu.stream();
8809 let mut b = __s_b.launch_builder(&f);
8810 b.arg(table)
8811 .arg(sel)
8812 .arg(w)
8813 .arg(aq2)
8814 .arg(ad2)
8815 .arg(dst)
8816 .arg(&inf)
8817 .arg(&outf)
8818 .arg(&nu)
8819 .arg(&ne)
8820 .arg(&qt)
8821 .arg(&rbi);
8822 unsafe {
8823 b.launch(cfg)?;
8824 }
8825 Ok(())
8826 }
8827
8828 #[allow(clippy::too_many_arguments)]
8832 pub fn moe_gate_up_silu8_dev_q8_csr(
8833 &self,
8834 table: &CudaSlice<u64>,
8835 sel: &CudaSlice<i32>,
8836 aq: &CudaSlice<i8>,
8837 ad: &CudaSlice<f32>,
8838 n_pairs: usize,
8839 in_f: usize,
8840 n_ff: usize,
8841 n_used: usize,
8842 n_expert: usize,
8843 qt_g: i32,
8844 qt_u: i32,
8845 rb_g: usize,
8846 rb_u: usize,
8847 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8848 let f = if qt_g == crate::QT_NVFP4 {
8851 self.func("moe_gate_up_silu8_dev_q8_csr_nvfp4")
8852 } else {
8853 self.func("moe_gate_up_silu8_dev_q8_csr_iq4")
8854 };
8855 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
8856 let cfg = LaunchConfig {
8857 grid_dim: (n_ff as u32, n_pairs as u32, 1),
8858 block_dim: (32, 1, 1),
8859 shared_mem_bytes: 0,
8860 };
8861 let (inf, nff, ne, nu, npi, rbg, rbu) = (
8862 in_f as i32,
8863 n_ff as i32,
8864 n_expert as i32,
8865 n_used as i32,
8866 n_pairs as i32,
8867 rb_g as i64,
8868 rb_u as i64,
8869 );
8870 let __s_b = self.gpu.stream();
8871 let mut b = __s_b.launch_builder(&f);
8872 b.arg(table)
8873 .arg(sel)
8874 .arg(aq)
8875 .arg(ad)
8876 .arg(&mut act)
8877 .arg(&inf)
8878 .arg(&nff)
8879 .arg(&ne)
8880 .arg(&qt_g)
8881 .arg(&qt_u)
8882 .arg(&rbg)
8883 .arg(&rbu)
8884 .arg(&nu)
8885 .arg(&npi);
8886 unsafe {
8887 b.launch(cfg)?;
8888 }
8889 Ok(act)
8890 }
8891
8892 #[allow(clippy::too_many_arguments)]
8896 pub fn moe_down8_fma_dev_q8_variant(
8897 &self,
8898 variant: &str,
8899 table: &CudaSlice<u64>,
8900 sel: &cudarc::driver::CudaView<i32>,
8901 w: &cudarc::driver::CudaView<f32>,
8902 aq2: &CudaSlice<i8>,
8903 ad2: &CudaSlice<f32>,
8904 dst: &mut cudarc::driver::CudaViewMut<f32>,
8905 in_f: usize,
8906 out_f: usize,
8907 n_used: usize,
8908 n_expert: usize,
8909 qt: i32,
8910 rb: usize,
8911 ) -> Result<(), Box<dyn std::error::Error>> {
8912 let (inf, outf, nu, ne, rbi) = (
8913 in_f as i32,
8914 out_f as i32,
8915 n_used as i32,
8916 n_expert as i32,
8917 rb as i64,
8918 );
8919 let (f, cfg) = match variant {
8920 "w8h2" | "w8h2v" => (
8921 self.func(if variant == "w8h2" {
8922 "moe_down8_fma_dev_q8_w8h2"
8923 } else {
8924 "moe_down8_fma_dev_q8_w8h2v"
8925 }),
8926 LaunchConfig {
8927 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8928 block_dim: (32, n_used as u32, 1),
8929 shared_mem_bytes: 0,
8930 },
8931 ),
8932 "w8h2r2" | "w8h2r2v" => (
8933 self.func(if variant == "w8h2r2" {
8934 "moe_down8_fma_dev_q8_w8h2r2"
8935 } else {
8936 "moe_down8_fma_dev_q8_w8h2r2v"
8937 }),
8938 LaunchConfig {
8939 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
8940 block_dim: (32, n_used as u32, 1),
8941 shared_mem_bytes: 0,
8942 },
8943 ),
8944 _ => (
8945 self.func("moe_down8_fma_dev_q8"),
8946 LaunchConfig {
8947 grid_dim: (out_f as u32, 1, 1),
8948 block_dim: (32, 1, 1),
8949 shared_mem_bytes: 0,
8950 },
8951 ),
8952 };
8953 let __s_b = self.gpu.stream();
8954 let mut b = __s_b.launch_builder(&f);
8955 b.arg(table)
8956 .arg(sel)
8957 .arg(w)
8958 .arg(aq2)
8959 .arg(ad2)
8960 .arg(dst)
8961 .arg(&inf)
8962 .arg(&outf)
8963 .arg(&nu)
8964 .arg(&ne)
8965 .arg(&qt)
8966 .arg(&rbi);
8967 unsafe {
8968 b.launch(cfg)?;
8969 }
8970 Ok(())
8971 }
8972
8973 #[allow(clippy::too_many_arguments)]
8975 pub fn moe_gate_up_silu8_dev_q8_variant(
8976 &self,
8977 variant: &str,
8978 table: &CudaSlice<u64>,
8979 sel: &cudarc::driver::CudaView<i32>,
8980 aq: &CudaSlice<i8>,
8981 ad: &CudaSlice<f32>,
8982 in_f: usize,
8983 n_ff: usize,
8984 n_used: usize,
8985 n_expert: usize,
8986 qt_g: i32,
8987 qt_u: i32,
8988 rb_g: usize,
8989 rb_u: usize,
8990 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8991 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
8992 let (inf, nff, ne, rbg, rbu) = (
8993 in_f as i32,
8994 n_ff as i32,
8995 n_expert as i32,
8996 rb_g as i64,
8997 rb_u as i64,
8998 );
8999 let f = self.func(if variant == "v" {
9000 "moe_gate_up_silu8_dev_q8_v"
9001 } else {
9002 "moe_gate_up_silu8_dev_q8"
9003 });
9004 let cfg = LaunchConfig {
9005 grid_dim: (n_ff as u32, n_used as u32, 1),
9006 block_dim: (32, 1, 1),
9007 shared_mem_bytes: 0,
9008 };
9009 let __s_b = self.gpu.stream();
9010 let mut b = __s_b.launch_builder(&f);
9011 b.arg(table)
9012 .arg(sel)
9013 .arg(aq)
9014 .arg(ad)
9015 .arg(&mut act)
9016 .arg(&inf)
9017 .arg(&nff)
9018 .arg(&ne)
9019 .arg(&qt_g)
9020 .arg(&qt_u)
9021 .arg(&rbg)
9022 .arg(&rbu);
9023 unsafe {
9024 b.launch(cfg)?;
9025 }
9026 Ok(act)
9027 }
9028
9029 #[allow(clippy::too_many_arguments)] pub fn moe_gate_up_silu8_dev(
9031 &self,
9032 table: &CudaSlice<u64>,
9033 sel: &cudarc::driver::CudaView<i32>,
9034 x: &cudarc::driver::CudaView<f32>,
9035 in_f: usize,
9036 n_ff: usize,
9037 n_used: usize,
9038 n_expert: usize,
9039 qt_g: i32,
9040 qt_u: i32,
9041 rb_g: usize,
9042 rb_u: usize,
9043 macros: &CudaSlice<f32>,
9044 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9045 let f = self.func("moe_gate_up_silu8_dev");
9046 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig {
9048 grid_dim: (n_ff as u32, n_used as u32, 1),
9049 block_dim: (256, 1, 1),
9050 shared_mem_bytes: 0,
9051 };
9052 let (inf, nff, ne, rbg, rbu) = (
9053 in_f as i32,
9054 n_ff as i32,
9055 n_expert as i32,
9056 rb_g as i64,
9057 rb_u as i64,
9058 );
9059 let __s_b = self.gpu.stream();
9060 let mut b = __s_b.launch_builder(&f);
9061 b.arg(table)
9062 .arg(sel)
9063 .arg(x)
9064 .arg(&mut act)
9065 .arg(&inf)
9066 .arg(&nff)
9067 .arg(&ne)
9068 .arg(&qt_g)
9069 .arg(&qt_u)
9070 .arg(&rbg)
9071 .arg(&rbu)
9072 .arg(macros);
9073 unsafe {
9074 b.launch(cfg)?;
9075 }
9076 Ok(act)
9077 }
9078
9079 #[allow(clippy::too_many_arguments)]
9082 pub fn moe_down8_fma_dev(
9083 &self,
9084 table: &CudaSlice<u64>,
9085 sel: &cudarc::driver::CudaView<i32>,
9086 w: &cudarc::driver::CudaView<f32>,
9087 act: &CudaSlice<f32>,
9088 dst: &mut cudarc::driver::CudaViewMut<f32>,
9089 in_f: usize,
9090 out_f: usize,
9091 n_used: usize,
9092 n_expert: usize,
9093 qt: i32,
9094 rb: usize,
9095 ) -> Result<(), Box<dyn std::error::Error>> {
9096 let f = self.func("moe_down8_fma_dev");
9097 let cfg = LaunchConfig {
9098 grid_dim: (out_f as u32, 1, 1),
9099 block_dim: (256, 1, 1),
9100 shared_mem_bytes: 0,
9101 };
9102 let (inf, outf, nu, ne, rbv) = (
9103 in_f as i32,
9104 out_f as i32,
9105 n_used as i32,
9106 n_expert as i32,
9107 rb as i64,
9108 );
9109 let __s_b = self.gpu.stream();
9110 let mut b = __s_b.launch_builder(&f);
9111 b.arg(table)
9112 .arg(sel)
9113 .arg(w)
9114 .arg(act)
9115 .arg(dst)
9116 .arg(&inf)
9117 .arg(&outf)
9118 .arg(&nu)
9119 .arg(&ne)
9120 .arg(&qt)
9121 .arg(&rbv);
9122 unsafe {
9123 b.launch(cfg)?;
9124 }
9125 Ok(())
9126 }
9127
9128 pub fn axpy_into(
9130 &self,
9131 src: &CudaSlice<f32>,
9132 alpha: f32,
9133 dst: &mut cudarc::driver::CudaViewMut<f32>,
9134 n: usize,
9135 ) -> Result<(), Box<dyn std::error::Error>> {
9136 let f = self.func("axpy_f32");
9137 let cfg = LaunchConfig::for_num_elems(n as u32);
9138 let (a, ni) = (alpha, n as i32);
9139 let __s_b = self.gpu.stream();
9140 let mut b = __s_b.launch_builder(&f);
9141 b.arg(src).arg(dst).arg(&a).arg(&ni);
9142 unsafe {
9143 b.launch(cfg)?;
9144 }
9145 Ok(())
9146 }
9147
9148 pub fn axpy_host_into(
9150 &self,
9151 src: &cudarc::driver::CudaView<'_, f32>,
9152 alpha: f32,
9153 dst: &mut cudarc::driver::CudaViewMut<f32>,
9154 n: usize,
9155 ) -> Result<(), Box<dyn std::error::Error>> {
9156 let f = self.func("axpy_host_f32");
9157 let cfg = LaunchConfig::for_num_elems(n as u32);
9158 let (a, ni) = (alpha, n as i32);
9159 let __s_b = self.gpu.stream();
9160 let mut b = __s_b.launch_builder(&f);
9161 b.arg(src).arg(dst).arg(&a).arg(&ni);
9162 unsafe {
9163 b.launch(cfg)?;
9164 }
9165 Ok(())
9166 }
9167
9168 pub fn add_scaled_rows(
9170 &self,
9171 src: &CudaSlice<f32>,
9172 scale: &CudaSlice<f32>,
9173 dst: &mut CudaSlice<f32>,
9174 ncols: usize,
9175 nrows: usize,
9176 ) -> Result<(), Box<dyn std::error::Error>> {
9177 let f = self.func("add_scaled_rows_f32");
9178 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
9179 let (nc, nr) = (ncols as i32, nrows as i32);
9180 let __s_b = self.gpu.stream();
9181 let mut b = __s_b.launch_builder(&f);
9182 b.arg(src).arg(scale).arg(dst).arg(&nc).arg(&nr);
9183 unsafe {
9184 b.launch(cfg)?;
9185 }
9186 Ok(())
9187 }
9188
9189 pub fn add_scaled_rows_ones(
9194 &self,
9195 src: &CudaSlice<f32>,
9196 dst: &mut CudaSlice<f32>,
9197 ncols: usize,
9198 nrows: usize,
9199 ) -> Result<(), Box<dyn std::error::Error>> {
9200 let mut guard = self
9201 .shexp_ones
9202 .lock()
9203 .map_err(|_| "shexp ones buffer is poisoned")?;
9204 if guard.as_ref().map(|b| b.len() < nrows).unwrap_or(true) {
9205 *guard = Some(self.htod(&vec![1.0f32; nrows.max(64)])?);
9208 }
9209 let ones = guard.as_ref().expect("just ensured");
9210 let f = self.func("add_scaled_rows_f32");
9211 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
9212 let (nc, nr) = (ncols as i32, nrows as i32);
9213 let __s_b = self.gpu.stream();
9214 let mut b = __s_b.launch_builder(&f);
9215 b.arg(src).arg(ones).arg(&mut *dst).arg(&nc).arg(&nr);
9216 unsafe {
9217 b.launch(cfg)?;
9218 }
9219 Ok(())
9220 }
9221
9222 pub fn i32_mirror_store(
9226 &self,
9227 dst: &mut CudaSlice<i32>,
9228 v: i32,
9229 ) -> Result<(), Box<dyn std::error::Error>> {
9230 if crate::htod_diet_on() {
9231 HTOD_DIET_AVOIDED.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
9232 return self.i32_set_k(dst, v);
9233 }
9234 self.gpu.stream().memcpy_htod(&[v], dst)?;
9235 Ok(())
9236 }
9237
9238 pub fn scale_rows(
9241 &self,
9242 y: &mut CudaSlice<f32>,
9243 s: &CudaSlice<f32>,
9244 ncols: usize,
9245 nrows: usize,
9246 ) -> Result<(), Box<dyn std::error::Error>> {
9247 let f = self.func("scale_rows_f32");
9248 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
9249 let (nc, nr) = (ncols as i32, nrows as i32);
9250 let __s_b = self.gpu.stream();
9251 let mut b = __s_b.launch_builder(&f);
9252 b.arg(&mut *y).arg(s).arg(&nc).arg(&nr);
9253 unsafe {
9254 b.launch(cfg)?;
9255 }
9256 Ok(())
9257 }
9258
9259 #[allow(clippy::too_many_arguments)]
9263 pub fn moe_prime_join_scatter(
9264 &self,
9265 y0: &CudaSlice<f32>,
9266 y1: &CudaSlice<f32>,
9267 inv: &CudaSlice<i32>,
9268 w: &CudaSlice<f32>,
9269 out: &mut CudaSlice<f32>,
9270 ncols: usize,
9271 n_used: usize,
9272 t: usize,
9273 ) -> Result<(), Box<dyn std::error::Error>> {
9274 let f = self.func("moe_prime_join_scatter_f32");
9275 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
9276 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
9277 let __s_b = self.gpu.stream();
9278 let mut b = __s_b.launch_builder(&f);
9279 b.arg(y0)
9280 .arg(y1)
9281 .arg(inv)
9282 .arg(w)
9283 .arg(&mut *out)
9284 .arg(&nc)
9285 .arg(&nu)
9286 .arg(&ti);
9287 unsafe {
9288 b.launch(cfg)?;
9289 }
9290 Ok(())
9291 }
9292
9293 pub fn moe_pairs_weighted_scatter(
9296 &self,
9297 y: &CudaSlice<f32>,
9298 w: &CudaSlice<f32>,
9299 out: &mut CudaSlice<f32>,
9300 ncols: usize,
9301 n_used: usize,
9302 t: usize,
9303 ) -> Result<(), Box<dyn std::error::Error>> {
9304 let f = self.func("moe_pairs_weighted_scatter_f32");
9305 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
9306 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
9307 let __s_b = self.gpu.stream();
9308 let mut b = __s_b.launch_builder(&f);
9309 b.arg(y).arg(w).arg(&mut *out).arg(&nc).arg(&nu).arg(&ti);
9310 unsafe {
9311 b.launch(cfg)?;
9312 }
9313 Ok(())
9314 }
9315
9316 pub fn gather_rows(
9320 &self,
9321 src: &CudaSlice<f32>,
9322 idx: &CudaSlice<i32>,
9323 dst: &mut CudaSlice<f32>,
9324 ncols: usize,
9325 m_e: usize,
9326 ) -> Result<(), Box<dyn std::error::Error>> {
9327 let f = self.func("gather_rows_f32");
9328 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
9329 let (nc, me) = (ncols as i32, m_e as i32);
9330 let __s_b = self.gpu.stream();
9331 let mut b = __s_b.launch_builder(&f);
9332 b.arg(src).arg(idx).arg(dst).arg(&nc).arg(&me);
9333 unsafe {
9334 b.launch(cfg)?;
9335 }
9336 Ok(())
9337 }
9338
9339 #[allow(clippy::too_many_arguments)] pub fn scatter_slot(
9345 &self,
9346 src: &CudaSlice<f32>,
9347 tok_idx: &CudaSlice<i32>,
9348 slot_idx: &CudaSlice<i32>,
9349 weight: &CudaSlice<f32>,
9350 dst: &mut CudaSlice<f32>,
9351 wbuf: &mut CudaSlice<f32>,
9352 ncols: usize,
9353 n_used: usize,
9354 m_e: usize,
9355 ) -> Result<(), Box<dyn std::error::Error>> {
9356 let f = self.func("scatter_add_slot_f32");
9357 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
9358 let (nc, nu, me) = (ncols as i32, n_used as i32, m_e as i32);
9359 let __s_b = self.gpu.stream();
9360 let mut b = __s_b.launch_builder(&f);
9361 b.arg(src)
9362 .arg(tok_idx)
9363 .arg(slot_idx)
9364 .arg(weight)
9365 .arg(dst)
9366 .arg(wbuf)
9367 .arg(&nc)
9368 .arg(&nu)
9369 .arg(&me);
9370 unsafe {
9371 b.launch(cfg)?;
9372 }
9373 Ok(())
9374 }
9375
9376 pub fn reduce_slots(
9380 &self,
9381 slots: &CudaSlice<f32>,
9382 wbuf: &CudaSlice<f32>,
9383 dst: &mut CudaSlice<f32>,
9384 ncols: usize,
9385 n_used: usize,
9386 t: usize,
9387 ) -> Result<(), Box<dyn std::error::Error>> {
9388 let f = self.func("reduce_slots_f32");
9389 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
9390 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
9391 let __s_b = self.gpu.stream();
9392 let mut b = __s_b.launch_builder(&f);
9393 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
9394 unsafe {
9395 b.launch(cfg)?;
9396 }
9397 Ok(())
9398 }
9399
9400 pub fn reduce_slots_host(
9405 &self,
9406 slots: &CudaSlice<f32>,
9407 wbuf: &CudaSlice<f32>,
9408 dst: &mut CudaSlice<f32>,
9409 ncols: usize,
9410 n_used: usize,
9411 t: usize,
9412 ) -> Result<(), Box<dyn std::error::Error>> {
9413 let f = self.func("reduce_slots_host_f32");
9414 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
9415 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
9416 let __s_b = self.gpu.stream();
9417 let mut b = __s_b.launch_builder(&f);
9418 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
9419 unsafe {
9420 b.launch(cfg)?;
9421 }
9422 Ok(())
9423 }
9424
9425 pub fn quantize_q8_1_view(
9432 &self,
9433 x: &cudarc::driver::CudaView<f32>,
9434 m: usize,
9435 in_f: usize,
9436 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9437 let f = self.func("quantize_q8_1");
9438 let nblk = in_f / 32;
9439 let mut q = self.alloc_uninit::<i8>(m * in_f)?;
9440 let mut d = self.alloc_uninit::<f32>(m * nblk)?;
9441 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
9442 let (inf, mi) = (in_f as i32, m as i32);
9443 let __s_b = self.gpu.stream();
9444 let mut b = __s_b.launch_builder(&f);
9445 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
9446 unsafe {
9447 b.launch(cfg)?;
9448 }
9449 Ok((q, d))
9450 }
9451
9452 pub fn quantize_q8_1(
9453 &self,
9454 x: &CudaSlice<f32>,
9455 m: usize,
9456 in_f: usize,
9457 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9458 let nblk = in_f / 32;
9459 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);
9463 let (inf, mi) = (in_f as i32, m as i32);
9464 if Self::pdl_on() && Self::pdl_wb_on() {
9465 {
9466 use cudarc::driver::{DevicePtr, DevicePtrMut};
9467 let s = &self.gpu.stream();
9468 let (px, _g0) = x.device_ptr(s);
9469 let (pq, _g1) = q.device_ptr_mut(s);
9470 let (pd, _g2) = d.device_ptr_mut(s);
9471 let mut ps = [
9472 &px as *const _ as *mut std::ffi::c_void,
9473 &pq as *const _ as *mut _,
9474 &pd as *const _ as *mut _,
9475 &inf as *const _ as *mut _,
9476 &mi as *const _ as *mut _,
9477 ];
9478 unsafe {
9479 self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?;
9480 }
9481 }
9482 return Ok((q, d));
9483 }
9484 let f = self.func("quantize_q8_1");
9485 let __s_b = self.gpu.stream();
9486 let mut b = __s_b.launch_builder(&f);
9487 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
9488 unsafe {
9489 b.launch(cfg)?;
9490 }
9491 Ok((q, d))
9492 }
9493
9494 pub fn quantize_fp4_act(
9498 &self,
9499 x: &CudaSlice<f32>,
9500 m: usize,
9501 in_f: usize,
9502 ) -> Result<(CudaSlice<u32>, CudaSlice<u8>), Box<dyn std::error::Error>> {
9503 let f = self.func("quantize_fp4_act");
9504 let nb16 = in_f / 16;
9505 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);
9508 let (inf, mi) = (in_f as i32, m as i32);
9509 let __s_b = self.gpu.stream();
9510 let mut b = __s_b.launch_builder(&f);
9511 b.arg(x).arg(&mut aq4).arg(&mut ad4).arg(&inf).arg(&mi);
9512 unsafe {
9513 b.launch(cfg)?;
9514 }
9515 Ok((aq4, ad4))
9516 }
9517
9518 #[allow(clippy::too_many_arguments)] pub fn qmatvec_gemm_nvfp4_fp4(
9524 &self,
9525 bytes: &CudaSlice<u8>,
9526 x: &CudaSlice<f32>,
9527 m: usize,
9528 in_f: usize,
9529 out_f: usize,
9530 row_bytes: usize,
9531 scale: f32,
9532 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9533 assert!(
9534 in_f.is_multiple_of(64),
9535 "FP4 GEMM requires in_f % 64 == 0, got {in_f}"
9536 );
9537 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
9538 let mut y = self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)?;
9539 if scale != 1.0 {
9540 self.scale_inplace(&mut y, scale, m * out_f)?;
9541 }
9542 Ok(y)
9543 }
9544
9545 #[allow(clippy::too_many_arguments)]
9548 #[allow(clippy::manual_div_ceil)] fn fp4_gemm_launch(
9551 &self,
9552 bytes: &CudaSlice<u8>,
9553 aq4: &CudaSlice<u32>,
9554 ad4: &CudaSlice<u8>,
9555 m: usize,
9556 in_f: usize,
9557 out_f: usize,
9558 row_bytes: usize,
9559 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9560 let f = self.func("qmatvec_gemm_nvfp4_fp4");
9561 let mut y = self.alloc_uninit::<f32>(m * out_f)?; const BM: u32 = 64;
9563 const BN: u32 = 256;
9564 let cfg = LaunchConfig {
9565 grid_dim: ((out_f as u32 + BM - 1) / BM, (m as u32 + BN - 1) / BN, 1),
9566 block_dim: (32, 4, 1),
9567 shared_mem_bytes: 0,
9568 };
9569 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9570 let __s_b = self.gpu.stream();
9571 let mut b = __s_b.launch_builder(&f);
9572 b.arg(bytes)
9573 .arg(aq4)
9574 .arg(ad4)
9575 .arg(&mut y)
9576 .arg(&inf)
9577 .arg(&outf)
9578 .arg(&mi)
9579 .arg(&rb);
9580 unsafe {
9581 b.launch(cfg)?;
9582 }
9583 Ok(y)
9584 }
9585
9586 pub fn qmatvec_gemm_nvfp4_fp4_raw(
9588 &self,
9589 bytes: &CudaSlice<u8>,
9590 x: &CudaSlice<f32>,
9591 m: usize,
9592 in_f: usize,
9593 out_f: usize,
9594 row_bytes: usize,
9595 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9596 assert!(
9597 in_f.is_multiple_of(64),
9598 "FP4 GEMM requires in_f % 64 == 0, got {in_f}"
9599 );
9600 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
9601 self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)
9602 }
9603
9604 pub fn qmatvec_q8_0_fast(
9606 &self,
9607 w: &CudaSlice<u8>,
9608 x: &CudaSlice<f32>,
9609 m: usize,
9610 in_f: usize,
9611 out_f: usize,
9612 row_bytes: usize,
9613 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9614 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
9615 let f = self.func("qmatvec_q8_0_dp4a");
9616 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
9618 grid_dim: (out_f as u32, m as u32, 1),
9619 block_dim: (128, 1, 1),
9620 shared_mem_bytes: 0,
9621 };
9622 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9623 let __s_b = self.gpu.stream();
9624 let mut b = __s_b.launch_builder(&f);
9625 b.arg(w)
9626 .arg(&aq)
9627 .arg(&ad)
9628 .arg(&mut y)
9629 .arg(&inf)
9630 .arg(&outf)
9631 .arg(&mi)
9632 .arg(&rb);
9633 unsafe {
9634 b.launch(cfg)?;
9635 }
9636 Ok(y)
9637 }
9638
9639 #[allow(non_snake_case)] pub fn qmatvec_q4_K_fast(
9642 &self,
9643 w: &CudaSlice<u8>,
9644 x: &CudaSlice<f32>,
9645 m: usize,
9646 in_f: usize,
9647 out_f: usize,
9648 row_bytes: usize,
9649 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9650 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
9651 let f = self.func("qmatvec_q4_K_dp4a");
9652 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
9654 grid_dim: (out_f as u32, m as u32, 1),
9655 block_dim: (128, 1, 1),
9656 shared_mem_bytes: 0,
9657 };
9658 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9659 let __s_b = self.gpu.stream();
9660 let mut b = __s_b.launch_builder(&f);
9661 b.arg(w)
9662 .arg(&aq)
9663 .arg(&ad)
9664 .arg(&mut y)
9665 .arg(&inf)
9666 .arg(&outf)
9667 .arg(&mi)
9668 .arg(&rb);
9669 unsafe {
9670 b.launch(cfg)?;
9671 }
9672 Ok(y)
9673 }
9674
9675 #[allow(non_snake_case)] pub fn qmatvec_q6_K_fast(
9678 &self,
9679 w: &CudaSlice<u8>,
9680 x: &CudaSlice<f32>,
9681 m: usize,
9682 in_f: usize,
9683 out_f: usize,
9684 row_bytes: usize,
9685 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9686 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
9687 let f = self.func("qmatvec_q6_K_dp4a");
9688 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
9690 grid_dim: (out_f as u32, m as u32, 1),
9691 block_dim: (128, 1, 1),
9692 shared_mem_bytes: 0,
9693 };
9694 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9695 let __s_b = self.gpu.stream();
9696 let mut b = __s_b.launch_builder(&f);
9697 b.arg(w)
9698 .arg(&aq)
9699 .arg(&ad)
9700 .arg(&mut y)
9701 .arg(&inf)
9702 .arg(&outf)
9703 .arg(&mi)
9704 .arg(&rb);
9705 unsafe {
9706 b.launch(cfg)?;
9707 }
9708 Ok(y)
9709 }
9710
9711 #[allow(non_snake_case)] pub fn qmatvec_q5_K_fast(
9714 &self,
9715 w: &CudaSlice<u8>,
9716 x: &CudaSlice<f32>,
9717 m: usize,
9718 in_f: usize,
9719 out_f: usize,
9720 row_bytes: usize,
9721 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9722 self.qmatvec_dp4a_named(
9723 "qmatvec_q5_K_dp4a",
9724 &w.slice(0..w.len()),
9725 x,
9726 m,
9727 in_f,
9728 out_f,
9729 row_bytes,
9730 )
9731 }
9732 #[allow(non_snake_case)] pub fn qmatvec_q3_K_fast(
9735 &self,
9736 w: &CudaSlice<u8>,
9737 x: &CudaSlice<f32>,
9738 m: usize,
9739 in_f: usize,
9740 out_f: usize,
9741 row_bytes: usize,
9742 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9743 self.qmatvec_dp4a_named(
9744 "qmatvec_q3_K_dp4a",
9745 &w.slice(0..w.len()),
9746 x,
9747 m,
9748 in_f,
9749 out_f,
9750 row_bytes,
9751 )
9752 }
9753 pub fn qmatvec_nvfp4_fast_rp(
9755 &self,
9756 w: &CudaSlice<u8>,
9757 x: &CudaSlice<f32>,
9758 m: usize,
9759 in_f: usize,
9760 out_f: usize,
9761 row_bytes: usize,
9762 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9763 assert!(
9764 in_f.is_multiple_of(64),
9765 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
9766 );
9767 self.qmatvec_dp4a_named(
9768 "qmatvec_nvfp4_dp4a_rp",
9769 &w.slice(0..w.len()),
9770 x,
9771 m,
9772 in_f,
9773 out_f,
9774 row_bytes,
9775 )
9776 }
9777 pub fn qmatvec_nvfp4_fast(
9779 &self,
9780 w: &cudarc::driver::CudaView<'_, u8>,
9781 x: &CudaSlice<f32>,
9782 m: usize,
9783 in_f: usize,
9784 out_f: usize,
9785 row_bytes: usize,
9786 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9787 assert!(
9790 in_f.is_multiple_of(64),
9791 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
9792 );
9793 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a", w, x, m, in_f, out_f, row_bytes)
9794 }
9795 pub fn qmatvec_nvfp4_fast_v2(
9800 &self,
9801 w: &cudarc::driver::CudaView<'_, u8>,
9802 x: &CudaSlice<f32>,
9803 m: usize,
9804 in_f: usize,
9805 out_f: usize,
9806 row_bytes: usize,
9807 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9808 assert!(
9809 in_f.is_multiple_of(64),
9810 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
9811 );
9812 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a_v2", w, x, m, in_f, out_f, row_bytes)
9813 }
9814 #[allow(non_snake_case)] pub fn qmatvec_iq4_XS_fast(
9817 &self,
9818 w: &CudaSlice<u8>,
9819 x: &CudaSlice<f32>,
9820 m: usize,
9821 in_f: usize,
9822 out_f: usize,
9823 row_bytes: usize,
9824 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9825 self.qmatvec_dp4a_named(
9826 "qmatvec_iq4_XS_dp4a",
9827 &w.slice(0..w.len()),
9828 x,
9829 m,
9830 in_f,
9831 out_f,
9832 row_bytes,
9833 )
9834 }
9835
9836 #[allow(clippy::too_many_arguments)] fn qmatvec_dp4a_named(
9839 &self,
9840 name: &str,
9841 w: &cudarc::driver::CudaView<'_, u8>,
9842 x: &CudaSlice<f32>,
9843 m: usize,
9844 in_f: usize,
9845 out_f: usize,
9846 row_bytes: usize,
9847 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9848 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
9849 let f = self.func(name);
9850 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
9852 grid_dim: (out_f as u32, m as u32, 1),
9853 block_dim: (128, 1, 1),
9854 shared_mem_bytes: 0,
9855 };
9856 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9857 let __s_b = self.gpu.stream();
9858 let mut b = __s_b.launch_builder(&f);
9859 b.arg(w)
9860 .arg(&aq)
9861 .arg(&ad)
9862 .arg(&mut y)
9863 .arg(&inf)
9864 .arg(&outf)
9865 .arg(&mi)
9866 .arg(&rb);
9867 unsafe {
9868 b.launch(cfg)?;
9869 }
9870 Ok(y)
9871 }
9872
9873 #[allow(clippy::too_many_arguments)]
9879 pub fn qmatvec_nvfp4_fast_prequant_into(
9880 &self,
9881 w: &CudaSlice<u8>,
9882 aq: &CudaSlice<i8>,
9883 ad: &CudaSlice<f32>,
9884 y: &mut CudaSlice<f32>,
9885 m: usize,
9886 in_f: usize,
9887 out_f: usize,
9888 row_bytes: usize,
9889 ) -> Result<(), Box<dyn std::error::Error>> {
9890 assert!(
9891 in_f.is_multiple_of(64),
9892 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
9893 );
9894 if y.len() < m * out_f {
9895 return Err(format!(
9896 "NVFP4 prequant output {} is shorter than {m}x{out_f}",
9897 y.len()
9898 )
9899 .into());
9900 }
9901 let f = self.func("qmatvec_nvfp4_dp4a");
9902 let cfg = LaunchConfig {
9903 grid_dim: (out_f as u32, m as u32, 1),
9904 block_dim: (128, 1, 1),
9905 shared_mem_bytes: 0,
9906 };
9907 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9908 let __s_b = self.gpu.stream();
9909 let mut b = __s_b.launch_builder(&f);
9910 b.arg(w)
9911 .arg(aq)
9912 .arg(ad)
9913 .arg(y)
9914 .arg(&inf)
9915 .arg(&outf)
9916 .arg(&mi)
9917 .arg(&rb);
9918 unsafe {
9919 b.launch(cfg)?;
9920 }
9921 Ok(())
9922 }
9923
9924 #[allow(clippy::too_many_arguments)]
9927 pub fn matvec_f32_qkv_into(
9928 &self,
9929 wq: &CudaSlice<f32>,
9930 wk: &CudaSlice<f32>,
9931 wv: &CudaSlice<f32>,
9932 wg: &CudaSlice<f32>,
9933 x: &CudaSlice<f32>,
9934 yq: &mut CudaSlice<f32>,
9935 yk: &mut CudaSlice<f32>,
9936 yv: &mut CudaSlice<f32>,
9937 yg: &mut CudaSlice<f32>,
9938 in_f: usize,
9939 out_q: usize,
9940 out_kv: usize,
9941 out_g: usize,
9942 ) -> Result<(), Box<dyn std::error::Error>> {
9943 if !in_f.is_multiple_of(4)
9944 || wq.len() != out_q * in_f
9945 || wk.len() != out_kv * in_f
9946 || wv.len() != out_kv * in_f
9947 || wg.len() < out_g * in_f
9948 || x.len() < in_f
9949 || yq.len() < out_q
9950 || yk.len() < out_kv
9951 || yv.len() < out_kv
9952 || (out_g > 0 && yg.len() < out_g)
9953 {
9954 return Err(format!(
9955 "fused QKV geometry in={in_f} out_q={out_q} out_kv={out_kv} out_g={out_g} \
9956 wq={} wk={} wv={} wg={}",
9957 wq.len(),
9958 wk.len(),
9959 wv.len(),
9960 wg.len()
9961 )
9962 .into());
9963 }
9964 let f = self.func("matvec_f32_qkv");
9965 let cfg = LaunchConfig {
9966 grid_dim: ((out_q + 2 * out_kv + out_g) as u32, 1, 1),
9967 block_dim: (128, 1, 1),
9968 shared_mem_bytes: 0,
9969 };
9970 let (inf, oq, okv, og) = (in_f as i32, out_q as i32, out_kv as i32, out_g as i32);
9971 let __s_b = self.gpu.stream();
9972 let mut b = __s_b.launch_builder(&f);
9973 b.arg(wq)
9974 .arg(wk)
9975 .arg(wv)
9976 .arg(wg)
9977 .arg(x)
9978 .arg(yq)
9979 .arg(yk)
9980 .arg(yv)
9981 .arg(yg)
9982 .arg(&inf)
9983 .arg(&oq)
9984 .arg(&okv)
9985 .arg(&og);
9986 unsafe {
9987 b.launch(cfg)?;
9988 }
9989 Ok(())
9990 }
9991
9992 #[allow(clippy::too_many_arguments)]
10004 pub fn qmatvec_nvfp4_sel_gu_into(
10005 &self,
10006 gate_bank: &CudaSlice<u8>,
10007 up_bank: &CudaSlice<u8>,
10008 sel: &CudaSlice<i32>,
10009 aq: &CudaSlice<i8>,
10010 ad: &CudaSlice<f32>,
10011 yg: &mut CudaSlice<f32>,
10012 yu: &mut CudaSlice<f32>,
10013 n_sel: usize,
10014 in_f: usize,
10015 out_f: usize,
10016 row_bytes: usize,
10017 expert_stride: usize,
10018 slot_major: bool,
10019 ) -> Result<(), Box<dyn std::error::Error>> {
10020 assert!(
10021 in_f.is_multiple_of(64),
10022 "NVFP4 dp4a requires in_f % 64 == 0"
10023 );
10024 if yg.len() < n_sel * out_f || yu.len() < n_sel * out_f || sel.len() < n_sel {
10025 return Err("NVFP4 gu sel geometry".into());
10026 }
10027 if !slot_major {
10028 return Err(
10029 "NVFP4 gu sel fusion reads slot-major rows: these banks are block_nvfp4 \
10030 v1 (arm MEMRA_NVFP4_BANK_SM to build slot-major TP banks)"
10031 .into(),
10032 );
10033 }
10034 static RPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10038 let rpw = *RPW.get_or_init(|| {
10039 std::env::var("MEMRA_NVFP4_SEL_GU_RPW")
10040 .ok()
10041 .and_then(|v| v.parse().ok())
10042 .filter(|r| *r == 2 || *r == 4)
10043 .unwrap_or(1)
10044 });
10045 let rpw = if out_f.is_multiple_of(rpw) { rpw } else { 1 };
10046 static WPR: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10051 let wpr =
10052 *WPR.get_or_init(|| std::env::var("MEMRA_NVFP4_SEL_GU_WPR").as_deref() == Ok("1"));
10053 let f = self.func(match (wpr, rpw) {
10054 (true, _) => "qmatvec_nvfp4_dp4a_sel_v2_gu_wpr",
10055 (_, 4) => "qmatvec_nvfp4_dp4a_sel_v2_gu_r4",
10056 (_, 2) => "qmatvec_nvfp4_dp4a_sel_v2_gu_r2",
10057 _ => "qmatvec_nvfp4_dp4a_sel_v2_gu",
10058 });
10059 let cfg = LaunchConfig {
10060 grid_dim: if wpr {
10061 (((2 * out_f) as u32).div_ceil(4), n_sel as u32, 1)
10062 } else if rpw == 1 {
10063 ((2 * out_f) as u32, n_sel as u32, 1)
10064 } else {
10065 ((out_f / rpw) as u32, n_sel as u32, 1)
10066 },
10067 block_dim: if wpr { (32, 4, 1) } else { (128, 1, 1) },
10068 shared_mem_bytes: 0,
10069 };
10070 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10071 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10072 let (ars, adrs) = (0i64, 0i64);
10073 let __s_b = self.gpu.stream();
10074 let mut b = __s_b.launch_builder(&f);
10075 b.arg(gate_bank)
10076 .arg(up_bank)
10077 .arg(sel)
10078 .arg(aq)
10079 .arg(ad)
10080 .arg(yg)
10081 .arg(yu)
10082 .arg(&inf)
10083 .arg(&outf)
10084 .arg(&ns)
10085 .arg(&rb)
10086 .arg(&es)
10087 .arg(&ars)
10088 .arg(&adrs);
10089 unsafe {
10090 b.launch(cfg)?;
10091 }
10092 Ok(())
10093 }
10094
10095 #[allow(clippy::too_many_arguments)]
10107 pub fn qmatvec_nvfp4_sel_down8_into(
10108 &self,
10109 bank: &CudaSlice<u8>,
10110 sel: &CudaSlice<i32>,
10111 aq: &CudaSlice<i8>,
10112 ad: &CudaSlice<f32>,
10113 route_w: &CudaSlice<f32>,
10114 md: &CudaSlice<f32>,
10115 dst: &mut CudaSlice<f32>,
10116 n_sel: usize,
10117 in_f: usize,
10118 out_f: usize,
10119 row_bytes: usize,
10120 expert_stride: usize,
10121 act_row_stride: usize,
10122 ad_row_stride: usize,
10123 slot_major: bool,
10124 ) -> Result<(), Box<dyn std::error::Error>> {
10125 if !in_f.is_multiple_of(64)
10126 || n_sel == 0
10127 || n_sel > 8
10128 || (in_f >> 5) > 32
10129 || dst.len() < out_f
10130 || sel.len() < n_sel
10131 || route_w.len() < n_sel
10132 {
10133 return Err(format!(
10134 "NVFP4 sel down8 geometry in_f={in_f} out_f={out_f} n_sel={n_sel} dst={}",
10135 dst.len()
10136 )
10137 .into());
10138 }
10139 if !slot_major {
10140 return Err(
10141 "NVFP4 sel down8 reads slot-major rows: this shard is block_nvfp4 v1 \
10142 (arm MEMRA_NVFP4_BANK_SM to build slot-major TP banks)"
10143 .into(),
10144 );
10145 }
10146 let f = self.func("qmatvec_nvfp4_dp4a_sel_v2_down8");
10147 let cfg = LaunchConfig {
10148 grid_dim: (out_f as u32, 1, 1),
10149 block_dim: (32, n_sel as u32, 1),
10150 shared_mem_bytes: 0,
10151 };
10152 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10153 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10154 let (ars, adrs) = (act_row_stride as i64, ad_row_stride as i64);
10155 let __s_b = self.gpu.stream();
10156 let mut b = __s_b.launch_builder(&f);
10157 b.arg(bank)
10158 .arg(sel)
10159 .arg(aq)
10160 .arg(ad)
10161 .arg(route_w)
10162 .arg(md)
10163 .arg(dst)
10164 .arg(&inf)
10165 .arg(&outf)
10166 .arg(&ns)
10167 .arg(&rb)
10168 .arg(&es)
10169 .arg(&ars)
10170 .arg(&adrs);
10171 unsafe {
10172 b.launch(cfg)?;
10173 }
10174 Ok(())
10175 }
10176
10177 #[allow(clippy::too_many_arguments)]
10180 pub fn qmatvec_nvfp4_sel_gu_ep_into(
10181 &self,
10182 gate_bank: &CudaSlice<u8>,
10183 up_bank: &CudaSlice<u8>,
10184 sel: &CudaSlice<i32>,
10185 aq: &CudaSlice<i8>,
10186 ad: &CudaSlice<f32>,
10187 yg: &mut CudaSlice<f32>,
10188 yu: &mut CudaSlice<f32>,
10189 n_sel: usize,
10190 in_f: usize,
10191 out_f: usize,
10192 row_bytes: usize,
10193 expert_stride: usize,
10194 owner: usize,
10195 ) -> Result<(), Box<dyn std::error::Error>> {
10196 assert!(
10197 in_f.is_multiple_of(64),
10198 "NVFP4 dp4a requires in_f % 64 == 0"
10199 );
10200 if yg.len() < n_sel * out_f || yu.len() < n_sel * out_f || sel.len() < n_sel {
10201 return Err("NVFP4 gu ep geometry".into());
10202 }
10203 let f = self.func("qmatvec_nvfp4_dp4a_sel_v2_gu_ep");
10204 let cfg = LaunchConfig {
10205 grid_dim: ((2 * out_f) as u32, n_sel as u32, 1),
10206 block_dim: (128, 1, 1),
10207 shared_mem_bytes: 0,
10208 };
10209 let (inf, outf, ns, own) = (in_f as i32, out_f as i32, n_sel as i32, owner as i32);
10210 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10211 let (ars, adrs) = (0i64, 0i64);
10212 let __s_b = self.gpu.stream();
10213 let mut b = __s_b.launch_builder(&f);
10214 b.arg(gate_bank)
10215 .arg(up_bank)
10216 .arg(sel)
10217 .arg(aq)
10218 .arg(ad)
10219 .arg(yg)
10220 .arg(yu)
10221 .arg(&inf)
10222 .arg(&outf)
10223 .arg(&ns)
10224 .arg(&rb)
10225 .arg(&es)
10226 .arg(&ars)
10227 .arg(&adrs)
10228 .arg(&own);
10229 unsafe {
10230 b.launch(cfg)?;
10231 }
10232 Ok(())
10233 }
10234
10235 #[allow(clippy::too_many_arguments)]
10237 pub fn silu_mul_scaled_q8_1_sel_ep_into(
10238 &self,
10239 gate: &CudaSlice<f32>,
10240 up: &CudaSlice<f32>,
10241 gmac: &CudaSlice<f32>,
10242 umac: &CudaSlice<f32>,
10243 sel: &CudaSlice<i32>,
10244 limit: Option<f32>,
10245 out_q: &mut CudaSlice<i8>,
10246 out_d: &mut CudaSlice<f32>,
10247 n_per: usize,
10248 n_sel: usize,
10249 owner: usize,
10250 ) -> Result<(), Box<dyn std::error::Error>> {
10251 if !n_per.is_multiple_of(32)
10252 || out_q.len() < n_sel * n_per
10253 || out_d.len() < n_sel * n_per / 32
10254 {
10255 return Err("NVFP4 silu ep geometry".into());
10256 }
10257 let f = self.func("silu_mul_scaled_q8_1_sel_ep");
10258 let warps = n_sel * n_per / 32;
10259 let cfg = LaunchConfig {
10260 grid_dim: ((warps as u32).div_ceil(4), 1, 1),
10261 block_dim: (128, 1, 1),
10262 shared_mem_bytes: 0,
10263 };
10264 let (np, ns, own) = (n_per as i32, n_sel as i32, owner as i32);
10265 let (lim, has) = match limit {
10266 Some(l) => (l, 1i32),
10267 None => (0.0f32, 0i32),
10268 };
10269 let __s_b = self.gpu.stream();
10270 let mut b = __s_b.launch_builder(&f);
10271 b.arg(gate)
10272 .arg(up)
10273 .arg(gmac)
10274 .arg(umac)
10275 .arg(sel)
10276 .arg(&lim)
10277 .arg(&has)
10278 .arg(out_q)
10279 .arg(out_d)
10280 .arg(&np)
10281 .arg(&ns)
10282 .arg(&own);
10283 unsafe {
10284 b.launch(cfg)?;
10285 }
10286 Ok(())
10287 }
10288
10289 #[allow(clippy::too_many_arguments)]
10291 pub fn qmatvec_nvfp4_sel_down8_ep_into(
10292 &self,
10293 bank: &CudaSlice<u8>,
10294 sel: &CudaSlice<i32>,
10295 aq: &CudaSlice<i8>,
10296 ad: &CudaSlice<f32>,
10297 route_w: &CudaSlice<f32>,
10298 md: &CudaSlice<f32>,
10299 dst: &mut CudaSlice<f32>,
10300 n_sel: usize,
10301 in_f: usize,
10302 out_f: usize,
10303 row_bytes: usize,
10304 expert_stride: usize,
10305 act_row_stride: usize,
10306 ad_row_stride: usize,
10307 owner: usize,
10308 ) -> Result<(), Box<dyn std::error::Error>> {
10309 if !in_f.is_multiple_of(64)
10310 || n_sel == 0
10311 || n_sel > 8
10312 || (in_f >> 5) > 64
10313 || dst.len() < out_f
10314 {
10315 return Err("NVFP4 down8 ep geometry".into());
10316 }
10317 let f = self.func("qmatvec_nvfp4_dp4a_sel_v2_down8_ep");
10318 let cfg = LaunchConfig {
10319 grid_dim: (out_f as u32, 1, 1),
10320 block_dim: (32, n_sel as u32, 1),
10321 shared_mem_bytes: 0,
10322 };
10323 let (inf, outf, ns, own) = (in_f as i32, out_f as i32, n_sel as i32, owner as i32);
10324 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10325 let (ars, adrs) = (act_row_stride as i64, ad_row_stride as i64);
10326 let __s_b = self.gpu.stream();
10327 let mut b = __s_b.launch_builder(&f);
10328 b.arg(bank)
10329 .arg(sel)
10330 .arg(aq)
10331 .arg(ad)
10332 .arg(route_w)
10333 .arg(md)
10334 .arg(dst)
10335 .arg(&inf)
10336 .arg(&outf)
10337 .arg(&ns)
10338 .arg(&rb)
10339 .arg(&es)
10340 .arg(&ars)
10341 .arg(&adrs)
10342 .arg(&own);
10343 unsafe {
10344 b.launch(cfg)?;
10345 }
10346 Ok(())
10347 }
10348
10349 #[allow(clippy::too_many_arguments)] pub fn qmatvec_nvfp4_sel_into(
10364 &self,
10365 bank: &CudaSlice<u8>,
10366 sel: &CudaSlice<i32>,
10367 aq: &CudaSlice<i8>,
10368 ad: &CudaSlice<f32>,
10369 y: &mut CudaSlice<f32>,
10370 n_sel: usize,
10371 in_f: usize,
10372 out_f: usize,
10373 row_bytes: usize,
10374 expert_stride: usize,
10375 act_row_stride: usize,
10376 ad_row_stride: usize,
10377 slot_major: bool,
10378 ) -> Result<(), Box<dyn std::error::Error>> {
10379 assert!(
10380 in_f.is_multiple_of(64),
10381 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
10382 );
10383 if y.len() < n_sel * out_f || sel.len() < n_sel {
10384 return Err(format!(
10385 "NVFP4 sel output {} / sel {} shorter than {n_sel}x{out_f}",
10386 y.len(),
10387 sel.len()
10388 )
10389 .into());
10390 }
10391 static MR: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
10398 let mode = *MR.get_or_init(|| {
10399 if std::env::var("MEMRA_SEL_STREAM").as_deref() == Ok("1") {
10400 2
10401 } else if std::env::var("MEMRA_SEL_MR").as_deref() == Ok("1") {
10402 1
10403 } else {
10404 0
10405 }
10406 });
10407 let mode = if mode == 2 && in_f > 4096 { 0 } else { mode };
10408 let mode = if slot_major { 3 } else { mode };
10412 static SM_STREAM: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10416 let sm_stream = mode == 3
10417 && *SM_STREAM
10418 .get_or_init(|| std::env::var("MEMRA_NVFP4_SEL_SM_STREAM").as_deref() == Ok("1"))
10419 && row_bytes.is_multiple_of(16)
10420 && in_f <= 4096;
10421 let kname = match (mode, sm_stream) {
10422 (3, true) => "qmatvec_nvfp4_dp4a_sel_v2s",
10423 (3, false) => "qmatvec_nvfp4_dp4a_sel_v2",
10424 (2, _) => "qmatvec_nvfp4_dp4a_sel_stream",
10425 (1, _) => "qmatvec_nvfp4_dp4a_sel_mr4",
10426 _ => "qmatvec_nvfp4_dp4a_sel",
10427 };
10428 {
10434 static SEEN_SEL: std::sync::Mutex<Vec<(&'static str, usize, usize)>> =
10435 std::sync::Mutex::new(Vec::new());
10436 let combo = (kname, in_f, out_f);
10437 let mut seen = SEEN_SEL.lock().unwrap();
10438 if !seen.contains(&combo) {
10439 seen.push(combo);
10440 eprintln!(
10441 "[nvfp4-sel] kernel={kname} slot_major={slot_major} in_f={in_f} \
10442 out_f={out_f} nsb={} row_bytes={row_bytes}",
10443 in_f >> 5
10444 );
10445 }
10446 }
10447 let f = self.func(kname);
10448 let nsb = in_f >> 5;
10453 let fit_block: u32 = if (mode == 0 || mode == 3) && !sm_stream && nsb <= 32 {
10454 32
10455 } else if mode == 1 {
10456 512
10457 } else {
10458 128
10459 };
10460 let cfg = LaunchConfig {
10461 grid_dim: (
10462 if sm_stream {
10463 (out_f as u32).div_ceil(8)
10464 } else {
10465 match mode {
10466 2 => (out_f as u32).div_ceil(16),
10467 1 => (out_f as u32).div_ceil(4),
10468 _ => out_f as u32,
10469 }
10470 },
10471 n_sel as u32,
10472 1,
10473 ),
10474 block_dim: (fit_block, 1, 1),
10475 shared_mem_bytes: 0,
10476 };
10477 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10478 let (rb, es, ars, adrs) = (
10479 row_bytes as i64,
10480 expert_stride as i64,
10481 act_row_stride as i64,
10482 ad_row_stride as i64,
10483 );
10484 let __s_b = self.gpu.stream();
10485 let mut b = __s_b.launch_builder(&f);
10486 b.arg(bank)
10487 .arg(sel)
10488 .arg(aq)
10489 .arg(ad)
10490 .arg(y)
10491 .arg(&inf)
10492 .arg(&outf)
10493 .arg(&ns)
10494 .arg(&rb)
10495 .arg(&es)
10496 .arg(&ars)
10497 .arg(&adrs);
10498 unsafe {
10499 b.launch(cfg)?;
10500 }
10501 Ok(())
10502 }
10503
10504 #[allow(clippy::too_many_arguments)]
10507 pub fn qmatvec_nvfp4_bf16_sel_dual_rows_into(
10508 &self,
10509 gate_bank: &CudaSlice<u8>,
10510 up_bank: &CudaSlice<u8>,
10511 sel: &CudaSlice<i32>,
10512 token_rows: &CudaSlice<i32>,
10513 x_bf16: &CudaSlice<u8>,
10514 gate_out: &mut CudaSlice<f32>,
10515 up_out: &mut CudaSlice<f32>,
10516 n_sel: usize,
10517 in_f: usize,
10518 out_f: usize,
10519 row_bytes: usize,
10520 expert_stride: usize,
10521 tokens: usize,
10522 ) -> Result<(), Box<dyn std::error::Error>> {
10523 if !in_f.is_multiple_of(64)
10524 || sel.len() < n_sel
10525 || token_rows.len() < n_sel
10526 || gate_out.len() < n_sel * out_f
10527 || up_out.len() < n_sel * out_f
10528 || x_bf16.len() < 2 * in_f * tokens
10529 {
10530 return Err(format!(
10531 "W4A16 NVFP4 dual selected rows geometry sel={} token_rows={} x={} gate={} up={} \
10532 n_sel={n_sel} tokens={tokens} in={in_f} out={out_f}",
10533 sel.len(),
10534 token_rows.len(),
10535 x_bf16.len(),
10536 gate_out.len(),
10537 up_out.len(),
10538 )
10539 .into());
10540 }
10541 let adjacent_rows = tokens > 1;
10542 let f = if adjacent_rows {
10543 self.func("qmatvec_nvfp4_bf16_sel_quad_rows")
10544 } else {
10545 self.func("qmatvec_nvfp4_bf16_sel_dual_rows")
10546 };
10547 let cfg = LaunchConfig {
10548 grid_dim: (
10549 if adjacent_rows {
10550 out_f.div_ceil(2) as u32
10551 } else {
10552 (2 * out_f) as u32
10553 },
10554 n_sel as u32,
10555 1,
10556 ),
10557 block_dim: (256, 1, 1),
10558 shared_mem_bytes: 0,
10559 };
10560 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10561 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10562 let __s_b = self.gpu.stream();
10563 let mut b = __s_b.launch_builder(&f);
10564 b.arg(gate_bank)
10565 .arg(up_bank)
10566 .arg(sel)
10567 .arg(token_rows)
10568 .arg(x_bf16)
10569 .arg(gate_out)
10570 .arg(up_out)
10571 .arg(&inf)
10572 .arg(&outf)
10573 .arg(&ns)
10574 .arg(&rb)
10575 .arg(&es);
10576 unsafe {
10577 b.launch(cfg)?;
10578 }
10579 Ok(())
10580 }
10581
10582 #[allow(clippy::too_many_arguments)]
10585 pub fn qmatvec_nvfp4_bf16_ep_dual_slots_into(
10586 &self,
10587 gate_bank: &CudaSlice<u8>,
10588 up_bank: &CudaSlice<u8>,
10589 sel: &CudaSlice<i32>,
10590 x_bf16: &CudaSlice<u8>,
10591 gate_out: &mut CudaSlice<f32>,
10592 up_out: &mut CudaSlice<f32>,
10593 n_pairs: usize,
10594 top_k: usize,
10595 in_f: usize,
10596 out_f: usize,
10597 owner_start: usize,
10598 owner_end: usize,
10599 row_bytes: usize,
10600 expert_stride: usize,
10601 ) -> Result<(), Box<dyn std::error::Error>> {
10602 let tokens = n_pairs.div_ceil(top_k);
10603 if top_k == 0
10604 || owner_start >= owner_end
10605 || !in_f.is_multiple_of(64)
10606 || sel.len() < n_pairs
10607 || gate_out.len() < n_pairs * out_f
10608 || up_out.len() < n_pairs * out_f
10609 || x_bf16.len() < 2 * in_f * tokens
10610 {
10611 return Err(format!(
10612 "W4A16 NVFP4 device EP dual-slot geometry sel={} x={} gate={} up={} \
10613 pairs={n_pairs} top_k={top_k} in={in_f} out={out_f} \
10614 owner={owner_start}..{owner_end}",
10615 sel.len(),
10616 x_bf16.len(),
10617 gate_out.len(),
10618 up_out.len(),
10619 )
10620 .into());
10621 }
10622 let pair_parallel = tokens > 1;
10623 let f = if pair_parallel {
10624 self.func("qmatvec_nvfp4_bf16_ep_quad_pairs")
10625 } else {
10626 self.func("qmatvec_nvfp4_bf16_ep_dual_slots")
10627 };
10628 let cfg = LaunchConfig {
10629 grid_dim: (
10630 if pair_parallel {
10631 out_f.div_ceil(2) as u32
10632 } else {
10633 (2 * out_f) as u32
10634 },
10635 if pair_parallel { n_pairs as u32 } else { 1 },
10636 1,
10637 ),
10638 block_dim: (256, 1, 1),
10639 shared_mem_bytes: 0,
10640 };
10641 let (inf, outf, np, tk) = (in_f as i32, out_f as i32, n_pairs as i32, top_k as i32);
10642 let (os, oe) = (owner_start as i32, owner_end as i32);
10643 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10644 let __s_b = self.gpu.stream();
10645 let mut b = __s_b.launch_builder(&f);
10646 b.arg(gate_bank)
10647 .arg(up_bank)
10648 .arg(sel)
10649 .arg(x_bf16)
10650 .arg(gate_out)
10651 .arg(up_out)
10652 .arg(&inf)
10653 .arg(&outf)
10654 .arg(&np)
10655 .arg(&tk)
10656 .arg(&os)
10657 .arg(&oe)
10658 .arg(&rb)
10659 .arg(&es);
10660 unsafe {
10661 b.launch(cfg)?;
10662 }
10663 Ok(())
10664 }
10665
10666 #[allow(clippy::too_many_arguments)]
10668 pub fn qmatvec_nvfp4_q8_ep_dual_slots_into(
10669 &self,
10670 gate_bank: &CudaSlice<u8>,
10671 up_bank: &CudaSlice<u8>,
10672 sel: &CudaSlice<i32>,
10673 aq: &CudaSlice<i8>,
10674 ad: &CudaSlice<f32>,
10675 gate_out: &mut CudaSlice<f32>,
10676 up_out: &mut CudaSlice<f32>,
10677 n_pairs: usize,
10678 top_k: usize,
10679 in_f: usize,
10680 out_f: usize,
10681 owner_start: usize,
10682 owner_end: usize,
10683 row_bytes: usize,
10684 expert_stride: usize,
10685 ) -> Result<(), Box<dyn std::error::Error>> {
10686 let tokens = n_pairs.div_ceil(top_k);
10687 if top_k == 0
10688 || owner_start >= owner_end
10689 || !in_f.is_multiple_of(64)
10690 || sel.len() < n_pairs
10691 || aq.len() < tokens * in_f
10692 || ad.len() < tokens * (in_f / 32)
10693 || gate_out.len() < n_pairs * out_f
10694 || up_out.len() < n_pairs * out_f
10695 {
10696 return Err(format!(
10697 "W4A8 NVFP4 device EP gate/up geometry sel={} aq={} ad={} gate={} up={} \
10698 pairs={n_pairs} top_k={top_k} in={in_f} out={out_f} \
10699 owner={owner_start}..{owner_end}",
10700 sel.len(),
10701 aq.len(),
10702 ad.len(),
10703 gate_out.len(),
10704 up_out.len(),
10705 )
10706 .into());
10707 }
10708 let f = self.func("qmatvec_nvfp4_q8_ep_dual_slots");
10709 let threads = ((in_f / 32).div_ceil(32) * 32).clamp(32, 256) as u32;
10710 let cfg = LaunchConfig {
10711 grid_dim: (out_f as u32, 1, 1),
10712 block_dim: (threads, 1, 1),
10713 shared_mem_bytes: 0,
10714 };
10715 let (inf, outf, np, tk) = (in_f as i32, out_f as i32, n_pairs as i32, top_k as i32);
10716 let (os, oe) = (owner_start as i32, owner_end as i32);
10717 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10718 let __s_b = self.gpu.stream();
10719 let mut b = __s_b.launch_builder(&f);
10720 b.arg(gate_bank)
10721 .arg(up_bank)
10722 .arg(sel)
10723 .arg(aq)
10724 .arg(ad)
10725 .arg(gate_out)
10726 .arg(up_out)
10727 .arg(&inf)
10728 .arg(&outf)
10729 .arg(&np)
10730 .arg(&tk)
10731 .arg(&os)
10732 .arg(&oe)
10733 .arg(&rb)
10734 .arg(&es);
10735 unsafe {
10736 b.launch(cfg)?;
10737 }
10738 Ok(())
10739 }
10740
10741 #[allow(clippy::too_many_arguments)]
10744 pub fn qmatvec_nvfp4_q8_ep_paired_slots_into(
10745 &self,
10746 gate_bank: &CudaSlice<u8>,
10747 up_bank: &CudaSlice<u8>,
10748 sel: &CudaSlice<i32>,
10749 aq: &CudaSlice<i8>,
10750 ad: &CudaSlice<f32>,
10751 gate_out: &mut CudaSlice<f32>,
10752 up_out: &mut CudaSlice<f32>,
10753 n_pairs: usize,
10754 top_k: usize,
10755 in_f: usize,
10756 out_f: usize,
10757 owner_start: usize,
10758 owner_end: usize,
10759 row_bytes: usize,
10760 expert_stride: usize,
10761 ) -> Result<(), Box<dyn std::error::Error>> {
10762 let tokens = n_pairs.div_ceil(top_k);
10763 if top_k == 0
10764 || owner_start >= owner_end
10765 || !in_f.is_multiple_of(64)
10766 || sel.len() < n_pairs
10767 || aq.len() < tokens * in_f
10768 || ad.len() < tokens * (in_f / 32)
10769 || gate_out.len() < n_pairs * out_f
10770 || up_out.len() < n_pairs * out_f
10771 {
10772 return Err(format!(
10773 "W4A8 NVFP4 paired gate/up geometry sel={} aq={} ad={} gate={} up={} \
10774 pairs={n_pairs} top_k={top_k} in={in_f} out={out_f} \
10775 owner={owner_start}..{owner_end}",
10776 sel.len(),
10777 aq.len(),
10778 ad.len(),
10779 gate_out.len(),
10780 up_out.len(),
10781 )
10782 .into());
10783 }
10784 let f = self.func("qmatvec_nvfp4_q8_ep_paired_slots");
10785 let threads = ((in_f / 32).div_ceil(32) * 32).clamp(32, 256) as u32;
10786 let cfg = LaunchConfig {
10787 grid_dim: (out_f as u32, 1, 1),
10788 block_dim: (threads, 1, 1),
10789 shared_mem_bytes: 0,
10790 };
10791 let (inf, outf, np, tk) = (in_f as i32, out_f as i32, n_pairs as i32, top_k as i32);
10792 let (os, oe) = (owner_start as i32, owner_end as i32);
10793 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10794 let __s_b = self.gpu.stream();
10795 let mut b = __s_b.launch_builder(&f);
10796 b.arg(gate_bank)
10797 .arg(up_bank)
10798 .arg(sel)
10799 .arg(aq)
10800 .arg(ad)
10801 .arg(gate_out)
10802 .arg(up_out)
10803 .arg(&inf)
10804 .arg(&outf)
10805 .arg(&np)
10806 .arg(&tk)
10807 .arg(&os)
10808 .arg(&oe)
10809 .arg(&rb)
10810 .arg(&es);
10811 unsafe {
10812 b.launch(cfg)?;
10813 }
10814 Ok(())
10815 }
10816
10817 #[allow(clippy::too_many_arguments)]
10819 pub fn qmatvec_nvfp4_bf16_sel_down_rows_raw(
10820 &self,
10821 bank: &CudaSlice<u8>,
10822 sel: &CudaSlice<i32>,
10823 global_pairs: &CudaSlice<i32>,
10824 activation_bf16: &CudaSlice<u8>,
10825 macros_down: &CudaSlice<f32>,
10826 dst_raw: u64,
10827 n_sel: usize,
10828 in_f: usize,
10829 out_f: usize,
10830 row_bytes: usize,
10831 expert_stride: usize,
10832 total_pairs: usize,
10833 ) -> Result<(), Box<dyn std::error::Error>> {
10834 if !in_f.is_multiple_of(64)
10835 || sel.len() < n_sel
10836 || global_pairs.len() < n_sel
10837 || activation_bf16.len() < 2 * n_sel * in_f
10838 || dst_raw == 0
10839 {
10840 return Err(format!(
10841 "W4A16 NVFP4 down rows geometry sel={} pairs={} act={} dst_raw={dst_raw:#x} \
10842 n_sel={n_sel} total_pairs={total_pairs} in={in_f} out={out_f}",
10843 sel.len(),
10844 global_pairs.len(),
10845 activation_bf16.len(),
10846 )
10847 .into());
10848 }
10849 let f = self.func("qmatvec_nvfp4_bf16_sel_down_rows");
10850 let cfg = LaunchConfig {
10851 grid_dim: (out_f.div_ceil(2) as u32, n_sel as u32, 1),
10852 block_dim: (256, 1, 1),
10853 shared_mem_bytes: 0,
10854 };
10855 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10856 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10857 let __s_b = self.gpu.stream();
10858 let mut b = __s_b.launch_builder(&f);
10859 b.arg(bank)
10860 .arg(sel)
10861 .arg(global_pairs)
10862 .arg(activation_bf16)
10863 .arg(macros_down)
10864 .arg(&dst_raw)
10865 .arg(&inf)
10866 .arg(&outf)
10867 .arg(&ns)
10868 .arg(&rb)
10869 .arg(&es);
10870 unsafe {
10871 b.launch(cfg)?;
10872 }
10873 Ok(())
10874 }
10875
10876 #[allow(clippy::too_many_arguments)]
10879 pub fn qmatvec_nvfp4_bf16_ep_down_slots_raw(
10880 &self,
10881 bank: &CudaSlice<u8>,
10882 sel: &CudaSlice<i32>,
10883 activation_bf16: &CudaSlice<u8>,
10884 macros_down: &CudaSlice<f32>,
10885 dst_raw: u64,
10886 n_pairs: usize,
10887 in_f: usize,
10888 out_f: usize,
10889 owner_start: usize,
10890 owner_end: usize,
10891 row_bytes: usize,
10892 expert_stride: usize,
10893 ) -> Result<(), Box<dyn std::error::Error>> {
10894 if owner_start >= owner_end
10895 || !in_f.is_multiple_of(64)
10896 || sel.len() < n_pairs
10897 || activation_bf16.len() < 2 * n_pairs * in_f
10898 || dst_raw == 0
10899 {
10900 return Err(format!(
10901 "W4A16 NVFP4 device EP down-slot geometry sel={} act={} dst_raw={dst_raw:#x} \
10902 pairs={n_pairs} in={in_f} out={out_f} owner={owner_start}..{owner_end}",
10903 sel.len(),
10904 activation_bf16.len(),
10905 )
10906 .into());
10907 }
10908 let f = self.func("qmatvec_nvfp4_bf16_ep_down_slots");
10909 let cfg = LaunchConfig {
10910 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
10911 block_dim: (256, 1, 1),
10912 shared_mem_bytes: 0,
10913 };
10914 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
10915 let (os, oe) = (owner_start as i32, owner_end as i32);
10916 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10917 let __s_b = self.gpu.stream();
10918 let mut b = __s_b.launch_builder(&f);
10919 b.arg(bank)
10920 .arg(sel)
10921 .arg(activation_bf16)
10922 .arg(macros_down)
10923 .arg(&dst_raw)
10924 .arg(&inf)
10925 .arg(&outf)
10926 .arg(&np)
10927 .arg(&os)
10928 .arg(&oe)
10929 .arg(&rb)
10930 .arg(&es);
10931 unsafe {
10932 b.launch(cfg)?;
10933 }
10934 Ok(())
10935 }
10936
10937 #[allow(clippy::too_many_arguments)]
10939 pub fn qmatvec_nvfp4_bf16_ep_down_pairs_raw(
10940 &self,
10941 bank: &CudaSlice<u8>,
10942 sel: &CudaSlice<i32>,
10943 activation_bf16: &CudaSlice<u8>,
10944 macros_down: &CudaSlice<f32>,
10945 dst_raw: u64,
10946 n_pairs: usize,
10947 in_f: usize,
10948 out_f: usize,
10949 owner_start: usize,
10950 owner_end: usize,
10951 row_bytes: usize,
10952 expert_stride: usize,
10953 ) -> Result<(), Box<dyn std::error::Error>> {
10954 if owner_start >= owner_end
10955 || !in_f.is_multiple_of(64)
10956 || sel.len() < n_pairs
10957 || activation_bf16.len() < 2 * n_pairs * in_f
10958 || dst_raw == 0
10959 {
10960 return Err(format!(
10961 "W4A16 NVFP4 device EP down-pair geometry sel={} act={} dst_raw={dst_raw:#x} \
10962 pairs={n_pairs} in={in_f} out={out_f} owner={owner_start}..{owner_end}",
10963 sel.len(),
10964 activation_bf16.len(),
10965 )
10966 .into());
10967 }
10968 let f = self.func("qmatvec_nvfp4_bf16_ep_down_pairs");
10969 let cfg = LaunchConfig {
10970 grid_dim: (out_f.div_ceil(2) as u32, n_pairs as u32, 1),
10971 block_dim: (256, 1, 1),
10972 shared_mem_bytes: 0,
10973 };
10974 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
10975 let (os, oe) = (owner_start as i32, owner_end as i32);
10976 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10977 let __s_b = self.gpu.stream();
10978 let mut b = __s_b.launch_builder(&f);
10979 b.arg(bank)
10980 .arg(sel)
10981 .arg(activation_bf16)
10982 .arg(macros_down)
10983 .arg(&dst_raw)
10984 .arg(&inf)
10985 .arg(&outf)
10986 .arg(&np)
10987 .arg(&os)
10988 .arg(&oe)
10989 .arg(&rb)
10990 .arg(&es);
10991 unsafe {
10992 b.launch(cfg)?;
10993 }
10994 Ok(())
10995 }
10996
10997 #[allow(clippy::too_many_arguments)]
10999 pub fn silu_mul_scaled_host_expf_bf16_sel_into(
11000 &self,
11001 gate: &CudaSlice<f32>,
11002 up: &CudaSlice<f32>,
11003 gate_macros: &CudaSlice<f32>,
11004 up_macros: &CudaSlice<f32>,
11005 sel: &CudaSlice<i32>,
11006 limit: Option<f32>,
11007 output_bf16: &mut CudaSlice<u8>,
11008 n_per: usize,
11009 n_sel: usize,
11010 ) -> Result<(), Box<dyn std::error::Error>> {
11011 let n = n_per * n_sel;
11012 if sel.len() < n_sel || gate.len() < n || up.len() < n || output_bf16.len() < 2 * n {
11013 return Err(format!(
11014 "W4A16 selected activation geometry sel={} gate={} up={} out={} \
11015 n_per={n_per} n_sel={n_sel}",
11016 sel.len(),
11017 gate.len(),
11018 up.len(),
11019 output_bf16.len(),
11020 )
11021 .into());
11022 }
11023 let (limit, has_limit) = match limit {
11024 Some(limit) if limit.is_finite() && limit > 1e-6 => (limit, 1i32),
11025 Some(limit) => {
11026 return Err(format!("W4A16 selected activation limit {limit} is invalid").into());
11027 }
11028 None => (0.0f32, 0i32),
11029 };
11030 let f = self.func("silu_mul_scaled_host_expf_bf16_sel");
11031 let cfg = LaunchConfig::for_num_elems(n as u32);
11032 let (np, ns) = (n_per as i32, n_sel as i32);
11033 let __s_b = self.gpu.stream();
11034 let mut b = __s_b.launch_builder(&f);
11035 b.arg(gate)
11036 .arg(up)
11037 .arg(gate_macros)
11038 .arg(up_macros)
11039 .arg(sel)
11040 .arg(&limit)
11041 .arg(&has_limit)
11042 .arg(output_bf16)
11043 .arg(&np)
11044 .arg(&ns);
11045 unsafe {
11046 b.launch(cfg)?;
11047 }
11048 Ok(())
11049 }
11050
11051 #[allow(clippy::too_many_arguments)]
11054 pub fn silu_mul_scaled_host_expf_bf16_ep_slots_into(
11055 &self,
11056 gate: &CudaSlice<f32>,
11057 up: &CudaSlice<f32>,
11058 gate_macros: &CudaSlice<f32>,
11059 up_macros: &CudaSlice<f32>,
11060 sel: &CudaSlice<i32>,
11061 owner_start: usize,
11062 owner_end: usize,
11063 limit: Option<f32>,
11064 output_bf16: &mut CudaSlice<u8>,
11065 n_per: usize,
11066 n_pairs: usize,
11067 ) -> Result<(), Box<dyn std::error::Error>> {
11068 let n = n_per * n_pairs;
11069 if owner_start >= owner_end
11070 || sel.len() < n_pairs
11071 || gate.len() < n
11072 || up.len() < n
11073 || output_bf16.len() < 2 * n
11074 {
11075 return Err(format!(
11076 "W4A16 device EP activation geometry sel={} gate={} up={} out={} \
11077 n_per={n_per} pairs={n_pairs} owner={owner_start}..{owner_end}",
11078 sel.len(),
11079 gate.len(),
11080 up.len(),
11081 output_bf16.len(),
11082 )
11083 .into());
11084 }
11085 let (limit, has_limit) = match limit {
11086 Some(limit) if limit.is_finite() && limit > 1e-6 => (limit, 1i32),
11087 Some(limit) => {
11088 return Err(format!("W4A16 selected activation limit {limit} is invalid").into());
11089 }
11090 None => (0.0f32, 0i32),
11091 };
11092 let f = self.func("silu_mul_scaled_host_expf_bf16_ep_slots");
11093 let cfg = LaunchConfig::for_num_elems(n as u32);
11094 let (np, pairs) = (n_per as i32, n_pairs as i32);
11095 let (os, oe) = (owner_start as i32, owner_end as i32);
11096 let __s_b = self.gpu.stream();
11097 let mut b = __s_b.launch_builder(&f);
11098 b.arg(gate)
11099 .arg(up)
11100 .arg(gate_macros)
11101 .arg(up_macros)
11102 .arg(sel)
11103 .arg(&limit)
11104 .arg(&has_limit)
11105 .arg(output_bf16)
11106 .arg(&np)
11107 .arg(&pairs)
11108 .arg(&os)
11109 .arg(&oe);
11110 unsafe {
11111 b.launch(cfg)?;
11112 }
11113 Ok(())
11114 }
11115
11116 #[allow(clippy::too_many_arguments)]
11118 pub fn silu_mul_scaled_host_expf_q8_ep_slots_into(
11119 &self,
11120 gate: &CudaSlice<f32>,
11121 up: &CudaSlice<f32>,
11122 gate_macros: &CudaSlice<f32>,
11123 up_macros: &CudaSlice<f32>,
11124 sel: &CudaSlice<i32>,
11125 owner_start: usize,
11126 owner_end: usize,
11127 limit: Option<f32>,
11128 output_q8: &mut CudaSlice<i8>,
11129 output_scales: &mut CudaSlice<f32>,
11130 n_per: usize,
11131 n_pairs: usize,
11132 ) -> Result<(), Box<dyn std::error::Error>> {
11133 let n = n_per * n_pairs;
11134 if owner_start >= owner_end
11135 || !n_per.is_multiple_of(32)
11136 || sel.len() < n_pairs
11137 || gate.len() < n
11138 || up.len() < n
11139 || output_q8.len() < n
11140 || output_scales.len() < n / 32
11141 {
11142 return Err(format!(
11143 "W4A8 device EP activation geometry sel={} gate={} up={} q8={} scales={} \
11144 n_per={n_per} pairs={n_pairs} owner={owner_start}..{owner_end}",
11145 sel.len(),
11146 gate.len(),
11147 up.len(),
11148 output_q8.len(),
11149 output_scales.len(),
11150 )
11151 .into());
11152 }
11153 let (limit, has_limit) = match limit {
11154 Some(limit) if limit.is_finite() && limit > 1e-6 => (limit, 1i32),
11155 Some(limit) => {
11156 return Err(format!("W4A8 selected activation limit {limit} is invalid").into());
11157 }
11158 None => (0.0f32, 0i32),
11159 };
11160 let f = self.func("silu_mul_scaled_host_expf_q8_ep_slots");
11161 let warps = n / 32;
11162 let cfg = LaunchConfig {
11163 grid_dim: ((warps as u32).div_ceil(4), 1, 1),
11164 block_dim: (128, 1, 1),
11165 shared_mem_bytes: 0,
11166 };
11167 let (np, pairs) = (n_per as i32, n_pairs as i32);
11168 let (os, oe) = (owner_start as i32, owner_end as i32);
11169 let __s_b = self.gpu.stream();
11170 let mut b = __s_b.launch_builder(&f);
11171 b.arg(gate)
11172 .arg(up)
11173 .arg(gate_macros)
11174 .arg(up_macros)
11175 .arg(sel)
11176 .arg(&limit)
11177 .arg(&has_limit)
11178 .arg(output_q8)
11179 .arg(output_scales)
11180 .arg(&np)
11181 .arg(&pairs)
11182 .arg(&os)
11183 .arg(&oe);
11184 unsafe {
11185 b.launch(cfg)?;
11186 }
11187 Ok(())
11188 }
11189
11190 #[allow(clippy::too_many_arguments)]
11193 pub fn qmatvec_nvfp4_bf16_sel_down_fma_into(
11194 &self,
11195 bank: &CudaSlice<u8>,
11196 sel: &CudaSlice<i32>,
11197 activation_bf16: &CudaSlice<u8>,
11198 route_weights: &CudaSlice<f32>,
11199 macros_down: &CudaSlice<f32>,
11200 dst: &mut cudarc::driver::CudaViewMut<f32>,
11201 n_sel: usize,
11202 in_f: usize,
11203 out_f: usize,
11204 row_bytes: usize,
11205 expert_stride: usize,
11206 ) -> Result<(), Box<dyn std::error::Error>> {
11207 if !in_f.is_multiple_of(64)
11208 || sel.len() < n_sel
11209 || route_weights.len() < n_sel
11210 || activation_bf16.len() < 2 * n_sel * in_f
11211 || dst.len() < out_f
11212 {
11213 return Err(format!(
11214 "W4A16 NVFP4 down selected geometry sel={} act={} weights={} dst={} \
11215 n_sel={n_sel} in={in_f} out={out_f}",
11216 sel.len(),
11217 activation_bf16.len(),
11218 route_weights.len(),
11219 dst.len(),
11220 )
11221 .into());
11222 }
11223 let f = self.func("qmatvec_nvfp4_bf16_sel_down_fma");
11224 let cfg = LaunchConfig {
11225 grid_dim: (out_f as u32, 1, 1),
11226 block_dim: (256, 1, 1),
11227 shared_mem_bytes: 0,
11228 };
11229 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
11230 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11231 let __s_b = self.gpu.stream();
11232 let mut b = __s_b.launch_builder(&f);
11233 b.arg(bank)
11234 .arg(sel)
11235 .arg(activation_bf16)
11236 .arg(route_weights)
11237 .arg(macros_down)
11238 .arg(dst)
11239 .arg(&inf)
11240 .arg(&outf)
11241 .arg(&ns)
11242 .arg(&rb)
11243 .arg(&es);
11244 unsafe {
11245 b.launch(cfg)?;
11246 }
11247 Ok(())
11248 }
11249
11250 #[allow(clippy::too_many_arguments)]
11252 pub fn qmatvec_nvfp4_bf16_ep_down_fma_into(
11253 &self,
11254 bank: &CudaSlice<u8>,
11255 sel: &CudaSlice<i32>,
11256 activation_bf16: &CudaSlice<u8>,
11257 route_weights: &CudaSlice<f32>,
11258 macros_down: &CudaSlice<f32>,
11259 dst: &mut cudarc::driver::CudaViewMut<f32>,
11260 n_pairs: usize,
11261 in_f: usize,
11262 out_f: usize,
11263 owner_start: usize,
11264 owner_end: usize,
11265 row_bytes: usize,
11266 expert_stride: usize,
11267 ) -> Result<(), Box<dyn std::error::Error>> {
11268 if owner_start >= owner_end
11269 || !in_f.is_multiple_of(64)
11270 || sel.len() < n_pairs
11271 || route_weights.len() < n_pairs
11272 || activation_bf16.len() < 2 * n_pairs * in_f
11273 || dst.len() < out_f
11274 {
11275 return Err(format!(
11276 "W4A16 device EP down-FMA geometry sel={} act={} weights={} dst={} \
11277 pairs={n_pairs} in={in_f} out={out_f} owner={owner_start}..{owner_end}",
11278 sel.len(),
11279 activation_bf16.len(),
11280 route_weights.len(),
11281 dst.len(),
11282 )
11283 .into());
11284 }
11285 let f = self.func("qmatvec_nvfp4_bf16_ep_down_fma");
11286 let cfg = LaunchConfig {
11287 grid_dim: (out_f as u32, 1, 1),
11288 block_dim: (256, 1, 1),
11289 shared_mem_bytes: 0,
11290 };
11291 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
11292 let (os, oe) = (owner_start as i32, owner_end as i32);
11293 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11294 let __s_b = self.gpu.stream();
11295 let mut b = __s_b.launch_builder(&f);
11296 b.arg(bank)
11297 .arg(sel)
11298 .arg(activation_bf16)
11299 .arg(route_weights)
11300 .arg(macros_down)
11301 .arg(dst)
11302 .arg(&inf)
11303 .arg(&outf)
11304 .arg(&np)
11305 .arg(&os)
11306 .arg(&oe)
11307 .arg(&rb)
11308 .arg(&es);
11309 unsafe {
11310 b.launch(cfg)?;
11311 }
11312 Ok(())
11313 }
11314
11315 #[allow(clippy::too_many_arguments)]
11318 pub fn qmatvec_nvfp4_bf16_ep_down_fma_raw(
11319 &self,
11320 bank: &CudaSlice<u8>,
11321 sel: &CudaSlice<i32>,
11322 activation_bf16: &CudaSlice<u8>,
11323 route_weights: &CudaSlice<f32>,
11324 macros_down: &CudaSlice<f32>,
11325 dst_raw: u64,
11326 n_pairs: usize,
11327 in_f: usize,
11328 out_f: usize,
11329 owner_start: usize,
11330 owner_end: usize,
11331 row_bytes: usize,
11332 expert_stride: usize,
11333 ) -> Result<(), Box<dyn std::error::Error>> {
11334 if dst_raw == 0
11335 || owner_start >= owner_end
11336 || !in_f.is_multiple_of(64)
11337 || sel.len() < n_pairs
11338 || route_weights.len() < n_pairs
11339 || activation_bf16.len() < 2 * n_pairs * in_f
11340 {
11341 return Err(format!(
11342 "W4A16 device EP raw down-FMA geometry sel={} act={} weights={} \
11343 dst={dst_raw:#x} pairs={n_pairs} in={in_f} out={out_f} \
11344 owner={owner_start}..{owner_end}",
11345 sel.len(),
11346 activation_bf16.len(),
11347 route_weights.len(),
11348 )
11349 .into());
11350 }
11351 let f = self.func("qmatvec_nvfp4_bf16_ep_down_fma");
11352 let cfg = LaunchConfig {
11353 grid_dim: (out_f as u32, 1, 1),
11354 block_dim: (256, 1, 1),
11355 shared_mem_bytes: 0,
11356 };
11357 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
11358 let (os, oe) = (owner_start as i32, owner_end as i32);
11359 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11360 let __s_b = self.gpu.stream();
11361 let mut b = __s_b.launch_builder(&f);
11362 b.arg(bank)
11363 .arg(sel)
11364 .arg(activation_bf16)
11365 .arg(route_weights)
11366 .arg(macros_down)
11367 .arg(&dst_raw)
11368 .arg(&inf)
11369 .arg(&outf)
11370 .arg(&np)
11371 .arg(&os)
11372 .arg(&oe)
11373 .arg(&rb)
11374 .arg(&es);
11375 unsafe {
11376 b.launch(cfg)?;
11377 }
11378 Ok(())
11379 }
11380
11381 #[allow(clippy::too_many_arguments)]
11384 pub fn qmatvec_nvfp4_q8_ep_down_slots_raw(
11385 &self,
11386 bank: &CudaSlice<u8>,
11387 sel: &CudaSlice<i32>,
11388 aq: &CudaSlice<i8>,
11389 ad: &CudaSlice<f32>,
11390 macros_down: &CudaSlice<f32>,
11391 dst_raw: u64,
11392 n_pairs: usize,
11393 in_f: usize,
11394 out_f: usize,
11395 owner_start: usize,
11396 owner_end: usize,
11397 row_bytes: usize,
11398 expert_stride: usize,
11399 ) -> Result<(), Box<dyn std::error::Error>> {
11400 if dst_raw == 0
11401 || owner_start >= owner_end
11402 || !in_f.is_multiple_of(64)
11403 || sel.len() < n_pairs
11404 || aq.len() < n_pairs * in_f
11405 || ad.len() < n_pairs * (in_f / 32)
11406 {
11407 return Err(format!(
11408 "W4A8 device EP raw down-slot geometry sel={} aq={} ad={} \
11409 dst={dst_raw:#x} pairs={n_pairs} in={in_f} out={out_f} \
11410 owner={owner_start}..{owner_end}",
11411 sel.len(),
11412 aq.len(),
11413 ad.len(),
11414 )
11415 .into());
11416 }
11417 let f = self.func("qmatvec_nvfp4_q8_ep_down_slots");
11418 let threads = ((in_f / 32).div_ceil(32) * 32).clamp(32, 256) as u32;
11419 let cfg = LaunchConfig {
11420 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
11421 block_dim: (threads, 1, 1),
11422 shared_mem_bytes: 0,
11423 };
11424 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
11425 let (os, oe) = (owner_start as i32, owner_end as i32);
11426 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11427 let __s_b = self.gpu.stream();
11428 let mut b = __s_b.launch_builder(&f);
11429 b.arg(bank)
11430 .arg(sel)
11431 .arg(aq)
11432 .arg(ad)
11433 .arg(macros_down)
11434 .arg(&dst_raw)
11435 .arg(&inf)
11436 .arg(&outf)
11437 .arg(&np)
11438 .arg(&os)
11439 .arg(&oe)
11440 .arg(&rb)
11441 .arg(&es);
11442 unsafe {
11443 b.launch(cfg)?;
11444 }
11445 Ok(())
11446 }
11447
11448 #[allow(clippy::too_many_arguments)]
11450 pub fn qmatvec_nvfp4_q8_ep_down_fma_raw(
11451 &self,
11452 bank: &CudaSlice<u8>,
11453 sel: &CudaSlice<i32>,
11454 aq: &CudaSlice<i8>,
11455 ad: &CudaSlice<f32>,
11456 route_weights: &CudaSlice<f32>,
11457 macros_down: &CudaSlice<f32>,
11458 dst_raw: u64,
11459 n_pairs: usize,
11460 in_f: usize,
11461 out_f: usize,
11462 owner_start: usize,
11463 owner_end: usize,
11464 row_bytes: usize,
11465 expert_stride: usize,
11466 ) -> Result<(), Box<dyn std::error::Error>> {
11467 if dst_raw == 0
11468 || owner_start >= owner_end
11469 || !in_f.is_multiple_of(64)
11470 || sel.len() < n_pairs
11471 || aq.len() < n_pairs * in_f
11472 || ad.len() < n_pairs * (in_f / 32)
11473 || route_weights.len() < n_pairs
11474 {
11475 return Err(format!(
11476 "W4A8 device EP raw down-FMA geometry sel={} aq={} ad={} weights={} \
11477 dst={dst_raw:#x} pairs={n_pairs} in={in_f} out={out_f} \
11478 owner={owner_start}..{owner_end}",
11479 sel.len(),
11480 aq.len(),
11481 ad.len(),
11482 route_weights.len(),
11483 )
11484 .into());
11485 }
11486 let f = self.func("qmatvec_nvfp4_q8_ep_down_fma");
11487 let cfg = LaunchConfig {
11488 grid_dim: (out_f as u32, 1, 1),
11489 block_dim: (256, 1, 1),
11490 shared_mem_bytes: 0,
11491 };
11492 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
11493 let (os, oe) = (owner_start as i32, owner_end as i32);
11494 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11495 let __s_b = self.gpu.stream();
11496 let mut b = __s_b.launch_builder(&f);
11497 b.arg(bank)
11498 .arg(sel)
11499 .arg(aq)
11500 .arg(ad)
11501 .arg(route_weights)
11502 .arg(macros_down)
11503 .arg(&dst_raw)
11504 .arg(&inf)
11505 .arg(&outf)
11506 .arg(&np)
11507 .arg(&os)
11508 .arg(&oe)
11509 .arg(&rb)
11510 .arg(&es);
11511 unsafe {
11512 b.launch(cfg)?;
11513 }
11514 Ok(())
11515 }
11516
11517 #[allow(clippy::too_many_arguments)]
11522 pub fn silu_mul_scaled_q8_1_sel_into(
11523 &self,
11524 gate: &CudaSlice<f32>,
11525 up: &CudaSlice<f32>,
11526 gmac: &CudaSlice<f32>,
11527 umac: &CudaSlice<f32>,
11528 sel: &CudaSlice<i32>,
11529 limit: Option<f32>,
11530 out_q: &mut CudaSlice<i8>,
11531 out_d: &mut CudaSlice<f32>,
11532 n_per: usize,
11533 n_sel: usize,
11534 ) -> Result<(), Box<dyn std::error::Error>> {
11535 let n = n_per * n_sel;
11536 if !n_per.is_multiple_of(32) || out_q.len() < n || out_d.len() < n / 32 {
11537 return Err(format!(
11538 "silu sel geometry n_per={n_per} n_sel={n_sel} q={} d={}",
11539 out_q.len(),
11540 out_d.len()
11541 )
11542 .into());
11543 }
11544 if let Some(limit) = limit {
11545 if limit <= 1e-6 {
11546 return Err(format!(
11547 "silu sel clamp limit {limit} is at or below the 1e-6 eps gate"
11548 )
11549 .into());
11550 }
11551 let f = self.func("silu_mul_scaled_q8_1_sel_clamp");
11552 let cfg = LaunchConfig::for_num_elems(n as u32);
11553 let (np, ns) = (n_per as i32, n_sel as i32);
11554 let __s_b = self.gpu.stream();
11555 let mut b = __s_b.launch_builder(&f);
11556 b.arg(gate)
11557 .arg(up)
11558 .arg(gmac)
11559 .arg(umac)
11560 .arg(sel)
11561 .arg(&limit)
11562 .arg(out_q)
11563 .arg(out_d)
11564 .arg(&np)
11565 .arg(&ns);
11566 unsafe {
11567 b.launch(cfg)?;
11568 }
11569 return Ok(());
11570 }
11571 let f = self.func("silu_mul_scaled_q8_1_sel");
11572 let cfg = LaunchConfig::for_num_elems(n as u32);
11573 let (np, ns) = (n_per as i32, n_sel as i32);
11574 let __s_b = self.gpu.stream();
11575 let mut b = __s_b.launch_builder(&f);
11576 b.arg(gate)
11577 .arg(up)
11578 .arg(gmac)
11579 .arg(umac)
11580 .arg(sel)
11581 .arg(out_q)
11582 .arg(out_d)
11583 .arg(&np)
11584 .arg(&ns);
11585 unsafe {
11586 b.launch(cfg)?;
11587 }
11588 Ok(())
11589 }
11590
11591 pub fn htod(&self, v: &[f32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11592 Ok(self.gpu.stream().clone_htod(v)?)
11593 }
11594 pub fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
11595 Ok(self.gpu.stream().clone_htod(v)?)
11596 }
11597 pub fn htod_i8(&self, v: &[i8]) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
11599 Ok(self.gpu.stream().clone_htod(v)?)
11600 }
11601 pub fn htod_u64(&self, v: &[u64]) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
11602 Ok(self.gpu.stream().clone_htod(v)?)
11603 }
11604 pub fn dtoh_view(
11606 &self,
11607 d: &cudarc::driver::CudaView<f32>,
11608 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
11609 let v = self.gpu.stream().clone_dtoh(d)?;
11610 self.gpu.stream().synchronize()?;
11611 Ok(v)
11612 }
11613 pub fn dtoh(&self, d: &CudaSlice<f32>) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
11614 let v = self.gpu.stream().clone_dtoh(d)?;
11615 self.gpu.stream().synchronize()?;
11616 Ok(v)
11617 }
11618 pub fn dtoh_pair(
11622 &self,
11623 a: &CudaSlice<f32>,
11624 b: &CudaSlice<f32>,
11625 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
11626 let av = self.gpu.stream().clone_dtoh(a)?;
11627 let bv = self.gpu.stream().clone_dtoh(b)?;
11628 self.gpu.stream().synchronize()?;
11629 Ok((av, bv))
11630 }
11631 pub fn dtoh_pair_views(
11634 &self,
11635 a: &cudarc::driver::CudaView<f32>,
11636 b: &cudarc::driver::CudaView<f32>,
11637 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
11638 let av = self.gpu.stream().clone_dtoh(a)?;
11639 let bv = self.gpu.stream().clone_dtoh(b)?;
11640 self.gpu.stream().synchronize()?;
11641 Ok((av, bv))
11642 }
11643 pub fn dtoh_i32(&self, d: &CudaSlice<i32>) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
11645 let v = self.gpu.stream().clone_dtoh(d)?;
11646 self.gpu.stream().synchronize()?;
11647 Ok(v)
11648 }
11649 pub fn dtoh_u8(&self, d: &CudaSlice<u8>) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
11651 let v = self.gpu.stream().clone_dtoh(d)?;
11652 self.gpu.stream().synchronize()?;
11653 Ok(v)
11654 }
11655 pub fn dtoh_u8_view(
11656 &self,
11657 d: &cudarc::driver::CudaView<u8>,
11658 ) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
11659 let v = self.gpu.stream().clone_dtoh(d)?;
11660 self.gpu.stream().synchronize()?;
11661 Ok(v)
11662 }
11663 pub fn dtoh_u8_into_pinned(
11671 &self,
11672 d: &CudaSlice<u8>,
11673 dst: &mut PinnedHostBuf,
11674 n: usize,
11675 ) -> Result<(), Box<dyn std::error::Error>> {
11676 if n > d.len() || n > dst.len() {
11677 return Err(format!(
11678 "dtoh_u8_into_pinned range {n} exceeds src {} or pinned dst {}",
11679 d.len(),
11680 dst.len(),
11681 )
11682 .into());
11683 }
11684 if n == 0 {
11685 return Ok(());
11686 }
11687 let host = &mut dst.as_mut_slice()[..n];
11688 self.gpu.stream().memcpy_dtoh(&d.slice(0..n), host)?;
11689 self.gpu.stream().synchronize()?;
11690 Ok(())
11691 }
11692 pub fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11693 SCRATCH_ALLOC_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11694 let s = self.gpu.stream().alloc_zeros::<f32>(n)?;
11695 self.keep_if_capturing(&s);
11696 Ok(s)
11697 }
11698
11699 pub(crate) fn hyper_ws_take(&self) -> Option<crate::hyper::HyperDecodeWs> {
11703 self.hyper_decode_ws.lock().unwrap().take()
11704 }
11705
11706 pub(crate) fn hyper_ws_put(&self, ws: crate::hyper::HyperDecodeWs) {
11707 *self.hyper_decode_ws.lock().unwrap() = Some(ws);
11708 }
11709
11710 pub(crate) fn vws_uninit(
11719 &self,
11720 n: usize,
11721 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11722 if verify_ws_on() {
11723 let mut ws = self.verify_ws.lock().unwrap();
11724 let ws = &mut *ws;
11725 if let Some(s) = VerifyWs::take(&mut ws.f32_pool, &mut ws.held_bytes, n) {
11726 if VERIFY_WS_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
11727 eprintln!(
11728 "[glm5-verify-ws] engaged: verify-walk buffers recycling through \
11729 the size-keyed pool (MEMRA_GLM5_VERIFY_WS=1)"
11730 );
11731 }
11732 return Ok(s);
11733 }
11734 }
11735 self.alloc_uninit::<f32>(n)
11736 }
11737
11738 pub(crate) fn vws_uninit_i8(
11740 &self,
11741 n: usize,
11742 ) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
11743 if verify_ws_on() {
11744 let mut ws = self.verify_ws.lock().unwrap();
11745 let ws = &mut *ws;
11746 if let Some(s) = VerifyWs::take(&mut ws.i8_pool, &mut ws.held_bytes, n) {
11747 VERIFY_WS_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11748 return Ok(s);
11749 }
11750 }
11751 self.alloc_uninit::<i8>(n)
11752 }
11753
11754 pub(crate) fn vws_uninit_u64(
11756 &self,
11757 n: usize,
11758 ) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
11759 if verify_ws_on() {
11760 let mut ws = self.verify_ws.lock().unwrap();
11761 let ws = &mut *ws;
11762 if let Some(s) = VerifyWs::take(&mut ws.u64_pool, &mut ws.held_bytes, n) {
11763 VERIFY_WS_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11764 return Ok(s);
11765 }
11766 }
11767 self.alloc_uninit::<u64>(n)
11768 }
11769
11770 pub(crate) fn vws_recycle(&self, s: CudaSlice<f32>) {
11773 if verify_ws_on() {
11774 let mut ws = self.verify_ws.lock().unwrap();
11775 let ws = &mut *ws;
11776 VerifyWs::put(&mut ws.f32_pool, &mut ws.held_bytes, s);
11777 }
11778 }
11779
11780 pub(crate) fn vws_recycle_i8(&self, s: CudaSlice<i8>) {
11782 if verify_ws_on() {
11783 let mut ws = self.verify_ws.lock().unwrap();
11784 let ws = &mut *ws;
11785 VerifyWs::put(&mut ws.i8_pool, &mut ws.held_bytes, s);
11786 }
11787 }
11788
11789 pub(crate) fn vws_recycle_u64(&self, s: CudaSlice<u64>) {
11791 if verify_ws_on() {
11792 let mut ws = self.verify_ws.lock().unwrap();
11793 let ws = &mut *ws;
11794 VerifyWs::put(&mut ws.u64_pool, &mut ws.held_bytes, s);
11795 }
11796 }
11797
11798 pub fn prob_of_token_device(
11807 &self,
11808 logits: &CudaSlice<f32>,
11809 tok: &CudaSlice<u32>,
11810 n_vocab: usize,
11811 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11812 let nb = ARGMAX_NB;
11813 let mut part = self.alloc_uninit::<f32>(nb)?;
11814 let mut p = self.alloc_uninit::<f32>(1)?;
11815 let f1 = self.func("prob_of_token_partial_f32");
11816 let cfg1 = LaunchConfig {
11817 grid_dim: (nb as u32, 1, 1),
11818 block_dim: (256, 1, 1),
11819 shared_mem_bytes: 0,
11820 };
11821 let nv = n_vocab as i32;
11822 let __s_b1 = self.gpu.stream();
11823 let mut b1 = __s_b1.launch_builder(&f1);
11824 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
11825 unsafe {
11826 b1.launch(cfg1)?;
11827 }
11828 let f2 = self.func("prob_of_token_final_f32");
11829 let cfg2 = LaunchConfig {
11830 grid_dim: (1, 1, 1),
11831 block_dim: (256, 1, 1),
11832 shared_mem_bytes: 0,
11833 };
11834 let nbi = nb as i32;
11835 let __s_b2 = self.gpu.stream();
11836 let mut b2 = __s_b2.launch_builder(&f2);
11837 b2.arg(&part).arg(&mut p).arg(&nbi);
11838 unsafe {
11839 b2.launch(cfg2)?;
11840 }
11841 Ok(p)
11842 }
11843
11844 pub fn prob_of_token_device_col(
11851 &self,
11852 logits: &CudaSlice<f32>,
11853 tok_all: &CudaSlice<u32>,
11854 tok_idx: usize,
11855 p_out: &mut CudaSlice<f32>,
11856 p_idx: usize,
11857 n_vocab: usize,
11858 ) -> Result<(), Box<dyn std::error::Error>> {
11859 let tok_v = tok_all.slice(tok_idx..tok_idx + 1);
11860 let mut p_v = p_out.slice_mut(p_idx..p_idx + 1);
11861 let nb = ARGMAX_NB;
11862 let mut part = self.alloc_uninit::<f32>(nb)?;
11863 let f1 = self.func("prob_of_token_partial_f32");
11864 let cfg1 = LaunchConfig {
11865 grid_dim: (nb as u32, 1, 1),
11866 block_dim: (256, 1, 1),
11867 shared_mem_bytes: 0,
11868 };
11869 let nv = n_vocab as i32;
11870 let __s_b1 = self.gpu.stream();
11871 let mut b1 = __s_b1.launch_builder(&f1);
11872 b1.arg(logits).arg(&tok_v).arg(&mut part).arg(&nv);
11873 unsafe {
11874 b1.launch(cfg1)?;
11875 }
11876 let f2 = self.func("prob_of_token_final_f32");
11877 let cfg2 = LaunchConfig {
11878 grid_dim: (1, 1, 1),
11879 block_dim: (256, 1, 1),
11880 shared_mem_bytes: 0,
11881 };
11882 let nbi = nb as i32;
11883 let __s_b2 = self.gpu.stream();
11884 let mut b2 = __s_b2.launch_builder(&f2);
11885 b2.arg(&part).arg(&mut p_v).arg(&nbi);
11886 unsafe {
11887 b2.launch(cfg2)?;
11888 }
11889 Ok(())
11890 }
11891
11892 pub fn prob_of_token_device_into(
11893 &self,
11894 logits: &CudaSlice<f32>,
11895 tok: &CudaSlice<u32>,
11896 p_out: &mut CudaSlice<f32>,
11897 n_vocab: usize,
11898 ) -> Result<(), Box<dyn std::error::Error>> {
11899 let nb = ARGMAX_NB;
11900 let mut part = self.alloc_uninit::<f32>(nb)?;
11901 let f1 = self.func("prob_of_token_partial_f32");
11902 let cfg1 = LaunchConfig {
11903 grid_dim: (nb as u32, 1, 1),
11904 block_dim: (256, 1, 1),
11905 shared_mem_bytes: 0,
11906 };
11907 let nv = n_vocab as i32;
11908 let __s_b1 = self.gpu.stream();
11909 let mut b1 = __s_b1.launch_builder(&f1);
11910 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
11911 unsafe {
11912 b1.launch(cfg1)?;
11913 }
11914 let f2 = self.func("prob_of_token_final_f32");
11915 let cfg2 = LaunchConfig {
11916 grid_dim: (1, 1, 1),
11917 block_dim: (256, 1, 1),
11918 shared_mem_bytes: 0,
11919 };
11920 let nbi = nb as i32;
11921 let __s_b2 = self.gpu.stream();
11922 let mut b2 = __s_b2.launch_builder(&f2);
11923 b2.arg(&part).arg(p_out).arg(&nbi);
11924 unsafe {
11925 b2.launch(cfg2)?;
11926 }
11927 Ok(())
11928 }
11929
11930 pub fn u32_hist_append(
11933 &self,
11934 tok: &CudaSlice<u32>,
11935 hist: &mut CudaSlice<u32>,
11936 idx: &mut CudaSlice<i32>,
11937 ) -> Result<(), Box<dyn std::error::Error>> {
11938 let f = self.func("u32_hist_append");
11939 let cfg = LaunchConfig {
11940 grid_dim: (1, 1, 1),
11941 block_dim: (32, 1, 1),
11942 shared_mem_bytes: 0,
11943 };
11944 let __s_b = self.gpu.stream();
11945 let mut b = __s_b.launch_builder(&f);
11946 b.arg(tok).arg(&mut *hist).arg(&mut *idx);
11947 unsafe {
11948 b.launch(cfg)?;
11949 }
11950 Ok(())
11951 }
11952
11953 pub fn argmax_token_device(
11954 &self,
11955 logits: &CudaSlice<f32>,
11956 n_vocab: usize,
11957 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
11958 let mut tok = unsafe { self.gpu.stream().alloc::<u32>(1)? };
11959 self.argmax_token_device_into(logits, &mut tok, n_vocab)?;
11960 Ok(tok)
11961 }
11962 pub fn argmax_token_device_into(
11969 &self,
11970 logits: &CudaSlice<f32>,
11971 tok: &mut CudaSlice<u32>,
11972 n_vocab: usize,
11973 ) -> Result<(), Box<dyn std::error::Error>> {
11974 let nb = ARGMAX_NB;
11975 let f1 = self.func("argmax_partial_f32");
11976 let f2 = self.func("argmax_final_f32");
11977 let mut guard = self.argmax_partials.lock().unwrap();
11978 if guard.is_none() {
11979 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
11982 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
11983 *guard = Some((pv, pi));
11984 }
11985 let (part_v, part_i) = guard.as_mut().unwrap();
11986 let nv = n_vocab as i32;
11987 let nbi = nb as i32;
11988 let cfg1 = LaunchConfig {
11990 grid_dim: (nb as u32, 1, 1),
11991 block_dim: (256, 1, 1),
11992 shared_mem_bytes: 0,
11993 };
11994 let __s_b1 = self.gpu.stream();
11995 let mut b1 = __s_b1.launch_builder(&f1);
11996 b1.arg(logits).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
11997 unsafe {
11998 b1.launch(cfg1)?;
11999 }
12000 let cfg2 = LaunchConfig {
12002 grid_dim: (1, 1, 1),
12003 block_dim: (256, 1, 1),
12004 shared_mem_bytes: 0,
12005 };
12006 let __s_b2 = self.gpu.stream();
12007 let mut b2 = __s_b2.launch_builder(&f2);
12008 b2.arg(&*part_v).arg(&*part_i).arg(tok).arg(&nbi);
12009 unsafe {
12010 b2.launch(cfg2)?;
12011 }
12012 Ok(())
12013 }
12014 pub fn argmax_token_device_col(
12020 &self,
12021 logits: &CudaSlice<f32>,
12022 col: usize,
12023 n_vocab: usize,
12024 toks: &mut CudaSlice<u32>,
12025 out_idx: usize,
12026 ) -> Result<(), Box<dyn std::error::Error>> {
12027 let nb = ARGMAX_NB;
12028 let f1 = self.func("argmax_partial_f32");
12029 let f2 = self.func("argmax_final_f32");
12030 let mut guard = self.argmax_partials.lock().unwrap();
12031 if guard.is_none() {
12032 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
12033 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
12034 *guard = Some((pv, pi));
12035 }
12036 let (part_v, part_i) = guard.as_mut().unwrap();
12037 let col_view = logits.slice(col * n_vocab..(col + 1) * n_vocab);
12038 let nv = n_vocab as i32;
12039 let nbi = nb as i32;
12040 let cfg1 = LaunchConfig {
12041 grid_dim: (nb as u32, 1, 1),
12042 block_dim: (256, 1, 1),
12043 shared_mem_bytes: 0,
12044 };
12045 let __s_b1 = self.gpu.stream();
12046 let mut b1 = __s_b1.launch_builder(&f1);
12047 b1.arg(&col_view)
12048 .arg(&mut *part_v)
12049 .arg(&mut *part_i)
12050 .arg(&nv);
12051 unsafe {
12052 b1.launch(cfg1)?;
12053 }
12054 let mut tok_view = toks.slice_mut(out_idx..out_idx + 1);
12055 let cfg2 = LaunchConfig {
12056 grid_dim: (1, 1, 1),
12057 block_dim: (256, 1, 1),
12058 shared_mem_bytes: 0,
12059 };
12060 let __s_b2 = self.gpu.stream();
12061 let mut b2 = __s_b2.launch_builder(&f2);
12062 b2.arg(&*part_v).arg(&*part_i).arg(&mut tok_view).arg(&nbi);
12063 unsafe {
12064 b2.launch(cfg2)?;
12065 }
12066 Ok(())
12067 }
12068 pub fn htod_u32_v(&self, v: &[u32]) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
12070 Ok(self.gpu.stream().clone_htod(v)?)
12071 }
12072 pub fn dtoh_u64(&self, d: &CudaSlice<u64>) -> Result<Vec<u64>, Box<dyn std::error::Error>> {
12073 let v = self.gpu.stream().clone_dtoh(d)?;
12074 self.gpu.stream().synchronize()?;
12075 Ok(v)
12076 }
12077
12078 pub fn dtoh_u32(&self, d: &CudaSlice<u32>) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
12079 let v = self.gpu.stream().clone_dtoh(d)?;
12080 self.gpu.stream().synchronize()?;
12081 Ok(v)
12082 }
12083 pub fn htod_u32_into(
12087 &self,
12088 dst: &mut CudaSlice<u32>,
12089 src: &[u32],
12090 ) -> Result<(), Box<dyn std::error::Error>> {
12091 let mut view = dst.slice_mut(0..src.len());
12092 self.gpu.stream().memcpy_htod(src, &mut view)?;
12093 Ok(())
12094 }
12095
12096 pub fn htod_i32_into(
12099 &self,
12100 dst: &mut CudaSlice<i32>,
12101 src: &[i32],
12102 ) -> Result<(), Box<dyn std::error::Error>> {
12103 let mut view = dst.slice_mut(0..src.len());
12104 self.gpu.stream().memcpy_htod(src, &mut view)?;
12105 Ok(())
12106 }
12107
12108 pub fn alloc_u32_zeroed(&self, n: usize) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
12109 let s = self.gpu.stream().alloc_zeros::<u32>(n)?;
12110 self.keep_if_capturing(&s);
12111 Ok(s)
12112 }
12113 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device_into(
12117 &self,
12118 embd: &CudaSlice<u8>,
12119 token_d: &CudaSlice<u32>,
12120 x_out: &mut CudaSlice<f32>,
12121 n_embd: usize,
12122 qtype: i32,
12123 row_bytes: usize,
12124 ) -> Result<(), Box<dyn std::error::Error>> {
12125 let f = self.func("embed_gather_u32");
12126 let cfg = LaunchConfig {
12127 grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
12128 block_dim: (256, 1, 1),
12129 shared_mem_bytes: 0,
12130 };
12131 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
12132 let __s_b = self.gpu.stream();
12133 let mut b = __s_b.launch_builder(&f);
12134 b.arg(embd)
12135 .arg(token_d)
12136 .arg(x_out)
12137 .arg(&ne)
12138 .arg(&qt)
12139 .arg(&rb);
12140 unsafe {
12141 b.launch(cfg)?;
12142 }
12143 Ok(())
12144 }
12145 pub fn dtoh_i32_one(&self, d: &CudaSlice<i32>) -> Result<i32, Box<dyn std::error::Error>> {
12147 let v = self.gpu.stream().clone_dtoh(d)?;
12148 self.gpu.stream().synchronize()?;
12149 Ok(v[0])
12150 }
12151 pub fn i32_set_k(
12158 &self,
12159 dst: &mut CudaSlice<i32>,
12160 v: i32,
12161 ) -> Result<(), Box<dyn std::error::Error>> {
12162 let f = self.func("i32_set_k");
12163 let cfg = LaunchConfig {
12164 grid_dim: (1, 1, 1),
12165 block_dim: (1, 1, 1),
12166 shared_mem_bytes: 0,
12167 };
12168 let idx = 0i32;
12169 let __s_b = self.gpu.stream();
12170 let mut b = __s_b.launch_builder(&f);
12171 b.arg(dst).arg(&v).arg(&idx);
12172 unsafe {
12173 b.launch(cfg)?;
12174 }
12175 Ok(())
12176 }
12177
12178 pub fn set_i32_one(
12179 &self,
12180 d: &mut CudaSlice<i32>,
12181 v: i32,
12182 ) -> Result<(), Box<dyn std::error::Error>> {
12183 self.gpu.stream().memcpy_htod(&[v], d)?;
12184 Ok(())
12185 }
12186 pub fn set_u32_one(
12189 &self,
12190 d: &mut CudaSlice<u32>,
12191 v: u32,
12192 ) -> Result<(), Box<dyn std::error::Error>> {
12193 self.gpu.stream().memcpy_htod(&[v], d)?;
12194 Ok(())
12195 }
12196 pub fn dtoh_u32_one(&self, d: &CudaSlice<u32>) -> Result<u32, Box<dyn std::error::Error>> {
12198 let v = self.gpu.stream().clone_dtoh(d)?;
12199 self.gpu.stream().synchronize()?;
12200 Ok(v[0])
12201 }
12202 pub fn upload_u8(&self, bytes: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
12204 Ok(self.gpu.stream().clone_htod(bytes)?)
12205 }
12206 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device(
12211 &self,
12212 embd: &CudaSlice<u8>,
12213 token_d: &CudaSlice<u32>,
12214 n_embd: usize,
12215 qtype: i32,
12216 row_bytes: usize,
12217 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12218 let f = self.func("embed_gather_u32");
12219 let mut x = self.alloc_uninit::<f32>(n_embd)?;
12220 let cfg = LaunchConfig {
12221 grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
12222 block_dim: (256, 1, 1),
12223 shared_mem_bytes: 0,
12224 };
12225 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
12226 let __s_b = self.gpu.stream();
12227 let mut b = __s_b.launch_builder(&f);
12228 b.arg(embd)
12229 .arg(token_d)
12230 .arg(&mut x)
12231 .arg(&ne)
12232 .arg(&qt)
12233 .arg(&rb);
12234 unsafe {
12235 b.launch(cfg)?;
12236 }
12237 Ok(x)
12238 }
12239
12240 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device_t(
12245 &self,
12246 embd: &CudaSlice<u8>,
12247 tokens: &[u32],
12248 n_embd: usize,
12249 qtype: i32,
12250 row_bytes: usize,
12251 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12252 let t = tokens.len();
12253 let tok_d = self.gpu.stream().clone_htod(tokens)?;
12254 let f = self.func("embed_gather_u32_t");
12255 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
12256 let cfg = LaunchConfig {
12257 grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
12258 block_dim: (256, 1, 1),
12259 shared_mem_bytes: 0,
12260 };
12261 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
12262 let __s_b = self.gpu.stream();
12263 let mut b = __s_b.launch_builder(&f);
12264 b.arg(embd)
12265 .arg(&tok_d)
12266 .arg(&mut x)
12267 .arg(&ne)
12268 .arg(&qt)
12269 .arg(&rb)
12270 .arg(&ti);
12271 unsafe {
12272 b.launch(cfg)?;
12273 }
12274 Ok(x)
12275 }
12276
12277 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device_tv(
12283 &self,
12284 embd: &CudaSlice<u8>,
12285 tok_v: &cudarc::driver::CudaView<u32>,
12286 t: usize,
12287 n_embd: usize,
12288 qtype: i32,
12289 row_bytes: usize,
12290 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12291 let f = self.func("embed_gather_u32_t");
12292 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
12293 let cfg = LaunchConfig {
12294 grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
12295 block_dim: (256, 1, 1),
12296 shared_mem_bytes: 0,
12297 };
12298 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
12299 let __s_b = self.gpu.stream();
12300 let mut b = __s_b.launch_builder(&f);
12301 b.arg(embd)
12302 .arg(tok_v)
12303 .arg(&mut x)
12304 .arg(&ne)
12305 .arg(&qt)
12306 .arg(&rb)
12307 .arg(&ti);
12308 unsafe {
12309 b.launch(cfg)?;
12310 }
12311 Ok(x)
12312 }
12313
12314 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device_td(
12316 &self,
12317 embd: &CudaSlice<u8>,
12318 tok_d: &CudaSlice<u32>,
12319 t: usize,
12320 n_embd: usize,
12321 qtype: i32,
12322 row_bytes: usize,
12323 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12324 let f = self.func("embed_gather_u32_t");
12325 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
12326 let cfg = LaunchConfig {
12327 grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
12328 block_dim: (256, 1, 1),
12329 shared_mem_bytes: 0,
12330 };
12331 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
12332 let __s_b = self.gpu.stream();
12333 let mut b = __s_b.launch_builder(&f);
12334 b.arg(embd)
12335 .arg(tok_d)
12336 .arg(&mut x)
12337 .arg(&ne)
12338 .arg(&qt)
12339 .arg(&rb)
12340 .arg(&ti);
12341 unsafe {
12342 b.launch(cfg)?;
12343 }
12344 Ok(x)
12345 }
12346
12347 #[inline]
12353 fn keep_if_capturing<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, s: &CudaSlice<T>) {
12355 if self
12356 .capture_keep_on
12357 .load(std::sync::atomic::Ordering::Relaxed)
12358 {
12359 self.capture_keep.lock().unwrap().push(Box::new(s.clone()));
12360 }
12361 }
12362
12363 fn alloc_uninit<T: cudarc::driver::DeviceRepr + Send + 'static>(
12364 &self,
12365 n: usize,
12366 ) -> Result<CudaSlice<T>, Box<dyn std::error::Error>> {
12367 SCRATCH_ALLOC_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12368 let mut s = unsafe { self.gpu.stream().alloc::<T>(n)? };
12369 {
12373 static Z: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12374 if *Z.get_or_init(|| std::env::var("MEMRA_DEBUG_ZERO_ALLOCS").as_deref() == Ok("1")) {
12375 use cudarc::driver::DevicePtrMut;
12377 let n_bytes = s.len() * std::mem::size_of::<T>();
12378 let stream = self.gpu.stream();
12379 let (p_, _g) = s.device_ptr_mut(&stream);
12380 unsafe {
12381 cudarc::driver::sys::cuMemsetD8Async(p_, 0, n_bytes, stream.cu_stream())
12382 .result()?;
12383 }
12384 }
12385 }
12386 self.keep_if_capturing(&s);
12387 Ok(s)
12388 }
12389
12390 pub fn uninit_q8_pair(
12395 &self,
12396 n: usize,
12397 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12398 Ok((
12399 self.alloc_uninit::<i8>(n)?,
12400 self.alloc_uninit::<f32>(n / 32)?,
12401 ))
12402 }
12403
12404 pub fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12405 self.alloc_uninit::<f32>(n)
12406 }
12407
12408 pub fn alloc_i8_uninit(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
12410 self.alloc_uninit::<i8>(n)
12411 }
12412
12413 pub fn uninit_i32(&self, n: usize) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
12415 self.alloc_uninit::<i32>(n)
12416 }
12417
12418 #[allow(clippy::too_many_arguments)]
12422 pub fn rms_norm3(
12423 &self,
12424 x: &CudaSlice<f32>,
12425 w0: &CudaSlice<f32>,
12426 w1: &CudaSlice<f32>,
12427 w2: &CudaSlice<f32>,
12428 d0: &mut CudaSlice<f32>,
12429 d1: &mut CudaSlice<f32>,
12430 d2: &mut CudaSlice<f32>,
12431 ncols: usize,
12432 nrows: usize,
12433 eps: f32,
12434 ) -> Result<(), Box<dyn std::error::Error>> {
12435 let f = self.func("rms_norm3_f32");
12436 let cfg = LaunchConfig {
12437 grid_dim: (nrows as u32, 1, 1),
12438 block_dim: (rms_block(), 1, 1),
12439 shared_mem_bytes: 0,
12440 };
12441 let (nc, e) = (ncols as i32, eps);
12442 let __s_b = self.gpu.stream();
12443 let mut b = __s_b.launch_builder(&f);
12444 b.arg(x)
12445 .arg(w0)
12446 .arg(w1)
12447 .arg(w2)
12448 .arg(d0)
12449 .arg(d1)
12450 .arg(d2)
12451 .arg(&nc)
12452 .arg(&e);
12453 unsafe {
12454 b.launch(cfg)?;
12455 }
12456 Ok(())
12457 }
12458
12459 #[allow(clippy::too_many_arguments)]
12461 pub fn qkvnorm_w_on_prefill(rows: usize, ncols: usize) -> bool {
12464 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12465 *WARP_ON.get_or_init(|| {
12466 std::env::var("MEMRA_QKVNORM_W")
12467 .map(|v| v != "0")
12468 .unwrap_or(true)
12469 }) && ncols.is_multiple_of(4)
12470 && rows >= 64
12471 }
12472
12473 #[allow(clippy::too_many_arguments)]
12476 pub fn rms_norm_qkv_w4b(
12477 &self,
12478 q: &CudaSlice<f32>,
12479 k: &CudaSlice<f32>,
12480 v: &CudaSlice<f32>,
12481 wq: &CudaSlice<f32>,
12482 wk: &CudaSlice<f32>,
12483 wv: &CudaSlice<f32>,
12484 dq: &mut CudaSlice<f32>,
12485 dk: &mut CudaSlice<f32>,
12486 dv: &mut CudaSlice<f32>,
12487 dvb: &mut CudaSlice<u8>,
12488 ncols: usize,
12489 rq: usize,
12490 rk: usize,
12491 eps: f32,
12492 vf16: bool,
12493 ) -> Result<(), Box<dyn std::error::Error>> {
12494 assert!(ncols.is_multiple_of(4) && rq + 2 * rk >= 64);
12495 let f = self.func("rms_norm_qkv_w4b_f32");
12496 let rows = (rq + 2 * rk) as u32;
12497 let cfg = LaunchConfig {
12498 grid_dim: (rows.div_ceil(8), 1, 1),
12499 block_dim: (256, 1, 1),
12500 shared_mem_bytes: 0,
12501 };
12502 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
12503 let vf = vf16 as i32;
12504 let __s_b = self.gpu.stream();
12505 let mut b = __s_b.launch_builder(&f);
12506 b.arg(q)
12507 .arg(k)
12508 .arg(v)
12509 .arg(wq)
12510 .arg(wk)
12511 .arg(wv)
12512 .arg(dq)
12513 .arg(dk)
12514 .arg(dv)
12515 .arg(&mut *dvb)
12516 .arg(&nc)
12517 .arg(&rqi)
12518 .arg(&rki)
12519 .arg(&rvi)
12520 .arg(&e)
12521 .arg(&vf);
12522 unsafe {
12523 b.launch(cfg)?;
12524 }
12525 Ok(())
12526 }
12527
12528 #[allow(clippy::too_many_arguments)] pub fn rms_norm_qkv(
12530 &self,
12531 q: &CudaSlice<f32>,
12532 k: &CudaSlice<f32>,
12533 v: &CudaSlice<f32>,
12534 wq: &CudaSlice<f32>,
12535 wk: &CudaSlice<f32>,
12536 wv: &CudaSlice<f32>,
12537 dq: &mut CudaSlice<f32>,
12538 dk: &mut CudaSlice<f32>,
12539 dv: &mut CudaSlice<f32>,
12540 ncols: usize,
12541 rq: usize,
12542 rk: usize,
12543 eps: f32,
12544 ) -> Result<(), Box<dyn std::error::Error>> {
12545 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12549 let warp_on = *WARP_ON.get_or_init(|| {
12550 std::env::var("MEMRA_QKVNORM_W")
12551 .map(|v| v != "0")
12552 .unwrap_or(true)
12553 });
12554 if warp_on && ncols.is_multiple_of(4) && rq + 2 * rk >= 64 {
12557 let f = self.func("rms_norm_qkv_w4_f32");
12558 let rows = (rq + 2 * rk) as u32;
12559 let cfg = LaunchConfig {
12560 grid_dim: (rows.div_ceil(8), 1, 1),
12561 block_dim: (256, 1, 1),
12562 shared_mem_bytes: 0,
12563 };
12564 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
12565 let __s_b = self.gpu.stream();
12566 let mut b = __s_b.launch_builder(&f);
12567 b.arg(q)
12568 .arg(k)
12569 .arg(v)
12570 .arg(wq)
12571 .arg(wk)
12572 .arg(wv)
12573 .arg(dq)
12574 .arg(dk)
12575 .arg(dv)
12576 .arg(&nc)
12577 .arg(&rqi)
12578 .arg(&rki)
12579 .arg(&rvi)
12580 .arg(&e);
12581 unsafe {
12582 b.launch(cfg)?;
12583 }
12584 return Ok(());
12585 }
12586 let f = self.func("rms_norm_qkv_f32");
12587 let grid = (rq + 2 * rk) as u32;
12588 let cfg = LaunchConfig {
12589 grid_dim: (grid, 1, 1),
12590 block_dim: (rms_block(), 1, 1),
12591 shared_mem_bytes: 0,
12592 };
12593 let (nc, rqi, rki, e) = (ncols as i32, rq as i32, rk as i32, eps);
12594 let __s_b = self.gpu.stream();
12595 let mut b = __s_b.launch_builder(&f);
12596 b.arg(q)
12597 .arg(k)
12598 .arg(v)
12599 .arg(wq)
12600 .arg(wk)
12601 .arg(wv)
12602 .arg(dq)
12603 .arg(dk)
12604 .arg(dv)
12605 .arg(&nc)
12606 .arg(&rqi)
12607 .arg(&rki)
12608 .arg(&e);
12609 unsafe {
12610 b.launch(cfg)?;
12611 }
12612 Ok(())
12613 }
12614
12615 #[allow(clippy::too_many_arguments)]
12617 pub fn rms_norm2x(
12618 &self,
12619 a: &CudaSlice<f32>,
12620 bb: &CudaSlice<f32>,
12621 wa: &CudaSlice<f32>,
12622 wb: &CudaSlice<f32>,
12623 da: &mut CudaSlice<f32>,
12624 db: &mut CudaSlice<f32>,
12625 ncols: usize,
12626 nrows: usize,
12627 eps: f32,
12628 ) -> Result<(), Box<dyn std::error::Error>> {
12629 let f = self.func("rms_norm2x_f32");
12630 let cfg = LaunchConfig {
12631 grid_dim: (2 * nrows as u32, 1, 1),
12632 block_dim: (rms_block(), 1, 1),
12633 shared_mem_bytes: 0,
12634 };
12635 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
12636 let __s_b = self.gpu.stream();
12637 let mut b = __s_b.launch_builder(&f);
12638 b.arg(a)
12639 .arg(bb)
12640 .arg(wa)
12641 .arg(wb)
12642 .arg(da)
12643 .arg(db)
12644 .arg(&nc)
12645 .arg(&nr)
12646 .arg(&e);
12647 unsafe {
12648 b.launch(cfg)?;
12649 }
12650 Ok(())
12651 }
12652
12653 pub fn softcap(
12655 &self,
12656 y: &mut CudaSlice<f32>,
12657 cap: f32,
12658 n: usize,
12659 ) -> Result<(), Box<dyn std::error::Error>> {
12660 let f = self.func("softcap_f32");
12661 let cfg = LaunchConfig::for_num_elems(n as u32);
12662 let ni = n as i32;
12663 let __s_b = self.gpu.stream();
12664 let mut b = __s_b.launch_builder(&f);
12665 b.arg(y).arg(&cap).arg(&ni);
12666 unsafe {
12667 b.launch(cfg)?;
12668 }
12669 Ok(())
12670 }
12671
12672 pub fn mask_ids_rows(
12675 &self,
12676 y: &mut CudaSlice<f32>,
12677 ids: &CudaSlice<i32>,
12678 n_ids: usize,
12679 n_vocab: usize,
12680 t: usize,
12681 ) -> Result<(), Box<dyn std::error::Error>> {
12682 let f = self.func("mask_ids_rows_f32");
12683 let cfg = LaunchConfig::for_num_elems((n_ids * t) as u32);
12684 let (ni, nv, ti) = (n_ids as i32, n_vocab as i32, t as i32);
12685 let __s_b = self.gpu.stream();
12686 let mut b = __s_b.launch_builder(&f);
12687 b.arg(y).arg(ids).arg(&ni).arg(&nv).arg(&ti);
12688 unsafe {
12689 b.launch(cfg)?;
12690 }
12691 Ok(())
12692 }
12693
12694 #[allow(clippy::too_many_arguments)]
12696 pub fn add_scale_rms_norm(
12697 &self,
12698 a: &CudaSlice<f32>,
12699 b_in: &CudaSlice<f32>,
12700 c: f32,
12701 w: &CudaSlice<f32>,
12702 res: &mut CudaSlice<f32>,
12703 dst: &mut CudaSlice<f32>,
12704 ncols: usize,
12705 nrows: usize,
12706 eps: f32,
12707 ) -> Result<(), Box<dyn std::error::Error>> {
12708 let f = self.func("add_scale_rms_norm_f32");
12709 let cfg = LaunchConfig {
12710 grid_dim: (nrows as u32, 1, 1),
12711 block_dim: (rms_block(), 1, 1),
12712 shared_mem_bytes: 0,
12713 };
12714 let (nc, e2) = (ncols as i32, eps);
12715 let __s_b = self.gpu.stream();
12716 let mut b = __s_b.launch_builder(&f);
12717 b.arg(a)
12718 .arg(b_in)
12719 .arg(&c)
12720 .arg(w)
12721 .arg(res)
12722 .arg(dst)
12723 .arg(&nc)
12724 .arg(&e2);
12725 unsafe {
12726 b.launch(cfg)?;
12727 }
12728 Ok(())
12729 }
12730
12731 #[allow(clippy::too_many_arguments)]
12734 pub fn add_scale_rms_norm_q8_1(
12735 &self,
12736 a: &CudaSlice<f32>,
12737 b_in: &CudaSlice<f32>,
12738 c: f32,
12739 w: &CudaSlice<f32>,
12740 res: &mut CudaSlice<f32>,
12741 ncols: usize,
12742 nrows: usize,
12743 eps: f32,
12744 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12745 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12746 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12747 let (nc, e2) = (ncols as i32, eps);
12748 if Self::pdl_on() && Self::pdl_wb_on() {
12749 {
12750 use cudarc::driver::{DevicePtr, DevicePtrMut};
12751 let s = &self.gpu.stream();
12752 let (pa, _g0) = a.device_ptr(s);
12753 let (pb, _g1) = b_in.device_ptr(s);
12754 let (pw, _g2) = w.device_ptr(s);
12755 let (pr, _g3) = res.device_ptr_mut(s);
12756 let (pq, _g4) = out_q.device_ptr_mut(s);
12757 let (pd, _g5) = out_d.device_ptr_mut(s);
12758 let mut ps = [
12759 &pa as *const _ as *mut std::ffi::c_void,
12760 &pb as *const _ as *mut _,
12761 &c as *const _ as *mut _,
12762 &pw as *const _ as *mut _,
12763 &pr as *const _ as *mut _,
12764 &pq as *const _ as *mut _,
12765 &pd as *const _ as *mut _,
12766 &nc as *const _ as *mut _,
12767 &e2 as *const _ as *mut _,
12768 ];
12769 unsafe {
12770 self.launch_pdl(
12771 "add_scale_rms_norm_q8_1",
12772 (nrows as u32, 1, 1),
12773 (rms_block(), 1, 1),
12774 &mut ps,
12775 )?;
12776 }
12777 }
12778 return Ok((out_q, out_d));
12779 }
12780 let f = self.func("add_scale_rms_norm_q8_1");
12781 let cfg = LaunchConfig {
12782 grid_dim: (nrows as u32, 1, 1),
12783 block_dim: (rms_block(), 1, 1),
12784 shared_mem_bytes: 0,
12785 };
12786 let __s_b = self.gpu.stream();
12787 let mut b = __s_b.launch_builder(&f);
12788 b.arg(a)
12789 .arg(b_in)
12790 .arg(&c)
12791 .arg(w)
12792 .arg(res)
12793 .arg(&mut out_q)
12794 .arg(&mut out_d)
12795 .arg(&nc)
12796 .arg(&e2);
12797 unsafe {
12798 b.launch(cfg)?;
12799 }
12800 Ok((out_q, out_d))
12801 }
12802
12803 #[allow(clippy::too_many_arguments)]
12805 pub fn add_scale_rms_norm_q8_1_into(
12806 &self,
12807 a: &CudaSlice<f32>,
12808 b_in: &CudaSlice<f32>,
12809 c: f32,
12810 w: &CudaSlice<f32>,
12811 res: &mut CudaSlice<f32>,
12812 ncols: usize,
12813 nrows: usize,
12814 eps: f32,
12815 out_q: &mut CudaSlice<i8>,
12816 out_d: &mut CudaSlice<f32>,
12817 ) -> Result<(), Box<dyn std::error::Error>> {
12818 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
12819 let (nc, e2) = (ncols as i32, eps);
12820 if Self::pdl_on() && Self::pdl_wb_on() {
12821 use cudarc::driver::{DevicePtr, DevicePtrMut};
12822 let s = &self.gpu.stream();
12823 let (pa, _g0) = a.device_ptr(s);
12824 let (pb, _g1) = b_in.device_ptr(s);
12825 let (pw, _g2) = w.device_ptr(s);
12826 let (pr, _g3) = res.device_ptr_mut(s);
12827 let (pq, _g4) = out_q.device_ptr_mut(s);
12828 let (pd, _g5) = out_d.device_ptr_mut(s);
12829 let mut ps = [
12830 &pa as *const _ as *mut std::ffi::c_void,
12831 &pb as *const _ as *mut _,
12832 &c as *const _ as *mut _,
12833 &pw as *const _ as *mut _,
12834 &pr as *const _ as *mut _,
12835 &pq as *const _ as *mut _,
12836 &pd as *const _ as *mut _,
12837 &nc as *const _ as *mut _,
12838 &e2 as *const _ as *mut _,
12839 ];
12840 unsafe {
12841 self.launch_pdl(
12842 "add_scale_rms_norm_q8_1",
12843 (nrows as u32, 1, 1),
12844 (rms_block(), 1, 1),
12845 &mut ps,
12846 )?;
12847 }
12848 return Ok(());
12849 }
12850 let f = self.func("add_scale_rms_norm_q8_1");
12851 let cfg = LaunchConfig {
12852 grid_dim: (nrows as u32, 1, 1),
12853 block_dim: (rms_block(), 1, 1),
12854 shared_mem_bytes: 0,
12855 };
12856 let __s_b = self.gpu.stream();
12857 let mut b = __s_b.launch_builder(&f);
12858 b.arg(a)
12859 .arg(b_in)
12860 .arg(&c)
12861 .arg(w)
12862 .arg(res)
12863 .arg(&mut *out_q)
12864 .arg(&mut *out_d)
12865 .arg(&nc)
12866 .arg(&e2);
12867 unsafe {
12868 b.launch(cfg)?;
12869 }
12870 Ok(())
12871 }
12872
12873 #[allow(clippy::too_many_arguments)]
12876 pub fn rms_pre_add_scale_rms_norm_q8_1(
12877 &self,
12878 a: &CudaSlice<f32>,
12879 wa: &CudaSlice<f32>,
12880 b_in: &CudaSlice<f32>,
12881 c: f32,
12882 w: &CudaSlice<f32>,
12883 res: &mut CudaSlice<f32>,
12884 ncols: usize,
12885 nrows: usize,
12886 eps: f32,
12887 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12888 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12889 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12890 let (nc, e2) = (ncols as i32, eps);
12891 if Self::pdl_on() {
12892 {
12893 use cudarc::driver::{DevicePtr, DevicePtrMut};
12894 let s = &self.gpu.stream();
12895 let (pa, _g0) = a.device_ptr(s);
12896 let (pwa, _g1) = wa.device_ptr(s);
12897 let (pb, _g2) = b_in.device_ptr(s);
12898 let (pw, _g3) = w.device_ptr(s);
12899 let (pr, _g4) = res.device_ptr_mut(s);
12900 let (pq, _g5) = out_q.device_ptr_mut(s);
12901 let (pd, _g6) = out_d.device_ptr_mut(s);
12902 let mut ps = [
12903 &pa as *const _ as *mut std::ffi::c_void,
12904 &pwa as *const _ as *mut _,
12905 &pb as *const _ as *mut _,
12906 &c as *const _ as *mut _,
12907 &pw as *const _ as *mut _,
12908 &pr as *const _ as *mut _,
12909 &pq as *const _ as *mut _,
12910 &pd as *const _ as *mut _,
12911 &nc as *const _ as *mut _,
12912 &e2 as *const _ as *mut _,
12913 ];
12914 unsafe {
12915 self.launch_pdl(
12916 "rms_pre_add_scale_rms_norm_q8_1",
12917 (nrows as u32, 1, 1),
12918 (rms_block(), 1, 1),
12919 &mut ps,
12920 )?;
12921 }
12922 }
12923 return Ok((out_q, out_d));
12924 }
12925 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
12926 let cfg = LaunchConfig {
12927 grid_dim: (nrows as u32, 1, 1),
12928 block_dim: (rms_block(), 1, 1),
12929 shared_mem_bytes: 0,
12930 };
12931 let __s_b = self.gpu.stream();
12932 let mut b = __s_b.launch_builder(&f);
12933 b.arg(a)
12934 .arg(wa)
12935 .arg(b_in)
12936 .arg(&c)
12937 .arg(w)
12938 .arg(res)
12939 .arg(&mut out_q)
12940 .arg(&mut out_d)
12941 .arg(&nc)
12942 .arg(&e2);
12943 unsafe {
12944 b.launch(cfg)?;
12945 }
12946 Ok((out_q, out_d))
12947 }
12948
12949 pub fn gelu_tanh_mul_q8_1(
12952 &self,
12953 gate: &CudaSlice<f32>,
12954 up: &cudarc::driver::CudaView<f32>,
12955 act: &mut CudaSlice<f32>,
12956 ncols: usize,
12957 nrows: usize,
12958 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12959 debug_assert!(ncols.is_multiple_of(128));
12960 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12961 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12962 let nc = ncols as i32;
12963 if Self::pdl_on() {
12964 {
12965 use cudarc::driver::{DevicePtr, DevicePtrMut};
12966 let s = &self.gpu.stream();
12967 let (pg, _g0) = gate.device_ptr(s);
12968 let (pu, _g1) = up.device_ptr(s);
12969 let (pact, _g2) = act.device_ptr_mut(s);
12970 let (pq, _g3) = out_q.device_ptr_mut(s);
12971 let (pd, _g4) = out_d.device_ptr_mut(s);
12972 let mut ps = [
12973 &pg as *const _ as *mut std::ffi::c_void,
12974 &pu as *const _ as *mut _,
12975 &pact as *const _ as *mut _,
12976 &pq as *const _ as *mut _,
12977 &pd as *const _ as *mut _,
12978 &nc as *const _ as *mut _,
12979 ];
12980 unsafe {
12981 self.launch_pdl(
12982 "gelu_tanh_mul_q8_1",
12983 (nrows as u32, 1, 1),
12984 (rms_block(), 1, 1),
12985 &mut ps,
12986 )?;
12987 }
12988 }
12989 return Ok((out_q, out_d));
12990 }
12991 let f = self.func("gelu_tanh_mul_q8_1");
12992 let cfg = LaunchConfig {
12993 grid_dim: (nrows as u32, 1, 1),
12994 block_dim: (rms_block(), 1, 1),
12995 shared_mem_bytes: 0,
12996 };
12997 let __s_b = self.gpu.stream();
12998 let mut b = __s_b.launch_builder(&f);
12999 b.arg(gate)
13000 .arg(up)
13001 .arg(act)
13002 .arg(&mut out_q)
13003 .arg(&mut out_d)
13004 .arg(&nc);
13005 unsafe {
13006 b.launch(cfg)?;
13007 }
13008 Ok((out_q, out_d))
13009 }
13010
13011 #[allow(clippy::too_many_arguments)]
13013 pub fn gelu_tanh_mul_q8_1_into(
13014 &self,
13015 gate: &CudaSlice<f32>,
13016 up: &cudarc::driver::CudaView<f32>,
13017 act: &mut CudaSlice<f32>,
13018 ncols: usize,
13019 nrows: usize,
13020 out_q: &mut CudaSlice<i8>,
13021 out_d: &mut CudaSlice<f32>,
13022 ) -> Result<(), Box<dyn std::error::Error>> {
13023 debug_assert!(ncols.is_multiple_of(128));
13024 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
13025 let nc = ncols as i32;
13026 if Self::pdl_on() {
13027 use cudarc::driver::{DevicePtr, DevicePtrMut};
13028 let s = &self.gpu.stream();
13029 let (pg, _g0) = gate.device_ptr(s);
13030 let (pu, _g1) = up.device_ptr(s);
13031 let (pact, _g2) = act.device_ptr_mut(s);
13032 let (pq, _g3) = out_q.device_ptr_mut(s);
13033 let (pd, _g4) = out_d.device_ptr_mut(s);
13034 let mut ps = [
13035 &pg as *const _ as *mut std::ffi::c_void,
13036 &pu as *const _ as *mut _,
13037 &pact as *const _ as *mut _,
13038 &pq as *const _ as *mut _,
13039 &pd as *const _ as *mut _,
13040 &nc as *const _ as *mut _,
13041 ];
13042 unsafe {
13043 self.launch_pdl(
13044 "gelu_tanh_mul_q8_1",
13045 (nrows as u32, 1, 1),
13046 (rms_block(), 1, 1),
13047 &mut ps,
13048 )?;
13049 }
13050 return Ok(());
13051 }
13052 let f = self.func("gelu_tanh_mul_q8_1");
13053 let cfg = LaunchConfig {
13054 grid_dim: (nrows as u32, 1, 1),
13055 block_dim: (rms_block(), 1, 1),
13056 shared_mem_bytes: 0,
13057 };
13058 let __s_b = self.gpu.stream();
13059 let mut b = __s_b.launch_builder(&f);
13060 b.arg(gate)
13061 .arg(up)
13062 .arg(&mut *act)
13063 .arg(&mut *out_q)
13064 .arg(&mut *out_d)
13065 .arg(&nc);
13066 unsafe {
13067 b.launch(cfg)?;
13068 }
13069 Ok(())
13070 }
13071
13072 #[allow(clippy::too_many_arguments)]
13074 #[allow(clippy::type_complexity)] pub fn add_rms_norm3_q8z(
13076 &self,
13077 a: &CudaSlice<f32>,
13078 b_in: &CudaSlice<f32>,
13079 w0: &CudaSlice<f32>,
13080 w1: &CudaSlice<f32>,
13081 w2: &CudaSlice<f32>,
13082 res: &mut CudaSlice<f32>,
13083 out1: &mut CudaSlice<f32>,
13084 ncols: usize,
13085 nrows: usize,
13086 eps: f32,
13087 ) -> Result<
13088 (
13089 (CudaSlice<i8>, CudaSlice<f32>),
13090 (CudaSlice<i8>, CudaSlice<f32>),
13091 ),
13092 Box<dyn std::error::Error>,
13093 > {
13094 let mut q0 = self.alloc_uninit::<i8>(nrows * ncols)?;
13095 let mut d0 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
13096 let mut q2 = self.alloc_uninit::<i8>(nrows * ncols)?;
13097 let mut d2 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
13098 let f = self.func("add_rms_norm3_q8z_f32");
13099 let cfg = LaunchConfig {
13100 grid_dim: (nrows as u32, 1, 1),
13101 block_dim: (rms_block(), 1, 1),
13102 shared_mem_bytes: 0,
13103 };
13104 let (nc, e2) = (ncols as i32, eps);
13105 let __s_b = self.gpu.stream();
13106 let mut b = __s_b.launch_builder(&f);
13107 b.arg(a)
13108 .arg(b_in)
13109 .arg(w0)
13110 .arg(w1)
13111 .arg(w2)
13112 .arg(res)
13113 .arg(&mut q0)
13114 .arg(&mut d0)
13115 .arg(out1)
13116 .arg(&mut q2)
13117 .arg(&mut d2)
13118 .arg(&nc)
13119 .arg(&e2);
13120 unsafe {
13121 b.launch(cfg)?;
13122 }
13123 Ok(((q0, d0), (q2, d2)))
13124 }
13125
13126 #[allow(clippy::too_many_arguments)]
13128 pub fn add_rms_norm3(
13129 &self,
13130 a: &CudaSlice<f32>,
13131 b_in: &CudaSlice<f32>,
13132 w0: &CudaSlice<f32>,
13133 w1: &CudaSlice<f32>,
13134 w2: &CudaSlice<f32>,
13135 res: &mut CudaSlice<f32>,
13136 d0: &mut CudaSlice<f32>,
13137 d1: &mut CudaSlice<f32>,
13138 d2: &mut CudaSlice<f32>,
13139 ncols: usize,
13140 nrows: usize,
13141 eps: f32,
13142 ) -> Result<(), Box<dyn std::error::Error>> {
13143 let f = self.func("add_rms_norm3_f32");
13144 let cfg = LaunchConfig {
13145 grid_dim: (nrows as u32, 1, 1),
13146 block_dim: (rms_block(), 1, 1),
13147 shared_mem_bytes: 0,
13148 };
13149 let (nc, e2) = (ncols as i32, eps);
13150 let __s_b = self.gpu.stream();
13151 let mut b = __s_b.launch_builder(&f);
13152 b.arg(a)
13153 .arg(b_in)
13154 .arg(w0)
13155 .arg(w1)
13156 .arg(w2)
13157 .arg(res)
13158 .arg(d0)
13159 .arg(d1)
13160 .arg(d2)
13161 .arg(&nc)
13162 .arg(&e2);
13163 unsafe {
13164 b.launch(cfg)?;
13165 }
13166 Ok(())
13167 }
13168
13169 pub fn add_scale(
13171 &self,
13172 a: &CudaSlice<f32>,
13173 b_in: &CudaSlice<f32>,
13174 c: f32,
13175 dst: &mut CudaSlice<f32>,
13176 n: usize,
13177 ) -> Result<(), Box<dyn std::error::Error>> {
13178 let f = self.func("add_scale_f32");
13179 let cfg = LaunchConfig::for_num_elems(n as u32);
13180 let ni = n as i32;
13181 let __s_b = self.gpu.stream();
13182 let mut b = __s_b.launch_builder(&f);
13183 b.arg(a).arg(b_in).arg(&c).arg(dst).arg(&ni);
13184 unsafe {
13185 b.launch(cfg)?;
13186 }
13187 Ok(())
13188 }
13189
13190 #[allow(clippy::too_many_arguments)] pub fn layer_norm_bias(
13193 &self,
13194 x: &CudaSlice<f32>,
13195 w: &CudaSlice<f32>,
13196 b: &CudaSlice<f32>,
13197 dst: &mut CudaSlice<f32>,
13198 ncols: usize,
13199 nrows: usize,
13200 eps: f32,
13201 ) -> Result<(), Box<dyn std::error::Error>> {
13202 let f = self.func("layer_norm_bias_f32");
13203 let (nc, e) = (ncols as i32, eps);
13204 let cfg = LaunchConfig {
13205 grid_dim: (nrows as u32, 1, 1),
13206 block_dim: (256, 1, 1),
13207 shared_mem_bytes: 0,
13208 };
13209 let __s_b = self.gpu.stream();
13210 let mut lb = __s_b.launch_builder(&f);
13211 lb.arg(x).arg(w).arg(b).arg(&mut *dst).arg(&nc).arg(&e);
13212 unsafe {
13213 lb.launch(cfg)?;
13214 }
13215 Ok(())
13216 }
13217
13218 pub fn gelu_tanh(
13220 &self,
13221 x: &CudaSlice<f32>,
13222 dst: &mut CudaSlice<f32>,
13223 n: usize,
13224 ) -> Result<(), Box<dyn std::error::Error>> {
13225 let f = self.func("gelu_tanh_f32");
13226 let ni = n as i64;
13227 let cfg = LaunchConfig {
13228 grid_dim: (n.div_ceil(256) as u32, 1, 1),
13229 block_dim: (256, 1, 1),
13230 shared_mem_bytes: 0,
13231 };
13232 let __s_b = self.gpu.stream();
13233 let mut lb = __s_b.launch_builder(&f);
13234 lb.arg(x).arg(&mut *dst).arg(&ni);
13235 unsafe {
13236 lb.launch(cfg)?;
13237 }
13238 Ok(())
13239 }
13240
13241 pub fn row_softmax(
13243 &self,
13244 x: &mut CudaSlice<f32>,
13245 ncols: usize,
13246 nrows: usize,
13247 ) -> Result<(), Box<dyn std::error::Error>> {
13248 let f = self.func("row_softmax_f32");
13249 let nc = ncols as i32;
13250 let cfg = LaunchConfig {
13251 grid_dim: (nrows as u32, 1, 1),
13252 block_dim: (256, 1, 1),
13253 shared_mem_bytes: 0,
13254 };
13255 let __s_b = self.gpu.stream();
13256 let mut lb = __s_b.launch_builder(&f);
13257 lb.arg(&mut *x).arg(&nc);
13258 unsafe {
13259 lb.launch(cfg)?;
13260 }
13261 Ok(())
13262 }
13263
13264 pub fn rms_norm(
13265 &self,
13266 x: &CudaSlice<f32>,
13267 w: &CudaSlice<f32>,
13268 dst: &mut CudaSlice<f32>,
13269 ncols: usize,
13270 nrows: usize,
13271 eps: f32,
13272 ) -> Result<(), Box<dyn std::error::Error>> {
13273 let (nc, e) = (ncols as i32, eps);
13274 let kname = if Self::norm_ilp_on() {
13275 "rms_norm_f32_v2"
13276 } else {
13277 "rms_norm_f32"
13278 };
13279 if Self::pdl_on() && Self::pdl_wb_on() {
13280 use cudarc::driver::{DevicePtr, DevicePtrMut};
13281 let s = &self.gpu.stream();
13282 let (px, _g0) = x.device_ptr(s);
13283 let (pw, _g1) = w.device_ptr(s);
13284 let (pd, _g2) = dst.device_ptr_mut(s);
13285 let mut ps = [
13286 &px as *const _ as *mut std::ffi::c_void,
13287 &pw as *const _ as *mut _,
13288 &pd as *const _ as *mut _,
13289 &nc as *const _ as *mut _,
13290 &e as *const _ as *mut _,
13291 ];
13292 unsafe {
13293 self.launch_pdl(kname, (nrows as u32, 1, 1), (rms_block(), 1, 1), &mut ps)?;
13294 }
13295 return Ok(());
13296 }
13297 let f = self.func(kname);
13298 let cfg = LaunchConfig {
13299 grid_dim: (nrows as u32, 1, 1),
13300 block_dim: (rms_block(), 1, 1),
13301 shared_mem_bytes: 0,
13302 };
13303 let __s_b = self.gpu.stream();
13304 let mut b = __s_b.launch_builder(&f);
13305 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
13306 unsafe {
13307 b.launch(cfg)?;
13308 }
13309 Ok(())
13310 }
13311
13312 pub fn rms_norm_decode(
13320 &self,
13321 x: &CudaSlice<f32>,
13322 w: &CudaSlice<f32>,
13323 dst: &mut CudaSlice<f32>,
13324 ncols: usize,
13325 nrows: usize,
13326 eps: f32,
13327 ) -> Result<(), Box<dyn std::error::Error>> {
13328 let f = self.func(if Self::norm_ilp_on() {
13329 "rms_norm_f32_v2"
13330 } else {
13331 "rms_norm_f32"
13332 });
13333 let cfg = LaunchConfig {
13334 grid_dim: (nrows as u32, 1, 1),
13335 block_dim: (1024, 1, 1),
13336 shared_mem_bytes: 0,
13337 };
13338 let (nc, e) = (ncols as i32, eps);
13339 let __s_b = self.gpu.stream();
13340 let mut b = __s_b.launch_builder(&f);
13341 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
13342 unsafe {
13343 b.launch(cfg)?;
13344 }
13345 Ok(())
13346 }
13347
13348 pub fn rms_norm_q8_1(
13352 &self,
13353 x: &CudaSlice<f32>,
13354 w: &CudaSlice<f32>,
13355 ncols: usize,
13356 nrows: usize,
13357 eps: f32,
13358 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13359 let nblk = ncols / 32;
13360 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
13361 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
13362 let (nc, e) = (ncols as i32, eps);
13363 if Self::pdl_on() {
13364 {
13365 use cudarc::driver::{DevicePtr, DevicePtrMut};
13366 let s = &self.gpu.stream();
13367 let (px, _g0) = x.device_ptr(s);
13368 let (pw, _g1) = w.device_ptr(s);
13369 let (pq, _g2) = q.device_ptr_mut(s);
13370 let (pd, _g3) = d.device_ptr_mut(s);
13371 let mut ps = [
13372 &px as *const _ as *mut std::ffi::c_void,
13373 &pw as *const _ as *mut _,
13374 &pq as *const _ as *mut _,
13375 &pd as *const _ as *mut _,
13376 &nc as *const _ as *mut _,
13377 &e as *const _ as *mut _,
13378 ];
13379 unsafe {
13380 self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1), &mut ps)?;
13381 }
13382 }
13383 return Ok((q, d));
13384 }
13385 let f = self.func("rms_norm_q8_1");
13386 let cfg = LaunchConfig {
13389 grid_dim: (nrows as u32, 1, 1),
13390 block_dim: (1024, 1, 1),
13391 shared_mem_bytes: 0,
13392 };
13393 let __s_b = self.gpu.stream();
13394 let mut b = __s_b.launch_builder(&f);
13395 b.arg(x).arg(w).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
13396 unsafe {
13397 b.launch(cfg)?;
13398 }
13399 Ok((q, d))
13400 }
13401
13402 #[allow(clippy::too_many_arguments)] pub fn rms_norm_q8_1_into(
13406 &self,
13407 x: &CudaSlice<f32>,
13408 w: &CudaSlice<f32>,
13409 ncols: usize,
13410 nrows: usize,
13411 eps: f32,
13412 q: &mut CudaSlice<i8>,
13413 d: &mut CudaSlice<f32>,
13414 ) -> Result<(), Box<dyn std::error::Error>> {
13415 let nblk = ncols / 32;
13416 debug_assert!(q.len() >= nrows * ncols && d.len() >= nrows * nblk);
13417 let (nc, e) = (ncols as i32, eps);
13418 if Self::pdl_on() {
13419 use cudarc::driver::{DevicePtr, DevicePtrMut};
13420 let s = &self.gpu.stream();
13421 let (px, _g0) = x.device_ptr(s);
13422 let (pw, _g1) = w.device_ptr(s);
13423 let (pq, _g2) = q.device_ptr_mut(s);
13424 let (pd, _g3) = d.device_ptr_mut(s);
13425 let mut ps = [
13426 &px as *const _ as *mut std::ffi::c_void,
13427 &pw as *const _ as *mut _,
13428 &pq as *const _ as *mut _,
13429 &pd as *const _ as *mut _,
13430 &nc as *const _ as *mut _,
13431 &e as *const _ as *mut _,
13432 ];
13433 unsafe {
13434 self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1), &mut ps)?;
13435 }
13436 return Ok(());
13437 }
13438 let f = self.func("rms_norm_q8_1");
13439 let cfg = LaunchConfig {
13440 grid_dim: (nrows as u32, 1, 1),
13441 block_dim: (1024, 1, 1),
13442 shared_mem_bytes: 0,
13443 };
13444 let __s_b = self.gpu.stream();
13445 let mut b = __s_b.launch_builder(&f);
13446 b.arg(x).arg(w).arg(&mut *q).arg(&mut *d).arg(&nc).arg(&e);
13447 unsafe {
13448 b.launch(cfg)?;
13449 }
13450 Ok(())
13451 }
13452
13453 pub fn quantize_q8_1_into(
13455 &self,
13456 x: &CudaSlice<f32>,
13457 m: usize,
13458 in_f: usize,
13459 q: &mut CudaSlice<i8>,
13460 d: &mut CudaSlice<f32>,
13461 ) -> Result<(), Box<dyn std::error::Error>> {
13462 let nblk = in_f / 32;
13463 debug_assert!(q.len() >= m * in_f && d.len() >= m * nblk);
13464 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
13465 let (inf, mi) = (in_f as i32, m as i32);
13466 if Self::pdl_on() && Self::pdl_wb_on() {
13467 use cudarc::driver::{DevicePtr, DevicePtrMut};
13468 let s = &self.gpu.stream();
13469 let (px, _g0) = x.device_ptr(s);
13470 let (pq, _g1) = q.device_ptr_mut(s);
13471 let (pd, _g2) = d.device_ptr_mut(s);
13472 let mut ps = [
13473 &px as *const _ as *mut std::ffi::c_void,
13474 &pq as *const _ as *mut _,
13475 &pd as *const _ as *mut _,
13476 &inf as *const _ as *mut _,
13477 &mi as *const _ as *mut _,
13478 ];
13479 unsafe {
13480 self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?;
13481 }
13482 return Ok(());
13483 }
13484 let f = self.func("quantize_q8_1");
13485 let __s_b = self.gpu.stream();
13486 let mut b = __s_b.launch_builder(&f);
13487 b.arg(x).arg(&mut *q).arg(&mut *d).arg(&inf).arg(&mi);
13488 unsafe {
13489 b.launch(cfg)?;
13490 }
13491 Ok(())
13492 }
13493
13494 #[allow(clippy::too_many_arguments)] pub fn add_rms_norm_q8_1(
13499 &self,
13500 a: &CudaSlice<f32>,
13501 b_in: &CudaSlice<f32>,
13502 w: &CudaSlice<f32>,
13503 res: &mut CudaSlice<f32>,
13504 ncols: usize,
13505 nrows: usize,
13506 eps: f32,
13507 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13508 let nblk = ncols / 32;
13509 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
13510 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
13511 let f = self.func("add_rms_norm_q8_1");
13512 let cfg = LaunchConfig {
13514 grid_dim: (nrows as u32, 1, 1),
13515 block_dim: (1024, 1, 1),
13516 shared_mem_bytes: 0,
13517 };
13518 let (nc, e) = (ncols as i32, eps);
13519 let __s_bld = self.gpu.stream();
13520 let mut bld = __s_bld.launch_builder(&f);
13521 bld.arg(a)
13522 .arg(b_in)
13523 .arg(w)
13524 .arg(res)
13525 .arg(&mut q)
13526 .arg(&mut d)
13527 .arg(&nc)
13528 .arg(&e);
13529 unsafe {
13530 bld.launch(cfg)?;
13531 }
13532 Ok((q, d))
13533 }
13534
13535 #[allow(clippy::too_many_arguments)]
13541 pub fn join_add_rms_norm_raw(
13542 &self,
13543 a0_raw: u64,
13544 a1_raw: u64,
13545 x: &CudaSlice<f32>,
13546 w: &CudaSlice<f32>,
13547 res: &mut CudaSlice<f32>,
13548 dst: &mut CudaSlice<f32>,
13549 ncols: usize,
13550 eps: f32,
13551 ) -> Result<(), Box<dyn std::error::Error>> {
13552 if a0_raw == 0 || a1_raw == 0 || x.len() < ncols || res.len() < ncols || dst.len() < ncols {
13553 return Err("join_add_rms_norm geometry".into());
13554 }
13555 let f = self.func("join_add_rms_norm_f32");
13556 let cfg = LaunchConfig {
13557 grid_dim: (1, 1, 1),
13558 block_dim: (rms_block(), 1, 1),
13559 shared_mem_bytes: 0,
13560 };
13561 let (nc, e) = (ncols as i32, eps);
13562 let __s_b = self.gpu.stream();
13563 let mut b = __s_b.launch_builder(&f);
13564 b.arg(&a0_raw)
13565 .arg(&a1_raw)
13566 .arg(x)
13567 .arg(w)
13568 .arg(&mut *res)
13569 .arg(&mut *dst)
13570 .arg(&nc)
13571 .arg(&e);
13572 unsafe {
13573 b.launch(cfg)?;
13574 }
13575 Ok(())
13576 }
13577
13578 #[allow(clippy::too_many_arguments)] pub fn add_rms_norm(
13580 &self,
13581 a: &CudaSlice<f32>,
13582 b: &CudaSlice<f32>,
13583 w: &CudaSlice<f32>,
13584 res: &mut CudaSlice<f32>,
13585 dst: &mut CudaSlice<f32>,
13586 ncols: usize,
13587 nrows: usize,
13588 eps: f32,
13589 ) -> Result<(), Box<dyn std::error::Error>> {
13590 let (nc, e) = (ncols as i32, eps);
13591 let kname = if Self::norm_ilp_on() {
13592 "add_rms_norm_f32_v2"
13593 } else {
13594 "add_rms_norm_f32"
13595 };
13596 if Self::pdl_on() && Self::pdl_wb_on() {
13597 use cudarc::driver::{DevicePtr, DevicePtrMut};
13598 let s = &self.gpu.stream();
13599 let (pa, _g0) = a.device_ptr(s);
13600 let (pb, _g1) = b.device_ptr(s);
13601 let (pw, _g2) = w.device_ptr(s);
13602 let (pr, _g3) = res.device_ptr_mut(s);
13603 let (pd, _g4) = dst.device_ptr_mut(s);
13604 let mut ps = [
13605 &pa as *const _ as *mut std::ffi::c_void,
13606 &pb as *const _ as *mut _,
13607 &pw as *const _ as *mut _,
13608 &pr as *const _ as *mut _,
13609 &pd as *const _ as *mut _,
13610 &nc as *const _ as *mut _,
13611 &e as *const _ as *mut _,
13612 ];
13613 unsafe {
13614 self.launch_pdl(kname, (nrows as u32, 1, 1), (rms_block(), 1, 1), &mut ps)?;
13615 }
13616 return Ok(());
13617 }
13618 let f = self.func(kname);
13619 let cfg = LaunchConfig {
13620 grid_dim: (nrows as u32, 1, 1),
13621 block_dim: (rms_block(), 1, 1),
13622 shared_mem_bytes: 0,
13623 };
13624 let __s_b2 = self.gpu.stream();
13625 let mut b2 = __s_b2.launch_builder(&f);
13626 b2.arg(a)
13627 .arg(b)
13628 .arg(w)
13629 .arg(&mut *res)
13630 .arg(&mut *dst)
13631 .arg(&nc)
13632 .arg(&e);
13633 unsafe {
13634 b2.launch(cfg)?;
13635 }
13636 Ok(())
13637 }
13638
13639 #[allow(clippy::too_many_arguments)]
13642 pub fn rms_pre_add_rms_norm(
13643 &self,
13644 a: &CudaSlice<f32>,
13645 wa: &CudaSlice<f32>,
13646 b: &CudaSlice<f32>,
13647 w: &CudaSlice<f32>,
13648 res: &mut CudaSlice<f32>,
13649 dst: &mut CudaSlice<f32>,
13650 ncols: usize,
13651 nrows: usize,
13652 eps: f32,
13653 ) -> Result<(), Box<dyn std::error::Error>> {
13654 let f = self.func("rms_pre_add_rms_norm_f32");
13655 let cfg = LaunchConfig {
13656 grid_dim: (nrows as u32, 1, 1),
13657 block_dim: (rms_block(), 1, 1),
13658 shared_mem_bytes: 0,
13659 };
13660 let (nc, e) = (ncols as i32, eps);
13661 let __s_b2 = self.gpu.stream();
13662 let mut b2 = __s_b2.launch_builder(&f);
13663 b2.arg(a)
13664 .arg(wa)
13665 .arg(b)
13666 .arg(w)
13667 .arg(&mut *res)
13668 .arg(&mut *dst)
13669 .arg(&nc)
13670 .arg(&e);
13671 unsafe {
13672 b2.launch(cfg)?;
13673 }
13674 Ok(())
13675 }
13676
13677 #[allow(clippy::too_many_arguments)]
13679 pub fn rms_pre_add_rms_norm_q8z(
13680 &self,
13681 a: &CudaSlice<f32>,
13682 wa: &CudaSlice<f32>,
13683 b: &CudaSlice<f32>,
13684 w: &CudaSlice<f32>,
13685 res: &mut CudaSlice<f32>,
13686 dst: &mut CudaSlice<f32>,
13687 ncols: usize,
13688 nrows: usize,
13689 eps: f32,
13690 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13691 debug_assert!(ncols.is_multiple_of(128));
13692 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
13693 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
13694 let (nc, e) = (ncols as i32, eps);
13695 if Self::pdl_on() {
13696 {
13697 use cudarc::driver::{DevicePtr, DevicePtrMut};
13698 let s = &self.gpu.stream();
13699 let (pa, _g0) = a.device_ptr(s);
13700 let (pwa, _g1) = wa.device_ptr(s);
13701 let (pb, _g2) = b.device_ptr(s);
13702 let (pw, _g3) = w.device_ptr(s);
13703 let (pr, _g4) = res.device_ptr_mut(s);
13704 let (pdst, _g5) = dst.device_ptr_mut(s);
13705 let (pq, _g6) = out_q.device_ptr_mut(s);
13706 let (pd, _g7) = out_d.device_ptr_mut(s);
13707 let mut ps = [
13708 &pa as *const _ as *mut std::ffi::c_void,
13709 &pwa as *const _ as *mut _,
13710 &pb as *const _ as *mut _,
13711 &pw as *const _ as *mut _,
13712 &pr as *const _ as *mut _,
13713 &pdst as *const _ as *mut _,
13714 &pq as *const _ as *mut _,
13715 &pd as *const _ as *mut _,
13716 &nc as *const _ as *mut _,
13717 &e as *const _ as *mut _,
13718 ];
13719 unsafe {
13720 self.launch_pdl(
13721 "rms_pre_add_rms_norm_q8z_f32",
13722 (nrows as u32, 1, 1),
13723 (rms_block(), 1, 1),
13724 &mut ps,
13725 )?;
13726 }
13727 }
13728 return Ok((out_q, out_d));
13729 }
13730 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
13731 let cfg = LaunchConfig {
13732 grid_dim: (nrows as u32, 1, 1),
13733 block_dim: (rms_block(), 1, 1),
13734 shared_mem_bytes: 0,
13735 };
13736 let __s_b2 = self.gpu.stream();
13737 let mut b2 = __s_b2.launch_builder(&f);
13738 b2.arg(a)
13739 .arg(wa)
13740 .arg(b)
13741 .arg(w)
13742 .arg(&mut *res)
13743 .arg(&mut *dst)
13744 .arg(&mut out_q)
13745 .arg(&mut out_d)
13746 .arg(&nc)
13747 .arg(&e);
13748 unsafe {
13749 b2.launch(cfg)?;
13750 }
13751 Ok((out_q, out_d))
13752 }
13753
13754 #[allow(clippy::too_many_arguments)]
13758 pub fn rms_pre_add_rms_norm_q8z_into(
13759 &self,
13760 a: &CudaSlice<f32>,
13761 wa: &CudaSlice<f32>,
13762 b: &CudaSlice<f32>,
13763 w: &CudaSlice<f32>,
13764 res: &mut CudaSlice<f32>,
13765 dst: &mut CudaSlice<f32>,
13766 ncols: usize,
13767 nrows: usize,
13768 eps: f32,
13769 out_q: &mut CudaSlice<i8>,
13770 out_d: &mut CudaSlice<f32>,
13771 ) -> Result<(), Box<dyn std::error::Error>> {
13772 debug_assert!(ncols.is_multiple_of(128));
13773 let (nc, e) = (ncols as i32, eps);
13774 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
13775 let cfg = LaunchConfig {
13776 grid_dim: (nrows as u32, 1, 1),
13777 block_dim: (rms_block(), 1, 1),
13778 shared_mem_bytes: 0,
13779 };
13780 let __s_b = self.gpu.stream();
13781 let mut b2 = __s_b.launch_builder(&f);
13782 b2.arg(a)
13783 .arg(wa)
13784 .arg(b)
13785 .arg(w)
13786 .arg(&mut *res)
13787 .arg(&mut *dst)
13788 .arg(&mut *out_q)
13789 .arg(&mut *out_d)
13790 .arg(&nc)
13791 .arg(&e);
13792 unsafe {
13793 b2.launch(cfg)?;
13794 }
13795 Ok(())
13796 }
13797
13798 #[allow(clippy::too_many_arguments)]
13801 pub fn rms_pre_add_scale_rms_norm_q8_1_into(
13802 &self,
13803 a: &CudaSlice<f32>,
13804 wa: &CudaSlice<f32>,
13805 b_in: &CudaSlice<f32>,
13806 c: f32,
13807 w: &CudaSlice<f32>,
13808 res: &mut CudaSlice<f32>,
13809 ncols: usize,
13810 nrows: usize,
13811 eps: f32,
13812 out_q: &mut CudaSlice<i8>,
13813 out_d: &mut CudaSlice<f32>,
13814 ) -> Result<(), Box<dyn std::error::Error>> {
13815 debug_assert!(ncols.is_multiple_of(128));
13816 let (nc, e2) = (ncols as i32, eps);
13817 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
13818 let cfg = LaunchConfig {
13819 grid_dim: (nrows as u32, 1, 1),
13820 block_dim: (rms_block(), 1, 1),
13821 shared_mem_bytes: 0,
13822 };
13823 let __s_b = self.gpu.stream();
13824 let mut b2 = __s_b.launch_builder(&f);
13825 b2.arg(a)
13826 .arg(wa)
13827 .arg(b_in)
13828 .arg(&c)
13829 .arg(w)
13830 .arg(&mut *res)
13831 .arg(&mut *out_q)
13832 .arg(&mut *out_d)
13833 .arg(&nc)
13834 .arg(&e2);
13835 unsafe {
13836 b2.launch(cfg)?;
13837 }
13838 Ok(())
13839 }
13840
13841 pub fn g4_pnfold_on() -> bool {
13849 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13850 *ON.get_or_init(|| {
13851 std::env::var("MEMRA_G4_PNFOLD")
13852 .map(|v| v != "0")
13853 .unwrap_or(true)
13854 })
13855 }
13856
13857 pub fn build_q4_out_concat3(
13861 &self,
13862 w0: &crate::model::GpuTensor,
13863 w1: &crate::model::GpuTensor,
13864 w2: &crate::model::GpuTensor,
13865 ) -> Result<Option<crate::model::GpuTensor>, Box<dyn std::error::Error>> {
13866 use crate::model::GpuTensor;
13867 let part = |w: &GpuTensor| -> Option<(usize, usize)> {
13868 match w {
13869 GpuTensor::Quant {
13870 qtype,
13871 row_bytes,
13872 rp,
13873 ..
13874 } if *qtype == QT_Q4_0 && !*rp => Some((*row_bytes, w.out_features())),
13875 _ => None,
13876 }
13877 };
13878 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (part(w0), part(w1), part(w2))
13879 else {
13880 return Ok(None);
13881 };
13882 if rb0 != rb1
13883 || rb0 != rb2
13884 || w0.in_features() != w1.in_features()
13885 || w0.in_features() != w2.in_features()
13886 {
13887 return Ok(None);
13888 }
13889 fn bytes_of(w: &crate::model::GpuTensor) -> &CudaSlice<u8> {
13890 match w {
13891 crate::model::GpuTensor::Quant { bytes, .. } => bytes,
13892 _ => unreachable!(),
13893 }
13894 }
13895 let (b0, b1, b2) = (bytes_of(w0), bytes_of(w1), bytes_of(w2));
13896 let total = rb0 * (o0 + o1 + o2);
13897 let mut cat = self.alloc_u8(total)?;
13898 self.copy_u8_into(&mut cat, 0, b0, rb0 * o0)?;
13899 self.copy_u8_into(&mut cat, rb0 * o0, b1, rb1 * o1)?;
13900 self.copy_u8_into(&mut cat, rb0 * (o0 + o1), b2, rb2 * o2)?;
13901 Ok(Some(GpuTensor::Quant {
13902 bytes: cat,
13903 qtype: QT_Q4_0,
13904 row_bytes: rb0,
13905 ne: vec![w0.in_features() as u64, (o0 + o1 + o2) as u64],
13906 scale: 1.0,
13907 rp: false,
13908 #[cfg(memra_cutlass)]
13909 cutlass: None,
13910 fp8: None,
13911 blk: None,
13912 rp4: None,
13913 f16: None,
13914 }))
13915 }
13916
13917 fn full_width_rope_only(
13935 kernel: &str,
13936 n_rot: usize,
13937 head_dim: usize,
13938 ) -> Result<(), Box<dyn std::error::Error>> {
13939 if n_rot == head_dim {
13940 return Ok(());
13941 }
13942 Err(format!(
13943 "{kernel}: PARTIAL ROTARY REFUSED — n_rot {n_rot} != head_dim {head_dim}. This fused \
13944 rms_norm+qkv+rope kernel carries no n_dims parameter and rotates the full head \
13945 width (half = ncols/2), so it would rotate dims {n_rot}..{head_dim} that must pass \
13946 through unrotated. Use the split path (rms_norm_qkv + rope_neox/rope_neox2 with \
13947 n_dims={n_rot}), or add an n_dims early-return to the kernel and widen this guard."
13948 )
13949 .into())
13950 }
13951
13952 #[allow(clippy::too_many_arguments)]
13956 pub fn rms_norm_qkv_rope_cat(
13957 &self,
13958 qkv: &CudaSlice<f32>,
13959 wq: &CudaSlice<f32>,
13960 wk: &CudaSlice<f32>,
13961 wv: &CudaSlice<f32>,
13962 q: &mut CudaSlice<f32>,
13963 k: &mut CudaSlice<f32>,
13964 v: &mut CudaSlice<f32>,
13965 head_dim: usize,
13966 n_rot: usize,
13967 rq: usize,
13968 rk: usize,
13969 pos: &CudaSlice<i32>,
13970 nh_q: usize,
13971 nh_k: usize,
13972 base: f32,
13973 freq_scale: f32,
13974 ff: Option<&CudaSlice<f32>>,
13975 eps: f32,
13976 ) -> Result<(), Box<dyn std::error::Error>> {
13977 Self::full_width_rope_only("rms_norm_qkv_rope_cat", n_rot, head_dim)?;
13978 let rows = rq + rk + rk;
13979 let theta_scale = base.powf(-2.0 / head_dim as f32);
13980 let (nc, rqi, rki, nhq, nhk) = (
13981 head_dim as i32,
13982 rq as i32,
13983 rk as i32,
13984 nh_q as i32,
13985 nh_k as i32,
13986 );
13987 if Self::pdl_on() {
13988 use cudarc::driver::{DevicePtr, DevicePtrMut};
13989 let s = &self.gpu.stream();
13990 let (pqkv, _g0) = qkv.device_ptr(s);
13991 let (pwq, _g1) = wq.device_ptr(s);
13992 let (pwk, _g2) = wk.device_ptr(s);
13993 let (pwv, _g3) = wv.device_ptr(s);
13994 let (pq, _g4) = q.device_ptr_mut(s);
13995 let (pk, _g5) = k.device_ptr_mut(s);
13996 let (pv, _g6) = v.device_ptr_mut(s);
13997 let (ppos, _g7) = pos.device_ptr(s);
13998 let (pff, _g8) = match ff {
13999 Some(t) => {
14000 let (p, g) = t.device_ptr(s);
14001 (p, Some(g))
14002 }
14003 None => (0, None),
14004 };
14005 let mut ps = [
14006 &pqkv as *const _ as *mut std::ffi::c_void,
14007 &pwq as *const _ as *mut _,
14008 &pwk as *const _ as *mut _,
14009 &pwv as *const _ as *mut _,
14010 &pq as *const _ as *mut _,
14011 &pk as *const _ as *mut _,
14012 &pv as *const _ as *mut _,
14013 &nc as *const _ as *mut _,
14014 &rqi as *const _ as *mut _,
14015 &rki as *const _ as *mut _,
14016 &ppos as *const _ as *mut _,
14017 &nhq as *const _ as *mut _,
14018 &nhk as *const _ as *mut _,
14019 &theta_scale as *const _ as *mut _,
14020 &freq_scale as *const _ as *mut _,
14021 &pff as *const _ as *mut _,
14022 &eps as *const _ as *mut _,
14023 ];
14024 unsafe {
14025 self.launch_pdl(
14026 "rms_norm_qkv_rope_cat_f32",
14027 (rows as u32, 1, 1),
14028 (rms_block(), 1, 1),
14029 &mut ps,
14030 )?;
14031 }
14032 return Ok(());
14033 }
14034 let f = self.func("rms_norm_qkv_rope_cat_f32");
14035 let cfg = LaunchConfig {
14036 grid_dim: (rows as u32, 1, 1),
14037 block_dim: (rms_block(), 1, 1),
14038 shared_mem_bytes: 0,
14039 };
14040 let __s_b = self.gpu.stream();
14041 let mut b = __s_b.launch_builder(&f);
14042 match ff {
14043 Some(t) => {
14044 b.arg(qkv)
14045 .arg(wq)
14046 .arg(wk)
14047 .arg(wv)
14048 .arg(&mut *q)
14049 .arg(&mut *k)
14050 .arg(&mut *v)
14051 .arg(&nc)
14052 .arg(&rqi)
14053 .arg(&rki)
14054 .arg(pos)
14055 .arg(&nhq)
14056 .arg(&nhk)
14057 .arg(&theta_scale)
14058 .arg(&freq_scale)
14059 .arg(t)
14060 .arg(&eps);
14061 unsafe {
14062 b.launch(cfg)?;
14063 }
14064 }
14065 None => {
14066 let null: u64 = 0;
14067 b.arg(qkv)
14068 .arg(wq)
14069 .arg(wk)
14070 .arg(wv)
14071 .arg(&mut *q)
14072 .arg(&mut *k)
14073 .arg(&mut *v)
14074 .arg(&nc)
14075 .arg(&rqi)
14076 .arg(&rki)
14077 .arg(pos)
14078 .arg(&nhq)
14079 .arg(&nhk)
14080 .arg(&theta_scale)
14081 .arg(&freq_scale)
14082 .arg(&null)
14083 .arg(&eps);
14084 unsafe {
14085 b.launch(cfg)?;
14086 }
14087 }
14088 }
14089 Ok(())
14090 }
14091
14092 #[allow(clippy::too_many_arguments)]
14096 pub fn rms_norm_qkv_rope(
14097 &self,
14098 q0: &CudaSlice<f32>,
14099 k0: &CudaSlice<f32>,
14100 v0: &CudaSlice<f32>,
14101 wq: &CudaSlice<f32>,
14102 wk: &CudaSlice<f32>,
14103 wv: &CudaSlice<f32>,
14104 q: &mut CudaSlice<f32>,
14105 k: &mut CudaSlice<f32>,
14106 v: &mut CudaSlice<f32>,
14107 head_dim: usize,
14108 n_rot: usize,
14109 rq: usize,
14110 rk: usize,
14111 pos: &CudaSlice<i32>,
14112 nh_q: usize,
14113 nh_k: usize,
14114 base: f32,
14115 freq_scale: f32,
14116 ff: Option<&CudaSlice<f32>>,
14117 eps: f32,
14118 ) -> Result<(), Box<dyn std::error::Error>> {
14119 Self::full_width_rope_only("rms_norm_qkv_rope", n_rot, head_dim)?;
14120 let f = self.func("rms_norm_qkv_rope_f32");
14121 let rows = rq + rk + rk; let cfg = LaunchConfig {
14123 grid_dim: (rows as u32, 1, 1),
14124 block_dim: (rms_block(), 1, 1),
14125 shared_mem_bytes: 0,
14126 };
14127 let theta_scale = base.powf(-2.0 / head_dim as f32);
14128 let (nc, rqi, rki, nhq, nhk) = (
14129 head_dim as i32,
14130 rq as i32,
14131 rk as i32,
14132 nh_q as i32,
14133 nh_k as i32,
14134 );
14135 let __s_b = self.gpu.stream();
14136 let mut b = __s_b.launch_builder(&f);
14137 match ff {
14138 Some(t) => {
14139 b.arg(q0)
14140 .arg(k0)
14141 .arg(v0)
14142 .arg(wq)
14143 .arg(wk)
14144 .arg(wv)
14145 .arg(&mut *q)
14146 .arg(&mut *k)
14147 .arg(&mut *v)
14148 .arg(&nc)
14149 .arg(&rqi)
14150 .arg(&rki)
14151 .arg(pos)
14152 .arg(&nhq)
14153 .arg(&nhk)
14154 .arg(&theta_scale)
14155 .arg(&freq_scale)
14156 .arg(t)
14157 .arg(&eps);
14158 unsafe {
14159 b.launch(cfg)?;
14160 }
14161 }
14162 None => {
14163 let null: u64 = 0;
14164 b.arg(q0)
14165 .arg(k0)
14166 .arg(v0)
14167 .arg(wq)
14168 .arg(wk)
14169 .arg(wv)
14170 .arg(&mut *q)
14171 .arg(&mut *k)
14172 .arg(&mut *v)
14173 .arg(&nc)
14174 .arg(&rqi)
14175 .arg(&rki)
14176 .arg(pos)
14177 .arg(&nhq)
14178 .arg(&nhk)
14179 .arg(&theta_scale)
14180 .arg(&freq_scale)
14181 .arg(&null)
14182 .arg(&eps);
14183 unsafe {
14184 b.launch(cfg)?;
14185 }
14186 }
14187 }
14188 Ok(())
14189 }
14190
14191 #[allow(clippy::too_many_arguments)]
14197 pub fn rms_norm_qkv_rope_append_dc(
14198 &self,
14199 q0: &CudaSlice<f32>,
14200 k0: &CudaSlice<f32>,
14201 v0: &CudaSlice<f32>,
14202 wq: &CudaSlice<f32>,
14203 wk: &CudaSlice<f32>,
14204 wv: &CudaSlice<f32>,
14205 q: &mut CudaSlice<f32>,
14206 k: &mut CudaSlice<f32>,
14207 v: &mut CudaSlice<f32>,
14208 head_dim: usize,
14209 n_rot: usize,
14210 rq: usize,
14211 rk: usize,
14212 pos: &CudaSlice<i32>,
14213 nh_q: usize,
14214 nh_k: usize,
14215 base: f32,
14216 freq_scale: f32,
14217 ff: Option<&CudaSlice<f32>>,
14218 eps: f32,
14219 kc: &mut CudaSlice<u8>,
14220 vc: &mut CudaSlice<u8>,
14221 t_dev: &CudaSlice<i32>,
14222 k_tok_bytes: usize,
14223 v_tok_bytes: usize,
14224 g: bool,
14225 ) -> Result<(), Box<dyn std::error::Error>> {
14226 Self::full_width_rope_only("rms_norm_qkv_rope_append_dc", n_rot, head_dim)?;
14227 let rows = rq + rk + rk;
14228 let theta_scale = base.powf(-2.0 / head_dim as f32);
14229 let (nc, rqi, rki, nhq, nhk) = (
14230 head_dim as i32,
14231 rq as i32,
14232 rk as i32,
14233 nh_q as i32,
14234 nh_k as i32,
14235 );
14236 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
14237 if Self::pdl_on() && Self::pdl_wb_on() {
14238 use cudarc::driver::{DevicePtr, DevicePtrMut};
14239 let s = &self.gpu.stream();
14240 let (p0, _a0) = q0.device_ptr(s);
14241 let (p1, _a1) = k0.device_ptr(s);
14242 let (p2, _a2) = v0.device_ptr(s);
14243 let (pwq, _a3) = wq.device_ptr(s);
14244 let (pwk, _a4) = wk.device_ptr(s);
14245 let (pwv, _a5) = wv.device_ptr(s);
14246 let (pq, _a6) = q.device_ptr_mut(s);
14247 let (pk, _a7) = k.device_ptr_mut(s);
14248 let (pv, _a8) = v.device_ptr_mut(s);
14249 let (pp, _a9) = pos.device_ptr(s);
14250 let pff: u64 = match ff {
14251 Some(t) => {
14252 let (p, _gg) = t.device_ptr(s);
14253 p
14254 }
14255 None => 0,
14256 };
14257 let (pkc, _a10) = kc.device_ptr_mut(s);
14258 let (pvc, _a11) = vc.device_ptr_mut(s);
14259 let (pt, _a12) = t_dev.device_ptr(s);
14260 let mut ps = [
14261 &p0 as *const _ as *mut std::ffi::c_void,
14262 &p1 as *const _ as *mut _,
14263 &p2 as *const _ as *mut _,
14264 &pwq as *const _ as *mut _,
14265 &pwk as *const _ as *mut _,
14266 &pwv as *const _ as *mut _,
14267 &pq as *const _ as *mut _,
14268 &pk as *const _ as *mut _,
14269 &pv as *const _ as *mut _,
14270 &nc as *const _ as *mut _,
14271 &rqi as *const _ as *mut _,
14272 &rki as *const _ as *mut _,
14273 &pp as *const _ as *mut _,
14274 &nhq as *const _ as *mut _,
14275 &nhk as *const _ as *mut _,
14276 &theta_scale as *const _ as *mut _,
14277 &freq_scale as *const _ as *mut _,
14278 &pff as *const _ as *mut _,
14279 &eps as *const _ as *mut _,
14280 &pkc as *const _ as *mut _,
14281 &pvc as *const _ as *mut _,
14282 &pt as *const _ as *mut _,
14283 &ktb as *const _ as *mut _,
14284 &vtb as *const _ as *mut _,
14285 ];
14286 unsafe {
14287 self.launch_pdl_flash(
14288 g,
14289 "rms_norm_qkv_rope_append_dc_f32",
14290 (rows as u32, 1, 1),
14291 (rms_block(), 1, 1),
14292 0,
14293 &mut ps,
14294 )?;
14295 }
14296 return Ok(());
14297 }
14298 let f = if g {
14299 self.func_g("rms_norm_qkv_rope_append_dc_f32")
14300 } else {
14301 self.func("rms_norm_qkv_rope_append_dc_f32")
14302 };
14303 let cfg = LaunchConfig {
14304 grid_dim: (rows as u32, 1, 1),
14305 block_dim: (rms_block(), 1, 1),
14306 shared_mem_bytes: 0,
14307 };
14308 let __s_b = self.gpu.stream();
14309 let mut b = __s_b.launch_builder(&f);
14310 match ff {
14311 Some(t) => {
14312 b.arg(q0)
14313 .arg(k0)
14314 .arg(v0)
14315 .arg(wq)
14316 .arg(wk)
14317 .arg(wv)
14318 .arg(&mut *q)
14319 .arg(&mut *k)
14320 .arg(&mut *v)
14321 .arg(&nc)
14322 .arg(&rqi)
14323 .arg(&rki)
14324 .arg(pos)
14325 .arg(&nhq)
14326 .arg(&nhk)
14327 .arg(&theta_scale)
14328 .arg(&freq_scale)
14329 .arg(t)
14330 .arg(&eps)
14331 .arg(&mut *kc)
14332 .arg(&mut *vc)
14333 .arg(t_dev)
14334 .arg(&ktb)
14335 .arg(&vtb);
14336 unsafe {
14337 b.launch(cfg)?;
14338 }
14339 }
14340 None => {
14341 let null: u64 = 0;
14342 b.arg(q0)
14343 .arg(k0)
14344 .arg(v0)
14345 .arg(wq)
14346 .arg(wk)
14347 .arg(wv)
14348 .arg(&mut *q)
14349 .arg(&mut *k)
14350 .arg(&mut *v)
14351 .arg(&nc)
14352 .arg(&rqi)
14353 .arg(&rki)
14354 .arg(pos)
14355 .arg(&nhq)
14356 .arg(&nhk)
14357 .arg(&theta_scale)
14358 .arg(&freq_scale)
14359 .arg(&null)
14360 .arg(&eps)
14361 .arg(&mut *kc)
14362 .arg(&mut *vc)
14363 .arg(t_dev)
14364 .arg(&ktb)
14365 .arg(&vtb);
14366 unsafe {
14367 b.launch(cfg)?;
14368 }
14369 }
14370 }
14371 Ok(())
14372 }
14373
14374 #[allow(clippy::too_many_arguments)]
14382 pub fn rms_norm_qkv_rope_append(
14383 &self,
14384 q0: &CudaSlice<f32>,
14385 k0: &CudaSlice<f32>,
14386 v0: &CudaSlice<f32>,
14387 wq: &CudaSlice<f32>,
14388 wk: &CudaSlice<f32>,
14389 wv: &CudaSlice<f32>,
14390 q: &mut CudaSlice<f32>,
14391 k: &mut CudaSlice<f32>,
14392 v: &mut CudaSlice<f32>,
14393 head_dim: usize,
14394 n_rot: usize,
14395 rq: usize,
14396 rk: usize,
14397 pos: &CudaSlice<i32>,
14398 nh_q: usize,
14399 nh_k: usize,
14400 base: f32,
14401 freq_scale: f32,
14402 ff: Option<&CudaSlice<f32>>,
14403 eps: f32,
14404 kc: &mut CudaSlice<u8>,
14405 vc: &mut CudaSlice<u8>,
14406 t: usize,
14407 k_tok_bytes: usize,
14408 v_tok_bytes: usize,
14409 g: bool,
14410 ) -> Result<(), Box<dyn std::error::Error>> {
14411 Self::full_width_rope_only("rms_norm_qkv_rope_append", n_rot, head_dim)?;
14412 let rows = rq + rk + rk;
14413 let theta_scale = base.powf(-2.0 / head_dim as f32);
14414 let (nc, rqi, rki, nhq, nhk) = (
14415 head_dim as i32,
14416 rq as i32,
14417 rk as i32,
14418 nh_q as i32,
14419 nh_k as i32,
14420 );
14421 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
14422 let ti = t as i32;
14423 if Self::pdl_on() && Self::pdl_wb_on() {
14424 use cudarc::driver::{DevicePtr, DevicePtrMut};
14425 let s = &self.gpu.stream();
14426 let (p0, _a0) = q0.device_ptr(s);
14427 let (p1, _a1) = k0.device_ptr(s);
14428 let (p2, _a2) = v0.device_ptr(s);
14429 let (pwq, _a3) = wq.device_ptr(s);
14430 let (pwk, _a4) = wk.device_ptr(s);
14431 let (pwv, _a5) = wv.device_ptr(s);
14432 let (pq, _a6) = q.device_ptr_mut(s);
14433 let (pk, _a7) = k.device_ptr_mut(s);
14434 let (pv, _a8) = v.device_ptr_mut(s);
14435 let (pp, _a9) = pos.device_ptr(s);
14436 let pff: u64 = match ff {
14437 Some(t) => {
14438 let (p, _gg) = t.device_ptr(s);
14439 p
14440 }
14441 None => 0,
14442 };
14443 let (pkc, _a10) = kc.device_ptr_mut(s);
14444 let (pvc, _a11) = vc.device_ptr_mut(s);
14445 let mut ps = [
14446 &p0 as *const _ as *mut std::ffi::c_void,
14447 &p1 as *const _ as *mut _,
14448 &p2 as *const _ as *mut _,
14449 &pwq as *const _ as *mut _,
14450 &pwk as *const _ as *mut _,
14451 &pwv as *const _ as *mut _,
14452 &pq as *const _ as *mut _,
14453 &pk as *const _ as *mut _,
14454 &pv as *const _ as *mut _,
14455 &nc as *const _ as *mut _,
14456 &rqi as *const _ as *mut _,
14457 &rki as *const _ as *mut _,
14458 &pp as *const _ as *mut _,
14459 &nhq as *const _ as *mut _,
14460 &nhk as *const _ as *mut _,
14461 &theta_scale as *const _ as *mut _,
14462 &freq_scale as *const _ as *mut _,
14463 &pff as *const _ as *mut _,
14464 &eps as *const _ as *mut _,
14465 &pkc as *const _ as *mut _,
14466 &pvc as *const _ as *mut _,
14467 &ti as *const _ as *mut _,
14468 &ktb as *const _ as *mut _,
14469 &vtb as *const _ as *mut _,
14470 ];
14471 unsafe {
14472 self.launch_pdl_flash(
14473 g,
14474 "rms_norm_qkv_rope_append_f32",
14475 (rows as u32, 1, 1),
14476 (rms_block(), 1, 1),
14477 0,
14478 &mut ps,
14479 )?;
14480 }
14481 return Ok(());
14482 }
14483 let f = if g {
14484 self.func_g("rms_norm_qkv_rope_append_f32")
14485 } else {
14486 self.func("rms_norm_qkv_rope_append_f32")
14487 };
14488 let cfg = LaunchConfig {
14489 grid_dim: (rows as u32, 1, 1),
14490 block_dim: (rms_block(), 1, 1),
14491 shared_mem_bytes: 0,
14492 };
14493 let __s_b = self.gpu.stream();
14494 let mut b = __s_b.launch_builder(&f);
14495 let null: u64 = 0;
14496 b.arg(q0)
14497 .arg(k0)
14498 .arg(v0)
14499 .arg(wq)
14500 .arg(wk)
14501 .arg(wv)
14502 .arg(&mut *q)
14503 .arg(&mut *k)
14504 .arg(&mut *v)
14505 .arg(&nc)
14506 .arg(&rqi)
14507 .arg(&rki)
14508 .arg(pos)
14509 .arg(&nhq)
14510 .arg(&nhk)
14511 .arg(&theta_scale)
14512 .arg(&freq_scale);
14513 match ff {
14514 Some(t) => {
14515 b.arg(t);
14516 }
14517 None => {
14518 b.arg(&null);
14519 }
14520 }
14521 b.arg(&eps)
14522 .arg(&mut *kc)
14523 .arg(&mut *vc)
14524 .arg(&ti)
14525 .arg(&ktb)
14526 .arg(&vtb);
14527 unsafe {
14528 b.launch(cfg)?;
14529 }
14530 Ok(())
14531 }
14532
14533 pub fn add_q8_1(
14534 &self,
14535 a: &CudaSlice<f32>,
14536 b: &CudaSlice<f32>,
14537 res: &mut CudaSlice<f32>,
14538 ncols: usize,
14539 nrows: usize,
14540 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14541 debug_assert!(ncols.is_multiple_of(128));
14542 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
14543 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
14544 let f = self.func("add_q8_1_f32");
14545 let cfg = LaunchConfig {
14546 grid_dim: (nrows as u32, 1, 1),
14547 block_dim: (rms_block(), 1, 1),
14548 shared_mem_bytes: 0,
14549 };
14550 let nc = ncols as i32;
14551 let __s_b2 = self.gpu.stream();
14552 let mut b2 = __s_b2.launch_builder(&f);
14553 b2.arg(a)
14554 .arg(b)
14555 .arg(&mut *res)
14556 .arg(&mut out_q)
14557 .arg(&mut out_d)
14558 .arg(&nc);
14559 unsafe {
14560 b2.launch(cfg)?;
14561 }
14562 Ok((out_q, out_d))
14563 }
14564
14565 #[allow(clippy::too_many_arguments)] pub fn rms_pre_add_q8_1(
14570 &self,
14571 a: &CudaSlice<f32>,
14572 wa: &CudaSlice<f32>,
14573 b: &CudaSlice<f32>,
14574 res: &mut CudaSlice<f32>,
14575 ncols: usize,
14576 nrows: usize,
14577 eps: f32,
14578 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14579 debug_assert!(ncols.is_multiple_of(128));
14580 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
14581 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
14582 let f = self.func("rms_pre_add_q8_1_f32");
14583 let cfg = LaunchConfig {
14584 grid_dim: (nrows as u32, 1, 1),
14585 block_dim: (rms_block(), 1, 1),
14586 shared_mem_bytes: 0,
14587 };
14588 let (nc, ep) = (ncols as i32, eps);
14589 let __s_b2 = self.gpu.stream();
14590 let mut b2 = __s_b2.launch_builder(&f);
14591 b2.arg(a)
14592 .arg(wa)
14593 .arg(b)
14594 .arg(&mut *res)
14595 .arg(&mut out_q)
14596 .arg(&mut out_d)
14597 .arg(&nc)
14598 .arg(&ep);
14599 unsafe {
14600 b2.launch(cfg)?;
14601 }
14602 Ok((out_q, out_d))
14603 }
14604
14605 pub fn l2_v2_on(ncols: usize) -> bool {
14609 ncols == 128 && std::env::var("MEMRA_L2_V2").as_deref() != Ok("0")
14610 }
14611
14612 pub fn l2_norm_pp(
14613 &self,
14614 x: &CudaSlice<f32>,
14615 dst: &mut CudaSlice<f32>,
14616 dst16: Option<&mut CudaSlice<u8>>,
14617 ncols: usize,
14618 nrows: usize,
14619 eps: f32,
14620 ) -> Result<(), Box<dyn std::error::Error>> {
14621 if Self::l2_v2_on(ncols) {
14622 let f = self.func("l2_norm_pp_v2_f32");
14623 let rows_per_block = 8u32; let cfg = LaunchConfig {
14625 grid_dim: ((nrows as u32).div_ceil(rows_per_block), 1, 1),
14626 block_dim: (256, 1, 1),
14627 shared_mem_bytes: 0,
14628 };
14629 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
14630 let d16: u64 = match dst16 {
14632 Some(d) => self.addr_u8(d),
14633 None => 0,
14634 };
14635 let __s_b = self.gpu.stream();
14636 let mut b = __s_b.launch_builder(&f);
14637 b.arg(x).arg(dst).arg(&d16).arg(&nc).arg(&nr).arg(&e);
14638 unsafe {
14639 b.launch(cfg)?;
14640 }
14641 return Ok(());
14642 }
14643 self.l2_norm(x, dst, ncols, nrows, eps)
14644 }
14645
14646 pub fn l2_norm(
14647 &self,
14648 x: &CudaSlice<f32>,
14649 dst: &mut CudaSlice<f32>,
14650 ncols: usize,
14651 nrows: usize,
14652 eps: f32,
14653 ) -> Result<(), Box<dyn std::error::Error>> {
14654 let f = self.func("l2_norm_f32");
14655 let cfg = LaunchConfig {
14656 grid_dim: (nrows as u32, 1, 1),
14657 block_dim: (256, 1, 1),
14658 shared_mem_bytes: 0,
14659 };
14660 let (nc, e) = (ncols as i32, eps);
14661 let __s_b = self.gpu.stream();
14662 let mut b = __s_b.launch_builder(&f);
14663 b.arg(x).arg(dst).arg(&nc).arg(&e);
14664 unsafe {
14665 b.launch(cfg)?;
14666 }
14667 Ok(())
14668 }
14669
14670 pub fn l2_norm_decode(
14676 &self,
14677 x: &CudaSlice<f32>,
14678 dst: &mut CudaSlice<f32>,
14679 ncols: usize,
14680 nrows: usize,
14681 eps: f32,
14682 ) -> Result<(), Box<dyn std::error::Error>> {
14683 let f = self.func("l2_norm_f32");
14684 let cfg = LaunchConfig {
14685 grid_dim: (nrows as u32, 1, 1),
14686 block_dim: (32, 1, 1),
14687 shared_mem_bytes: 0,
14688 };
14689 let (nc, e) = (ncols as i32, eps);
14690 let __s_b = self.gpu.stream();
14691 let mut b = __s_b.launch_builder(&f);
14692 b.arg(x).arg(dst).arg(&nc).arg(&e);
14693 unsafe {
14694 b.launch(cfg)?;
14695 }
14696 Ok(())
14697 }
14698
14699 #[allow(clippy::too_many_arguments)] pub fn rope_neox(
14702 &self,
14703 x: &mut CudaSlice<f32>,
14704 pos: &CudaSlice<i32>,
14705 head_dim: usize,
14706 n_dims: usize,
14707 n_heads: usize,
14708 n_tokens: usize,
14709 freq_base: f32,
14710 freq_scale: f32,
14711 ) -> Result<(), Box<dyn std::error::Error>> {
14712 let f = self.func("rope_neox_f32");
14713 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
14714 let grid = (n_heads * n_tokens) as u32;
14715 let cfg = LaunchConfig {
14716 grid_dim: (grid, 1, 1),
14717 block_dim: ((head_dim / 2) as u32, 1, 1),
14718 shared_mem_bytes: 0,
14719 };
14720 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
14721 let __s_b = self.gpu.stream();
14722 let mut b = __s_b.launch_builder(&f);
14723 b.arg(x)
14724 .arg(pos)
14725 .arg(&hd)
14726 .arg(&nd)
14727 .arg(&nh)
14728 .arg(&theta_scale)
14729 .arg(&freq_scale);
14730 unsafe {
14731 b.launch(cfg)?;
14732 }
14733 Ok(())
14734 }
14735
14736 #[allow(clippy::too_many_arguments)] pub fn rope_neox_ff(
14739 &self,
14740 x: &mut CudaSlice<f32>,
14741 pos: &CudaSlice<i32>,
14742 head_dim: usize,
14743 n_dims: usize,
14744 n_heads: usize,
14745 n_tokens: usize,
14746 freq_base: f32,
14747 freq_scale: f32,
14748 ff: &CudaSlice<f32>,
14749 ) -> Result<(), Box<dyn std::error::Error>> {
14750 let f = self.func("rope_neox_ff_f32");
14751 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
14752 let grid = (n_heads * n_tokens) as u32;
14753 let cfg = LaunchConfig {
14754 grid_dim: (grid, 1, 1),
14755 block_dim: ((head_dim / 2) as u32, 1, 1),
14756 shared_mem_bytes: 0,
14757 };
14758 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
14759 let __s_b = self.gpu.stream();
14760 let mut b = __s_b.launch_builder(&f);
14761 b.arg(x)
14762 .arg(pos)
14763 .arg(&hd)
14764 .arg(&nd)
14765 .arg(&nh)
14766 .arg(&theta_scale)
14767 .arg(&freq_scale)
14768 .arg(ff);
14769 unsafe {
14770 b.launch(cfg)?;
14771 }
14772 Ok(())
14773 }
14774
14775 #[allow(clippy::too_many_arguments)]
14779 pub fn rope_neox_ffm(
14780 &self,
14781 x: &mut CudaSlice<f32>,
14782 pos: &CudaSlice<i32>,
14783 head_dim: usize,
14784 n_dims: usize,
14785 n_heads: usize,
14786 n_tokens: usize,
14787 freq_base: f32,
14788 freq_scale: f32,
14789 ff: &CudaSlice<f32>,
14790 mscale: f32,
14791 ) -> Result<(), Box<dyn std::error::Error>> {
14792 let f = self.func("rope_neox_ffm_f32");
14793 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
14794 let grid = (n_heads * n_tokens) as u32;
14795 let cfg = LaunchConfig {
14796 grid_dim: (grid, 1, 1),
14797 block_dim: ((head_dim / 2) as u32, 1, 1),
14798 shared_mem_bytes: 0,
14799 };
14800 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
14801 let __s_b = self.gpu.stream();
14802 let mut b = __s_b.launch_builder(&f);
14803 b.arg(x)
14804 .arg(pos)
14805 .arg(&hd)
14806 .arg(&nd)
14807 .arg(&nh)
14808 .arg(&theta_scale)
14809 .arg(&freq_scale)
14810 .arg(ff)
14811 .arg(&mscale);
14812 unsafe {
14813 b.launch(cfg)?;
14814 }
14815 Ok(())
14816 }
14817
14818 #[allow(clippy::too_many_arguments)]
14820 pub fn rope_neox2(
14821 &self,
14822 q: &mut CudaSlice<f32>,
14823 k: &mut CudaSlice<f32>,
14824 pos: &CudaSlice<i32>,
14825 head_dim: usize,
14826 n_dims: usize,
14827 nh_q: usize,
14828 nh_k: usize,
14829 n_tokens: usize,
14830 freq_base: f32,
14831 freq_scale: f32,
14832 ff: Option<&CudaSlice<f32>>,
14833 ) -> Result<(), Box<dyn std::error::Error>> {
14834 let f = self.func("rope_neox2_f32");
14835 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
14836 let grid = ((nh_q + nh_k) * n_tokens) as u32;
14837 let cfg = LaunchConfig {
14838 grid_dim: (grid, 1, 1),
14839 block_dim: ((head_dim / 2) as u32, 1, 1),
14840 shared_mem_bytes: 0,
14841 };
14842 let (hd, nd, nq, nk, nt) = (
14843 head_dim as i32,
14844 n_dims as i32,
14845 nh_q as i32,
14846 nh_k as i32,
14847 n_tokens as i32,
14848 );
14849 let __s_b = self.gpu.stream();
14850 let mut b = __s_b.launch_builder(&f);
14851 b.arg(q)
14852 .arg(k)
14853 .arg(pos)
14854 .arg(&hd)
14855 .arg(&nd)
14856 .arg(&nq)
14857 .arg(&nk)
14858 .arg(&nt)
14859 .arg(&theta_scale)
14860 .arg(&freq_scale);
14861 match ff {
14862 Some(ffv) => {
14863 b.arg(ffv);
14864 unsafe {
14865 b.launch(cfg)?;
14866 }
14867 }
14868 None => {
14869 let null: u64 = 0;
14870 b.arg(&null);
14871 unsafe {
14872 b.launch(cfg)?;
14873 }
14874 }
14875 }
14876 Ok(())
14877 }
14878
14879 pub fn gelu_tanh_mul(
14881 &self,
14882 gate: &CudaSlice<f32>,
14883 up: &CudaSlice<f32>,
14884 dst: &mut CudaSlice<f32>,
14885 n: usize,
14886 ) -> Result<(), Box<dyn std::error::Error>> {
14887 let f = self.func("gelu_tanh_mul_f32");
14888 let cfg = LaunchConfig::for_num_elems(n as u32);
14889 let ni = n as i32;
14890 let __s_b = self.gpu.stream();
14891 let mut b = __s_b.launch_builder(&f);
14892 b.arg(gate).arg(up).arg(dst).arg(&ni);
14893 unsafe {
14894 b.launch(cfg)?;
14895 }
14896 Ok(())
14897 }
14898
14899 pub fn silu_mul(
14900 &self,
14901 gate: &CudaSlice<f32>,
14902 up: &CudaSlice<f32>,
14903 dst: &mut CudaSlice<f32>,
14904 n: usize,
14905 ) -> Result<(), Box<dyn std::error::Error>> {
14906 let f = self.func("silu_mul_f32");
14907 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
14909 let ni = n as i32;
14910 let __s_b = self.gpu.stream();
14911 let mut b = __s_b.launch_builder(&f);
14912 b.arg(gate).arg(up).arg(dst).arg(&ni);
14913 unsafe {
14914 b.launch(cfg)?;
14915 }
14916 Ok(())
14917 }
14918
14919 pub fn silu_mul_host_expf(
14921 &self,
14922 gate: &CudaSlice<f32>,
14923 up: &CudaSlice<f32>,
14924 dst: &mut CudaSlice<f32>,
14925 n: usize,
14926 ) -> Result<(), Box<dyn std::error::Error>> {
14927 let f = self.func("silu_mul_host_expf_f32");
14928 let cfg = LaunchConfig::for_num_elems(n as u32);
14929 let ni = n as i32;
14930 let __s_b = self.gpu.stream();
14931 let mut b = __s_b.launch_builder(&f);
14932 b.arg(gate).arg(up).arg(dst).arg(&ni);
14933 unsafe {
14934 b.launch(cfg)?;
14935 }
14936 Ok(())
14937 }
14938
14939 pub fn silu_clamped_mul_host_expf(
14941 &self,
14942 gate: &CudaSlice<f32>,
14943 up: &CudaSlice<f32>,
14944 limit: f32,
14945 dst: &mut CudaSlice<f32>,
14946 n: usize,
14947 ) -> Result<(), Box<dyn std::error::Error>> {
14948 if !limit.is_finite() || limit <= 0.0 {
14949 return Err(
14950 format!("Step routed-expert clamp limit must be positive, got {limit}").into(),
14951 );
14952 }
14953 let f = self.func("silu_clamped_mul_host_expf_f32");
14954 let cfg = LaunchConfig::for_num_elems(n as u32);
14955 let ni = n as i32;
14956 let __s_b = self.gpu.stream();
14957 let mut b = __s_b.launch_builder(&f);
14958 b.arg(gate).arg(up).arg(&limit).arg(dst).arg(&ni);
14959 unsafe {
14960 b.launch(cfg)?;
14961 }
14962 Ok(())
14963 }
14964
14965 pub fn silu_mul_f16out(
14968 &self,
14969 gate: &CudaSlice<f32>,
14970 up: &CudaSlice<f32>,
14971 dst: &mut CudaSlice<f32>,
14972 dst16: &mut CudaSlice<u8>,
14973 n: usize,
14974 ) -> Result<(), Box<dyn std::error::Error>> {
14975 let f = self.func("silu_mul_f16out_f32");
14976 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
14977 let ni = n as i32;
14978 let __s_b = self.gpu.stream();
14979 let mut b = __s_b.launch_builder(&f);
14980 b.arg(gate).arg(up).arg(dst).arg(dst16).arg(&ni);
14981 unsafe {
14982 b.launch(cfg)?;
14983 }
14984 Ok(())
14985 }
14986
14987 pub fn silu_mul_scaled(
14994 &self,
14995 gate: &CudaSlice<f32>,
14996 up: &CudaSlice<f32>,
14997 gs: f32,
14998 us: f32,
14999 dst: &mut CudaSlice<f32>,
15000 n: usize,
15001 ) -> Result<(), Box<dyn std::error::Error>> {
15002 let f = self.func("silu_mul_scaled_f32");
15003 let cfg = LaunchConfig::for_num_elems(n as u32);
15004 let ni = n as i32;
15005 let (gsf, usf) = (gs, us);
15006 let __s_b = self.gpu.stream();
15007 let mut b = __s_b.launch_builder(&f);
15008 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(dst).arg(&ni);
15009 unsafe {
15010 b.launch(cfg)?;
15011 }
15012 Ok(())
15013 }
15014
15015 #[allow(clippy::too_many_arguments)]
15019 pub fn swigluoai_mul_scaled(
15020 &self,
15021 gate: &CudaSlice<f32>,
15022 up: &CudaSlice<f32>,
15023 gs: f32,
15024 us: f32,
15025 alpha: f32,
15026 limit: f32,
15027 dst: &mut CudaSlice<f32>,
15028 n: usize,
15029 ) -> Result<(), Box<dyn std::error::Error>> {
15030 let f = self.func("swigluoai_mul_scaled_f32");
15031 let cfg = LaunchConfig::for_num_elems(n as u32);
15032 let ni = n as i32;
15033 let __s_b = self.gpu.stream();
15034 let mut b = __s_b.launch_builder(&f);
15035 b.arg(gate)
15036 .arg(up)
15037 .arg(&gs)
15038 .arg(&us)
15039 .arg(&alpha)
15040 .arg(&limit)
15041 .arg(dst)
15042 .arg(&ni);
15043 unsafe {
15044 b.launch(cfg)?;
15045 }
15046 Ok(())
15047 }
15048
15049 pub fn silu_mul_scaled_q8_1(
15057 &self,
15058 gate: &CudaSlice<f32>,
15059 up: &CudaSlice<f32>,
15060 gs: f32,
15061 us: f32,
15062 n: usize,
15063 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15064 let f = self.func("silu_mul_scaled_q8_1");
15065 let nblk = n / 32;
15066 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);
15070 let (gsf, usf, ni) = (gs, us, n as i32);
15071 let __s_b = self.gpu.stream();
15072 let mut b = __s_b.launch_builder(&f);
15073 b.arg(gate)
15074 .arg(up)
15075 .arg(&gsf)
15076 .arg(&usf)
15077 .arg(&mut aq)
15078 .arg(&mut ad)
15079 .arg(&ni);
15080 unsafe {
15081 b.launch(cfg)?;
15082 }
15083 Ok((aq, ad))
15084 }
15085
15086 pub fn add(
15087 &self,
15088 a: &CudaSlice<f32>,
15089 b_in: &CudaSlice<f32>,
15090 dst: &mut CudaSlice<f32>,
15091 n: usize,
15092 ) -> Result<(), Box<dyn std::error::Error>> {
15093 let f = self.func("add_f32");
15094 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
15096 let ni = n as i32;
15097 let __s_bld = self.gpu.stream();
15098 let mut bld = __s_bld.launch_builder(&f);
15099 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
15100 unsafe {
15101 bld.launch(cfg)?;
15102 }
15103 Ok(())
15104 }
15105
15106 pub fn mul(
15107 &self,
15108 a: &CudaSlice<f32>,
15109 b_in: &CudaSlice<f32>,
15110 dst: &mut CudaSlice<f32>,
15111 n: usize,
15112 ) -> Result<(), Box<dyn std::error::Error>> {
15113 let f = self.func("mul_f32");
15114 let cfg = LaunchConfig::for_num_elems(n as u32);
15115 let ni = n as i32;
15116 let __s_bld = self.gpu.stream();
15117 let mut bld = __s_bld.launch_builder(&f);
15118 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
15119 unsafe {
15120 bld.launch(cfg)?;
15121 }
15122 Ok(())
15123 }
15124
15125 pub fn matmul(
15128 &self,
15129 w: &crate::model::GpuTensor,
15130 x: &CudaSlice<f32>,
15131 m: usize,
15132 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15133 use crate::model::GpuTensor;
15134 let in_f = w.in_features();
15135 let out_f = w.out_features();
15136 #[allow(non_snake_case)]
15144 let GEMM_M_THRESHOLD = if self.verify_exact_on() {
15147 usize::MAX
15148 } else {
15149 16usize
15150 };
15151
15152 const GEMM_MIN_OUT_F: usize = 128; if m >= GEMM_M_THRESHOLD {
15177 if let Some(y) = self.try_fp8_gemm(w, x, m)? {
15178 return Ok(y);
15179 }
15180 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? {
15187 return Ok(y);
15188 }
15189 if let Some(y) = self.try_f16_gemm(w, x, m)? {
15192 return Ok(y);
15193 }
15194 }
15195 if let GpuTensor::Quant { qtype, .. } = w
15210 && *qtype == QT_F8_E4M3_BLK
15211 {
15212 if m >= GEMM_M_THRESHOLD
15213 && let Some(y) = self.try_e4m3_blk_prefill(w, x, m)?
15214 {
15215 return Ok(y);
15216 }
15217 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15218 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? {
15219 return Ok(y);
15220 }
15221 }
15222 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.mmq_supports(w) {
15223 return self.qmatvec_mmq(w, x, m);
15224 }
15225 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.gemm_supports(w) {
15226 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15227 return self.qmatvec_gemm(w, &aq, &ad, m);
15228 }
15229 if m >= GEMM_M_THRESHOLD
15232 && let Some(y) = self.try_fp4_gemm(w, x, m, in_f, out_f)?
15233 {
15234 return Ok(y);
15235 }
15236 let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
15240 if m == 1
15245 && fast
15246 && let GpuTensor::Quant {
15247 bytes,
15248 qtype,
15249 row_bytes,
15250 rp,
15251 rp4,
15252 scale,
15253 ..
15254 } = w
15255 && self.mmvq_supports(*qtype)
15256 {
15257 let (bytes, rp) = match rp4 {
15261 Some(m4) => (m4, true),
15262 None => (bytes, *rp),
15263 };
15264 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15265 return self.qmatvec_mmvq(
15266 bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, rp,
15267 );
15268 }
15269 if (2..=16).contains(&m)
15285 && fast
15286 && std::env::var("MEMRA_NO_BATCHED").is_err()
15287 && (m <= 4 || Self::b8_enabled())
15288 {
15289 let m_ok = m <= 8
15299 || matches!(w, GpuTensor::Quant { qtype, .. }
15300 if *qtype == QT_Q4_0 || *qtype == QT_Q6_K || *qtype == QT_F8_E4M3
15301 || *qtype == QT_NVFP4 || *qtype == QT_Q4_K || *qtype == QT_Q5_K || *qtype == QT_Q8_0);
15302 if m_ok
15303 && let GpuTensor::Quant {
15304 bytes,
15305 qtype,
15306 row_bytes,
15307 rp,
15308 rp4,
15309 ..
15310 } = w
15311 && self.batched_supports(*qtype)
15312 && self.mmvq_supports(*qtype)
15313 {
15314 let (bytes, rp) = match rp4 {
15315 Some(m4) => (m4, true),
15316 None => (bytes, *rp),
15317 };
15318 let mcols = Self::batched_mcols(m);
15319 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15320 let mut y = self.qmatvec_mmvq_batched(
15321 bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, mcols, 1.0, rp,
15322 )?;
15323 if let GpuTensor::Quant { scale, .. } = w
15324 && *scale != 1.0
15325 {
15326 self.scale_inplace(&mut y, *scale, m * out_f)?;
15327 }
15328 return Ok(y);
15329 }
15330 }
15331 if fast
15337 && let GpuTensor::Quant {
15338 bytes,
15339 qtype,
15340 row_bytes,
15341 scale,
15342 ..
15343 } = w
15344 && *qtype == QT_F8_E4M3
15345 {
15346 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15347 return self.qmatvec_mmvq(
15348 bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, false,
15349 );
15350 }
15351 let mut y = match w {
15352 GpuTensor::Quant {
15353 bytes,
15354 qtype,
15355 row_bytes,
15356 ..
15357 } if fast && *qtype == QT_Q8_0 => {
15358 self.qmatvec_q8_0_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15359 }
15360 GpuTensor::Quant {
15361 bytes,
15362 qtype,
15363 row_bytes,
15364 ..
15365 } if fast && *qtype == QT_Q4_K => {
15366 self.qmatvec_q4_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15367 }
15368 GpuTensor::Quant {
15369 bytes,
15370 qtype,
15371 row_bytes,
15372 ..
15373 } if fast && *qtype == QT_Q6_K => {
15374 self.qmatvec_q6_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15375 }
15376 GpuTensor::Quant {
15377 bytes,
15378 qtype,
15379 row_bytes,
15380 ..
15381 } if fast && *qtype == QT_Q5_K => {
15382 self.qmatvec_q5_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15383 }
15384 GpuTensor::Quant {
15385 bytes,
15386 qtype,
15387 row_bytes,
15388 ..
15389 } if fast && *qtype == QT_Q3_K => {
15390 self.qmatvec_q3_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15391 }
15392 GpuTensor::Quant {
15393 bytes,
15394 qtype,
15395 row_bytes,
15396 rp,
15397 ..
15398 } if fast && *qtype == QT_NVFP4 => self.qmatvec_dp4a_named(
15399 if *rp {
15400 "qmatvec_nvfp4_dp4a_rp"
15401 } else {
15402 "qmatvec_nvfp4_dp4a"
15403 },
15404 &bytes.slice(0..bytes.len()),
15405 x,
15406 m,
15407 in_f,
15408 out_f,
15409 *row_bytes,
15410 )?,
15411 GpuTensor::Quant {
15415 bytes,
15416 qtype,
15417 row_bytes,
15418 ..
15419 } if fast && *qtype == QT_IQ4_XS && Self::iq_fast_enabled() => {
15420 self.qmatvec_iq4_XS_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15421 }
15422 GpuTensor::Quant {
15427 bytes,
15428 qtype,
15429 row_bytes,
15430 rp,
15431 ..
15432 } =>
15433 {
15436 self.qmatvec(
15437 bytes,
15438 x,
15439 m,
15440 in_f,
15441 out_f,
15442 if *rp && *qtype == QT_NVFP4 {
15443 QT_NVFP4_RP
15444 } else {
15445 *qtype
15446 },
15447 *row_bytes,
15448 )?
15449 }
15450 GpuTensor::Float { data, .. } => self.linear(x, data, m, in_f, out_f)?,
15451 GpuTensor::FloatBf16 { data, .. } => {
15454 if (1..=32).contains(&m) && Self::bf16_mmv_on() && in_f.is_multiple_of(8) {
15461 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
15462 self.matvec_bf16_rows_into(data, x, &mut y, in_f, out_f, m)?;
15463 y
15464 } else {
15465 self.linear_bf16_chunked(x, data, m, in_f, out_f, false, None)?
15466 }
15467 }
15468 };
15469 if let GpuTensor::Quant { scale, .. } = w
15471 && *scale != 1.0
15472 {
15473 self.scale_inplace(&mut y, *scale, m * out_f)?;
15474 }
15475 Ok(y)
15476 }
15477
15478 pub fn stage_a_raw_needed() -> bool {
15488 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15489 *ON.get_or_init(|| std::env::var("MEMRA_FAST").as_deref() == Ok("0"))
15490 }
15491
15492 pub fn uses_q8_1_fast(&self, w: &crate::model::GpuTensor) -> bool {
15495 use crate::model::GpuTensor;
15496 if std::env::var("MEMRA_FAST").as_deref() == Ok("0") {
15497 return false;
15498 }
15499 match w {
15500 GpuTensor::Quant { qtype, .. } => {
15507 matches!(
15508 *qtype,
15509 QT_Q8_0
15510 | QT_Q4_K
15511 | QT_Q6_K
15512 | QT_Q5_K
15513 | QT_Q3_K
15514 | QT_NVFP4
15515 | QT_F8_E4M3
15516 | QT_F8_E4M3_BLK
15517 | QT_Q4_0
15518 ) || (*qtype == QT_IQ4_XS && Self::iq_fast_enabled())
15519 }
15520 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
15521 }
15522 }
15523
15524 pub fn matmul_pre(
15529 &self,
15530 w: &crate::model::GpuTensor,
15531 aq: &CudaSlice<i8>,
15532 ad: &CudaSlice<f32>,
15533 x_fallback: &CudaSlice<f32>,
15534 m: usize,
15535 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15536 use crate::model::GpuTensor;
15537 let x_raw_ok = x_fallback.len() >= m * w.in_features();
15543 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
15546 if let Some(y) = self.try_fp8_gemm(w, x_fallback, m)? {
15547 return Ok(y);
15548 }
15549 if let Some(y) = self.try_fp8_blk_mmq(w, x_fallback, m)? {
15552 return Ok(y);
15553 }
15554 if let Some(y) = self.try_f16_gemm(w, x_fallback, m)? {
15556 return Ok(y);
15557 }
15558 }
15559 if m >= 16
15565 && x_raw_ok
15566 && !self.verify_exact_on()
15567 && let Some(y) = self.try_e4m3_blk_prefill(w, x_fallback, m)?
15568 {
15569 return Ok(y);
15570 }
15571 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
15572 return Ok(y);
15573 }
15574 if m >= 16
15579 && w.out_features() >= 128
15580 && self.mmq_supports(w)
15581 && !self.verify_exact_on()
15582 && x_raw_ok
15583 {
15584 return self.qmatvec_mmq(w, x_fallback, m);
15585 }
15586 if m >= 16
15589 && x_raw_ok
15590 && !self.verify_exact_on()
15591 && let Some(y) =
15592 self.try_fp4_gemm(w, x_fallback, m, w.in_features(), w.out_features())?
15593 {
15594 return Ok(y);
15595 }
15596 if m >= 16 && self.gemm_supports(w) && !self.verify_exact_on() {
15599 return self.qmatvec_gemm(w, aq, ad, m);
15600 }
15601 if !self.uses_q8_1_fast(w) {
15620 if !x_raw_ok {
15621 return Err(format!(
15622 "matmul_pre: q8_1-fast is off for this weight but x_fallback holds {} f32 \
15623 (need m*in_f = {}*{} = {}). This call site pre-quantized its activation and \
15624 dropped the f32, so there is nothing to fall back to — pass the real f32 \
15625 activation (see Engine::rms_norm_decode, which is bit-identical to \
15626 rms_norm_q8_1's reduction) or keep the weight on the q8_1 path.",
15627 x_fallback.len(),
15628 m,
15629 w.in_features(),
15630 m * w.in_features()
15631 )
15632 .into());
15633 }
15634 return self.matmul(w, x_fallback, m);
15635 }
15636 let in_f = w.in_features();
15637 let out_f = w.out_features();
15638 let (bytes, qtype, row_bytes, scale, rp) = match w {
15639 GpuTensor::Quant {
15640 bytes,
15641 qtype,
15642 row_bytes,
15643 scale,
15644 rp,
15645 ..
15646 } => (bytes, *qtype, *row_bytes, *scale, *rp),
15647 _ => unreachable!("uses_q8_1_fast guaranteed Quant"),
15648 };
15649 let (mbytes, mrp) = match w {
15652 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
15653 _ => (bytes, rp),
15654 };
15655 if m == 1 && self.mmvq_supports(qtype) {
15659 return self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, mrp);
15660 }
15661 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
15674 && std::env::var("MEMRA_NO_BATCHED").is_err()
15675 && (m <= 4 || Self::b8_enabled())
15676 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_NVFP4
15680 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_F8_E4M3 || qtype == QT_Q8_0)
15681 {
15682 let mcols = Self::batched_mcols(m);
15683 return self.qmatvec_mmvq_batched(
15684 mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, mrp,
15685 );
15686 }
15687 if qtype == QT_F8_E4M3 || qtype == QT_Q4_0 {
15693 let (b2, r2) = if qtype == QT_Q4_0 {
15694 (mbytes, mrp)
15695 } else {
15696 (bytes, rp)
15697 };
15698 return self.qmatvec_mmvq(b2, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, r2);
15699 }
15700 let name = match qtype {
15701 QT_Q8_0 => "qmatvec_q8_0_dp4a",
15702 QT_Q4_K => "qmatvec_q4_K_dp4a",
15703 QT_Q6_K => "qmatvec_q6_K_dp4a",
15704 QT_Q5_K => "qmatvec_q5_K_dp4a",
15705 QT_Q3_K => "qmatvec_q3_K_dp4a",
15706 QT_NVFP4 => {
15707 if rp {
15708 "qmatvec_nvfp4_dp4a_rp"
15709 } else {
15710 "qmatvec_nvfp4_dp4a"
15711 }
15712 }
15713 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
15714 _ => unreachable!(),
15715 };
15716 let f = self.func(name);
15717 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
15719 grid_dim: (out_f as u32, m as u32, 1),
15720 block_dim: (128, 1, 1),
15721 shared_mem_bytes: 0,
15722 };
15723 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
15724 let __s_b = self.gpu.stream();
15725 let mut b = __s_b.launch_builder(&f);
15726 b.arg(bytes)
15727 .arg(aq)
15728 .arg(ad)
15729 .arg(&mut y)
15730 .arg(&inf)
15731 .arg(&outf)
15732 .arg(&mi)
15733 .arg(&rb);
15734 unsafe {
15735 b.launch(cfg)?;
15736 }
15737 if scale != 1.0 {
15738 self.scale_inplace(&mut y, scale, m * out_f)?;
15739 }
15740 Ok(y)
15741 }
15742
15743 pub fn matmul_decode_exact(
15751 &self,
15752 w: &crate::model::GpuTensor,
15753 x: &CudaSlice<f32>,
15754 m: usize,
15755 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15756 use crate::model::GpuTensor;
15757 if let GpuTensor::Float { data, .. } = w {
15765 return self.linear_decode_exact(x, data, m, w.in_features(), w.out_features());
15766 }
15767 if let GpuTensor::FloatBf16 { data, .. } = w {
15770 let (in_f, out_f) = (w.in_features(), w.out_features());
15771 if (1..=32).contains(&m) && Self::bf16_mmv_on() && in_f % 8 == 0 {
15774 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
15775 self.matvec_bf16_rows_into(data, x, &mut y, in_f, out_f, m)?;
15776 return Ok(y);
15777 }
15778 return self.linear_bf16_chunked(x, data, m, in_f, out_f, true, None);
15779 }
15780 if !self.uses_q8_1_fast(w) {
15781 return self.matmul(w, x, m);
15782 }
15783 let in_f = w.in_features();
15784 let out_f = w.out_features();
15785 let (bytes, qtype, row_bytes, scale, rp) = match w {
15786 GpuTensor::Quant {
15787 bytes,
15788 qtype,
15789 row_bytes,
15790 scale,
15791 rp,
15792 ..
15793 } => (bytes, *qtype, *row_bytes, *scale, *rp),
15794 _ => return self.matmul(w, x, m),
15795 };
15796 let (bytes, rp) = match w {
15799 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
15800 _ => (bytes, rp),
15801 };
15802 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15803 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? {
15807 return Ok(y);
15808 }
15809 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
15818 && std::env::var("MEMRA_NO_BATCHED").is_err()
15819 && (m <= 4 || Self::b8_enabled())
15820 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
15823 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0)
15824 {
15825 let mcols = Self::batched_mcols(m);
15826 return self.qmatvec_mmvq_batched(
15827 bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp,
15828 );
15829 }
15830 if self.mmvq_supports(qtype) {
15831 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
15834 }
15835 self.matmul_pre(w, &aq, &ad, x, m)
15838 }
15839
15840 pub fn matmul_decode_exact_pre(
15850 &self,
15851 w: &crate::model::GpuTensor,
15852 aq: &CudaSlice<i8>,
15853 ad: &CudaSlice<f32>,
15854 m: usize,
15855 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15856 use crate::model::GpuTensor;
15857 debug_assert!(
15858 self.uses_q8_1_fast(w),
15859 "matmul_decode_exact_pre: caller must guarantee q8_1-fast"
15860 );
15861 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
15863 return Ok(y);
15864 }
15865 let in_f = w.in_features();
15866 let out_f = w.out_features();
15867 let (bytes, qtype, row_bytes, scale, rp) = match w {
15868 GpuTensor::Quant {
15869 bytes,
15870 qtype,
15871 row_bytes,
15872 scale,
15873 rp,
15874 ..
15875 } => (bytes, *qtype, *row_bytes, *scale, *rp),
15876 _ => {
15877 return Err(
15878 "matmul_decode_exact_pre: Quant tensor required (q8_1-fast contract)".into(),
15879 );
15880 }
15881 };
15882 let (bytes, rp) = match w {
15884 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
15885 _ => (bytes, rp),
15886 };
15887 if (2..=16).contains(&m)
15889 && self.batched_supports(qtype)
15890 && self.mmvq_supports(qtype)
15891 && std::env::var("MEMRA_NO_BATCHED").is_err()
15892 && (m <= 4 || Self::b8_enabled())
15893 && (m <= 8
15894 || qtype == QT_Q4_0
15895 || qtype == QT_Q6_K
15896 || qtype == QT_F8_E4M3
15897 || qtype == QT_NVFP4
15898 || qtype == QT_Q4_K
15899 || qtype == QT_Q5_K
15900 || qtype == QT_Q8_0)
15901 {
15902 let mcols = Self::batched_mcols(m);
15903 return self.qmatvec_mmvq_batched(
15904 bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp,
15905 );
15906 }
15907 if self.mmvq_supports(qtype) {
15908 return self.qmatvec_mmvq(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
15909 }
15910 let x0 = self.zeros(0)?;
15913 self.matmul_pre(w, aq, ad, &x0, m)
15914 }
15915
15916 #[allow(clippy::type_complexity)] pub fn matmul_decode_exact_dual_pre(
15926 &self,
15927 w0: &crate::model::GpuTensor,
15928 w1: &crate::model::GpuTensor,
15929 aq: &CudaSlice<i8>,
15930 ad: &CudaSlice<f32>,
15931 m: usize,
15932 ) -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>>
15933 {
15934 use crate::model::GpuTensor;
15935 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15936 let on = *ON.get_or_init(|| {
15937 std::env::var("MEMRA_SPEC_DUAL_T")
15938 .map(|v| v != "0")
15939 .unwrap_or(true)
15940 });
15941 if !on
15942 || !(2..=7).contains(&m)
15943 || std::env::var("MEMRA_NO_BATCHED").is_ok()
15944 || !self.uses_q8_1_fast(w0)
15945 || !self.uses_q8_1_fast(w1)
15946 {
15947 return Ok(None);
15948 }
15949 if !self.mmvq_supports(QT_NVFP4) {
15954 return Ok(None);
15955 }
15956 let (in_f, out_f) = (w0.in_features(), w0.out_features());
15957 if w1.in_features() != in_f || w1.out_features() != out_f {
15958 return Ok(None);
15959 }
15960 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
15961 (
15962 GpuTensor::Quant {
15963 bytes: b0,
15964 qtype: q0,
15965 row_bytes: rb0,
15966 scale: s0,
15967 rp: rp0,
15968 rp4: None,
15969 ..
15970 },
15971 GpuTensor::Quant {
15972 bytes: b1,
15973 qtype: q1,
15974 row_bytes: rb1,
15975 scale: s1,
15976 rp: rp1,
15977 rp4: None,
15978 ..
15979 },
15980 ) if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 => {
15981 (b0, b1, *rb0, *s0, *s1, *rp0)
15982 }
15983 _ => return Ok(None),
15984 };
15985 if m > 4 && !(rp && Self::b8_enabled() && std::env::var("MEMRA_B567").as_deref() != Ok("0"))
15988 {
15989 return Ok(None);
15990 }
15991 let (y0, y1) =
15992 self.qmatvec_batched_dual_raw(b0, b1, aq, ad, m, in_f, out_f, row_bytes, rp)?;
15993 Ok(Some(((y0, s0), (y1, s1))))
15994 }
15995
15996 pub fn matmul_decode_exact_group4_pre(
16008 &self,
16009 ws: [&crate::model::GpuTensor; 4],
16010 aq: &CudaSlice<i8>,
16011 ad: &CudaSlice<f32>,
16012 m: usize,
16013 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
16014 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16015 let on = *ON.get_or_init(|| {
16016 std::env::var("MEMRA_TK_GDN_GROUP")
16017 .map(|v| v != "0")
16018 .unwrap_or(true)
16019 });
16020 self.matmul_decode_exact_group_pre(&ws, aq, ad, m, on, "GDN group4")
16021 }
16022
16023 pub fn matmul_decode_exact_group3_pre(
16028 &self,
16029 ws: [&crate::model::GpuTensor; 3],
16030 aq: &CudaSlice<i8>,
16031 ad: &CudaSlice<f32>,
16032 m: usize,
16033 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
16034 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16035 let on = *ON.get_or_init(|| {
16036 std::env::var("MEMRA_TK_FA_GROUP")
16037 .map(|v| v != "0")
16038 .unwrap_or(true)
16039 });
16040 self.matmul_decode_exact_group_pre(&ws, aq, ad, m, on, "FA group3")
16041 }
16042
16043 #[allow(clippy::manual_div_ceil)] fn matmul_decode_exact_group_pre(
16048 &self,
16049 ws: &[&crate::model::GpuTensor],
16050 aq: &CudaSlice<i8>,
16051 ad: &CudaSlice<f32>,
16052 m: usize,
16053 on: bool,
16054 tag: &'static str,
16055 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
16056 use crate::model::GpuTensor;
16057 if !on
16058 || !(2..=16).contains(&m)
16059 || std::env::var("MEMRA_NO_BATCHED").is_ok()
16060 || (m > 4 && !Self::b8_enabled())
16061 || !self.mmvq_supports(QT_NVFP4)
16062 || !self.batched_supports(QT_NVFP4)
16063 {
16064 return Ok(None);
16065 }
16066 let in_f = ws[0].in_features();
16067 let mut parts: Vec<(&CudaSlice<u8>, usize, f32)> = Vec::with_capacity(4);
16068 for w in ws {
16069 if !self.uses_q8_1_fast(w) || w.in_features() != in_f {
16070 return Ok(None);
16071 }
16072 match w {
16073 GpuTensor::Quant {
16074 bytes,
16075 qtype,
16076 scale,
16077 rp: true,
16078 rp4: None,
16079 ..
16080 } if *qtype == QT_NVFP4 && w.out_features() % 8 == 0 => {
16081 parts.push((bytes, w.out_features(), *scale));
16082 }
16083 _ => return Ok(None),
16084 }
16085 }
16086 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16088 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
16089 let mcols = if (5..=7).contains(&m) && b567 {
16090 m
16091 } else {
16092 Self::batched_mcols(m)
16093 };
16094 let kname: &'static str = match mcols {
16095 2 => "qmatvec_nvfp4_mmvq_group4_b2_rp",
16096 4 => "qmatvec_nvfp4_mmvq_group4_b4_rp",
16097 5 => "qmatvec_nvfp4_mmvq_group4_b5_rp",
16098 6 => "qmatvec_nvfp4_mmvq_group4_b6_rp",
16099 7 => "qmatvec_nvfp4_mmvq_group4_b7_rp",
16100 8 => "qmatvec_nvfp4_mmvq_group4_b8_rp",
16101 16 => "qmatvec_nvfp4_mmvq_group4_b16_rp",
16102 _ => return Ok(None),
16103 };
16104 if std::env::var("MEMRA_DEBUG").is_ok() {
16107 use std::sync::Mutex;
16108 static SEEN: Mutex<Vec<&'static str>> = Mutex::new(Vec::new());
16109 let mut seen = SEEN.lock().unwrap();
16110 if !seen.contains(&tag) {
16111 seen.push(tag);
16112 eprintln!("[memra] {tag} batched ENGAGED (m={m})");
16113 }
16114 }
16115 const ROWS_PER_BLOCK: u32 = 4; let rows_per_block = ROWS_PER_BLOCK * 2; let total: usize = parts.iter().map(|p| p.1).sum();
16118 let three = parts.len() == 3;
16119 let mut y0 = self.alloc_uninit::<f32>(m * parts[0].1)?;
16120 let mut y1 = self.alloc_uninit::<f32>(m * parts[1].1)?;
16121 let mut y2 = self.alloc_uninit::<f32>(m * parts[2].1)?;
16122 let mut y3 = self.alloc_uninit::<f32>(if three { 1 } else { m * parts[3].1 })?;
16125 let cfg = LaunchConfig {
16126 grid_dim: ((total as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
16127 block_dim: (32, ROWS_PER_BLOCK, 1),
16128 shared_mem_bytes: 0,
16129 };
16130 let (inf, mi) = (in_f as i32, m as i32);
16131 let (n0, n1, n2) = (parts[0].1 as i32, parts[1].1 as i32, parts[2].1 as i32);
16132 let n3 = if three { 0i32 } else { parts[3].1 as i32 };
16133 let (s0, s1, s2) = (parts[0].2, parts[1].2, parts[2].2);
16134 let s3 = if three { 1.0f32 } else { parts[3].2 };
16135 let w3 = if three { parts[0].0 } else { parts[3].0 };
16136 let f = self.func(kname);
16137 let __s_b = self.gpu.stream();
16138 let mut b = __s_b.launch_builder(&f);
16139 b.arg(parts[0].0)
16140 .arg(parts[1].0)
16141 .arg(parts[2].0)
16142 .arg(w3)
16143 .arg(aq)
16144 .arg(ad)
16145 .arg(&mut y0)
16146 .arg(&mut y1)
16147 .arg(&mut y2)
16148 .arg(&mut y3)
16149 .arg(&inf)
16150 .arg(&n0)
16151 .arg(&n1)
16152 .arg(&n2)
16153 .arg(&n3)
16154 .arg(&mi)
16155 .arg(&s0)
16156 .arg(&s1)
16157 .arg(&s2)
16158 .arg(&s3);
16159 unsafe {
16160 b.launch(cfg)?;
16161 }
16162 Ok(Some(if three {
16163 vec![y0, y1, y2]
16164 } else {
16165 vec![y0, y1, y2, y3]
16166 }))
16167 }
16168
16169 #[allow(clippy::type_complexity)] pub fn matmul_decode_exact_dual(
16186 &self,
16187 w0: &crate::model::GpuTensor,
16188 w1: &crate::model::GpuTensor,
16189 x: &CudaSlice<f32>,
16190 m: usize,
16191 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
16192 use crate::model::GpuTensor;
16193 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16194 let on = *ON.get_or_init(|| {
16195 std::env::var("MEMRA_SPEC_DUAL_T")
16196 .map(|v| v != "0")
16197 .unwrap_or(true)
16198 });
16199 if !on
16200 || !(2..=4).contains(&m)
16201 || std::env::var("MEMRA_NO_BATCHED").is_ok()
16202 || !self.uses_q8_1_fast(w0)
16203 || !self.uses_q8_1_fast(w1)
16204 {
16205 return Ok(None);
16206 }
16207 if !self.mmvq_supports(QT_NVFP4) {
16212 return Ok(None);
16213 }
16214 let (in_f, out_f) = (w0.in_features(), w0.out_features());
16215 if w1.in_features() != in_f || w1.out_features() != out_f {
16216 return Ok(None);
16217 }
16218 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
16219 (
16220 GpuTensor::Quant {
16221 bytes: b0,
16222 qtype: q0,
16223 row_bytes: rb0,
16224 scale: s0,
16225 rp: rp0,
16226 rp4: None,
16227 ..
16228 },
16229 GpuTensor::Quant {
16230 bytes: b1,
16231 qtype: q1,
16232 row_bytes: rb1,
16233 scale: s1,
16234 rp: rp1,
16235 rp4: None,
16236 ..
16237 },
16238 ) if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 => {
16239 (b0, b1, *rb0, *s0, *s1, *rp0)
16240 }
16241 _ => return Ok(None),
16242 };
16243 if std::env::var("MEMRA_DEBUG").is_ok() {
16246 static ONCE: std::sync::Once = std::sync::Once::new();
16247 ONCE.call_once(|| eprintln!("[memra] dual gate+up batched ENGAGED (m={m} rp={rp})"));
16248 }
16249 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
16250 let (y0, y1) =
16251 self.qmatvec_batched_dual_raw(b0, b1, &aq, &ad, m, in_f, out_f, row_bytes, rp)?;
16252 let mut y0 = y0;
16253 let mut y1 = y1;
16254 if s0 != 1.0 {
16255 self.scale_inplace(&mut y0, s0, m * out_f)?;
16256 }
16257 if s1 != 1.0 {
16258 self.scale_inplace(&mut y1, s1, m * out_f)?;
16259 }
16260 Ok(Some((y0, y1)))
16261 }
16262
16263 #[allow(clippy::too_many_arguments)]
16268 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_batched_dual_raw(
16270 &self,
16271 b0: &CudaSlice<u8>,
16272 b1: &CudaSlice<u8>,
16273 aq: &CudaSlice<i8>,
16274 ad: &CudaSlice<f32>,
16275 m: usize,
16276 in_f: usize,
16277 out_f: usize,
16278 row_bytes: usize,
16279 rp: bool,
16280 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
16281 const ROWS_PER_BLOCK: u32 = 4;
16282 let mcols = Self::batched_mcols(m);
16283 let tiny_rp1 = rp
16286 && mcols == 4
16287 && out_f <= 128
16288 && std::env::var("MEMRA_NVFP4_AUX_DUAL").as_deref() != Ok("0");
16289 let (name, rows_per_block) = if tiny_rp1 {
16290 ("qmatvec_nvfp4_mmvq_dual_b4_rp", ROWS_PER_BLOCK)
16291 } else {
16292 match (mcols, rp, m) {
16293 (2, false, _) => ("qmatvec_nvfp4_mmvq_dual_b2", ROWS_PER_BLOCK),
16294 (4, false, _) => ("qmatvec_nvfp4_mmvq_dual_b4_r2", ROWS_PER_BLOCK * 2),
16295 (2, true, _) => ("qmatvec_nvfp4_mmvq_dual_b2_rp", ROWS_PER_BLOCK),
16296 (4, true, _) => ("qmatvec_nvfp4_mmvq_dual_b4_rpr2", ROWS_PER_BLOCK * 2),
16297 (8, true, 5) => ("qmatvec_nvfp4_mmvq_dual_b5_rpr2", ROWS_PER_BLOCK * 2),
16298 (8, true, 6) => ("qmatvec_nvfp4_mmvq_dual_b6_rpr2", ROWS_PER_BLOCK * 2),
16299 (8, true, 7) => ("qmatvec_nvfp4_mmvq_dual_b7_rpr2", ROWS_PER_BLOCK * 2),
16300 _ => {
16301 return Err(
16302 format!("qmatvec_batched_dual_raw: no dual kernel for m {m}").into(),
16303 );
16304 }
16305 }
16306 };
16307 let f = self.func(name);
16308 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
16309 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
16310 let cfg = LaunchConfig {
16311 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
16312 block_dim: (32, ROWS_PER_BLOCK, 1),
16313 shared_mem_bytes: 0,
16314 };
16315 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
16316 let __s_b = self.gpu.stream();
16317 let mut b = __s_b.launch_builder(&f);
16318 b.arg(b0)
16319 .arg(b1)
16320 .arg(aq)
16321 .arg(ad)
16322 .arg(&mut y0)
16323 .arg(&mut y1)
16324 .arg(&inf)
16325 .arg(&outf)
16326 .arg(&mi)
16327 .arg(&rb);
16328 unsafe {
16329 b.launch(cfg)?;
16330 }
16331 Ok((y0, y1))
16332 }
16333
16334 #[allow(clippy::type_complexity)] #[allow(clippy::manual_div_ceil)] pub fn matmul_pre_dual_noscale(
16348 &self,
16349 w0: &crate::model::GpuTensor,
16350 w1: &crate::model::GpuTensor,
16351 aq: &CudaSlice<i8>,
16352 ad: &CudaSlice<f32>,
16353 m: usize,
16354 ) -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>>
16355 {
16356 use crate::model::GpuTensor;
16357 if m != 1 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
16358 return Ok(None);
16359 }
16360 if !self.mmvq_supports(QT_NVFP4) {
16370 return Ok(None);
16371 }
16372 let (in_f, out_f) = (w0.in_features(), w0.out_features());
16373 if w1.in_features() != in_f || w1.out_features() != out_f {
16374 return Ok(None);
16375 }
16376 let no_mirror =
16389 |w: &crate::model::GpuTensor| !matches!(w, GpuTensor::Quant { rp4: Some(_), .. });
16390 if self.q8_ffn_fuse2_on()
16391 && no_mirror(w0)
16392 && no_mirror(w1)
16393 && let Some([p0, p1]) = self.q8_fused_params(&[w0, w1])
16394 {
16395 let (y0, y1) = self.q8_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2)?;
16396 return Ok(Some(((y0, 1.0), (y1, 1.0))));
16397 }
16398 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
16408 let (y0, y1) =
16409 self.e4m3_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2, 1.0, 1.0)?;
16410 return Ok(Some(((y0, p0.3), (y1, p1.3))));
16411 }
16412 let (b0, q0, rb0, s0, rp0) = match w0 {
16413 GpuTensor::Quant {
16414 bytes,
16415 qtype,
16416 row_bytes,
16417 scale,
16418 rp,
16419 ..
16420 } => (bytes, *qtype, *row_bytes, *scale, *rp),
16421 _ => return Ok(None),
16422 };
16423 let (b1, q1, rb1, s1, rp1) = match w1 {
16424 GpuTensor::Quant {
16425 bytes,
16426 qtype,
16427 row_bytes,
16428 scale,
16429 rp,
16430 ..
16431 } => (bytes, *qtype, *row_bytes, *scale, *rp),
16432 _ => return Ok(None),
16433 };
16434 if q0 != QT_NVFP4 || q1 != QT_NVFP4 || rb0 != rb1 || rp0 != rp1 {
16435 return Ok(None);
16436 }
16437 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16439 let rows_per_block = ROWS_PER_BLOCK * RPW;
16440 let f = self.func(if rp0 {
16441 "qmatvec_nvfp4_mmvq_dual_mr2_rp"
16442 } else {
16443 "qmatvec_nvfp4_mmvq_dual_mr2"
16444 });
16445 let mut y0 = self.alloc_uninit::<f32>(out_f)?;
16446 let mut y1 = self.alloc_uninit::<f32>(out_f)?;
16447 let cfg = LaunchConfig {
16448 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
16449 block_dim: (32, ROWS_PER_BLOCK, 1),
16450 shared_mem_bytes: 0,
16451 };
16452 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, rb0 as i64);
16453 let one = 1.0f32;
16456 let __s_b = self.gpu.stream();
16457 let mut b = __s_b.launch_builder(&f);
16458 b.arg(b0)
16459 .arg(b1)
16460 .arg(aq)
16461 .arg(ad)
16462 .arg(&mut y0)
16463 .arg(&mut y1)
16464 .arg(&inf)
16465 .arg(&outf)
16466 .arg(&mi)
16467 .arg(&rb)
16468 .arg(&one)
16469 .arg(&one);
16470 unsafe {
16471 b.launch(cfg)?;
16472 }
16473 Ok(Some(((y0, s0), (y1, s1))))
16474 }
16475
16476 #[allow(clippy::too_many_arguments)]
16484 #[allow(clippy::type_complexity)] pub fn matmul_nvfp4_fused3(
16486 &self,
16487 w0: &crate::model::GpuTensor,
16488 w1: &crate::model::GpuTensor,
16489 w2: &crate::model::GpuTensor,
16490 aq: &CudaSlice<i8>,
16491 ad: &CudaSlice<f32>,
16492 m: usize,
16493 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
16494 {
16495 use crate::model::GpuTensor;
16496 if !self.mmvq_supports(QT_NVFP4)
16503 || !self.uses_q8_1_fast(w0)
16504 || !self.uses_q8_1_fast(w1)
16505 || !self.uses_q8_1_fast(w2)
16506 {
16507 return Ok(None);
16508 }
16509 if (9..=16).contains(&m) {
16512 return Ok(
16513 match self.matmul_decode_exact_group3_pre([w0, w1, w2], aq, ad, m)? {
16514 Some(mut ys) => {
16515 let y2 = ys.pop().unwrap();
16516 let y1 = ys.pop().unwrap();
16517 let y0 = ys.pop().unwrap();
16518 Some((y0, y1, y2))
16519 }
16520 None => None,
16521 },
16522 );
16523 }
16524 if !(1..=8).contains(&m) {
16525 return Ok(None);
16526 }
16527 if m > 1 {
16528 let in_f = w0.in_features();
16529 if std::env::var("MEMRA_NVFP4_FUSED3B").as_deref() == Ok("0")
16530 || !self.batched_supports(QT_NVFP4)
16531 || std::env::var("MEMRA_NO_BATCHED").is_ok()
16532 || (m > 4 && !Self::b8_enabled())
16533 || !in_f.is_multiple_of(512)
16534 || in_f / 64 > 272
16535 {
16536 return Ok(None);
16537 }
16538 }
16539 let unpack = |w: &crate::model::GpuTensor| match w {
16540 GpuTensor::Quant {
16541 bytes,
16542 qtype,
16543 scale,
16544 rp,
16545 ..
16546 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
16547 _ => None,
16548 };
16549 let (Some(p0), Some(p1), Some(p2)) = (unpack(w0), unpack(w1), unpack(w2)) else {
16550 return Ok(None);
16551 };
16552 let in_f = w0.in_features();
16553 if w1.in_features() != in_f || w2.in_features() != in_f {
16554 return Ok(None);
16555 }
16556 let (o0, o1, o2) = (w0.out_features(), w1.out_features(), w2.out_features());
16557 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16559 let rows_pb = ROWS_PER_BLOCK * RPW;
16560 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
16561 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
16562 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
16563 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
16564 let (inf, oi0, oi1, oi2, mi) = (in_f as i32, o0 as i32, o1 as i32, o2 as i32, m as i32);
16565 let (b0, b1, b2) = unsafe { (&*p0.0, &*p1.0, &*p2.0) };
16568 if m > 1 {
16569 if p0.1 != 1.0 || p1.1 != 1.0 || p2.1 != 1.0 {
16571 return Ok(None);
16572 }
16573 let f = self.func("qmatvec_nvfp4_mmvq_fused3_b8_rpsc");
16574 let cfg = LaunchConfig {
16575 grid_dim: (nb(o0) + nb(o1) + nb(o2), 1, 1),
16576 block_dim: (32, ROWS_PER_BLOCK, 1),
16577 shared_mem_bytes: 0,
16578 };
16579 let __s_b = self.gpu.stream();
16580 let mut b = __s_b.launch_builder(&f);
16581 b.arg(b0)
16582 .arg(b1)
16583 .arg(b2)
16584 .arg(aq)
16585 .arg(ad)
16586 .arg(&mut y0)
16587 .arg(&mut y1)
16588 .arg(&mut y2)
16589 .arg(&inf)
16590 .arg(&oi0)
16591 .arg(&oi1)
16592 .arg(&oi2)
16593 .arg(&mi);
16594 unsafe {
16595 b.launch(cfg)?;
16596 }
16597 return Ok(Some((y0, y1, y2)));
16598 }
16599 let f = self.func("qmatvec_nvfp4_mmvq_fused3_rp");
16600 let cfg = LaunchConfig {
16601 grid_dim: (nb(o0) + nb(o1) + nb(o2), m as u32, 1),
16602 block_dim: (32, ROWS_PER_BLOCK, 1),
16603 shared_mem_bytes: 0,
16604 };
16605 let __s_b = self.gpu.stream();
16606 let mut b = __s_b.launch_builder(&f);
16607 b.arg(b0)
16608 .arg(b1)
16609 .arg(b2)
16610 .arg(aq)
16611 .arg(ad)
16612 .arg(&mut y0)
16613 .arg(&mut y1)
16614 .arg(&mut y2)
16615 .arg(&inf)
16616 .arg(&oi0)
16617 .arg(&oi1)
16618 .arg(&oi2)
16619 .arg(&mi)
16620 .arg(&p0.1)
16621 .arg(&p1.1)
16622 .arg(&p2.1);
16623 unsafe {
16624 b.launch(cfg)?;
16625 }
16626 Ok(Some((y0, y1, y2)))
16627 }
16628
16629 #[allow(clippy::type_complexity)] pub fn matmul_nvfp4_fused2(
16639 &self,
16640 w0: &crate::model::GpuTensor,
16641 w1: &crate::model::GpuTensor,
16642 aq: &CudaSlice<i8>,
16643 ad: &CudaSlice<f32>,
16644 m: usize,
16645 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
16646 use crate::model::GpuTensor;
16647 static FUSED2_OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16648 let off =
16649 *FUSED2_OFF.get_or_init(|| std::env::var("MEMRA_NVFP4_FUSED2").as_deref() == Ok("0"));
16650 if off
16653 || m != 1
16654 || !self.mmvq_supports(QT_NVFP4)
16655 || !self.uses_q8_1_fast(w0)
16656 || !self.uses_q8_1_fast(w1)
16657 {
16658 return Ok(None);
16659 }
16660 let unpack = |w: &crate::model::GpuTensor| match w {
16661 GpuTensor::Quant {
16662 bytes,
16663 qtype,
16664 scale,
16665 rp,
16666 ..
16667 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
16668 _ => None,
16669 };
16670 let (Some(p0), Some(p1)) = (unpack(w0), unpack(w1)) else {
16671 return Ok(None);
16672 };
16673 let in_f = w0.in_features();
16674 if w1.in_features() != in_f {
16675 return Ok(None);
16676 }
16677 let (o0, o1) = (w0.out_features(), w1.out_features());
16678 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16680 let rows_pb = ROWS_PER_BLOCK * RPW;
16681 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
16682 let f = self.func("qmatvec_nvfp4_mmvq_fused2_rp");
16683 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
16684 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
16685 let cfg = LaunchConfig {
16686 grid_dim: (nb(o0) + nb(o1), m as u32, 1),
16687 block_dim: (32, ROWS_PER_BLOCK, 1),
16688 shared_mem_bytes: 0,
16689 };
16690 let (inf, oi0, oi1, mi) = (in_f as i32, o0 as i32, o1 as i32, m as i32);
16691 let (b0, b1) = unsafe { (&*p0.0, &*p1.0) };
16694 if Self::pdl_on() && Self::pdl_mmvq_on() && Self::pdl_nvfp4q8_on() {
16697 {
16698 use cudarc::driver::{DevicePtr, DevicePtrMut};
16699 let s = &self.gpu.stream();
16700 let (pw0, _g0) = b0.device_ptr(s);
16701 let (pw1, _g1) = b1.device_ptr(s);
16702 let (paq, _g2) = aq.device_ptr(s);
16703 let (pad, _g3) = ad.device_ptr(s);
16704 let (py0, _g4) = y0.device_ptr_mut(s);
16705 let (py1, _g5) = y1.device_ptr_mut(s);
16706 let (s0, s1) = (p0.1, p1.1);
16707 let mut ps = [
16708 &pw0 as *const _ as *mut std::ffi::c_void,
16709 &pw1 as *const _ as *mut _,
16710 &paq as *const _ as *mut _,
16711 &pad as *const _ as *mut _,
16712 &py0 as *const _ as *mut _,
16713 &py1 as *const _ as *mut _,
16714 &inf as *const _ as *mut _,
16715 &oi0 as *const _ as *mut _,
16716 &oi1 as *const _ as *mut _,
16717 &mi as *const _ as *mut _,
16718 &s0 as *const _ as *mut _,
16719 &s1 as *const _ as *mut _,
16720 ];
16721 unsafe {
16722 self.launch_pdl(
16723 "qmatvec_nvfp4_mmvq_fused2_rp",
16724 cfg.grid_dim,
16725 cfg.block_dim,
16726 &mut ps,
16727 )?;
16728 }
16729 }
16730 return Ok(Some((y0, y1)));
16731 }
16732 let __s_b = self.gpu.stream();
16733 let mut b = __s_b.launch_builder(&f);
16734 b.arg(b0)
16735 .arg(b1)
16736 .arg(aq)
16737 .arg(ad)
16738 .arg(&mut y0)
16739 .arg(&mut y1)
16740 .arg(&inf)
16741 .arg(&oi0)
16742 .arg(&oi1)
16743 .arg(&mi)
16744 .arg(&p0.1)
16745 .arg(&p1.1);
16746 unsafe {
16747 b.launch(cfg)?;
16748 }
16749 Ok(Some((y0, y1)))
16750 }
16751
16752 pub fn matmul_nvfp4_fused2_into(
16757 &self,
16758 w0: &crate::model::GpuTensor,
16759 w1: &crate::model::GpuTensor,
16760 aq: &CudaSlice<i8>,
16761 ad: &CudaSlice<f32>,
16762 y0: &mut CudaSlice<f32>,
16763 y1: &mut CudaSlice<f32>,
16764 ) -> Result<bool, Box<dyn std::error::Error>> {
16765 use crate::model::GpuTensor;
16766 static FUSED2_OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16767 let off =
16768 *FUSED2_OFF.get_or_init(|| std::env::var("MEMRA_NVFP4_FUSED2").as_deref() == Ok("0"));
16769 if off
16770 || !self.mmvq_supports(QT_NVFP4)
16771 || !self.uses_q8_1_fast(w0)
16772 || !self.uses_q8_1_fast(w1)
16773 {
16774 return Ok(false);
16775 }
16776 let unpack = |w: &crate::model::GpuTensor| match w {
16777 GpuTensor::Quant {
16778 bytes,
16779 qtype,
16780 scale,
16781 rp,
16782 ..
16783 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
16784 _ => None,
16785 };
16786 let (Some(p0), Some(p1)) = (unpack(w0), unpack(w1)) else {
16787 return Ok(false);
16788 };
16789 let in_f = w0.in_features();
16790 if w1.in_features() != in_f {
16791 return Ok(false);
16792 }
16793 let (o0, o1) = (w0.out_features(), w1.out_features());
16794 if y0.len() < o0 || y1.len() < o1 {
16795 return Ok(false);
16796 }
16797 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16799 let rows_pb = ROWS_PER_BLOCK * RPW;
16800 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
16801 let f = self.func("qmatvec_nvfp4_mmvq_fused2_rp");
16802 let cfg = LaunchConfig {
16803 grid_dim: (nb(o0) + nb(o1), 1, 1),
16804 block_dim: (32, ROWS_PER_BLOCK, 1),
16805 shared_mem_bytes: 0,
16806 };
16807 let (inf, oi0, oi1, mi) = (in_f as i32, o0 as i32, o1 as i32, 1i32);
16808 let (b0, b1) = unsafe { (&*p0.0, &*p1.0) };
16811 let __s_b = self.gpu.stream();
16812 let mut b = __s_b.launch_builder(&f);
16813 b.arg(b0)
16814 .arg(b1)
16815 .arg(aq)
16816 .arg(ad)
16817 .arg(&mut *y0)
16818 .arg(&mut *y1)
16819 .arg(&inf)
16820 .arg(&oi0)
16821 .arg(&oi1)
16822 .arg(&mi)
16823 .arg(&p0.1)
16824 .arg(&p1.1);
16825 unsafe {
16826 b.launch(cfg)?;
16827 }
16828 Ok(true)
16829 }
16830
16831 #[allow(clippy::type_complexity)]
16836 #[allow(clippy::too_many_arguments)] pub fn matmul_nvfp4_fused4(
16838 &self,
16839 w0: &crate::model::GpuTensor,
16840 w1: &crate::model::GpuTensor,
16841 w2: &crate::model::GpuTensor,
16842 w3: &crate::model::GpuTensor,
16843 aq: &CudaSlice<i8>,
16844 ad: &CudaSlice<f32>,
16845 m: usize,
16846 ) -> Result<
16847 Option<(
16848 CudaSlice<f32>,
16849 CudaSlice<f32>,
16850 CudaSlice<f32>,
16851 CudaSlice<f32>,
16852 )>,
16853 Box<dyn std::error::Error>,
16854 > {
16855 use crate::model::GpuTensor;
16856 if std::env::var("MEMRA_NVFP4_FUSED4").as_deref() == Ok("0")
16863 || !self.mmvq_supports(QT_NVFP4)
16864 || !self.uses_q8_1_fast(w0)
16865 || !self.uses_q8_1_fast(w1)
16866 || !self.uses_q8_1_fast(w2)
16867 || !self.uses_q8_1_fast(w3)
16868 {
16869 return Ok(None);
16870 }
16871 if (9..=16).contains(&m) {
16876 return Ok(
16877 match self.matmul_decode_exact_group4_pre([w0, w1, w2, w3], aq, ad, m)? {
16878 Some(mut ys) => {
16879 let y3 = ys.pop().unwrap();
16880 let y2 = ys.pop().unwrap();
16881 let y1 = ys.pop().unwrap();
16882 let y0 = ys.pop().unwrap();
16883 Some((y0, y1, y2, y3))
16884 }
16885 None => None,
16886 },
16887 );
16888 }
16889 if !(1..=8).contains(&m) {
16890 return Ok(None);
16891 }
16892 if m > 1 {
16893 let in_f = w0.in_features();
16896 if !self.batched_supports(QT_NVFP4)
16897 || std::env::var("MEMRA_NO_BATCHED").is_ok()
16898 || (m > 4 && !Self::b8_enabled())
16899 || !in_f.is_multiple_of(512)
16900 || in_f / 64 > 272
16901 {
16902 return Ok(None);
16903 }
16904 }
16905 let unpack = |w: &crate::model::GpuTensor| match w {
16906 GpuTensor::Quant {
16907 bytes,
16908 qtype,
16909 scale,
16910 rp,
16911 ..
16912 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
16913 _ => None,
16914 };
16915 let (Some(p0), Some(p1), Some(p2), Some(p3)) =
16916 (unpack(w0), unpack(w1), unpack(w2), unpack(w3))
16917 else {
16918 return Ok(None);
16919 };
16920 let in_f = w0.in_features();
16921 if w1.in_features() != in_f || w2.in_features() != in_f || w3.in_features() != in_f {
16922 return Ok(None);
16923 }
16924 let (o0, o1, o2, o3) = (
16925 w0.out_features(),
16926 w1.out_features(),
16927 w2.out_features(),
16928 w3.out_features(),
16929 );
16930 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16932 let rows_pb = ROWS_PER_BLOCK * RPW;
16933 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
16934 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
16935 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
16936 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
16937 let mut y3 = self.alloc_uninit::<f32>(m * o3)?;
16938 let (inf, oi0, oi1, oi2, oi3, mi) = (
16939 in_f as i32,
16940 o0 as i32,
16941 o1 as i32,
16942 o2 as i32,
16943 o3 as i32,
16944 m as i32,
16945 );
16946 let (b0, b1, b2, b3) = unsafe { (&*p0.0, &*p1.0, &*p2.0, &*p3.0) };
16949 if m > 1 {
16950 if p0.1 != 1.0 || p1.1 != 1.0 || p2.1 != 1.0 || p3.1 != 1.0 {
16953 return Ok(None);
16954 }
16955 let f = self.func("qmatvec_nvfp4_mmvq_fused4_b8_rpsc");
16956 let cfg = LaunchConfig {
16957 grid_dim: (nb(o0) + nb(o1) + nb(o2) + nb(o3), 1, 1),
16958 block_dim: (32, ROWS_PER_BLOCK, 1),
16959 shared_mem_bytes: 0,
16960 };
16961 let __s_b = self.gpu.stream();
16962 let mut b = __s_b.launch_builder(&f);
16963 b.arg(b0)
16964 .arg(b1)
16965 .arg(b2)
16966 .arg(b3)
16967 .arg(aq)
16968 .arg(ad)
16969 .arg(&mut y0)
16970 .arg(&mut y1)
16971 .arg(&mut y2)
16972 .arg(&mut y3)
16973 .arg(&inf)
16974 .arg(&oi0)
16975 .arg(&oi1)
16976 .arg(&oi2)
16977 .arg(&oi3)
16978 .arg(&mi);
16979 unsafe {
16980 b.launch(cfg)?;
16981 }
16982 return Ok(Some((y0, y1, y2, y3)));
16983 }
16984 let f = self.func("qmatvec_nvfp4_mmvq_fused4_rp");
16985 let cfg = LaunchConfig {
16986 grid_dim: (nb(o0) + nb(o1) + nb(o2) + nb(o3), m as u32, 1),
16987 block_dim: (32, ROWS_PER_BLOCK, 1),
16988 shared_mem_bytes: 0,
16989 };
16990 let __s_b = self.gpu.stream();
16991 let mut b = __s_b.launch_builder(&f);
16992 b.arg(b0)
16993 .arg(b1)
16994 .arg(b2)
16995 .arg(b3)
16996 .arg(aq)
16997 .arg(ad)
16998 .arg(&mut y0)
16999 .arg(&mut y1)
17000 .arg(&mut y2)
17001 .arg(&mut y3)
17002 .arg(&inf)
17003 .arg(&oi0)
17004 .arg(&oi1)
17005 .arg(&oi2)
17006 .arg(&oi3)
17007 .arg(&mi)
17008 .arg(&p0.1)
17009 .arg(&p1.1)
17010 .arg(&p2.1)
17011 .arg(&p3.1);
17012 unsafe {
17013 b.launch(cfg)?;
17014 }
17015 Ok(Some((y0, y1, y2, y3)))
17016 }
17017
17018 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused2(
17027 &self,
17028 w0: &crate::model::GpuTensor,
17029 w1: &crate::model::GpuTensor,
17030 aq: &CudaSlice<i8>,
17031 ad: &CudaSlice<f32>,
17032 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
17033 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
17039 return Ok(Some(self.e4m3_fused2_core(
17040 p0.0,
17041 p1.0,
17042 aq,
17043 ad,
17044 w0.in_features(),
17045 p0.1,
17046 p1.1,
17047 p0.2,
17048 p0.3,
17049 p1.3,
17050 )?));
17051 }
17052 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
17053 return Ok(None);
17054 };
17055 Ok(Some(self.q8_fused2_core(
17056 p0.0,
17057 p1.0,
17058 aq,
17059 ad,
17060 w0.in_features(),
17061 p0.1,
17062 p1.1,
17063 p0.2,
17064 )?))
17065 }
17066
17067 #[allow(clippy::too_many_arguments)]
17068 fn q8_fused2_core(
17069 &self,
17070 b0: &CudaSlice<u8>,
17071 b1: &CudaSlice<u8>,
17072 aq: &CudaSlice<i8>,
17073 ad: &CudaSlice<f32>,
17074 in_f: usize,
17075 out0: usize,
17076 out1: usize,
17077 row_bytes: usize,
17078 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
17079 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
17081 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
17082 let f = self.func("qmatvec_q8_0_mmvq_fused2");
17083 let mut y0 = self.alloc_uninit::<f32>(out0)?;
17084 let mut y1 = self.alloc_uninit::<f32>(out1)?;
17085 let cfg = LaunchConfig {
17086 grid_dim: (nb0 + nb1, 1, 1),
17087 block_dim: (32, ROWS_PER_BLOCK, 1),
17088 shared_mem_bytes: 0,
17089 };
17090 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
17091 let __s_b = self.gpu.stream();
17092 let mut b = __s_b.launch_builder(&f);
17093 b.arg(b0)
17094 .arg(b1)
17095 .arg(aq)
17096 .arg(ad)
17097 .arg(&mut y0)
17098 .arg(&mut y1)
17099 .arg(&inf)
17100 .arg(&o0)
17101 .arg(&o1)
17102 .arg(&rbl);
17103 unsafe {
17104 b.launch(cfg)?;
17105 }
17106 Ok((y0, y1))
17107 }
17108
17109 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused2_x(
17116 &self,
17117 w0: &crate::model::GpuTensor,
17118 w1: &crate::model::GpuTensor,
17119 x: &CudaSlice<f32>,
17120 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
17121 if !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
17122 return Ok(None);
17123 }
17124 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
17125 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
17126 return Ok(Some(self.e4m3_fused2_core(
17127 p0.0,
17128 p1.0,
17129 &aq,
17130 &ad,
17131 w0.in_features(),
17132 p0.1,
17133 p1.1,
17134 p0.2,
17135 p0.3,
17136 p1.3,
17137 )?));
17138 }
17139 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
17140 return Ok(None);
17141 };
17142 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
17143 Ok(Some(self.q8_fused2_core(
17144 p0.0,
17145 p1.0,
17146 &aq,
17147 &ad,
17148 w0.in_features(),
17149 p0.1,
17150 p1.1,
17151 p0.2,
17152 )?))
17153 }
17154
17155 #[allow(clippy::too_many_arguments)]
17158 pub fn qmatvec_q8_fused2_raw(
17159 &self,
17160 b0: &CudaSlice<u8>,
17161 b1: &CudaSlice<u8>,
17162 x: &CudaSlice<f32>,
17163 in_f: usize,
17164 out0: usize,
17165 out1: usize,
17166 row_bytes: usize,
17167 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
17168 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
17169 self.q8_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes)
17170 }
17171
17172 #[allow(clippy::type_complexity)] pub fn matmul_q4_fused3(
17179 &self,
17180 w0: &crate::model::GpuTensor,
17181 w1: &crate::model::GpuTensor,
17182 w2: &crate::model::GpuTensor,
17183 aq: &CudaSlice<i8>,
17184 ad: &CudaSlice<f32>,
17185 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
17186 {
17187 use crate::model::GpuTensor;
17188 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17189 match w {
17190 GpuTensor::Quant {
17191 qtype, row_bytes, ..
17192 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17193 _ => None,
17194 }
17195 };
17196 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2)) else {
17197 return Ok(None);
17198 };
17199 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
17200 return Ok(None);
17201 }
17202 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17206 match w {
17207 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17208 Some(m) => (m, true),
17209 None => (bytes, *rp),
17210 },
17211 _ => unreachable!(),
17212 }
17213 }
17214 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
17215 if rp0 != rp1 || rp1 != rp2 {
17216 return Ok(None);
17217 }
17218 let rp = rp0;
17219 let rpb: u32 = 4;
17220 let mr1 = rp && Self::q40_mr1_on();
17224 let nb = |o: usize| {
17225 if mr1 {
17226 (o as u32).div_ceil(rpb)
17227 } else {
17228 (o as u32).div_ceil(2).div_ceil(rpb)
17229 }
17230 };
17231 let grid = nb(o0) + nb(o1) + nb(o2);
17232 let mut y0 = self.alloc_uninit::<f32>(o0)?;
17233 let mut y1 = self.alloc_uninit::<f32>(o1)?;
17234 let mut y2 = self.alloc_uninit::<f32>(o2)?;
17235 let f = self.func(if mr1 {
17236 "qmatvec_q4_0_mmvq_fused3_mr1_rp"
17237 } else if rp {
17238 "qmatvec_q4_0_mmvq_fused3_rp"
17239 } else {
17240 "qmatvec_q4_0_mmvq_fused3"
17241 });
17242 let cfg = LaunchConfig {
17243 grid_dim: (grid, 1, 1),
17244 block_dim: (32, rpb, 1),
17245 shared_mem_bytes: 0,
17246 };
17247 let inf = w0.in_features() as i32;
17248 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
17249 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
17250 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
17253 {
17254 use cudarc::driver::{DevicePtr, DevicePtrMut};
17255 let s = &self.gpu.stream();
17256 let (p0, _g0) = b0.device_ptr(s);
17257 let (p1, _g1) = b1.device_ptr(s);
17258 let (p2, _g2) = b2.device_ptr(s);
17259 let (paq, _g3) = aq.device_ptr(s);
17260 let (pad, _g4) = ad.device_ptr(s);
17261 let (py0, _g5) = y0.device_ptr_mut(s);
17262 let (py1, _g6) = y1.device_ptr_mut(s);
17263 let (py2, _g7) = y2.device_ptr_mut(s);
17264 let mut ps = [
17265 &p0 as *const _ as *mut std::ffi::c_void,
17266 &p1 as *const _ as *mut _,
17267 &p2 as *const _ as *mut _,
17268 &paq as *const _ as *mut _,
17269 &pad as *const _ as *mut _,
17270 &py0 as *const _ as *mut _,
17271 &py1 as *const _ as *mut _,
17272 &py2 as *const _ as *mut _,
17273 &inf as *const _ as *mut _,
17274 &oo0 as *const _ as *mut _,
17275 &oo1 as *const _ as *mut _,
17276 &oo2 as *const _ as *mut _,
17277 &r0 as *const _ as *mut _,
17278 &r1 as *const _ as *mut _,
17279 &r2 as *const _ as *mut _,
17280 ];
17281 unsafe {
17282 self.launch_pdl(
17283 "qmatvec_q4_0_mmvq_fused3_mr1_rp",
17284 (grid, 1, 1),
17285 (32, rpb, 1),
17286 &mut ps,
17287 )?;
17288 }
17289 }
17290 return Ok(Some((y0, y1, y2)));
17291 }
17292 let __s_b = self.gpu.stream();
17293 let mut b = __s_b.launch_builder(&f);
17294 b.arg(b0)
17295 .arg(b1)
17296 .arg(b2)
17297 .arg(aq)
17298 .arg(ad)
17299 .arg(&mut y0)
17300 .arg(&mut y1)
17301 .arg(&mut y2)
17302 .arg(&inf)
17303 .arg(&oo0)
17304 .arg(&oo1)
17305 .arg(&oo2)
17306 .arg(&r0)
17307 .arg(&r1)
17308 .arg(&r2);
17309 unsafe {
17310 b.launch(cfg)?;
17311 }
17312 Ok(Some((y0, y1, y2)))
17313 }
17314
17315 #[allow(clippy::too_many_arguments)]
17318 pub fn matmul_q4_fused3_into(
17319 &self,
17320 w0: &crate::model::GpuTensor,
17321 w1: &crate::model::GpuTensor,
17322 w2: &crate::model::GpuTensor,
17323 aq: &CudaSlice<i8>,
17324 ad: &CudaSlice<f32>,
17325 y0: &mut CudaSlice<f32>,
17326 y1: &mut CudaSlice<f32>,
17327 y2: &mut CudaSlice<f32>,
17328 ) -> Result<bool, Box<dyn std::error::Error>> {
17329 use crate::model::GpuTensor;
17330 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17331 match w {
17332 GpuTensor::Quant {
17333 qtype, row_bytes, ..
17334 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17335 _ => None,
17336 }
17337 };
17338 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2)) else {
17339 return Ok(false);
17340 };
17341 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
17342 return Ok(false);
17343 }
17344 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17345 match w {
17346 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17347 Some(m) => (m, true),
17348 None => (bytes, *rp),
17349 },
17350 _ => unreachable!(),
17351 }
17352 }
17353 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
17354 if rp0 != rp1 || rp1 != rp2 {
17355 return Ok(false);
17356 }
17357 let rp = rp0;
17358 let rpb: u32 = 4;
17359 let mr1 = rp && Self::q40_mr1_on();
17360 let nb = |o: usize| {
17361 if mr1 {
17362 (o as u32).div_ceil(rpb)
17363 } else {
17364 (o as u32).div_ceil(2).div_ceil(rpb)
17365 }
17366 };
17367 let grid = nb(o0) + nb(o1) + nb(o2);
17368 debug_assert!(y0.len() >= o0 && y1.len() >= o1 && y2.len() >= o2);
17369 let f = self.func(if mr1 {
17370 "qmatvec_q4_0_mmvq_fused3_mr1_rp"
17371 } else if rp {
17372 "qmatvec_q4_0_mmvq_fused3_rp"
17373 } else {
17374 "qmatvec_q4_0_mmvq_fused3"
17375 });
17376 let cfg = LaunchConfig {
17377 grid_dim: (grid, 1, 1),
17378 block_dim: (32, rpb, 1),
17379 shared_mem_bytes: 0,
17380 };
17381 let inf = w0.in_features() as i32;
17382 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
17383 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
17384 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
17386 use cudarc::driver::{DevicePtr, DevicePtrMut};
17387 let s = &self.gpu.stream();
17388 let (p0, _g0) = b0.device_ptr(s);
17389 let (p1, _g1) = b1.device_ptr(s);
17390 let (p2, _g2) = b2.device_ptr(s);
17391 let (paq, _g3) = aq.device_ptr(s);
17392 let (pad, _g4) = ad.device_ptr(s);
17393 let (py0, _g5) = y0.device_ptr_mut(s);
17394 let (py1, _g6) = y1.device_ptr_mut(s);
17395 let (py2, _g7) = y2.device_ptr_mut(s);
17396 let mut ps = [
17397 &p0 as *const _ as *mut std::ffi::c_void,
17398 &p1 as *const _ as *mut _,
17399 &p2 as *const _ as *mut _,
17400 &paq as *const _ as *mut _,
17401 &pad as *const _ as *mut _,
17402 &py0 as *const _ as *mut _,
17403 &py1 as *const _ as *mut _,
17404 &py2 as *const _ as *mut _,
17405 &inf as *const _ as *mut _,
17406 &oo0 as *const _ as *mut _,
17407 &oo1 as *const _ as *mut _,
17408 &oo2 as *const _ as *mut _,
17409 &r0 as *const _ as *mut _,
17410 &r1 as *const _ as *mut _,
17411 &r2 as *const _ as *mut _,
17412 ];
17413 unsafe {
17414 self.launch_pdl(
17415 "qmatvec_q4_0_mmvq_fused3_mr1_rp",
17416 (grid, 1, 1),
17417 (32, rpb, 1),
17418 &mut ps,
17419 )?;
17420 }
17421 return Ok(true);
17422 }
17423 let __s_b = self.gpu.stream();
17424 let mut b = __s_b.launch_builder(&f);
17425 b.arg(b0)
17426 .arg(b1)
17427 .arg(b2)
17428 .arg(aq)
17429 .arg(ad)
17430 .arg(&mut *y0)
17431 .arg(&mut *y1)
17432 .arg(&mut *y2)
17433 .arg(&inf)
17434 .arg(&oo0)
17435 .arg(&oo1)
17436 .arg(&oo2)
17437 .arg(&r0)
17438 .arg(&r1)
17439 .arg(&r2);
17440 unsafe {
17441 b.launch(cfg)?;
17442 }
17443 Ok(true)
17444 }
17445
17446 #[allow(clippy::type_complexity)] pub fn matmul_q4_fused2(
17449 &self,
17450 w0: &crate::model::GpuTensor,
17451 w1: &crate::model::GpuTensor,
17452 aq: &CudaSlice<i8>,
17453 ad: &CudaSlice<f32>,
17454 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
17455 use crate::model::GpuTensor;
17456 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17457 match w {
17458 GpuTensor::Quant {
17459 qtype, row_bytes, ..
17460 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17461 _ => None,
17462 }
17463 };
17464 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else {
17465 return Ok(None);
17466 };
17467 if w0.in_features() != w1.in_features() {
17468 return Ok(None);
17469 }
17470 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17472 match w {
17473 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17474 Some(m) => (m, true),
17475 None => (bytes, *rp),
17476 },
17477 _ => unreachable!(),
17478 }
17479 }
17480 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
17481 if rp0 != rp1 {
17482 return Ok(None);
17483 }
17484 let rp = rp0;
17485 let rpb: u32 = 4;
17486 let mr1 = rp && Self::q40_mr1_on();
17488 let nb = |o: usize| {
17489 if mr1 {
17490 (o as u32).div_ceil(rpb)
17491 } else {
17492 (o as u32).div_ceil(2).div_ceil(rpb)
17493 }
17494 };
17495 let grid = nb(o0) + nb(o1);
17496 let mut y0 = self.alloc_uninit::<f32>(o0)?;
17497 let mut y1 = self.alloc_uninit::<f32>(o1)?;
17498 let f = self.func(if mr1 {
17499 "qmatvec_q4_0_mmvq_fused2_mr1_rp"
17500 } else if rp {
17501 "qmatvec_q4_0_mmvq_fused2_rp"
17502 } else {
17503 "qmatvec_q4_0_mmvq_fused2"
17504 });
17505 let cfg = LaunchConfig {
17506 grid_dim: (grid, 1, 1),
17507 block_dim: (32, rpb, 1),
17508 shared_mem_bytes: 0,
17509 };
17510 let inf = w0.in_features() as i32;
17511 let (oo0, oo1) = (o0 as i32, o1 as i32);
17512 let (r0, r1) = (rb0 as i64, rb1 as i64);
17513 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
17515 {
17516 use cudarc::driver::{DevicePtr, DevicePtrMut};
17517 let s = &self.gpu.stream();
17518 let (p0, _g0) = b0.device_ptr(s);
17519 let (p1, _g1) = b1.device_ptr(s);
17520 let (paq, _g2) = aq.device_ptr(s);
17521 let (pad, _g3) = ad.device_ptr(s);
17522 let (py0, _g4) = y0.device_ptr_mut(s);
17523 let (py1, _g5) = y1.device_ptr_mut(s);
17524 let mut ps = [
17525 &p0 as *const _ as *mut std::ffi::c_void,
17526 &p1 as *const _ as *mut _,
17527 &paq as *const _ as *mut _,
17528 &pad as *const _ as *mut _,
17529 &py0 as *const _ as *mut _,
17530 &py1 as *const _ as *mut _,
17531 &inf as *const _ as *mut _,
17532 &oo0 as *const _ as *mut _,
17533 &oo1 as *const _ as *mut _,
17534 &r0 as *const _ as *mut _,
17535 &r1 as *const _ as *mut _,
17536 ];
17537 unsafe {
17538 self.launch_pdl(
17539 "qmatvec_q4_0_mmvq_fused2_mr1_rp",
17540 (grid, 1, 1),
17541 (32, rpb, 1),
17542 &mut ps,
17543 )?;
17544 }
17545 }
17546 return Ok(Some((y0, y1)));
17547 }
17548 let __s_b = self.gpu.stream();
17549 let mut b = __s_b.launch_builder(&f);
17550 b.arg(b0)
17551 .arg(b1)
17552 .arg(aq)
17553 .arg(ad)
17554 .arg(&mut y0)
17555 .arg(&mut y1)
17556 .arg(&inf)
17557 .arg(&oo0)
17558 .arg(&oo1)
17559 .arg(&r0)
17560 .arg(&r1);
17561 unsafe {
17562 b.launch(cfg)?;
17563 }
17564 Ok(Some((y0, y1)))
17565 }
17566
17567 pub fn matmul_q4_fused2_into(
17569 &self,
17570 w0: &crate::model::GpuTensor,
17571 w1: &crate::model::GpuTensor,
17572 aq: &CudaSlice<i8>,
17573 ad: &CudaSlice<f32>,
17574 y0: &mut CudaSlice<f32>,
17575 y1: &mut CudaSlice<f32>,
17576 ) -> Result<bool, Box<dyn std::error::Error>> {
17577 use crate::model::GpuTensor;
17578 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17579 match w {
17580 GpuTensor::Quant {
17581 qtype, row_bytes, ..
17582 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17583 _ => None,
17584 }
17585 };
17586 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else {
17587 return Ok(false);
17588 };
17589 if w0.in_features() != w1.in_features() {
17590 return Ok(false);
17591 }
17592 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17593 match w {
17594 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17595 Some(m) => (m, true),
17596 None => (bytes, *rp),
17597 },
17598 _ => unreachable!(),
17599 }
17600 }
17601 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
17602 if rp0 != rp1 {
17603 return Ok(false);
17604 }
17605 let rp = rp0;
17606 let rpb: u32 = 4;
17607 let mr1 = rp && Self::q40_mr1_on();
17608 let nb = |o: usize| {
17609 if mr1 {
17610 (o as u32).div_ceil(rpb)
17611 } else {
17612 (o as u32).div_ceil(2).div_ceil(rpb)
17613 }
17614 };
17615 let grid = nb(o0) + nb(o1);
17616 debug_assert!(y0.len() >= o0 && y1.len() >= o1);
17617 let f = self.func(if mr1 {
17618 "qmatvec_q4_0_mmvq_fused2_mr1_rp"
17619 } else if rp {
17620 "qmatvec_q4_0_mmvq_fused2_rp"
17621 } else {
17622 "qmatvec_q4_0_mmvq_fused2"
17623 });
17624 let cfg = LaunchConfig {
17625 grid_dim: (grid, 1, 1),
17626 block_dim: (32, rpb, 1),
17627 shared_mem_bytes: 0,
17628 };
17629 let inf = w0.in_features() as i32;
17630 let (oo0, oo1) = (o0 as i32, o1 as i32);
17631 let (r0, r1) = (rb0 as i64, rb1 as i64);
17632 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
17634 use cudarc::driver::{DevicePtr, DevicePtrMut};
17635 let s = &self.gpu.stream();
17636 let (p0, _g0) = b0.device_ptr(s);
17637 let (p1, _g1) = b1.device_ptr(s);
17638 let (paq, _g2) = aq.device_ptr(s);
17639 let (pad, _g3) = ad.device_ptr(s);
17640 let (py0, _g4) = y0.device_ptr_mut(s);
17641 let (py1, _g5) = y1.device_ptr_mut(s);
17642 let mut ps = [
17643 &p0 as *const _ as *mut std::ffi::c_void,
17644 &p1 as *const _ as *mut _,
17645 &paq as *const _ as *mut _,
17646 &pad as *const _ as *mut _,
17647 &py0 as *const _ as *mut _,
17648 &py1 as *const _ as *mut _,
17649 &inf as *const _ as *mut _,
17650 &oo0 as *const _ as *mut _,
17651 &oo1 as *const _ as *mut _,
17652 &r0 as *const _ as *mut _,
17653 &r1 as *const _ as *mut _,
17654 ];
17655 unsafe {
17656 self.launch_pdl(
17657 "qmatvec_q4_0_mmvq_fused2_mr1_rp",
17658 (grid, 1, 1),
17659 (32, rpb, 1),
17660 &mut ps,
17661 )?;
17662 }
17663 return Ok(true);
17664 }
17665 let __s_b = self.gpu.stream();
17666 let mut b = __s_b.launch_builder(&f);
17667 b.arg(b0)
17668 .arg(b1)
17669 .arg(aq)
17670 .arg(ad)
17671 .arg(&mut *y0)
17672 .arg(&mut *y1)
17673 .arg(&inf)
17674 .arg(&oo0)
17675 .arg(&oo1)
17676 .arg(&r0)
17677 .arg(&r1);
17678 unsafe {
17679 b.launch(cfg)?;
17680 }
17681 Ok(true)
17682 }
17683
17684 #[allow(clippy::type_complexity)] pub fn matmul_q4_fused2_batched(
17690 &self,
17691 w0: &crate::model::GpuTensor,
17692 w1: &crate::model::GpuTensor,
17693 aq: &CudaSlice<i8>,
17694 ad: &CudaSlice<f32>,
17695 m: usize,
17696 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
17697 use crate::model::GpuTensor;
17698 if !(2..=8).contains(&m) {
17699 return Ok(None);
17700 }
17701 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17702 match w {
17703 GpuTensor::Quant {
17704 qtype, row_bytes, ..
17705 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17706 _ => None,
17707 }
17708 };
17709 let (Some((rb0, o0)), Some((_rb1, o1))) = (q4(w0), q4(w1)) else {
17710 return Ok(None);
17711 };
17712 if w0.in_features() != w1.in_features() {
17713 return Ok(None);
17714 }
17715 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17716 match w {
17717 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17718 Some(mr) => (mr, true),
17719 None => (bytes, *rp),
17720 },
17721 _ => unreachable!(),
17722 }
17723 }
17724 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
17725 if !rp0 || !rp1 {
17726 return Ok(None);
17727 }
17728 let mcols = Self::batched_mcols(m);
17729 let rpb: u32 = 4;
17730 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
17731 let grid = nb(o0) + nb(o1);
17732 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
17733 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
17734 let f = self.func(match mcols {
17735 2 => "qmatvec_q4_0_mmvq_b2_f2_rp",
17736 4 => "qmatvec_q4_0_mmvq_b4_f2_rp",
17737 _ => "qmatvec_q4_0_mmvq_b8_f2_rp",
17738 });
17739 let cfg = LaunchConfig {
17740 grid_dim: (grid, 1, 1),
17741 block_dim: (32, rpb, 1),
17742 shared_mem_bytes: 0,
17743 };
17744 let inf = w0.in_features() as i32;
17745 let (oo0, oo1, mi) = (o0 as i32, o1 as i32, m as i32);
17746 let rb = rb0 as i64;
17747 let __s_b = self.gpu.stream();
17748 let mut b = __s_b.launch_builder(&f);
17749 b.arg(b0)
17750 .arg(b1)
17751 .arg(aq)
17752 .arg(ad)
17753 .arg(&mut y0)
17754 .arg(&mut y1)
17755 .arg(&inf)
17756 .arg(&oo0)
17757 .arg(&oo1)
17758 .arg(&mi)
17759 .arg(&rb);
17760 unsafe {
17761 b.launch(cfg)?;
17762 }
17763 Ok(Some((y0, y1)))
17764 }
17765
17766 #[allow(clippy::too_many_arguments)]
17769 #[allow(clippy::type_complexity)] pub fn matmul_q4_fused3_batched(
17771 &self,
17772 w0: &crate::model::GpuTensor,
17773 w1: &crate::model::GpuTensor,
17774 w2: &crate::model::GpuTensor,
17775 aq: &CudaSlice<i8>,
17776 ad: &CudaSlice<f32>,
17777 m: usize,
17778 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
17779 {
17780 use crate::model::GpuTensor;
17781 if !(2..=8).contains(&m) {
17782 return Ok(None);
17783 }
17784 let q4 = |w: &GpuTensor| -> Option<usize> {
17785 match w {
17786 GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0 => Some(w.out_features()),
17787 _ => None,
17788 }
17789 };
17790 let (Some(o0), Some(o1), Some(o2)) = (q4(w0), q4(w1), q4(w2)) else {
17791 return Ok(None);
17792 };
17793 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
17794 return Ok(None);
17795 }
17796 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17797 match w {
17798 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17799 Some(mr) => (mr, true),
17800 None => (bytes, *rp),
17801 },
17802 _ => unreachable!(),
17803 }
17804 }
17805 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
17806 if !rp0 || !rp1 || !rp2 {
17807 return Ok(None);
17808 }
17809 let mcols = Self::batched_mcols(m);
17810 let rpb: u32 = 4;
17811 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
17812 let grid = nb(o0) + nb(o1) + nb(o2);
17813 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
17814 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
17815 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
17816 let f = self.func(match mcols {
17817 2 => "qmatvec_q4_0_mmvq_b2_f3_rp",
17818 4 => "qmatvec_q4_0_mmvq_b4_f3_rp",
17819 _ => "qmatvec_q4_0_mmvq_b8_f3_rp",
17820 });
17821 let cfg = LaunchConfig {
17822 grid_dim: (grid, 1, 1),
17823 block_dim: (32, rpb, 1),
17824 shared_mem_bytes: 0,
17825 };
17826 let inf = w0.in_features() as i32;
17827 let (oo0, oo1, oo2, mi) = (o0 as i32, o1 as i32, o2 as i32, m as i32);
17828 let rb = 0i64;
17829 let __s_b = self.gpu.stream();
17830 let mut b = __s_b.launch_builder(&f);
17831 b.arg(b0)
17832 .arg(b1)
17833 .arg(b2)
17834 .arg(aq)
17835 .arg(ad)
17836 .arg(&mut y0)
17837 .arg(&mut y1)
17838 .arg(&mut y2)
17839 .arg(&inf)
17840 .arg(&oo0)
17841 .arg(&oo1)
17842 .arg(&oo2)
17843 .arg(&mi)
17844 .arg(&rb);
17845 unsafe {
17846 b.launch(cfg)?;
17847 }
17848 Ok(Some((y0, y1, y2)))
17849 }
17850
17851 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused3(
17853 &self,
17854 w0: &crate::model::GpuTensor,
17855 w1: &crate::model::GpuTensor,
17856 w2: &crate::model::GpuTensor,
17857 aq: &CudaSlice<i8>,
17858 ad: &CudaSlice<f32>,
17859 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
17860 {
17861 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
17864 return Ok(Some(self.e4m3_fused3_core(
17865 p0.0,
17866 p1.0,
17867 p2.0,
17868 aq,
17869 ad,
17870 w0.in_features(),
17871 p0.1,
17872 p1.1,
17873 p2.1,
17874 p0.2,
17875 p0.3,
17876 p1.3,
17877 p2.3,
17878 )?));
17879 }
17880 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else {
17881 return Ok(None);
17882 };
17883 Ok(Some(self.q8_fused3_core(
17884 p0.0,
17885 p1.0,
17886 p2.0,
17887 aq,
17888 ad,
17889 w0.in_features(),
17890 p0.1,
17891 p1.1,
17892 p2.1,
17893 p0.2,
17894 )?))
17895 }
17896
17897 #[allow(clippy::too_many_arguments)]
17898 #[allow(clippy::type_complexity)] fn q8_fused3_core(
17900 &self,
17901 b0: &CudaSlice<u8>,
17902 b1: &CudaSlice<u8>,
17903 b2: &CudaSlice<u8>,
17904 aq: &CudaSlice<i8>,
17905 ad: &CudaSlice<f32>,
17906 in_f: usize,
17907 out0: usize,
17908 out1: usize,
17909 out2: usize,
17910 row_bytes: usize,
17911 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
17912 const ROWS_PER_BLOCK: u32 = 4;
17913 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
17914 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
17915 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
17916 let f = self.func("qmatvec_q8_0_mmvq_fused3");
17917 let mut y0 = self.alloc_uninit::<f32>(out0)?;
17918 let mut y1 = self.alloc_uninit::<f32>(out1)?;
17919 let mut y2 = self.alloc_uninit::<f32>(out2)?;
17920 let cfg = LaunchConfig {
17921 grid_dim: (nb0 + nb1 + nb2, 1, 1),
17922 block_dim: (32, ROWS_PER_BLOCK, 1),
17923 shared_mem_bytes: 0,
17924 };
17925 let (inf, o0, o1, o2, rbl) = (
17926 in_f as i32,
17927 out0 as i32,
17928 out1 as i32,
17929 out2 as i32,
17930 row_bytes as i64,
17931 );
17932 let __s_b = self.gpu.stream();
17933 let mut b = __s_b.launch_builder(&f);
17934 b.arg(b0)
17935 .arg(b1)
17936 .arg(b2)
17937 .arg(aq)
17938 .arg(ad)
17939 .arg(&mut y0)
17940 .arg(&mut y1)
17941 .arg(&mut y2)
17942 .arg(&inf)
17943 .arg(&o0)
17944 .arg(&o1)
17945 .arg(&o2)
17946 .arg(&rbl);
17947 unsafe {
17948 b.launch(cfg)?;
17949 }
17950 Ok((y0, y1, y2))
17951 }
17952
17953 #[allow(clippy::too_many_arguments)]
17955 #[allow(clippy::type_complexity)] pub fn qmatvec_q8_fused3_raw(
17957 &self,
17958 b0: &CudaSlice<u8>,
17959 b1: &CudaSlice<u8>,
17960 b2: &CudaSlice<u8>,
17961 x: &CudaSlice<f32>,
17962 in_f: usize,
17963 out0: usize,
17964 out1: usize,
17965 out2: usize,
17966 row_bytes: usize,
17967 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
17968 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
17969 self.q8_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes)
17970 }
17971
17972 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused2_t(
17984 &self,
17985 w0: &crate::model::GpuTensor,
17986 w1: &crate::model::GpuTensor,
17987 aq: &CudaSlice<i8>,
17988 ad: &CudaSlice<f32>,
17989 m: usize,
17990 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
17991 if !(2..=8).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() {
17995 return Ok(None);
17996 }
17997 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
18000 if m > 4 && !Self::b8_enabled() {
18001 return Ok(None);
18002 }
18003 return Ok(Some(self.e4m3_fused2_t_core(
18004 p0.0,
18005 p1.0,
18006 aq,
18007 ad,
18008 m,
18009 w0.in_features(),
18010 p0.1,
18011 p1.1,
18012 p0.2,
18013 p0.3,
18014 p1.3,
18015 )?));
18016 }
18017 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
18018 return Ok(None);
18019 };
18020 Ok(Some(self.q8_fused2_t_core(
18021 p0.0,
18022 p1.0,
18023 aq,
18024 ad,
18025 m,
18026 w0.in_features(),
18027 p0.1,
18028 p1.1,
18029 p0.2,
18030 )?))
18031 }
18032
18033 #[allow(clippy::too_many_arguments)]
18034 fn q8_fused2_t_core(
18035 &self,
18036 b0: &CudaSlice<u8>,
18037 b1: &CudaSlice<u8>,
18038 aq: &CudaSlice<i8>,
18039 ad: &CudaSlice<f32>,
18040 m: usize,
18041 in_f: usize,
18042 out0: usize,
18043 out1: usize,
18044 row_bytes: usize,
18045 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18046 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18048 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18049 let f = self.func(match Self::batched_mcols(m) {
18050 2 => "qmatvec_q8_0_mmvq_fused2_b2",
18051 4 => "qmatvec_q8_0_mmvq_fused2_b4",
18052 _ => "qmatvec_q8_0_mmvq_fused2_b8",
18054 });
18055 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
18056 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
18057 let cfg = LaunchConfig {
18058 grid_dim: (nb0 + nb1, 1, 1),
18059 block_dim: (32, ROWS_PER_BLOCK, 1),
18060 shared_mem_bytes: 0,
18061 };
18062 let (inf, o0, o1, mi, rbl) = (
18063 in_f as i32,
18064 out0 as i32,
18065 out1 as i32,
18066 m as i32,
18067 row_bytes as i64,
18068 );
18069 let __s_b = self.gpu.stream();
18070 let mut b = __s_b.launch_builder(&f);
18071 b.arg(b0)
18072 .arg(b1)
18073 .arg(aq)
18074 .arg(ad)
18075 .arg(&mut y0)
18076 .arg(&mut y1)
18077 .arg(&inf)
18078 .arg(&o0)
18079 .arg(&o1)
18080 .arg(&mi)
18081 .arg(&rbl);
18082 unsafe {
18083 b.launch(cfg)?;
18084 }
18085 Ok((y0, y1))
18086 }
18087
18088 #[allow(clippy::too_many_arguments)]
18091 pub fn qmatvec_q8_fused2_t_raw(
18092 &self,
18093 b0: &CudaSlice<u8>,
18094 b1: &CudaSlice<u8>,
18095 x: &CudaSlice<f32>,
18096 m: usize,
18097 in_f: usize,
18098 out0: usize,
18099 out1: usize,
18100 row_bytes: usize,
18101 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18102 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18103 self.q8_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes)
18104 }
18105
18106 #[allow(clippy::too_many_arguments)]
18109 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused3_t(
18111 &self,
18112 w0: &crate::model::GpuTensor,
18113 w1: &crate::model::GpuTensor,
18114 w2: &crate::model::GpuTensor,
18115 aq: &CudaSlice<i8>,
18116 ad: &CudaSlice<f32>,
18117 m: usize,
18118 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
18119 {
18120 if !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() {
18121 return Ok(None);
18122 }
18123 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
18124 return Ok(Some(self.e4m3_fused3_t_core(
18125 p0.0,
18126 p1.0,
18127 p2.0,
18128 aq,
18129 ad,
18130 m,
18131 w0.in_features(),
18132 p0.1,
18133 p1.1,
18134 p2.1,
18135 p0.2,
18136 p0.3,
18137 p1.3,
18138 p2.3,
18139 )?));
18140 }
18141 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else {
18142 return Ok(None);
18143 };
18144 Ok(Some(self.q8_fused3_t_core(
18145 p0.0,
18146 p1.0,
18147 p2.0,
18148 aq,
18149 ad,
18150 m,
18151 w0.in_features(),
18152 p0.1,
18153 p1.1,
18154 p2.1,
18155 p0.2,
18156 )?))
18157 }
18158
18159 #[allow(clippy::too_many_arguments)]
18160 #[allow(clippy::type_complexity)] fn q8_fused3_t_core(
18162 &self,
18163 b0: &CudaSlice<u8>,
18164 b1: &CudaSlice<u8>,
18165 b2: &CudaSlice<u8>,
18166 aq: &CudaSlice<i8>,
18167 ad: &CudaSlice<f32>,
18168 m: usize,
18169 in_f: usize,
18170 out0: usize,
18171 out1: usize,
18172 out2: usize,
18173 row_bytes: usize,
18174 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18175 const ROWS_PER_BLOCK: u32 = 4;
18176 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18177 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18178 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
18179 let f = self.func(if Self::batched_mcols(m) == 2 {
18180 "qmatvec_q8_0_mmvq_fused3_b2"
18181 } else {
18182 "qmatvec_q8_0_mmvq_fused3_b4"
18183 });
18184 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
18185 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
18186 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
18187 let cfg = LaunchConfig {
18188 grid_dim: (nb0 + nb1 + nb2, 1, 1),
18189 block_dim: (32, ROWS_PER_BLOCK, 1),
18190 shared_mem_bytes: 0,
18191 };
18192 let (inf, o0, o1, o2, mi, rbl) = (
18193 in_f as i32,
18194 out0 as i32,
18195 out1 as i32,
18196 out2 as i32,
18197 m as i32,
18198 row_bytes as i64,
18199 );
18200 let __s_b = self.gpu.stream();
18201 let mut b = __s_b.launch_builder(&f);
18202 b.arg(b0)
18203 .arg(b1)
18204 .arg(b2)
18205 .arg(aq)
18206 .arg(ad)
18207 .arg(&mut y0)
18208 .arg(&mut y1)
18209 .arg(&mut y2)
18210 .arg(&inf)
18211 .arg(&o0)
18212 .arg(&o1)
18213 .arg(&o2)
18214 .arg(&mi)
18215 .arg(&rbl);
18216 unsafe {
18217 b.launch(cfg)?;
18218 }
18219 Ok((y0, y1, y2))
18220 }
18221
18222 #[allow(clippy::too_many_arguments)]
18224 #[allow(clippy::type_complexity)] pub fn qmatvec_q8_fused3_t_raw(
18226 &self,
18227 b0: &CudaSlice<u8>,
18228 b1: &CudaSlice<u8>,
18229 b2: &CudaSlice<u8>,
18230 x: &CudaSlice<f32>,
18231 m: usize,
18232 in_f: usize,
18233 out0: usize,
18234 out1: usize,
18235 out2: usize,
18236 row_bytes: usize,
18237 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18238 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18239 self.q8_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes)
18240 }
18241
18242 pub fn q8_ffn_fuse2_on(&self) -> bool {
18246 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18247 *ON.get_or_init(|| std::env::var("MEMRA_Q8_FFN_FUSE2").as_deref() != Ok("0"))
18248 }
18249
18250 #[allow(clippy::type_complexity)]
18256 fn q8_fused_params<'w, const N: usize>(
18257 &self,
18258 ws: &[&'w crate::model::GpuTensor; N],
18259 ) -> Option<[(&'w CudaSlice<u8>, usize, usize); N]> {
18260 use crate::model::GpuTensor;
18261 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") {
18262 return None;
18263 }
18264 if std::env::var("MEMRA_Q8_DUAL").is_ok_and(|v| v == "0") {
18265 return None;
18266 }
18267 let in_f = ws[0].in_features();
18268 let mut out: [Option<(&CudaSlice<u8>, usize, usize)>; N] = [None; N];
18269 for (i, w) in ws.iter().enumerate() {
18270 match w {
18271 GpuTensor::Quant {
18272 bytes,
18273 qtype,
18274 row_bytes,
18275 scale,
18276 ..
18277 } if *qtype == QT_Q8_0 && *scale == 1.0 && w.in_features() == in_f => {
18278 out[i] = Some((bytes, w.out_features(), *row_bytes))
18279 }
18280 _ => return None,
18281 }
18282 }
18283 Some(out.map(|o| o.unwrap()))
18284 }
18285
18286 pub fn e4m3_dual_on(&self) -> bool {
18289 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18290 *ON.get_or_init(|| std::env::var("MEMRA_E4M3_DUAL").as_deref() != Ok("0"))
18291 }
18292
18293 #[allow(clippy::type_complexity)]
18305 fn e4m3_fused_params<'w, const N: usize>(
18306 &self,
18307 ws: &[&'w crate::model::GpuTensor; N],
18308 ) -> Option<[(&'w CudaSlice<u8>, usize, usize, f32); N]> {
18309 use crate::model::GpuTensor;
18310 if !self.e4m3_dual_on() {
18311 return None;
18312 }
18313 let in_f = ws[0].in_features();
18314 let mut out: [Option<(&CudaSlice<u8>, usize, usize, f32)>; N] = [None; N];
18315 for (i, w) in ws.iter().enumerate() {
18316 match w {
18317 GpuTensor::Quant {
18318 bytes,
18319 qtype,
18320 row_bytes,
18321 scale,
18322 rp,
18323 rp4,
18324 ..
18325 } if *qtype == QT_F8_E4M3
18326 && w.in_features() == in_f
18327 && *row_bytes == in_f
18328 && !*rp
18329 && rp4.is_none() =>
18330 {
18331 out[i] = Some((bytes, w.out_features(), *row_bytes, *scale))
18332 }
18333 _ => return None,
18334 }
18335 }
18336 Some(out.map(|o| o.unwrap()))
18337 }
18338
18339 #[allow(clippy::too_many_arguments)]
18343 fn e4m3_fused2_core(
18344 &self,
18345 b0: &CudaSlice<u8>,
18346 b1: &CudaSlice<u8>,
18347 aq: &CudaSlice<i8>,
18348 ad: &CudaSlice<f32>,
18349 in_f: usize,
18350 out0: usize,
18351 out1: usize,
18352 row_bytes: usize,
18353 ws0: f32,
18354 ws1: f32,
18355 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18356 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18358 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18359 let f = self.func("qmatvec_e4m3_mmvq_fused2");
18360 let mut y0 = self.alloc_uninit::<f32>(out0)?;
18361 let mut y1 = self.alloc_uninit::<f32>(out1)?;
18362 let cfg = LaunchConfig {
18363 grid_dim: (nb0 + nb1, 1, 1),
18364 block_dim: (32, ROWS_PER_BLOCK, 1),
18365 shared_mem_bytes: 0,
18366 };
18367 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
18368 let __s_b = self.gpu.stream();
18369 let mut b = __s_b.launch_builder(&f);
18370 b.arg(b0)
18371 .arg(b1)
18372 .arg(aq)
18373 .arg(ad)
18374 .arg(&mut y0)
18375 .arg(&mut y1)
18376 .arg(&inf)
18377 .arg(&o0)
18378 .arg(&o1)
18379 .arg(&rbl)
18380 .arg(&ws0)
18381 .arg(&ws1);
18382 unsafe {
18383 b.launch(cfg)?;
18384 }
18385 Ok((y0, y1))
18386 }
18387
18388 #[allow(clippy::too_many_arguments)]
18390 #[allow(clippy::type_complexity)] fn e4m3_fused3_core(
18392 &self,
18393 b0: &CudaSlice<u8>,
18394 b1: &CudaSlice<u8>,
18395 b2: &CudaSlice<u8>,
18396 aq: &CudaSlice<i8>,
18397 ad: &CudaSlice<f32>,
18398 in_f: usize,
18399 out0: usize,
18400 out1: usize,
18401 out2: usize,
18402 row_bytes: usize,
18403 ws0: f32,
18404 ws1: f32,
18405 ws2: f32,
18406 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18407 const ROWS_PER_BLOCK: u32 = 4;
18408 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18409 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18410 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
18411 let f = self.func("qmatvec_e4m3_mmvq_fused3");
18412 let mut y0 = self.alloc_uninit::<f32>(out0)?;
18413 let mut y1 = self.alloc_uninit::<f32>(out1)?;
18414 let mut y2 = self.alloc_uninit::<f32>(out2)?;
18415 let cfg = LaunchConfig {
18416 grid_dim: (nb0 + nb1 + nb2, 1, 1),
18417 block_dim: (32, ROWS_PER_BLOCK, 1),
18418 shared_mem_bytes: 0,
18419 };
18420 let (inf, o0, o1, o2, rbl) = (
18421 in_f as i32,
18422 out0 as i32,
18423 out1 as i32,
18424 out2 as i32,
18425 row_bytes as i64,
18426 );
18427 let __s_b = self.gpu.stream();
18428 let mut b = __s_b.launch_builder(&f);
18429 b.arg(b0)
18430 .arg(b1)
18431 .arg(b2)
18432 .arg(aq)
18433 .arg(ad)
18434 .arg(&mut y0)
18435 .arg(&mut y1)
18436 .arg(&mut y2)
18437 .arg(&inf)
18438 .arg(&o0)
18439 .arg(&o1)
18440 .arg(&o2)
18441 .arg(&rbl)
18442 .arg(&ws0)
18443 .arg(&ws1)
18444 .arg(&ws2);
18445 unsafe {
18446 b.launch(cfg)?;
18447 }
18448 Ok((y0, y1, y2))
18449 }
18450
18451 #[allow(clippy::too_many_arguments)]
18455 fn e4m3_fused2_t_core(
18456 &self,
18457 b0: &CudaSlice<u8>,
18458 b1: &CudaSlice<u8>,
18459 aq: &CudaSlice<i8>,
18460 ad: &CudaSlice<f32>,
18461 m: usize,
18462 in_f: usize,
18463 out0: usize,
18464 out1: usize,
18465 row_bytes: usize,
18466 ws0: f32,
18467 ws1: f32,
18468 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18469 const ROWS_PER_BLOCK: u32 = 4;
18470 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18471 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18472 let f = self.func(match Self::batched_mcols(m) {
18473 2 => "qmatvec_e4m3_mmvq_fused2_b2",
18474 4 => "qmatvec_e4m3_mmvq_fused2_b4",
18475 _ => "qmatvec_e4m3_mmvq_fused2_b8",
18476 });
18477 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
18478 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
18479 let cfg = LaunchConfig {
18480 grid_dim: (nb0 + nb1, 1, 1),
18481 block_dim: (32, ROWS_PER_BLOCK, 1),
18482 shared_mem_bytes: 0,
18483 };
18484 let (inf, o0, o1, mi, rbl) = (
18485 in_f as i32,
18486 out0 as i32,
18487 out1 as i32,
18488 m as i32,
18489 row_bytes as i64,
18490 );
18491 let __s_b = self.gpu.stream();
18492 let mut b = __s_b.launch_builder(&f);
18493 b.arg(b0)
18494 .arg(b1)
18495 .arg(aq)
18496 .arg(ad)
18497 .arg(&mut y0)
18498 .arg(&mut y1)
18499 .arg(&inf)
18500 .arg(&o0)
18501 .arg(&o1)
18502 .arg(&mi)
18503 .arg(&rbl);
18504 unsafe {
18505 b.launch(cfg)?;
18506 }
18507 if ws0 != 1.0 {
18508 self.scale_inplace(&mut y0, ws0, m * out0)?;
18509 }
18510 if ws1 != 1.0 {
18511 self.scale_inplace(&mut y1, ws1, m * out1)?;
18512 }
18513 Ok((y0, y1))
18514 }
18515
18516 #[allow(clippy::too_many_arguments)]
18518 #[allow(clippy::type_complexity)] fn e4m3_fused3_t_core(
18520 &self,
18521 b0: &CudaSlice<u8>,
18522 b1: &CudaSlice<u8>,
18523 b2: &CudaSlice<u8>,
18524 aq: &CudaSlice<i8>,
18525 ad: &CudaSlice<f32>,
18526 m: usize,
18527 in_f: usize,
18528 out0: usize,
18529 out1: usize,
18530 out2: usize,
18531 row_bytes: usize,
18532 ws0: f32,
18533 ws1: f32,
18534 ws2: f32,
18535 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18536 const ROWS_PER_BLOCK: u32 = 4;
18537 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18538 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18539 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
18540 let f = self.func(if Self::batched_mcols(m) == 2 {
18541 "qmatvec_e4m3_mmvq_fused3_b2"
18542 } else {
18543 "qmatvec_e4m3_mmvq_fused3_b4"
18544 });
18545 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
18546 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
18547 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
18548 let cfg = LaunchConfig {
18549 grid_dim: (nb0 + nb1 + nb2, 1, 1),
18550 block_dim: (32, ROWS_PER_BLOCK, 1),
18551 shared_mem_bytes: 0,
18552 };
18553 let (inf, o0, o1, o2, mi, rbl) = (
18554 in_f as i32,
18555 out0 as i32,
18556 out1 as i32,
18557 out2 as i32,
18558 m as i32,
18559 row_bytes as i64,
18560 );
18561 let __s_b = self.gpu.stream();
18562 let mut b = __s_b.launch_builder(&f);
18563 b.arg(b0)
18564 .arg(b1)
18565 .arg(b2)
18566 .arg(aq)
18567 .arg(ad)
18568 .arg(&mut y0)
18569 .arg(&mut y1)
18570 .arg(&mut y2)
18571 .arg(&inf)
18572 .arg(&o0)
18573 .arg(&o1)
18574 .arg(&o2)
18575 .arg(&mi)
18576 .arg(&rbl);
18577 unsafe {
18578 b.launch(cfg)?;
18579 }
18580 if ws0 != 1.0 {
18581 self.scale_inplace(&mut y0, ws0, m * out0)?;
18582 }
18583 if ws1 != 1.0 {
18584 self.scale_inplace(&mut y1, ws1, m * out1)?;
18585 }
18586 if ws2 != 1.0 {
18587 self.scale_inplace(&mut y2, ws2, m * out2)?;
18588 }
18589 Ok((y0, y1, y2))
18590 }
18591
18592 #[allow(clippy::too_many_arguments)] pub fn qmatvec_e4m3_blk_mmvq(
18603 &self,
18604 bytes: &CudaSlice<u8>,
18605 aq: &CudaSlice<i8>,
18606 ad: &CudaSlice<f32>,
18607 scales: &CudaSlice<f32>,
18608 m: usize,
18609 in_f: usize,
18610 out_f: usize,
18611 row_bytes: usize,
18612 scale_cols: usize,
18613 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18614 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_e4m3_blk_mmvq_into(
18616 bytes, aq, ad, scales, m, in_f, out_f, row_bytes, scale_cols, &mut y,
18617 )?;
18618 Ok(y)
18619 }
18620
18621 #[allow(clippy::too_many_arguments)]
18623 pub fn qmatvec_e4m3_blk_mmvq_into(
18624 &self,
18625 bytes: &CudaSlice<u8>,
18626 aq: &CudaSlice<i8>,
18627 ad: &CudaSlice<f32>,
18628 scales: &CudaSlice<f32>,
18629 m: usize,
18630 in_f: usize,
18631 out_f: usize,
18632 row_bytes: usize,
18633 scale_cols: usize,
18634 y: &mut CudaSlice<f32>,
18635 ) -> Result<(), Box<dyn std::error::Error>> {
18636 const ROWS_PER_BLOCK: u32 = 4; let f = self.func("qmatvec_e4m3_blk_mmvq");
18638 let cfg = LaunchConfig {
18639 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), m as u32, 1),
18640 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
18643 let (inf, outf, mi, rb, sc) = (
18644 in_f as i32,
18645 out_f as i32,
18646 m as i32,
18647 row_bytes as i64,
18648 scale_cols as i32,
18649 );
18650 let __s_b = self.gpu.stream();
18651 let mut b = __s_b.launch_builder(&f);
18652 b.arg(bytes)
18653 .arg(aq)
18654 .arg(ad)
18655 .arg(scales)
18656 .arg(&mut *y)
18657 .arg(&inf)
18658 .arg(&outf)
18659 .arg(&mi)
18660 .arg(&rb)
18661 .arg(&sc);
18662 unsafe {
18663 b.launch(cfg)?;
18664 }
18665 Ok(())
18666 }
18667
18668 #[allow(clippy::too_many_arguments)]
18674 pub fn qmatvec_e4m3_blk_mmvq_batched(
18675 &self,
18676 bytes: &CudaSlice<u8>,
18677 aq: &CudaSlice<i8>,
18678 ad: &CudaSlice<f32>,
18679 scales: &CudaSlice<f32>,
18680 m: usize,
18681 in_f: usize,
18682 out_f: usize,
18683 row_bytes: usize,
18684 scale_cols: usize,
18685 mcols: usize,
18686 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18687 const ROWS_PER_BLOCK: u32 = 4; debug_assert!(mcols >= m, "blk batched: mcols {mcols} < m {m}");
18689 let name = match mcols {
18690 2 => "qmatvec_e4m3_blk_mmvq_b2",
18691 4 => "qmatvec_e4m3_blk_mmvq_b4",
18692 8 => "qmatvec_e4m3_blk_mmvq_b8",
18693 16 => "qmatvec_e4m3_blk_mmvq_b16",
18694 _ => {
18695 return Err(
18696 format!("qmatvec_e4m3_blk_mmvq_batched: no kernel for mcols {mcols}").into(),
18697 );
18698 }
18699 };
18700 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
18701 let f = self.func(name);
18702 let cfg = LaunchConfig {
18703 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
18704 block_dim: (32, ROWS_PER_BLOCK, 1),
18705 shared_mem_bytes: 0,
18706 };
18707 let (inf, outf, mi, rb, sc) = (
18708 in_f as i32,
18709 out_f as i32,
18710 m as i32,
18711 row_bytes as i64,
18712 scale_cols as i32,
18713 );
18714 let __s_b = self.gpu.stream();
18715 let mut b = __s_b.launch_builder(&f);
18716 b.arg(bytes)
18717 .arg(aq)
18718 .arg(ad)
18719 .arg(scales)
18720 .arg(&mut y)
18721 .arg(&inf)
18722 .arg(&outf)
18723 .arg(&mi)
18724 .arg(&rb)
18725 .arg(&sc);
18726 unsafe {
18727 b.launch(cfg)?;
18728 }
18729 Ok(y)
18730 }
18731
18732 #[allow(clippy::too_many_arguments)]
18735 pub fn qmatvec_e4m3_blk_batched_raw(
18736 &self,
18737 bytes: &CudaSlice<u8>,
18738 x: &CudaSlice<f32>,
18739 scales: &CudaSlice<f32>,
18740 m: usize,
18741 in_f: usize,
18742 out_f: usize,
18743 row_bytes: usize,
18744 scale_cols: usize,
18745 mcols: usize,
18746 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18747 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18748 self.qmatvec_e4m3_blk_mmvq_batched(
18749 bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols, mcols,
18750 )
18751 }
18752
18753 #[allow(clippy::too_many_arguments)]
18756 pub fn qmatvec_e4m3_blk_mmvq_raw(
18757 &self,
18758 bytes: &CudaSlice<u8>,
18759 x: &CudaSlice<f32>,
18760 scales: &CudaSlice<f32>,
18761 m: usize,
18762 in_f: usize,
18763 out_f: usize,
18764 row_bytes: usize,
18765 scale_cols: usize,
18766 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18767 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18768 self.qmatvec_e4m3_blk_mmvq(
18769 bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols,
18770 )
18771 }
18772
18773 #[allow(clippy::too_many_arguments)]
18776 pub fn qmatvec_e4m3_fused2_raw(
18777 &self,
18778 b0: &CudaSlice<u8>,
18779 b1: &CudaSlice<u8>,
18780 x: &CudaSlice<f32>,
18781 in_f: usize,
18782 out0: usize,
18783 out1: usize,
18784 row_bytes: usize,
18785 ws0: f32,
18786 ws1: f32,
18787 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18788 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
18789 self.e4m3_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes, ws0, ws1)
18790 }
18791
18792 #[allow(clippy::too_many_arguments)]
18793 #[allow(clippy::type_complexity)] pub fn qmatvec_e4m3_fused3_raw(
18795 &self,
18796 b0: &CudaSlice<u8>,
18797 b1: &CudaSlice<u8>,
18798 b2: &CudaSlice<u8>,
18799 x: &CudaSlice<f32>,
18800 in_f: usize,
18801 out0: usize,
18802 out1: usize,
18803 out2: usize,
18804 row_bytes: usize,
18805 ws0: f32,
18806 ws1: f32,
18807 ws2: f32,
18808 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18809 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
18810 self.e4m3_fused3_core(
18811 b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes, ws0, ws1, ws2,
18812 )
18813 }
18814
18815 #[allow(clippy::too_many_arguments)]
18816 pub fn qmatvec_e4m3_fused2_t_raw(
18817 &self,
18818 b0: &CudaSlice<u8>,
18819 b1: &CudaSlice<u8>,
18820 x: &CudaSlice<f32>,
18821 m: usize,
18822 in_f: usize,
18823 out0: usize,
18824 out1: usize,
18825 row_bytes: usize,
18826 ws0: f32,
18827 ws1: f32,
18828 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18829 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18830 self.e4m3_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes, ws0, ws1)
18831 }
18832
18833 #[allow(clippy::too_many_arguments)]
18834 #[allow(clippy::type_complexity)] pub fn qmatvec_e4m3_fused3_t_raw(
18836 &self,
18837 b0: &CudaSlice<u8>,
18838 b1: &CudaSlice<u8>,
18839 b2: &CudaSlice<u8>,
18840 x: &CudaSlice<f32>,
18841 m: usize,
18842 in_f: usize,
18843 out0: usize,
18844 out1: usize,
18845 out2: usize,
18846 row_bytes: usize,
18847 ws0: f32,
18848 ws1: f32,
18849 ws2: f32,
18850 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18851 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18852 self.e4m3_fused3_t_core(
18853 b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes, ws0, ws1, ws2,
18854 )
18855 }
18856
18857 fn try_e4m3_blk_pre(
18868 &self,
18869 w: &crate::model::GpuTensor,
18870 aq: &CudaSlice<i8>,
18871 ad: &CudaSlice<f32>,
18872 m: usize,
18873 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
18874 use crate::model::GpuTensor;
18875 if let GpuTensor::Quant {
18876 bytes,
18877 qtype,
18878 row_bytes,
18879 blk: Some(g),
18880 ..
18881 } = w
18882 && *qtype == QT_F8_E4M3_BLK
18883 {
18884 if (2..=16).contains(&m)
18890 && std::env::var("MEMRA_NO_BATCHED").is_err()
18891 && (m <= 4 || Self::b8_enabled())
18892 {
18893 let mcols = Self::batched_mcols(m);
18894 return Ok(Some(self.qmatvec_e4m3_blk_mmvq_batched(
18895 bytes,
18896 aq,
18897 ad,
18898 &g.scales,
18899 m,
18900 w.in_features(),
18901 w.out_features(),
18902 *row_bytes,
18903 g.cols,
18904 mcols,
18905 )?));
18906 }
18907 return Ok(Some(self.qmatvec_e4m3_blk_mmvq(
18908 bytes,
18909 aq,
18910 ad,
18911 &g.scales,
18912 m,
18913 w.in_features(),
18914 w.out_features(),
18915 *row_bytes,
18916 g.cols,
18917 )?));
18918 }
18919 Ok(None)
18920 }
18921
18922 fn try_e4m3_blk_prefill(
18969 &self,
18970 w: &crate::model::GpuTensor,
18971 x: &CudaSlice<f32>,
18972 m: usize,
18973 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
18974 use crate::model::GpuTensor;
18975 let GpuTensor::Quant {
18976 bytes,
18977 qtype,
18978 blk: Some(g),
18979 ..
18980 } = w
18981 else {
18982 return Ok(None);
18983 };
18984 if *qtype != QT_F8_E4M3_BLK {
18985 return Ok(None);
18986 }
18987 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? {
18992 return Ok(Some(y));
18993 }
18994 let (in_f, out_f) = (w.in_features(), w.out_features());
18995 let slab = self.fp8_blk_dequant_q8_0_dev(bytes, &g.scales, out_f, in_f)?;
18996 let tmp = GpuTensor::Quant {
18997 bytes: slab,
18998 qtype: QT_Q8_0,
18999 row_bytes: in_f / 32 * 34,
19000 ne: vec![in_f as u64, out_f as u64],
19001 scale: 1.0,
19002 rp: false,
19003 #[cfg(memra_cutlass)]
19004 cutlass: None,
19005 fp8: None,
19006 blk: None,
19007 f16: None,
19008 rp4: None,
19009 };
19010 Ok(Some(self.matmul(&tmp, x, m)?))
19012 }
19013
19014 #[allow(clippy::type_complexity)] pub fn matmul_pre_noscale(
19016 &self,
19017 w: &crate::model::GpuTensor,
19018 aq: &CudaSlice<i8>,
19019 ad: &CudaSlice<f32>,
19020 m: usize,
19021 ) -> Result<Option<(CudaSlice<f32>, f32)>, Box<dyn std::error::Error>> {
19022 use crate::model::GpuTensor;
19023 if m == 1
19027 && let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)?
19028 {
19029 return Ok(Some((y, 1.0)));
19030 }
19031 if m != 1 || !self.uses_q8_1_fast(w) {
19033 return Ok(None);
19034 }
19035 let in_f = w.in_features();
19036 let out_f = w.out_features();
19037 let (bytes, qtype, row_bytes, scale, rp) = match w {
19038 GpuTensor::Quant {
19039 bytes,
19040 qtype,
19041 row_bytes,
19042 scale,
19043 rp,
19044 ..
19045 } => (bytes, *qtype, *row_bytes, *scale, *rp),
19046 _ => return Ok(None),
19047 };
19048 if self.mmvq_supports(qtype) {
19050 let (mbytes, mrp) = match w {
19052 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
19053 _ => (bytes, rp),
19054 };
19055 let y = self.qmatvec_mmvq(
19056 mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, 1.0, mrp,
19057 )?;
19058 return Ok(Some((y, scale)));
19059 }
19060 let name = match qtype {
19062 QT_Q8_0 => "qmatvec_q8_0_dp4a",
19063 QT_Q4_K => "qmatvec_q4_K_dp4a",
19064 QT_Q6_K => "qmatvec_q6_K_dp4a",
19065 QT_Q5_K => "qmatvec_q5_K_dp4a",
19066 QT_Q3_K => "qmatvec_q3_K_dp4a",
19067 QT_NVFP4 => {
19068 if rp {
19069 "qmatvec_nvfp4_dp4a_rp"
19070 } else {
19071 "qmatvec_nvfp4_dp4a"
19072 }
19073 }
19074 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
19075 _ => return Ok(None),
19076 };
19077 let f = self.func(name);
19078 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
19079 let cfg = LaunchConfig {
19080 grid_dim: (out_f as u32, m as u32, 1),
19081 block_dim: (128, 1, 1),
19082 shared_mem_bytes: 0,
19083 };
19084 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
19085 let __s_b = self.gpu.stream();
19086 let mut b = __s_b.launch_builder(&f);
19087 b.arg(bytes)
19088 .arg(aq)
19089 .arg(ad)
19090 .arg(&mut y)
19091 .arg(&inf)
19092 .arg(&outf)
19093 .arg(&mi)
19094 .arg(&rb);
19095 unsafe {
19096 b.launch(cfg)?;
19097 }
19098 Ok(Some((y, scale)))
19099 }
19100
19101 pub fn mmvq_supports(&self, qtype: i32) -> bool {
19104 if qtype == QT_F8_E4M3 {
19109 return true;
19110 }
19111 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") {
19112 return false;
19113 }
19114 matches!(
19115 qtype,
19116 QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_Q4_0
19117 )
19118 }
19119
19120 #[allow(clippy::too_many_arguments)] pub fn qmatvec_mmvq(
19126 &self,
19127 bytes: &CudaSlice<u8>,
19128 aq: &CudaSlice<i8>,
19129 ad: &CudaSlice<f32>,
19130 m: usize,
19131 in_f: usize,
19132 out_f: usize,
19133 qtype: i32,
19134 row_bytes: usize,
19135 scale: f32,
19136 rp: bool,
19137 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19138 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_mmvq_into(
19140 bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp, &mut y,
19141 )?;
19142 Ok(y)
19143 }
19144
19145 #[allow(clippy::too_many_arguments)]
19147 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_mmvq_into(
19149 &self,
19150 bytes: &CudaSlice<u8>,
19151 aq: &CudaSlice<i8>,
19152 ad: &CudaSlice<f32>,
19153 m: usize,
19154 in_f: usize,
19155 out_f: usize,
19156 qtype: i32,
19157 row_bytes: usize,
19158 scale: f32,
19159 rp: bool,
19160 y: &mut CudaSlice<f32>,
19161 ) -> Result<(), Box<dyn std::error::Error>> {
19162 debug_assert!(y.len() >= m * out_f);
19163 const ROWS_PER_BLOCK: u32 = 4; if qtype == QT_Q8_0
19169 && rp
19170 && m == 1
19171 && out_f >= 64
19172 && (out_f as u32).div_ceil(ROWS_PER_BLOCK) < 4 * self.sm_count() as u32
19173 && {
19174 static G2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19175 *G2.get_or_init(|| std::env::var("MEMRA_Q80_G2").as_deref() != Ok("0"))
19176 }
19177 {
19178 let f = self.func("qmatvec_q8_0_mmvq_rp_g2");
19179 let cfg = LaunchConfig {
19180 grid_dim: ((out_f as u32).div_ceil(2), 1, 1),
19181 block_dim: (32, 2, 1),
19182 shared_mem_bytes: 0,
19183 };
19184 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, row_bytes as i64);
19185 let __s_b = self.gpu.stream();
19186 let mut b = __s_b.launch_builder(&f);
19187 b.arg(bytes)
19188 .arg(aq)
19189 .arg(ad)
19190 .arg(&mut *y)
19191 .arg(&inf)
19192 .arg(&outf)
19193 .arg(&mi)
19194 .arg(&rb);
19195 unsafe {
19196 b.launch(cfg)?;
19197 }
19198 if scale != 1.0 {
19199 self.scale_inplace(y, scale, out_f)?;
19200 }
19201 return Ok(());
19202 }
19203 let mut mr: u32 = if m == 1 && (qtype == QT_NVFP4 || qtype == QT_Q5_K) {
19212 2
19213 } else {
19214 1
19215 };
19216 if m == 1 && qtype == QT_Q4_0 {
19221 static Q40MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
19222 mr = *Q40MR.get_or_init(|| {
19225 std::env::var("MEMRA_Q40_MR")
19226 .ok()
19227 .and_then(|v| v.parse().ok())
19228 .unwrap_or(1)
19229 });
19230 }
19231 let q5_mode = std::env::var("MEMRA_Q5K_ISSUE").ok();
19242 let q5_force = q5_mode.as_deref() == Some("2");
19243 let q5_il = qtype == QT_Q5_K
19246 && m == 1
19247 && (q5_force || q5_mode.as_deref().map(|v| v != "0").unwrap_or(true));
19248 if q5_il && !q5_force && out_f > 65536 {
19249 mr = 1;
19250 }
19251 if qtype == QT_Q4_0 && rp && mr != 1 {
19254 mr = 2;
19255 }
19256 if qtype == QT_Q8_0 && rp {
19260 static Q80MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
19261 mr = *Q80MR.get_or_init(|| {
19262 std::env::var("MEMRA_Q80_MR")
19263 .ok()
19264 .and_then(|v| v.parse().ok())
19265 .unwrap_or(1)
19266 });
19267 }
19268 let name = match (qtype, mr, rp) {
19269 (QT_NVFP4, 2, false) => "qmatvec_nvfp4_mmvq_mr2",
19270 (QT_NVFP4, 2, true) => "qmatvec_nvfp4_mmvq_mr2_rp",
19271 (QT_NVFP4, _, true) => "qmatvec_nvfp4_mmvq_rp",
19272 (QT_Q4_0, 1, true) => "qmatvec_q4_0_mmvq_rp",
19273 (QT_Q4_0, _, true) => "qmatvec_q4_0_mmvq_mr2_rp",
19274 (QT_Q5_K, 2, _) => {
19275 if q5_il {
19276 "qmatvec_q5_K_mmvq_mr2_il"
19277 } else {
19278 "qmatvec_q5_K_mmvq_mr2"
19279 }
19280 }
19281 (QT_Q8_0, 2, true) => "qmatvec_q8_0_mmvq_mr2_rp",
19282 (QT_Q8_0, _, true)
19287 if in_f.is_multiple_of(1024) && {
19288 static CA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19289 *CA.get_or_init(|| std::env::var("MEMRA_Q80_CA").as_deref() == Ok("1"))
19290 } =>
19291 {
19292 "qmatvec_q8_0_mmvq_rpca"
19293 }
19294 (QT_Q8_0, _, true) => "qmatvec_q8_0_mmvq_rp",
19295 (QT_Q8_0, _, _) => "qmatvec_q8_0_mmvq",
19296 (QT_Q4_K, _, true) => "qmatvec_q4_K_mmvq_rp",
19300 (QT_Q6_K, _, true) => "qmatvec_q6_K_mmvq_rp",
19301 (QT_Q4_K, _, _) => "qmatvec_q4_K_mmvq",
19302 (QT_Q4_0, 2, false) => "qmatvec_q4_0_mmvq_mr2",
19303 (QT_Q4_0, _, false) => "qmatvec_q4_0_mmvq",
19304 (QT_Q5_K, _, _) => {
19305 if q5_il {
19306 "qmatvec_q5_K_mmvq_il"
19307 } else {
19308 "qmatvec_q5_K_mmvq"
19309 }
19310 }
19311 (QT_Q6_K, _, _) => "qmatvec_q6_K_mmvq",
19312 (QT_NVFP4, _, false) => "qmatvec_nvfp4_mmvq",
19313 (QT_F8_E4M3, _, _) => "qmatvec_e4m3_mmvq",
19314 _ => panic!("qmatvec_mmvq: qtype {qtype} has no MMVQ kernel"),
19315 };
19316 let f = self.func(name);
19317 let rows_per_block = ROWS_PER_BLOCK * mr;
19319 let cfg = LaunchConfig {
19320 grid_dim: (
19321 (out_f as u32 + rows_per_block - 1) / rows_per_block,
19322 m as u32,
19323 1,
19324 ),
19325 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
19328 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
19329 let __s_b = self.gpu.stream();
19330 let mut b = __s_b.launch_builder(&f);
19331 if qtype == QT_NVFP4 || qtype == QT_F8_E4M3 {
19336 if Self::pdl_on()
19339 && Self::pdl_mmvq_on()
19340 && Self::pdl_nvfp4q8_on()
19341 && name == "qmatvec_nvfp4_mmvq_mr2_rp"
19342 {
19343 use cudarc::driver::{DevicePtr, DevicePtrMut};
19344 let s = &self.gpu.stream();
19345 let (pw, _g0) = bytes.device_ptr(s);
19346 let (paq, _g1) = aq.device_ptr(s);
19347 let (pad, _g2) = ad.device_ptr(s);
19348 let (py, _g3) = y.device_ptr_mut(s);
19349 let mut ps = [
19350 &pw as *const _ as *mut std::ffi::c_void,
19351 &paq as *const _ as *mut _,
19352 &pad as *const _ as *mut _,
19353 &py as *const _ as *mut _,
19354 &inf as *const _ as *mut _,
19355 &outf as *const _ as *mut _,
19356 &mi as *const _ as *mut _,
19357 &rb as *const _ as *mut _,
19358 &scale as *const _ as *mut _,
19359 ];
19360 unsafe {
19361 self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?;
19362 }
19363 return Ok(());
19364 }
19365 b.arg(bytes)
19366 .arg(aq)
19367 .arg(ad)
19368 .arg(&mut *y)
19369 .arg(&inf)
19370 .arg(&outf)
19371 .arg(&mi)
19372 .arg(&rb)
19373 .arg(&scale);
19374 unsafe {
19375 b.launch(cfg)?;
19376 }
19377 } else if Self::pdl_on()
19378 && Self::pdl_mmvq_on()
19379 && (matches!(
19380 name,
19381 "qmatvec_q4_0_mmvq_rp" | "qmatvec_q6_K_mmvq" | "qmatvec_q6_K_mmvq_rp"
19382 ) || (Self::pdl_nvfp4q8_on()
19383 && matches!(name, "qmatvec_q8_0_mmvq_rp" | "qmatvec_q8_0_mmvq_mr2_rp")))
19384 {
19385 {
19389 use cudarc::driver::{DevicePtr, DevicePtrMut};
19390 let s = &self.gpu.stream();
19391 let (pw, _g0) = bytes.device_ptr(s);
19392 let (paq, _g1) = aq.device_ptr(s);
19393 let (pad, _g2) = ad.device_ptr(s);
19394 let (py, _g3) = y.device_ptr_mut(s);
19395 let mut ps = [
19396 &pw as *const _ as *mut std::ffi::c_void,
19397 &paq as *const _ as *mut _,
19398 &pad as *const _ as *mut _,
19399 &py as *const _ as *mut _,
19400 &inf as *const _ as *mut _,
19401 &outf as *const _ as *mut _,
19402 &mi as *const _ as *mut _,
19403 &rb as *const _ as *mut _,
19404 ];
19405 unsafe {
19406 self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?;
19407 }
19408 }
19409 if scale != 1.0 {
19410 self.scale_inplace(y, scale, m * out_f)?;
19411 }
19412 } else {
19413 b.arg(bytes)
19414 .arg(aq)
19415 .arg(ad)
19416 .arg(&mut *y)
19417 .arg(&inf)
19418 .arg(&outf)
19419 .arg(&mi)
19420 .arg(&rb);
19421 unsafe {
19422 b.launch(cfg)?;
19423 }
19424 if scale != 1.0 {
19425 self.scale_inplace(y, scale, m * out_f)?;
19426 }
19427 }
19428 Ok(())
19429 }
19430
19431 #[allow(clippy::too_many_arguments)] pub fn qmatvec_mmvq_raw(
19436 &self,
19437 bytes: &CudaSlice<u8>,
19438 x: &CudaSlice<f32>,
19439 m: usize,
19440 in_f: usize,
19441 out_f: usize,
19442 qtype: i32,
19443 row_bytes: usize,
19444 rp: bool,
19445 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19446 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
19447 self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, 1.0, rp)
19448 }
19449
19450 pub fn batched_supports(&self, qtype: i32) -> bool {
19454 matches!(
19455 qtype,
19456 QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_F8_E4M3 | QT_Q4_0
19457 )
19458 }
19459
19460 pub fn iq_fast_enabled() -> bool {
19468 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19469 *ON.get_or_init(|| {
19470 std::env::var("MEMRA_IQ_FAST")
19471 .map(|v| v != "0")
19472 .unwrap_or(true)
19473 })
19474 }
19475
19476 pub fn b8_enabled() -> bool {
19479 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19480 *ON.get_or_init(|| std::env::var("MEMRA_B8").map(|v| v != "0").unwrap_or(true))
19481 }
19482
19483 pub fn batched_mcols(m: usize) -> usize {
19485 if m == 2 {
19486 2
19487 } else if m <= 4 {
19488 4
19489 } else if m <= 8 {
19490 8
19491 } else {
19492 16
19493 }
19494 }
19495
19496 fn batched_kernel_name(qtype: i32, mcols: usize) -> Option<&'static str> {
19501 Some(match (qtype, mcols) {
19502 (QT_Q8_0, 2) => "qmatvec_q8_0_mmvq_b2",
19503 (QT_Q8_0, 4) => "qmatvec_q8_0_mmvq_b4",
19504 (QT_Q8_0, 8) => "qmatvec_q8_0_mmvq_b8",
19505 (QT_Q8_0, 16) => "qmatvec_q8_0_mmvq_b16",
19511 (QT_Q4_K, 2) => "qmatvec_q4_K_mmvq_b2",
19512 (QT_Q4_K, 4) => "qmatvec_q4_K_mmvq_b4",
19513 (QT_Q4_K, 8) => "qmatvec_q4_K_mmvq_b8",
19514 (QT_Q4_K, 16) => "qmatvec_q4_K_mmvq_b16",
19517 (QT_Q5_K, 2) => "qmatvec_q5_K_mmvq_b2",
19518 (QT_Q5_K, 4) => "qmatvec_q5_K_mmvq_b4",
19519 (QT_Q5_K, 8) => "qmatvec_q5_K_mmvq_b8",
19520 (QT_Q5_K, 16) => "qmatvec_q5_K_mmvq_b16",
19523 (QT_Q6_K, 2) => "qmatvec_q6_K_mmvq_b2",
19524 (QT_Q6_K, 4) => "qmatvec_q6_K_mmvq_b4",
19525 (QT_Q6_K, 8) => "qmatvec_q6_K_mmvq_b8",
19526 (QT_Q6_K, 16) => "qmatvec_q6_K_mmvq_b16",
19527 (QT_NVFP4, 2) => "qmatvec_nvfp4_mmvq_b2",
19528 (QT_NVFP4, 4) => "qmatvec_nvfp4_mmvq_b4",
19529 (QT_NVFP4, 8) => "qmatvec_nvfp4_mmvq_b8",
19530 (QT_NVFP4, 16) => "qmatvec_nvfp4_mmvq_b16",
19534 (QT_F8_E4M3, 2) => "qmatvec_e4m3_mmvq_b2",
19535 (QT_F8_E4M3, 4) => "qmatvec_e4m3_mmvq_b4",
19536 (QT_F8_E4M3, 8) => "qmatvec_e4m3_mmvq_b8",
19537 (QT_F8_E4M3, 16) => "qmatvec_e4m3_mmvq_b16",
19541 (QT_Q4_0, 2) => "qmatvec_q4_0_mmvq_b2",
19542 (QT_Q4_0, 4) => "qmatvec_q4_0_mmvq_b4",
19543 (QT_Q4_0, 8) => "qmatvec_q4_0_mmvq_b8",
19544 (QT_Q4_0, 16) => "qmatvec_q4_0_mmvq_b16",
19545 _ => return None,
19546 })
19547 }
19548
19549 pub fn sm_count(&self) -> i32 {
19584 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
19585 *SMS.get_or_init(|| {
19586 use cudarc::driver::sys::CUdevice_attribute_enum as A;
19587 self.gpu
19588 .ctx
19589 .attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT)
19590 .unwrap_or(82)
19591 })
19592 }
19593
19594 #[allow(clippy::too_many_arguments)]
19595 #[allow(clippy::if_same_then_else)] pub fn batched_variant(
19598 &self,
19599 _m: usize,
19600 in_f: usize,
19601 out_f: usize,
19602 qtype: i32,
19603 row_bytes: usize,
19604 mcols: usize,
19605 rp: bool,
19606 ) -> &'static str {
19607 if qtype == QT_Q8_0 {
19612 return if rp { "rp" } else { "base" };
19613 }
19614 static BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
19615 let bv = *BV.get_or_init(|| match std::env::var("MEMRA_MMVQ_BV").as_deref() {
19616 Ok("base") => "base",
19617 Ok("pf") => "pf",
19618 Ok("r2") => "r2",
19619 Ok("r2w8") => "r2w8",
19620 Ok("pfr2") => "pfr2",
19621 Ok("ca") => "ca",
19622 Ok("car2") => "car2",
19623 Ok("rp") => "rp",
19626 Ok("rpr2") => "rpr2",
19627 Ok("rpr2w8") => "rpr2w8",
19628 Ok("rpca") => "rpca",
19631 Ok("rpcar2") => "rpcar2",
19632 Ok("rpsc") => "rpsc",
19639 Ok("rpms") => "rpms",
19640 Ok("rpmsc") => "rpmsc",
19641 Ok("rpks") => "rpks",
19642 Ok("rpksc") => "rpksc",
19643 _ => "auto",
19644 });
19645 let ca_ok = qtype == QT_NVFP4 && row_bytes.is_multiple_of(16) && in_f.is_multiple_of(1024);
19649 static KS_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19654 let ks_on = *KS_ON.get_or_init(|| std::env::var("MEMRA_KS").as_deref() != Ok("0"));
19655 let sc_ok = ks_on && qtype == QT_NVFP4 && in_f.is_multiple_of(256) && (in_f / 64 <= 272);
19656 let ks_ok = ks_on && qtype == QT_NVFP4 && in_f.is_multiple_of(512) && (in_f / 64 <= 272);
19657 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
19658 let sms = *SMS.get_or_init(|| {
19659 use cudarc::driver::sys::CUdevice_attribute_enum as A;
19660 self.gpu
19661 .ctx
19662 .attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT)
19663 .unwrap_or(82)
19664 });
19665 let kq_r2 = matches!(qtype, QT_Q4_K | QT_Q5_K | QT_Q6_K);
19685 static KQBV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
19688 let kq_bv = *KQBV.get_or_init(|| match std::env::var("MEMRA_KQ_BV").as_deref() {
19689 Ok("base") => "base",
19690 Ok("r2") => "r2",
19691 Ok("r2w8") => "r2w8",
19692 _ => "auto",
19693 });
19694 let variant: &'static str = if qtype == QT_Q4_0 {
19695 static Q40BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
19699 let q40 = *Q40BV.get_or_init(|| match std::env::var("MEMRA_Q40_BV").as_deref() {
19700 Ok("base") => "base",
19706 Ok("r2") => "r2",
19707 Ok("ms") => "ms",
19708 Ok("sm") => "sm",
19709 Ok("la") => "la",
19710 _ => "auto",
19711 });
19712 let v = if q40 != "auto" {
19713 q40
19714 } else if (out_f as u32).div_ceil(8) >= 4 * sms as u32 {
19715 "r2"
19716 } else {
19717 "base"
19718 };
19719 if rp {
19724 match v {
19725 "ms" => "r2ms_rp",
19726 "sm" => "r2sm_rp",
19727 "la" => "r2la_rp",
19728 "r2" => "r2_rp",
19729 _ => "rp",
19730 }
19731 } else if matches!(v, "ms" | "sm" | "la") {
19732 "r2"
19733 } else {
19734 v
19735 }
19736 } else if qtype != QT_NVFP4 && !kq_r2 {
19737 "base"
19738 } else if kq_r2 && rp {
19739 "rp"
19743 } else if kq_r2 {
19744 if kq_bv != "auto" {
19747 if kq_bv == "r2w8" && mcols != 4 {
19748 "r2"
19749 } else {
19750 kq_bv
19751 }
19752 } else if bv != "auto" {
19753 match bv {
19754 "r2" | "pfr2" | "rpr2" | "car2" => "r2",
19755 "r2w8" | "rpr2w8" => {
19756 if mcols != 4 {
19757 "r2"
19758 } else {
19759 "r2w8"
19760 }
19761 }
19762 _ => "base", }
19764 } else {
19765 #[allow(clippy::manual_div_ceil)]
19766 let blocks = (out_f + 7) / 8;
19768 let waves = blocks as f64 / (7 * sms as usize) as f64;
19769 let filled = blocks >= 4 * sms as usize;
19770 let use_r2 = if qtype == QT_Q4_K {
19771 filled
19772 } else {
19773 waves >= 2.0
19774 };
19775 if use_r2 { "r2" } else { "base" }
19776 }
19777 } else if bv != "auto" {
19778 let v = if bv == "r2w8" && mcols == 2 {
19783 "r2"
19784 } else if bv == "ca" && (!ca_ok || mcols == 8) {
19785 "pf"
19786 } else if bv == "car2" && (!ca_ok || mcols == 8) {
19787 "r2"
19788 } else if bv == "pfr2" && mcols == 8 {
19789 "r2"
19790 } else if (bv == "rpr2w8" || bv == "rpr2") && mcols == 2 {
19791 "rpr2"
19792 }
19793 else if (bv == "rpca" || bv == "rpcar2") && (!ca_ok || mcols == 8) {
19795 if mcols == 8 { "rpr2w8" } else { "rpr2" }
19796 } else if bv == "rpcar2" && mcols == 2 {
19797 "rpca"
19798 }
19799 else if (bv == "rpsc" || bv == "rpmsc") && !sc_ok {
19802 "rpr2"
19803 } else if (bv == "rpks" || bv == "rpksc") && !ks_ok {
19804 "rpr2"
19805 } else {
19806 bv
19807 };
19808 if rp {
19809 match v {
19810 "base" | "pf" | "ca" | "rp" => "rp",
19811 "r2" | "pfr2" | "car2" | "rpr2" => "rpr2",
19812 "r2w8" | "rpr2w8" => {
19813 if mcols == 2 {
19814 "rpr2"
19815 } else {
19816 "rpr2w8"
19817 }
19818 }
19819 other => other, }
19821 } else {
19822 v
19823 }
19824 } else if mcols == 8 {
19825 if rp {
19836 if sc_ok { "rpsc" } else { "rpr2w8" }
19837 } else {
19838 "r2w8"
19839 }
19840 } else if mcols >= 4 {
19841 #[allow(clippy::manual_div_ceil)]
19845 let blocks = (out_f + 7) / 8;
19847 let r7 = 7 * sms as usize;
19848 let r8 = 8 * sms as usize;
19849 let waves = blocks as f64 / r7 as f64;
19850 let filled = blocks >= 4 * sms as usize;
19851 if filled && blocks.div_ceil(r8) < blocks.div_ceil(r7) {
19855 if rp { "rpr2w8" } else { "r2w8" }
19859 } else if waves >= 2.0 || (waves <= 1.0 && filled) {
19860 if rp { "rpr2" } else { "r2" }
19863 } else {
19864 if rp { "rp" } else { "pf" }
19868 }
19869 } else if in_f >= 6144 {
19870 if rp { "rpr2" } else { "r2" }
19874 } else if rp {
19875 #[allow(clippy::manual_div_ceil)]
19880 let waves = ((out_f + 7) / 8) as f64 / (7 * sms as usize) as f64;
19882 if sc_ok && (0.9..=1.1).contains(&waves) {
19883 "rpsc"
19884 } else {
19885 "rp"
19886 }
19887 } else {
19888 "base"
19889 };
19890 variant
19891 }
19892
19893 #[allow(clippy::too_many_arguments)]
19894 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_mmvq_batched(
19897 &self,
19898 bytes: &CudaSlice<u8>,
19899 aq: &CudaSlice<i8>,
19900 ad: &CudaSlice<f32>,
19901 m: usize,
19902 in_f: usize,
19903 out_f: usize,
19904 qtype: i32,
19905 row_bytes: usize,
19906 mcols: usize,
19907 scale: f32,
19908 rp: bool,
19909 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19910 const ROWS_PER_BLOCK: u32 = 4;
19911 let forced: Option<&'static str> = {
19916 static V: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
19917 V.get_or_init(|| std::env::var("MEMRA_BVAR").ok())
19918 .as_deref()
19919 .map(|s| Box::leak(s.to_string().into_boxed_str()) as &'static str)
19920 };
19921 let variant = match forced {
19922 Some(v) if !rp || v.contains("rp") => v,
19923 _ => self.batched_variant(m, in_f, out_f, qtype, row_bytes, mcols, rp),
19924 };
19925 let base_name = Self::batched_kernel_name(qtype, mcols).ok_or_else(|| {
19926 format!("qmatvec_mmvq_batched: no kernel for qtype {qtype} mcols {mcols}")
19927 })?;
19928 let variant = if mcols == 16 {
19932 if rp { "rp" } else { "base" }
19933 } else {
19934 variant
19935 };
19936 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19943 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
19944 if b567
19945 && qtype == QT_NVFP4
19946 && rp
19947 && mcols == 8
19948 && (5..=7).contains(&m)
19949 && matches!(variant, "rpsc" | "rpr2w8")
19950 {
19951 let f = self.func(&format!("qmatvec_nvfp4_mmvq_b{m}_{variant}"));
19952 let rows_per_block = ROWS_PER_BLOCK * 2; let mut y = self.alloc_uninit::<f32>(m * out_f)?;
19954 let cfg = LaunchConfig {
19955 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
19956 block_dim: (32, ROWS_PER_BLOCK, 1),
19957 shared_mem_bytes: 0,
19958 };
19959 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
19960 let __s_b = self.gpu.stream();
19961 let mut b = __s_b.launch_builder(&f);
19962 b.arg(bytes)
19963 .arg(aq)
19964 .arg(ad)
19965 .arg(&mut y)
19966 .arg(&inf)
19967 .arg(&outf)
19968 .arg(&mi)
19969 .arg(&rb);
19970 unsafe {
19971 b.launch(cfg)?;
19972 }
19973 if scale != 1.0 {
19974 self.scale_inplace(&mut y, scale, m * out_f)?;
19975 }
19976 return Ok(y);
19977 }
19978 let (name, rows_per_block): (std::borrow::Cow<'static, str>, u32) = match variant {
19979 "base" => (base_name.into(), ROWS_PER_BLOCK),
19980 "pf" => (format!("{base_name}_pf").into(), ROWS_PER_BLOCK),
19981 "ca" => (format!("{base_name}_ca").into(), ROWS_PER_BLOCK),
19982 "rp" => (format!("{base_name}_rp").into(), ROWS_PER_BLOCK),
19983 "rpca" => (format!("{base_name}_rpca").into(), ROWS_PER_BLOCK), "rpks" => (format!("{base_name}_rpks").into(), ROWS_PER_BLOCK),
19987 "rpksc" => (format!("{base_name}_rpksc").into(), ROWS_PER_BLOCK),
19988 "rpms" => (format!("{base_name}_rpms").into(), ROWS_PER_BLOCK),
19989 "rpmsc" => (format!("{base_name}_rpmsc").into(), ROWS_PER_BLOCK),
19990 "r2ms_rp" => (format!("{base_name}_r2ms_rp").into(), ROWS_PER_BLOCK),
19991 "r2sm_rp" => (format!("{base_name}_r2sm_rp").into(), ROWS_PER_BLOCK * 2),
19992 "r2la_rp" => (format!("{base_name}_r2la_rp").into(), ROWS_PER_BLOCK * 2),
19993 v => (format!("{base_name}_{v}").into(), ROWS_PER_BLOCK * 2), };
19995 debug_assert!(
19996 !rp || name.contains("_rp"),
19997 "rp weight dispatched to a GGUF-layout kernel"
19998 );
19999 let f = self.func(&name);
20000 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
20001 let smem = if name.contains("_r2sm_rp") {
20003 (mcols * 32 * 9 * 4 + mcols * 32 * 4) as u32
20004 } else {
20005 0
20006 };
20007 let cfg = LaunchConfig {
20008 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
20009 block_dim: (32, ROWS_PER_BLOCK, 1),
20010 shared_mem_bytes: smem,
20011 };
20012 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
20013 let __s_b = self.gpu.stream();
20014 let mut b = __s_b.launch_builder(&f);
20015 b.arg(bytes)
20016 .arg(aq)
20017 .arg(ad)
20018 .arg(&mut y)
20019 .arg(&inf)
20020 .arg(&outf)
20021 .arg(&mi)
20022 .arg(&rb);
20023 unsafe {
20024 b.launch(cfg)?;
20025 }
20026 if scale != 1.0 {
20027 self.scale_inplace(&mut y, scale, m * out_f)?;
20028 }
20029 Ok(y)
20030 }
20031
20032 #[allow(clippy::too_many_arguments)] pub fn qmatvec_batched_raw(
20037 &self,
20038 bytes: &CudaSlice<u8>,
20039 x: &CudaSlice<f32>,
20040 m: usize,
20041 in_f: usize,
20042 out_f: usize,
20043 qtype: i32,
20044 row_bytes: usize,
20045 mcols: usize,
20046 rp: bool,
20047 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20048 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
20049 self.qmatvec_mmvq_batched(
20050 bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, 1.0, rp,
20051 )
20052 }
20053
20054 #[allow(clippy::too_many_arguments)] pub fn qmatvec_nvfp4_batched_raw(
20057 &self,
20058 bytes: &CudaSlice<u8>,
20059 x: &CudaSlice<f32>,
20060 m: usize,
20061 in_f: usize,
20062 out_f: usize,
20063 row_bytes: usize,
20064 mcols: usize,
20065 rp: bool,
20066 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20067 self.qmatvec_batched_raw(bytes, x, m, in_f, out_f, QT_NVFP4, row_bytes, mcols, rp)
20068 }
20069
20070 fn try_fp4_gemm(
20074 &self,
20075 w: &crate::model::GpuTensor,
20076 x: &CudaSlice<f32>,
20077 m: usize,
20078 in_f: usize,
20079 out_f: usize,
20080 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
20081 use crate::model::GpuTensor;
20082 if cfg!(memra_portable_cuda) {
20083 return Ok(None);
20084 }
20085 if std::env::var("MEMRA_FP4").is_ok() {
20094 refuse_portable_force("MEMRA_FP4", "the sm_120a mxf4 block-scale MMA");
20095 assert!(
20096 konst_eq(env!("MEMRA_BUILT_CUDA_ARCH"), "120a"),
20097 "MEMRA_FP4 forces the native mxf4 block-scale GEMM (qmatvec_gemm_nvfp4_fp4), \
20098 which only the sm_120a fatbin contains — this is an sm_{} build. Unset \
20099 MEMRA_FP4; the W4A8 int8 path is the correct default for NVFP4 weights.",
20100 env!("MEMRA_BUILT_CUDA_ARCH")
20101 );
20102 }
20103 if std::env::var("MEMRA_FP4").is_err() {
20104 return Ok(None);
20105 }
20106 #[cfg(memra_cutlass)]
20115 if m >= 128 && std::env::var("MEMRA_FP4_CUTLASS").is_ok() {
20116 if let GpuTensor::Quant {
20117 bytes,
20118 qtype,
20119 scale,
20120 row_bytes,
20121 cutlass,
20122 ..
20123 } = w
20124 {
20125 if *qtype == QT_NVFP4 && in_f % 64 == 0 {
20126 if let Some(cw) = cutlass {
20127 let y = self.cutlass_fp4_gemm(
20129 &cw.b_packed,
20130 &cw.sfb_swizzled,
20131 x,
20132 *scale,
20133 m,
20134 out_f,
20135 in_f,
20136 )?;
20137 return Ok(Some(y));
20138 } else if std::env::var("MEMRA_FP4_CUTLASS_OTF").is_ok() {
20139 let (b_packed, sfb_sw) =
20144 self.build_cutlass_weight(bytes, out_f, in_f, *row_bytes)?;
20145 let y =
20146 self.cutlass_fp4_gemm(&b_packed, &sfb_sw, x, *scale, m, out_f, in_f)?;
20147 return Ok(Some(y));
20148 }
20149 }
20150 }
20151 }
20152 if let GpuTensor::Quant {
20153 bytes,
20154 qtype,
20155 row_bytes,
20156 scale,
20157 rp,
20158 ..
20159 } = w
20160 {
20161 if *qtype == QT_NVFP4 && in_f.is_multiple_of(64) && !*rp {
20164 let y =
20165 self.qmatvec_gemm_nvfp4_fp4(bytes, x, m, in_f, out_f, *row_bytes, *scale)?;
20166 return Ok(Some(y));
20167 }
20168 }
20169 Ok(None)
20170 }
20171
20172 #[allow(clippy::too_many_arguments)] pub fn rms_norm_f16out(
20177 &self,
20178 x: &CudaSlice<f32>,
20179 w: &CudaSlice<f32>,
20180 dst: &mut CudaSlice<f32>,
20181 dst16: &mut CudaSlice<u8>,
20182 ncols: usize,
20183 nrows: usize,
20184 eps: f32,
20185 ) -> Result<(), Box<dyn std::error::Error>> {
20186 let f = self.func("rms_norm_f16out_f32");
20187 let cfg = LaunchConfig {
20188 grid_dim: (nrows as u32, 1, 1),
20189 block_dim: (rms_block(), 1, 1),
20190 shared_mem_bytes: 0,
20191 };
20192 let (nc, e) = (ncols as i32, eps);
20193 let __s_b = self.gpu.stream();
20194 let mut b = __s_b.launch_builder(&f);
20195 b.arg(x).arg(w).arg(dst).arg(dst16).arg(&nc).arg(&e);
20196 unsafe {
20197 b.launch(cfg)?;
20198 }
20199 Ok(())
20200 }
20201
20202 #[allow(clippy::too_many_arguments)]
20205 pub fn add_rms_norm_f16out(
20206 &self,
20207 a: &CudaSlice<f32>,
20208 b: &CudaSlice<f32>,
20209 w: &CudaSlice<f32>,
20210 res: &mut CudaSlice<f32>,
20211 dst: &mut CudaSlice<f32>,
20212 dst16: &mut CudaSlice<u8>,
20213 ncols: usize,
20214 nrows: usize,
20215 eps: f32,
20216 ) -> Result<(), Box<dyn std::error::Error>> {
20217 let f = self.func("add_rms_norm_f16out_f32");
20218 let cfg = LaunchConfig {
20219 grid_dim: (nrows as u32, 1, 1),
20220 block_dim: (rms_block(), 1, 1),
20221 shared_mem_bytes: 0,
20222 };
20223 let (nc, e) = (ncols as i32, eps);
20224 let __s_lb = self.gpu.stream();
20225 let mut lb = __s_lb.launch_builder(&f);
20226 lb.arg(a)
20227 .arg(b)
20228 .arg(w)
20229 .arg(res)
20230 .arg(dst)
20231 .arg(dst16)
20232 .arg(&nc)
20233 .arg(&e);
20234 unsafe {
20235 lb.launch(cfg)?;
20236 }
20237 Ok(())
20238 }
20239
20240 pub fn matmul_group_xh(
20243 &self,
20244 ws: &[&crate::model::GpuTensor],
20245 x: &CudaSlice<f32>,
20246 xh: &CudaSlice<u8>,
20247 m: usize,
20248 ) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
20249 let mut out = Vec::with_capacity(ws.len());
20250 let in_f = ws[0].in_features();
20251 for w in ws {
20252 if w.in_features() == in_f
20253 && m >= 16
20254 && !self.verify_exact_on()
20255 && let Some(y) = self.try_f16_gemm_pre(w, xh, m)?
20256 {
20257 out.push(y);
20258 continue;
20259 }
20260 out.push(self.matmul(w, x, m)?);
20261 }
20262 Ok(out)
20263 }
20264
20265 pub fn gdn_pad_mask(
20268 &self,
20269 beta: &mut CudaSlice<f32>,
20270 g_log: &mut CudaSlice<f32>,
20271 len_d: &CudaSlice<i32>,
20272 h: usize,
20273 t: usize,
20274 ) -> Result<(), Box<dyn std::error::Error>> {
20275 let f = self.func("gdn_pad_mask_f32");
20276 let cfg = LaunchConfig::for_num_elems((t * h) as u32);
20277 let (hi, ti) = (h as i32, t as i32);
20278 let __s_b = self.gpu.stream();
20279 let mut b = __s_b.launch_builder(&f);
20280 b.arg(beta).arg(g_log).arg(len_d).arg(&hi).arg(&ti);
20281 unsafe {
20282 b.launch(cfg)?;
20283 }
20284 Ok(())
20285 }
20286
20287 pub fn row_gather_dev(
20290 &self,
20291 src: &CudaSlice<f32>,
20292 dst: &mut CudaSlice<f32>,
20293 len_d: &CudaSlice<i32>,
20294 ncols: usize,
20295 ) -> Result<(), Box<dyn std::error::Error>> {
20296 let f = self.func("row_gather_dev_f32");
20297 let cfg = LaunchConfig::for_num_elems(ncols as u32);
20298 let nc = ncols as i32;
20299 let __s_b = self.gpu.stream();
20300 let mut b = __s_b.launch_builder(&f);
20301 b.arg(src).arg(dst).arg(len_d).arg(&nc);
20302 unsafe {
20303 b.launch(cfg)?;
20304 }
20305 Ok(())
20306 }
20307
20308 pub fn matmul_group(
20315 &self,
20316 ws: &[&crate::model::GpuTensor],
20317 x: &CudaSlice<f32>,
20318 m: usize,
20319 ) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
20320 use crate::model::GpuTensor;
20321 let mut out = Vec::with_capacity(ws.len());
20322 let any_mirror = ws
20323 .iter()
20324 .any(|w| matches!(w, GpuTensor::Quant { f16: Some(_), .. }));
20325 if m >= 16 && any_mirror && !self.verify_exact_on() {
20326 let in_f = ws[0].in_features();
20327 let xh = self.f16_act(x, m * in_f, in_f)?;
20328 for w in ws {
20329 if w.in_features() == in_f
20330 && let Some(y) = self.try_f16_gemm_pre(w, &xh, m)?
20331 {
20332 out.push(y);
20333 continue;
20334 }
20335 out.push(self.matmul(w, x, m)?);
20336 }
20337 return Ok(out);
20338 }
20339 for w in ws {
20340 out.push(self.matmul(w, x, m)?);
20341 }
20342 Ok(out)
20343 }
20344
20345 pub fn matmul_group_multi(
20352 &self,
20353 ws: &[&crate::model::GpuTensor],
20354 xs: &[&CudaSlice<f32>],
20355 ms: &[usize],
20356 ) -> Result<Vec<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
20357 assert_eq!(xs.len(), ms.len());
20358 let in_f = ws[0].in_features();
20359 let total: usize = ms.iter().sum();
20360 let mut xcat = self.uninit(total * in_f)?;
20361 let mut off = 0usize;
20362 for (x, &m) in xs.iter().zip(ms) {
20363 self.copy_into(&mut xcat, off * in_f, x, m * in_f)?;
20364 off += m;
20365 }
20366 let ys = self.matmul_group(ws, &xcat, total)?;
20367 let mut out: Vec<Vec<CudaSlice<f32>>> = (0..xs.len()).map(|_| Vec::new()).collect();
20368 for (w, y) in ws.iter().zip(ys) {
20369 let out_f = w.out_features();
20370 let mut off = 0usize;
20371 for (s, &m) in ms.iter().enumerate() {
20372 let mut ys_s = self.uninit(m * out_f)?;
20373 let src = y.slice(off * out_f..(off + m) * out_f);
20374 self.gpu.stream().memcpy_dtod(&src, &mut ys_s)?;
20375 out[s].push(ys_s);
20376 off += m;
20377 }
20378 }
20379 Ok(out)
20380 }
20381
20382 pub fn gemm_supports(&self, w: &crate::model::GpuTensor) -> bool {
20392 use crate::model::GpuTensor;
20393 if !legacy_quant_gemm_allowed(
20394 cfg!(memra_portable_cuda),
20395 cfg!(memra_hopper_mma),
20396 std::env::var_os("MEMRA_NO_GEMM").is_some(),
20397 ) {
20398 return false;
20399 }
20400 match w {
20401 GpuTensor::Quant { qtype, .. } => {
20402 matches!(*qtype, QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q4_0)
20403 || (*qtype == QT_NVFP4 && w.in_features().is_multiple_of(64))
20404 }
20405 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
20406 }
20407 }
20408
20409 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_gemm(
20417 &self,
20418 w: &crate::model::GpuTensor,
20419 aq: &CudaSlice<i8>,
20420 ad: &CudaSlice<f32>,
20421 m: usize,
20422 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20423 use crate::model::GpuTensor;
20424 let in_f = w.in_features();
20425 let out_f = w.out_features();
20426 let (bytes, qtype, row_bytes, scale, rp) = match w {
20427 GpuTensor::Quant {
20428 bytes,
20429 qtype,
20430 row_bytes,
20431 scale,
20432 rp,
20433 ..
20434 } => (bytes, *qtype, *row_bytes, *scale, *rp),
20435 _ => unreachable!("gemm_supports guaranteed Quant"),
20436 };
20437 if cfg!(memra_hopper_mma)
20443 && qtype == QT_Q8_0
20444 && out_f.is_multiple_of(64)
20445 && wgmma_gemm_enabled()
20446 && let GpuTensor::Quant { rp4: Some(m4), .. } = w
20447 {
20448 let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, aq, ad, m, in_f, out_f)?;
20449 if scale != 1.0 {
20450 self.scale_inplace(&mut y, scale, m * out_f)?;
20451 }
20452 return Ok(y);
20453 }
20454 let name = match qtype {
20455 QT_Q8_0 => "qmatvec_gemm_q8_0",
20456 QT_Q4_K => "qmatvec_gemm_q4_K",
20457 QT_Q4_0 => {
20458 if rp {
20459 "qmatvec_gemm_q4_0_rp"
20460 } else {
20461 "qmatvec_gemm_q4_0"
20462 }
20463 }
20464 QT_Q5_K => "qmatvec_gemm_q5_K",
20465 QT_Q6_K => "qmatvec_gemm_q6_K",
20466 QT_NVFP4 => {
20467 if rp {
20468 "qmatvec_gemm_nvfp4_rp"
20469 } else {
20470 "qmatvec_gemm_nvfp4"
20471 }
20472 }
20473 _ => unreachable!(),
20474 };
20475 let f = self.func(name);
20476 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);
20481 let k1_tile = if is_k1 {
20483 k1_launch_override().unwrap_or((128, 128, 8))
20484 } else {
20485 (128, 128, 8)
20486 };
20487 let (bm, bn): (u32, u32) = if is_k1 {
20488 (k1_tile.0, k1_tile.1)
20489 } else {
20490 (64, 256)
20491 };
20492 let warps: u32 = if is_k1 {
20493 k1_tile.2
20494 } else {
20495 match qtype {
20496 QT_NVFP4 => 8,
20497 _ => 4,
20498 }
20499 };
20500 let cfg = LaunchConfig {
20501 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
20502 block_dim: (32, warps, 1),
20503 shared_mem_bytes: 0,
20504 };
20505 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
20506 let __s_b = self.gpu.stream();
20507 let mut b = __s_b.launch_builder(&f);
20508 b.arg(bytes)
20509 .arg(aq)
20510 .arg(ad)
20511 .arg(&mut y)
20512 .arg(&inf)
20513 .arg(&outf)
20514 .arg(&mi)
20515 .arg(&rb);
20516 unsafe {
20517 b.launch(cfg)?;
20518 }
20519 if scale != 1.0 {
20520 self.scale_inplace(&mut y, scale, m * out_f)?;
20521 }
20522 Ok(y)
20523 }
20524
20525 #[allow(clippy::too_many_arguments)]
20530 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_gemm_raw(
20533 &self,
20534 bytes: &CudaSlice<u8>,
20535 x: &CudaSlice<f32>,
20536 m: usize,
20537 in_f: usize,
20538 out_f: usize,
20539 qtype: i32,
20540 row_bytes: usize,
20541 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20542 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
20543 let name = match qtype {
20544 QT_Q8_0 => "qmatvec_gemm_q8_0",
20545 QT_Q4_K => "qmatvec_gemm_q4_K",
20546 QT_Q4_0 => "qmatvec_gemm_q4_0",
20547 QT_Q5_K => "qmatvec_gemm_q5_K",
20548 QT_Q6_K => "qmatvec_gemm_q6_K",
20549 QT_NVFP4 => "qmatvec_gemm_nvfp4",
20550 QT_NVFP4_RP => "qmatvec_gemm_nvfp4_rp",
20551 _ => panic!("qmatvec_gemm_raw: qtype {qtype} has no GEMM kernel"),
20552 };
20553 let f = self.func(name);
20554 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);
20558 let k1_tile = if is_k1 {
20560 k1_launch_override().unwrap_or((128, 128, 8))
20561 } else {
20562 (128, 128, 8)
20563 };
20564 let (bm, bn): (u32, u32) = if is_k1 {
20565 (k1_tile.0, k1_tile.1)
20566 } else {
20567 (64, 256)
20568 };
20569 let warps: u32 = if is_k1 {
20570 k1_tile.2
20571 } else {
20572 match qtype {
20573 QT_NVFP4 | QT_NVFP4_RP => 8,
20574 _ => 4,
20575 }
20576 };
20577 let cfg = LaunchConfig {
20578 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
20579 block_dim: (32, warps, 1),
20580 shared_mem_bytes: 0,
20581 };
20582 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
20583 let __s_b = self.gpu.stream();
20584 let mut b = __s_b.launch_builder(&f);
20585 b.arg(bytes)
20586 .arg(&aq)
20587 .arg(&ad)
20588 .arg(&mut y)
20589 .arg(&inf)
20590 .arg(&outf)
20591 .arg(&mi)
20592 .arg(&rb);
20593 unsafe {
20594 b.launch(cfg)?;
20595 }
20596 Ok(y)
20597 }
20598
20599 pub fn qmatvec_gemm_q8_0_wgmma_raw(
20606 &self,
20607 rp4: &CudaSlice<u8>,
20608 aq: &CudaSlice<i8>,
20609 ad: &CudaSlice<f32>,
20610 m: usize,
20611 in_f: usize,
20612 out_f: usize,
20613 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20614 assert!(
20615 out_f.is_multiple_of(64) && in_f.is_multiple_of(32),
20616 "wgmma GEMM needs out_f%64==0, in_f%32==0"
20617 );
20618 let f = self.func("qmatvec_gemm_q8_0_wgmma");
20619 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
20621 grid_dim: ((out_f / 64) as u32, (m as u32).div_ceil(64), 1),
20622 block_dim: (128, 1, 1),
20623 shared_mem_bytes: 0,
20624 };
20625 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
20626 let __s_b = self.gpu.stream();
20627 let mut b = __s_b.launch_builder(&f);
20628 b.arg(rp4)
20629 .arg(aq)
20630 .arg(ad)
20631 .arg(&mut y)
20632 .arg(&inf)
20633 .arg(&outf)
20634 .arg(&mi);
20635 unsafe {
20636 b.launch(cfg)?;
20637 }
20638 Ok(y)
20639 }
20640
20641 pub fn scale_inplace(
20643 &self,
20644 y: &mut CudaSlice<f32>,
20645 s: f32,
20646 n: usize,
20647 ) -> Result<(), Box<dyn std::error::Error>> {
20648 let f = self.func("scale_f32");
20649 let cfg = LaunchConfig::for_num_elems(n as u32);
20650 let (sf, ni) = (s, n as i32);
20651 let __s_b = self.gpu.stream();
20652 let mut b = __s_b.launch_builder(&f);
20653 b.arg(y).arg(&sf).arg(&ni);
20654 unsafe {
20655 b.launch(cfg)?;
20656 }
20657 Ok(())
20658 }
20659
20660 pub fn bf16_to_f32(
20665 &self,
20666 data: &cudarc::driver::CudaView<'_, u8>,
20667 n: usize,
20668 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20669 let mut out = self.alloc_uninit::<f32>(n)?;
20670 let f = self.func("bf16_to_f32");
20671 let cfg = LaunchConfig::for_num_elems(n as u32);
20672 let ni = n as i32;
20673 let __s_b = self.gpu.stream();
20674 let mut b = __s_b.launch_builder(&f);
20675 b.arg(data).arg(&mut out).arg(&ni);
20676 unsafe {
20677 b.launch(cfg)?;
20678 }
20679 Ok(out)
20680 }
20681
20682 #[allow(clippy::too_many_arguments)] fn linear_bf16_chunked(
20690 &self,
20691 x: &CudaSlice<f32>,
20692 data: &CudaSlice<u8>,
20693 m: usize,
20694 in_f: usize,
20695 out_f: usize,
20696 exact: bool,
20697 canonical_chunk_rows: Option<usize>,
20698 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20699 static EXP_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
20703 static EXP_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
20704 static EXP_WBYTES: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
20705 let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
20706 let started = timing.then(std::time::Instant::now);
20707 let result =
20708 self.linear_bf16_chunked_inner(x, data, m, in_f, out_f, exact, canonical_chunk_rows);
20709 if let Some(started) = started {
20710 use std::sync::atomic::Ordering;
20711 self.stream().synchronize()?;
20712 let ns = EXP_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
20713 + started.elapsed().as_nanos() as u64;
20714 let wb = EXP_WBYTES.fetch_add((in_f * out_f * 2) as u64, Ordering::Relaxed)
20715 + (in_f * out_f * 2) as u64;
20716 let calls = EXP_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
20717 if calls.is_multiple_of(1024) {
20718 eprintln!(
20719 "[bf16-expand-timing] calls={calls} total_ms={:.1} avg_us={:.1} \
20720 weight_gb={:.2}",
20721 ns as f64 / 1.0e6,
20722 ns as f64 / calls as f64 / 1.0e3,
20723 wb as f64 / 1.0e9,
20724 );
20725 }
20726 }
20727 result
20728 }
20729
20730 pub(crate) fn bf16_mmv_on() -> bool {
20735 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20736 *ON.get_or_init(|| std::env::var("MEMRA_BF16_MMV").as_deref() == Ok("1"))
20737 }
20738
20739 fn matvec_bf16(
20742 &self,
20743 data: &CudaSlice<u8>,
20744 x: &CudaSlice<f32>,
20745 in_f: usize,
20746 out_f: usize,
20747 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20748 if data.len() != in_f * out_f * 2 || x.len() < in_f || !in_f.is_multiple_of(8) {
20749 return Err(format!(
20750 "matvec_bf16 geometry bytes={} x={} in={in_f} out={out_f}",
20751 data.len(),
20752 x.len()
20753 )
20754 .into());
20755 }
20756 let mut y = self.alloc_uninit::<f32>(out_f)?;
20757 let f = self.func("matvec_bf16_f32acc");
20758 let cfg = LaunchConfig {
20759 grid_dim: (out_f as u32, 1, 1),
20760 block_dim: (mmv_block(), 1, 1),
20761 shared_mem_bytes: 0,
20762 };
20763 let ini = in_f as i32;
20764 let __s_bld = self.gpu.stream();
20765 let mut bld = __s_bld.launch_builder(&f);
20766 bld.arg(data).arg(x).arg(&mut y).arg(&ini);
20767 unsafe {
20768 bld.launch(cfg)?;
20769 }
20770 Ok(y)
20771 }
20772
20773 #[allow(clippy::too_many_arguments)]
20777 #[allow(clippy::too_many_arguments)]
20782 #[allow(clippy::too_many_arguments)]
20787 pub fn qk_norm_rope_append_inc_dcw_rows(
20788 &self,
20789 q_raw_t: &CudaSlice<f32>,
20790 k_raw_t: &CudaSlice<f32>,
20791 v_raw_t: &CudaSlice<f32>,
20792 qw: &CudaSlice<f32>,
20793 kw: &CudaSlice<f32>,
20794 q_out_t: &mut CudaSlice<f32>,
20795 k_out_t: &mut CudaSlice<f32>,
20796 tab: &CudaSlice<u64>,
20797 pos_t: &CudaSlice<i32>,
20798 same_session: bool,
20799 t: usize,
20800 kv_dim_k: usize,
20801 kv_dim_v: usize,
20802 k_tok_bytes: usize,
20803 v_tok_bytes: usize,
20804 head_dim: usize,
20805 n_dims: usize,
20806 nh_q: usize,
20807 nh_k: usize,
20808 eps: f32,
20809 freq_base: f32,
20810 freq_scale: f32,
20811 ff: Option<&CudaSlice<f32>>,
20812 ) -> Result<(), Box<dyn std::error::Error>> {
20813 if head_dim != 128
20814 || kv_dim_v != kv_dim_k
20815 || kv_dim_k != nh_k * head_dim
20816 || t == 0
20817 || t > 32
20818 || tab.len() < t * 6
20819 || pos_t.len() < t
20820 || q_raw_t.len() < t * nh_q * head_dim
20821 || k_raw_t.len() < t * nh_k * head_dim
20822 || v_raw_t.len() < t * kv_dim_v
20823 || q_out_t.len() < t * nh_q * head_dim
20824 || k_out_t.len() < t * nh_k * head_dim
20825 {
20826 return Err(format!(
20827 "qk_norm_rope_append_inc_rows geometry head_dim={head_dim} t={t} \
20828 nh_q={nh_q} nh_k={nh_k}"
20829 )
20830 .into());
20831 }
20832 let f = self.func("qk_norm_rope_append_inc_dcw_rows");
20833 let same_t: i32 = if same_session { t as i32 } else { 0 };
20834 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
20835 let cfg = LaunchConfig {
20836 grid_dim: ((nh_q + nh_k) as u32, 1, t as u32),
20837 block_dim: (128, 1, 1),
20838 shared_mem_bytes: 0,
20839 };
20840 let (kvk, kvv) = (kv_dim_k as i32, kv_dim_v as i32);
20841 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
20842 let (hd, nd, nq, nk) = (head_dim as i32, n_dims as i32, nh_q as i32, nh_k as i32);
20843 let null: u64 = 0;
20844 let __s_b = self.gpu.stream();
20845 let mut b = __s_b.launch_builder(&f);
20846 b.arg(q_raw_t)
20847 .arg(k_raw_t)
20848 .arg(v_raw_t)
20849 .arg(qw)
20850 .arg(kw)
20851 .arg(q_out_t)
20852 .arg(k_out_t)
20853 .arg(tab)
20854 .arg(pos_t)
20855 .arg(&same_t)
20856 .arg(&kvk)
20857 .arg(&kvv)
20858 .arg(&ktb)
20859 .arg(&vtb)
20860 .arg(&hd)
20861 .arg(&nd)
20862 .arg(&nq)
20863 .arg(&nk)
20864 .arg(&eps)
20865 .arg(&theta_scale)
20866 .arg(&freq_scale);
20867 match ff {
20868 Some(freqs) => {
20869 b.arg(freqs);
20870 }
20871 None => {
20872 b.arg(&null);
20873 }
20874 }
20875 unsafe {
20876 b.launch(cfg)?;
20877 }
20878 Ok(())
20879 }
20880
20881 #[allow(clippy::too_many_arguments)] pub fn qk_norm_rope_append_inc_dcw(
20883 &self,
20884 q_raw: &CudaSlice<f32>,
20885 k_raw: &CudaSlice<f32>,
20886 v_raw: &CudaSlice<f32>,
20887 qw: &CudaSlice<f32>,
20888 kw: &CudaSlice<f32>,
20889 q_out: &mut CudaSlice<f32>,
20890 k_out: &mut CudaSlice<f32>,
20891 pos: &CudaSlice<i32>,
20892 k_plane: &mut CudaSlice<u8>,
20893 v_plane: &mut CudaSlice<u8>,
20894 len_dev: &CudaSlice<i32>,
20897 base_dev: Option<&CudaSlice<i32>>,
20898 done_ctr: &mut CudaSlice<u32>,
20899 kv_dim_k: usize,
20900 kv_dim_v: usize,
20901 k_tok_bytes: usize,
20902 v_tok_bytes: usize,
20903 head_dim: usize,
20904 n_dims: usize,
20905 nh_q: usize,
20906 nh_k: usize,
20907 eps: f32,
20908 freq_base: f32,
20909 freq_scale: f32,
20910 ff: Option<&CudaSlice<f32>>,
20911 ) -> Result<(), Box<dyn std::error::Error>> {
20912 if head_dim != 128
20913 || kv_dim_v != kv_dim_k
20914 || kv_dim_k != nh_k * head_dim
20915 || q_raw.len() < nh_q * head_dim
20916 || k_raw.len() < nh_k * head_dim
20917 || v_raw.len() < kv_dim_v
20918 || q_out.len() < nh_q * head_dim
20919 || k_out.len() < nh_k * head_dim
20920 || pos.is_empty()
20921 || done_ctr.is_empty()
20922 {
20923 return Err(format!(
20924 "qk_norm_rope_append_inc geometry head_dim={head_dim} nh_q={nh_q} nh_k={nh_k} kv_k={kv_dim_k} kv_v={kv_dim_v}"
20925 )
20926 .into());
20927 }
20928 let f = self.func("qk_norm_rope_append_inc_dcw");
20929 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
20930 let cfg = LaunchConfig {
20931 grid_dim: ((nh_q + nh_k) as u32, 1, 1),
20932 block_dim: (128, 1, 1),
20933 shared_mem_bytes: 0,
20934 };
20935 let (kvk, kvv) = (kv_dim_k as i32, kv_dim_v as i32);
20936 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
20937 let (hd, nd, nq) = (head_dim as i32, n_dims as i32, nh_q as i32);
20938 let null: u64 = 0;
20939 let __s_b = self.gpu.stream();
20940 let mut b = __s_b.launch_builder(&f);
20941 b.arg(q_raw)
20942 .arg(k_raw)
20943 .arg(v_raw)
20944 .arg(qw)
20945 .arg(kw)
20946 .arg(q_out)
20947 .arg(k_out)
20948 .arg(pos)
20949 .arg(&mut *k_plane)
20950 .arg(&mut *v_plane)
20951 .arg(len_dev);
20952 match base_dev {
20953 Some(base) => {
20954 b.arg(base);
20955 }
20956 None => {
20957 b.arg(&null);
20958 }
20959 }
20960 b.arg(&mut *done_ctr)
20961 .arg(&kvk)
20962 .arg(&kvv)
20963 .arg(&ktb)
20964 .arg(&vtb)
20965 .arg(&hd)
20966 .arg(&nd)
20967 .arg(&nq)
20968 .arg(&eps)
20969 .arg(&theta_scale)
20970 .arg(&freq_scale);
20971 match ff {
20972 Some(freqs) => {
20973 b.arg(freqs);
20974 }
20975 None => {
20976 b.arg(&null);
20977 }
20978 }
20979 unsafe {
20980 b.launch(cfg)?;
20981 }
20982 Ok(())
20983 }
20984
20985 #[allow(clippy::too_many_arguments)] pub fn qk_norm_rope_into(
20987 &self,
20988 q_raw: &CudaSlice<f32>,
20989 k_raw: &CudaSlice<f32>,
20990 qw: &CudaSlice<f32>,
20991 kw: &CudaSlice<f32>,
20992 q_out: &mut CudaSlice<f32>,
20993 k_out: &mut CudaSlice<f32>,
20994 pos: &CudaSlice<i32>,
20995 head_dim: usize,
20996 n_dims: usize,
20997 nh_q: usize,
20998 nh_k: usize,
20999 eps: f32,
21000 freq_base: f32,
21001 freq_scale: f32,
21002 ff: Option<&CudaSlice<f32>>,
21003 ) -> Result<(), Box<dyn std::error::Error>> {
21004 if head_dim > 512
21005 || q_raw.len() < nh_q * head_dim
21006 || k_raw.len() < nh_k * head_dim
21007 || q_out.len() < nh_q * head_dim
21008 || k_out.len() < nh_k * head_dim
21009 || qw.len() < head_dim
21010 || kw.len() < head_dim
21011 || pos.is_empty()
21012 {
21013 return Err(format!(
21014 "qk_norm_rope geometry head_dim={head_dim} nh_q={nh_q} nh_k={nh_k}"
21015 )
21016 .into());
21017 }
21018 let f = self.func("qk_norm_rope_f32");
21019 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
21020 let cfg = LaunchConfig {
21021 grid_dim: ((nh_q + nh_k) as u32, 1, 1),
21022 block_dim: (128, 1, 1),
21023 shared_mem_bytes: 0,
21024 };
21025 let (hd, nd, nq) = (head_dim as i32, n_dims as i32, nh_q as i32);
21026 let __s_b = self.gpu.stream();
21027 let mut b = __s_b.launch_builder(&f);
21028 b.arg(q_raw)
21029 .arg(k_raw)
21030 .arg(qw)
21031 .arg(kw)
21032 .arg(q_out)
21033 .arg(k_out)
21034 .arg(pos)
21035 .arg(&hd)
21036 .arg(&nd)
21037 .arg(&nq)
21038 .arg(&eps)
21039 .arg(&theta_scale)
21040 .arg(&freq_scale);
21041 match ff {
21042 Some(ffv) => {
21043 b.arg(ffv);
21044 unsafe {
21045 b.launch(cfg)?;
21046 }
21047 }
21048 None => {
21049 let null: u64 = 0;
21050 b.arg(&null);
21051 unsafe {
21052 b.launch(cfg)?;
21053 }
21054 }
21055 }
21056 Ok(())
21057 }
21058
21059 #[allow(clippy::too_many_arguments)]
21062 pub fn matvec_f32_b4_into(
21063 &self,
21064 w: [&CudaSlice<f32>; 4],
21065 x: &CudaSlice<f32>,
21066 y: &mut CudaSlice<f32>,
21067 block_cols: usize,
21068 out_f: usize,
21069 ) -> Result<(), Box<dyn std::error::Error>> {
21070 if !block_cols.is_multiple_of(4)
21071 || x.len() < 4 * block_cols
21072 || y.len() < out_f
21073 || w.iter().any(|w| w.len() != out_f * block_cols)
21074 {
21075 return Err(format!(
21076 "matvec_f32_b4 geometry block_cols={block_cols} out={out_f} x={}",
21077 x.len()
21078 )
21079 .into());
21080 }
21081 let f = self.func("matvec_f32_b4");
21082 let cfg = LaunchConfig {
21083 grid_dim: (out_f as u32, 1, 1),
21084 block_dim: (128, 1, 1),
21085 shared_mem_bytes: 0,
21086 };
21087 let (bc, of) = (block_cols as i32, out_f as i32);
21088 let __s_b = self.gpu.stream();
21089 let mut b = __s_b.launch_builder(&f);
21090 b.arg(w[0])
21091 .arg(w[1])
21092 .arg(w[2])
21093 .arg(w[3])
21094 .arg(x)
21095 .arg(y)
21096 .arg(&bc)
21097 .arg(&of);
21098 unsafe {
21099 b.launch(cfg)?;
21100 }
21101 Ok(())
21102 }
21103
21104 pub fn axpy_rows_seq_into(
21107 &self,
21108 x: &CudaSlice<f32>,
21109 w: &CudaSlice<f32>,
21110 y: &mut CudaSlice<f32>,
21111 width: usize,
21112 n_rows: usize,
21113 ) -> Result<(), Box<dyn std::error::Error>> {
21114 if x.len() < n_rows * width || w.len() < n_rows || y.len() < width {
21115 return Err(format!(
21116 "axpy_rows_seq geometry x={} w={} y={} width={width} rows={n_rows}",
21117 x.len(),
21118 w.len(),
21119 y.len()
21120 )
21121 .into());
21122 }
21123 let f = self.func("axpy_rows_seq_f32");
21124 let cfg = LaunchConfig::for_num_elems(width as u32);
21125 let (wi, nr) = (width as i32, n_rows as i32);
21126 let __s_b = self.gpu.stream();
21127 let mut b = __s_b.launch_builder(&f);
21128 b.arg(x).arg(w).arg(y).arg(&wi).arg(&nr);
21129 unsafe {
21130 b.launch(cfg)?;
21131 }
21132 Ok(())
21133 }
21134
21135 pub fn axpy_rows_seq_tokens_into(
21138 &self,
21139 x: &CudaSlice<f32>,
21140 w: &CudaSlice<f32>,
21141 y: &mut CudaSlice<f32>,
21142 width: usize,
21143 slots: usize,
21144 tokens: usize,
21145 ) -> Result<(), Box<dyn std::error::Error>> {
21146 let rows = slots
21147 .checked_mul(tokens)
21148 .ok_or("axpy_rows_seq_tokens row count overflow")?;
21149 if x.len() < rows * width || w.len() < rows || y.len() < tokens * width {
21150 return Err(format!(
21151 "axpy_rows_seq_tokens geometry x={} w={} y={} width={width} \
21152 slots={slots} tokens={tokens}",
21153 x.len(),
21154 w.len(),
21155 y.len()
21156 )
21157 .into());
21158 }
21159 let f = self.func("axpy_rows_seq_tokens_f32");
21160 let block = 256u32;
21161 let cfg = LaunchConfig {
21162 grid_dim: ((width as u32).div_ceil(block), tokens as u32, 1),
21163 block_dim: (block, 1, 1),
21164 shared_mem_bytes: 0,
21165 };
21166 let (wi, sl, tk) = (width as i32, slots as i32, tokens as i32);
21167 let __s_b = self.gpu.stream();
21168 let mut b = __s_b.launch_builder(&f);
21169 b.arg(x).arg(w).arg(y).arg(&wi).arg(&sl).arg(&tk);
21170 unsafe {
21171 b.launch(cfg)?;
21172 }
21173 Ok(())
21174 }
21175
21176 #[allow(clippy::too_many_arguments)]
21180 pub fn axpy_rows_seq_md_off_into(
21181 &self,
21182 x: &CudaSlice<f32>,
21183 w_route: &CudaSlice<f32>,
21184 md: &CudaSlice<f32>,
21185 sel: &CudaSlice<i32>,
21186 y: &mut CudaSlice<f32>,
21187 width: usize,
21188 n_rows: usize,
21189 row0: usize,
21190 ) -> Result<(), Box<dyn std::error::Error>> {
21191 if x.len() < (row0 + n_rows) * width
21192 || w_route.len() < row0 + n_rows
21193 || sel.len() < row0 + n_rows
21194 || y.len() < width
21195 {
21196 return Err(format!(
21197 "axpy_rows_seq_md_off geometry x={} w={} sel={} y={} width={width} \
21198 rows={n_rows} row0={row0}",
21199 x.len(),
21200 w_route.len(),
21201 sel.len(),
21202 y.len()
21203 )
21204 .into());
21205 }
21206 let f = self.func("axpy_rows_seq_md_off_f32");
21207 let cfg = LaunchConfig::for_num_elems(width as u32);
21208 let (wi, nr, r0) = (width as i32, n_rows as i32, row0 as i32);
21209 let __s_b = self.gpu.stream();
21210 let mut b = __s_b.launch_builder(&f);
21211 b.arg(x)
21212 .arg(w_route)
21213 .arg(md)
21214 .arg(sel)
21215 .arg(y)
21216 .arg(&wi)
21217 .arg(&nr)
21218 .arg(&r0);
21219 unsafe {
21220 b.launch(cfg)?;
21221 }
21222 Ok(())
21223 }
21224
21225 #[allow(clippy::too_many_arguments)]
21228 pub fn axpy_rows_seq_md_into(
21229 &self,
21230 x: &CudaSlice<f32>,
21231 w_route: &CudaSlice<f32>,
21232 md: &CudaSlice<f32>,
21233 sel: &CudaSlice<i32>,
21234 y: &mut CudaSlice<f32>,
21235 width: usize,
21236 n_rows: usize,
21237 ) -> Result<(), Box<dyn std::error::Error>> {
21238 if x.len() < n_rows * width
21239 || w_route.len() < n_rows
21240 || sel.len() < n_rows
21241 || y.len() < width
21242 {
21243 return Err(format!(
21244 "axpy_rows_seq_md geometry x={} w={} sel={} y={} width={width} rows={n_rows}",
21245 x.len(),
21246 w_route.len(),
21247 sel.len(),
21248 y.len()
21249 )
21250 .into());
21251 }
21252 let f = self.func("axpy_rows_seq_md_f32");
21253 let cfg = LaunchConfig::for_num_elems(width as u32);
21254 let (wi, nr) = (width as i32, n_rows as i32);
21255 let __s_b = self.gpu.stream();
21256 let mut b = __s_b.launch_builder(&f);
21257 b.arg(x)
21258 .arg(w_route)
21259 .arg(md)
21260 .arg(sel)
21261 .arg(y)
21262 .arg(&wi)
21263 .arg(&nr);
21264 unsafe {
21265 b.launch(cfg)?;
21266 }
21267 Ok(())
21268 }
21269
21270 #[allow(clippy::too_many_arguments)]
21272 #[allow(clippy::too_many_arguments)]
21276 pub fn matvec_bf16_qkvg_tcol_into(
21277 &self,
21278 wq: &CudaSlice<u8>,
21279 wk: &CudaSlice<u8>,
21280 wv: &CudaSlice<u8>,
21281 wg: &CudaSlice<u8>,
21282 x_t: &CudaSlice<f32>,
21283 yq: &mut CudaSlice<f32>,
21284 yk: &mut CudaSlice<f32>,
21285 yv: &mut CudaSlice<f32>,
21286 yg: &mut CudaSlice<f32>,
21287 in_f: usize,
21288 out_q: usize,
21289 out_kv: usize,
21290 out_g: usize,
21291 t: usize,
21292 ) -> Result<(), Box<dyn std::error::Error>> {
21293 if t == 0
21294 || t > 8
21295 || !in_f.is_multiple_of(8)
21296 || x_t.len() < t * in_f
21297 || yq.len() < t * out_q
21298 || yk.len() < t * out_kv
21299 || yv.len() < t * out_kv
21300 || (out_g > 0 && yg.len() < t * out_g)
21301 {
21302 return Err("matvec_bf16_qkvg_tcol geometry".into());
21303 }
21304 let grid = out_q + 2 * out_kv + out_g;
21305 let cfg = LaunchConfig {
21306 grid_dim: (grid as u32, 1, 1),
21307 block_dim: (mmv_block(), 1, 1),
21308 shared_mem_bytes: 0,
21309 };
21310 let (ini, oq, okv, og, ti) = (
21311 in_f as i32,
21312 out_q as i32,
21313 out_kv as i32,
21314 out_g as i32,
21315 t as i32,
21316 );
21317 let __s_b = self.gpu.stream();
21318 let f = self.func("matvec_bf16_qkvg_tcol");
21324 let mut b = __s_b.launch_builder(&f);
21325 b.arg(wq)
21326 .arg(wk)
21327 .arg(wv)
21328 .arg(wg)
21329 .arg(x_t)
21330 .arg(yq)
21331 .arg(yk)
21332 .arg(yv)
21333 .arg(yg)
21334 .arg(&ini)
21335 .arg(&oq)
21336 .arg(&okv)
21337 .arg(&og)
21338 .arg(&ti);
21339 unsafe {
21340 b.launch(cfg)?;
21341 }
21342 Ok(())
21343 }
21344
21345 #[allow(clippy::too_many_arguments)] pub fn matvec_bf16_qkvg_into(
21347 &self,
21348 wq: &CudaSlice<u8>,
21349 wk: &CudaSlice<u8>,
21350 wv: &CudaSlice<u8>,
21351 wg: &CudaSlice<u8>,
21352 x: &CudaSlice<f32>,
21353 yq: &mut CudaSlice<f32>,
21354 yk: &mut CudaSlice<f32>,
21355 yv: &mut CudaSlice<f32>,
21356 yg: &mut CudaSlice<f32>,
21357 in_f: usize,
21358 out_q: usize,
21359 out_kv: usize,
21360 out_g: usize,
21361 ) -> Result<(), Box<dyn std::error::Error>> {
21362 if !in_f.is_multiple_of(8)
21363 || wq.len() != out_q * in_f * 2
21364 || wk.len() != out_kv * in_f * 2
21365 || wv.len() != out_kv * in_f * 2
21366 || wg.len() < out_g * in_f * 2
21367 || x.len() < in_f
21368 || yq.len() < out_q
21369 || yk.len() < out_kv
21370 || yv.len() < out_kv
21371 || (out_g > 0 && yg.len() < out_g)
21372 {
21373 return Err(format!(
21374 "fused bf16 QKV geometry in={in_f} out_q={out_q} out_kv={out_kv} out_g={out_g}"
21375 )
21376 .into());
21377 }
21378 let f = self.func("matvec_bf16_qkvg");
21379 let cfg = LaunchConfig {
21380 grid_dim: ((out_q + 2 * out_kv + out_g) as u32, 1, 1),
21381 block_dim: (mmv_block(), 1, 1),
21382 shared_mem_bytes: 0,
21383 };
21384 let (inf, oq, okv, og) = (in_f as i32, out_q as i32, out_kv as i32, out_g as i32);
21385 let __s_b = self.gpu.stream();
21386 let mut b = __s_b.launch_builder(&f);
21387 b.arg(wq)
21388 .arg(wk)
21389 .arg(wv)
21390 .arg(wg)
21391 .arg(x)
21392 .arg(yq)
21393 .arg(yk)
21394 .arg(yv)
21395 .arg(yg)
21396 .arg(&inf)
21397 .arg(&oq)
21398 .arg(&okv)
21399 .arg(&og);
21400 unsafe {
21401 b.launch(cfg)?;
21402 }
21403 Ok(())
21404 }
21405
21406 pub fn matvec_bf16_b4_into(
21408 &self,
21409 w: [&CudaSlice<u8>; 4],
21410 x: &CudaSlice<f32>,
21411 y: &mut CudaSlice<f32>,
21412 block_cols: usize,
21413 out_f: usize,
21414 ) -> Result<(), Box<dyn std::error::Error>> {
21415 if !block_cols.is_multiple_of(8)
21416 || x.len() < 4 * block_cols
21417 || y.len() < out_f
21418 || w.iter().any(|w| w.len() != out_f * block_cols * 2)
21419 {
21420 return Err(format!(
21421 "bf16 b4 geometry block_cols={block_cols} out={out_f} x={}",
21422 x.len()
21423 )
21424 .into());
21425 }
21426 static B4_X2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
21429 let x2 = *B4_X2.get_or_init(|| std::env::var("MEMRA_B4_X2").as_deref() == Ok("1"));
21430 let f = self.func(if x2 {
21431 "matvec_bf16_b4_x2"
21432 } else {
21433 "matvec_bf16_b4"
21434 });
21435 let grid = if x2 { out_f.div_ceil(2) } else { out_f };
21436 let cfg = LaunchConfig {
21437 grid_dim: (grid as u32, 1, 1),
21438 block_dim: (mmv_block(), 1, 1),
21439 shared_mem_bytes: 0,
21440 };
21441 let (bc, of) = (block_cols as i32, out_f as i32);
21442 let __s_b = self.gpu.stream();
21443 let mut b = __s_b.launch_builder(&f);
21444 b.arg(w[0])
21445 .arg(w[1])
21446 .arg(w[2])
21447 .arg(w[3])
21448 .arg(x)
21449 .arg(y)
21450 .arg(&bc)
21451 .arg(&of);
21452 unsafe {
21453 b.launch(cfg)?;
21454 }
21455 Ok(())
21456 }
21457
21458 pub fn matvec_bf16_b4_tcol_into(
21464 &self,
21465 w: [&CudaSlice<u8>; 4],
21466 x_t: &CudaSlice<f32>,
21467 y_t: &mut CudaSlice<f32>,
21468 block_cols: usize,
21469 out_f: usize,
21470 t: usize,
21471 ) -> Result<(), Box<dyn std::error::Error>> {
21472 if !block_cols.is_multiple_of(8)
21473 || t == 0
21474 || t > 8
21475 || x_t.len() < t * 4 * block_cols
21476 || y_t.len() < t * out_f
21477 || w.iter().any(|w| w.len() != out_f * block_cols * 2)
21478 {
21479 return Err(format!(
21480 "bf16 b4 tcol geometry block_cols={block_cols} out={out_f} t={t} x={}",
21481 x_t.len()
21482 )
21483 .into());
21484 }
21485 if std::env::var("MEMRA_B4_X2").as_deref() == Ok("1") {
21486 return Err(
21487 "b4 tcol verify is qualified against the plain b4 kernel only \
21488 (MEMRA_B4_X2=1 is a different t=1 program)"
21489 .into(),
21490 );
21491 }
21492 let cfg = LaunchConfig {
21496 grid_dim: (out_f as u32, 1, 1),
21497 block_dim: (mmv_block(), 1, 1),
21498 shared_mem_bytes: 0,
21499 };
21500 let (bc, of, ti) = (block_cols as i32, out_f as i32, t as i32);
21501 let __s_b = self.gpu.stream();
21502 let f = self.func("matvec_bf16_b4_tcol");
21503 let mut b = __s_b.launch_builder(&f);
21504 b.arg(w[0])
21505 .arg(w[1])
21506 .arg(w[2])
21507 .arg(w[3])
21508 .arg(x_t)
21509 .arg(y_t)
21510 .arg(&bc)
21511 .arg(&of)
21512 .arg(&ti);
21513 unsafe {
21514 b.launch(cfg)?;
21515 }
21516 Ok(())
21517 }
21518
21519 pub fn q8_0_row_bytes(in_f: usize) -> usize {
21522 in_f / 32 * 34
21523 }
21524
21525 pub fn encode_q8_0_from_bf16(
21529 &self,
21530 w_bf16: &CudaSlice<u8>,
21531 out: &mut CudaSlice<u8>,
21532 in_f: usize,
21533 out_f: usize,
21534 ) -> Result<(), Box<dyn std::error::Error>> {
21535 if !in_f.is_multiple_of(32)
21536 || w_bf16.len() < in_f * out_f * 2
21537 || out.len() < out_f * Self::q8_0_row_bytes(in_f)
21538 {
21539 return Err(format!(
21540 "encode_q8_0_from_bf16 geometry in={in_f} out={out_f} src={} dst={}",
21541 w_bf16.len(),
21542 out.len()
21543 )
21544 .into());
21545 }
21546 let f = self.func("encode_q8_0_rows_from_bf16");
21547 const PAIRS_PER_BLOCK: u32 = 4;
21550 let pairs = (out_f * (in_f / 32)) as u64;
21551 let cfg = LaunchConfig {
21552 grid_dim: ((pairs.div_ceil(PAIRS_PER_BLOCK as u64)) as u32, 1, 1),
21553 block_dim: (32, PAIRS_PER_BLOCK, 1),
21554 shared_mem_bytes: 0,
21555 };
21556 let (ini, outi) = (in_f as i32, out_f as i32);
21557 let __s_b = self.gpu.stream();
21558 let mut b = __s_b.launch_builder(&f);
21559 b.arg(w_bf16).arg(out).arg(&ini).arg(&outi);
21560 unsafe {
21561 b.launch(cfg)?;
21562 }
21563 Ok(())
21564 }
21565
21566 pub fn encode_q8_0_from_bf16_view(
21570 &self,
21571 w_bf16: &cudarc::driver::CudaView<'_, u8>,
21572 out: &mut CudaSlice<u8>,
21573 in_f: usize,
21574 out_f: usize,
21575 ) -> Result<(), Box<dyn std::error::Error>> {
21576 if !in_f.is_multiple_of(32)
21577 || w_bf16.len() < in_f * out_f * 2
21578 || out.len() < out_f * Self::q8_0_row_bytes(in_f)
21579 {
21580 return Err(format!(
21581 "encode_q8_0_from_bf16_view geometry in={in_f} out={out_f} src={} dst={}",
21582 w_bf16.len(),
21583 out.len()
21584 )
21585 .into());
21586 }
21587 let f = self.func("encode_q8_0_rows_from_bf16");
21588 const PAIRS_PER_BLOCK: u32 = 4;
21589 let pairs = (out_f * (in_f / 32)) as u64;
21590 let cfg = LaunchConfig {
21591 grid_dim: ((pairs.div_ceil(PAIRS_PER_BLOCK as u64)) as u32, 1, 1),
21592 block_dim: (32, PAIRS_PER_BLOCK, 1),
21593 shared_mem_bytes: 0,
21594 };
21595 let (ini, outi) = (in_f as i32, out_f as i32);
21596 let __s_b = self.gpu.stream();
21597 let mut b = __s_b.launch_builder(&f);
21598 b.arg(w_bf16).arg(out).arg(&ini).arg(&outi);
21599 unsafe {
21600 b.launch(cfg)?;
21601 }
21602 Ok(())
21603 }
21604
21605 #[allow(clippy::too_many_arguments)]
21610 pub fn qmatvec_q8_0_qkv_rp_into(
21611 &self,
21612 wq: &CudaSlice<u8>,
21613 wk: &CudaSlice<u8>,
21614 wv: &CudaSlice<u8>,
21615 aq: &CudaSlice<i8>,
21616 ad: &CudaSlice<f32>,
21617 yq: &mut CudaSlice<f32>,
21618 yk: &mut CudaSlice<f32>,
21619 yv: &mut CudaSlice<f32>,
21620 in_f: usize,
21621 out_q: usize,
21622 out_kv: usize,
21623 ) -> Result<(), Box<dyn std::error::Error>> {
21624 const ROWS_PER_BLOCK: u32 = 4; let rows = out_q + 2 * out_kv;
21626 let nblk = in_f / 32;
21627 if !in_f.is_multiple_of(32)
21628 || aq.len() < in_f
21629 || ad.len() < nblk
21630 || yq.len() < out_q
21631 || yk.len() < out_kv
21632 || yv.len() < out_kv
21633 || wq.len() < out_q * nblk * 34
21634 || wk.len() < out_kv * nblk * 34
21635 || wv.len() < out_kv * nblk * 34
21636 {
21637 return Err(
21638 format!("q8_0 qkv rp geometry in={in_f} out_q={out_q} out_kv={out_kv}").into(),
21639 );
21640 }
21641 let f = self.func("qmatvec_q8_0_qkv_rp");
21642 let cfg = LaunchConfig {
21643 grid_dim: ((rows as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
21644 block_dim: (32, ROWS_PER_BLOCK, 1),
21645 shared_mem_bytes: 0,
21646 };
21647 let (ini, oq, okv) = (in_f as i32, out_q as i32, out_kv as i32);
21648 let __s_b = self.gpu.stream();
21649 let mut b = __s_b.launch_builder(&f);
21650 b.arg(wq)
21651 .arg(wk)
21652 .arg(wv)
21653 .arg(aq)
21654 .arg(ad)
21655 .arg(yq)
21656 .arg(yk)
21657 .arg(yv)
21658 .arg(&ini)
21659 .arg(&oq)
21660 .arg(&okv);
21661 unsafe {
21662 b.launch(cfg)?;
21663 }
21664 Ok(())
21665 }
21666
21667 #[allow(clippy::too_many_arguments)]
21671 pub fn qmatvec_q8_0_b4_rp_into(
21672 &self,
21673 w: [&CudaSlice<u8>; 4],
21674 aq: &CudaSlice<i8>,
21675 ad: &CudaSlice<f32>,
21676 y: &mut CudaSlice<f32>,
21677 block_cols: usize,
21678 out_f: usize,
21679 ) -> Result<(), Box<dyn std::error::Error>> {
21680 const ROWS_PER_BLOCK: u32 = 4; let nblk = block_cols / 32;
21682 if !block_cols.is_multiple_of(32)
21683 || aq.len() < 4 * block_cols
21684 || ad.len() < 4 * nblk
21685 || y.len() < out_f
21686 || w.iter().any(|p| p.len() < out_f * nblk * 34)
21687 {
21688 return Err(format!("q8_0 b4 rp geometry block_cols={block_cols} out={out_f}").into());
21689 }
21690 let f = self.func("qmatvec_q8_0_b4_rp");
21691 let cfg = LaunchConfig {
21692 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
21693 block_dim: (32, ROWS_PER_BLOCK, 1),
21694 shared_mem_bytes: 0,
21695 };
21696 let (bc, of) = (block_cols as i32, out_f as i32);
21697 let __s_b = self.gpu.stream();
21698 let mut b = __s_b.launch_builder(&f);
21699 b.arg(w[0])
21700 .arg(w[1])
21701 .arg(w[2])
21702 .arg(w[3])
21703 .arg(aq)
21704 .arg(ad)
21705 .arg(y)
21706 .arg(&bc)
21707 .arg(&of);
21708 unsafe {
21709 b.launch(cfg)?;
21710 }
21711 Ok(())
21712 }
21713
21714 #[allow(clippy::map_entry)] fn matvec_bf16_via_q8_mirror_t(
21718 &self,
21719 data: &CudaSlice<u8>,
21720 x: &CudaSlice<f32>,
21721 y: &mut CudaSlice<f32>,
21722 in_f: usize,
21723 out_f: usize,
21724 t: usize,
21725 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
21726 use cudarc::driver::DevicePtr;
21727 let key = {
21728 let s = self.gpu.stream();
21729 let (p, _g) = data.device_ptr(&s);
21730 (p, in_f as u32, out_f as u32)
21731 };
21732 {
21733 let mut mirrors = self
21734 .w8_mirrors
21735 .lock()
21736 .map_err(|_| "w8 mirror map is poisoned")?;
21737 if !mirrors.contains_key(&key) {
21738 let mut interleaved = self.alloc_u8_uninit(out_f * Self::q8_0_row_bytes(in_f))?;
21739 self.encode_q8_0_from_bf16(data, &mut interleaved, in_f, out_f)?;
21740 let planar = self.build_q8_rp4_raw(&interleaved, in_f, out_f)?;
21741 mirrors.insert(key, planar);
21742 }
21743 }
21744 let nblk = in_f / 32;
21745 let akey = in_f * 64 + t.min(32);
21747 {
21748 let mut act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
21749 if !act.contains_key(&akey) {
21750 let aq = self.alloc_i8_uninit(32 * in_f)?;
21751 let ad = self.alloc_uninit::<f32>(32 * nblk)?;
21752 act.insert(akey, (aq, ad));
21753 }
21754 let (aq, ad) = act.get_mut(&akey).expect("just inserted");
21755 self.quantize_q8_1_into(x, t, in_f, aq, ad)?;
21756 }
21757 let mirrors = self
21758 .w8_mirrors
21759 .lock()
21760 .map_err(|_| "w8 mirror map is poisoned")?;
21761 let act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
21762 let mirror = mirrors.get(&key).expect("built above");
21763 let (aq, ad) = act.get(&akey).expect("built above");
21764 const ROWS_PER_BLOCK: u32 = 4;
21765 let (ini, of) = (in_f as i32, out_f as i32);
21766 if q8t_wonce_on() && t <= 32 {
21770 let f = self.func(if t <= 8 {
21771 "qmatvec_q8_0_rows_tw"
21772 } else {
21773 "qmatvec_q8_0_rows_tw32"
21774 });
21775 let cfg = LaunchConfig {
21776 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
21777 block_dim: (32, ROWS_PER_BLOCK, 1),
21778 shared_mem_bytes: 0,
21779 };
21780 let ti = t as i32;
21781 let __s_b = self.gpu.stream();
21782 let mut b = __s_b.launch_builder(&f);
21783 b.arg(mirror)
21784 .arg(aq)
21785 .arg(ad)
21786 .arg(&mut *y)
21787 .arg(&ini)
21788 .arg(&of)
21789 .arg(&ti);
21790 unsafe {
21791 b.launch(cfg)?;
21792 }
21793 return Ok(Some(()));
21794 }
21795 let f = self.func("qmatvec_q8_0_rows_t");
21796 let cfg = LaunchConfig {
21797 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), t as u32, 1),
21798 block_dim: (32, ROWS_PER_BLOCK, 1),
21799 shared_mem_bytes: 0,
21800 };
21801 let __s_b = self.gpu.stream();
21802 let mut b = __s_b.launch_builder(&f);
21803 b.arg(mirror)
21804 .arg(aq)
21805 .arg(ad)
21806 .arg(&mut *y)
21807 .arg(&ini)
21808 .arg(&of);
21809 unsafe {
21810 b.launch(cfg)?;
21811 }
21812 Ok(Some(()))
21813 }
21814
21815 #[allow(clippy::map_entry)] fn matvec_bf16_via_q8_mirror(
21819 &self,
21820 data: &CudaSlice<u8>,
21821 x: &CudaSlice<f32>,
21822 y: &mut CudaSlice<f32>,
21823 in_f: usize,
21824 out_f: usize,
21825 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
21826 use cudarc::driver::DevicePtr;
21827 let key = {
21828 let s = self.gpu.stream();
21829 let (p, _g) = data.device_ptr(&s);
21830 (p, in_f as u32, out_f as u32)
21831 };
21832 {
21833 let mut mirrors = self
21834 .w8_mirrors
21835 .lock()
21836 .map_err(|_| "w8 mirror map is poisoned")?;
21837 if !mirrors.contains_key(&key) {
21838 let mut interleaved = self.alloc_u8_uninit(out_f * Self::q8_0_row_bytes(in_f))?;
21839 self.encode_q8_0_from_bf16(data, &mut interleaved, in_f, out_f)?;
21840 let planar = self.build_q8_rp4_raw(&interleaved, in_f, out_f)?;
21841 mirrors.insert(key, planar);
21842 if std::env::var("MEMRA_W8_TRACE").as_deref() == Ok("1") {
21848 eprintln!(
21849 "[w8-mirror] built in_f={in_f} out_f={out_f} mirrors={}",
21850 mirrors.len()
21851 );
21852 }
21853 }
21854 }
21855 let nblk = in_f / 32;
21856 {
21857 let mut act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
21858 if !act.contains_key(&in_f) {
21859 let aq = self.alloc_uninit::<i8>(in_f)?;
21860 let ad = self.alloc_uninit::<f32>(nblk)?;
21861 act.insert(in_f, (aq, ad));
21862 }
21863 let (aq, ad) = act.get_mut(&in_f).expect("just inserted");
21864 self.quantize_q8_1_into(x, 1, in_f, aq, ad)?;
21865 }
21866 let mirrors = self
21867 .w8_mirrors
21868 .lock()
21869 .map_err(|_| "w8 mirror map is poisoned")?;
21870 let act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
21871 let mirror = mirrors.get(&key).expect("built above");
21872 let (aq, ad) = act.get(&in_f).expect("built above");
21873 self.qmatvec_mmvq_into(
21874 mirror,
21875 aq,
21876 ad,
21877 1,
21878 in_f,
21879 out_f,
21880 QT_Q8_0,
21881 Self::q8_0_row_bytes(in_f),
21882 1.0,
21883 true,
21884 y,
21885 )?;
21886 Ok(Some(()))
21887 }
21888
21889 #[allow(clippy::too_many_arguments)]
21894 pub fn qmatvec_q8_0_qkv_rp_t_into(
21895 &self,
21896 wq: &CudaSlice<u8>,
21897 wk: &CudaSlice<u8>,
21898 wv: &CudaSlice<u8>,
21899 aq: &CudaSlice<i8>,
21900 ad: &CudaSlice<f32>,
21901 yq: &mut CudaSlice<f32>,
21902 yk: &mut CudaSlice<f32>,
21903 yv: &mut CudaSlice<f32>,
21904 in_f: usize,
21905 out_q: usize,
21906 out_kv: usize,
21907 t: usize,
21908 ) -> Result<(), Box<dyn std::error::Error>> {
21909 const ROWS_PER_BLOCK: u32 = 4;
21910 let rows = out_q + 2 * out_kv;
21911 let nblk = in_f / 32;
21912 if !in_f.is_multiple_of(32)
21913 || t == 0
21914 || aq.len() < t * in_f
21915 || ad.len() < t * nblk
21916 || yq.len() < t * out_q
21917 || yk.len() < t * out_kv
21918 || yv.len() < t * out_kv
21919 {
21920 return Err(format!("q8_0 qkv rp_t geometry in={in_f} t={t}").into());
21921 }
21922 let (ini, oq, okv) = (in_f as i32, out_q as i32, out_kv as i32);
21923 if q8t_wonce_on() && t <= 32 {
21927 let f = self.func(if t <= 8 {
21928 "qmatvec_q8_0_qkv_rp_tw"
21929 } else {
21930 "qmatvec_q8_0_qkv_rp_tw32"
21931 });
21932 let cfg = LaunchConfig {
21933 grid_dim: ((rows as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
21934 block_dim: (32, ROWS_PER_BLOCK, 1),
21935 shared_mem_bytes: 0,
21936 };
21937 let ti = t as i32;
21938 let __s_b = self.gpu.stream();
21939 let mut b = __s_b.launch_builder(&f);
21940 b.arg(wq)
21941 .arg(wk)
21942 .arg(wv)
21943 .arg(aq)
21944 .arg(ad)
21945 .arg(yq)
21946 .arg(yk)
21947 .arg(yv)
21948 .arg(&ini)
21949 .arg(&oq)
21950 .arg(&okv)
21951 .arg(&ti);
21952 unsafe {
21953 b.launch(cfg)?;
21954 }
21955 return Ok(());
21956 }
21957 let f = self.func("qmatvec_q8_0_qkv_rp_t");
21958 let cfg = LaunchConfig {
21959 grid_dim: ((rows as u32).div_ceil(ROWS_PER_BLOCK), t as u32, 1),
21960 block_dim: (32, ROWS_PER_BLOCK, 1),
21961 shared_mem_bytes: 0,
21962 };
21963 let __s_b = self.gpu.stream();
21964 let mut b = __s_b.launch_builder(&f);
21965 b.arg(wq)
21966 .arg(wk)
21967 .arg(wv)
21968 .arg(aq)
21969 .arg(ad)
21970 .arg(yq)
21971 .arg(yk)
21972 .arg(yv)
21973 .arg(&ini)
21974 .arg(&oq)
21975 .arg(&okv);
21976 unsafe {
21977 b.launch(cfg)?;
21978 }
21979 Ok(())
21980 }
21981
21982 #[allow(clippy::too_many_arguments)]
21985 pub fn qmatvec_q8_0_b4_rp_t_into(
21986 &self,
21987 w: [&CudaSlice<u8>; 4],
21988 aq: &CudaSlice<i8>,
21989 ad: &CudaSlice<f32>,
21990 y: &mut CudaSlice<f32>,
21991 block_cols: usize,
21992 out_f: usize,
21993 t: usize,
21994 ) -> Result<(), Box<dyn std::error::Error>> {
21995 const ROWS_PER_BLOCK: u32 = 4;
21996 let nblk = block_cols / 32;
21997 if !block_cols.is_multiple_of(32)
21998 || t == 0
21999 || aq.len() < t * 4 * block_cols
22000 || ad.len() < t * 4 * nblk
22001 || y.len() < t * out_f
22002 {
22003 return Err(format!("q8_0 b4 rp_t geometry cols={block_cols} t={t}").into());
22004 }
22005 let (bc, of) = (block_cols as i32, out_f as i32);
22006 if q8t_wonce_on() && t <= 32 {
22009 let f = self.func(if t <= 8 {
22010 "qmatvec_q8_0_b4_rp_tw"
22011 } else {
22012 "qmatvec_q8_0_b4_rp_tw32"
22013 });
22014 let cfg = LaunchConfig {
22015 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
22016 block_dim: (32, ROWS_PER_BLOCK, 1),
22017 shared_mem_bytes: 0,
22018 };
22019 let ti = t as i32;
22020 let __s_b = self.gpu.stream();
22021 let mut b = __s_b.launch_builder(&f);
22022 b.arg(w[0])
22023 .arg(w[1])
22024 .arg(w[2])
22025 .arg(w[3])
22026 .arg(aq)
22027 .arg(ad)
22028 .arg(y)
22029 .arg(&bc)
22030 .arg(&of)
22031 .arg(&ti);
22032 unsafe {
22033 b.launch(cfg)?;
22034 }
22035 return Ok(());
22036 }
22037 let f = self.func("qmatvec_q8_0_b4_rp_t");
22038 let cfg = LaunchConfig {
22039 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), t as u32, 1),
22040 block_dim: (32, ROWS_PER_BLOCK, 1),
22041 shared_mem_bytes: 0,
22042 };
22043 let __s_b = self.gpu.stream();
22044 let mut b = __s_b.launch_builder(&f);
22045 b.arg(w[0])
22046 .arg(w[1])
22047 .arg(w[2])
22048 .arg(w[3])
22049 .arg(aq)
22050 .arg(ad)
22051 .arg(y)
22052 .arg(&bc)
22053 .arg(&of);
22054 unsafe {
22055 b.launch(cfg)?;
22056 }
22057 Ok(())
22058 }
22059
22060 #[allow(clippy::map_entry)] fn matvec_bf16_view_via_q8_mirror(
22071 &self,
22072 data: &cudarc::driver::CudaView<'_, u8>,
22073 x: &CudaSlice<f32>,
22074 y: &mut CudaSlice<f32>,
22075 in_f: usize,
22076 out_f: usize,
22077 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
22078 use cudarc::driver::DevicePtr;
22079 let key = {
22080 let s = self.gpu.stream();
22081 let (p, _g) = data.device_ptr(&s);
22082 (p, in_f as u32, out_f as u32)
22083 };
22084 {
22085 let mut mirrors = self
22086 .w8_mirrors
22087 .lock()
22088 .map_err(|_| "w8 mirror map is poisoned")?;
22089 if !mirrors.contains_key(&key) {
22090 let mut interleaved = self.alloc_u8_uninit(out_f * Self::q8_0_row_bytes(in_f))?;
22091 self.encode_q8_0_from_bf16_view(data, &mut interleaved, in_f, out_f)?;
22092 let planar = self.build_q8_rp4_raw(&interleaved, in_f, out_f)?;
22093 mirrors.insert(key, planar);
22094 eprintln!("[w8-view] mirror built in_f={in_f} out_f={out_f}");
22098 }
22099 }
22100 let nblk = in_f / 32;
22101 {
22102 let mut act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
22103 if !act.contains_key(&in_f) {
22104 let aq = self.alloc_uninit::<i8>(in_f)?;
22105 let ad = self.alloc_uninit::<f32>(nblk)?;
22106 act.insert(in_f, (aq, ad));
22107 }
22108 let (aq, ad) = act.get_mut(&in_f).expect("just inserted");
22109 self.quantize_q8_1_into(x, 1, in_f, aq, ad)?;
22110 }
22111 let mirrors = self
22112 .w8_mirrors
22113 .lock()
22114 .map_err(|_| "w8 mirror map is poisoned")?;
22115 let act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
22116 let mirror = mirrors.get(&key).expect("built above");
22117 let (aq, ad) = act.get(&in_f).expect("built above");
22118 self.qmatvec_mmvq_into(
22119 mirror,
22120 aq,
22121 ad,
22122 1,
22123 in_f,
22124 out_f,
22125 QT_Q8_0,
22126 Self::q8_0_row_bytes(in_f),
22127 1.0,
22128 true,
22129 y,
22130 )?;
22131 Ok(Some(()))
22132 }
22133
22134 pub fn matvec_bf16_into(
22135 &self,
22136 data: &CudaSlice<u8>,
22137 x: &CudaSlice<f32>,
22138 y: &mut CudaSlice<f32>,
22139 in_f: usize,
22140 out_f: usize,
22141 ) -> Result<(), Box<dyn std::error::Error>> {
22142 if data.len() != in_f * out_f * 2
22143 || x.len() < in_f
22144 || !in_f.is_multiple_of(8)
22145 || y.len() < out_f
22146 {
22147 return Err(format!(
22148 "matvec_bf16_into geometry bytes={} x={} y={} in={in_f} out={out_f}",
22149 data.len(),
22150 x.len(),
22151 y.len()
22152 )
22153 .into());
22154 }
22155 if step_tp_w8_on()
22162 && w8_hybrid_on()
22163 && in_f.is_multiple_of(32)
22164 && out_f >= 64
22165 && let Some(()) = self.matvec_bf16_via_q8_mirror(data, x, y, in_f, out_f)?
22166 {
22167 return Ok(());
22168 }
22169 static X4: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
22173 let x4 = *X4.get_or_init(|| std::env::var("MEMRA_DOWN_X4").as_deref() == Ok("1"))
22174 && in_f <= 2048;
22175 if x4 {
22176 let f = self.func("matvec_bf16_f32acc_x4");
22177 let cfg = LaunchConfig {
22178 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
22179 block_dim: (mmv_block(), 1, 1),
22180 shared_mem_bytes: 0,
22181 };
22182 let (ini, outi) = (in_f as i32, out_f as i32);
22183 let __s_b = self.gpu.stream();
22184 let mut b = __s_b.launch_builder(&f);
22185 b.arg(data).arg(x).arg(y).arg(&ini).arg(&outi);
22186 unsafe {
22187 b.launch(cfg)?;
22188 }
22189 return Ok(());
22190 }
22191 let f = self.func("matvec_bf16_f32acc");
22192 let cfg = LaunchConfig {
22193 grid_dim: (out_f as u32, 1, 1),
22194 block_dim: (mmv_block(), 1, 1),
22195 shared_mem_bytes: 0,
22196 };
22197 let ini = in_f as i32;
22198 let __s_b = self.gpu.stream();
22199 let mut b = __s_b.launch_builder(&f);
22200 b.arg(data).arg(x).arg(y).arg(&ini);
22201 unsafe {
22202 b.launch(cfg)?;
22203 }
22204 Ok(())
22205 }
22206
22207 pub fn matvec_bf16_views_into(
22212 &self,
22213 data: &CudaSlice<u8>,
22214 x: &cudarc::driver::CudaView<'_, f32>,
22215 y: &mut cudarc::driver::CudaViewMut<'_, f32>,
22216 in_f: usize,
22217 out_f: usize,
22218 ) -> Result<(), Box<dyn std::error::Error>> {
22219 if data.len() != in_f * out_f * 2
22220 || x.len() < in_f
22221 || !in_f.is_multiple_of(8)
22222 || y.len() < out_f
22223 {
22224 return Err(format!(
22225 "matvec_bf16_views_into geometry bytes={} x={} y={} in={in_f} out={out_f}",
22226 data.len(),
22227 x.len(),
22228 y.len()
22229 )
22230 .into());
22231 }
22232 let f = self.func("matvec_bf16_f32acc");
22233 let cfg = LaunchConfig {
22234 grid_dim: (out_f as u32, 1, 1),
22235 block_dim: (mmv_block(), 1, 1),
22236 shared_mem_bytes: 0,
22237 };
22238 let ini = in_f as i32;
22239 let __s_b = self.gpu.stream();
22240 let mut b = __s_b.launch_builder(&f);
22241 b.arg(data).arg(x).arg(y).arg(&ini);
22242 unsafe {
22243 b.launch(cfg)?;
22244 }
22245 Ok(())
22246 }
22247
22248 pub fn matvec_bf16_view_into(
22251 &self,
22252 data: &cudarc::driver::CudaView<'_, u8>,
22253 x: &CudaSlice<f32>,
22254 y: &mut CudaSlice<f32>,
22255 in_f: usize,
22256 out_f: usize,
22257 ) -> Result<(), Box<dyn std::error::Error>> {
22258 if data.len() != in_f * out_f * 2
22259 || x.len() < in_f
22260 || !in_f.is_multiple_of(8)
22261 || y.len() < out_f
22262 {
22263 return Err(format!(
22264 "matvec_bf16_view_into geometry bytes={} x={} y={} in={in_f} out={out_f}",
22265 data.len(),
22266 x.len(),
22267 y.len()
22268 )
22269 .into());
22270 }
22271 if w8_view_on()
22272 && step_tp_w8_on()
22273 && w8_hybrid_on()
22274 && in_f.is_multiple_of(32)
22275 && out_f >= 64
22276 && let Some(()) = self.matvec_bf16_view_via_q8_mirror(data, x, y, in_f, out_f)?
22277 {
22278 return Ok(());
22279 }
22280 let f = self.func("matvec_bf16_f32acc");
22281 let cfg = LaunchConfig {
22282 grid_dim: (out_f as u32, 1, 1),
22283 block_dim: (mmv_block(), 1, 1),
22284 shared_mem_bytes: 0,
22285 };
22286 let ini = in_f as i32;
22287 let __s_b = self.gpu.stream();
22288 let mut b = __s_b.launch_builder(&f);
22289 b.arg(data).arg(x).arg(y).arg(&ini);
22290 unsafe {
22291 b.launch(cfg)?;
22292 }
22293 Ok(())
22294 }
22295
22296 pub fn matvec_bf16_raw_out(
22299 &self,
22300 w: &CudaSlice<u8>,
22301 x: &CudaSlice<f32>,
22302 y_raw: u64,
22303 in_f: usize,
22304 out_f: usize,
22305 ) -> Result<(), Box<dyn std::error::Error>> {
22306 if w.len() != in_f * out_f * 2 || x.len() < in_f || !in_f.is_multiple_of(8) || y_raw == 0 {
22307 return Err("matvec_bf16_raw_out geometry".into());
22308 }
22309 let f = self.func("matvec_bf16_f32acc");
22310 let cfg = LaunchConfig {
22311 grid_dim: (out_f as u32, 1, 1),
22312 block_dim: (mmv_block(), 1, 1),
22313 shared_mem_bytes: 0,
22314 };
22315 let ini = in_f as i32;
22316 let __s_b = self.gpu.stream();
22317 let mut b = __s_b.launch_builder(&f);
22318 b.arg(w).arg(x).arg(&y_raw).arg(&ini);
22319 unsafe {
22320 b.launch(cfg)?;
22321 }
22322 Ok(())
22323 }
22324
22325 pub fn add3_raw(
22329 &self,
22330 a: &CudaSlice<f32>,
22331 b: &CudaSlice<f32>,
22332 sh_raw: u64,
22333 scale_raw: u64,
22334 dst: &mut CudaSlice<f32>,
22335 n: usize,
22336 ) -> Result<(), Box<dyn std::error::Error>> {
22337 if a.len() < n || b.len() < n || dst.len() < n || sh_raw == 0 || scale_raw == 0 {
22338 return Err("add3_raw geometry".into());
22339 }
22340 let f = self.func("add3_f32");
22341 let cfg = LaunchConfig {
22342 grid_dim: ((n as u32).div_ceil(256), 1, 1),
22343 block_dim: (256, 1, 1),
22344 shared_mem_bytes: 0,
22345 };
22346 let ni = n as i32;
22347 let __s_b = self.gpu.stream();
22348 let mut bld = __s_b.launch_builder(&f);
22349 bld.arg(a)
22350 .arg(b)
22351 .arg(&sh_raw)
22352 .arg(&scale_raw)
22353 .arg(dst)
22354 .arg(&ni);
22355 unsafe {
22356 bld.launch(cfg)?;
22357 }
22358 Ok(())
22359 }
22360
22361 pub fn matvec_bf16_down_addscale_into(
22364 &self,
22365 w: &CudaSlice<u8>,
22366 x: &CudaSlice<f32>,
22367 scale: &CudaSlice<f32>,
22368 dst: &mut CudaSlice<f32>,
22369 in_f: usize,
22370 out_f: usize,
22371 ) -> Result<(), Box<dyn std::error::Error>> {
22372 if w.len() != in_f * out_f * 2
22373 || x.len() < in_f
22374 || !in_f.is_multiple_of(8)
22375 || dst.len() < out_f
22376 || scale.is_empty()
22377 {
22378 return Err("matvec_bf16_down_addscale geometry".into());
22379 }
22380 let f = self.func("matvec_bf16_down_addscale");
22381 let cfg = LaunchConfig {
22382 grid_dim: (out_f as u32, 1, 1),
22383 block_dim: (mmv_block(), 1, 1),
22384 shared_mem_bytes: 0,
22385 };
22386 let ini = in_f as i32;
22387 let __s_b = self.gpu.stream();
22388 let mut b = __s_b.launch_builder(&f);
22389 b.arg(w).arg(x).arg(scale).arg(dst).arg(&ini);
22390 unsafe {
22391 b.launch(cfg)?;
22392 }
22393 Ok(())
22394 }
22395
22396 #[allow(clippy::too_many_arguments)]
22400 pub fn matvec_bf16_dual_silu_rows_into(
22401 &self,
22402 wg: &CudaSlice<u8>,
22403 wu: &CudaSlice<u8>,
22404 x: &CudaSlice<f32>,
22405 act: &mut CudaSlice<f32>,
22406 in_f: usize,
22407 out_f: usize,
22408 limit: Option<f32>,
22409 t: usize,
22410 ) -> Result<(), Box<dyn std::error::Error>> {
22411 if x.len() < t * in_f || act.len() < t * out_f || t == 0 || t > 32 {
22412 return Err("matvec_bf16_dual_silu_rows geometry".into());
22413 }
22414 let f = self.func("matvec_bf16_dual_silu_rows");
22415 let cfg = LaunchConfig {
22416 grid_dim: (out_f as u32, t as u32, 1),
22417 block_dim: (mmv_block(), 1, 1),
22418 shared_mem_bytes: 0,
22419 };
22420 let (ini, outi) = (in_f as i32, out_f as i32);
22421 let lim = limit.unwrap_or(0.0);
22422 let __s_b = self.gpu.stream();
22423 let mut b = __s_b.launch_builder(&f);
22424 b.arg(wg)
22425 .arg(wu)
22426 .arg(x)
22427 .arg(&mut *act)
22428 .arg(&ini)
22429 .arg(&outi)
22430 .arg(&lim);
22431 unsafe {
22432 b.launch(cfg)?;
22433 }
22434 Ok(())
22435 }
22436
22437 pub fn matvec_bf16_rows_into(
22439 &self,
22440 w: &CudaSlice<u8>,
22441 x: &CudaSlice<f32>,
22442 y: &mut CudaSlice<f32>,
22443 in_f: usize,
22444 out_f: usize,
22445 t: usize,
22446 ) -> Result<(), Box<dyn std::error::Error>> {
22447 if x.len() < t * in_f || y.len() < t * out_f || t == 0 || t > 32 || !in_f.is_multiple_of(8)
22448 {
22449 return Err("matvec_bf16_rows geometry".into());
22450 }
22451 if (2..=32).contains(&t)
22456 && step_tp_w8_on()
22457 && w8_hybrid_on()
22458 && in_f.is_multiple_of(32)
22459 && out_f >= 64
22460 && let Some(()) = self.matvec_bf16_via_q8_mirror_t(w, x, y, in_f, out_f, t)?
22461 {
22462 return Ok(());
22463 }
22464 if t == 1
22470 && step_tp_w8_on()
22471 && w8_hybrid_on()
22472 && in_f.is_multiple_of(32)
22473 && out_f >= 64
22474 && let Some(()) = self.matvec_bf16_via_q8_mirror(w, x, y, in_f, out_f)?
22475 {
22476 return Ok(());
22477 }
22478 if (2..=16).contains(&t) && bf16_tcols_wide_on() {
22486 if BF16_TCOLS_WIDE_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
22487 eprintln!(
22488 "[bf16-tcols-wide] engaged: t={t} in_f={in_f} out_f={out_f} rides the \
22489 weight-once tcols class (MEMRA_BF16_TCOLS_WIDE=1)"
22490 );
22491 }
22492 if t <= 8 {
22493 return self.matvec_bf16_tcols_into(w, x, y, in_f, out_f, t);
22494 }
22495 return self.matvec_bf16_tcols16_into(w, x, y, in_f, out_f, t);
22496 }
22497 let f = self.func("matvec_bf16_f32acc_x4_rows");
22498 let cfg = LaunchConfig {
22499 grid_dim: (out_f.div_ceil(4) as u32, t as u32, 1),
22500 block_dim: (mmv_block(), 1, 1),
22501 shared_mem_bytes: 0,
22502 };
22503 let (ini, outi) = (in_f as i32, out_f as i32);
22504 let __s_b = self.gpu.stream();
22505 let mut b = __s_b.launch_builder(&f);
22506 b.arg(w).arg(x).arg(&mut *y).arg(&ini).arg(&outi);
22507 unsafe {
22508 b.launch(cfg)?;
22509 }
22510 Ok(())
22511 }
22512
22513 #[allow(clippy::too_many_arguments)]
22519 pub fn matvec_bf16_col_range_into(
22520 &self,
22521 w: &CudaSlice<u8>,
22522 x: &CudaSlice<f32>,
22523 y: &mut CudaSlice<f32>,
22524 in_f: usize,
22525 out_f: usize,
22526 k_start: usize,
22527 k_len: usize,
22528 ) -> Result<(), Box<dyn std::error::Error>> {
22529 let weight_bytes = in_f
22530 .checked_mul(out_f)
22531 .and_then(|elements| elements.checked_mul(2))
22532 .ok_or("BF16 column-range matvec geometry")?;
22533 let k_end = k_start
22534 .checked_add(k_len)
22535 .ok_or("BF16 column-range matvec geometry")?;
22536 if w.len() < weight_bytes
22537 || x.len() < k_len
22538 || y.len() < out_f
22539 || in_f > i32::MAX as usize
22540 || out_f > i32::MAX as usize
22541 || k_start > i32::MAX as usize
22542 || k_len > i32::MAX as usize
22543 || in_f == 0
22544 || out_f == 0
22545 || k_len == 0
22546 || !in_f.is_multiple_of(8)
22547 || !k_start.is_multiple_of(8)
22548 || !k_len.is_multiple_of(8)
22549 || k_end > in_f
22550 {
22551 return Err("BF16 column-range matvec geometry".into());
22552 }
22553 let function = self.func("matvec_bf16_f32acc_x4_range");
22554 let config = LaunchConfig {
22555 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
22556 block_dim: (mmv_block(), 1, 1),
22557 shared_mem_bytes: 0,
22558 };
22559 let (in_f, out_f, k_start, k_len) =
22560 (in_f as i32, out_f as i32, k_start as i32, k_len as i32);
22561 let stream = self.gpu.stream();
22562 let mut launch = stream.launch_builder(&function);
22563 launch
22564 .arg(w)
22565 .arg(x)
22566 .arg(y)
22567 .arg(&in_f)
22568 .arg(&out_f)
22569 .arg(&k_start)
22570 .arg(&k_len);
22571 unsafe {
22572 launch.launch(config)?;
22573 }
22574 Ok(())
22575 }
22576
22577 pub fn matvec_bf16_tcols_into(
22585 &self,
22586 w: &CudaSlice<u8>,
22587 x: &CudaSlice<f32>,
22588 y: &mut CudaSlice<f32>,
22589 in_f: usize,
22590 out_f: usize,
22591 t: usize,
22592 ) -> Result<(), Box<dyn std::error::Error>> {
22593 if x.len() < t * in_f
22594 || y.len() < t * out_f
22595 || !(2..=8).contains(&t)
22596 || !in_f.is_multiple_of(8)
22597 {
22598 return Err("matvec_bf16_tcols geometry".into());
22599 }
22600 let rf = bf16_tcols_red_fused_on() && mmv_block().is_power_of_two();
22611 if rf
22612 && BF16_TCOLS_RED_FUSED_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
22613 == 0
22614 {
22615 eprintln!(
22616 "[bf16-tcols-red-fused] engaged: fused-t reduce tail, one barrier sequence \
22617 shared across the t token columns + intra-warp shuffles at the identical \
22618 pairing (MEMRA_BF16_TCOLS_RED_FUSED=1)"
22619 );
22620 }
22621 let x1 = bf16_tcols_x1_on();
22622 let (fname, grid_x) = match (x1, rf) {
22623 (true, true) => ("matvec_bf16_f32acc_x1_tcols_rf", out_f as u32),
22624 (true, false) => ("matvec_bf16_f32acc_x1_tcols", out_f as u32),
22625 (false, true) => ("matvec_bf16_f32acc_x4_tcols_rf", out_f.div_ceil(4) as u32),
22626 (false, false) => ("matvec_bf16_f32acc_x4_tcols", out_f.div_ceil(4) as u32),
22627 };
22628 if x1 && BF16_TCOLS_X1_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
22629 eprintln!(
22630 "[bf16-tcols-x1] engaged: one-row-per-block tcols grid \
22631 (MEMRA_BF16_TCOLS_X1=1)"
22632 );
22633 }
22634 let f = self.func(fname);
22635 let cfg = LaunchConfig {
22636 grid_dim: (grid_x, 1, 1),
22637 block_dim: (mmv_block(), 1, 1),
22638 shared_mem_bytes: if rf { (t as u32) * mmv_block() * 4 } else { 0 },
22639 };
22640 let (ini, outi, ti) = (in_f as i32, out_f as i32, t as i32);
22641 let __s_b = self.gpu.stream();
22642 let mut b = __s_b.launch_builder(&f);
22643 b.arg(w).arg(x).arg(&mut *y).arg(&ini).arg(&outi).arg(&ti);
22644 unsafe {
22645 b.launch(cfg)?;
22646 }
22647 Ok(())
22648 }
22649
22650 pub fn matvec_bf16_tcols16_into(
22656 &self,
22657 w: &CudaSlice<u8>,
22658 x: &CudaSlice<f32>,
22659 y: &mut CudaSlice<f32>,
22660 in_f: usize,
22661 out_f: usize,
22662 t: usize,
22663 ) -> Result<(), Box<dyn std::error::Error>> {
22664 if x.len() < t * in_f
22665 || y.len() < t * out_f
22666 || !(9..=16).contains(&t)
22667 || !in_f.is_multiple_of(8)
22668 {
22669 return Err("matvec_bf16_tcols16 geometry".into());
22670 }
22671 let rf = bf16_tcols_red_fused_on() && mmv_block().is_power_of_two();
22675 if rf
22676 && BF16_TCOLS_RED_FUSED_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
22677 == 0
22678 {
22679 eprintln!(
22680 "[bf16-tcols-red-fused] engaged: fused-t reduce tail, one barrier sequence \
22681 shared across the t token columns + intra-warp shuffles at the identical \
22682 pairing (MEMRA_BF16_TCOLS_RED_FUSED=1)"
22683 );
22684 }
22685 let f = self.func(if rf {
22686 "matvec_bf16_f32acc_x4_tcols16_rf"
22687 } else {
22688 "matvec_bf16_f32acc_x4_tcols16"
22689 });
22690 let cfg = LaunchConfig {
22691 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
22692 block_dim: (mmv_block(), 1, 1),
22693 shared_mem_bytes: if rf { (t as u32) * mmv_block() * 4 } else { 0 },
22694 };
22695 let (ini, outi, ti) = (in_f as i32, out_f as i32, t as i32);
22696 let __s_b = self.gpu.stream();
22697 let mut b = __s_b.launch_builder(&f);
22698 b.arg(w).arg(x).arg(&mut *y).arg(&ini).arg(&outi).arg(&ti);
22699 unsafe {
22700 b.launch(cfg)?;
22701 }
22702 Ok(())
22703 }
22704
22705 #[allow(clippy::too_many_arguments)]
22711 pub fn matvec_bf16_tcols_gate_kernel_into(
22712 &self,
22713 kernel: &str,
22714 w: &CudaSlice<u8>,
22715 x: &CudaSlice<f32>,
22716 y: &mut CudaSlice<f32>,
22717 in_f: usize,
22718 out_f: usize,
22719 t: usize,
22720 ) -> Result<(), Box<dyn std::error::Error>> {
22721 let (grid_x, t_max) = match kernel {
22722 "matvec_bf16_f32acc_x1_tcols_rf" | "matvec_bf16_f32acc_x1_tcols_rf_redshift" => {
22723 (out_f as u32, 8usize)
22724 }
22725 "matvec_bf16_f32acc_x4_tcols_rf" => (out_f.div_ceil(4) as u32, 8usize),
22726 "matvec_bf16_f32acc_x4_tcols16_rf" => (out_f.div_ceil(4) as u32, 16usize),
22727 _ => return Err("matvec_bf16_tcols_gate_kernel_into: unknown kernel".into()),
22728 };
22729 if x.len() < t * in_f
22730 || y.len() < t * out_f
22731 || !(1..=t_max).contains(&t)
22732 || !in_f.is_multiple_of(8)
22733 || !mmv_block().is_power_of_two()
22734 {
22735 return Err("matvec_bf16_tcols_gate_kernel geometry".into());
22736 }
22737 let f = self.func(kernel);
22738 let cfg = LaunchConfig {
22739 grid_dim: (grid_x, 1, 1),
22740 block_dim: (mmv_block(), 1, 1),
22741 shared_mem_bytes: (t as u32) * mmv_block() * 4,
22742 };
22743 let (ini, outi, ti) = (in_f as i32, out_f as i32, t as i32);
22744 let __s_b = self.gpu.stream();
22745 let mut b = __s_b.launch_builder(&f);
22746 b.arg(w).arg(x).arg(&mut *y).arg(&ini).arg(&outi).arg(&ti);
22747 unsafe {
22748 b.launch(cfg)?;
22749 }
22750 Ok(())
22751 }
22752
22753 pub fn matmul_rows_exact(
22761 &self,
22762 w: &crate::model::GpuTensor,
22763 x: &CudaSlice<f32>,
22764 m: usize,
22765 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22766 use crate::model::GpuTensor;
22767 if let GpuTensor::FloatBf16 { data, .. } = w
22768 && (2..=8).contains(&m)
22769 && Self::bf16_mmv_on()
22770 && w.in_features().is_multiple_of(8)
22771 && !(step_tp_w8_on() && w8_hybrid_on())
22772 {
22773 let (in_f, out_f) = (w.in_features(), w.out_features());
22774 let mut y = self.vws_uninit(m * out_f)?;
22777 self.matvec_bf16_tcols_into(data, x, &mut y, in_f, out_f, m)?;
22778 return Ok(y);
22779 }
22780 self.matmul_decode_exact(w, x, m)
22781 }
22782
22783 #[allow(clippy::too_many_arguments)] pub fn matvec_bf16_dual_silu_into(
22785 &self,
22786 wg: &CudaSlice<u8>,
22787 wu: &CudaSlice<u8>,
22788 x: &CudaSlice<f32>,
22789 act: &mut CudaSlice<f32>,
22790 in_f: usize,
22791 out_f: usize,
22792 limit: Option<f32>,
22793 ) -> Result<(), Box<dyn std::error::Error>> {
22794 if wg.len() != in_f * out_f * 2
22795 || wu.len() != in_f * out_f * 2
22796 || x.len() < in_f
22797 || !in_f.is_multiple_of(8)
22798 || act.len() < out_f
22799 {
22800 return Err("matvec_bf16_dual_silu geometry".into());
22801 }
22802 let f = self.func("matvec_bf16_dual_silu");
22803 let cfg = LaunchConfig {
22804 grid_dim: (out_f as u32, 1, 1),
22805 block_dim: (mmv_block(), 1, 1),
22806 shared_mem_bytes: 0,
22807 };
22808 let (ini, outi) = (in_f as i32, out_f as i32);
22809 let lim = limit.unwrap_or(0.0);
22810 let __s_b = self.gpu.stream();
22811 let mut b = __s_b.launch_builder(&f);
22812 b.arg(wg)
22813 .arg(wu)
22814 .arg(x)
22815 .arg(act)
22816 .arg(&ini)
22817 .arg(&outi)
22818 .arg(&lim);
22819 unsafe {
22820 b.launch(cfg)?;
22821 }
22822 Ok(())
22823 }
22824
22825 #[allow(clippy::too_many_arguments)]
22828 pub fn matvec_bf16_dual_view_into(
22829 &self,
22830 wg: &cudarc::driver::CudaView<'_, u8>,
22831 wu: &cudarc::driver::CudaView<'_, u8>,
22832 x: &CudaSlice<f32>,
22833 yg: &mut CudaSlice<f32>,
22834 yu: &mut CudaSlice<f32>,
22835 in_f: usize,
22836 out_f: usize,
22837 ) -> Result<(), Box<dyn std::error::Error>> {
22838 if wg.len() != in_f * out_f * 2
22839 || wu.len() != in_f * out_f * 2
22840 || x.len() < in_f
22841 || !in_f.is_multiple_of(8)
22842 || yg.len() < out_f
22843 || yu.len() < out_f
22844 {
22845 return Err(format!(
22846 "matvec_bf16_dual_view_into geometry wg={} wu={} x={} in={in_f} out={out_f}",
22847 wg.len(),
22848 wu.len(),
22849 x.len()
22850 )
22851 .into());
22852 }
22853 let f = self.func("matvec_bf16_dual");
22854 let cfg = LaunchConfig {
22855 grid_dim: ((2 * out_f) as u32, 1, 1),
22856 block_dim: (mmv_block(), 1, 1),
22857 shared_mem_bytes: 0,
22858 };
22859 let (ini, outi) = (in_f as i32, out_f as i32);
22860 let __s_b = self.gpu.stream();
22861 let mut b = __s_b.launch_builder(&f);
22862 b.arg(wg)
22863 .arg(wu)
22864 .arg(x)
22865 .arg(yg)
22866 .arg(yu)
22867 .arg(&ini)
22868 .arg(&outi);
22869 unsafe {
22870 b.launch(cfg)?;
22871 }
22872 Ok(())
22873 }
22874
22875 #[allow(clippy::too_many_arguments)]
22877 pub fn matvec_bf16_dual_into(
22878 &self,
22879 wg: &CudaSlice<u8>,
22880 wu: &CudaSlice<u8>,
22881 x: &CudaSlice<f32>,
22882 yg: &mut CudaSlice<f32>,
22883 yu: &mut CudaSlice<f32>,
22884 in_f: usize,
22885 out_f: usize,
22886 ) -> Result<(), Box<dyn std::error::Error>> {
22887 if wg.len() != in_f * out_f * 2
22888 || wu.len() != in_f * out_f * 2
22889 || x.len() < in_f
22890 || !in_f.is_multiple_of(8)
22891 || yg.len() < out_f
22892 || yu.len() < out_f
22893 {
22894 return Err(format!(
22895 "matvec_bf16_dual_into geometry wg={} wu={} x={} in={in_f} out={out_f}",
22896 wg.len(),
22897 wu.len(),
22898 x.len()
22899 )
22900 .into());
22901 }
22902 let f = self.func("matvec_bf16_dual");
22903 let cfg = LaunchConfig {
22904 grid_dim: ((2 * out_f) as u32, 1, 1),
22905 block_dim: (mmv_block(), 1, 1),
22906 shared_mem_bytes: 0,
22907 };
22908 let (ini, outi) = (in_f as i32, out_f as i32);
22909 let __s_b = self.gpu.stream();
22910 let mut b = __s_b.launch_builder(&f);
22911 b.arg(wg)
22912 .arg(wu)
22913 .arg(x)
22914 .arg(yg)
22915 .arg(yu)
22916 .arg(&ini)
22917 .arg(&outi);
22918 unsafe {
22919 b.launch(cfg)?;
22920 }
22921 Ok(())
22922 }
22923
22924 #[allow(dead_code)] pub(crate) fn matvec_bf16_dual(
22928 &self,
22929 wg: &CudaSlice<u8>,
22930 wu: &CudaSlice<u8>,
22931 x: &CudaSlice<f32>,
22932 in_f: usize,
22933 out_f: usize,
22934 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
22935 if wg.len() != in_f * out_f * 2
22936 || wu.len() != in_f * out_f * 2
22937 || x.len() < in_f
22938 || !in_f.is_multiple_of(8)
22939 {
22940 return Err(format!(
22941 "matvec_bf16_dual geometry wg={} wu={} x={} in={in_f} out={out_f}",
22942 wg.len(),
22943 wu.len(),
22944 x.len()
22945 )
22946 .into());
22947 }
22948 let mut yg = self.alloc_uninit::<f32>(out_f)?;
22949 let mut yu = self.alloc_uninit::<f32>(out_f)?;
22950 let f = self.func("matvec_bf16_dual");
22951 let cfg = LaunchConfig {
22952 grid_dim: ((2 * out_f) as u32, 1, 1),
22953 block_dim: (mmv_block(), 1, 1),
22954 shared_mem_bytes: 0,
22955 };
22956 let (ini, outi) = (in_f as i32, out_f as i32);
22957 let __s_b = self.gpu.stream();
22958 let mut b = __s_b.launch_builder(&f);
22959 b.arg(wg)
22960 .arg(wu)
22961 .arg(x)
22962 .arg(&mut yg)
22963 .arg(&mut yu)
22964 .arg(&ini)
22965 .arg(&outi);
22966 unsafe {
22967 b.launch(cfg)?;
22968 }
22969 Ok((yg, yu))
22970 }
22971
22972 #[allow(clippy::too_many_arguments)]
22973 #[allow(clippy::manual_is_multiple_of)] fn linear_bf16_chunked_inner(
22975 &self,
22976 x: &CudaSlice<f32>,
22977 data: &CudaSlice<u8>,
22978 m: usize,
22979 in_f: usize,
22980 out_f: usize,
22981 exact: bool,
22982 canonical_chunk_rows: Option<usize>,
22983 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22984 const CHUNK_BYTES: usize = 256 << 20;
22985 if m == 1
22988 && !exact
22989 && canonical_chunk_rows.is_none()
22990 && in_f.is_multiple_of(8)
22991 && Self::bf16_mmv_on()
22992 {
22993 return self.matvec_bf16(data, x, in_f, out_f);
22994 }
22995 if m >= 16
23000 && !exact
23001 && canonical_chunk_rows.is_none()
23002 && data.len() == in_f * out_f * 2
23003 && crate::f16_ffi::pp_bf16_enabled()
23004 {
23005 if let Some(y) = self.bf16_tc_gemm(data, x, m, in_f, out_f)? {
23008 return Ok(y);
23009 }
23010 }
23011 let row_bytes = in_f
23012 .checked_mul(std::mem::size_of::<f32>())
23013 .ok_or("BF16 chunk row byte count overflow")?;
23014 if row_bytes == 0 || out_f == 0 {
23015 return Err("BF16 chunk dimensions must be nonzero".into());
23016 }
23017 let max_chunk_rows = (CHUNK_BYTES / row_bytes).max(1).min(out_f);
23018 let chunk_rows = match canonical_chunk_rows {
23019 Some(0) => {
23020 return Err("canonical BF16 chunk rows must be nonzero".into());
23021 }
23022 Some(rows) if rows > max_chunk_rows => {
23023 return Err(format!(
23024 "canonical BF16 chunk rows {rows} exceed the {max_chunk_rows}-row scratch limit"
23025 )
23026 .into());
23027 }
23028 Some(rows) if out_f % rows != 0 => {
23029 return Err(format!(
23030 "BF16 output width {out_f} is not divisible by canonical {rows}-row chunks"
23031 )
23032 .into());
23033 }
23034 Some(rows) => rows,
23035 None => max_chunk_rows,
23036 };
23037 if chunk_rows >= out_f {
23038 let wf32 = self.bf16_to_f32(&data.slice(0..in_f * out_f * 2), in_f * out_f)?;
23039 return if exact {
23040 self.linear_decode_exact(x, &wf32, m, in_f, out_f)
23041 } else {
23042 self.linear(x, &wf32, m, in_f, out_f)
23043 };
23044 }
23045 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
23046 let mut r0 = 0usize;
23047 while r0 < out_f {
23048 let rows = chunk_rows.min(out_f - r0);
23049 let wslice = data.slice(r0 * in_f * 2..(r0 + rows) * in_f * 2);
23050 let wf32 = self.bf16_to_f32(&wslice, in_f * rows)?;
23051 let yc = if exact {
23052 self.linear_decode_exact(x, &wf32, m, in_f, rows)?
23053 } else {
23054 self.linear(x, &wf32, m, in_f, rows)?
23055 };
23056 for mi in 0..m {
23058 let src = yc.slice(mi * rows..(mi + 1) * rows);
23059 let mut dst = y.slice_mut(mi * out_f + r0..mi * out_f + r0 + rows);
23060 self.gpu.stream().memcpy_dtod(&src, &mut dst)?;
23061 }
23062 r0 += rows;
23063 }
23064 Ok(y)
23065 }
23066
23067 pub fn linear_bf16_resident(
23071 &self,
23072 x: &CudaSlice<f32>,
23073 data: &CudaSlice<u8>,
23074 m: usize,
23075 in_f: usize,
23076 out_f: usize,
23077 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23078 if data.len() != in_f * out_f * 2 {
23079 return Err(format!("resident BF16 bytes {} != {out_f}x{in_f}x2", data.len()).into());
23080 }
23081 self.linear_bf16_chunked(x, data, m, in_f, out_f, false, None)
23082 }
23083
23084 pub fn linear_bf16_resident_canonical_rows(
23090 &self,
23091 x: &CudaSlice<f32>,
23092 data: &CudaSlice<u8>,
23093 m: usize,
23094 in_f: usize,
23095 out_f: usize,
23096 canonical_chunk_rows: usize,
23097 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23098 if data.len() != in_f * out_f * 2 {
23099 return Err(format!("resident BF16 bytes {} != {out_f}x{in_f}x2", data.len()).into());
23100 }
23101 self.linear_bf16_chunked(x, data, m, in_f, out_f, false, Some(canonical_chunk_rows))
23102 }
23103
23104 pub fn linear_f32_resident_canonical_rows(
23109 &self,
23110 x: &CudaSlice<f32>,
23111 data: &CudaSlice<f32>,
23112 m: usize,
23113 in_f: usize,
23114 out_f: usize,
23115 canonical_chunk_rows: usize,
23116 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23117 self.linear_f32_resident_canonical_rows_inner(
23118 x,
23119 data,
23120 m,
23121 in_f,
23122 out_f,
23123 canonical_chunk_rows,
23124 false,
23125 )
23126 }
23127
23128 pub fn linear_f32_resident_canonical_rows_strided(
23134 &self,
23135 x: &CudaSlice<f32>,
23136 data: &CudaSlice<f32>,
23137 m: usize,
23138 in_f: usize,
23139 out_f: usize,
23140 canonical_chunk_rows: usize,
23141 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23142 self.linear_f32_resident_canonical_rows_inner(
23143 x,
23144 data,
23145 m,
23146 in_f,
23147 out_f,
23148 canonical_chunk_rows,
23149 true,
23150 )
23151 }
23152
23153 #[allow(clippy::too_many_arguments)]
23154 #[allow(clippy::manual_is_multiple_of)] fn linear_f32_resident_canonical_rows_inner(
23157 &self,
23158 x: &CudaSlice<f32>,
23159 data: &CudaSlice<f32>,
23160 m: usize,
23161 in_f: usize,
23162 out_f: usize,
23163 canonical_chunk_rows: usize,
23164 strided_output: bool,
23165 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23166 if data.len() != in_f * out_f {
23167 return Err(format!("resident F32 values {} != {out_f}x{in_f}", data.len()).into());
23168 }
23169 if canonical_chunk_rows == 0
23170 || canonical_chunk_rows > out_f
23171 || out_f % canonical_chunk_rows != 0
23172 {
23173 return Err(format!(
23174 "invalid canonical F32 chunk rows {canonical_chunk_rows} for output width {out_f}"
23175 )
23176 .into());
23177 }
23178 if canonical_chunk_rows == out_f {
23179 return self.linear(x, data, m, in_f, out_f);
23180 }
23181
23182 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
23183 let input = x.slice(0..x.len());
23184 for r0 in (0..out_f).step_by(canonical_chunk_rows) {
23185 let weights = data.slice(r0 * in_f..(r0 + canonical_chunk_rows) * in_f);
23186 if m == 1 {
23187 let mut destination = y.slice_mut(r0..r0 + canonical_chunk_rows);
23188 self.linear_device_into(
23189 &input,
23190 &weights,
23191 &mut destination,
23192 1,
23193 in_f,
23194 canonical_chunk_rows,
23195 )?;
23196 continue;
23197 }
23198 let chunk = self.linear_device(&input, &weights, m, in_f, canonical_chunk_rows)?;
23199 if strided_output {
23200 self.place_rows_strided(&chunk, &mut y, canonical_chunk_rows, m, out_f, r0)?;
23201 } else {
23202 for token in 0..m {
23203 let source = chunk
23204 .slice(token * canonical_chunk_rows..(token + 1) * canonical_chunk_rows);
23205 let mut destination =
23206 y.slice_mut(token * out_f + r0..token * out_f + r0 + canonical_chunk_rows);
23207 self.gpu.stream().memcpy_dtod(&source, &mut destination)?;
23208 }
23209 }
23210 }
23211 Ok(y)
23212 }
23213
23214 #[allow(clippy::manual_is_multiple_of)] pub fn linear_f32_resident_canonical_rows_t1_into(
23220 &self,
23221 x: &CudaSlice<f32>,
23222 data: &CudaSlice<f32>,
23223 y: &mut CudaSlice<f32>,
23224 in_f: usize,
23225 out_f: usize,
23226 canonical_chunk_rows: usize,
23227 ) -> Result<(), Box<dyn std::error::Error>> {
23228 if data.len() != in_f * out_f {
23229 return Err(format!("resident F32 values {} != {out_f}x{in_f}", data.len()).into());
23230 }
23231 if y.len() != out_f || x.len() != in_f {
23232 return Err(format!(
23233 "resident F32 t1 shapes x={} y={} != in {in_f} out {out_f}",
23234 x.len(),
23235 y.len()
23236 )
23237 .into());
23238 }
23239 if canonical_chunk_rows == 0
23240 || canonical_chunk_rows > out_f
23241 || out_f % canonical_chunk_rows != 0
23242 {
23243 return Err(format!(
23244 "invalid canonical F32 chunk rows {canonical_chunk_rows} for output width {out_f}"
23245 )
23246 .into());
23247 }
23248 let input = x.slice(0..x.len());
23249 for r0 in (0..out_f).step_by(canonical_chunk_rows) {
23250 let weights = data.slice(r0 * in_f..(r0 + canonical_chunk_rows) * in_f);
23251 let mut destination = y.slice_mut(r0..r0 + canonical_chunk_rows);
23252 self.linear_device_into(
23253 &input,
23254 &weights,
23255 &mut destination,
23256 1,
23257 in_f,
23258 canonical_chunk_rows,
23259 )?;
23260 }
23261 Ok(())
23262 }
23263
23264 pub fn linear_t1_into(
23267 &self,
23268 x: &cudarc::driver::CudaView<'_, f32>,
23269 w: &cudarc::driver::CudaView<'_, f32>,
23270 y: &mut cudarc::driver::CudaViewMut<'_, f32>,
23271 in_f: usize,
23272 out_f: usize,
23273 ) -> Result<(), Box<dyn std::error::Error>> {
23274 self.linear_device_into(x, w, y, 1, in_f, out_f)
23275 }
23276
23277 pub fn linear_decode_exact(
23284 &self,
23285 x: &CudaSlice<f32>,
23286 w: &CudaSlice<f32>,
23287 m_tokens: usize,
23288 in_f: usize,
23289 out_f: usize,
23290 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23291 if m_tokens == 1 {
23292 return self.linear(x, w, 1, in_f, out_f);
23293 }
23294 let xv = self.view(x, m_tokens * in_f);
23295 let mut y = self.alloc_uninit::<f32>(m_tokens * out_f)?;
23296 for t in 0..m_tokens {
23297 let row = xv.slice(t * in_f..(t + 1) * in_f);
23298 let mut xr = self.alloc_uninit::<f32>(in_f)?;
23299 self.copy_view_into(&mut xr, 0, &row, in_f)?;
23300 let yr = self.linear(&xr, w, 1, in_f, out_f)?;
23301 self.copy_into(&mut y, t * out_f, &yr, out_f)?;
23302 }
23303 Ok(y)
23304 }
23305
23306 pub fn linear(
23307 &self,
23308 x: &CudaSlice<f32>,
23309 w: &CudaSlice<f32>,
23310 m_tokens: usize,
23311 in_f: usize,
23312 out_f: usize,
23313 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23314 self.linear_device(x, w, m_tokens, in_f, out_f)
23315 }
23316
23317 fn linear_device<I>(
23318 &self,
23319 x: &I,
23320 w: &I,
23321 m_tokens: usize,
23322 in_f: usize,
23323 out_f: usize,
23324 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>>
23325 where
23326 I: cudarc::driver::DevicePtr<f32>,
23327 {
23328 let mut c = self.alloc_uninit::<f32>(m_tokens * out_f)?; self.linear_device_into(x, w, &mut c, m_tokens, in_f, out_f)?;
23330 Ok(c)
23331 }
23332
23333 fn linear_device_into<I, O>(
23334 &self,
23335 x: &I,
23336 w: &I,
23337 c: &mut O,
23338 m_tokens: usize,
23339 in_f: usize,
23340 out_f: usize,
23341 ) -> Result<(), Box<dyn std::error::Error>>
23342 where
23343 I: cudarc::driver::DevicePtr<f32>,
23344 O: cudarc::driver::DevicePtrMut<f32>,
23345 {
23346 use cudarc::cublaslt::{Matmul, MatmulConfig};
23347 let cfg = MatmulConfig {
23348 transa: true,
23349 transb: false,
23350 transc: false,
23351 m: out_f as u64,
23352 n: m_tokens as u64,
23353 k: in_f as u64,
23354 alpha: 1.0,
23355 lda: in_f as i64,
23356 ldb: in_f as i64,
23357 beta: 0.0,
23358 ldc: out_f as i64,
23359 stride_a: None,
23360 stride_b: None,
23361 stride_c: None,
23362 stride_bias: None,
23363 batch_size: None,
23364 };
23365 let blas = self.gpu.blas();
23366 unsafe {
23367 blas.matmul(cfg, w, x, c, None, None)?;
23368 }
23369 Ok(())
23370 }
23371
23372 #[allow(clippy::too_many_arguments)] pub fn sdpa_naive(
23382 &self,
23383 q: &CudaSlice<f32>,
23384 k: &CudaSlice<f32>,
23385 v: &CudaSlice<f32>,
23386 o: &mut CudaSlice<f32>,
23387 head_dim: usize,
23388 n_head: usize,
23389 n_head_kv: usize,
23390 t: usize,
23391 t_kv: usize,
23392 scale: f32,
23393 causal: bool,
23394 ) -> Result<(), Box<dyn std::error::Error>> {
23395 if t_kv * 4 > SDPA_NAIVE_SMEM_MAX {
23396 return self.sdpa_naive_gmem(
23397 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
23398 );
23399 }
23400 let f = self.func("sdpa_naive_f32");
23401 let cfg = LaunchConfig {
23402 grid_dim: (n_head as u32, t as u32, 1),
23403 block_dim: (128, 1, 1),
23404 shared_mem_bytes: (t_kv * 4) as u32,
23405 };
23406 let (hd, nh, nhkv, ti, tkvi, cz) = (
23407 head_dim as i32,
23408 n_head as i32,
23409 n_head_kv as i32,
23410 t as i32,
23411 t_kv as i32,
23412 causal as i32,
23413 );
23414 let __s_b = self.gpu.stream();
23415 let mut b = __s_b.launch_builder(&f);
23416 b.arg(q)
23417 .arg(k)
23418 .arg(v)
23419 .arg(o)
23420 .arg(&hd)
23421 .arg(&nh)
23422 .arg(&nhkv)
23423 .arg(&ti)
23424 .arg(&tkvi)
23425 .arg(&scale)
23426 .arg(&cz);
23427 unsafe {
23428 b.launch(cfg)?;
23429 }
23430 Ok(())
23431 }
23432
23433 #[allow(clippy::too_many_arguments)]
23442 pub fn sdpa_naive_gmem(
23443 &self,
23444 q: &CudaSlice<f32>,
23445 k: &CudaSlice<f32>,
23446 v: &CudaSlice<f32>,
23447 o: &mut CudaSlice<f32>,
23448 head_dim: usize,
23449 n_head: usize,
23450 n_head_kv: usize,
23451 t: usize,
23452 t_kv: usize,
23453 scale: f32,
23454 causal: bool,
23455 ) -> Result<(), Box<dyn std::error::Error>> {
23456 let ws_len = n_head
23457 .checked_mul(t)
23458 .and_then(|x| x.checked_mul(t_kv))
23459 .ok_or("sdpa_naive_gmem: scores workspace size overflow")?;
23460 let ws_bytes = ws_len
23461 .checked_mul(std::mem::size_of::<f32>())
23462 .ok_or("sdpa_naive_gmem: scores workspace byte count overflow")?;
23463 if ws_bytes > SDPA_NAIVE_GMEM_WS_MAX {
23464 return Err(format!(
23465 "sdpa_naive_gmem: scores workspace {ws_bytes} bytes (heads {n_head} x T {t} x \
23466 T_kv {t_kv}) exceeds the {SDPA_NAIVE_GMEM_WS_MAX}-byte guard — this shape \
23467 needs a tiled/flash kernel, not the naive oracle"
23468 )
23469 .into());
23470 }
23471 let mut scores = self.uninit(ws_len)?;
23472 let f = self.func("sdpa_naive_gmem_f32");
23473 let cfg = LaunchConfig {
23474 grid_dim: (n_head as u32, t as u32, 1),
23475 block_dim: (128, 1, 1),
23476 shared_mem_bytes: 0,
23477 };
23478 let (hd, nh, nhkv, ti, tkvi, cz) = (
23479 head_dim as i32,
23480 n_head as i32,
23481 n_head_kv as i32,
23482 t as i32,
23483 t_kv as i32,
23484 causal as i32,
23485 );
23486 let __s_b = self.gpu.stream();
23487 let mut b = __s_b.launch_builder(&f);
23488 b.arg(q)
23489 .arg(k)
23490 .arg(v)
23491 .arg(o)
23492 .arg(&mut scores)
23493 .arg(&hd)
23494 .arg(&nh)
23495 .arg(&nhkv)
23496 .arg(&ti)
23497 .arg(&tkvi)
23498 .arg(&scale)
23499 .arg(&cz);
23500 unsafe {
23501 b.launch(cfg)?;
23502 }
23503 Ok(())
23504 }
23505
23506 #[allow(clippy::too_many_arguments)]
23511 pub fn sdpa_naive_island(
23512 &self,
23513 q: &CudaSlice<f32>,
23514 k: &CudaSlice<f32>,
23515 v: &CudaSlice<f32>,
23516 o: &mut CudaSlice<f32>,
23517 span_id: &CudaSlice<i32>,
23518 head_dim: usize,
23519 n_head: usize,
23520 n_head_kv: usize,
23521 t: usize,
23522 t_kv: usize,
23523 scale: f32,
23524 window: usize,
23525 ) -> Result<(), Box<dyn std::error::Error>> {
23526 let f = self.func("sdpa_naive_island_f32");
23527 let cfg = LaunchConfig {
23528 grid_dim: (n_head as u32, t as u32, 1),
23529 block_dim: (128, 1, 1),
23530 shared_mem_bytes: (t_kv * 4) as u32,
23531 };
23532 let (hd, nh, nhkv, ti, tkvi, wi) = (
23533 head_dim as i32,
23534 n_head as i32,
23535 n_head_kv as i32,
23536 t as i32,
23537 t_kv as i32,
23538 window as i32,
23539 );
23540 let __s_b = self.gpu.stream();
23541 let mut b = __s_b.launch_builder(&f);
23542 b.arg(q)
23543 .arg(k)
23544 .arg(v)
23545 .arg(o)
23546 .arg(span_id)
23547 .arg(&hd)
23548 .arg(&nh)
23549 .arg(&nhkv)
23550 .arg(&ti)
23551 .arg(&tkvi)
23552 .arg(&scale)
23553 .arg(&wi);
23554 unsafe {
23555 b.launch(cfg)?;
23556 }
23557 Ok(())
23558 }
23559
23560 #[allow(clippy::too_many_arguments)]
23562 pub fn sdpa_naive_w(
23563 &self,
23564 q: &CudaSlice<f32>,
23565 k: &CudaSlice<f32>,
23566 v: &CudaSlice<f32>,
23567 o: &mut CudaSlice<f32>,
23568 head_dim: usize,
23569 n_head: usize,
23570 n_head_kv: usize,
23571 t: usize,
23572 t_kv: usize,
23573 scale: f32,
23574 causal: bool,
23575 window: usize,
23576 ) -> Result<(), Box<dyn std::error::Error>> {
23577 let f = self.func("sdpa_naive_w_f32");
23578 let cfg = LaunchConfig {
23579 grid_dim: (n_head as u32, t as u32, 1),
23580 block_dim: (128, 1, 1),
23581 shared_mem_bytes: (t_kv * 4) as u32,
23582 };
23583 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
23584 head_dim as i32,
23585 n_head as i32,
23586 n_head_kv as i32,
23587 t as i32,
23588 t_kv as i32,
23589 causal as i32,
23590 window as i32,
23591 );
23592 let __s_b = self.gpu.stream();
23593 let mut b = __s_b.launch_builder(&f);
23594 b.arg(q)
23595 .arg(k)
23596 .arg(v)
23597 .arg(o)
23598 .arg(&hd)
23599 .arg(&nh)
23600 .arg(&nhkv)
23601 .arg(&ti)
23602 .arg(&tkvi)
23603 .arg(&scale)
23604 .arg(&cz)
23605 .arg(&wi);
23606 unsafe {
23607 b.launch(cfg)?;
23608 }
23609 Ok(())
23610 }
23611
23612 #[allow(clippy::too_many_arguments)]
23622 pub fn sdpa_naive_w_lo(
23623 &self,
23624 q: &CudaSlice<f32>,
23625 k: &CudaSlice<f32>,
23626 v: &CudaSlice<f32>,
23627 o: &mut CudaSlice<f32>,
23628 head_dim: usize,
23629 n_head: usize,
23630 n_head_kv: usize,
23631 t: usize,
23632 t_kv: usize,
23633 scale: f32,
23634 causal: bool,
23635 window: usize,
23636 ) -> Result<(), Box<dyn std::error::Error>> {
23637 let kv_lo = if window > 0 {
23638 (t_kv - t + 1).saturating_sub(window)
23639 } else {
23640 0
23641 };
23642 let smem = (t_kv - kv_lo) * 4;
23643 if smem > 48 * 1024 {
23644 return Err(format!(
23645 "sdpa_naive_w_lo: window {window} + T {t} rows need {smem} bytes of dynamic \
23646 shared memory (> 48KB launch bound) — this kernel clips the OLD side only; \
23647 a window this wide needs the multi-pass long-ctx kernel"
23648 )
23649 .into());
23650 }
23651 let f = self.func("sdpa_naive_w_lo_f32");
23652 let cfg = LaunchConfig {
23653 grid_dim: (n_head as u32, t as u32, 1),
23654 block_dim: (128, 1, 1),
23655 shared_mem_bytes: smem as u32,
23656 };
23657 let (hd, nh, nhkv, ti, tkvi, cz, wi, lo) = (
23658 head_dim as i32,
23659 n_head as i32,
23660 n_head_kv as i32,
23661 t as i32,
23662 t_kv as i32,
23663 causal as i32,
23664 window as i32,
23665 kv_lo as i32,
23666 );
23667 let __s_b = self.gpu.stream();
23668 let mut b = __s_b.launch_builder(&f);
23669 b.arg(q)
23670 .arg(k)
23671 .arg(v)
23672 .arg(o)
23673 .arg(&hd)
23674 .arg(&nh)
23675 .arg(&nhkv)
23676 .arg(&ti)
23677 .arg(&tkvi)
23678 .arg(&scale)
23679 .arg(&cz)
23680 .arg(&wi)
23681 .arg(&lo);
23682 unsafe {
23683 b.launch(cfg)?;
23684 }
23685 Ok(())
23686 }
23687
23688 #[allow(clippy::too_many_arguments)] pub fn sdpa_naive_view(
23691 &self,
23692 q: &CudaSlice<f32>,
23693 k: &cudarc::driver::CudaView<f32>,
23694 v: &cudarc::driver::CudaView<f32>,
23695 o: &mut CudaSlice<f32>,
23696 head_dim: usize,
23697 n_head: usize,
23698 n_head_kv: usize,
23699 t: usize,
23700 t_kv: usize,
23701 scale: f32,
23702 causal: bool,
23703 ) -> Result<(), Box<dyn std::error::Error>> {
23704 let f = self.func("sdpa_naive_f32");
23705 let cfg = LaunchConfig {
23706 grid_dim: (n_head as u32, t as u32, 1),
23707 block_dim: (128, 1, 1),
23708 shared_mem_bytes: (t_kv * 4) as u32,
23709 };
23710 let (hd, nh, nhkv, ti, tkvi, cz) = (
23711 head_dim as i32,
23712 n_head as i32,
23713 n_head_kv as i32,
23714 t as i32,
23715 t_kv as i32,
23716 causal as i32,
23717 );
23718 let __s_b = self.gpu.stream();
23719 let mut b = __s_b.launch_builder(&f);
23720 b.arg(q)
23721 .arg(k)
23722 .arg(v)
23723 .arg(o)
23724 .arg(&hd)
23725 .arg(&nh)
23726 .arg(&nhkv)
23727 .arg(&ti)
23728 .arg(&tkvi)
23729 .arg(&scale)
23730 .arg(&cz);
23731 unsafe {
23732 b.launch(cfg)?;
23733 }
23734 Ok(())
23735 }
23736
23737 #[allow(clippy::too_many_arguments)]
23745 pub fn fa_dequant_kv_view_f32(
23746 &self,
23747 k: &cudarc::driver::CudaView<u8>,
23748 v: &cudarc::driver::CudaView<u8>,
23749 kf: &mut CudaSlice<f32>,
23750 vf: &mut CudaSlice<f32>,
23751 kv_dim_k: usize,
23752 kv_dim_v: usize,
23753 t_kv: usize,
23754 k_tok_bytes: usize,
23755 v_tok_bytes: usize,
23756 g: bool,
23757 ) -> Result<(), Box<dyn std::error::Error>> {
23758 let f = if g {
23759 self.func_g("fa_dequant_kv_ws_f32")
23760 } else {
23761 self.func("fa_dequant_kv_ws_f32")
23762 };
23763 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
23764 #[allow(clippy::manual_div_ceil)]
23765 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
23767 let cfg = LaunchConfig {
23768 grid_dim: (nblk.max(1), 1, 1),
23769 block_dim: (256, 1, 1),
23770 shared_mem_bytes: 0,
23771 };
23772 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
23773 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
23774 let __s_b = self.gpu.stream();
23775 let mut b = __s_b.launch_builder(&f);
23776 b.arg(k)
23777 .arg(v)
23778 .arg(&mut *kf)
23779 .arg(&mut *vf)
23780 .arg(&kdk)
23781 .arg(&kdv)
23782 .arg(&tkvi)
23783 .arg(&ktb)
23784 .arg(&vtb);
23785 unsafe {
23786 b.launch(cfg)?;
23787 }
23788 Ok(())
23789 }
23790
23791 #[allow(clippy::too_many_arguments)]
23792 pub fn sdpa_naive_quantized_view(
23793 &self,
23794 q: &CudaSlice<f32>,
23795 k: &cudarc::driver::CudaView<u8>,
23796 v: &cudarc::driver::CudaView<u8>,
23797 o: &mut CudaSlice<f32>,
23798 head_dim: usize,
23799 n_head: usize,
23800 n_head_kv: usize,
23801 t: usize,
23802 t_kv: usize,
23803 scale: f32,
23804 causal: bool,
23805 k_tok_bytes: usize,
23806 v_tok_bytes: usize,
23807 ) -> Result<(), Box<dyn std::error::Error>> {
23808 let kv_dim = n_head_kv * head_dim;
23809 let mut kf = self.uninit(t_kv * kv_dim)?;
23810 let mut vf = self.uninit(t_kv * kv_dim)?;
23811 let f = self.func("fa_dequant_kv_ws_f32");
23812 let total = (2 * t_kv * kv_dim) as u64;
23813 #[allow(clippy::manual_div_ceil)]
23814 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
23816 let cfg = LaunchConfig {
23817 grid_dim: (nblk.max(1), 1, 1),
23818 block_dim: (256, 1, 1),
23819 shared_mem_bytes: 0,
23820 };
23821 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
23822 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
23823 let __s_b = self.gpu.stream();
23824 let mut b = __s_b.launch_builder(&f);
23825 b.arg(k)
23826 .arg(v)
23827 .arg(&mut kf)
23828 .arg(&mut vf)
23829 .arg(&kv_dim_i)
23830 .arg(&kv_dim_i)
23831 .arg(&t_kv_i)
23832 .arg(&k_tok_bytes_i)
23833 .arg(&v_tok_bytes_i);
23834 unsafe { b.launch(cfg)? };
23835 self.sdpa_naive(
23836 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
23837 )
23838 }
23839
23840 #[allow(clippy::too_many_arguments)]
23852 pub fn sdpa_naive_w_quantized_view(
23853 &self,
23854 q: &CudaSlice<f32>,
23855 k: &cudarc::driver::CudaView<u8>,
23856 v: &cudarc::driver::CudaView<u8>,
23857 o: &mut CudaSlice<f32>,
23858 head_dim: usize,
23859 n_head: usize,
23860 n_head_kv: usize,
23861 t: usize,
23862 t_kv: usize,
23863 scale: f32,
23864 causal: bool,
23865 window: usize,
23866 k_tok_bytes: usize,
23867 v_tok_bytes: usize,
23868 ) -> Result<(), Box<dyn std::error::Error>> {
23869 let kv_dim = n_head_kv * head_dim;
23870 let mut kf = self.uninit(t_kv * kv_dim)?;
23871 let mut vf = self.uninit(t_kv * kv_dim)?;
23872 let f = self.func("fa_dequant_kv_ws_f32");
23873 let total = (2 * t_kv * kv_dim) as u64;
23874 #[allow(clippy::manual_div_ceil)]
23875 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
23877 let cfg = LaunchConfig {
23878 grid_dim: (nblk.max(1), 1, 1),
23879 block_dim: (256, 1, 1),
23880 shared_mem_bytes: 0,
23881 };
23882 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
23883 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
23884 let __s_b = self.gpu.stream();
23885 let mut b = __s_b.launch_builder(&f);
23886 b.arg(k)
23887 .arg(v)
23888 .arg(&mut kf)
23889 .arg(&mut vf)
23890 .arg(&kv_dim_i)
23891 .arg(&kv_dim_i)
23892 .arg(&t_kv_i)
23893 .arg(&k_tok_bytes_i)
23894 .arg(&v_tok_bytes_i);
23895 unsafe { b.launch(cfg)? };
23896 self.sdpa_naive_w(
23897 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
23898 )
23899 }
23900
23901 #[allow(clippy::too_many_arguments)]
23905 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill(
23908 &self,
23909 q: &CudaSlice<f32>,
23910 k: &CudaSlice<f32>,
23911 v: &CudaSlice<f32>,
23912 o: &mut CudaSlice<f32>,
23913 head_dim: usize,
23914 n_head: usize,
23915 n_head_kv: usize,
23916 t: usize,
23917 t_kv: usize,
23918 scale: f32,
23919 causal: bool,
23920 ) -> Result<(), Box<dyn std::error::Error>> {
23921 if portable_mma_gated() {
23922 return self.sdpa_naive(
23923 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
23924 );
23925 }
23926 let fa3_on = head_dim == 256
23934 && causal
23935 && t == t_kv
23936 && match std::env::var("MEMRA_FA3").as_deref() {
23937 Ok("0") => false,
23938 Ok("1") => {
23942 refuse_portable_force("MEMRA_FA3=1", "the sm_90a fa3/bf16 kernels");
23943 true
23944 }
23945 _ => cfg!(memra_hopper_mma),
23946 };
23947 if fa3_on {
23948 let n = t * n_head * head_dim;
23949 let nkv = t * n_head_kv * head_dim;
23950 let mut q16 = self.alloc_u8_uninit(n * 2)?;
23951 let mut k16 = self.alloc_u8_uninit(nkv * 2)?;
23952 let mut v16 = self.alloc_u8_uninit(nkv * 2)?;
23953 self.f32_to_bf16_into(q, &mut q16, n)?;
23954 self.f32_to_bf16_into(k, &mut k16, nkv)?;
23955 self.f32_to_bf16_into(v, &mut v16, nkv)?;
23956 let rc = {
23957 use cudarc::driver::{DevicePtr, DevicePtrMut};
23958 let stream = self.gpu.stream();
23959 let (qp, _g1) = q16.device_ptr(&stream);
23960 let (kp, _g2) = k16.device_ptr(&stream);
23961 let (vp, _g3) = v16.device_ptr(&stream);
23962 let (op, _g4) = o.device_ptr_mut(&stream);
23963 unsafe {
23964 memra_fa3_prefill(
23965 qp as *const core::ffi::c_void,
23966 kp as *const core::ffi::c_void,
23967 vp as *const core::ffi::c_void,
23968 op as *mut f32,
23969 t as i32,
23970 n_head as i32,
23971 n_head_kv as i32,
23972 head_dim as i32,
23973 scale,
23974 stream.cu_stream() as *mut core::ffi::c_void,
23975 )
23976 }
23977 };
23978 if rc != 0 {
23979 return Err(format!("memra_fa3_prefill rc={rc}").into());
23980 }
23981 return Ok(());
23982 }
23983 static FA_P1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
23988 let fa_p1 = *FA_P1.get_or_init(|| std::env::var("MEMRA_FA_P1").as_deref() == Ok("1"));
23989 if fa_p1 && head_dim == 256 && !std::env::var("MEMRA_FA_FLOOR").is_ok() {
23990 const BLOCK_Q: usize = 64;
23991 const BKX: usize = 32;
23992 let f = self.func("fa_prefill_bf16_p1");
23993 let shmem = (2 * (2 * BKX * head_dim + BLOCK_Q * BKX)
23994 + 4 * (BLOCK_Q * BKX + 2 * BLOCK_Q)) as u32;
23995 use cudarc::driver::sys::CUfunction_attribute_enum as A;
23996 f.set_attribute(
23997 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
23998 shmem as i32,
23999 )?;
24000 let cfg = LaunchConfig {
24001 grid_dim: (
24002 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
24003 n_head as u32,
24004 1,
24005 ),
24006 block_dim: (32, 4, 1),
24007 shared_mem_bytes: shmem,
24008 };
24009 let (hd, nh, nhkv, ti, tkvi, cz) = (
24010 head_dim as i32,
24011 n_head as i32,
24012 n_head_kv as i32,
24013 t as i32,
24014 t_kv as i32,
24015 causal as i32,
24016 );
24017 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24018 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24019 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
24020 let __s_b = self.gpu.stream();
24021 let mut b = __s_b.launch_builder(&f);
24022 b.arg(&qb)
24023 .arg(&kb)
24024 .arg(&vb)
24025 .arg(o)
24026 .arg(&hd)
24027 .arg(&nh)
24028 .arg(&nhkv)
24029 .arg(&ti)
24030 .arg(&tkvi)
24031 .arg(&scale)
24032 .arg(&cz);
24033 unsafe {
24034 b.launch(cfg)?;
24035 }
24036 return Ok(());
24037 }
24038 const BK: usize = 32;
24044 let w2 = std::env::var("MEMRA_FA_PP_W2").as_deref() == Ok("1");
24047 let (block_q, warps, w2_sfx): (usize, u32, &str) =
24048 if w2 { (32, 2, "_w2") } else { (64, 4, "") };
24049 let hd_sfx = fa_hd_suffix(head_dim)?;
24053 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
24054 let bf16kv = !floor && !w2 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0");
24059 let (kb16, vb16) = if bf16kv {
24060 let n = t_kv * n_head_kv * head_dim;
24061 let mut kb = self.alloc_u8_uninit(n * 2)?;
24062 let mut vb = self.alloc_u8_uninit(n * 2)?;
24063 let fcv = self.func("f32_to_bf16_bulk");
24064 let ni = n as i64;
24065 let cfgc = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
24066 let __s_b = self.gpu.stream();
24067 let mut b = __s_b.launch_builder(&fcv);
24068 b.arg(k).arg(&mut kb).arg(&ni);
24069 unsafe {
24070 b.launch(cfgc)?;
24071 }
24072 let __s_b = self.gpu.stream();
24073 let mut b = __s_b.launch_builder(&fcv);
24074 b.arg(v).arg(&mut vb).arg(&ni);
24075 unsafe {
24076 b.launch(cfgc)?;
24077 }
24078 (Some(kb), Some(vb))
24079 } else {
24080 (None, None)
24081 };
24082 let f = self.func(&if bf16kv {
24083 format!("fa_prefill_bf16kv_pp{hd_sfx}")
24084 } else {
24085 format!(
24086 "fa_prefill_f32{}{}{hd_sfx}",
24087 if floor { "" } else { "_pp" },
24088 if floor { "" } else { w2_sfx }
24089 )
24090 });
24091 let kv_stages = if bf16kv { 2 } else { 1 };
24094 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
24095 + 4 * (block_q * BK + 2 * block_q)) as u32;
24096 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24097 f.set_attribute(
24098 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24099 shmem as i32,
24100 )?;
24101 let cfg = LaunchConfig {
24102 grid_dim: (
24103 (t as u32 + block_q as u32 - 1) / block_q as u32,
24104 n_head as u32,
24105 1,
24106 ),
24107 block_dim: (32, warps, 1),
24108 shared_mem_bytes: shmem,
24109 };
24110 let (hd, nh, nhkv, ti, tkvi, cz) = (
24111 head_dim as i32,
24112 n_head as i32,
24113 n_head_kv as i32,
24114 t as i32,
24115 t_kv as i32,
24116 causal as i32,
24117 );
24118 let __s_b = self.gpu.stream();
24119 let mut b = __s_b.launch_builder(&f);
24120 b.arg(q);
24121 match (&kb16, &vb16) {
24122 (Some(kb), Some(vb)) => {
24123 b.arg(kb).arg(vb);
24124 }
24125 _ => {
24126 b.arg(k).arg(v);
24127 }
24128 }
24129 b.arg(o)
24130 .arg(&hd)
24131 .arg(&nh)
24132 .arg(&nhkv)
24133 .arg(&ti)
24134 .arg(&tkvi)
24135 .arg(&scale)
24136 .arg(&cz);
24137 unsafe {
24138 b.launch(cfg)?;
24139 }
24140 Ok(())
24141 }
24142
24143 #[allow(clippy::too_many_arguments)]
24147 pub fn fa_prefill_w(
24148 &self,
24149 q: &CudaSlice<f32>,
24150 k: &CudaSlice<f32>,
24151 v: &CudaSlice<f32>,
24152 o: &mut CudaSlice<f32>,
24153 head_dim: usize,
24154 n_head: usize,
24155 n_head_kv: usize,
24156 t: usize,
24157 t_kv: usize,
24158 scale: f32,
24159 causal: bool,
24160 window: usize,
24161 ) -> Result<(), Box<dyn std::error::Error>> {
24162 if portable_mma_gated() {
24165 return self.sdpa_naive_w(
24166 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
24167 );
24168 }
24169 static FAW_F32: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24173 let faw_f32 =
24174 *FAW_F32.get_or_init(|| std::env::var("MEMRA_FAW_STAGE").as_deref() == Ok("f32"));
24175 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
24176 self.fa_prefill_w_arm(
24177 q,
24178 k,
24179 v,
24180 o,
24181 head_dim,
24182 n_head,
24183 n_head_kv,
24184 t,
24185 t_kv,
24186 scale,
24187 causal,
24188 window,
24189 floor || faw_f32,
24190 floor,
24191 )
24192 }
24193
24194 #[allow(clippy::too_many_arguments)]
24197 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_w_pre(
24199 &self,
24200 qb: &CudaSlice<u8>,
24201 kb: &CudaSlice<u8>,
24202 vb: &CudaSlice<u8>,
24203 o: &mut CudaSlice<f32>,
24204 head_dim: usize,
24205 n_head: usize,
24206 n_head_kv: usize,
24207 t: usize,
24208 t_kv: usize,
24209 scale: f32,
24210 causal: bool,
24211 window: usize,
24212 v_f16: bool,
24213 ) -> Result<(), Box<dyn std::error::Error>> {
24214 const BLOCK_Q: usize = 64;
24215 const BK: usize = 32;
24216 debug_assert_eq!(head_dim, 256);
24217 let hp = fa_f16pv_on()
24218 && faw_hp_on()
24219 && n_head.is_multiple_of(2)
24220 && (n_head / n_head_kv).is_multiple_of(2);
24221 debug_assert!(!v_f16 || hp, "f16 V emitted but the SWA hp arm is off");
24222 if hp {
24223 const BLOCK_QH: usize = 32;
24224 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
24227 let vh: &CudaSlice<u8> = if v_f16 {
24228 vb
24229 } else {
24230 let n = t_kv * n_head_kv * head_dim;
24231 if vguard.as_ref().map(|b| b.len() < n * 2).unwrap_or(true) {
24232 *vguard = Some(self.alloc_uninit::<u8>(n * 2)?);
24233 }
24234 self.bf16_to_f16_into(vb, n, vguard.as_mut().unwrap())?;
24235 vguard.as_ref().unwrap()
24236 };
24237 let f = self.func("fa_prefill_w_bf16_p1h2");
24238 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK) + 4 * (2 * BLOCK_QH)) as u32;
24239 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24240 f.set_attribute(
24241 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24242 shmem as i32,
24243 )?;
24244 let cfg = LaunchConfig {
24245 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
24246 block_dim: (32, 4, 1),
24247 shared_mem_bytes: shmem,
24248 };
24249 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24250 head_dim as i32,
24251 n_head as i32,
24252 n_head_kv as i32,
24253 t as i32,
24254 t_kv as i32,
24255 causal as i32,
24256 window as i32,
24257 );
24258 let __s_b = self.gpu.stream();
24259 let mut b = __s_b.launch_builder(&f);
24260 b.arg(qb)
24261 .arg(kb)
24262 .arg(vh)
24263 .arg(o)
24264 .arg(&hd)
24265 .arg(&nh)
24266 .arg(&nhkv)
24267 .arg(&ti)
24268 .arg(&tkvi)
24269 .arg(&scale)
24270 .arg(&cz)
24271 .arg(&wi);
24272 unsafe {
24273 b.launch(cfg)?;
24274 }
24275 return Ok(());
24276 }
24277 let f = self.func("fa_prefill_w_bf16_p1");
24278 let shmem =
24279 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
24280 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24281 f.set_attribute(
24282 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24283 shmem as i32,
24284 )?;
24285 let cfg = LaunchConfig {
24286 grid_dim: (
24287 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
24288 n_head as u32,
24289 1,
24290 ),
24291 block_dim: (32, 4, 1),
24292 shared_mem_bytes: shmem,
24293 };
24294 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24295 head_dim as i32,
24296 n_head as i32,
24297 n_head_kv as i32,
24298 t as i32,
24299 t_kv as i32,
24300 causal as i32,
24301 window as i32,
24302 );
24303 let __s_b = self.gpu.stream();
24304 let mut b = __s_b.launch_builder(&f);
24305 b.arg(qb)
24306 .arg(kb)
24307 .arg(vb)
24308 .arg(o)
24309 .arg(&hd)
24310 .arg(&nh)
24311 .arg(&nhkv)
24312 .arg(&ti)
24313 .arg(&tkvi)
24314 .arg(&scale)
24315 .arg(&cz)
24316 .arg(&wi);
24317 unsafe {
24318 b.launch(cfg)?;
24319 }
24320 Ok(())
24321 }
24322
24323 #[allow(clippy::too_many_arguments)]
24325 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_w_arm(
24327 &self,
24328 q: &CudaSlice<f32>,
24329 k: &CudaSlice<f32>,
24330 v: &CudaSlice<f32>,
24331 o: &mut CudaSlice<f32>,
24332 head_dim: usize,
24333 n_head: usize,
24334 n_head_kv: usize,
24335 t: usize,
24336 t_kv: usize,
24337 scale: f32,
24338 causal: bool,
24339 window: usize,
24340 f32_stage: bool,
24341 floor: bool,
24342 ) -> Result<(), Box<dyn std::error::Error>> {
24343 const BLOCK_Q: usize = 64;
24344 const BK: usize = 32;
24345 debug_assert_eq!(head_dim, 256, "fa_prefill_w is stamped hd256 only");
24346 static P1_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24350 let p1 = !floor
24351 && !f32_stage
24352 && *P1_ON.get_or_init(|| {
24353 std::env::var("MEMRA_FAW_P1")
24354 .map(|v| v != "0")
24355 .unwrap_or(true)
24356 });
24357 let hp = p1
24358 && fa_f16pv_on()
24359 && faw_hp_on()
24360 && n_head.is_multiple_of(2)
24361 && (n_head / n_head_kv).is_multiple_of(2);
24362 if hp {
24363 const BLOCK_QH: usize = 32;
24364 let f = self.func("fa_prefill_w_bf16_p1h2");
24365 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK) + 4 * (2 * BLOCK_QH)) as u32;
24366 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24367 f.set_attribute(
24368 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24369 shmem as i32,
24370 )?;
24371 let cfg = LaunchConfig {
24372 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
24373 block_dim: (32, 4, 1),
24374 shared_mem_bytes: shmem,
24375 };
24376 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24377 head_dim as i32,
24378 n_head as i32,
24379 n_head_kv as i32,
24380 t as i32,
24381 t_kv as i32,
24382 causal as i32,
24383 window as i32,
24384 );
24385 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24386 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24387 let vh = self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?;
24388 let __s_b = self.gpu.stream();
24389 let mut b = __s_b.launch_builder(&f);
24390 b.arg(&qb)
24391 .arg(&kb)
24392 .arg(&vh)
24393 .arg(o)
24394 .arg(&hd)
24395 .arg(&nh)
24396 .arg(&nhkv)
24397 .arg(&ti)
24398 .arg(&tkvi)
24399 .arg(&scale)
24400 .arg(&cz)
24401 .arg(&wi);
24402 unsafe {
24403 b.launch(cfg)?;
24404 }
24405 return Ok(());
24406 }
24407 if p1 {
24408 let f = self.func("fa_prefill_w_bf16_p1");
24409 let shmem =
24410 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
24411 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24412 f.set_attribute(
24413 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24414 shmem as i32,
24415 )?;
24416 let cfg = LaunchConfig {
24417 grid_dim: (
24418 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
24419 n_head as u32,
24420 1,
24421 ),
24422 block_dim: (32, 4, 1),
24423 shared_mem_bytes: shmem,
24424 };
24425 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24426 head_dim as i32,
24427 n_head as i32,
24428 n_head_kv as i32,
24429 t as i32,
24430 t_kv as i32,
24431 causal as i32,
24432 window as i32,
24433 );
24434 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24435 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24436 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
24437 let __s_b = self.gpu.stream();
24438 let mut b = __s_b.launch_builder(&f);
24439 b.arg(&qb)
24440 .arg(&kb)
24441 .arg(&vb)
24442 .arg(o)
24443 .arg(&hd)
24444 .arg(&nh)
24445 .arg(&nhkv)
24446 .arg(&ti)
24447 .arg(&tkvi)
24448 .arg(&scale)
24449 .arg(&cz)
24450 .arg(&wi);
24451 unsafe {
24452 b.launch(cfg)?;
24453 }
24454 return Ok(());
24455 }
24456 static G4_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24459 let g4 = !floor
24460 && !f32_stage
24461 && n_head_kv == 1
24462 && n_head.is_multiple_of(4)
24463 && *G4_ON.get_or_init(|| {
24464 std::env::var("MEMRA_FAW_G4")
24465 .map(|v| v != "0")
24466 .unwrap_or(true)
24467 });
24468 if g4 {
24469 const SP_M: usize = 16;
24470 static O2_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24473 let o2 = *O2_ON.get_or_init(|| {
24474 std::env::var("MEMRA_FAW_O2")
24475 .map(|v| v != "0")
24476 .unwrap_or(true)
24477 });
24478 let f = self.func(if o2 {
24479 "fa_prefill_w_bf16_g4o2"
24480 } else {
24481 "fa_prefill_w_bf16_g4"
24482 });
24483 let shmem = if o2 {
24484 (2 * (4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M)) as u32
24485 } else {
24486 (2 * (2 * BK * head_dim + 4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M))
24487 as u32
24488 };
24489 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24490 f.set_attribute(
24491 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24492 shmem as i32,
24493 )?;
24494 let cfg = LaunchConfig {
24495 grid_dim: ((t as u32).div_ceil(SP_M as u32), (n_head / 4) as u32, 1),
24496 block_dim: (32, 4, 1),
24497 shared_mem_bytes: shmem,
24498 };
24499 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24500 head_dim as i32,
24501 n_head as i32,
24502 n_head_kv as i32,
24503 t as i32,
24504 t_kv as i32,
24505 causal as i32,
24506 window as i32,
24507 );
24508 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24509 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24510 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
24511 let __s_b = self.gpu.stream();
24512 let mut b = __s_b.launch_builder(&f);
24513 b.arg(&qb)
24514 .arg(&kb)
24515 .arg(&vb)
24516 .arg(o)
24517 .arg(&hd)
24518 .arg(&nh)
24519 .arg(&nhkv)
24520 .arg(&ti)
24521 .arg(&tkvi)
24522 .arg(&scale)
24523 .arg(&cz)
24524 .arg(&wi);
24525 unsafe {
24526 b.launch(cfg)?;
24527 }
24528 return Ok(());
24529 }
24530 let f = self.func(if floor {
24531 "fa_prefill_w_f32"
24532 } else if f32_stage {
24533 "fa_prefill_w_f32_pp"
24534 } else {
24535 "fa_prefill_w_bf16_pp"
24536 });
24537 let shmem =
24538 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
24539 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24540 f.set_attribute(
24541 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24542 shmem as i32,
24543 )?;
24544 let cfg = LaunchConfig {
24545 grid_dim: (
24546 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
24547 n_head as u32,
24548 1,
24549 ),
24550 block_dim: (32, 4, 1),
24551 shared_mem_bytes: shmem,
24552 };
24553 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24554 head_dim as i32,
24555 n_head as i32,
24556 n_head_kv as i32,
24557 t as i32,
24558 t_kv as i32,
24559 causal as i32,
24560 window as i32,
24561 );
24562 if f32_stage {
24563 let __s_b = self.gpu.stream();
24564 let mut b = __s_b.launch_builder(&f);
24565 b.arg(q)
24566 .arg(k)
24567 .arg(v)
24568 .arg(o)
24569 .arg(&hd)
24570 .arg(&nh)
24571 .arg(&nhkv)
24572 .arg(&ti)
24573 .arg(&tkvi)
24574 .arg(&scale)
24575 .arg(&cz)
24576 .arg(&wi);
24577 unsafe {
24578 b.launch(cfg)?;
24579 }
24580 } else {
24581 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24582 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24583 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
24584 let __s_b = self.gpu.stream();
24585 let mut b = __s_b.launch_builder(&f);
24586 b.arg(&qb)
24587 .arg(&kb)
24588 .arg(&vb)
24589 .arg(o)
24590 .arg(&hd)
24591 .arg(&nh)
24592 .arg(&nhkv)
24593 .arg(&ti)
24594 .arg(&tkvi)
24595 .arg(&scale)
24596 .arg(&cz)
24597 .arg(&wi);
24598 unsafe {
24599 b.launch(cfg)?;
24600 }
24601 }
24602 Ok(())
24603 }
24604
24605 #[allow(clippy::too_many_arguments)]
24609 pub fn fa_prefill_hd512(
24610 &self,
24611 q: &CudaSlice<f32>,
24612 k: &CudaSlice<f32>,
24613 v: &CudaSlice<f32>,
24614 o: &mut CudaSlice<f32>,
24615 head_dim: usize,
24616 n_head: usize,
24617 n_head_kv: usize,
24618 t: usize,
24619 t_kv: usize,
24620 scale: f32,
24621 causal: bool,
24622 ) -> Result<(), Box<dyn std::error::Error>> {
24623 if portable_mma_gated() {
24625 return self.sdpa_naive(
24626 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
24627 );
24628 }
24629 static F32_STAGE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24635 let f32_stage =
24636 *F32_STAGE.get_or_init(|| std::env::var("MEMRA_FA512_STAGE").as_deref() == Ok("f32"));
24637 static SP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24641 let sp = !f32_stage
24642 && *SP_ON.get_or_init(|| {
24643 std::env::var("MEMRA_FA512_SP")
24644 .map(|v| v != "0")
24645 .unwrap_or(true)
24646 });
24647 self.fa_prefill_hd512_arm(
24648 q,
24649 k,
24650 v,
24651 o,
24652 head_dim,
24653 n_head,
24654 n_head_kv,
24655 t,
24656 t_kv,
24657 scale,
24658 causal,
24659 f32_stage,
24660 sp,
24661 sp && fa_f16pv_on(),
24662 )
24663 }
24664
24665 #[allow(clippy::too_many_arguments)]
24667 pub fn fa_prefill_hd512_pre(
24668 &self,
24669 qb: &CudaSlice<u8>,
24670 kb: &CudaSlice<u8>,
24671 vb: &CudaSlice<u8>,
24672 o: &mut CudaSlice<f32>,
24673 head_dim: usize,
24674 n_head: usize,
24675 n_head_kv: usize,
24676 t: usize,
24677 t_kv: usize,
24678 scale: f32,
24679 causal: bool,
24680 v_f16: bool,
24681 ) -> Result<(), Box<dyn std::error::Error>> {
24682 debug_assert_eq!(head_dim, 512);
24683 const SP_M: usize = 16;
24684 const BKS: usize = 32;
24685 let f16pv = fa_f16pv_on();
24689 let nw = if f16pv { fa512_wide_warps() } else { 2 };
24690 let hp = f16pv
24691 && fa512_hp_on()
24692 && n_head.is_multiple_of(2)
24693 && (n_head / n_head_kv).is_multiple_of(2);
24694 debug_assert!(!v_f16 || f16pv, "f16 V emitted without the door on");
24695 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
24696 let vref: &CudaSlice<u8> = if f16pv && !v_f16 {
24697 let n = t_kv * n_head_kv * head_dim;
24699 let need = n * 2;
24700 if vguard.as_ref().map(|b| b.len() < need).unwrap_or(true) {
24701 *vguard = Some(self.alloc_uninit::<u8>(need)?);
24702 }
24703 let dst = vguard.as_mut().unwrap();
24704 self.bf16_to_f16_into(vb, n, dst)?;
24705 vguard.as_ref().unwrap()
24706 } else {
24707 vb
24708 };
24709 let f = self.func(if hp {
24710 "fa_prefill_bf16_hd512_sp16h2"
24711 } else {
24712 match (f16pv, nw) {
24713 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
24714 (true, _) => "fa_prefill_bf16_hd512_sp16",
24715 _ => "fa_prefill_bf16_hd512_sp",
24716 }
24717 });
24718 let (nwarp, npart) = if hp {
24719 (4usize, 4usize)
24720 } else if nw > 2 {
24721 (nw, nw)
24722 } else {
24723 (2, 1)
24724 };
24725 let shmem = if hp {
24727 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS) + 4 * (2 * npart * SP_M * BKS + 2 * SP_M))
24728 as u32
24729 } else {
24730 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
24731 + 4 * (npart * SP_M * BKS + SP_M)) as u32
24732 };
24733 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24734 f.set_attribute(
24735 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24736 shmem as i32,
24737 )?;
24738 let grid_y = if hp {
24739 (n_head / 2) as u32
24740 } else {
24741 n_head as u32
24742 };
24743 let cfg = LaunchConfig {
24744 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
24745 block_dim: (32, nwarp as u32, 1),
24746 shared_mem_bytes: shmem,
24747 };
24748 let (hd, nh, nhkv, ti, tkvi, cz) = (
24749 head_dim as i32,
24750 n_head as i32,
24751 n_head_kv as i32,
24752 t as i32,
24753 t_kv as i32,
24754 causal as i32,
24755 );
24756 let __s_b = self.gpu.stream();
24757 let mut b = __s_b.launch_builder(&f);
24758 b.arg(qb)
24759 .arg(kb)
24760 .arg(vref)
24761 .arg(o)
24762 .arg(&hd)
24763 .arg(&nh)
24764 .arg(&nhkv)
24765 .arg(&ti)
24766 .arg(&tkvi)
24767 .arg(&scale)
24768 .arg(&cz);
24769 unsafe {
24770 b.launch(cfg)?;
24771 }
24772 Ok(())
24773 }
24774
24775 #[allow(clippy::too_many_arguments)]
24782 pub fn mla_attn_gathered_tc(
24783 &self,
24784 q_lat_bf: &CudaSlice<u8>, cache_bf: &CudaSlice<u8>, idx: &CudaSlice<i32>, o_lat: &mut CudaSlice<f32>, n_head: usize,
24789 kv_rank: usize,
24790 t_q: usize,
24791 width: usize,
24792 scale: f32,
24793 ) -> Result<(), Box<dyn std::error::Error>> {
24794 if kv_rank != 512 {
24795 return Err(format!(
24796 "mla_attn_gathered_tc is stamped at kv_rank 512 (the glm5_next latent width); \
24797 got {kv_rank} — the caller's door must fall back to the f32 gathered kernel"
24798 )
24799 .into());
24800 }
24801 if t_q == 0 || n_head == 0 {
24802 return Ok(());
24803 }
24804 const SP_M: usize = 16;
24805 const BKS: usize = 32;
24806 const HD: usize = 512;
24807 let f = self.func("fa_mla_gathered_bf16");
24808 let shmem =
24810 (2 * (SP_M * HD + BKS * HD + SP_M * BKS) + 4 * (SP_M * BKS + SP_M) + 4 * BKS) as u32;
24811 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24812 f.set_attribute(
24813 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24814 shmem as i32,
24815 )?;
24816 let cfg = LaunchConfig {
24817 grid_dim: (t_q as u32, (n_head as u32).div_ceil(SP_M as u32), 1),
24818 block_dim: (32, 2, 1),
24819 shared_mem_bytes: shmem,
24820 };
24821 let (nh, tq, w) = (n_head as i32, t_q as i32, width as i32);
24822 let __s_b = self.gpu.stream();
24823 let mut b = __s_b.launch_builder(&f);
24824 b.arg(q_lat_bf)
24825 .arg(cache_bf)
24826 .arg(idx)
24827 .arg(o_lat)
24828 .arg(&nh)
24829 .arg(&tq)
24830 .arg(&w)
24831 .arg(&scale);
24832 unsafe {
24833 b.launch(cfg)?;
24834 }
24835 Ok(())
24836 }
24837
24838 #[allow(clippy::too_many_arguments)]
24841 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_hd512_arm(
24843 &self,
24844 q: &CudaSlice<f32>,
24845 k: &CudaSlice<f32>,
24846 v: &CudaSlice<f32>,
24847 o: &mut CudaSlice<f32>,
24848 head_dim: usize,
24849 n_head: usize,
24850 n_head_kv: usize,
24851 t: usize,
24852 t_kv: usize,
24853 scale: f32,
24854 causal: bool,
24855 f32_stage: bool,
24856 sp: bool,
24857 f16pv: bool,
24858 ) -> Result<(), Box<dyn std::error::Error>> {
24859 debug_assert_eq!(head_dim, 512, "fa_prefill_hd512 is hd512 only");
24860 if sp && !f32_stage {
24861 const SP_M: usize = 16;
24865 const BKS: usize = 32;
24866 let nw = if f16pv { fa512_wide_warps() } else { 2 };
24867 let hp = f16pv
24868 && fa512_hp_on()
24869 && n_head.is_multiple_of(2)
24870 && (n_head / n_head_kv).is_multiple_of(2);
24871 let f = self.func(if hp {
24872 "fa_prefill_bf16_hd512_sp16h2"
24873 } else {
24874 match (f16pv, nw) {
24875 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
24876 (true, _) => "fa_prefill_bf16_hd512_sp16",
24877 _ => "fa_prefill_bf16_hd512_sp",
24878 }
24879 });
24880 let (nwarp, npart) = if hp {
24881 (4usize, 4usize)
24882 } else if nw > 2 {
24883 (nw, nw)
24884 } else {
24885 (2, 1)
24886 };
24887 let shmem = if hp {
24888 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
24889 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
24890 } else {
24891 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
24892 + 4 * (npart * SP_M * BKS + SP_M)) as u32
24893 };
24894 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24895 f.set_attribute(
24896 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24897 shmem as i32,
24898 )?;
24899 let grid_y = if hp {
24900 (n_head / 2) as u32
24901 } else {
24902 n_head as u32
24903 };
24904 let cfg = LaunchConfig {
24905 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
24906 block_dim: (32, nwarp as u32, 1),
24907 shared_mem_bytes: shmem,
24908 };
24909 let (hd, nh, nhkv, ti, tkvi, cz) = (
24910 head_dim as i32,
24911 n_head as i32,
24912 n_head_kv as i32,
24913 t as i32,
24914 t_kv as i32,
24915 causal as i32,
24916 );
24917 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24918 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24919 let vb = if f16pv {
24920 self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?
24921 } else {
24922 self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?
24923 };
24924 let __s_b = self.gpu.stream();
24925 let mut b = __s_b.launch_builder(&f);
24926 b.arg(&qb)
24927 .arg(&kb)
24928 .arg(&vb)
24929 .arg(o)
24930 .arg(&hd)
24931 .arg(&nh)
24932 .arg(&nhkv)
24933 .arg(&ti)
24934 .arg(&tkvi)
24935 .arg(&scale)
24936 .arg(&cz);
24937 unsafe {
24938 b.launch(cfg)?;
24939 }
24940 return Ok(());
24941 }
24942 const BLOCK_Q: usize = 32;
24943 const BK: usize = 32;
24944 const HALF: usize = 256;
24945 let f = self.func(if f32_stage {
24946 "fa_prefill_f32_hd512"
24947 } else {
24948 "fa_prefill_bf16_hd512"
24949 });
24950 let shmem = (2 * (BLOCK_Q * head_dim + BK * head_dim + BK * HALF + BLOCK_Q * BK)
24952 + 4 * BLOCK_Q) as u32;
24953 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24954 f.set_attribute(
24955 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24956 shmem as i32,
24957 )?;
24958 let cfg = LaunchConfig {
24959 grid_dim: (
24960 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
24961 n_head as u32,
24962 2,
24963 ),
24964 block_dim: (32, 2, 1),
24965 shared_mem_bytes: shmem,
24966 };
24967 let (hd, nh, nhkv, ti, tkvi, cz) = (
24968 head_dim as i32,
24969 n_head as i32,
24970 n_head_kv as i32,
24971 t as i32,
24972 t_kv as i32,
24973 causal as i32,
24974 );
24975 if f32_stage {
24976 let __s_b = self.gpu.stream();
24977 let mut b = __s_b.launch_builder(&f);
24978 b.arg(q)
24979 .arg(k)
24980 .arg(v)
24981 .arg(o)
24982 .arg(&hd)
24983 .arg(&nh)
24984 .arg(&nhkv)
24985 .arg(&ti)
24986 .arg(&tkvi)
24987 .arg(&scale)
24988 .arg(&cz);
24989 unsafe {
24990 b.launch(cfg)?;
24991 }
24992 } else {
24993 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24994 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24995 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
24996 let __s_b = self.gpu.stream();
24997 let mut b = __s_b.launch_builder(&f);
24998 b.arg(&qb)
24999 .arg(&kb)
25000 .arg(&vb)
25001 .arg(o)
25002 .arg(&hd)
25003 .arg(&nh)
25004 .arg(&nhkv)
25005 .arg(&ti)
25006 .arg(&tkvi)
25007 .arg(&scale)
25008 .arg(&cz);
25009 unsafe {
25010 b.launch(cfg)?;
25011 }
25012 }
25013 Ok(())
25014 }
25015
25016 #[allow(clippy::too_many_arguments)]
25020 pub fn rope_neox2_bf16e(
25021 &self,
25022 q: &mut CudaSlice<f32>,
25023 k: &mut CudaSlice<f32>,
25024 qb: &mut CudaSlice<u8>,
25025 kb: &mut CudaSlice<u8>,
25026 pos: &CudaSlice<i32>,
25027 head_dim: usize,
25028 n_dims: usize,
25029 nh_q: usize,
25030 nh_k: usize,
25031 n_tokens: usize,
25032 base: f32,
25033 freq_scale: f32,
25034 ff: Option<&CudaSlice<f32>>,
25035 ) -> Result<(), Box<dyn std::error::Error>> {
25036 let f = self.func("rope_neox2_bf16e_f32");
25037 let rows = ((nh_q + nh_k) * n_tokens) as u32;
25038 let cfg = LaunchConfig {
25039 grid_dim: (rows, 1, 1),
25040 block_dim: ((head_dim / 2) as u32, 1, 1),
25041 shared_mem_bytes: 0,
25042 };
25043 let theta_scale = base.powf(-2.0 / n_dims as f32);
25044 let (hd, nd, nhq, nhk, nt) = (
25045 head_dim as i32,
25046 n_dims as i32,
25047 nh_q as i32,
25048 nh_k as i32,
25049 n_tokens as i32,
25050 );
25051 let __s_b = self.gpu.stream();
25052 let mut b = __s_b.launch_builder(&f);
25053 match ff {
25054 Some(t) => {
25055 b.arg(&mut *q)
25056 .arg(&mut *k)
25057 .arg(&mut *qb)
25058 .arg(&mut *kb)
25059 .arg(pos)
25060 .arg(&hd)
25061 .arg(&nd)
25062 .arg(&nhq)
25063 .arg(&nhk)
25064 .arg(&nt)
25065 .arg(&theta_scale)
25066 .arg(&freq_scale)
25067 .arg(t);
25068 unsafe {
25069 b.launch(cfg)?;
25070 }
25071 }
25072 None => {
25073 let null: u64 = 0;
25074 b.arg(&mut *q)
25075 .arg(&mut *k)
25076 .arg(&mut *qb)
25077 .arg(&mut *kb)
25078 .arg(pos)
25079 .arg(&hd)
25080 .arg(&nd)
25081 .arg(&nhq)
25082 .arg(&nhk)
25083 .arg(&nt)
25084 .arg(&theta_scale)
25085 .arg(&freq_scale)
25086 .arg(&null);
25087 unsafe {
25088 b.launch(cfg)?;
25089 }
25090 }
25091 }
25092 Ok(())
25093 }
25094
25095 pub fn f32_to_bf16(
25098 &self,
25099 x: &CudaSlice<f32>,
25100 n: usize,
25101 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
25102 assert!(
25103 n.is_multiple_of(4),
25104 "f32_to_bf16 requires n % 4 == 0, got {n}"
25105 );
25106 let mut y = self.alloc_uninit::<u8>(n * 2)?;
25107 let f = self.func("f32_to_bf16_flat");
25108 let n_i = n as i64;
25109 let cfg = LaunchConfig {
25110 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
25111 block_dim: (256, 1, 1),
25112 shared_mem_bytes: 0,
25113 };
25114 let __s_b = self.gpu.stream();
25115 let mut b = __s_b.launch_builder(&f);
25116 b.arg(x).arg(&mut y).arg(&n_i);
25117 unsafe {
25118 b.launch(cfg)?;
25119 }
25120 Ok(y)
25121 }
25122
25123 pub fn f32_to_f16(
25124 &self,
25125 x: &CudaSlice<f32>,
25126 n: usize,
25127 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
25128 assert!(
25129 n.is_multiple_of(4),
25130 "f32_to_f16 requires n % 4 == 0, got {n}"
25131 );
25132 let mut y = self.alloc_uninit::<u8>(n * 2)?;
25133 let f = self.func("f32_to_f16_flat");
25134 let n_i = n as i64;
25135 let cfg = LaunchConfig {
25136 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
25137 block_dim: (256, 1, 1),
25138 shared_mem_bytes: 0,
25139 };
25140 let __s_b = self.gpu.stream();
25141 let mut b = __s_b.launch_builder(&f);
25142 b.arg(x).arg(&mut y).arg(&n_i);
25143 unsafe {
25144 b.launch(cfg)?;
25145 }
25146 Ok(y)
25147 }
25148
25149 pub fn bf16_to_f16(
25151 &self,
25152 xb: &CudaSlice<u8>,
25153 n: usize,
25154 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
25155 let mut y = self.alloc_uninit::<u8>(n * 2)?;
25156 self.bf16_to_f16_into(xb, n, &mut y)?;
25157 Ok(y)
25158 }
25159
25160 pub fn bf16_to_f16_into(
25162 &self,
25163 xb: &CudaSlice<u8>,
25164 n: usize,
25165 y: &mut CudaSlice<u8>,
25166 ) -> Result<(), Box<dyn std::error::Error>> {
25167 assert!(
25168 n.is_multiple_of(2),
25169 "bf16_to_f16 requires n % 2 == 0, got {n}"
25170 );
25171 assert!(y.len() >= n * 2);
25172 let f = self.func("bf16_to_f16_flat");
25173 let n2 = (n / 2) as i64;
25174 let cfg = LaunchConfig {
25175 grid_dim: (((n / 2) as u32).div_ceil(256), 1, 1),
25176 block_dim: (256, 1, 1),
25177 shared_mem_bytes: 0,
25178 };
25179 let __s_b = self.gpu.stream();
25180 let mut b = __s_b.launch_builder(&f);
25181 b.arg(xb).arg(y).arg(&n2);
25182 unsafe {
25183 b.launch(cfg)?;
25184 }
25185 Ok(())
25186 }
25187
25188 #[allow(clippy::too_many_arguments)]
25193 pub fn fa_prefill_vl8(
25194 &self,
25195 seqs: &[FaSeqVl],
25196 head_dim: usize,
25197 n_head: usize,
25198 n_head_kv: usize,
25199 scale: f32,
25200 ) -> Result<(), Box<dyn std::error::Error>> {
25201 const BK: usize = 32;
25202 let b = seqs.len();
25203 assert!((1..=8).contains(&b));
25204 let mut packed = [FaSeqVl::default(); 8];
25205 packed[..b].copy_from_slice(seqs);
25206 let v = FaVl8(packed);
25207 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
25208 let ept = (n_head_kv * head_dim) as i32;
25209 {
25210 let f = self.func("fa_mirror_vl");
25211 let max_n = (max_t as i64) * ept as i64;
25212 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
25213 for which in 0..2i32 {
25214 let cfg = LaunchConfig {
25215 grid_dim: (blocks, 1, b as u32),
25216 block_dim: (256, 1, 1),
25217 shared_mem_bytes: 0,
25218 };
25219 let __s_lb = self.gpu.stream();
25220 let mut lb = __s_lb.launch_builder(&f);
25221 lb.arg(&v).arg(&ept).arg(&which);
25222 unsafe {
25223 lb.launch(cfg)?;
25224 }
25225 }
25226 }
25227 let hd_sfx = fa_hd_suffix(head_dim)?;
25228 let f = self.func(&format!("fa_prefill_bf16kv_vl{hd_sfx}"));
25229 let block_q = 64usize;
25230 let kv_stages = 2usize;
25231 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
25232 + 4 * (block_q * BK + 2 * block_q)) as u32;
25233 use cudarc::driver::sys::CUfunction_attribute_enum as A;
25234 f.set_attribute(
25235 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
25236 shmem as i32,
25237 )?;
25238 let cfg = LaunchConfig {
25239 grid_dim: (max_t.div_ceil(block_q as u32), n_head as u32, b as u32),
25240 block_dim: (32, 4, 1),
25241 shared_mem_bytes: shmem,
25242 };
25243 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
25244 let __s_lb = self.gpu.stream();
25245 let mut lb = __s_lb.launch_builder(&f);
25246 lb.arg(&v).arg(&hd).arg(&nh).arg(&nhkv).arg(&scale);
25247 unsafe {
25248 lb.launch(cfg)?;
25249 }
25250 Ok(())
25251 }
25252
25253 #[allow(clippy::too_many_arguments)]
25257 pub fn attn_pre_vl8(
25258 &self,
25259 seqs: &[AttnPreVl],
25260 wq: &CudaSlice<f32>,
25261 wk: &CudaSlice<f32>,
25262 head_dim: usize,
25263 rope_dims: usize,
25264 n_head: usize,
25265 n_head_kv: usize,
25266 eps: f32,
25267 freq_base: f32,
25268 freq_scale: f32,
25269 kv_dim_k: usize,
25270 kv_dim_v: usize,
25271 k_tok_bytes: usize,
25272 v_tok_bytes: usize,
25273 ) -> Result<(), Box<dyn std::error::Error>> {
25274 let b = seqs.len();
25275 assert!((1..=8).contains(&b));
25276 let mut packed = [AttnPreVl::default(); 8];
25277 packed[..b].copy_from_slice(seqs);
25278 let v = AttnPreVl8(packed);
25279 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
25280 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
25281 {
25282 let f = self.func("q_gate_split_vl");
25283 let n = max_t * (n_head * head_dim) as u32;
25284 let cfg = LaunchConfig {
25285 grid_dim: (n.div_ceil(256), 1, b as u32),
25286 block_dim: (256, 1, 1),
25287 shared_mem_bytes: 0,
25288 };
25289 let __s_lb = self.gpu.stream();
25290 let mut lb = __s_lb.launch_builder(&f);
25291 lb.arg(&v).arg(&hd).arg(&nh);
25292 unsafe {
25293 lb.launch(cfg)?;
25294 }
25295 }
25296 {
25297 let f = self.func("attn_rms_vl");
25298 let cfg = LaunchConfig {
25299 grid_dim: (max_t * n_head as u32, 2, b as u32),
25300 block_dim: (rms_block(), 1, 1),
25301 shared_mem_bytes: 0,
25302 };
25303 let __s_lb = self.gpu.stream();
25304 let mut lb = __s_lb.launch_builder(&f);
25305 lb.arg(&v)
25306 .arg(wq)
25307 .arg(wk)
25308 .arg(&hd)
25309 .arg(&nh)
25310 .arg(&nhkv)
25311 .arg(&eps);
25312 unsafe {
25313 lb.launch(cfg)?;
25314 }
25315 }
25316 {
25317 let f = self.func("attn_rope_vl");
25318 let theta_scale = freq_base.powf(-2.0 / rope_dims as f32);
25319 let nd = rope_dims as i32;
25320 let cfg = LaunchConfig {
25321 grid_dim: (max_t * n_head as u32, 2, b as u32),
25322 block_dim: ((head_dim / 2) as u32, 1, 1),
25323 shared_mem_bytes: 0,
25324 };
25325 let __s_lb = self.gpu.stream();
25326 let mut lb = __s_lb.launch_builder(&f);
25327 lb.arg(&v)
25328 .arg(&hd)
25329 .arg(&nd)
25330 .arg(&nh)
25331 .arg(&nhkv)
25332 .arg(&theta_scale)
25333 .arg(&freq_scale);
25334 unsafe {
25335 lb.launch(cfg)?;
25336 }
25337 }
25338 {
25339 let f = self.func("append_kv_vl");
25340 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
25341 let cfg = LaunchConfig {
25342 grid_dim: (nblk, max_t, b as u32),
25343 block_dim: (32, 1, 1),
25344 shared_mem_bytes: 0,
25345 };
25346 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
25347 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
25348 let __s_lb = self.gpu.stream();
25349 let mut lb = __s_lb.launch_builder(&f);
25350 lb.arg(&v).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
25351 unsafe {
25352 lb.launch(cfg)?;
25353 }
25354 }
25355 Ok(())
25356 }
25357
25358 #[allow(clippy::too_many_arguments)]
25363 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_view(
25366 &self,
25367 q: &CudaSlice<f32>,
25368 k: &cudarc::driver::CudaView<u8>,
25369 v: &cudarc::driver::CudaView<u8>,
25370 o: &mut CudaSlice<f32>,
25371 head_dim: usize,
25372 n_head: usize,
25373 n_head_kv: usize,
25374 t: usize,
25375 t_kv: usize,
25376 scale: f32,
25377 causal: bool,
25378 k_tok_bytes: usize,
25379 v_tok_bytes: usize,
25380 g: bool,
25381 ) -> Result<(), Box<dyn std::error::Error>> {
25382 if portable_mma_gated() {
25383 return self.sdpa_naive_quantized_view(
25384 q,
25385 k,
25386 v,
25387 o,
25388 head_dim,
25389 n_head,
25390 n_head_kv,
25391 t,
25392 t_kv,
25393 scale,
25394 causal,
25395 k_tok_bytes,
25396 v_tok_bytes,
25397 );
25398 }
25399 const BLOCK_Q: usize = 64;
25400 const BK: usize = 32;
25401 let name = format!("fa_prefill_q{}", fa_hd_suffix(head_dim)?);
25404 let f = if g {
25405 self.func_g(&name)
25406 } else {
25407 self.func(&name)
25408 };
25409 let shmem =
25410 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
25411 use cudarc::driver::sys::CUfunction_attribute_enum as A;
25412 f.set_attribute(
25413 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
25414 shmem as i32,
25415 )?;
25416 let cfg = LaunchConfig {
25417 grid_dim: (
25418 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
25419 n_head as u32,
25420 1,
25421 ),
25422 block_dim: (32, 4, 1),
25423 shared_mem_bytes: shmem,
25424 };
25425 let (hd, nh, nhkv, ti, tkvi, cz) = (
25426 head_dim as i32,
25427 n_head as i32,
25428 n_head_kv as i32,
25429 t as i32,
25430 t_kv as i32,
25431 causal as i32,
25432 );
25433 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
25434 let __s_b = self.gpu.stream();
25435 let mut b = __s_b.launch_builder(&f);
25436 b.arg(q)
25437 .arg(k)
25438 .arg(v)
25439 .arg(o)
25440 .arg(&hd)
25441 .arg(&nh)
25442 .arg(&nhkv)
25443 .arg(&ti)
25444 .arg(&tkvi)
25445 .arg(&scale)
25446 .arg(&cz)
25447 .arg(&ktb)
25448 .arg(&vtb);
25449 unsafe {
25450 b.launch(cfg)?;
25451 }
25452 Ok(())
25453 }
25454
25455 #[allow(clippy::too_many_arguments)]
25465 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_view_ws(
25467 &self,
25468 q: &CudaSlice<f32>,
25469 k: &cudarc::driver::CudaView<u8>,
25470 v: &cudarc::driver::CudaView<u8>,
25471 o: &mut CudaSlice<f32>,
25472 head_dim: usize,
25473 n_head: usize,
25474 n_head_kv: usize,
25475 t: usize,
25476 t_kv: usize,
25477 scale: f32,
25478 causal: bool,
25479 k_tok_bytes: usize,
25480 v_tok_bytes: usize,
25481 g: bool,
25482 ) -> Result<(), Box<dyn std::error::Error>> {
25483 if portable_mma_gated() {
25484 return self.sdpa_naive_quantized_view(
25485 q,
25486 k,
25487 v,
25488 o,
25489 head_dim,
25490 n_head,
25491 n_head_kv,
25492 t,
25493 t_kv,
25494 scale,
25495 causal,
25496 k_tok_bytes,
25497 v_tok_bytes,
25498 );
25499 }
25500 const BLOCK_Q: usize = 64;
25501 const BK: usize = 32;
25502 let kv_dim_k = n_head_kv * head_dim;
25503 let kv_dim_v = n_head_kv * head_dim;
25504 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
25506 let mut guard = self.prime_deqw_ws.lock().unwrap();
25508 let need_grow = match guard.as_ref() {
25509 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
25510 None => true,
25511 };
25512 if need_grow {
25513 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
25514 let (ck, cv) = guard
25515 .as_ref()
25516 .map(|(a, b)| (a.len(), b.len()))
25517 .unwrap_or((0, 0));
25518 *guard = Some((
25519 self.alloc_u8(grow(ck, k_ws_bytes))?,
25520 self.alloc_u8(grow(cv, v_ws_bytes))?,
25521 ));
25522 }
25523 let (kw, vw) = guard.as_mut().unwrap();
25524 {
25526 let f = if g {
25528 self.func_g("fa_dequant_kv_ws_bf16")
25529 } else {
25530 self.func("fa_dequant_kv_ws_bf16")
25531 };
25532 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
25533 #[allow(clippy::manual_div_ceil)]
25534 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
25536 let cfg = LaunchConfig {
25537 grid_dim: (nblk.max(1), 1, 1),
25538 block_dim: (256, 1, 1),
25539 shared_mem_bytes: 0,
25540 };
25541 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
25542 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
25543 let __s_b = self.gpu.stream();
25544 let mut b = __s_b.launch_builder(&f);
25545 b.arg(k)
25546 .arg(v)
25547 .arg(&mut *kw)
25548 .arg(&mut *vw)
25549 .arg(&kdk)
25550 .arg(&kdv)
25551 .arg(&tkvi)
25552 .arg(&ktb)
25553 .arg(&vtb);
25554 unsafe {
25555 b.launch(cfg)?;
25556 }
25557 }
25558 let db = std::env::var("MEMRA_PRIME_DEQW_DB")
25566 .map(|v| v != "0")
25567 .unwrap_or(true);
25568 {
25569 let hd_sfx = fa_hd_suffix(head_dim)?;
25570 let f = self.func(&format!(
25571 "fa_prefill_qw{}{hd_sfx}",
25572 if db { "_db" } else { "" }
25573 ));
25574 let shmem = if db {
25575 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
25577 } else {
25578 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
25579 };
25580 use cudarc::driver::sys::CUfunction_attribute_enum as A;
25581 f.set_attribute(
25582 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
25583 shmem as i32,
25584 )?;
25585 let cfg = LaunchConfig {
25586 grid_dim: (
25587 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
25588 n_head as u32,
25589 1,
25590 ),
25591 block_dim: (32, 4, 1),
25592 shared_mem_bytes: shmem,
25593 };
25594 let (hd, nh, nhkv, ti, tkvi, cz) = (
25595 head_dim as i32,
25596 n_head as i32,
25597 n_head_kv as i32,
25598 t as i32,
25599 t_kv as i32,
25600 causal as i32,
25601 );
25602 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
25603 let __s_b = self.gpu.stream();
25604 let mut b = __s_b.launch_builder(&f);
25605 b.arg(q)
25606 .arg(&*kw)
25607 .arg(&*vw)
25608 .arg(o)
25609 .arg(&hd)
25610 .arg(&nh)
25611 .arg(&nhkv)
25612 .arg(&ti)
25613 .arg(&tkvi)
25614 .arg(&scale)
25615 .arg(&cz)
25616 .arg(&kdk)
25617 .arg(&kdv);
25618 unsafe {
25619 b.launch(cfg)?;
25620 }
25621 }
25622 Ok(())
25623 }
25624
25625 #[allow(clippy::too_many_arguments)]
25641 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_view_ws_w_hd128(
25643 &self,
25644 q: &CudaSlice<f32>,
25645 k: &cudarc::driver::CudaView<u8>,
25646 v: &cudarc::driver::CudaView<u8>,
25647 o: &mut CudaSlice<f32>,
25648 head_dim: usize,
25649 n_head: usize,
25650 n_head_kv: usize,
25651 t: usize,
25652 t_kv: usize,
25653 scale: f32,
25654 causal: bool,
25655 window: usize,
25656 k_tok_bytes: usize,
25657 v_tok_bytes: usize,
25658 ) -> Result<(), Box<dyn std::error::Error>> {
25659 assert_eq!(
25660 head_dim, 128,
25661 "fa_prefill_view_ws_w_hd128: only the hd128 twin is stamped"
25662 );
25663 if portable_mma_gated() {
25664 return self.sdpa_naive_w_quantized_view(
25665 q,
25666 k,
25667 v,
25668 o,
25669 head_dim,
25670 n_head,
25671 n_head_kv,
25672 t,
25673 t_kv,
25674 scale,
25675 causal,
25676 window,
25677 k_tok_bytes,
25678 v_tok_bytes,
25679 );
25680 }
25681 const BLOCK_Q: usize = 64;
25682 const BK: usize = 32;
25683 let kv_dim_k = n_head_kv * head_dim;
25684 let kv_dim_v = n_head_kv * head_dim;
25685 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
25687 let mut guard = self.prime_deqw_ws.lock().unwrap();
25688 let need_grow = match guard.as_ref() {
25689 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
25690 None => true,
25691 };
25692 if need_grow {
25693 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
25694 let (ck, cv) = guard
25695 .as_ref()
25696 .map(|(a, b)| (a.len(), b.len()))
25697 .unwrap_or((0, 0));
25698 *guard = Some((
25699 self.alloc_u8(grow(ck, k_ws_bytes))?,
25700 self.alloc_u8(grow(cv, v_ws_bytes))?,
25701 ));
25702 }
25703 let (kw, vw) = guard.as_mut().unwrap();
25704 {
25707 let f = self.func("fa_dequant_kv_ws_bf16");
25708 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
25709 #[allow(clippy::manual_div_ceil)]
25710 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
25712 let cfg = LaunchConfig {
25713 grid_dim: (nblk.max(1), 1, 1),
25714 block_dim: (256, 1, 1),
25715 shared_mem_bytes: 0,
25716 };
25717 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
25718 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
25719 let __s_b = self.gpu.stream();
25720 let mut b = __s_b.launch_builder(&f);
25721 b.arg(k)
25722 .arg(v)
25723 .arg(&mut *kw)
25724 .arg(&mut *vw)
25725 .arg(&kdk)
25726 .arg(&kdv)
25727 .arg(&tkvi)
25728 .arg(&ktb)
25729 .arg(&vtb);
25730 unsafe {
25731 b.launch(cfg)?;
25732 }
25733 }
25734 let db = std::env::var("MEMRA_PRIME_DEQW_DB")
25736 .map(|v| v != "0")
25737 .unwrap_or(true);
25738 {
25739 let f = self.func(if db {
25740 "fa_prefill_qw_db_w_hd128"
25741 } else {
25742 "fa_prefill_qw_w_hd128"
25743 });
25744 let shmem = if db {
25745 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
25746 } else {
25747 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
25748 };
25749 use cudarc::driver::sys::CUfunction_attribute_enum as A;
25750 f.set_attribute(
25751 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
25752 shmem as i32,
25753 )?;
25754 let cfg = LaunchConfig {
25755 grid_dim: (
25756 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
25757 n_head as u32,
25758 1,
25759 ),
25760 block_dim: (32, 4, 1),
25761 shared_mem_bytes: shmem,
25762 };
25763 let (hd, nh, nhkv, ti, tkvi, cz) = (
25764 head_dim as i32,
25765 n_head as i32,
25766 n_head_kv as i32,
25767 t as i32,
25768 t_kv as i32,
25769 causal as i32,
25770 );
25771 let (kdk, kdv, wnd) = (kv_dim_k as i32, kv_dim_v as i32, window as i32);
25772 let __s_b = self.gpu.stream();
25773 let mut b = __s_b.launch_builder(&f);
25774 b.arg(q)
25775 .arg(&*kw)
25776 .arg(&*vw)
25777 .arg(o)
25778 .arg(&hd)
25779 .arg(&nh)
25780 .arg(&nhkv)
25781 .arg(&ti)
25782 .arg(&tkvi)
25783 .arg(&scale)
25784 .arg(&cz)
25785 .arg(&kdk)
25786 .arg(&kdv)
25787 .arg(&wnd);
25788 unsafe {
25789 b.launch(cfg)?;
25790 }
25791 }
25792 Ok(())
25793 }
25794
25795 #[allow(clippy::too_many_arguments)] pub fn fa_decode(
25800 &self,
25801 q: &CudaSlice<f32>,
25802 k: &cudarc::driver::CudaView<u8>,
25803 v: &cudarc::driver::CudaView<u8>,
25804 o: &mut CudaSlice<f32>,
25805 head_dim: usize,
25806 n_head: usize,
25807 n_head_kv: usize,
25808 t_kv: usize,
25809 scale: f32,
25810 k_tok_bytes: usize,
25811 v_tok_bytes: usize,
25812 ) -> Result<(), Box<dyn std::error::Error>> {
25813 self.fa_decode_kvmod(
25814 q,
25815 k,
25816 v,
25817 o,
25818 head_dim,
25819 n_head,
25820 n_head_kv,
25821 t_kv,
25822 scale,
25823 k_tok_bytes,
25824 v_tok_bytes,
25825 false,
25826 )
25827 }
25828
25829 #[allow(clippy::too_many_arguments)]
25833 #[allow(clippy::too_many_arguments)]
25837 #[allow(clippy::too_many_arguments)]
25838 fn fa_decode_scalar_unified(
25839 &self,
25840 q: &cudarc::driver::CudaView<f32>,
25841 k: &cudarc::driver::CudaView<u8>,
25842 v: &cudarc::driver::CudaView<u8>,
25843 o: &mut cudarc::driver::CudaViewMut<f32>,
25844 head_dim: usize,
25845 n_head: usize,
25846 n_head_kv: usize,
25847 t_kv_host: usize,
25848 t_kv_dev: Option<&CudaSlice<i32>>,
25849 scale: f32,
25850 n_splits: usize,
25851 split_keys: usize,
25852 k_tok_bytes: usize,
25853 v_tok_bytes: usize,
25854 g: bool,
25855 part_o: &mut CudaSlice<f32>,
25856 part_m: &mut CudaSlice<f32>,
25857 part_l: &mut CudaSlice<f32>,
25858 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
25859 ) -> Result<(), Box<dyn std::error::Error>> {
25860 let f = if g {
25861 self.func_g("fa_decode_f32")
25862 } else {
25863 self.fa_func("fa_decode_f32", head_dim)
25864 };
25865 let cfg = LaunchConfig {
25866 grid_dim: (n_head as u32, n_splits as u32, 1),
25867 block_dim: (head_dim as u32, 1, 1),
25868 shared_mem_bytes: (4 * (head_dim + 32)) as u32,
25869 };
25870 let (hd, nh, nhkv, nsp) = (
25871 head_dim as i32,
25872 n_head as i32,
25873 n_head_kv as i32,
25874 n_splits as i32,
25875 );
25876 let (ktb, vtb, tkvi, ski) = (
25877 k_tok_bytes as i64,
25878 v_tok_bytes as i64,
25879 t_kv_host as i32,
25880 split_keys as i32,
25881 );
25882 let __s_b = self.gpu.stream();
25883 let mut b = __s_b.launch_builder(&f);
25884 match t_kv_dev {
25885 Some(d) => {
25886 b.arg(q)
25887 .arg(k)
25888 .arg(v)
25889 .arg(&mut *part_o)
25890 .arg(&mut *part_m)
25891 .arg(&mut *part_l)
25892 .arg(&hd)
25893 .arg(&nh)
25894 .arg(&nhkv)
25895 .arg(&tkvi)
25896 .arg(d)
25897 .arg(&scale)
25898 .arg(&nsp)
25899 .arg(&ski)
25900 .arg(&ktb)
25901 .arg(&vtb);
25902 unsafe {
25903 b.launch(cfg)?;
25904 }
25905 }
25906 None => {
25907 let null: u64 = 0;
25908 b.arg(q)
25909 .arg(k)
25910 .arg(v)
25911 .arg(&mut *part_o)
25912 .arg(&mut *part_m)
25913 .arg(&mut *part_l)
25914 .arg(&hd)
25915 .arg(&nh)
25916 .arg(&nhkv)
25917 .arg(&tkvi)
25918 .arg(&null)
25919 .arg(&scale)
25920 .arg(&nsp)
25921 .arg(&ski)
25922 .arg(&ktb)
25923 .arg(&vtb);
25924 unsafe {
25925 b.launch(cfg)?;
25926 }
25927 }
25928 }
25929 let cfg2 = LaunchConfig {
25930 grid_dim: (n_head as u32, 1, 1),
25931 block_dim: (head_dim as u32, 1, 1),
25932 shared_mem_bytes: 0,
25933 };
25934 if let Some((oq, od)) = q8_out {
25935 let fc = if g {
25937 self.func_g("fa_decode_combine_q8_1")
25938 } else {
25939 self.fa_func("fa_decode_combine_q8_1", head_dim)
25940 };
25941 let __s_b2 = self.gpu.stream();
25942 let mut b2 = __s_b2.launch_builder(&fc);
25943 b2.arg(&*part_o)
25944 .arg(&*part_m)
25945 .arg(&*part_l)
25946 .arg(oq)
25947 .arg(od)
25948 .arg(&hd)
25949 .arg(&nh)
25950 .arg(&nsp);
25951 unsafe {
25952 b2.launch(cfg2)?;
25953 }
25954 return Ok(());
25955 }
25956 let fc = if g {
25957 self.func_g("fa_decode_combine_f32")
25958 } else {
25959 self.fa_func("fa_decode_combine_f32", head_dim)
25960 };
25961 let __s_b2 = self.gpu.stream();
25962 let mut b2 = __s_b2.launch_builder(&fc);
25963 b2.arg(&*part_o)
25964 .arg(&*part_m)
25965 .arg(&*part_l)
25966 .arg(o)
25967 .arg(&hd)
25968 .arg(&nh)
25969 .arg(&nsp);
25970 unsafe {
25971 b2.launch(cfg2)?;
25972 }
25973 Ok(())
25974 }
25975
25976 #[allow(clippy::too_many_arguments)] pub fn fa_decode_kvmod(
25978 &self,
25979 q: &CudaSlice<f32>,
25980 k: &cudarc::driver::CudaView<u8>,
25981 v: &cudarc::driver::CudaView<u8>,
25982 o: &mut CudaSlice<f32>,
25983 head_dim: usize,
25984 n_head: usize,
25985 n_head_kv: usize,
25986 t_kv: usize,
25987 scale: f32,
25988 k_tok_bytes: usize,
25989 v_tok_bytes: usize,
25990 g: bool,
25991 ) -> Result<(), Box<dyn std::error::Error>> {
25992 let q_view = q.as_view();
25993 let mut o_view = o.as_view_mut();
25994 self.fa_decode_kvmod_view(
25995 &q_view,
25996 k,
25997 v,
25998 &mut o_view,
25999 head_dim,
26000 n_head,
26001 n_head_kv,
26002 t_kv,
26003 scale,
26004 k_tok_bytes,
26005 v_tok_bytes,
26006 g,
26007 )
26008 }
26009
26010 #[allow(clippy::too_many_arguments)]
26015 #[allow(clippy::manual_div_ceil)] pub fn fa_decode_kvmod_view(
26017 &self,
26018 q: &cudarc::driver::CudaView<f32>,
26019 k: &cudarc::driver::CudaView<u8>,
26020 v: &cudarc::driver::CudaView<u8>,
26021 o: &mut cudarc::driver::CudaViewMut<f32>,
26022 head_dim: usize,
26023 n_head: usize,
26024 n_head_kv: usize,
26025 t_kv: usize,
26026 scale: f32,
26027 k_tok_bytes: usize,
26028 v_tok_bytes: usize,
26029 g: bool,
26030 ) -> Result<(), Box<dyn std::error::Error>> {
26031 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
26052 if g && head_dim == 256 && !fa_v4_at(t_kv) {
26056 fa_vec = false;
26057 }
26058 let sp = fa_split_keys(t_kv, n_head_kv);
26059 let n_splits = if fa_vec {
26060 ((t_kv + sp - 1) / sp).max(1)
26061 } else {
26062 ((t_kv + 255) / 256).max(1)
26063 };
26064 let o_len = n_head * n_splits * head_dim;
26065 let ml_len = n_head * n_splits;
26066 let mut part_guard = self.fa_part_pool.lock().unwrap();
26067 if part_guard
26068 .as_ref()
26069 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
26070 .unwrap_or(true)
26071 {
26072 let old = part_guard.take();
26083 let (co, cm) = old
26084 .as_ref()
26085 .map(|pp| (pp.0.len(), pp.1.len()))
26086 .unwrap_or((0, 0));
26087 if let Some(old) = old {
26088 self.fa_part_retired.lock().unwrap().push(old);
26089 }
26090 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
26091 eprintln!(
26092 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
26093 co, o_len, cm, ml_len
26094 );
26095 }
26096 *part_guard =
26097 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
26098 }
26099 let pg = part_guard.as_mut().unwrap();
26100 self.gpu
26101 .stream()
26102 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
26103 self.gpu
26104 .stream()
26105 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
26106 self.gpu
26107 .stream()
26108 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
26109 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
26110 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
26111 let (hd, nh, nhkv, tkvi, nsp) = (
26112 head_dim as i32,
26113 n_head as i32,
26114 n_head_kv as i32,
26115 t_kv as i32,
26116 n_splits as i32,
26117 );
26118 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
26119 let fa_vec = fa_vec && head_dim <= 512 && head_dim.is_multiple_of(32);
26123 let fa512_min = fa512_min_tkv();
26128 let deep = fa_vec
26131 && head_dim == 256
26132 && fa_v4_at(t_kv)
26133 && !g
26134 && fa_deep_at(t_kv)
26135 && !matches!(fa_v4_mode(), "noB3" | "stage");
26136 let (f, cfg) = if fa_vec && head_dim == 512 && t_kv >= fa512_min {
26137 let gqa = (n_head / n_head_kv).max(1) as u32;
26140 let fv = self.fa_func("fa_decode_vec_q_dpl16", head_dim);
26141 (
26142 fv,
26143 LaunchConfig {
26144 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26145 block_dim: (32, gqa, 1),
26146 shared_mem_bytes: 0,
26147 },
26148 )
26149 } else if fa_vec && head_dim <= 256 {
26150 let gqa = (n_head / n_head_kv).max(1) as u32;
26151 static SMEM_TKV: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
26162 let smem_tkv = *SMEM_TKV.get_or_init(|| {
26163 std::env::var("MEMRA_FA_SMEM_TKV")
26164 .ok()
26165 .and_then(|v| v.parse().ok())
26166 .unwrap_or_else(|| {
26167 FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
26168 })
26169 });
26170 if fa_v4_at(t_kv) && head_dim == 256 {
26171 let v4name = match fa_v4_mode() {
26175 "noB3" => "fa_decode_vec_q_v4_noB3", "stage" => "fa_decode_vec_q_v4_stage", _ if deep => "fa_decode_vec_q_v4_deep",
26178 _ => "fa_decode_vec_q_v4",
26179 };
26180 let fv = if g {
26181 self.func_g(v4name)
26182 } else {
26183 self.func(v4name)
26184 };
26185 let shmem = (if deep { 12160 } else { 11520 }
26188 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
26189 use cudarc::driver::sys::CUfunction_attribute_enum as A;
26190 fv.set_attribute(
26191 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26192 shmem as i32,
26193 )?;
26194 (
26195 fv,
26196 LaunchConfig {
26197 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26198 block_dim: (32, gqa, 1),
26199 shared_mem_bytes: shmem,
26200 },
26201 )
26202 } else if fa_v3_active(head_dim) {
26203 let fv = if g {
26206 self.func_g("fa_decode_vec_q_v3")
26207 } else {
26208 self.func("fa_decode_vec_q_v3")
26209 };
26210 let shmem = (32 * head_dim * 2) as u32; (
26212 fv,
26213 LaunchConfig {
26214 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26215 block_dim: (32, gqa, 1),
26216 shared_mem_bytes: shmem,
26217 },
26218 )
26219 } else if fa_v2_on() {
26220 let fv = if g {
26224 self.func_g("fa_decode_vec_q_v2")
26225 } else {
26226 self.func("fa_decode_vec_q_v2")
26227 };
26228 let shmem = (2 * 32 * head_dim * 2) as u32; (
26230 fv,
26231 LaunchConfig {
26232 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26233 block_dim: (32, gqa, 1),
26234 shared_mem_bytes: shmem,
26235 },
26236 )
26237 } else if smem_tkv > 0 && t_kv >= smem_tkv && !g && !(head_dim == 512 && Self::gkv_on())
26238 {
26239 let fv = if g {
26243 self.func_g("fa_decode_vec_q_smem")
26244 } else {
26245 self.func("fa_decode_vec_q_smem")
26246 };
26247 let shmem = (2 * 32 * head_dim * 2) as u32; use cudarc::driver::sys::CUfunction_attribute_enum as A;
26249 fv.set_attribute(
26250 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26251 shmem as i32,
26252 )?;
26253 (
26254 fv,
26255 LaunchConfig {
26256 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26257 block_dim: (32, gqa, 1),
26258 shared_mem_bytes: shmem,
26259 },
26260 )
26261 } else {
26262 let fv = if g {
26265 self.func_g("fa_decode_vec_q")
26266 } else {
26267 self.func("fa_decode_vec_q")
26268 };
26269 (
26270 fv,
26271 LaunchConfig {
26272 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26273 block_dim: (32, gqa, 1),
26274 shared_mem_bytes: 0,
26275 },
26276 )
26277 }
26278 } else {
26279 return self.fa_decode_scalar_unified(
26282 q,
26283 k,
26284 v,
26285 o,
26286 head_dim,
26287 n_head,
26288 n_head_kv,
26289 t_kv,
26290 None,
26291 scale,
26292 n_splits,
26293 if fa_vec { sp } else { 256 },
26294 k_tok_bytes,
26295 v_tok_bytes,
26296 g,
26297 part_o,
26298 part_m,
26299 part_l,
26300 None,
26301 );
26302 };
26303 let __s_b = self.gpu.stream();
26304 let mut b = __s_b.launch_builder(&f);
26305 b.arg(q)
26306 .arg(k)
26307 .arg(v)
26308 .arg(&mut *part_o)
26309 .arg(&mut *part_m)
26310 .arg(&mut *part_l)
26311 .arg(&hd)
26312 .arg(&nh)
26313 .arg(&nhkv)
26314 .arg(&tkvi)
26315 .arg(&scale)
26316 .arg(&nsp)
26317 .arg(&ktb)
26318 .arg(&vtb);
26319 unsafe {
26320 b.launch(cfg)?;
26321 }
26322 let (fc, cfg2) = (
26325 if g {
26326 self.func_g("fa_decode_combine_f32")
26327 } else {
26328 self.fa_func("fa_decode_combine_f32", head_dim)
26329 },
26330 LaunchConfig {
26331 grid_dim: (n_head as u32, 1, 1),
26332 block_dim: (head_dim as u32, 1, 1),
26333 shared_mem_bytes: 0,
26334 },
26335 );
26336 let __s_b2 = self.gpu.stream();
26337 let mut b2 = __s_b2.launch_builder(&fc);
26338 b2.arg(&*part_o)
26339 .arg(&*part_m)
26340 .arg(&*part_l)
26341 .arg(o)
26342 .arg(&hd)
26343 .arg(&nh)
26344 .arg(&nsp);
26345 unsafe {
26346 b2.launch(cfg2)?;
26347 }
26348 Ok(())
26349 }
26350
26351 #[allow(clippy::too_many_arguments)]
26362 pub fn fa_decode_batch_seqs_v4(
26363 &self,
26364 q: &CudaSlice<f32>,
26365 kv_ptrs: &cudarc::driver::CudaView<u64>,
26366 pos_seq: &CudaSlice<i32>,
26367 o: &mut CudaSlice<f32>,
26368 head_dim: usize,
26369 n_head: usize,
26370 n_head_kv: usize,
26371 b_n: usize,
26372 t_kv_max: usize,
26373 scale: f32,
26374 split_keys: usize,
26375 k_tok_bytes: usize,
26376 v_tok_bytes: usize,
26377 ) -> Result<(), Box<dyn std::error::Error>> {
26378 debug_assert!(head_dim == 256, "seqs twin is v4-stamped (hd256 only)");
26379 #[allow(clippy::manual_div_ceil)]
26380 let n_splits_max = (t_kv_max + split_keys - 1) / split_keys;
26382 let o_len = b_n * n_head * n_splits_max * head_dim;
26383 let ml_len = b_n * n_head * n_splits_max;
26384 let mut part_guard = self.fa_part_pool.lock().unwrap();
26385 if part_guard
26386 .as_ref()
26387 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
26388 .unwrap_or(true)
26389 {
26390 let old = part_guard.take();
26401 let (co, cm) = old
26402 .as_ref()
26403 .map(|pp| (pp.0.len(), pp.1.len()))
26404 .unwrap_or((0, 0));
26405 if let Some(old) = old {
26406 self.fa_part_retired.lock().unwrap().push(old);
26407 }
26408 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
26409 eprintln!(
26410 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
26411 co, o_len, cm, ml_len
26412 );
26413 }
26414 *part_guard =
26415 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
26416 }
26417 let pg = part_guard.as_mut().unwrap();
26418 self.gpu
26419 .stream()
26420 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
26421 self.gpu
26422 .stream()
26423 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
26424 self.gpu
26425 .stream()
26426 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
26427 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
26428 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
26429 let (nspm, spk) = (n_splits_max as i32, split_keys as i32);
26430 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
26431 let gqa = (n_head / n_head_kv).max(1) as u32;
26432 let f = self.func("fa_decode_vec_q_seqs_v4");
26433 let shmem = (11520 + 32 * head_dim * 2) as u32;
26435 use cudarc::driver::sys::CUfunction_attribute_enum as A;
26436 f.set_attribute(
26437 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26438 shmem as i32,
26439 )?;
26440 let cfg = LaunchConfig {
26441 grid_dim: (n_head_kv as u32, n_splits_max as u32, b_n as u32),
26442 block_dim: (32, gqa, 1),
26443 shared_mem_bytes: shmem,
26444 };
26445 {
26446 let __s_b = self.gpu.stream();
26447 let mut b = __s_b.launch_builder(&f);
26448 b.arg(q)
26449 .arg(kv_ptrs)
26450 .arg(pos_seq)
26451 .arg(&mut *part_o)
26452 .arg(&mut *part_m)
26453 .arg(&mut *part_l)
26454 .arg(&hd)
26455 .arg(&nh)
26456 .arg(&nhkv)
26457 .arg(&scale)
26458 .arg(&nspm)
26459 .arg(&spk)
26460 .arg(&ktb)
26461 .arg(&vtb);
26462 unsafe {
26463 b.launch(cfg)?;
26464 }
26465 }
26466 let fc = self.func("fa_decode_combine_seqs");
26467 let cfg2 = LaunchConfig {
26468 grid_dim: (n_head as u32, b_n as u32, 1),
26469 block_dim: (head_dim as u32, 1, 1),
26470 shared_mem_bytes: 0,
26471 };
26472 let __s_b2 = self.gpu.stream();
26473 let mut b2 = __s_b2.launch_builder(&fc);
26474 b2.arg(&*part_o)
26475 .arg(&*part_m)
26476 .arg(&*part_l)
26477 .arg(o)
26478 .arg(&hd)
26479 .arg(&nh)
26480 .arg(pos_seq)
26481 .arg(&nspm)
26482 .arg(&spk);
26483 unsafe {
26484 b2.launch(cfg2)?;
26485 }
26486 Ok(())
26487 }
26488
26489 #[allow(clippy::too_many_arguments)]
26496 pub fn append_kv_quantized_seqs(
26497 &self,
26498 k_rows: &CudaSlice<f32>,
26499 v_rows: &CudaSlice<f32>,
26500 kv_ptrs: &cudarc::driver::CudaView<u64>,
26501 pos_seq: &CudaSlice<i32>,
26502 b_n: usize,
26503 kv_dim_k: usize,
26504 kv_dim_v: usize,
26505 k_tok_bytes: usize,
26506 v_tok_bytes: usize,
26507 ) -> Result<(), Box<dyn std::error::Error>> {
26508 let f = self.func("append_quantize_kv_q8_0_q5_1_seqs");
26509 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
26510 let cfg = LaunchConfig {
26511 grid_dim: (nblk, b_n as u32, 1),
26512 block_dim: (32, 1, 1),
26513 shared_mem_bytes: 0,
26514 };
26515 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
26516 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
26517 let __s_b = self.gpu.stream();
26518 let mut b = __s_b.launch_builder(&f);
26519 b.arg(k_rows)
26520 .arg(v_rows)
26521 .arg(kv_ptrs)
26522 .arg(pos_seq)
26523 .arg(&kdk)
26524 .arg(&kdv)
26525 .arg(&ktb)
26526 .arg(&vtb);
26527 unsafe {
26528 b.launch(cfg)?;
26529 }
26530 Ok(())
26531 }
26532
26533 pub fn fa_rows_eligible(&self, base_len: usize, head_dim: usize) -> bool {
26539 std::env::var("MEMRA_NO_FA_VEC").is_err()
26540 && std::env::var("MEMRA_FA_ROWS_OFF").is_err()
26541 && base_len + 1 >= fa_vec_min_tkv()
26542 && head_dim <= 256
26543 && head_dim.is_multiple_of(32)
26544 }
26545
26546 #[allow(clippy::too_many_arguments)]
26555 pub fn fa_decode_rows(
26556 &self,
26557 q: &CudaSlice<f32>,
26558 k: &cudarc::driver::CudaView<u8>,
26559 v: &cudarc::driver::CudaView<u8>,
26560 o: &mut CudaSlice<f32>,
26561 head_dim: usize,
26562 n_head: usize,
26563 n_head_kv: usize,
26564 base_len: usize,
26565 t: usize,
26566 scale: f32,
26567 k_tok_bytes: usize,
26568 v_tok_bytes: usize,
26569 base_dev: Option<(&CudaSlice<i32>, i32)>,
26573 kv_shared: bool,
26576 g: bool,
26580 mut q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
26583 ) -> Result<(), Box<dyn std::error::Error>> {
26584 debug_assert!(
26585 base_len + 1 >= fa_vec_min_tkv() && head_dim <= 512 && head_dim.is_multiple_of(32)
26586 );
26587 let t_kv_max = base_len + t; let mut sp = fa_split_keys(t_kv_max, n_head_kv); if head_dim == 512 {
26594 static SP512: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
26595 let v = *SP512.get_or_init(|| {
26598 std::env::var("MEMRA_FA_SP512")
26599 .ok()
26600 .and_then(|x| x.parse().ok())
26601 .unwrap_or(0)
26602 });
26603 sp = if v >= 8 {
26604 v
26605 } else {
26606 FA_SP512_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
26607 };
26608 }
26609 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
26610 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
26611 let gqa = (n_head / n_head_kv).max(1) as u32;
26612 let mut groups: Vec<(usize, usize, usize)> = Vec::new(); if head_dim == 512 || fa_split_keys(base_len + 1, n_head_kv) == sp {
26623 groups.push((0, t, sp));
26624 } else {
26625 let mut r0 = 0usize;
26626 while r0 < t {
26627 let sp_g = fa_split_keys(base_len + r0 + 1, n_head_kv);
26628 let mut r1 = r0 + 1;
26629 while r1 < t && fa_split_keys(base_len + r1 + 1, n_head_kv) == sp_g {
26630 r1 += 1;
26631 }
26632 groups.push((r0, r1 - r0, sp_g));
26633 r0 = r1;
26634 }
26635 }
26636 static SMEM_TKV_R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
26640 let smem_tkv = *SMEM_TKV_R.get_or_init(|| {
26641 std::env::var("MEMRA_FA_SMEM_TKV")
26642 .ok()
26643 .and_then(|v| v.parse().ok())
26644 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
26645 });
26646 let v4 = fa_v4_at(base_len + t) && head_dim == 256;
26647 let v3 = fa_v3_active(head_dim);
26648 let smem_rows =
26649 head_dim <= 256 && !v3 && !fa_v2_on() && smem_tkv > 0 && t_kv_max >= smem_tkv;
26650 let _ = kv_shared;
26655 let i2 = head_dim == 512 && std::env::var("MEMRA_FA_I2").as_deref() != Ok("0");
26658 static TB512: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
26672 let tb512 = head_dim == 512
26674 && sp <= 32
26675 && n_head / n_head_kv.max(1) <= 16
26676 && *TB512.get_or_init(|| std::env::var("MEMRA_FA_TB512").as_deref() != Ok("0"));
26677 let fname = if tb512 {
26678 "fa_decode_vec_q_rows_v4_512_tb"
26679 } else if i2 {
26680 "fa_decode_vec_q_rows_dpl16_i2"
26681 } else if head_dim == 512 {
26682 "fa_decode_vec_q_rows_dpl16"
26683 }
26684 else if v4 {
26686 "fa_decode_vec_q_rows_v4"
26687 } else if v3 {
26688 "fa_decode_vec_q_rows_v3"
26689 } else if fa_v2_on() {
26690 "fa_decode_vec_q_rows_v2"
26691 } else if smem_rows {
26692 "fa_decode_vec_q_rows_smem"
26693 } else {
26694 "fa_decode_vec_q_rows"
26695 };
26696 let f = if head_dim == 512 {
26697 self.fa_func(fname, head_dim)
26698 } else if g {
26699 self.func_g(if smem_rows {
26707 "fa_decode_vec_q_rows"
26708 } else {
26709 fname
26710 })
26711 } else {
26712 self.func(fname)
26713 };
26714 let shmem = if tb512 {
26715 let gk = Self::gkv_on();
26717 let sh =
26718 (8192 + 1024 + 32 * 512 + 32 * 64 + 32 * head_dim * if gk { 1 } else { 2 }) as u32;
26719 use cudarc::driver::sys::CUfunction_attribute_enum as A;
26720 f.set_attribute(
26721 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26722 sh as i32,
26723 )?;
26724 sh
26725 } else if v4 || v3 || smem_rows || fa_v2_on() {
26726 let sh = (if v4 {
26728 11520 + 32 * head_dim * if g { 1 } else { 2 }
26729 } else if v3 {
26730 32 * head_dim * 2
26731 } else {
26732 2 * 32 * head_dim * 2
26733 }) as u32;
26734 use cudarc::driver::sys::CUfunction_attribute_enum as A;
26735 f.set_attribute(
26736 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26737 sh as i32,
26738 )?;
26739 sh
26740 } else {
26741 0
26742 };
26743 for &(r0, t_g, sp_g) in &groups {
26747 let n_splits_g = (base_len + r0 + t_g).div_ceil(sp_g);
26748 let (nspm, spk) = (n_splits_g as i32, sp_g as i32);
26749 let base_i = (base_len + r0) as i32;
26750 let o_len = t_g * n_head * n_splits_g * head_dim;
26751 let ml_len = t_g * n_head * n_splits_g;
26752 let mut part_guard = self.fa_part_pool.lock().unwrap();
26753 if part_guard
26754 .as_ref()
26755 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
26756 .unwrap_or(true)
26757 {
26758 let old = part_guard.take();
26769 let (co, cm) = old
26770 .as_ref()
26771 .map(|pp| (pp.0.len(), pp.1.len()))
26772 .unwrap_or((0, 0));
26773 if let Some(old) = old {
26774 self.fa_part_retired.lock().unwrap().push(old);
26775 }
26776 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
26777 eprintln!(
26778 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
26779 co, o_len, cm, ml_len
26780 );
26781 }
26782 *part_guard =
26783 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
26784 }
26785 let pg = part_guard.as_mut().unwrap();
26786 self.gpu
26787 .stream()
26788 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
26789 self.gpu
26790 .stream()
26791 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
26792 self.gpu
26793 .stream()
26794 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
26795 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
26796 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
26797 let qv = self.view(q, t * n_head * head_dim);
26798 let q_g = qv.slice(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
26799 let cfg = LaunchConfig {
26800 grid_dim: (n_head_kv as u32, n_splits_g as u32, t_g as u32),
26801 block_dim: (32, gqa, 1),
26802 shared_mem_bytes: shmem,
26803 };
26804 {
26805 let __s_b = self.gpu.stream();
26806 let mut b = __s_b.launch_builder(&f);
26807 if tb512 {
26808 let (bd, plus) =
26810 base_dev.expect("hd512 rows twin requires a device base counter");
26811 let plus_g = plus + r0 as i32;
26812 let nr = t_g as i32;
26813 if Self::pdl_on() && Self::pdl_wb_on() {
26814 use cudarc::driver::{DevicePtr, DevicePtrMut};
26816 let s = &self.gpu.stream();
26817 let (pq, _b0) = q_g.device_ptr(s);
26818 let (pk, _b1) = k.device_ptr(s);
26819 let (pv, _b2) = v.device_ptr(s);
26820 let (po, _b3) = part_o.device_ptr_mut(s);
26821 let (pm, _b4) = part_m.device_ptr_mut(s);
26822 let (pl, _b5) = part_l.device_ptr_mut(s);
26823 let (pb, _b6) = bd.device_ptr(s);
26824 let mut ps = [
26825 &pq as *const _ as *mut std::ffi::c_void,
26826 &pk as *const _ as *mut _,
26827 &pv as *const _ as *mut _,
26828 &po as *const _ as *mut _,
26829 &pm as *const _ as *mut _,
26830 &pl as *const _ as *mut _,
26831 &hd as *const _ as *mut _,
26832 &nh as *const _ as *mut _,
26833 &nhkv as *const _ as *mut _,
26834 &pb as *const _ as *mut _,
26835 &plus_g as *const _ as *mut _,
26836 &scale as *const _ as *mut _,
26837 &nspm as *const _ as *mut _,
26838 &spk as *const _ as *mut _,
26839 &ktb as *const _ as *mut _,
26840 &vtb as *const _ as *mut _,
26841 &nr as *const _ as *mut _,
26842 ];
26843 unsafe {
26844 self.launch_pdl_flash(
26845 Self::gkv_on(),
26846 "fa_decode_vec_q_rows_v4_512_tb",
26847 (n_head_kv as u32, n_splits_g as u32, 1),
26848 (32, gqa, 1),
26849 shmem,
26850 &mut ps,
26851 )?;
26852 }
26853 } else {
26854 let cfg_tb = LaunchConfig {
26855 grid_dim: (n_head_kv as u32, n_splits_g as u32, 1),
26856 block_dim: (32, gqa, 1),
26857 shared_mem_bytes: shmem,
26858 };
26859 b.arg(&q_g)
26860 .arg(k)
26861 .arg(v)
26862 .arg(&mut *part_o)
26863 .arg(&mut *part_m)
26864 .arg(&mut *part_l)
26865 .arg(&hd)
26866 .arg(&nh)
26867 .arg(&nhkv)
26868 .arg(bd)
26869 .arg(&plus_g)
26870 .arg(&scale)
26871 .arg(&nspm)
26872 .arg(&spk)
26873 .arg(&ktb)
26874 .arg(&vtb)
26875 .arg(&nr);
26876 unsafe {
26877 b.launch(cfg_tb)?;
26878 }
26879 }
26880 } else if head_dim == 512 {
26881 let (bd, plus) =
26882 base_dev.expect("hd512 rows twin requires a device base counter");
26883 let plus_g = plus + r0 as i32;
26884 b.arg(&q_g)
26885 .arg(k)
26886 .arg(v)
26887 .arg(&mut *part_o)
26888 .arg(&mut *part_m)
26889 .arg(&mut *part_l)
26890 .arg(&hd)
26891 .arg(&nh)
26892 .arg(&nhkv)
26893 .arg(bd)
26894 .arg(&plus_g)
26895 .arg(&scale)
26896 .arg(&nspm)
26897 .arg(&spk)
26898 .arg(&ktb)
26899 .arg(&vtb);
26900 unsafe {
26901 b.launch(cfg)?;
26902 }
26903 } else {
26904 b.arg(&q_g)
26905 .arg(k)
26906 .arg(v)
26907 .arg(&mut *part_o)
26908 .arg(&mut *part_m)
26909 .arg(&mut *part_l)
26910 .arg(&hd)
26911 .arg(&nh)
26912 .arg(&nhkv)
26913 .arg(&base_i)
26914 .arg(&scale)
26915 .arg(&nspm)
26916 .arg(&spk)
26917 .arg(&ktb)
26918 .arg(&vtb);
26919 unsafe {
26920 b.launch(cfg)?;
26921 }
26922 }
26923 }
26924 let cfg2 = LaunchConfig {
26925 grid_dim: (n_head as u32, t_g as u32, 1),
26926 block_dim: (head_dim as u32, 1, 1),
26927 shared_mem_bytes: 0,
26928 };
26929 let mut o_g = o.slice_mut(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
26930 if head_dim == 512 {
26931 let (bd, plus) = base_dev.unwrap();
26934 let plus_g = plus + r0 as i32;
26935 if let Some((oq, od)) = q8_out.as_mut() {
26936 debug_assert!(t == 1, "rows q8 emit is a t=1 decode arm");
26938 if Self::pdl_on() && Self::pdl_wb_on() {
26939 use cudarc::driver::{DevicePtr, DevicePtrMut};
26941 let s = &self.gpu.stream();
26942 let (po, _g0) = part_o.device_ptr(s);
26943 let (pm, _g1) = part_m.device_ptr(s);
26944 let (pl, _g2) = part_l.device_ptr(s);
26945 let (pq, _g3) = oq.device_ptr_mut(s);
26946 let (pd, _g4) = od.device_ptr_mut(s);
26947 let (pb, _g5) = bd.device_ptr(s);
26948 let mut ps = [
26949 &po as *const _ as *mut std::ffi::c_void,
26950 &pm as *const _ as *mut _,
26951 &pl as *const _ as *mut _,
26952 &pq as *const _ as *mut _,
26953 &pd as *const _ as *mut _,
26954 &hd as *const _ as *mut _,
26955 &nh as *const _ as *mut _,
26956 &pb as *const _ as *mut _,
26957 &plus_g as *const _ as *mut _,
26958 &nspm as *const _ as *mut _,
26959 &spk as *const _ as *mut _,
26960 ];
26961 unsafe {
26962 self.launch_pdl_flash(
26963 Self::gkv_on(),
26964 "fa_decode_combine_rows_dc_q8_1",
26965 cfg2.grid_dim,
26966 cfg2.block_dim,
26967 0,
26968 &mut ps,
26969 )?;
26970 }
26971 continue;
26972 }
26973 let fc = self.fa_func("fa_decode_combine_rows_dc_q8_1", head_dim);
26974 let __s_b2 = self.gpu.stream();
26975 let mut b2 = __s_b2.launch_builder(&fc);
26976 b2.arg(&*part_o)
26977 .arg(&*part_m)
26978 .arg(&*part_l)
26979 .arg(&mut **oq)
26980 .arg(&mut **od)
26981 .arg(&hd)
26982 .arg(&nh)
26983 .arg(bd)
26984 .arg(&plus_g)
26985 .arg(&nspm)
26986 .arg(&spk);
26987 unsafe {
26988 b2.launch(cfg2)?;
26989 }
26990 continue;
26991 }
26992 let fc = self.fa_func("fa_decode_combine_rows_dc", head_dim);
26993 let __s_b2 = self.gpu.stream();
26994 let mut b2 = __s_b2.launch_builder(&fc);
26995 b2.arg(&*part_o)
26996 .arg(&*part_m)
26997 .arg(&*part_l)
26998 .arg(&mut o_g)
26999 .arg(&hd)
27000 .arg(&nh)
27001 .arg(bd)
27002 .arg(&plus_g)
27003 .arg(&nspm)
27004 .arg(&spk);
27005 unsafe {
27006 b2.launch(cfg2)?;
27007 }
27008 } else {
27009 assert!(
27012 q8_out.is_none(),
27013 "rows q8 emit requires the hd512 dc combine"
27014 );
27015 let fc = self.func("fa_decode_combine_rows");
27016 let __s_b2 = self.gpu.stream();
27017 let mut b2 = __s_b2.launch_builder(&fc);
27018 b2.arg(&*part_o)
27019 .arg(&*part_m)
27020 .arg(&*part_l)
27021 .arg(&mut o_g)
27022 .arg(&hd)
27023 .arg(&nh)
27024 .arg(&base_i)
27025 .arg(&nspm)
27026 .arg(&spk);
27027 unsafe {
27028 b2.launch(cfg2)?;
27029 }
27030 }
27031 }
27032 Ok(())
27033 }
27034
27035 #[allow(clippy::too_many_arguments)]
27039 pub fn fa_decode_rows_w(
27040 &self,
27041 q: &CudaSlice<f32>,
27042 k: &cudarc::driver::CudaView<u8>,
27043 v: &cudarc::driver::CudaView<u8>,
27044 o: &mut CudaSlice<f32>,
27045 head_dim: usize,
27046 n_head: usize,
27047 n_head_kv: usize,
27048 base_dev: &CudaSlice<i32>,
27049 base_plus: i32,
27050 t: usize,
27051 scale: f32,
27052 window: usize,
27053 k_tok_bytes: usize,
27054 v_tok_bytes: usize,
27055 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
27056 ) -> Result<(), Box<dyn std::error::Error>> {
27057 debug_assert!(head_dim == 256);
27062 let sp = {
27070 static SPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
27071 let v = *SPW.get_or_init(|| {
27072 std::env::var("MEMRA_FA_SPW")
27073 .ok()
27074 .and_then(|x| x.parse().ok())
27075 .unwrap_or(0)
27076 });
27077 if v >= 8 {
27078 v
27079 } else {
27080 FA_SPW_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
27081 }
27082 };
27083 #[allow(clippy::manual_div_ceil)]
27084 let n_splits_max = (window + sp - 1) / sp;
27086 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
27087 let (nspm, spk, wini) = (n_splits_max as i32, sp as i32, window as i32);
27088 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
27089 let gqa = (n_head / n_head_kv).max(1) as u32;
27090 let o_len = t * n_head * n_splits_max * head_dim;
27091 let ml_len = t * n_head * n_splits_max;
27092 let mut part_guard = self.fa_part_pool.lock().unwrap();
27093 if part_guard
27094 .as_ref()
27095 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
27096 .unwrap_or(true)
27097 {
27098 let old = part_guard.take();
27109 let (co, cm) = old
27110 .as_ref()
27111 .map(|pp| (pp.0.len(), pp.1.len()))
27112 .unwrap_or((0, 0));
27113 if let Some(old) = old {
27114 self.fa_part_retired.lock().unwrap().push(old);
27115 }
27116 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
27117 eprintln!(
27118 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
27119 co, o_len, cm, ml_len
27120 );
27121 }
27122 *part_guard =
27123 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
27124 }
27125 let pg = part_guard.as_mut().unwrap();
27126 self.gpu
27127 .stream()
27128 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
27129 self.gpu
27130 .stream()
27131 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
27132 self.gpu
27133 .stream()
27134 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
27135 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
27136 static SMEM_TKV_W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
27142 let smem_tkv = *SMEM_TKV_W.get_or_init(|| {
27143 std::env::var("MEMRA_FA_SMEM_TKV")
27144 .ok()
27145 .and_then(|v| v.parse().ok())
27146 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
27147 });
27148 use cudarc::driver::sys::CUfunction_attribute_enum as A;
27154 let wg = Self::wkv_on();
27159 let sp2 =
27162 gqa <= 4 && fa_v4_at(window) && std::env::var("MEMRA_FA_SPW2").as_deref() != Ok("0");
27163 if sp2 {
27164 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
27165 if Self::pdl_on() && Self::pdl_wb_on() {
27166 use cudarc::driver::{DevicePtr, DevicePtrMut};
27168 let s = &self.gpu.stream();
27169 let (pq, _b0) = q.device_ptr(s);
27170 let (pk, _b1) = k.device_ptr(s);
27171 let (pv, _b2) = v.device_ptr(s);
27172 let (po, _b3) = part_o.device_ptr_mut(s);
27173 let (pm, _b4) = part_m.device_ptr_mut(s);
27174 let (pl, _b5) = part_l.device_ptr_mut(s);
27175 let (pb, _b6) = base_dev.device_ptr(s);
27176 let mut ps = [
27177 &pq as *const _ as *mut std::ffi::c_void,
27178 &pk as *const _ as *mut _,
27179 &pv as *const _ as *mut _,
27180 &po as *const _ as *mut _,
27181 &pm as *const _ as *mut _,
27182 &pl as *const _ as *mut _,
27183 &hd as *const _ as *mut _,
27184 &nh as *const _ as *mut _,
27185 &nhkv as *const _ as *mut _,
27186 &pb as *const _ as *mut _,
27187 &base_plus as *const _ as *mut _,
27188 &scale as *const _ as *mut _,
27189 &nspm as *const _ as *mut _,
27190 &spk as *const _ as *mut _,
27191 &ktb as *const _ as *mut _,
27192 &vtb as *const _ as *mut _,
27193 &wini as *const _ as *mut _,
27194 ];
27195 unsafe {
27196 self.launch_pdl_flash(
27197 wg,
27198 "fa_decode_vec_q_rows_v4_w_sp",
27199 (n_head_kv as u32, n_splits_max as u32, t as u32),
27200 (32, gqa + 1, 1),
27201 sh,
27202 &mut ps,
27203 )?;
27204 }
27205 } else {
27206 let f = if wg {
27207 self.func_g("fa_decode_vec_q_rows_v4_w_sp")
27208 } else {
27209 self.func("fa_decode_vec_q_rows_v4_w_sp")
27210 };
27211 f.set_attribute(
27212 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27213 sh as i32,
27214 )?;
27215 let cfg = LaunchConfig {
27216 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
27217 block_dim: (32, gqa + 1, 1),
27218 shared_mem_bytes: sh,
27219 };
27220 let __s_b = self.gpu.stream();
27221 let mut b = __s_b.launch_builder(&f);
27222 b.arg(q)
27223 .arg(k)
27224 .arg(v)
27225 .arg(&mut *part_o)
27226 .arg(&mut *part_m)
27227 .arg(&mut *part_l)
27228 .arg(&hd)
27229 .arg(&nh)
27230 .arg(&nhkv)
27231 .arg(base_dev)
27232 .arg(&base_plus)
27233 .arg(&scale)
27234 .arg(&nspm)
27235 .arg(&spk)
27236 .arg(&ktb)
27237 .arg(&vtb)
27238 .arg(&wini);
27239 unsafe {
27240 b.launch(cfg)?;
27241 }
27242 }
27243 } else {
27244 if fa_v4_at(window) && Self::pdl_on() && Self::pdl_wb_on() {
27245 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
27247 use cudarc::driver::{DevicePtr, DevicePtrMut};
27248 let s = &self.gpu.stream();
27249 let (pq, _b0) = q.device_ptr(s);
27250 let (pk, _b1) = k.device_ptr(s);
27251 let (pv, _b2) = v.device_ptr(s);
27252 let (po, _b3) = part_o.device_ptr_mut(s);
27253 let (pm, _b4) = part_m.device_ptr_mut(s);
27254 let (pl, _b5) = part_l.device_ptr_mut(s);
27255 let (pb, _b6) = base_dev.device_ptr(s);
27256 let mut ps = [
27257 &pq as *const _ as *mut std::ffi::c_void,
27258 &pk as *const _ as *mut _,
27259 &pv as *const _ as *mut _,
27260 &po as *const _ as *mut _,
27261 &pm as *const _ as *mut _,
27262 &pl as *const _ as *mut _,
27263 &hd as *const _ as *mut _,
27264 &nh as *const _ as *mut _,
27265 &nhkv as *const _ as *mut _,
27266 &pb as *const _ as *mut _,
27267 &base_plus as *const _ as *mut _,
27268 &scale as *const _ as *mut _,
27269 &nspm as *const _ as *mut _,
27270 &spk as *const _ as *mut _,
27271 &ktb as *const _ as *mut _,
27272 &vtb as *const _ as *mut _,
27273 &wini as *const _ as *mut _,
27274 ];
27275 unsafe {
27276 self.launch_pdl_flash(
27277 wg,
27278 "fa_decode_vec_q_rows_v4_w",
27279 (n_head_kv as u32, n_splits_max as u32, t as u32),
27280 (32, gqa, 1),
27281 sh,
27282 &mut ps,
27283 )?;
27284 }
27285 } else {
27286 let pick = |name: &str| {
27287 if wg {
27288 self.func_g(name)
27289 } else {
27290 self.func(name)
27291 }
27292 };
27293 let (f, sh) = if fa_v4_at(window) {
27294 let f = pick("fa_decode_vec_q_rows_v4_w");
27295 (f, (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32)
27296 } else if smem_tkv > 0 && window >= smem_tkv {
27297 (
27300 pick("fa_decode_vec_q_rows_smem_w"),
27301 (2 * 32 * head_dim * 2) as u32,
27302 )
27303 } else {
27304 (pick("fa_decode_vec_q_rows_reg_w"), 0u32)
27305 };
27306 f.set_attribute(
27307 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27308 sh as i32,
27309 )?;
27310 let cfg = LaunchConfig {
27311 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
27312 block_dim: (32, gqa, 1),
27313 shared_mem_bytes: sh,
27314 };
27315 let __s_b = self.gpu.stream();
27316 let mut b = __s_b.launch_builder(&f);
27317 b.arg(q)
27318 .arg(k)
27319 .arg(v)
27320 .arg(&mut *part_o)
27321 .arg(&mut *part_m)
27322 .arg(&mut *part_l)
27323 .arg(&hd)
27324 .arg(&nh)
27325 .arg(&nhkv)
27326 .arg(base_dev)
27327 .arg(&base_plus)
27328 .arg(&scale)
27329 .arg(&nspm)
27330 .arg(&spk)
27331 .arg(&ktb)
27332 .arg(&vtb)
27333 .arg(&wini);
27334 unsafe {
27335 b.launch(cfg)?;
27336 }
27337 }
27338 }
27339 let cfg2 = LaunchConfig {
27340 grid_dim: (n_head as u32, t as u32, 1),
27341 block_dim: (head_dim as u32, 1, 1),
27342 shared_mem_bytes: 0,
27343 };
27344 if let Some((oq, od)) = q8_out {
27345 if Self::pdl_on() && Self::pdl_wb_on() {
27348 use cudarc::driver::{DevicePtr, DevicePtrMut};
27350 let s = &self.gpu.stream();
27351 let (po, _g0) = part_o.device_ptr(s);
27352 let (pm, _g1) = part_m.device_ptr(s);
27353 let (pl, _g2) = part_l.device_ptr(s);
27354 let (pq, _g3) = oq.device_ptr_mut(s);
27355 let (pd, _g4) = od.device_ptr_mut(s);
27356 let mut ps = [
27357 &po as *const _ as *mut std::ffi::c_void,
27358 &pm as *const _ as *mut _,
27359 &pl as *const _ as *mut _,
27360 &pq as *const _ as *mut _,
27361 &pd as *const _ as *mut _,
27362 &hd as *const _ as *mut _,
27363 &nh as *const _ as *mut _,
27364 &nspm as *const _ as *mut _,
27365 &spk as *const _ as *mut _,
27366 &wini as *const _ as *mut _,
27367 ];
27368 unsafe {
27369 self.launch_pdl_flash(
27370 wg,
27371 "fa_decode_combine_rows_w_q8_1",
27372 cfg2.grid_dim,
27373 cfg2.block_dim,
27374 0,
27375 &mut ps,
27376 )?;
27377 }
27378 return Ok(());
27379 }
27380 let fc = if wg {
27381 self.func_g("fa_decode_combine_rows_w_q8_1")
27382 } else {
27383 self.func("fa_decode_combine_rows_w_q8_1")
27384 };
27385 let __s_b2 = self.gpu.stream();
27386 let mut b2 = __s_b2.launch_builder(&fc);
27387 b2.arg(&*part_o)
27388 .arg(&*part_m)
27389 .arg(&*part_l)
27390 .arg(oq)
27391 .arg(od)
27392 .arg(&hd)
27393 .arg(&nh)
27394 .arg(&nspm)
27395 .arg(&spk)
27396 .arg(&wini);
27397 unsafe {
27398 b2.launch(cfg2)?;
27399 }
27400 return Ok(());
27401 }
27402 let fc = if wg {
27403 self.func_g("fa_decode_combine_rows_w")
27404 } else {
27405 self.func("fa_decode_combine_rows_w")
27406 };
27407 let __s_b2 = self.gpu.stream();
27408 let mut b2 = __s_b2.launch_builder(&fc);
27409 b2.arg(&*part_o)
27410 .arg(&*part_m)
27411 .arg(&*part_l)
27412 .arg(o)
27413 .arg(&hd)
27414 .arg(&nh)
27415 .arg(&nspm)
27416 .arg(&spk)
27417 .arg(&wini);
27418 unsafe {
27419 b2.launch(cfg2)?;
27420 }
27421 Ok(())
27422 }
27423
27424 #[allow(clippy::too_many_arguments)]
27430 pub fn fa_decode_rows_dc(
27431 &self,
27432 q: &CudaSlice<f32>,
27433 k: &cudarc::driver::CudaView<u8>,
27434 v: &cudarc::driver::CudaView<u8>,
27435 o: &mut CudaSlice<f32>,
27436 head_dim: usize,
27437 n_head: usize,
27438 n_head_kv: usize,
27439 base_dev: &CudaSlice<i32>,
27440 t_kv_upper: usize,
27441 t: usize,
27442 scale: f32,
27443 k_tok_bytes: usize,
27444 v_tok_bytes: usize,
27445 base_plus: i32,
27446 g: bool,
27447 ) -> Result<(), Box<dyn std::error::Error>> {
27448 let v4 = head_dim == 256 && fa_v4_at(t_kv_upper);
27449 assert!(
27450 v4 || fa_v3_active(head_dim),
27451 "stream fa rows requires the v3 or v4 lane"
27452 );
27453 assert!(v4 || base_plus == 0, "v3_dc kernel takes no plus arg");
27454 if v4 {
27455 let sp = fa_split_keys(t_kv_upper, n_head_kv);
27456 #[allow(clippy::manual_div_ceil)]
27457 let n_splits_max = (t_kv_upper + sp - 1) / sp;
27459 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
27460 let (nspm, spk) = (n_splits_max as i32, sp as i32);
27461 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
27462 let gqa = (n_head / n_head_kv).max(1) as u32;
27463 let o_len = t * n_head * n_splits_max * head_dim;
27464 let ml_len = t * n_head * n_splits_max;
27465 let mut part_guard = self.fa_part_pool.lock().unwrap();
27466 if part_guard
27467 .as_ref()
27468 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
27469 .unwrap_or(true)
27470 {
27471 let old = part_guard.take();
27482 let (co, cm) = old
27483 .as_ref()
27484 .map(|pp| (pp.0.len(), pp.1.len()))
27485 .unwrap_or((0, 0));
27486 if let Some(old) = old {
27487 self.fa_part_retired.lock().unwrap().push(old);
27488 }
27489 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
27490 eprintln!(
27491 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
27492 co, o_len, cm, ml_len
27493 );
27494 }
27495 *part_guard =
27496 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
27497 }
27498 let pg = part_guard.as_mut().unwrap();
27499 self.gpu
27500 .stream()
27501 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
27502 self.gpu
27503 .stream()
27504 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
27505 self.gpu
27506 .stream()
27507 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
27508 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
27509 let f = if g {
27510 self.func_g("fa_decode_vec_q_rows_v4_dc")
27511 } else {
27512 self.func("fa_decode_vec_q_rows_v4_dc")
27513 };
27514 let sh = (11520 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
27515 use cudarc::driver::sys::CUfunction_attribute_enum as A;
27516 f.set_attribute(
27517 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27518 sh as i32,
27519 )?;
27520 let cfg = LaunchConfig {
27521 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
27522 block_dim: (32, gqa, 1),
27523 shared_mem_bytes: sh,
27524 };
27525 let __s_b = self.gpu.stream();
27526 let mut b = __s_b.launch_builder(&f);
27527 b.arg(q)
27528 .arg(k)
27529 .arg(v)
27530 .arg(&mut *part_o)
27531 .arg(&mut *part_m)
27532 .arg(&mut *part_l)
27533 .arg(&hd)
27534 .arg(&nh)
27535 .arg(&nhkv)
27536 .arg(base_dev)
27537 .arg(&base_plus)
27538 .arg(&scale)
27539 .arg(&nspm)
27540 .arg(&spk)
27541 .arg(&ktb)
27542 .arg(&vtb);
27543 unsafe {
27544 b.launch(cfg)?;
27545 }
27546 let fc = self.func("fa_decode_combine_rows_dc");
27547 let cfg2 = LaunchConfig {
27548 grid_dim: (n_head as u32, t as u32, 1),
27549 block_dim: (head_dim as u32, 1, 1),
27550 shared_mem_bytes: 0,
27551 };
27552 let __s_b2 = self.gpu.stream();
27553 let mut b2 = __s_b2.launch_builder(&fc);
27554 b2.arg(&*part_o)
27555 .arg(&*part_m)
27556 .arg(&*part_l)
27557 .arg(o)
27558 .arg(&hd)
27559 .arg(&nh)
27560 .arg(base_dev)
27561 .arg(&base_plus)
27562 .arg(&nspm)
27563 .arg(&spk);
27564 unsafe {
27565 b2.launch(cfg2)?;
27566 }
27567 return Ok(());
27568 }
27569 let sp = fa_split_keys(t_kv_upper, n_head_kv);
27570 #[allow(clippy::manual_div_ceil)]
27571 let n_splits_max = (t_kv_upper + sp - 1) / sp;
27573 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
27574 let (nspm, spk) = (n_splits_max as i32, sp as i32);
27575 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
27576 let gqa = (n_head / n_head_kv).max(1) as u32;
27577 let o_len = t * n_head * n_splits_max * head_dim;
27578 let ml_len = t * n_head * n_splits_max;
27579 let mut part_guard = self.fa_part_pool.lock().unwrap();
27580 if part_guard
27581 .as_ref()
27582 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
27583 .unwrap_or(true)
27584 {
27585 let old = part_guard.take();
27596 let (co, cm) = old
27597 .as_ref()
27598 .map(|pp| (pp.0.len(), pp.1.len()))
27599 .unwrap_or((0, 0));
27600 if let Some(old) = old {
27601 self.fa_part_retired.lock().unwrap().push(old);
27602 }
27603 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
27604 eprintln!(
27605 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
27606 co, o_len, cm, ml_len
27607 );
27608 }
27609 *part_guard =
27610 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
27611 }
27612 let pg = part_guard.as_mut().unwrap();
27613 self.gpu
27614 .stream()
27615 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
27616 self.gpu
27617 .stream()
27618 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
27619 self.gpu
27620 .stream()
27621 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
27622 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
27623 let f = self.func("fa_decode_vec_q_rows_v3_dc");
27624 let sh = (32 * head_dim * 2) as u32;
27625 use cudarc::driver::sys::CUfunction_attribute_enum as A;
27626 f.set_attribute(
27627 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27628 sh as i32,
27629 )?;
27630 let cfg = LaunchConfig {
27631 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
27632 block_dim: (32, gqa, 1),
27633 shared_mem_bytes: sh,
27634 };
27635 let __s_b = self.gpu.stream();
27636 let mut b = __s_b.launch_builder(&f);
27637 b.arg(q)
27638 .arg(k)
27639 .arg(v)
27640 .arg(&mut *part_o)
27641 .arg(&mut *part_m)
27642 .arg(&mut *part_l)
27643 .arg(&hd)
27644 .arg(&nh)
27645 .arg(&nhkv)
27646 .arg(base_dev)
27647 .arg(&scale)
27648 .arg(&nspm)
27649 .arg(&spk)
27650 .arg(&ktb)
27651 .arg(&vtb);
27652 unsafe {
27653 b.launch(cfg)?;
27654 }
27655 let fc = self.func("fa_decode_combine_rows_dc");
27656 let cfg2 = LaunchConfig {
27657 grid_dim: (n_head as u32, t as u32, 1),
27658 block_dim: (head_dim as u32, 1, 1),
27659 shared_mem_bytes: 0,
27660 };
27661 let plus0 = 0i32;
27662 let __s_b2 = self.gpu.stream();
27663 let mut b2 = __s_b2.launch_builder(&fc);
27664 b2.arg(&*part_o)
27665 .arg(&*part_m)
27666 .arg(&*part_l)
27667 .arg(o)
27668 .arg(&hd)
27669 .arg(&nh)
27670 .arg(base_dev)
27671 .arg(&plus0)
27672 .arg(&nspm)
27673 .arg(&spk);
27674 unsafe {
27675 b2.launch(cfg2)?;
27676 }
27677 Ok(())
27678 }
27679
27680 #[allow(clippy::too_many_arguments)] pub fn fa_decode_dc(
27692 &self,
27693 q: &CudaSlice<f32>,
27694 k: &cudarc::driver::CudaView<u8>,
27695 v: &cudarc::driver::CudaView<u8>,
27696 o: &mut CudaSlice<f32>,
27697 head_dim: usize,
27698 n_head: usize,
27699 n_head_kv: usize,
27700 t_kv_dev: &CudaSlice<i32>,
27701 bucket_max: usize,
27702 scale: f32,
27703 k_tok_bytes: usize,
27704 v_tok_bytes: usize,
27705 g: bool,
27706 ) -> Result<(), Box<dyn std::error::Error>> {
27707 self.fa_decode_dc_q8(
27708 q,
27709 k,
27710 v,
27711 o,
27712 head_dim,
27713 n_head,
27714 n_head_kv,
27715 t_kv_dev,
27716 bucket_max,
27717 scale,
27718 k_tok_bytes,
27719 v_tok_bytes,
27720 g,
27721 None,
27722 )
27723 }
27724
27725 #[allow(clippy::too_many_arguments)]
27728 #[allow(clippy::manual_div_ceil)] pub fn fa_decode_dc_q8(
27730 &self,
27731 q: &CudaSlice<f32>,
27732 k: &cudarc::driver::CudaView<u8>,
27733 v: &cudarc::driver::CudaView<u8>,
27734 o: &mut CudaSlice<f32>,
27735 head_dim: usize,
27736 n_head: usize,
27737 n_head_kv: usize,
27738 t_kv_dev: &CudaSlice<i32>,
27739 bucket_max: usize,
27740 scale: f32,
27741 k_tok_bytes: usize,
27742 v_tok_bytes: usize,
27743 g: bool,
27744 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
27745 ) -> Result<(), Box<dyn std::error::Error>> {
27746 let mut fa_vec =
27754 std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
27755 if g && head_dim == 256 && !fa_v4_at(bucket_max) {
27756 fa_vec = false;
27757 } let sp = fa_split_keys(bucket_max, n_head_kv);
27759 let n_splits = if fa_vec {
27760 ((bucket_max + sp - 1) / sp).max(1)
27761 } else {
27762 ((bucket_max + 255) / 256).max(1)
27763 };
27764 let o_len = n_head * n_splits * head_dim;
27765 let ml_len = n_head * n_splits;
27766 let mut part_guard = self.fa_part_pool.lock().unwrap();
27767 if part_guard
27768 .as_ref()
27769 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
27770 .unwrap_or(true)
27771 {
27772 let old = part_guard.take();
27783 let (co, cm) = old
27784 .as_ref()
27785 .map(|pp| (pp.0.len(), pp.1.len()))
27786 .unwrap_or((0, 0));
27787 if let Some(old) = old {
27788 self.fa_part_retired.lock().unwrap().push(old);
27789 }
27790 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
27791 eprintln!(
27792 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
27793 co, o_len, cm, ml_len
27794 );
27795 }
27796 *part_guard =
27797 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
27798 }
27799 let pg = part_guard.as_mut().unwrap();
27800 self.gpu
27801 .stream()
27802 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
27803 self.gpu
27804 .stream()
27805 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
27806 self.gpu
27807 .stream()
27808 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
27809 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
27810 let (hd, nh, nhkv, nsp) = (
27811 head_dim as i32,
27812 n_head as i32,
27813 n_head_kv as i32,
27814 n_splits as i32,
27815 );
27816 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
27817 let fa_vec = fa_vec && head_dim <= 512 && head_dim.is_multiple_of(32);
27818 let deep = fa_vec
27821 && head_dim == 256
27822 && fa_v4_at(bucket_max)
27823 && !g
27824 && fa_deep_at(bucket_max)
27825 && !matches!(fa_v4_mode(), "noB3" | "stage");
27826 let (f, cfg) = if fa_vec
27827 && head_dim == 512
27828 && bucket_max >= {
27829 static FA512_MIN_DC: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
27830 *FA512_MIN_DC.get_or_init(|| {
27831 std::env::var("MEMRA_FA512_MIN")
27832 .ok()
27833 .and_then(|v| v.parse().ok())
27834 .unwrap_or(512)
27835 })
27836 } {
27837 let gqa = (n_head / n_head_kv).max(1) as u32;
27839 (
27840 self.fa_func("fa_decode_vec_q_dpl16_dc", head_dim),
27841 LaunchConfig {
27842 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27843 block_dim: (32, gqa, 1),
27844 shared_mem_bytes: 0,
27845 },
27846 )
27847 } else if fa_vec && head_dim == 512 {
27848 let q_view = q.as_view();
27851 let mut o_view = o.as_view_mut();
27852 return self.fa_decode_scalar_unified(
27853 &q_view,
27854 k,
27855 v,
27856 &mut o_view,
27857 head_dim,
27858 n_head,
27859 n_head_kv,
27860 0,
27861 Some(t_kv_dev),
27862 scale,
27863 n_splits,
27864 sp,
27865 k_tok_bytes,
27866 v_tok_bytes,
27867 g,
27868 &mut *part_o,
27869 &mut *part_m,
27870 &mut *part_l,
27871 q8_out,
27872 );
27873 } else if fa_vec && head_dim == 256 && fa_v4_at(bucket_max) {
27874 let gqa = (n_head / n_head_kv).max(1) as u32;
27877 let fv = if g {
27878 self.func_g("fa_decode_vec_q_v4_dc")
27879 } else if deep {
27880 self.func("fa_decode_vec_q_v4_deep_dc")
27881 } else {
27882 self.func("fa_decode_vec_q_v4_dc")
27883 };
27884 let shmem =
27885 (if deep { 12160 } else { 11520 } + 32 * head_dim * if g { 1 } else { 2 }) as u32;
27886 use cudarc::driver::sys::CUfunction_attribute_enum as A;
27887 fv.set_attribute(
27888 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27889 shmem as i32,
27890 )?;
27891 (
27892 fv,
27893 LaunchConfig {
27894 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27895 block_dim: (32, gqa, 1),
27896 shared_mem_bytes: shmem,
27897 },
27898 )
27899 } else if fa_vec && fa_v3_active(head_dim) {
27900 let gqa = (n_head / n_head_kv).max(1) as u32;
27903 let fv = if g {
27904 self.func_g("fa_decode_vec_q_v3_dc")
27905 } else {
27906 self.func("fa_decode_vec_q_v3_dc")
27907 };
27908 let shmem = (32 * head_dim * 2) as u32; (
27910 fv,
27911 LaunchConfig {
27912 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27913 block_dim: (32, gqa, 1),
27914 shared_mem_bytes: shmem,
27915 },
27916 )
27917 } else if fa_vec && fa_v2_on() {
27918 let gqa = (n_head / n_head_kv).max(1) as u32;
27922 let fv = if g {
27923 self.func_g("fa_decode_vec_q_v2_dc")
27924 } else {
27925 self.func("fa_decode_vec_q_v2_dc")
27926 };
27927 let shmem = (2 * 32 * head_dim * 2) as u32; (
27929 fv,
27930 LaunchConfig {
27931 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27932 block_dim: (32, gqa, 1),
27933 shared_mem_bytes: shmem,
27934 },
27935 )
27936 } else if fa_vec {
27937 let gqa = (n_head / n_head_kv).max(1) as u32;
27938 let fv = if g {
27940 self.func_g("fa_decode_vec_q_dc")
27941 } else {
27942 self.func("fa_decode_vec_q_dc")
27943 };
27944 (
27945 fv,
27946 LaunchConfig {
27947 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27948 block_dim: (32, gqa, 1),
27949 shared_mem_bytes: 0,
27950 },
27951 )
27952 } else {
27953 let q_view = q.as_view();
27954 let mut o_view = o.as_view_mut();
27955 return self.fa_decode_scalar_unified(
27956 &q_view,
27957 k,
27958 v,
27959 &mut o_view,
27960 head_dim,
27961 n_head,
27962 n_head_kv,
27963 0,
27964 Some(t_kv_dev),
27965 scale,
27966 n_splits,
27967 if fa_vec { sp } else { 256 },
27968 k_tok_bytes,
27969 v_tok_bytes,
27970 g,
27971 &mut *part_o,
27972 &mut *part_m,
27973 &mut *part_l,
27974 q8_out,
27975 );
27976 };
27977 let ski = sp as i32; let __s_b = self.gpu.stream();
27979 let mut b = __s_b.launch_builder(&f);
27980 b.arg(q)
27981 .arg(k)
27982 .arg(v)
27983 .arg(&mut *part_o)
27984 .arg(&mut *part_m)
27985 .arg(&mut *part_l)
27986 .arg(&hd)
27987 .arg(&nh)
27988 .arg(&nhkv)
27989 .arg(t_kv_dev)
27990 .arg(&scale)
27991 .arg(&nsp)
27992 .arg(&ski)
27993 .arg(&ktb)
27994 .arg(&vtb);
27995 unsafe {
27996 b.launch(cfg)?;
27997 }
27998 let cfg2 = LaunchConfig {
27999 grid_dim: (n_head as u32, 1, 1),
28000 block_dim: (head_dim as u32, 1, 1),
28001 shared_mem_bytes: 0,
28002 };
28003 if let Some((oq, od)) = q8_out {
28004 let fc = if g {
28005 self.func_g("fa_decode_combine_q8_1")
28006 } else {
28007 self.fa_func("fa_decode_combine_q8_1", head_dim)
28008 };
28009 let __s_b2 = self.gpu.stream();
28010 let mut b2 = __s_b2.launch_builder(&fc);
28011 b2.arg(&*part_o)
28012 .arg(&*part_m)
28013 .arg(&*part_l)
28014 .arg(oq)
28015 .arg(od)
28016 .arg(&hd)
28017 .arg(&nh)
28018 .arg(&nsp);
28019 unsafe {
28020 b2.launch(cfg2)?;
28021 }
28022 return Ok(());
28023 }
28024 let fc = if g {
28025 self.func_g("fa_decode_combine_f32")
28026 } else {
28027 self.fa_func("fa_decode_combine_f32", head_dim)
28028 };
28029 let __s_b2 = self.gpu.stream();
28030 let mut b2 = __s_b2.launch_builder(&fc);
28031 b2.arg(&*part_o)
28032 .arg(&*part_m)
28033 .arg(&*part_l)
28034 .arg(o)
28035 .arg(&hd)
28036 .arg(&nh)
28037 .arg(&nsp);
28038 unsafe {
28039 b2.launch(cfg2)?;
28040 }
28041 Ok(())
28042 }
28043
28044 #[allow(clippy::too_many_arguments)]
28048 pub fn append_kv_quantized_dcw(
28049 &self,
28050 k_row: &CudaSlice<f32>,
28051 v_row: &CudaSlice<f32>,
28052 kc: &mut CudaSlice<u8>,
28053 vc: &mut CudaSlice<u8>,
28054 len_dev: &CudaSlice<i32>,
28055 base_dev: Option<&CudaSlice<i32>>,
28056 kv_dim_k: usize,
28057 kv_dim_v: usize,
28058 k_tok_bytes: usize,
28059 v_tok_bytes: usize,
28060 ) -> Result<(), Box<dyn std::error::Error>> {
28061 let f = self.func("append_quantize_kv_q8_0_q5_1_dcw");
28062 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
28063 let cfg = LaunchConfig {
28064 grid_dim: (nblk, 1, 1),
28065 block_dim: (32, 1, 1),
28066 shared_mem_bytes: 0,
28067 };
28068 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
28069 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
28070 let null: u64 = 0;
28071 let __s_b = self.gpu.stream();
28072 let mut b = __s_b.launch_builder(&f);
28073 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(len_dev);
28074 match base_dev {
28075 Some(base) => {
28076 b.arg(base);
28077 }
28078 None => {
28079 b.arg(&null);
28080 }
28081 }
28082 b.arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
28083 unsafe {
28084 b.launch(cfg)?;
28085 }
28086 Ok(())
28087 }
28088
28089 pub fn inc_i32(&self, counter: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
28091 let f = self.func("inc_i32");
28092 let cfg = LaunchConfig {
28093 grid_dim: (1, 1, 1),
28094 block_dim: (1, 1, 1),
28095 shared_mem_bytes: 0,
28096 };
28097 let __s_b = self.gpu.stream();
28098 let mut b = __s_b.launch_builder(&f);
28099 b.arg(counter);
28100 unsafe {
28101 b.launch(cfg)?;
28102 }
28103 Ok(())
28104 }
28105
28106 #[allow(clippy::too_many_arguments)]
28115 #[allow(clippy::type_complexity)] fn fa_part_alloc(
28140 &self,
28141 o_len: usize,
28142 ml_len: usize,
28143 co: usize,
28144 cm: usize,
28145 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
28146 static GROWS: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
28147 let n = GROWS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
28148 if n < 64 {
28149 eprintln!(
28150 "[fa-pool] grow #{n} dev={} o_len {co} -> {o_len} ml_len {cm} -> {ml_len} (retired kept, zero={})",
28151 self.ctx().ordinal(),
28152 fa_part_zero_on()
28153 );
28154 }
28155 let mut po = self.alloc_uninit::<f32>(o_len)?;
28156 let mut pm = self.alloc_uninit::<f32>(ml_len)?;
28157 let mut pl = self.alloc_uninit::<f32>(ml_len)?;
28158 if fa_part_zero_on() {
28159 self.gpu.stream().memset_zeros(&mut po)?;
28160 self.gpu.stream().memset_zeros(&mut pm)?;
28161 self.gpu.stream().memset_zeros(&mut pl)?;
28162 }
28163 Ok((po, pm, pl))
28164 }
28165
28166 fn fa_part_pool_grow(
28167 &self,
28168 part_guard: &mut Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>,
28169 o_len: usize,
28170 ml_len: usize,
28171 ) -> Result<(), Box<dyn std::error::Error>> {
28172 if part_guard
28173 .as_ref()
28174 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
28175 .unwrap_or(true)
28176 {
28177 let old = part_guard.take();
28178 let (co, cm) = old
28179 .as_ref()
28180 .map(|pp| (pp.0.len(), pp.1.len()))
28181 .unwrap_or((0, 0));
28182 if let Some(old) = old {
28183 self.fa_part_retired.lock().unwrap().push(old);
28184 }
28185 *part_guard =
28198 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
28199 }
28200 Ok(())
28201 }
28202
28203 pub fn fa_dcw_pool_ensure(
28206 &self,
28207 head_dim: usize,
28208 n_head: usize,
28209 n_head_kv: usize,
28210 bucket_max: usize,
28211 ) -> Result<(), Box<dyn std::error::Error>> {
28212 let sp = fa_split_keys(bucket_max, n_head_kv);
28213 #[allow(clippy::manual_div_ceil)]
28214 let n_splits = ((bucket_max + sp - 1) / sp).max(1);
28216 let o_len = n_head * n_splits * head_dim;
28217 let ml_len = n_head * n_splits;
28218 let mut part_guard = self.fa_part_pool.lock().unwrap();
28219 self.fa_part_pool_grow(&mut part_guard, o_len, ml_len)
28220 }
28221
28222 #[allow(clippy::too_many_arguments)]
28230 pub fn fa_decode_dcw2(
28231 &self,
28232 q2: &CudaSlice<f32>,
28233 k_ring: &cudarc::driver::CudaView<u8>,
28234 v_ring: &cudarc::driver::CudaView<u8>,
28235 o2: &mut CudaSlice<f32>,
28236 head_dim: usize,
28237 n_head: usize,
28238 n_head_kv: usize,
28239 len_dev: &CudaSlice<i32>,
28240 base_dev: Option<&CudaSlice<i32>>,
28241 window: usize,
28242 bucket_max: usize,
28243 scale: f32,
28244 k_tok_bytes: usize,
28245 v_tok_bytes: usize,
28246 gate2: &CudaSlice<f32>,
28247 ) -> Result<(), Box<dyn std::error::Error>> {
28248 let fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
28249 if !fa_vec || head_dim > 256 || !head_dim.is_multiple_of(32) || !fa_v3_on() {
28250 return Err("fa_decode_dcw2 supports the default v3-vec class only".into());
28251 }
28252 let sp = fa_split_keys(bucket_max, n_head_kv);
28253 #[allow(clippy::manual_div_ceil)]
28254 let n_splits = ((bucket_max + sp - 1) / sp).max(1);
28256 let o_len = 2 * n_head * n_splits * head_dim;
28258 let ml_len = 2 * n_head * n_splits;
28259 let mut part_guard = self.fa_part_pool.lock().unwrap();
28260 Self::fa_part_pool_grow(self, &mut part_guard, o_len, ml_len)?;
28261 let pg = part_guard.as_mut().unwrap();
28262 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
28263 let (hd, nh, nhkv, nsp) = (
28264 head_dim as i32,
28265 n_head as i32,
28266 n_head_kv as i32,
28267 n_splits as i32,
28268 );
28269 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
28270 let (ski, win) = (sp as i32, window as i32);
28271 let gqa = (n_head / n_head_kv).max(1) as u32;
28272 let smem = (32 * head_dim * 2) as u32;
28273 let f = self.func("fa_decode_vec_q_v3_dcw2");
28274 let cfg = LaunchConfig {
28275 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
28276 block_dim: (32, gqa, 1),
28277 shared_mem_bytes: smem,
28278 };
28279 let null: u64 = 0;
28280 {
28281 let __s_b = self.gpu.stream();
28282 let mut b = __s_b.launch_builder(&f);
28283 b.arg(q2)
28284 .arg(k_ring)
28285 .arg(v_ring)
28286 .arg(&mut *part_o)
28287 .arg(&mut *part_m)
28288 .arg(&mut *part_l)
28289 .arg(&hd)
28290 .arg(&nh)
28291 .arg(&nhkv)
28292 .arg(len_dev);
28293 match base_dev {
28294 Some(base) => {
28295 b.arg(base);
28296 }
28297 None => {
28298 b.arg(&null);
28299 }
28300 }
28301 b.arg(&win)
28302 .arg(&scale)
28303 .arg(&nsp)
28304 .arg(&ski)
28305 .arg(&ktb)
28306 .arg(&vtb);
28307 unsafe {
28308 b.launch(cfg)?;
28309 }
28310 }
28311 let fc = {
28315 static CS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28316 if *CS.get_or_init(|| std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1")) {
28317 self.func("fa_decode_combine_gate_f32_s")
28318 } else {
28319 self.func("fa_decode_combine_gate_f32")
28320 }
28321 };
28322 let combine_shared = std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1");
28323 let nh2 = (2 * n_head) as i32;
28324 let cfg2 = LaunchConfig {
28325 grid_dim: ((2 * n_head) as u32, 1, 1),
28326 block_dim: (head_dim as u32, 1, 1),
28327 shared_mem_bytes: if combine_shared {
28328 (2 * n_splits * 4) as u32
28329 } else {
28330 0
28331 },
28332 };
28333 let __s_b2 = self.gpu.stream();
28334 let mut b2 = __s_b2.launch_builder(&fc);
28335 b2.arg(&*part_o)
28336 .arg(&*part_m)
28337 .arg(&*part_l)
28338 .arg(gate2)
28339 .arg(o2)
28340 .arg(&hd)
28341 .arg(&nh2)
28342 .arg(&nsp);
28343 unsafe {
28344 b2.launch(cfg2)?;
28345 }
28346 Ok(())
28347 }
28348
28349 #[allow(clippy::too_many_arguments)]
28358 pub fn fa_decode_dcw_rows(
28359 &self,
28360 q_rows: &CudaSlice<f32>,
28361 tab: &CudaSlice<u64>,
28362 o_rows: &mut CudaSlice<f32>,
28363 t: usize,
28364 head_dim: usize,
28365 n_head: usize,
28366 n_head_kv: usize,
28367 window: usize,
28368 max_ns: usize,
28369 scale: f32,
28370 k_tok_bytes: usize,
28371 v_tok_bytes: usize,
28372 gate_rows: &CudaSlice<f32>,
28373 ) -> Result<(), Box<dyn std::error::Error>> {
28374 if std::env::var("MEMRA_NO_FA_VEC").is_ok()
28375 || head_dim > 256
28376 || !head_dim.is_multiple_of(32)
28377 || !fa_v3_on()
28378 {
28379 return Err("fa_decode_dcw_rows supports the default v3-vec class only".into());
28380 }
28381 if fa_sm_count() < 128
28382 || std::env::var("MEMRA_FA_SPLIT").is_ok()
28383 || std::env::var("MEMRA_FA_SP_SHORT").is_ok()
28384 || std::env::var("MEMRA_FA_SP16").is_ok()
28385 {
28386 return Err(
28387 "fa_decode_dcw_rows embeds the big-rig split ladder; env split overrides \
28388 (or a <128-SM rig) keep the per-row path"
28389 .into(),
28390 );
28391 }
28392 if t == 0 || t > 32 || max_ns == 0 || tab.len() < t * 6 {
28393 return Err("fa_decode_dcw_rows geometry".into());
28394 }
28395 let o_len = t * n_head * max_ns * head_dim;
28396 let ml_len = t * n_head * max_ns;
28397 let mut part_guard = self.fa_part_pool.lock().unwrap();
28398 Self::fa_part_pool_grow(self, &mut part_guard, o_len, ml_len)?;
28399 let pg = part_guard.as_mut().unwrap();
28400 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
28401 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
28402 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
28403 let (win, mns) = (window as i32, max_ns as i32);
28404 let gqa = (n_head / n_head_kv).max(1) as u32;
28405 let smem = (32 * head_dim * 2) as u32;
28406 let f = self.func("fa_decode_vec_q_v3_dcw_rows");
28407 let cfg = LaunchConfig {
28408 grid_dim: (n_head_kv as u32, max_ns as u32, t as u32),
28409 block_dim: (32, gqa, 1),
28410 shared_mem_bytes: smem,
28411 };
28412 {
28413 let __s_b = self.gpu.stream();
28414 let mut b = __s_b.launch_builder(&f);
28415 b.arg(q_rows)
28416 .arg(tab)
28417 .arg(&mut *part_o)
28418 .arg(&mut *part_m)
28419 .arg(&mut *part_l)
28420 .arg(&hd)
28421 .arg(&nh)
28422 .arg(&nhkv)
28423 .arg(&win)
28424 .arg(&scale)
28425 .arg(&mns)
28426 .arg(&ktb)
28427 .arg(&vtb);
28428 unsafe {
28429 b.launch(cfg)?;
28430 }
28431 }
28432 let fc = {
28436 static CS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28437 if *CS.get_or_init(|| std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1")) {
28438 self.func("fa_decode_combine_gate_f32_s")
28439 } else {
28440 self.func("fa_decode_combine_gate_f32")
28441 }
28442 };
28443 let combine_shared = std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1");
28444 let nht = (t * n_head) as i32;
28445 let cfg2 = LaunchConfig {
28446 grid_dim: ((t * n_head) as u32, 1, 1),
28447 block_dim: (head_dim as u32, 1, 1),
28448 shared_mem_bytes: if combine_shared {
28449 (2 * max_ns * 4) as u32
28450 } else {
28451 0
28452 },
28453 };
28454 let __s_b2 = self.gpu.stream();
28455 let mut b2 = __s_b2.launch_builder(&fc);
28456 b2.arg(&*part_o)
28457 .arg(&*part_m)
28458 .arg(&*part_l)
28459 .arg(gate_rows)
28460 .arg(o_rows)
28461 .arg(&hd)
28462 .arg(&nht)
28463 .arg(&mns);
28464 unsafe {
28465 b2.launch(cfg2)?;
28466 }
28467 Ok(())
28468 }
28469
28470 #[allow(clippy::too_many_arguments)] pub fn fa_decode_dcw(
28472 &self,
28473 q: &CudaSlice<f32>,
28474 k_ring: &cudarc::driver::CudaView<u8>,
28475 v_ring: &cudarc::driver::CudaView<u8>,
28476 o: &mut CudaSlice<f32>,
28477 head_dim: usize,
28478 n_head: usize,
28479 n_head_kv: usize,
28480 len_dev: &CudaSlice<i32>,
28481 base_dev: Option<&CudaSlice<i32>>,
28482 window: usize,
28483 bucket_max: usize,
28484 scale: f32,
28485 k_tok_bytes: usize,
28486 v_tok_bytes: usize,
28487 fused_gate: Option<&CudaSlice<f32>>,
28491 ) -> Result<(), Box<dyn std::error::Error>> {
28492 let fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
28493 if !fa_vec || head_dim > 256 || !head_dim.is_multiple_of(32) || !fa_v3_on() {
28494 return Err("fa_decode_dcw supports the default v3-vec class only (bucket >= vec floor, head_dim <= 256, MEMRA_FA_V3 on); keep eager outside it"
28495 .into());
28496 }
28497 let sp = fa_split_keys(bucket_max, n_head_kv);
28498 #[allow(clippy::manual_div_ceil)]
28499 let n_splits = ((bucket_max + sp - 1) / sp).max(1);
28501 let o_len = n_head * n_splits * head_dim;
28502 let ml_len = n_head * n_splits;
28503 let mut part_guard = self.fa_part_pool.lock().unwrap();
28504 Self::fa_part_pool_grow(self, &mut part_guard, o_len, ml_len)?;
28505 let pg = part_guard.as_mut().unwrap();
28506 static MEMSET_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28511 let memset_on = *MEMSET_ON
28516 .get_or_init(|| std::env::var("MEMRA_FA_DCW_MEMSET").as_deref() != Ok("0"))
28517 || crate::tp::token_graph_building();
28518 if memset_on {
28519 self.gpu
28520 .stream()
28521 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
28522 self.gpu
28523 .stream()
28524 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
28525 self.gpu
28526 .stream()
28527 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
28528 }
28529 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
28530 let (hd, nh, nhkv, nsp) = (
28531 head_dim as i32,
28532 n_head as i32,
28533 n_head_kv as i32,
28534 n_splits as i32,
28535 );
28536 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
28537 let (ski, win) = (sp as i32, window as i32);
28538 let gqa = (n_head / n_head_kv).max(1) as u32;
28539 let smem = (32 * head_dim * 2) as u32; static U8: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28543 static HOIST: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
28544 let hoist = *HOIST.get_or_init(|| match std::env::var("MEMRA_FA_HOIST").as_deref() {
28545 Ok("2") => 2,
28546 Ok("1") => 1,
28547 _ => 0,
28548 });
28549 static FPROF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28554 let fprof = *FPROF.get_or_init(|| std::env::var("MEMRA_FA_PROF").as_deref() == Ok("1"));
28555 static PROF_BUF: std::sync::Mutex<Option<(usize, CudaSlice<u64>)>> =
28556 std::sync::Mutex::new(None);
28557 static HS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28561 let hs2 = *HS.get_or_init(|| std::env::var("MEMRA_FA_HSPLIT").as_deref() == Ok("2"))
28562 && (n_head / n_head_kv).is_multiple_of(2)
28563 && (n_head / n_head_kv) >= 2;
28564 let f = if fprof {
28565 self.func("fa_decode_vec_q_v3_dcw_prof")
28566 } else if hs2 {
28567 self.func("fa_decode_vec_q_v3_dcw_hs2")
28568 } else if hoist == 2 {
28569 self.func("fa_decode_vec_q_v3_dcw_hc")
28571 } else if hoist == 1 {
28572 self.func("fa_decode_vec_q_v3_dcw_h")
28574 } else if *U8.get_or_init(|| std::env::var("MEMRA_FA_UNROLL").as_deref() == Ok("8")) {
28575 self.func("fa_decode_vec_q_v3_dcw_u8")
28576 } else {
28577 self.func("fa_decode_vec_q_v3_dcw")
28578 };
28579 let cfg = LaunchConfig {
28580 grid_dim: if hs2 {
28581 ((2 * n_head_kv) as u32, n_splits as u32, 1)
28582 } else {
28583 (n_head_kv as u32, n_splits as u32, 1)
28584 },
28585 block_dim: if hs2 { (32, gqa / 2, 1) } else { (32, gqa, 1) },
28586 shared_mem_bytes: smem,
28587 };
28588 let null: u64 = 0;
28589 let __s_b = self.gpu.stream();
28590 let mut b = __s_b.launch_builder(&f);
28591 b.arg(q)
28592 .arg(k_ring)
28593 .arg(v_ring)
28594 .arg(&mut *part_o)
28595 .arg(&mut *part_m)
28596 .arg(&mut *part_l)
28597 .arg(&hd)
28598 .arg(&nh)
28599 .arg(&nhkv)
28600 .arg(len_dev);
28601 match base_dev {
28602 Some(base) => {
28603 b.arg(base);
28604 }
28605 None => {
28606 b.arg(&null);
28607 }
28608 }
28609 b.arg(&win)
28610 .arg(&scale)
28611 .arg(&nsp)
28612 .arg(&ski)
28613 .arg(&ktb)
28614 .arg(&vtb);
28615 if fprof {
28616 let mut guard = PROF_BUF.lock().map_err(|_| "fa prof buffer lock")?;
28617 if guard
28618 .as_ref()
28619 .is_none_or(|(d, _)| *d != self.ctx().ordinal())
28620 {
28621 *guard = Some((self.ctx().ordinal(), self.htod_u64(&[0u64; 8])?));
28622 }
28623 let (_, buf) = guard.as_mut().expect("armed above");
28624 b.arg(&*buf);
28625 unsafe {
28626 b.launch(cfg)?;
28627 }
28628 static CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
28629 let n = CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1;
28630 if n.is_multiple_of(430) {
28631 self.stream().synchronize()?;
28632 let h = self.dtoh_u64(buf)?;
28633 let phases = ["setup", "stageV", "b1_klo", "b2_soft", "sync", "b3_vacc"];
28634 let tot: u64 = h[..6].iter().sum();
28635 let mut line = format!("[fa-prof] calls={n} keys={} cycles={tot}", h[6]);
28636 for (i, name) in phases.iter().enumerate() {
28637 let pct = if tot > 0 {
28638 h[i] as f64 / tot as f64 * 100.0
28639 } else {
28640 0.0
28641 };
28642 line.push_str(&format!(" {name}={pct:.1}%"));
28643 }
28644 if h[6] > 0 {
28645 line.push_str(&format!(" cyc/key={:.0}", tot as f64 / h[6] as f64));
28646 }
28647 eprintln!("{line}");
28648 }
28649 } else {
28650 unsafe {
28651 b.launch(cfg)?;
28652 }
28653 }
28654 let mut combine_shared = false;
28655 let fc = if fused_gate.is_some() {
28656 static CS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28659 if *CS.get_or_init(|| std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1")) {
28660 combine_shared = true;
28661 self.func("fa_decode_combine_gate_f32_s")
28662 } else {
28663 self.func("fa_decode_combine_gate_f32")
28664 }
28665 } else {
28666 self.fa_func("fa_decode_combine_f32", head_dim)
28667 };
28668 let cfg2 = LaunchConfig {
28669 grid_dim: (n_head as u32, 1, 1),
28670 block_dim: (head_dim as u32, 1, 1),
28671 shared_mem_bytes: if combine_shared {
28672 (2 * n_splits * 4) as u32
28673 } else {
28674 0
28675 },
28676 };
28677 let __s_b2 = self.gpu.stream();
28678 let mut b2 = __s_b2.launch_builder(&fc);
28679 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l);
28680 if let Some(gate_row) = fused_gate {
28681 b2.arg(gate_row);
28682 }
28683 b2.arg(o).arg(&hd).arg(&nh).arg(&nsp);
28684 unsafe {
28685 b2.launch(cfg2)?;
28686 }
28687 Ok(())
28688 }
28689
28690 #[allow(clippy::manual_div_ceil)] pub fn fa_geom_eager(
28697 &self,
28698 t_kv: usize,
28699 head_dim: usize,
28700 n_head_kv: usize,
28701 g: bool,
28702 ) -> (bool, usize) {
28703 let fa_ok = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
28707 let vec512 = fa_ok && head_dim == 512 && t_kv >= fa512_min_tkv();
28713 let mut fa_vec = vec512 || (fa_ok && head_dim <= 256 && head_dim.is_multiple_of(32));
28714 if g && head_dim == 256 && !fa_v4_at(t_kv) {
28720 fa_vec = false;
28721 }
28722 let sp = fa_split_keys(t_kv, n_head_kv);
28723 let n_splits = if fa_vec {
28724 ((t_kv + sp - 1) / sp).max(1)
28725 } else {
28726 ((t_kv + 255) / 256).max(1)
28727 };
28728 (fa_vec, n_splits)
28729 }
28730
28731 pub fn fa_bucket_key(
28737 &self,
28738 t_kv: usize,
28739 head_dim: usize,
28740 n_head_kv: usize,
28741 g: bool,
28742 ) -> (bool, usize) {
28743 self.fa_geom_eager(t_kv, head_dim, n_head_kv, g)
28744 }
28745
28746 #[allow(clippy::type_complexity)] pub fn capture_graph_retained<F>(
28759 &self,
28760 step: F,
28761 ) -> Result<
28762 (
28763 cudarc::driver::CudaGraph,
28764 Vec<Box<dyn std::any::Any + Send>>,
28765 ),
28766 Box<dyn std::error::Error>,
28767 >
28768 where
28769 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
28770 {
28771 use cudarc::driver::sys::CUgraphInstantiate_flags;
28772 self.capture_graph_retained_flags(
28773 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
28774 step,
28775 )
28776 }
28777
28778 #[allow(clippy::type_complexity)] pub fn capture_graph_retained_flags<F>(
28784 &self,
28785 flags: cudarc::driver::sys::CUgraphInstantiate_flags,
28786 mut step: F,
28787 ) -> Result<
28788 (
28789 cudarc::driver::CudaGraph,
28790 Vec<Box<dyn std::any::Any + Send>>,
28791 ),
28792 Box<dyn std::error::Error>,
28793 >
28794 where
28795 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
28796 {
28797 use cudarc::driver::sys::CUstreamCaptureMode;
28798 self.capture_keep.lock().unwrap().clear();
28806 let was_tracking = self.gpu.ctx.is_event_tracking();
28807 if was_tracking {
28808 unsafe {
28809 self.gpu.ctx.disable_event_tracking();
28810 }
28811 }
28812 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
28813 self.capture_keep_on
28814 .store(true, std::sync::atomic::Ordering::Relaxed);
28815 let w = (|| {
28816 step(self)?;
28817 step(self)
28818 })();
28819 self.capture_keep_on
28820 .store(false, std::sync::atomic::Ordering::Relaxed);
28821 w?;
28822 self.gpu.stream().synchronize()?;
28823 self.gpu
28824 .stream()
28825 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
28826 let r = step(self);
28827 let g = self.gpu.stream().end_capture(flags);
28828 r?;
28829 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
28830 graph.upload()?;
28831 Ok(graph)
28832 };
28833 let result = run();
28834 self.capture_keep_on
28835 .store(false, std::sync::atomic::Ordering::Relaxed);
28836 if was_tracking {
28837 unsafe {
28838 self.gpu.ctx.enable_event_tracking();
28839 }
28840 }
28841 let keeper = std::mem::take(&mut *self.capture_keep.lock().unwrap());
28842 Ok((result?, keeper))
28843 }
28844
28845 #[allow(clippy::type_complexity)] pub fn capture_graph_retained_nowarm<F>(
28852 &self,
28853 mut step: F,
28854 ) -> Result<
28855 (
28856 cudarc::driver::CudaGraph,
28857 Vec<Box<dyn std::any::Any + Send>>,
28858 ),
28859 Box<dyn std::error::Error>,
28860 >
28861 where
28862 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
28863 {
28864 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
28865 let was_tracking = self.gpu.ctx.is_event_tracking();
28866 if was_tracking {
28867 unsafe {
28868 self.gpu.ctx.disable_event_tracking();
28869 }
28870 }
28871 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
28872 self.gpu.stream().synchronize()?;
28873 self.gpu
28874 .stream()
28875 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
28876 let r = step(self);
28877 let g = self.gpu.stream().end_capture(
28878 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
28879 );
28880 r?;
28881 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
28882 graph.upload()?;
28883 Ok(graph)
28884 };
28885 let result = run();
28886 if was_tracking {
28887 unsafe {
28888 self.gpu.ctx.enable_event_tracking();
28889 }
28890 }
28891 Ok((result?, Vec::new()))
28892 }
28893
28894 pub fn capture_graph<F>(
28895 &self,
28896 mut step: F,
28897 ) -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>>
28898 where
28899 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
28900 {
28901 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
28902 let was_tracking = self.gpu.ctx.is_event_tracking();
28910 if was_tracking {
28911 unsafe {
28912 self.gpu.ctx.disable_event_tracking();
28913 }
28914 }
28915 let iflag = {
28922 static F: std::sync::OnceLock<CUgraphInstantiate_flags> = std::sync::OnceLock::new();
28923 *F.get_or_init(|| match std::env::var("MEMRA_GRAPH_IFLAG").as_deref() {
28924 Ok("upload") => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_UPLOAD,
28927 Ok("priority") => {
28928 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY
28929 }
28930 _ => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
28931 })
28932 };
28933 let ct = {
28940 static T: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28941 *T.get_or_init(|| std::env::var("MEMRA_GRAPH_CAPTIME").as_deref() == Ok("1"))
28942 };
28943 let warmups = {
28966 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
28967 *W.get_or_init(|| {
28968 std::env::var("MEMRA_GRAPH_WARMUPS")
28969 .ok()
28970 .and_then(|v| v.parse().ok())
28971 .filter(|n| *n >= 1)
28972 .unwrap_or(1)
28973 })
28974 };
28975 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
28976 let t_w = std::time::Instant::now();
28977 for _ in 0..warmups {
28979 step(self)?;
28980 }
28981 self.gpu.stream().synchronize()?;
28982 let ms_warm = t_w.elapsed().as_secs_f64() * 1e3;
28983 let t_c = std::time::Instant::now();
28985 self.gpu
28986 .stream()
28987 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
28988 let r = step(self);
28991 let ms_body = t_c.elapsed().as_secs_f64() * 1e3;
28992 let t_i = std::time::Instant::now();
28993 let g = self.gpu.stream().end_capture(iflag);
28994 let ms_inst = t_i.elapsed().as_secs_f64() * 1e3;
28995 r?;
28996 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
28997 let t_u = std::time::Instant::now();
28998 graph.upload()?;
28999 if ct {
29000 println!(
29001 "[graph-captime] warmup2x {ms_warm:.2} ms capture-body {ms_body:.2} ms \
29002 instantiate {ms_inst:.2} ms upload {:.2} ms",
29003 t_u.elapsed().as_secs_f64() * 1e3
29004 );
29005 }
29006 Ok(graph)
29007 };
29008 let result = run();
29009 if was_tracking {
29010 unsafe {
29011 self.gpu.ctx.enable_event_tracking();
29012 }
29013 }
29014 result
29015 }
29016
29017 #[allow(clippy::too_many_arguments)] pub fn gdn_scan_s128_view(
29020 &self,
29021 q: &CudaSlice<f32>,
29022 k: &CudaSlice<f32>,
29023 v: &CudaSlice<f32>,
29024 g: &CudaSlice<f32>,
29025 beta: &CudaSlice<f32>,
29026 state_in: &cudarc::driver::CudaView<f32>,
29027 state_out: &mut cudarc::driver::CudaViewMut<f32>,
29028 o: &mut CudaSlice<f32>,
29029 n_head: usize,
29030 t: usize,
29031 scale: f32,
29032 ) -> Result<(), Box<dyn std::error::Error>> {
29033 let f = self.func("gdn_scan_s128");
29034 const S_V: u32 = 128;
29035 const WARP: u32 = 32;
29036 const COLS: u32 = 4;
29037 let cfg = LaunchConfig {
29038 grid_dim: (n_head as u32, 1, S_V / COLS),
29039 block_dim: (WARP, COLS, 1),
29040 shared_mem_bytes: 0,
29041 };
29042 let (h, ti) = (n_head as i32, t as i32);
29043 let __s_b = self.gpu.stream();
29044 let mut b = __s_b.launch_builder(&f);
29045 b.arg(q)
29046 .arg(k)
29047 .arg(v)
29048 .arg(g)
29049 .arg(beta)
29050 .arg(state_in)
29051 .arg(state_out)
29052 .arg(o)
29053 .arg(&h)
29054 .arg(&ti)
29055 .arg(&scale);
29056 unsafe {
29057 b.launch(cfg)?;
29058 }
29059 Ok(())
29060 }
29061
29062 #[allow(clippy::too_many_arguments)]
29064 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_view(
29067 &self,
29068 x: &cudarc::driver::CudaView<f32>,
29069 w: &CudaSlice<f32>,
29070 y: &mut CudaSlice<f32>,
29071 conv_dim: usize,
29072 t: usize,
29073 d_conv: usize,
29074 silu: bool,
29075 ) -> Result<(), Box<dyn std::error::Error>> {
29076 let f = self.func("ssm_conv1d_silu_f32");
29077 let cfg = LaunchConfig {
29079 grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
29080 block_dim: (256, 1, 1),
29081 shared_mem_bytes: 0,
29082 };
29083 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
29084 let __s_b = self.gpu.stream();
29085 let mut b = __s_b.launch_builder(&f);
29086 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
29087 unsafe {
29088 b.launch(cfg)?;
29089 }
29090 Ok(())
29091 }
29092
29093 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_tm(
29101 &self,
29102 qkv_tm: &CudaSlice<f32>,
29103 w: &CudaSlice<f32>,
29104 y: &mut CudaSlice<f32>,
29105 conv_dim: usize,
29106 t: usize,
29107 d_conv: usize,
29108 ) -> Result<(), Box<dyn std::error::Error>> {
29109 let f = self.func("ssm_conv1d_tm_f32");
29110 let cfg = LaunchConfig {
29111 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
29112 block_dim: (256, 1, 1),
29113 shared_mem_bytes: 0,
29114 };
29115 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29116 let __s_b = self.gpu.stream();
29117 let mut b = __s_b.launch_builder(&f);
29118 b.arg(qkv_tm).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
29119 unsafe {
29120 b.launch(cfg)?;
29121 }
29122 Ok(())
29123 }
29124
29125 #[allow(clippy::too_many_arguments)] pub fn ssm_conv1d_tm_state(
29134 &self,
29135 qkv_tm: &CudaSlice<f32>,
29136 conv_state: &mut CudaSlice<f32>,
29137 w: &CudaSlice<f32>,
29138 y: &mut CudaSlice<f32>,
29139 conv_dim: usize,
29140 t: usize,
29141 d_conv: usize,
29142 ) -> Result<(), Box<dyn std::error::Error>> {
29143 self.ssm_conv1d_tm_state_pad(qkv_tm, conv_state, w, y, conv_dim, t, d_conv, None)
29144 }
29145
29146 #[allow(clippy::too_many_arguments)]
29149 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_tm_state_pad(
29151 &self,
29152 qkv_tm: &CudaSlice<f32>,
29153 conv_state: &mut CudaSlice<f32>,
29154 w: &CudaSlice<f32>,
29155 y: &mut CudaSlice<f32>,
29156 conv_dim: usize,
29157 t: usize,
29158 d_conv: usize,
29159 pad_len: Option<&CudaSlice<i32>>,
29160 ) -> Result<(), Box<dyn std::error::Error>> {
29161 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
29162 let ring_old = if t < d_conv - 1 {
29166 Some(self.clone_dtod(conv_state)?)
29167 } else {
29168 None
29169 };
29170 {
29171 let f = self.func("ssm_conv1d_tm_state_f32");
29172 let cfg = LaunchConfig {
29173 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
29174 block_dim: (256, 1, 1),
29175 shared_mem_bytes: 0,
29176 };
29177 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29178 let __s_b = self.gpu.stream();
29179 let mut b = __s_b.launch_builder(&f);
29180 b.arg(qkv_tm)
29181 .arg(&*conv_state)
29182 .arg(w)
29183 .arg(y)
29184 .arg(&cd)
29185 .arg(&ti)
29186 .arg(&dc);
29187 unsafe {
29188 b.launch(cfg)?;
29189 }
29190 }
29191 match (ring_old, pad_len) {
29192 (None, Some(len_d)) => {
29193 let f = self.func("ssm_conv_ring_update_dev_f32");
29194 let n = conv_dim * (d_conv - 1);
29195 let cfg = LaunchConfig::for_num_elems(n as u32);
29196 let (cd, dc) = (conv_dim as i32, d_conv as i32);
29197 let __s_b = self.gpu.stream();
29198 let mut b = __s_b.launch_builder(&f);
29199 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
29200 unsafe {
29201 b.launch(cfg)?;
29202 }
29203 }
29204 (None, None) => {
29205 let f = self.func("ssm_conv_ring_update_f32");
29206 let n = conv_dim * (d_conv - 1);
29207 let cfg = LaunchConfig::for_num_elems(n as u32);
29208 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29209 let __s_b = self.gpu.stream();
29210 let mut b = __s_b.launch_builder(&f);
29211 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
29212 unsafe {
29213 b.launch(cfg)?;
29214 }
29215 }
29216 (Some(old), _) => {
29217 self.ssm_conv_ring_rebuild(qkv_tm, &old, conv_state, conv_dim, t, d_conv)?
29218 }
29219 }
29220 Ok(())
29221 }
29222
29223 #[allow(clippy::too_many_arguments)]
29225 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_tm_state_pad_v(
29228 &self,
29229 qkv_tm: &cudarc::driver::CudaView<f32>,
29230 conv_state: &mut CudaSlice<f32>,
29231 w: &CudaSlice<f32>,
29232 y: &mut CudaSlice<f32>,
29233 conv_dim: usize,
29234 t: usize,
29235 d_conv: usize,
29236 pad_len: Option<&CudaSlice<i32>>,
29237 ) -> Result<(), Box<dyn std::error::Error>> {
29238 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
29239 let ring_old = if t < d_conv - 1 {
29243 Some(self.clone_dtod(conv_state)?)
29244 } else {
29245 None
29246 };
29247 {
29248 let f = self.func("ssm_conv1d_tm_state_f32");
29249 let cfg = LaunchConfig {
29250 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
29251 block_dim: (256, 1, 1),
29252 shared_mem_bytes: 0,
29253 };
29254 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29255 let __s_b = self.gpu.stream();
29256 let mut b = __s_b.launch_builder(&f);
29257 b.arg(qkv_tm)
29258 .arg(&*conv_state)
29259 .arg(w)
29260 .arg(y)
29261 .arg(&cd)
29262 .arg(&ti)
29263 .arg(&dc);
29264 unsafe {
29265 b.launch(cfg)?;
29266 }
29267 }
29268 match (ring_old, pad_len) {
29269 (None, Some(len_d)) => {
29270 let f = self.func("ssm_conv_ring_update_dev_f32");
29271 let n = conv_dim * (d_conv - 1);
29272 let cfg = LaunchConfig::for_num_elems(n as u32);
29273 let (cd, dc) = (conv_dim as i32, d_conv as i32);
29274 let __s_b = self.gpu.stream();
29275 let mut b = __s_b.launch_builder(&f);
29276 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
29277 unsafe {
29278 b.launch(cfg)?;
29279 }
29280 }
29281 (None, None) => {
29282 let f = self.func("ssm_conv_ring_update_f32");
29283 let n = conv_dim * (d_conv - 1);
29284 let cfg = LaunchConfig::for_num_elems(n as u32);
29285 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29286 let __s_b = self.gpu.stream();
29287 let mut b = __s_b.launch_builder(&f);
29288 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
29289 unsafe {
29290 b.launch(cfg)?;
29291 }
29292 }
29293 (Some(_), _) => unreachable!(
29294 "ssm_conv1d_tm_state_pad_v: T < d_conv-1 has no view path (PRIME_MIN_T gates it)"
29295 ),
29296 }
29297 Ok(())
29298 }
29299
29300 pub fn ssm_conv_ring_rebuild(
29305 &self,
29306 qkv_tm: &CudaSlice<f32>,
29307 ring_old: &CudaSlice<f32>,
29308 conv_state: &mut CudaSlice<f32>,
29309 conv_dim: usize,
29310 tc: usize,
29311 d_conv: usize,
29312 ) -> Result<(), Box<dyn std::error::Error>> {
29313 let f = self.func("ssm_conv_ring_rebuild_f32");
29314 let n = conv_dim * (d_conv - 1);
29315 let cfg = LaunchConfig::for_num_elems(n as u32);
29316 let (cd, ti, dc) = (conv_dim as i32, tc as i32, d_conv as i32);
29317 let __s_b = self.gpu.stream();
29318 let mut b = __s_b.launch_builder(&f);
29319 b.arg(qkv_tm)
29320 .arg(ring_old)
29321 .arg(conv_state)
29322 .arg(&cd)
29323 .arg(&ti)
29324 .arg(&dc);
29325 unsafe {
29326 b.launch(cfg)?;
29327 }
29328 Ok(())
29329 }
29330
29331 #[allow(clippy::too_many_arguments)]
29336 pub fn gdn_prep_decode(
29337 &self,
29338 conv_out: &CudaSlice<f32>,
29339 beta_raw: &CudaSlice<f32>,
29340 alpha: &CudaSlice<f32>,
29341 dt_bias: &CudaSlice<f32>,
29342 a: &CudaSlice<f32>,
29343 q_l2: &mut CudaSlice<f32>,
29344 k_l2: &mut CudaSlice<f32>,
29345 v_g: &mut CudaSlice<f32>,
29346 beta: &mut CudaSlice<f32>,
29347 g_log: &mut CudaSlice<f32>,
29348 d_state: usize,
29349 num_v: usize,
29350 num_k: usize,
29351 key_dim: usize,
29352 eps: f32,
29353 ) -> Result<(), Box<dyn std::error::Error>> {
29354 let f = self.func("gdn_prep_decode_f32");
29355 let cfg = LaunchConfig {
29356 grid_dim: (num_v as u32, 1, 1),
29357 block_dim: (32, 4, 1),
29358 shared_mem_bytes: 0,
29359 };
29360 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
29361 let __s_b = self.gpu.stream();
29362 let mut b = __s_b.launch_builder(&f);
29363 b.arg(conv_out)
29364 .arg(beta_raw)
29365 .arg(alpha)
29366 .arg(dt_bias)
29367 .arg(a)
29368 .arg(q_l2)
29369 .arg(k_l2)
29370 .arg(v_g)
29371 .arg(beta)
29372 .arg(g_log)
29373 .arg(&ds)
29374 .arg(&nv)
29375 .arg(&nk)
29376 .arg(&kd)
29377 .arg(&eps);
29378 unsafe {
29379 b.launch(cfg)?;
29380 }
29381 Ok(())
29382 }
29383
29384 #[allow(clippy::too_many_arguments)]
29388 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_gdn(
29390 &self,
29391 qkv_tm: &CudaSlice<f32>,
29392 w: &CudaSlice<f32>,
29393 q_g: &mut CudaSlice<f32>,
29394 k_g: &mut CudaSlice<f32>,
29395 v_g: &mut CudaSlice<f32>,
29396 conv_dim: usize,
29397 t: usize,
29398 d_conv: usize,
29399 d_state: usize,
29400 num_v: usize,
29401 num_k: usize,
29402 key_dim: usize,
29403 ) -> Result<(), Box<dyn std::error::Error>> {
29404 let f = self.func("ssm_conv1d_gdn_f32");
29405 let cfg = LaunchConfig {
29406 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
29407 block_dim: (256, 1, 1),
29408 shared_mem_bytes: 0,
29409 };
29410 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29411 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
29412 let __s_b = self.gpu.stream();
29413 let mut b = __s_b.launch_builder(&f);
29414 b.arg(qkv_tm)
29415 .arg(w)
29416 .arg(q_g)
29417 .arg(k_g)
29418 .arg(v_g)
29419 .arg(&cd)
29420 .arg(&ti)
29421 .arg(&dc)
29422 .arg(&ds)
29423 .arg(&nv)
29424 .arg(&nk)
29425 .arg(&kd);
29426 unsafe {
29427 b.launch(cfg)?;
29428 }
29429 Ok(())
29430 }
29431
29432 #[allow(clippy::too_many_arguments)]
29433 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d(
29436 &self,
29437 x: &CudaSlice<f32>,
29438 w: &CudaSlice<f32>,
29439 y: &mut CudaSlice<f32>,
29440 conv_dim: usize,
29441 t: usize,
29442 d_conv: usize,
29443 silu: bool,
29444 ) -> Result<(), Box<dyn std::error::Error>> {
29445 let f = self.func("ssm_conv1d_silu_f32");
29446 let cfg = LaunchConfig {
29447 grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
29448 block_dim: (256, 1, 1),
29449 shared_mem_bytes: 0,
29450 };
29451 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
29452 let __s_b = self.gpu.stream();
29453 let mut b = __s_b.launch_builder(&f);
29454 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
29455 unsafe {
29456 b.launch(cfg)?;
29457 }
29458 Ok(())
29459 }
29460
29461 #[allow(clippy::too_many_arguments)] pub fn gdn_scan_s128(
29465 &self,
29466 q: &CudaSlice<f32>,
29467 k: &CudaSlice<f32>,
29468 v: &CudaSlice<f32>,
29469 g: &CudaSlice<f32>,
29470 beta: &CudaSlice<f32>,
29471 state_in: &CudaSlice<f32>,
29472 state_out: &mut CudaSlice<f32>,
29473 o: &mut CudaSlice<f32>,
29474 n_head: usize,
29475 t: usize,
29476 scale: f32,
29477 ) -> Result<(), Box<dyn std::error::Error>> {
29478 let f = self.func("gdn_scan_s128");
29479 const S_V: u32 = 128;
29480 const WARP: u32 = 32;
29481 const COLS_PER_BLOCK: u32 = 4;
29482 let cfg = LaunchConfig {
29483 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
29484 block_dim: (WARP, COLS_PER_BLOCK, 1),
29485 shared_mem_bytes: 0,
29486 };
29487 let (h, ti) = (n_head as i32, t as i32);
29488 let __s_b = self.gpu.stream();
29489 let mut b = __s_b.launch_builder(&f);
29490 b.arg(q)
29491 .arg(k)
29492 .arg(v)
29493 .arg(g)
29494 .arg(beta)
29495 .arg(state_in)
29496 .arg(state_out)
29497 .arg(o)
29498 .arg(&h)
29499 .arg(&ti)
29500 .arg(&scale);
29501 unsafe {
29502 b.launch(cfg)?;
29503 }
29504 Ok(())
29505 }
29506
29507 #[allow(clippy::too_many_arguments)]
29512 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_fused_decode_b(
29514 &self,
29515 qkv_cols: &CudaSlice<f32>,
29516 conv_state_ptrs: &cudarc::driver::CudaView<u64>,
29517 w: &CudaSlice<f32>,
29518 conv_outs: &mut CudaSlice<f32>,
29519 conv_dim: usize,
29520 d_conv: usize,
29521 b_n: usize,
29522 ) -> Result<(), Box<dyn std::error::Error>> {
29523 let f = self.func("ssm_conv1d_fused_decode_b_f32");
29524 let cfg = LaunchConfig {
29525 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
29526 block_dim: (256, 1, 1),
29527 shared_mem_bytes: 0,
29528 };
29529 let (cd, dc) = (conv_dim as i32, d_conv as i32);
29530 let __s_b = self.gpu.stream();
29531 let mut b = __s_b.launch_builder(&f);
29532 b.arg(qkv_cols)
29533 .arg(conv_state_ptrs)
29534 .arg(w)
29535 .arg(conv_outs)
29536 .arg(&cd)
29537 .arg(&dc);
29538 unsafe {
29539 b.launch(cfg)?;
29540 }
29541 Ok(())
29542 }
29543
29544 #[allow(clippy::too_many_arguments)]
29545 pub fn gdn_prep_decode_b(
29546 &self,
29547 conv_outs: &CudaSlice<f32>,
29548 beta_raws: &CudaSlice<f32>,
29549 alphas: &CudaSlice<f32>,
29550 dt_bias: &CudaSlice<f32>,
29551 a: &CudaSlice<f32>,
29552 q_l2: &mut CudaSlice<f32>,
29553 k_l2: &mut CudaSlice<f32>,
29554 v_g: &mut CudaSlice<f32>,
29555 beta: &mut CudaSlice<f32>,
29556 g_log: &mut CudaSlice<f32>,
29557 d_state: usize,
29558 num_v: usize,
29559 num_k: usize,
29560 key_dim: usize,
29561 eps: f32,
29562 conv_dim: usize,
29563 b_n: usize,
29564 ) -> Result<(), Box<dyn std::error::Error>> {
29565 let f = self.func("gdn_prep_decode_b_f32");
29566 let cfg = LaunchConfig {
29567 grid_dim: (num_v as u32, 1, b_n as u32),
29568 block_dim: (32, 4, 1),
29569 shared_mem_bytes: 0,
29570 };
29571 let (ds, nv, nk, kd, cd) = (
29572 d_state as i32,
29573 num_v as i32,
29574 num_k as i32,
29575 key_dim as i32,
29576 conv_dim as i32,
29577 );
29578 let __s_b = self.gpu.stream();
29579 let mut b = __s_b.launch_builder(&f);
29580 b.arg(conv_outs)
29581 .arg(beta_raws)
29582 .arg(alphas)
29583 .arg(dt_bias)
29584 .arg(a)
29585 .arg(q_l2)
29586 .arg(k_l2)
29587 .arg(v_g)
29588 .arg(beta)
29589 .arg(g_log)
29590 .arg(&ds)
29591 .arg(&nv)
29592 .arg(&nk)
29593 .arg(&kd)
29594 .arg(&eps)
29595 .arg(&cd);
29596 unsafe {
29597 b.launch(cfg)?;
29598 }
29599 Ok(())
29600 }
29601
29602 #[allow(clippy::too_many_arguments)]
29603 pub fn gdn_scan_s128_batched(
29604 &self,
29605 q: &CudaSlice<f32>,
29606 k: &CudaSlice<f32>,
29607 v: &CudaSlice<f32>,
29608 g: &CudaSlice<f32>,
29609 beta: &CudaSlice<f32>,
29610 state_in_ptrs: &cudarc::driver::CudaView<u64>,
29611 state_out_ptrs: &cudarc::driver::CudaView<u64>,
29612 o: &mut CudaSlice<f32>,
29613 n_head: usize,
29614 b_n: usize,
29615 scale: f32,
29616 ) -> Result<(), Box<dyn std::error::Error>> {
29617 let f = self.func("gdn_scan_s128_b");
29618 const S_V: u32 = 128;
29619 const WARP: u32 = 32;
29620 const COLS_PER_BLOCK: u32 = 4;
29621 let cfg = LaunchConfig {
29622 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
29623 block_dim: (WARP, COLS_PER_BLOCK, 1),
29624 shared_mem_bytes: 0,
29625 };
29626 let h = n_head as i32;
29627 let __s_b = self.gpu.stream();
29628 let mut b = __s_b.launch_builder(&f);
29629 b.arg(q)
29630 .arg(k)
29631 .arg(v)
29632 .arg(g)
29633 .arg(beta)
29634 .arg(state_in_ptrs)
29635 .arg(state_out_ptrs)
29636 .arg(o)
29637 .arg(&h)
29638 .arg(&scale);
29639 unsafe {
29640 b.launch(cfg)?;
29641 }
29642 Ok(())
29643 }
29644
29645 #[allow(clippy::too_many_arguments)]
29651 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_fused_decode_b_view(
29653 &self,
29654 qkv_cols: &cudarc::driver::CudaView<f32>,
29655 conv_state_ptrs: &cudarc::driver::CudaView<u64>,
29656 w: &CudaSlice<f32>,
29657 conv_outs: &mut CudaSlice<f32>,
29658 conv_dim: usize,
29659 d_conv: usize,
29660 b_n: usize,
29661 ) -> Result<(), Box<dyn std::error::Error>> {
29662 let f = self.func("ssm_conv1d_fused_decode_b_f32");
29663 let cfg = LaunchConfig {
29664 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
29665 block_dim: (256, 1, 1),
29666 shared_mem_bytes: 0,
29667 };
29668 let (cd, dc) = (conv_dim as i32, d_conv as i32);
29669 let __s_b = self.gpu.stream();
29670 let mut b = __s_b.launch_builder(&f);
29671 b.arg(qkv_cols)
29672 .arg(conv_state_ptrs)
29673 .arg(w)
29674 .arg(conv_outs)
29675 .arg(&cd)
29676 .arg(&dc);
29677 unsafe {
29678 b.launch(cfg)?;
29679 }
29680 Ok(())
29681 }
29682
29683 #[allow(clippy::too_many_arguments)]
29684 pub fn gdn_prep_decode_b_view(
29685 &self,
29686 conv_outs: &CudaSlice<f32>,
29687 beta_raws: &cudarc::driver::CudaView<f32>,
29688 alphas: &cudarc::driver::CudaView<f32>,
29689 dt_bias: &CudaSlice<f32>,
29690 a: &CudaSlice<f32>,
29691 q_l2: &mut CudaSlice<f32>,
29692 k_l2: &mut CudaSlice<f32>,
29693 v_g: &mut CudaSlice<f32>,
29694 beta: &mut CudaSlice<f32>,
29695 g_log: &mut CudaSlice<f32>,
29696 d_state: usize,
29697 num_v: usize,
29698 num_k: usize,
29699 key_dim: usize,
29700 eps: f32,
29701 conv_dim: usize,
29702 b_n: usize,
29703 ) -> Result<(), Box<dyn std::error::Error>> {
29704 let f = self.func("gdn_prep_decode_b_f32");
29705 let cfg = LaunchConfig {
29706 grid_dim: (num_v as u32, 1, b_n as u32),
29707 block_dim: (32, 4, 1),
29708 shared_mem_bytes: 0,
29709 };
29710 let (ds, nv, nk, kd, cd) = (
29711 d_state as i32,
29712 num_v as i32,
29713 num_k as i32,
29714 key_dim as i32,
29715 conv_dim as i32,
29716 );
29717 let __s_b = self.gpu.stream();
29718 let mut b = __s_b.launch_builder(&f);
29719 b.arg(conv_outs)
29720 .arg(beta_raws)
29721 .arg(alphas)
29722 .arg(dt_bias)
29723 .arg(a)
29724 .arg(q_l2)
29725 .arg(k_l2)
29726 .arg(v_g)
29727 .arg(beta)
29728 .arg(g_log)
29729 .arg(&ds)
29730 .arg(&nv)
29731 .arg(&nk)
29732 .arg(&kd)
29733 .arg(&eps)
29734 .arg(&cd);
29735 unsafe {
29736 b.launch(cfg)?;
29737 }
29738 Ok(())
29739 }
29740
29741 #[allow(clippy::too_many_arguments)]
29742 pub fn gdn_scan_s128_batched_view(
29743 &self,
29744 q: &CudaSlice<f32>,
29745 k: &CudaSlice<f32>,
29746 v: &CudaSlice<f32>,
29747 g: &CudaSlice<f32>,
29748 beta: &CudaSlice<f32>,
29749 state_in_ptrs: &cudarc::driver::CudaView<u64>,
29750 state_out_ptrs: &cudarc::driver::CudaView<u64>,
29751 o: &mut cudarc::driver::CudaViewMut<f32>,
29752 n_head: usize,
29753 b_n: usize,
29754 scale: f32,
29755 ) -> Result<(), Box<dyn std::error::Error>> {
29756 let f = self.func("gdn_scan_s128_b");
29757 const S_V: u32 = 128;
29758 const WARP: u32 = 32;
29759 const COLS_PER_BLOCK: u32 = 4;
29760 let cfg = LaunchConfig {
29761 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
29762 block_dim: (WARP, COLS_PER_BLOCK, 1),
29763 shared_mem_bytes: 0,
29764 };
29765 let h = n_head as i32;
29766 let __s_b = self.gpu.stream();
29767 let mut b = __s_b.launch_builder(&f);
29768 b.arg(q)
29769 .arg(k)
29770 .arg(v)
29771 .arg(g)
29772 .arg(beta)
29773 .arg(state_in_ptrs)
29774 .arg(state_out_ptrs)
29775 .arg(o)
29776 .arg(&h)
29777 .arg(&scale);
29778 unsafe {
29779 b.launch(cfg)?;
29780 }
29781 Ok(())
29782 }
29783
29784 pub fn gdn_chunked_enabled() -> bool {
29793 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
29794 *E.get_or_init(|| {
29795 std::env::var("MEMRA_GDN_CHUNKED")
29796 .map(|v| v != "0")
29797 .unwrap_or(true)
29798 })
29799 }
29800
29801 pub fn gdn_chunk_size() -> usize {
29806 static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
29807 *C.get_or_init(|| {
29808 let c: usize = std::env::var("MEMRA_GDN_CHUNK")
29809 .ok()
29810 .and_then(|v| v.parse().ok())
29811 .unwrap_or(32);
29812 c.clamp(32, 128) / 32 * 32
29813 })
29814 }
29815
29816 #[allow(clippy::too_many_arguments)]
29821 #[allow(clippy::too_many_arguments, clippy::type_complexity)]
29824 #[allow(clippy::too_many_arguments)]
29825 pub fn gdn_chunk_k123(
29826 &self,
29827 q: &CudaSlice<f32>,
29828 k: &CudaSlice<f32>,
29829 v: &CudaSlice<f32>,
29830 g: &CudaSlice<f32>,
29831 beta: &CudaSlice<f32>,
29832 wb16: Option<&mut CudaSlice<u8>>,
29833 n_head: usize,
29834 t: usize,
29835 c: usize,
29836 hk: usize,
29837 k2w: Option<(&CudaSlice<u8>, &CudaSlice<u8>, &mut CudaSlice<u8>)>,
29838 ) -> Result<
29839 (
29840 CudaSlice<f32>,
29841 CudaSlice<f32>,
29842 CudaSlice<f32>,
29843 CudaSlice<f32>,
29844 ),
29845 Box<dyn std::error::Error>,
29846 > {
29847 const D: usize = 128;
29848 let h = n_head;
29849 #[allow(clippy::manual_div_ceil)]
29850 let nc = (t + c - 1) / c;
29852 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
29853 let mut gcum = self.uninit(t * h)?;
29854 let mut a = self.uninit(nc * h * c * c)?;
29855 let mut p = self.uninit(nc * h * c * c)?;
29856 let mut u = self.uninit(nc * h * c * D)?;
29857 let mut w = self.uninit(nc * h * c * D)?;
29858 {
29859 let f = self.func("gdn_chunk_cumgate_f32");
29861 let cfg = LaunchConfig {
29862 grid_dim: (nc as u32, h as u32, 1),
29863 block_dim: (32, 1, 1),
29864 shared_mem_bytes: 0,
29865 };
29866 let __s_b = self.gpu.stream();
29867 let mut b = __s_b.launch_builder(&f);
29868 b.arg(g).arg(&mut gcum).arg(&hi).arg(&ti).arg(&ci);
29869 unsafe {
29870 b.launch(cfg)?;
29871 }
29872 }
29873 if let Some((qb, kb, pb)) = k2w {
29874 assert!(c == 32, "gdn_k2_wgmma is a C==32 tile");
29877 let f = self.func("gdn_k2_wgmma");
29878 let cfg = LaunchConfig {
29879 grid_dim: (nc as u32, h as u32, 1),
29880 block_dim: (128, 1, 1),
29881 shared_mem_bytes: 0,
29882 };
29883 let hki = hk as i32;
29884 let __s_b = self.gpu.stream();
29885 let mut b = __s_b.launch_builder(&f);
29886 b.arg(qb)
29887 .arg(kb)
29888 .arg(&gcum)
29889 .arg(beta)
29890 .arg(&mut a)
29891 .arg(&mut *pb)
29892 .arg(&hi)
29893 .arg(&ti)
29894 .arg(&ci)
29895 .arg(&hki);
29896 unsafe {
29897 b.launch(cfg)?;
29898 }
29899 } else if c <= 64 && !portable_mma_gated() {
29900 let f = self.func("gdn_chunk_attn_f32");
29902 f.set_attribute(
29903 CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
29904 GDN_K2_DYNAMIC_SHARED_BYTES as i32,
29905 )?;
29906 #[allow(clippy::manual_div_ceil)]
29907 let jt = ((c + 31) / 32) as u32;
29909 let cfg = LaunchConfig {
29910 grid_dim: (nc as u32, h as u32, jt),
29911 block_dim: (256, 1, 1),
29912 shared_mem_bytes: GDN_K2_DYNAMIC_SHARED_BYTES,
29913 };
29914 let hki = hk as i32;
29915 let __s_b = self.gpu.stream();
29916 let mut b = __s_b.launch_builder(&f);
29917 b.arg(q)
29918 .arg(k)
29919 .arg(&gcum)
29920 .arg(beta)
29921 .arg(&mut a)
29922 .arg(&mut p)
29923 .arg(&hi)
29924 .arg(&ti)
29925 .arg(&ci)
29926 .arg(&hki);
29927 unsafe {
29928 b.launch(cfg)?;
29929 }
29930 } else {
29931 assert!(
29933 hk == h,
29934 "generic K2 is broadcast-only (de-broadcast rides C==32)"
29935 );
29936 let f = self.func("gdn_chunk_attn_g_f32");
29937 let cfg = LaunchConfig {
29938 grid_dim: (nc as u32, h as u32, 1),
29939 block_dim: (32, 8, 1),
29940 shared_mem_bytes: 0,
29941 };
29942 let __s_b = self.gpu.stream();
29943 let mut b = __s_b.launch_builder(&f);
29944 b.arg(q)
29945 .arg(k)
29946 .arg(&gcum)
29947 .arg(beta)
29948 .arg(&mut a)
29949 .arg(&mut p)
29950 .arg(&hi)
29951 .arg(&ti)
29952 .arg(&ci);
29953 unsafe {
29954 b.launch(cfg)?;
29955 }
29956 }
29957 {
29958 let cfg = LaunchConfig {
29960 grid_dim: (nc as u32, h as u32, 1),
29961 block_dim: (256, 1, 1),
29962 shared_mem_bytes: 0,
29963 };
29964 match c {
29965 32 | 64 => {
29966 let f = self.func(if c == 32 {
29967 "gdn_chunk_solve32_f32"
29968 } else {
29969 "gdn_chunk_solve64_f32"
29970 });
29971 let wb: u64 = match wb16 {
29973 Some(d) => self.addr_u8(d),
29974 None => 0,
29975 };
29976 let hki = hk as i32;
29977 let __s_b = self.gpu.stream();
29978 let mut b = __s_b.launch_builder(&f);
29979 b.arg(v)
29980 .arg(k)
29981 .arg(&a)
29982 .arg(&gcum)
29983 .arg(&mut u)
29984 .arg(&mut w)
29985 .arg(&wb)
29986 .arg(&hi)
29987 .arg(&ti)
29988 .arg(&hki);
29989 unsafe {
29990 b.launch(cfg)?;
29991 }
29992 }
29993 _ => {
29994 assert!(hk == h, "generic K3 is broadcast-only");
29995 let f = self.func("gdn_chunk_solve_f32");
29996 let __s_b = self.gpu.stream();
29997 let mut b = __s_b.launch_builder(&f);
29998 b.arg(v)
29999 .arg(k)
30000 .arg(&a)
30001 .arg(&gcum)
30002 .arg(&mut u)
30003 .arg(&mut w)
30004 .arg(&hi)
30005 .arg(&ti)
30006 .arg(&ci);
30007 unsafe {
30008 b.launch(cfg)?;
30009 }
30010 }
30011 }
30012 }
30013 Ok((gcum, p, u, w))
30014 }
30015
30016 pub fn gdn_db_on() -> bool {
30020 std::env::var("MEMRA_GDN_DB").as_deref() != Ok("0")
30021 }
30022
30023 pub fn gdn_mma_enabled(&self, c: usize) -> bool {
30032 !portable_mma_gated()
30033 && c == 32
30034 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
30035 Ok("1") => true,
30036 Ok("0") => false,
30037 _ => gdn_mma_default_on(),
30038 }
30039 }
30040
30041 pub fn gdn_wgmma_on(&self, c: usize) -> bool {
30048 cfg!(memra_hopper_mma)
30049 && self.gdn_mma_enabled(c)
30050 && std::env::var("MEMRA_GDN_WGMMA").as_deref() != Ok("0")
30051 }
30052
30053 #[allow(clippy::too_many_arguments)]
30058 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_gdn_state_pad(
30060 &self,
30061 qkv_tm: &cudarc::driver::CudaView<f32>,
30062 conv_state: &mut CudaSlice<f32>,
30063 w: &CudaSlice<f32>,
30064 q_g: &mut CudaSlice<f32>,
30065 k_g: &mut CudaSlice<f32>,
30066 v_g: &mut CudaSlice<f32>,
30067 conv_dim: usize,
30068 t: usize,
30069 d_conv: usize,
30070 d_state: usize,
30071 num_v: usize,
30072 num_k: usize,
30073 key_dim: usize,
30074 hk: usize,
30075 pad_len: Option<&CudaSlice<i32>>,
30076 ) -> Result<(), Box<dyn std::error::Error>> {
30077 assert!(
30078 t >= d_conv - 1,
30079 "fused state conv requires T >= pad (PRIME_MIN_T gates)"
30080 );
30081 {
30082 let f = self.func("ssm_conv1d_gdn_state_f32");
30083 let cfg = LaunchConfig {
30084 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
30085 block_dim: (256, 1, 1),
30086 shared_mem_bytes: 0,
30087 };
30088 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
30089 let (ds, nv, nk, kd, hki) = (
30090 d_state as i32,
30091 num_v as i32,
30092 num_k as i32,
30093 key_dim as i32,
30094 hk as i32,
30095 );
30096 let __s_b = self.gpu.stream();
30097 let mut b = __s_b.launch_builder(&f);
30098 b.arg(qkv_tm)
30099 .arg(&*conv_state)
30100 .arg(w)
30101 .arg(q_g)
30102 .arg(k_g)
30103 .arg(v_g)
30104 .arg(&cd)
30105 .arg(&ti)
30106 .arg(&dc)
30107 .arg(&ds)
30108 .arg(&nv)
30109 .arg(&nk)
30110 .arg(&kd)
30111 .arg(&hki);
30112 unsafe {
30113 b.launch(cfg)?;
30114 }
30115 }
30116 match pad_len {
30117 Some(len_d) => {
30118 let f = self.func("ssm_conv_ring_update_dev_f32");
30119 let n = conv_dim * (d_conv - 1);
30120 let cfg = LaunchConfig::for_num_elems(n as u32);
30121 let (cd, dc) = (conv_dim as i32, d_conv as i32);
30122 let __s_b = self.gpu.stream();
30123 let mut b = __s_b.launch_builder(&f);
30124 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
30125 unsafe {
30126 b.launch(cfg)?;
30127 }
30128 }
30129 None => {
30130 let f = self.func("ssm_conv_ring_update_f32");
30131 let n = conv_dim * (d_conv - 1);
30132 let cfg = LaunchConfig::for_num_elems(n as u32);
30133 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
30134 let __s_b = self.gpu.stream();
30135 let mut b = __s_b.launch_builder(&f);
30136 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
30137 unsafe {
30138 b.launch(cfg)?;
30139 }
30140 }
30141 }
30142 Ok(())
30143 }
30144
30145 pub fn gdn_chunk_alloc(
30149 &self,
30150 n_head: usize,
30151 t: usize,
30152 c: usize,
30153 hk: usize,
30154 ) -> Result<GdnChunkBufs, Box<dyn std::error::Error>> {
30155 const D: usize = 128;
30156 assert!(
30157 c == 32,
30158 "gdn_chunk_alloc: varlen chain is the C==32 mma pair"
30159 );
30160 let h = n_head;
30161 #[allow(clippy::manual_div_ceil)]
30162 let nc = (t + c - 1) / c;
30164 Ok(GdnChunkBufs {
30165 gcum: self.uninit(t * h)?,
30166 a: self.uninit(nc * h * c * c)?,
30167 p: self.uninit(nc * h * c * c)?,
30168 u: self.uninit(nc * h * c * D)?,
30169 w: self.uninit(nc * h * c * D)?,
30170 kb16: self.alloc_u8_uninit(t * hk * D * 2)?,
30171 wb16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
30172 y16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
30173 ssnap16: self.alloc_u8_uninit(nc * h * D * D * 2)?,
30174 qb16: self.alloc_u8_uninit(t * hk * D * 2)?,
30175 pb16: self.alloc_u8_uninit(nc * h * c * c * 2)?,
30176 o: self.uninit(D * h * t)?,
30177 t,
30178 nc,
30179 })
30180 }
30181
30182 pub fn f32_to_bf16_v(
30184 &self,
30185 x: &cudarc::driver::CudaView<f32>,
30186 dst: &mut CudaSlice<u8>,
30187 n: usize,
30188 ) -> Result<(), Box<dyn std::error::Error>> {
30189 let f = self.func("f32_to_bf16_bulk");
30190 let ni = n as i64;
30191 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
30192 let __s_b = self.gpu.stream();
30193 let mut b = __s_b.launch_builder(&f);
30194 b.arg(x).arg(dst).arg(&ni);
30195 unsafe {
30196 b.launch(cfg)?;
30197 }
30198 Ok(())
30199 }
30200
30201 pub fn f32_to_bf16_into(
30203 &self,
30204 x: &CudaSlice<f32>,
30205 dst: &mut CudaSlice<u8>,
30206 n: usize,
30207 ) -> Result<(), Box<dyn std::error::Error>> {
30208 let f = self.func("f32_to_bf16_bulk");
30209 let ni = n as i64;
30210 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
30211 let __s_b = self.gpu.stream();
30212 let mut b = __s_b.launch_builder(&f);
30213 b.arg(x).arg(dst).arg(&ni);
30214 unsafe {
30215 b.launch(cfg)?;
30216 }
30217 Ok(())
30218 }
30219
30220 pub fn gdn_chunk_k123_vl8(
30223 &self,
30224 seqs: &[GdnSeqVl],
30225 n_head: usize,
30226 hk: usize,
30227 wq: Option<&GdnWVl8>,
30228 ) -> Result<(), Box<dyn std::error::Error>> {
30229 let b = seqs.len();
30230 assert!((1..=8).contains(&b), "gdn_chunk_k123_vl8: 1..=8 sequences");
30231 let mut packed = [GdnSeqVl::default(); 8];
30232 packed[..b].copy_from_slice(seqs);
30233 let v = GdnVl8(packed);
30234 let (hi, ci) = (n_head as i32, 32i32);
30235 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
30236 {
30237 let f = self.func("gdn_chunk_cumgate_vl");
30238 let cfg = LaunchConfig {
30239 grid_dim: (max_nc, n_head as u32, b as u32),
30240 block_dim: (32, 1, 1),
30241 shared_mem_bytes: 0,
30242 };
30243 let __s_lb = self.gpu.stream();
30244 let mut lb = __s_lb.launch_builder(&f);
30245 lb.arg(&v).arg(&hi).arg(&ci);
30246 unsafe {
30247 lb.launch(cfg)?;
30248 }
30249 }
30250 let hki = hk as i32;
30251 if let Some(w) = wq {
30252 let f = self.func("gdn_k2_wgmma_vl");
30254 let cfg = LaunchConfig {
30255 grid_dim: (max_nc, n_head as u32, b as u32),
30256 block_dim: (128, 1, 1),
30257 shared_mem_bytes: 0,
30258 };
30259 let __s_lb = self.gpu.stream();
30260 let mut lb = __s_lb.launch_builder(&f);
30261 lb.arg(&v).arg(w).arg(&hi).arg(&ci).arg(&hki);
30262 unsafe {
30263 lb.launch(cfg)?;
30264 }
30265 } else {
30266 let f = self.func("gdn_chunk_attn_vl");
30267 f.set_attribute(
30268 CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
30269 GDN_K2_DYNAMIC_SHARED_BYTES as i32,
30270 )?;
30271 let cfg = LaunchConfig {
30272 grid_dim: (max_nc, n_head as u32, b as u32),
30273 block_dim: (256, 1, 1),
30274 shared_mem_bytes: GDN_K2_DYNAMIC_SHARED_BYTES,
30275 };
30276 let __s_lb = self.gpu.stream();
30277 let mut lb = __s_lb.launch_builder(&f);
30278 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
30279 unsafe {
30280 lb.launch(cfg)?;
30281 }
30282 }
30283 {
30284 let f = self.func("gdn_chunk_solve32_vl");
30285 let cfg = LaunchConfig {
30286 grid_dim: (max_nc, n_head as u32, b as u32),
30287 block_dim: (256, 1, 1),
30288 shared_mem_bytes: 0,
30289 };
30290 let __s_lb = self.gpu.stream();
30291 let mut lb = __s_lb.launch_builder(&f);
30292 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
30293 unsafe {
30294 lb.launch(cfg)?;
30295 }
30296 }
30297 Ok(())
30298 }
30299
30300 #[allow(clippy::too_many_arguments)]
30304 pub fn gdn_prep_vl8(
30305 &self,
30306 seqs: &[GdnPrepVl],
30307 conv_w: &CudaSlice<f32>,
30308 dt_bias: &CudaSlice<f32>,
30309 a: &CudaSlice<f32>,
30310 conv_dim: usize,
30311 d_conv: usize,
30312 d_state: usize,
30313 num_v: usize,
30314 num_k: usize,
30315 key_dim: usize,
30316 hk: usize,
30317 eps: f32,
30318 ) -> Result<(), Box<dyn std::error::Error>> {
30319 let b = seqs.len();
30320 assert!((1..=8).contains(&b));
30321 let mut packed = [GdnPrepVl::default(); 8];
30322 packed[..b].copy_from_slice(seqs);
30323 let v = GdnPrepVl8(packed);
30324 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
30325 let (cdi, dci) = (conv_dim as i32, d_conv as i32);
30326 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
30327 assert!(
30328 conv_fuse || hk == num_v,
30329 "de-broadcast requires the fused conv"
30330 );
30331 if conv_fuse {
30332 let f = self.func("ssm_conv1d_gdn_state_vl");
30333 let cfg = LaunchConfig {
30334 grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32),
30335 block_dim: (256, 1, 1),
30336 shared_mem_bytes: 0,
30337 };
30338 let (dsi, nvi, nki, kdi, hki) = (
30339 d_state as i32,
30340 num_v as i32,
30341 num_k as i32,
30342 key_dim as i32,
30343 hk as i32,
30344 );
30345 let __s_lb = self.gpu.stream();
30346 let mut lb = __s_lb.launch_builder(&f);
30347 lb.arg(&v)
30348 .arg(conv_w)
30349 .arg(&cdi)
30350 .arg(&dci)
30351 .arg(&dsi)
30352 .arg(&nvi)
30353 .arg(&nki)
30354 .arg(&kdi)
30355 .arg(&hki);
30356 unsafe {
30357 lb.launch(cfg)?;
30358 }
30359 } else {
30360 let f = self.func("ssm_conv1d_tm_state_vl");
30361 let cfg = LaunchConfig {
30362 grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32),
30363 block_dim: (256, 1, 1),
30364 shared_mem_bytes: 0,
30365 };
30366 let __s_lb = self.gpu.stream();
30367 let mut lb = __s_lb.launch_builder(&f);
30368 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci);
30369 unsafe {
30370 lb.launch(cfg)?;
30371 }
30372 }
30373 {
30374 let f = self.func("ssm_conv_ring_update_vl");
30375 let n = (conv_dim * (d_conv - 1)) as u32;
30376 let cfg = LaunchConfig {
30377 grid_dim: (n.div_ceil(256), 1, b as u32),
30378 block_dim: (256, 1, 1),
30379 shared_mem_bytes: 0,
30380 };
30381 let __s_lb = self.gpu.stream();
30382 let mut lb = __s_lb.launch_builder(&f);
30383 lb.arg(&v).arg(&cdi).arg(&dci);
30384 unsafe {
30385 lb.launch(cfg)?;
30386 }
30387 }
30388 if !conv_fuse {
30389 let f = self.func("qkv_to_gdn_repack_vl");
30390 let n = max_t * (num_v * d_state) as u32;
30391 let cfg = LaunchConfig {
30392 grid_dim: (n.div_ceil(256), 1, b as u32),
30393 block_dim: (256, 1, 1),
30394 shared_mem_bytes: 0,
30395 };
30396 let (dsi, nvi, nki, kdi) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
30397 let __s_lb = self.gpu.stream();
30398 let mut lb = __s_lb.launch_builder(&f);
30399 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi);
30400 unsafe {
30401 lb.launch(cfg)?;
30402 }
30403 }
30404 if Self::l2_v2_on(d_state) {
30405 let f = self.func("gdn_l2_v2_vl");
30406 let cfg = LaunchConfig {
30407 grid_dim: ((max_t * hk as u32).div_ceil(8), 2, b as u32),
30408 block_dim: (256, 1, 1),
30409 shared_mem_bytes: 0,
30410 };
30411 let (dsi, nvi) = (d_state as i32, hk as i32);
30412 let __s_lb = self.gpu.stream();
30413 let mut lb = __s_lb.launch_builder(&f);
30414 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
30415 unsafe {
30416 lb.launch(cfg)?;
30417 }
30418 } else {
30419 let f = self.func("gdn_l2_vl");
30420 let cfg = LaunchConfig {
30421 grid_dim: (max_t * hk as u32, 2, b as u32),
30422 block_dim: (256, 1, 1),
30423 shared_mem_bytes: 0,
30424 };
30425 let (dsi, nvi) = (d_state as i32, hk as i32);
30426 let __s_lb = self.gpu.stream();
30427 let mut lb = __s_lb.launch_builder(&f);
30428 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
30429 unsafe {
30430 lb.launch(cfg)?;
30431 }
30432 }
30433 {
30434 let f = self.func("gdn_gate_prep_vl");
30435 let n = max_t * num_v as u32;
30436 let cfg = LaunchConfig {
30437 grid_dim: (n.div_ceil(256), 1, b as u32),
30438 block_dim: (256, 1, 1),
30439 shared_mem_bytes: 0,
30440 };
30441 let nvi = num_v as i32;
30442 let __s_lb = self.gpu.stream();
30443 let mut lb = __s_lb.launch_builder(&f);
30444 lb.arg(&v).arg(dt_bias).arg(a).arg(&nvi);
30445 unsafe {
30446 lb.launch(cfg)?;
30447 }
30448 }
30449 Ok(())
30450 }
30451
30452 pub fn gdn_mirror_vl8(
30454 &self,
30455 seqs: &[GdnSeqVl],
30456 n_head: usize,
30457 which: i32,
30458 hk: usize,
30459 ) -> Result<(), Box<dyn std::error::Error>> {
30460 let b = seqs.len();
30461 assert!((1..=8).contains(&b));
30462 let mut packed = [GdnSeqVl::default(); 8];
30463 packed[..b].copy_from_slice(seqs);
30464 let v = GdnVl8(packed);
30465 let ept = (if which == 0 { hk } else { n_head } * 128) as i32;
30466 let max_n = seqs
30467 .iter()
30468 .map(|s| {
30469 if which == 0 {
30470 s.t as i64 * ept as i64
30471 } else {
30472 s.nc as i64 * ept as i64 * 32
30473 }
30474 })
30475 .max()
30476 .unwrap();
30477 let f = self.func("gdn_mirror_vl");
30478 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
30479 let cfg = LaunchConfig {
30480 grid_dim: (blocks, 1, b as u32),
30481 block_dim: (256, 1, 1),
30482 shared_mem_bytes: 0,
30483 };
30484 let __s_lb = self.gpu.stream();
30485 let mut lb = __s_lb.launch_builder(&f);
30486 lb.arg(&v).arg(&ept).arg(&which);
30487 unsafe {
30488 lb.launch(cfg)?;
30489 }
30490 Ok(())
30491 }
30492
30493 pub fn gdn_tail_vl8(
30495 &self,
30496 seqs: &[GdnPrepVl],
30497 norm_w: &CudaSlice<f32>,
30498 d_state: usize,
30499 num_v: usize,
30500 eps: f32,
30501 ) -> Result<(), Box<dyn std::error::Error>> {
30502 let b = seqs.len();
30503 assert!((1..=8).contains(&b));
30504 let mut packed = [GdnPrepVl::default(); 8];
30505 packed[..b].copy_from_slice(seqs);
30506 let v = GdnPrepVl8(packed);
30507 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
30508 let f = self.func("gated_rmsnorm_f16out_vl");
30509 let cfg = LaunchConfig {
30511 grid_dim: (max_t * num_v as u32, 1, b as u32),
30512 block_dim: (128, 1, 1),
30513 shared_mem_bytes: 0,
30514 };
30515 let (dsi, nvi) = (d_state as i32, num_v as i32);
30516 let __s_lb = self.gpu.stream();
30517 let mut lb = __s_lb.launch_builder(&f);
30518 lb.arg(&v).arg(norm_w).arg(&dsi).arg(&nvi).arg(&eps);
30519 unsafe {
30520 lb.launch(cfg)?;
30521 }
30522 Ok(())
30523 }
30524
30525 pub fn addr_f32(&self, x: &CudaSlice<f32>) -> u64 {
30528 use cudarc::driver::DevicePtr;
30529 let s = self.gpu.stream();
30530 let (p, _g) = x.device_ptr(&s);
30531 p
30532 }
30533 pub fn addr_f32_mut(&self, x: &mut CudaSlice<f32>) -> u64 {
30534 use cudarc::driver::DevicePtrMut;
30535 let s = self.gpu.stream();
30536 let (p, _g) = x.device_ptr_mut(&s);
30537 p
30538 }
30539 pub fn addr_f32v(&self, x: &cudarc::driver::CudaView<f32>) -> u64 {
30540 use cudarc::driver::DevicePtr;
30541 let s = self.gpu.stream();
30542 let (p, _g) = x.device_ptr(&s);
30543 p
30544 }
30545 pub fn addr_u8(&self, x: &CudaSlice<u8>) -> u64 {
30546 use cudarc::driver::DevicePtr;
30547 let s = self.gpu.stream();
30548 let (p, _g) = x.device_ptr(&s);
30549 p
30550 }
30551
30552 pub fn gdn_chunk_vl8(
30556 &self,
30557 seqs: &[GdnSeqVl],
30558 n_head: usize,
30559 scale: f32,
30560 hk: usize,
30561 wq: Option<&GdnWVl8>,
30562 ) -> Result<(), Box<dyn std::error::Error>> {
30563 const NSPLIT: u32 = 4;
30564 let b = seqs.len();
30565 assert!((1..=8).contains(&b), "gdn_chunk_vl8: 1..=8 sequences");
30566 let mut packed = [GdnSeqVl::default(); 8];
30567 packed[..b].copy_from_slice(seqs);
30568 let v = GdnVl8(packed);
30569 let (hi, ci) = (n_head as i32, 32i32);
30570 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
30571 let hki = hk as i32;
30572 if let Some(w) = wq {
30573 let f = self.func("gdn_k45_wgmma_vl");
30575 let cfg = LaunchConfig {
30576 grid_dim: (n_head as u32, NSPLIT, b as u32),
30577 block_dim: (256, 1, 1),
30578 shared_mem_bytes: 0,
30579 };
30580 let __s_lb = self.gpu.stream();
30581 let mut lb = __s_lb.launch_builder(&f);
30582 lb.arg(&v).arg(w).arg(&scale).arg(&hi).arg(&ci).arg(&hki);
30583 unsafe {
30584 lb.launch(cfg)?;
30585 }
30586 let _ = max_nc;
30587 return Ok(());
30588 }
30589 {
30590 let f = self.func("gdn_chunk_state_mma_vl");
30591 let cfg = LaunchConfig {
30592 grid_dim: (n_head as u32, NSPLIT, b as u32),
30593 block_dim: (256, 1, 1),
30594 shared_mem_bytes: 0,
30595 };
30596 let __s_lb = self.gpu.stream();
30597 let mut lb = __s_lb.launch_builder(&f);
30598 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
30599 unsafe {
30600 lb.launch(cfg)?;
30601 }
30602 }
30603 {
30604 let f = self.func("gdn_chunk_output_mma_vl");
30605 let cfg = LaunchConfig {
30606 grid_dim: (max_nc, n_head as u32, b as u32),
30607 block_dim: (256, 1, 1),
30608 shared_mem_bytes: 0,
30609 };
30610 let __s_lb = self.gpu.stream();
30611 let mut lb = __s_lb.launch_builder(&f);
30612 lb.arg(&v).arg(&hi).arg(&ci).arg(&scale).arg(&hki);
30613 unsafe {
30614 lb.launch(cfg)?;
30615 }
30616 }
30617 Ok(())
30618 }
30619 #[allow(clippy::too_many_arguments)] pub fn gdn_scan_chunked(
30621 &self,
30622 q: &CudaSlice<f32>,
30623 k: &CudaSlice<f32>,
30624 v: &CudaSlice<f32>,
30625 g: &CudaSlice<f32>,
30626 beta: &CudaSlice<f32>,
30627 kb16_pre: Option<&CudaSlice<u8>>,
30628 qb16_pre: Option<&CudaSlice<u8>>,
30629 state_in: &CudaSlice<f32>,
30630 state_out: &mut CudaSlice<f32>,
30631 o: &mut CudaSlice<f32>,
30632 n_head: usize,
30633 t: usize,
30634 scale: f32,
30635 c: usize,
30636 hk: usize,
30637 ) -> Result<(), Box<dyn std::error::Error>> {
30638 const D: usize = 128;
30639 const NSPLIT: u32 = 4;
30640 assert!(
30641 (1..=128).contains(&c),
30642 "gdn_scan_chunked: C must be in 1..=128"
30643 );
30644 let h = n_head;
30645 #[allow(clippy::manual_div_ceil)]
30646 let nc = (t + c - 1) / c;
30648 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
30649 let gdn_mma_pre = !portable_mma_gated()
30654 && c == 32
30655 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
30656 Ok("1") => true,
30657 Ok("0") => false,
30658 _ => gdn_mma_default_on(),
30659 };
30660 let mut wb16_pre: Option<CudaSlice<u8>> = if gdn_mma_pre {
30661 Some(self.alloc_u8_uninit(nc * h * c * D * 2)?)
30662 } else {
30663 None
30664 };
30665 let gdn_wgmma_pre = cfg!(memra_hopper_mma)
30670 && gdn_mma_pre
30671 && std::env::var("MEMRA_GDN_WGMMA").as_deref() != Ok("0");
30672 let nk = t * hk * D;
30673 let mut kb16_local: Option<CudaSlice<u8>> = None;
30674 if gdn_mma_pre && kb16_pre.is_none() {
30675 let mut kb = self.alloc_u8_uninit(nk * 2)?;
30676 let f = self.func("f32_to_bf16_bulk");
30677 let n2 = nk as i64;
30678 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
30679 let __s_b = self.gpu.stream();
30680 let mut b = __s_b.launch_builder(&f);
30681 b.arg(k).arg(&mut kb).arg(&n2);
30682 unsafe {
30683 b.launch(cfg2)?;
30684 }
30685 kb16_local = Some(kb);
30686 }
30687 let kb16_ref0: Option<&CudaSlice<u8>> = kb16_local.as_ref().or(kb16_pre);
30688 if let Some(kb) = kb16_pre {
30689 assert!(kb.len() >= nk * 2, "kb16_pre too small");
30690 }
30691 let mut qb16: Option<CudaSlice<u8>> = None;
30692 let mut pb16: Option<CudaSlice<u8>> = None;
30693 if gdn_wgmma_pre {
30694 if qb16_pre.is_none() {
30697 let mut qb = self.alloc_u8_uninit(nk * 2)?;
30698 let f = self.func("f32_to_bf16_bulk");
30699 let n2 = nk as i64;
30700 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
30701 let __s_b = self.gpu.stream();
30702 let mut b = __s_b.launch_builder(&f);
30703 b.arg(q).arg(&mut qb).arg(&n2);
30704 unsafe {
30705 b.launch(cfg2)?;
30706 }
30707 qb16 = Some(qb);
30708 } else if let Some(qb) = qb16_pre {
30709 assert!(qb.len() >= nk * 2, "qb16_pre too small");
30710 }
30711 pb16 = Some(self.alloc_u8_uninit(nc * h * c * c * 2)?);
30712 }
30713 let qb16_ref0: Option<&CudaSlice<u8>> = qb16.as_ref().or(qb16_pre);
30714 let k2w = if gdn_wgmma_pre {
30715 Some((
30716 *qb16_ref0.as_ref().unwrap(),
30717 *kb16_ref0.as_ref().unwrap(),
30718 pb16.as_mut().unwrap(),
30719 ))
30720 } else {
30721 None
30722 };
30723 let (gcum, p, u, w) =
30724 self.gdn_chunk_k123(q, k, v, g, beta, wb16_pre.as_mut(), n_head, t, c, hk, k2w)?;
30725 let _ = &w;
30726 let mut y = self.uninit(nc * h * c * D)?;
30727 let mut ssnap = self.uninit(nc * h * D * D)?; let gdn_mma = !portable_mma_gated()
30741 && c == 32
30742 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
30743 Ok("1") => true,
30744 Ok("0") => false,
30745 _ => gdn_mma_default_on(),
30746 };
30747 if gdn_mma {
30748 let wb16 = wb16_pre
30749 .take()
30750 .expect("mma path pre-allocates wb16 (K3 store fold)");
30751 let kb16_ref: &CudaSlice<u8> = kb16_ref0.expect("mma path pre-builds kb16 above K123");
30752 if gdn_wgmma_pre {
30764 let qb16 = qb16_ref0.unwrap();
30766 let pb16 = pb16.as_ref().unwrap();
30767 {
30768 let f = self.func("gdn_k45_wgmma");
30769 let cfg = LaunchConfig {
30770 grid_dim: (h as u32, 4, 1),
30771 block_dim: (256, 1, 1),
30772 shared_mem_bytes: 0,
30773 };
30774 let hki = hk as i32;
30775 let __s_b = self.gpu.stream();
30776 let mut b = __s_b.launch_builder(&f);
30777 b.arg(kb16_ref)
30778 .arg(&gcum)
30779 .arg(beta)
30780 .arg(&u)
30781 .arg(&wb16)
30782 .arg(qb16)
30783 .arg(pb16)
30784 .arg(o)
30785 .arg(&scale)
30786 .arg(state_in)
30787 .arg(&mut *state_out)
30788 .arg(&hi)
30789 .arg(&ti)
30790 .arg(&ci)
30791 .arg(&hki);
30792 unsafe {
30793 b.launch(cfg)?;
30794 }
30795 }
30796 return Ok(());
30797 }
30798 let mut y16 = self.alloc_u8_uninit(nc * h * c * D * 2)?;
30802 let mut ssnap16 = self.alloc_u8_uninit(nc * h * D * D * 2)?;
30803 {
30804 let f = self.func("gdn_chunk_state_mma");
30805 let cfg = LaunchConfig {
30806 grid_dim: (h as u32, NSPLIT, 1),
30807 block_dim: (256, 1, 1),
30808 shared_mem_bytes: 0,
30809 };
30810 let hki = hk as i32;
30811 let __s_b = self.gpu.stream();
30812 let mut b = __s_b.launch_builder(&f);
30813 b.arg(kb16_ref)
30814 .arg(&gcum)
30815 .arg(beta)
30816 .arg(&u)
30817 .arg(&wb16)
30818 .arg(&mut y16)
30819 .arg(&mut ssnap16)
30820 .arg(state_in)
30821 .arg(&mut *state_out)
30822 .arg(&hi)
30823 .arg(&ti)
30824 .arg(&ci)
30825 .arg(&hki);
30826 unsafe {
30827 b.launch(cfg)?;
30828 }
30829 }
30830 {
30831 let f = self.func("gdn_chunk_output_mma");
30833 #[allow(clippy::manual_div_ceil)]
30834 let jt = ((c + 31) / 32) as u32;
30836 let cfg = LaunchConfig {
30837 grid_dim: (nc as u32, h as u32, jt),
30838 block_dim: (256, 1, 1),
30839 shared_mem_bytes: 0,
30840 };
30841 let hki = hk as i32;
30842 let __s_b = self.gpu.stream();
30843 let mut b = __s_b.launch_builder(&f);
30844 b.arg(q)
30845 .arg(&gcum)
30846 .arg(&p)
30847 .arg(&y16)
30848 .arg(&ssnap16)
30849 .arg(o)
30850 .arg(&hi)
30851 .arg(&ti)
30852 .arg(&ci)
30853 .arg(&scale)
30854 .arg(&hki);
30855 unsafe {
30856 b.launch(cfg)?;
30857 }
30858 }
30859 return Ok(());
30860 }
30861 {
30862 let f = self.func("gdn_chunk_state_f32");
30864 let cfg = LaunchConfig {
30865 grid_dim: (h as u32, NSPLIT, 1),
30866 block_dim: (256, 1, 1),
30867 shared_mem_bytes: 0,
30868 };
30869 let __s_b = self.gpu.stream();
30870 let mut b = __s_b.launch_builder(&f);
30871 b.arg(k)
30872 .arg(&gcum)
30873 .arg(beta)
30874 .arg(&u)
30875 .arg(&w)
30876 .arg(&mut y)
30877 .arg(&mut ssnap)
30878 .arg(state_in)
30879 .arg(&mut *state_out)
30880 .arg(&hi)
30881 .arg(&ti)
30882 .arg(&ci);
30883 unsafe {
30884 b.launch(cfg)?;
30885 }
30886 }
30887 {
30888 let f = self.func("gdn_chunk_output_f32");
30890 #[allow(clippy::manual_div_ceil)]
30891 let jt = ((c + 31) / 32) as u32;
30893 let cfg = LaunchConfig {
30894 grid_dim: (nc as u32, h as u32, jt),
30895 block_dim: (256, 1, 1),
30896 shared_mem_bytes: 0,
30897 };
30898 let __s_b = self.gpu.stream();
30899 let mut b = __s_b.launch_builder(&f);
30900 b.arg(q)
30901 .arg(&gcum)
30902 .arg(&p)
30903 .arg(&y)
30904 .arg(&ssnap)
30905 .arg(o)
30906 .arg(&hi)
30907 .arg(&ti)
30908 .arg(&ci)
30909 .arg(&scale);
30910 unsafe {
30911 b.launch(cfg)?;
30912 }
30913 }
30914 Ok(())
30915 }
30916
30917 #[allow(clippy::too_many_arguments)]
30926 #[allow(clippy::too_many_arguments)]
30927 pub fn gdn_scan_prefill(
30928 &self,
30929 q: &CudaSlice<f32>,
30930 k: &CudaSlice<f32>,
30931 v: &CudaSlice<f32>,
30932 g: &CudaSlice<f32>,
30933 beta: &CudaSlice<f32>,
30934 kb16_pre: Option<&CudaSlice<u8>>,
30935 qb16_pre: Option<&CudaSlice<u8>>,
30936 state_in: &CudaSlice<f32>,
30937 state_out: &mut CudaSlice<f32>,
30938 o: &mut CudaSlice<f32>,
30939 n_head: usize,
30940 t: usize,
30941 scale: f32,
30942 hk: usize,
30943 ) -> Result<(), Box<dyn std::error::Error>> {
30944 if std::env::var("MEMRA_GDN_DIFF").is_ok() && t >= 16 {
30945 assert!(hk == n_head, "GDN_DIFF oracle is broadcast-only");
30946 return self.gdn_scan_diff(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale);
30947 }
30948 if Self::gdn_chunked_enabled() && t >= 16 {
30949 self.gdn_scan_chunked(
30950 q,
30951 k,
30952 v,
30953 g,
30954 beta,
30955 kb16_pre,
30956 qb16_pre,
30957 state_in,
30958 state_out,
30959 o,
30960 n_head,
30961 t,
30962 scale,
30963 Self::gdn_chunk_size(),
30964 hk,
30965 )
30966 } else {
30967 assert!(
30968 hk == n_head,
30969 "s128 scan is broadcast-only (prep guarantees by predicate)"
30970 );
30971 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)
30972 }
30973 }
30974
30975 #[allow(clippy::too_many_arguments)]
30977 fn gdn_scan_diff(
30978 &self,
30979 q: &CudaSlice<f32>,
30980 k: &CudaSlice<f32>,
30981 v: &CudaSlice<f32>,
30982 g: &CudaSlice<f32>,
30983 beta: &CudaSlice<f32>,
30984 state_in: &CudaSlice<f32>,
30985 state_out: &mut CudaSlice<f32>,
30986 o: &mut CudaSlice<f32>,
30987 n_head: usize,
30988 t: usize,
30989 scale: f32,
30990 ) -> Result<(), Box<dyn std::error::Error>> {
30991 static CALL: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
30992 let call = CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
30993 let mut o_c = self.uninit(o.len())?;
30994 let mut st_c = self.uninit(state_out.len())?;
30995 self.gdn_scan_chunked(
30996 q,
30997 k,
30998 v,
30999 g,
31000 beta,
31001 None,
31002 None,
31003 state_in,
31004 &mut st_c,
31005 &mut o_c,
31006 n_head,
31007 t,
31008 scale,
31009 Self::gdn_chunk_size(),
31010 n_head,
31011 )?;
31012 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)?;
31013 let (oh_s, oh_c) = (self.dtoh(o)?, self.dtoh(&o_c)?);
31014 let (sh_s, sh_c) = (self.dtoh(state_out)?, self.dtoh(&st_c)?);
31015 let stats = |a: &[f32], b: &[f32]| -> (f32, f32, f64) {
31016 let mut max_abs = 0f32;
31017 let mut max_rel = 0f32;
31018 let mut sum_rel = 0f64;
31019 for (x, y) in a.iter().zip(b) {
31020 let ad = (x - y).abs();
31021 let rel = ad / x.abs().max(y.abs()).max(1e-3);
31022 if ad > max_abs {
31023 max_abs = ad;
31024 }
31025 if rel > max_rel {
31026 max_rel = rel;
31027 }
31028 sum_rel += rel as f64;
31029 }
31030 (max_abs, max_rel, sum_rel / a.len() as f64)
31031 };
31032 let (o_ma, o_mr, o_mean) = stats(&oh_s, &oh_c);
31033 let (s_ma, s_mr, s_mean) = stats(&sh_s, &sh_c);
31034 println!(
31035 "[gdn-diff call {call:3} T={t} C={}] out: max_abs={o_ma:.3e} max_rel={o_mr:.3e} mean_rel={o_mean:.3e} | \
31036 state: max_abs={s_ma:.3e} max_rel={s_mr:.3e} mean_rel={s_mean:.3e}",
31037 Self::gdn_chunk_size()
31038 );
31039 Ok(())
31040 }
31041
31042 pub fn gdn_glog(
31044 &self,
31045 alpha: &CudaSlice<f32>,
31046 dt_bias: &CudaSlice<f32>,
31047 a: &CudaSlice<f32>,
31048 g_log: &mut CudaSlice<f32>,
31049 n_head: usize,
31050 t: usize,
31051 ) -> Result<(), Box<dyn std::error::Error>> {
31052 let f = self.func("gdn_glog_f32");
31053 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
31054 let (h, ti) = (n_head as i32, t as i32);
31055 let __s_b = self.gpu.stream();
31056 let mut b = __s_b.launch_builder(&f);
31057 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
31058 unsafe {
31059 b.launch(cfg)?;
31060 }
31061 Ok(())
31062 }
31063
31064 pub fn sigmoid_v(
31067 &self,
31068 x: &cudarc::driver::CudaView<f32>,
31069 y: &mut CudaSlice<f32>,
31070 n: usize,
31071 ) -> Result<(), Box<dyn std::error::Error>> {
31072 let f = self.func("sigmoid_f32");
31073 let cfg = LaunchConfig::for_num_elems(n as u32);
31074 let ni = n as i32;
31075 let __s_b = self.gpu.stream();
31076 let mut b = __s_b.launch_builder(&f);
31077 b.arg(x).arg(y).arg(&ni);
31078 unsafe {
31079 b.launch(cfg)?;
31080 }
31081 Ok(())
31082 }
31083
31084 pub fn gdn_glog_v(
31085 &self,
31086 alpha: &cudarc::driver::CudaView<f32>,
31087 dt_bias: &CudaSlice<f32>,
31088 a: &CudaSlice<f32>,
31089 g_log: &mut CudaSlice<f32>,
31090 n_head: usize,
31091 t: usize,
31092 ) -> Result<(), Box<dyn std::error::Error>> {
31093 let f = self.func("gdn_glog_f32");
31094 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
31095 let (h, ti) = (n_head as i32, t as i32);
31096 let __s_b = self.gpu.stream();
31097 let mut b = __s_b.launch_builder(&f);
31098 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
31099 unsafe {
31100 b.launch(cfg)?;
31101 }
31102 Ok(())
31103 }
31104
31105 pub fn sigmoid(
31106 &self,
31107 x: &CudaSlice<f32>,
31108 y: &mut CudaSlice<f32>,
31109 n: usize,
31110 ) -> Result<(), Box<dyn std::error::Error>> {
31111 let f = self.func("sigmoid_f32");
31112 let cfg = LaunchConfig::for_num_elems(n as u32);
31113 let ni = n as i32;
31114 let __s_b = self.gpu.stream();
31115 let mut b = __s_b.launch_builder(&f);
31116 b.arg(x).arg(y).arg(&ni);
31117 unsafe {
31118 b.launch(cfg)?;
31119 }
31120 Ok(())
31121 }
31122
31123 pub fn sig_mul_f16out(
31126 &self,
31127 a: &CudaSlice<f32>,
31128 g: &CudaSlice<f32>,
31129 dst: &mut CudaSlice<f32>,
31130 dst16: &mut CudaSlice<u8>,
31131 n: usize,
31132 ) -> Result<(), Box<dyn std::error::Error>> {
31133 let f = self.func("sig_mul_f16out_f32");
31134 let cfg = LaunchConfig::for_num_elems(n as u32);
31135 let ni = n as i32;
31136 let __s_b = self.gpu.stream();
31137 let mut b = __s_b.launch_builder(&f);
31138 b.arg(a).arg(g).arg(dst).arg(dst16).arg(&ni);
31139 unsafe {
31140 b.launch(cfg)?;
31141 }
31142 Ok(())
31143 }
31144
31145 #[allow(clippy::too_many_arguments)]
31154 pub fn attn_head_gate(
31155 &self,
31156 a: &CudaSlice<f32>,
31157 g: &CudaSlice<f32>,
31158 dst: &mut CudaSlice<f32>,
31159 dst16: Option<&mut CudaSlice<u8>>,
31160 head_dim: usize,
31161 n_head: usize,
31162 t: usize,
31163 ) -> Result<(), Box<dyn std::error::Error>> {
31164 let f = self.func("attn_head_gate_f32");
31165 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
31166 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
31167 let d16: u64 = match dst16 {
31169 Some(d) => self.addr_u8(d),
31170 None => 0,
31171 };
31172 let __s_b = self.gpu.stream();
31173 let mut b = __s_b.launch_builder(&f);
31174 b.arg(a)
31175 .arg(g)
31176 .arg(dst)
31177 .arg(&d16)
31178 .arg(&hd)
31179 .arg(&nh)
31180 .arg(&ti);
31181 unsafe {
31182 b.launch(cfg)?;
31183 }
31184 Ok(())
31185 }
31186
31187 #[allow(clippy::too_many_arguments)]
31196 pub fn swiglu_clamped_mul_scaled(
31197 &self,
31198 gate: &CudaSlice<f32>,
31199 up: &CudaSlice<f32>,
31200 gs: f32,
31201 us: f32,
31202 limit: f32,
31203 dst: &mut CudaSlice<f32>,
31204 n: usize,
31205 ) -> Result<(), Box<dyn std::error::Error>> {
31206 debug_assert!(
31207 limit > 1e-6,
31208 "swiglu_clamped needs a live limit; use silu_mul_scaled"
31209 );
31210 let f = self.func("swiglu_clamped_mul_scaled_f32");
31211 let cfg = LaunchConfig::for_num_elems(n as u32);
31212 let ni = n as i32;
31213 let __s_b = self.gpu.stream();
31214 let mut b = __s_b.launch_builder(&f);
31215 b.arg(gate)
31216 .arg(up)
31217 .arg(&gs)
31218 .arg(&us)
31219 .arg(&limit)
31220 .arg(dst)
31221 .arg(&ni);
31222 unsafe {
31223 b.launch(cfg)?;
31224 }
31225 Ok(())
31226 }
31227
31228 #[allow(clippy::too_many_arguments)]
31237 pub fn swiglu_preclamped_mul_scaled(
31238 &self,
31239 gate: &CudaSlice<f32>,
31240 up: &CudaSlice<f32>,
31241 gs: f32,
31242 us: f32,
31243 limit: f32,
31244 dst: &mut CudaSlice<f32>,
31245 n: usize,
31246 ) -> Result<(), Box<dyn std::error::Error>> {
31247 debug_assert!(
31248 limit > 1e-6,
31249 "swiglu_preclamped needs a live limit; use silu_mul_scaled"
31250 );
31251 let f = self.func("swiglu_preclamped_mul_scaled_f32");
31252 let cfg = LaunchConfig::for_num_elems(n as u32);
31253 let ni = n as i32;
31254 let __s_b = self.gpu.stream();
31255 let mut b = __s_b.launch_builder(&f);
31256 b.arg(gate)
31257 .arg(up)
31258 .arg(&gs)
31259 .arg(&us)
31260 .arg(&limit)
31261 .arg(dst)
31262 .arg(&ni);
31263 unsafe {
31264 b.launch(cfg)?;
31265 }
31266 Ok(())
31267 }
31268
31269 #[allow(clippy::too_many_arguments)] pub fn gated_rmsnorm(
31272 &self,
31273 o: &CudaSlice<f32>,
31274 w: &CudaSlice<f32>,
31275 z: &CudaSlice<f32>,
31276 dst: &mut CudaSlice<f32>,
31277 ncols: usize,
31278 nrows: usize,
31279 eps: f32,
31280 ) -> Result<(), Box<dyn std::error::Error>> {
31281 let f = self.func("gated_rmsnorm_f32");
31282 let cfg = LaunchConfig {
31283 grid_dim: (nrows as u32, 1, 1),
31284 block_dim: (128, 1, 1),
31285 shared_mem_bytes: 0,
31286 };
31287 let (nc, e) = (ncols as i32, eps);
31288 let __s_b = self.gpu.stream();
31289 let mut b = __s_b.launch_builder(&f);
31290 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
31291 unsafe {
31292 b.launch(cfg)?;
31293 }
31294 Ok(())
31295 }
31296
31297 #[allow(clippy::too_many_arguments)] pub fn gated_rmsnorm_f16out(
31301 &self,
31302 o: &CudaSlice<f32>,
31303 w: &CudaSlice<f32>,
31304 z: &CudaSlice<f32>,
31305 dst: &mut CudaSlice<f32>,
31306 dst16: &mut CudaSlice<u8>,
31307 ncols: usize,
31308 nrows: usize,
31309 eps: f32,
31310 ) -> Result<(), Box<dyn std::error::Error>> {
31311 let f = self.func("gated_rmsnorm_f16out_f32");
31312 let cfg = LaunchConfig {
31314 grid_dim: (nrows as u32, 1, 1),
31315 block_dim: (128, 1, 1),
31316 shared_mem_bytes: 0,
31317 };
31318 let (nc, e) = (ncols as i32, eps);
31319 let __s_b = self.gpu.stream();
31320 let mut b = __s_b.launch_builder(&f);
31321 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
31322 unsafe {
31323 b.launch(cfg)?;
31324 }
31325 Ok(())
31326 }
31327
31328 #[allow(clippy::too_many_arguments)]
31332 pub fn add_rms_norm_zq8(
31333 &self,
31334 a: &CudaSlice<f32>,
31335 b_in: &CudaSlice<f32>,
31336 w: &CudaSlice<f32>,
31337 res: &mut CudaSlice<f32>,
31338 z: &mut CudaSlice<f32>,
31339 ncols: usize,
31340 nrows: usize,
31341 eps: f32,
31342 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
31343 assert!(ncols.is_multiple_of(32));
31344 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
31345 let mut d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
31346 let f = self.func("add_rms_norm_zq8");
31347 let cfg = LaunchConfig {
31348 grid_dim: (nrows as u32, 1, 1),
31349 block_dim: (1024, 1, 1),
31350 shared_mem_bytes: 0,
31351 };
31352 let (nc, ep) = (ncols as i32, eps);
31353 let __s_b = self.gpu.stream();
31354 let mut b = __s_b.launch_builder(&f);
31355 b.arg(a)
31356 .arg(b_in)
31357 .arg(w)
31358 .arg(res)
31359 .arg(z)
31360 .arg(&mut q)
31361 .arg(&mut d)
31362 .arg(&nc)
31363 .arg(&ep);
31364 unsafe {
31365 b.launch(cfg)?;
31366 }
31367 Ok((q, d))
31368 }
31369
31370 #[allow(clippy::too_many_arguments)] pub fn gated_rmsnorm_zv(
31376 &self,
31377 o: &CudaSlice<f32>,
31378 w: &CudaSlice<f32>,
31379 z: &cudarc::driver::CudaView<f32>,
31380 dst: &mut CudaSlice<f32>,
31381 ncols: usize,
31382 nrows: usize,
31383 eps: f32,
31384 ) -> Result<(), Box<dyn std::error::Error>> {
31385 let f = self.func("gated_rmsnorm_f32");
31386 let cfg = LaunchConfig {
31387 grid_dim: (nrows as u32, 1, 1),
31388 block_dim: (128, 1, 1),
31389 shared_mem_bytes: 0,
31390 };
31391 let (nc, e) = (ncols as i32, eps);
31392 let __s_b = self.gpu.stream();
31393 let mut b = __s_b.launch_builder(&f);
31394 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
31395 unsafe {
31396 b.launch(cfg)?;
31397 }
31398 Ok(())
31399 }
31400
31401 #[allow(clippy::too_many_arguments)] pub fn gated_rmsnorm_f16out_zv(
31403 &self,
31404 o: &CudaSlice<f32>,
31405 w: &CudaSlice<f32>,
31406 z: &cudarc::driver::CudaView<f32>,
31407 dst: &mut CudaSlice<f32>,
31408 dst16: &mut CudaSlice<u8>,
31409 ncols: usize,
31410 nrows: usize,
31411 eps: f32,
31412 ) -> Result<(), Box<dyn std::error::Error>> {
31413 let f = self.func("gated_rmsnorm_f16out_f32");
31414 let cfg = LaunchConfig {
31416 grid_dim: (nrows as u32, 1, 1),
31417 block_dim: (128, 1, 1),
31418 shared_mem_bytes: 0,
31419 };
31420 let (nc, e) = (ncols as i32, eps);
31421 let __s_b = self.gpu.stream();
31422 let mut b = __s_b.launch_builder(&f);
31423 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
31424 unsafe {
31425 b.launch(cfg)?;
31426 }
31427 Ok(())
31428 }
31429
31430 pub fn gated_rmsnorm_q8_1(
31431 &self,
31432 o: &CudaSlice<f32>,
31433 w: &CudaSlice<f32>,
31434 z: &CudaSlice<f32>,
31435 ncols: usize,
31436 nrows: usize,
31437 eps: f32,
31438 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
31439 assert!(ncols.is_multiple_of(32));
31440 let f = self.func("gated_rmsnorm_q8_1");
31441 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
31442 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
31443 let cfg = LaunchConfig {
31444 grid_dim: (nrows as u32, 1, 1),
31445 block_dim: (128, 1, 1),
31446 shared_mem_bytes: 0,
31447 };
31448 let (nc, ep) = (ncols as i32, eps);
31449 let __s_b = self.gpu.stream();
31450 let mut b = __s_b.launch_builder(&f);
31451 b.arg(o)
31452 .arg(w)
31453 .arg(z)
31454 .arg(&mut out_q)
31455 .arg(&mut out_d)
31456 .arg(&nc)
31457 .arg(&ep);
31458 unsafe {
31459 b.launch(cfg)?;
31460 }
31461 Ok((out_q, out_d))
31462 }
31463
31464 pub fn transpose(
31466 &self,
31467 inp: &CudaSlice<f32>,
31468 rows: usize,
31469 cols: usize,
31470 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31471 let f = self.func("transpose_f32");
31472 let mut out = self.zeros(rows * cols)?;
31473 let cfg = LaunchConfig::for_num_elems((rows * cols) as u32);
31474 let (r, c) = (rows as i32, cols as i32);
31475 let __s_b = self.gpu.stream();
31476 let mut b = __s_b.launch_builder(&f);
31477 b.arg(inp).arg(&mut out).arg(&r).arg(&c);
31478 unsafe {
31479 b.launch(cfg)?;
31480 }
31481 Ok(out)
31482 }
31483
31484 pub fn repeat_heads(
31486 &self,
31487 inp: &CudaSlice<f32>,
31488 out: &mut CudaSlice<f32>,
31489 head_dim: usize,
31490 n_in: usize,
31491 n_out: usize,
31492 t: usize,
31493 ) -> Result<(), Box<dyn std::error::Error>> {
31494 let f = self.func("repeat_heads_f32");
31495 let cfg = LaunchConfig::for_num_elems((head_dim * n_out * t) as u32);
31496 let (hd, ni, no, ti) = (head_dim as i32, n_in as i32, n_out as i32, t as i32);
31497 let __s_b = self.gpu.stream();
31498 let mut b = __s_b.launch_builder(&f);
31499 b.arg(inp).arg(out).arg(&hd).arg(&ni).arg(&no).arg(&ti);
31500 unsafe {
31501 b.launch(cfg)?;
31502 }
31503 Ok(())
31504 }
31505
31506 pub fn q_gate_split(
31513 &self,
31514 qf: &CudaSlice<f32>,
31515 q_out: &mut CudaSlice<f32>,
31516 gate_out: &mut CudaSlice<f32>,
31517 head_dim: usize,
31518 n_head: usize,
31519 t: usize,
31520 ) -> Result<(), Box<dyn std::error::Error>> {
31521 memra_gguf::config::check_fused_q_gate_extent(qf.len(), head_dim, n_head, t)?;
31522 let out_need = head_dim * n_head * t;
31523 if q_out.len() < out_need || gate_out.len() < out_need {
31524 return Err(format!(
31525 "q_gate_split destinations too small: need {out_need} each, have q={} gate={}",
31526 q_out.len(),
31527 gate_out.len()
31528 )
31529 .into());
31530 }
31531 let f = self.func("q_gate_split_f32");
31532 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
31533 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
31534 let __s_b = self.gpu.stream();
31535 let mut b = __s_b.launch_builder(&f);
31536 b.arg(qf)
31537 .arg(q_out)
31538 .arg(gate_out)
31539 .arg(&hd)
31540 .arg(&nh)
31541 .arg(&ti);
31542 unsafe {
31543 b.launch(cfg)?;
31544 }
31545 Ok(())
31546 }
31547
31548 #[allow(clippy::too_many_arguments)] pub fn qkv_to_gdn_repack(
31553 &self,
31554 conv_out: &CudaSlice<f32>,
31555 q_g: &mut CudaSlice<f32>,
31556 k_g: &mut CudaSlice<f32>,
31557 v_g: &mut CudaSlice<f32>,
31558 d_state: usize,
31559 num_v: usize,
31560 num_k: usize,
31561 key_dim: usize,
31562 t: usize,
31563 ) -> Result<(), Box<dyn std::error::Error>> {
31564 let f = self.func("qkv_to_gdn_repack_f32");
31565 let cfg = LaunchConfig::for_num_elems((d_state * num_v * t) as u32);
31566 let (ds, nv, nk, kd, ti) = (
31567 d_state as i32,
31568 num_v as i32,
31569 num_k as i32,
31570 key_dim as i32,
31571 t as i32,
31572 );
31573 let __s_b = self.gpu.stream();
31574 let mut b = __s_b.launch_builder(&f);
31575 b.arg(conv_out)
31576 .arg(q_g)
31577 .arg(k_g)
31578 .arg(v_g)
31579 .arg(&ds)
31580 .arg(&nv)
31581 .arg(&nk)
31582 .arg(&kd)
31583 .arg(&ti);
31584 unsafe {
31585 b.launch(cfg)?;
31586 }
31587 Ok(())
31588 }
31589
31590 pub fn conv_left_pad(
31593 &self,
31594 src: &CudaSlice<f32>,
31595 dst: &mut CudaSlice<f32>,
31596 conv_dim: usize,
31597 t: usize,
31598 pad: usize,
31599 ) -> Result<(), Box<dyn std::error::Error>> {
31600 let f = self.func("conv_left_pad_f32");
31601 let cfg = LaunchConfig::for_num_elems((conv_dim * t) as u32);
31602 let (cd, ti, p) = (conv_dim as i32, t as i32, pad as i32);
31603 let __s_b = self.gpu.stream();
31604 let mut b = __s_b.launch_builder(&f);
31605 b.arg(src).arg(dst).arg(&cd).arg(&ti).arg(&p);
31606 unsafe {
31607 b.launch(cfg)?;
31608 }
31609 Ok(())
31610 }
31611
31612 pub fn conv_assemble_and_roll(
31616 &self,
31617 qkv_col: &CudaSlice<f32>,
31618 conv_state: &mut CudaSlice<f32>,
31619 conv_in: &mut CudaSlice<f32>,
31620 conv_dim: usize,
31621 pad: usize,
31622 ) -> Result<(), Box<dyn std::error::Error>> {
31623 let f = self.func("conv_assemble_and_roll_f32");
31624 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
31625 let (cd, p) = (conv_dim as i32, pad as i32);
31626 let __s_b = self.gpu.stream();
31627 let mut b = __s_b.launch_builder(&f);
31628 b.arg(qkv_col).arg(conv_state).arg(conv_in).arg(&cd).arg(&p);
31629 unsafe {
31630 b.launch(cfg)?;
31631 }
31632 Ok(())
31633 }
31634
31635 pub fn ssm_conv1d_fused_decode(
31641 &self,
31642 qkv_col: &CudaSlice<f32>,
31643 conv_state: &mut CudaSlice<f32>,
31644 w: &CudaSlice<f32>,
31645 conv_out: &mut CudaSlice<f32>,
31646 conv_dim: usize,
31647 d_conv: usize,
31648 ) -> Result<(), Box<dyn std::error::Error>> {
31649 let f = self.func("ssm_conv1d_fused_decode_f32");
31650 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
31651 let (cd, dc) = (conv_dim as i32, d_conv as i32);
31652 let __s_b = self.gpu.stream();
31653 let mut b = __s_b.launch_builder(&f);
31654 b.arg(qkv_col)
31655 .arg(conv_state)
31656 .arg(w)
31657 .arg(conv_out)
31658 .arg(&cd)
31659 .arg(&dc);
31660 unsafe {
31661 b.launch(cfg)?;
31662 }
31663 Ok(())
31664 }
31665
31666 pub fn slice_range(
31669 &self,
31670 src: &CudaSlice<f32>,
31671 start: usize,
31672 len: usize,
31673 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31674 let host = self.gpu.stream().clone_dtoh(src)?;
31675 self.gpu.stream().synchronize()?;
31676 self.htod(&host[start..start + len])
31677 }
31678}
31679
31680#[cfg(test)]
31681mod target_dispatch_tests {
31682 use super::legacy_quant_gemm_allowed;
31683
31684 #[test]
31685 fn legacy_quant_gemm_arch_policy_honors_the_escape_hatch() {
31686 assert!(legacy_quant_gemm_allowed(false, false, false));
31688 assert!(!legacy_quant_gemm_allowed(false, false, true));
31689 assert!(!legacy_quant_gemm_allowed(true, false, false));
31691 assert!(!legacy_quant_gemm_allowed(true, false, true));
31692 assert!(legacy_quant_gemm_allowed(true, true, false));
31694 assert!(!legacy_quant_gemm_allowed(true, true, true));
31695 }
31696
31697 #[cfg(all(memra_portable_cuda, not(memra_hopper_mma)))]
31698 #[test]
31699 fn portable_build_disables_legacy_quant_gemm_without_an_env_override() {
31700 assert!(!legacy_quant_gemm_allowed(
31701 cfg!(memra_portable_cuda),
31702 cfg!(memra_hopper_mma),
31703 false
31704 ));
31705 }
31706
31707 #[cfg(memra_hopper_mma)]
31708 #[test]
31709 fn hopper_mma_build_re_admits_legacy_quant_gemm() {
31710 assert!(legacy_quant_gemm_allowed(
31711 cfg!(memra_portable_cuda),
31712 cfg!(memra_hopper_mma),
31713 false
31714 ));
31715 assert!(super::portable_mma_gated() == false);
31716 }
31717}
31718
31719impl memra_kv::KvDev for Engine {
31722 fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31723 Engine::zeros(self, n)
31724 }
31725 fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31726 Engine::uninit(self, n)
31727 }
31728 fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
31729 Engine::alloc_u8(self, n)
31730 }
31731 fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
31732 Engine::htod_i32(self, v)
31733 }
31734 fn clone_dtod(
31735 &self,
31736 src: &CudaSlice<f32>,
31737 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31738 Engine::clone_dtod(self, src)
31739 }
31740 fn copy_into(
31741 &self,
31742 dst: &mut CudaSlice<f32>,
31743 off: usize,
31744 src: &CudaSlice<f32>,
31745 len: usize,
31746 ) -> Result<(), Box<dyn std::error::Error>> {
31747 Engine::copy_into(self, dst, off, src, len)
31748 }
31749 fn copy_range_into(
31750 &self,
31751 dst: &mut CudaSlice<f32>,
31752 dst_off: usize,
31753 src: &CudaSlice<f32>,
31754 src_off: usize,
31755 len: usize,
31756 ) -> Result<(), Box<dyn std::error::Error>> {
31757 Engine::copy_range_into(self, dst, dst_off, src, src_off, len)
31758 }
31759 fn set_i32_one(
31760 &self,
31761 d: &mut CudaSlice<i32>,
31762 v: i32,
31763 ) -> Result<(), Box<dyn std::error::Error>> {
31764 Engine::set_i32_one(self, d, v)
31765 }
31766}
31767
31768#[cfg(test)]
31769mod fused_gate_bounds_tests {
31770 use super::*;
31771
31772 #[test]
31785 #[ignore = "requires a CUDA GPU"]
31786 fn q_gate_split_refuses_a_separate_gate_wq_instead_of_reading_past_it() {
31787 let e = Engine::new(0).unwrap();
31788 let (head_dim, n_head, t) = (8usize, 4usize, 2usize);
31789 let fused = 2 * head_dim * n_head * t;
31790 let out_n = head_dim * n_head * t;
31791
31792 let narrow = e.htod(&vec![1.0f32; out_n]).unwrap();
31794 let mut q = e.uninit(out_n).unwrap();
31795 let mut gate = e.uninit(out_n).unwrap();
31796 let err = e
31797 .q_gate_split(&narrow, &mut q, &mut gate, head_dim, n_head, t)
31798 .expect_err("half-width wq must be refused, not read past")
31799 .to_string();
31800 assert!(err.contains("NO fused gate"), "{err}");
31801 assert!(err.contains(&format!("{fused}")), "{err}");
31802
31803 let host: Vec<f32> = (0..fused).map(|i| i as f32).collect();
31806 let wide = e.htod(&host).unwrap();
31807 e.q_gate_split(&wide, &mut q, &mut gate, head_dim, n_head, t)
31808 .expect("full-width wq splits");
31809 let (qh, gh) = (e.dtoh(&q).unwrap(), e.dtoh(&gate).unwrap());
31810 for tok in 0..t {
31811 for hh in 0..n_head {
31812 for d in 0..head_dim {
31813 let base = tok * (n_head * 2 * head_dim) + hh * (2 * head_dim);
31814 let idx = tok * (n_head * head_dim) + hh * head_dim + d;
31815 assert_eq!(qh[idx], host[base + d], "q t{tok} h{hh} d{d}");
31816 assert_eq!(gh[idx], host[base + head_dim + d], "gate t{tok} h{hh} d{d}");
31817 }
31818 }
31819 }
31820
31821 let mut small = e.uninit(out_n - 1).unwrap();
31823 assert!(
31824 e.q_gate_split(&wide, &mut small, &mut gate, head_dim, n_head, t)
31825 .is_err()
31826 );
31827 }
31828}
31829
31830#[cfg(test)]
31834mod fused_rope_width_tests {
31835 use super::Engine;
31836
31837 #[test]
31840 fn full_width_is_accepted() {
31841 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope", 256, 256).is_ok());
31842 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_cat", 512, 512).is_ok());
31843 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_append", 128, 128).is_ok());
31844 }
31845
31846 #[test]
31859 fn gemma4_official_artifact_widths_pass() {
31860 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope", 512, 512).is_ok());
31861 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_append_dc", 256, 256).is_ok());
31862 }
31863
31864 #[test]
31867 fn partial_rotary_is_refused_with_the_geometry_named() {
31868 let err = Engine::full_width_rope_only("rms_norm_qkv_rope", 64, 256)
31870 .expect_err("partial rotary must refuse");
31871 let msg = err.to_string();
31872 assert!(msg.contains("PARTIAL ROTARY REFUSED"), "{msg}");
31873 assert!(msg.contains("n_rot 64"), "{msg}");
31874 assert!(msg.contains("head_dim 256"), "{msg}");
31875 assert!(
31876 msg.contains("64..256"),
31877 "names the band it would corrupt: {msg}"
31878 );
31879 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_append_dc", 64, 128).is_err());
31881 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope", 256, 128).is_err());
31883 }
31884}