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 pub fn htod_u64_into(
5856 &self,
5857 v: &[u64],
5858 dst: &mut CudaSlice<u64>,
5859 ) -> Result<(), Box<dyn std::error::Error>> {
5860 let mut view = dst.slice_mut(0..v.len());
5861 self.gpu.stream().memcpy_htod(v, &mut view)?;
5862 Ok(())
5863 }
5864
5865 pub fn htod_f32_into(
5868 &self,
5869 v: &[f32],
5870 dst: &mut CudaSlice<f32>,
5871 ) -> Result<(), Box<dyn std::error::Error>> {
5872 let mut view = dst.slice_mut(0..v.len());
5873 self.gpu.stream().memcpy_htod(v, &mut view)?;
5874 Ok(())
5875 }
5876
5877 pub fn htod_f32_into_at(
5881 &self,
5882 v: &[f32],
5883 dst: &mut CudaSlice<f32>,
5884 off: usize,
5885 ) -> Result<(), Box<dyn std::error::Error>> {
5886 if off + v.len() > dst.len() {
5887 return Err(format!(
5888 "htod_f32_into_at range {}..{} exceeds dst {}",
5889 off,
5890 off + v.len(),
5891 dst.len()
5892 )
5893 .into());
5894 }
5895 let mut view = dst.slice_mut(off..off + v.len());
5896 self.gpu.stream().memcpy_htod(v, &mut view)?;
5897 Ok(())
5898 }
5899
5900 pub fn copy_indirect_src_f32(
5905 &self,
5906 src_entry: &cudarc::driver::CudaView<u64>,
5907 dst: &mut CudaSlice<f32>,
5908 dst_off: usize,
5909 words: usize,
5910 ) -> Result<(), Box<dyn std::error::Error>> {
5911 let f = self.func("copy_indirect_src_f32");
5912 let chunks = (words / 4).max(1).div_ceil(256).min(48) as u32;
5913 let wi = words as i32;
5914 let cfg = LaunchConfig {
5915 grid_dim: (chunks, 1, 1),
5916 block_dim: (256, 1, 1),
5917 shared_mem_bytes: 0,
5918 };
5919 let mut dv = dst.slice_mut(dst_off..dst_off + words);
5920 let __s = self.gpu.stream();
5921 let mut b = __s.launch_builder(&f);
5922 b.arg(src_entry).arg(&mut dv).arg(&wi);
5923 unsafe {
5924 b.launch(cfg)?;
5925 }
5926 Ok(())
5927 }
5928
5929 pub fn uninit_i8(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
5931 self.alloc_uninit::<i8>(n)
5932 }
5933
5934 #[allow(clippy::too_many_arguments)] pub fn qmatvec(
5937 &self,
5938 w: &CudaSlice<u8>,
5939 x: &CudaSlice<f32>,
5940 m: usize,
5941 in_f: usize,
5942 out_f: usize,
5943 qtype: i32,
5944 row_bytes: usize,
5945 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5946 let f = self.func("qmatvec_f32");
5947 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
5949 grid_dim: (out_f as u32, m as u32, 1),
5950 block_dim: (256, 1, 1),
5951 shared_mem_bytes: 0,
5952 };
5953 let (inf, outf, mi, qt, rb) =
5954 (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
5955 let __s_b = self.gpu.stream();
5956 let mut b = __s_b.launch_builder(&f);
5957 b.arg(w)
5958 .arg(x)
5959 .arg(&mut y)
5960 .arg(&inf)
5961 .arg(&outf)
5962 .arg(&mi)
5963 .arg(&qt)
5964 .arg(&rb);
5965 unsafe {
5966 b.launch(cfg)?;
5967 }
5968 Ok(y)
5969 }
5970
5971 pub fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
5973 let s = self.gpu.stream().alloc_zeros::<u8>(n)?;
5974 self.keep_if_capturing(&s);
5975 Ok(s)
5976 }
5977
5978 pub fn alloc_u8_uninit(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
5982 let s = unsafe { self.gpu.stream().alloc::<u8>(n)? };
5983 self.keep_if_capturing(&s);
5984 Ok(s)
5985 }
5986
5987 pub fn memset_zeros_view(
5990 &self,
5991 dst: &mut cudarc::driver::CudaViewMut<f32>,
5992 ) -> Result<(), Box<dyn std::error::Error>> {
5993 self.gpu.stream().memset_zeros(dst)?;
5994 Ok(())
5995 }
5996
5997 pub fn stage_expert(
6003 &self,
6004 host_bytes: &[u8],
6005 scratch: &mut CudaSlice<u8>,
6006 off: usize,
6007 ) -> Result<(), Box<dyn std::error::Error>> {
6008 let mut dst = scratch.slice_mut(off..off + host_bytes.len()); self.gpu.stream().memcpy_htod(host_bytes, &mut dst)?; Ok(())
6011 }
6012
6013 pub fn moe_router_topk(
6019 &self,
6020 logits: &CudaSlice<f32>,
6021 t: usize,
6022 n_expert: usize,
6023 n_used: usize,
6024 ) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6025 let f = self.func("moe_router_topk_f32");
6026 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 {
6029 grid_dim: (t as u32, 1, 1),
6030 block_dim: (n_expert as u32, 1, 1),
6031 shared_mem_bytes: 0,
6032 };
6033 let (ne, nu) = (n_expert as i32, n_used as i32);
6034 let __s_b = self.gpu.stream();
6035 let mut b = __s_b.launch_builder(&f);
6036 b.arg(logits)
6037 .arg(&mut sel_idx)
6038 .arg(&mut sel_w)
6039 .arg(&ne)
6040 .arg(&nu);
6041 unsafe {
6042 b.launch(cfg)?;
6043 }
6044 Ok((sel_idx, sel_w))
6045 }
6046
6047 pub fn moe_router_topk_scaled(
6050 &self,
6051 logits: &CudaSlice<f32>,
6052 t: usize,
6053 n_expert: usize,
6054 n_used: usize,
6055 ex_scale: &CudaSlice<f32>,
6056 ) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6057 let f = self.func("moe_router_topk_scaled_f32");
6062 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
6063 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
6064 let cfg = LaunchConfig {
6065 grid_dim: (t as u32, 1, 1),
6066 block_dim: (n_expert as u32, 1, 1),
6067 shared_mem_bytes: 0,
6068 };
6069 let (ne, nu) = (n_expert as i32, n_used as i32);
6070 let __s_b = self.gpu.stream();
6071 let mut b = __s_b.launch_builder(&f);
6072 b.arg(logits)
6073 .arg(&mut sel_idx)
6074 .arg(&mut sel_w)
6075 .arg(&ne)
6076 .arg(&nu)
6077 .arg(ex_scale);
6078 unsafe {
6079 b.launch(cfg)?;
6080 }
6081 Ok((sel_idx, sel_w))
6082 }
6083
6084 pub fn moe_router_topk_host(
6092 &self,
6093 logits: &CudaSlice<f32>,
6094 t: usize,
6095 n_expert: usize,
6096 n_used: usize,
6097 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
6098 let f = self.func("moe_router_topk_f32");
6099 let n = t * n_used;
6100 let mut sel_idx = self.alloc_uninit::<i32>(n)?;
6101 let mut sel_w = self.alloc_uninit::<f32>(n)?;
6102 let cfg = LaunchConfig {
6103 grid_dim: (t as u32, 1, 1),
6104 block_dim: (n_expert as u32, 1, 1),
6105 shared_mem_bytes: 0,
6106 };
6107 let (ne, nu) = (n_expert as i32, n_used as i32);
6108 let __s_b = self.gpu.stream();
6109 let mut b = __s_b.launch_builder(&f);
6110 b.arg(logits)
6111 .arg(&mut sel_idx)
6112 .arg(&mut sel_w)
6113 .arg(&ne)
6114 .arg(&nu);
6115 unsafe {
6116 b.launch(cfg)?;
6117 }
6118 let bytes = n * 8;
6120 let mut guard = self.router_stage.lock().unwrap();
6121 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
6122 *guard = Some(PinnedStage::new(bytes.max(4096))?);
6123 }
6124 let stage = guard.as_mut().unwrap();
6125 let (si, sw) = unsafe {
6126 (
6127 std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
6128 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n),
6129 )
6130 };
6131 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()))
6135 }
6136
6137 #[allow(clippy::too_many_arguments)]
6141 pub fn moe_router_sigmoid_topk(
6142 &self,
6143 logits: &CudaSlice<f32>,
6144 t: usize,
6145 n_expert: usize,
6146 n_used: usize,
6147 active_count: usize,
6148 correction_bias: &CudaSlice<f32>,
6149 active: &CudaSlice<u8>,
6150 scaling_factor: f32,
6151 route_norm: bool,
6152 ) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6153 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
6154 if n_expert == 0 || n_expert > 1024 || n_used == 0 || n_used > n_expert {
6155 return Err(format!(
6156 "sigmoid router shape unsupported: n_expert={n_expert}, n_used={n_used}",
6157 )
6158 .into());
6159 }
6160 if logits.len() < t * n_expert
6161 || correction_bias.len() != n_expert
6162 || active.len() != n_expert
6163 {
6164 return Err(format!(
6165 "sigmoid router buffer mismatch: logits={} bias={} active={} expected logits>={} row={}",
6166 logits.len(), correction_bias.len(), active.len(), t * n_expert, n_expert,
6167 ).into());
6168 }
6169 let f = self.func(crate::sigmoid_topk_kernel(
6170 crate::sig_expf_dev_on(),
6171 crate::topk_fast_on(),
6172 n_used,
6173 ));
6174 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
6175 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
6176 let threads = n_expert.div_ceil(32) * 32;
6177 let cfg = LaunchConfig {
6178 grid_dim: (t as u32, 1, 1),
6179 block_dim: (threads as u32, 1, 1),
6180 shared_mem_bytes: 0,
6181 };
6182 let (ne, nu, rn) = (n_expert as i32, n_used as i32, i32::from(route_norm));
6183 let __s_b = self.gpu.stream();
6184 let mut b = __s_b.launch_builder(&f);
6185 b.arg(logits)
6186 .arg(correction_bias)
6187 .arg(active)
6188 .arg(&mut sel_idx)
6189 .arg(&mut sel_w)
6190 .arg(&ne)
6191 .arg(&nu)
6192 .arg(&scaling_factor)
6193 .arg(&rn);
6194 unsafe {
6195 b.launch(cfg)?;
6196 }
6197 Ok((sel_idx, sel_w))
6198 }
6199
6200 #[allow(clippy::too_many_arguments)]
6203 pub fn ring_flag_raw(&self, ptr: u64, value: u32) -> Result<(), Box<dyn std::error::Error>> {
6207 if ptr == 0 {
6208 return Err("ring_flag_raw: unarmed flag".into());
6209 }
6210 let f = self.func("memra_ring_flag");
6211 let cfg = LaunchConfig {
6212 grid_dim: (1, 1, 1),
6213 block_dim: (32, 1, 1),
6214 shared_mem_bytes: 0,
6215 };
6216 let __s_b = self.gpu.stream();
6217 let mut b = __s_b.launch_builder(&f);
6218 b.arg(&ptr).arg(&value);
6219 unsafe {
6220 b.launch(cfg)?;
6221 }
6222 Ok(())
6223 }
6224
6225 pub fn moe_sel_w_mirror(
6228 &self,
6229 sel_src: &CudaSlice<i32>,
6230 w_src: &CudaSlice<f32>,
6231 sel_dst: &mut CudaSlice<i32>,
6232 w_dst: &mut CudaSlice<f32>,
6233 n: usize,
6234 ) -> Result<(), Box<dyn std::error::Error>> {
6235 if n == 0
6236 || n > i32::MAX as usize
6237 || sel_src.len() < n
6238 || w_src.len() < n
6239 || sel_dst.len() < n
6240 || w_dst.len() < n
6241 {
6242 return Err(format!("moe_sel_w_mirror geometry n={n}").into());
6243 }
6244 let f = self.func("moe_sel_w_mirror");
6245 let threads = if n <= 32 { 32 } else { 128 };
6246 let cfg = LaunchConfig {
6247 grid_dim: ((n as u32).div_ceil(threads), 1, 1),
6248 block_dim: (threads, 1, 1),
6249 shared_mem_bytes: 0,
6250 };
6251 let ni = n as i32;
6252 let __s_b = self.gpu.stream();
6253 let mut b = __s_b.launch_builder(&f);
6254 b.arg(sel_src).arg(w_src).arg(sel_dst).arg(w_dst).arg(&ni);
6255 unsafe {
6256 b.launch(cfg)?;
6257 }
6258 Ok(())
6259 }
6260
6261 #[allow(clippy::too_many_arguments)]
6265 pub fn nvfp4_ep_stage_inputs(
6266 &self,
6267 input_src: &CudaSlice<f32>,
6268 sel_src: &CudaSlice<i32>,
6269 w_src: &CudaSlice<f32>,
6270 input_bf16_dst: &mut CudaSlice<u8>,
6271 sel_dst: &mut CudaSlice<i32>,
6272 w_dst: &mut CudaSlice<f32>,
6273 input_values: usize,
6274 pairs: usize,
6275 copy_weights: bool,
6276 ) -> Result<(), Box<dyn std::error::Error>> {
6277 if input_values == 0
6278 || pairs == 0
6279 || input_src.len() < input_values
6280 || sel_src.len() < pairs
6281 || w_src.len() < pairs
6282 || input_bf16_dst.len() < 2 * input_values
6283 || sel_dst.len() < pairs
6284 || w_dst.len() < pairs
6285 {
6286 return Err(format!(
6287 "W4A16 EP stage geometry input={} sel={} weights={} input_bf16={} \
6288 sel_dst={} weights_dst={} active={input_values} pairs={pairs}",
6289 input_src.len(),
6290 sel_src.len(),
6291 w_src.len(),
6292 input_bf16_dst.len(),
6293 sel_dst.len(),
6294 w_dst.len(),
6295 )
6296 .into());
6297 }
6298 let f = self.func("nvfp4_ep_stage_inputs");
6299 let n = input_values.max(pairs);
6300 let cfg = LaunchConfig::for_num_elems(n as u32);
6301 let (input_values, pairs, copy_weights) =
6302 (input_values as i32, pairs as i32, i32::from(copy_weights));
6303 let __s_b = self.gpu.stream();
6304 let mut b = __s_b.launch_builder(&f);
6305 b.arg(input_src)
6306 .arg(sel_src)
6307 .arg(w_src)
6308 .arg(input_bf16_dst)
6309 .arg(sel_dst)
6310 .arg(w_dst)
6311 .arg(&input_values)
6312 .arg(&pairs)
6313 .arg(©_weights);
6314 unsafe {
6315 b.launch(cfg)?;
6316 }
6317 Ok(())
6318 }
6319
6320 #[allow(clippy::too_many_arguments)]
6323 pub fn nvfp4_ep_stage_inputs_raw(
6324 &self,
6325 input_src: u64,
6326 sel_src: u64,
6327 w_src: u64,
6328 input_bf16_dst: &mut CudaSlice<u8>,
6329 sel_dst: &mut CudaSlice<i32>,
6330 w_dst: &mut CudaSlice<f32>,
6331 input_values: usize,
6332 pairs: usize,
6333 copy_weights: bool,
6334 ) -> Result<(), Box<dyn std::error::Error>> {
6335 if input_src == 0
6336 || sel_src == 0
6337 || w_src == 0
6338 || input_values == 0
6339 || pairs == 0
6340 || input_bf16_dst.len() < 2 * input_values
6341 || sel_dst.len() < pairs
6342 || w_dst.len() < pairs
6343 {
6344 return Err(format!(
6345 "W4A16 EP raw stage geometry input={input_src:#x} sel={sel_src:#x} \
6346 weights={w_src:#x} input_bf16={} sel_dst={} weights_dst={} \
6347 active={input_values} pairs={pairs}",
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)] pub fn moe_router_sigmoid_topk_into(
6378 &self,
6379 logits: &CudaSlice<f32>,
6380 t: usize,
6381 n_expert: usize,
6382 n_used: usize,
6383 active_count: usize,
6384 correction_bias: &CudaSlice<f32>,
6385 active: &CudaSlice<u8>,
6386 scaling_factor: f32,
6387 route_norm: bool,
6388 sel_idx: &mut CudaSlice<i32>,
6389 sel_w: &mut CudaSlice<f32>,
6390 ) -> Result<(), Box<dyn std::error::Error>> {
6391 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
6392 if n_expert == 0
6393 || n_expert > 1024
6394 || n_used == 0
6395 || n_used > 32 || n_used > n_expert
6397 || logits.len() < t * n_expert
6398 || correction_bias.len() != n_expert
6399 || active.len() != n_expert
6400 || sel_idx.len() < t * n_used
6401 || sel_w.len() < t * n_used
6402 {
6403 return Err("sigmoid router _into geometry mismatch".into());
6404 }
6405 let f = self.func(crate::sigmoid_topk_kernel(
6406 crate::sig_expf_dev_on(),
6407 crate::topk_fast_on(),
6408 n_used,
6409 ));
6410 let threads = n_expert.div_ceil(32) * 32;
6411 let cfg = LaunchConfig {
6412 grid_dim: (t as u32, 1, 1),
6413 block_dim: (threads as u32, 1, 1),
6414 shared_mem_bytes: 0,
6415 };
6416 let (ne, nu, rn) = (n_expert as i32, n_used as i32, i32::from(route_norm));
6417 let __s_b = self.gpu.stream();
6418 let mut b = __s_b.launch_builder(&f);
6419 b.arg(logits)
6420 .arg(correction_bias)
6421 .arg(active)
6422 .arg(&mut *sel_idx)
6423 .arg(&mut *sel_w)
6424 .arg(&ne)
6425 .arg(&nu)
6426 .arg(&scaling_factor)
6427 .arg(&rn);
6428 unsafe {
6429 b.launch(cfg)?;
6430 }
6431 Ok(())
6432 }
6433
6434 #[allow(clippy::too_many_arguments)]
6437 pub fn moe_router_sigmoid_topk_host(
6438 &self,
6439 logits: &CudaSlice<f32>,
6440 t: usize,
6441 n_expert: usize,
6442 n_used: usize,
6443 active_count: usize,
6444 correction_bias: &CudaSlice<f32>,
6445 active: &CudaSlice<u8>,
6446 scaling_factor: f32,
6447 route_norm: bool,
6448 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
6449 let (sel_idx, sel_w) = self.moe_router_sigmoid_topk(
6450 logits,
6451 t,
6452 n_expert,
6453 n_used,
6454 active_count,
6455 correction_bias,
6456 active,
6457 scaling_factor,
6458 route_norm,
6459 )?;
6460 let n = t * n_used;
6461 let bytes = n * 8;
6462 let mut guard = self.router_stage.lock().unwrap();
6463 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
6464 *guard = Some(PinnedStage::new(bytes.max(4096))?);
6465 }
6466 let stage = guard.as_mut().unwrap();
6467 let (si, sw) = unsafe {
6468 (
6469 std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
6470 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n),
6471 )
6472 };
6473 self.gpu.stream().memcpy_dtoh(&sel_idx, si)?;
6474 self.gpu.stream().memcpy_dtoh(&sel_w, sw)?;
6475 self.gpu.stream().synchronize()?;
6476 Ok((si.iter().map(|&i| i as u32).collect(), sw.to_vec()))
6477 }
6478
6479 pub fn stage_expert_async(
6483 &self,
6484 host_bytes: &[u8],
6485 scratch: &mut CudaSlice<u8>,
6486 off: usize,
6487 ) -> Result<cudarc::driver::CudaEvent, Box<dyn std::error::Error>> {
6488 let mut dst = scratch.slice_mut(off..off + host_bytes.len());
6489 self.copy_stream.memcpy_htod(host_bytes, &mut dst)?;
6490 Ok(self.copy_stream.record_event(None)?)
6491 }
6492
6493 pub fn compute_wait(
6495 &self,
6496 ev: &cudarc::driver::CudaEvent,
6497 ) -> Result<(), Box<dyn std::error::Error>> {
6498 self.gpu.stream().wait(ev)?;
6499 Ok(())
6500 }
6501
6502 #[allow(clippy::too_many_arguments)] pub fn qmatvec_view(
6508 &self,
6509 w: &CudaSlice<u8>,
6510 range: std::ops::Range<usize>,
6511 x: &cudarc::driver::CudaView<f32>,
6512 m: usize,
6513 in_f: usize,
6514 out_f: usize,
6515 qtype: i32,
6516 row_bytes: usize,
6517 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6518 self.qmatvec_view_inner(w, range, x, m, in_f, out_f, qtype, row_bytes)
6519 }
6520
6521 #[allow(clippy::too_many_arguments)]
6525 pub fn qmatvec_view_bf16_activation(
6526 &self,
6527 w: &CudaSlice<u8>,
6528 range: std::ops::Range<usize>,
6529 x: &cudarc::driver::CudaView<f32>,
6530 m: usize,
6531 in_f: usize,
6532 out_f: usize,
6533 qtype: i32,
6534 row_bytes: usize,
6535 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6536 let n = m * in_f;
6537 if x.len() != n {
6538 return Err(format!(
6539 "W4A16 BF16 activation input length {} != {m}x{in_f}",
6540 x.len()
6541 )
6542 .into());
6543 }
6544 let mut x_bf16 = self.alloc_u8_uninit(n * 2)?;
6545 self.f32_to_bf16_v(x, &mut x_bf16, n)?;
6546 let x_f32 = self.bf16_to_f32(&x_bf16.slice(0..n * 2), n)?;
6547 self.qmatvec_view_inner(
6548 w,
6549 range,
6550 &x_f32.slice(0..n),
6551 m,
6552 in_f,
6553 out_f,
6554 qtype,
6555 row_bytes,
6556 )
6557 }
6558
6559 #[allow(clippy::too_many_arguments)]
6560 fn qmatvec_view_inner(
6561 &self,
6562 w: &CudaSlice<u8>,
6563 range: std::ops::Range<usize>,
6564 x: &cudarc::driver::CudaView<f32>,
6565 m: usize,
6566 in_f: usize,
6567 out_f: usize,
6568 qtype: i32,
6569 row_bytes: usize,
6570 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6571 let f = self.func("qmatvec_f32");
6572 let wv = w.slice(range); let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
6575 grid_dim: (out_f as u32, m as u32, 1),
6576 block_dim: (256, 1, 1),
6577 shared_mem_bytes: 0,
6578 };
6579 let (inf, outf, mi, qt, rb) =
6580 (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
6581 let __s_b = self.gpu.stream();
6582 let mut b = __s_b.launch_builder(&f);
6583 b.arg(&wv)
6584 .arg(x)
6585 .arg(&mut y)
6586 .arg(&inf)
6587 .arg(&outf)
6588 .arg(&mi)
6589 .arg(&qt)
6590 .arg(&rb);
6591 unsafe {
6592 b.launch(cfg)?;
6593 }
6594 Ok(y)
6595 }
6596
6597 #[allow(clippy::too_many_arguments)]
6604 pub fn moe_gate_up_silu8_q8(
6608 &self,
6609 gp: WPtr8,
6610 up: WPtr8,
6611 aq: &CudaSlice<i8>,
6612 ad: &CudaSlice<f32>,
6613 in_f: usize,
6614 n_ff: usize,
6615 n_used: usize,
6616 qt_g: i32,
6617 qt_u: i32,
6618 rb_g: usize,
6619 rb_u: usize,
6620 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6621 let f = self.func("moe_gate_up_silu8_q8");
6622 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
6623 let cfg = LaunchConfig {
6624 grid_dim: (n_ff as u32, n_used as u32, 1),
6625 block_dim: (32, 1, 1),
6626 shared_mem_bytes: 0,
6627 };
6628 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
6629 let __s_b = self.gpu.stream();
6630 let mut b = __s_b.launch_builder(&f);
6631 b.arg(&gp)
6632 .arg(&up)
6633 .arg(aq)
6634 .arg(ad)
6635 .arg(&mut act)
6636 .arg(&inf)
6637 .arg(&nff)
6638 .arg(&qt_g)
6639 .arg(&qt_u)
6640 .arg(&rbg)
6641 .arg(&rbu);
6642 unsafe {
6643 b.launch(cfg)?;
6644 }
6645 Ok(act)
6646 }
6647
6648 #[allow(clippy::too_many_arguments)]
6661 pub fn moe_gate_up_preclamp8_q8(
6662 &self,
6663 gp: WPtr8,
6664 up: WPtr8,
6665 aq: &CudaSlice<i8>,
6666 ad: &CudaSlice<f32>,
6667 gs: F32x8,
6668 us: F32x8,
6669 limit: f32,
6670 in_f: usize,
6671 n_ff: usize,
6672 n_used: usize,
6673 qt_g: i32,
6674 qt_u: i32,
6675 rb_g: usize,
6676 rb_u: usize,
6677 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6678 debug_assert!(
6679 limit > 1e-6,
6680 "moe_gate_up_preclamp8_q8 needs a live limit; use moe_gate_up_silu8_q8"
6681 );
6682 let f = self.func("moe_gate_up_preclamp8_q8");
6683 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
6684 let cfg = LaunchConfig {
6685 grid_dim: (n_ff as u32, n_used as u32, 1),
6686 block_dim: (32, 1, 1),
6687 shared_mem_bytes: 0,
6688 };
6689 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
6690 let __s_b = self.gpu.stream();
6691 let mut b = __s_b.launch_builder(&f);
6692 b.arg(&gp)
6693 .arg(&up)
6694 .arg(aq)
6695 .arg(ad)
6696 .arg(&gs)
6697 .arg(&us)
6698 .arg(&limit)
6699 .arg(&mut act)
6700 .arg(&inf)
6701 .arg(&nff)
6702 .arg(&qt_g)
6703 .arg(&qt_u)
6704 .arg(&rbg)
6705 .arg(&rbu);
6706 unsafe {
6707 b.launch(cfg)?;
6708 }
6709 Ok(act)
6710 }
6711
6712 #[allow(clippy::too_many_arguments)]
6713 pub fn moe_down8_fma_q8(
6714 &self,
6715 dp: WPtr8,
6716 w: F32x8,
6717 aq2: &CudaSlice<i8>,
6718 ad2: &CudaSlice<f32>,
6719 dst: &mut cudarc::driver::CudaViewMut<f32>,
6720 in_f: usize,
6721 out_f: usize,
6722 n_used: usize,
6723 qt: i32,
6724 rb: usize,
6725 ) -> Result<(), Box<dyn std::error::Error>> {
6726 let f = self.func("moe_down8_fma_q8");
6727 let cfg = LaunchConfig {
6728 grid_dim: (out_f as u32, 1, 1),
6729 block_dim: (32, 1, 1),
6730 shared_mem_bytes: 0,
6731 };
6732 let (inf, outf, nu, rbi) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
6733 let __s_b = self.gpu.stream();
6734 let mut b = __s_b.launch_builder(&f);
6735 b.arg(&dp)
6736 .arg(&w)
6737 .arg(aq2)
6738 .arg(ad2)
6739 .arg(dst)
6740 .arg(&inf)
6741 .arg(&outf)
6742 .arg(&nu)
6743 .arg(&qt)
6744 .arg(&rbi);
6745 unsafe {
6746 b.launch(cfg)?;
6747 }
6748 Ok(())
6749 }
6750
6751 #[allow(clippy::too_many_arguments)]
6763 pub fn moe_vrows_tables_from_sel(
6765 &self,
6766 sel: &CudaSlice<i32>,
6767 selw: &CudaSlice<f32>,
6768 il: u16,
6769 macros: Option<(&[f32], &[f32], &[f32])>,
6770 (pg, pu, pd): (u64, u64, u64),
6771 (sg, su, sd): (usize, usize, usize),
6772 n_pairs: usize,
6773 ptrs: &mut CudaSlice<u64>,
6774 scl: &mut CudaSlice<f32>,
6775 ) -> Result<(), Box<dyn std::error::Error>> {
6776 debug_assert!(sel.len() >= n_pairs && selw.len() >= n_pairs);
6777 debug_assert!(ptrs.len() >= 3 * n_pairs);
6779 debug_assert_eq!(scl.len(), 3 * n_pairs);
6780 let mut mac = self
6783 .vrows_macro_dev
6784 .lock()
6785 .map_err(|_| "vrows macro mirror map is poisoned")?;
6786 if let Some((hg, hu, hd)) = macros {
6787 for (plane, host) in [(0u8, hg), (1u8, hu), (2u8, hd)] {
6788 if let std::collections::hash_map::Entry::Vacant(slot) = mac.entry((il, plane)) {
6791 slot.insert(self.htod(host)?);
6792 }
6793 }
6794 }
6795 let (mg, mu, md, have) = match macros {
6798 Some(_) => (
6799 mac.get(&(il, 0)).expect("gate macro mirror built above"),
6800 mac.get(&(il, 1)).expect("up macro mirror built above"),
6801 mac.get(&(il, 2)).expect("down macro mirror built above"),
6802 1i32,
6803 ),
6804 None => (selw, selw, selw, 0i32),
6805 };
6806 let f = self.func("moe_vrows_tables_from_sel");
6807 let threads = 128u32;
6808 let cfg = LaunchConfig {
6809 grid_dim: ((n_pairs as u32).div_ceil(threads), 1, 1),
6810 block_dim: (threads, 1, 1),
6811 shared_mem_bytes: 0,
6812 };
6813 let (sgi, sui, sdi) = (sg as i64, su as i64, sd as i64);
6814 let (npi, havei) = (n_pairs as i32, have);
6815 let __s_b = self.gpu.stream();
6816 let mut b = __s_b.launch_builder(&f);
6817 b.arg(sel)
6818 .arg(selw)
6819 .arg(mg)
6820 .arg(mu)
6821 .arg(md)
6822 .arg(&mut *ptrs)
6823 .arg(&mut *scl)
6824 .arg(&pg)
6825 .arg(&pu)
6826 .arg(&pd)
6827 .arg(&sgi)
6828 .arg(&sui)
6829 .arg(&sdi)
6830 .arg(&npi)
6831 .arg(&havei);
6832 unsafe {
6833 b.launch(cfg)?;
6834 }
6835 Ok(())
6836 }
6837
6838 pub fn moe_vrows_order_from_sel(
6850 &self,
6851 sel: &CudaSlice<i32>,
6852 n_pairs: usize,
6853 ptrs: &mut CudaSlice<u64>,
6854 ) -> Result<(), Box<dyn std::error::Error>> {
6855 debug_assert!(sel.len() >= n_pairs);
6856 debug_assert!(
6857 ptrs.len() >= 4 * n_pairs,
6858 "the order plane lives at ptrs[3*n_pairs .. 4*n_pairs)"
6859 );
6860 let f = self.func("moe_vrows_order_from_sel");
6861 let threads = 128u32;
6862 let cfg = LaunchConfig {
6863 grid_dim: ((n_pairs as u32).div_ceil(threads), 1, 1),
6864 block_dim: (threads, 1, 1),
6865 shared_mem_bytes: 0,
6866 };
6867 let np = n_pairs as i32;
6868 let __s_b = self.gpu.stream();
6869 let mut b = __s_b.launch_builder(&f);
6870 b.arg(sel).arg(&mut *ptrs).arg(&np);
6871 unsafe {
6872 b.launch(cfg)?;
6873 }
6874 Ok(())
6875 }
6876
6877 #[allow(clippy::too_many_arguments)]
6883 pub fn moe_gate_up_preclamp8_q8_rows(
6885 &self,
6886 ptrs: &CudaSlice<u64>,
6887 scl: &CudaSlice<f32>,
6888 aq: &CudaSlice<i8>,
6889 ad: &CudaSlice<f32>,
6890 limit: f32,
6891 in_f: usize,
6892 n_ff: usize,
6893 n_used: usize,
6894 n_pairs: usize,
6895 qt_g: i32,
6896 qt_u: i32,
6897 rb_g: usize,
6898 rb_u: usize,
6899 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6900 debug_assert!(
6901 limit > 1e-6,
6902 "moe_gate_up_preclamp8_q8_rows needs a live limit; the kernel collapses every gate \
6903 to silu(0) at limit 0"
6904 );
6905 debug_assert!(ptrs.len() >= 3 * n_pairs);
6906 debug_assert_eq!(scl.len(), 3 * n_pairs);
6907 let packed = moe_vrows_pack_on();
6915 let ordered =
6916 !packed && moe_vrows_dedup_order_on() && ptrs.len() >= 4 * n_pairs && n_ff <= 65535;
6917 let (f, cfg) = if packed {
6918 if MOE_VROWS_PACK_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
6919 eprintln!(
6920 "[moe-vrows-pack] engaged: 4-warp blocks on the verify-rows MoE pair \
6921 (MEMRA_MOE_VROWS_PACK=1)"
6922 );
6923 }
6924 (
6925 self.func("moe_gate_up_preclamp8_q8_rows_w4"),
6926 LaunchConfig {
6927 grid_dim: ((n_ff as u32).div_ceil(4), n_pairs as u32, 1),
6928 block_dim: (32, 4, 1),
6929 shared_mem_bytes: 0,
6930 },
6931 )
6932 } else if ordered {
6933 if MOE_VROWS_DEDUP_ORDER_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
6934 == 0
6935 {
6936 eprintln!(
6937 "[moe-vrows-dedup-order] engaged: verify-rows gate/up walks the pair union \
6938 EXPERT-MAJOR with the pair index as the fastest grid dimension, so the \
6939 21.96%-measured repeat visits read a shared expert slab's rows in adjacent \
6940 blocks (MEMRA_MOE_VROWS_DEDUP_ORDER=1)"
6941 );
6942 }
6943 (
6944 self.func("moe_gate_up_preclamp8_q8_rows_ord"),
6945 LaunchConfig {
6946 grid_dim: (n_pairs as u32, n_ff as u32, 1),
6947 block_dim: (32, 1, 1),
6948 shared_mem_bytes: 0,
6949 },
6950 )
6951 } else {
6952 (
6953 self.func("moe_gate_up_preclamp8_q8_rows"),
6954 LaunchConfig {
6955 grid_dim: (n_ff as u32, n_pairs as u32, 1),
6956 block_dim: (32, 1, 1),
6957 shared_mem_bytes: 0,
6958 },
6959 )
6960 };
6961 let mut act = self.vws_uninit(n_pairs * n_ff)?;
6963 let (inf, nff, nu, np) = (in_f as i32, n_ff as i32, n_used as i32, n_pairs as i32);
6964 let (rbg, rbu) = (rb_g as i64, rb_u as i64);
6965 let __s_b = self.gpu.stream();
6966 let mut b = __s_b.launch_builder(&f);
6967 b.arg(ptrs)
6968 .arg(scl)
6969 .arg(aq)
6970 .arg(ad)
6971 .arg(&limit)
6972 .arg(&mut act)
6973 .arg(&inf)
6974 .arg(&nff)
6975 .arg(&nu)
6976 .arg(&np)
6977 .arg(&qt_g)
6978 .arg(&qt_u)
6979 .arg(&rbg)
6980 .arg(&rbu);
6981 unsafe {
6982 b.launch(cfg)?;
6983 }
6984 Ok(act)
6985 }
6986
6987 #[allow(clippy::too_many_arguments)]
6991 pub fn moe_down8_fma_q8_rows(
6993 &self,
6994 ptrs: &CudaSlice<u64>,
6995 scl: &CudaSlice<f32>,
6996 aq2: &CudaSlice<i8>,
6997 ad2: &CudaSlice<f32>,
6998 dst: &mut CudaSlice<f32>,
6999 in_f: usize,
7000 out_f: usize,
7001 n_used: usize,
7002 n_pairs: usize,
7003 qt: i32,
7004 rb: usize,
7005 ) -> Result<(), Box<dyn std::error::Error>> {
7006 debug_assert!(ptrs.len() >= 3 * n_pairs);
7007 debug_assert_eq!(scl.len(), 3 * n_pairs);
7008 debug_assert_eq!(n_pairs % n_used, 0, "pairs are dense slot-major");
7009 let t = n_pairs / n_used;
7010 debug_assert!(dst.len() >= t * out_f);
7011 let packed = moe_vrows_pack_on();
7013 let tmaj = !packed && moe_vrows_down_tmaj_on() && out_f <= 65535;
7019 let (f, cfg) = if packed {
7020 (
7021 self.func("moe_down8_fma_q8_rows_w4"),
7022 LaunchConfig {
7023 grid_dim: ((out_f as u32).div_ceil(4), t as u32, 1),
7024 block_dim: (32, 4, 1),
7025 shared_mem_bytes: 0,
7026 },
7027 )
7028 } else if tmaj {
7029 if MOE_VROWS_DOWN_TMAJ_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
7030 == 0
7031 {
7032 eprintln!(
7033 "[moe-vrows-down-tmaj] engaged: verify-rows down/FMA grid transposed to \
7034 (t, out_f) so the verify rows at one output row are adjacent blocks; the \
7035 slot-ordered FMA chain is unchanged (MEMRA_MOE_VROWS_DOWN_TMAJ=1)"
7036 );
7037 }
7038 (
7039 self.func("moe_down8_fma_q8_rows_tmaj"),
7040 LaunchConfig {
7041 grid_dim: (t as u32, out_f as u32, 1),
7042 block_dim: (32, 1, 1),
7043 shared_mem_bytes: 0,
7044 },
7045 )
7046 } else {
7047 (
7048 self.func("moe_down8_fma_q8_rows"),
7049 LaunchConfig {
7050 grid_dim: (out_f as u32, t as u32, 1),
7051 block_dim: (32, 1, 1),
7052 shared_mem_bytes: 0,
7053 },
7054 )
7055 };
7056 let (inf, outf, nu, np, rbi) = (
7057 in_f as i32,
7058 out_f as i32,
7059 n_used as i32,
7060 n_pairs as i32,
7061 rb as i64,
7062 );
7063 let __s_b = self.gpu.stream();
7064 let mut b = __s_b.launch_builder(&f);
7065 b.arg(ptrs)
7066 .arg(scl)
7067 .arg(aq2)
7068 .arg(ad2)
7069 .arg(dst)
7070 .arg(&inf)
7071 .arg(&outf)
7072 .arg(&nu)
7073 .arg(&np)
7074 .arg(&qt)
7075 .arg(&rbi);
7076 unsafe {
7077 b.launch(cfg)?;
7078 }
7079 Ok(())
7080 }
7081
7082 #[allow(clippy::too_many_arguments)]
7084 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_expert_q8(
7087 &self,
7088 w: &CudaSlice<u8>,
7089 range: std::ops::Range<usize>,
7090 aq: &CudaSlice<i8>,
7091 ad: &CudaSlice<f32>,
7092 m: usize,
7093 in_f: usize,
7094 out_f: usize,
7095 qtype: i32,
7096 row_bytes: usize,
7097 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7098 let f = self.func("qmatvec_expert_q8");
7099 let wv = w.slice(range);
7100 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7101 const ROWS: u32 = 4; let cfg = LaunchConfig {
7103 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, m as u32, 1),
7104 block_dim: (32, ROWS, 1),
7105 shared_mem_bytes: 0,
7106 };
7107 let (inf, outf, mi, rbi) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7108 let __s_b = self.gpu.stream();
7109 let mut b = __s_b.launch_builder(&f);
7110 b.arg(&wv)
7111 .arg(aq)
7112 .arg(ad)
7113 .arg(&mut y)
7114 .arg(&inf)
7115 .arg(&outf)
7116 .arg(&mi)
7117 .arg(&qtype)
7118 .arg(&rbi);
7119 unsafe {
7120 b.launch(cfg)?;
7121 }
7122 Ok(y)
7123 }
7124
7125 #[allow(clippy::too_many_arguments)] pub fn moe_gate_up_silu8(
7127 &self,
7128 gp: WPtr8,
7129 up: WPtr8,
7130 x: &cudarc::driver::CudaView<f32>,
7131 in_f: usize,
7132 n_ff: usize,
7133 n_used: usize,
7134 qt_g: i32,
7135 qt_u: i32,
7136 rb_g: usize,
7137 rb_u: usize,
7138 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7139 let f = self.func("moe_gate_up_silu8_f32");
7140 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig {
7142 grid_dim: (n_ff as u32, n_used as u32, 1),
7143 block_dim: (256, 1, 1),
7144 shared_mem_bytes: 0,
7145 };
7146 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
7147 let __s_b = self.gpu.stream();
7148 let mut b = __s_b.launch_builder(&f);
7149 b.arg(&gp)
7150 .arg(&up)
7151 .arg(x)
7152 .arg(&mut act)
7153 .arg(&inf)
7154 .arg(&nff)
7155 .arg(&qt_g)
7156 .arg(&qt_u)
7157 .arg(&rbg)
7158 .arg(&rbu);
7159 unsafe {
7160 b.launch(cfg)?;
7161 }
7162 Ok(act)
7163 }
7164
7165 #[allow(clippy::too_many_arguments)]
7171 pub fn moe_down8_fma_into(
7172 &self,
7173 dp: WPtr8,
7174 w: F32x8,
7175 act: &CudaSlice<f32>,
7176 dst: &mut cudarc::driver::CudaViewMut<f32>,
7177 in_f: usize,
7178 out_f: usize,
7179 n_used: usize,
7180 qt: i32,
7181 rb: usize,
7182 ) -> Result<(), Box<dyn std::error::Error>> {
7183 let f = self.func("moe_down8_fma_f32");
7184 let cfg = LaunchConfig {
7185 grid_dim: (out_f as u32, 1, 1),
7186 block_dim: (256, 1, 1),
7187 shared_mem_bytes: 0,
7188 };
7189 let (inf, outf, nu, rbv) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
7190 let __s_b = self.gpu.stream();
7191 let mut b = __s_b.launch_builder(&f);
7192 b.arg(&dp)
7193 .arg(&w)
7194 .arg(act)
7195 .arg(dst)
7196 .arg(&inf)
7197 .arg(&outf)
7198 .arg(&nu)
7199 .arg(&qt)
7200 .arg(&rbv);
7201 unsafe {
7202 b.launch(cfg)?;
7203 }
7204 Ok(())
7205 }
7206
7207 #[allow(clippy::too_many_arguments)]
7212 #[allow(clippy::too_many_arguments)]
7227 #[allow(clippy::too_many_arguments)]
7229 #[allow(clippy::manual_div_ceil)] pub fn moe_pairs_matvec_q8(
7231 &self,
7232 table: &CudaSlice<u64>,
7233 proj: i32,
7234 pair_tok: &CudaSlice<i32>,
7235 pair_ex: &CudaSlice<i32>,
7236 aq: &CudaSlice<i8>,
7237 ad: &CudaSlice<f32>,
7238 in_f: usize,
7239 out_f: usize,
7240 n_expert: usize,
7241 n_pairs: usize,
7242 qtype: i32,
7243 row_bytes: usize,
7244 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7245 let f = self.func("moe_pairs_matvec_q8");
7246 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
7247 const ROWS: u32 = 4;
7248 let cfg = LaunchConfig {
7249 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_pairs as u32, 1),
7250 block_dim: (32, ROWS, 1),
7251 shared_mem_bytes: 0,
7252 };
7253 let (inf, outf, ne, np, rbi) = (
7254 in_f as i32,
7255 out_f as i32,
7256 n_expert as i32,
7257 n_pairs as i32,
7258 row_bytes as i64,
7259 );
7260 let __s_b = self.gpu.stream();
7261 let mut b = __s_b.launch_builder(&f);
7262 b.arg(table)
7263 .arg(&proj)
7264 .arg(pair_tok)
7265 .arg(pair_ex)
7266 .arg(aq)
7267 .arg(ad)
7268 .arg(&mut y)
7269 .arg(&inf)
7270 .arg(&outf)
7271 .arg(&ne)
7272 .arg(&np)
7273 .arg(&qtype)
7274 .arg(&rbi);
7275 unsafe {
7276 b.launch(cfg)?;
7277 }
7278 Ok(y)
7279 }
7280
7281 #[allow(clippy::too_many_arguments)]
7283 #[allow(clippy::manual_div_ceil)] pub fn moe_pairs_matvec_q8_em(
7285 &self,
7286 table: &CudaSlice<u64>,
7287 proj: i32,
7288 ex_ids: &CudaSlice<i32>,
7289 ex_off: &CudaSlice<i32>,
7290 ex_pairs: &CudaSlice<i32>,
7291 pair_tok: &CudaSlice<i32>,
7292 aq: &CudaSlice<i8>,
7293 ad: &CudaSlice<f32>,
7294 in_f: usize,
7295 out_f: usize,
7296 n_expert: usize,
7297 n_active: usize,
7298 n_pairs: usize,
7299 qtype: i32,
7300 row_bytes: usize,
7301 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7302 let f = self.func("moe_pairs_matvec_q8_em");
7303 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
7304 const ROWS: u32 = 4;
7305 let cfg = LaunchConfig {
7306 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
7307 block_dim: (32, ROWS, 1),
7308 shared_mem_bytes: 0,
7309 };
7310 let (inf, outf, ne, na, rbi) = (
7311 in_f as i32,
7312 out_f as i32,
7313 n_expert as i32,
7314 n_active as i32,
7315 row_bytes as i64,
7316 );
7317 let __s_b = self.gpu.stream();
7318 let mut b = __s_b.launch_builder(&f);
7319 b.arg(table)
7320 .arg(&proj)
7321 .arg(ex_ids)
7322 .arg(ex_off)
7323 .arg(ex_pairs)
7324 .arg(pair_tok)
7325 .arg(aq)
7326 .arg(ad)
7327 .arg(&mut y)
7328 .arg(&inf)
7329 .arg(&outf)
7330 .arg(&ne)
7331 .arg(&na)
7332 .arg(&qtype)
7333 .arg(&rbi);
7334 unsafe {
7335 b.launch(cfg)?;
7336 }
7337 Ok(y)
7338 }
7339
7340 #[allow(clippy::too_many_arguments)]
7343 #[allow(clippy::manual_div_ceil)] pub fn moe_pairs_matvec_q8_dec(
7345 &self,
7346 table: &CudaSlice<u64>,
7347 proj: i32,
7348 ex_ids: &CudaSlice<i32>,
7349 ex_off: &CudaSlice<i32>,
7350 ex_pairs: &CudaSlice<i32>,
7351 pair_tok: &CudaSlice<i32>,
7352 aq: &CudaSlice<i8>,
7353 ad: &CudaSlice<f32>,
7354 in_f: usize,
7355 out_f: usize,
7356 n_expert: usize,
7357 n_active: usize,
7358 n_pairs: usize,
7359 qtype: i32,
7360 row_bytes: usize,
7361 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7362 let f = self.func("moe_pairs_matvec_q8_dec");
7363 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
7364 const ROWS: u32 = 4;
7365 let cfg = LaunchConfig {
7366 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
7367 block_dim: (32, ROWS, 1),
7368 shared_mem_bytes: 0,
7369 };
7370 let (inf, outf, ne, na, rbi) = (
7371 in_f as i32,
7372 out_f as i32,
7373 n_expert as i32,
7374 n_active as i32,
7375 row_bytes as i64,
7376 );
7377 let __s_b = self.gpu.stream();
7378 let mut b = __s_b.launch_builder(&f);
7379 b.arg(table)
7380 .arg(&proj)
7381 .arg(ex_ids)
7382 .arg(ex_off)
7383 .arg(ex_pairs)
7384 .arg(pair_tok)
7385 .arg(aq)
7386 .arg(ad)
7387 .arg(&mut y)
7388 .arg(&inf)
7389 .arg(&outf)
7390 .arg(&ne)
7391 .arg(&na)
7392 .arg(&qtype)
7393 .arg(&rbi);
7394 unsafe {
7395 b.launch(cfg)?;
7396 }
7397 Ok(y)
7398 }
7399
7400 pub fn moe_pairs_gelu_mul(
7401 &self,
7402 gate: &CudaSlice<f32>,
7403 up: &CudaSlice<f32>,
7404 n: usize,
7405 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7406 let f = self.func("moe_pairs_gelu_mul");
7407 let mut act = self.alloc_uninit::<f32>(n)?;
7408 let cfg = LaunchConfig::for_num_elems(n as u32);
7409 let nl = n as i64;
7410 let __s_b = self.gpu.stream();
7411 let mut b = __s_b.launch_builder(&f);
7412 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
7413 unsafe {
7414 b.launch(cfg)?;
7415 }
7416 Ok(act)
7417 }
7418
7419 pub fn moe_pairs_silu_mul(
7420 &self,
7421 gate: &CudaSlice<f32>,
7422 up: &CudaSlice<f32>,
7423 n: usize,
7424 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7425 let f = self.func("moe_pairs_silu_mul");
7426 let mut act = self.alloc_uninit::<f32>(n)?;
7427 let cfg = LaunchConfig::for_num_elems(n as u32);
7428 let nl = n as i64;
7429 let __s_b = self.gpu.stream();
7430 let mut b = __s_b.launch_builder(&f);
7431 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
7432 unsafe {
7433 b.launch(cfg)?;
7434 }
7435 Ok(act)
7436 }
7437
7438 #[allow(clippy::too_many_arguments)]
7439 #[allow(clippy::manual_div_ceil)] pub fn moe_pairs_scatter(
7441 &self,
7442 y_down: &CudaSlice<f32>,
7443 pair_w: &CudaSlice<f32>,
7444 tok_pair_off: &CudaSlice<i32>,
7445 tok_pair_ids: &CudaSlice<i32>,
7446 moe_out: &mut CudaSlice<f32>,
7447 t: usize,
7448 n_embd: usize,
7449 ) -> Result<(), Box<dyn std::error::Error>> {
7450 let f = self.func("moe_pairs_scatter");
7451 let cfg = LaunchConfig {
7452 grid_dim: (((n_embd + 255) / 256) as u32, t as u32, 1),
7453 block_dim: (256, 1, 1),
7454 shared_mem_bytes: 0,
7455 };
7456 let ne = n_embd as i32;
7457 let __s_b = self.gpu.stream();
7458 let mut b = __s_b.launch_builder(&f);
7459 b.arg(y_down)
7460 .arg(pair_w)
7461 .arg(tok_pair_off)
7462 .arg(tok_pair_ids)
7463 .arg(moe_out)
7464 .arg(&ne);
7465 unsafe {
7466 b.launch(cfg)?;
7467 }
7468 Ok(())
7469 }
7470
7471 #[allow(clippy::too_many_arguments)]
7475 pub fn moe_gate_up_gelu8_dev_q8(
7476 &self,
7477 table: &CudaSlice<u64>,
7478 sel: &cudarc::driver::CudaView<i32>,
7479 aq: &CudaSlice<i8>,
7480 ad: &CudaSlice<f32>,
7481 in_f: usize,
7482 n_ff: usize,
7483 n_used: usize,
7484 n_expert: usize,
7485 qt_g: i32,
7486 qt_u: i32,
7487 rb_g: usize,
7488 rb_u: usize,
7489 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7490 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
7491 let (inf, nff, ne, rbg, rbu) = (
7492 in_f as i32,
7493 n_ff as i32,
7494 n_expert as i32,
7495 rb_g as i64,
7496 rb_u as i64,
7497 );
7498 let f = self.func("moe_gate_up_gelu8_dev_q8");
7499 let cfg = LaunchConfig {
7500 grid_dim: (n_ff as u32, n_used as u32, 1),
7501 block_dim: (32, 1, 1),
7502 shared_mem_bytes: 0,
7503 };
7504 let __s_b = self.gpu.stream();
7505 let mut b = __s_b.launch_builder(&f);
7506 b.arg(table)
7507 .arg(sel)
7508 .arg(aq)
7509 .arg(ad)
7510 .arg(&mut act)
7511 .arg(&inf)
7512 .arg(&nff)
7513 .arg(&ne)
7514 .arg(&qt_g)
7515 .arg(&qt_u)
7516 .arg(&rbg)
7517 .arg(&rbu);
7518 unsafe {
7519 b.launch(cfg)?;
7520 }
7521 Ok(act)
7522 }
7523
7524 #[allow(clippy::too_many_arguments)]
7526 pub fn moe_gate_up_gelu8_dev_q8_rows(
7527 &self,
7528 table: &CudaSlice<u64>,
7529 sel: &CudaSlice<i32>,
7530 aq: &CudaSlice<i8>,
7531 ad: &CudaSlice<f32>,
7532 t: usize,
7533 in_f: usize,
7534 n_ff: usize,
7535 n_used: usize,
7536 n_expert: usize,
7537 qt_g: i32,
7538 qt_u: i32,
7539 rb_g: usize,
7540 rb_u: usize,
7541 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7542 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
7543 let (inf, nff, ne, rbg, rbu, nu) = (
7544 in_f as i32,
7545 n_ff as i32,
7546 n_expert as i32,
7547 rb_g as i64,
7548 rb_u as i64,
7549 n_used as i32,
7550 );
7551 let f = self.func("moe_gate_up_gelu8_dev_q8_rows");
7552 let cfg = LaunchConfig {
7553 grid_dim: (n_ff as u32, n_used as u32, t as u32),
7554 block_dim: (32, 1, 1),
7555 shared_mem_bytes: 0,
7556 };
7557 let __s_b = self.gpu.stream();
7558 let mut b = __s_b.launch_builder(&f);
7559 b.arg(table)
7560 .arg(sel)
7561 .arg(aq)
7562 .arg(ad)
7563 .arg(&mut act)
7564 .arg(&inf)
7565 .arg(&nff)
7566 .arg(&ne)
7567 .arg(&qt_g)
7568 .arg(&qt_u)
7569 .arg(&rbg)
7570 .arg(&rbu)
7571 .arg(&nu);
7572 unsafe {
7573 b.launch(cfg)?;
7574 }
7575 Ok(act)
7576 }
7577
7578 #[allow(clippy::too_many_arguments)]
7580 pub fn moe_gate_up_gelu8_dev_q8_csr(
7581 &self,
7582 table: &CudaSlice<u64>,
7583 sel: &CudaSlice<i32>,
7584 aq: &CudaSlice<i8>,
7585 ad: &CudaSlice<f32>,
7586 n_pairs: usize,
7587 in_f: usize,
7588 n_ff: usize,
7589 n_used: usize,
7590 n_expert: usize,
7591 qt_g: i32,
7592 qt_u: i32,
7593 rb_g: usize,
7594 rb_u: usize,
7595 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7596 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
7597 let (inf, nff, ne, rbg, rbu, nu, npi) = (
7598 in_f as i32,
7599 n_ff as i32,
7600 n_expert as i32,
7601 rb_g as i64,
7602 rb_u as i64,
7603 n_used as i32,
7604 n_pairs as i32,
7605 );
7606 let f = self.func("moe_gate_up_gelu8_dev_q8_csr");
7607 let cfg = LaunchConfig {
7608 grid_dim: (n_ff as u32, n_pairs as u32, 1),
7609 block_dim: (32, 1, 1),
7610 shared_mem_bytes: 0,
7611 };
7612 let __s_b = self.gpu.stream();
7613 let mut b = __s_b.launch_builder(&f);
7614 b.arg(table)
7615 .arg(sel)
7616 .arg(aq)
7617 .arg(ad)
7618 .arg(&mut act)
7619 .arg(&inf)
7620 .arg(&nff)
7621 .arg(&ne)
7622 .arg(&qt_g)
7623 .arg(&qt_u)
7624 .arg(&rbg)
7625 .arg(&rbu)
7626 .arg(&nu)
7627 .arg(&npi);
7628 unsafe {
7629 b.launch(cfg)?;
7630 }
7631 Ok(act)
7632 }
7633
7634 #[allow(clippy::too_many_arguments)]
7636 pub fn moe_down8_fma_dev_q8_rows_g(
7637 &self,
7638 table: &CudaSlice<u64>,
7639 sel: &CudaSlice<i32>,
7640 w: &CudaSlice<f32>,
7641 aq2: &CudaSlice<i8>,
7642 ad2: &CudaSlice<f32>,
7643 dst: &mut CudaSlice<f32>,
7644 t: usize,
7645 in_f: usize,
7646 out_f: usize,
7647 n_used: usize,
7648 n_expert: usize,
7649 qt: i32,
7650 rb: usize,
7651 ) -> Result<(), Box<dyn std::error::Error>> {
7652 let (inf, outf, nu, ne, rbi) = (
7653 in_f as i32,
7654 out_f as i32,
7655 n_used as i32,
7656 n_expert as i32,
7657 rb as i64,
7658 );
7659 let step_b1_w8 = t == 1 && in_f == 1280 && out_f == 4096 && n_used == 8 && qt == QT_IQ4_XS;
7663 let f = self.func(if step_b1_w8 {
7664 "moe_down8_fma_dev_q8_rows_w8"
7665 } else {
7666 "moe_down8_fma_dev_q8_rows_g"
7667 });
7668 let cfg = LaunchConfig {
7669 grid_dim: (out_f as u32, 1, t as u32),
7670 block_dim: (32, if step_b1_w8 { 8 } else { 1 }, 1),
7671 shared_mem_bytes: 0,
7672 };
7673 let __s_b = self.gpu.stream();
7674 let mut b = __s_b.launch_builder(&f);
7675 b.arg(table)
7676 .arg(sel)
7677 .arg(w)
7678 .arg(aq2)
7679 .arg(ad2)
7680 .arg(dst)
7681 .arg(&inf)
7682 .arg(&outf)
7683 .arg(&nu)
7684 .arg(&ne)
7685 .arg(&qt)
7686 .arg(&rbi);
7687 unsafe {
7688 b.launch(cfg)?;
7689 }
7690 Ok(())
7691 }
7692
7693 pub fn rp_probe_q4(&self, m: usize) -> Result<(f64, f64), Box<dyn std::error::Error>> {
7697 let (out_f, in_f) = (2048usize, 2816usize);
7698 let nblk = in_f / 32;
7699 let mut seed = 0x9E3779B97F4A7C15u64;
7700 let mut rng = move || {
7701 seed = seed
7702 .wrapping_mul(6364136223846793005)
7703 .wrapping_add(1442695040888963407);
7704 (seed >> 33) as u8
7705 };
7706 let mut w = vec![0u8; out_f * nblk * 18];
7707 for b in w.iter_mut() {
7708 *b = rng();
7709 }
7710 for r in 0..out_f {
7711 for g in 0..nblk {
7712 let off = (r * nblk + g) * 18;
7713 w[off] = 0x00;
7714 w[off + 1] = 0x2C; }
7716 }
7717 let qplane = out_f * nblk * 16;
7718 let mut wrp = vec![0u8; w.len()];
7719 for r in 0..out_f {
7720 for g in 0..nblk {
7721 let src = &w[(r * nblk + g) * 18..(r * nblk + g) * 18 + 18];
7722 wrp[qplane + (r * nblk + g) * 2..qplane + (r * nblk + g) * 2 + 2]
7723 .copy_from_slice(&src[0..2]);
7724 wrp[(r * nblk + g) * 16..(r * nblk + g) * 16 + 16].copy_from_slice(&src[2..18]);
7725 }
7726 }
7727 let w_d = self.htod_bytes(&w)?;
7728 let wrp_d = self.htod_bytes(&wrp)?;
7729 let mut aq = vec![0i8; m * in_f];
7730 for v in aq.iter_mut() {
7731 *v = rng() as i8;
7732 }
7733 let aq_d = self.htod_i8(&aq)?;
7734 let ad_d = self.htod(&vec![0.03125f32; m * nblk])?;
7735 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
7736 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
7737 const RPB: u32 = 4;
7738 let cfg = LaunchConfig {
7739 grid_dim: ((out_f as u32).div_ceil(RPB), 1, 1),
7740 block_dim: (32, RPB, 1),
7741 shared_mem_bytes: 0,
7742 };
7743 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
7744 let (rb, qp) = ((nblk * 18) as i64, qplane as i64);
7745 let fb = self.func("qmatvec_q4_0_mmvq_b4");
7746 let fr = self.func("qmatvec_q4_0_mmvq_b4_rp");
7747 {
7748 let __s_b = self.gpu.stream();
7749 let mut b = __s_b.launch_builder(&fb);
7750 b.arg(&w_d)
7751 .arg(&aq_d)
7752 .arg(&ad_d)
7753 .arg(&mut y0)
7754 .arg(&inf)
7755 .arg(&outf)
7756 .arg(&mi)
7757 .arg(&rb);
7758 unsafe {
7759 b.launch(cfg)?;
7760 }
7761 let __s_b = self.gpu.stream();
7762 let mut b = __s_b.launch_builder(&fr);
7763 b.arg(&wrp_d)
7764 .arg(&aq_d)
7765 .arg(&ad_d)
7766 .arg(&mut y1)
7767 .arg(&inf)
7768 .arg(&outf)
7769 .arg(&mi)
7770 .arg(&qp);
7771 unsafe {
7772 b.launch(cfg)?;
7773 }
7774 }
7775 self.gpu.stream().synchronize()?;
7776 let (h0, h1) = (self.dtoh(&y0)?, self.dtoh(&y1)?);
7777 let nd = h0
7778 .iter()
7779 .zip(&h1)
7780 .filter(|(a, b)| a.to_bits() != b.to_bits())
7781 .count();
7782 if nd != 0 {
7783 return Err(format!("rp twin not bitwise: {nd}/{} diffs", h0.len()).into());
7784 }
7785 let mut time = |rp: bool| -> Result<f64, Box<dyn std::error::Error>> {
7786 self.gpu.stream().synchronize()?;
7787 let t0 = std::time::Instant::now();
7788 for _ in 0..500 {
7789 if rp {
7790 let __s_b = self.gpu.stream();
7791 let mut b = __s_b.launch_builder(&fr);
7792 b.arg(&wrp_d)
7793 .arg(&aq_d)
7794 .arg(&ad_d)
7795 .arg(&mut y1)
7796 .arg(&inf)
7797 .arg(&outf)
7798 .arg(&mi)
7799 .arg(&qp);
7800 unsafe {
7801 b.launch(cfg)?;
7802 }
7803 } else {
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 }
7818 }
7819 self.gpu.stream().synchronize()?;
7820 Ok(t0.elapsed().as_secs_f64() * 1e6 / 500.0)
7821 };
7822 let _ = time(false)?;
7823 let _ = time(true)?; Ok((time(false)?, time(true)?))
7825 }
7826
7827 pub fn build_q4_rp4(
7832 &self,
7833 t: &mut crate::model::GpuTensor,
7834 ) -> Result<(), Box<dyn std::error::Error>> {
7835 use crate::model::GpuTensor;
7836 let GpuTensor::Quant {
7837 bytes,
7838 qtype,
7839 row_bytes,
7840 ne,
7841 rp4,
7842 ..
7843 } = t
7844 else {
7845 return Ok(());
7846 };
7847 if *qtype != QT_Q4_0 || rp4.is_some() || ne.len() != 2 {
7848 return Ok(());
7849 }
7850 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
7851 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 18 {
7852 return Ok(());
7853 }
7854 let nblk = in_f / 32;
7855 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 18)?;
7856 let f = self.func("q4_0_split_rp_build");
7857 let n = (out_f * nblk) as i32;
7858 let cfg = LaunchConfig {
7859 grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
7860 block_dim: (256, 1, 1),
7861 shared_mem_bytes: 0,
7862 };
7863 let (of, nb) = (out_f as i32, nblk as i32);
7864 let _ = n;
7865 let __s_b = self.gpu.stream();
7866 let mut b = __s_b.launch_builder(&f);
7867 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
7868 unsafe {
7869 b.launch(cfg)?;
7870 }
7871 *rp4 = Some(dst);
7872 Ok(())
7873 }
7874
7875 pub fn build_q8_rp4(
7880 &self,
7881 t: &mut crate::model::GpuTensor,
7882 ) -> Result<(), Box<dyn std::error::Error>> {
7883 use crate::model::GpuTensor;
7884 let GpuTensor::Quant {
7885 bytes,
7886 qtype,
7887 row_bytes,
7888 ne,
7889 rp4,
7890 ..
7891 } = t
7892 else {
7893 return Ok(());
7894 };
7895 if *qtype != QT_Q8_0 || rp4.is_some() || ne.len() != 2 {
7896 return Ok(());
7897 }
7898 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
7899 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 34 {
7900 return Ok(());
7901 }
7902 *rp4 = Some(self.build_q8_rp4_raw(bytes, in_f, out_f)?);
7903 Ok(())
7904 }
7905
7906 pub fn build_q8_rp4_raw(
7909 &self,
7910 bytes: &CudaSlice<u8>,
7911 in_f: usize,
7912 out_f: usize,
7913 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
7914 assert!(in_f.is_multiple_of(32));
7915 let nblk = in_f / 32;
7916 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 34)?;
7917 let f = self.func("q8_0_split_rp_build");
7918 let cfg = LaunchConfig {
7919 grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
7920 block_dim: (256, 1, 1),
7921 shared_mem_bytes: 0,
7922 };
7923 let (of, nb) = (out_f as i32, nblk as i32);
7924 let __s_b = self.gpu.stream();
7925 let mut b = __s_b.launch_builder(&f);
7926 b.arg(bytes).arg(&mut dst).arg(&of).arg(&nb);
7927 unsafe {
7928 b.launch(cfg)?;
7929 }
7930 Ok(dst)
7931 }
7932
7933 pub fn build_q4k_rp4(
7941 &self,
7942 t: &mut crate::model::GpuTensor,
7943 ) -> Result<(), Box<dyn std::error::Error>> {
7944 use crate::model::GpuTensor;
7945 let GpuTensor::Quant {
7946 bytes,
7947 qtype,
7948 row_bytes,
7949 ne,
7950 rp4,
7951 ..
7952 } = t
7953 else {
7954 return Ok(());
7955 };
7956 if *qtype != QT_Q4_K || rp4.is_some() || ne.len() != 2 {
7957 return Ok(());
7958 }
7959 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
7960 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 144 {
7961 return Ok(());
7962 }
7963 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q4_K)?);
7964 Ok(())
7965 }
7966
7967 pub fn build_q6k_rp4(
7968 &self,
7969 t: &mut crate::model::GpuTensor,
7970 ) -> Result<(), Box<dyn std::error::Error>> {
7971 use crate::model::GpuTensor;
7972 let GpuTensor::Quant {
7973 bytes,
7974 qtype,
7975 row_bytes,
7976 ne,
7977 rp4,
7978 ..
7979 } = t
7980 else {
7981 return Ok(());
7982 };
7983 if *qtype != QT_Q6_K || rp4.is_some() || ne.len() != 2 {
7984 return Ok(());
7985 }
7986 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
7987 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 210 {
7988 return Ok(());
7989 }
7990 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q6_K)?);
7991 Ok(())
7992 }
7993
7994 pub fn build_kq_rp4_raw(
7996 &self,
7997 bytes: &CudaSlice<u8>,
7998 in_f: usize,
7999 out_f: usize,
8000 qtype: i32,
8001 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
8002 assert!(in_f.is_multiple_of(256));
8003 let nsbk = in_f / 256;
8004 let (sb_bytes, kname) = match qtype {
8005 QT_Q4_K => (144usize, "q4_K_split_rp_build"),
8006 QT_Q6_K => (210usize, "q6_K_split_rp_build"),
8007 _ => return Err(format!("build_kq_rp4_raw: qtype {qtype} has no rp mirror").into()),
8008 };
8009 let mut dst = self.alloc_uninit::<u8>(out_f * nsbk * sb_bytes)?;
8010 let f = self.func(kname);
8011 let cfg = LaunchConfig {
8012 grid_dim: (((out_f * nsbk) as u32).div_ceil(256), 1, 1),
8013 block_dim: (256, 1, 1),
8014 shared_mem_bytes: 0,
8015 };
8016 let (of, nb) = (out_f as i32, nsbk as i32);
8017 let __s_b = self.gpu.stream();
8018 let mut b = __s_b.launch_builder(&f);
8019 b.arg(bytes).arg(&mut dst).arg(&of).arg(&nb);
8020 unsafe {
8021 b.launch(cfg)?;
8022 }
8023 Ok(dst)
8024 }
8025
8026 pub fn kqrp_enabled() -> bool {
8030 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8031 *ON.get_or_init(|| match std::env::var("MEMRA_KQRP").as_deref() {
8032 Ok("0") => false,
8033 Ok(_) => true,
8034 Err(_) => cfg!(memra_hopper_mma),
8035 })
8036 }
8037
8038 pub fn build_q4_rp_swap(
8044 &self,
8045 t: &mut crate::model::GpuTensor,
8046 ) -> Result<bool, Box<dyn std::error::Error>> {
8047 use crate::model::GpuTensor;
8048 if !matches!(t, GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0) {
8058 return Ok(false);
8059 }
8060 self.build_q4_rp4(t)?;
8061 self.gpu.stream().synchronize()?; let GpuTensor::Quant { bytes, rp4, rp, .. } = t else {
8063 return Ok(false);
8064 };
8065 match rp4.take() {
8066 Some(split) => {
8067 *bytes = split; *rp = true;
8069 Ok(true)
8070 }
8071 None => Ok(false),
8072 }
8073 }
8074
8075 pub fn q4rp_enabled() -> bool {
8077 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8078 *ON.get_or_init(|| {
8079 std::env::var("MEMRA_Q4RP")
8080 .map(|v| v != "0")
8081 .unwrap_or(true)
8082 })
8083 }
8084
8085 #[allow(clippy::manual_div_ceil)] pub fn copy_rows_strided(
8089 &self,
8090 src: &CudaSlice<f32>,
8091 dst: &mut CudaSlice<f32>,
8092 row_elems: usize,
8093 n_rows: usize,
8094 src_stride: usize,
8095 src_off: usize,
8096 ) -> Result<(), Box<dyn std::error::Error>> {
8097 let f = self.func("copy_rows_strided_f32");
8098 let cfg = LaunchConfig {
8099 grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
8100 block_dim: (256, 1, 1),
8101 shared_mem_bytes: 0,
8102 };
8103 let (re, nr) = (row_elems as i32, n_rows as i32);
8104 let (st, off) = (src_stride as i64, src_off as i64);
8105 let __s_b = self.gpu.stream();
8106 let mut b = __s_b.launch_builder(&f);
8107 b.arg(src)
8108 .arg(&mut *dst)
8109 .arg(&re)
8110 .arg(&nr)
8111 .arg(&st)
8112 .arg(&off);
8113 unsafe {
8114 b.launch(cfg)?;
8115 }
8116 Ok(())
8117 }
8118
8119 #[allow(clippy::manual_div_ceil)] pub fn place_rows_strided(
8126 &self,
8127 src: &CudaSlice<f32>,
8128 dst: &mut CudaSlice<f32>,
8129 row_elems: usize,
8130 n_rows: usize,
8131 dst_stride: usize,
8132 dst_off: usize,
8133 ) -> Result<(), Box<dyn std::error::Error>> {
8134 if row_elems == 0 || n_rows == 0 {
8135 return Err("strided row placement requires nonzero rows and row width".into());
8136 }
8137 let src_len = n_rows
8138 .checked_mul(row_elems)
8139 .ok_or("strided row placement source size overflow")?;
8140 let dst_len = n_rows
8141 .checked_sub(1)
8142 .and_then(|rows| rows.checked_mul(dst_stride))
8143 .and_then(|base| base.checked_add(dst_off))
8144 .and_then(|base| base.checked_add(row_elems))
8145 .ok_or("strided row placement destination size overflow")?;
8146 let row_end = dst_off
8147 .checked_add(row_elems)
8148 .ok_or("strided row placement row size overflow")?;
8149 if src.len() < src_len || dst.len() < dst_len || row_end > dst_stride {
8150 return Err(format!(
8151 "strided row placement geometry mismatch: src={} need_src={src_len} \
8152 dst={} need_dst={dst_len} row_elems={row_elems} rows={n_rows} \
8153 dst_stride={dst_stride} dst_off={dst_off}",
8154 src.len(),
8155 dst.len(),
8156 )
8157 .into());
8158 }
8159 if row_elems > i32::MAX as usize || n_rows > i32::MAX as usize {
8160 return Err("strided row placement exceeds CUDA kernel geometry".into());
8161 }
8162 let f = self.func("place_rows_strided_f32");
8163 let cfg = LaunchConfig {
8164 grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
8165 block_dim: (256, 1, 1),
8166 shared_mem_bytes: 0,
8167 };
8168 let (re, nr) = (row_elems as i32, n_rows as i32);
8169 let (st, off) = (dst_stride as i64, dst_off as i64);
8170 let __s_b = self.gpu.stream();
8171 let mut b = __s_b.launch_builder(&f);
8172 b.arg(src)
8173 .arg(&mut *dst)
8174 .arg(&re)
8175 .arg(&nr)
8176 .arg(&st)
8177 .arg(&off);
8178 unsafe {
8179 b.launch(cfg)?;
8180 }
8181 Ok(())
8182 }
8183
8184 pub fn u32_set_k(
8186 &self,
8187 dst: &mut CudaSlice<u32>,
8188 v: u32,
8189 idx: usize,
8190 ) -> Result<(), Box<dyn std::error::Error>> {
8191 let f = self.func("u32_set_k");
8192 let cfg = LaunchConfig {
8193 grid_dim: (1, 1, 1),
8194 block_dim: (1, 1, 1),
8195 shared_mem_bytes: 0,
8196 };
8197 let ii = idx as i32;
8198 let __s_b = self.gpu.stream();
8199 let mut b = __s_b.launch_builder(&f);
8200 b.arg(dst).arg(&v).arg(&ii);
8201 unsafe {
8202 b.launch(cfg)?;
8203 }
8204 Ok(())
8205 }
8206
8207 pub fn i32_add_k(
8209 &self,
8210 d: &mut CudaSlice<i32>,
8211 v: i32,
8212 ) -> Result<(), Box<dyn std::error::Error>> {
8213 let f = self.func("i32_add_k");
8214 let cfg = LaunchConfig {
8215 grid_dim: (1, 1, 1),
8216 block_dim: (32, 1, 1),
8217 shared_mem_bytes: 0,
8218 };
8219 let __s_b = self.gpu.stream();
8220 let mut b = __s_b.launch_builder(&f);
8221 b.arg(d).arg(&v);
8222 unsafe {
8223 b.launch(cfg)?;
8224 }
8225 Ok(())
8226 }
8227
8228 pub fn i32_iota_from(
8230 &self,
8231 ctr: &CudaSlice<i32>,
8232 dst: &mut CudaSlice<i32>,
8233 n: usize,
8234 ) -> Result<(), Box<dyn std::error::Error>> {
8235 let f = self.func("i32_iota_from");
8236 let cfg = LaunchConfig::for_num_elems(n as u32);
8237 let ni = n as i32;
8238 let __s_b = self.gpu.stream();
8239 let mut b = __s_b.launch_builder(&f);
8240 b.arg(ctr).arg(dst).arg(&ni);
8241 unsafe {
8242 b.launch(cfg)?;
8243 }
8244 Ok(())
8245 }
8246
8247 pub fn u32_map_k(
8249 &self,
8250 buf: &mut CudaSlice<u32>,
8251 map: &CudaSlice<u32>,
8252 idx: usize,
8253 ) -> Result<(), Box<dyn std::error::Error>> {
8254 let f = self.func("u32_map_k");
8255 let cfg = LaunchConfig {
8256 grid_dim: (1, 1, 1),
8257 block_dim: (1, 1, 1),
8258 shared_mem_bytes: 0,
8259 };
8260 let ii = idx as i32;
8261 let __s_b = self.gpu.stream();
8262 let mut b = __s_b.launch_builder(&f);
8263 b.arg(buf).arg(map).arg(&ii);
8264 unsafe {
8265 b.launch(cfg)?;
8266 }
8267 Ok(())
8268 }
8269
8270 #[allow(clippy::too_many_arguments)]
8272 pub fn u32_pack2(
8273 &self,
8274 a: &CudaSlice<u32>,
8275 off_a: usize,
8276 n1: usize,
8277 b_in: &CudaSlice<u32>,
8278 n2: usize,
8279 out: &mut CudaSlice<u32>,
8280 ) -> Result<(), Box<dyn std::error::Error>> {
8281 let f = self.func("u32_pack2");
8282 let cfg = LaunchConfig::for_num_elems((n1 + n2) as u32);
8283 let (oa, i1, i2) = (off_a as i32, n1 as i32, n2 as i32);
8284 let __s_b = self.gpu.stream();
8285 let mut b = __s_b.launch_builder(&f);
8286 b.arg(a).arg(&oa).arg(&i1).arg(b_in).arg(&i2).arg(out);
8287 unsafe {
8288 b.launch(cfg)?;
8289 }
8290 Ok(())
8291 }
8292
8293 pub fn moe_w_exscale(
8295 &self,
8296 w: &mut CudaSlice<f32>,
8297 sel: &CudaSlice<i32>,
8298 s: &CudaSlice<f32>,
8299 n: usize,
8300 ) -> Result<(), Box<dyn std::error::Error>> {
8301 let f = self.func("moe_w_exscale");
8302 let cfg = LaunchConfig::for_num_elems(n as u32);
8303 let ni = n as i32;
8304 let __s_b = self.gpu.stream();
8305 let mut b = __s_b.launch_builder(&f);
8306 b.arg(w).arg(sel).arg(s).arg(&ni);
8307 unsafe {
8308 b.launch(cfg)?;
8309 }
8310 Ok(())
8311 }
8312
8313 pub fn moe_w_scale_by_expert(
8316 &self,
8317 w: &mut CudaSlice<f32>,
8318 sel: &CudaSlice<i32>,
8319 macros: &CudaSlice<f32>,
8320 n_expert: usize,
8321 n: usize,
8322 ) -> Result<(), Box<dyn std::error::Error>> {
8323 let f = self.func("moe_w_scale_by_expert");
8324 let cfg = LaunchConfig {
8325 grid_dim: (n.div_ceil(64) as u32, 1, 1),
8326 block_dim: (64, 1, 1),
8327 shared_mem_bytes: 0,
8328 };
8329 let (ne, nn) = (n_expert as i32, n as i32);
8330 let __s_b = self.gpu.stream();
8331 let mut b = __s_b.launch_builder(&f);
8332 b.arg(w).arg(sel).arg(macros).arg(&ne).arg(&nn);
8333 unsafe {
8334 b.launch(cfg)?;
8335 }
8336 Ok(())
8337 }
8338
8339 #[allow(clippy::too_many_arguments)] pub fn moe_gate_up_silu8_dev_q8(
8341 &self,
8342 table: &CudaSlice<u64>,
8343 sel: &cudarc::driver::CudaView<i32>,
8344 aq: &CudaSlice<i8>,
8345 ad: &CudaSlice<f32>,
8346 in_f: usize,
8347 n_ff: usize,
8348 n_used: usize,
8349 n_expert: usize,
8350 qt_g: i32,
8351 qt_u: i32,
8352 rb_g: usize,
8353 rb_u: usize,
8354 macros: &CudaSlice<f32>,
8355 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8356 static GU: std::sync::OnceLock<(String, u32)> = std::sync::OnceLock::new();
8357 let (mode, wpb) = GU.get_or_init(|| {
8358 let mode = std::env::var("MEMRA_MOE_DEVQ8_GU").unwrap_or_default();
8359 let wpb = std::env::var("MEMRA_MOE_DEVQ8_WPB")
8360 .ok()
8361 .and_then(|v| v.parse().ok())
8362 .unwrap_or(4u32)
8363 .clamp(1, 16);
8364 (mode, wpb)
8365 });
8366 let (mode, wpb) = (mode.as_str(), *wpb);
8367 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
8368 let (inf, nff, ne, rbg, rbu) = (
8369 in_f as i32,
8370 n_ff as i32,
8371 n_expert as i32,
8372 rb_g as i64,
8373 rb_u as i64,
8374 );
8375 let (f, cfg) = match mode {
8376 "1" | "2" | "4" => {
8377 let rpw: u32 = mode.parse().unwrap();
8378 let f = self.func(match rpw {
8379 1 => "moe_gate_up_silu8_dev_q8_r1",
8380 2 => "moe_gate_up_silu8_dev_q8_r2",
8381 _ => "moe_gate_up_silu8_dev_q8_r4",
8382 });
8383 let rows_per_block = (rpw * wpb) as usize;
8384 let gx = n_ff.div_ceil(rows_per_block) as u32;
8385 (
8386 f,
8387 LaunchConfig {
8388 grid_dim: (gx, n_used as u32, 1),
8389 block_dim: (32, wpb, 1),
8390 shared_mem_bytes: 0,
8391 },
8392 )
8393 }
8394 "j8" if n_used <= 32 => (
8395 self.func("moe_gate_up_silu8_dev_q8_j8"),
8396 LaunchConfig {
8397 grid_dim: (n_ff as u32, 1, 1),
8398 block_dim: (32, n_used as u32, 1),
8399 shared_mem_bytes: 0,
8400 },
8401 ),
8402 "vsm2" => {
8404 let f = self.func("moe_gate_up_silu8_dev_q8_vsm2");
8405 let sh = (rb_g + rb_u) as u32;
8406 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8407 f.set_attribute(
8408 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
8409 sh as i32,
8410 )?;
8411 (
8412 f,
8413 LaunchConfig {
8414 grid_dim: (n_ff as u32, n_used as u32, 1),
8415 block_dim: (32, 1, 1),
8416 shared_mem_bytes: sh,
8417 },
8418 )
8419 }
8420 "vsm" => {
8421 let f = self.func("moe_gate_up_silu8_dev_q8_vsm");
8422 let sh = (rb_g + rb_u) as u32;
8423 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8424 f.set_attribute(
8425 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
8426 sh as i32,
8427 )?;
8428 (
8429 f,
8430 LaunchConfig {
8431 grid_dim: (n_ff as u32, n_used as u32, 1),
8432 block_dim: (32, 1, 1),
8433 shared_mem_bytes: sh,
8434 },
8435 )
8436 }
8437 "sg" => (
8438 self.func("moe_gate_up_silu8_dev_q8_sg"),
8439 LaunchConfig {
8440 grid_dim: (n_ff as u32, n_used as u32, 1),
8441 block_dim: (32, 1, 1),
8442 shared_mem_bytes: 0,
8443 },
8444 ),
8445 "j8sg" if n_used <= 32 => (
8446 self.func("moe_gate_up_silu8_dev_q8_j8sg"),
8447 LaunchConfig {
8448 grid_dim: (n_ff as u32, 1, 1),
8449 block_dim: (32, n_used as u32, 1),
8450 shared_mem_bytes: 0,
8451 },
8452 ),
8453 "u64" if in_f == 2048 => (
8454 self.func("moe_gate_up_silu8_dev_q8_u64"),
8455 LaunchConfig {
8456 grid_dim: (n_ff as u32, n_used as u32, 1),
8457 block_dim: (32, 1, 1),
8458 shared_mem_bytes: 0,
8459 },
8460 ),
8461 "gs4" if in_f == 2048 => (
8462 self.func("moe_gate_up_silu8_dev_q8_gs4"),
8463 LaunchConfig {
8464 grid_dim: (n_ff as u32, n_used as u32, 1),
8465 block_dim: (32, 4, 1),
8466 shared_mem_bytes: 0,
8467 },
8468 ),
8469 "v" | "" => (
8471 self.func("moe_gate_up_silu8_dev_q8_v"),
8472 LaunchConfig {
8473 grid_dim: (n_ff as u32, n_used as u32, 1),
8474 block_dim: (32, 1, 1),
8475 shared_mem_bytes: 0,
8476 },
8477 ),
8478 "s2" => (
8479 self.func("moe_gate_up_silu8_dev_q8_s2"),
8480 LaunchConfig {
8481 grid_dim: (n_ff as u32, n_used as u32, 1),
8482 block_dim: (32, 2, 1),
8483 shared_mem_bytes: 0,
8484 },
8485 ),
8486 "s2z" => {
8487 let rz = wpb.min(16); (
8489 self.func("moe_gate_up_silu8_dev_q8_s2z"),
8490 LaunchConfig {
8491 grid_dim: (n_ff.div_ceil(rz as usize) as u32, n_used as u32, 1),
8492 block_dim: (32, 2, rz),
8493 shared_mem_bytes: 0,
8494 },
8495 )
8496 }
8497 _ => (
8498 self.func("moe_gate_up_silu8_dev_q8"),
8499 LaunchConfig {
8500 grid_dim: (n_ff as u32, n_used as u32, 1),
8501 block_dim: (32, 1, 1),
8502 shared_mem_bytes: 0,
8503 },
8504 ),
8505 };
8506 let __s_b = self.gpu.stream();
8507 let mut b = __s_b.launch_builder(&f);
8508 b.arg(table)
8509 .arg(sel)
8510 .arg(aq)
8511 .arg(ad)
8512 .arg(&mut act)
8513 .arg(&inf)
8514 .arg(&nff)
8515 .arg(&ne)
8516 .arg(&qt_g)
8517 .arg(&qt_u)
8518 .arg(&rbg)
8519 .arg(&rbu)
8520 .arg(macros);
8521 unsafe {
8522 b.launch(cfg)?;
8523 }
8524 Ok(act)
8525 }
8526
8527 #[allow(clippy::too_many_arguments)]
8528 pub fn moe_down8_fma_dev_q8(
8529 &self,
8530 table: &CudaSlice<u64>,
8531 sel: &cudarc::driver::CudaView<i32>,
8532 w: &cudarc::driver::CudaView<f32>,
8533 aq2: &CudaSlice<i8>,
8534 ad2: &CudaSlice<f32>,
8535 dst: &mut cudarc::driver::CudaViewMut<f32>,
8536 in_f: usize,
8537 out_f: usize,
8538 n_used: usize,
8539 n_expert: usize,
8540 qt: i32,
8541 rb: usize,
8542 ) -> Result<(), Box<dyn std::error::Error>> {
8543 static DOWN: std::sync::OnceLock<String> = std::sync::OnceLock::new();
8544 let mode = DOWN.get_or_init(|| std::env::var("MEMRA_MOE_DEVQ8_DOWN").unwrap_or_default());
8545 let (inf, outf, nu, ne, rbi) = (
8546 in_f as i32,
8547 out_f as i32,
8548 n_used as i32,
8549 n_expert as i32,
8550 rb as i64,
8551 );
8552 let (f, cfg) = match mode.as_str() {
8555 m @ ("1" | "2" | "4") if n_used <= 8 => {
8556 let rpw: usize = m.parse().unwrap();
8557 let f = self.func(match rpw {
8558 1 => "moe_down8_fma_dev_q8_w8r1",
8559 2 => "moe_down8_fma_dev_q8_w8r2",
8560 _ => "moe_down8_fma_dev_q8_w8r4",
8561 });
8562 (
8563 f,
8564 LaunchConfig {
8565 grid_dim: (out_f.div_ceil(rpw) as u32, 1, 1),
8566 block_dim: (32, n_used as u32, 1),
8567 shared_mem_bytes: 0,
8568 },
8569 )
8570 }
8571 "h2" if in_f == 512 => (
8572 self.func("moe_down8_fma_dev_q8_h2"),
8573 LaunchConfig {
8574 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8575 block_dim: (32, 1, 1),
8576 shared_mem_bytes: 0,
8577 },
8578 ),
8579 "" if in_f == 704 && n_used <= 8 => (
8582 self.func("moe_down8_fma_dev_q8_w8r2"),
8583 LaunchConfig {
8584 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8585 block_dim: (32, n_used as u32, 1),
8586 shared_mem_bytes: 0,
8587 },
8588 ),
8589 "w8h2v" | "" if in_f == 512 && n_used <= 8 => (
8593 self.func("moe_down8_fma_dev_q8_w8h2v"),
8594 LaunchConfig {
8595 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8596 block_dim: (32, n_used as u32, 1),
8597 shared_mem_bytes: 0,
8598 },
8599 ),
8600 "w8h2r2v" if in_f == 512 && n_used <= 8 => (
8601 self.func("moe_down8_fma_dev_q8_w8h2r2v"),
8602 LaunchConfig {
8603 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
8604 block_dim: (32, n_used as u32, 1),
8605 shared_mem_bytes: 0,
8606 },
8607 ),
8608 "w8h2r2" if in_f == 512 && n_used <= 8 => (
8609 self.func("moe_down8_fma_dev_q8_w8h2r2"),
8610 LaunchConfig {
8611 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
8612 block_dim: (32, n_used as u32, 1),
8613 shared_mem_bytes: 0,
8614 },
8615 ),
8616 "w8h2" if in_f == 512 && n_used <= 8 => (
8617 self.func("moe_down8_fma_dev_q8_w8h2"),
8618 LaunchConfig {
8619 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8620 block_dim: (32, n_used as u32, 1),
8621 shared_mem_bytes: 0,
8622 },
8623 ),
8624 _ => (
8625 self.func("moe_down8_fma_dev_q8"),
8626 LaunchConfig {
8627 grid_dim: (out_f as u32, 1, 1),
8628 block_dim: (32, 1, 1),
8629 shared_mem_bytes: 0,
8630 },
8631 ),
8632 };
8633 let __s_b = self.gpu.stream();
8634 let mut b = __s_b.launch_builder(&f);
8635 b.arg(table)
8636 .arg(sel)
8637 .arg(w)
8638 .arg(aq2)
8639 .arg(ad2)
8640 .arg(dst)
8641 .arg(&inf)
8642 .arg(&outf)
8643 .arg(&nu)
8644 .arg(&ne)
8645 .arg(&qt)
8646 .arg(&rbi);
8647 unsafe {
8648 b.launch(cfg)?;
8649 }
8650 Ok(())
8651 }
8652
8653 #[allow(clippy::too_many_arguments)]
8660 pub fn moe_gate_up_silu8_dev_q8_rows(
8661 &self,
8662 table: &CudaSlice<u64>,
8663 sel: &CudaSlice<i32>,
8664 aq: &CudaSlice<i8>,
8665 ad: &CudaSlice<f32>,
8666 t: usize,
8667 in_f: usize,
8668 n_ff: usize,
8669 n_used: usize,
8670 n_expert: usize,
8671 qt_g: i32,
8672 qt_u: i32,
8673 rb_g: usize,
8674 rb_u: usize,
8675 macros: &CudaSlice<f32>,
8676 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8677 let f = self.func("moe_gate_up_silu8_dev_q8_v_rows");
8678 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
8679 let cfg = LaunchConfig {
8680 grid_dim: (n_ff as u32, n_used as u32, t as u32),
8681 block_dim: (32, 1, 1),
8682 shared_mem_bytes: 0,
8683 };
8684 let (inf, nff, ne, nu, rbg, rbu) = (
8685 in_f as i32,
8686 n_ff as i32,
8687 n_expert as i32,
8688 n_used as i32,
8689 rb_g as i64,
8690 rb_u as i64,
8691 );
8692 let __s_b = self.gpu.stream();
8693 let mut b = __s_b.launch_builder(&f);
8694 b.arg(table)
8695 .arg(sel)
8696 .arg(aq)
8697 .arg(ad)
8698 .arg(&mut act)
8699 .arg(&inf)
8700 .arg(&nff)
8701 .arg(&ne)
8702 .arg(&qt_g)
8703 .arg(&qt_u)
8704 .arg(&rbg)
8705 .arg(&rbu)
8706 .arg(&nu)
8707 .arg(macros);
8708 unsafe {
8709 b.launch(cfg)?;
8710 }
8711 Ok(act)
8712 }
8713
8714 #[allow(clippy::too_many_arguments)]
8719 pub fn moe_down8_fma_dev_q8_rows(
8720 &self,
8721 table: &CudaSlice<u64>,
8722 sel: &CudaSlice<i32>,
8723 w: &CudaSlice<f32>,
8724 aq2: &CudaSlice<i8>,
8725 ad2: &CudaSlice<f32>,
8726 dst: &mut CudaSlice<f32>,
8727 t: usize,
8728 in_f: usize,
8729 out_f: usize,
8730 n_used: usize,
8731 n_expert: usize,
8732 qt: i32,
8733 rb: usize,
8734 ) -> Result<(), Box<dyn std::error::Error>> {
8735 assert!(
8736 in_f == 512 && n_used <= 8,
8737 "down rows twin is w8h2v shape-gated"
8738 );
8739 let f = self.func("moe_down8_fma_dev_q8_w8h2v_rows");
8740 let cfg = LaunchConfig {
8741 grid_dim: (out_f.div_ceil(2) as u32, 1, t as u32),
8742 block_dim: (32, n_used as u32, 1),
8743 shared_mem_bytes: 0,
8744 };
8745 let (inf, outf, nu, ne, rbi) = (
8746 in_f as i32,
8747 out_f as i32,
8748 n_used as i32,
8749 n_expert as i32,
8750 rb as i64,
8751 );
8752 let __s_b = self.gpu.stream();
8753 let mut b = __s_b.launch_builder(&f);
8754 b.arg(table)
8755 .arg(sel)
8756 .arg(w)
8757 .arg(aq2)
8758 .arg(ad2)
8759 .arg(dst)
8760 .arg(&inf)
8761 .arg(&outf)
8762 .arg(&nu)
8763 .arg(&ne)
8764 .arg(&qt)
8765 .arg(&rbi);
8766 unsafe {
8767 b.launch(cfg)?;
8768 }
8769 Ok(())
8770 }
8771
8772 #[allow(clippy::too_many_arguments)]
8776 pub fn moe_gate_up_silu8_dev_q8_csr(
8777 &self,
8778 table: &CudaSlice<u64>,
8779 sel: &CudaSlice<i32>,
8780 aq: &CudaSlice<i8>,
8781 ad: &CudaSlice<f32>,
8782 n_pairs: usize,
8783 in_f: usize,
8784 n_ff: usize,
8785 n_used: usize,
8786 n_expert: usize,
8787 qt_g: i32,
8788 qt_u: i32,
8789 rb_g: usize,
8790 rb_u: usize,
8791 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8792 let f = if qt_g == crate::QT_NVFP4 {
8795 self.func("moe_gate_up_silu8_dev_q8_csr_nvfp4")
8796 } else {
8797 self.func("moe_gate_up_silu8_dev_q8_csr_iq4")
8798 };
8799 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
8800 let cfg = LaunchConfig {
8801 grid_dim: (n_ff as u32, n_pairs as u32, 1),
8802 block_dim: (32, 1, 1),
8803 shared_mem_bytes: 0,
8804 };
8805 let (inf, nff, ne, nu, npi, rbg, rbu) = (
8806 in_f as i32,
8807 n_ff as i32,
8808 n_expert as i32,
8809 n_used as i32,
8810 n_pairs as i32,
8811 rb_g as i64,
8812 rb_u as i64,
8813 );
8814 let __s_b = self.gpu.stream();
8815 let mut b = __s_b.launch_builder(&f);
8816 b.arg(table)
8817 .arg(sel)
8818 .arg(aq)
8819 .arg(ad)
8820 .arg(&mut act)
8821 .arg(&inf)
8822 .arg(&nff)
8823 .arg(&ne)
8824 .arg(&qt_g)
8825 .arg(&qt_u)
8826 .arg(&rbg)
8827 .arg(&rbu)
8828 .arg(&nu)
8829 .arg(&npi);
8830 unsafe {
8831 b.launch(cfg)?;
8832 }
8833 Ok(act)
8834 }
8835
8836 #[allow(clippy::too_many_arguments)]
8840 pub fn moe_down8_fma_dev_q8_variant(
8841 &self,
8842 variant: &str,
8843 table: &CudaSlice<u64>,
8844 sel: &cudarc::driver::CudaView<i32>,
8845 w: &cudarc::driver::CudaView<f32>,
8846 aq2: &CudaSlice<i8>,
8847 ad2: &CudaSlice<f32>,
8848 dst: &mut cudarc::driver::CudaViewMut<f32>,
8849 in_f: usize,
8850 out_f: usize,
8851 n_used: usize,
8852 n_expert: usize,
8853 qt: i32,
8854 rb: usize,
8855 ) -> Result<(), Box<dyn std::error::Error>> {
8856 let (inf, outf, nu, ne, rbi) = (
8857 in_f as i32,
8858 out_f as i32,
8859 n_used as i32,
8860 n_expert as i32,
8861 rb as i64,
8862 );
8863 let (f, cfg) = match variant {
8864 "w8h2" | "w8h2v" => (
8865 self.func(if variant == "w8h2" {
8866 "moe_down8_fma_dev_q8_w8h2"
8867 } else {
8868 "moe_down8_fma_dev_q8_w8h2v"
8869 }),
8870 LaunchConfig {
8871 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
8872 block_dim: (32, n_used as u32, 1),
8873 shared_mem_bytes: 0,
8874 },
8875 ),
8876 "w8h2r2" | "w8h2r2v" => (
8877 self.func(if variant == "w8h2r2" {
8878 "moe_down8_fma_dev_q8_w8h2r2"
8879 } else {
8880 "moe_down8_fma_dev_q8_w8h2r2v"
8881 }),
8882 LaunchConfig {
8883 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
8884 block_dim: (32, n_used as u32, 1),
8885 shared_mem_bytes: 0,
8886 },
8887 ),
8888 _ => (
8889 self.func("moe_down8_fma_dev_q8"),
8890 LaunchConfig {
8891 grid_dim: (out_f as u32, 1, 1),
8892 block_dim: (32, 1, 1),
8893 shared_mem_bytes: 0,
8894 },
8895 ),
8896 };
8897 let __s_b = self.gpu.stream();
8898 let mut b = __s_b.launch_builder(&f);
8899 b.arg(table)
8900 .arg(sel)
8901 .arg(w)
8902 .arg(aq2)
8903 .arg(ad2)
8904 .arg(dst)
8905 .arg(&inf)
8906 .arg(&outf)
8907 .arg(&nu)
8908 .arg(&ne)
8909 .arg(&qt)
8910 .arg(&rbi);
8911 unsafe {
8912 b.launch(cfg)?;
8913 }
8914 Ok(())
8915 }
8916
8917 #[allow(clippy::too_many_arguments)]
8919 pub fn moe_gate_up_silu8_dev_q8_variant(
8920 &self,
8921 variant: &str,
8922 table: &CudaSlice<u64>,
8923 sel: &cudarc::driver::CudaView<i32>,
8924 aq: &CudaSlice<i8>,
8925 ad: &CudaSlice<f32>,
8926 in_f: usize,
8927 n_ff: usize,
8928 n_used: usize,
8929 n_expert: usize,
8930 qt_g: i32,
8931 qt_u: i32,
8932 rb_g: usize,
8933 rb_u: usize,
8934 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8935 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
8936 let (inf, nff, ne, rbg, rbu) = (
8937 in_f as i32,
8938 n_ff as i32,
8939 n_expert as i32,
8940 rb_g as i64,
8941 rb_u as i64,
8942 );
8943 let f = self.func(if variant == "v" {
8944 "moe_gate_up_silu8_dev_q8_v"
8945 } else {
8946 "moe_gate_up_silu8_dev_q8"
8947 });
8948 let cfg = LaunchConfig {
8949 grid_dim: (n_ff as u32, n_used as u32, 1),
8950 block_dim: (32, 1, 1),
8951 shared_mem_bytes: 0,
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(aq)
8958 .arg(ad)
8959 .arg(&mut act)
8960 .arg(&inf)
8961 .arg(&nff)
8962 .arg(&ne)
8963 .arg(&qt_g)
8964 .arg(&qt_u)
8965 .arg(&rbg)
8966 .arg(&rbu);
8967 unsafe {
8968 b.launch(cfg)?;
8969 }
8970 Ok(act)
8971 }
8972
8973 #[allow(clippy::too_many_arguments)] pub fn moe_gate_up_silu8_dev(
8975 &self,
8976 table: &CudaSlice<u64>,
8977 sel: &cudarc::driver::CudaView<i32>,
8978 x: &cudarc::driver::CudaView<f32>,
8979 in_f: usize,
8980 n_ff: usize,
8981 n_used: usize,
8982 n_expert: usize,
8983 qt_g: i32,
8984 qt_u: i32,
8985 rb_g: usize,
8986 rb_u: usize,
8987 macros: &CudaSlice<f32>,
8988 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8989 let f = self.func("moe_gate_up_silu8_dev");
8990 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig {
8992 grid_dim: (n_ff as u32, n_used as u32, 1),
8993 block_dim: (256, 1, 1),
8994 shared_mem_bytes: 0,
8995 };
8996 let (inf, nff, ne, rbg, rbu) = (
8997 in_f as i32,
8998 n_ff as i32,
8999 n_expert as i32,
9000 rb_g as i64,
9001 rb_u as i64,
9002 );
9003 let __s_b = self.gpu.stream();
9004 let mut b = __s_b.launch_builder(&f);
9005 b.arg(table)
9006 .arg(sel)
9007 .arg(x)
9008 .arg(&mut act)
9009 .arg(&inf)
9010 .arg(&nff)
9011 .arg(&ne)
9012 .arg(&qt_g)
9013 .arg(&qt_u)
9014 .arg(&rbg)
9015 .arg(&rbu)
9016 .arg(macros);
9017 unsafe {
9018 b.launch(cfg)?;
9019 }
9020 Ok(act)
9021 }
9022
9023 #[allow(clippy::too_many_arguments)]
9026 pub fn moe_down8_fma_dev(
9027 &self,
9028 table: &CudaSlice<u64>,
9029 sel: &cudarc::driver::CudaView<i32>,
9030 w: &cudarc::driver::CudaView<f32>,
9031 act: &CudaSlice<f32>,
9032 dst: &mut cudarc::driver::CudaViewMut<f32>,
9033 in_f: usize,
9034 out_f: usize,
9035 n_used: usize,
9036 n_expert: usize,
9037 qt: i32,
9038 rb: usize,
9039 ) -> Result<(), Box<dyn std::error::Error>> {
9040 let f = self.func("moe_down8_fma_dev");
9041 let cfg = LaunchConfig {
9042 grid_dim: (out_f as u32, 1, 1),
9043 block_dim: (256, 1, 1),
9044 shared_mem_bytes: 0,
9045 };
9046 let (inf, outf, nu, ne, rbv) = (
9047 in_f as i32,
9048 out_f as i32,
9049 n_used as i32,
9050 n_expert as i32,
9051 rb as i64,
9052 );
9053 let __s_b = self.gpu.stream();
9054 let mut b = __s_b.launch_builder(&f);
9055 b.arg(table)
9056 .arg(sel)
9057 .arg(w)
9058 .arg(act)
9059 .arg(dst)
9060 .arg(&inf)
9061 .arg(&outf)
9062 .arg(&nu)
9063 .arg(&ne)
9064 .arg(&qt)
9065 .arg(&rbv);
9066 unsafe {
9067 b.launch(cfg)?;
9068 }
9069 Ok(())
9070 }
9071
9072 pub fn axpy_into(
9074 &self,
9075 src: &CudaSlice<f32>,
9076 alpha: f32,
9077 dst: &mut cudarc::driver::CudaViewMut<f32>,
9078 n: usize,
9079 ) -> Result<(), Box<dyn std::error::Error>> {
9080 let f = self.func("axpy_f32");
9081 let cfg = LaunchConfig::for_num_elems(n as u32);
9082 let (a, ni) = (alpha, n as i32);
9083 let __s_b = self.gpu.stream();
9084 let mut b = __s_b.launch_builder(&f);
9085 b.arg(src).arg(dst).arg(&a).arg(&ni);
9086 unsafe {
9087 b.launch(cfg)?;
9088 }
9089 Ok(())
9090 }
9091
9092 pub fn axpy_host_into(
9094 &self,
9095 src: &cudarc::driver::CudaView<'_, f32>,
9096 alpha: f32,
9097 dst: &mut cudarc::driver::CudaViewMut<f32>,
9098 n: usize,
9099 ) -> Result<(), Box<dyn std::error::Error>> {
9100 let f = self.func("axpy_host_f32");
9101 let cfg = LaunchConfig::for_num_elems(n as u32);
9102 let (a, ni) = (alpha, n as i32);
9103 let __s_b = self.gpu.stream();
9104 let mut b = __s_b.launch_builder(&f);
9105 b.arg(src).arg(dst).arg(&a).arg(&ni);
9106 unsafe {
9107 b.launch(cfg)?;
9108 }
9109 Ok(())
9110 }
9111
9112 pub fn add_scaled_rows(
9114 &self,
9115 src: &CudaSlice<f32>,
9116 scale: &CudaSlice<f32>,
9117 dst: &mut CudaSlice<f32>,
9118 ncols: usize,
9119 nrows: usize,
9120 ) -> Result<(), Box<dyn std::error::Error>> {
9121 let f = self.func("add_scaled_rows_f32");
9122 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
9123 let (nc, nr) = (ncols as i32, nrows as i32);
9124 let __s_b = self.gpu.stream();
9125 let mut b = __s_b.launch_builder(&f);
9126 b.arg(src).arg(scale).arg(dst).arg(&nc).arg(&nr);
9127 unsafe {
9128 b.launch(cfg)?;
9129 }
9130 Ok(())
9131 }
9132
9133 pub fn add_scaled_rows_ones(
9138 &self,
9139 src: &CudaSlice<f32>,
9140 dst: &mut CudaSlice<f32>,
9141 ncols: usize,
9142 nrows: usize,
9143 ) -> Result<(), Box<dyn std::error::Error>> {
9144 let mut guard = self
9145 .shexp_ones
9146 .lock()
9147 .map_err(|_| "shexp ones buffer is poisoned")?;
9148 if guard.as_ref().map(|b| b.len() < nrows).unwrap_or(true) {
9149 *guard = Some(self.htod(&vec![1.0f32; nrows.max(64)])?);
9152 }
9153 let ones = guard.as_ref().expect("just ensured");
9154 let f = self.func("add_scaled_rows_f32");
9155 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
9156 let (nc, nr) = (ncols as i32, nrows as i32);
9157 let __s_b = self.gpu.stream();
9158 let mut b = __s_b.launch_builder(&f);
9159 b.arg(src).arg(ones).arg(&mut *dst).arg(&nc).arg(&nr);
9160 unsafe {
9161 b.launch(cfg)?;
9162 }
9163 Ok(())
9164 }
9165
9166 pub fn i32_mirror_store(
9170 &self,
9171 dst: &mut CudaSlice<i32>,
9172 v: i32,
9173 ) -> Result<(), Box<dyn std::error::Error>> {
9174 if crate::htod_diet_on() {
9175 HTOD_DIET_AVOIDED.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
9176 return self.i32_set_k(dst, v);
9177 }
9178 self.gpu.stream().memcpy_htod(&[v], dst)?;
9179 Ok(())
9180 }
9181
9182 pub fn scale_rows(
9185 &self,
9186 y: &mut CudaSlice<f32>,
9187 s: &CudaSlice<f32>,
9188 ncols: usize,
9189 nrows: usize,
9190 ) -> Result<(), Box<dyn std::error::Error>> {
9191 let f = self.func("scale_rows_f32");
9192 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
9193 let (nc, nr) = (ncols as i32, nrows as i32);
9194 let __s_b = self.gpu.stream();
9195 let mut b = __s_b.launch_builder(&f);
9196 b.arg(&mut *y).arg(s).arg(&nc).arg(&nr);
9197 unsafe {
9198 b.launch(cfg)?;
9199 }
9200 Ok(())
9201 }
9202
9203 #[allow(clippy::too_many_arguments)]
9207 pub fn moe_prime_join_scatter(
9208 &self,
9209 y0: &CudaSlice<f32>,
9210 y1: &CudaSlice<f32>,
9211 inv: &CudaSlice<i32>,
9212 w: &CudaSlice<f32>,
9213 out: &mut CudaSlice<f32>,
9214 ncols: usize,
9215 n_used: usize,
9216 t: usize,
9217 ) -> Result<(), Box<dyn std::error::Error>> {
9218 let f = self.func("moe_prime_join_scatter_f32");
9219 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
9220 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
9221 let __s_b = self.gpu.stream();
9222 let mut b = __s_b.launch_builder(&f);
9223 b.arg(y0)
9224 .arg(y1)
9225 .arg(inv)
9226 .arg(w)
9227 .arg(&mut *out)
9228 .arg(&nc)
9229 .arg(&nu)
9230 .arg(&ti);
9231 unsafe {
9232 b.launch(cfg)?;
9233 }
9234 Ok(())
9235 }
9236
9237 pub fn moe_pairs_weighted_scatter(
9240 &self,
9241 y: &CudaSlice<f32>,
9242 w: &CudaSlice<f32>,
9243 out: &mut CudaSlice<f32>,
9244 ncols: usize,
9245 n_used: usize,
9246 t: usize,
9247 ) -> Result<(), Box<dyn std::error::Error>> {
9248 let f = self.func("moe_pairs_weighted_scatter_f32");
9249 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
9250 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
9251 let __s_b = self.gpu.stream();
9252 let mut b = __s_b.launch_builder(&f);
9253 b.arg(y).arg(w).arg(&mut *out).arg(&nc).arg(&nu).arg(&ti);
9254 unsafe {
9255 b.launch(cfg)?;
9256 }
9257 Ok(())
9258 }
9259
9260 pub fn gather_rows(
9264 &self,
9265 src: &CudaSlice<f32>,
9266 idx: &CudaSlice<i32>,
9267 dst: &mut CudaSlice<f32>,
9268 ncols: usize,
9269 m_e: usize,
9270 ) -> Result<(), Box<dyn std::error::Error>> {
9271 let f = self.func("gather_rows_f32");
9272 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
9273 let (nc, me) = (ncols as i32, m_e as i32);
9274 let __s_b = self.gpu.stream();
9275 let mut b = __s_b.launch_builder(&f);
9276 b.arg(src).arg(idx).arg(dst).arg(&nc).arg(&me);
9277 unsafe {
9278 b.launch(cfg)?;
9279 }
9280 Ok(())
9281 }
9282
9283 #[allow(clippy::too_many_arguments)] pub fn scatter_slot(
9289 &self,
9290 src: &CudaSlice<f32>,
9291 tok_idx: &CudaSlice<i32>,
9292 slot_idx: &CudaSlice<i32>,
9293 weight: &CudaSlice<f32>,
9294 dst: &mut CudaSlice<f32>,
9295 wbuf: &mut CudaSlice<f32>,
9296 ncols: usize,
9297 n_used: usize,
9298 m_e: usize,
9299 ) -> Result<(), Box<dyn std::error::Error>> {
9300 let f = self.func("scatter_add_slot_f32");
9301 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
9302 let (nc, nu, me) = (ncols as i32, n_used as i32, m_e as i32);
9303 let __s_b = self.gpu.stream();
9304 let mut b = __s_b.launch_builder(&f);
9305 b.arg(src)
9306 .arg(tok_idx)
9307 .arg(slot_idx)
9308 .arg(weight)
9309 .arg(dst)
9310 .arg(wbuf)
9311 .arg(&nc)
9312 .arg(&nu)
9313 .arg(&me);
9314 unsafe {
9315 b.launch(cfg)?;
9316 }
9317 Ok(())
9318 }
9319
9320 pub fn reduce_slots(
9324 &self,
9325 slots: &CudaSlice<f32>,
9326 wbuf: &CudaSlice<f32>,
9327 dst: &mut CudaSlice<f32>,
9328 ncols: usize,
9329 n_used: usize,
9330 t: usize,
9331 ) -> Result<(), Box<dyn std::error::Error>> {
9332 let f = self.func("reduce_slots_f32");
9333 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
9334 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
9335 let __s_b = self.gpu.stream();
9336 let mut b = __s_b.launch_builder(&f);
9337 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
9338 unsafe {
9339 b.launch(cfg)?;
9340 }
9341 Ok(())
9342 }
9343
9344 pub fn reduce_slots_host(
9349 &self,
9350 slots: &CudaSlice<f32>,
9351 wbuf: &CudaSlice<f32>,
9352 dst: &mut CudaSlice<f32>,
9353 ncols: usize,
9354 n_used: usize,
9355 t: usize,
9356 ) -> Result<(), Box<dyn std::error::Error>> {
9357 let f = self.func("reduce_slots_host_f32");
9358 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
9359 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
9360 let __s_b = self.gpu.stream();
9361 let mut b = __s_b.launch_builder(&f);
9362 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
9363 unsafe {
9364 b.launch(cfg)?;
9365 }
9366 Ok(())
9367 }
9368
9369 pub fn quantize_q8_1_view(
9376 &self,
9377 x: &cudarc::driver::CudaView<f32>,
9378 m: usize,
9379 in_f: usize,
9380 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9381 let f = self.func("quantize_q8_1");
9382 let nblk = in_f / 32;
9383 let mut q = self.alloc_uninit::<i8>(m * in_f)?;
9384 let mut d = self.alloc_uninit::<f32>(m * nblk)?;
9385 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
9386 let (inf, mi) = (in_f as i32, m as i32);
9387 let __s_b = self.gpu.stream();
9388 let mut b = __s_b.launch_builder(&f);
9389 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
9390 unsafe {
9391 b.launch(cfg)?;
9392 }
9393 Ok((q, d))
9394 }
9395
9396 pub fn quantize_q8_1(
9397 &self,
9398 x: &CudaSlice<f32>,
9399 m: usize,
9400 in_f: usize,
9401 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9402 let nblk = in_f / 32;
9403 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);
9407 let (inf, mi) = (in_f as i32, m as i32);
9408 if Self::pdl_on() && Self::pdl_wb_on() {
9409 {
9410 use cudarc::driver::{DevicePtr, DevicePtrMut};
9411 let s = &self.gpu.stream();
9412 let (px, _g0) = x.device_ptr(s);
9413 let (pq, _g1) = q.device_ptr_mut(s);
9414 let (pd, _g2) = d.device_ptr_mut(s);
9415 let mut ps = [
9416 &px as *const _ as *mut std::ffi::c_void,
9417 &pq as *const _ as *mut _,
9418 &pd as *const _ as *mut _,
9419 &inf as *const _ as *mut _,
9420 &mi as *const _ as *mut _,
9421 ];
9422 unsafe {
9423 self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?;
9424 }
9425 }
9426 return Ok((q, d));
9427 }
9428 let f = self.func("quantize_q8_1");
9429 let __s_b = self.gpu.stream();
9430 let mut b = __s_b.launch_builder(&f);
9431 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
9432 unsafe {
9433 b.launch(cfg)?;
9434 }
9435 Ok((q, d))
9436 }
9437
9438 pub fn quantize_fp4_act(
9442 &self,
9443 x: &CudaSlice<f32>,
9444 m: usize,
9445 in_f: usize,
9446 ) -> Result<(CudaSlice<u32>, CudaSlice<u8>), Box<dyn std::error::Error>> {
9447 let f = self.func("quantize_fp4_act");
9448 let nb16 = in_f / 16;
9449 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);
9452 let (inf, mi) = (in_f as i32, m as i32);
9453 let __s_b = self.gpu.stream();
9454 let mut b = __s_b.launch_builder(&f);
9455 b.arg(x).arg(&mut aq4).arg(&mut ad4).arg(&inf).arg(&mi);
9456 unsafe {
9457 b.launch(cfg)?;
9458 }
9459 Ok((aq4, ad4))
9460 }
9461
9462 #[allow(clippy::too_many_arguments)] pub fn qmatvec_gemm_nvfp4_fp4(
9468 &self,
9469 bytes: &CudaSlice<u8>,
9470 x: &CudaSlice<f32>,
9471 m: usize,
9472 in_f: usize,
9473 out_f: usize,
9474 row_bytes: usize,
9475 scale: f32,
9476 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9477 assert!(
9478 in_f.is_multiple_of(64),
9479 "FP4 GEMM requires in_f % 64 == 0, got {in_f}"
9480 );
9481 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
9482 let mut y = self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)?;
9483 if scale != 1.0 {
9484 self.scale_inplace(&mut y, scale, m * out_f)?;
9485 }
9486 Ok(y)
9487 }
9488
9489 #[allow(clippy::too_many_arguments)]
9492 #[allow(clippy::manual_div_ceil)] fn fp4_gemm_launch(
9495 &self,
9496 bytes: &CudaSlice<u8>,
9497 aq4: &CudaSlice<u32>,
9498 ad4: &CudaSlice<u8>,
9499 m: usize,
9500 in_f: usize,
9501 out_f: usize,
9502 row_bytes: usize,
9503 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9504 let f = self.func("qmatvec_gemm_nvfp4_fp4");
9505 let mut y = self.alloc_uninit::<f32>(m * out_f)?; const BM: u32 = 64;
9507 const BN: u32 = 256;
9508 let cfg = LaunchConfig {
9509 grid_dim: ((out_f as u32 + BM - 1) / BM, (m as u32 + BN - 1) / BN, 1),
9510 block_dim: (32, 4, 1),
9511 shared_mem_bytes: 0,
9512 };
9513 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9514 let __s_b = self.gpu.stream();
9515 let mut b = __s_b.launch_builder(&f);
9516 b.arg(bytes)
9517 .arg(aq4)
9518 .arg(ad4)
9519 .arg(&mut y)
9520 .arg(&inf)
9521 .arg(&outf)
9522 .arg(&mi)
9523 .arg(&rb);
9524 unsafe {
9525 b.launch(cfg)?;
9526 }
9527 Ok(y)
9528 }
9529
9530 pub fn qmatvec_gemm_nvfp4_fp4_raw(
9532 &self,
9533 bytes: &CudaSlice<u8>,
9534 x: &CudaSlice<f32>,
9535 m: usize,
9536 in_f: usize,
9537 out_f: usize,
9538 row_bytes: usize,
9539 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9540 assert!(
9541 in_f.is_multiple_of(64),
9542 "FP4 GEMM requires in_f % 64 == 0, got {in_f}"
9543 );
9544 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
9545 self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)
9546 }
9547
9548 pub fn qmatvec_q8_0_fast(
9550 &self,
9551 w: &CudaSlice<u8>,
9552 x: &CudaSlice<f32>,
9553 m: usize,
9554 in_f: usize,
9555 out_f: usize,
9556 row_bytes: usize,
9557 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9558 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
9559 let f = self.func("qmatvec_q8_0_dp4a");
9560 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
9562 grid_dim: (out_f as u32, m as u32, 1),
9563 block_dim: (128, 1, 1),
9564 shared_mem_bytes: 0,
9565 };
9566 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9567 let __s_b = self.gpu.stream();
9568 let mut b = __s_b.launch_builder(&f);
9569 b.arg(w)
9570 .arg(&aq)
9571 .arg(&ad)
9572 .arg(&mut y)
9573 .arg(&inf)
9574 .arg(&outf)
9575 .arg(&mi)
9576 .arg(&rb);
9577 unsafe {
9578 b.launch(cfg)?;
9579 }
9580 Ok(y)
9581 }
9582
9583 #[allow(non_snake_case)] pub fn qmatvec_q4_K_fast(
9586 &self,
9587 w: &CudaSlice<u8>,
9588 x: &CudaSlice<f32>,
9589 m: usize,
9590 in_f: usize,
9591 out_f: usize,
9592 row_bytes: usize,
9593 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9594 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
9595 let f = self.func("qmatvec_q4_K_dp4a");
9596 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
9598 grid_dim: (out_f as u32, m as u32, 1),
9599 block_dim: (128, 1, 1),
9600 shared_mem_bytes: 0,
9601 };
9602 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9603 let __s_b = self.gpu.stream();
9604 let mut b = __s_b.launch_builder(&f);
9605 b.arg(w)
9606 .arg(&aq)
9607 .arg(&ad)
9608 .arg(&mut y)
9609 .arg(&inf)
9610 .arg(&outf)
9611 .arg(&mi)
9612 .arg(&rb);
9613 unsafe {
9614 b.launch(cfg)?;
9615 }
9616 Ok(y)
9617 }
9618
9619 #[allow(non_snake_case)] pub fn qmatvec_q6_K_fast(
9622 &self,
9623 w: &CudaSlice<u8>,
9624 x: &CudaSlice<f32>,
9625 m: usize,
9626 in_f: usize,
9627 out_f: usize,
9628 row_bytes: usize,
9629 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9630 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
9631 let f = self.func("qmatvec_q6_K_dp4a");
9632 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
9634 grid_dim: (out_f as u32, m as u32, 1),
9635 block_dim: (128, 1, 1),
9636 shared_mem_bytes: 0,
9637 };
9638 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9639 let __s_b = self.gpu.stream();
9640 let mut b = __s_b.launch_builder(&f);
9641 b.arg(w)
9642 .arg(&aq)
9643 .arg(&ad)
9644 .arg(&mut y)
9645 .arg(&inf)
9646 .arg(&outf)
9647 .arg(&mi)
9648 .arg(&rb);
9649 unsafe {
9650 b.launch(cfg)?;
9651 }
9652 Ok(y)
9653 }
9654
9655 #[allow(non_snake_case)] pub fn qmatvec_q5_K_fast(
9658 &self,
9659 w: &CudaSlice<u8>,
9660 x: &CudaSlice<f32>,
9661 m: usize,
9662 in_f: usize,
9663 out_f: usize,
9664 row_bytes: usize,
9665 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9666 self.qmatvec_dp4a_named(
9667 "qmatvec_q5_K_dp4a",
9668 &w.slice(0..w.len()),
9669 x,
9670 m,
9671 in_f,
9672 out_f,
9673 row_bytes,
9674 )
9675 }
9676 #[allow(non_snake_case)] pub fn qmatvec_q3_K_fast(
9679 &self,
9680 w: &CudaSlice<u8>,
9681 x: &CudaSlice<f32>,
9682 m: usize,
9683 in_f: usize,
9684 out_f: usize,
9685 row_bytes: usize,
9686 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9687 self.qmatvec_dp4a_named(
9688 "qmatvec_q3_K_dp4a",
9689 &w.slice(0..w.len()),
9690 x,
9691 m,
9692 in_f,
9693 out_f,
9694 row_bytes,
9695 )
9696 }
9697 pub fn qmatvec_nvfp4_fast_rp(
9699 &self,
9700 w: &CudaSlice<u8>,
9701 x: &CudaSlice<f32>,
9702 m: usize,
9703 in_f: usize,
9704 out_f: usize,
9705 row_bytes: usize,
9706 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9707 assert!(
9708 in_f.is_multiple_of(64),
9709 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
9710 );
9711 self.qmatvec_dp4a_named(
9712 "qmatvec_nvfp4_dp4a_rp",
9713 &w.slice(0..w.len()),
9714 x,
9715 m,
9716 in_f,
9717 out_f,
9718 row_bytes,
9719 )
9720 }
9721 pub fn qmatvec_nvfp4_fast(
9723 &self,
9724 w: &cudarc::driver::CudaView<'_, u8>,
9725 x: &CudaSlice<f32>,
9726 m: usize,
9727 in_f: usize,
9728 out_f: usize,
9729 row_bytes: usize,
9730 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9731 assert!(
9734 in_f.is_multiple_of(64),
9735 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
9736 );
9737 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a", w, x, m, in_f, out_f, row_bytes)
9738 }
9739 pub fn qmatvec_nvfp4_fast_v2(
9744 &self,
9745 w: &cudarc::driver::CudaView<'_, u8>,
9746 x: &CudaSlice<f32>,
9747 m: usize,
9748 in_f: usize,
9749 out_f: usize,
9750 row_bytes: usize,
9751 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9752 assert!(
9753 in_f.is_multiple_of(64),
9754 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
9755 );
9756 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a_v2", w, x, m, in_f, out_f, row_bytes)
9757 }
9758 #[allow(non_snake_case)] pub fn qmatvec_iq4_XS_fast(
9761 &self,
9762 w: &CudaSlice<u8>,
9763 x: &CudaSlice<f32>,
9764 m: usize,
9765 in_f: usize,
9766 out_f: usize,
9767 row_bytes: usize,
9768 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9769 self.qmatvec_dp4a_named(
9770 "qmatvec_iq4_XS_dp4a",
9771 &w.slice(0..w.len()),
9772 x,
9773 m,
9774 in_f,
9775 out_f,
9776 row_bytes,
9777 )
9778 }
9779
9780 #[allow(clippy::too_many_arguments)] fn qmatvec_dp4a_named(
9783 &self,
9784 name: &str,
9785 w: &cudarc::driver::CudaView<'_, u8>,
9786 x: &CudaSlice<f32>,
9787 m: usize,
9788 in_f: usize,
9789 out_f: usize,
9790 row_bytes: usize,
9791 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9792 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
9793 let f = self.func(name);
9794 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
9796 grid_dim: (out_f as u32, m as u32, 1),
9797 block_dim: (128, 1, 1),
9798 shared_mem_bytes: 0,
9799 };
9800 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9801 let __s_b = self.gpu.stream();
9802 let mut b = __s_b.launch_builder(&f);
9803 b.arg(w)
9804 .arg(&aq)
9805 .arg(&ad)
9806 .arg(&mut y)
9807 .arg(&inf)
9808 .arg(&outf)
9809 .arg(&mi)
9810 .arg(&rb);
9811 unsafe {
9812 b.launch(cfg)?;
9813 }
9814 Ok(y)
9815 }
9816
9817 #[allow(clippy::too_many_arguments)]
9823 pub fn qmatvec_nvfp4_fast_prequant_into(
9824 &self,
9825 w: &CudaSlice<u8>,
9826 aq: &CudaSlice<i8>,
9827 ad: &CudaSlice<f32>,
9828 y: &mut CudaSlice<f32>,
9829 m: usize,
9830 in_f: usize,
9831 out_f: usize,
9832 row_bytes: usize,
9833 ) -> Result<(), Box<dyn std::error::Error>> {
9834 assert!(
9835 in_f.is_multiple_of(64),
9836 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
9837 );
9838 if y.len() < m * out_f {
9839 return Err(format!(
9840 "NVFP4 prequant output {} is shorter than {m}x{out_f}",
9841 y.len()
9842 )
9843 .into());
9844 }
9845 let f = self.func("qmatvec_nvfp4_dp4a");
9846 let cfg = LaunchConfig {
9847 grid_dim: (out_f as u32, m as u32, 1),
9848 block_dim: (128, 1, 1),
9849 shared_mem_bytes: 0,
9850 };
9851 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
9852 let __s_b = self.gpu.stream();
9853 let mut b = __s_b.launch_builder(&f);
9854 b.arg(w)
9855 .arg(aq)
9856 .arg(ad)
9857 .arg(y)
9858 .arg(&inf)
9859 .arg(&outf)
9860 .arg(&mi)
9861 .arg(&rb);
9862 unsafe {
9863 b.launch(cfg)?;
9864 }
9865 Ok(())
9866 }
9867
9868 #[allow(clippy::too_many_arguments)]
9871 pub fn matvec_f32_qkv_into(
9872 &self,
9873 wq: &CudaSlice<f32>,
9874 wk: &CudaSlice<f32>,
9875 wv: &CudaSlice<f32>,
9876 wg: &CudaSlice<f32>,
9877 x: &CudaSlice<f32>,
9878 yq: &mut CudaSlice<f32>,
9879 yk: &mut CudaSlice<f32>,
9880 yv: &mut CudaSlice<f32>,
9881 yg: &mut CudaSlice<f32>,
9882 in_f: usize,
9883 out_q: usize,
9884 out_kv: usize,
9885 out_g: usize,
9886 ) -> Result<(), Box<dyn std::error::Error>> {
9887 if !in_f.is_multiple_of(4)
9888 || wq.len() != out_q * in_f
9889 || wk.len() != out_kv * in_f
9890 || wv.len() != out_kv * in_f
9891 || wg.len() < out_g * in_f
9892 || x.len() < in_f
9893 || yq.len() < out_q
9894 || yk.len() < out_kv
9895 || yv.len() < out_kv
9896 || (out_g > 0 && yg.len() < out_g)
9897 {
9898 return Err(format!(
9899 "fused QKV geometry in={in_f} out_q={out_q} out_kv={out_kv} out_g={out_g} \
9900 wq={} wk={} wv={} wg={}",
9901 wq.len(),
9902 wk.len(),
9903 wv.len(),
9904 wg.len()
9905 )
9906 .into());
9907 }
9908 let f = self.func("matvec_f32_qkv");
9909 let cfg = LaunchConfig {
9910 grid_dim: ((out_q + 2 * out_kv + out_g) as u32, 1, 1),
9911 block_dim: (128, 1, 1),
9912 shared_mem_bytes: 0,
9913 };
9914 let (inf, oq, okv, og) = (in_f as i32, out_q as i32, out_kv as i32, out_g as i32);
9915 let __s_b = self.gpu.stream();
9916 let mut b = __s_b.launch_builder(&f);
9917 b.arg(wq)
9918 .arg(wk)
9919 .arg(wv)
9920 .arg(wg)
9921 .arg(x)
9922 .arg(yq)
9923 .arg(yk)
9924 .arg(yv)
9925 .arg(yg)
9926 .arg(&inf)
9927 .arg(&oq)
9928 .arg(&okv)
9929 .arg(&og);
9930 unsafe {
9931 b.launch(cfg)?;
9932 }
9933 Ok(())
9934 }
9935
9936 #[allow(clippy::too_many_arguments)]
9948 pub fn qmatvec_nvfp4_sel_gu_into(
9949 &self,
9950 gate_bank: &CudaSlice<u8>,
9951 up_bank: &CudaSlice<u8>,
9952 sel: &CudaSlice<i32>,
9953 aq: &CudaSlice<i8>,
9954 ad: &CudaSlice<f32>,
9955 yg: &mut CudaSlice<f32>,
9956 yu: &mut CudaSlice<f32>,
9957 n_sel: usize,
9958 in_f: usize,
9959 out_f: usize,
9960 row_bytes: usize,
9961 expert_stride: usize,
9962 slot_major: bool,
9963 ) -> Result<(), Box<dyn std::error::Error>> {
9964 assert!(
9965 in_f.is_multiple_of(64),
9966 "NVFP4 dp4a requires in_f % 64 == 0"
9967 );
9968 if yg.len() < n_sel * out_f || yu.len() < n_sel * out_f || sel.len() < n_sel {
9969 return Err("NVFP4 gu sel geometry".into());
9970 }
9971 if !slot_major {
9972 return Err(
9973 "NVFP4 gu sel fusion reads slot-major rows: these banks are block_nvfp4 \
9974 v1 (arm MEMRA_NVFP4_BANK_SM to build slot-major TP banks)"
9975 .into(),
9976 );
9977 }
9978 static RPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9982 let rpw = *RPW.get_or_init(|| {
9983 std::env::var("MEMRA_NVFP4_SEL_GU_RPW")
9984 .ok()
9985 .and_then(|v| v.parse().ok())
9986 .filter(|r| *r == 2 || *r == 4)
9987 .unwrap_or(1)
9988 });
9989 let rpw = if out_f.is_multiple_of(rpw) { rpw } else { 1 };
9990 static WPR: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9995 let wpr =
9996 *WPR.get_or_init(|| std::env::var("MEMRA_NVFP4_SEL_GU_WPR").as_deref() == Ok("1"));
9997 let f = self.func(match (wpr, rpw) {
9998 (true, _) => "qmatvec_nvfp4_dp4a_sel_v2_gu_wpr",
9999 (_, 4) => "qmatvec_nvfp4_dp4a_sel_v2_gu_r4",
10000 (_, 2) => "qmatvec_nvfp4_dp4a_sel_v2_gu_r2",
10001 _ => "qmatvec_nvfp4_dp4a_sel_v2_gu",
10002 });
10003 let cfg = LaunchConfig {
10004 grid_dim: if wpr {
10005 (((2 * out_f) as u32).div_ceil(4), n_sel as u32, 1)
10006 } else if rpw == 1 {
10007 ((2 * out_f) as u32, n_sel as u32, 1)
10008 } else {
10009 ((out_f / rpw) as u32, n_sel as u32, 1)
10010 },
10011 block_dim: if wpr { (32, 4, 1) } else { (128, 1, 1) },
10012 shared_mem_bytes: 0,
10013 };
10014 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10015 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10016 let (ars, adrs) = (0i64, 0i64);
10017 let __s_b = self.gpu.stream();
10018 let mut b = __s_b.launch_builder(&f);
10019 b.arg(gate_bank)
10020 .arg(up_bank)
10021 .arg(sel)
10022 .arg(aq)
10023 .arg(ad)
10024 .arg(yg)
10025 .arg(yu)
10026 .arg(&inf)
10027 .arg(&outf)
10028 .arg(&ns)
10029 .arg(&rb)
10030 .arg(&es)
10031 .arg(&ars)
10032 .arg(&adrs);
10033 unsafe {
10034 b.launch(cfg)?;
10035 }
10036 Ok(())
10037 }
10038
10039 #[allow(clippy::too_many_arguments)]
10051 pub fn qmatvec_nvfp4_sel_down8_into(
10052 &self,
10053 bank: &CudaSlice<u8>,
10054 sel: &CudaSlice<i32>,
10055 aq: &CudaSlice<i8>,
10056 ad: &CudaSlice<f32>,
10057 route_w: &CudaSlice<f32>,
10058 md: &CudaSlice<f32>,
10059 dst: &mut CudaSlice<f32>,
10060 n_sel: usize,
10061 in_f: usize,
10062 out_f: usize,
10063 row_bytes: usize,
10064 expert_stride: usize,
10065 act_row_stride: usize,
10066 ad_row_stride: usize,
10067 slot_major: bool,
10068 ) -> Result<(), Box<dyn std::error::Error>> {
10069 if !in_f.is_multiple_of(64)
10070 || n_sel == 0
10071 || n_sel > 8
10072 || (in_f >> 5) > 32
10073 || dst.len() < out_f
10074 || sel.len() < n_sel
10075 || route_w.len() < n_sel
10076 {
10077 return Err(format!(
10078 "NVFP4 sel down8 geometry in_f={in_f} out_f={out_f} n_sel={n_sel} dst={}",
10079 dst.len()
10080 )
10081 .into());
10082 }
10083 if !slot_major {
10084 return Err(
10085 "NVFP4 sel down8 reads slot-major rows: this shard is block_nvfp4 v1 \
10086 (arm MEMRA_NVFP4_BANK_SM to build slot-major TP banks)"
10087 .into(),
10088 );
10089 }
10090 let f = self.func("qmatvec_nvfp4_dp4a_sel_v2_down8");
10091 let cfg = LaunchConfig {
10092 grid_dim: (out_f as u32, 1, 1),
10093 block_dim: (32, n_sel as u32, 1),
10094 shared_mem_bytes: 0,
10095 };
10096 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10097 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10098 let (ars, adrs) = (act_row_stride as i64, ad_row_stride as i64);
10099 let __s_b = self.gpu.stream();
10100 let mut b = __s_b.launch_builder(&f);
10101 b.arg(bank)
10102 .arg(sel)
10103 .arg(aq)
10104 .arg(ad)
10105 .arg(route_w)
10106 .arg(md)
10107 .arg(dst)
10108 .arg(&inf)
10109 .arg(&outf)
10110 .arg(&ns)
10111 .arg(&rb)
10112 .arg(&es)
10113 .arg(&ars)
10114 .arg(&adrs);
10115 unsafe {
10116 b.launch(cfg)?;
10117 }
10118 Ok(())
10119 }
10120
10121 #[allow(clippy::too_many_arguments)]
10124 pub fn qmatvec_nvfp4_sel_gu_ep_into(
10125 &self,
10126 gate_bank: &CudaSlice<u8>,
10127 up_bank: &CudaSlice<u8>,
10128 sel: &CudaSlice<i32>,
10129 aq: &CudaSlice<i8>,
10130 ad: &CudaSlice<f32>,
10131 yg: &mut CudaSlice<f32>,
10132 yu: &mut CudaSlice<f32>,
10133 n_sel: usize,
10134 in_f: usize,
10135 out_f: usize,
10136 row_bytes: usize,
10137 expert_stride: usize,
10138 owner: usize,
10139 ) -> Result<(), Box<dyn std::error::Error>> {
10140 assert!(
10141 in_f.is_multiple_of(64),
10142 "NVFP4 dp4a requires in_f % 64 == 0"
10143 );
10144 if yg.len() < n_sel * out_f || yu.len() < n_sel * out_f || sel.len() < n_sel {
10145 return Err("NVFP4 gu ep geometry".into());
10146 }
10147 let f = self.func("qmatvec_nvfp4_dp4a_sel_v2_gu_ep");
10148 let cfg = LaunchConfig {
10149 grid_dim: ((2 * out_f) as u32, n_sel as u32, 1),
10150 block_dim: (128, 1, 1),
10151 shared_mem_bytes: 0,
10152 };
10153 let (inf, outf, ns, own) = (in_f as i32, out_f as i32, n_sel as i32, owner as i32);
10154 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10155 let (ars, adrs) = (0i64, 0i64);
10156 let __s_b = self.gpu.stream();
10157 let mut b = __s_b.launch_builder(&f);
10158 b.arg(gate_bank)
10159 .arg(up_bank)
10160 .arg(sel)
10161 .arg(aq)
10162 .arg(ad)
10163 .arg(yg)
10164 .arg(yu)
10165 .arg(&inf)
10166 .arg(&outf)
10167 .arg(&ns)
10168 .arg(&rb)
10169 .arg(&es)
10170 .arg(&ars)
10171 .arg(&adrs)
10172 .arg(&own);
10173 unsafe {
10174 b.launch(cfg)?;
10175 }
10176 Ok(())
10177 }
10178
10179 #[allow(clippy::too_many_arguments)]
10181 pub fn silu_mul_scaled_q8_1_sel_ep_into(
10182 &self,
10183 gate: &CudaSlice<f32>,
10184 up: &CudaSlice<f32>,
10185 gmac: &CudaSlice<f32>,
10186 umac: &CudaSlice<f32>,
10187 sel: &CudaSlice<i32>,
10188 limit: Option<f32>,
10189 out_q: &mut CudaSlice<i8>,
10190 out_d: &mut CudaSlice<f32>,
10191 n_per: usize,
10192 n_sel: usize,
10193 owner: usize,
10194 ) -> Result<(), Box<dyn std::error::Error>> {
10195 if !n_per.is_multiple_of(32)
10196 || out_q.len() < n_sel * n_per
10197 || out_d.len() < n_sel * n_per / 32
10198 {
10199 return Err("NVFP4 silu ep geometry".into());
10200 }
10201 let f = self.func("silu_mul_scaled_q8_1_sel_ep");
10202 let warps = n_sel * n_per / 32;
10203 let cfg = LaunchConfig {
10204 grid_dim: ((warps as u32).div_ceil(4), 1, 1),
10205 block_dim: (128, 1, 1),
10206 shared_mem_bytes: 0,
10207 };
10208 let (np, ns, own) = (n_per as i32, n_sel as i32, owner as i32);
10209 let (lim, has) = match limit {
10210 Some(l) => (l, 1i32),
10211 None => (0.0f32, 0i32),
10212 };
10213 let __s_b = self.gpu.stream();
10214 let mut b = __s_b.launch_builder(&f);
10215 b.arg(gate)
10216 .arg(up)
10217 .arg(gmac)
10218 .arg(umac)
10219 .arg(sel)
10220 .arg(&lim)
10221 .arg(&has)
10222 .arg(out_q)
10223 .arg(out_d)
10224 .arg(&np)
10225 .arg(&ns)
10226 .arg(&own);
10227 unsafe {
10228 b.launch(cfg)?;
10229 }
10230 Ok(())
10231 }
10232
10233 #[allow(clippy::too_many_arguments)]
10235 pub fn qmatvec_nvfp4_sel_down8_ep_into(
10236 &self,
10237 bank: &CudaSlice<u8>,
10238 sel: &CudaSlice<i32>,
10239 aq: &CudaSlice<i8>,
10240 ad: &CudaSlice<f32>,
10241 route_w: &CudaSlice<f32>,
10242 md: &CudaSlice<f32>,
10243 dst: &mut CudaSlice<f32>,
10244 n_sel: usize,
10245 in_f: usize,
10246 out_f: usize,
10247 row_bytes: usize,
10248 expert_stride: usize,
10249 act_row_stride: usize,
10250 ad_row_stride: usize,
10251 owner: usize,
10252 ) -> Result<(), Box<dyn std::error::Error>> {
10253 if !in_f.is_multiple_of(64)
10254 || n_sel == 0
10255 || n_sel > 8
10256 || (in_f >> 5) > 64
10257 || dst.len() < out_f
10258 {
10259 return Err("NVFP4 down8 ep geometry".into());
10260 }
10261 let f = self.func("qmatvec_nvfp4_dp4a_sel_v2_down8_ep");
10262 let cfg = LaunchConfig {
10263 grid_dim: (out_f as u32, 1, 1),
10264 block_dim: (32, n_sel as u32, 1),
10265 shared_mem_bytes: 0,
10266 };
10267 let (inf, outf, ns, own) = (in_f as i32, out_f as i32, n_sel as i32, owner as i32);
10268 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10269 let (ars, adrs) = (act_row_stride as i64, ad_row_stride as i64);
10270 let __s_b = self.gpu.stream();
10271 let mut b = __s_b.launch_builder(&f);
10272 b.arg(bank)
10273 .arg(sel)
10274 .arg(aq)
10275 .arg(ad)
10276 .arg(route_w)
10277 .arg(md)
10278 .arg(dst)
10279 .arg(&inf)
10280 .arg(&outf)
10281 .arg(&ns)
10282 .arg(&rb)
10283 .arg(&es)
10284 .arg(&ars)
10285 .arg(&adrs)
10286 .arg(&own);
10287 unsafe {
10288 b.launch(cfg)?;
10289 }
10290 Ok(())
10291 }
10292
10293 #[allow(clippy::too_many_arguments)] pub fn qmatvec_nvfp4_sel_into(
10308 &self,
10309 bank: &CudaSlice<u8>,
10310 sel: &CudaSlice<i32>,
10311 aq: &CudaSlice<i8>,
10312 ad: &CudaSlice<f32>,
10313 y: &mut CudaSlice<f32>,
10314 n_sel: usize,
10315 in_f: usize,
10316 out_f: usize,
10317 row_bytes: usize,
10318 expert_stride: usize,
10319 act_row_stride: usize,
10320 ad_row_stride: usize,
10321 slot_major: bool,
10322 ) -> Result<(), Box<dyn std::error::Error>> {
10323 assert!(
10324 in_f.is_multiple_of(64),
10325 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
10326 );
10327 if y.len() < n_sel * out_f || sel.len() < n_sel {
10328 return Err(format!(
10329 "NVFP4 sel output {} / sel {} shorter than {n_sel}x{out_f}",
10330 y.len(),
10331 sel.len()
10332 )
10333 .into());
10334 }
10335 static MR: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
10342 let mode = *MR.get_or_init(|| {
10343 if std::env::var("MEMRA_SEL_STREAM").as_deref() == Ok("1") {
10344 2
10345 } else if std::env::var("MEMRA_SEL_MR").as_deref() == Ok("1") {
10346 1
10347 } else {
10348 0
10349 }
10350 });
10351 let mode = if mode == 2 && in_f > 4096 { 0 } else { mode };
10352 let mode = if slot_major { 3 } else { mode };
10356 static SM_STREAM: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10360 let sm_stream = mode == 3
10361 && *SM_STREAM
10362 .get_or_init(|| std::env::var("MEMRA_NVFP4_SEL_SM_STREAM").as_deref() == Ok("1"))
10363 && row_bytes.is_multiple_of(16)
10364 && in_f <= 4096;
10365 let kname = match (mode, sm_stream) {
10366 (3, true) => "qmatvec_nvfp4_dp4a_sel_v2s",
10367 (3, false) => "qmatvec_nvfp4_dp4a_sel_v2",
10368 (2, _) => "qmatvec_nvfp4_dp4a_sel_stream",
10369 (1, _) => "qmatvec_nvfp4_dp4a_sel_mr4",
10370 _ => "qmatvec_nvfp4_dp4a_sel",
10371 };
10372 {
10378 static SEEN_SEL: std::sync::Mutex<Vec<(&'static str, usize, usize)>> =
10379 std::sync::Mutex::new(Vec::new());
10380 let combo = (kname, in_f, out_f);
10381 let mut seen = SEEN_SEL.lock().unwrap();
10382 if !seen.contains(&combo) {
10383 seen.push(combo);
10384 eprintln!(
10385 "[nvfp4-sel] kernel={kname} slot_major={slot_major} in_f={in_f} \
10386 out_f={out_f} nsb={} row_bytes={row_bytes}",
10387 in_f >> 5
10388 );
10389 }
10390 }
10391 let f = self.func(kname);
10392 let nsb = in_f >> 5;
10397 let fit_block: u32 = if (mode == 0 || mode == 3) && !sm_stream && nsb <= 32 {
10398 32
10399 } else if mode == 1 {
10400 512
10401 } else {
10402 128
10403 };
10404 let cfg = LaunchConfig {
10405 grid_dim: (
10406 if sm_stream {
10407 (out_f as u32).div_ceil(8)
10408 } else {
10409 match mode {
10410 2 => (out_f as u32).div_ceil(16),
10411 1 => (out_f as u32).div_ceil(4),
10412 _ => out_f as u32,
10413 }
10414 },
10415 n_sel as u32,
10416 1,
10417 ),
10418 block_dim: (fit_block, 1, 1),
10419 shared_mem_bytes: 0,
10420 };
10421 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10422 let (rb, es, ars, adrs) = (
10423 row_bytes as i64,
10424 expert_stride as i64,
10425 act_row_stride as i64,
10426 ad_row_stride as i64,
10427 );
10428 let __s_b = self.gpu.stream();
10429 let mut b = __s_b.launch_builder(&f);
10430 b.arg(bank)
10431 .arg(sel)
10432 .arg(aq)
10433 .arg(ad)
10434 .arg(y)
10435 .arg(&inf)
10436 .arg(&outf)
10437 .arg(&ns)
10438 .arg(&rb)
10439 .arg(&es)
10440 .arg(&ars)
10441 .arg(&adrs);
10442 unsafe {
10443 b.launch(cfg)?;
10444 }
10445 Ok(())
10446 }
10447
10448 #[allow(clippy::too_many_arguments)]
10451 pub fn qmatvec_nvfp4_bf16_sel_dual_rows_into(
10452 &self,
10453 gate_bank: &CudaSlice<u8>,
10454 up_bank: &CudaSlice<u8>,
10455 sel: &CudaSlice<i32>,
10456 token_rows: &CudaSlice<i32>,
10457 x_bf16: &CudaSlice<u8>,
10458 gate_out: &mut CudaSlice<f32>,
10459 up_out: &mut CudaSlice<f32>,
10460 n_sel: usize,
10461 in_f: usize,
10462 out_f: usize,
10463 row_bytes: usize,
10464 expert_stride: usize,
10465 tokens: usize,
10466 ) -> Result<(), Box<dyn std::error::Error>> {
10467 if !in_f.is_multiple_of(64)
10468 || sel.len() < n_sel
10469 || token_rows.len() < n_sel
10470 || gate_out.len() < n_sel * out_f
10471 || up_out.len() < n_sel * out_f
10472 || x_bf16.len() < 2 * in_f * tokens
10473 {
10474 return Err(format!(
10475 "W4A16 NVFP4 dual selected rows geometry sel={} token_rows={} x={} gate={} up={} \
10476 n_sel={n_sel} tokens={tokens} in={in_f} out={out_f}",
10477 sel.len(),
10478 token_rows.len(),
10479 x_bf16.len(),
10480 gate_out.len(),
10481 up_out.len(),
10482 )
10483 .into());
10484 }
10485 let adjacent_rows = tokens > 1;
10486 let f = if adjacent_rows {
10487 self.func("qmatvec_nvfp4_bf16_sel_quad_rows")
10488 } else {
10489 self.func("qmatvec_nvfp4_bf16_sel_dual_rows")
10490 };
10491 let cfg = LaunchConfig {
10492 grid_dim: (
10493 if adjacent_rows {
10494 out_f.div_ceil(2) as u32
10495 } else {
10496 (2 * out_f) as u32
10497 },
10498 n_sel as u32,
10499 1,
10500 ),
10501 block_dim: (256, 1, 1),
10502 shared_mem_bytes: 0,
10503 };
10504 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10505 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10506 let __s_b = self.gpu.stream();
10507 let mut b = __s_b.launch_builder(&f);
10508 b.arg(gate_bank)
10509 .arg(up_bank)
10510 .arg(sel)
10511 .arg(token_rows)
10512 .arg(x_bf16)
10513 .arg(gate_out)
10514 .arg(up_out)
10515 .arg(&inf)
10516 .arg(&outf)
10517 .arg(&ns)
10518 .arg(&rb)
10519 .arg(&es);
10520 unsafe {
10521 b.launch(cfg)?;
10522 }
10523 Ok(())
10524 }
10525
10526 #[allow(clippy::too_many_arguments)]
10529 pub fn qmatvec_nvfp4_bf16_ep_dual_slots_into(
10530 &self,
10531 gate_bank: &CudaSlice<u8>,
10532 up_bank: &CudaSlice<u8>,
10533 sel: &CudaSlice<i32>,
10534 x_bf16: &CudaSlice<u8>,
10535 gate_out: &mut CudaSlice<f32>,
10536 up_out: &mut CudaSlice<f32>,
10537 n_pairs: usize,
10538 top_k: usize,
10539 in_f: usize,
10540 out_f: usize,
10541 owner_start: usize,
10542 owner_end: usize,
10543 row_bytes: usize,
10544 expert_stride: usize,
10545 ) -> Result<(), Box<dyn std::error::Error>> {
10546 let tokens = n_pairs.div_ceil(top_k);
10547 if top_k == 0
10548 || owner_start >= owner_end
10549 || !in_f.is_multiple_of(64)
10550 || sel.len() < n_pairs
10551 || gate_out.len() < n_pairs * out_f
10552 || up_out.len() < n_pairs * out_f
10553 || x_bf16.len() < 2 * in_f * tokens
10554 {
10555 return Err(format!(
10556 "W4A16 NVFP4 device EP dual-slot geometry sel={} x={} gate={} up={} \
10557 pairs={n_pairs} top_k={top_k} in={in_f} out={out_f} \
10558 owner={owner_start}..{owner_end}",
10559 sel.len(),
10560 x_bf16.len(),
10561 gate_out.len(),
10562 up_out.len(),
10563 )
10564 .into());
10565 }
10566 let pair_parallel = tokens > 1;
10567 let f = if pair_parallel {
10568 self.func("qmatvec_nvfp4_bf16_ep_quad_pairs")
10569 } else {
10570 self.func("qmatvec_nvfp4_bf16_ep_dual_slots")
10571 };
10572 let cfg = LaunchConfig {
10573 grid_dim: (
10574 if pair_parallel {
10575 out_f.div_ceil(2) as u32
10576 } else {
10577 (2 * out_f) as u32
10578 },
10579 if pair_parallel { n_pairs as u32 } else { 1 },
10580 1,
10581 ),
10582 block_dim: (256, 1, 1),
10583 shared_mem_bytes: 0,
10584 };
10585 let (inf, outf, np, tk) = (in_f as i32, out_f as i32, n_pairs as i32, top_k as i32);
10586 let (os, oe) = (owner_start as i32, owner_end as i32);
10587 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10588 let __s_b = self.gpu.stream();
10589 let mut b = __s_b.launch_builder(&f);
10590 b.arg(gate_bank)
10591 .arg(up_bank)
10592 .arg(sel)
10593 .arg(x_bf16)
10594 .arg(gate_out)
10595 .arg(up_out)
10596 .arg(&inf)
10597 .arg(&outf)
10598 .arg(&np)
10599 .arg(&tk)
10600 .arg(&os)
10601 .arg(&oe)
10602 .arg(&rb)
10603 .arg(&es);
10604 unsafe {
10605 b.launch(cfg)?;
10606 }
10607 Ok(())
10608 }
10609
10610 #[allow(clippy::too_many_arguments)]
10612 pub fn qmatvec_nvfp4_q8_ep_dual_slots_into(
10613 &self,
10614 gate_bank: &CudaSlice<u8>,
10615 up_bank: &CudaSlice<u8>,
10616 sel: &CudaSlice<i32>,
10617 aq: &CudaSlice<i8>,
10618 ad: &CudaSlice<f32>,
10619 gate_out: &mut CudaSlice<f32>,
10620 up_out: &mut CudaSlice<f32>,
10621 n_pairs: usize,
10622 top_k: usize,
10623 in_f: usize,
10624 out_f: usize,
10625 owner_start: usize,
10626 owner_end: usize,
10627 row_bytes: usize,
10628 expert_stride: usize,
10629 ) -> Result<(), Box<dyn std::error::Error>> {
10630 let tokens = n_pairs.div_ceil(top_k);
10631 if top_k == 0
10632 || owner_start >= owner_end
10633 || !in_f.is_multiple_of(64)
10634 || sel.len() < n_pairs
10635 || aq.len() < tokens * in_f
10636 || ad.len() < tokens * (in_f / 32)
10637 || gate_out.len() < n_pairs * out_f
10638 || up_out.len() < n_pairs * out_f
10639 {
10640 return Err(format!(
10641 "W4A8 NVFP4 device EP gate/up geometry sel={} aq={} ad={} gate={} up={} \
10642 pairs={n_pairs} top_k={top_k} in={in_f} out={out_f} \
10643 owner={owner_start}..{owner_end}",
10644 sel.len(),
10645 aq.len(),
10646 ad.len(),
10647 gate_out.len(),
10648 up_out.len(),
10649 )
10650 .into());
10651 }
10652 let f = self.func("qmatvec_nvfp4_q8_ep_dual_slots");
10653 let threads = ((in_f / 32).div_ceil(32) * 32).clamp(32, 256) as u32;
10654 let cfg = LaunchConfig {
10655 grid_dim: (out_f as u32, 1, 1),
10656 block_dim: (threads, 1, 1),
10657 shared_mem_bytes: 0,
10658 };
10659 let (inf, outf, np, tk) = (in_f as i32, out_f as i32, n_pairs as i32, top_k as i32);
10660 let (os, oe) = (owner_start as i32, owner_end as i32);
10661 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10662 let __s_b = self.gpu.stream();
10663 let mut b = __s_b.launch_builder(&f);
10664 b.arg(gate_bank)
10665 .arg(up_bank)
10666 .arg(sel)
10667 .arg(aq)
10668 .arg(ad)
10669 .arg(gate_out)
10670 .arg(up_out)
10671 .arg(&inf)
10672 .arg(&outf)
10673 .arg(&np)
10674 .arg(&tk)
10675 .arg(&os)
10676 .arg(&oe)
10677 .arg(&rb)
10678 .arg(&es);
10679 unsafe {
10680 b.launch(cfg)?;
10681 }
10682 Ok(())
10683 }
10684
10685 #[allow(clippy::too_many_arguments)]
10688 pub fn qmatvec_nvfp4_q8_ep_paired_slots_into(
10689 &self,
10690 gate_bank: &CudaSlice<u8>,
10691 up_bank: &CudaSlice<u8>,
10692 sel: &CudaSlice<i32>,
10693 aq: &CudaSlice<i8>,
10694 ad: &CudaSlice<f32>,
10695 gate_out: &mut CudaSlice<f32>,
10696 up_out: &mut CudaSlice<f32>,
10697 n_pairs: usize,
10698 top_k: usize,
10699 in_f: usize,
10700 out_f: usize,
10701 owner_start: usize,
10702 owner_end: usize,
10703 row_bytes: usize,
10704 expert_stride: usize,
10705 ) -> Result<(), Box<dyn std::error::Error>> {
10706 let tokens = n_pairs.div_ceil(top_k);
10707 if top_k == 0
10708 || owner_start >= owner_end
10709 || !in_f.is_multiple_of(64)
10710 || sel.len() < n_pairs
10711 || aq.len() < tokens * in_f
10712 || ad.len() < tokens * (in_f / 32)
10713 || gate_out.len() < n_pairs * out_f
10714 || up_out.len() < n_pairs * out_f
10715 {
10716 return Err(format!(
10717 "W4A8 NVFP4 paired gate/up geometry sel={} aq={} ad={} gate={} up={} \
10718 pairs={n_pairs} top_k={top_k} in={in_f} out={out_f} \
10719 owner={owner_start}..{owner_end}",
10720 sel.len(),
10721 aq.len(),
10722 ad.len(),
10723 gate_out.len(),
10724 up_out.len(),
10725 )
10726 .into());
10727 }
10728 let f = self.func("qmatvec_nvfp4_q8_ep_paired_slots");
10729 let threads = ((in_f / 32).div_ceil(32) * 32).clamp(32, 256) as u32;
10730 let cfg = LaunchConfig {
10731 grid_dim: (out_f as u32, 1, 1),
10732 block_dim: (threads, 1, 1),
10733 shared_mem_bytes: 0,
10734 };
10735 let (inf, outf, np, tk) = (in_f as i32, out_f as i32, n_pairs as i32, top_k as i32);
10736 let (os, oe) = (owner_start as i32, owner_end as i32);
10737 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10738 let __s_b = self.gpu.stream();
10739 let mut b = __s_b.launch_builder(&f);
10740 b.arg(gate_bank)
10741 .arg(up_bank)
10742 .arg(sel)
10743 .arg(aq)
10744 .arg(ad)
10745 .arg(gate_out)
10746 .arg(up_out)
10747 .arg(&inf)
10748 .arg(&outf)
10749 .arg(&np)
10750 .arg(&tk)
10751 .arg(&os)
10752 .arg(&oe)
10753 .arg(&rb)
10754 .arg(&es);
10755 unsafe {
10756 b.launch(cfg)?;
10757 }
10758 Ok(())
10759 }
10760
10761 #[allow(clippy::too_many_arguments)]
10763 pub fn qmatvec_nvfp4_bf16_sel_down_rows_raw(
10764 &self,
10765 bank: &CudaSlice<u8>,
10766 sel: &CudaSlice<i32>,
10767 global_pairs: &CudaSlice<i32>,
10768 activation_bf16: &CudaSlice<u8>,
10769 macros_down: &CudaSlice<f32>,
10770 dst_raw: u64,
10771 n_sel: usize,
10772 in_f: usize,
10773 out_f: usize,
10774 row_bytes: usize,
10775 expert_stride: usize,
10776 total_pairs: usize,
10777 ) -> Result<(), Box<dyn std::error::Error>> {
10778 if !in_f.is_multiple_of(64)
10779 || sel.len() < n_sel
10780 || global_pairs.len() < n_sel
10781 || activation_bf16.len() < 2 * n_sel * in_f
10782 || dst_raw == 0
10783 {
10784 return Err(format!(
10785 "W4A16 NVFP4 down rows geometry sel={} pairs={} act={} dst_raw={dst_raw:#x} \
10786 n_sel={n_sel} total_pairs={total_pairs} in={in_f} out={out_f}",
10787 sel.len(),
10788 global_pairs.len(),
10789 activation_bf16.len(),
10790 )
10791 .into());
10792 }
10793 let f = self.func("qmatvec_nvfp4_bf16_sel_down_rows");
10794 let cfg = LaunchConfig {
10795 grid_dim: (out_f.div_ceil(2) as u32, n_sel as u32, 1),
10796 block_dim: (256, 1, 1),
10797 shared_mem_bytes: 0,
10798 };
10799 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
10800 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10801 let __s_b = self.gpu.stream();
10802 let mut b = __s_b.launch_builder(&f);
10803 b.arg(bank)
10804 .arg(sel)
10805 .arg(global_pairs)
10806 .arg(activation_bf16)
10807 .arg(macros_down)
10808 .arg(&dst_raw)
10809 .arg(&inf)
10810 .arg(&outf)
10811 .arg(&ns)
10812 .arg(&rb)
10813 .arg(&es);
10814 unsafe {
10815 b.launch(cfg)?;
10816 }
10817 Ok(())
10818 }
10819
10820 #[allow(clippy::too_many_arguments)]
10823 pub fn qmatvec_nvfp4_bf16_ep_down_slots_raw(
10824 &self,
10825 bank: &CudaSlice<u8>,
10826 sel: &CudaSlice<i32>,
10827 activation_bf16: &CudaSlice<u8>,
10828 macros_down: &CudaSlice<f32>,
10829 dst_raw: u64,
10830 n_pairs: usize,
10831 in_f: usize,
10832 out_f: usize,
10833 owner_start: usize,
10834 owner_end: usize,
10835 row_bytes: usize,
10836 expert_stride: usize,
10837 ) -> Result<(), Box<dyn std::error::Error>> {
10838 if owner_start >= owner_end
10839 || !in_f.is_multiple_of(64)
10840 || sel.len() < n_pairs
10841 || activation_bf16.len() < 2 * n_pairs * in_f
10842 || dst_raw == 0
10843 {
10844 return Err(format!(
10845 "W4A16 NVFP4 device EP down-slot geometry sel={} act={} dst_raw={dst_raw:#x} \
10846 pairs={n_pairs} in={in_f} out={out_f} owner={owner_start}..{owner_end}",
10847 sel.len(),
10848 activation_bf16.len(),
10849 )
10850 .into());
10851 }
10852 let f = self.func("qmatvec_nvfp4_bf16_ep_down_slots");
10853 let cfg = LaunchConfig {
10854 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
10855 block_dim: (256, 1, 1),
10856 shared_mem_bytes: 0,
10857 };
10858 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
10859 let (os, oe) = (owner_start as i32, owner_end as i32);
10860 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10861 let __s_b = self.gpu.stream();
10862 let mut b = __s_b.launch_builder(&f);
10863 b.arg(bank)
10864 .arg(sel)
10865 .arg(activation_bf16)
10866 .arg(macros_down)
10867 .arg(&dst_raw)
10868 .arg(&inf)
10869 .arg(&outf)
10870 .arg(&np)
10871 .arg(&os)
10872 .arg(&oe)
10873 .arg(&rb)
10874 .arg(&es);
10875 unsafe {
10876 b.launch(cfg)?;
10877 }
10878 Ok(())
10879 }
10880
10881 #[allow(clippy::too_many_arguments)]
10883 pub fn qmatvec_nvfp4_bf16_ep_down_pairs_raw(
10884 &self,
10885 bank: &CudaSlice<u8>,
10886 sel: &CudaSlice<i32>,
10887 activation_bf16: &CudaSlice<u8>,
10888 macros_down: &CudaSlice<f32>,
10889 dst_raw: u64,
10890 n_pairs: usize,
10891 in_f: usize,
10892 out_f: usize,
10893 owner_start: usize,
10894 owner_end: usize,
10895 row_bytes: usize,
10896 expert_stride: usize,
10897 ) -> Result<(), Box<dyn std::error::Error>> {
10898 if owner_start >= owner_end
10899 || !in_f.is_multiple_of(64)
10900 || sel.len() < n_pairs
10901 || activation_bf16.len() < 2 * n_pairs * in_f
10902 || dst_raw == 0
10903 {
10904 return Err(format!(
10905 "W4A16 NVFP4 device EP down-pair geometry sel={} act={} dst_raw={dst_raw:#x} \
10906 pairs={n_pairs} in={in_f} out={out_f} owner={owner_start}..{owner_end}",
10907 sel.len(),
10908 activation_bf16.len(),
10909 )
10910 .into());
10911 }
10912 let f = self.func("qmatvec_nvfp4_bf16_ep_down_pairs");
10913 let cfg = LaunchConfig {
10914 grid_dim: (out_f.div_ceil(2) as u32, n_pairs as u32, 1),
10915 block_dim: (256, 1, 1),
10916 shared_mem_bytes: 0,
10917 };
10918 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
10919 let (os, oe) = (owner_start as i32, owner_end as i32);
10920 let (rb, es) = (row_bytes as i64, expert_stride as i64);
10921 let __s_b = self.gpu.stream();
10922 let mut b = __s_b.launch_builder(&f);
10923 b.arg(bank)
10924 .arg(sel)
10925 .arg(activation_bf16)
10926 .arg(macros_down)
10927 .arg(&dst_raw)
10928 .arg(&inf)
10929 .arg(&outf)
10930 .arg(&np)
10931 .arg(&os)
10932 .arg(&oe)
10933 .arg(&rb)
10934 .arg(&es);
10935 unsafe {
10936 b.launch(cfg)?;
10937 }
10938 Ok(())
10939 }
10940
10941 #[allow(clippy::too_many_arguments)]
10943 pub fn silu_mul_scaled_host_expf_bf16_sel_into(
10944 &self,
10945 gate: &CudaSlice<f32>,
10946 up: &CudaSlice<f32>,
10947 gate_macros: &CudaSlice<f32>,
10948 up_macros: &CudaSlice<f32>,
10949 sel: &CudaSlice<i32>,
10950 limit: Option<f32>,
10951 output_bf16: &mut CudaSlice<u8>,
10952 n_per: usize,
10953 n_sel: usize,
10954 ) -> Result<(), Box<dyn std::error::Error>> {
10955 let n = n_per * n_sel;
10956 if sel.len() < n_sel || gate.len() < n || up.len() < n || output_bf16.len() < 2 * n {
10957 return Err(format!(
10958 "W4A16 selected activation geometry sel={} gate={} up={} out={} \
10959 n_per={n_per} n_sel={n_sel}",
10960 sel.len(),
10961 gate.len(),
10962 up.len(),
10963 output_bf16.len(),
10964 )
10965 .into());
10966 }
10967 let (limit, has_limit) = match limit {
10968 Some(limit) if limit.is_finite() && limit > 1e-6 => (limit, 1i32),
10969 Some(limit) => {
10970 return Err(format!("W4A16 selected activation limit {limit} is invalid").into());
10971 }
10972 None => (0.0f32, 0i32),
10973 };
10974 let f = self.func("silu_mul_scaled_host_expf_bf16_sel");
10975 let cfg = LaunchConfig::for_num_elems(n as u32);
10976 let (np, ns) = (n_per as i32, n_sel as i32);
10977 let __s_b = self.gpu.stream();
10978 let mut b = __s_b.launch_builder(&f);
10979 b.arg(gate)
10980 .arg(up)
10981 .arg(gate_macros)
10982 .arg(up_macros)
10983 .arg(sel)
10984 .arg(&limit)
10985 .arg(&has_limit)
10986 .arg(output_bf16)
10987 .arg(&np)
10988 .arg(&ns);
10989 unsafe {
10990 b.launch(cfg)?;
10991 }
10992 Ok(())
10993 }
10994
10995 #[allow(clippy::too_many_arguments)]
10998 pub fn silu_mul_scaled_host_expf_bf16_ep_slots_into(
10999 &self,
11000 gate: &CudaSlice<f32>,
11001 up: &CudaSlice<f32>,
11002 gate_macros: &CudaSlice<f32>,
11003 up_macros: &CudaSlice<f32>,
11004 sel: &CudaSlice<i32>,
11005 owner_start: usize,
11006 owner_end: usize,
11007 limit: Option<f32>,
11008 output_bf16: &mut CudaSlice<u8>,
11009 n_per: usize,
11010 n_pairs: usize,
11011 ) -> Result<(), Box<dyn std::error::Error>> {
11012 let n = n_per * n_pairs;
11013 if owner_start >= owner_end
11014 || sel.len() < n_pairs
11015 || gate.len() < n
11016 || up.len() < n
11017 || output_bf16.len() < 2 * n
11018 {
11019 return Err(format!(
11020 "W4A16 device EP activation geometry sel={} gate={} up={} out={} \
11021 n_per={n_per} pairs={n_pairs} owner={owner_start}..{owner_end}",
11022 sel.len(),
11023 gate.len(),
11024 up.len(),
11025 output_bf16.len(),
11026 )
11027 .into());
11028 }
11029 let (limit, has_limit) = match limit {
11030 Some(limit) if limit.is_finite() && limit > 1e-6 => (limit, 1i32),
11031 Some(limit) => {
11032 return Err(format!("W4A16 selected activation limit {limit} is invalid").into());
11033 }
11034 None => (0.0f32, 0i32),
11035 };
11036 let f = self.func("silu_mul_scaled_host_expf_bf16_ep_slots");
11037 let cfg = LaunchConfig::for_num_elems(n as u32);
11038 let (np, pairs) = (n_per as i32, n_pairs as i32);
11039 let (os, oe) = (owner_start as i32, owner_end as i32);
11040 let __s_b = self.gpu.stream();
11041 let mut b = __s_b.launch_builder(&f);
11042 b.arg(gate)
11043 .arg(up)
11044 .arg(gate_macros)
11045 .arg(up_macros)
11046 .arg(sel)
11047 .arg(&limit)
11048 .arg(&has_limit)
11049 .arg(output_bf16)
11050 .arg(&np)
11051 .arg(&pairs)
11052 .arg(&os)
11053 .arg(&oe);
11054 unsafe {
11055 b.launch(cfg)?;
11056 }
11057 Ok(())
11058 }
11059
11060 #[allow(clippy::too_many_arguments)]
11062 pub fn silu_mul_scaled_host_expf_q8_ep_slots_into(
11063 &self,
11064 gate: &CudaSlice<f32>,
11065 up: &CudaSlice<f32>,
11066 gate_macros: &CudaSlice<f32>,
11067 up_macros: &CudaSlice<f32>,
11068 sel: &CudaSlice<i32>,
11069 owner_start: usize,
11070 owner_end: usize,
11071 limit: Option<f32>,
11072 output_q8: &mut CudaSlice<i8>,
11073 output_scales: &mut CudaSlice<f32>,
11074 n_per: usize,
11075 n_pairs: usize,
11076 ) -> Result<(), Box<dyn std::error::Error>> {
11077 let n = n_per * n_pairs;
11078 if owner_start >= owner_end
11079 || !n_per.is_multiple_of(32)
11080 || sel.len() < n_pairs
11081 || gate.len() < n
11082 || up.len() < n
11083 || output_q8.len() < n
11084 || output_scales.len() < n / 32
11085 {
11086 return Err(format!(
11087 "W4A8 device EP activation geometry sel={} gate={} up={} q8={} scales={} \
11088 n_per={n_per} pairs={n_pairs} owner={owner_start}..{owner_end}",
11089 sel.len(),
11090 gate.len(),
11091 up.len(),
11092 output_q8.len(),
11093 output_scales.len(),
11094 )
11095 .into());
11096 }
11097 let (limit, has_limit) = match limit {
11098 Some(limit) if limit.is_finite() && limit > 1e-6 => (limit, 1i32),
11099 Some(limit) => {
11100 return Err(format!("W4A8 selected activation limit {limit} is invalid").into());
11101 }
11102 None => (0.0f32, 0i32),
11103 };
11104 let f = self.func("silu_mul_scaled_host_expf_q8_ep_slots");
11105 let warps = n / 32;
11106 let cfg = LaunchConfig {
11107 grid_dim: ((warps as u32).div_ceil(4), 1, 1),
11108 block_dim: (128, 1, 1),
11109 shared_mem_bytes: 0,
11110 };
11111 let (np, pairs) = (n_per as i32, n_pairs as i32);
11112 let (os, oe) = (owner_start as i32, owner_end as i32);
11113 let __s_b = self.gpu.stream();
11114 let mut b = __s_b.launch_builder(&f);
11115 b.arg(gate)
11116 .arg(up)
11117 .arg(gate_macros)
11118 .arg(up_macros)
11119 .arg(sel)
11120 .arg(&limit)
11121 .arg(&has_limit)
11122 .arg(output_q8)
11123 .arg(output_scales)
11124 .arg(&np)
11125 .arg(&pairs)
11126 .arg(&os)
11127 .arg(&oe);
11128 unsafe {
11129 b.launch(cfg)?;
11130 }
11131 Ok(())
11132 }
11133
11134 #[allow(clippy::too_many_arguments)]
11137 pub fn qmatvec_nvfp4_bf16_sel_down_fma_into(
11138 &self,
11139 bank: &CudaSlice<u8>,
11140 sel: &CudaSlice<i32>,
11141 activation_bf16: &CudaSlice<u8>,
11142 route_weights: &CudaSlice<f32>,
11143 macros_down: &CudaSlice<f32>,
11144 dst: &mut cudarc::driver::CudaViewMut<f32>,
11145 n_sel: usize,
11146 in_f: usize,
11147 out_f: usize,
11148 row_bytes: usize,
11149 expert_stride: usize,
11150 ) -> Result<(), Box<dyn std::error::Error>> {
11151 if !in_f.is_multiple_of(64)
11152 || sel.len() < n_sel
11153 || route_weights.len() < n_sel
11154 || activation_bf16.len() < 2 * n_sel * in_f
11155 || dst.len() < out_f
11156 {
11157 return Err(format!(
11158 "W4A16 NVFP4 down selected geometry sel={} act={} weights={} dst={} \
11159 n_sel={n_sel} in={in_f} out={out_f}",
11160 sel.len(),
11161 activation_bf16.len(),
11162 route_weights.len(),
11163 dst.len(),
11164 )
11165 .into());
11166 }
11167 let f = self.func("qmatvec_nvfp4_bf16_sel_down_fma");
11168 let cfg = LaunchConfig {
11169 grid_dim: (out_f as u32, 1, 1),
11170 block_dim: (256, 1, 1),
11171 shared_mem_bytes: 0,
11172 };
11173 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
11174 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11175 let __s_b = self.gpu.stream();
11176 let mut b = __s_b.launch_builder(&f);
11177 b.arg(bank)
11178 .arg(sel)
11179 .arg(activation_bf16)
11180 .arg(route_weights)
11181 .arg(macros_down)
11182 .arg(dst)
11183 .arg(&inf)
11184 .arg(&outf)
11185 .arg(&ns)
11186 .arg(&rb)
11187 .arg(&es);
11188 unsafe {
11189 b.launch(cfg)?;
11190 }
11191 Ok(())
11192 }
11193
11194 #[allow(clippy::too_many_arguments)]
11196 pub fn qmatvec_nvfp4_bf16_ep_down_fma_into(
11197 &self,
11198 bank: &CudaSlice<u8>,
11199 sel: &CudaSlice<i32>,
11200 activation_bf16: &CudaSlice<u8>,
11201 route_weights: &CudaSlice<f32>,
11202 macros_down: &CudaSlice<f32>,
11203 dst: &mut cudarc::driver::CudaViewMut<f32>,
11204 n_pairs: usize,
11205 in_f: usize,
11206 out_f: usize,
11207 owner_start: usize,
11208 owner_end: usize,
11209 row_bytes: usize,
11210 expert_stride: usize,
11211 ) -> Result<(), Box<dyn std::error::Error>> {
11212 if owner_start >= owner_end
11213 || !in_f.is_multiple_of(64)
11214 || sel.len() < n_pairs
11215 || route_weights.len() < n_pairs
11216 || activation_bf16.len() < 2 * n_pairs * in_f
11217 || dst.len() < out_f
11218 {
11219 return Err(format!(
11220 "W4A16 device EP down-FMA geometry sel={} act={} weights={} dst={} \
11221 pairs={n_pairs} in={in_f} out={out_f} owner={owner_start}..{owner_end}",
11222 sel.len(),
11223 activation_bf16.len(),
11224 route_weights.len(),
11225 dst.len(),
11226 )
11227 .into());
11228 }
11229 let f = self.func("qmatvec_nvfp4_bf16_ep_down_fma");
11230 let cfg = LaunchConfig {
11231 grid_dim: (out_f as u32, 1, 1),
11232 block_dim: (256, 1, 1),
11233 shared_mem_bytes: 0,
11234 };
11235 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
11236 let (os, oe) = (owner_start as i32, owner_end as i32);
11237 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11238 let __s_b = self.gpu.stream();
11239 let mut b = __s_b.launch_builder(&f);
11240 b.arg(bank)
11241 .arg(sel)
11242 .arg(activation_bf16)
11243 .arg(route_weights)
11244 .arg(macros_down)
11245 .arg(dst)
11246 .arg(&inf)
11247 .arg(&outf)
11248 .arg(&np)
11249 .arg(&os)
11250 .arg(&oe)
11251 .arg(&rb)
11252 .arg(&es);
11253 unsafe {
11254 b.launch(cfg)?;
11255 }
11256 Ok(())
11257 }
11258
11259 #[allow(clippy::too_many_arguments)]
11262 pub fn qmatvec_nvfp4_bf16_ep_down_fma_raw(
11263 &self,
11264 bank: &CudaSlice<u8>,
11265 sel: &CudaSlice<i32>,
11266 activation_bf16: &CudaSlice<u8>,
11267 route_weights: &CudaSlice<f32>,
11268 macros_down: &CudaSlice<f32>,
11269 dst_raw: u64,
11270 n_pairs: usize,
11271 in_f: usize,
11272 out_f: usize,
11273 owner_start: usize,
11274 owner_end: usize,
11275 row_bytes: usize,
11276 expert_stride: usize,
11277 ) -> Result<(), Box<dyn std::error::Error>> {
11278 if dst_raw == 0
11279 || owner_start >= owner_end
11280 || !in_f.is_multiple_of(64)
11281 || sel.len() < n_pairs
11282 || route_weights.len() < n_pairs
11283 || activation_bf16.len() < 2 * n_pairs * in_f
11284 {
11285 return Err(format!(
11286 "W4A16 device EP raw down-FMA geometry sel={} act={} weights={} \
11287 dst={dst_raw:#x} pairs={n_pairs} in={in_f} out={out_f} \
11288 owner={owner_start}..{owner_end}",
11289 sel.len(),
11290 activation_bf16.len(),
11291 route_weights.len(),
11292 )
11293 .into());
11294 }
11295 let f = self.func("qmatvec_nvfp4_bf16_ep_down_fma");
11296 let cfg = LaunchConfig {
11297 grid_dim: (out_f as u32, 1, 1),
11298 block_dim: (256, 1, 1),
11299 shared_mem_bytes: 0,
11300 };
11301 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
11302 let (os, oe) = (owner_start as i32, owner_end as i32);
11303 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11304 let __s_b = self.gpu.stream();
11305 let mut b = __s_b.launch_builder(&f);
11306 b.arg(bank)
11307 .arg(sel)
11308 .arg(activation_bf16)
11309 .arg(route_weights)
11310 .arg(macros_down)
11311 .arg(&dst_raw)
11312 .arg(&inf)
11313 .arg(&outf)
11314 .arg(&np)
11315 .arg(&os)
11316 .arg(&oe)
11317 .arg(&rb)
11318 .arg(&es);
11319 unsafe {
11320 b.launch(cfg)?;
11321 }
11322 Ok(())
11323 }
11324
11325 #[allow(clippy::too_many_arguments)]
11328 pub fn qmatvec_nvfp4_q8_ep_down_slots_raw(
11329 &self,
11330 bank: &CudaSlice<u8>,
11331 sel: &CudaSlice<i32>,
11332 aq: &CudaSlice<i8>,
11333 ad: &CudaSlice<f32>,
11334 macros_down: &CudaSlice<f32>,
11335 dst_raw: u64,
11336 n_pairs: usize,
11337 in_f: usize,
11338 out_f: usize,
11339 owner_start: usize,
11340 owner_end: usize,
11341 row_bytes: usize,
11342 expert_stride: usize,
11343 ) -> Result<(), Box<dyn std::error::Error>> {
11344 if dst_raw == 0
11345 || owner_start >= owner_end
11346 || !in_f.is_multiple_of(64)
11347 || sel.len() < n_pairs
11348 || aq.len() < n_pairs * in_f
11349 || ad.len() < n_pairs * (in_f / 32)
11350 {
11351 return Err(format!(
11352 "W4A8 device EP raw down-slot geometry sel={} aq={} ad={} \
11353 dst={dst_raw:#x} pairs={n_pairs} in={in_f} out={out_f} \
11354 owner={owner_start}..{owner_end}",
11355 sel.len(),
11356 aq.len(),
11357 ad.len(),
11358 )
11359 .into());
11360 }
11361 let f = self.func("qmatvec_nvfp4_q8_ep_down_slots");
11362 let threads = ((in_f / 32).div_ceil(32) * 32).clamp(32, 256) as u32;
11363 let cfg = LaunchConfig {
11364 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
11365 block_dim: (threads, 1, 1),
11366 shared_mem_bytes: 0,
11367 };
11368 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
11369 let (os, oe) = (owner_start as i32, owner_end as i32);
11370 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11371 let __s_b = self.gpu.stream();
11372 let mut b = __s_b.launch_builder(&f);
11373 b.arg(bank)
11374 .arg(sel)
11375 .arg(aq)
11376 .arg(ad)
11377 .arg(macros_down)
11378 .arg(&dst_raw)
11379 .arg(&inf)
11380 .arg(&outf)
11381 .arg(&np)
11382 .arg(&os)
11383 .arg(&oe)
11384 .arg(&rb)
11385 .arg(&es);
11386 unsafe {
11387 b.launch(cfg)?;
11388 }
11389 Ok(())
11390 }
11391
11392 #[allow(clippy::too_many_arguments)]
11394 pub fn qmatvec_nvfp4_q8_ep_down_fma_raw(
11395 &self,
11396 bank: &CudaSlice<u8>,
11397 sel: &CudaSlice<i32>,
11398 aq: &CudaSlice<i8>,
11399 ad: &CudaSlice<f32>,
11400 route_weights: &CudaSlice<f32>,
11401 macros_down: &CudaSlice<f32>,
11402 dst_raw: u64,
11403 n_pairs: usize,
11404 in_f: usize,
11405 out_f: usize,
11406 owner_start: usize,
11407 owner_end: usize,
11408 row_bytes: usize,
11409 expert_stride: usize,
11410 ) -> Result<(), Box<dyn std::error::Error>> {
11411 if dst_raw == 0
11412 || owner_start >= owner_end
11413 || !in_f.is_multiple_of(64)
11414 || sel.len() < n_pairs
11415 || aq.len() < n_pairs * in_f
11416 || ad.len() < n_pairs * (in_f / 32)
11417 || route_weights.len() < n_pairs
11418 {
11419 return Err(format!(
11420 "W4A8 device EP raw down-FMA geometry sel={} aq={} ad={} weights={} \
11421 dst={dst_raw:#x} pairs={n_pairs} in={in_f} out={out_f} \
11422 owner={owner_start}..{owner_end}",
11423 sel.len(),
11424 aq.len(),
11425 ad.len(),
11426 route_weights.len(),
11427 )
11428 .into());
11429 }
11430 let f = self.func("qmatvec_nvfp4_q8_ep_down_fma");
11431 let cfg = LaunchConfig {
11432 grid_dim: (out_f as u32, 1, 1),
11433 block_dim: (256, 1, 1),
11434 shared_mem_bytes: 0,
11435 };
11436 let (inf, outf, np) = (in_f as i32, out_f as i32, n_pairs as i32);
11437 let (os, oe) = (owner_start as i32, owner_end as i32);
11438 let (rb, es) = (row_bytes as i64, expert_stride as i64);
11439 let __s_b = self.gpu.stream();
11440 let mut b = __s_b.launch_builder(&f);
11441 b.arg(bank)
11442 .arg(sel)
11443 .arg(aq)
11444 .arg(ad)
11445 .arg(route_weights)
11446 .arg(macros_down)
11447 .arg(&dst_raw)
11448 .arg(&inf)
11449 .arg(&outf)
11450 .arg(&np)
11451 .arg(&os)
11452 .arg(&oe)
11453 .arg(&rb)
11454 .arg(&es);
11455 unsafe {
11456 b.launch(cfg)?;
11457 }
11458 Ok(())
11459 }
11460
11461 #[allow(clippy::too_many_arguments)]
11466 pub fn silu_mul_scaled_q8_1_sel_into(
11467 &self,
11468 gate: &CudaSlice<f32>,
11469 up: &CudaSlice<f32>,
11470 gmac: &CudaSlice<f32>,
11471 umac: &CudaSlice<f32>,
11472 sel: &CudaSlice<i32>,
11473 limit: Option<f32>,
11474 out_q: &mut CudaSlice<i8>,
11475 out_d: &mut CudaSlice<f32>,
11476 n_per: usize,
11477 n_sel: usize,
11478 ) -> Result<(), Box<dyn std::error::Error>> {
11479 let n = n_per * n_sel;
11480 if !n_per.is_multiple_of(32) || out_q.len() < n || out_d.len() < n / 32 {
11481 return Err(format!(
11482 "silu sel geometry n_per={n_per} n_sel={n_sel} q={} d={}",
11483 out_q.len(),
11484 out_d.len()
11485 )
11486 .into());
11487 }
11488 if let Some(limit) = limit {
11489 if limit <= 1e-6 {
11490 return Err(format!(
11491 "silu sel clamp limit {limit} is at or below the 1e-6 eps gate"
11492 )
11493 .into());
11494 }
11495 let f = self.func("silu_mul_scaled_q8_1_sel_clamp");
11496 let cfg = LaunchConfig::for_num_elems(n as u32);
11497 let (np, ns) = (n_per as i32, n_sel as i32);
11498 let __s_b = self.gpu.stream();
11499 let mut b = __s_b.launch_builder(&f);
11500 b.arg(gate)
11501 .arg(up)
11502 .arg(gmac)
11503 .arg(umac)
11504 .arg(sel)
11505 .arg(&limit)
11506 .arg(out_q)
11507 .arg(out_d)
11508 .arg(&np)
11509 .arg(&ns);
11510 unsafe {
11511 b.launch(cfg)?;
11512 }
11513 return Ok(());
11514 }
11515 let f = self.func("silu_mul_scaled_q8_1_sel");
11516 let cfg = LaunchConfig::for_num_elems(n as u32);
11517 let (np, ns) = (n_per as i32, n_sel as i32);
11518 let __s_b = self.gpu.stream();
11519 let mut b = __s_b.launch_builder(&f);
11520 b.arg(gate)
11521 .arg(up)
11522 .arg(gmac)
11523 .arg(umac)
11524 .arg(sel)
11525 .arg(out_q)
11526 .arg(out_d)
11527 .arg(&np)
11528 .arg(&ns);
11529 unsafe {
11530 b.launch(cfg)?;
11531 }
11532 Ok(())
11533 }
11534
11535 pub fn htod(&self, v: &[f32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11536 Ok(self.gpu.stream().clone_htod(v)?)
11537 }
11538 pub fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
11539 Ok(self.gpu.stream().clone_htod(v)?)
11540 }
11541 pub fn htod_i8(&self, v: &[i8]) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
11543 Ok(self.gpu.stream().clone_htod(v)?)
11544 }
11545 pub fn htod_u64(&self, v: &[u64]) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
11546 Ok(self.gpu.stream().clone_htod(v)?)
11547 }
11548 pub fn dtoh_view(
11550 &self,
11551 d: &cudarc::driver::CudaView<f32>,
11552 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
11553 let v = self.gpu.stream().clone_dtoh(d)?;
11554 self.gpu.stream().synchronize()?;
11555 Ok(v)
11556 }
11557 pub fn dtoh(&self, d: &CudaSlice<f32>) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
11558 let v = self.gpu.stream().clone_dtoh(d)?;
11559 self.gpu.stream().synchronize()?;
11560 Ok(v)
11561 }
11562 pub fn dtoh_pair(
11566 &self,
11567 a: &CudaSlice<f32>,
11568 b: &CudaSlice<f32>,
11569 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
11570 let av = self.gpu.stream().clone_dtoh(a)?;
11571 let bv = self.gpu.stream().clone_dtoh(b)?;
11572 self.gpu.stream().synchronize()?;
11573 Ok((av, bv))
11574 }
11575 pub fn dtoh_pair_views(
11578 &self,
11579 a: &cudarc::driver::CudaView<f32>,
11580 b: &cudarc::driver::CudaView<f32>,
11581 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
11582 let av = self.gpu.stream().clone_dtoh(a)?;
11583 let bv = self.gpu.stream().clone_dtoh(b)?;
11584 self.gpu.stream().synchronize()?;
11585 Ok((av, bv))
11586 }
11587 pub fn dtoh_i32(&self, d: &CudaSlice<i32>) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
11589 let v = self.gpu.stream().clone_dtoh(d)?;
11590 self.gpu.stream().synchronize()?;
11591 Ok(v)
11592 }
11593 pub fn dtoh_u8(&self, d: &CudaSlice<u8>) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
11595 let v = self.gpu.stream().clone_dtoh(d)?;
11596 self.gpu.stream().synchronize()?;
11597 Ok(v)
11598 }
11599 pub fn dtoh_u8_view(
11600 &self,
11601 d: &cudarc::driver::CudaView<u8>,
11602 ) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
11603 let v = self.gpu.stream().clone_dtoh(d)?;
11604 self.gpu.stream().synchronize()?;
11605 Ok(v)
11606 }
11607 pub fn dtoh_u8_into_pinned(
11615 &self,
11616 d: &CudaSlice<u8>,
11617 dst: &mut PinnedHostBuf,
11618 n: usize,
11619 ) -> Result<(), Box<dyn std::error::Error>> {
11620 if n > d.len() || n > dst.len() {
11621 return Err(format!(
11622 "dtoh_u8_into_pinned range {n} exceeds src {} or pinned dst {}",
11623 d.len(),
11624 dst.len(),
11625 )
11626 .into());
11627 }
11628 if n == 0 {
11629 return Ok(());
11630 }
11631 let host = &mut dst.as_mut_slice()[..n];
11632 self.gpu.stream().memcpy_dtoh(&d.slice(0..n), host)?;
11633 self.gpu.stream().synchronize()?;
11634 Ok(())
11635 }
11636 pub fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11637 SCRATCH_ALLOC_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11638 let s = self.gpu.stream().alloc_zeros::<f32>(n)?;
11639 self.keep_if_capturing(&s);
11640 Ok(s)
11641 }
11642
11643 pub(crate) fn hyper_ws_take(&self) -> Option<crate::hyper::HyperDecodeWs> {
11647 self.hyper_decode_ws.lock().unwrap().take()
11648 }
11649
11650 pub(crate) fn hyper_ws_put(&self, ws: crate::hyper::HyperDecodeWs) {
11651 *self.hyper_decode_ws.lock().unwrap() = Some(ws);
11652 }
11653
11654 pub(crate) fn vws_uninit(
11663 &self,
11664 n: usize,
11665 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11666 if verify_ws_on() {
11667 let mut ws = self.verify_ws.lock().unwrap();
11668 let ws = &mut *ws;
11669 if let Some(s) = VerifyWs::take(&mut ws.f32_pool, &mut ws.held_bytes, n) {
11670 if VERIFY_WS_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
11671 eprintln!(
11672 "[glm5-verify-ws] engaged: verify-walk buffers recycling through \
11673 the size-keyed pool (MEMRA_GLM5_VERIFY_WS=1)"
11674 );
11675 }
11676 return Ok(s);
11677 }
11678 }
11679 self.alloc_uninit::<f32>(n)
11680 }
11681
11682 pub(crate) fn vws_uninit_i8(
11684 &self,
11685 n: usize,
11686 ) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
11687 if verify_ws_on() {
11688 let mut ws = self.verify_ws.lock().unwrap();
11689 let ws = &mut *ws;
11690 if let Some(s) = VerifyWs::take(&mut ws.i8_pool, &mut ws.held_bytes, n) {
11691 VERIFY_WS_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11692 return Ok(s);
11693 }
11694 }
11695 self.alloc_uninit::<i8>(n)
11696 }
11697
11698 pub(crate) fn vws_uninit_u64(
11700 &self,
11701 n: usize,
11702 ) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
11703 if verify_ws_on() {
11704 let mut ws = self.verify_ws.lock().unwrap();
11705 let ws = &mut *ws;
11706 if let Some(s) = VerifyWs::take(&mut ws.u64_pool, &mut ws.held_bytes, n) {
11707 VERIFY_WS_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11708 return Ok(s);
11709 }
11710 }
11711 self.alloc_uninit::<u64>(n)
11712 }
11713
11714 pub(crate) fn vws_recycle(&self, s: CudaSlice<f32>) {
11717 if verify_ws_on() {
11718 let mut ws = self.verify_ws.lock().unwrap();
11719 let ws = &mut *ws;
11720 VerifyWs::put(&mut ws.f32_pool, &mut ws.held_bytes, s);
11721 }
11722 }
11723
11724 pub(crate) fn vws_recycle_i8(&self, s: CudaSlice<i8>) {
11726 if verify_ws_on() {
11727 let mut ws = self.verify_ws.lock().unwrap();
11728 let ws = &mut *ws;
11729 VerifyWs::put(&mut ws.i8_pool, &mut ws.held_bytes, s);
11730 }
11731 }
11732
11733 pub(crate) fn vws_recycle_u64(&self, s: CudaSlice<u64>) {
11735 if verify_ws_on() {
11736 let mut ws = self.verify_ws.lock().unwrap();
11737 let ws = &mut *ws;
11738 VerifyWs::put(&mut ws.u64_pool, &mut ws.held_bytes, s);
11739 }
11740 }
11741
11742 pub fn prob_of_token_device(
11751 &self,
11752 logits: &CudaSlice<f32>,
11753 tok: &CudaSlice<u32>,
11754 n_vocab: usize,
11755 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11756 let nb = ARGMAX_NB;
11757 let mut part = self.alloc_uninit::<f32>(nb)?;
11758 let mut p = self.alloc_uninit::<f32>(1)?;
11759 let f1 = self.func("prob_of_token_partial_f32");
11760 let cfg1 = LaunchConfig {
11761 grid_dim: (nb as u32, 1, 1),
11762 block_dim: (256, 1, 1),
11763 shared_mem_bytes: 0,
11764 };
11765 let nv = n_vocab as i32;
11766 let __s_b1 = self.gpu.stream();
11767 let mut b1 = __s_b1.launch_builder(&f1);
11768 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
11769 unsafe {
11770 b1.launch(cfg1)?;
11771 }
11772 let f2 = self.func("prob_of_token_final_f32");
11773 let cfg2 = LaunchConfig {
11774 grid_dim: (1, 1, 1),
11775 block_dim: (256, 1, 1),
11776 shared_mem_bytes: 0,
11777 };
11778 let nbi = nb as i32;
11779 let __s_b2 = self.gpu.stream();
11780 let mut b2 = __s_b2.launch_builder(&f2);
11781 b2.arg(&part).arg(&mut p).arg(&nbi);
11782 unsafe {
11783 b2.launch(cfg2)?;
11784 }
11785 Ok(p)
11786 }
11787
11788 pub fn prob_of_token_device_col(
11795 &self,
11796 logits: &CudaSlice<f32>,
11797 tok_all: &CudaSlice<u32>,
11798 tok_idx: usize,
11799 p_out: &mut CudaSlice<f32>,
11800 p_idx: usize,
11801 n_vocab: usize,
11802 ) -> Result<(), Box<dyn std::error::Error>> {
11803 let tok_v = tok_all.slice(tok_idx..tok_idx + 1);
11804 let mut p_v = p_out.slice_mut(p_idx..p_idx + 1);
11805 let nb = ARGMAX_NB;
11806 let mut part = self.alloc_uninit::<f32>(nb)?;
11807 let f1 = self.func("prob_of_token_partial_f32");
11808 let cfg1 = LaunchConfig {
11809 grid_dim: (nb as u32, 1, 1),
11810 block_dim: (256, 1, 1),
11811 shared_mem_bytes: 0,
11812 };
11813 let nv = n_vocab as i32;
11814 let __s_b1 = self.gpu.stream();
11815 let mut b1 = __s_b1.launch_builder(&f1);
11816 b1.arg(logits).arg(&tok_v).arg(&mut part).arg(&nv);
11817 unsafe {
11818 b1.launch(cfg1)?;
11819 }
11820 let f2 = self.func("prob_of_token_final_f32");
11821 let cfg2 = LaunchConfig {
11822 grid_dim: (1, 1, 1),
11823 block_dim: (256, 1, 1),
11824 shared_mem_bytes: 0,
11825 };
11826 let nbi = nb as i32;
11827 let __s_b2 = self.gpu.stream();
11828 let mut b2 = __s_b2.launch_builder(&f2);
11829 b2.arg(&part).arg(&mut p_v).arg(&nbi);
11830 unsafe {
11831 b2.launch(cfg2)?;
11832 }
11833 Ok(())
11834 }
11835
11836 pub fn prob_of_token_device_into(
11837 &self,
11838 logits: &CudaSlice<f32>,
11839 tok: &CudaSlice<u32>,
11840 p_out: &mut CudaSlice<f32>,
11841 n_vocab: usize,
11842 ) -> Result<(), Box<dyn std::error::Error>> {
11843 let nb = ARGMAX_NB;
11844 let mut part = self.alloc_uninit::<f32>(nb)?;
11845 let f1 = self.func("prob_of_token_partial_f32");
11846 let cfg1 = LaunchConfig {
11847 grid_dim: (nb as u32, 1, 1),
11848 block_dim: (256, 1, 1),
11849 shared_mem_bytes: 0,
11850 };
11851 let nv = n_vocab as i32;
11852 let __s_b1 = self.gpu.stream();
11853 let mut b1 = __s_b1.launch_builder(&f1);
11854 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
11855 unsafe {
11856 b1.launch(cfg1)?;
11857 }
11858 let f2 = self.func("prob_of_token_final_f32");
11859 let cfg2 = LaunchConfig {
11860 grid_dim: (1, 1, 1),
11861 block_dim: (256, 1, 1),
11862 shared_mem_bytes: 0,
11863 };
11864 let nbi = nb as i32;
11865 let __s_b2 = self.gpu.stream();
11866 let mut b2 = __s_b2.launch_builder(&f2);
11867 b2.arg(&part).arg(p_out).arg(&nbi);
11868 unsafe {
11869 b2.launch(cfg2)?;
11870 }
11871 Ok(())
11872 }
11873
11874 pub fn u32_hist_append(
11877 &self,
11878 tok: &CudaSlice<u32>,
11879 hist: &mut CudaSlice<u32>,
11880 idx: &mut CudaSlice<i32>,
11881 ) -> Result<(), Box<dyn std::error::Error>> {
11882 let f = self.func("u32_hist_append");
11883 let cfg = LaunchConfig {
11884 grid_dim: (1, 1, 1),
11885 block_dim: (32, 1, 1),
11886 shared_mem_bytes: 0,
11887 };
11888 let __s_b = self.gpu.stream();
11889 let mut b = __s_b.launch_builder(&f);
11890 b.arg(tok).arg(&mut *hist).arg(&mut *idx);
11891 unsafe {
11892 b.launch(cfg)?;
11893 }
11894 Ok(())
11895 }
11896
11897 pub fn argmax_token_device(
11898 &self,
11899 logits: &CudaSlice<f32>,
11900 n_vocab: usize,
11901 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
11902 let mut tok = unsafe { self.gpu.stream().alloc::<u32>(1)? };
11903 self.argmax_token_device_into(logits, &mut tok, n_vocab)?;
11904 Ok(tok)
11905 }
11906 pub fn argmax_token_device_into(
11913 &self,
11914 logits: &CudaSlice<f32>,
11915 tok: &mut CudaSlice<u32>,
11916 n_vocab: usize,
11917 ) -> Result<(), Box<dyn std::error::Error>> {
11918 let nb = ARGMAX_NB;
11919 let f1 = self.func("argmax_partial_f32");
11920 let f2 = self.func("argmax_final_f32");
11921 let mut guard = self.argmax_partials.lock().unwrap();
11922 if guard.is_none() {
11923 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
11926 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
11927 *guard = Some((pv, pi));
11928 }
11929 let (part_v, part_i) = guard.as_mut().unwrap();
11930 let nv = n_vocab as i32;
11931 let nbi = nb as i32;
11932 let cfg1 = LaunchConfig {
11934 grid_dim: (nb as u32, 1, 1),
11935 block_dim: (256, 1, 1),
11936 shared_mem_bytes: 0,
11937 };
11938 let __s_b1 = self.gpu.stream();
11939 let mut b1 = __s_b1.launch_builder(&f1);
11940 b1.arg(logits).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
11941 unsafe {
11942 b1.launch(cfg1)?;
11943 }
11944 let cfg2 = LaunchConfig {
11946 grid_dim: (1, 1, 1),
11947 block_dim: (256, 1, 1),
11948 shared_mem_bytes: 0,
11949 };
11950 let __s_b2 = self.gpu.stream();
11951 let mut b2 = __s_b2.launch_builder(&f2);
11952 b2.arg(&*part_v).arg(&*part_i).arg(tok).arg(&nbi);
11953 unsafe {
11954 b2.launch(cfg2)?;
11955 }
11956 Ok(())
11957 }
11958 pub fn argmax_token_device_col(
11964 &self,
11965 logits: &CudaSlice<f32>,
11966 col: usize,
11967 n_vocab: usize,
11968 toks: &mut CudaSlice<u32>,
11969 out_idx: usize,
11970 ) -> Result<(), Box<dyn std::error::Error>> {
11971 let nb = ARGMAX_NB;
11972 let f1 = self.func("argmax_partial_f32");
11973 let f2 = self.func("argmax_final_f32");
11974 let mut guard = self.argmax_partials.lock().unwrap();
11975 if guard.is_none() {
11976 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
11977 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
11978 *guard = Some((pv, pi));
11979 }
11980 let (part_v, part_i) = guard.as_mut().unwrap();
11981 let col_view = logits.slice(col * n_vocab..(col + 1) * n_vocab);
11982 let nv = n_vocab as i32;
11983 let nbi = nb as i32;
11984 let cfg1 = LaunchConfig {
11985 grid_dim: (nb as u32, 1, 1),
11986 block_dim: (256, 1, 1),
11987 shared_mem_bytes: 0,
11988 };
11989 let __s_b1 = self.gpu.stream();
11990 let mut b1 = __s_b1.launch_builder(&f1);
11991 b1.arg(&col_view)
11992 .arg(&mut *part_v)
11993 .arg(&mut *part_i)
11994 .arg(&nv);
11995 unsafe {
11996 b1.launch(cfg1)?;
11997 }
11998 let mut tok_view = toks.slice_mut(out_idx..out_idx + 1);
11999 let cfg2 = LaunchConfig {
12000 grid_dim: (1, 1, 1),
12001 block_dim: (256, 1, 1),
12002 shared_mem_bytes: 0,
12003 };
12004 let __s_b2 = self.gpu.stream();
12005 let mut b2 = __s_b2.launch_builder(&f2);
12006 b2.arg(&*part_v).arg(&*part_i).arg(&mut tok_view).arg(&nbi);
12007 unsafe {
12008 b2.launch(cfg2)?;
12009 }
12010 Ok(())
12011 }
12012 pub fn htod_u32_v(&self, v: &[u32]) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
12014 Ok(self.gpu.stream().clone_htod(v)?)
12015 }
12016 pub fn dtoh_u64(&self, d: &CudaSlice<u64>) -> Result<Vec<u64>, Box<dyn std::error::Error>> {
12017 let v = self.gpu.stream().clone_dtoh(d)?;
12018 self.gpu.stream().synchronize()?;
12019 Ok(v)
12020 }
12021
12022 pub fn dtoh_u32(&self, d: &CudaSlice<u32>) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
12023 let v = self.gpu.stream().clone_dtoh(d)?;
12024 self.gpu.stream().synchronize()?;
12025 Ok(v)
12026 }
12027 pub fn htod_u32_into(
12031 &self,
12032 dst: &mut CudaSlice<u32>,
12033 src: &[u32],
12034 ) -> Result<(), Box<dyn std::error::Error>> {
12035 let mut view = dst.slice_mut(0..src.len());
12036 self.gpu.stream().memcpy_htod(src, &mut view)?;
12037 Ok(())
12038 }
12039
12040 pub fn htod_i32_into(
12043 &self,
12044 dst: &mut CudaSlice<i32>,
12045 src: &[i32],
12046 ) -> Result<(), Box<dyn std::error::Error>> {
12047 let mut view = dst.slice_mut(0..src.len());
12048 self.gpu.stream().memcpy_htod(src, &mut view)?;
12049 Ok(())
12050 }
12051
12052 pub fn alloc_u32_zeroed(&self, n: usize) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
12053 let s = self.gpu.stream().alloc_zeros::<u32>(n)?;
12054 self.keep_if_capturing(&s);
12055 Ok(s)
12056 }
12057 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device_into(
12061 &self,
12062 embd: &CudaSlice<u8>,
12063 token_d: &CudaSlice<u32>,
12064 x_out: &mut CudaSlice<f32>,
12065 n_embd: usize,
12066 qtype: i32,
12067 row_bytes: usize,
12068 ) -> Result<(), Box<dyn std::error::Error>> {
12069 let f = self.func("embed_gather_u32");
12070 let cfg = LaunchConfig {
12071 grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
12072 block_dim: (256, 1, 1),
12073 shared_mem_bytes: 0,
12074 };
12075 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
12076 let __s_b = self.gpu.stream();
12077 let mut b = __s_b.launch_builder(&f);
12078 b.arg(embd)
12079 .arg(token_d)
12080 .arg(x_out)
12081 .arg(&ne)
12082 .arg(&qt)
12083 .arg(&rb);
12084 unsafe {
12085 b.launch(cfg)?;
12086 }
12087 Ok(())
12088 }
12089 pub fn dtoh_i32_one(&self, d: &CudaSlice<i32>) -> Result<i32, Box<dyn std::error::Error>> {
12091 let v = self.gpu.stream().clone_dtoh(d)?;
12092 self.gpu.stream().synchronize()?;
12093 Ok(v[0])
12094 }
12095 pub fn i32_set_k(
12102 &self,
12103 dst: &mut CudaSlice<i32>,
12104 v: i32,
12105 ) -> Result<(), Box<dyn std::error::Error>> {
12106 let f = self.func("i32_set_k");
12107 let cfg = LaunchConfig {
12108 grid_dim: (1, 1, 1),
12109 block_dim: (1, 1, 1),
12110 shared_mem_bytes: 0,
12111 };
12112 let idx = 0i32;
12113 let __s_b = self.gpu.stream();
12114 let mut b = __s_b.launch_builder(&f);
12115 b.arg(dst).arg(&v).arg(&idx);
12116 unsafe {
12117 b.launch(cfg)?;
12118 }
12119 Ok(())
12120 }
12121
12122 pub fn set_i32_one(
12123 &self,
12124 d: &mut CudaSlice<i32>,
12125 v: i32,
12126 ) -> Result<(), Box<dyn std::error::Error>> {
12127 self.gpu.stream().memcpy_htod(&[v], d)?;
12128 Ok(())
12129 }
12130 pub fn set_u32_one(
12133 &self,
12134 d: &mut CudaSlice<u32>,
12135 v: u32,
12136 ) -> Result<(), Box<dyn std::error::Error>> {
12137 self.gpu.stream().memcpy_htod(&[v], d)?;
12138 Ok(())
12139 }
12140 pub fn dtoh_u32_one(&self, d: &CudaSlice<u32>) -> Result<u32, Box<dyn std::error::Error>> {
12142 let v = self.gpu.stream().clone_dtoh(d)?;
12143 self.gpu.stream().synchronize()?;
12144 Ok(v[0])
12145 }
12146 pub fn upload_u8(&self, bytes: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
12148 Ok(self.gpu.stream().clone_htod(bytes)?)
12149 }
12150 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device(
12155 &self,
12156 embd: &CudaSlice<u8>,
12157 token_d: &CudaSlice<u32>,
12158 n_embd: usize,
12159 qtype: i32,
12160 row_bytes: usize,
12161 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12162 let f = self.func("embed_gather_u32");
12163 let mut x = self.alloc_uninit::<f32>(n_embd)?;
12164 let cfg = LaunchConfig {
12165 grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
12166 block_dim: (256, 1, 1),
12167 shared_mem_bytes: 0,
12168 };
12169 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
12170 let __s_b = self.gpu.stream();
12171 let mut b = __s_b.launch_builder(&f);
12172 b.arg(embd)
12173 .arg(token_d)
12174 .arg(&mut x)
12175 .arg(&ne)
12176 .arg(&qt)
12177 .arg(&rb);
12178 unsafe {
12179 b.launch(cfg)?;
12180 }
12181 Ok(x)
12182 }
12183
12184 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device_t(
12189 &self,
12190 embd: &CudaSlice<u8>,
12191 tokens: &[u32],
12192 n_embd: usize,
12193 qtype: i32,
12194 row_bytes: usize,
12195 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12196 let t = tokens.len();
12197 let tok_d = self.gpu.stream().clone_htod(tokens)?;
12198 let f = self.func("embed_gather_u32_t");
12199 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
12200 let cfg = LaunchConfig {
12201 grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
12202 block_dim: (256, 1, 1),
12203 shared_mem_bytes: 0,
12204 };
12205 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
12206 let __s_b = self.gpu.stream();
12207 let mut b = __s_b.launch_builder(&f);
12208 b.arg(embd)
12209 .arg(&tok_d)
12210 .arg(&mut x)
12211 .arg(&ne)
12212 .arg(&qt)
12213 .arg(&rb)
12214 .arg(&ti);
12215 unsafe {
12216 b.launch(cfg)?;
12217 }
12218 Ok(x)
12219 }
12220
12221 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device_tv(
12227 &self,
12228 embd: &CudaSlice<u8>,
12229 tok_v: &cudarc::driver::CudaView<u32>,
12230 t: usize,
12231 n_embd: usize,
12232 qtype: i32,
12233 row_bytes: usize,
12234 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12235 let f = self.func("embed_gather_u32_t");
12236 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
12237 let cfg = LaunchConfig {
12238 grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
12239 block_dim: (256, 1, 1),
12240 shared_mem_bytes: 0,
12241 };
12242 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
12243 let __s_b = self.gpu.stream();
12244 let mut b = __s_b.launch_builder(&f);
12245 b.arg(embd)
12246 .arg(tok_v)
12247 .arg(&mut x)
12248 .arg(&ne)
12249 .arg(&qt)
12250 .arg(&rb)
12251 .arg(&ti);
12252 unsafe {
12253 b.launch(cfg)?;
12254 }
12255 Ok(x)
12256 }
12257
12258 #[allow(clippy::manual_div_ceil)] pub fn embed_gather_device_td(
12260 &self,
12261 embd: &CudaSlice<u8>,
12262 tok_d: &CudaSlice<u32>,
12263 t: usize,
12264 n_embd: usize,
12265 qtype: i32,
12266 row_bytes: usize,
12267 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12268 let f = self.func("embed_gather_u32_t");
12269 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
12270 let cfg = LaunchConfig {
12271 grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
12272 block_dim: (256, 1, 1),
12273 shared_mem_bytes: 0,
12274 };
12275 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
12276 let __s_b = self.gpu.stream();
12277 let mut b = __s_b.launch_builder(&f);
12278 b.arg(embd)
12279 .arg(tok_d)
12280 .arg(&mut x)
12281 .arg(&ne)
12282 .arg(&qt)
12283 .arg(&rb)
12284 .arg(&ti);
12285 unsafe {
12286 b.launch(cfg)?;
12287 }
12288 Ok(x)
12289 }
12290
12291 #[inline]
12297 fn keep_if_capturing<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, s: &CudaSlice<T>) {
12299 if self
12300 .capture_keep_on
12301 .load(std::sync::atomic::Ordering::Relaxed)
12302 {
12303 self.capture_keep.lock().unwrap().push(Box::new(s.clone()));
12304 }
12305 }
12306
12307 fn alloc_uninit<T: cudarc::driver::DeviceRepr + Send + 'static>(
12308 &self,
12309 n: usize,
12310 ) -> Result<CudaSlice<T>, Box<dyn std::error::Error>> {
12311 SCRATCH_ALLOC_CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12312 let mut s = unsafe { self.gpu.stream().alloc::<T>(n)? };
12313 {
12317 static Z: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12318 if *Z.get_or_init(|| std::env::var("MEMRA_DEBUG_ZERO_ALLOCS").as_deref() == Ok("1")) {
12319 use cudarc::driver::DevicePtrMut;
12321 let n_bytes = s.len() * std::mem::size_of::<T>();
12322 let stream = self.gpu.stream();
12323 let (p_, _g) = s.device_ptr_mut(&stream);
12324 unsafe {
12325 cudarc::driver::sys::cuMemsetD8Async(p_, 0, n_bytes, stream.cu_stream())
12326 .result()?;
12327 }
12328 }
12329 }
12330 self.keep_if_capturing(&s);
12331 Ok(s)
12332 }
12333
12334 pub fn uninit_q8_pair(
12339 &self,
12340 n: usize,
12341 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12342 Ok((
12343 self.alloc_uninit::<i8>(n)?,
12344 self.alloc_uninit::<f32>(n / 32)?,
12345 ))
12346 }
12347
12348 pub fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12349 self.alloc_uninit::<f32>(n)
12350 }
12351
12352 pub fn alloc_i8_uninit(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
12354 self.alloc_uninit::<i8>(n)
12355 }
12356
12357 pub fn uninit_i32(&self, n: usize) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
12359 self.alloc_uninit::<i32>(n)
12360 }
12361
12362 #[allow(clippy::too_many_arguments)]
12366 pub fn rms_norm3(
12367 &self,
12368 x: &CudaSlice<f32>,
12369 w0: &CudaSlice<f32>,
12370 w1: &CudaSlice<f32>,
12371 w2: &CudaSlice<f32>,
12372 d0: &mut CudaSlice<f32>,
12373 d1: &mut CudaSlice<f32>,
12374 d2: &mut CudaSlice<f32>,
12375 ncols: usize,
12376 nrows: usize,
12377 eps: f32,
12378 ) -> Result<(), Box<dyn std::error::Error>> {
12379 let f = self.func("rms_norm3_f32");
12380 let cfg = LaunchConfig {
12381 grid_dim: (nrows as u32, 1, 1),
12382 block_dim: (rms_block(), 1, 1),
12383 shared_mem_bytes: 0,
12384 };
12385 let (nc, e) = (ncols as i32, eps);
12386 let __s_b = self.gpu.stream();
12387 let mut b = __s_b.launch_builder(&f);
12388 b.arg(x)
12389 .arg(w0)
12390 .arg(w1)
12391 .arg(w2)
12392 .arg(d0)
12393 .arg(d1)
12394 .arg(d2)
12395 .arg(&nc)
12396 .arg(&e);
12397 unsafe {
12398 b.launch(cfg)?;
12399 }
12400 Ok(())
12401 }
12402
12403 #[allow(clippy::too_many_arguments)]
12405 pub fn qkvnorm_w_on_prefill(rows: usize, ncols: usize) -> bool {
12408 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12409 *WARP_ON.get_or_init(|| {
12410 std::env::var("MEMRA_QKVNORM_W")
12411 .map(|v| v != "0")
12412 .unwrap_or(true)
12413 }) && ncols.is_multiple_of(4)
12414 && rows >= 64
12415 }
12416
12417 #[allow(clippy::too_many_arguments)]
12420 pub fn rms_norm_qkv_w4b(
12421 &self,
12422 q: &CudaSlice<f32>,
12423 k: &CudaSlice<f32>,
12424 v: &CudaSlice<f32>,
12425 wq: &CudaSlice<f32>,
12426 wk: &CudaSlice<f32>,
12427 wv: &CudaSlice<f32>,
12428 dq: &mut CudaSlice<f32>,
12429 dk: &mut CudaSlice<f32>,
12430 dv: &mut CudaSlice<f32>,
12431 dvb: &mut CudaSlice<u8>,
12432 ncols: usize,
12433 rq: usize,
12434 rk: usize,
12435 eps: f32,
12436 vf16: bool,
12437 ) -> Result<(), Box<dyn std::error::Error>> {
12438 assert!(ncols.is_multiple_of(4) && rq + 2 * rk >= 64);
12439 let f = self.func("rms_norm_qkv_w4b_f32");
12440 let rows = (rq + 2 * rk) as u32;
12441 let cfg = LaunchConfig {
12442 grid_dim: (rows.div_ceil(8), 1, 1),
12443 block_dim: (256, 1, 1),
12444 shared_mem_bytes: 0,
12445 };
12446 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
12447 let vf = vf16 as i32;
12448 let __s_b = self.gpu.stream();
12449 let mut b = __s_b.launch_builder(&f);
12450 b.arg(q)
12451 .arg(k)
12452 .arg(v)
12453 .arg(wq)
12454 .arg(wk)
12455 .arg(wv)
12456 .arg(dq)
12457 .arg(dk)
12458 .arg(dv)
12459 .arg(&mut *dvb)
12460 .arg(&nc)
12461 .arg(&rqi)
12462 .arg(&rki)
12463 .arg(&rvi)
12464 .arg(&e)
12465 .arg(&vf);
12466 unsafe {
12467 b.launch(cfg)?;
12468 }
12469 Ok(())
12470 }
12471
12472 #[allow(clippy::too_many_arguments)] pub fn rms_norm_qkv(
12474 &self,
12475 q: &CudaSlice<f32>,
12476 k: &CudaSlice<f32>,
12477 v: &CudaSlice<f32>,
12478 wq: &CudaSlice<f32>,
12479 wk: &CudaSlice<f32>,
12480 wv: &CudaSlice<f32>,
12481 dq: &mut CudaSlice<f32>,
12482 dk: &mut CudaSlice<f32>,
12483 dv: &mut CudaSlice<f32>,
12484 ncols: usize,
12485 rq: usize,
12486 rk: usize,
12487 eps: f32,
12488 ) -> Result<(), Box<dyn std::error::Error>> {
12489 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12493 let warp_on = *WARP_ON.get_or_init(|| {
12494 std::env::var("MEMRA_QKVNORM_W")
12495 .map(|v| v != "0")
12496 .unwrap_or(true)
12497 });
12498 if warp_on && ncols.is_multiple_of(4) && rq + 2 * rk >= 64 {
12501 let f = self.func("rms_norm_qkv_w4_f32");
12502 let rows = (rq + 2 * rk) as u32;
12503 let cfg = LaunchConfig {
12504 grid_dim: (rows.div_ceil(8), 1, 1),
12505 block_dim: (256, 1, 1),
12506 shared_mem_bytes: 0,
12507 };
12508 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
12509 let __s_b = self.gpu.stream();
12510 let mut b = __s_b.launch_builder(&f);
12511 b.arg(q)
12512 .arg(k)
12513 .arg(v)
12514 .arg(wq)
12515 .arg(wk)
12516 .arg(wv)
12517 .arg(dq)
12518 .arg(dk)
12519 .arg(dv)
12520 .arg(&nc)
12521 .arg(&rqi)
12522 .arg(&rki)
12523 .arg(&rvi)
12524 .arg(&e);
12525 unsafe {
12526 b.launch(cfg)?;
12527 }
12528 return Ok(());
12529 }
12530 let f = self.func("rms_norm_qkv_f32");
12531 let grid = (rq + 2 * rk) as u32;
12532 let cfg = LaunchConfig {
12533 grid_dim: (grid, 1, 1),
12534 block_dim: (rms_block(), 1, 1),
12535 shared_mem_bytes: 0,
12536 };
12537 let (nc, rqi, rki, e) = (ncols as i32, rq as i32, rk as i32, eps);
12538 let __s_b = self.gpu.stream();
12539 let mut b = __s_b.launch_builder(&f);
12540 b.arg(q)
12541 .arg(k)
12542 .arg(v)
12543 .arg(wq)
12544 .arg(wk)
12545 .arg(wv)
12546 .arg(dq)
12547 .arg(dk)
12548 .arg(dv)
12549 .arg(&nc)
12550 .arg(&rqi)
12551 .arg(&rki)
12552 .arg(&e);
12553 unsafe {
12554 b.launch(cfg)?;
12555 }
12556 Ok(())
12557 }
12558
12559 #[allow(clippy::too_many_arguments)]
12561 pub fn rms_norm2x(
12562 &self,
12563 a: &CudaSlice<f32>,
12564 bb: &CudaSlice<f32>,
12565 wa: &CudaSlice<f32>,
12566 wb: &CudaSlice<f32>,
12567 da: &mut CudaSlice<f32>,
12568 db: &mut CudaSlice<f32>,
12569 ncols: usize,
12570 nrows: usize,
12571 eps: f32,
12572 ) -> Result<(), Box<dyn std::error::Error>> {
12573 let f = self.func("rms_norm2x_f32");
12574 let cfg = LaunchConfig {
12575 grid_dim: (2 * nrows as u32, 1, 1),
12576 block_dim: (rms_block(), 1, 1),
12577 shared_mem_bytes: 0,
12578 };
12579 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
12580 let __s_b = self.gpu.stream();
12581 let mut b = __s_b.launch_builder(&f);
12582 b.arg(a)
12583 .arg(bb)
12584 .arg(wa)
12585 .arg(wb)
12586 .arg(da)
12587 .arg(db)
12588 .arg(&nc)
12589 .arg(&nr)
12590 .arg(&e);
12591 unsafe {
12592 b.launch(cfg)?;
12593 }
12594 Ok(())
12595 }
12596
12597 pub fn softcap(
12599 &self,
12600 y: &mut CudaSlice<f32>,
12601 cap: f32,
12602 n: usize,
12603 ) -> Result<(), Box<dyn std::error::Error>> {
12604 let f = self.func("softcap_f32");
12605 let cfg = LaunchConfig::for_num_elems(n as u32);
12606 let ni = n as i32;
12607 let __s_b = self.gpu.stream();
12608 let mut b = __s_b.launch_builder(&f);
12609 b.arg(y).arg(&cap).arg(&ni);
12610 unsafe {
12611 b.launch(cfg)?;
12612 }
12613 Ok(())
12614 }
12615
12616 pub fn mask_ids_rows(
12619 &self,
12620 y: &mut CudaSlice<f32>,
12621 ids: &CudaSlice<i32>,
12622 n_ids: usize,
12623 n_vocab: usize,
12624 t: usize,
12625 ) -> Result<(), Box<dyn std::error::Error>> {
12626 let f = self.func("mask_ids_rows_f32");
12627 let cfg = LaunchConfig::for_num_elems((n_ids * t) as u32);
12628 let (ni, nv, ti) = (n_ids as i32, n_vocab as i32, t as i32);
12629 let __s_b = self.gpu.stream();
12630 let mut b = __s_b.launch_builder(&f);
12631 b.arg(y).arg(ids).arg(&ni).arg(&nv).arg(&ti);
12632 unsafe {
12633 b.launch(cfg)?;
12634 }
12635 Ok(())
12636 }
12637
12638 #[allow(clippy::too_many_arguments)]
12640 pub fn add_scale_rms_norm(
12641 &self,
12642 a: &CudaSlice<f32>,
12643 b_in: &CudaSlice<f32>,
12644 c: f32,
12645 w: &CudaSlice<f32>,
12646 res: &mut CudaSlice<f32>,
12647 dst: &mut CudaSlice<f32>,
12648 ncols: usize,
12649 nrows: usize,
12650 eps: f32,
12651 ) -> Result<(), Box<dyn std::error::Error>> {
12652 let f = self.func("add_scale_rms_norm_f32");
12653 let cfg = LaunchConfig {
12654 grid_dim: (nrows as u32, 1, 1),
12655 block_dim: (rms_block(), 1, 1),
12656 shared_mem_bytes: 0,
12657 };
12658 let (nc, e2) = (ncols as i32, eps);
12659 let __s_b = self.gpu.stream();
12660 let mut b = __s_b.launch_builder(&f);
12661 b.arg(a)
12662 .arg(b_in)
12663 .arg(&c)
12664 .arg(w)
12665 .arg(res)
12666 .arg(dst)
12667 .arg(&nc)
12668 .arg(&e2);
12669 unsafe {
12670 b.launch(cfg)?;
12671 }
12672 Ok(())
12673 }
12674
12675 #[allow(clippy::too_many_arguments)]
12678 pub fn add_scale_rms_norm_q8_1(
12679 &self,
12680 a: &CudaSlice<f32>,
12681 b_in: &CudaSlice<f32>,
12682 c: f32,
12683 w: &CudaSlice<f32>,
12684 res: &mut CudaSlice<f32>,
12685 ncols: usize,
12686 nrows: usize,
12687 eps: f32,
12688 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12689 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12690 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12691 let (nc, e2) = (ncols as i32, eps);
12692 if Self::pdl_on() && Self::pdl_wb_on() {
12693 {
12694 use cudarc::driver::{DevicePtr, DevicePtrMut};
12695 let s = &self.gpu.stream();
12696 let (pa, _g0) = a.device_ptr(s);
12697 let (pb, _g1) = b_in.device_ptr(s);
12698 let (pw, _g2) = w.device_ptr(s);
12699 let (pr, _g3) = res.device_ptr_mut(s);
12700 let (pq, _g4) = out_q.device_ptr_mut(s);
12701 let (pd, _g5) = out_d.device_ptr_mut(s);
12702 let mut ps = [
12703 &pa as *const _ as *mut std::ffi::c_void,
12704 &pb as *const _ as *mut _,
12705 &c as *const _ as *mut _,
12706 &pw as *const _ as *mut _,
12707 &pr as *const _ as *mut _,
12708 &pq as *const _ as *mut _,
12709 &pd as *const _ as *mut _,
12710 &nc as *const _ as *mut _,
12711 &e2 as *const _ as *mut _,
12712 ];
12713 unsafe {
12714 self.launch_pdl(
12715 "add_scale_rms_norm_q8_1",
12716 (nrows as u32, 1, 1),
12717 (rms_block(), 1, 1),
12718 &mut ps,
12719 )?;
12720 }
12721 }
12722 return Ok((out_q, out_d));
12723 }
12724 let f = self.func("add_scale_rms_norm_q8_1");
12725 let cfg = LaunchConfig {
12726 grid_dim: (nrows as u32, 1, 1),
12727 block_dim: (rms_block(), 1, 1),
12728 shared_mem_bytes: 0,
12729 };
12730 let __s_b = self.gpu.stream();
12731 let mut b = __s_b.launch_builder(&f);
12732 b.arg(a)
12733 .arg(b_in)
12734 .arg(&c)
12735 .arg(w)
12736 .arg(res)
12737 .arg(&mut out_q)
12738 .arg(&mut out_d)
12739 .arg(&nc)
12740 .arg(&e2);
12741 unsafe {
12742 b.launch(cfg)?;
12743 }
12744 Ok((out_q, out_d))
12745 }
12746
12747 #[allow(clippy::too_many_arguments)]
12749 pub fn add_scale_rms_norm_q8_1_into(
12750 &self,
12751 a: &CudaSlice<f32>,
12752 b_in: &CudaSlice<f32>,
12753 c: f32,
12754 w: &CudaSlice<f32>,
12755 res: &mut CudaSlice<f32>,
12756 ncols: usize,
12757 nrows: usize,
12758 eps: f32,
12759 out_q: &mut CudaSlice<i8>,
12760 out_d: &mut CudaSlice<f32>,
12761 ) -> Result<(), Box<dyn std::error::Error>> {
12762 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
12763 let (nc, e2) = (ncols as i32, eps);
12764 if Self::pdl_on() && Self::pdl_wb_on() {
12765 use cudarc::driver::{DevicePtr, DevicePtrMut};
12766 let s = &self.gpu.stream();
12767 let (pa, _g0) = a.device_ptr(s);
12768 let (pb, _g1) = b_in.device_ptr(s);
12769 let (pw, _g2) = w.device_ptr(s);
12770 let (pr, _g3) = res.device_ptr_mut(s);
12771 let (pq, _g4) = out_q.device_ptr_mut(s);
12772 let (pd, _g5) = out_d.device_ptr_mut(s);
12773 let mut ps = [
12774 &pa as *const _ as *mut std::ffi::c_void,
12775 &pb as *const _ as *mut _,
12776 &c as *const _ as *mut _,
12777 &pw as *const _ as *mut _,
12778 &pr as *const _ as *mut _,
12779 &pq as *const _ as *mut _,
12780 &pd as *const _ as *mut _,
12781 &nc as *const _ as *mut _,
12782 &e2 as *const _ as *mut _,
12783 ];
12784 unsafe {
12785 self.launch_pdl(
12786 "add_scale_rms_norm_q8_1",
12787 (nrows as u32, 1, 1),
12788 (rms_block(), 1, 1),
12789 &mut ps,
12790 )?;
12791 }
12792 return Ok(());
12793 }
12794 let f = self.func("add_scale_rms_norm_q8_1");
12795 let cfg = LaunchConfig {
12796 grid_dim: (nrows as u32, 1, 1),
12797 block_dim: (rms_block(), 1, 1),
12798 shared_mem_bytes: 0,
12799 };
12800 let __s_b = self.gpu.stream();
12801 let mut b = __s_b.launch_builder(&f);
12802 b.arg(a)
12803 .arg(b_in)
12804 .arg(&c)
12805 .arg(w)
12806 .arg(res)
12807 .arg(&mut *out_q)
12808 .arg(&mut *out_d)
12809 .arg(&nc)
12810 .arg(&e2);
12811 unsafe {
12812 b.launch(cfg)?;
12813 }
12814 Ok(())
12815 }
12816
12817 #[allow(clippy::too_many_arguments)]
12820 pub fn rms_pre_add_scale_rms_norm_q8_1(
12821 &self,
12822 a: &CudaSlice<f32>,
12823 wa: &CudaSlice<f32>,
12824 b_in: &CudaSlice<f32>,
12825 c: f32,
12826 w: &CudaSlice<f32>,
12827 res: &mut CudaSlice<f32>,
12828 ncols: usize,
12829 nrows: usize,
12830 eps: f32,
12831 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12832 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12833 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12834 let (nc, e2) = (ncols as i32, eps);
12835 if Self::pdl_on() {
12836 {
12837 use cudarc::driver::{DevicePtr, DevicePtrMut};
12838 let s = &self.gpu.stream();
12839 let (pa, _g0) = a.device_ptr(s);
12840 let (pwa, _g1) = wa.device_ptr(s);
12841 let (pb, _g2) = b_in.device_ptr(s);
12842 let (pw, _g3) = w.device_ptr(s);
12843 let (pr, _g4) = res.device_ptr_mut(s);
12844 let (pq, _g5) = out_q.device_ptr_mut(s);
12845 let (pd, _g6) = out_d.device_ptr_mut(s);
12846 let mut ps = [
12847 &pa as *const _ as *mut std::ffi::c_void,
12848 &pwa as *const _ as *mut _,
12849 &pb as *const _ as *mut _,
12850 &c as *const _ as *mut _,
12851 &pw as *const _ as *mut _,
12852 &pr as *const _ as *mut _,
12853 &pq as *const _ as *mut _,
12854 &pd as *const _ as *mut _,
12855 &nc as *const _ as *mut _,
12856 &e2 as *const _ as *mut _,
12857 ];
12858 unsafe {
12859 self.launch_pdl(
12860 "rms_pre_add_scale_rms_norm_q8_1",
12861 (nrows as u32, 1, 1),
12862 (rms_block(), 1, 1),
12863 &mut ps,
12864 )?;
12865 }
12866 }
12867 return Ok((out_q, out_d));
12868 }
12869 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
12870 let cfg = LaunchConfig {
12871 grid_dim: (nrows as u32, 1, 1),
12872 block_dim: (rms_block(), 1, 1),
12873 shared_mem_bytes: 0,
12874 };
12875 let __s_b = self.gpu.stream();
12876 let mut b = __s_b.launch_builder(&f);
12877 b.arg(a)
12878 .arg(wa)
12879 .arg(b_in)
12880 .arg(&c)
12881 .arg(w)
12882 .arg(res)
12883 .arg(&mut out_q)
12884 .arg(&mut out_d)
12885 .arg(&nc)
12886 .arg(&e2);
12887 unsafe {
12888 b.launch(cfg)?;
12889 }
12890 Ok((out_q, out_d))
12891 }
12892
12893 pub fn gelu_tanh_mul_q8_1(
12896 &self,
12897 gate: &CudaSlice<f32>,
12898 up: &cudarc::driver::CudaView<f32>,
12899 act: &mut CudaSlice<f32>,
12900 ncols: usize,
12901 nrows: usize,
12902 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12903 debug_assert!(ncols.is_multiple_of(128));
12904 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12905 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12906 let nc = ncols as i32;
12907 if Self::pdl_on() {
12908 {
12909 use cudarc::driver::{DevicePtr, DevicePtrMut};
12910 let s = &self.gpu.stream();
12911 let (pg, _g0) = gate.device_ptr(s);
12912 let (pu, _g1) = up.device_ptr(s);
12913 let (pact, _g2) = act.device_ptr_mut(s);
12914 let (pq, _g3) = out_q.device_ptr_mut(s);
12915 let (pd, _g4) = out_d.device_ptr_mut(s);
12916 let mut ps = [
12917 &pg as *const _ as *mut std::ffi::c_void,
12918 &pu as *const _ as *mut _,
12919 &pact as *const _ as *mut _,
12920 &pq as *const _ as *mut _,
12921 &pd as *const _ as *mut _,
12922 &nc as *const _ as *mut _,
12923 ];
12924 unsafe {
12925 self.launch_pdl(
12926 "gelu_tanh_mul_q8_1",
12927 (nrows as u32, 1, 1),
12928 (rms_block(), 1, 1),
12929 &mut ps,
12930 )?;
12931 }
12932 }
12933 return Ok((out_q, out_d));
12934 }
12935 let f = self.func("gelu_tanh_mul_q8_1");
12936 let cfg = LaunchConfig {
12937 grid_dim: (nrows as u32, 1, 1),
12938 block_dim: (rms_block(), 1, 1),
12939 shared_mem_bytes: 0,
12940 };
12941 let __s_b = self.gpu.stream();
12942 let mut b = __s_b.launch_builder(&f);
12943 b.arg(gate)
12944 .arg(up)
12945 .arg(act)
12946 .arg(&mut out_q)
12947 .arg(&mut out_d)
12948 .arg(&nc);
12949 unsafe {
12950 b.launch(cfg)?;
12951 }
12952 Ok((out_q, out_d))
12953 }
12954
12955 #[allow(clippy::too_many_arguments)]
12957 pub fn gelu_tanh_mul_q8_1_into(
12958 &self,
12959 gate: &CudaSlice<f32>,
12960 up: &cudarc::driver::CudaView<f32>,
12961 act: &mut CudaSlice<f32>,
12962 ncols: usize,
12963 nrows: usize,
12964 out_q: &mut CudaSlice<i8>,
12965 out_d: &mut CudaSlice<f32>,
12966 ) -> Result<(), Box<dyn std::error::Error>> {
12967 debug_assert!(ncols.is_multiple_of(128));
12968 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
12969 let nc = ncols as i32;
12970 if Self::pdl_on() {
12971 use cudarc::driver::{DevicePtr, DevicePtrMut};
12972 let s = &self.gpu.stream();
12973 let (pg, _g0) = gate.device_ptr(s);
12974 let (pu, _g1) = up.device_ptr(s);
12975 let (pact, _g2) = act.device_ptr_mut(s);
12976 let (pq, _g3) = out_q.device_ptr_mut(s);
12977 let (pd, _g4) = out_d.device_ptr_mut(s);
12978 let mut ps = [
12979 &pg as *const _ as *mut std::ffi::c_void,
12980 &pu as *const _ as *mut _,
12981 &pact as *const _ as *mut _,
12982 &pq as *const _ as *mut _,
12983 &pd as *const _ as *mut _,
12984 &nc as *const _ as *mut _,
12985 ];
12986 unsafe {
12987 self.launch_pdl(
12988 "gelu_tanh_mul_q8_1",
12989 (nrows as u32, 1, 1),
12990 (rms_block(), 1, 1),
12991 &mut ps,
12992 )?;
12993 }
12994 return Ok(());
12995 }
12996 let f = self.func("gelu_tanh_mul_q8_1");
12997 let cfg = LaunchConfig {
12998 grid_dim: (nrows as u32, 1, 1),
12999 block_dim: (rms_block(), 1, 1),
13000 shared_mem_bytes: 0,
13001 };
13002 let __s_b = self.gpu.stream();
13003 let mut b = __s_b.launch_builder(&f);
13004 b.arg(gate)
13005 .arg(up)
13006 .arg(&mut *act)
13007 .arg(&mut *out_q)
13008 .arg(&mut *out_d)
13009 .arg(&nc);
13010 unsafe {
13011 b.launch(cfg)?;
13012 }
13013 Ok(())
13014 }
13015
13016 #[allow(clippy::too_many_arguments)]
13018 #[allow(clippy::type_complexity)] pub fn add_rms_norm3_q8z(
13020 &self,
13021 a: &CudaSlice<f32>,
13022 b_in: &CudaSlice<f32>,
13023 w0: &CudaSlice<f32>,
13024 w1: &CudaSlice<f32>,
13025 w2: &CudaSlice<f32>,
13026 res: &mut CudaSlice<f32>,
13027 out1: &mut CudaSlice<f32>,
13028 ncols: usize,
13029 nrows: usize,
13030 eps: f32,
13031 ) -> Result<
13032 (
13033 (CudaSlice<i8>, CudaSlice<f32>),
13034 (CudaSlice<i8>, CudaSlice<f32>),
13035 ),
13036 Box<dyn std::error::Error>,
13037 > {
13038 let mut q0 = self.alloc_uninit::<i8>(nrows * ncols)?;
13039 let mut d0 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
13040 let mut q2 = self.alloc_uninit::<i8>(nrows * ncols)?;
13041 let mut d2 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
13042 let f = self.func("add_rms_norm3_q8z_f32");
13043 let cfg = LaunchConfig {
13044 grid_dim: (nrows as u32, 1, 1),
13045 block_dim: (rms_block(), 1, 1),
13046 shared_mem_bytes: 0,
13047 };
13048 let (nc, e2) = (ncols as i32, eps);
13049 let __s_b = self.gpu.stream();
13050 let mut b = __s_b.launch_builder(&f);
13051 b.arg(a)
13052 .arg(b_in)
13053 .arg(w0)
13054 .arg(w1)
13055 .arg(w2)
13056 .arg(res)
13057 .arg(&mut q0)
13058 .arg(&mut d0)
13059 .arg(out1)
13060 .arg(&mut q2)
13061 .arg(&mut d2)
13062 .arg(&nc)
13063 .arg(&e2);
13064 unsafe {
13065 b.launch(cfg)?;
13066 }
13067 Ok(((q0, d0), (q2, d2)))
13068 }
13069
13070 #[allow(clippy::too_many_arguments)]
13072 pub fn add_rms_norm3(
13073 &self,
13074 a: &CudaSlice<f32>,
13075 b_in: &CudaSlice<f32>,
13076 w0: &CudaSlice<f32>,
13077 w1: &CudaSlice<f32>,
13078 w2: &CudaSlice<f32>,
13079 res: &mut CudaSlice<f32>,
13080 d0: &mut CudaSlice<f32>,
13081 d1: &mut CudaSlice<f32>,
13082 d2: &mut CudaSlice<f32>,
13083 ncols: usize,
13084 nrows: usize,
13085 eps: f32,
13086 ) -> Result<(), Box<dyn std::error::Error>> {
13087 let f = self.func("add_rms_norm3_f32");
13088 let cfg = LaunchConfig {
13089 grid_dim: (nrows as u32, 1, 1),
13090 block_dim: (rms_block(), 1, 1),
13091 shared_mem_bytes: 0,
13092 };
13093 let (nc, e2) = (ncols as i32, eps);
13094 let __s_b = self.gpu.stream();
13095 let mut b = __s_b.launch_builder(&f);
13096 b.arg(a)
13097 .arg(b_in)
13098 .arg(w0)
13099 .arg(w1)
13100 .arg(w2)
13101 .arg(res)
13102 .arg(d0)
13103 .arg(d1)
13104 .arg(d2)
13105 .arg(&nc)
13106 .arg(&e2);
13107 unsafe {
13108 b.launch(cfg)?;
13109 }
13110 Ok(())
13111 }
13112
13113 pub fn add_scale(
13115 &self,
13116 a: &CudaSlice<f32>,
13117 b_in: &CudaSlice<f32>,
13118 c: f32,
13119 dst: &mut CudaSlice<f32>,
13120 n: usize,
13121 ) -> Result<(), Box<dyn std::error::Error>> {
13122 let f = self.func("add_scale_f32");
13123 let cfg = LaunchConfig::for_num_elems(n as u32);
13124 let ni = n as i32;
13125 let __s_b = self.gpu.stream();
13126 let mut b = __s_b.launch_builder(&f);
13127 b.arg(a).arg(b_in).arg(&c).arg(dst).arg(&ni);
13128 unsafe {
13129 b.launch(cfg)?;
13130 }
13131 Ok(())
13132 }
13133
13134 #[allow(clippy::too_many_arguments)] pub fn layer_norm_bias(
13137 &self,
13138 x: &CudaSlice<f32>,
13139 w: &CudaSlice<f32>,
13140 b: &CudaSlice<f32>,
13141 dst: &mut CudaSlice<f32>,
13142 ncols: usize,
13143 nrows: usize,
13144 eps: f32,
13145 ) -> Result<(), Box<dyn std::error::Error>> {
13146 let f = self.func("layer_norm_bias_f32");
13147 let (nc, e) = (ncols as i32, eps);
13148 let cfg = LaunchConfig {
13149 grid_dim: (nrows as u32, 1, 1),
13150 block_dim: (256, 1, 1),
13151 shared_mem_bytes: 0,
13152 };
13153 let __s_b = self.gpu.stream();
13154 let mut lb = __s_b.launch_builder(&f);
13155 lb.arg(x).arg(w).arg(b).arg(&mut *dst).arg(&nc).arg(&e);
13156 unsafe {
13157 lb.launch(cfg)?;
13158 }
13159 Ok(())
13160 }
13161
13162 pub fn gelu_tanh(
13164 &self,
13165 x: &CudaSlice<f32>,
13166 dst: &mut CudaSlice<f32>,
13167 n: usize,
13168 ) -> Result<(), Box<dyn std::error::Error>> {
13169 let f = self.func("gelu_tanh_f32");
13170 let ni = n as i64;
13171 let cfg = LaunchConfig {
13172 grid_dim: (n.div_ceil(256) as u32, 1, 1),
13173 block_dim: (256, 1, 1),
13174 shared_mem_bytes: 0,
13175 };
13176 let __s_b = self.gpu.stream();
13177 let mut lb = __s_b.launch_builder(&f);
13178 lb.arg(x).arg(&mut *dst).arg(&ni);
13179 unsafe {
13180 lb.launch(cfg)?;
13181 }
13182 Ok(())
13183 }
13184
13185 pub fn row_softmax(
13187 &self,
13188 x: &mut CudaSlice<f32>,
13189 ncols: usize,
13190 nrows: usize,
13191 ) -> Result<(), Box<dyn std::error::Error>> {
13192 let f = self.func("row_softmax_f32");
13193 let nc = ncols as i32;
13194 let cfg = LaunchConfig {
13195 grid_dim: (nrows as u32, 1, 1),
13196 block_dim: (256, 1, 1),
13197 shared_mem_bytes: 0,
13198 };
13199 let __s_b = self.gpu.stream();
13200 let mut lb = __s_b.launch_builder(&f);
13201 lb.arg(&mut *x).arg(&nc);
13202 unsafe {
13203 lb.launch(cfg)?;
13204 }
13205 Ok(())
13206 }
13207
13208 pub fn rms_norm(
13209 &self,
13210 x: &CudaSlice<f32>,
13211 w: &CudaSlice<f32>,
13212 dst: &mut CudaSlice<f32>,
13213 ncols: usize,
13214 nrows: usize,
13215 eps: f32,
13216 ) -> Result<(), Box<dyn std::error::Error>> {
13217 let (nc, e) = (ncols as i32, eps);
13218 let kname = if Self::norm_ilp_on() {
13219 "rms_norm_f32_v2"
13220 } else {
13221 "rms_norm_f32"
13222 };
13223 if Self::pdl_on() && Self::pdl_wb_on() {
13224 use cudarc::driver::{DevicePtr, DevicePtrMut};
13225 let s = &self.gpu.stream();
13226 let (px, _g0) = x.device_ptr(s);
13227 let (pw, _g1) = w.device_ptr(s);
13228 let (pd, _g2) = dst.device_ptr_mut(s);
13229 let mut ps = [
13230 &px as *const _ as *mut std::ffi::c_void,
13231 &pw as *const _ as *mut _,
13232 &pd as *const _ as *mut _,
13233 &nc as *const _ as *mut _,
13234 &e as *const _ as *mut _,
13235 ];
13236 unsafe {
13237 self.launch_pdl(kname, (nrows as u32, 1, 1), (rms_block(), 1, 1), &mut ps)?;
13238 }
13239 return Ok(());
13240 }
13241 let f = self.func(kname);
13242 let cfg = LaunchConfig {
13243 grid_dim: (nrows as u32, 1, 1),
13244 block_dim: (rms_block(), 1, 1),
13245 shared_mem_bytes: 0,
13246 };
13247 let __s_b = self.gpu.stream();
13248 let mut b = __s_b.launch_builder(&f);
13249 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
13250 unsafe {
13251 b.launch(cfg)?;
13252 }
13253 Ok(())
13254 }
13255
13256 pub fn rms_norm_decode(
13264 &self,
13265 x: &CudaSlice<f32>,
13266 w: &CudaSlice<f32>,
13267 dst: &mut CudaSlice<f32>,
13268 ncols: usize,
13269 nrows: usize,
13270 eps: f32,
13271 ) -> Result<(), Box<dyn std::error::Error>> {
13272 let f = self.func(if Self::norm_ilp_on() {
13273 "rms_norm_f32_v2"
13274 } else {
13275 "rms_norm_f32"
13276 });
13277 let cfg = LaunchConfig {
13278 grid_dim: (nrows as u32, 1, 1),
13279 block_dim: (1024, 1, 1),
13280 shared_mem_bytes: 0,
13281 };
13282 let (nc, e) = (ncols as i32, eps);
13283 let __s_b = self.gpu.stream();
13284 let mut b = __s_b.launch_builder(&f);
13285 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
13286 unsafe {
13287 b.launch(cfg)?;
13288 }
13289 Ok(())
13290 }
13291
13292 pub fn rms_norm_q8_1(
13296 &self,
13297 x: &CudaSlice<f32>,
13298 w: &CudaSlice<f32>,
13299 ncols: usize,
13300 nrows: usize,
13301 eps: f32,
13302 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13303 let nblk = ncols / 32;
13304 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
13305 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
13306 let (nc, e) = (ncols as i32, eps);
13307 if Self::pdl_on() {
13308 {
13309 use cudarc::driver::{DevicePtr, DevicePtrMut};
13310 let s = &self.gpu.stream();
13311 let (px, _g0) = x.device_ptr(s);
13312 let (pw, _g1) = w.device_ptr(s);
13313 let (pq, _g2) = q.device_ptr_mut(s);
13314 let (pd, _g3) = d.device_ptr_mut(s);
13315 let mut ps = [
13316 &px as *const _ as *mut std::ffi::c_void,
13317 &pw as *const _ as *mut _,
13318 &pq as *const _ as *mut _,
13319 &pd as *const _ as *mut _,
13320 &nc as *const _ as *mut _,
13321 &e as *const _ as *mut _,
13322 ];
13323 unsafe {
13324 self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1), &mut ps)?;
13325 }
13326 }
13327 return Ok((q, d));
13328 }
13329 let f = self.func("rms_norm_q8_1");
13330 let cfg = LaunchConfig {
13333 grid_dim: (nrows as u32, 1, 1),
13334 block_dim: (1024, 1, 1),
13335 shared_mem_bytes: 0,
13336 };
13337 let __s_b = self.gpu.stream();
13338 let mut b = __s_b.launch_builder(&f);
13339 b.arg(x).arg(w).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
13340 unsafe {
13341 b.launch(cfg)?;
13342 }
13343 Ok((q, d))
13344 }
13345
13346 #[allow(clippy::too_many_arguments)] pub fn rms_norm_q8_1_into(
13350 &self,
13351 x: &CudaSlice<f32>,
13352 w: &CudaSlice<f32>,
13353 ncols: usize,
13354 nrows: usize,
13355 eps: f32,
13356 q: &mut CudaSlice<i8>,
13357 d: &mut CudaSlice<f32>,
13358 ) -> Result<(), Box<dyn std::error::Error>> {
13359 let nblk = ncols / 32;
13360 debug_assert!(q.len() >= nrows * ncols && d.len() >= nrows * nblk);
13361 let (nc, e) = (ncols as i32, eps);
13362 if Self::pdl_on() {
13363 use cudarc::driver::{DevicePtr, DevicePtrMut};
13364 let s = &self.gpu.stream();
13365 let (px, _g0) = x.device_ptr(s);
13366 let (pw, _g1) = w.device_ptr(s);
13367 let (pq, _g2) = q.device_ptr_mut(s);
13368 let (pd, _g3) = d.device_ptr_mut(s);
13369 let mut ps = [
13370 &px as *const _ as *mut std::ffi::c_void,
13371 &pw as *const _ as *mut _,
13372 &pq as *const _ as *mut _,
13373 &pd as *const _ as *mut _,
13374 &nc as *const _ as *mut _,
13375 &e as *const _ as *mut _,
13376 ];
13377 unsafe {
13378 self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1), &mut ps)?;
13379 }
13380 return Ok(());
13381 }
13382 let f = self.func("rms_norm_q8_1");
13383 let cfg = LaunchConfig {
13384 grid_dim: (nrows as u32, 1, 1),
13385 block_dim: (1024, 1, 1),
13386 shared_mem_bytes: 0,
13387 };
13388 let __s_b = self.gpu.stream();
13389 let mut b = __s_b.launch_builder(&f);
13390 b.arg(x).arg(w).arg(&mut *q).arg(&mut *d).arg(&nc).arg(&e);
13391 unsafe {
13392 b.launch(cfg)?;
13393 }
13394 Ok(())
13395 }
13396
13397 pub fn quantize_q8_1_into(
13399 &self,
13400 x: &CudaSlice<f32>,
13401 m: usize,
13402 in_f: usize,
13403 q: &mut CudaSlice<i8>,
13404 d: &mut CudaSlice<f32>,
13405 ) -> Result<(), Box<dyn std::error::Error>> {
13406 let nblk = in_f / 32;
13407 debug_assert!(q.len() >= m * in_f && d.len() >= m * nblk);
13408 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
13409 let (inf, mi) = (in_f as i32, m as i32);
13410 if Self::pdl_on() && Self::pdl_wb_on() {
13411 use cudarc::driver::{DevicePtr, DevicePtrMut};
13412 let s = &self.gpu.stream();
13413 let (px, _g0) = x.device_ptr(s);
13414 let (pq, _g1) = q.device_ptr_mut(s);
13415 let (pd, _g2) = d.device_ptr_mut(s);
13416 let mut ps = [
13417 &px as *const _ as *mut std::ffi::c_void,
13418 &pq as *const _ as *mut _,
13419 &pd as *const _ as *mut _,
13420 &inf as *const _ as *mut _,
13421 &mi as *const _ as *mut _,
13422 ];
13423 unsafe {
13424 self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?;
13425 }
13426 return Ok(());
13427 }
13428 let f = self.func("quantize_q8_1");
13429 let __s_b = self.gpu.stream();
13430 let mut b = __s_b.launch_builder(&f);
13431 b.arg(x).arg(&mut *q).arg(&mut *d).arg(&inf).arg(&mi);
13432 unsafe {
13433 b.launch(cfg)?;
13434 }
13435 Ok(())
13436 }
13437
13438 #[allow(clippy::too_many_arguments)] pub fn add_rms_norm_q8_1(
13443 &self,
13444 a: &CudaSlice<f32>,
13445 b_in: &CudaSlice<f32>,
13446 w: &CudaSlice<f32>,
13447 res: &mut CudaSlice<f32>,
13448 ncols: usize,
13449 nrows: usize,
13450 eps: f32,
13451 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13452 let nblk = ncols / 32;
13453 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
13454 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
13455 let f = self.func("add_rms_norm_q8_1");
13456 let cfg = LaunchConfig {
13458 grid_dim: (nrows as u32, 1, 1),
13459 block_dim: (1024, 1, 1),
13460 shared_mem_bytes: 0,
13461 };
13462 let (nc, e) = (ncols as i32, eps);
13463 let __s_bld = self.gpu.stream();
13464 let mut bld = __s_bld.launch_builder(&f);
13465 bld.arg(a)
13466 .arg(b_in)
13467 .arg(w)
13468 .arg(res)
13469 .arg(&mut q)
13470 .arg(&mut d)
13471 .arg(&nc)
13472 .arg(&e);
13473 unsafe {
13474 bld.launch(cfg)?;
13475 }
13476 Ok((q, d))
13477 }
13478
13479 #[allow(clippy::too_many_arguments)]
13485 pub fn join_add_rms_norm_raw(
13486 &self,
13487 a0_raw: u64,
13488 a1_raw: u64,
13489 x: &CudaSlice<f32>,
13490 w: &CudaSlice<f32>,
13491 res: &mut CudaSlice<f32>,
13492 dst: &mut CudaSlice<f32>,
13493 ncols: usize,
13494 eps: f32,
13495 ) -> Result<(), Box<dyn std::error::Error>> {
13496 if a0_raw == 0 || a1_raw == 0 || x.len() < ncols || res.len() < ncols || dst.len() < ncols {
13497 return Err("join_add_rms_norm geometry".into());
13498 }
13499 let f = self.func("join_add_rms_norm_f32");
13500 let cfg = LaunchConfig {
13501 grid_dim: (1, 1, 1),
13502 block_dim: (rms_block(), 1, 1),
13503 shared_mem_bytes: 0,
13504 };
13505 let (nc, e) = (ncols as i32, eps);
13506 let __s_b = self.gpu.stream();
13507 let mut b = __s_b.launch_builder(&f);
13508 b.arg(&a0_raw)
13509 .arg(&a1_raw)
13510 .arg(x)
13511 .arg(w)
13512 .arg(&mut *res)
13513 .arg(&mut *dst)
13514 .arg(&nc)
13515 .arg(&e);
13516 unsafe {
13517 b.launch(cfg)?;
13518 }
13519 Ok(())
13520 }
13521
13522 #[allow(clippy::too_many_arguments)] pub fn add_rms_norm(
13524 &self,
13525 a: &CudaSlice<f32>,
13526 b: &CudaSlice<f32>,
13527 w: &CudaSlice<f32>,
13528 res: &mut CudaSlice<f32>,
13529 dst: &mut CudaSlice<f32>,
13530 ncols: usize,
13531 nrows: usize,
13532 eps: f32,
13533 ) -> Result<(), Box<dyn std::error::Error>> {
13534 let (nc, e) = (ncols as i32, eps);
13535 let kname = if Self::norm_ilp_on() {
13536 "add_rms_norm_f32_v2"
13537 } else {
13538 "add_rms_norm_f32"
13539 };
13540 if Self::pdl_on() && Self::pdl_wb_on() {
13541 use cudarc::driver::{DevicePtr, DevicePtrMut};
13542 let s = &self.gpu.stream();
13543 let (pa, _g0) = a.device_ptr(s);
13544 let (pb, _g1) = b.device_ptr(s);
13545 let (pw, _g2) = w.device_ptr(s);
13546 let (pr, _g3) = res.device_ptr_mut(s);
13547 let (pd, _g4) = dst.device_ptr_mut(s);
13548 let mut ps = [
13549 &pa as *const _ as *mut std::ffi::c_void,
13550 &pb as *const _ as *mut _,
13551 &pw as *const _ as *mut _,
13552 &pr as *const _ as *mut _,
13553 &pd as *const _ as *mut _,
13554 &nc as *const _ as *mut _,
13555 &e as *const _ as *mut _,
13556 ];
13557 unsafe {
13558 self.launch_pdl(kname, (nrows as u32, 1, 1), (rms_block(), 1, 1), &mut ps)?;
13559 }
13560 return Ok(());
13561 }
13562 let f = self.func(kname);
13563 let cfg = LaunchConfig {
13564 grid_dim: (nrows as u32, 1, 1),
13565 block_dim: (rms_block(), 1, 1),
13566 shared_mem_bytes: 0,
13567 };
13568 let __s_b2 = self.gpu.stream();
13569 let mut b2 = __s_b2.launch_builder(&f);
13570 b2.arg(a)
13571 .arg(b)
13572 .arg(w)
13573 .arg(&mut *res)
13574 .arg(&mut *dst)
13575 .arg(&nc)
13576 .arg(&e);
13577 unsafe {
13578 b2.launch(cfg)?;
13579 }
13580 Ok(())
13581 }
13582
13583 #[allow(clippy::too_many_arguments)]
13586 pub fn rms_pre_add_rms_norm(
13587 &self,
13588 a: &CudaSlice<f32>,
13589 wa: &CudaSlice<f32>,
13590 b: &CudaSlice<f32>,
13591 w: &CudaSlice<f32>,
13592 res: &mut CudaSlice<f32>,
13593 dst: &mut CudaSlice<f32>,
13594 ncols: usize,
13595 nrows: usize,
13596 eps: f32,
13597 ) -> Result<(), Box<dyn std::error::Error>> {
13598 let f = self.func("rms_pre_add_rms_norm_f32");
13599 let cfg = LaunchConfig {
13600 grid_dim: (nrows as u32, 1, 1),
13601 block_dim: (rms_block(), 1, 1),
13602 shared_mem_bytes: 0,
13603 };
13604 let (nc, e) = (ncols as i32, eps);
13605 let __s_b2 = self.gpu.stream();
13606 let mut b2 = __s_b2.launch_builder(&f);
13607 b2.arg(a)
13608 .arg(wa)
13609 .arg(b)
13610 .arg(w)
13611 .arg(&mut *res)
13612 .arg(&mut *dst)
13613 .arg(&nc)
13614 .arg(&e);
13615 unsafe {
13616 b2.launch(cfg)?;
13617 }
13618 Ok(())
13619 }
13620
13621 #[allow(clippy::too_many_arguments)]
13623 pub fn rms_pre_add_rms_norm_q8z(
13624 &self,
13625 a: &CudaSlice<f32>,
13626 wa: &CudaSlice<f32>,
13627 b: &CudaSlice<f32>,
13628 w: &CudaSlice<f32>,
13629 res: &mut CudaSlice<f32>,
13630 dst: &mut CudaSlice<f32>,
13631 ncols: usize,
13632 nrows: usize,
13633 eps: f32,
13634 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13635 debug_assert!(ncols.is_multiple_of(128));
13636 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
13637 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
13638 let (nc, e) = (ncols as i32, eps);
13639 if Self::pdl_on() {
13640 {
13641 use cudarc::driver::{DevicePtr, DevicePtrMut};
13642 let s = &self.gpu.stream();
13643 let (pa, _g0) = a.device_ptr(s);
13644 let (pwa, _g1) = wa.device_ptr(s);
13645 let (pb, _g2) = b.device_ptr(s);
13646 let (pw, _g3) = w.device_ptr(s);
13647 let (pr, _g4) = res.device_ptr_mut(s);
13648 let (pdst, _g5) = dst.device_ptr_mut(s);
13649 let (pq, _g6) = out_q.device_ptr_mut(s);
13650 let (pd, _g7) = out_d.device_ptr_mut(s);
13651 let mut ps = [
13652 &pa as *const _ as *mut std::ffi::c_void,
13653 &pwa as *const _ as *mut _,
13654 &pb as *const _ as *mut _,
13655 &pw as *const _ as *mut _,
13656 &pr as *const _ as *mut _,
13657 &pdst as *const _ as *mut _,
13658 &pq as *const _ as *mut _,
13659 &pd as *const _ as *mut _,
13660 &nc as *const _ as *mut _,
13661 &e as *const _ as *mut _,
13662 ];
13663 unsafe {
13664 self.launch_pdl(
13665 "rms_pre_add_rms_norm_q8z_f32",
13666 (nrows as u32, 1, 1),
13667 (rms_block(), 1, 1),
13668 &mut ps,
13669 )?;
13670 }
13671 }
13672 return Ok((out_q, out_d));
13673 }
13674 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
13675 let cfg = LaunchConfig {
13676 grid_dim: (nrows as u32, 1, 1),
13677 block_dim: (rms_block(), 1, 1),
13678 shared_mem_bytes: 0,
13679 };
13680 let __s_b2 = self.gpu.stream();
13681 let mut b2 = __s_b2.launch_builder(&f);
13682 b2.arg(a)
13683 .arg(wa)
13684 .arg(b)
13685 .arg(w)
13686 .arg(&mut *res)
13687 .arg(&mut *dst)
13688 .arg(&mut out_q)
13689 .arg(&mut out_d)
13690 .arg(&nc)
13691 .arg(&e);
13692 unsafe {
13693 b2.launch(cfg)?;
13694 }
13695 Ok((out_q, out_d))
13696 }
13697
13698 #[allow(clippy::too_many_arguments)]
13702 pub fn rms_pre_add_rms_norm_q8z_into(
13703 &self,
13704 a: &CudaSlice<f32>,
13705 wa: &CudaSlice<f32>,
13706 b: &CudaSlice<f32>,
13707 w: &CudaSlice<f32>,
13708 res: &mut CudaSlice<f32>,
13709 dst: &mut CudaSlice<f32>,
13710 ncols: usize,
13711 nrows: usize,
13712 eps: f32,
13713 out_q: &mut CudaSlice<i8>,
13714 out_d: &mut CudaSlice<f32>,
13715 ) -> Result<(), Box<dyn std::error::Error>> {
13716 debug_assert!(ncols.is_multiple_of(128));
13717 let (nc, e) = (ncols as i32, eps);
13718 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
13719 let cfg = LaunchConfig {
13720 grid_dim: (nrows as u32, 1, 1),
13721 block_dim: (rms_block(), 1, 1),
13722 shared_mem_bytes: 0,
13723 };
13724 let __s_b = self.gpu.stream();
13725 let mut b2 = __s_b.launch_builder(&f);
13726 b2.arg(a)
13727 .arg(wa)
13728 .arg(b)
13729 .arg(w)
13730 .arg(&mut *res)
13731 .arg(&mut *dst)
13732 .arg(&mut *out_q)
13733 .arg(&mut *out_d)
13734 .arg(&nc)
13735 .arg(&e);
13736 unsafe {
13737 b2.launch(cfg)?;
13738 }
13739 Ok(())
13740 }
13741
13742 #[allow(clippy::too_many_arguments)]
13745 pub fn rms_pre_add_scale_rms_norm_q8_1_into(
13746 &self,
13747 a: &CudaSlice<f32>,
13748 wa: &CudaSlice<f32>,
13749 b_in: &CudaSlice<f32>,
13750 c: f32,
13751 w: &CudaSlice<f32>,
13752 res: &mut CudaSlice<f32>,
13753 ncols: usize,
13754 nrows: usize,
13755 eps: f32,
13756 out_q: &mut CudaSlice<i8>,
13757 out_d: &mut CudaSlice<f32>,
13758 ) -> Result<(), Box<dyn std::error::Error>> {
13759 debug_assert!(ncols.is_multiple_of(128));
13760 let (nc, e2) = (ncols as i32, eps);
13761 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
13762 let cfg = LaunchConfig {
13763 grid_dim: (nrows as u32, 1, 1),
13764 block_dim: (rms_block(), 1, 1),
13765 shared_mem_bytes: 0,
13766 };
13767 let __s_b = self.gpu.stream();
13768 let mut b2 = __s_b.launch_builder(&f);
13769 b2.arg(a)
13770 .arg(wa)
13771 .arg(b_in)
13772 .arg(&c)
13773 .arg(w)
13774 .arg(&mut *res)
13775 .arg(&mut *out_q)
13776 .arg(&mut *out_d)
13777 .arg(&nc)
13778 .arg(&e2);
13779 unsafe {
13780 b2.launch(cfg)?;
13781 }
13782 Ok(())
13783 }
13784
13785 pub fn g4_pnfold_on() -> bool {
13793 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13794 *ON.get_or_init(|| {
13795 std::env::var("MEMRA_G4_PNFOLD")
13796 .map(|v| v != "0")
13797 .unwrap_or(true)
13798 })
13799 }
13800
13801 pub fn build_q4_out_concat3(
13805 &self,
13806 w0: &crate::model::GpuTensor,
13807 w1: &crate::model::GpuTensor,
13808 w2: &crate::model::GpuTensor,
13809 ) -> Result<Option<crate::model::GpuTensor>, Box<dyn std::error::Error>> {
13810 use crate::model::GpuTensor;
13811 let part = |w: &GpuTensor| -> Option<(usize, usize)> {
13812 match w {
13813 GpuTensor::Quant {
13814 qtype,
13815 row_bytes,
13816 rp,
13817 ..
13818 } if *qtype == QT_Q4_0 && !*rp => Some((*row_bytes, w.out_features())),
13819 _ => None,
13820 }
13821 };
13822 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (part(w0), part(w1), part(w2))
13823 else {
13824 return Ok(None);
13825 };
13826 if rb0 != rb1
13827 || rb0 != rb2
13828 || w0.in_features() != w1.in_features()
13829 || w0.in_features() != w2.in_features()
13830 {
13831 return Ok(None);
13832 }
13833 fn bytes_of(w: &crate::model::GpuTensor) -> &CudaSlice<u8> {
13834 match w {
13835 crate::model::GpuTensor::Quant { bytes, .. } => bytes,
13836 _ => unreachable!(),
13837 }
13838 }
13839 let (b0, b1, b2) = (bytes_of(w0), bytes_of(w1), bytes_of(w2));
13840 let total = rb0 * (o0 + o1 + o2);
13841 let mut cat = self.alloc_u8(total)?;
13842 self.copy_u8_into(&mut cat, 0, b0, rb0 * o0)?;
13843 self.copy_u8_into(&mut cat, rb0 * o0, b1, rb1 * o1)?;
13844 self.copy_u8_into(&mut cat, rb0 * (o0 + o1), b2, rb2 * o2)?;
13845 Ok(Some(GpuTensor::Quant {
13846 bytes: cat,
13847 qtype: QT_Q4_0,
13848 row_bytes: rb0,
13849 ne: vec![w0.in_features() as u64, (o0 + o1 + o2) as u64],
13850 scale: 1.0,
13851 rp: false,
13852 #[cfg(memra_cutlass)]
13853 cutlass: None,
13854 fp8: None,
13855 blk: None,
13856 rp4: None,
13857 f16: None,
13858 }))
13859 }
13860
13861 fn full_width_rope_only(
13879 kernel: &str,
13880 n_rot: usize,
13881 head_dim: usize,
13882 ) -> Result<(), Box<dyn std::error::Error>> {
13883 if n_rot == head_dim {
13884 return Ok(());
13885 }
13886 Err(format!(
13887 "{kernel}: PARTIAL ROTARY REFUSED — n_rot {n_rot} != head_dim {head_dim}. This fused \
13888 rms_norm+qkv+rope kernel carries no n_dims parameter and rotates the full head \
13889 width (half = ncols/2), so it would rotate dims {n_rot}..{head_dim} that must pass \
13890 through unrotated. Use the split path (rms_norm_qkv + rope_neox/rope_neox2 with \
13891 n_dims={n_rot}), or add an n_dims early-return to the kernel and widen this guard."
13892 )
13893 .into())
13894 }
13895
13896 #[allow(clippy::too_many_arguments)]
13900 pub fn rms_norm_qkv_rope_cat(
13901 &self,
13902 qkv: &CudaSlice<f32>,
13903 wq: &CudaSlice<f32>,
13904 wk: &CudaSlice<f32>,
13905 wv: &CudaSlice<f32>,
13906 q: &mut CudaSlice<f32>,
13907 k: &mut CudaSlice<f32>,
13908 v: &mut CudaSlice<f32>,
13909 head_dim: usize,
13910 n_rot: usize,
13911 rq: usize,
13912 rk: usize,
13913 pos: &CudaSlice<i32>,
13914 nh_q: usize,
13915 nh_k: usize,
13916 base: f32,
13917 freq_scale: f32,
13918 ff: Option<&CudaSlice<f32>>,
13919 eps: f32,
13920 ) -> Result<(), Box<dyn std::error::Error>> {
13921 Self::full_width_rope_only("rms_norm_qkv_rope_cat", n_rot, head_dim)?;
13922 let rows = rq + rk + rk;
13923 let theta_scale = base.powf(-2.0 / head_dim as f32);
13924 let (nc, rqi, rki, nhq, nhk) = (
13925 head_dim as i32,
13926 rq as i32,
13927 rk as i32,
13928 nh_q as i32,
13929 nh_k as i32,
13930 );
13931 if Self::pdl_on() {
13932 use cudarc::driver::{DevicePtr, DevicePtrMut};
13933 let s = &self.gpu.stream();
13934 let (pqkv, _g0) = qkv.device_ptr(s);
13935 let (pwq, _g1) = wq.device_ptr(s);
13936 let (pwk, _g2) = wk.device_ptr(s);
13937 let (pwv, _g3) = wv.device_ptr(s);
13938 let (pq, _g4) = q.device_ptr_mut(s);
13939 let (pk, _g5) = k.device_ptr_mut(s);
13940 let (pv, _g6) = v.device_ptr_mut(s);
13941 let (ppos, _g7) = pos.device_ptr(s);
13942 let (pff, _g8) = match ff {
13943 Some(t) => {
13944 let (p, g) = t.device_ptr(s);
13945 (p, Some(g))
13946 }
13947 None => (0, None),
13948 };
13949 let mut ps = [
13950 &pqkv as *const _ as *mut std::ffi::c_void,
13951 &pwq as *const _ as *mut _,
13952 &pwk as *const _ as *mut _,
13953 &pwv as *const _ as *mut _,
13954 &pq as *const _ as *mut _,
13955 &pk as *const _ as *mut _,
13956 &pv as *const _ as *mut _,
13957 &nc as *const _ as *mut _,
13958 &rqi as *const _ as *mut _,
13959 &rki as *const _ as *mut _,
13960 &ppos as *const _ as *mut _,
13961 &nhq as *const _ as *mut _,
13962 &nhk as *const _ as *mut _,
13963 &theta_scale as *const _ as *mut _,
13964 &freq_scale as *const _ as *mut _,
13965 &pff as *const _ as *mut _,
13966 &eps as *const _ as *mut _,
13967 ];
13968 unsafe {
13969 self.launch_pdl(
13970 "rms_norm_qkv_rope_cat_f32",
13971 (rows as u32, 1, 1),
13972 (rms_block(), 1, 1),
13973 &mut ps,
13974 )?;
13975 }
13976 return Ok(());
13977 }
13978 let f = self.func("rms_norm_qkv_rope_cat_f32");
13979 let cfg = LaunchConfig {
13980 grid_dim: (rows as u32, 1, 1),
13981 block_dim: (rms_block(), 1, 1),
13982 shared_mem_bytes: 0,
13983 };
13984 let __s_b = self.gpu.stream();
13985 let mut b = __s_b.launch_builder(&f);
13986 match ff {
13987 Some(t) => {
13988 b.arg(qkv)
13989 .arg(wq)
13990 .arg(wk)
13991 .arg(wv)
13992 .arg(&mut *q)
13993 .arg(&mut *k)
13994 .arg(&mut *v)
13995 .arg(&nc)
13996 .arg(&rqi)
13997 .arg(&rki)
13998 .arg(pos)
13999 .arg(&nhq)
14000 .arg(&nhk)
14001 .arg(&theta_scale)
14002 .arg(&freq_scale)
14003 .arg(t)
14004 .arg(&eps);
14005 unsafe {
14006 b.launch(cfg)?;
14007 }
14008 }
14009 None => {
14010 let null: u64 = 0;
14011 b.arg(qkv)
14012 .arg(wq)
14013 .arg(wk)
14014 .arg(wv)
14015 .arg(&mut *q)
14016 .arg(&mut *k)
14017 .arg(&mut *v)
14018 .arg(&nc)
14019 .arg(&rqi)
14020 .arg(&rki)
14021 .arg(pos)
14022 .arg(&nhq)
14023 .arg(&nhk)
14024 .arg(&theta_scale)
14025 .arg(&freq_scale)
14026 .arg(&null)
14027 .arg(&eps);
14028 unsafe {
14029 b.launch(cfg)?;
14030 }
14031 }
14032 }
14033 Ok(())
14034 }
14035
14036 #[allow(clippy::too_many_arguments)]
14040 pub fn rms_norm_qkv_rope(
14041 &self,
14042 q0: &CudaSlice<f32>,
14043 k0: &CudaSlice<f32>,
14044 v0: &CudaSlice<f32>,
14045 wq: &CudaSlice<f32>,
14046 wk: &CudaSlice<f32>,
14047 wv: &CudaSlice<f32>,
14048 q: &mut CudaSlice<f32>,
14049 k: &mut CudaSlice<f32>,
14050 v: &mut CudaSlice<f32>,
14051 head_dim: usize,
14052 n_rot: usize,
14053 rq: usize,
14054 rk: usize,
14055 pos: &CudaSlice<i32>,
14056 nh_q: usize,
14057 nh_k: usize,
14058 base: f32,
14059 freq_scale: f32,
14060 ff: Option<&CudaSlice<f32>>,
14061 eps: f32,
14062 ) -> Result<(), Box<dyn std::error::Error>> {
14063 Self::full_width_rope_only("rms_norm_qkv_rope", n_rot, head_dim)?;
14064 let f = self.func("rms_norm_qkv_rope_f32");
14065 let rows = rq + rk + rk; let cfg = LaunchConfig {
14067 grid_dim: (rows as u32, 1, 1),
14068 block_dim: (rms_block(), 1, 1),
14069 shared_mem_bytes: 0,
14070 };
14071 let theta_scale = base.powf(-2.0 / head_dim as f32);
14072 let (nc, rqi, rki, nhq, nhk) = (
14073 head_dim as i32,
14074 rq as i32,
14075 rk as i32,
14076 nh_q as i32,
14077 nh_k as i32,
14078 );
14079 let __s_b = self.gpu.stream();
14080 let mut b = __s_b.launch_builder(&f);
14081 match ff {
14082 Some(t) => {
14083 b.arg(q0)
14084 .arg(k0)
14085 .arg(v0)
14086 .arg(wq)
14087 .arg(wk)
14088 .arg(wv)
14089 .arg(&mut *q)
14090 .arg(&mut *k)
14091 .arg(&mut *v)
14092 .arg(&nc)
14093 .arg(&rqi)
14094 .arg(&rki)
14095 .arg(pos)
14096 .arg(&nhq)
14097 .arg(&nhk)
14098 .arg(&theta_scale)
14099 .arg(&freq_scale)
14100 .arg(t)
14101 .arg(&eps);
14102 unsafe {
14103 b.launch(cfg)?;
14104 }
14105 }
14106 None => {
14107 let null: u64 = 0;
14108 b.arg(q0)
14109 .arg(k0)
14110 .arg(v0)
14111 .arg(wq)
14112 .arg(wk)
14113 .arg(wv)
14114 .arg(&mut *q)
14115 .arg(&mut *k)
14116 .arg(&mut *v)
14117 .arg(&nc)
14118 .arg(&rqi)
14119 .arg(&rki)
14120 .arg(pos)
14121 .arg(&nhq)
14122 .arg(&nhk)
14123 .arg(&theta_scale)
14124 .arg(&freq_scale)
14125 .arg(&null)
14126 .arg(&eps);
14127 unsafe {
14128 b.launch(cfg)?;
14129 }
14130 }
14131 }
14132 Ok(())
14133 }
14134
14135 #[allow(clippy::too_many_arguments)]
14141 pub fn rms_norm_qkv_rope_append_dc(
14142 &self,
14143 q0: &CudaSlice<f32>,
14144 k0: &CudaSlice<f32>,
14145 v0: &CudaSlice<f32>,
14146 wq: &CudaSlice<f32>,
14147 wk: &CudaSlice<f32>,
14148 wv: &CudaSlice<f32>,
14149 q: &mut CudaSlice<f32>,
14150 k: &mut CudaSlice<f32>,
14151 v: &mut CudaSlice<f32>,
14152 head_dim: usize,
14153 n_rot: usize,
14154 rq: usize,
14155 rk: usize,
14156 pos: &CudaSlice<i32>,
14157 nh_q: usize,
14158 nh_k: usize,
14159 base: f32,
14160 freq_scale: f32,
14161 ff: Option<&CudaSlice<f32>>,
14162 eps: f32,
14163 kc: &mut CudaSlice<u8>,
14164 vc: &mut CudaSlice<u8>,
14165 t_dev: &CudaSlice<i32>,
14166 k_tok_bytes: usize,
14167 v_tok_bytes: usize,
14168 g: bool,
14169 ) -> Result<(), Box<dyn std::error::Error>> {
14170 Self::full_width_rope_only("rms_norm_qkv_rope_append_dc", n_rot, head_dim)?;
14171 let rows = rq + rk + rk;
14172 let theta_scale = base.powf(-2.0 / head_dim as f32);
14173 let (nc, rqi, rki, nhq, nhk) = (
14174 head_dim as i32,
14175 rq as i32,
14176 rk as i32,
14177 nh_q as i32,
14178 nh_k as i32,
14179 );
14180 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
14181 if Self::pdl_on() && Self::pdl_wb_on() {
14182 use cudarc::driver::{DevicePtr, DevicePtrMut};
14183 let s = &self.gpu.stream();
14184 let (p0, _a0) = q0.device_ptr(s);
14185 let (p1, _a1) = k0.device_ptr(s);
14186 let (p2, _a2) = v0.device_ptr(s);
14187 let (pwq, _a3) = wq.device_ptr(s);
14188 let (pwk, _a4) = wk.device_ptr(s);
14189 let (pwv, _a5) = wv.device_ptr(s);
14190 let (pq, _a6) = q.device_ptr_mut(s);
14191 let (pk, _a7) = k.device_ptr_mut(s);
14192 let (pv, _a8) = v.device_ptr_mut(s);
14193 let (pp, _a9) = pos.device_ptr(s);
14194 let pff: u64 = match ff {
14195 Some(t) => {
14196 let (p, _gg) = t.device_ptr(s);
14197 p
14198 }
14199 None => 0,
14200 };
14201 let (pkc, _a10) = kc.device_ptr_mut(s);
14202 let (pvc, _a11) = vc.device_ptr_mut(s);
14203 let (pt, _a12) = t_dev.device_ptr(s);
14204 let mut ps = [
14205 &p0 as *const _ as *mut std::ffi::c_void,
14206 &p1 as *const _ as *mut _,
14207 &p2 as *const _ as *mut _,
14208 &pwq as *const _ as *mut _,
14209 &pwk as *const _ as *mut _,
14210 &pwv as *const _ as *mut _,
14211 &pq as *const _ as *mut _,
14212 &pk as *const _ as *mut _,
14213 &pv as *const _ as *mut _,
14214 &nc as *const _ as *mut _,
14215 &rqi as *const _ as *mut _,
14216 &rki as *const _ as *mut _,
14217 &pp as *const _ as *mut _,
14218 &nhq as *const _ as *mut _,
14219 &nhk as *const _ as *mut _,
14220 &theta_scale as *const _ as *mut _,
14221 &freq_scale as *const _ as *mut _,
14222 &pff as *const _ as *mut _,
14223 &eps as *const _ as *mut _,
14224 &pkc as *const _ as *mut _,
14225 &pvc as *const _ as *mut _,
14226 &pt as *const _ as *mut _,
14227 &ktb as *const _ as *mut _,
14228 &vtb as *const _ as *mut _,
14229 ];
14230 unsafe {
14231 self.launch_pdl_flash(
14232 g,
14233 "rms_norm_qkv_rope_append_dc_f32",
14234 (rows as u32, 1, 1),
14235 (rms_block(), 1, 1),
14236 0,
14237 &mut ps,
14238 )?;
14239 }
14240 return Ok(());
14241 }
14242 let f = if g {
14243 self.func_g("rms_norm_qkv_rope_append_dc_f32")
14244 } else {
14245 self.func("rms_norm_qkv_rope_append_dc_f32")
14246 };
14247 let cfg = LaunchConfig {
14248 grid_dim: (rows as u32, 1, 1),
14249 block_dim: (rms_block(), 1, 1),
14250 shared_mem_bytes: 0,
14251 };
14252 let __s_b = self.gpu.stream();
14253 let mut b = __s_b.launch_builder(&f);
14254 match ff {
14255 Some(t) => {
14256 b.arg(q0)
14257 .arg(k0)
14258 .arg(v0)
14259 .arg(wq)
14260 .arg(wk)
14261 .arg(wv)
14262 .arg(&mut *q)
14263 .arg(&mut *k)
14264 .arg(&mut *v)
14265 .arg(&nc)
14266 .arg(&rqi)
14267 .arg(&rki)
14268 .arg(pos)
14269 .arg(&nhq)
14270 .arg(&nhk)
14271 .arg(&theta_scale)
14272 .arg(&freq_scale)
14273 .arg(t)
14274 .arg(&eps)
14275 .arg(&mut *kc)
14276 .arg(&mut *vc)
14277 .arg(t_dev)
14278 .arg(&ktb)
14279 .arg(&vtb);
14280 unsafe {
14281 b.launch(cfg)?;
14282 }
14283 }
14284 None => {
14285 let null: u64 = 0;
14286 b.arg(q0)
14287 .arg(k0)
14288 .arg(v0)
14289 .arg(wq)
14290 .arg(wk)
14291 .arg(wv)
14292 .arg(&mut *q)
14293 .arg(&mut *k)
14294 .arg(&mut *v)
14295 .arg(&nc)
14296 .arg(&rqi)
14297 .arg(&rki)
14298 .arg(pos)
14299 .arg(&nhq)
14300 .arg(&nhk)
14301 .arg(&theta_scale)
14302 .arg(&freq_scale)
14303 .arg(&null)
14304 .arg(&eps)
14305 .arg(&mut *kc)
14306 .arg(&mut *vc)
14307 .arg(t_dev)
14308 .arg(&ktb)
14309 .arg(&vtb);
14310 unsafe {
14311 b.launch(cfg)?;
14312 }
14313 }
14314 }
14315 Ok(())
14316 }
14317
14318 #[allow(clippy::too_many_arguments)]
14326 pub fn rms_norm_qkv_rope_append(
14327 &self,
14328 q0: &CudaSlice<f32>,
14329 k0: &CudaSlice<f32>,
14330 v0: &CudaSlice<f32>,
14331 wq: &CudaSlice<f32>,
14332 wk: &CudaSlice<f32>,
14333 wv: &CudaSlice<f32>,
14334 q: &mut CudaSlice<f32>,
14335 k: &mut CudaSlice<f32>,
14336 v: &mut CudaSlice<f32>,
14337 head_dim: usize,
14338 n_rot: usize,
14339 rq: usize,
14340 rk: usize,
14341 pos: &CudaSlice<i32>,
14342 nh_q: usize,
14343 nh_k: usize,
14344 base: f32,
14345 freq_scale: f32,
14346 ff: Option<&CudaSlice<f32>>,
14347 eps: f32,
14348 kc: &mut CudaSlice<u8>,
14349 vc: &mut CudaSlice<u8>,
14350 t: usize,
14351 k_tok_bytes: usize,
14352 v_tok_bytes: usize,
14353 g: bool,
14354 ) -> Result<(), Box<dyn std::error::Error>> {
14355 Self::full_width_rope_only("rms_norm_qkv_rope_append", n_rot, head_dim)?;
14356 let rows = rq + rk + rk;
14357 let theta_scale = base.powf(-2.0 / head_dim as f32);
14358 let (nc, rqi, rki, nhq, nhk) = (
14359 head_dim as i32,
14360 rq as i32,
14361 rk as i32,
14362 nh_q as i32,
14363 nh_k as i32,
14364 );
14365 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
14366 let ti = t as i32;
14367 if Self::pdl_on() && Self::pdl_wb_on() {
14368 use cudarc::driver::{DevicePtr, DevicePtrMut};
14369 let s = &self.gpu.stream();
14370 let (p0, _a0) = q0.device_ptr(s);
14371 let (p1, _a1) = k0.device_ptr(s);
14372 let (p2, _a2) = v0.device_ptr(s);
14373 let (pwq, _a3) = wq.device_ptr(s);
14374 let (pwk, _a4) = wk.device_ptr(s);
14375 let (pwv, _a5) = wv.device_ptr(s);
14376 let (pq, _a6) = q.device_ptr_mut(s);
14377 let (pk, _a7) = k.device_ptr_mut(s);
14378 let (pv, _a8) = v.device_ptr_mut(s);
14379 let (pp, _a9) = pos.device_ptr(s);
14380 let pff: u64 = match ff {
14381 Some(t) => {
14382 let (p, _gg) = t.device_ptr(s);
14383 p
14384 }
14385 None => 0,
14386 };
14387 let (pkc, _a10) = kc.device_ptr_mut(s);
14388 let (pvc, _a11) = vc.device_ptr_mut(s);
14389 let mut ps = [
14390 &p0 as *const _ as *mut std::ffi::c_void,
14391 &p1 as *const _ as *mut _,
14392 &p2 as *const _ as *mut _,
14393 &pwq as *const _ as *mut _,
14394 &pwk as *const _ as *mut _,
14395 &pwv as *const _ as *mut _,
14396 &pq as *const _ as *mut _,
14397 &pk as *const _ as *mut _,
14398 &pv as *const _ as *mut _,
14399 &nc as *const _ as *mut _,
14400 &rqi as *const _ as *mut _,
14401 &rki as *const _ as *mut _,
14402 &pp as *const _ as *mut _,
14403 &nhq as *const _ as *mut _,
14404 &nhk as *const _ as *mut _,
14405 &theta_scale as *const _ as *mut _,
14406 &freq_scale as *const _ as *mut _,
14407 &pff as *const _ as *mut _,
14408 &eps as *const _ as *mut _,
14409 &pkc as *const _ as *mut _,
14410 &pvc as *const _ as *mut _,
14411 &ti as *const _ as *mut _,
14412 &ktb as *const _ as *mut _,
14413 &vtb as *const _ as *mut _,
14414 ];
14415 unsafe {
14416 self.launch_pdl_flash(
14417 g,
14418 "rms_norm_qkv_rope_append_f32",
14419 (rows as u32, 1, 1),
14420 (rms_block(), 1, 1),
14421 0,
14422 &mut ps,
14423 )?;
14424 }
14425 return Ok(());
14426 }
14427 let f = if g {
14428 self.func_g("rms_norm_qkv_rope_append_f32")
14429 } else {
14430 self.func("rms_norm_qkv_rope_append_f32")
14431 };
14432 let cfg = LaunchConfig {
14433 grid_dim: (rows as u32, 1, 1),
14434 block_dim: (rms_block(), 1, 1),
14435 shared_mem_bytes: 0,
14436 };
14437 let __s_b = self.gpu.stream();
14438 let mut b = __s_b.launch_builder(&f);
14439 let null: u64 = 0;
14440 b.arg(q0)
14441 .arg(k0)
14442 .arg(v0)
14443 .arg(wq)
14444 .arg(wk)
14445 .arg(wv)
14446 .arg(&mut *q)
14447 .arg(&mut *k)
14448 .arg(&mut *v)
14449 .arg(&nc)
14450 .arg(&rqi)
14451 .arg(&rki)
14452 .arg(pos)
14453 .arg(&nhq)
14454 .arg(&nhk)
14455 .arg(&theta_scale)
14456 .arg(&freq_scale);
14457 match ff {
14458 Some(t) => {
14459 b.arg(t);
14460 }
14461 None => {
14462 b.arg(&null);
14463 }
14464 }
14465 b.arg(&eps)
14466 .arg(&mut *kc)
14467 .arg(&mut *vc)
14468 .arg(&ti)
14469 .arg(&ktb)
14470 .arg(&vtb);
14471 unsafe {
14472 b.launch(cfg)?;
14473 }
14474 Ok(())
14475 }
14476
14477 pub fn add_q8_1(
14478 &self,
14479 a: &CudaSlice<f32>,
14480 b: &CudaSlice<f32>,
14481 res: &mut CudaSlice<f32>,
14482 ncols: usize,
14483 nrows: usize,
14484 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14485 debug_assert!(ncols.is_multiple_of(128));
14486 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
14487 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
14488 let f = self.func("add_q8_1_f32");
14489 let cfg = LaunchConfig {
14490 grid_dim: (nrows as u32, 1, 1),
14491 block_dim: (rms_block(), 1, 1),
14492 shared_mem_bytes: 0,
14493 };
14494 let nc = ncols as i32;
14495 let __s_b2 = self.gpu.stream();
14496 let mut b2 = __s_b2.launch_builder(&f);
14497 b2.arg(a)
14498 .arg(b)
14499 .arg(&mut *res)
14500 .arg(&mut out_q)
14501 .arg(&mut out_d)
14502 .arg(&nc);
14503 unsafe {
14504 b2.launch(cfg)?;
14505 }
14506 Ok((out_q, out_d))
14507 }
14508
14509 #[allow(clippy::too_many_arguments)] pub fn rms_pre_add_q8_1(
14514 &self,
14515 a: &CudaSlice<f32>,
14516 wa: &CudaSlice<f32>,
14517 b: &CudaSlice<f32>,
14518 res: &mut CudaSlice<f32>,
14519 ncols: usize,
14520 nrows: usize,
14521 eps: f32,
14522 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14523 debug_assert!(ncols.is_multiple_of(128));
14524 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
14525 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
14526 let f = self.func("rms_pre_add_q8_1_f32");
14527 let cfg = LaunchConfig {
14528 grid_dim: (nrows as u32, 1, 1),
14529 block_dim: (rms_block(), 1, 1),
14530 shared_mem_bytes: 0,
14531 };
14532 let (nc, ep) = (ncols as i32, eps);
14533 let __s_b2 = self.gpu.stream();
14534 let mut b2 = __s_b2.launch_builder(&f);
14535 b2.arg(a)
14536 .arg(wa)
14537 .arg(b)
14538 .arg(&mut *res)
14539 .arg(&mut out_q)
14540 .arg(&mut out_d)
14541 .arg(&nc)
14542 .arg(&ep);
14543 unsafe {
14544 b2.launch(cfg)?;
14545 }
14546 Ok((out_q, out_d))
14547 }
14548
14549 pub fn l2_v2_on(ncols: usize) -> bool {
14553 ncols == 128 && std::env::var("MEMRA_L2_V2").as_deref() != Ok("0")
14554 }
14555
14556 pub fn l2_norm_pp(
14557 &self,
14558 x: &CudaSlice<f32>,
14559 dst: &mut CudaSlice<f32>,
14560 dst16: Option<&mut CudaSlice<u8>>,
14561 ncols: usize,
14562 nrows: usize,
14563 eps: f32,
14564 ) -> Result<(), Box<dyn std::error::Error>> {
14565 if Self::l2_v2_on(ncols) {
14566 let f = self.func("l2_norm_pp_v2_f32");
14567 let rows_per_block = 8u32; let cfg = LaunchConfig {
14569 grid_dim: ((nrows as u32).div_ceil(rows_per_block), 1, 1),
14570 block_dim: (256, 1, 1),
14571 shared_mem_bytes: 0,
14572 };
14573 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
14574 let d16: u64 = match dst16 {
14576 Some(d) => self.addr_u8(d),
14577 None => 0,
14578 };
14579 let __s_b = self.gpu.stream();
14580 let mut b = __s_b.launch_builder(&f);
14581 b.arg(x).arg(dst).arg(&d16).arg(&nc).arg(&nr).arg(&e);
14582 unsafe {
14583 b.launch(cfg)?;
14584 }
14585 return Ok(());
14586 }
14587 self.l2_norm(x, dst, ncols, nrows, eps)
14588 }
14589
14590 pub fn l2_norm(
14591 &self,
14592 x: &CudaSlice<f32>,
14593 dst: &mut CudaSlice<f32>,
14594 ncols: usize,
14595 nrows: usize,
14596 eps: f32,
14597 ) -> Result<(), Box<dyn std::error::Error>> {
14598 let f = self.func("l2_norm_f32");
14599 let cfg = LaunchConfig {
14600 grid_dim: (nrows as u32, 1, 1),
14601 block_dim: (256, 1, 1),
14602 shared_mem_bytes: 0,
14603 };
14604 let (nc, e) = (ncols as i32, eps);
14605 let __s_b = self.gpu.stream();
14606 let mut b = __s_b.launch_builder(&f);
14607 b.arg(x).arg(dst).arg(&nc).arg(&e);
14608 unsafe {
14609 b.launch(cfg)?;
14610 }
14611 Ok(())
14612 }
14613
14614 pub fn l2_norm_decode(
14620 &self,
14621 x: &CudaSlice<f32>,
14622 dst: &mut CudaSlice<f32>,
14623 ncols: usize,
14624 nrows: usize,
14625 eps: f32,
14626 ) -> Result<(), Box<dyn std::error::Error>> {
14627 let f = self.func("l2_norm_f32");
14628 let cfg = LaunchConfig {
14629 grid_dim: (nrows as u32, 1, 1),
14630 block_dim: (32, 1, 1),
14631 shared_mem_bytes: 0,
14632 };
14633 let (nc, e) = (ncols as i32, eps);
14634 let __s_b = self.gpu.stream();
14635 let mut b = __s_b.launch_builder(&f);
14636 b.arg(x).arg(dst).arg(&nc).arg(&e);
14637 unsafe {
14638 b.launch(cfg)?;
14639 }
14640 Ok(())
14641 }
14642
14643 #[allow(clippy::too_many_arguments)] pub fn rope_neox(
14646 &self,
14647 x: &mut CudaSlice<f32>,
14648 pos: &CudaSlice<i32>,
14649 head_dim: usize,
14650 n_dims: usize,
14651 n_heads: usize,
14652 n_tokens: usize,
14653 freq_base: f32,
14654 freq_scale: f32,
14655 ) -> Result<(), Box<dyn std::error::Error>> {
14656 let f = self.func("rope_neox_f32");
14657 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
14658 let grid = (n_heads * n_tokens) as u32;
14659 let cfg = LaunchConfig {
14660 grid_dim: (grid, 1, 1),
14661 block_dim: ((head_dim / 2) as u32, 1, 1),
14662 shared_mem_bytes: 0,
14663 };
14664 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
14665 let __s_b = self.gpu.stream();
14666 let mut b = __s_b.launch_builder(&f);
14667 b.arg(x)
14668 .arg(pos)
14669 .arg(&hd)
14670 .arg(&nd)
14671 .arg(&nh)
14672 .arg(&theta_scale)
14673 .arg(&freq_scale);
14674 unsafe {
14675 b.launch(cfg)?;
14676 }
14677 Ok(())
14678 }
14679
14680 #[allow(clippy::too_many_arguments)] pub fn rope_neox_ff(
14683 &self,
14684 x: &mut CudaSlice<f32>,
14685 pos: &CudaSlice<i32>,
14686 head_dim: usize,
14687 n_dims: usize,
14688 n_heads: usize,
14689 n_tokens: usize,
14690 freq_base: f32,
14691 freq_scale: f32,
14692 ff: &CudaSlice<f32>,
14693 ) -> Result<(), Box<dyn std::error::Error>> {
14694 let f = self.func("rope_neox_ff_f32");
14695 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
14696 let grid = (n_heads * n_tokens) as u32;
14697 let cfg = LaunchConfig {
14698 grid_dim: (grid, 1, 1),
14699 block_dim: ((head_dim / 2) as u32, 1, 1),
14700 shared_mem_bytes: 0,
14701 };
14702 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
14703 let __s_b = self.gpu.stream();
14704 let mut b = __s_b.launch_builder(&f);
14705 b.arg(x)
14706 .arg(pos)
14707 .arg(&hd)
14708 .arg(&nd)
14709 .arg(&nh)
14710 .arg(&theta_scale)
14711 .arg(&freq_scale)
14712 .arg(ff);
14713 unsafe {
14714 b.launch(cfg)?;
14715 }
14716 Ok(())
14717 }
14718
14719 #[allow(clippy::too_many_arguments)]
14723 pub fn rope_neox_ffm(
14724 &self,
14725 x: &mut CudaSlice<f32>,
14726 pos: &CudaSlice<i32>,
14727 head_dim: usize,
14728 n_dims: usize,
14729 n_heads: usize,
14730 n_tokens: usize,
14731 freq_base: f32,
14732 freq_scale: f32,
14733 ff: &CudaSlice<f32>,
14734 mscale: f32,
14735 ) -> Result<(), Box<dyn std::error::Error>> {
14736 let f = self.func("rope_neox_ffm_f32");
14737 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
14738 let grid = (n_heads * n_tokens) as u32;
14739 let cfg = LaunchConfig {
14740 grid_dim: (grid, 1, 1),
14741 block_dim: ((head_dim / 2) as u32, 1, 1),
14742 shared_mem_bytes: 0,
14743 };
14744 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
14745 let __s_b = self.gpu.stream();
14746 let mut b = __s_b.launch_builder(&f);
14747 b.arg(x)
14748 .arg(pos)
14749 .arg(&hd)
14750 .arg(&nd)
14751 .arg(&nh)
14752 .arg(&theta_scale)
14753 .arg(&freq_scale)
14754 .arg(ff)
14755 .arg(&mscale);
14756 unsafe {
14757 b.launch(cfg)?;
14758 }
14759 Ok(())
14760 }
14761
14762 #[allow(clippy::too_many_arguments)]
14764 pub fn rope_neox2(
14765 &self,
14766 q: &mut CudaSlice<f32>,
14767 k: &mut CudaSlice<f32>,
14768 pos: &CudaSlice<i32>,
14769 head_dim: usize,
14770 n_dims: usize,
14771 nh_q: usize,
14772 nh_k: usize,
14773 n_tokens: usize,
14774 freq_base: f32,
14775 freq_scale: f32,
14776 ff: Option<&CudaSlice<f32>>,
14777 ) -> Result<(), Box<dyn std::error::Error>> {
14778 let f = self.func("rope_neox2_f32");
14779 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
14780 let grid = ((nh_q + nh_k) * n_tokens) as u32;
14781 let cfg = LaunchConfig {
14782 grid_dim: (grid, 1, 1),
14783 block_dim: ((head_dim / 2) as u32, 1, 1),
14784 shared_mem_bytes: 0,
14785 };
14786 let (hd, nd, nq, nk, nt) = (
14787 head_dim as i32,
14788 n_dims as i32,
14789 nh_q as i32,
14790 nh_k as i32,
14791 n_tokens as i32,
14792 );
14793 let __s_b = self.gpu.stream();
14794 let mut b = __s_b.launch_builder(&f);
14795 b.arg(q)
14796 .arg(k)
14797 .arg(pos)
14798 .arg(&hd)
14799 .arg(&nd)
14800 .arg(&nq)
14801 .arg(&nk)
14802 .arg(&nt)
14803 .arg(&theta_scale)
14804 .arg(&freq_scale);
14805 match ff {
14806 Some(ffv) => {
14807 b.arg(ffv);
14808 unsafe {
14809 b.launch(cfg)?;
14810 }
14811 }
14812 None => {
14813 let null: u64 = 0;
14814 b.arg(&null);
14815 unsafe {
14816 b.launch(cfg)?;
14817 }
14818 }
14819 }
14820 Ok(())
14821 }
14822
14823 pub fn gelu_tanh_mul(
14825 &self,
14826 gate: &CudaSlice<f32>,
14827 up: &CudaSlice<f32>,
14828 dst: &mut CudaSlice<f32>,
14829 n: usize,
14830 ) -> Result<(), Box<dyn std::error::Error>> {
14831 let f = self.func("gelu_tanh_mul_f32");
14832 let cfg = LaunchConfig::for_num_elems(n as u32);
14833 let ni = n as i32;
14834 let __s_b = self.gpu.stream();
14835 let mut b = __s_b.launch_builder(&f);
14836 b.arg(gate).arg(up).arg(dst).arg(&ni);
14837 unsafe {
14838 b.launch(cfg)?;
14839 }
14840 Ok(())
14841 }
14842
14843 pub fn silu_mul(
14844 &self,
14845 gate: &CudaSlice<f32>,
14846 up: &CudaSlice<f32>,
14847 dst: &mut CudaSlice<f32>,
14848 n: usize,
14849 ) -> Result<(), Box<dyn std::error::Error>> {
14850 let f = self.func("silu_mul_f32");
14851 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
14853 let ni = n as i32;
14854 let __s_b = self.gpu.stream();
14855 let mut b = __s_b.launch_builder(&f);
14856 b.arg(gate).arg(up).arg(dst).arg(&ni);
14857 unsafe {
14858 b.launch(cfg)?;
14859 }
14860 Ok(())
14861 }
14862
14863 pub fn silu_mul_host_expf(
14865 &self,
14866 gate: &CudaSlice<f32>,
14867 up: &CudaSlice<f32>,
14868 dst: &mut CudaSlice<f32>,
14869 n: usize,
14870 ) -> Result<(), Box<dyn std::error::Error>> {
14871 let f = self.func("silu_mul_host_expf_f32");
14872 let cfg = LaunchConfig::for_num_elems(n as u32);
14873 let ni = n as i32;
14874 let __s_b = self.gpu.stream();
14875 let mut b = __s_b.launch_builder(&f);
14876 b.arg(gate).arg(up).arg(dst).arg(&ni);
14877 unsafe {
14878 b.launch(cfg)?;
14879 }
14880 Ok(())
14881 }
14882
14883 pub fn silu_clamped_mul_host_expf(
14885 &self,
14886 gate: &CudaSlice<f32>,
14887 up: &CudaSlice<f32>,
14888 limit: f32,
14889 dst: &mut CudaSlice<f32>,
14890 n: usize,
14891 ) -> Result<(), Box<dyn std::error::Error>> {
14892 if !limit.is_finite() || limit <= 0.0 {
14893 return Err(
14894 format!("Step routed-expert clamp limit must be positive, got {limit}").into(),
14895 );
14896 }
14897 let f = self.func("silu_clamped_mul_host_expf_f32");
14898 let cfg = LaunchConfig::for_num_elems(n as u32);
14899 let ni = n as i32;
14900 let __s_b = self.gpu.stream();
14901 let mut b = __s_b.launch_builder(&f);
14902 b.arg(gate).arg(up).arg(&limit).arg(dst).arg(&ni);
14903 unsafe {
14904 b.launch(cfg)?;
14905 }
14906 Ok(())
14907 }
14908
14909 pub fn silu_mul_f16out(
14912 &self,
14913 gate: &CudaSlice<f32>,
14914 up: &CudaSlice<f32>,
14915 dst: &mut CudaSlice<f32>,
14916 dst16: &mut CudaSlice<u8>,
14917 n: usize,
14918 ) -> Result<(), Box<dyn std::error::Error>> {
14919 let f = self.func("silu_mul_f16out_f32");
14920 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
14921 let ni = n as i32;
14922 let __s_b = self.gpu.stream();
14923 let mut b = __s_b.launch_builder(&f);
14924 b.arg(gate).arg(up).arg(dst).arg(dst16).arg(&ni);
14925 unsafe {
14926 b.launch(cfg)?;
14927 }
14928 Ok(())
14929 }
14930
14931 pub fn silu_mul_scaled(
14938 &self,
14939 gate: &CudaSlice<f32>,
14940 up: &CudaSlice<f32>,
14941 gs: f32,
14942 us: f32,
14943 dst: &mut CudaSlice<f32>,
14944 n: usize,
14945 ) -> Result<(), Box<dyn std::error::Error>> {
14946 let f = self.func("silu_mul_scaled_f32");
14947 let cfg = LaunchConfig::for_num_elems(n as u32);
14948 let ni = n as i32;
14949 let (gsf, usf) = (gs, us);
14950 let __s_b = self.gpu.stream();
14951 let mut b = __s_b.launch_builder(&f);
14952 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(dst).arg(&ni);
14953 unsafe {
14954 b.launch(cfg)?;
14955 }
14956 Ok(())
14957 }
14958
14959 #[allow(clippy::too_many_arguments)]
14963 pub fn swigluoai_mul_scaled(
14964 &self,
14965 gate: &CudaSlice<f32>,
14966 up: &CudaSlice<f32>,
14967 gs: f32,
14968 us: f32,
14969 alpha: f32,
14970 limit: f32,
14971 dst: &mut CudaSlice<f32>,
14972 n: usize,
14973 ) -> Result<(), Box<dyn std::error::Error>> {
14974 let f = self.func("swigluoai_mul_scaled_f32");
14975 let cfg = LaunchConfig::for_num_elems(n as u32);
14976 let ni = n as i32;
14977 let __s_b = self.gpu.stream();
14978 let mut b = __s_b.launch_builder(&f);
14979 b.arg(gate)
14980 .arg(up)
14981 .arg(&gs)
14982 .arg(&us)
14983 .arg(&alpha)
14984 .arg(&limit)
14985 .arg(dst)
14986 .arg(&ni);
14987 unsafe {
14988 b.launch(cfg)?;
14989 }
14990 Ok(())
14991 }
14992
14993 pub fn silu_mul_scaled_q8_1(
15001 &self,
15002 gate: &CudaSlice<f32>,
15003 up: &CudaSlice<f32>,
15004 gs: f32,
15005 us: f32,
15006 n: usize,
15007 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15008 let f = self.func("silu_mul_scaled_q8_1");
15009 let nblk = n / 32;
15010 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);
15014 let (gsf, usf, ni) = (gs, us, n as i32);
15015 let __s_b = self.gpu.stream();
15016 let mut b = __s_b.launch_builder(&f);
15017 b.arg(gate)
15018 .arg(up)
15019 .arg(&gsf)
15020 .arg(&usf)
15021 .arg(&mut aq)
15022 .arg(&mut ad)
15023 .arg(&ni);
15024 unsafe {
15025 b.launch(cfg)?;
15026 }
15027 Ok((aq, ad))
15028 }
15029
15030 pub fn add(
15031 &self,
15032 a: &CudaSlice<f32>,
15033 b_in: &CudaSlice<f32>,
15034 dst: &mut CudaSlice<f32>,
15035 n: usize,
15036 ) -> Result<(), Box<dyn std::error::Error>> {
15037 let f = self.func("add_f32");
15038 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
15040 let ni = n as i32;
15041 let __s_bld = self.gpu.stream();
15042 let mut bld = __s_bld.launch_builder(&f);
15043 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
15044 unsafe {
15045 bld.launch(cfg)?;
15046 }
15047 Ok(())
15048 }
15049
15050 pub fn mul(
15051 &self,
15052 a: &CudaSlice<f32>,
15053 b_in: &CudaSlice<f32>,
15054 dst: &mut CudaSlice<f32>,
15055 n: usize,
15056 ) -> Result<(), Box<dyn std::error::Error>> {
15057 let f = self.func("mul_f32");
15058 let cfg = LaunchConfig::for_num_elems(n as u32);
15059 let ni = n as i32;
15060 let __s_bld = self.gpu.stream();
15061 let mut bld = __s_bld.launch_builder(&f);
15062 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
15063 unsafe {
15064 bld.launch(cfg)?;
15065 }
15066 Ok(())
15067 }
15068
15069 pub fn matmul(
15072 &self,
15073 w: &crate::model::GpuTensor,
15074 x: &CudaSlice<f32>,
15075 m: usize,
15076 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15077 use crate::model::GpuTensor;
15078 let in_f = w.in_features();
15079 let out_f = w.out_features();
15080 #[allow(non_snake_case)]
15088 let GEMM_M_THRESHOLD = if self.verify_exact_on() {
15091 usize::MAX
15092 } else {
15093 16usize
15094 };
15095
15096 const GEMM_MIN_OUT_F: usize = 128; if m >= GEMM_M_THRESHOLD {
15121 if let Some(y) = self.try_fp8_gemm(w, x, m)? {
15122 return Ok(y);
15123 }
15124 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? {
15131 return Ok(y);
15132 }
15133 if let Some(y) = self.try_f16_gemm(w, x, m)? {
15136 return Ok(y);
15137 }
15138 }
15139 if let GpuTensor::Quant { qtype, .. } = w
15154 && *qtype == QT_F8_E4M3_BLK
15155 {
15156 if m >= GEMM_M_THRESHOLD
15157 && let Some(y) = self.try_e4m3_blk_prefill(w, x, m)?
15158 {
15159 return Ok(y);
15160 }
15161 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15162 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? {
15163 return Ok(y);
15164 }
15165 }
15166 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.mmq_supports(w) {
15167 return self.qmatvec_mmq(w, x, m);
15168 }
15169 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.gemm_supports(w) {
15170 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15171 return self.qmatvec_gemm(w, &aq, &ad, m);
15172 }
15173 if m >= GEMM_M_THRESHOLD
15176 && let Some(y) = self.try_fp4_gemm(w, x, m, in_f, out_f)?
15177 {
15178 return Ok(y);
15179 }
15180 let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
15184 if m == 1
15189 && fast
15190 && let GpuTensor::Quant {
15191 bytes,
15192 qtype,
15193 row_bytes,
15194 rp,
15195 rp4,
15196 scale,
15197 ..
15198 } = w
15199 && self.mmvq_supports(*qtype)
15200 {
15201 let (bytes, rp) = match rp4 {
15205 Some(m4) => (m4, true),
15206 None => (bytes, *rp),
15207 };
15208 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15209 return self.qmatvec_mmvq(
15210 bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, rp,
15211 );
15212 }
15213 if (2..=16).contains(&m)
15229 && fast
15230 && std::env::var("MEMRA_NO_BATCHED").is_err()
15231 && (m <= 4 || Self::b8_enabled())
15232 {
15233 let m_ok = m <= 8
15243 || matches!(w, GpuTensor::Quant { qtype, .. }
15244 if *qtype == QT_Q4_0 || *qtype == QT_Q6_K || *qtype == QT_F8_E4M3
15245 || *qtype == QT_NVFP4 || *qtype == QT_Q4_K || *qtype == QT_Q5_K || *qtype == QT_Q8_0);
15246 if m_ok
15247 && let GpuTensor::Quant {
15248 bytes,
15249 qtype,
15250 row_bytes,
15251 rp,
15252 rp4,
15253 ..
15254 } = w
15255 && self.batched_supports(*qtype)
15256 && self.mmvq_supports(*qtype)
15257 {
15258 let (bytes, rp) = match rp4 {
15259 Some(m4) => (m4, true),
15260 None => (bytes, *rp),
15261 };
15262 let mcols = Self::batched_mcols(m);
15263 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15264 let mut y = self.qmatvec_mmvq_batched(
15265 bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, mcols, 1.0, rp,
15266 )?;
15267 if let GpuTensor::Quant { scale, .. } = w
15268 && *scale != 1.0
15269 {
15270 self.scale_inplace(&mut y, *scale, m * out_f)?;
15271 }
15272 return Ok(y);
15273 }
15274 }
15275 if fast
15281 && let GpuTensor::Quant {
15282 bytes,
15283 qtype,
15284 row_bytes,
15285 scale,
15286 ..
15287 } = w
15288 && *qtype == QT_F8_E4M3
15289 {
15290 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15291 return self.qmatvec_mmvq(
15292 bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, false,
15293 );
15294 }
15295 let mut y = match w {
15296 GpuTensor::Quant {
15297 bytes,
15298 qtype,
15299 row_bytes,
15300 ..
15301 } if fast && *qtype == QT_Q8_0 => {
15302 self.qmatvec_q8_0_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15303 }
15304 GpuTensor::Quant {
15305 bytes,
15306 qtype,
15307 row_bytes,
15308 ..
15309 } if fast && *qtype == QT_Q4_K => {
15310 self.qmatvec_q4_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15311 }
15312 GpuTensor::Quant {
15313 bytes,
15314 qtype,
15315 row_bytes,
15316 ..
15317 } if fast && *qtype == QT_Q6_K => {
15318 self.qmatvec_q6_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15319 }
15320 GpuTensor::Quant {
15321 bytes,
15322 qtype,
15323 row_bytes,
15324 ..
15325 } if fast && *qtype == QT_Q5_K => {
15326 self.qmatvec_q5_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15327 }
15328 GpuTensor::Quant {
15329 bytes,
15330 qtype,
15331 row_bytes,
15332 ..
15333 } if fast && *qtype == QT_Q3_K => {
15334 self.qmatvec_q3_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15335 }
15336 GpuTensor::Quant {
15337 bytes,
15338 qtype,
15339 row_bytes,
15340 rp,
15341 ..
15342 } if fast && *qtype == QT_NVFP4 => self.qmatvec_dp4a_named(
15343 if *rp {
15344 "qmatvec_nvfp4_dp4a_rp"
15345 } else {
15346 "qmatvec_nvfp4_dp4a"
15347 },
15348 &bytes.slice(0..bytes.len()),
15349 x,
15350 m,
15351 in_f,
15352 out_f,
15353 *row_bytes,
15354 )?,
15355 GpuTensor::Quant {
15359 bytes,
15360 qtype,
15361 row_bytes,
15362 ..
15363 } if fast && *qtype == QT_IQ4_XS && Self::iq_fast_enabled() => {
15364 self.qmatvec_iq4_XS_fast(bytes, x, m, in_f, out_f, *row_bytes)?
15365 }
15366 GpuTensor::Quant {
15371 bytes,
15372 qtype,
15373 row_bytes,
15374 rp,
15375 ..
15376 } =>
15377 {
15380 self.qmatvec(
15381 bytes,
15382 x,
15383 m,
15384 in_f,
15385 out_f,
15386 if *rp && *qtype == QT_NVFP4 {
15387 QT_NVFP4_RP
15388 } else {
15389 *qtype
15390 },
15391 *row_bytes,
15392 )?
15393 }
15394 GpuTensor::Float { data, .. } => self.linear(x, data, m, in_f, out_f)?,
15395 GpuTensor::FloatBf16 { data, .. } => {
15398 if (1..=32).contains(&m) && Self::bf16_mmv_on() && in_f.is_multiple_of(8) {
15405 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
15406 self.matvec_bf16_rows_into(data, x, &mut y, in_f, out_f, m)?;
15407 y
15408 } else {
15409 self.linear_bf16_chunked(x, data, m, in_f, out_f, false, None)?
15410 }
15411 }
15412 };
15413 if let GpuTensor::Quant { scale, .. } = w
15415 && *scale != 1.0
15416 {
15417 self.scale_inplace(&mut y, *scale, m * out_f)?;
15418 }
15419 Ok(y)
15420 }
15421
15422 pub fn stage_a_raw_needed() -> bool {
15432 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15433 *ON.get_or_init(|| std::env::var("MEMRA_FAST").as_deref() == Ok("0"))
15434 }
15435
15436 pub fn uses_q8_1_fast(&self, w: &crate::model::GpuTensor) -> bool {
15439 use crate::model::GpuTensor;
15440 if std::env::var("MEMRA_FAST").as_deref() == Ok("0") {
15441 return false;
15442 }
15443 match w {
15444 GpuTensor::Quant { qtype, .. } => {
15451 matches!(
15452 *qtype,
15453 QT_Q8_0
15454 | QT_Q4_K
15455 | QT_Q6_K
15456 | QT_Q5_K
15457 | QT_Q3_K
15458 | QT_NVFP4
15459 | QT_F8_E4M3
15460 | QT_F8_E4M3_BLK
15461 | QT_Q4_0
15462 ) || (*qtype == QT_IQ4_XS && Self::iq_fast_enabled())
15463 }
15464 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
15465 }
15466 }
15467
15468 pub fn matmul_pre(
15473 &self,
15474 w: &crate::model::GpuTensor,
15475 aq: &CudaSlice<i8>,
15476 ad: &CudaSlice<f32>,
15477 x_fallback: &CudaSlice<f32>,
15478 m: usize,
15479 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15480 use crate::model::GpuTensor;
15481 let x_raw_ok = x_fallback.len() >= m * w.in_features();
15487 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
15490 if let Some(y) = self.try_fp8_gemm(w, x_fallback, m)? {
15491 return Ok(y);
15492 }
15493 if let Some(y) = self.try_fp8_blk_mmq(w, x_fallback, m)? {
15496 return Ok(y);
15497 }
15498 if let Some(y) = self.try_f16_gemm(w, x_fallback, m)? {
15500 return Ok(y);
15501 }
15502 }
15503 if m >= 16
15509 && x_raw_ok
15510 && !self.verify_exact_on()
15511 && let Some(y) = self.try_e4m3_blk_prefill(w, x_fallback, m)?
15512 {
15513 return Ok(y);
15514 }
15515 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
15516 return Ok(y);
15517 }
15518 if m >= 16
15523 && w.out_features() >= 128
15524 && self.mmq_supports(w)
15525 && !self.verify_exact_on()
15526 && x_raw_ok
15527 {
15528 return self.qmatvec_mmq(w, x_fallback, m);
15529 }
15530 if m >= 16
15533 && x_raw_ok
15534 && !self.verify_exact_on()
15535 && let Some(y) =
15536 self.try_fp4_gemm(w, x_fallback, m, w.in_features(), w.out_features())?
15537 {
15538 return Ok(y);
15539 }
15540 if m >= 16 && self.gemm_supports(w) && !self.verify_exact_on() {
15543 return self.qmatvec_gemm(w, aq, ad, m);
15544 }
15545 if !self.uses_q8_1_fast(w) {
15564 if !x_raw_ok {
15565 return Err(format!(
15566 "matmul_pre: q8_1-fast is off for this weight but x_fallback holds {} f32 \
15567 (need m*in_f = {}*{} = {}). This call site pre-quantized its activation and \
15568 dropped the f32, so there is nothing to fall back to — pass the real f32 \
15569 activation (see Engine::rms_norm_decode, which is bit-identical to \
15570 rms_norm_q8_1's reduction) or keep the weight on the q8_1 path.",
15571 x_fallback.len(),
15572 m,
15573 w.in_features(),
15574 m * w.in_features()
15575 )
15576 .into());
15577 }
15578 return self.matmul(w, x_fallback, m);
15579 }
15580 let in_f = w.in_features();
15581 let out_f = w.out_features();
15582 let (bytes, qtype, row_bytes, scale, rp) = match w {
15583 GpuTensor::Quant {
15584 bytes,
15585 qtype,
15586 row_bytes,
15587 scale,
15588 rp,
15589 ..
15590 } => (bytes, *qtype, *row_bytes, *scale, *rp),
15591 _ => unreachable!("uses_q8_1_fast guaranteed Quant"),
15592 };
15593 let (mbytes, mrp) = match w {
15596 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
15597 _ => (bytes, rp),
15598 };
15599 if m == 1 && self.mmvq_supports(qtype) {
15603 return self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, mrp);
15604 }
15605 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
15618 && std::env::var("MEMRA_NO_BATCHED").is_err()
15619 && (m <= 4 || Self::b8_enabled())
15620 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_NVFP4
15624 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_F8_E4M3 || qtype == QT_Q8_0)
15625 {
15626 let mcols = Self::batched_mcols(m);
15627 return self.qmatvec_mmvq_batched(
15628 mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, mrp,
15629 );
15630 }
15631 if qtype == QT_F8_E4M3 || qtype == QT_Q4_0 {
15637 let (b2, r2) = if qtype == QT_Q4_0 {
15638 (mbytes, mrp)
15639 } else {
15640 (bytes, rp)
15641 };
15642 return self.qmatvec_mmvq(b2, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, r2);
15643 }
15644 let name = match qtype {
15645 QT_Q8_0 => "qmatvec_q8_0_dp4a",
15646 QT_Q4_K => "qmatvec_q4_K_dp4a",
15647 QT_Q6_K => "qmatvec_q6_K_dp4a",
15648 QT_Q5_K => "qmatvec_q5_K_dp4a",
15649 QT_Q3_K => "qmatvec_q3_K_dp4a",
15650 QT_NVFP4 => {
15651 if rp {
15652 "qmatvec_nvfp4_dp4a_rp"
15653 } else {
15654 "qmatvec_nvfp4_dp4a"
15655 }
15656 }
15657 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
15658 _ => unreachable!(),
15659 };
15660 let f = self.func(name);
15661 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
15663 grid_dim: (out_f as u32, m as u32, 1),
15664 block_dim: (128, 1, 1),
15665 shared_mem_bytes: 0,
15666 };
15667 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
15668 let __s_b = self.gpu.stream();
15669 let mut b = __s_b.launch_builder(&f);
15670 b.arg(bytes)
15671 .arg(aq)
15672 .arg(ad)
15673 .arg(&mut y)
15674 .arg(&inf)
15675 .arg(&outf)
15676 .arg(&mi)
15677 .arg(&rb);
15678 unsafe {
15679 b.launch(cfg)?;
15680 }
15681 if scale != 1.0 {
15682 self.scale_inplace(&mut y, scale, m * out_f)?;
15683 }
15684 Ok(y)
15685 }
15686
15687 pub fn matmul_decode_exact(
15695 &self,
15696 w: &crate::model::GpuTensor,
15697 x: &CudaSlice<f32>,
15698 m: usize,
15699 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15700 use crate::model::GpuTensor;
15701 if let GpuTensor::Float { data, .. } = w {
15709 return self.linear_decode_exact(x, data, m, w.in_features(), w.out_features());
15710 }
15711 if let GpuTensor::FloatBf16 { data, .. } = w {
15714 let (in_f, out_f) = (w.in_features(), w.out_features());
15715 if (1..=32).contains(&m) && Self::bf16_mmv_on() && in_f % 8 == 0 {
15718 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
15719 self.matvec_bf16_rows_into(data, x, &mut y, in_f, out_f, m)?;
15720 return Ok(y);
15721 }
15722 return self.linear_bf16_chunked(x, data, m, in_f, out_f, true, None);
15723 }
15724 if !self.uses_q8_1_fast(w) {
15725 return self.matmul(w, x, m);
15726 }
15727 let in_f = w.in_features();
15728 let out_f = w.out_features();
15729 let (bytes, qtype, row_bytes, scale, rp) = match w {
15730 GpuTensor::Quant {
15731 bytes,
15732 qtype,
15733 row_bytes,
15734 scale,
15735 rp,
15736 ..
15737 } => (bytes, *qtype, *row_bytes, *scale, *rp),
15738 _ => return self.matmul(w, x, m),
15739 };
15740 let (bytes, rp) = match w {
15743 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
15744 _ => (bytes, rp),
15745 };
15746 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15747 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? {
15751 return Ok(y);
15752 }
15753 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
15762 && std::env::var("MEMRA_NO_BATCHED").is_err()
15763 && (m <= 4 || Self::b8_enabled())
15764 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
15767 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0)
15768 {
15769 let mcols = Self::batched_mcols(m);
15770 return self.qmatvec_mmvq_batched(
15771 bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp,
15772 );
15773 }
15774 if self.mmvq_supports(qtype) {
15775 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
15778 }
15779 self.matmul_pre(w, &aq, &ad, x, m)
15782 }
15783
15784 pub fn matmul_decode_exact_pre(
15794 &self,
15795 w: &crate::model::GpuTensor,
15796 aq: &CudaSlice<i8>,
15797 ad: &CudaSlice<f32>,
15798 m: usize,
15799 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15800 use crate::model::GpuTensor;
15801 debug_assert!(
15802 self.uses_q8_1_fast(w),
15803 "matmul_decode_exact_pre: caller must guarantee q8_1-fast"
15804 );
15805 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
15807 return Ok(y);
15808 }
15809 let in_f = w.in_features();
15810 let out_f = w.out_features();
15811 let (bytes, qtype, row_bytes, scale, rp) = match w {
15812 GpuTensor::Quant {
15813 bytes,
15814 qtype,
15815 row_bytes,
15816 scale,
15817 rp,
15818 ..
15819 } => (bytes, *qtype, *row_bytes, *scale, *rp),
15820 _ => {
15821 return Err(
15822 "matmul_decode_exact_pre: Quant tensor required (q8_1-fast contract)".into(),
15823 );
15824 }
15825 };
15826 let (bytes, rp) = match w {
15828 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
15829 _ => (bytes, rp),
15830 };
15831 if (2..=16).contains(&m)
15833 && self.batched_supports(qtype)
15834 && self.mmvq_supports(qtype)
15835 && std::env::var("MEMRA_NO_BATCHED").is_err()
15836 && (m <= 4 || Self::b8_enabled())
15837 && (m <= 8
15838 || qtype == QT_Q4_0
15839 || qtype == QT_Q6_K
15840 || qtype == QT_F8_E4M3
15841 || qtype == QT_NVFP4
15842 || qtype == QT_Q4_K
15843 || qtype == QT_Q5_K
15844 || qtype == QT_Q8_0)
15845 {
15846 let mcols = Self::batched_mcols(m);
15847 return self.qmatvec_mmvq_batched(
15848 bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp,
15849 );
15850 }
15851 if self.mmvq_supports(qtype) {
15852 return self.qmatvec_mmvq(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
15853 }
15854 let x0 = self.zeros(0)?;
15857 self.matmul_pre(w, aq, ad, &x0, m)
15858 }
15859
15860 #[allow(clippy::type_complexity)] pub fn matmul_decode_exact_dual_pre(
15870 &self,
15871 w0: &crate::model::GpuTensor,
15872 w1: &crate::model::GpuTensor,
15873 aq: &CudaSlice<i8>,
15874 ad: &CudaSlice<f32>,
15875 m: usize,
15876 ) -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>>
15877 {
15878 use crate::model::GpuTensor;
15879 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15880 let on = *ON.get_or_init(|| {
15881 std::env::var("MEMRA_SPEC_DUAL_T")
15882 .map(|v| v != "0")
15883 .unwrap_or(true)
15884 });
15885 if !on
15886 || !(2..=7).contains(&m)
15887 || std::env::var("MEMRA_NO_BATCHED").is_ok()
15888 || !self.uses_q8_1_fast(w0)
15889 || !self.uses_q8_1_fast(w1)
15890 {
15891 return Ok(None);
15892 }
15893 if !self.mmvq_supports(QT_NVFP4) {
15898 return Ok(None);
15899 }
15900 let (in_f, out_f) = (w0.in_features(), w0.out_features());
15901 if w1.in_features() != in_f || w1.out_features() != out_f {
15902 return Ok(None);
15903 }
15904 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
15905 (
15906 GpuTensor::Quant {
15907 bytes: b0,
15908 qtype: q0,
15909 row_bytes: rb0,
15910 scale: s0,
15911 rp: rp0,
15912 rp4: None,
15913 ..
15914 },
15915 GpuTensor::Quant {
15916 bytes: b1,
15917 qtype: q1,
15918 row_bytes: rb1,
15919 scale: s1,
15920 rp: rp1,
15921 rp4: None,
15922 ..
15923 },
15924 ) if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 => {
15925 (b0, b1, *rb0, *s0, *s1, *rp0)
15926 }
15927 _ => return Ok(None),
15928 };
15929 if m > 4 && !(rp && Self::b8_enabled() && std::env::var("MEMRA_B567").as_deref() != Ok("0"))
15932 {
15933 return Ok(None);
15934 }
15935 let (y0, y1) =
15936 self.qmatvec_batched_dual_raw(b0, b1, aq, ad, m, in_f, out_f, row_bytes, rp)?;
15937 Ok(Some(((y0, s0), (y1, s1))))
15938 }
15939
15940 pub fn matmul_decode_exact_group4_pre(
15952 &self,
15953 ws: [&crate::model::GpuTensor; 4],
15954 aq: &CudaSlice<i8>,
15955 ad: &CudaSlice<f32>,
15956 m: usize,
15957 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
15958 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15959 let on = *ON.get_or_init(|| {
15960 std::env::var("MEMRA_TK_GDN_GROUP")
15961 .map(|v| v != "0")
15962 .unwrap_or(true)
15963 });
15964 self.matmul_decode_exact_group_pre(&ws, aq, ad, m, on, "GDN group4")
15965 }
15966
15967 pub fn matmul_decode_exact_group3_pre(
15972 &self,
15973 ws: [&crate::model::GpuTensor; 3],
15974 aq: &CudaSlice<i8>,
15975 ad: &CudaSlice<f32>,
15976 m: usize,
15977 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
15978 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15979 let on = *ON.get_or_init(|| {
15980 std::env::var("MEMRA_TK_FA_GROUP")
15981 .map(|v| v != "0")
15982 .unwrap_or(true)
15983 });
15984 self.matmul_decode_exact_group_pre(&ws, aq, ad, m, on, "FA group3")
15985 }
15986
15987 #[allow(clippy::manual_div_ceil)] fn matmul_decode_exact_group_pre(
15992 &self,
15993 ws: &[&crate::model::GpuTensor],
15994 aq: &CudaSlice<i8>,
15995 ad: &CudaSlice<f32>,
15996 m: usize,
15997 on: bool,
15998 tag: &'static str,
15999 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
16000 use crate::model::GpuTensor;
16001 if !on
16002 || !(2..=16).contains(&m)
16003 || std::env::var("MEMRA_NO_BATCHED").is_ok()
16004 || (m > 4 && !Self::b8_enabled())
16005 || !self.mmvq_supports(QT_NVFP4)
16006 || !self.batched_supports(QT_NVFP4)
16007 {
16008 return Ok(None);
16009 }
16010 let in_f = ws[0].in_features();
16011 let mut parts: Vec<(&CudaSlice<u8>, usize, f32)> = Vec::with_capacity(4);
16012 for w in ws {
16013 if !self.uses_q8_1_fast(w) || w.in_features() != in_f {
16014 return Ok(None);
16015 }
16016 match w {
16017 GpuTensor::Quant {
16018 bytes,
16019 qtype,
16020 scale,
16021 rp: true,
16022 rp4: None,
16023 ..
16024 } if *qtype == QT_NVFP4 && w.out_features() % 8 == 0 => {
16025 parts.push((bytes, w.out_features(), *scale));
16026 }
16027 _ => return Ok(None),
16028 }
16029 }
16030 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16032 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
16033 let mcols = if (5..=7).contains(&m) && b567 {
16034 m
16035 } else {
16036 Self::batched_mcols(m)
16037 };
16038 let kname: &'static str = match mcols {
16039 2 => "qmatvec_nvfp4_mmvq_group4_b2_rp",
16040 4 => "qmatvec_nvfp4_mmvq_group4_b4_rp",
16041 5 => "qmatvec_nvfp4_mmvq_group4_b5_rp",
16042 6 => "qmatvec_nvfp4_mmvq_group4_b6_rp",
16043 7 => "qmatvec_nvfp4_mmvq_group4_b7_rp",
16044 8 => "qmatvec_nvfp4_mmvq_group4_b8_rp",
16045 16 => "qmatvec_nvfp4_mmvq_group4_b16_rp",
16046 _ => return Ok(None),
16047 };
16048 if std::env::var("MEMRA_DEBUG").is_ok() {
16051 use std::sync::Mutex;
16052 static SEEN: Mutex<Vec<&'static str>> = Mutex::new(Vec::new());
16053 let mut seen = SEEN.lock().unwrap();
16054 if !seen.contains(&tag) {
16055 seen.push(tag);
16056 eprintln!("[memra] {tag} batched ENGAGED (m={m})");
16057 }
16058 }
16059 const ROWS_PER_BLOCK: u32 = 4; let rows_per_block = ROWS_PER_BLOCK * 2; let total: usize = parts.iter().map(|p| p.1).sum();
16062 let three = parts.len() == 3;
16063 let mut y0 = self.alloc_uninit::<f32>(m * parts[0].1)?;
16064 let mut y1 = self.alloc_uninit::<f32>(m * parts[1].1)?;
16065 let mut y2 = self.alloc_uninit::<f32>(m * parts[2].1)?;
16066 let mut y3 = self.alloc_uninit::<f32>(if three { 1 } else { m * parts[3].1 })?;
16069 let cfg = LaunchConfig {
16070 grid_dim: ((total as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
16071 block_dim: (32, ROWS_PER_BLOCK, 1),
16072 shared_mem_bytes: 0,
16073 };
16074 let (inf, mi) = (in_f as i32, m as i32);
16075 let (n0, n1, n2) = (parts[0].1 as i32, parts[1].1 as i32, parts[2].1 as i32);
16076 let n3 = if three { 0i32 } else { parts[3].1 as i32 };
16077 let (s0, s1, s2) = (parts[0].2, parts[1].2, parts[2].2);
16078 let s3 = if three { 1.0f32 } else { parts[3].2 };
16079 let w3 = if three { parts[0].0 } else { parts[3].0 };
16080 let f = self.func(kname);
16081 let __s_b = self.gpu.stream();
16082 let mut b = __s_b.launch_builder(&f);
16083 b.arg(parts[0].0)
16084 .arg(parts[1].0)
16085 .arg(parts[2].0)
16086 .arg(w3)
16087 .arg(aq)
16088 .arg(ad)
16089 .arg(&mut y0)
16090 .arg(&mut y1)
16091 .arg(&mut y2)
16092 .arg(&mut y3)
16093 .arg(&inf)
16094 .arg(&n0)
16095 .arg(&n1)
16096 .arg(&n2)
16097 .arg(&n3)
16098 .arg(&mi)
16099 .arg(&s0)
16100 .arg(&s1)
16101 .arg(&s2)
16102 .arg(&s3);
16103 unsafe {
16104 b.launch(cfg)?;
16105 }
16106 Ok(Some(if three {
16107 vec![y0, y1, y2]
16108 } else {
16109 vec![y0, y1, y2, y3]
16110 }))
16111 }
16112
16113 #[allow(clippy::type_complexity)] pub fn matmul_decode_exact_dual(
16130 &self,
16131 w0: &crate::model::GpuTensor,
16132 w1: &crate::model::GpuTensor,
16133 x: &CudaSlice<f32>,
16134 m: usize,
16135 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
16136 use crate::model::GpuTensor;
16137 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16138 let on = *ON.get_or_init(|| {
16139 std::env::var("MEMRA_SPEC_DUAL_T")
16140 .map(|v| v != "0")
16141 .unwrap_or(true)
16142 });
16143 if !on
16144 || !(2..=4).contains(&m)
16145 || std::env::var("MEMRA_NO_BATCHED").is_ok()
16146 || !self.uses_q8_1_fast(w0)
16147 || !self.uses_q8_1_fast(w1)
16148 {
16149 return Ok(None);
16150 }
16151 if !self.mmvq_supports(QT_NVFP4) {
16156 return Ok(None);
16157 }
16158 let (in_f, out_f) = (w0.in_features(), w0.out_features());
16159 if w1.in_features() != in_f || w1.out_features() != out_f {
16160 return Ok(None);
16161 }
16162 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
16163 (
16164 GpuTensor::Quant {
16165 bytes: b0,
16166 qtype: q0,
16167 row_bytes: rb0,
16168 scale: s0,
16169 rp: rp0,
16170 rp4: None,
16171 ..
16172 },
16173 GpuTensor::Quant {
16174 bytes: b1,
16175 qtype: q1,
16176 row_bytes: rb1,
16177 scale: s1,
16178 rp: rp1,
16179 rp4: None,
16180 ..
16181 },
16182 ) if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 => {
16183 (b0, b1, *rb0, *s0, *s1, *rp0)
16184 }
16185 _ => return Ok(None),
16186 };
16187 if std::env::var("MEMRA_DEBUG").is_ok() {
16190 static ONCE: std::sync::Once = std::sync::Once::new();
16191 ONCE.call_once(|| eprintln!("[memra] dual gate+up batched ENGAGED (m={m} rp={rp})"));
16192 }
16193 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
16194 let (y0, y1) =
16195 self.qmatvec_batched_dual_raw(b0, b1, &aq, &ad, m, in_f, out_f, row_bytes, rp)?;
16196 let mut y0 = y0;
16197 let mut y1 = y1;
16198 if s0 != 1.0 {
16199 self.scale_inplace(&mut y0, s0, m * out_f)?;
16200 }
16201 if s1 != 1.0 {
16202 self.scale_inplace(&mut y1, s1, m * out_f)?;
16203 }
16204 Ok(Some((y0, y1)))
16205 }
16206
16207 #[allow(clippy::too_many_arguments)]
16212 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_batched_dual_raw(
16214 &self,
16215 b0: &CudaSlice<u8>,
16216 b1: &CudaSlice<u8>,
16217 aq: &CudaSlice<i8>,
16218 ad: &CudaSlice<f32>,
16219 m: usize,
16220 in_f: usize,
16221 out_f: usize,
16222 row_bytes: usize,
16223 rp: bool,
16224 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
16225 const ROWS_PER_BLOCK: u32 = 4;
16226 let mcols = Self::batched_mcols(m);
16227 let tiny_rp1 = rp
16230 && mcols == 4
16231 && out_f <= 128
16232 && std::env::var("MEMRA_NVFP4_AUX_DUAL").as_deref() != Ok("0");
16233 let (name, rows_per_block) = if tiny_rp1 {
16234 ("qmatvec_nvfp4_mmvq_dual_b4_rp", ROWS_PER_BLOCK)
16235 } else {
16236 match (mcols, rp, m) {
16237 (2, false, _) => ("qmatvec_nvfp4_mmvq_dual_b2", ROWS_PER_BLOCK),
16238 (4, false, _) => ("qmatvec_nvfp4_mmvq_dual_b4_r2", ROWS_PER_BLOCK * 2),
16239 (2, true, _) => ("qmatvec_nvfp4_mmvq_dual_b2_rp", ROWS_PER_BLOCK),
16240 (4, true, _) => ("qmatvec_nvfp4_mmvq_dual_b4_rpr2", ROWS_PER_BLOCK * 2),
16241 (8, true, 5) => ("qmatvec_nvfp4_mmvq_dual_b5_rpr2", ROWS_PER_BLOCK * 2),
16242 (8, true, 6) => ("qmatvec_nvfp4_mmvq_dual_b6_rpr2", ROWS_PER_BLOCK * 2),
16243 (8, true, 7) => ("qmatvec_nvfp4_mmvq_dual_b7_rpr2", ROWS_PER_BLOCK * 2),
16244 _ => {
16245 return Err(
16246 format!("qmatvec_batched_dual_raw: no dual kernel for m {m}").into(),
16247 );
16248 }
16249 }
16250 };
16251 let f = self.func(name);
16252 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
16253 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
16254 let cfg = LaunchConfig {
16255 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
16256 block_dim: (32, ROWS_PER_BLOCK, 1),
16257 shared_mem_bytes: 0,
16258 };
16259 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
16260 let __s_b = self.gpu.stream();
16261 let mut b = __s_b.launch_builder(&f);
16262 b.arg(b0)
16263 .arg(b1)
16264 .arg(aq)
16265 .arg(ad)
16266 .arg(&mut y0)
16267 .arg(&mut y1)
16268 .arg(&inf)
16269 .arg(&outf)
16270 .arg(&mi)
16271 .arg(&rb);
16272 unsafe {
16273 b.launch(cfg)?;
16274 }
16275 Ok((y0, y1))
16276 }
16277
16278 #[allow(clippy::type_complexity)] #[allow(clippy::manual_div_ceil)] pub fn matmul_pre_dual_noscale(
16292 &self,
16293 w0: &crate::model::GpuTensor,
16294 w1: &crate::model::GpuTensor,
16295 aq: &CudaSlice<i8>,
16296 ad: &CudaSlice<f32>,
16297 m: usize,
16298 ) -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>>
16299 {
16300 use crate::model::GpuTensor;
16301 if m != 1 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
16302 return Ok(None);
16303 }
16304 if !self.mmvq_supports(QT_NVFP4) {
16314 return Ok(None);
16315 }
16316 let (in_f, out_f) = (w0.in_features(), w0.out_features());
16317 if w1.in_features() != in_f || w1.out_features() != out_f {
16318 return Ok(None);
16319 }
16320 let no_mirror =
16333 |w: &crate::model::GpuTensor| !matches!(w, GpuTensor::Quant { rp4: Some(_), .. });
16334 if self.q8_ffn_fuse2_on()
16335 && no_mirror(w0)
16336 && no_mirror(w1)
16337 && let Some([p0, p1]) = self.q8_fused_params(&[w0, w1])
16338 {
16339 let (y0, y1) = self.q8_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2)?;
16340 return Ok(Some(((y0, 1.0), (y1, 1.0))));
16341 }
16342 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
16352 let (y0, y1) =
16353 self.e4m3_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2, 1.0, 1.0)?;
16354 return Ok(Some(((y0, p0.3), (y1, p1.3))));
16355 }
16356 let (b0, q0, rb0, s0, rp0) = match w0 {
16357 GpuTensor::Quant {
16358 bytes,
16359 qtype,
16360 row_bytes,
16361 scale,
16362 rp,
16363 ..
16364 } => (bytes, *qtype, *row_bytes, *scale, *rp),
16365 _ => return Ok(None),
16366 };
16367 let (b1, q1, rb1, s1, rp1) = match w1 {
16368 GpuTensor::Quant {
16369 bytes,
16370 qtype,
16371 row_bytes,
16372 scale,
16373 rp,
16374 ..
16375 } => (bytes, *qtype, *row_bytes, *scale, *rp),
16376 _ => return Ok(None),
16377 };
16378 if q0 != QT_NVFP4 || q1 != QT_NVFP4 || rb0 != rb1 || rp0 != rp1 {
16379 return Ok(None);
16380 }
16381 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16383 let rows_per_block = ROWS_PER_BLOCK * RPW;
16384 let f = self.func(if rp0 {
16385 "qmatvec_nvfp4_mmvq_dual_mr2_rp"
16386 } else {
16387 "qmatvec_nvfp4_mmvq_dual_mr2"
16388 });
16389 let mut y0 = self.alloc_uninit::<f32>(out_f)?;
16390 let mut y1 = self.alloc_uninit::<f32>(out_f)?;
16391 let cfg = LaunchConfig {
16392 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
16393 block_dim: (32, ROWS_PER_BLOCK, 1),
16394 shared_mem_bytes: 0,
16395 };
16396 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, rb0 as i64);
16397 let one = 1.0f32;
16400 let __s_b = self.gpu.stream();
16401 let mut b = __s_b.launch_builder(&f);
16402 b.arg(b0)
16403 .arg(b1)
16404 .arg(aq)
16405 .arg(ad)
16406 .arg(&mut y0)
16407 .arg(&mut y1)
16408 .arg(&inf)
16409 .arg(&outf)
16410 .arg(&mi)
16411 .arg(&rb)
16412 .arg(&one)
16413 .arg(&one);
16414 unsafe {
16415 b.launch(cfg)?;
16416 }
16417 Ok(Some(((y0, s0), (y1, s1))))
16418 }
16419
16420 #[allow(clippy::too_many_arguments)]
16428 #[allow(clippy::type_complexity)] pub fn matmul_nvfp4_fused3(
16430 &self,
16431 w0: &crate::model::GpuTensor,
16432 w1: &crate::model::GpuTensor,
16433 w2: &crate::model::GpuTensor,
16434 aq: &CudaSlice<i8>,
16435 ad: &CudaSlice<f32>,
16436 m: usize,
16437 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
16438 {
16439 use crate::model::GpuTensor;
16440 if !self.mmvq_supports(QT_NVFP4)
16447 || !self.uses_q8_1_fast(w0)
16448 || !self.uses_q8_1_fast(w1)
16449 || !self.uses_q8_1_fast(w2)
16450 {
16451 return Ok(None);
16452 }
16453 if (9..=16).contains(&m) {
16456 return Ok(
16457 match self.matmul_decode_exact_group3_pre([w0, w1, w2], aq, ad, m)? {
16458 Some(mut ys) => {
16459 let y2 = ys.pop().unwrap();
16460 let y1 = ys.pop().unwrap();
16461 let y0 = ys.pop().unwrap();
16462 Some((y0, y1, y2))
16463 }
16464 None => None,
16465 },
16466 );
16467 }
16468 if !(1..=8).contains(&m) {
16469 return Ok(None);
16470 }
16471 if m > 1 {
16472 let in_f = w0.in_features();
16473 if std::env::var("MEMRA_NVFP4_FUSED3B").as_deref() == Ok("0")
16474 || !self.batched_supports(QT_NVFP4)
16475 || std::env::var("MEMRA_NO_BATCHED").is_ok()
16476 || (m > 4 && !Self::b8_enabled())
16477 || !in_f.is_multiple_of(512)
16478 || in_f / 64 > 272
16479 {
16480 return Ok(None);
16481 }
16482 }
16483 let unpack = |w: &crate::model::GpuTensor| match w {
16484 GpuTensor::Quant {
16485 bytes,
16486 qtype,
16487 scale,
16488 rp,
16489 ..
16490 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
16491 _ => None,
16492 };
16493 let (Some(p0), Some(p1), Some(p2)) = (unpack(w0), unpack(w1), unpack(w2)) else {
16494 return Ok(None);
16495 };
16496 let in_f = w0.in_features();
16497 if w1.in_features() != in_f || w2.in_features() != in_f {
16498 return Ok(None);
16499 }
16500 let (o0, o1, o2) = (w0.out_features(), w1.out_features(), w2.out_features());
16501 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16503 let rows_pb = ROWS_PER_BLOCK * RPW;
16504 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
16505 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
16506 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
16507 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
16508 let (inf, oi0, oi1, oi2, mi) = (in_f as i32, o0 as i32, o1 as i32, o2 as i32, m as i32);
16509 let (b0, b1, b2) = unsafe { (&*p0.0, &*p1.0, &*p2.0) };
16512 if m > 1 {
16513 if p0.1 != 1.0 || p1.1 != 1.0 || p2.1 != 1.0 {
16515 return Ok(None);
16516 }
16517 let f = self.func("qmatvec_nvfp4_mmvq_fused3_b8_rpsc");
16518 let cfg = LaunchConfig {
16519 grid_dim: (nb(o0) + nb(o1) + nb(o2), 1, 1),
16520 block_dim: (32, ROWS_PER_BLOCK, 1),
16521 shared_mem_bytes: 0,
16522 };
16523 let __s_b = self.gpu.stream();
16524 let mut b = __s_b.launch_builder(&f);
16525 b.arg(b0)
16526 .arg(b1)
16527 .arg(b2)
16528 .arg(aq)
16529 .arg(ad)
16530 .arg(&mut y0)
16531 .arg(&mut y1)
16532 .arg(&mut y2)
16533 .arg(&inf)
16534 .arg(&oi0)
16535 .arg(&oi1)
16536 .arg(&oi2)
16537 .arg(&mi);
16538 unsafe {
16539 b.launch(cfg)?;
16540 }
16541 return Ok(Some((y0, y1, y2)));
16542 }
16543 let f = self.func("qmatvec_nvfp4_mmvq_fused3_rp");
16544 let cfg = LaunchConfig {
16545 grid_dim: (nb(o0) + nb(o1) + nb(o2), m as u32, 1),
16546 block_dim: (32, ROWS_PER_BLOCK, 1),
16547 shared_mem_bytes: 0,
16548 };
16549 let __s_b = self.gpu.stream();
16550 let mut b = __s_b.launch_builder(&f);
16551 b.arg(b0)
16552 .arg(b1)
16553 .arg(b2)
16554 .arg(aq)
16555 .arg(ad)
16556 .arg(&mut y0)
16557 .arg(&mut y1)
16558 .arg(&mut y2)
16559 .arg(&inf)
16560 .arg(&oi0)
16561 .arg(&oi1)
16562 .arg(&oi2)
16563 .arg(&mi)
16564 .arg(&p0.1)
16565 .arg(&p1.1)
16566 .arg(&p2.1);
16567 unsafe {
16568 b.launch(cfg)?;
16569 }
16570 Ok(Some((y0, y1, y2)))
16571 }
16572
16573 #[allow(clippy::type_complexity)] pub fn matmul_nvfp4_fused2(
16583 &self,
16584 w0: &crate::model::GpuTensor,
16585 w1: &crate::model::GpuTensor,
16586 aq: &CudaSlice<i8>,
16587 ad: &CudaSlice<f32>,
16588 m: usize,
16589 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
16590 use crate::model::GpuTensor;
16591 static FUSED2_OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16592 let off =
16593 *FUSED2_OFF.get_or_init(|| std::env::var("MEMRA_NVFP4_FUSED2").as_deref() == Ok("0"));
16594 if off
16597 || m != 1
16598 || !self.mmvq_supports(QT_NVFP4)
16599 || !self.uses_q8_1_fast(w0)
16600 || !self.uses_q8_1_fast(w1)
16601 {
16602 return Ok(None);
16603 }
16604 let unpack = |w: &crate::model::GpuTensor| match w {
16605 GpuTensor::Quant {
16606 bytes,
16607 qtype,
16608 scale,
16609 rp,
16610 ..
16611 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
16612 _ => None,
16613 };
16614 let (Some(p0), Some(p1)) = (unpack(w0), unpack(w1)) else {
16615 return Ok(None);
16616 };
16617 let in_f = w0.in_features();
16618 if w1.in_features() != in_f {
16619 return Ok(None);
16620 }
16621 let (o0, o1) = (w0.out_features(), w1.out_features());
16622 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16624 let rows_pb = ROWS_PER_BLOCK * RPW;
16625 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
16626 let f = self.func("qmatvec_nvfp4_mmvq_fused2_rp");
16627 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
16628 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
16629 let cfg = LaunchConfig {
16630 grid_dim: (nb(o0) + nb(o1), m as u32, 1),
16631 block_dim: (32, ROWS_PER_BLOCK, 1),
16632 shared_mem_bytes: 0,
16633 };
16634 let (inf, oi0, oi1, mi) = (in_f as i32, o0 as i32, o1 as i32, m as i32);
16635 let (b0, b1) = unsafe { (&*p0.0, &*p1.0) };
16638 if Self::pdl_on() && Self::pdl_mmvq_on() && Self::pdl_nvfp4q8_on() {
16641 {
16642 use cudarc::driver::{DevicePtr, DevicePtrMut};
16643 let s = &self.gpu.stream();
16644 let (pw0, _g0) = b0.device_ptr(s);
16645 let (pw1, _g1) = b1.device_ptr(s);
16646 let (paq, _g2) = aq.device_ptr(s);
16647 let (pad, _g3) = ad.device_ptr(s);
16648 let (py0, _g4) = y0.device_ptr_mut(s);
16649 let (py1, _g5) = y1.device_ptr_mut(s);
16650 let (s0, s1) = (p0.1, p1.1);
16651 let mut ps = [
16652 &pw0 as *const _ as *mut std::ffi::c_void,
16653 &pw1 as *const _ as *mut _,
16654 &paq as *const _ as *mut _,
16655 &pad as *const _ as *mut _,
16656 &py0 as *const _ as *mut _,
16657 &py1 as *const _ as *mut _,
16658 &inf as *const _ as *mut _,
16659 &oi0 as *const _ as *mut _,
16660 &oi1 as *const _ as *mut _,
16661 &mi as *const _ as *mut _,
16662 &s0 as *const _ as *mut _,
16663 &s1 as *const _ as *mut _,
16664 ];
16665 unsafe {
16666 self.launch_pdl(
16667 "qmatvec_nvfp4_mmvq_fused2_rp",
16668 cfg.grid_dim,
16669 cfg.block_dim,
16670 &mut ps,
16671 )?;
16672 }
16673 }
16674 return Ok(Some((y0, y1)));
16675 }
16676 let __s_b = self.gpu.stream();
16677 let mut b = __s_b.launch_builder(&f);
16678 b.arg(b0)
16679 .arg(b1)
16680 .arg(aq)
16681 .arg(ad)
16682 .arg(&mut y0)
16683 .arg(&mut y1)
16684 .arg(&inf)
16685 .arg(&oi0)
16686 .arg(&oi1)
16687 .arg(&mi)
16688 .arg(&p0.1)
16689 .arg(&p1.1);
16690 unsafe {
16691 b.launch(cfg)?;
16692 }
16693 Ok(Some((y0, y1)))
16694 }
16695
16696 pub fn matmul_nvfp4_fused2_into(
16701 &self,
16702 w0: &crate::model::GpuTensor,
16703 w1: &crate::model::GpuTensor,
16704 aq: &CudaSlice<i8>,
16705 ad: &CudaSlice<f32>,
16706 y0: &mut CudaSlice<f32>,
16707 y1: &mut CudaSlice<f32>,
16708 ) -> Result<bool, Box<dyn std::error::Error>> {
16709 use crate::model::GpuTensor;
16710 static FUSED2_OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16711 let off =
16712 *FUSED2_OFF.get_or_init(|| std::env::var("MEMRA_NVFP4_FUSED2").as_deref() == Ok("0"));
16713 if off
16714 || !self.mmvq_supports(QT_NVFP4)
16715 || !self.uses_q8_1_fast(w0)
16716 || !self.uses_q8_1_fast(w1)
16717 {
16718 return Ok(false);
16719 }
16720 let unpack = |w: &crate::model::GpuTensor| match w {
16721 GpuTensor::Quant {
16722 bytes,
16723 qtype,
16724 scale,
16725 rp,
16726 ..
16727 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
16728 _ => None,
16729 };
16730 let (Some(p0), Some(p1)) = (unpack(w0), unpack(w1)) else {
16731 return Ok(false);
16732 };
16733 let in_f = w0.in_features();
16734 if w1.in_features() != in_f {
16735 return Ok(false);
16736 }
16737 let (o0, o1) = (w0.out_features(), w1.out_features());
16738 if y0.len() < o0 || y1.len() < o1 {
16739 return Ok(false);
16740 }
16741 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16743 let rows_pb = ROWS_PER_BLOCK * RPW;
16744 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
16745 let f = self.func("qmatvec_nvfp4_mmvq_fused2_rp");
16746 let cfg = LaunchConfig {
16747 grid_dim: (nb(o0) + nb(o1), 1, 1),
16748 block_dim: (32, ROWS_PER_BLOCK, 1),
16749 shared_mem_bytes: 0,
16750 };
16751 let (inf, oi0, oi1, mi) = (in_f as i32, o0 as i32, o1 as i32, 1i32);
16752 let (b0, b1) = unsafe { (&*p0.0, &*p1.0) };
16755 let __s_b = self.gpu.stream();
16756 let mut b = __s_b.launch_builder(&f);
16757 b.arg(b0)
16758 .arg(b1)
16759 .arg(aq)
16760 .arg(ad)
16761 .arg(&mut *y0)
16762 .arg(&mut *y1)
16763 .arg(&inf)
16764 .arg(&oi0)
16765 .arg(&oi1)
16766 .arg(&mi)
16767 .arg(&p0.1)
16768 .arg(&p1.1);
16769 unsafe {
16770 b.launch(cfg)?;
16771 }
16772 Ok(true)
16773 }
16774
16775 #[allow(clippy::type_complexity)]
16780 #[allow(clippy::too_many_arguments)] pub fn matmul_nvfp4_fused4(
16782 &self,
16783 w0: &crate::model::GpuTensor,
16784 w1: &crate::model::GpuTensor,
16785 w2: &crate::model::GpuTensor,
16786 w3: &crate::model::GpuTensor,
16787 aq: &CudaSlice<i8>,
16788 ad: &CudaSlice<f32>,
16789 m: usize,
16790 ) -> Result<
16791 Option<(
16792 CudaSlice<f32>,
16793 CudaSlice<f32>,
16794 CudaSlice<f32>,
16795 CudaSlice<f32>,
16796 )>,
16797 Box<dyn std::error::Error>,
16798 > {
16799 use crate::model::GpuTensor;
16800 if std::env::var("MEMRA_NVFP4_FUSED4").as_deref() == Ok("0")
16807 || !self.mmvq_supports(QT_NVFP4)
16808 || !self.uses_q8_1_fast(w0)
16809 || !self.uses_q8_1_fast(w1)
16810 || !self.uses_q8_1_fast(w2)
16811 || !self.uses_q8_1_fast(w3)
16812 {
16813 return Ok(None);
16814 }
16815 if (9..=16).contains(&m) {
16820 return Ok(
16821 match self.matmul_decode_exact_group4_pre([w0, w1, w2, w3], aq, ad, m)? {
16822 Some(mut ys) => {
16823 let y3 = ys.pop().unwrap();
16824 let y2 = ys.pop().unwrap();
16825 let y1 = ys.pop().unwrap();
16826 let y0 = ys.pop().unwrap();
16827 Some((y0, y1, y2, y3))
16828 }
16829 None => None,
16830 },
16831 );
16832 }
16833 if !(1..=8).contains(&m) {
16834 return Ok(None);
16835 }
16836 if m > 1 {
16837 let in_f = w0.in_features();
16840 if !self.batched_supports(QT_NVFP4)
16841 || std::env::var("MEMRA_NO_BATCHED").is_ok()
16842 || (m > 4 && !Self::b8_enabled())
16843 || !in_f.is_multiple_of(512)
16844 || in_f / 64 > 272
16845 {
16846 return Ok(None);
16847 }
16848 }
16849 let unpack = |w: &crate::model::GpuTensor| match w {
16850 GpuTensor::Quant {
16851 bytes,
16852 qtype,
16853 scale,
16854 rp,
16855 ..
16856 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
16857 _ => None,
16858 };
16859 let (Some(p0), Some(p1), Some(p2), Some(p3)) =
16860 (unpack(w0), unpack(w1), unpack(w2), unpack(w3))
16861 else {
16862 return Ok(None);
16863 };
16864 let in_f = w0.in_features();
16865 if w1.in_features() != in_f || w2.in_features() != in_f || w3.in_features() != in_f {
16866 return Ok(None);
16867 }
16868 let (o0, o1, o2, o3) = (
16869 w0.out_features(),
16870 w1.out_features(),
16871 w2.out_features(),
16872 w3.out_features(),
16873 );
16874 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
16876 let rows_pb = ROWS_PER_BLOCK * RPW;
16877 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
16878 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
16879 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
16880 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
16881 let mut y3 = self.alloc_uninit::<f32>(m * o3)?;
16882 let (inf, oi0, oi1, oi2, oi3, mi) = (
16883 in_f as i32,
16884 o0 as i32,
16885 o1 as i32,
16886 o2 as i32,
16887 o3 as i32,
16888 m as i32,
16889 );
16890 let (b0, b1, b2, b3) = unsafe { (&*p0.0, &*p1.0, &*p2.0, &*p3.0) };
16893 if m > 1 {
16894 if p0.1 != 1.0 || p1.1 != 1.0 || p2.1 != 1.0 || p3.1 != 1.0 {
16897 return Ok(None);
16898 }
16899 let f = self.func("qmatvec_nvfp4_mmvq_fused4_b8_rpsc");
16900 let cfg = LaunchConfig {
16901 grid_dim: (nb(o0) + nb(o1) + nb(o2) + nb(o3), 1, 1),
16902 block_dim: (32, ROWS_PER_BLOCK, 1),
16903 shared_mem_bytes: 0,
16904 };
16905 let __s_b = self.gpu.stream();
16906 let mut b = __s_b.launch_builder(&f);
16907 b.arg(b0)
16908 .arg(b1)
16909 .arg(b2)
16910 .arg(b3)
16911 .arg(aq)
16912 .arg(ad)
16913 .arg(&mut y0)
16914 .arg(&mut y1)
16915 .arg(&mut y2)
16916 .arg(&mut y3)
16917 .arg(&inf)
16918 .arg(&oi0)
16919 .arg(&oi1)
16920 .arg(&oi2)
16921 .arg(&oi3)
16922 .arg(&mi);
16923 unsafe {
16924 b.launch(cfg)?;
16925 }
16926 return Ok(Some((y0, y1, y2, y3)));
16927 }
16928 let f = self.func("qmatvec_nvfp4_mmvq_fused4_rp");
16929 let cfg = LaunchConfig {
16930 grid_dim: (nb(o0) + nb(o1) + nb(o2) + nb(o3), m as u32, 1),
16931 block_dim: (32, ROWS_PER_BLOCK, 1),
16932 shared_mem_bytes: 0,
16933 };
16934 let __s_b = self.gpu.stream();
16935 let mut b = __s_b.launch_builder(&f);
16936 b.arg(b0)
16937 .arg(b1)
16938 .arg(b2)
16939 .arg(b3)
16940 .arg(aq)
16941 .arg(ad)
16942 .arg(&mut y0)
16943 .arg(&mut y1)
16944 .arg(&mut y2)
16945 .arg(&mut y3)
16946 .arg(&inf)
16947 .arg(&oi0)
16948 .arg(&oi1)
16949 .arg(&oi2)
16950 .arg(&oi3)
16951 .arg(&mi)
16952 .arg(&p0.1)
16953 .arg(&p1.1)
16954 .arg(&p2.1)
16955 .arg(&p3.1);
16956 unsafe {
16957 b.launch(cfg)?;
16958 }
16959 Ok(Some((y0, y1, y2, y3)))
16960 }
16961
16962 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused2(
16971 &self,
16972 w0: &crate::model::GpuTensor,
16973 w1: &crate::model::GpuTensor,
16974 aq: &CudaSlice<i8>,
16975 ad: &CudaSlice<f32>,
16976 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
16977 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
16983 return Ok(Some(self.e4m3_fused2_core(
16984 p0.0,
16985 p1.0,
16986 aq,
16987 ad,
16988 w0.in_features(),
16989 p0.1,
16990 p1.1,
16991 p0.2,
16992 p0.3,
16993 p1.3,
16994 )?));
16995 }
16996 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
16997 return Ok(None);
16998 };
16999 Ok(Some(self.q8_fused2_core(
17000 p0.0,
17001 p1.0,
17002 aq,
17003 ad,
17004 w0.in_features(),
17005 p0.1,
17006 p1.1,
17007 p0.2,
17008 )?))
17009 }
17010
17011 #[allow(clippy::too_many_arguments)]
17012 fn q8_fused2_core(
17013 &self,
17014 b0: &CudaSlice<u8>,
17015 b1: &CudaSlice<u8>,
17016 aq: &CudaSlice<i8>,
17017 ad: &CudaSlice<f32>,
17018 in_f: usize,
17019 out0: usize,
17020 out1: usize,
17021 row_bytes: usize,
17022 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
17023 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
17025 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
17026 let f = self.func("qmatvec_q8_0_mmvq_fused2");
17027 let mut y0 = self.alloc_uninit::<f32>(out0)?;
17028 let mut y1 = self.alloc_uninit::<f32>(out1)?;
17029 let cfg = LaunchConfig {
17030 grid_dim: (nb0 + nb1, 1, 1),
17031 block_dim: (32, ROWS_PER_BLOCK, 1),
17032 shared_mem_bytes: 0,
17033 };
17034 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
17035 let __s_b = self.gpu.stream();
17036 let mut b = __s_b.launch_builder(&f);
17037 b.arg(b0)
17038 .arg(b1)
17039 .arg(aq)
17040 .arg(ad)
17041 .arg(&mut y0)
17042 .arg(&mut y1)
17043 .arg(&inf)
17044 .arg(&o0)
17045 .arg(&o1)
17046 .arg(&rbl);
17047 unsafe {
17048 b.launch(cfg)?;
17049 }
17050 Ok((y0, y1))
17051 }
17052
17053 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused2_x(
17060 &self,
17061 w0: &crate::model::GpuTensor,
17062 w1: &crate::model::GpuTensor,
17063 x: &CudaSlice<f32>,
17064 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
17065 if !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
17066 return Ok(None);
17067 }
17068 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
17069 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
17070 return Ok(Some(self.e4m3_fused2_core(
17071 p0.0,
17072 p1.0,
17073 &aq,
17074 &ad,
17075 w0.in_features(),
17076 p0.1,
17077 p1.1,
17078 p0.2,
17079 p0.3,
17080 p1.3,
17081 )?));
17082 }
17083 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
17084 return Ok(None);
17085 };
17086 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
17087 Ok(Some(self.q8_fused2_core(
17088 p0.0,
17089 p1.0,
17090 &aq,
17091 &ad,
17092 w0.in_features(),
17093 p0.1,
17094 p1.1,
17095 p0.2,
17096 )?))
17097 }
17098
17099 #[allow(clippy::too_many_arguments)]
17102 pub fn qmatvec_q8_fused2_raw(
17103 &self,
17104 b0: &CudaSlice<u8>,
17105 b1: &CudaSlice<u8>,
17106 x: &CudaSlice<f32>,
17107 in_f: usize,
17108 out0: usize,
17109 out1: usize,
17110 row_bytes: usize,
17111 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
17112 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
17113 self.q8_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes)
17114 }
17115
17116 #[allow(clippy::type_complexity)] pub fn matmul_q4_fused3(
17123 &self,
17124 w0: &crate::model::GpuTensor,
17125 w1: &crate::model::GpuTensor,
17126 w2: &crate::model::GpuTensor,
17127 aq: &CudaSlice<i8>,
17128 ad: &CudaSlice<f32>,
17129 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
17130 {
17131 use crate::model::GpuTensor;
17132 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17133 match w {
17134 GpuTensor::Quant {
17135 qtype, row_bytes, ..
17136 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17137 _ => None,
17138 }
17139 };
17140 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2)) else {
17141 return Ok(None);
17142 };
17143 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
17144 return Ok(None);
17145 }
17146 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17150 match w {
17151 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17152 Some(m) => (m, true),
17153 None => (bytes, *rp),
17154 },
17155 _ => unreachable!(),
17156 }
17157 }
17158 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
17159 if rp0 != rp1 || rp1 != rp2 {
17160 return Ok(None);
17161 }
17162 let rp = rp0;
17163 let rpb: u32 = 4;
17164 let mr1 = rp && Self::q40_mr1_on();
17168 let nb = |o: usize| {
17169 if mr1 {
17170 (o as u32).div_ceil(rpb)
17171 } else {
17172 (o as u32).div_ceil(2).div_ceil(rpb)
17173 }
17174 };
17175 let grid = nb(o0) + nb(o1) + nb(o2);
17176 let mut y0 = self.alloc_uninit::<f32>(o0)?;
17177 let mut y1 = self.alloc_uninit::<f32>(o1)?;
17178 let mut y2 = self.alloc_uninit::<f32>(o2)?;
17179 let f = self.func(if mr1 {
17180 "qmatvec_q4_0_mmvq_fused3_mr1_rp"
17181 } else if rp {
17182 "qmatvec_q4_0_mmvq_fused3_rp"
17183 } else {
17184 "qmatvec_q4_0_mmvq_fused3"
17185 });
17186 let cfg = LaunchConfig {
17187 grid_dim: (grid, 1, 1),
17188 block_dim: (32, rpb, 1),
17189 shared_mem_bytes: 0,
17190 };
17191 let inf = w0.in_features() as i32;
17192 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
17193 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
17194 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
17197 {
17198 use cudarc::driver::{DevicePtr, DevicePtrMut};
17199 let s = &self.gpu.stream();
17200 let (p0, _g0) = b0.device_ptr(s);
17201 let (p1, _g1) = b1.device_ptr(s);
17202 let (p2, _g2) = b2.device_ptr(s);
17203 let (paq, _g3) = aq.device_ptr(s);
17204 let (pad, _g4) = ad.device_ptr(s);
17205 let (py0, _g5) = y0.device_ptr_mut(s);
17206 let (py1, _g6) = y1.device_ptr_mut(s);
17207 let (py2, _g7) = y2.device_ptr_mut(s);
17208 let mut ps = [
17209 &p0 as *const _ as *mut std::ffi::c_void,
17210 &p1 as *const _ as *mut _,
17211 &p2 as *const _ as *mut _,
17212 &paq as *const _ as *mut _,
17213 &pad as *const _ as *mut _,
17214 &py0 as *const _ as *mut _,
17215 &py1 as *const _ as *mut _,
17216 &py2 as *const _ as *mut _,
17217 &inf as *const _ as *mut _,
17218 &oo0 as *const _ as *mut _,
17219 &oo1 as *const _ as *mut _,
17220 &oo2 as *const _ as *mut _,
17221 &r0 as *const _ as *mut _,
17222 &r1 as *const _ as *mut _,
17223 &r2 as *const _ as *mut _,
17224 ];
17225 unsafe {
17226 self.launch_pdl(
17227 "qmatvec_q4_0_mmvq_fused3_mr1_rp",
17228 (grid, 1, 1),
17229 (32, rpb, 1),
17230 &mut ps,
17231 )?;
17232 }
17233 }
17234 return Ok(Some((y0, y1, y2)));
17235 }
17236 let __s_b = self.gpu.stream();
17237 let mut b = __s_b.launch_builder(&f);
17238 b.arg(b0)
17239 .arg(b1)
17240 .arg(b2)
17241 .arg(aq)
17242 .arg(ad)
17243 .arg(&mut y0)
17244 .arg(&mut y1)
17245 .arg(&mut y2)
17246 .arg(&inf)
17247 .arg(&oo0)
17248 .arg(&oo1)
17249 .arg(&oo2)
17250 .arg(&r0)
17251 .arg(&r1)
17252 .arg(&r2);
17253 unsafe {
17254 b.launch(cfg)?;
17255 }
17256 Ok(Some((y0, y1, y2)))
17257 }
17258
17259 #[allow(clippy::too_many_arguments)]
17262 pub fn matmul_q4_fused3_into(
17263 &self,
17264 w0: &crate::model::GpuTensor,
17265 w1: &crate::model::GpuTensor,
17266 w2: &crate::model::GpuTensor,
17267 aq: &CudaSlice<i8>,
17268 ad: &CudaSlice<f32>,
17269 y0: &mut CudaSlice<f32>,
17270 y1: &mut CudaSlice<f32>,
17271 y2: &mut CudaSlice<f32>,
17272 ) -> Result<bool, Box<dyn std::error::Error>> {
17273 use crate::model::GpuTensor;
17274 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17275 match w {
17276 GpuTensor::Quant {
17277 qtype, row_bytes, ..
17278 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17279 _ => None,
17280 }
17281 };
17282 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2)) else {
17283 return Ok(false);
17284 };
17285 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
17286 return Ok(false);
17287 }
17288 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17289 match w {
17290 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17291 Some(m) => (m, true),
17292 None => (bytes, *rp),
17293 },
17294 _ => unreachable!(),
17295 }
17296 }
17297 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
17298 if rp0 != rp1 || rp1 != rp2 {
17299 return Ok(false);
17300 }
17301 let rp = rp0;
17302 let rpb: u32 = 4;
17303 let mr1 = rp && Self::q40_mr1_on();
17304 let nb = |o: usize| {
17305 if mr1 {
17306 (o as u32).div_ceil(rpb)
17307 } else {
17308 (o as u32).div_ceil(2).div_ceil(rpb)
17309 }
17310 };
17311 let grid = nb(o0) + nb(o1) + nb(o2);
17312 debug_assert!(y0.len() >= o0 && y1.len() >= o1 && y2.len() >= o2);
17313 let f = self.func(if mr1 {
17314 "qmatvec_q4_0_mmvq_fused3_mr1_rp"
17315 } else if rp {
17316 "qmatvec_q4_0_mmvq_fused3_rp"
17317 } else {
17318 "qmatvec_q4_0_mmvq_fused3"
17319 });
17320 let cfg = LaunchConfig {
17321 grid_dim: (grid, 1, 1),
17322 block_dim: (32, rpb, 1),
17323 shared_mem_bytes: 0,
17324 };
17325 let inf = w0.in_features() as i32;
17326 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
17327 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
17328 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
17330 use cudarc::driver::{DevicePtr, DevicePtrMut};
17331 let s = &self.gpu.stream();
17332 let (p0, _g0) = b0.device_ptr(s);
17333 let (p1, _g1) = b1.device_ptr(s);
17334 let (p2, _g2) = b2.device_ptr(s);
17335 let (paq, _g3) = aq.device_ptr(s);
17336 let (pad, _g4) = ad.device_ptr(s);
17337 let (py0, _g5) = y0.device_ptr_mut(s);
17338 let (py1, _g6) = y1.device_ptr_mut(s);
17339 let (py2, _g7) = y2.device_ptr_mut(s);
17340 let mut ps = [
17341 &p0 as *const _ as *mut std::ffi::c_void,
17342 &p1 as *const _ as *mut _,
17343 &p2 as *const _ as *mut _,
17344 &paq as *const _ as *mut _,
17345 &pad as *const _ as *mut _,
17346 &py0 as *const _ as *mut _,
17347 &py1 as *const _ as *mut _,
17348 &py2 as *const _ as *mut _,
17349 &inf as *const _ as *mut _,
17350 &oo0 as *const _ as *mut _,
17351 &oo1 as *const _ as *mut _,
17352 &oo2 as *const _ as *mut _,
17353 &r0 as *const _ as *mut _,
17354 &r1 as *const _ as *mut _,
17355 &r2 as *const _ as *mut _,
17356 ];
17357 unsafe {
17358 self.launch_pdl(
17359 "qmatvec_q4_0_mmvq_fused3_mr1_rp",
17360 (grid, 1, 1),
17361 (32, rpb, 1),
17362 &mut ps,
17363 )?;
17364 }
17365 return Ok(true);
17366 }
17367 let __s_b = self.gpu.stream();
17368 let mut b = __s_b.launch_builder(&f);
17369 b.arg(b0)
17370 .arg(b1)
17371 .arg(b2)
17372 .arg(aq)
17373 .arg(ad)
17374 .arg(&mut *y0)
17375 .arg(&mut *y1)
17376 .arg(&mut *y2)
17377 .arg(&inf)
17378 .arg(&oo0)
17379 .arg(&oo1)
17380 .arg(&oo2)
17381 .arg(&r0)
17382 .arg(&r1)
17383 .arg(&r2);
17384 unsafe {
17385 b.launch(cfg)?;
17386 }
17387 Ok(true)
17388 }
17389
17390 #[allow(clippy::type_complexity)] pub fn matmul_q4_fused2(
17393 &self,
17394 w0: &crate::model::GpuTensor,
17395 w1: &crate::model::GpuTensor,
17396 aq: &CudaSlice<i8>,
17397 ad: &CudaSlice<f32>,
17398 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
17399 use crate::model::GpuTensor;
17400 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17401 match w {
17402 GpuTensor::Quant {
17403 qtype, row_bytes, ..
17404 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17405 _ => None,
17406 }
17407 };
17408 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else {
17409 return Ok(None);
17410 };
17411 if w0.in_features() != w1.in_features() {
17412 return Ok(None);
17413 }
17414 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17416 match w {
17417 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17418 Some(m) => (m, true),
17419 None => (bytes, *rp),
17420 },
17421 _ => unreachable!(),
17422 }
17423 }
17424 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
17425 if rp0 != rp1 {
17426 return Ok(None);
17427 }
17428 let rp = rp0;
17429 let rpb: u32 = 4;
17430 let mr1 = rp && Self::q40_mr1_on();
17432 let nb = |o: usize| {
17433 if mr1 {
17434 (o as u32).div_ceil(rpb)
17435 } else {
17436 (o as u32).div_ceil(2).div_ceil(rpb)
17437 }
17438 };
17439 let grid = nb(o0) + nb(o1);
17440 let mut y0 = self.alloc_uninit::<f32>(o0)?;
17441 let mut y1 = self.alloc_uninit::<f32>(o1)?;
17442 let f = self.func(if mr1 {
17443 "qmatvec_q4_0_mmvq_fused2_mr1_rp"
17444 } else if rp {
17445 "qmatvec_q4_0_mmvq_fused2_rp"
17446 } else {
17447 "qmatvec_q4_0_mmvq_fused2"
17448 });
17449 let cfg = LaunchConfig {
17450 grid_dim: (grid, 1, 1),
17451 block_dim: (32, rpb, 1),
17452 shared_mem_bytes: 0,
17453 };
17454 let inf = w0.in_features() as i32;
17455 let (oo0, oo1) = (o0 as i32, o1 as i32);
17456 let (r0, r1) = (rb0 as i64, rb1 as i64);
17457 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
17459 {
17460 use cudarc::driver::{DevicePtr, DevicePtrMut};
17461 let s = &self.gpu.stream();
17462 let (p0, _g0) = b0.device_ptr(s);
17463 let (p1, _g1) = b1.device_ptr(s);
17464 let (paq, _g2) = aq.device_ptr(s);
17465 let (pad, _g3) = ad.device_ptr(s);
17466 let (py0, _g4) = y0.device_ptr_mut(s);
17467 let (py1, _g5) = y1.device_ptr_mut(s);
17468 let mut ps = [
17469 &p0 as *const _ as *mut std::ffi::c_void,
17470 &p1 as *const _ as *mut _,
17471 &paq as *const _ as *mut _,
17472 &pad as *const _ as *mut _,
17473 &py0 as *const _ as *mut _,
17474 &py1 as *const _ as *mut _,
17475 &inf as *const _ as *mut _,
17476 &oo0 as *const _ as *mut _,
17477 &oo1 as *const _ as *mut _,
17478 &r0 as *const _ as *mut _,
17479 &r1 as *const _ as *mut _,
17480 ];
17481 unsafe {
17482 self.launch_pdl(
17483 "qmatvec_q4_0_mmvq_fused2_mr1_rp",
17484 (grid, 1, 1),
17485 (32, rpb, 1),
17486 &mut ps,
17487 )?;
17488 }
17489 }
17490 return Ok(Some((y0, y1)));
17491 }
17492 let __s_b = self.gpu.stream();
17493 let mut b = __s_b.launch_builder(&f);
17494 b.arg(b0)
17495 .arg(b1)
17496 .arg(aq)
17497 .arg(ad)
17498 .arg(&mut y0)
17499 .arg(&mut y1)
17500 .arg(&inf)
17501 .arg(&oo0)
17502 .arg(&oo1)
17503 .arg(&r0)
17504 .arg(&r1);
17505 unsafe {
17506 b.launch(cfg)?;
17507 }
17508 Ok(Some((y0, y1)))
17509 }
17510
17511 pub fn matmul_q4_fused2_into(
17513 &self,
17514 w0: &crate::model::GpuTensor,
17515 w1: &crate::model::GpuTensor,
17516 aq: &CudaSlice<i8>,
17517 ad: &CudaSlice<f32>,
17518 y0: &mut CudaSlice<f32>,
17519 y1: &mut CudaSlice<f32>,
17520 ) -> Result<bool, Box<dyn std::error::Error>> {
17521 use crate::model::GpuTensor;
17522 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17523 match w {
17524 GpuTensor::Quant {
17525 qtype, row_bytes, ..
17526 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17527 _ => None,
17528 }
17529 };
17530 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else {
17531 return Ok(false);
17532 };
17533 if w0.in_features() != w1.in_features() {
17534 return Ok(false);
17535 }
17536 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17537 match w {
17538 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17539 Some(m) => (m, true),
17540 None => (bytes, *rp),
17541 },
17542 _ => unreachable!(),
17543 }
17544 }
17545 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
17546 if rp0 != rp1 {
17547 return Ok(false);
17548 }
17549 let rp = rp0;
17550 let rpb: u32 = 4;
17551 let mr1 = rp && Self::q40_mr1_on();
17552 let nb = |o: usize| {
17553 if mr1 {
17554 (o as u32).div_ceil(rpb)
17555 } else {
17556 (o as u32).div_ceil(2).div_ceil(rpb)
17557 }
17558 };
17559 let grid = nb(o0) + nb(o1);
17560 debug_assert!(y0.len() >= o0 && y1.len() >= o1);
17561 let f = self.func(if mr1 {
17562 "qmatvec_q4_0_mmvq_fused2_mr1_rp"
17563 } else if rp {
17564 "qmatvec_q4_0_mmvq_fused2_rp"
17565 } else {
17566 "qmatvec_q4_0_mmvq_fused2"
17567 });
17568 let cfg = LaunchConfig {
17569 grid_dim: (grid, 1, 1),
17570 block_dim: (32, rpb, 1),
17571 shared_mem_bytes: 0,
17572 };
17573 let inf = w0.in_features() as i32;
17574 let (oo0, oo1) = (o0 as i32, o1 as i32);
17575 let (r0, r1) = (rb0 as i64, rb1 as i64);
17576 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
17578 use cudarc::driver::{DevicePtr, DevicePtrMut};
17579 let s = &self.gpu.stream();
17580 let (p0, _g0) = b0.device_ptr(s);
17581 let (p1, _g1) = b1.device_ptr(s);
17582 let (paq, _g2) = aq.device_ptr(s);
17583 let (pad, _g3) = ad.device_ptr(s);
17584 let (py0, _g4) = y0.device_ptr_mut(s);
17585 let (py1, _g5) = y1.device_ptr_mut(s);
17586 let mut ps = [
17587 &p0 as *const _ as *mut std::ffi::c_void,
17588 &p1 as *const _ as *mut _,
17589 &paq as *const _ as *mut _,
17590 &pad as *const _ as *mut _,
17591 &py0 as *const _ as *mut _,
17592 &py1 as *const _ as *mut _,
17593 &inf as *const _ as *mut _,
17594 &oo0 as *const _ as *mut _,
17595 &oo1 as *const _ as *mut _,
17596 &r0 as *const _ as *mut _,
17597 &r1 as *const _ as *mut _,
17598 ];
17599 unsafe {
17600 self.launch_pdl(
17601 "qmatvec_q4_0_mmvq_fused2_mr1_rp",
17602 (grid, 1, 1),
17603 (32, rpb, 1),
17604 &mut ps,
17605 )?;
17606 }
17607 return Ok(true);
17608 }
17609 let __s_b = self.gpu.stream();
17610 let mut b = __s_b.launch_builder(&f);
17611 b.arg(b0)
17612 .arg(b1)
17613 .arg(aq)
17614 .arg(ad)
17615 .arg(&mut *y0)
17616 .arg(&mut *y1)
17617 .arg(&inf)
17618 .arg(&oo0)
17619 .arg(&oo1)
17620 .arg(&r0)
17621 .arg(&r1);
17622 unsafe {
17623 b.launch(cfg)?;
17624 }
17625 Ok(true)
17626 }
17627
17628 #[allow(clippy::type_complexity)] pub fn matmul_q4_fused2_batched(
17634 &self,
17635 w0: &crate::model::GpuTensor,
17636 w1: &crate::model::GpuTensor,
17637 aq: &CudaSlice<i8>,
17638 ad: &CudaSlice<f32>,
17639 m: usize,
17640 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
17641 use crate::model::GpuTensor;
17642 if !(2..=8).contains(&m) {
17643 return Ok(None);
17644 }
17645 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
17646 match w {
17647 GpuTensor::Quant {
17648 qtype, row_bytes, ..
17649 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
17650 _ => None,
17651 }
17652 };
17653 let (Some((rb0, o0)), Some((_rb1, o1))) = (q4(w0), q4(w1)) else {
17654 return Ok(None);
17655 };
17656 if w0.in_features() != w1.in_features() {
17657 return Ok(None);
17658 }
17659 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17660 match w {
17661 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17662 Some(mr) => (mr, true),
17663 None => (bytes, *rp),
17664 },
17665 _ => unreachable!(),
17666 }
17667 }
17668 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
17669 if !rp0 || !rp1 {
17670 return Ok(None);
17671 }
17672 let mcols = Self::batched_mcols(m);
17673 let rpb: u32 = 4;
17674 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
17675 let grid = nb(o0) + nb(o1);
17676 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
17677 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
17678 let f = self.func(match mcols {
17679 2 => "qmatvec_q4_0_mmvq_b2_f2_rp",
17680 4 => "qmatvec_q4_0_mmvq_b4_f2_rp",
17681 _ => "qmatvec_q4_0_mmvq_b8_f2_rp",
17682 });
17683 let cfg = LaunchConfig {
17684 grid_dim: (grid, 1, 1),
17685 block_dim: (32, rpb, 1),
17686 shared_mem_bytes: 0,
17687 };
17688 let inf = w0.in_features() as i32;
17689 let (oo0, oo1, mi) = (o0 as i32, o1 as i32, m as i32);
17690 let rb = rb0 as i64;
17691 let __s_b = self.gpu.stream();
17692 let mut b = __s_b.launch_builder(&f);
17693 b.arg(b0)
17694 .arg(b1)
17695 .arg(aq)
17696 .arg(ad)
17697 .arg(&mut y0)
17698 .arg(&mut y1)
17699 .arg(&inf)
17700 .arg(&oo0)
17701 .arg(&oo1)
17702 .arg(&mi)
17703 .arg(&rb);
17704 unsafe {
17705 b.launch(cfg)?;
17706 }
17707 Ok(Some((y0, y1)))
17708 }
17709
17710 #[allow(clippy::too_many_arguments)]
17713 #[allow(clippy::type_complexity)] pub fn matmul_q4_fused3_batched(
17715 &self,
17716 w0: &crate::model::GpuTensor,
17717 w1: &crate::model::GpuTensor,
17718 w2: &crate::model::GpuTensor,
17719 aq: &CudaSlice<i8>,
17720 ad: &CudaSlice<f32>,
17721 m: usize,
17722 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
17723 {
17724 use crate::model::GpuTensor;
17725 if !(2..=8).contains(&m) {
17726 return Ok(None);
17727 }
17728 let q4 = |w: &GpuTensor| -> Option<usize> {
17729 match w {
17730 GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0 => Some(w.out_features()),
17731 _ => None,
17732 }
17733 };
17734 let (Some(o0), Some(o1), Some(o2)) = (q4(w0), q4(w1), q4(w2)) else {
17735 return Ok(None);
17736 };
17737 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
17738 return Ok(None);
17739 }
17740 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
17741 match w {
17742 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
17743 Some(mr) => (mr, true),
17744 None => (bytes, *rp),
17745 },
17746 _ => unreachable!(),
17747 }
17748 }
17749 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
17750 if !rp0 || !rp1 || !rp2 {
17751 return Ok(None);
17752 }
17753 let mcols = Self::batched_mcols(m);
17754 let rpb: u32 = 4;
17755 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
17756 let grid = nb(o0) + nb(o1) + nb(o2);
17757 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
17758 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
17759 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
17760 let f = self.func(match mcols {
17761 2 => "qmatvec_q4_0_mmvq_b2_f3_rp",
17762 4 => "qmatvec_q4_0_mmvq_b4_f3_rp",
17763 _ => "qmatvec_q4_0_mmvq_b8_f3_rp",
17764 });
17765 let cfg = LaunchConfig {
17766 grid_dim: (grid, 1, 1),
17767 block_dim: (32, rpb, 1),
17768 shared_mem_bytes: 0,
17769 };
17770 let inf = w0.in_features() as i32;
17771 let (oo0, oo1, oo2, mi) = (o0 as i32, o1 as i32, o2 as i32, m as i32);
17772 let rb = 0i64;
17773 let __s_b = self.gpu.stream();
17774 let mut b = __s_b.launch_builder(&f);
17775 b.arg(b0)
17776 .arg(b1)
17777 .arg(b2)
17778 .arg(aq)
17779 .arg(ad)
17780 .arg(&mut y0)
17781 .arg(&mut y1)
17782 .arg(&mut y2)
17783 .arg(&inf)
17784 .arg(&oo0)
17785 .arg(&oo1)
17786 .arg(&oo2)
17787 .arg(&mi)
17788 .arg(&rb);
17789 unsafe {
17790 b.launch(cfg)?;
17791 }
17792 Ok(Some((y0, y1, y2)))
17793 }
17794
17795 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused3(
17797 &self,
17798 w0: &crate::model::GpuTensor,
17799 w1: &crate::model::GpuTensor,
17800 w2: &crate::model::GpuTensor,
17801 aq: &CudaSlice<i8>,
17802 ad: &CudaSlice<f32>,
17803 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
17804 {
17805 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
17808 return Ok(Some(self.e4m3_fused3_core(
17809 p0.0,
17810 p1.0,
17811 p2.0,
17812 aq,
17813 ad,
17814 w0.in_features(),
17815 p0.1,
17816 p1.1,
17817 p2.1,
17818 p0.2,
17819 p0.3,
17820 p1.3,
17821 p2.3,
17822 )?));
17823 }
17824 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else {
17825 return Ok(None);
17826 };
17827 Ok(Some(self.q8_fused3_core(
17828 p0.0,
17829 p1.0,
17830 p2.0,
17831 aq,
17832 ad,
17833 w0.in_features(),
17834 p0.1,
17835 p1.1,
17836 p2.1,
17837 p0.2,
17838 )?))
17839 }
17840
17841 #[allow(clippy::too_many_arguments)]
17842 #[allow(clippy::type_complexity)] fn q8_fused3_core(
17844 &self,
17845 b0: &CudaSlice<u8>,
17846 b1: &CudaSlice<u8>,
17847 b2: &CudaSlice<u8>,
17848 aq: &CudaSlice<i8>,
17849 ad: &CudaSlice<f32>,
17850 in_f: usize,
17851 out0: usize,
17852 out1: usize,
17853 out2: usize,
17854 row_bytes: usize,
17855 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
17856 const ROWS_PER_BLOCK: u32 = 4;
17857 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
17858 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
17859 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
17860 let f = self.func("qmatvec_q8_0_mmvq_fused3");
17861 let mut y0 = self.alloc_uninit::<f32>(out0)?;
17862 let mut y1 = self.alloc_uninit::<f32>(out1)?;
17863 let mut y2 = self.alloc_uninit::<f32>(out2)?;
17864 let cfg = LaunchConfig {
17865 grid_dim: (nb0 + nb1 + nb2, 1, 1),
17866 block_dim: (32, ROWS_PER_BLOCK, 1),
17867 shared_mem_bytes: 0,
17868 };
17869 let (inf, o0, o1, o2, rbl) = (
17870 in_f as i32,
17871 out0 as i32,
17872 out1 as i32,
17873 out2 as i32,
17874 row_bytes as i64,
17875 );
17876 let __s_b = self.gpu.stream();
17877 let mut b = __s_b.launch_builder(&f);
17878 b.arg(b0)
17879 .arg(b1)
17880 .arg(b2)
17881 .arg(aq)
17882 .arg(ad)
17883 .arg(&mut y0)
17884 .arg(&mut y1)
17885 .arg(&mut y2)
17886 .arg(&inf)
17887 .arg(&o0)
17888 .arg(&o1)
17889 .arg(&o2)
17890 .arg(&rbl);
17891 unsafe {
17892 b.launch(cfg)?;
17893 }
17894 Ok((y0, y1, y2))
17895 }
17896
17897 #[allow(clippy::too_many_arguments)]
17899 #[allow(clippy::type_complexity)] pub fn qmatvec_q8_fused3_raw(
17901 &self,
17902 b0: &CudaSlice<u8>,
17903 b1: &CudaSlice<u8>,
17904 b2: &CudaSlice<u8>,
17905 x: &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 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
17913 self.q8_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes)
17914 }
17915
17916 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused2_t(
17928 &self,
17929 w0: &crate::model::GpuTensor,
17930 w1: &crate::model::GpuTensor,
17931 aq: &CudaSlice<i8>,
17932 ad: &CudaSlice<f32>,
17933 m: usize,
17934 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
17935 if !(2..=8).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() {
17939 return Ok(None);
17940 }
17941 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
17944 if m > 4 && !Self::b8_enabled() {
17945 return Ok(None);
17946 }
17947 return Ok(Some(self.e4m3_fused2_t_core(
17948 p0.0,
17949 p1.0,
17950 aq,
17951 ad,
17952 m,
17953 w0.in_features(),
17954 p0.1,
17955 p1.1,
17956 p0.2,
17957 p0.3,
17958 p1.3,
17959 )?));
17960 }
17961 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
17962 return Ok(None);
17963 };
17964 Ok(Some(self.q8_fused2_t_core(
17965 p0.0,
17966 p1.0,
17967 aq,
17968 ad,
17969 m,
17970 w0.in_features(),
17971 p0.1,
17972 p1.1,
17973 p0.2,
17974 )?))
17975 }
17976
17977 #[allow(clippy::too_many_arguments)]
17978 fn q8_fused2_t_core(
17979 &self,
17980 b0: &CudaSlice<u8>,
17981 b1: &CudaSlice<u8>,
17982 aq: &CudaSlice<i8>,
17983 ad: &CudaSlice<f32>,
17984 m: usize,
17985 in_f: usize,
17986 out0: usize,
17987 out1: usize,
17988 row_bytes: usize,
17989 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
17990 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
17992 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
17993 let f = self.func(match Self::batched_mcols(m) {
17994 2 => "qmatvec_q8_0_mmvq_fused2_b2",
17995 4 => "qmatvec_q8_0_mmvq_fused2_b4",
17996 _ => "qmatvec_q8_0_mmvq_fused2_b8",
17998 });
17999 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
18000 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
18001 let cfg = LaunchConfig {
18002 grid_dim: (nb0 + nb1, 1, 1),
18003 block_dim: (32, ROWS_PER_BLOCK, 1),
18004 shared_mem_bytes: 0,
18005 };
18006 let (inf, o0, o1, mi, rbl) = (
18007 in_f as i32,
18008 out0 as i32,
18009 out1 as i32,
18010 m as i32,
18011 row_bytes as i64,
18012 );
18013 let __s_b = self.gpu.stream();
18014 let mut b = __s_b.launch_builder(&f);
18015 b.arg(b0)
18016 .arg(b1)
18017 .arg(aq)
18018 .arg(ad)
18019 .arg(&mut y0)
18020 .arg(&mut y1)
18021 .arg(&inf)
18022 .arg(&o0)
18023 .arg(&o1)
18024 .arg(&mi)
18025 .arg(&rbl);
18026 unsafe {
18027 b.launch(cfg)?;
18028 }
18029 Ok((y0, y1))
18030 }
18031
18032 #[allow(clippy::too_many_arguments)]
18035 pub fn qmatvec_q8_fused2_t_raw(
18036 &self,
18037 b0: &CudaSlice<u8>,
18038 b1: &CudaSlice<u8>,
18039 x: &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 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18047 self.q8_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes)
18048 }
18049
18050 #[allow(clippy::too_many_arguments)]
18053 #[allow(clippy::type_complexity)] pub fn matmul_q8_fused3_t(
18055 &self,
18056 w0: &crate::model::GpuTensor,
18057 w1: &crate::model::GpuTensor,
18058 w2: &crate::model::GpuTensor,
18059 aq: &CudaSlice<i8>,
18060 ad: &CudaSlice<f32>,
18061 m: usize,
18062 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
18063 {
18064 if !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() {
18065 return Ok(None);
18066 }
18067 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
18068 return Ok(Some(self.e4m3_fused3_t_core(
18069 p0.0,
18070 p1.0,
18071 p2.0,
18072 aq,
18073 ad,
18074 m,
18075 w0.in_features(),
18076 p0.1,
18077 p1.1,
18078 p2.1,
18079 p0.2,
18080 p0.3,
18081 p1.3,
18082 p2.3,
18083 )?));
18084 }
18085 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else {
18086 return Ok(None);
18087 };
18088 Ok(Some(self.q8_fused3_t_core(
18089 p0.0,
18090 p1.0,
18091 p2.0,
18092 aq,
18093 ad,
18094 m,
18095 w0.in_features(),
18096 p0.1,
18097 p1.1,
18098 p2.1,
18099 p0.2,
18100 )?))
18101 }
18102
18103 #[allow(clippy::too_many_arguments)]
18104 #[allow(clippy::type_complexity)] fn q8_fused3_t_core(
18106 &self,
18107 b0: &CudaSlice<u8>,
18108 b1: &CudaSlice<u8>,
18109 b2: &CudaSlice<u8>,
18110 aq: &CudaSlice<i8>,
18111 ad: &CudaSlice<f32>,
18112 m: usize,
18113 in_f: usize,
18114 out0: usize,
18115 out1: usize,
18116 out2: usize,
18117 row_bytes: usize,
18118 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18119 const ROWS_PER_BLOCK: u32 = 4;
18120 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18121 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18122 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
18123 let f = self.func(if Self::batched_mcols(m) == 2 {
18124 "qmatvec_q8_0_mmvq_fused3_b2"
18125 } else {
18126 "qmatvec_q8_0_mmvq_fused3_b4"
18127 });
18128 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
18129 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
18130 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
18131 let cfg = LaunchConfig {
18132 grid_dim: (nb0 + nb1 + nb2, 1, 1),
18133 block_dim: (32, ROWS_PER_BLOCK, 1),
18134 shared_mem_bytes: 0,
18135 };
18136 let (inf, o0, o1, o2, mi, rbl) = (
18137 in_f as i32,
18138 out0 as i32,
18139 out1 as i32,
18140 out2 as i32,
18141 m as i32,
18142 row_bytes as i64,
18143 );
18144 let __s_b = self.gpu.stream();
18145 let mut b = __s_b.launch_builder(&f);
18146 b.arg(b0)
18147 .arg(b1)
18148 .arg(b2)
18149 .arg(aq)
18150 .arg(ad)
18151 .arg(&mut y0)
18152 .arg(&mut y1)
18153 .arg(&mut y2)
18154 .arg(&inf)
18155 .arg(&o0)
18156 .arg(&o1)
18157 .arg(&o2)
18158 .arg(&mi)
18159 .arg(&rbl);
18160 unsafe {
18161 b.launch(cfg)?;
18162 }
18163 Ok((y0, y1, y2))
18164 }
18165
18166 #[allow(clippy::too_many_arguments)]
18168 #[allow(clippy::type_complexity)] pub fn qmatvec_q8_fused3_t_raw(
18170 &self,
18171 b0: &CudaSlice<u8>,
18172 b1: &CudaSlice<u8>,
18173 b2: &CudaSlice<u8>,
18174 x: &CudaSlice<f32>,
18175 m: usize,
18176 in_f: usize,
18177 out0: usize,
18178 out1: usize,
18179 out2: usize,
18180 row_bytes: usize,
18181 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18182 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18183 self.q8_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes)
18184 }
18185
18186 pub fn q8_ffn_fuse2_on(&self) -> bool {
18190 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18191 *ON.get_or_init(|| std::env::var("MEMRA_Q8_FFN_FUSE2").as_deref() != Ok("0"))
18192 }
18193
18194 #[allow(clippy::type_complexity)]
18200 fn q8_fused_params<'w, const N: usize>(
18201 &self,
18202 ws: &[&'w crate::model::GpuTensor; N],
18203 ) -> Option<[(&'w CudaSlice<u8>, usize, usize); N]> {
18204 use crate::model::GpuTensor;
18205 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") {
18206 return None;
18207 }
18208 if std::env::var("MEMRA_Q8_DUAL").is_ok_and(|v| v == "0") {
18209 return None;
18210 }
18211 let in_f = ws[0].in_features();
18212 let mut out: [Option<(&CudaSlice<u8>, usize, usize)>; N] = [None; N];
18213 for (i, w) in ws.iter().enumerate() {
18214 match w {
18215 GpuTensor::Quant {
18216 bytes,
18217 qtype,
18218 row_bytes,
18219 scale,
18220 ..
18221 } if *qtype == QT_Q8_0 && *scale == 1.0 && w.in_features() == in_f => {
18222 out[i] = Some((bytes, w.out_features(), *row_bytes))
18223 }
18224 _ => return None,
18225 }
18226 }
18227 Some(out.map(|o| o.unwrap()))
18228 }
18229
18230 pub fn e4m3_dual_on(&self) -> bool {
18233 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18234 *ON.get_or_init(|| std::env::var("MEMRA_E4M3_DUAL").as_deref() != Ok("0"))
18235 }
18236
18237 #[allow(clippy::type_complexity)]
18249 fn e4m3_fused_params<'w, const N: usize>(
18250 &self,
18251 ws: &[&'w crate::model::GpuTensor; N],
18252 ) -> Option<[(&'w CudaSlice<u8>, usize, usize, f32); N]> {
18253 use crate::model::GpuTensor;
18254 if !self.e4m3_dual_on() {
18255 return None;
18256 }
18257 let in_f = ws[0].in_features();
18258 let mut out: [Option<(&CudaSlice<u8>, usize, usize, f32)>; N] = [None; N];
18259 for (i, w) in ws.iter().enumerate() {
18260 match w {
18261 GpuTensor::Quant {
18262 bytes,
18263 qtype,
18264 row_bytes,
18265 scale,
18266 rp,
18267 rp4,
18268 ..
18269 } if *qtype == QT_F8_E4M3
18270 && w.in_features() == in_f
18271 && *row_bytes == in_f
18272 && !*rp
18273 && rp4.is_none() =>
18274 {
18275 out[i] = Some((bytes, w.out_features(), *row_bytes, *scale))
18276 }
18277 _ => return None,
18278 }
18279 }
18280 Some(out.map(|o| o.unwrap()))
18281 }
18282
18283 #[allow(clippy::too_many_arguments)]
18287 fn e4m3_fused2_core(
18288 &self,
18289 b0: &CudaSlice<u8>,
18290 b1: &CudaSlice<u8>,
18291 aq: &CudaSlice<i8>,
18292 ad: &CudaSlice<f32>,
18293 in_f: usize,
18294 out0: usize,
18295 out1: usize,
18296 row_bytes: usize,
18297 ws0: f32,
18298 ws1: f32,
18299 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18300 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18302 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18303 let f = self.func("qmatvec_e4m3_mmvq_fused2");
18304 let mut y0 = self.alloc_uninit::<f32>(out0)?;
18305 let mut y1 = self.alloc_uninit::<f32>(out1)?;
18306 let cfg = LaunchConfig {
18307 grid_dim: (nb0 + nb1, 1, 1),
18308 block_dim: (32, ROWS_PER_BLOCK, 1),
18309 shared_mem_bytes: 0,
18310 };
18311 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
18312 let __s_b = self.gpu.stream();
18313 let mut b = __s_b.launch_builder(&f);
18314 b.arg(b0)
18315 .arg(b1)
18316 .arg(aq)
18317 .arg(ad)
18318 .arg(&mut y0)
18319 .arg(&mut y1)
18320 .arg(&inf)
18321 .arg(&o0)
18322 .arg(&o1)
18323 .arg(&rbl)
18324 .arg(&ws0)
18325 .arg(&ws1);
18326 unsafe {
18327 b.launch(cfg)?;
18328 }
18329 Ok((y0, y1))
18330 }
18331
18332 #[allow(clippy::too_many_arguments)]
18334 #[allow(clippy::type_complexity)] fn e4m3_fused3_core(
18336 &self,
18337 b0: &CudaSlice<u8>,
18338 b1: &CudaSlice<u8>,
18339 b2: &CudaSlice<u8>,
18340 aq: &CudaSlice<i8>,
18341 ad: &CudaSlice<f32>,
18342 in_f: usize,
18343 out0: usize,
18344 out1: usize,
18345 out2: usize,
18346 row_bytes: usize,
18347 ws0: f32,
18348 ws1: f32,
18349 ws2: f32,
18350 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18351 const ROWS_PER_BLOCK: u32 = 4;
18352 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18353 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18354 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
18355 let f = self.func("qmatvec_e4m3_mmvq_fused3");
18356 let mut y0 = self.alloc_uninit::<f32>(out0)?;
18357 let mut y1 = self.alloc_uninit::<f32>(out1)?;
18358 let mut y2 = self.alloc_uninit::<f32>(out2)?;
18359 let cfg = LaunchConfig {
18360 grid_dim: (nb0 + nb1 + nb2, 1, 1),
18361 block_dim: (32, ROWS_PER_BLOCK, 1),
18362 shared_mem_bytes: 0,
18363 };
18364 let (inf, o0, o1, o2, rbl) = (
18365 in_f as i32,
18366 out0 as i32,
18367 out1 as i32,
18368 out2 as i32,
18369 row_bytes as i64,
18370 );
18371 let __s_b = self.gpu.stream();
18372 let mut b = __s_b.launch_builder(&f);
18373 b.arg(b0)
18374 .arg(b1)
18375 .arg(b2)
18376 .arg(aq)
18377 .arg(ad)
18378 .arg(&mut y0)
18379 .arg(&mut y1)
18380 .arg(&mut y2)
18381 .arg(&inf)
18382 .arg(&o0)
18383 .arg(&o1)
18384 .arg(&o2)
18385 .arg(&rbl)
18386 .arg(&ws0)
18387 .arg(&ws1)
18388 .arg(&ws2);
18389 unsafe {
18390 b.launch(cfg)?;
18391 }
18392 Ok((y0, y1, y2))
18393 }
18394
18395 #[allow(clippy::too_many_arguments)]
18399 fn e4m3_fused2_t_core(
18400 &self,
18401 b0: &CudaSlice<u8>,
18402 b1: &CudaSlice<u8>,
18403 aq: &CudaSlice<i8>,
18404 ad: &CudaSlice<f32>,
18405 m: usize,
18406 in_f: usize,
18407 out0: usize,
18408 out1: usize,
18409 row_bytes: usize,
18410 ws0: f32,
18411 ws1: f32,
18412 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18413 const ROWS_PER_BLOCK: u32 = 4;
18414 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18415 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18416 let f = self.func(match Self::batched_mcols(m) {
18417 2 => "qmatvec_e4m3_mmvq_fused2_b2",
18418 4 => "qmatvec_e4m3_mmvq_fused2_b4",
18419 _ => "qmatvec_e4m3_mmvq_fused2_b8",
18420 });
18421 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
18422 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
18423 let cfg = LaunchConfig {
18424 grid_dim: (nb0 + nb1, 1, 1),
18425 block_dim: (32, ROWS_PER_BLOCK, 1),
18426 shared_mem_bytes: 0,
18427 };
18428 let (inf, o0, o1, mi, rbl) = (
18429 in_f as i32,
18430 out0 as i32,
18431 out1 as i32,
18432 m as i32,
18433 row_bytes as i64,
18434 );
18435 let __s_b = self.gpu.stream();
18436 let mut b = __s_b.launch_builder(&f);
18437 b.arg(b0)
18438 .arg(b1)
18439 .arg(aq)
18440 .arg(ad)
18441 .arg(&mut y0)
18442 .arg(&mut y1)
18443 .arg(&inf)
18444 .arg(&o0)
18445 .arg(&o1)
18446 .arg(&mi)
18447 .arg(&rbl);
18448 unsafe {
18449 b.launch(cfg)?;
18450 }
18451 if ws0 != 1.0 {
18452 self.scale_inplace(&mut y0, ws0, m * out0)?;
18453 }
18454 if ws1 != 1.0 {
18455 self.scale_inplace(&mut y1, ws1, m * out1)?;
18456 }
18457 Ok((y0, y1))
18458 }
18459
18460 #[allow(clippy::too_many_arguments)]
18462 #[allow(clippy::type_complexity)] fn e4m3_fused3_t_core(
18464 &self,
18465 b0: &CudaSlice<u8>,
18466 b1: &CudaSlice<u8>,
18467 b2: &CudaSlice<u8>,
18468 aq: &CudaSlice<i8>,
18469 ad: &CudaSlice<f32>,
18470 m: usize,
18471 in_f: usize,
18472 out0: usize,
18473 out1: usize,
18474 out2: usize,
18475 row_bytes: usize,
18476 ws0: f32,
18477 ws1: f32,
18478 ws2: f32,
18479 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18480 const ROWS_PER_BLOCK: u32 = 4;
18481 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
18482 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
18483 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
18484 let f = self.func(if Self::batched_mcols(m) == 2 {
18485 "qmatvec_e4m3_mmvq_fused3_b2"
18486 } else {
18487 "qmatvec_e4m3_mmvq_fused3_b4"
18488 });
18489 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
18490 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
18491 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
18492 let cfg = LaunchConfig {
18493 grid_dim: (nb0 + nb1 + nb2, 1, 1),
18494 block_dim: (32, ROWS_PER_BLOCK, 1),
18495 shared_mem_bytes: 0,
18496 };
18497 let (inf, o0, o1, o2, mi, rbl) = (
18498 in_f as i32,
18499 out0 as i32,
18500 out1 as i32,
18501 out2 as i32,
18502 m as i32,
18503 row_bytes as i64,
18504 );
18505 let __s_b = self.gpu.stream();
18506 let mut b = __s_b.launch_builder(&f);
18507 b.arg(b0)
18508 .arg(b1)
18509 .arg(b2)
18510 .arg(aq)
18511 .arg(ad)
18512 .arg(&mut y0)
18513 .arg(&mut y1)
18514 .arg(&mut y2)
18515 .arg(&inf)
18516 .arg(&o0)
18517 .arg(&o1)
18518 .arg(&o2)
18519 .arg(&mi)
18520 .arg(&rbl);
18521 unsafe {
18522 b.launch(cfg)?;
18523 }
18524 if ws0 != 1.0 {
18525 self.scale_inplace(&mut y0, ws0, m * out0)?;
18526 }
18527 if ws1 != 1.0 {
18528 self.scale_inplace(&mut y1, ws1, m * out1)?;
18529 }
18530 if ws2 != 1.0 {
18531 self.scale_inplace(&mut y2, ws2, m * out2)?;
18532 }
18533 Ok((y0, y1, y2))
18534 }
18535
18536 #[allow(clippy::too_many_arguments)] pub fn qmatvec_e4m3_blk_mmvq(
18547 &self,
18548 bytes: &CudaSlice<u8>,
18549 aq: &CudaSlice<i8>,
18550 ad: &CudaSlice<f32>,
18551 scales: &CudaSlice<f32>,
18552 m: usize,
18553 in_f: usize,
18554 out_f: usize,
18555 row_bytes: usize,
18556 scale_cols: usize,
18557 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18558 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_e4m3_blk_mmvq_into(
18560 bytes, aq, ad, scales, m, in_f, out_f, row_bytes, scale_cols, &mut y,
18561 )?;
18562 Ok(y)
18563 }
18564
18565 #[allow(clippy::too_many_arguments)]
18567 pub fn qmatvec_e4m3_blk_mmvq_into(
18568 &self,
18569 bytes: &CudaSlice<u8>,
18570 aq: &CudaSlice<i8>,
18571 ad: &CudaSlice<f32>,
18572 scales: &CudaSlice<f32>,
18573 m: usize,
18574 in_f: usize,
18575 out_f: usize,
18576 row_bytes: usize,
18577 scale_cols: usize,
18578 y: &mut CudaSlice<f32>,
18579 ) -> Result<(), Box<dyn std::error::Error>> {
18580 const ROWS_PER_BLOCK: u32 = 4; let f = self.func("qmatvec_e4m3_blk_mmvq");
18582 let cfg = LaunchConfig {
18583 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), m as u32, 1),
18584 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
18587 let (inf, outf, mi, rb, sc) = (
18588 in_f as i32,
18589 out_f as i32,
18590 m as i32,
18591 row_bytes as i64,
18592 scale_cols as i32,
18593 );
18594 let __s_b = self.gpu.stream();
18595 let mut b = __s_b.launch_builder(&f);
18596 b.arg(bytes)
18597 .arg(aq)
18598 .arg(ad)
18599 .arg(scales)
18600 .arg(&mut *y)
18601 .arg(&inf)
18602 .arg(&outf)
18603 .arg(&mi)
18604 .arg(&rb)
18605 .arg(&sc);
18606 unsafe {
18607 b.launch(cfg)?;
18608 }
18609 Ok(())
18610 }
18611
18612 #[allow(clippy::too_many_arguments)]
18618 pub fn qmatvec_e4m3_blk_mmvq_batched(
18619 &self,
18620 bytes: &CudaSlice<u8>,
18621 aq: &CudaSlice<i8>,
18622 ad: &CudaSlice<f32>,
18623 scales: &CudaSlice<f32>,
18624 m: usize,
18625 in_f: usize,
18626 out_f: usize,
18627 row_bytes: usize,
18628 scale_cols: usize,
18629 mcols: usize,
18630 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18631 const ROWS_PER_BLOCK: u32 = 4; debug_assert!(mcols >= m, "blk batched: mcols {mcols} < m {m}");
18633 let name = match mcols {
18634 2 => "qmatvec_e4m3_blk_mmvq_b2",
18635 4 => "qmatvec_e4m3_blk_mmvq_b4",
18636 8 => "qmatvec_e4m3_blk_mmvq_b8",
18637 16 => "qmatvec_e4m3_blk_mmvq_b16",
18638 _ => {
18639 return Err(
18640 format!("qmatvec_e4m3_blk_mmvq_batched: no kernel for mcols {mcols}").into(),
18641 );
18642 }
18643 };
18644 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
18645 let f = self.func(name);
18646 let cfg = LaunchConfig {
18647 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
18648 block_dim: (32, ROWS_PER_BLOCK, 1),
18649 shared_mem_bytes: 0,
18650 };
18651 let (inf, outf, mi, rb, sc) = (
18652 in_f as i32,
18653 out_f as i32,
18654 m as i32,
18655 row_bytes as i64,
18656 scale_cols as i32,
18657 );
18658 let __s_b = self.gpu.stream();
18659 let mut b = __s_b.launch_builder(&f);
18660 b.arg(bytes)
18661 .arg(aq)
18662 .arg(ad)
18663 .arg(scales)
18664 .arg(&mut y)
18665 .arg(&inf)
18666 .arg(&outf)
18667 .arg(&mi)
18668 .arg(&rb)
18669 .arg(&sc);
18670 unsafe {
18671 b.launch(cfg)?;
18672 }
18673 Ok(y)
18674 }
18675
18676 #[allow(clippy::too_many_arguments)]
18679 pub fn qmatvec_e4m3_blk_batched_raw(
18680 &self,
18681 bytes: &CudaSlice<u8>,
18682 x: &CudaSlice<f32>,
18683 scales: &CudaSlice<f32>,
18684 m: usize,
18685 in_f: usize,
18686 out_f: usize,
18687 row_bytes: usize,
18688 scale_cols: usize,
18689 mcols: usize,
18690 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18691 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18692 self.qmatvec_e4m3_blk_mmvq_batched(
18693 bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols, mcols,
18694 )
18695 }
18696
18697 #[allow(clippy::too_many_arguments)]
18700 pub fn qmatvec_e4m3_blk_mmvq_raw(
18701 &self,
18702 bytes: &CudaSlice<u8>,
18703 x: &CudaSlice<f32>,
18704 scales: &CudaSlice<f32>,
18705 m: usize,
18706 in_f: usize,
18707 out_f: usize,
18708 row_bytes: usize,
18709 scale_cols: usize,
18710 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18711 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18712 self.qmatvec_e4m3_blk_mmvq(
18713 bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols,
18714 )
18715 }
18716
18717 #[allow(clippy::too_many_arguments)]
18720 pub fn qmatvec_e4m3_fused2_raw(
18721 &self,
18722 b0: &CudaSlice<u8>,
18723 b1: &CudaSlice<u8>,
18724 x: &CudaSlice<f32>,
18725 in_f: usize,
18726 out0: usize,
18727 out1: usize,
18728 row_bytes: usize,
18729 ws0: f32,
18730 ws1: f32,
18731 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18732 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
18733 self.e4m3_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes, ws0, ws1)
18734 }
18735
18736 #[allow(clippy::too_many_arguments)]
18737 #[allow(clippy::type_complexity)] pub fn qmatvec_e4m3_fused3_raw(
18739 &self,
18740 b0: &CudaSlice<u8>,
18741 b1: &CudaSlice<u8>,
18742 b2: &CudaSlice<u8>,
18743 x: &CudaSlice<f32>,
18744 in_f: usize,
18745 out0: usize,
18746 out1: usize,
18747 out2: usize,
18748 row_bytes: usize,
18749 ws0: f32,
18750 ws1: f32,
18751 ws2: f32,
18752 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18753 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
18754 self.e4m3_fused3_core(
18755 b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes, ws0, ws1, ws2,
18756 )
18757 }
18758
18759 #[allow(clippy::too_many_arguments)]
18760 pub fn qmatvec_e4m3_fused2_t_raw(
18761 &self,
18762 b0: &CudaSlice<u8>,
18763 b1: &CudaSlice<u8>,
18764 x: &CudaSlice<f32>,
18765 m: usize,
18766 in_f: usize,
18767 out0: usize,
18768 out1: usize,
18769 row_bytes: usize,
18770 ws0: f32,
18771 ws1: f32,
18772 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18773 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18774 self.e4m3_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes, ws0, ws1)
18775 }
18776
18777 #[allow(clippy::too_many_arguments)]
18778 #[allow(clippy::type_complexity)] pub fn qmatvec_e4m3_fused3_t_raw(
18780 &self,
18781 b0: &CudaSlice<u8>,
18782 b1: &CudaSlice<u8>,
18783 b2: &CudaSlice<u8>,
18784 x: &CudaSlice<f32>,
18785 m: usize,
18786 in_f: usize,
18787 out0: usize,
18788 out1: usize,
18789 out2: usize,
18790 row_bytes: usize,
18791 ws0: f32,
18792 ws1: f32,
18793 ws2: f32,
18794 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18795 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
18796 self.e4m3_fused3_t_core(
18797 b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes, ws0, ws1, ws2,
18798 )
18799 }
18800
18801 fn try_e4m3_blk_pre(
18812 &self,
18813 w: &crate::model::GpuTensor,
18814 aq: &CudaSlice<i8>,
18815 ad: &CudaSlice<f32>,
18816 m: usize,
18817 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
18818 use crate::model::GpuTensor;
18819 if let GpuTensor::Quant {
18820 bytes,
18821 qtype,
18822 row_bytes,
18823 blk: Some(g),
18824 ..
18825 } = w
18826 && *qtype == QT_F8_E4M3_BLK
18827 {
18828 if (2..=16).contains(&m)
18834 && std::env::var("MEMRA_NO_BATCHED").is_err()
18835 && (m <= 4 || Self::b8_enabled())
18836 {
18837 let mcols = Self::batched_mcols(m);
18838 return Ok(Some(self.qmatvec_e4m3_blk_mmvq_batched(
18839 bytes,
18840 aq,
18841 ad,
18842 &g.scales,
18843 m,
18844 w.in_features(),
18845 w.out_features(),
18846 *row_bytes,
18847 g.cols,
18848 mcols,
18849 )?));
18850 }
18851 return Ok(Some(self.qmatvec_e4m3_blk_mmvq(
18852 bytes,
18853 aq,
18854 ad,
18855 &g.scales,
18856 m,
18857 w.in_features(),
18858 w.out_features(),
18859 *row_bytes,
18860 g.cols,
18861 )?));
18862 }
18863 Ok(None)
18864 }
18865
18866 fn try_e4m3_blk_prefill(
18913 &self,
18914 w: &crate::model::GpuTensor,
18915 x: &CudaSlice<f32>,
18916 m: usize,
18917 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
18918 use crate::model::GpuTensor;
18919 let GpuTensor::Quant {
18920 bytes,
18921 qtype,
18922 blk: Some(g),
18923 ..
18924 } = w
18925 else {
18926 return Ok(None);
18927 };
18928 if *qtype != QT_F8_E4M3_BLK {
18929 return Ok(None);
18930 }
18931 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? {
18936 return Ok(Some(y));
18937 }
18938 let (in_f, out_f) = (w.in_features(), w.out_features());
18939 let slab = self.fp8_blk_dequant_q8_0_dev(bytes, &g.scales, out_f, in_f)?;
18940 let tmp = GpuTensor::Quant {
18941 bytes: slab,
18942 qtype: QT_Q8_0,
18943 row_bytes: in_f / 32 * 34,
18944 ne: vec![in_f as u64, out_f as u64],
18945 scale: 1.0,
18946 rp: false,
18947 #[cfg(memra_cutlass)]
18948 cutlass: None,
18949 fp8: None,
18950 blk: None,
18951 f16: None,
18952 rp4: None,
18953 };
18954 Ok(Some(self.matmul(&tmp, x, m)?))
18956 }
18957
18958 #[allow(clippy::type_complexity)] pub fn matmul_pre_noscale(
18960 &self,
18961 w: &crate::model::GpuTensor,
18962 aq: &CudaSlice<i8>,
18963 ad: &CudaSlice<f32>,
18964 m: usize,
18965 ) -> Result<Option<(CudaSlice<f32>, f32)>, Box<dyn std::error::Error>> {
18966 use crate::model::GpuTensor;
18967 if m == 1
18971 && let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)?
18972 {
18973 return Ok(Some((y, 1.0)));
18974 }
18975 if m != 1 || !self.uses_q8_1_fast(w) {
18977 return Ok(None);
18978 }
18979 let in_f = w.in_features();
18980 let out_f = w.out_features();
18981 let (bytes, qtype, row_bytes, scale, rp) = match w {
18982 GpuTensor::Quant {
18983 bytes,
18984 qtype,
18985 row_bytes,
18986 scale,
18987 rp,
18988 ..
18989 } => (bytes, *qtype, *row_bytes, *scale, *rp),
18990 _ => return Ok(None),
18991 };
18992 if self.mmvq_supports(qtype) {
18994 let (mbytes, mrp) = match w {
18996 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
18997 _ => (bytes, rp),
18998 };
18999 let y = self.qmatvec_mmvq(
19000 mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, 1.0, mrp,
19001 )?;
19002 return Ok(Some((y, scale)));
19003 }
19004 let name = match qtype {
19006 QT_Q8_0 => "qmatvec_q8_0_dp4a",
19007 QT_Q4_K => "qmatvec_q4_K_dp4a",
19008 QT_Q6_K => "qmatvec_q6_K_dp4a",
19009 QT_Q5_K => "qmatvec_q5_K_dp4a",
19010 QT_Q3_K => "qmatvec_q3_K_dp4a",
19011 QT_NVFP4 => {
19012 if rp {
19013 "qmatvec_nvfp4_dp4a_rp"
19014 } else {
19015 "qmatvec_nvfp4_dp4a"
19016 }
19017 }
19018 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
19019 _ => return Ok(None),
19020 };
19021 let f = self.func(name);
19022 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
19023 let cfg = LaunchConfig {
19024 grid_dim: (out_f as u32, m as u32, 1),
19025 block_dim: (128, 1, 1),
19026 shared_mem_bytes: 0,
19027 };
19028 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
19029 let __s_b = self.gpu.stream();
19030 let mut b = __s_b.launch_builder(&f);
19031 b.arg(bytes)
19032 .arg(aq)
19033 .arg(ad)
19034 .arg(&mut y)
19035 .arg(&inf)
19036 .arg(&outf)
19037 .arg(&mi)
19038 .arg(&rb);
19039 unsafe {
19040 b.launch(cfg)?;
19041 }
19042 Ok(Some((y, scale)))
19043 }
19044
19045 pub fn mmvq_supports(&self, qtype: i32) -> bool {
19048 if qtype == QT_F8_E4M3 {
19053 return true;
19054 }
19055 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") {
19056 return false;
19057 }
19058 matches!(
19059 qtype,
19060 QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_Q4_0
19061 )
19062 }
19063
19064 #[allow(clippy::too_many_arguments)] pub fn qmatvec_mmvq(
19070 &self,
19071 bytes: &CudaSlice<u8>,
19072 aq: &CudaSlice<i8>,
19073 ad: &CudaSlice<f32>,
19074 m: usize,
19075 in_f: usize,
19076 out_f: usize,
19077 qtype: i32,
19078 row_bytes: usize,
19079 scale: f32,
19080 rp: bool,
19081 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19082 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_mmvq_into(
19084 bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp, &mut y,
19085 )?;
19086 Ok(y)
19087 }
19088
19089 #[allow(clippy::too_many_arguments)]
19091 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_mmvq_into(
19093 &self,
19094 bytes: &CudaSlice<u8>,
19095 aq: &CudaSlice<i8>,
19096 ad: &CudaSlice<f32>,
19097 m: usize,
19098 in_f: usize,
19099 out_f: usize,
19100 qtype: i32,
19101 row_bytes: usize,
19102 scale: f32,
19103 rp: bool,
19104 y: &mut CudaSlice<f32>,
19105 ) -> Result<(), Box<dyn std::error::Error>> {
19106 debug_assert!(y.len() >= m * out_f);
19107 const ROWS_PER_BLOCK: u32 = 4; if qtype == QT_Q8_0
19113 && rp
19114 && m == 1
19115 && out_f >= 64
19116 && (out_f as u32).div_ceil(ROWS_PER_BLOCK) < 4 * self.sm_count() as u32
19117 && {
19118 static G2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19119 *G2.get_or_init(|| std::env::var("MEMRA_Q80_G2").as_deref() != Ok("0"))
19120 }
19121 {
19122 let f = self.func("qmatvec_q8_0_mmvq_rp_g2");
19123 let cfg = LaunchConfig {
19124 grid_dim: ((out_f as u32).div_ceil(2), 1, 1),
19125 block_dim: (32, 2, 1),
19126 shared_mem_bytes: 0,
19127 };
19128 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, row_bytes as i64);
19129 let __s_b = self.gpu.stream();
19130 let mut b = __s_b.launch_builder(&f);
19131 b.arg(bytes)
19132 .arg(aq)
19133 .arg(ad)
19134 .arg(&mut *y)
19135 .arg(&inf)
19136 .arg(&outf)
19137 .arg(&mi)
19138 .arg(&rb);
19139 unsafe {
19140 b.launch(cfg)?;
19141 }
19142 if scale != 1.0 {
19143 self.scale_inplace(y, scale, out_f)?;
19144 }
19145 return Ok(());
19146 }
19147 let mut mr: u32 = if m == 1 && (qtype == QT_NVFP4 || qtype == QT_Q5_K) {
19156 2
19157 } else {
19158 1
19159 };
19160 if m == 1 && qtype == QT_Q4_0 {
19165 static Q40MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
19166 mr = *Q40MR.get_or_init(|| {
19169 std::env::var("MEMRA_Q40_MR")
19170 .ok()
19171 .and_then(|v| v.parse().ok())
19172 .unwrap_or(1)
19173 });
19174 }
19175 let q5_mode = std::env::var("MEMRA_Q5K_ISSUE").ok();
19186 let q5_force = q5_mode.as_deref() == Some("2");
19187 let q5_il = qtype == QT_Q5_K
19190 && m == 1
19191 && (q5_force || q5_mode.as_deref().map(|v| v != "0").unwrap_or(true));
19192 if q5_il && !q5_force && out_f > 65536 {
19193 mr = 1;
19194 }
19195 if qtype == QT_Q4_0 && rp && mr != 1 {
19198 mr = 2;
19199 }
19200 if qtype == QT_Q8_0 && rp {
19204 static Q80MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
19205 mr = *Q80MR.get_or_init(|| {
19206 std::env::var("MEMRA_Q80_MR")
19207 .ok()
19208 .and_then(|v| v.parse().ok())
19209 .unwrap_or(1)
19210 });
19211 }
19212 let name = match (qtype, mr, rp) {
19213 (QT_NVFP4, 2, false) => "qmatvec_nvfp4_mmvq_mr2",
19214 (QT_NVFP4, 2, true) => "qmatvec_nvfp4_mmvq_mr2_rp",
19215 (QT_NVFP4, _, true) => "qmatvec_nvfp4_mmvq_rp",
19216 (QT_Q4_0, 1, true) => "qmatvec_q4_0_mmvq_rp",
19217 (QT_Q4_0, _, true) => "qmatvec_q4_0_mmvq_mr2_rp",
19218 (QT_Q5_K, 2, _) => {
19219 if q5_il {
19220 "qmatvec_q5_K_mmvq_mr2_il"
19221 } else {
19222 "qmatvec_q5_K_mmvq_mr2"
19223 }
19224 }
19225 (QT_Q8_0, 2, true) => "qmatvec_q8_0_mmvq_mr2_rp",
19226 (QT_Q8_0, _, true)
19231 if in_f.is_multiple_of(1024) && {
19232 static CA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19233 *CA.get_or_init(|| std::env::var("MEMRA_Q80_CA").as_deref() == Ok("1"))
19234 } =>
19235 {
19236 "qmatvec_q8_0_mmvq_rpca"
19237 }
19238 (QT_Q8_0, _, true) => "qmatvec_q8_0_mmvq_rp",
19239 (QT_Q8_0, _, _) => "qmatvec_q8_0_mmvq",
19240 (QT_Q4_K, _, true) => "qmatvec_q4_K_mmvq_rp",
19244 (QT_Q6_K, _, true) => "qmatvec_q6_K_mmvq_rp",
19245 (QT_Q4_K, _, _) => "qmatvec_q4_K_mmvq",
19246 (QT_Q4_0, 2, false) => "qmatvec_q4_0_mmvq_mr2",
19247 (QT_Q4_0, _, false) => "qmatvec_q4_0_mmvq",
19248 (QT_Q5_K, _, _) => {
19249 if q5_il {
19250 "qmatvec_q5_K_mmvq_il"
19251 } else {
19252 "qmatvec_q5_K_mmvq"
19253 }
19254 }
19255 (QT_Q6_K, _, _) => "qmatvec_q6_K_mmvq",
19256 (QT_NVFP4, _, false) => "qmatvec_nvfp4_mmvq",
19257 (QT_F8_E4M3, _, _) => "qmatvec_e4m3_mmvq",
19258 _ => panic!("qmatvec_mmvq: qtype {qtype} has no MMVQ kernel"),
19259 };
19260 let f = self.func(name);
19261 let rows_per_block = ROWS_PER_BLOCK * mr;
19263 let cfg = LaunchConfig {
19264 grid_dim: (
19265 (out_f as u32 + rows_per_block - 1) / rows_per_block,
19266 m as u32,
19267 1,
19268 ),
19269 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
19272 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
19273 let __s_b = self.gpu.stream();
19274 let mut b = __s_b.launch_builder(&f);
19275 if qtype == QT_NVFP4 || qtype == QT_F8_E4M3 {
19280 if Self::pdl_on()
19283 && Self::pdl_mmvq_on()
19284 && Self::pdl_nvfp4q8_on()
19285 && name == "qmatvec_nvfp4_mmvq_mr2_rp"
19286 {
19287 use cudarc::driver::{DevicePtr, DevicePtrMut};
19288 let s = &self.gpu.stream();
19289 let (pw, _g0) = bytes.device_ptr(s);
19290 let (paq, _g1) = aq.device_ptr(s);
19291 let (pad, _g2) = ad.device_ptr(s);
19292 let (py, _g3) = y.device_ptr_mut(s);
19293 let mut ps = [
19294 &pw as *const _ as *mut std::ffi::c_void,
19295 &paq as *const _ as *mut _,
19296 &pad as *const _ as *mut _,
19297 &py as *const _ as *mut _,
19298 &inf as *const _ as *mut _,
19299 &outf as *const _ as *mut _,
19300 &mi as *const _ as *mut _,
19301 &rb as *const _ as *mut _,
19302 &scale as *const _ as *mut _,
19303 ];
19304 unsafe {
19305 self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?;
19306 }
19307 return Ok(());
19308 }
19309 b.arg(bytes)
19310 .arg(aq)
19311 .arg(ad)
19312 .arg(&mut *y)
19313 .arg(&inf)
19314 .arg(&outf)
19315 .arg(&mi)
19316 .arg(&rb)
19317 .arg(&scale);
19318 unsafe {
19319 b.launch(cfg)?;
19320 }
19321 } else if Self::pdl_on()
19322 && Self::pdl_mmvq_on()
19323 && (matches!(
19324 name,
19325 "qmatvec_q4_0_mmvq_rp" | "qmatvec_q6_K_mmvq" | "qmatvec_q6_K_mmvq_rp"
19326 ) || (Self::pdl_nvfp4q8_on()
19327 && matches!(name, "qmatvec_q8_0_mmvq_rp" | "qmatvec_q8_0_mmvq_mr2_rp")))
19328 {
19329 {
19333 use cudarc::driver::{DevicePtr, DevicePtrMut};
19334 let s = &self.gpu.stream();
19335 let (pw, _g0) = bytes.device_ptr(s);
19336 let (paq, _g1) = aq.device_ptr(s);
19337 let (pad, _g2) = ad.device_ptr(s);
19338 let (py, _g3) = y.device_ptr_mut(s);
19339 let mut ps = [
19340 &pw as *const _ as *mut std::ffi::c_void,
19341 &paq as *const _ as *mut _,
19342 &pad as *const _ as *mut _,
19343 &py as *const _ as *mut _,
19344 &inf as *const _ as *mut _,
19345 &outf as *const _ as *mut _,
19346 &mi as *const _ as *mut _,
19347 &rb as *const _ as *mut _,
19348 ];
19349 unsafe {
19350 self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?;
19351 }
19352 }
19353 if scale != 1.0 {
19354 self.scale_inplace(y, scale, m * out_f)?;
19355 }
19356 } else {
19357 b.arg(bytes)
19358 .arg(aq)
19359 .arg(ad)
19360 .arg(&mut *y)
19361 .arg(&inf)
19362 .arg(&outf)
19363 .arg(&mi)
19364 .arg(&rb);
19365 unsafe {
19366 b.launch(cfg)?;
19367 }
19368 if scale != 1.0 {
19369 self.scale_inplace(y, scale, m * out_f)?;
19370 }
19371 }
19372 Ok(())
19373 }
19374
19375 #[allow(clippy::too_many_arguments)] pub fn qmatvec_mmvq_raw(
19380 &self,
19381 bytes: &CudaSlice<u8>,
19382 x: &CudaSlice<f32>,
19383 m: usize,
19384 in_f: usize,
19385 out_f: usize,
19386 qtype: i32,
19387 row_bytes: usize,
19388 rp: bool,
19389 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19390 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
19391 self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, 1.0, rp)
19392 }
19393
19394 pub fn batched_supports(&self, qtype: i32) -> bool {
19398 matches!(
19399 qtype,
19400 QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_F8_E4M3 | QT_Q4_0
19401 )
19402 }
19403
19404 pub fn iq_fast_enabled() -> bool {
19412 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19413 *ON.get_or_init(|| {
19414 std::env::var("MEMRA_IQ_FAST")
19415 .map(|v| v != "0")
19416 .unwrap_or(true)
19417 })
19418 }
19419
19420 pub fn b8_enabled() -> bool {
19423 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19424 *ON.get_or_init(|| std::env::var("MEMRA_B8").map(|v| v != "0").unwrap_or(true))
19425 }
19426
19427 pub fn batched_mcols(m: usize) -> usize {
19429 if m == 2 {
19430 2
19431 } else if m <= 4 {
19432 4
19433 } else if m <= 8 {
19434 8
19435 } else {
19436 16
19437 }
19438 }
19439
19440 fn batched_kernel_name(qtype: i32, mcols: usize) -> Option<&'static str> {
19445 Some(match (qtype, mcols) {
19446 (QT_Q8_0, 2) => "qmatvec_q8_0_mmvq_b2",
19447 (QT_Q8_0, 4) => "qmatvec_q8_0_mmvq_b4",
19448 (QT_Q8_0, 8) => "qmatvec_q8_0_mmvq_b8",
19449 (QT_Q8_0, 16) => "qmatvec_q8_0_mmvq_b16",
19455 (QT_Q4_K, 2) => "qmatvec_q4_K_mmvq_b2",
19456 (QT_Q4_K, 4) => "qmatvec_q4_K_mmvq_b4",
19457 (QT_Q4_K, 8) => "qmatvec_q4_K_mmvq_b8",
19458 (QT_Q4_K, 16) => "qmatvec_q4_K_mmvq_b16",
19461 (QT_Q5_K, 2) => "qmatvec_q5_K_mmvq_b2",
19462 (QT_Q5_K, 4) => "qmatvec_q5_K_mmvq_b4",
19463 (QT_Q5_K, 8) => "qmatvec_q5_K_mmvq_b8",
19464 (QT_Q5_K, 16) => "qmatvec_q5_K_mmvq_b16",
19467 (QT_Q6_K, 2) => "qmatvec_q6_K_mmvq_b2",
19468 (QT_Q6_K, 4) => "qmatvec_q6_K_mmvq_b4",
19469 (QT_Q6_K, 8) => "qmatvec_q6_K_mmvq_b8",
19470 (QT_Q6_K, 16) => "qmatvec_q6_K_mmvq_b16",
19471 (QT_NVFP4, 2) => "qmatvec_nvfp4_mmvq_b2",
19472 (QT_NVFP4, 4) => "qmatvec_nvfp4_mmvq_b4",
19473 (QT_NVFP4, 8) => "qmatvec_nvfp4_mmvq_b8",
19474 (QT_NVFP4, 16) => "qmatvec_nvfp4_mmvq_b16",
19478 (QT_F8_E4M3, 2) => "qmatvec_e4m3_mmvq_b2",
19479 (QT_F8_E4M3, 4) => "qmatvec_e4m3_mmvq_b4",
19480 (QT_F8_E4M3, 8) => "qmatvec_e4m3_mmvq_b8",
19481 (QT_F8_E4M3, 16) => "qmatvec_e4m3_mmvq_b16",
19485 (QT_Q4_0, 2) => "qmatvec_q4_0_mmvq_b2",
19486 (QT_Q4_0, 4) => "qmatvec_q4_0_mmvq_b4",
19487 (QT_Q4_0, 8) => "qmatvec_q4_0_mmvq_b8",
19488 (QT_Q4_0, 16) => "qmatvec_q4_0_mmvq_b16",
19489 _ => return None,
19490 })
19491 }
19492
19493 pub fn sm_count(&self) -> i32 {
19528 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
19529 *SMS.get_or_init(|| {
19530 use cudarc::driver::sys::CUdevice_attribute_enum as A;
19531 self.gpu
19532 .ctx
19533 .attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT)
19534 .unwrap_or(82)
19535 })
19536 }
19537
19538 #[allow(clippy::too_many_arguments)]
19539 #[allow(clippy::if_same_then_else)] pub fn batched_variant(
19542 &self,
19543 _m: usize,
19544 in_f: usize,
19545 out_f: usize,
19546 qtype: i32,
19547 row_bytes: usize,
19548 mcols: usize,
19549 rp: bool,
19550 ) -> &'static str {
19551 if qtype == QT_Q8_0 {
19556 return if rp { "rp" } else { "base" };
19557 }
19558 static BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
19559 let bv = *BV.get_or_init(|| match std::env::var("MEMRA_MMVQ_BV").as_deref() {
19560 Ok("base") => "base",
19561 Ok("pf") => "pf",
19562 Ok("r2") => "r2",
19563 Ok("r2w8") => "r2w8",
19564 Ok("pfr2") => "pfr2",
19565 Ok("ca") => "ca",
19566 Ok("car2") => "car2",
19567 Ok("rp") => "rp",
19570 Ok("rpr2") => "rpr2",
19571 Ok("rpr2w8") => "rpr2w8",
19572 Ok("rpca") => "rpca",
19575 Ok("rpcar2") => "rpcar2",
19576 Ok("rpsc") => "rpsc",
19583 Ok("rpms") => "rpms",
19584 Ok("rpmsc") => "rpmsc",
19585 Ok("rpks") => "rpks",
19586 Ok("rpksc") => "rpksc",
19587 _ => "auto",
19588 });
19589 let ca_ok = qtype == QT_NVFP4 && row_bytes.is_multiple_of(16) && in_f.is_multiple_of(1024);
19593 static KS_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19598 let ks_on = *KS_ON.get_or_init(|| std::env::var("MEMRA_KS").as_deref() != Ok("0"));
19599 let sc_ok = ks_on && qtype == QT_NVFP4 && in_f.is_multiple_of(256) && (in_f / 64 <= 272);
19600 let ks_ok = ks_on && qtype == QT_NVFP4 && in_f.is_multiple_of(512) && (in_f / 64 <= 272);
19601 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
19602 let sms = *SMS.get_or_init(|| {
19603 use cudarc::driver::sys::CUdevice_attribute_enum as A;
19604 self.gpu
19605 .ctx
19606 .attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT)
19607 .unwrap_or(82)
19608 });
19609 let kq_r2 = matches!(qtype, QT_Q4_K | QT_Q5_K | QT_Q6_K);
19629 static KQBV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
19632 let kq_bv = *KQBV.get_or_init(|| match std::env::var("MEMRA_KQ_BV").as_deref() {
19633 Ok("base") => "base",
19634 Ok("r2") => "r2",
19635 Ok("r2w8") => "r2w8",
19636 _ => "auto",
19637 });
19638 let variant: &'static str = if qtype == QT_Q4_0 {
19639 static Q40BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
19643 let q40 = *Q40BV.get_or_init(|| match std::env::var("MEMRA_Q40_BV").as_deref() {
19644 Ok("base") => "base",
19650 Ok("r2") => "r2",
19651 Ok("ms") => "ms",
19652 Ok("sm") => "sm",
19653 Ok("la") => "la",
19654 _ => "auto",
19655 });
19656 let v = if q40 != "auto" {
19657 q40
19658 } else if (out_f as u32).div_ceil(8) >= 4 * sms as u32 {
19659 "r2"
19660 } else {
19661 "base"
19662 };
19663 if rp {
19668 match v {
19669 "ms" => "r2ms_rp",
19670 "sm" => "r2sm_rp",
19671 "la" => "r2la_rp",
19672 "r2" => "r2_rp",
19673 _ => "rp",
19674 }
19675 } else if matches!(v, "ms" | "sm" | "la") {
19676 "r2"
19677 } else {
19678 v
19679 }
19680 } else if qtype != QT_NVFP4 && !kq_r2 {
19681 "base"
19682 } else if kq_r2 && rp {
19683 "rp"
19687 } else if kq_r2 {
19688 if kq_bv != "auto" {
19691 if kq_bv == "r2w8" && mcols != 4 {
19692 "r2"
19693 } else {
19694 kq_bv
19695 }
19696 } else if bv != "auto" {
19697 match bv {
19698 "r2" | "pfr2" | "rpr2" | "car2" => "r2",
19699 "r2w8" | "rpr2w8" => {
19700 if mcols != 4 {
19701 "r2"
19702 } else {
19703 "r2w8"
19704 }
19705 }
19706 _ => "base", }
19708 } else {
19709 #[allow(clippy::manual_div_ceil)]
19710 let blocks = (out_f + 7) / 8;
19712 let waves = blocks as f64 / (7 * sms as usize) as f64;
19713 let filled = blocks >= 4 * sms as usize;
19714 let use_r2 = if qtype == QT_Q4_K {
19715 filled
19716 } else {
19717 waves >= 2.0
19718 };
19719 if use_r2 { "r2" } else { "base" }
19720 }
19721 } else if bv != "auto" {
19722 let v = if bv == "r2w8" && mcols == 2 {
19727 "r2"
19728 } else if bv == "ca" && (!ca_ok || mcols == 8) {
19729 "pf"
19730 } else if bv == "car2" && (!ca_ok || mcols == 8) {
19731 "r2"
19732 } else if bv == "pfr2" && mcols == 8 {
19733 "r2"
19734 } else if (bv == "rpr2w8" || bv == "rpr2") && mcols == 2 {
19735 "rpr2"
19736 }
19737 else if (bv == "rpca" || bv == "rpcar2") && (!ca_ok || mcols == 8) {
19739 if mcols == 8 { "rpr2w8" } else { "rpr2" }
19740 } else if bv == "rpcar2" && mcols == 2 {
19741 "rpca"
19742 }
19743 else if (bv == "rpsc" || bv == "rpmsc") && !sc_ok {
19746 "rpr2"
19747 } else if (bv == "rpks" || bv == "rpksc") && !ks_ok {
19748 "rpr2"
19749 } else {
19750 bv
19751 };
19752 if rp {
19753 match v {
19754 "base" | "pf" | "ca" | "rp" => "rp",
19755 "r2" | "pfr2" | "car2" | "rpr2" => "rpr2",
19756 "r2w8" | "rpr2w8" => {
19757 if mcols == 2 {
19758 "rpr2"
19759 } else {
19760 "rpr2w8"
19761 }
19762 }
19763 other => other, }
19765 } else {
19766 v
19767 }
19768 } else if mcols == 8 {
19769 if rp {
19780 if sc_ok { "rpsc" } else { "rpr2w8" }
19781 } else {
19782 "r2w8"
19783 }
19784 } else if mcols >= 4 {
19785 #[allow(clippy::manual_div_ceil)]
19789 let blocks = (out_f + 7) / 8;
19791 let r7 = 7 * sms as usize;
19792 let r8 = 8 * sms as usize;
19793 let waves = blocks as f64 / r7 as f64;
19794 let filled = blocks >= 4 * sms as usize;
19795 if filled && blocks.div_ceil(r8) < blocks.div_ceil(r7) {
19799 if rp { "rpr2w8" } else { "r2w8" }
19803 } else if waves >= 2.0 || (waves <= 1.0 && filled) {
19804 if rp { "rpr2" } else { "r2" }
19807 } else {
19808 if rp { "rp" } else { "pf" }
19812 }
19813 } else if in_f >= 6144 {
19814 if rp { "rpr2" } else { "r2" }
19818 } else if rp {
19819 #[allow(clippy::manual_div_ceil)]
19824 let waves = ((out_f + 7) / 8) as f64 / (7 * sms as usize) as f64;
19826 if sc_ok && (0.9..=1.1).contains(&waves) {
19827 "rpsc"
19828 } else {
19829 "rp"
19830 }
19831 } else {
19832 "base"
19833 };
19834 variant
19835 }
19836
19837 #[allow(clippy::too_many_arguments)]
19838 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_mmvq_batched(
19841 &self,
19842 bytes: &CudaSlice<u8>,
19843 aq: &CudaSlice<i8>,
19844 ad: &CudaSlice<f32>,
19845 m: usize,
19846 in_f: usize,
19847 out_f: usize,
19848 qtype: i32,
19849 row_bytes: usize,
19850 mcols: usize,
19851 scale: f32,
19852 rp: bool,
19853 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19854 const ROWS_PER_BLOCK: u32 = 4;
19855 let forced: Option<&'static str> = {
19860 static V: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
19861 V.get_or_init(|| std::env::var("MEMRA_BVAR").ok())
19862 .as_deref()
19863 .map(|s| Box::leak(s.to_string().into_boxed_str()) as &'static str)
19864 };
19865 let variant = match forced {
19866 Some(v) if !rp || v.contains("rp") => v,
19867 _ => self.batched_variant(m, in_f, out_f, qtype, row_bytes, mcols, rp),
19868 };
19869 let base_name = Self::batched_kernel_name(qtype, mcols).ok_or_else(|| {
19870 format!("qmatvec_mmvq_batched: no kernel for qtype {qtype} mcols {mcols}")
19871 })?;
19872 let variant = if mcols == 16 {
19876 if rp { "rp" } else { "base" }
19877 } else {
19878 variant
19879 };
19880 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19887 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
19888 if b567
19889 && qtype == QT_NVFP4
19890 && rp
19891 && mcols == 8
19892 && (5..=7).contains(&m)
19893 && matches!(variant, "rpsc" | "rpr2w8")
19894 {
19895 let f = self.func(&format!("qmatvec_nvfp4_mmvq_b{m}_{variant}"));
19896 let rows_per_block = ROWS_PER_BLOCK * 2; let mut y = self.alloc_uninit::<f32>(m * out_f)?;
19898 let cfg = LaunchConfig {
19899 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
19900 block_dim: (32, ROWS_PER_BLOCK, 1),
19901 shared_mem_bytes: 0,
19902 };
19903 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
19904 let __s_b = self.gpu.stream();
19905 let mut b = __s_b.launch_builder(&f);
19906 b.arg(bytes)
19907 .arg(aq)
19908 .arg(ad)
19909 .arg(&mut y)
19910 .arg(&inf)
19911 .arg(&outf)
19912 .arg(&mi)
19913 .arg(&rb);
19914 unsafe {
19915 b.launch(cfg)?;
19916 }
19917 if scale != 1.0 {
19918 self.scale_inplace(&mut y, scale, m * out_f)?;
19919 }
19920 return Ok(y);
19921 }
19922 let (name, rows_per_block): (std::borrow::Cow<'static, str>, u32) = match variant {
19923 "base" => (base_name.into(), ROWS_PER_BLOCK),
19924 "pf" => (format!("{base_name}_pf").into(), ROWS_PER_BLOCK),
19925 "ca" => (format!("{base_name}_ca").into(), ROWS_PER_BLOCK),
19926 "rp" => (format!("{base_name}_rp").into(), ROWS_PER_BLOCK),
19927 "rpca" => (format!("{base_name}_rpca").into(), ROWS_PER_BLOCK), "rpks" => (format!("{base_name}_rpks").into(), ROWS_PER_BLOCK),
19931 "rpksc" => (format!("{base_name}_rpksc").into(), ROWS_PER_BLOCK),
19932 "rpms" => (format!("{base_name}_rpms").into(), ROWS_PER_BLOCK),
19933 "rpmsc" => (format!("{base_name}_rpmsc").into(), ROWS_PER_BLOCK),
19934 "r2ms_rp" => (format!("{base_name}_r2ms_rp").into(), ROWS_PER_BLOCK),
19935 "r2sm_rp" => (format!("{base_name}_r2sm_rp").into(), ROWS_PER_BLOCK * 2),
19936 "r2la_rp" => (format!("{base_name}_r2la_rp").into(), ROWS_PER_BLOCK * 2),
19937 v => (format!("{base_name}_{v}").into(), ROWS_PER_BLOCK * 2), };
19939 debug_assert!(
19940 !rp || name.contains("_rp"),
19941 "rp weight dispatched to a GGUF-layout kernel"
19942 );
19943 let f = self.func(&name);
19944 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
19945 let smem = if name.contains("_r2sm_rp") {
19947 (mcols * 32 * 9 * 4 + mcols * 32 * 4) as u32
19948 } else {
19949 0
19950 };
19951 let cfg = LaunchConfig {
19952 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
19953 block_dim: (32, ROWS_PER_BLOCK, 1),
19954 shared_mem_bytes: smem,
19955 };
19956 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
19957 let __s_b = self.gpu.stream();
19958 let mut b = __s_b.launch_builder(&f);
19959 b.arg(bytes)
19960 .arg(aq)
19961 .arg(ad)
19962 .arg(&mut y)
19963 .arg(&inf)
19964 .arg(&outf)
19965 .arg(&mi)
19966 .arg(&rb);
19967 unsafe {
19968 b.launch(cfg)?;
19969 }
19970 if scale != 1.0 {
19971 self.scale_inplace(&mut y, scale, m * out_f)?;
19972 }
19973 Ok(y)
19974 }
19975
19976 #[allow(clippy::too_many_arguments)] pub fn qmatvec_batched_raw(
19981 &self,
19982 bytes: &CudaSlice<u8>,
19983 x: &CudaSlice<f32>,
19984 m: usize,
19985 in_f: usize,
19986 out_f: usize,
19987 qtype: i32,
19988 row_bytes: usize,
19989 mcols: usize,
19990 rp: bool,
19991 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19992 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
19993 self.qmatvec_mmvq_batched(
19994 bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, 1.0, rp,
19995 )
19996 }
19997
19998 #[allow(clippy::too_many_arguments)] pub fn qmatvec_nvfp4_batched_raw(
20001 &self,
20002 bytes: &CudaSlice<u8>,
20003 x: &CudaSlice<f32>,
20004 m: usize,
20005 in_f: usize,
20006 out_f: usize,
20007 row_bytes: usize,
20008 mcols: usize,
20009 rp: bool,
20010 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20011 self.qmatvec_batched_raw(bytes, x, m, in_f, out_f, QT_NVFP4, row_bytes, mcols, rp)
20012 }
20013
20014 fn try_fp4_gemm(
20018 &self,
20019 w: &crate::model::GpuTensor,
20020 x: &CudaSlice<f32>,
20021 m: usize,
20022 in_f: usize,
20023 out_f: usize,
20024 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
20025 use crate::model::GpuTensor;
20026 if cfg!(memra_portable_cuda) {
20027 return Ok(None);
20028 }
20029 if std::env::var("MEMRA_FP4").is_ok() {
20038 refuse_portable_force("MEMRA_FP4", "the sm_120a mxf4 block-scale MMA");
20039 assert!(
20040 konst_eq(env!("MEMRA_BUILT_CUDA_ARCH"), "120a"),
20041 "MEMRA_FP4 forces the native mxf4 block-scale GEMM (qmatvec_gemm_nvfp4_fp4), \
20042 which only the sm_120a fatbin contains — this is an sm_{} build. Unset \
20043 MEMRA_FP4; the W4A8 int8 path is the correct default for NVFP4 weights.",
20044 env!("MEMRA_BUILT_CUDA_ARCH")
20045 );
20046 }
20047 if std::env::var("MEMRA_FP4").is_err() {
20048 return Ok(None);
20049 }
20050 #[cfg(memra_cutlass)]
20059 if m >= 128 && std::env::var("MEMRA_FP4_CUTLASS").is_ok() {
20060 if let GpuTensor::Quant {
20061 bytes,
20062 qtype,
20063 scale,
20064 row_bytes,
20065 cutlass,
20066 ..
20067 } = w
20068 {
20069 if *qtype == QT_NVFP4 && in_f % 64 == 0 {
20070 if let Some(cw) = cutlass {
20071 let y = self.cutlass_fp4_gemm(
20073 &cw.b_packed,
20074 &cw.sfb_swizzled,
20075 x,
20076 *scale,
20077 m,
20078 out_f,
20079 in_f,
20080 )?;
20081 return Ok(Some(y));
20082 } else if std::env::var("MEMRA_FP4_CUTLASS_OTF").is_ok() {
20083 let (b_packed, sfb_sw) =
20088 self.build_cutlass_weight(bytes, out_f, in_f, *row_bytes)?;
20089 let y =
20090 self.cutlass_fp4_gemm(&b_packed, &sfb_sw, x, *scale, m, out_f, in_f)?;
20091 return Ok(Some(y));
20092 }
20093 }
20094 }
20095 }
20096 if let GpuTensor::Quant {
20097 bytes,
20098 qtype,
20099 row_bytes,
20100 scale,
20101 rp,
20102 ..
20103 } = w
20104 {
20105 if *qtype == QT_NVFP4 && in_f.is_multiple_of(64) && !*rp {
20108 let y =
20109 self.qmatvec_gemm_nvfp4_fp4(bytes, x, m, in_f, out_f, *row_bytes, *scale)?;
20110 return Ok(Some(y));
20111 }
20112 }
20113 Ok(None)
20114 }
20115
20116 #[allow(clippy::too_many_arguments)] pub fn rms_norm_f16out(
20121 &self,
20122 x: &CudaSlice<f32>,
20123 w: &CudaSlice<f32>,
20124 dst: &mut CudaSlice<f32>,
20125 dst16: &mut CudaSlice<u8>,
20126 ncols: usize,
20127 nrows: usize,
20128 eps: f32,
20129 ) -> Result<(), Box<dyn std::error::Error>> {
20130 let f = self.func("rms_norm_f16out_f32");
20131 let cfg = LaunchConfig {
20132 grid_dim: (nrows as u32, 1, 1),
20133 block_dim: (rms_block(), 1, 1),
20134 shared_mem_bytes: 0,
20135 };
20136 let (nc, e) = (ncols as i32, eps);
20137 let __s_b = self.gpu.stream();
20138 let mut b = __s_b.launch_builder(&f);
20139 b.arg(x).arg(w).arg(dst).arg(dst16).arg(&nc).arg(&e);
20140 unsafe {
20141 b.launch(cfg)?;
20142 }
20143 Ok(())
20144 }
20145
20146 #[allow(clippy::too_many_arguments)]
20149 pub fn add_rms_norm_f16out(
20150 &self,
20151 a: &CudaSlice<f32>,
20152 b: &CudaSlice<f32>,
20153 w: &CudaSlice<f32>,
20154 res: &mut CudaSlice<f32>,
20155 dst: &mut CudaSlice<f32>,
20156 dst16: &mut CudaSlice<u8>,
20157 ncols: usize,
20158 nrows: usize,
20159 eps: f32,
20160 ) -> Result<(), Box<dyn std::error::Error>> {
20161 let f = self.func("add_rms_norm_f16out_f32");
20162 let cfg = LaunchConfig {
20163 grid_dim: (nrows as u32, 1, 1),
20164 block_dim: (rms_block(), 1, 1),
20165 shared_mem_bytes: 0,
20166 };
20167 let (nc, e) = (ncols as i32, eps);
20168 let __s_lb = self.gpu.stream();
20169 let mut lb = __s_lb.launch_builder(&f);
20170 lb.arg(a)
20171 .arg(b)
20172 .arg(w)
20173 .arg(res)
20174 .arg(dst)
20175 .arg(dst16)
20176 .arg(&nc)
20177 .arg(&e);
20178 unsafe {
20179 lb.launch(cfg)?;
20180 }
20181 Ok(())
20182 }
20183
20184 pub fn matmul_group_xh(
20187 &self,
20188 ws: &[&crate::model::GpuTensor],
20189 x: &CudaSlice<f32>,
20190 xh: &CudaSlice<u8>,
20191 m: usize,
20192 ) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
20193 let mut out = Vec::with_capacity(ws.len());
20194 let in_f = ws[0].in_features();
20195 for w in ws {
20196 if w.in_features() == in_f
20197 && m >= 16
20198 && !self.verify_exact_on()
20199 && let Some(y) = self.try_f16_gemm_pre(w, xh, m)?
20200 {
20201 out.push(y);
20202 continue;
20203 }
20204 out.push(self.matmul(w, x, m)?);
20205 }
20206 Ok(out)
20207 }
20208
20209 pub fn gdn_pad_mask(
20212 &self,
20213 beta: &mut CudaSlice<f32>,
20214 g_log: &mut CudaSlice<f32>,
20215 len_d: &CudaSlice<i32>,
20216 h: usize,
20217 t: usize,
20218 ) -> Result<(), Box<dyn std::error::Error>> {
20219 let f = self.func("gdn_pad_mask_f32");
20220 let cfg = LaunchConfig::for_num_elems((t * h) as u32);
20221 let (hi, ti) = (h as i32, t as i32);
20222 let __s_b = self.gpu.stream();
20223 let mut b = __s_b.launch_builder(&f);
20224 b.arg(beta).arg(g_log).arg(len_d).arg(&hi).arg(&ti);
20225 unsafe {
20226 b.launch(cfg)?;
20227 }
20228 Ok(())
20229 }
20230
20231 pub fn row_gather_dev(
20234 &self,
20235 src: &CudaSlice<f32>,
20236 dst: &mut CudaSlice<f32>,
20237 len_d: &CudaSlice<i32>,
20238 ncols: usize,
20239 ) -> Result<(), Box<dyn std::error::Error>> {
20240 let f = self.func("row_gather_dev_f32");
20241 let cfg = LaunchConfig::for_num_elems(ncols as u32);
20242 let nc = ncols as i32;
20243 let __s_b = self.gpu.stream();
20244 let mut b = __s_b.launch_builder(&f);
20245 b.arg(src).arg(dst).arg(len_d).arg(&nc);
20246 unsafe {
20247 b.launch(cfg)?;
20248 }
20249 Ok(())
20250 }
20251
20252 pub fn matmul_group(
20259 &self,
20260 ws: &[&crate::model::GpuTensor],
20261 x: &CudaSlice<f32>,
20262 m: usize,
20263 ) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
20264 use crate::model::GpuTensor;
20265 let mut out = Vec::with_capacity(ws.len());
20266 let any_mirror = ws
20267 .iter()
20268 .any(|w| matches!(w, GpuTensor::Quant { f16: Some(_), .. }));
20269 if m >= 16 && any_mirror && !self.verify_exact_on() {
20270 let in_f = ws[0].in_features();
20271 let xh = self.f16_act(x, m * in_f, in_f)?;
20272 for w in ws {
20273 if w.in_features() == in_f
20274 && let Some(y) = self.try_f16_gemm_pre(w, &xh, m)?
20275 {
20276 out.push(y);
20277 continue;
20278 }
20279 out.push(self.matmul(w, x, m)?);
20280 }
20281 return Ok(out);
20282 }
20283 for w in ws {
20284 out.push(self.matmul(w, x, m)?);
20285 }
20286 Ok(out)
20287 }
20288
20289 pub fn matmul_group_multi(
20296 &self,
20297 ws: &[&crate::model::GpuTensor],
20298 xs: &[&CudaSlice<f32>],
20299 ms: &[usize],
20300 ) -> Result<Vec<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
20301 assert_eq!(xs.len(), ms.len());
20302 let in_f = ws[0].in_features();
20303 let total: usize = ms.iter().sum();
20304 let mut xcat = self.uninit(total * in_f)?;
20305 let mut off = 0usize;
20306 for (x, &m) in xs.iter().zip(ms) {
20307 self.copy_into(&mut xcat, off * in_f, x, m * in_f)?;
20308 off += m;
20309 }
20310 let ys = self.matmul_group(ws, &xcat, total)?;
20311 let mut out: Vec<Vec<CudaSlice<f32>>> = (0..xs.len()).map(|_| Vec::new()).collect();
20312 for (w, y) in ws.iter().zip(ys) {
20313 let out_f = w.out_features();
20314 let mut off = 0usize;
20315 for (s, &m) in ms.iter().enumerate() {
20316 let mut ys_s = self.uninit(m * out_f)?;
20317 let src = y.slice(off * out_f..(off + m) * out_f);
20318 self.gpu.stream().memcpy_dtod(&src, &mut ys_s)?;
20319 out[s].push(ys_s);
20320 off += m;
20321 }
20322 }
20323 Ok(out)
20324 }
20325
20326 pub fn gemm_supports(&self, w: &crate::model::GpuTensor) -> bool {
20336 use crate::model::GpuTensor;
20337 if !legacy_quant_gemm_allowed(
20338 cfg!(memra_portable_cuda),
20339 cfg!(memra_hopper_mma),
20340 std::env::var_os("MEMRA_NO_GEMM").is_some(),
20341 ) {
20342 return false;
20343 }
20344 match w {
20345 GpuTensor::Quant { qtype, .. } => {
20346 matches!(*qtype, QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q4_0)
20347 || (*qtype == QT_NVFP4 && w.in_features().is_multiple_of(64))
20348 }
20349 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
20350 }
20351 }
20352
20353 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_gemm(
20361 &self,
20362 w: &crate::model::GpuTensor,
20363 aq: &CudaSlice<i8>,
20364 ad: &CudaSlice<f32>,
20365 m: usize,
20366 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20367 use crate::model::GpuTensor;
20368 let in_f = w.in_features();
20369 let out_f = w.out_features();
20370 let (bytes, qtype, row_bytes, scale, rp) = match w {
20371 GpuTensor::Quant {
20372 bytes,
20373 qtype,
20374 row_bytes,
20375 scale,
20376 rp,
20377 ..
20378 } => (bytes, *qtype, *row_bytes, *scale, *rp),
20379 _ => unreachable!("gemm_supports guaranteed Quant"),
20380 };
20381 if cfg!(memra_hopper_mma)
20387 && qtype == QT_Q8_0
20388 && out_f.is_multiple_of(64)
20389 && wgmma_gemm_enabled()
20390 && let GpuTensor::Quant { rp4: Some(m4), .. } = w
20391 {
20392 let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, aq, ad, m, in_f, out_f)?;
20393 if scale != 1.0 {
20394 self.scale_inplace(&mut y, scale, m * out_f)?;
20395 }
20396 return Ok(y);
20397 }
20398 let name = match qtype {
20399 QT_Q8_0 => "qmatvec_gemm_q8_0",
20400 QT_Q4_K => "qmatvec_gemm_q4_K",
20401 QT_Q4_0 => {
20402 if rp {
20403 "qmatvec_gemm_q4_0_rp"
20404 } else {
20405 "qmatvec_gemm_q4_0"
20406 }
20407 }
20408 QT_Q5_K => "qmatvec_gemm_q5_K",
20409 QT_Q6_K => "qmatvec_gemm_q6_K",
20410 QT_NVFP4 => {
20411 if rp {
20412 "qmatvec_gemm_nvfp4_rp"
20413 } else {
20414 "qmatvec_gemm_nvfp4"
20415 }
20416 }
20417 _ => unreachable!(),
20418 };
20419 let f = self.func(name);
20420 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);
20425 let k1_tile = if is_k1 {
20427 k1_launch_override().unwrap_or((128, 128, 8))
20428 } else {
20429 (128, 128, 8)
20430 };
20431 let (bm, bn): (u32, u32) = if is_k1 {
20432 (k1_tile.0, k1_tile.1)
20433 } else {
20434 (64, 256)
20435 };
20436 let warps: u32 = if is_k1 {
20437 k1_tile.2
20438 } else {
20439 match qtype {
20440 QT_NVFP4 => 8,
20441 _ => 4,
20442 }
20443 };
20444 let cfg = LaunchConfig {
20445 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
20446 block_dim: (32, warps, 1),
20447 shared_mem_bytes: 0,
20448 };
20449 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
20450 let __s_b = self.gpu.stream();
20451 let mut b = __s_b.launch_builder(&f);
20452 b.arg(bytes)
20453 .arg(aq)
20454 .arg(ad)
20455 .arg(&mut y)
20456 .arg(&inf)
20457 .arg(&outf)
20458 .arg(&mi)
20459 .arg(&rb);
20460 unsafe {
20461 b.launch(cfg)?;
20462 }
20463 if scale != 1.0 {
20464 self.scale_inplace(&mut y, scale, m * out_f)?;
20465 }
20466 Ok(y)
20467 }
20468
20469 #[allow(clippy::too_many_arguments)]
20474 #[allow(clippy::manual_div_ceil)] pub fn qmatvec_gemm_raw(
20477 &self,
20478 bytes: &CudaSlice<u8>,
20479 x: &CudaSlice<f32>,
20480 m: usize,
20481 in_f: usize,
20482 out_f: usize,
20483 qtype: i32,
20484 row_bytes: usize,
20485 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20486 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
20487 let name = match qtype {
20488 QT_Q8_0 => "qmatvec_gemm_q8_0",
20489 QT_Q4_K => "qmatvec_gemm_q4_K",
20490 QT_Q4_0 => "qmatvec_gemm_q4_0",
20491 QT_Q5_K => "qmatvec_gemm_q5_K",
20492 QT_Q6_K => "qmatvec_gemm_q6_K",
20493 QT_NVFP4 => "qmatvec_gemm_nvfp4",
20494 QT_NVFP4_RP => "qmatvec_gemm_nvfp4_rp",
20495 _ => panic!("qmatvec_gemm_raw: qtype {qtype} has no GEMM kernel"),
20496 };
20497 let f = self.func(name);
20498 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);
20502 let k1_tile = if is_k1 {
20504 k1_launch_override().unwrap_or((128, 128, 8))
20505 } else {
20506 (128, 128, 8)
20507 };
20508 let (bm, bn): (u32, u32) = if is_k1 {
20509 (k1_tile.0, k1_tile.1)
20510 } else {
20511 (64, 256)
20512 };
20513 let warps: u32 = if is_k1 {
20514 k1_tile.2
20515 } else {
20516 match qtype {
20517 QT_NVFP4 | QT_NVFP4_RP => 8,
20518 _ => 4,
20519 }
20520 };
20521 let cfg = LaunchConfig {
20522 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
20523 block_dim: (32, warps, 1),
20524 shared_mem_bytes: 0,
20525 };
20526 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
20527 let __s_b = self.gpu.stream();
20528 let mut b = __s_b.launch_builder(&f);
20529 b.arg(bytes)
20530 .arg(&aq)
20531 .arg(&ad)
20532 .arg(&mut y)
20533 .arg(&inf)
20534 .arg(&outf)
20535 .arg(&mi)
20536 .arg(&rb);
20537 unsafe {
20538 b.launch(cfg)?;
20539 }
20540 Ok(y)
20541 }
20542
20543 pub fn qmatvec_gemm_q8_0_wgmma_raw(
20550 &self,
20551 rp4: &CudaSlice<u8>,
20552 aq: &CudaSlice<i8>,
20553 ad: &CudaSlice<f32>,
20554 m: usize,
20555 in_f: usize,
20556 out_f: usize,
20557 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20558 assert!(
20559 out_f.is_multiple_of(64) && in_f.is_multiple_of(32),
20560 "wgmma GEMM needs out_f%64==0, in_f%32==0"
20561 );
20562 let f = self.func("qmatvec_gemm_q8_0_wgmma");
20563 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
20565 grid_dim: ((out_f / 64) as u32, (m as u32).div_ceil(64), 1),
20566 block_dim: (128, 1, 1),
20567 shared_mem_bytes: 0,
20568 };
20569 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
20570 let __s_b = self.gpu.stream();
20571 let mut b = __s_b.launch_builder(&f);
20572 b.arg(rp4)
20573 .arg(aq)
20574 .arg(ad)
20575 .arg(&mut y)
20576 .arg(&inf)
20577 .arg(&outf)
20578 .arg(&mi);
20579 unsafe {
20580 b.launch(cfg)?;
20581 }
20582 Ok(y)
20583 }
20584
20585 pub fn scale_inplace(
20587 &self,
20588 y: &mut CudaSlice<f32>,
20589 s: f32,
20590 n: usize,
20591 ) -> Result<(), Box<dyn std::error::Error>> {
20592 let f = self.func("scale_f32");
20593 let cfg = LaunchConfig::for_num_elems(n as u32);
20594 let (sf, ni) = (s, n as i32);
20595 let __s_b = self.gpu.stream();
20596 let mut b = __s_b.launch_builder(&f);
20597 b.arg(y).arg(&sf).arg(&ni);
20598 unsafe {
20599 b.launch(cfg)?;
20600 }
20601 Ok(())
20602 }
20603
20604 pub fn bf16_to_f32(
20609 &self,
20610 data: &cudarc::driver::CudaView<'_, u8>,
20611 n: usize,
20612 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20613 let mut out = self.alloc_uninit::<f32>(n)?;
20614 let f = self.func("bf16_to_f32");
20615 let cfg = LaunchConfig::for_num_elems(n as u32);
20616 let ni = n as i32;
20617 let __s_b = self.gpu.stream();
20618 let mut b = __s_b.launch_builder(&f);
20619 b.arg(data).arg(&mut out).arg(&ni);
20620 unsafe {
20621 b.launch(cfg)?;
20622 }
20623 Ok(out)
20624 }
20625
20626 #[allow(clippy::too_many_arguments)] fn linear_bf16_chunked(
20634 &self,
20635 x: &CudaSlice<f32>,
20636 data: &CudaSlice<u8>,
20637 m: usize,
20638 in_f: usize,
20639 out_f: usize,
20640 exact: bool,
20641 canonical_chunk_rows: Option<usize>,
20642 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20643 static EXP_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
20647 static EXP_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
20648 static EXP_WBYTES: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
20649 let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
20650 let started = timing.then(std::time::Instant::now);
20651 let result =
20652 self.linear_bf16_chunked_inner(x, data, m, in_f, out_f, exact, canonical_chunk_rows);
20653 if let Some(started) = started {
20654 use std::sync::atomic::Ordering;
20655 self.stream().synchronize()?;
20656 let ns = EXP_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
20657 + started.elapsed().as_nanos() as u64;
20658 let wb = EXP_WBYTES.fetch_add((in_f * out_f * 2) as u64, Ordering::Relaxed)
20659 + (in_f * out_f * 2) as u64;
20660 let calls = EXP_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
20661 if calls.is_multiple_of(1024) {
20662 eprintln!(
20663 "[bf16-expand-timing] calls={calls} total_ms={:.1} avg_us={:.1} \
20664 weight_gb={:.2}",
20665 ns as f64 / 1.0e6,
20666 ns as f64 / calls as f64 / 1.0e3,
20667 wb as f64 / 1.0e9,
20668 );
20669 }
20670 }
20671 result
20672 }
20673
20674 pub(crate) fn bf16_mmv_on() -> bool {
20679 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20680 *ON.get_or_init(|| std::env::var("MEMRA_BF16_MMV").as_deref() == Ok("1"))
20681 }
20682
20683 fn matvec_bf16(
20686 &self,
20687 data: &CudaSlice<u8>,
20688 x: &CudaSlice<f32>,
20689 in_f: usize,
20690 out_f: usize,
20691 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20692 if data.len() != in_f * out_f * 2 || x.len() < in_f || !in_f.is_multiple_of(8) {
20693 return Err(format!(
20694 "matvec_bf16 geometry bytes={} x={} in={in_f} out={out_f}",
20695 data.len(),
20696 x.len()
20697 )
20698 .into());
20699 }
20700 let mut y = self.alloc_uninit::<f32>(out_f)?;
20701 let f = self.func("matvec_bf16_f32acc");
20702 let cfg = LaunchConfig {
20703 grid_dim: (out_f as u32, 1, 1),
20704 block_dim: (mmv_block(), 1, 1),
20705 shared_mem_bytes: 0,
20706 };
20707 let ini = in_f as i32;
20708 let __s_bld = self.gpu.stream();
20709 let mut bld = __s_bld.launch_builder(&f);
20710 bld.arg(data).arg(x).arg(&mut y).arg(&ini);
20711 unsafe {
20712 bld.launch(cfg)?;
20713 }
20714 Ok(y)
20715 }
20716
20717 #[allow(clippy::too_many_arguments)]
20721 #[allow(clippy::too_many_arguments)]
20726 #[allow(clippy::too_many_arguments)]
20731 pub fn qk_norm_rope_append_inc_dcw_rows(
20732 &self,
20733 q_raw_t: &CudaSlice<f32>,
20734 k_raw_t: &CudaSlice<f32>,
20735 v_raw_t: &CudaSlice<f32>,
20736 qw: &CudaSlice<f32>,
20737 kw: &CudaSlice<f32>,
20738 q_out_t: &mut CudaSlice<f32>,
20739 k_out_t: &mut CudaSlice<f32>,
20740 tab: &CudaSlice<u64>,
20741 pos_t: &CudaSlice<i32>,
20742 same_session: bool,
20743 t: usize,
20744 kv_dim_k: usize,
20745 kv_dim_v: usize,
20746 k_tok_bytes: usize,
20747 v_tok_bytes: usize,
20748 head_dim: usize,
20749 n_dims: usize,
20750 nh_q: usize,
20751 nh_k: usize,
20752 eps: f32,
20753 freq_base: f32,
20754 freq_scale: f32,
20755 ff: Option<&CudaSlice<f32>>,
20756 ) -> Result<(), Box<dyn std::error::Error>> {
20757 if head_dim != 128
20758 || kv_dim_v != kv_dim_k
20759 || kv_dim_k != nh_k * head_dim
20760 || t == 0
20761 || t > 32
20762 || tab.len() < t * 6
20763 || pos_t.len() < t
20764 || q_raw_t.len() < t * nh_q * head_dim
20765 || k_raw_t.len() < t * nh_k * head_dim
20766 || v_raw_t.len() < t * kv_dim_v
20767 || q_out_t.len() < t * nh_q * head_dim
20768 || k_out_t.len() < t * nh_k * head_dim
20769 {
20770 return Err(format!(
20771 "qk_norm_rope_append_inc_rows geometry head_dim={head_dim} t={t} \
20772 nh_q={nh_q} nh_k={nh_k}"
20773 )
20774 .into());
20775 }
20776 let f = self.func("qk_norm_rope_append_inc_dcw_rows");
20777 let same_t: i32 = if same_session { t as i32 } else { 0 };
20778 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
20779 let cfg = LaunchConfig {
20780 grid_dim: ((nh_q + nh_k) as u32, 1, t as u32),
20781 block_dim: (128, 1, 1),
20782 shared_mem_bytes: 0,
20783 };
20784 let (kvk, kvv) = (kv_dim_k as i32, kv_dim_v as i32);
20785 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
20786 let (hd, nd, nq, nk) = (head_dim as i32, n_dims as i32, nh_q as i32, nh_k as i32);
20787 let null: u64 = 0;
20788 let __s_b = self.gpu.stream();
20789 let mut b = __s_b.launch_builder(&f);
20790 b.arg(q_raw_t)
20791 .arg(k_raw_t)
20792 .arg(v_raw_t)
20793 .arg(qw)
20794 .arg(kw)
20795 .arg(q_out_t)
20796 .arg(k_out_t)
20797 .arg(tab)
20798 .arg(pos_t)
20799 .arg(&same_t)
20800 .arg(&kvk)
20801 .arg(&kvv)
20802 .arg(&ktb)
20803 .arg(&vtb)
20804 .arg(&hd)
20805 .arg(&nd)
20806 .arg(&nq)
20807 .arg(&nk)
20808 .arg(&eps)
20809 .arg(&theta_scale)
20810 .arg(&freq_scale);
20811 match ff {
20812 Some(freqs) => {
20813 b.arg(freqs);
20814 }
20815 None => {
20816 b.arg(&null);
20817 }
20818 }
20819 unsafe {
20820 b.launch(cfg)?;
20821 }
20822 Ok(())
20823 }
20824
20825 #[allow(clippy::too_many_arguments)] pub fn qk_norm_rope_append_inc_dcw(
20827 &self,
20828 q_raw: &CudaSlice<f32>,
20829 k_raw: &CudaSlice<f32>,
20830 v_raw: &CudaSlice<f32>,
20831 qw: &CudaSlice<f32>,
20832 kw: &CudaSlice<f32>,
20833 q_out: &mut CudaSlice<f32>,
20834 k_out: &mut CudaSlice<f32>,
20835 pos: &CudaSlice<i32>,
20836 k_plane: &mut CudaSlice<u8>,
20837 v_plane: &mut CudaSlice<u8>,
20838 len_dev: &CudaSlice<i32>,
20841 base_dev: Option<&CudaSlice<i32>>,
20842 done_ctr: &mut CudaSlice<u32>,
20843 kv_dim_k: usize,
20844 kv_dim_v: usize,
20845 k_tok_bytes: usize,
20846 v_tok_bytes: usize,
20847 head_dim: usize,
20848 n_dims: usize,
20849 nh_q: usize,
20850 nh_k: usize,
20851 eps: f32,
20852 freq_base: f32,
20853 freq_scale: f32,
20854 ff: Option<&CudaSlice<f32>>,
20855 ) -> Result<(), Box<dyn std::error::Error>> {
20856 if head_dim != 128
20857 || kv_dim_v != kv_dim_k
20858 || kv_dim_k != nh_k * head_dim
20859 || q_raw.len() < nh_q * head_dim
20860 || k_raw.len() < nh_k * head_dim
20861 || v_raw.len() < kv_dim_v
20862 || q_out.len() < nh_q * head_dim
20863 || k_out.len() < nh_k * head_dim
20864 || pos.is_empty()
20865 || done_ctr.is_empty()
20866 {
20867 return Err(format!(
20868 "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}"
20869 )
20870 .into());
20871 }
20872 let f = self.func("qk_norm_rope_append_inc_dcw");
20873 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
20874 let cfg = LaunchConfig {
20875 grid_dim: ((nh_q + nh_k) as u32, 1, 1),
20876 block_dim: (128, 1, 1),
20877 shared_mem_bytes: 0,
20878 };
20879 let (kvk, kvv) = (kv_dim_k as i32, kv_dim_v as i32);
20880 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
20881 let (hd, nd, nq) = (head_dim as i32, n_dims as i32, nh_q as i32);
20882 let null: u64 = 0;
20883 let __s_b = self.gpu.stream();
20884 let mut b = __s_b.launch_builder(&f);
20885 b.arg(q_raw)
20886 .arg(k_raw)
20887 .arg(v_raw)
20888 .arg(qw)
20889 .arg(kw)
20890 .arg(q_out)
20891 .arg(k_out)
20892 .arg(pos)
20893 .arg(&mut *k_plane)
20894 .arg(&mut *v_plane)
20895 .arg(len_dev);
20896 match base_dev {
20897 Some(base) => {
20898 b.arg(base);
20899 }
20900 None => {
20901 b.arg(&null);
20902 }
20903 }
20904 b.arg(&mut *done_ctr)
20905 .arg(&kvk)
20906 .arg(&kvv)
20907 .arg(&ktb)
20908 .arg(&vtb)
20909 .arg(&hd)
20910 .arg(&nd)
20911 .arg(&nq)
20912 .arg(&eps)
20913 .arg(&theta_scale)
20914 .arg(&freq_scale);
20915 match ff {
20916 Some(freqs) => {
20917 b.arg(freqs);
20918 }
20919 None => {
20920 b.arg(&null);
20921 }
20922 }
20923 unsafe {
20924 b.launch(cfg)?;
20925 }
20926 Ok(())
20927 }
20928
20929 #[allow(clippy::too_many_arguments)] pub fn qk_norm_rope_into(
20931 &self,
20932 q_raw: &CudaSlice<f32>,
20933 k_raw: &CudaSlice<f32>,
20934 qw: &CudaSlice<f32>,
20935 kw: &CudaSlice<f32>,
20936 q_out: &mut CudaSlice<f32>,
20937 k_out: &mut CudaSlice<f32>,
20938 pos: &CudaSlice<i32>,
20939 head_dim: usize,
20940 n_dims: usize,
20941 nh_q: usize,
20942 nh_k: usize,
20943 eps: f32,
20944 freq_base: f32,
20945 freq_scale: f32,
20946 ff: Option<&CudaSlice<f32>>,
20947 ) -> Result<(), Box<dyn std::error::Error>> {
20948 if head_dim > 512
20949 || q_raw.len() < nh_q * head_dim
20950 || k_raw.len() < nh_k * head_dim
20951 || q_out.len() < nh_q * head_dim
20952 || k_out.len() < nh_k * head_dim
20953 || qw.len() < head_dim
20954 || kw.len() < head_dim
20955 || pos.is_empty()
20956 {
20957 return Err(format!(
20958 "qk_norm_rope geometry head_dim={head_dim} nh_q={nh_q} nh_k={nh_k}"
20959 )
20960 .into());
20961 }
20962 let f = self.func("qk_norm_rope_f32");
20963 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
20964 let cfg = LaunchConfig {
20965 grid_dim: ((nh_q + nh_k) as u32, 1, 1),
20966 block_dim: (128, 1, 1),
20967 shared_mem_bytes: 0,
20968 };
20969 let (hd, nd, nq) = (head_dim as i32, n_dims as i32, nh_q as i32);
20970 let __s_b = self.gpu.stream();
20971 let mut b = __s_b.launch_builder(&f);
20972 b.arg(q_raw)
20973 .arg(k_raw)
20974 .arg(qw)
20975 .arg(kw)
20976 .arg(q_out)
20977 .arg(k_out)
20978 .arg(pos)
20979 .arg(&hd)
20980 .arg(&nd)
20981 .arg(&nq)
20982 .arg(&eps)
20983 .arg(&theta_scale)
20984 .arg(&freq_scale);
20985 match ff {
20986 Some(ffv) => {
20987 b.arg(ffv);
20988 unsafe {
20989 b.launch(cfg)?;
20990 }
20991 }
20992 None => {
20993 let null: u64 = 0;
20994 b.arg(&null);
20995 unsafe {
20996 b.launch(cfg)?;
20997 }
20998 }
20999 }
21000 Ok(())
21001 }
21002
21003 #[allow(clippy::too_many_arguments)]
21006 pub fn matvec_f32_b4_into(
21007 &self,
21008 w: [&CudaSlice<f32>; 4],
21009 x: &CudaSlice<f32>,
21010 y: &mut CudaSlice<f32>,
21011 block_cols: usize,
21012 out_f: usize,
21013 ) -> Result<(), Box<dyn std::error::Error>> {
21014 if !block_cols.is_multiple_of(4)
21015 || x.len() < 4 * block_cols
21016 || y.len() < out_f
21017 || w.iter().any(|w| w.len() != out_f * block_cols)
21018 {
21019 return Err(format!(
21020 "matvec_f32_b4 geometry block_cols={block_cols} out={out_f} x={}",
21021 x.len()
21022 )
21023 .into());
21024 }
21025 let f = self.func("matvec_f32_b4");
21026 let cfg = LaunchConfig {
21027 grid_dim: (out_f as u32, 1, 1),
21028 block_dim: (128, 1, 1),
21029 shared_mem_bytes: 0,
21030 };
21031 let (bc, of) = (block_cols as i32, out_f as i32);
21032 let __s_b = self.gpu.stream();
21033 let mut b = __s_b.launch_builder(&f);
21034 b.arg(w[0])
21035 .arg(w[1])
21036 .arg(w[2])
21037 .arg(w[3])
21038 .arg(x)
21039 .arg(y)
21040 .arg(&bc)
21041 .arg(&of);
21042 unsafe {
21043 b.launch(cfg)?;
21044 }
21045 Ok(())
21046 }
21047
21048 pub fn axpy_rows_seq_into(
21051 &self,
21052 x: &CudaSlice<f32>,
21053 w: &CudaSlice<f32>,
21054 y: &mut CudaSlice<f32>,
21055 width: usize,
21056 n_rows: usize,
21057 ) -> Result<(), Box<dyn std::error::Error>> {
21058 if x.len() < n_rows * width || w.len() < n_rows || y.len() < width {
21059 return Err(format!(
21060 "axpy_rows_seq geometry x={} w={} y={} width={width} rows={n_rows}",
21061 x.len(),
21062 w.len(),
21063 y.len()
21064 )
21065 .into());
21066 }
21067 let f = self.func("axpy_rows_seq_f32");
21068 let cfg = LaunchConfig::for_num_elems(width as u32);
21069 let (wi, nr) = (width as i32, n_rows as i32);
21070 let __s_b = self.gpu.stream();
21071 let mut b = __s_b.launch_builder(&f);
21072 b.arg(x).arg(w).arg(y).arg(&wi).arg(&nr);
21073 unsafe {
21074 b.launch(cfg)?;
21075 }
21076 Ok(())
21077 }
21078
21079 pub fn axpy_rows_seq_tokens_into(
21082 &self,
21083 x: &CudaSlice<f32>,
21084 w: &CudaSlice<f32>,
21085 y: &mut CudaSlice<f32>,
21086 width: usize,
21087 slots: usize,
21088 tokens: usize,
21089 ) -> Result<(), Box<dyn std::error::Error>> {
21090 let rows = slots
21091 .checked_mul(tokens)
21092 .ok_or("axpy_rows_seq_tokens row count overflow")?;
21093 if x.len() < rows * width || w.len() < rows || y.len() < tokens * width {
21094 return Err(format!(
21095 "axpy_rows_seq_tokens geometry x={} w={} y={} width={width} \
21096 slots={slots} tokens={tokens}",
21097 x.len(),
21098 w.len(),
21099 y.len()
21100 )
21101 .into());
21102 }
21103 let f = self.func("axpy_rows_seq_tokens_f32");
21104 let block = 256u32;
21105 let cfg = LaunchConfig {
21106 grid_dim: ((width as u32).div_ceil(block), tokens as u32, 1),
21107 block_dim: (block, 1, 1),
21108 shared_mem_bytes: 0,
21109 };
21110 let (wi, sl, tk) = (width as i32, slots as i32, tokens as i32);
21111 let __s_b = self.gpu.stream();
21112 let mut b = __s_b.launch_builder(&f);
21113 b.arg(x).arg(w).arg(y).arg(&wi).arg(&sl).arg(&tk);
21114 unsafe {
21115 b.launch(cfg)?;
21116 }
21117 Ok(())
21118 }
21119
21120 #[allow(clippy::too_many_arguments)]
21124 pub fn axpy_rows_seq_md_off_into(
21125 &self,
21126 x: &CudaSlice<f32>,
21127 w_route: &CudaSlice<f32>,
21128 md: &CudaSlice<f32>,
21129 sel: &CudaSlice<i32>,
21130 y: &mut CudaSlice<f32>,
21131 width: usize,
21132 n_rows: usize,
21133 row0: usize,
21134 ) -> Result<(), Box<dyn std::error::Error>> {
21135 if x.len() < (row0 + n_rows) * width
21136 || w_route.len() < row0 + n_rows
21137 || sel.len() < row0 + n_rows
21138 || y.len() < width
21139 {
21140 return Err(format!(
21141 "axpy_rows_seq_md_off geometry x={} w={} sel={} y={} width={width} \
21142 rows={n_rows} row0={row0}",
21143 x.len(),
21144 w_route.len(),
21145 sel.len(),
21146 y.len()
21147 )
21148 .into());
21149 }
21150 let f = self.func("axpy_rows_seq_md_off_f32");
21151 let cfg = LaunchConfig::for_num_elems(width as u32);
21152 let (wi, nr, r0) = (width as i32, n_rows as i32, row0 as i32);
21153 let __s_b = self.gpu.stream();
21154 let mut b = __s_b.launch_builder(&f);
21155 b.arg(x)
21156 .arg(w_route)
21157 .arg(md)
21158 .arg(sel)
21159 .arg(y)
21160 .arg(&wi)
21161 .arg(&nr)
21162 .arg(&r0);
21163 unsafe {
21164 b.launch(cfg)?;
21165 }
21166 Ok(())
21167 }
21168
21169 #[allow(clippy::too_many_arguments)]
21172 pub fn axpy_rows_seq_md_into(
21173 &self,
21174 x: &CudaSlice<f32>,
21175 w_route: &CudaSlice<f32>,
21176 md: &CudaSlice<f32>,
21177 sel: &CudaSlice<i32>,
21178 y: &mut CudaSlice<f32>,
21179 width: usize,
21180 n_rows: usize,
21181 ) -> Result<(), Box<dyn std::error::Error>> {
21182 if x.len() < n_rows * width
21183 || w_route.len() < n_rows
21184 || sel.len() < n_rows
21185 || y.len() < width
21186 {
21187 return Err(format!(
21188 "axpy_rows_seq_md geometry x={} w={} sel={} y={} width={width} rows={n_rows}",
21189 x.len(),
21190 w_route.len(),
21191 sel.len(),
21192 y.len()
21193 )
21194 .into());
21195 }
21196 let f = self.func("axpy_rows_seq_md_f32");
21197 let cfg = LaunchConfig::for_num_elems(width as u32);
21198 let (wi, nr) = (width as i32, n_rows as i32);
21199 let __s_b = self.gpu.stream();
21200 let mut b = __s_b.launch_builder(&f);
21201 b.arg(x)
21202 .arg(w_route)
21203 .arg(md)
21204 .arg(sel)
21205 .arg(y)
21206 .arg(&wi)
21207 .arg(&nr);
21208 unsafe {
21209 b.launch(cfg)?;
21210 }
21211 Ok(())
21212 }
21213
21214 #[allow(clippy::too_many_arguments)]
21216 #[allow(clippy::too_many_arguments)]
21220 pub fn matvec_bf16_qkvg_tcol_into(
21221 &self,
21222 wq: &CudaSlice<u8>,
21223 wk: &CudaSlice<u8>,
21224 wv: &CudaSlice<u8>,
21225 wg: &CudaSlice<u8>,
21226 x_t: &CudaSlice<f32>,
21227 yq: &mut CudaSlice<f32>,
21228 yk: &mut CudaSlice<f32>,
21229 yv: &mut CudaSlice<f32>,
21230 yg: &mut CudaSlice<f32>,
21231 in_f: usize,
21232 out_q: usize,
21233 out_kv: usize,
21234 out_g: usize,
21235 t: usize,
21236 ) -> Result<(), Box<dyn std::error::Error>> {
21237 if t == 0
21238 || t > 8
21239 || !in_f.is_multiple_of(8)
21240 || x_t.len() < t * in_f
21241 || yq.len() < t * out_q
21242 || yk.len() < t * out_kv
21243 || yv.len() < t * out_kv
21244 || (out_g > 0 && yg.len() < t * out_g)
21245 {
21246 return Err("matvec_bf16_qkvg_tcol geometry".into());
21247 }
21248 let grid = out_q + 2 * out_kv + out_g;
21249 let cfg = LaunchConfig {
21250 grid_dim: (grid as u32, 1, 1),
21251 block_dim: (mmv_block(), 1, 1),
21252 shared_mem_bytes: 0,
21253 };
21254 let (ini, oq, okv, og, ti) = (
21255 in_f as i32,
21256 out_q as i32,
21257 out_kv as i32,
21258 out_g as i32,
21259 t as i32,
21260 );
21261 let __s_b = self.gpu.stream();
21262 let f = self.func("matvec_bf16_qkvg_tcol");
21268 let mut b = __s_b.launch_builder(&f);
21269 b.arg(wq)
21270 .arg(wk)
21271 .arg(wv)
21272 .arg(wg)
21273 .arg(x_t)
21274 .arg(yq)
21275 .arg(yk)
21276 .arg(yv)
21277 .arg(yg)
21278 .arg(&ini)
21279 .arg(&oq)
21280 .arg(&okv)
21281 .arg(&og)
21282 .arg(&ti);
21283 unsafe {
21284 b.launch(cfg)?;
21285 }
21286 Ok(())
21287 }
21288
21289 #[allow(clippy::too_many_arguments)] pub fn matvec_bf16_qkvg_into(
21291 &self,
21292 wq: &CudaSlice<u8>,
21293 wk: &CudaSlice<u8>,
21294 wv: &CudaSlice<u8>,
21295 wg: &CudaSlice<u8>,
21296 x: &CudaSlice<f32>,
21297 yq: &mut CudaSlice<f32>,
21298 yk: &mut CudaSlice<f32>,
21299 yv: &mut CudaSlice<f32>,
21300 yg: &mut CudaSlice<f32>,
21301 in_f: usize,
21302 out_q: usize,
21303 out_kv: usize,
21304 out_g: usize,
21305 ) -> Result<(), Box<dyn std::error::Error>> {
21306 if !in_f.is_multiple_of(8)
21307 || wq.len() != out_q * in_f * 2
21308 || wk.len() != out_kv * in_f * 2
21309 || wv.len() != out_kv * in_f * 2
21310 || wg.len() < out_g * in_f * 2
21311 || x.len() < in_f
21312 || yq.len() < out_q
21313 || yk.len() < out_kv
21314 || yv.len() < out_kv
21315 || (out_g > 0 && yg.len() < out_g)
21316 {
21317 return Err(format!(
21318 "fused bf16 QKV geometry in={in_f} out_q={out_q} out_kv={out_kv} out_g={out_g}"
21319 )
21320 .into());
21321 }
21322 let f = self.func("matvec_bf16_qkvg");
21323 let cfg = LaunchConfig {
21324 grid_dim: ((out_q + 2 * out_kv + out_g) as u32, 1, 1),
21325 block_dim: (mmv_block(), 1, 1),
21326 shared_mem_bytes: 0,
21327 };
21328 let (inf, oq, okv, og) = (in_f as i32, out_q as i32, out_kv as i32, out_g as i32);
21329 let __s_b = self.gpu.stream();
21330 let mut b = __s_b.launch_builder(&f);
21331 b.arg(wq)
21332 .arg(wk)
21333 .arg(wv)
21334 .arg(wg)
21335 .arg(x)
21336 .arg(yq)
21337 .arg(yk)
21338 .arg(yv)
21339 .arg(yg)
21340 .arg(&inf)
21341 .arg(&oq)
21342 .arg(&okv)
21343 .arg(&og);
21344 unsafe {
21345 b.launch(cfg)?;
21346 }
21347 Ok(())
21348 }
21349
21350 pub fn matvec_bf16_b4_into(
21352 &self,
21353 w: [&CudaSlice<u8>; 4],
21354 x: &CudaSlice<f32>,
21355 y: &mut CudaSlice<f32>,
21356 block_cols: usize,
21357 out_f: usize,
21358 ) -> Result<(), Box<dyn std::error::Error>> {
21359 if !block_cols.is_multiple_of(8)
21360 || x.len() < 4 * block_cols
21361 || y.len() < out_f
21362 || w.iter().any(|w| w.len() != out_f * block_cols * 2)
21363 {
21364 return Err(format!(
21365 "bf16 b4 geometry block_cols={block_cols} out={out_f} x={}",
21366 x.len()
21367 )
21368 .into());
21369 }
21370 static B4_X2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
21373 let x2 = *B4_X2.get_or_init(|| std::env::var("MEMRA_B4_X2").as_deref() == Ok("1"));
21374 let f = self.func(if x2 {
21375 "matvec_bf16_b4_x2"
21376 } else {
21377 "matvec_bf16_b4"
21378 });
21379 let grid = if x2 { out_f.div_ceil(2) } else { out_f };
21380 let cfg = LaunchConfig {
21381 grid_dim: (grid as u32, 1, 1),
21382 block_dim: (mmv_block(), 1, 1),
21383 shared_mem_bytes: 0,
21384 };
21385 let (bc, of) = (block_cols as i32, out_f as i32);
21386 let __s_b = self.gpu.stream();
21387 let mut b = __s_b.launch_builder(&f);
21388 b.arg(w[0])
21389 .arg(w[1])
21390 .arg(w[2])
21391 .arg(w[3])
21392 .arg(x)
21393 .arg(y)
21394 .arg(&bc)
21395 .arg(&of);
21396 unsafe {
21397 b.launch(cfg)?;
21398 }
21399 Ok(())
21400 }
21401
21402 pub fn matvec_bf16_b4_tcol_into(
21408 &self,
21409 w: [&CudaSlice<u8>; 4],
21410 x_t: &CudaSlice<f32>,
21411 y_t: &mut CudaSlice<f32>,
21412 block_cols: usize,
21413 out_f: usize,
21414 t: usize,
21415 ) -> Result<(), Box<dyn std::error::Error>> {
21416 if !block_cols.is_multiple_of(8)
21417 || t == 0
21418 || t > 8
21419 || x_t.len() < t * 4 * block_cols
21420 || y_t.len() < t * out_f
21421 || w.iter().any(|w| w.len() != out_f * block_cols * 2)
21422 {
21423 return Err(format!(
21424 "bf16 b4 tcol geometry block_cols={block_cols} out={out_f} t={t} x={}",
21425 x_t.len()
21426 )
21427 .into());
21428 }
21429 if std::env::var("MEMRA_B4_X2").as_deref() == Ok("1") {
21430 return Err(
21431 "b4 tcol verify is qualified against the plain b4 kernel only \
21432 (MEMRA_B4_X2=1 is a different t=1 program)"
21433 .into(),
21434 );
21435 }
21436 let cfg = LaunchConfig {
21440 grid_dim: (out_f as u32, 1, 1),
21441 block_dim: (mmv_block(), 1, 1),
21442 shared_mem_bytes: 0,
21443 };
21444 let (bc, of, ti) = (block_cols as i32, out_f as i32, t as i32);
21445 let __s_b = self.gpu.stream();
21446 let f = self.func("matvec_bf16_b4_tcol");
21447 let mut b = __s_b.launch_builder(&f);
21448 b.arg(w[0])
21449 .arg(w[1])
21450 .arg(w[2])
21451 .arg(w[3])
21452 .arg(x_t)
21453 .arg(y_t)
21454 .arg(&bc)
21455 .arg(&of)
21456 .arg(&ti);
21457 unsafe {
21458 b.launch(cfg)?;
21459 }
21460 Ok(())
21461 }
21462
21463 pub fn q8_0_row_bytes(in_f: usize) -> usize {
21466 in_f / 32 * 34
21467 }
21468
21469 pub fn encode_q8_0_from_bf16(
21473 &self,
21474 w_bf16: &CudaSlice<u8>,
21475 out: &mut CudaSlice<u8>,
21476 in_f: usize,
21477 out_f: usize,
21478 ) -> Result<(), Box<dyn std::error::Error>> {
21479 if !in_f.is_multiple_of(32)
21480 || w_bf16.len() < in_f * out_f * 2
21481 || out.len() < out_f * Self::q8_0_row_bytes(in_f)
21482 {
21483 return Err(format!(
21484 "encode_q8_0_from_bf16 geometry in={in_f} out={out_f} src={} dst={}",
21485 w_bf16.len(),
21486 out.len()
21487 )
21488 .into());
21489 }
21490 let f = self.func("encode_q8_0_rows_from_bf16");
21491 const PAIRS_PER_BLOCK: u32 = 4;
21494 let pairs = (out_f * (in_f / 32)) as u64;
21495 let cfg = LaunchConfig {
21496 grid_dim: ((pairs.div_ceil(PAIRS_PER_BLOCK as u64)) as u32, 1, 1),
21497 block_dim: (32, PAIRS_PER_BLOCK, 1),
21498 shared_mem_bytes: 0,
21499 };
21500 let (ini, outi) = (in_f as i32, out_f as i32);
21501 let __s_b = self.gpu.stream();
21502 let mut b = __s_b.launch_builder(&f);
21503 b.arg(w_bf16).arg(out).arg(&ini).arg(&outi);
21504 unsafe {
21505 b.launch(cfg)?;
21506 }
21507 Ok(())
21508 }
21509
21510 pub fn encode_q8_0_from_bf16_view(
21514 &self,
21515 w_bf16: &cudarc::driver::CudaView<'_, u8>,
21516 out: &mut CudaSlice<u8>,
21517 in_f: usize,
21518 out_f: usize,
21519 ) -> Result<(), Box<dyn std::error::Error>> {
21520 if !in_f.is_multiple_of(32)
21521 || w_bf16.len() < in_f * out_f * 2
21522 || out.len() < out_f * Self::q8_0_row_bytes(in_f)
21523 {
21524 return Err(format!(
21525 "encode_q8_0_from_bf16_view geometry in={in_f} out={out_f} src={} dst={}",
21526 w_bf16.len(),
21527 out.len()
21528 )
21529 .into());
21530 }
21531 let f = self.func("encode_q8_0_rows_from_bf16");
21532 const PAIRS_PER_BLOCK: u32 = 4;
21533 let pairs = (out_f * (in_f / 32)) as u64;
21534 let cfg = LaunchConfig {
21535 grid_dim: ((pairs.div_ceil(PAIRS_PER_BLOCK as u64)) as u32, 1, 1),
21536 block_dim: (32, PAIRS_PER_BLOCK, 1),
21537 shared_mem_bytes: 0,
21538 };
21539 let (ini, outi) = (in_f as i32, out_f as i32);
21540 let __s_b = self.gpu.stream();
21541 let mut b = __s_b.launch_builder(&f);
21542 b.arg(w_bf16).arg(out).arg(&ini).arg(&outi);
21543 unsafe {
21544 b.launch(cfg)?;
21545 }
21546 Ok(())
21547 }
21548
21549 #[allow(clippy::too_many_arguments)]
21554 pub fn qmatvec_q8_0_qkv_rp_into(
21555 &self,
21556 wq: &CudaSlice<u8>,
21557 wk: &CudaSlice<u8>,
21558 wv: &CudaSlice<u8>,
21559 aq: &CudaSlice<i8>,
21560 ad: &CudaSlice<f32>,
21561 yq: &mut CudaSlice<f32>,
21562 yk: &mut CudaSlice<f32>,
21563 yv: &mut CudaSlice<f32>,
21564 in_f: usize,
21565 out_q: usize,
21566 out_kv: usize,
21567 ) -> Result<(), Box<dyn std::error::Error>> {
21568 const ROWS_PER_BLOCK: u32 = 4; let rows = out_q + 2 * out_kv;
21570 let nblk = in_f / 32;
21571 if !in_f.is_multiple_of(32)
21572 || aq.len() < in_f
21573 || ad.len() < nblk
21574 || yq.len() < out_q
21575 || yk.len() < out_kv
21576 || yv.len() < out_kv
21577 || wq.len() < out_q * nblk * 34
21578 || wk.len() < out_kv * nblk * 34
21579 || wv.len() < out_kv * nblk * 34
21580 {
21581 return Err(
21582 format!("q8_0 qkv rp geometry in={in_f} out_q={out_q} out_kv={out_kv}").into(),
21583 );
21584 }
21585 let f = self.func("qmatvec_q8_0_qkv_rp");
21586 let cfg = LaunchConfig {
21587 grid_dim: ((rows as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
21588 block_dim: (32, ROWS_PER_BLOCK, 1),
21589 shared_mem_bytes: 0,
21590 };
21591 let (ini, oq, okv) = (in_f as i32, out_q as i32, out_kv as i32);
21592 let __s_b = self.gpu.stream();
21593 let mut b = __s_b.launch_builder(&f);
21594 b.arg(wq)
21595 .arg(wk)
21596 .arg(wv)
21597 .arg(aq)
21598 .arg(ad)
21599 .arg(yq)
21600 .arg(yk)
21601 .arg(yv)
21602 .arg(&ini)
21603 .arg(&oq)
21604 .arg(&okv);
21605 unsafe {
21606 b.launch(cfg)?;
21607 }
21608 Ok(())
21609 }
21610
21611 #[allow(clippy::too_many_arguments)]
21615 pub fn qmatvec_q8_0_b4_rp_into(
21616 &self,
21617 w: [&CudaSlice<u8>; 4],
21618 aq: &CudaSlice<i8>,
21619 ad: &CudaSlice<f32>,
21620 y: &mut CudaSlice<f32>,
21621 block_cols: usize,
21622 out_f: usize,
21623 ) -> Result<(), Box<dyn std::error::Error>> {
21624 const ROWS_PER_BLOCK: u32 = 4; let nblk = block_cols / 32;
21626 if !block_cols.is_multiple_of(32)
21627 || aq.len() < 4 * block_cols
21628 || ad.len() < 4 * nblk
21629 || y.len() < out_f
21630 || w.iter().any(|p| p.len() < out_f * nblk * 34)
21631 {
21632 return Err(format!("q8_0 b4 rp geometry block_cols={block_cols} out={out_f}").into());
21633 }
21634 let f = self.func("qmatvec_q8_0_b4_rp");
21635 let cfg = LaunchConfig {
21636 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
21637 block_dim: (32, ROWS_PER_BLOCK, 1),
21638 shared_mem_bytes: 0,
21639 };
21640 let (bc, of) = (block_cols as i32, out_f as i32);
21641 let __s_b = self.gpu.stream();
21642 let mut b = __s_b.launch_builder(&f);
21643 b.arg(w[0])
21644 .arg(w[1])
21645 .arg(w[2])
21646 .arg(w[3])
21647 .arg(aq)
21648 .arg(ad)
21649 .arg(y)
21650 .arg(&bc)
21651 .arg(&of);
21652 unsafe {
21653 b.launch(cfg)?;
21654 }
21655 Ok(())
21656 }
21657
21658 #[allow(clippy::map_entry)] fn matvec_bf16_via_q8_mirror_t(
21662 &self,
21663 data: &CudaSlice<u8>,
21664 x: &CudaSlice<f32>,
21665 y: &mut CudaSlice<f32>,
21666 in_f: usize,
21667 out_f: usize,
21668 t: usize,
21669 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
21670 use cudarc::driver::DevicePtr;
21671 let key = {
21672 let s = self.gpu.stream();
21673 let (p, _g) = data.device_ptr(&s);
21674 (p, in_f as u32, out_f as u32)
21675 };
21676 {
21677 let mut mirrors = self
21678 .w8_mirrors
21679 .lock()
21680 .map_err(|_| "w8 mirror map is poisoned")?;
21681 if !mirrors.contains_key(&key) {
21682 let mut interleaved = self.alloc_u8_uninit(out_f * Self::q8_0_row_bytes(in_f))?;
21683 self.encode_q8_0_from_bf16(data, &mut interleaved, in_f, out_f)?;
21684 let planar = self.build_q8_rp4_raw(&interleaved, in_f, out_f)?;
21685 mirrors.insert(key, planar);
21686 }
21687 }
21688 let nblk = in_f / 32;
21689 let akey = in_f * 64 + t.min(32);
21691 {
21692 let mut act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
21693 if !act.contains_key(&akey) {
21694 let aq = self.alloc_i8_uninit(32 * in_f)?;
21695 let ad = self.alloc_uninit::<f32>(32 * nblk)?;
21696 act.insert(akey, (aq, ad));
21697 }
21698 let (aq, ad) = act.get_mut(&akey).expect("just inserted");
21699 self.quantize_q8_1_into(x, t, in_f, aq, ad)?;
21700 }
21701 let mirrors = self
21702 .w8_mirrors
21703 .lock()
21704 .map_err(|_| "w8 mirror map is poisoned")?;
21705 let act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
21706 let mirror = mirrors.get(&key).expect("built above");
21707 let (aq, ad) = act.get(&akey).expect("built above");
21708 const ROWS_PER_BLOCK: u32 = 4;
21709 let (ini, of) = (in_f as i32, out_f as i32);
21710 if q8t_wonce_on() && t <= 32 {
21714 let f = self.func(if t <= 8 {
21715 "qmatvec_q8_0_rows_tw"
21716 } else {
21717 "qmatvec_q8_0_rows_tw32"
21718 });
21719 let cfg = LaunchConfig {
21720 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
21721 block_dim: (32, ROWS_PER_BLOCK, 1),
21722 shared_mem_bytes: 0,
21723 };
21724 let ti = t as i32;
21725 let __s_b = self.gpu.stream();
21726 let mut b = __s_b.launch_builder(&f);
21727 b.arg(mirror)
21728 .arg(aq)
21729 .arg(ad)
21730 .arg(&mut *y)
21731 .arg(&ini)
21732 .arg(&of)
21733 .arg(&ti);
21734 unsafe {
21735 b.launch(cfg)?;
21736 }
21737 return Ok(Some(()));
21738 }
21739 let f = self.func("qmatvec_q8_0_rows_t");
21740 let cfg = LaunchConfig {
21741 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), t as u32, 1),
21742 block_dim: (32, ROWS_PER_BLOCK, 1),
21743 shared_mem_bytes: 0,
21744 };
21745 let __s_b = self.gpu.stream();
21746 let mut b = __s_b.launch_builder(&f);
21747 b.arg(mirror)
21748 .arg(aq)
21749 .arg(ad)
21750 .arg(&mut *y)
21751 .arg(&ini)
21752 .arg(&of);
21753 unsafe {
21754 b.launch(cfg)?;
21755 }
21756 Ok(Some(()))
21757 }
21758
21759 #[allow(clippy::map_entry)] fn matvec_bf16_via_q8_mirror(
21763 &self,
21764 data: &CudaSlice<u8>,
21765 x: &CudaSlice<f32>,
21766 y: &mut CudaSlice<f32>,
21767 in_f: usize,
21768 out_f: usize,
21769 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
21770 use cudarc::driver::DevicePtr;
21771 let key = {
21772 let s = self.gpu.stream();
21773 let (p, _g) = data.device_ptr(&s);
21774 (p, in_f as u32, out_f as u32)
21775 };
21776 {
21777 let mut mirrors = self
21778 .w8_mirrors
21779 .lock()
21780 .map_err(|_| "w8 mirror map is poisoned")?;
21781 if !mirrors.contains_key(&key) {
21782 let mut interleaved = self.alloc_u8_uninit(out_f * Self::q8_0_row_bytes(in_f))?;
21783 self.encode_q8_0_from_bf16(data, &mut interleaved, in_f, out_f)?;
21784 let planar = self.build_q8_rp4_raw(&interleaved, in_f, out_f)?;
21785 mirrors.insert(key, planar);
21786 if std::env::var("MEMRA_W8_TRACE").as_deref() == Ok("1") {
21792 eprintln!(
21793 "[w8-mirror] built in_f={in_f} out_f={out_f} mirrors={}",
21794 mirrors.len()
21795 );
21796 }
21797 }
21798 }
21799 let nblk = in_f / 32;
21800 {
21801 let mut act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
21802 if !act.contains_key(&in_f) {
21803 let aq = self.alloc_uninit::<i8>(in_f)?;
21804 let ad = self.alloc_uninit::<f32>(nblk)?;
21805 act.insert(in_f, (aq, ad));
21806 }
21807 let (aq, ad) = act.get_mut(&in_f).expect("just inserted");
21808 self.quantize_q8_1_into(x, 1, in_f, aq, ad)?;
21809 }
21810 let mirrors = self
21811 .w8_mirrors
21812 .lock()
21813 .map_err(|_| "w8 mirror map is poisoned")?;
21814 let act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
21815 let mirror = mirrors.get(&key).expect("built above");
21816 let (aq, ad) = act.get(&in_f).expect("built above");
21817 self.qmatvec_mmvq_into(
21818 mirror,
21819 aq,
21820 ad,
21821 1,
21822 in_f,
21823 out_f,
21824 QT_Q8_0,
21825 Self::q8_0_row_bytes(in_f),
21826 1.0,
21827 true,
21828 y,
21829 )?;
21830 Ok(Some(()))
21831 }
21832
21833 #[allow(clippy::too_many_arguments)]
21838 pub fn qmatvec_q8_0_qkv_rp_t_into(
21839 &self,
21840 wq: &CudaSlice<u8>,
21841 wk: &CudaSlice<u8>,
21842 wv: &CudaSlice<u8>,
21843 aq: &CudaSlice<i8>,
21844 ad: &CudaSlice<f32>,
21845 yq: &mut CudaSlice<f32>,
21846 yk: &mut CudaSlice<f32>,
21847 yv: &mut CudaSlice<f32>,
21848 in_f: usize,
21849 out_q: usize,
21850 out_kv: usize,
21851 t: usize,
21852 ) -> Result<(), Box<dyn std::error::Error>> {
21853 const ROWS_PER_BLOCK: u32 = 4;
21854 let rows = out_q + 2 * out_kv;
21855 let nblk = in_f / 32;
21856 if !in_f.is_multiple_of(32)
21857 || t == 0
21858 || aq.len() < t * in_f
21859 || ad.len() < t * nblk
21860 || yq.len() < t * out_q
21861 || yk.len() < t * out_kv
21862 || yv.len() < t * out_kv
21863 {
21864 return Err(format!("q8_0 qkv rp_t geometry in={in_f} t={t}").into());
21865 }
21866 let (ini, oq, okv) = (in_f as i32, out_q as i32, out_kv as i32);
21867 if q8t_wonce_on() && t <= 32 {
21871 let f = self.func(if t <= 8 {
21872 "qmatvec_q8_0_qkv_rp_tw"
21873 } else {
21874 "qmatvec_q8_0_qkv_rp_tw32"
21875 });
21876 let cfg = LaunchConfig {
21877 grid_dim: ((rows as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
21878 block_dim: (32, ROWS_PER_BLOCK, 1),
21879 shared_mem_bytes: 0,
21880 };
21881 let ti = t as i32;
21882 let __s_b = self.gpu.stream();
21883 let mut b = __s_b.launch_builder(&f);
21884 b.arg(wq)
21885 .arg(wk)
21886 .arg(wv)
21887 .arg(aq)
21888 .arg(ad)
21889 .arg(yq)
21890 .arg(yk)
21891 .arg(yv)
21892 .arg(&ini)
21893 .arg(&oq)
21894 .arg(&okv)
21895 .arg(&ti);
21896 unsafe {
21897 b.launch(cfg)?;
21898 }
21899 return Ok(());
21900 }
21901 let f = self.func("qmatvec_q8_0_qkv_rp_t");
21902 let cfg = LaunchConfig {
21903 grid_dim: ((rows as u32).div_ceil(ROWS_PER_BLOCK), t as u32, 1),
21904 block_dim: (32, ROWS_PER_BLOCK, 1),
21905 shared_mem_bytes: 0,
21906 };
21907 let __s_b = self.gpu.stream();
21908 let mut b = __s_b.launch_builder(&f);
21909 b.arg(wq)
21910 .arg(wk)
21911 .arg(wv)
21912 .arg(aq)
21913 .arg(ad)
21914 .arg(yq)
21915 .arg(yk)
21916 .arg(yv)
21917 .arg(&ini)
21918 .arg(&oq)
21919 .arg(&okv);
21920 unsafe {
21921 b.launch(cfg)?;
21922 }
21923 Ok(())
21924 }
21925
21926 #[allow(clippy::too_many_arguments)]
21929 pub fn qmatvec_q8_0_b4_rp_t_into(
21930 &self,
21931 w: [&CudaSlice<u8>; 4],
21932 aq: &CudaSlice<i8>,
21933 ad: &CudaSlice<f32>,
21934 y: &mut CudaSlice<f32>,
21935 block_cols: usize,
21936 out_f: usize,
21937 t: usize,
21938 ) -> Result<(), Box<dyn std::error::Error>> {
21939 const ROWS_PER_BLOCK: u32 = 4;
21940 let nblk = block_cols / 32;
21941 if !block_cols.is_multiple_of(32)
21942 || t == 0
21943 || aq.len() < t * 4 * block_cols
21944 || ad.len() < t * 4 * nblk
21945 || y.len() < t * out_f
21946 {
21947 return Err(format!("q8_0 b4 rp_t geometry cols={block_cols} t={t}").into());
21948 }
21949 let (bc, of) = (block_cols as i32, out_f as i32);
21950 if q8t_wonce_on() && t <= 32 {
21953 let f = self.func(if t <= 8 {
21954 "qmatvec_q8_0_b4_rp_tw"
21955 } else {
21956 "qmatvec_q8_0_b4_rp_tw32"
21957 });
21958 let cfg = LaunchConfig {
21959 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
21960 block_dim: (32, ROWS_PER_BLOCK, 1),
21961 shared_mem_bytes: 0,
21962 };
21963 let ti = t as i32;
21964 let __s_b = self.gpu.stream();
21965 let mut b = __s_b.launch_builder(&f);
21966 b.arg(w[0])
21967 .arg(w[1])
21968 .arg(w[2])
21969 .arg(w[3])
21970 .arg(aq)
21971 .arg(ad)
21972 .arg(y)
21973 .arg(&bc)
21974 .arg(&of)
21975 .arg(&ti);
21976 unsafe {
21977 b.launch(cfg)?;
21978 }
21979 return Ok(());
21980 }
21981 let f = self.func("qmatvec_q8_0_b4_rp_t");
21982 let cfg = LaunchConfig {
21983 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), t as u32, 1),
21984 block_dim: (32, ROWS_PER_BLOCK, 1),
21985 shared_mem_bytes: 0,
21986 };
21987 let __s_b = self.gpu.stream();
21988 let mut b = __s_b.launch_builder(&f);
21989 b.arg(w[0])
21990 .arg(w[1])
21991 .arg(w[2])
21992 .arg(w[3])
21993 .arg(aq)
21994 .arg(ad)
21995 .arg(y)
21996 .arg(&bc)
21997 .arg(&of);
21998 unsafe {
21999 b.launch(cfg)?;
22000 }
22001 Ok(())
22002 }
22003
22004 #[allow(clippy::map_entry)] fn matvec_bf16_view_via_q8_mirror(
22015 &self,
22016 data: &cudarc::driver::CudaView<'_, u8>,
22017 x: &CudaSlice<f32>,
22018 y: &mut CudaSlice<f32>,
22019 in_f: usize,
22020 out_f: usize,
22021 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
22022 use cudarc::driver::DevicePtr;
22023 let key = {
22024 let s = self.gpu.stream();
22025 let (p, _g) = data.device_ptr(&s);
22026 (p, in_f as u32, out_f as u32)
22027 };
22028 {
22029 let mut mirrors = self
22030 .w8_mirrors
22031 .lock()
22032 .map_err(|_| "w8 mirror map is poisoned")?;
22033 if !mirrors.contains_key(&key) {
22034 let mut interleaved = self.alloc_u8_uninit(out_f * Self::q8_0_row_bytes(in_f))?;
22035 self.encode_q8_0_from_bf16_view(data, &mut interleaved, in_f, out_f)?;
22036 let planar = self.build_q8_rp4_raw(&interleaved, in_f, out_f)?;
22037 mirrors.insert(key, planar);
22038 eprintln!("[w8-view] mirror built in_f={in_f} out_f={out_f}");
22042 }
22043 }
22044 let nblk = in_f / 32;
22045 {
22046 let mut act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
22047 if !act.contains_key(&in_f) {
22048 let aq = self.alloc_uninit::<i8>(in_f)?;
22049 let ad = self.alloc_uninit::<f32>(nblk)?;
22050 act.insert(in_f, (aq, ad));
22051 }
22052 let (aq, ad) = act.get_mut(&in_f).expect("just inserted");
22053 self.quantize_q8_1_into(x, 1, in_f, aq, ad)?;
22054 }
22055 let mirrors = self
22056 .w8_mirrors
22057 .lock()
22058 .map_err(|_| "w8 mirror map is poisoned")?;
22059 let act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
22060 let mirror = mirrors.get(&key).expect("built above");
22061 let (aq, ad) = act.get(&in_f).expect("built above");
22062 self.qmatvec_mmvq_into(
22063 mirror,
22064 aq,
22065 ad,
22066 1,
22067 in_f,
22068 out_f,
22069 QT_Q8_0,
22070 Self::q8_0_row_bytes(in_f),
22071 1.0,
22072 true,
22073 y,
22074 )?;
22075 Ok(Some(()))
22076 }
22077
22078 pub fn matvec_bf16_into(
22079 &self,
22080 data: &CudaSlice<u8>,
22081 x: &CudaSlice<f32>,
22082 y: &mut CudaSlice<f32>,
22083 in_f: usize,
22084 out_f: usize,
22085 ) -> Result<(), Box<dyn std::error::Error>> {
22086 if data.len() != in_f * out_f * 2
22087 || x.len() < in_f
22088 || !in_f.is_multiple_of(8)
22089 || y.len() < out_f
22090 {
22091 return Err(format!(
22092 "matvec_bf16_into geometry bytes={} x={} y={} in={in_f} out={out_f}",
22093 data.len(),
22094 x.len(),
22095 y.len()
22096 )
22097 .into());
22098 }
22099 if step_tp_w8_on()
22106 && w8_hybrid_on()
22107 && in_f.is_multiple_of(32)
22108 && out_f >= 64
22109 && let Some(()) = self.matvec_bf16_via_q8_mirror(data, x, y, in_f, out_f)?
22110 {
22111 return Ok(());
22112 }
22113 static X4: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
22117 let x4 = *X4.get_or_init(|| std::env::var("MEMRA_DOWN_X4").as_deref() == Ok("1"))
22118 && in_f <= 2048;
22119 if x4 {
22120 let f = self.func("matvec_bf16_f32acc_x4");
22121 let cfg = LaunchConfig {
22122 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
22123 block_dim: (mmv_block(), 1, 1),
22124 shared_mem_bytes: 0,
22125 };
22126 let (ini, outi) = (in_f as i32, out_f as i32);
22127 let __s_b = self.gpu.stream();
22128 let mut b = __s_b.launch_builder(&f);
22129 b.arg(data).arg(x).arg(y).arg(&ini).arg(&outi);
22130 unsafe {
22131 b.launch(cfg)?;
22132 }
22133 return Ok(());
22134 }
22135 let f = self.func("matvec_bf16_f32acc");
22136 let cfg = LaunchConfig {
22137 grid_dim: (out_f as u32, 1, 1),
22138 block_dim: (mmv_block(), 1, 1),
22139 shared_mem_bytes: 0,
22140 };
22141 let ini = in_f as i32;
22142 let __s_b = self.gpu.stream();
22143 let mut b = __s_b.launch_builder(&f);
22144 b.arg(data).arg(x).arg(y).arg(&ini);
22145 unsafe {
22146 b.launch(cfg)?;
22147 }
22148 Ok(())
22149 }
22150
22151 pub fn matvec_bf16_views_into(
22156 &self,
22157 data: &CudaSlice<u8>,
22158 x: &cudarc::driver::CudaView<'_, f32>,
22159 y: &mut cudarc::driver::CudaViewMut<'_, f32>,
22160 in_f: usize,
22161 out_f: usize,
22162 ) -> Result<(), Box<dyn std::error::Error>> {
22163 if data.len() != in_f * out_f * 2
22164 || x.len() < in_f
22165 || !in_f.is_multiple_of(8)
22166 || y.len() < out_f
22167 {
22168 return Err(format!(
22169 "matvec_bf16_views_into geometry bytes={} x={} y={} in={in_f} out={out_f}",
22170 data.len(),
22171 x.len(),
22172 y.len()
22173 )
22174 .into());
22175 }
22176 let f = self.func("matvec_bf16_f32acc");
22177 let cfg = LaunchConfig {
22178 grid_dim: (out_f as u32, 1, 1),
22179 block_dim: (mmv_block(), 1, 1),
22180 shared_mem_bytes: 0,
22181 };
22182 let ini = in_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);
22186 unsafe {
22187 b.launch(cfg)?;
22188 }
22189 Ok(())
22190 }
22191
22192 pub fn matvec_bf16_view_into(
22195 &self,
22196 data: &cudarc::driver::CudaView<'_, u8>,
22197 x: &CudaSlice<f32>,
22198 y: &mut CudaSlice<f32>,
22199 in_f: usize,
22200 out_f: usize,
22201 ) -> Result<(), Box<dyn std::error::Error>> {
22202 if data.len() != in_f * out_f * 2
22203 || x.len() < in_f
22204 || !in_f.is_multiple_of(8)
22205 || y.len() < out_f
22206 {
22207 return Err(format!(
22208 "matvec_bf16_view_into geometry bytes={} x={} y={} in={in_f} out={out_f}",
22209 data.len(),
22210 x.len(),
22211 y.len()
22212 )
22213 .into());
22214 }
22215 if w8_view_on()
22216 && step_tp_w8_on()
22217 && w8_hybrid_on()
22218 && in_f.is_multiple_of(32)
22219 && out_f >= 64
22220 && let Some(()) = self.matvec_bf16_view_via_q8_mirror(data, x, y, in_f, out_f)?
22221 {
22222 return Ok(());
22223 }
22224 let f = self.func("matvec_bf16_f32acc");
22225 let cfg = LaunchConfig {
22226 grid_dim: (out_f as u32, 1, 1),
22227 block_dim: (mmv_block(), 1, 1),
22228 shared_mem_bytes: 0,
22229 };
22230 let ini = in_f as i32;
22231 let __s_b = self.gpu.stream();
22232 let mut b = __s_b.launch_builder(&f);
22233 b.arg(data).arg(x).arg(y).arg(&ini);
22234 unsafe {
22235 b.launch(cfg)?;
22236 }
22237 Ok(())
22238 }
22239
22240 pub fn matvec_bf16_raw_out(
22243 &self,
22244 w: &CudaSlice<u8>,
22245 x: &CudaSlice<f32>,
22246 y_raw: u64,
22247 in_f: usize,
22248 out_f: usize,
22249 ) -> Result<(), Box<dyn std::error::Error>> {
22250 if w.len() != in_f * out_f * 2 || x.len() < in_f || !in_f.is_multiple_of(8) || y_raw == 0 {
22251 return Err("matvec_bf16_raw_out geometry".into());
22252 }
22253 let f = self.func("matvec_bf16_f32acc");
22254 let cfg = LaunchConfig {
22255 grid_dim: (out_f as u32, 1, 1),
22256 block_dim: (mmv_block(), 1, 1),
22257 shared_mem_bytes: 0,
22258 };
22259 let ini = in_f as i32;
22260 let __s_b = self.gpu.stream();
22261 let mut b = __s_b.launch_builder(&f);
22262 b.arg(w).arg(x).arg(&y_raw).arg(&ini);
22263 unsafe {
22264 b.launch(cfg)?;
22265 }
22266 Ok(())
22267 }
22268
22269 pub fn add3_raw(
22273 &self,
22274 a: &CudaSlice<f32>,
22275 b: &CudaSlice<f32>,
22276 sh_raw: u64,
22277 scale_raw: u64,
22278 dst: &mut CudaSlice<f32>,
22279 n: usize,
22280 ) -> Result<(), Box<dyn std::error::Error>> {
22281 if a.len() < n || b.len() < n || dst.len() < n || sh_raw == 0 || scale_raw == 0 {
22282 return Err("add3_raw geometry".into());
22283 }
22284 let f = self.func("add3_f32");
22285 let cfg = LaunchConfig {
22286 grid_dim: ((n as u32).div_ceil(256), 1, 1),
22287 block_dim: (256, 1, 1),
22288 shared_mem_bytes: 0,
22289 };
22290 let ni = n as i32;
22291 let __s_b = self.gpu.stream();
22292 let mut bld = __s_b.launch_builder(&f);
22293 bld.arg(a)
22294 .arg(b)
22295 .arg(&sh_raw)
22296 .arg(&scale_raw)
22297 .arg(dst)
22298 .arg(&ni);
22299 unsafe {
22300 bld.launch(cfg)?;
22301 }
22302 Ok(())
22303 }
22304
22305 pub fn matvec_bf16_down_addscale_into(
22308 &self,
22309 w: &CudaSlice<u8>,
22310 x: &CudaSlice<f32>,
22311 scale: &CudaSlice<f32>,
22312 dst: &mut CudaSlice<f32>,
22313 in_f: usize,
22314 out_f: usize,
22315 ) -> Result<(), Box<dyn std::error::Error>> {
22316 if w.len() != in_f * out_f * 2
22317 || x.len() < in_f
22318 || !in_f.is_multiple_of(8)
22319 || dst.len() < out_f
22320 || scale.is_empty()
22321 {
22322 return Err("matvec_bf16_down_addscale geometry".into());
22323 }
22324 let f = self.func("matvec_bf16_down_addscale");
22325 let cfg = LaunchConfig {
22326 grid_dim: (out_f as u32, 1, 1),
22327 block_dim: (mmv_block(), 1, 1),
22328 shared_mem_bytes: 0,
22329 };
22330 let ini = in_f as i32;
22331 let __s_b = self.gpu.stream();
22332 let mut b = __s_b.launch_builder(&f);
22333 b.arg(w).arg(x).arg(scale).arg(dst).arg(&ini);
22334 unsafe {
22335 b.launch(cfg)?;
22336 }
22337 Ok(())
22338 }
22339
22340 #[allow(clippy::too_many_arguments)]
22344 pub fn matvec_bf16_dual_silu_rows_into(
22345 &self,
22346 wg: &CudaSlice<u8>,
22347 wu: &CudaSlice<u8>,
22348 x: &CudaSlice<f32>,
22349 act: &mut CudaSlice<f32>,
22350 in_f: usize,
22351 out_f: usize,
22352 limit: Option<f32>,
22353 t: usize,
22354 ) -> Result<(), Box<dyn std::error::Error>> {
22355 if x.len() < t * in_f || act.len() < t * out_f || t == 0 || t > 32 {
22356 return Err("matvec_bf16_dual_silu_rows geometry".into());
22357 }
22358 let f = self.func("matvec_bf16_dual_silu_rows");
22359 let cfg = LaunchConfig {
22360 grid_dim: (out_f as u32, t as u32, 1),
22361 block_dim: (mmv_block(), 1, 1),
22362 shared_mem_bytes: 0,
22363 };
22364 let (ini, outi) = (in_f as i32, out_f as i32);
22365 let lim = limit.unwrap_or(0.0);
22366 let __s_b = self.gpu.stream();
22367 let mut b = __s_b.launch_builder(&f);
22368 b.arg(wg)
22369 .arg(wu)
22370 .arg(x)
22371 .arg(&mut *act)
22372 .arg(&ini)
22373 .arg(&outi)
22374 .arg(&lim);
22375 unsafe {
22376 b.launch(cfg)?;
22377 }
22378 Ok(())
22379 }
22380
22381 pub fn matvec_bf16_rows_into(
22383 &self,
22384 w: &CudaSlice<u8>,
22385 x: &CudaSlice<f32>,
22386 y: &mut CudaSlice<f32>,
22387 in_f: usize,
22388 out_f: usize,
22389 t: usize,
22390 ) -> Result<(), Box<dyn std::error::Error>> {
22391 if x.len() < t * in_f || y.len() < t * out_f || t == 0 || t > 32 || !in_f.is_multiple_of(8)
22392 {
22393 return Err("matvec_bf16_rows geometry".into());
22394 }
22395 if (2..=32).contains(&t)
22400 && step_tp_w8_on()
22401 && w8_hybrid_on()
22402 && in_f.is_multiple_of(32)
22403 && out_f >= 64
22404 && let Some(()) = self.matvec_bf16_via_q8_mirror_t(w, x, y, in_f, out_f, t)?
22405 {
22406 return Ok(());
22407 }
22408 if t == 1
22414 && step_tp_w8_on()
22415 && w8_hybrid_on()
22416 && in_f.is_multiple_of(32)
22417 && out_f >= 64
22418 && let Some(()) = self.matvec_bf16_via_q8_mirror(w, x, y, in_f, out_f)?
22419 {
22420 return Ok(());
22421 }
22422 if (2..=16).contains(&t) && bf16_tcols_wide_on() {
22430 if BF16_TCOLS_WIDE_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
22431 eprintln!(
22432 "[bf16-tcols-wide] engaged: t={t} in_f={in_f} out_f={out_f} rides the \
22433 weight-once tcols class (MEMRA_BF16_TCOLS_WIDE=1)"
22434 );
22435 }
22436 if t <= 8 {
22437 return self.matvec_bf16_tcols_into(w, x, y, in_f, out_f, t);
22438 }
22439 return self.matvec_bf16_tcols16_into(w, x, y, in_f, out_f, t);
22440 }
22441 let f = self.func("matvec_bf16_f32acc_x4_rows");
22442 let cfg = LaunchConfig {
22443 grid_dim: (out_f.div_ceil(4) as u32, t as u32, 1),
22444 block_dim: (mmv_block(), 1, 1),
22445 shared_mem_bytes: 0,
22446 };
22447 let (ini, outi) = (in_f as i32, out_f as i32);
22448 let __s_b = self.gpu.stream();
22449 let mut b = __s_b.launch_builder(&f);
22450 b.arg(w).arg(x).arg(&mut *y).arg(&ini).arg(&outi);
22451 unsafe {
22452 b.launch(cfg)?;
22453 }
22454 Ok(())
22455 }
22456
22457 pub fn matvec_bf16_tcols_into(
22465 &self,
22466 w: &CudaSlice<u8>,
22467 x: &CudaSlice<f32>,
22468 y: &mut CudaSlice<f32>,
22469 in_f: usize,
22470 out_f: usize,
22471 t: usize,
22472 ) -> Result<(), Box<dyn std::error::Error>> {
22473 if x.len() < t * in_f
22474 || y.len() < t * out_f
22475 || !(2..=8).contains(&t)
22476 || !in_f.is_multiple_of(8)
22477 {
22478 return Err("matvec_bf16_tcols geometry".into());
22479 }
22480 let rf = bf16_tcols_red_fused_on() && mmv_block().is_power_of_two();
22491 if rf
22492 && BF16_TCOLS_RED_FUSED_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
22493 == 0
22494 {
22495 eprintln!(
22496 "[bf16-tcols-red-fused] engaged: fused-t reduce tail, one barrier sequence \
22497 shared across the t token columns + intra-warp shuffles at the identical \
22498 pairing (MEMRA_BF16_TCOLS_RED_FUSED=1)"
22499 );
22500 }
22501 let x1 = bf16_tcols_x1_on();
22502 let (fname, grid_x) = match (x1, rf) {
22503 (true, true) => ("matvec_bf16_f32acc_x1_tcols_rf", out_f as u32),
22504 (true, false) => ("matvec_bf16_f32acc_x1_tcols", out_f as u32),
22505 (false, true) => ("matvec_bf16_f32acc_x4_tcols_rf", out_f.div_ceil(4) as u32),
22506 (false, false) => ("matvec_bf16_f32acc_x4_tcols", out_f.div_ceil(4) as u32),
22507 };
22508 if x1 && BF16_TCOLS_X1_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
22509 eprintln!(
22510 "[bf16-tcols-x1] engaged: one-row-per-block tcols grid \
22511 (MEMRA_BF16_TCOLS_X1=1)"
22512 );
22513 }
22514 let f = self.func(fname);
22515 let cfg = LaunchConfig {
22516 grid_dim: (grid_x, 1, 1),
22517 block_dim: (mmv_block(), 1, 1),
22518 shared_mem_bytes: if rf { (t as u32) * mmv_block() * 4 } else { 0 },
22519 };
22520 let (ini, outi, ti) = (in_f as i32, out_f as i32, t as i32);
22521 let __s_b = self.gpu.stream();
22522 let mut b = __s_b.launch_builder(&f);
22523 b.arg(w).arg(x).arg(&mut *y).arg(&ini).arg(&outi).arg(&ti);
22524 unsafe {
22525 b.launch(cfg)?;
22526 }
22527 Ok(())
22528 }
22529
22530 pub fn matvec_bf16_tcols16_into(
22536 &self,
22537 w: &CudaSlice<u8>,
22538 x: &CudaSlice<f32>,
22539 y: &mut CudaSlice<f32>,
22540 in_f: usize,
22541 out_f: usize,
22542 t: usize,
22543 ) -> Result<(), Box<dyn std::error::Error>> {
22544 if x.len() < t * in_f
22545 || y.len() < t * out_f
22546 || !(9..=16).contains(&t)
22547 || !in_f.is_multiple_of(8)
22548 {
22549 return Err("matvec_bf16_tcols16 geometry".into());
22550 }
22551 let rf = bf16_tcols_red_fused_on() && mmv_block().is_power_of_two();
22555 if rf
22556 && BF16_TCOLS_RED_FUSED_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
22557 == 0
22558 {
22559 eprintln!(
22560 "[bf16-tcols-red-fused] engaged: fused-t reduce tail, one barrier sequence \
22561 shared across the t token columns + intra-warp shuffles at the identical \
22562 pairing (MEMRA_BF16_TCOLS_RED_FUSED=1)"
22563 );
22564 }
22565 let f = self.func(if rf {
22566 "matvec_bf16_f32acc_x4_tcols16_rf"
22567 } else {
22568 "matvec_bf16_f32acc_x4_tcols16"
22569 });
22570 let cfg = LaunchConfig {
22571 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
22572 block_dim: (mmv_block(), 1, 1),
22573 shared_mem_bytes: if rf { (t as u32) * mmv_block() * 4 } else { 0 },
22574 };
22575 let (ini, outi, ti) = (in_f as i32, out_f as i32, t as i32);
22576 let __s_b = self.gpu.stream();
22577 let mut b = __s_b.launch_builder(&f);
22578 b.arg(w).arg(x).arg(&mut *y).arg(&ini).arg(&outi).arg(&ti);
22579 unsafe {
22580 b.launch(cfg)?;
22581 }
22582 Ok(())
22583 }
22584
22585 #[allow(clippy::too_many_arguments)]
22591 pub fn matvec_bf16_tcols_gate_kernel_into(
22592 &self,
22593 kernel: &str,
22594 w: &CudaSlice<u8>,
22595 x: &CudaSlice<f32>,
22596 y: &mut CudaSlice<f32>,
22597 in_f: usize,
22598 out_f: usize,
22599 t: usize,
22600 ) -> Result<(), Box<dyn std::error::Error>> {
22601 let (grid_x, t_max) = match kernel {
22602 "matvec_bf16_f32acc_x1_tcols_rf" | "matvec_bf16_f32acc_x1_tcols_rf_redshift" => {
22603 (out_f as u32, 8usize)
22604 }
22605 "matvec_bf16_f32acc_x4_tcols_rf" => (out_f.div_ceil(4) as u32, 8usize),
22606 "matvec_bf16_f32acc_x4_tcols16_rf" => (out_f.div_ceil(4) as u32, 16usize),
22607 _ => return Err("matvec_bf16_tcols_gate_kernel_into: unknown kernel".into()),
22608 };
22609 if x.len() < t * in_f
22610 || y.len() < t * out_f
22611 || !(1..=t_max).contains(&t)
22612 || !in_f.is_multiple_of(8)
22613 || !mmv_block().is_power_of_two()
22614 {
22615 return Err("matvec_bf16_tcols_gate_kernel geometry".into());
22616 }
22617 let f = self.func(kernel);
22618 let cfg = LaunchConfig {
22619 grid_dim: (grid_x, 1, 1),
22620 block_dim: (mmv_block(), 1, 1),
22621 shared_mem_bytes: (t as u32) * mmv_block() * 4,
22622 };
22623 let (ini, outi, ti) = (in_f as i32, out_f as i32, t as i32);
22624 let __s_b = self.gpu.stream();
22625 let mut b = __s_b.launch_builder(&f);
22626 b.arg(w).arg(x).arg(&mut *y).arg(&ini).arg(&outi).arg(&ti);
22627 unsafe {
22628 b.launch(cfg)?;
22629 }
22630 Ok(())
22631 }
22632
22633 pub fn matmul_rows_exact(
22641 &self,
22642 w: &crate::model::GpuTensor,
22643 x: &CudaSlice<f32>,
22644 m: usize,
22645 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22646 use crate::model::GpuTensor;
22647 if let GpuTensor::FloatBf16 { data, .. } = w
22648 && (2..=8).contains(&m)
22649 && Self::bf16_mmv_on()
22650 && w.in_features().is_multiple_of(8)
22651 && !(step_tp_w8_on() && w8_hybrid_on())
22652 {
22653 let (in_f, out_f) = (w.in_features(), w.out_features());
22654 let mut y = self.vws_uninit(m * out_f)?;
22657 self.matvec_bf16_tcols_into(data, x, &mut y, in_f, out_f, m)?;
22658 return Ok(y);
22659 }
22660 self.matmul_decode_exact(w, x, m)
22661 }
22662
22663 #[allow(clippy::too_many_arguments)] pub fn matvec_bf16_dual_silu_into(
22665 &self,
22666 wg: &CudaSlice<u8>,
22667 wu: &CudaSlice<u8>,
22668 x: &CudaSlice<f32>,
22669 act: &mut CudaSlice<f32>,
22670 in_f: usize,
22671 out_f: usize,
22672 limit: Option<f32>,
22673 ) -> Result<(), Box<dyn std::error::Error>> {
22674 if wg.len() != in_f * out_f * 2
22675 || wu.len() != in_f * out_f * 2
22676 || x.len() < in_f
22677 || !in_f.is_multiple_of(8)
22678 || act.len() < out_f
22679 {
22680 return Err("matvec_bf16_dual_silu geometry".into());
22681 }
22682 let f = self.func("matvec_bf16_dual_silu");
22683 let cfg = LaunchConfig {
22684 grid_dim: (out_f as u32, 1, 1),
22685 block_dim: (mmv_block(), 1, 1),
22686 shared_mem_bytes: 0,
22687 };
22688 let (ini, outi) = (in_f as i32, out_f as i32);
22689 let lim = limit.unwrap_or(0.0);
22690 let __s_b = self.gpu.stream();
22691 let mut b = __s_b.launch_builder(&f);
22692 b.arg(wg)
22693 .arg(wu)
22694 .arg(x)
22695 .arg(act)
22696 .arg(&ini)
22697 .arg(&outi)
22698 .arg(&lim);
22699 unsafe {
22700 b.launch(cfg)?;
22701 }
22702 Ok(())
22703 }
22704
22705 #[allow(clippy::too_many_arguments)]
22708 pub fn matvec_bf16_dual_view_into(
22709 &self,
22710 wg: &cudarc::driver::CudaView<'_, u8>,
22711 wu: &cudarc::driver::CudaView<'_, u8>,
22712 x: &CudaSlice<f32>,
22713 yg: &mut CudaSlice<f32>,
22714 yu: &mut CudaSlice<f32>,
22715 in_f: usize,
22716 out_f: usize,
22717 ) -> Result<(), Box<dyn std::error::Error>> {
22718 if wg.len() != in_f * out_f * 2
22719 || wu.len() != in_f * out_f * 2
22720 || x.len() < in_f
22721 || !in_f.is_multiple_of(8)
22722 || yg.len() < out_f
22723 || yu.len() < out_f
22724 {
22725 return Err(format!(
22726 "matvec_bf16_dual_view_into geometry wg={} wu={} x={} in={in_f} out={out_f}",
22727 wg.len(),
22728 wu.len(),
22729 x.len()
22730 )
22731 .into());
22732 }
22733 let f = self.func("matvec_bf16_dual");
22734 let cfg = LaunchConfig {
22735 grid_dim: ((2 * out_f) as u32, 1, 1),
22736 block_dim: (mmv_block(), 1, 1),
22737 shared_mem_bytes: 0,
22738 };
22739 let (ini, outi) = (in_f as i32, out_f as i32);
22740 let __s_b = self.gpu.stream();
22741 let mut b = __s_b.launch_builder(&f);
22742 b.arg(wg)
22743 .arg(wu)
22744 .arg(x)
22745 .arg(yg)
22746 .arg(yu)
22747 .arg(&ini)
22748 .arg(&outi);
22749 unsafe {
22750 b.launch(cfg)?;
22751 }
22752 Ok(())
22753 }
22754
22755 #[allow(clippy::too_many_arguments)]
22757 pub fn matvec_bf16_dual_into(
22758 &self,
22759 wg: &CudaSlice<u8>,
22760 wu: &CudaSlice<u8>,
22761 x: &CudaSlice<f32>,
22762 yg: &mut CudaSlice<f32>,
22763 yu: &mut CudaSlice<f32>,
22764 in_f: usize,
22765 out_f: usize,
22766 ) -> Result<(), Box<dyn std::error::Error>> {
22767 if wg.len() != in_f * out_f * 2
22768 || wu.len() != in_f * out_f * 2
22769 || x.len() < in_f
22770 || !in_f.is_multiple_of(8)
22771 || yg.len() < out_f
22772 || yu.len() < out_f
22773 {
22774 return Err(format!(
22775 "matvec_bf16_dual_into geometry wg={} wu={} x={} in={in_f} out={out_f}",
22776 wg.len(),
22777 wu.len(),
22778 x.len()
22779 )
22780 .into());
22781 }
22782 let f = self.func("matvec_bf16_dual");
22783 let cfg = LaunchConfig {
22784 grid_dim: ((2 * out_f) as u32, 1, 1),
22785 block_dim: (mmv_block(), 1, 1),
22786 shared_mem_bytes: 0,
22787 };
22788 let (ini, outi) = (in_f as i32, out_f as i32);
22789 let __s_b = self.gpu.stream();
22790 let mut b = __s_b.launch_builder(&f);
22791 b.arg(wg)
22792 .arg(wu)
22793 .arg(x)
22794 .arg(yg)
22795 .arg(yu)
22796 .arg(&ini)
22797 .arg(&outi);
22798 unsafe {
22799 b.launch(cfg)?;
22800 }
22801 Ok(())
22802 }
22803
22804 #[allow(dead_code)] pub(crate) fn matvec_bf16_dual(
22808 &self,
22809 wg: &CudaSlice<u8>,
22810 wu: &CudaSlice<u8>,
22811 x: &CudaSlice<f32>,
22812 in_f: usize,
22813 out_f: usize,
22814 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
22815 if wg.len() != in_f * out_f * 2
22816 || wu.len() != in_f * out_f * 2
22817 || x.len() < in_f
22818 || !in_f.is_multiple_of(8)
22819 {
22820 return Err(format!(
22821 "matvec_bf16_dual geometry wg={} wu={} x={} in={in_f} out={out_f}",
22822 wg.len(),
22823 wu.len(),
22824 x.len()
22825 )
22826 .into());
22827 }
22828 let mut yg = self.alloc_uninit::<f32>(out_f)?;
22829 let mut yu = self.alloc_uninit::<f32>(out_f)?;
22830 let f = self.func("matvec_bf16_dual");
22831 let cfg = LaunchConfig {
22832 grid_dim: ((2 * out_f) as u32, 1, 1),
22833 block_dim: (mmv_block(), 1, 1),
22834 shared_mem_bytes: 0,
22835 };
22836 let (ini, outi) = (in_f as i32, out_f as i32);
22837 let __s_b = self.gpu.stream();
22838 let mut b = __s_b.launch_builder(&f);
22839 b.arg(wg)
22840 .arg(wu)
22841 .arg(x)
22842 .arg(&mut yg)
22843 .arg(&mut yu)
22844 .arg(&ini)
22845 .arg(&outi);
22846 unsafe {
22847 b.launch(cfg)?;
22848 }
22849 Ok((yg, yu))
22850 }
22851
22852 #[allow(clippy::too_many_arguments)]
22853 #[allow(clippy::manual_is_multiple_of)] fn linear_bf16_chunked_inner(
22855 &self,
22856 x: &CudaSlice<f32>,
22857 data: &CudaSlice<u8>,
22858 m: usize,
22859 in_f: usize,
22860 out_f: usize,
22861 exact: bool,
22862 canonical_chunk_rows: Option<usize>,
22863 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22864 const CHUNK_BYTES: usize = 256 << 20;
22865 if m == 1
22868 && !exact
22869 && canonical_chunk_rows.is_none()
22870 && in_f.is_multiple_of(8)
22871 && Self::bf16_mmv_on()
22872 {
22873 return self.matvec_bf16(data, x, in_f, out_f);
22874 }
22875 if m >= 16
22880 && !exact
22881 && canonical_chunk_rows.is_none()
22882 && data.len() == in_f * out_f * 2
22883 && crate::f16_ffi::pp_bf16_enabled()
22884 {
22885 if let Some(y) = self.bf16_tc_gemm(data, x, m, in_f, out_f)? {
22888 return Ok(y);
22889 }
22890 }
22891 let row_bytes = in_f
22892 .checked_mul(std::mem::size_of::<f32>())
22893 .ok_or("BF16 chunk row byte count overflow")?;
22894 if row_bytes == 0 || out_f == 0 {
22895 return Err("BF16 chunk dimensions must be nonzero".into());
22896 }
22897 let max_chunk_rows = (CHUNK_BYTES / row_bytes).max(1).min(out_f);
22898 let chunk_rows = match canonical_chunk_rows {
22899 Some(0) => {
22900 return Err("canonical BF16 chunk rows must be nonzero".into());
22901 }
22902 Some(rows) if rows > max_chunk_rows => {
22903 return Err(format!(
22904 "canonical BF16 chunk rows {rows} exceed the {max_chunk_rows}-row scratch limit"
22905 )
22906 .into());
22907 }
22908 Some(rows) if out_f % rows != 0 => {
22909 return Err(format!(
22910 "BF16 output width {out_f} is not divisible by canonical {rows}-row chunks"
22911 )
22912 .into());
22913 }
22914 Some(rows) => rows,
22915 None => max_chunk_rows,
22916 };
22917 if chunk_rows >= out_f {
22918 let wf32 = self.bf16_to_f32(&data.slice(0..in_f * out_f * 2), in_f * out_f)?;
22919 return if exact {
22920 self.linear_decode_exact(x, &wf32, m, in_f, out_f)
22921 } else {
22922 self.linear(x, &wf32, m, in_f, out_f)
22923 };
22924 }
22925 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
22926 let mut r0 = 0usize;
22927 while r0 < out_f {
22928 let rows = chunk_rows.min(out_f - r0);
22929 let wslice = data.slice(r0 * in_f * 2..(r0 + rows) * in_f * 2);
22930 let wf32 = self.bf16_to_f32(&wslice, in_f * rows)?;
22931 let yc = if exact {
22932 self.linear_decode_exact(x, &wf32, m, in_f, rows)?
22933 } else {
22934 self.linear(x, &wf32, m, in_f, rows)?
22935 };
22936 for mi in 0..m {
22938 let src = yc.slice(mi * rows..(mi + 1) * rows);
22939 let mut dst = y.slice_mut(mi * out_f + r0..mi * out_f + r0 + rows);
22940 self.gpu.stream().memcpy_dtod(&src, &mut dst)?;
22941 }
22942 r0 += rows;
22943 }
22944 Ok(y)
22945 }
22946
22947 pub fn linear_bf16_resident(
22951 &self,
22952 x: &CudaSlice<f32>,
22953 data: &CudaSlice<u8>,
22954 m: usize,
22955 in_f: usize,
22956 out_f: usize,
22957 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22958 if data.len() != in_f * out_f * 2 {
22959 return Err(format!("resident BF16 bytes {} != {out_f}x{in_f}x2", data.len()).into());
22960 }
22961 self.linear_bf16_chunked(x, data, m, in_f, out_f, false, None)
22962 }
22963
22964 pub fn linear_bf16_resident_canonical_rows(
22970 &self,
22971 x: &CudaSlice<f32>,
22972 data: &CudaSlice<u8>,
22973 m: usize,
22974 in_f: usize,
22975 out_f: usize,
22976 canonical_chunk_rows: usize,
22977 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22978 if data.len() != in_f * out_f * 2 {
22979 return Err(format!("resident BF16 bytes {} != {out_f}x{in_f}x2", data.len()).into());
22980 }
22981 self.linear_bf16_chunked(x, data, m, in_f, out_f, false, Some(canonical_chunk_rows))
22982 }
22983
22984 pub fn linear_f32_resident_canonical_rows(
22989 &self,
22990 x: &CudaSlice<f32>,
22991 data: &CudaSlice<f32>,
22992 m: usize,
22993 in_f: usize,
22994 out_f: usize,
22995 canonical_chunk_rows: usize,
22996 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22997 self.linear_f32_resident_canonical_rows_inner(
22998 x,
22999 data,
23000 m,
23001 in_f,
23002 out_f,
23003 canonical_chunk_rows,
23004 false,
23005 )
23006 }
23007
23008 pub fn linear_f32_resident_canonical_rows_strided(
23014 &self,
23015 x: &CudaSlice<f32>,
23016 data: &CudaSlice<f32>,
23017 m: usize,
23018 in_f: usize,
23019 out_f: usize,
23020 canonical_chunk_rows: usize,
23021 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23022 self.linear_f32_resident_canonical_rows_inner(
23023 x,
23024 data,
23025 m,
23026 in_f,
23027 out_f,
23028 canonical_chunk_rows,
23029 true,
23030 )
23031 }
23032
23033 #[allow(clippy::too_many_arguments)]
23034 #[allow(clippy::manual_is_multiple_of)] fn linear_f32_resident_canonical_rows_inner(
23037 &self,
23038 x: &CudaSlice<f32>,
23039 data: &CudaSlice<f32>,
23040 m: usize,
23041 in_f: usize,
23042 out_f: usize,
23043 canonical_chunk_rows: usize,
23044 strided_output: bool,
23045 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23046 if data.len() != in_f * out_f {
23047 return Err(format!("resident F32 values {} != {out_f}x{in_f}", data.len()).into());
23048 }
23049 if canonical_chunk_rows == 0
23050 || canonical_chunk_rows > out_f
23051 || out_f % canonical_chunk_rows != 0
23052 {
23053 return Err(format!(
23054 "invalid canonical F32 chunk rows {canonical_chunk_rows} for output width {out_f}"
23055 )
23056 .into());
23057 }
23058 if canonical_chunk_rows == out_f {
23059 return self.linear(x, data, m, in_f, out_f);
23060 }
23061
23062 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
23063 let input = x.slice(0..x.len());
23064 for r0 in (0..out_f).step_by(canonical_chunk_rows) {
23065 let weights = data.slice(r0 * in_f..(r0 + canonical_chunk_rows) * in_f);
23066 if m == 1 {
23067 let mut destination = y.slice_mut(r0..r0 + canonical_chunk_rows);
23068 self.linear_device_into(
23069 &input,
23070 &weights,
23071 &mut destination,
23072 1,
23073 in_f,
23074 canonical_chunk_rows,
23075 )?;
23076 continue;
23077 }
23078 let chunk = self.linear_device(&input, &weights, m, in_f, canonical_chunk_rows)?;
23079 if strided_output {
23080 self.place_rows_strided(&chunk, &mut y, canonical_chunk_rows, m, out_f, r0)?;
23081 } else {
23082 for token in 0..m {
23083 let source = chunk
23084 .slice(token * canonical_chunk_rows..(token + 1) * canonical_chunk_rows);
23085 let mut destination =
23086 y.slice_mut(token * out_f + r0..token * out_f + r0 + canonical_chunk_rows);
23087 self.gpu.stream().memcpy_dtod(&source, &mut destination)?;
23088 }
23089 }
23090 }
23091 Ok(y)
23092 }
23093
23094 #[allow(clippy::manual_is_multiple_of)] pub fn linear_f32_resident_canonical_rows_t1_into(
23100 &self,
23101 x: &CudaSlice<f32>,
23102 data: &CudaSlice<f32>,
23103 y: &mut CudaSlice<f32>,
23104 in_f: usize,
23105 out_f: usize,
23106 canonical_chunk_rows: usize,
23107 ) -> Result<(), Box<dyn std::error::Error>> {
23108 if data.len() != in_f * out_f {
23109 return Err(format!("resident F32 values {} != {out_f}x{in_f}", data.len()).into());
23110 }
23111 if y.len() != out_f || x.len() != in_f {
23112 return Err(format!(
23113 "resident F32 t1 shapes x={} y={} != in {in_f} out {out_f}",
23114 x.len(),
23115 y.len()
23116 )
23117 .into());
23118 }
23119 if canonical_chunk_rows == 0
23120 || canonical_chunk_rows > out_f
23121 || out_f % canonical_chunk_rows != 0
23122 {
23123 return Err(format!(
23124 "invalid canonical F32 chunk rows {canonical_chunk_rows} for output width {out_f}"
23125 )
23126 .into());
23127 }
23128 let input = x.slice(0..x.len());
23129 for r0 in (0..out_f).step_by(canonical_chunk_rows) {
23130 let weights = data.slice(r0 * in_f..(r0 + canonical_chunk_rows) * in_f);
23131 let mut destination = y.slice_mut(r0..r0 + canonical_chunk_rows);
23132 self.linear_device_into(
23133 &input,
23134 &weights,
23135 &mut destination,
23136 1,
23137 in_f,
23138 canonical_chunk_rows,
23139 )?;
23140 }
23141 Ok(())
23142 }
23143
23144 pub fn linear_t1_into(
23147 &self,
23148 x: &cudarc::driver::CudaView<'_, f32>,
23149 w: &cudarc::driver::CudaView<'_, f32>,
23150 y: &mut cudarc::driver::CudaViewMut<'_, f32>,
23151 in_f: usize,
23152 out_f: usize,
23153 ) -> Result<(), Box<dyn std::error::Error>> {
23154 self.linear_device_into(x, w, y, 1, in_f, out_f)
23155 }
23156
23157 pub fn linear_decode_exact(
23164 &self,
23165 x: &CudaSlice<f32>,
23166 w: &CudaSlice<f32>,
23167 m_tokens: usize,
23168 in_f: usize,
23169 out_f: usize,
23170 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23171 if m_tokens == 1 {
23172 return self.linear(x, w, 1, in_f, out_f);
23173 }
23174 let xv = self.view(x, m_tokens * in_f);
23175 let mut y = self.alloc_uninit::<f32>(m_tokens * out_f)?;
23176 for t in 0..m_tokens {
23177 let row = xv.slice(t * in_f..(t + 1) * in_f);
23178 let mut xr = self.alloc_uninit::<f32>(in_f)?;
23179 self.copy_view_into(&mut xr, 0, &row, in_f)?;
23180 let yr = self.linear(&xr, w, 1, in_f, out_f)?;
23181 self.copy_into(&mut y, t * out_f, &yr, out_f)?;
23182 }
23183 Ok(y)
23184 }
23185
23186 pub fn linear(
23187 &self,
23188 x: &CudaSlice<f32>,
23189 w: &CudaSlice<f32>,
23190 m_tokens: usize,
23191 in_f: usize,
23192 out_f: usize,
23193 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23194 self.linear_device(x, w, m_tokens, in_f, out_f)
23195 }
23196
23197 fn linear_device<I>(
23198 &self,
23199 x: &I,
23200 w: &I,
23201 m_tokens: usize,
23202 in_f: usize,
23203 out_f: usize,
23204 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>>
23205 where
23206 I: cudarc::driver::DevicePtr<f32>,
23207 {
23208 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)?;
23210 Ok(c)
23211 }
23212
23213 fn linear_device_into<I, O>(
23214 &self,
23215 x: &I,
23216 w: &I,
23217 c: &mut O,
23218 m_tokens: usize,
23219 in_f: usize,
23220 out_f: usize,
23221 ) -> Result<(), Box<dyn std::error::Error>>
23222 where
23223 I: cudarc::driver::DevicePtr<f32>,
23224 O: cudarc::driver::DevicePtrMut<f32>,
23225 {
23226 use cudarc::cublaslt::{Matmul, MatmulConfig};
23227 let cfg = MatmulConfig {
23228 transa: true,
23229 transb: false,
23230 transc: false,
23231 m: out_f as u64,
23232 n: m_tokens as u64,
23233 k: in_f as u64,
23234 alpha: 1.0,
23235 lda: in_f as i64,
23236 ldb: in_f as i64,
23237 beta: 0.0,
23238 ldc: out_f as i64,
23239 stride_a: None,
23240 stride_b: None,
23241 stride_c: None,
23242 stride_bias: None,
23243 batch_size: None,
23244 };
23245 let blas = self.gpu.blas();
23246 unsafe {
23247 blas.matmul(cfg, w, x, c, None, None)?;
23248 }
23249 Ok(())
23250 }
23251
23252 #[allow(clippy::too_many_arguments)] pub fn sdpa_naive(
23262 &self,
23263 q: &CudaSlice<f32>,
23264 k: &CudaSlice<f32>,
23265 v: &CudaSlice<f32>,
23266 o: &mut CudaSlice<f32>,
23267 head_dim: usize,
23268 n_head: usize,
23269 n_head_kv: usize,
23270 t: usize,
23271 t_kv: usize,
23272 scale: f32,
23273 causal: bool,
23274 ) -> Result<(), Box<dyn std::error::Error>> {
23275 if t_kv * 4 > SDPA_NAIVE_SMEM_MAX {
23276 return self.sdpa_naive_gmem(
23277 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
23278 );
23279 }
23280 let f = self.func("sdpa_naive_f32");
23281 let cfg = LaunchConfig {
23282 grid_dim: (n_head as u32, t as u32, 1),
23283 block_dim: (128, 1, 1),
23284 shared_mem_bytes: (t_kv * 4) as u32,
23285 };
23286 let (hd, nh, nhkv, ti, tkvi, cz) = (
23287 head_dim as i32,
23288 n_head as i32,
23289 n_head_kv as i32,
23290 t as i32,
23291 t_kv as i32,
23292 causal as i32,
23293 );
23294 let __s_b = self.gpu.stream();
23295 let mut b = __s_b.launch_builder(&f);
23296 b.arg(q)
23297 .arg(k)
23298 .arg(v)
23299 .arg(o)
23300 .arg(&hd)
23301 .arg(&nh)
23302 .arg(&nhkv)
23303 .arg(&ti)
23304 .arg(&tkvi)
23305 .arg(&scale)
23306 .arg(&cz);
23307 unsafe {
23308 b.launch(cfg)?;
23309 }
23310 Ok(())
23311 }
23312
23313 #[allow(clippy::too_many_arguments)]
23322 pub fn sdpa_naive_gmem(
23323 &self,
23324 q: &CudaSlice<f32>,
23325 k: &CudaSlice<f32>,
23326 v: &CudaSlice<f32>,
23327 o: &mut CudaSlice<f32>,
23328 head_dim: usize,
23329 n_head: usize,
23330 n_head_kv: usize,
23331 t: usize,
23332 t_kv: usize,
23333 scale: f32,
23334 causal: bool,
23335 ) -> Result<(), Box<dyn std::error::Error>> {
23336 let ws_len = n_head
23337 .checked_mul(t)
23338 .and_then(|x| x.checked_mul(t_kv))
23339 .ok_or("sdpa_naive_gmem: scores workspace size overflow")?;
23340 let ws_bytes = ws_len
23341 .checked_mul(std::mem::size_of::<f32>())
23342 .ok_or("sdpa_naive_gmem: scores workspace byte count overflow")?;
23343 if ws_bytes > SDPA_NAIVE_GMEM_WS_MAX {
23344 return Err(format!(
23345 "sdpa_naive_gmem: scores workspace {ws_bytes} bytes (heads {n_head} x T {t} x \
23346 T_kv {t_kv}) exceeds the {SDPA_NAIVE_GMEM_WS_MAX}-byte guard — this shape \
23347 needs a tiled/flash kernel, not the naive oracle"
23348 )
23349 .into());
23350 }
23351 let mut scores = self.uninit(ws_len)?;
23352 let f = self.func("sdpa_naive_gmem_f32");
23353 let cfg = LaunchConfig {
23354 grid_dim: (n_head as u32, t as u32, 1),
23355 block_dim: (128, 1, 1),
23356 shared_mem_bytes: 0,
23357 };
23358 let (hd, nh, nhkv, ti, tkvi, cz) = (
23359 head_dim as i32,
23360 n_head as i32,
23361 n_head_kv as i32,
23362 t as i32,
23363 t_kv as i32,
23364 causal as i32,
23365 );
23366 let __s_b = self.gpu.stream();
23367 let mut b = __s_b.launch_builder(&f);
23368 b.arg(q)
23369 .arg(k)
23370 .arg(v)
23371 .arg(o)
23372 .arg(&mut scores)
23373 .arg(&hd)
23374 .arg(&nh)
23375 .arg(&nhkv)
23376 .arg(&ti)
23377 .arg(&tkvi)
23378 .arg(&scale)
23379 .arg(&cz);
23380 unsafe {
23381 b.launch(cfg)?;
23382 }
23383 Ok(())
23384 }
23385
23386 #[allow(clippy::too_many_arguments)]
23391 pub fn sdpa_naive_island(
23392 &self,
23393 q: &CudaSlice<f32>,
23394 k: &CudaSlice<f32>,
23395 v: &CudaSlice<f32>,
23396 o: &mut CudaSlice<f32>,
23397 span_id: &CudaSlice<i32>,
23398 head_dim: usize,
23399 n_head: usize,
23400 n_head_kv: usize,
23401 t: usize,
23402 t_kv: usize,
23403 scale: f32,
23404 window: usize,
23405 ) -> Result<(), Box<dyn std::error::Error>> {
23406 let f = self.func("sdpa_naive_island_f32");
23407 let cfg = LaunchConfig {
23408 grid_dim: (n_head as u32, t as u32, 1),
23409 block_dim: (128, 1, 1),
23410 shared_mem_bytes: (t_kv * 4) as u32,
23411 };
23412 let (hd, nh, nhkv, ti, tkvi, wi) = (
23413 head_dim as i32,
23414 n_head as i32,
23415 n_head_kv as i32,
23416 t as i32,
23417 t_kv as i32,
23418 window as i32,
23419 );
23420 let __s_b = self.gpu.stream();
23421 let mut b = __s_b.launch_builder(&f);
23422 b.arg(q)
23423 .arg(k)
23424 .arg(v)
23425 .arg(o)
23426 .arg(span_id)
23427 .arg(&hd)
23428 .arg(&nh)
23429 .arg(&nhkv)
23430 .arg(&ti)
23431 .arg(&tkvi)
23432 .arg(&scale)
23433 .arg(&wi);
23434 unsafe {
23435 b.launch(cfg)?;
23436 }
23437 Ok(())
23438 }
23439
23440 #[allow(clippy::too_many_arguments)]
23442 pub fn sdpa_naive_w(
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 window: usize,
23456 ) -> Result<(), Box<dyn std::error::Error>> {
23457 let f = self.func("sdpa_naive_w_f32");
23458 let cfg = LaunchConfig {
23459 grid_dim: (n_head as u32, t as u32, 1),
23460 block_dim: (128, 1, 1),
23461 shared_mem_bytes: (t_kv * 4) as u32,
23462 };
23463 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
23464 head_dim as i32,
23465 n_head as i32,
23466 n_head_kv as i32,
23467 t as i32,
23468 t_kv as i32,
23469 causal as i32,
23470 window as i32,
23471 );
23472 let __s_b = self.gpu.stream();
23473 let mut b = __s_b.launch_builder(&f);
23474 b.arg(q)
23475 .arg(k)
23476 .arg(v)
23477 .arg(o)
23478 .arg(&hd)
23479 .arg(&nh)
23480 .arg(&nhkv)
23481 .arg(&ti)
23482 .arg(&tkvi)
23483 .arg(&scale)
23484 .arg(&cz)
23485 .arg(&wi);
23486 unsafe {
23487 b.launch(cfg)?;
23488 }
23489 Ok(())
23490 }
23491
23492 #[allow(clippy::too_many_arguments)]
23502 pub fn sdpa_naive_w_lo(
23503 &self,
23504 q: &CudaSlice<f32>,
23505 k: &CudaSlice<f32>,
23506 v: &CudaSlice<f32>,
23507 o: &mut CudaSlice<f32>,
23508 head_dim: usize,
23509 n_head: usize,
23510 n_head_kv: usize,
23511 t: usize,
23512 t_kv: usize,
23513 scale: f32,
23514 causal: bool,
23515 window: usize,
23516 ) -> Result<(), Box<dyn std::error::Error>> {
23517 let kv_lo = if window > 0 {
23518 (t_kv - t + 1).saturating_sub(window)
23519 } else {
23520 0
23521 };
23522 let smem = (t_kv - kv_lo) * 4;
23523 if smem > 48 * 1024 {
23524 return Err(format!(
23525 "sdpa_naive_w_lo: window {window} + T {t} rows need {smem} bytes of dynamic \
23526 shared memory (> 48KB launch bound) — this kernel clips the OLD side only; \
23527 a window this wide needs the multi-pass long-ctx kernel"
23528 )
23529 .into());
23530 }
23531 let f = self.func("sdpa_naive_w_lo_f32");
23532 let cfg = LaunchConfig {
23533 grid_dim: (n_head as u32, t as u32, 1),
23534 block_dim: (128, 1, 1),
23535 shared_mem_bytes: smem as u32,
23536 };
23537 let (hd, nh, nhkv, ti, tkvi, cz, wi, lo) = (
23538 head_dim as i32,
23539 n_head as i32,
23540 n_head_kv as i32,
23541 t as i32,
23542 t_kv as i32,
23543 causal as i32,
23544 window as i32,
23545 kv_lo as i32,
23546 );
23547 let __s_b = self.gpu.stream();
23548 let mut b = __s_b.launch_builder(&f);
23549 b.arg(q)
23550 .arg(k)
23551 .arg(v)
23552 .arg(o)
23553 .arg(&hd)
23554 .arg(&nh)
23555 .arg(&nhkv)
23556 .arg(&ti)
23557 .arg(&tkvi)
23558 .arg(&scale)
23559 .arg(&cz)
23560 .arg(&wi)
23561 .arg(&lo);
23562 unsafe {
23563 b.launch(cfg)?;
23564 }
23565 Ok(())
23566 }
23567
23568 #[allow(clippy::too_many_arguments)] pub fn sdpa_naive_view(
23571 &self,
23572 q: &CudaSlice<f32>,
23573 k: &cudarc::driver::CudaView<f32>,
23574 v: &cudarc::driver::CudaView<f32>,
23575 o: &mut CudaSlice<f32>,
23576 head_dim: usize,
23577 n_head: usize,
23578 n_head_kv: usize,
23579 t: usize,
23580 t_kv: usize,
23581 scale: f32,
23582 causal: bool,
23583 ) -> Result<(), Box<dyn std::error::Error>> {
23584 let f = self.func("sdpa_naive_f32");
23585 let cfg = LaunchConfig {
23586 grid_dim: (n_head as u32, t as u32, 1),
23587 block_dim: (128, 1, 1),
23588 shared_mem_bytes: (t_kv * 4) as u32,
23589 };
23590 let (hd, nh, nhkv, ti, tkvi, cz) = (
23591 head_dim as i32,
23592 n_head as i32,
23593 n_head_kv as i32,
23594 t as i32,
23595 t_kv as i32,
23596 causal as i32,
23597 );
23598 let __s_b = self.gpu.stream();
23599 let mut b = __s_b.launch_builder(&f);
23600 b.arg(q)
23601 .arg(k)
23602 .arg(v)
23603 .arg(o)
23604 .arg(&hd)
23605 .arg(&nh)
23606 .arg(&nhkv)
23607 .arg(&ti)
23608 .arg(&tkvi)
23609 .arg(&scale)
23610 .arg(&cz);
23611 unsafe {
23612 b.launch(cfg)?;
23613 }
23614 Ok(())
23615 }
23616
23617 #[allow(clippy::too_many_arguments)]
23625 pub fn fa_dequant_kv_view_f32(
23626 &self,
23627 k: &cudarc::driver::CudaView<u8>,
23628 v: &cudarc::driver::CudaView<u8>,
23629 kf: &mut CudaSlice<f32>,
23630 vf: &mut CudaSlice<f32>,
23631 kv_dim_k: usize,
23632 kv_dim_v: usize,
23633 t_kv: usize,
23634 k_tok_bytes: usize,
23635 v_tok_bytes: usize,
23636 g: bool,
23637 ) -> Result<(), Box<dyn std::error::Error>> {
23638 let f = if g {
23639 self.func_g("fa_dequant_kv_ws_f32")
23640 } else {
23641 self.func("fa_dequant_kv_ws_f32")
23642 };
23643 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
23644 #[allow(clippy::manual_div_ceil)]
23645 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
23647 let cfg = LaunchConfig {
23648 grid_dim: (nblk.max(1), 1, 1),
23649 block_dim: (256, 1, 1),
23650 shared_mem_bytes: 0,
23651 };
23652 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
23653 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
23654 let __s_b = self.gpu.stream();
23655 let mut b = __s_b.launch_builder(&f);
23656 b.arg(k)
23657 .arg(v)
23658 .arg(&mut *kf)
23659 .arg(&mut *vf)
23660 .arg(&kdk)
23661 .arg(&kdv)
23662 .arg(&tkvi)
23663 .arg(&ktb)
23664 .arg(&vtb);
23665 unsafe {
23666 b.launch(cfg)?;
23667 }
23668 Ok(())
23669 }
23670
23671 #[allow(clippy::too_many_arguments)]
23672 pub fn sdpa_naive_quantized_view(
23673 &self,
23674 q: &CudaSlice<f32>,
23675 k: &cudarc::driver::CudaView<u8>,
23676 v: &cudarc::driver::CudaView<u8>,
23677 o: &mut CudaSlice<f32>,
23678 head_dim: usize,
23679 n_head: usize,
23680 n_head_kv: usize,
23681 t: usize,
23682 t_kv: usize,
23683 scale: f32,
23684 causal: bool,
23685 k_tok_bytes: usize,
23686 v_tok_bytes: usize,
23687 ) -> Result<(), Box<dyn std::error::Error>> {
23688 let kv_dim = n_head_kv * head_dim;
23689 let mut kf = self.uninit(t_kv * kv_dim)?;
23690 let mut vf = self.uninit(t_kv * kv_dim)?;
23691 let f = self.func("fa_dequant_kv_ws_f32");
23692 let total = (2 * t_kv * kv_dim) as u64;
23693 #[allow(clippy::manual_div_ceil)]
23694 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
23696 let cfg = LaunchConfig {
23697 grid_dim: (nblk.max(1), 1, 1),
23698 block_dim: (256, 1, 1),
23699 shared_mem_bytes: 0,
23700 };
23701 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
23702 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
23703 let __s_b = self.gpu.stream();
23704 let mut b = __s_b.launch_builder(&f);
23705 b.arg(k)
23706 .arg(v)
23707 .arg(&mut kf)
23708 .arg(&mut vf)
23709 .arg(&kv_dim_i)
23710 .arg(&kv_dim_i)
23711 .arg(&t_kv_i)
23712 .arg(&k_tok_bytes_i)
23713 .arg(&v_tok_bytes_i);
23714 unsafe { b.launch(cfg)? };
23715 self.sdpa_naive(
23716 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
23717 )
23718 }
23719
23720 #[allow(clippy::too_many_arguments)]
23732 pub fn sdpa_naive_w_quantized_view(
23733 &self,
23734 q: &CudaSlice<f32>,
23735 k: &cudarc::driver::CudaView<u8>,
23736 v: &cudarc::driver::CudaView<u8>,
23737 o: &mut CudaSlice<f32>,
23738 head_dim: usize,
23739 n_head: usize,
23740 n_head_kv: usize,
23741 t: usize,
23742 t_kv: usize,
23743 scale: f32,
23744 causal: bool,
23745 window: usize,
23746 k_tok_bytes: usize,
23747 v_tok_bytes: usize,
23748 ) -> Result<(), Box<dyn std::error::Error>> {
23749 let kv_dim = n_head_kv * head_dim;
23750 let mut kf = self.uninit(t_kv * kv_dim)?;
23751 let mut vf = self.uninit(t_kv * kv_dim)?;
23752 let f = self.func("fa_dequant_kv_ws_f32");
23753 let total = (2 * t_kv * kv_dim) as u64;
23754 #[allow(clippy::manual_div_ceil)]
23755 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
23757 let cfg = LaunchConfig {
23758 grid_dim: (nblk.max(1), 1, 1),
23759 block_dim: (256, 1, 1),
23760 shared_mem_bytes: 0,
23761 };
23762 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
23763 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
23764 let __s_b = self.gpu.stream();
23765 let mut b = __s_b.launch_builder(&f);
23766 b.arg(k)
23767 .arg(v)
23768 .arg(&mut kf)
23769 .arg(&mut vf)
23770 .arg(&kv_dim_i)
23771 .arg(&kv_dim_i)
23772 .arg(&t_kv_i)
23773 .arg(&k_tok_bytes_i)
23774 .arg(&v_tok_bytes_i);
23775 unsafe { b.launch(cfg)? };
23776 self.sdpa_naive_w(
23777 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
23778 )
23779 }
23780
23781 #[allow(clippy::too_many_arguments)]
23785 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill(
23788 &self,
23789 q: &CudaSlice<f32>,
23790 k: &CudaSlice<f32>,
23791 v: &CudaSlice<f32>,
23792 o: &mut CudaSlice<f32>,
23793 head_dim: usize,
23794 n_head: usize,
23795 n_head_kv: usize,
23796 t: usize,
23797 t_kv: usize,
23798 scale: f32,
23799 causal: bool,
23800 ) -> Result<(), Box<dyn std::error::Error>> {
23801 if portable_mma_gated() {
23802 return self.sdpa_naive(
23803 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
23804 );
23805 }
23806 let fa3_on = head_dim == 256
23814 && causal
23815 && t == t_kv
23816 && match std::env::var("MEMRA_FA3").as_deref() {
23817 Ok("0") => false,
23818 Ok("1") => {
23822 refuse_portable_force("MEMRA_FA3=1", "the sm_90a fa3/bf16 kernels");
23823 true
23824 }
23825 _ => cfg!(memra_hopper_mma),
23826 };
23827 if fa3_on {
23828 let n = t * n_head * head_dim;
23829 let nkv = t * n_head_kv * head_dim;
23830 let mut q16 = self.alloc_u8_uninit(n * 2)?;
23831 let mut k16 = self.alloc_u8_uninit(nkv * 2)?;
23832 let mut v16 = self.alloc_u8_uninit(nkv * 2)?;
23833 self.f32_to_bf16_into(q, &mut q16, n)?;
23834 self.f32_to_bf16_into(k, &mut k16, nkv)?;
23835 self.f32_to_bf16_into(v, &mut v16, nkv)?;
23836 let rc = {
23837 use cudarc::driver::{DevicePtr, DevicePtrMut};
23838 let stream = self.gpu.stream();
23839 let (qp, _g1) = q16.device_ptr(&stream);
23840 let (kp, _g2) = k16.device_ptr(&stream);
23841 let (vp, _g3) = v16.device_ptr(&stream);
23842 let (op, _g4) = o.device_ptr_mut(&stream);
23843 unsafe {
23844 memra_fa3_prefill(
23845 qp as *const core::ffi::c_void,
23846 kp as *const core::ffi::c_void,
23847 vp as *const core::ffi::c_void,
23848 op as *mut f32,
23849 t as i32,
23850 n_head as i32,
23851 n_head_kv as i32,
23852 head_dim as i32,
23853 scale,
23854 stream.cu_stream() as *mut core::ffi::c_void,
23855 )
23856 }
23857 };
23858 if rc != 0 {
23859 return Err(format!("memra_fa3_prefill rc={rc}").into());
23860 }
23861 return Ok(());
23862 }
23863 static FA_P1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
23868 let fa_p1 = *FA_P1.get_or_init(|| std::env::var("MEMRA_FA_P1").as_deref() == Ok("1"));
23869 if fa_p1 && head_dim == 256 && !std::env::var("MEMRA_FA_FLOOR").is_ok() {
23870 const BLOCK_Q: usize = 64;
23871 const BKX: usize = 32;
23872 let f = self.func("fa_prefill_bf16_p1");
23873 let shmem = (2 * (2 * BKX * head_dim + BLOCK_Q * BKX)
23874 + 4 * (BLOCK_Q * BKX + 2 * BLOCK_Q)) as u32;
23875 use cudarc::driver::sys::CUfunction_attribute_enum as A;
23876 f.set_attribute(
23877 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
23878 shmem as i32,
23879 )?;
23880 let cfg = LaunchConfig {
23881 grid_dim: (
23882 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
23883 n_head as u32,
23884 1,
23885 ),
23886 block_dim: (32, 4, 1),
23887 shared_mem_bytes: shmem,
23888 };
23889 let (hd, nh, nhkv, ti, tkvi, cz) = (
23890 head_dim as i32,
23891 n_head as i32,
23892 n_head_kv as i32,
23893 t as i32,
23894 t_kv as i32,
23895 causal as i32,
23896 );
23897 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
23898 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
23899 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
23900 let __s_b = self.gpu.stream();
23901 let mut b = __s_b.launch_builder(&f);
23902 b.arg(&qb)
23903 .arg(&kb)
23904 .arg(&vb)
23905 .arg(o)
23906 .arg(&hd)
23907 .arg(&nh)
23908 .arg(&nhkv)
23909 .arg(&ti)
23910 .arg(&tkvi)
23911 .arg(&scale)
23912 .arg(&cz);
23913 unsafe {
23914 b.launch(cfg)?;
23915 }
23916 return Ok(());
23917 }
23918 const BK: usize = 32;
23924 let w2 = std::env::var("MEMRA_FA_PP_W2").as_deref() == Ok("1");
23927 let (block_q, warps, w2_sfx): (usize, u32, &str) =
23928 if w2 { (32, 2, "_w2") } else { (64, 4, "") };
23929 let hd_sfx = fa_hd_suffix(head_dim)?;
23933 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
23934 let bf16kv = !floor && !w2 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0");
23939 let (kb16, vb16) = if bf16kv {
23940 let n = t_kv * n_head_kv * head_dim;
23941 let mut kb = self.alloc_u8_uninit(n * 2)?;
23942 let mut vb = self.alloc_u8_uninit(n * 2)?;
23943 let fcv = self.func("f32_to_bf16_bulk");
23944 let ni = n as i64;
23945 let cfgc = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
23946 let __s_b = self.gpu.stream();
23947 let mut b = __s_b.launch_builder(&fcv);
23948 b.arg(k).arg(&mut kb).arg(&ni);
23949 unsafe {
23950 b.launch(cfgc)?;
23951 }
23952 let __s_b = self.gpu.stream();
23953 let mut b = __s_b.launch_builder(&fcv);
23954 b.arg(v).arg(&mut vb).arg(&ni);
23955 unsafe {
23956 b.launch(cfgc)?;
23957 }
23958 (Some(kb), Some(vb))
23959 } else {
23960 (None, None)
23961 };
23962 let f = self.func(&if bf16kv {
23963 format!("fa_prefill_bf16kv_pp{hd_sfx}")
23964 } else {
23965 format!(
23966 "fa_prefill_f32{}{}{hd_sfx}",
23967 if floor { "" } else { "_pp" },
23968 if floor { "" } else { w2_sfx }
23969 )
23970 });
23971 let kv_stages = if bf16kv { 2 } else { 1 };
23974 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
23975 + 4 * (block_q * BK + 2 * block_q)) as u32;
23976 use cudarc::driver::sys::CUfunction_attribute_enum as A;
23977 f.set_attribute(
23978 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
23979 shmem as i32,
23980 )?;
23981 let cfg = LaunchConfig {
23982 grid_dim: (
23983 (t as u32 + block_q as u32 - 1) / block_q as u32,
23984 n_head as u32,
23985 1,
23986 ),
23987 block_dim: (32, warps, 1),
23988 shared_mem_bytes: shmem,
23989 };
23990 let (hd, nh, nhkv, ti, tkvi, cz) = (
23991 head_dim as i32,
23992 n_head as i32,
23993 n_head_kv as i32,
23994 t as i32,
23995 t_kv as i32,
23996 causal as i32,
23997 );
23998 let __s_b = self.gpu.stream();
23999 let mut b = __s_b.launch_builder(&f);
24000 b.arg(q);
24001 match (&kb16, &vb16) {
24002 (Some(kb), Some(vb)) => {
24003 b.arg(kb).arg(vb);
24004 }
24005 _ => {
24006 b.arg(k).arg(v);
24007 }
24008 }
24009 b.arg(o)
24010 .arg(&hd)
24011 .arg(&nh)
24012 .arg(&nhkv)
24013 .arg(&ti)
24014 .arg(&tkvi)
24015 .arg(&scale)
24016 .arg(&cz);
24017 unsafe {
24018 b.launch(cfg)?;
24019 }
24020 Ok(())
24021 }
24022
24023 #[allow(clippy::too_many_arguments)]
24027 pub fn fa_prefill_w(
24028 &self,
24029 q: &CudaSlice<f32>,
24030 k: &CudaSlice<f32>,
24031 v: &CudaSlice<f32>,
24032 o: &mut CudaSlice<f32>,
24033 head_dim: usize,
24034 n_head: usize,
24035 n_head_kv: usize,
24036 t: usize,
24037 t_kv: usize,
24038 scale: f32,
24039 causal: bool,
24040 window: usize,
24041 ) -> Result<(), Box<dyn std::error::Error>> {
24042 if portable_mma_gated() {
24045 return self.sdpa_naive_w(
24046 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
24047 );
24048 }
24049 static FAW_F32: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24053 let faw_f32 =
24054 *FAW_F32.get_or_init(|| std::env::var("MEMRA_FAW_STAGE").as_deref() == Ok("f32"));
24055 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
24056 self.fa_prefill_w_arm(
24057 q,
24058 k,
24059 v,
24060 o,
24061 head_dim,
24062 n_head,
24063 n_head_kv,
24064 t,
24065 t_kv,
24066 scale,
24067 causal,
24068 window,
24069 floor || faw_f32,
24070 floor,
24071 )
24072 }
24073
24074 #[allow(clippy::too_many_arguments)]
24077 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_w_pre(
24079 &self,
24080 qb: &CudaSlice<u8>,
24081 kb: &CudaSlice<u8>,
24082 vb: &CudaSlice<u8>,
24083 o: &mut CudaSlice<f32>,
24084 head_dim: usize,
24085 n_head: usize,
24086 n_head_kv: usize,
24087 t: usize,
24088 t_kv: usize,
24089 scale: f32,
24090 causal: bool,
24091 window: usize,
24092 v_f16: bool,
24093 ) -> Result<(), Box<dyn std::error::Error>> {
24094 const BLOCK_Q: usize = 64;
24095 const BK: usize = 32;
24096 debug_assert_eq!(head_dim, 256);
24097 let hp = fa_f16pv_on()
24098 && faw_hp_on()
24099 && n_head.is_multiple_of(2)
24100 && (n_head / n_head_kv).is_multiple_of(2);
24101 debug_assert!(!v_f16 || hp, "f16 V emitted but the SWA hp arm is off");
24102 if hp {
24103 const BLOCK_QH: usize = 32;
24104 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
24107 let vh: &CudaSlice<u8> = if v_f16 {
24108 vb
24109 } else {
24110 let n = t_kv * n_head_kv * head_dim;
24111 if vguard.as_ref().map(|b| b.len() < n * 2).unwrap_or(true) {
24112 *vguard = Some(self.alloc_uninit::<u8>(n * 2)?);
24113 }
24114 self.bf16_to_f16_into(vb, n, vguard.as_mut().unwrap())?;
24115 vguard.as_ref().unwrap()
24116 };
24117 let f = self.func("fa_prefill_w_bf16_p1h2");
24118 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK) + 4 * (2 * BLOCK_QH)) as u32;
24119 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24120 f.set_attribute(
24121 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24122 shmem as i32,
24123 )?;
24124 let cfg = LaunchConfig {
24125 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
24126 block_dim: (32, 4, 1),
24127 shared_mem_bytes: shmem,
24128 };
24129 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24130 head_dim as i32,
24131 n_head as i32,
24132 n_head_kv as i32,
24133 t as i32,
24134 t_kv as i32,
24135 causal as i32,
24136 window as i32,
24137 );
24138 let __s_b = self.gpu.stream();
24139 let mut b = __s_b.launch_builder(&f);
24140 b.arg(qb)
24141 .arg(kb)
24142 .arg(vh)
24143 .arg(o)
24144 .arg(&hd)
24145 .arg(&nh)
24146 .arg(&nhkv)
24147 .arg(&ti)
24148 .arg(&tkvi)
24149 .arg(&scale)
24150 .arg(&cz)
24151 .arg(&wi);
24152 unsafe {
24153 b.launch(cfg)?;
24154 }
24155 return Ok(());
24156 }
24157 let f = self.func("fa_prefill_w_bf16_p1");
24158 let shmem =
24159 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
24160 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24161 f.set_attribute(
24162 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24163 shmem as i32,
24164 )?;
24165 let cfg = LaunchConfig {
24166 grid_dim: (
24167 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
24168 n_head as u32,
24169 1,
24170 ),
24171 block_dim: (32, 4, 1),
24172 shared_mem_bytes: shmem,
24173 };
24174 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24175 head_dim as i32,
24176 n_head as i32,
24177 n_head_kv as i32,
24178 t as i32,
24179 t_kv as i32,
24180 causal as i32,
24181 window as i32,
24182 );
24183 let __s_b = self.gpu.stream();
24184 let mut b = __s_b.launch_builder(&f);
24185 b.arg(qb)
24186 .arg(kb)
24187 .arg(vb)
24188 .arg(o)
24189 .arg(&hd)
24190 .arg(&nh)
24191 .arg(&nhkv)
24192 .arg(&ti)
24193 .arg(&tkvi)
24194 .arg(&scale)
24195 .arg(&cz)
24196 .arg(&wi);
24197 unsafe {
24198 b.launch(cfg)?;
24199 }
24200 Ok(())
24201 }
24202
24203 #[allow(clippy::too_many_arguments)]
24205 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_w_arm(
24207 &self,
24208 q: &CudaSlice<f32>,
24209 k: &CudaSlice<f32>,
24210 v: &CudaSlice<f32>,
24211 o: &mut CudaSlice<f32>,
24212 head_dim: usize,
24213 n_head: usize,
24214 n_head_kv: usize,
24215 t: usize,
24216 t_kv: usize,
24217 scale: f32,
24218 causal: bool,
24219 window: usize,
24220 f32_stage: bool,
24221 floor: bool,
24222 ) -> Result<(), Box<dyn std::error::Error>> {
24223 const BLOCK_Q: usize = 64;
24224 const BK: usize = 32;
24225 debug_assert_eq!(head_dim, 256, "fa_prefill_w is stamped hd256 only");
24226 static P1_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24230 let p1 = !floor
24231 && !f32_stage
24232 && *P1_ON.get_or_init(|| {
24233 std::env::var("MEMRA_FAW_P1")
24234 .map(|v| v != "0")
24235 .unwrap_or(true)
24236 });
24237 let hp = p1
24238 && fa_f16pv_on()
24239 && faw_hp_on()
24240 && n_head.is_multiple_of(2)
24241 && (n_head / n_head_kv).is_multiple_of(2);
24242 if hp {
24243 const BLOCK_QH: usize = 32;
24244 let f = self.func("fa_prefill_w_bf16_p1h2");
24245 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK) + 4 * (2 * BLOCK_QH)) as u32;
24246 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24247 f.set_attribute(
24248 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24249 shmem as i32,
24250 )?;
24251 let cfg = LaunchConfig {
24252 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
24253 block_dim: (32, 4, 1),
24254 shared_mem_bytes: shmem,
24255 };
24256 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24257 head_dim as i32,
24258 n_head as i32,
24259 n_head_kv as i32,
24260 t as i32,
24261 t_kv as i32,
24262 causal as i32,
24263 window as i32,
24264 );
24265 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24266 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24267 let vh = self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?;
24268 let __s_b = self.gpu.stream();
24269 let mut b = __s_b.launch_builder(&f);
24270 b.arg(&qb)
24271 .arg(&kb)
24272 .arg(&vh)
24273 .arg(o)
24274 .arg(&hd)
24275 .arg(&nh)
24276 .arg(&nhkv)
24277 .arg(&ti)
24278 .arg(&tkvi)
24279 .arg(&scale)
24280 .arg(&cz)
24281 .arg(&wi);
24282 unsafe {
24283 b.launch(cfg)?;
24284 }
24285 return Ok(());
24286 }
24287 if p1 {
24288 let f = self.func("fa_prefill_w_bf16_p1");
24289 let shmem =
24290 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
24291 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24292 f.set_attribute(
24293 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24294 shmem as i32,
24295 )?;
24296 let cfg = LaunchConfig {
24297 grid_dim: (
24298 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
24299 n_head as u32,
24300 1,
24301 ),
24302 block_dim: (32, 4, 1),
24303 shared_mem_bytes: shmem,
24304 };
24305 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24306 head_dim as i32,
24307 n_head as i32,
24308 n_head_kv as i32,
24309 t as i32,
24310 t_kv as i32,
24311 causal as i32,
24312 window as i32,
24313 );
24314 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24315 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24316 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
24317 let __s_b = self.gpu.stream();
24318 let mut b = __s_b.launch_builder(&f);
24319 b.arg(&qb)
24320 .arg(&kb)
24321 .arg(&vb)
24322 .arg(o)
24323 .arg(&hd)
24324 .arg(&nh)
24325 .arg(&nhkv)
24326 .arg(&ti)
24327 .arg(&tkvi)
24328 .arg(&scale)
24329 .arg(&cz)
24330 .arg(&wi);
24331 unsafe {
24332 b.launch(cfg)?;
24333 }
24334 return Ok(());
24335 }
24336 static G4_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24339 let g4 = !floor
24340 && !f32_stage
24341 && n_head_kv == 1
24342 && n_head.is_multiple_of(4)
24343 && *G4_ON.get_or_init(|| {
24344 std::env::var("MEMRA_FAW_G4")
24345 .map(|v| v != "0")
24346 .unwrap_or(true)
24347 });
24348 if g4 {
24349 const SP_M: usize = 16;
24350 static O2_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24353 let o2 = *O2_ON.get_or_init(|| {
24354 std::env::var("MEMRA_FAW_O2")
24355 .map(|v| v != "0")
24356 .unwrap_or(true)
24357 });
24358 let f = self.func(if o2 {
24359 "fa_prefill_w_bf16_g4o2"
24360 } else {
24361 "fa_prefill_w_bf16_g4"
24362 });
24363 let shmem = if o2 {
24364 (2 * (4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M)) as u32
24365 } else {
24366 (2 * (2 * BK * head_dim + 4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M))
24367 as u32
24368 };
24369 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24370 f.set_attribute(
24371 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24372 shmem as i32,
24373 )?;
24374 let cfg = LaunchConfig {
24375 grid_dim: ((t as u32).div_ceil(SP_M as u32), (n_head / 4) as u32, 1),
24376 block_dim: (32, 4, 1),
24377 shared_mem_bytes: shmem,
24378 };
24379 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24380 head_dim as i32,
24381 n_head as i32,
24382 n_head_kv as i32,
24383 t as i32,
24384 t_kv as i32,
24385 causal as i32,
24386 window as i32,
24387 );
24388 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24389 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24390 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
24391 let __s_b = self.gpu.stream();
24392 let mut b = __s_b.launch_builder(&f);
24393 b.arg(&qb)
24394 .arg(&kb)
24395 .arg(&vb)
24396 .arg(o)
24397 .arg(&hd)
24398 .arg(&nh)
24399 .arg(&nhkv)
24400 .arg(&ti)
24401 .arg(&tkvi)
24402 .arg(&scale)
24403 .arg(&cz)
24404 .arg(&wi);
24405 unsafe {
24406 b.launch(cfg)?;
24407 }
24408 return Ok(());
24409 }
24410 let f = self.func(if floor {
24411 "fa_prefill_w_f32"
24412 } else if f32_stage {
24413 "fa_prefill_w_f32_pp"
24414 } else {
24415 "fa_prefill_w_bf16_pp"
24416 });
24417 let shmem =
24418 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
24419 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24420 f.set_attribute(
24421 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24422 shmem as i32,
24423 )?;
24424 let cfg = LaunchConfig {
24425 grid_dim: (
24426 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
24427 n_head as u32,
24428 1,
24429 ),
24430 block_dim: (32, 4, 1),
24431 shared_mem_bytes: shmem,
24432 };
24433 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
24434 head_dim as i32,
24435 n_head as i32,
24436 n_head_kv as i32,
24437 t as i32,
24438 t_kv as i32,
24439 causal as i32,
24440 window as i32,
24441 );
24442 if f32_stage {
24443 let __s_b = self.gpu.stream();
24444 let mut b = __s_b.launch_builder(&f);
24445 b.arg(q)
24446 .arg(k)
24447 .arg(v)
24448 .arg(o)
24449 .arg(&hd)
24450 .arg(&nh)
24451 .arg(&nhkv)
24452 .arg(&ti)
24453 .arg(&tkvi)
24454 .arg(&scale)
24455 .arg(&cz)
24456 .arg(&wi);
24457 unsafe {
24458 b.launch(cfg)?;
24459 }
24460 } else {
24461 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24462 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24463 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
24464 let __s_b = self.gpu.stream();
24465 let mut b = __s_b.launch_builder(&f);
24466 b.arg(&qb)
24467 .arg(&kb)
24468 .arg(&vb)
24469 .arg(o)
24470 .arg(&hd)
24471 .arg(&nh)
24472 .arg(&nhkv)
24473 .arg(&ti)
24474 .arg(&tkvi)
24475 .arg(&scale)
24476 .arg(&cz)
24477 .arg(&wi);
24478 unsafe {
24479 b.launch(cfg)?;
24480 }
24481 }
24482 Ok(())
24483 }
24484
24485 #[allow(clippy::too_many_arguments)]
24489 pub fn fa_prefill_hd512(
24490 &self,
24491 q: &CudaSlice<f32>,
24492 k: &CudaSlice<f32>,
24493 v: &CudaSlice<f32>,
24494 o: &mut CudaSlice<f32>,
24495 head_dim: usize,
24496 n_head: usize,
24497 n_head_kv: usize,
24498 t: usize,
24499 t_kv: usize,
24500 scale: f32,
24501 causal: bool,
24502 ) -> Result<(), Box<dyn std::error::Error>> {
24503 if portable_mma_gated() {
24505 return self.sdpa_naive(
24506 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
24507 );
24508 }
24509 static F32_STAGE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24515 let f32_stage =
24516 *F32_STAGE.get_or_init(|| std::env::var("MEMRA_FA512_STAGE").as_deref() == Ok("f32"));
24517 static SP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24521 let sp = !f32_stage
24522 && *SP_ON.get_or_init(|| {
24523 std::env::var("MEMRA_FA512_SP")
24524 .map(|v| v != "0")
24525 .unwrap_or(true)
24526 });
24527 self.fa_prefill_hd512_arm(
24528 q,
24529 k,
24530 v,
24531 o,
24532 head_dim,
24533 n_head,
24534 n_head_kv,
24535 t,
24536 t_kv,
24537 scale,
24538 causal,
24539 f32_stage,
24540 sp,
24541 sp && fa_f16pv_on(),
24542 )
24543 }
24544
24545 #[allow(clippy::too_many_arguments)]
24547 pub fn fa_prefill_hd512_pre(
24548 &self,
24549 qb: &CudaSlice<u8>,
24550 kb: &CudaSlice<u8>,
24551 vb: &CudaSlice<u8>,
24552 o: &mut CudaSlice<f32>,
24553 head_dim: usize,
24554 n_head: usize,
24555 n_head_kv: usize,
24556 t: usize,
24557 t_kv: usize,
24558 scale: f32,
24559 causal: bool,
24560 v_f16: bool,
24561 ) -> Result<(), Box<dyn std::error::Error>> {
24562 debug_assert_eq!(head_dim, 512);
24563 const SP_M: usize = 16;
24564 const BKS: usize = 32;
24565 let f16pv = fa_f16pv_on();
24569 let nw = if f16pv { fa512_wide_warps() } else { 2 };
24570 let hp = f16pv
24571 && fa512_hp_on()
24572 && n_head.is_multiple_of(2)
24573 && (n_head / n_head_kv).is_multiple_of(2);
24574 debug_assert!(!v_f16 || f16pv, "f16 V emitted without the door on");
24575 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
24576 let vref: &CudaSlice<u8> = if f16pv && !v_f16 {
24577 let n = t_kv * n_head_kv * head_dim;
24579 let need = n * 2;
24580 if vguard.as_ref().map(|b| b.len() < need).unwrap_or(true) {
24581 *vguard = Some(self.alloc_uninit::<u8>(need)?);
24582 }
24583 let dst = vguard.as_mut().unwrap();
24584 self.bf16_to_f16_into(vb, n, dst)?;
24585 vguard.as_ref().unwrap()
24586 } else {
24587 vb
24588 };
24589 let f = self.func(if hp {
24590 "fa_prefill_bf16_hd512_sp16h2"
24591 } else {
24592 match (f16pv, nw) {
24593 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
24594 (true, _) => "fa_prefill_bf16_hd512_sp16",
24595 _ => "fa_prefill_bf16_hd512_sp",
24596 }
24597 });
24598 let (nwarp, npart) = if hp {
24599 (4usize, 4usize)
24600 } else if nw > 2 {
24601 (nw, nw)
24602 } else {
24603 (2, 1)
24604 };
24605 let shmem = if hp {
24607 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS) + 4 * (2 * npart * SP_M * BKS + 2 * SP_M))
24608 as u32
24609 } else {
24610 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
24611 + 4 * (npart * SP_M * BKS + SP_M)) as u32
24612 };
24613 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24614 f.set_attribute(
24615 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24616 shmem as i32,
24617 )?;
24618 let grid_y = if hp {
24619 (n_head / 2) as u32
24620 } else {
24621 n_head as u32
24622 };
24623 let cfg = LaunchConfig {
24624 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
24625 block_dim: (32, nwarp as u32, 1),
24626 shared_mem_bytes: shmem,
24627 };
24628 let (hd, nh, nhkv, ti, tkvi, cz) = (
24629 head_dim as i32,
24630 n_head as i32,
24631 n_head_kv as i32,
24632 t as i32,
24633 t_kv as i32,
24634 causal as i32,
24635 );
24636 let __s_b = self.gpu.stream();
24637 let mut b = __s_b.launch_builder(&f);
24638 b.arg(qb)
24639 .arg(kb)
24640 .arg(vref)
24641 .arg(o)
24642 .arg(&hd)
24643 .arg(&nh)
24644 .arg(&nhkv)
24645 .arg(&ti)
24646 .arg(&tkvi)
24647 .arg(&scale)
24648 .arg(&cz);
24649 unsafe {
24650 b.launch(cfg)?;
24651 }
24652 Ok(())
24653 }
24654
24655 #[allow(clippy::too_many_arguments)]
24662 pub fn mla_attn_gathered_tc(
24663 &self,
24664 q_lat_bf: &CudaSlice<u8>, cache_bf: &CudaSlice<u8>, idx: &CudaSlice<i32>, o_lat: &mut CudaSlice<f32>, n_head: usize,
24669 kv_rank: usize,
24670 t_q: usize,
24671 width: usize,
24672 scale: f32,
24673 ) -> Result<(), Box<dyn std::error::Error>> {
24674 if kv_rank != 512 {
24675 return Err(format!(
24676 "mla_attn_gathered_tc is stamped at kv_rank 512 (the glm5_next latent width); \
24677 got {kv_rank} — the caller's door must fall back to the f32 gathered kernel"
24678 )
24679 .into());
24680 }
24681 if t_q == 0 || n_head == 0 {
24682 return Ok(());
24683 }
24684 const SP_M: usize = 16;
24685 const BKS: usize = 32;
24686 const HD: usize = 512;
24687 let f = self.func("fa_mla_gathered_bf16");
24688 let shmem =
24690 (2 * (SP_M * HD + BKS * HD + SP_M * BKS) + 4 * (SP_M * BKS + SP_M) + 4 * BKS) as u32;
24691 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24692 f.set_attribute(
24693 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24694 shmem as i32,
24695 )?;
24696 let cfg = LaunchConfig {
24697 grid_dim: (t_q as u32, (n_head as u32).div_ceil(SP_M as u32), 1),
24698 block_dim: (32, 2, 1),
24699 shared_mem_bytes: shmem,
24700 };
24701 let (nh, tq, w) = (n_head as i32, t_q as i32, width as i32);
24702 let __s_b = self.gpu.stream();
24703 let mut b = __s_b.launch_builder(&f);
24704 b.arg(q_lat_bf)
24705 .arg(cache_bf)
24706 .arg(idx)
24707 .arg(o_lat)
24708 .arg(&nh)
24709 .arg(&tq)
24710 .arg(&w)
24711 .arg(&scale);
24712 unsafe {
24713 b.launch(cfg)?;
24714 }
24715 Ok(())
24716 }
24717
24718 #[allow(clippy::too_many_arguments)]
24721 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_hd512_arm(
24723 &self,
24724 q: &CudaSlice<f32>,
24725 k: &CudaSlice<f32>,
24726 v: &CudaSlice<f32>,
24727 o: &mut CudaSlice<f32>,
24728 head_dim: usize,
24729 n_head: usize,
24730 n_head_kv: usize,
24731 t: usize,
24732 t_kv: usize,
24733 scale: f32,
24734 causal: bool,
24735 f32_stage: bool,
24736 sp: bool,
24737 f16pv: bool,
24738 ) -> Result<(), Box<dyn std::error::Error>> {
24739 debug_assert_eq!(head_dim, 512, "fa_prefill_hd512 is hd512 only");
24740 if sp && !f32_stage {
24741 const SP_M: usize = 16;
24745 const BKS: usize = 32;
24746 let nw = if f16pv { fa512_wide_warps() } else { 2 };
24747 let hp = f16pv
24748 && fa512_hp_on()
24749 && n_head.is_multiple_of(2)
24750 && (n_head / n_head_kv).is_multiple_of(2);
24751 let f = self.func(if hp {
24752 "fa_prefill_bf16_hd512_sp16h2"
24753 } else {
24754 match (f16pv, nw) {
24755 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
24756 (true, _) => "fa_prefill_bf16_hd512_sp16",
24757 _ => "fa_prefill_bf16_hd512_sp",
24758 }
24759 });
24760 let (nwarp, npart) = if hp {
24761 (4usize, 4usize)
24762 } else if nw > 2 {
24763 (nw, nw)
24764 } else {
24765 (2, 1)
24766 };
24767 let shmem = if hp {
24768 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
24769 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
24770 } else {
24771 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
24772 + 4 * (npart * SP_M * BKS + SP_M)) as u32
24773 };
24774 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24775 f.set_attribute(
24776 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24777 shmem as i32,
24778 )?;
24779 let grid_y = if hp {
24780 (n_head / 2) as u32
24781 } else {
24782 n_head as u32
24783 };
24784 let cfg = LaunchConfig {
24785 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
24786 block_dim: (32, nwarp as u32, 1),
24787 shared_mem_bytes: shmem,
24788 };
24789 let (hd, nh, nhkv, ti, tkvi, cz) = (
24790 head_dim as i32,
24791 n_head as i32,
24792 n_head_kv as i32,
24793 t as i32,
24794 t_kv as i32,
24795 causal as i32,
24796 );
24797 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24798 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24799 let vb = if f16pv {
24800 self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?
24801 } else {
24802 self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?
24803 };
24804 let __s_b = self.gpu.stream();
24805 let mut b = __s_b.launch_builder(&f);
24806 b.arg(&qb)
24807 .arg(&kb)
24808 .arg(&vb)
24809 .arg(o)
24810 .arg(&hd)
24811 .arg(&nh)
24812 .arg(&nhkv)
24813 .arg(&ti)
24814 .arg(&tkvi)
24815 .arg(&scale)
24816 .arg(&cz);
24817 unsafe {
24818 b.launch(cfg)?;
24819 }
24820 return Ok(());
24821 }
24822 const BLOCK_Q: usize = 32;
24823 const BK: usize = 32;
24824 const HALF: usize = 256;
24825 let f = self.func(if f32_stage {
24826 "fa_prefill_f32_hd512"
24827 } else {
24828 "fa_prefill_bf16_hd512"
24829 });
24830 let shmem = (2 * (BLOCK_Q * head_dim + BK * head_dim + BK * HALF + BLOCK_Q * BK)
24832 + 4 * BLOCK_Q) as u32;
24833 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24834 f.set_attribute(
24835 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24836 shmem as i32,
24837 )?;
24838 let cfg = LaunchConfig {
24839 grid_dim: (
24840 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
24841 n_head as u32,
24842 2,
24843 ),
24844 block_dim: (32, 2, 1),
24845 shared_mem_bytes: shmem,
24846 };
24847 let (hd, nh, nhkv, ti, tkvi, cz) = (
24848 head_dim as i32,
24849 n_head as i32,
24850 n_head_kv as i32,
24851 t as i32,
24852 t_kv as i32,
24853 causal as i32,
24854 );
24855 if f32_stage {
24856 let __s_b = self.gpu.stream();
24857 let mut b = __s_b.launch_builder(&f);
24858 b.arg(q)
24859 .arg(k)
24860 .arg(v)
24861 .arg(o)
24862 .arg(&hd)
24863 .arg(&nh)
24864 .arg(&nhkv)
24865 .arg(&ti)
24866 .arg(&tkvi)
24867 .arg(&scale)
24868 .arg(&cz);
24869 unsafe {
24870 b.launch(cfg)?;
24871 }
24872 } else {
24873 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
24874 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
24875 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
24876 let __s_b = self.gpu.stream();
24877 let mut b = __s_b.launch_builder(&f);
24878 b.arg(&qb)
24879 .arg(&kb)
24880 .arg(&vb)
24881 .arg(o)
24882 .arg(&hd)
24883 .arg(&nh)
24884 .arg(&nhkv)
24885 .arg(&ti)
24886 .arg(&tkvi)
24887 .arg(&scale)
24888 .arg(&cz);
24889 unsafe {
24890 b.launch(cfg)?;
24891 }
24892 }
24893 Ok(())
24894 }
24895
24896 #[allow(clippy::too_many_arguments)]
24900 pub fn rope_neox2_bf16e(
24901 &self,
24902 q: &mut CudaSlice<f32>,
24903 k: &mut CudaSlice<f32>,
24904 qb: &mut CudaSlice<u8>,
24905 kb: &mut CudaSlice<u8>,
24906 pos: &CudaSlice<i32>,
24907 head_dim: usize,
24908 n_dims: usize,
24909 nh_q: usize,
24910 nh_k: usize,
24911 n_tokens: usize,
24912 base: f32,
24913 freq_scale: f32,
24914 ff: Option<&CudaSlice<f32>>,
24915 ) -> Result<(), Box<dyn std::error::Error>> {
24916 let f = self.func("rope_neox2_bf16e_f32");
24917 let rows = ((nh_q + nh_k) * n_tokens) as u32;
24918 let cfg = LaunchConfig {
24919 grid_dim: (rows, 1, 1),
24920 block_dim: ((head_dim / 2) as u32, 1, 1),
24921 shared_mem_bytes: 0,
24922 };
24923 let theta_scale = base.powf(-2.0 / n_dims as f32);
24924 let (hd, nd, nhq, nhk, nt) = (
24925 head_dim as i32,
24926 n_dims as i32,
24927 nh_q as i32,
24928 nh_k as i32,
24929 n_tokens as i32,
24930 );
24931 let __s_b = self.gpu.stream();
24932 let mut b = __s_b.launch_builder(&f);
24933 match ff {
24934 Some(t) => {
24935 b.arg(&mut *q)
24936 .arg(&mut *k)
24937 .arg(&mut *qb)
24938 .arg(&mut *kb)
24939 .arg(pos)
24940 .arg(&hd)
24941 .arg(&nd)
24942 .arg(&nhq)
24943 .arg(&nhk)
24944 .arg(&nt)
24945 .arg(&theta_scale)
24946 .arg(&freq_scale)
24947 .arg(t);
24948 unsafe {
24949 b.launch(cfg)?;
24950 }
24951 }
24952 None => {
24953 let null: u64 = 0;
24954 b.arg(&mut *q)
24955 .arg(&mut *k)
24956 .arg(&mut *qb)
24957 .arg(&mut *kb)
24958 .arg(pos)
24959 .arg(&hd)
24960 .arg(&nd)
24961 .arg(&nhq)
24962 .arg(&nhk)
24963 .arg(&nt)
24964 .arg(&theta_scale)
24965 .arg(&freq_scale)
24966 .arg(&null);
24967 unsafe {
24968 b.launch(cfg)?;
24969 }
24970 }
24971 }
24972 Ok(())
24973 }
24974
24975 pub fn f32_to_bf16(
24978 &self,
24979 x: &CudaSlice<f32>,
24980 n: usize,
24981 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
24982 assert!(
24983 n.is_multiple_of(4),
24984 "f32_to_bf16 requires n % 4 == 0, got {n}"
24985 );
24986 let mut y = self.alloc_uninit::<u8>(n * 2)?;
24987 let f = self.func("f32_to_bf16_flat");
24988 let n_i = n as i64;
24989 let cfg = LaunchConfig {
24990 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
24991 block_dim: (256, 1, 1),
24992 shared_mem_bytes: 0,
24993 };
24994 let __s_b = self.gpu.stream();
24995 let mut b = __s_b.launch_builder(&f);
24996 b.arg(x).arg(&mut y).arg(&n_i);
24997 unsafe {
24998 b.launch(cfg)?;
24999 }
25000 Ok(y)
25001 }
25002
25003 pub fn f32_to_f16(
25004 &self,
25005 x: &CudaSlice<f32>,
25006 n: usize,
25007 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
25008 assert!(
25009 n.is_multiple_of(4),
25010 "f32_to_f16 requires n % 4 == 0, got {n}"
25011 );
25012 let mut y = self.alloc_uninit::<u8>(n * 2)?;
25013 let f = self.func("f32_to_f16_flat");
25014 let n_i = n as i64;
25015 let cfg = LaunchConfig {
25016 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
25017 block_dim: (256, 1, 1),
25018 shared_mem_bytes: 0,
25019 };
25020 let __s_b = self.gpu.stream();
25021 let mut b = __s_b.launch_builder(&f);
25022 b.arg(x).arg(&mut y).arg(&n_i);
25023 unsafe {
25024 b.launch(cfg)?;
25025 }
25026 Ok(y)
25027 }
25028
25029 pub fn bf16_to_f16(
25031 &self,
25032 xb: &CudaSlice<u8>,
25033 n: usize,
25034 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
25035 let mut y = self.alloc_uninit::<u8>(n * 2)?;
25036 self.bf16_to_f16_into(xb, n, &mut y)?;
25037 Ok(y)
25038 }
25039
25040 pub fn bf16_to_f16_into(
25042 &self,
25043 xb: &CudaSlice<u8>,
25044 n: usize,
25045 y: &mut CudaSlice<u8>,
25046 ) -> Result<(), Box<dyn std::error::Error>> {
25047 assert!(
25048 n.is_multiple_of(2),
25049 "bf16_to_f16 requires n % 2 == 0, got {n}"
25050 );
25051 assert!(y.len() >= n * 2);
25052 let f = self.func("bf16_to_f16_flat");
25053 let n2 = (n / 2) as i64;
25054 let cfg = LaunchConfig {
25055 grid_dim: (((n / 2) as u32).div_ceil(256), 1, 1),
25056 block_dim: (256, 1, 1),
25057 shared_mem_bytes: 0,
25058 };
25059 let __s_b = self.gpu.stream();
25060 let mut b = __s_b.launch_builder(&f);
25061 b.arg(xb).arg(y).arg(&n2);
25062 unsafe {
25063 b.launch(cfg)?;
25064 }
25065 Ok(())
25066 }
25067
25068 #[allow(clippy::too_many_arguments)]
25073 pub fn fa_prefill_vl8(
25074 &self,
25075 seqs: &[FaSeqVl],
25076 head_dim: usize,
25077 n_head: usize,
25078 n_head_kv: usize,
25079 scale: f32,
25080 ) -> Result<(), Box<dyn std::error::Error>> {
25081 const BK: usize = 32;
25082 let b = seqs.len();
25083 assert!((1..=8).contains(&b));
25084 let mut packed = [FaSeqVl::default(); 8];
25085 packed[..b].copy_from_slice(seqs);
25086 let v = FaVl8(packed);
25087 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
25088 let ept = (n_head_kv * head_dim) as i32;
25089 {
25090 let f = self.func("fa_mirror_vl");
25091 let max_n = (max_t as i64) * ept as i64;
25092 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
25093 for which in 0..2i32 {
25094 let cfg = LaunchConfig {
25095 grid_dim: (blocks, 1, b as u32),
25096 block_dim: (256, 1, 1),
25097 shared_mem_bytes: 0,
25098 };
25099 let __s_lb = self.gpu.stream();
25100 let mut lb = __s_lb.launch_builder(&f);
25101 lb.arg(&v).arg(&ept).arg(&which);
25102 unsafe {
25103 lb.launch(cfg)?;
25104 }
25105 }
25106 }
25107 let hd_sfx = fa_hd_suffix(head_dim)?;
25108 let f = self.func(&format!("fa_prefill_bf16kv_vl{hd_sfx}"));
25109 let block_q = 64usize;
25110 let kv_stages = 2usize;
25111 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
25112 + 4 * (block_q * BK + 2 * block_q)) as u32;
25113 use cudarc::driver::sys::CUfunction_attribute_enum as A;
25114 f.set_attribute(
25115 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
25116 shmem as i32,
25117 )?;
25118 let cfg = LaunchConfig {
25119 grid_dim: (max_t.div_ceil(block_q as u32), n_head as u32, b as u32),
25120 block_dim: (32, 4, 1),
25121 shared_mem_bytes: shmem,
25122 };
25123 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
25124 let __s_lb = self.gpu.stream();
25125 let mut lb = __s_lb.launch_builder(&f);
25126 lb.arg(&v).arg(&hd).arg(&nh).arg(&nhkv).arg(&scale);
25127 unsafe {
25128 lb.launch(cfg)?;
25129 }
25130 Ok(())
25131 }
25132
25133 #[allow(clippy::too_many_arguments)]
25137 pub fn attn_pre_vl8(
25138 &self,
25139 seqs: &[AttnPreVl],
25140 wq: &CudaSlice<f32>,
25141 wk: &CudaSlice<f32>,
25142 head_dim: usize,
25143 rope_dims: usize,
25144 n_head: usize,
25145 n_head_kv: usize,
25146 eps: f32,
25147 freq_base: f32,
25148 freq_scale: f32,
25149 kv_dim_k: usize,
25150 kv_dim_v: usize,
25151 k_tok_bytes: usize,
25152 v_tok_bytes: usize,
25153 ) -> Result<(), Box<dyn std::error::Error>> {
25154 let b = seqs.len();
25155 assert!((1..=8).contains(&b));
25156 let mut packed = [AttnPreVl::default(); 8];
25157 packed[..b].copy_from_slice(seqs);
25158 let v = AttnPreVl8(packed);
25159 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
25160 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
25161 {
25162 let f = self.func("q_gate_split_vl");
25163 let n = max_t * (n_head * head_dim) as u32;
25164 let cfg = LaunchConfig {
25165 grid_dim: (n.div_ceil(256), 1, b as u32),
25166 block_dim: (256, 1, 1),
25167 shared_mem_bytes: 0,
25168 };
25169 let __s_lb = self.gpu.stream();
25170 let mut lb = __s_lb.launch_builder(&f);
25171 lb.arg(&v).arg(&hd).arg(&nh);
25172 unsafe {
25173 lb.launch(cfg)?;
25174 }
25175 }
25176 {
25177 let f = self.func("attn_rms_vl");
25178 let cfg = LaunchConfig {
25179 grid_dim: (max_t * n_head as u32, 2, b as u32),
25180 block_dim: (rms_block(), 1, 1),
25181 shared_mem_bytes: 0,
25182 };
25183 let __s_lb = self.gpu.stream();
25184 let mut lb = __s_lb.launch_builder(&f);
25185 lb.arg(&v)
25186 .arg(wq)
25187 .arg(wk)
25188 .arg(&hd)
25189 .arg(&nh)
25190 .arg(&nhkv)
25191 .arg(&eps);
25192 unsafe {
25193 lb.launch(cfg)?;
25194 }
25195 }
25196 {
25197 let f = self.func("attn_rope_vl");
25198 let theta_scale = freq_base.powf(-2.0 / rope_dims as f32);
25199 let nd = rope_dims as i32;
25200 let cfg = LaunchConfig {
25201 grid_dim: (max_t * n_head as u32, 2, b as u32),
25202 block_dim: ((head_dim / 2) as u32, 1, 1),
25203 shared_mem_bytes: 0,
25204 };
25205 let __s_lb = self.gpu.stream();
25206 let mut lb = __s_lb.launch_builder(&f);
25207 lb.arg(&v)
25208 .arg(&hd)
25209 .arg(&nd)
25210 .arg(&nh)
25211 .arg(&nhkv)
25212 .arg(&theta_scale)
25213 .arg(&freq_scale);
25214 unsafe {
25215 lb.launch(cfg)?;
25216 }
25217 }
25218 {
25219 let f = self.func("append_kv_vl");
25220 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
25221 let cfg = LaunchConfig {
25222 grid_dim: (nblk, max_t, b as u32),
25223 block_dim: (32, 1, 1),
25224 shared_mem_bytes: 0,
25225 };
25226 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
25227 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
25228 let __s_lb = self.gpu.stream();
25229 let mut lb = __s_lb.launch_builder(&f);
25230 lb.arg(&v).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
25231 unsafe {
25232 lb.launch(cfg)?;
25233 }
25234 }
25235 Ok(())
25236 }
25237
25238 #[allow(clippy::too_many_arguments)]
25243 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_view(
25246 &self,
25247 q: &CudaSlice<f32>,
25248 k: &cudarc::driver::CudaView<u8>,
25249 v: &cudarc::driver::CudaView<u8>,
25250 o: &mut CudaSlice<f32>,
25251 head_dim: usize,
25252 n_head: usize,
25253 n_head_kv: usize,
25254 t: usize,
25255 t_kv: usize,
25256 scale: f32,
25257 causal: bool,
25258 k_tok_bytes: usize,
25259 v_tok_bytes: usize,
25260 g: bool,
25261 ) -> Result<(), Box<dyn std::error::Error>> {
25262 if portable_mma_gated() {
25263 return self.sdpa_naive_quantized_view(
25264 q,
25265 k,
25266 v,
25267 o,
25268 head_dim,
25269 n_head,
25270 n_head_kv,
25271 t,
25272 t_kv,
25273 scale,
25274 causal,
25275 k_tok_bytes,
25276 v_tok_bytes,
25277 );
25278 }
25279 const BLOCK_Q: usize = 64;
25280 const BK: usize = 32;
25281 let name = format!("fa_prefill_q{}", fa_hd_suffix(head_dim)?);
25284 let f = if g {
25285 self.func_g(&name)
25286 } else {
25287 self.func(&name)
25288 };
25289 let shmem =
25290 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
25291 use cudarc::driver::sys::CUfunction_attribute_enum as A;
25292 f.set_attribute(
25293 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
25294 shmem as i32,
25295 )?;
25296 let cfg = LaunchConfig {
25297 grid_dim: (
25298 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
25299 n_head as u32,
25300 1,
25301 ),
25302 block_dim: (32, 4, 1),
25303 shared_mem_bytes: shmem,
25304 };
25305 let (hd, nh, nhkv, ti, tkvi, cz) = (
25306 head_dim as i32,
25307 n_head as i32,
25308 n_head_kv as i32,
25309 t as i32,
25310 t_kv as i32,
25311 causal as i32,
25312 );
25313 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
25314 let __s_b = self.gpu.stream();
25315 let mut b = __s_b.launch_builder(&f);
25316 b.arg(q)
25317 .arg(k)
25318 .arg(v)
25319 .arg(o)
25320 .arg(&hd)
25321 .arg(&nh)
25322 .arg(&nhkv)
25323 .arg(&ti)
25324 .arg(&tkvi)
25325 .arg(&scale)
25326 .arg(&cz)
25327 .arg(&ktb)
25328 .arg(&vtb);
25329 unsafe {
25330 b.launch(cfg)?;
25331 }
25332 Ok(())
25333 }
25334
25335 #[allow(clippy::too_many_arguments)]
25345 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_view_ws(
25347 &self,
25348 q: &CudaSlice<f32>,
25349 k: &cudarc::driver::CudaView<u8>,
25350 v: &cudarc::driver::CudaView<u8>,
25351 o: &mut CudaSlice<f32>,
25352 head_dim: usize,
25353 n_head: usize,
25354 n_head_kv: usize,
25355 t: usize,
25356 t_kv: usize,
25357 scale: f32,
25358 causal: bool,
25359 k_tok_bytes: usize,
25360 v_tok_bytes: usize,
25361 g: bool,
25362 ) -> Result<(), Box<dyn std::error::Error>> {
25363 if portable_mma_gated() {
25364 return self.sdpa_naive_quantized_view(
25365 q,
25366 k,
25367 v,
25368 o,
25369 head_dim,
25370 n_head,
25371 n_head_kv,
25372 t,
25373 t_kv,
25374 scale,
25375 causal,
25376 k_tok_bytes,
25377 v_tok_bytes,
25378 );
25379 }
25380 const BLOCK_Q: usize = 64;
25381 const BK: usize = 32;
25382 let kv_dim_k = n_head_kv * head_dim;
25383 let kv_dim_v = n_head_kv * head_dim;
25384 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
25386 let mut guard = self.prime_deqw_ws.lock().unwrap();
25388 let need_grow = match guard.as_ref() {
25389 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
25390 None => true,
25391 };
25392 if need_grow {
25393 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
25394 let (ck, cv) = guard
25395 .as_ref()
25396 .map(|(a, b)| (a.len(), b.len()))
25397 .unwrap_or((0, 0));
25398 *guard = Some((
25399 self.alloc_u8(grow(ck, k_ws_bytes))?,
25400 self.alloc_u8(grow(cv, v_ws_bytes))?,
25401 ));
25402 }
25403 let (kw, vw) = guard.as_mut().unwrap();
25404 {
25406 let f = if g {
25408 self.func_g("fa_dequant_kv_ws_bf16")
25409 } else {
25410 self.func("fa_dequant_kv_ws_bf16")
25411 };
25412 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
25413 #[allow(clippy::manual_div_ceil)]
25414 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
25416 let cfg = LaunchConfig {
25417 grid_dim: (nblk.max(1), 1, 1),
25418 block_dim: (256, 1, 1),
25419 shared_mem_bytes: 0,
25420 };
25421 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
25422 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
25423 let __s_b = self.gpu.stream();
25424 let mut b = __s_b.launch_builder(&f);
25425 b.arg(k)
25426 .arg(v)
25427 .arg(&mut *kw)
25428 .arg(&mut *vw)
25429 .arg(&kdk)
25430 .arg(&kdv)
25431 .arg(&tkvi)
25432 .arg(&ktb)
25433 .arg(&vtb);
25434 unsafe {
25435 b.launch(cfg)?;
25436 }
25437 }
25438 let db = std::env::var("MEMRA_PRIME_DEQW_DB")
25446 .map(|v| v != "0")
25447 .unwrap_or(true);
25448 {
25449 let hd_sfx = fa_hd_suffix(head_dim)?;
25450 let f = self.func(&format!(
25451 "fa_prefill_qw{}{hd_sfx}",
25452 if db { "_db" } else { "" }
25453 ));
25454 let shmem = if db {
25455 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
25457 } else {
25458 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
25459 };
25460 use cudarc::driver::sys::CUfunction_attribute_enum as A;
25461 f.set_attribute(
25462 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
25463 shmem as i32,
25464 )?;
25465 let cfg = LaunchConfig {
25466 grid_dim: (
25467 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
25468 n_head as u32,
25469 1,
25470 ),
25471 block_dim: (32, 4, 1),
25472 shared_mem_bytes: shmem,
25473 };
25474 let (hd, nh, nhkv, ti, tkvi, cz) = (
25475 head_dim as i32,
25476 n_head as i32,
25477 n_head_kv as i32,
25478 t as i32,
25479 t_kv as i32,
25480 causal as i32,
25481 );
25482 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
25483 let __s_b = self.gpu.stream();
25484 let mut b = __s_b.launch_builder(&f);
25485 b.arg(q)
25486 .arg(&*kw)
25487 .arg(&*vw)
25488 .arg(o)
25489 .arg(&hd)
25490 .arg(&nh)
25491 .arg(&nhkv)
25492 .arg(&ti)
25493 .arg(&tkvi)
25494 .arg(&scale)
25495 .arg(&cz)
25496 .arg(&kdk)
25497 .arg(&kdv);
25498 unsafe {
25499 b.launch(cfg)?;
25500 }
25501 }
25502 Ok(())
25503 }
25504
25505 #[allow(clippy::too_many_arguments)]
25521 #[allow(clippy::manual_div_ceil)] pub fn fa_prefill_view_ws_w_hd128(
25523 &self,
25524 q: &CudaSlice<f32>,
25525 k: &cudarc::driver::CudaView<u8>,
25526 v: &cudarc::driver::CudaView<u8>,
25527 o: &mut CudaSlice<f32>,
25528 head_dim: usize,
25529 n_head: usize,
25530 n_head_kv: usize,
25531 t: usize,
25532 t_kv: usize,
25533 scale: f32,
25534 causal: bool,
25535 window: usize,
25536 k_tok_bytes: usize,
25537 v_tok_bytes: usize,
25538 ) -> Result<(), Box<dyn std::error::Error>> {
25539 assert_eq!(
25540 head_dim, 128,
25541 "fa_prefill_view_ws_w_hd128: only the hd128 twin is stamped"
25542 );
25543 if portable_mma_gated() {
25544 return self.sdpa_naive_w_quantized_view(
25545 q,
25546 k,
25547 v,
25548 o,
25549 head_dim,
25550 n_head,
25551 n_head_kv,
25552 t,
25553 t_kv,
25554 scale,
25555 causal,
25556 window,
25557 k_tok_bytes,
25558 v_tok_bytes,
25559 );
25560 }
25561 const BLOCK_Q: usize = 64;
25562 const BK: usize = 32;
25563 let kv_dim_k = n_head_kv * head_dim;
25564 let kv_dim_v = n_head_kv * head_dim;
25565 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
25567 let mut guard = self.prime_deqw_ws.lock().unwrap();
25568 let need_grow = match guard.as_ref() {
25569 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
25570 None => true,
25571 };
25572 if need_grow {
25573 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
25574 let (ck, cv) = guard
25575 .as_ref()
25576 .map(|(a, b)| (a.len(), b.len()))
25577 .unwrap_or((0, 0));
25578 *guard = Some((
25579 self.alloc_u8(grow(ck, k_ws_bytes))?,
25580 self.alloc_u8(grow(cv, v_ws_bytes))?,
25581 ));
25582 }
25583 let (kw, vw) = guard.as_mut().unwrap();
25584 {
25587 let f = self.func("fa_dequant_kv_ws_bf16");
25588 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
25589 #[allow(clippy::manual_div_ceil)]
25590 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
25592 let cfg = LaunchConfig {
25593 grid_dim: (nblk.max(1), 1, 1),
25594 block_dim: (256, 1, 1),
25595 shared_mem_bytes: 0,
25596 };
25597 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
25598 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
25599 let __s_b = self.gpu.stream();
25600 let mut b = __s_b.launch_builder(&f);
25601 b.arg(k)
25602 .arg(v)
25603 .arg(&mut *kw)
25604 .arg(&mut *vw)
25605 .arg(&kdk)
25606 .arg(&kdv)
25607 .arg(&tkvi)
25608 .arg(&ktb)
25609 .arg(&vtb);
25610 unsafe {
25611 b.launch(cfg)?;
25612 }
25613 }
25614 let db = std::env::var("MEMRA_PRIME_DEQW_DB")
25616 .map(|v| v != "0")
25617 .unwrap_or(true);
25618 {
25619 let f = self.func(if db {
25620 "fa_prefill_qw_db_w_hd128"
25621 } else {
25622 "fa_prefill_qw_w_hd128"
25623 });
25624 let shmem = if db {
25625 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
25626 } else {
25627 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
25628 };
25629 use cudarc::driver::sys::CUfunction_attribute_enum as A;
25630 f.set_attribute(
25631 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
25632 shmem as i32,
25633 )?;
25634 let cfg = LaunchConfig {
25635 grid_dim: (
25636 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
25637 n_head as u32,
25638 1,
25639 ),
25640 block_dim: (32, 4, 1),
25641 shared_mem_bytes: shmem,
25642 };
25643 let (hd, nh, nhkv, ti, tkvi, cz) = (
25644 head_dim as i32,
25645 n_head as i32,
25646 n_head_kv as i32,
25647 t as i32,
25648 t_kv as i32,
25649 causal as i32,
25650 );
25651 let (kdk, kdv, wnd) = (kv_dim_k as i32, kv_dim_v as i32, window as i32);
25652 let __s_b = self.gpu.stream();
25653 let mut b = __s_b.launch_builder(&f);
25654 b.arg(q)
25655 .arg(&*kw)
25656 .arg(&*vw)
25657 .arg(o)
25658 .arg(&hd)
25659 .arg(&nh)
25660 .arg(&nhkv)
25661 .arg(&ti)
25662 .arg(&tkvi)
25663 .arg(&scale)
25664 .arg(&cz)
25665 .arg(&kdk)
25666 .arg(&kdv)
25667 .arg(&wnd);
25668 unsafe {
25669 b.launch(cfg)?;
25670 }
25671 }
25672 Ok(())
25673 }
25674
25675 #[allow(clippy::too_many_arguments)] pub fn fa_decode(
25680 &self,
25681 q: &CudaSlice<f32>,
25682 k: &cudarc::driver::CudaView<u8>,
25683 v: &cudarc::driver::CudaView<u8>,
25684 o: &mut CudaSlice<f32>,
25685 head_dim: usize,
25686 n_head: usize,
25687 n_head_kv: usize,
25688 t_kv: usize,
25689 scale: f32,
25690 k_tok_bytes: usize,
25691 v_tok_bytes: usize,
25692 ) -> Result<(), Box<dyn std::error::Error>> {
25693 self.fa_decode_kvmod(
25694 q,
25695 k,
25696 v,
25697 o,
25698 head_dim,
25699 n_head,
25700 n_head_kv,
25701 t_kv,
25702 scale,
25703 k_tok_bytes,
25704 v_tok_bytes,
25705 false,
25706 )
25707 }
25708
25709 #[allow(clippy::too_many_arguments)]
25713 #[allow(clippy::too_many_arguments)]
25717 #[allow(clippy::too_many_arguments)]
25718 fn fa_decode_scalar_unified(
25719 &self,
25720 q: &cudarc::driver::CudaView<f32>,
25721 k: &cudarc::driver::CudaView<u8>,
25722 v: &cudarc::driver::CudaView<u8>,
25723 o: &mut cudarc::driver::CudaViewMut<f32>,
25724 head_dim: usize,
25725 n_head: usize,
25726 n_head_kv: usize,
25727 t_kv_host: usize,
25728 t_kv_dev: Option<&CudaSlice<i32>>,
25729 scale: f32,
25730 n_splits: usize,
25731 split_keys: usize,
25732 k_tok_bytes: usize,
25733 v_tok_bytes: usize,
25734 g: bool,
25735 part_o: &mut CudaSlice<f32>,
25736 part_m: &mut CudaSlice<f32>,
25737 part_l: &mut CudaSlice<f32>,
25738 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
25739 ) -> Result<(), Box<dyn std::error::Error>> {
25740 let f = if g {
25741 self.func_g("fa_decode_f32")
25742 } else {
25743 self.fa_func("fa_decode_f32", head_dim)
25744 };
25745 let cfg = LaunchConfig {
25746 grid_dim: (n_head as u32, n_splits as u32, 1),
25747 block_dim: (head_dim as u32, 1, 1),
25748 shared_mem_bytes: (4 * (head_dim + 32)) as u32,
25749 };
25750 let (hd, nh, nhkv, nsp) = (
25751 head_dim as i32,
25752 n_head as i32,
25753 n_head_kv as i32,
25754 n_splits as i32,
25755 );
25756 let (ktb, vtb, tkvi, ski) = (
25757 k_tok_bytes as i64,
25758 v_tok_bytes as i64,
25759 t_kv_host as i32,
25760 split_keys as i32,
25761 );
25762 let __s_b = self.gpu.stream();
25763 let mut b = __s_b.launch_builder(&f);
25764 match t_kv_dev {
25765 Some(d) => {
25766 b.arg(q)
25767 .arg(k)
25768 .arg(v)
25769 .arg(&mut *part_o)
25770 .arg(&mut *part_m)
25771 .arg(&mut *part_l)
25772 .arg(&hd)
25773 .arg(&nh)
25774 .arg(&nhkv)
25775 .arg(&tkvi)
25776 .arg(d)
25777 .arg(&scale)
25778 .arg(&nsp)
25779 .arg(&ski)
25780 .arg(&ktb)
25781 .arg(&vtb);
25782 unsafe {
25783 b.launch(cfg)?;
25784 }
25785 }
25786 None => {
25787 let null: u64 = 0;
25788 b.arg(q)
25789 .arg(k)
25790 .arg(v)
25791 .arg(&mut *part_o)
25792 .arg(&mut *part_m)
25793 .arg(&mut *part_l)
25794 .arg(&hd)
25795 .arg(&nh)
25796 .arg(&nhkv)
25797 .arg(&tkvi)
25798 .arg(&null)
25799 .arg(&scale)
25800 .arg(&nsp)
25801 .arg(&ski)
25802 .arg(&ktb)
25803 .arg(&vtb);
25804 unsafe {
25805 b.launch(cfg)?;
25806 }
25807 }
25808 }
25809 let cfg2 = LaunchConfig {
25810 grid_dim: (n_head as u32, 1, 1),
25811 block_dim: (head_dim as u32, 1, 1),
25812 shared_mem_bytes: 0,
25813 };
25814 if let Some((oq, od)) = q8_out {
25815 let fc = if g {
25817 self.func_g("fa_decode_combine_q8_1")
25818 } else {
25819 self.fa_func("fa_decode_combine_q8_1", head_dim)
25820 };
25821 let __s_b2 = self.gpu.stream();
25822 let mut b2 = __s_b2.launch_builder(&fc);
25823 b2.arg(&*part_o)
25824 .arg(&*part_m)
25825 .arg(&*part_l)
25826 .arg(oq)
25827 .arg(od)
25828 .arg(&hd)
25829 .arg(&nh)
25830 .arg(&nsp);
25831 unsafe {
25832 b2.launch(cfg2)?;
25833 }
25834 return Ok(());
25835 }
25836 let fc = if g {
25837 self.func_g("fa_decode_combine_f32")
25838 } else {
25839 self.fa_func("fa_decode_combine_f32", head_dim)
25840 };
25841 let __s_b2 = self.gpu.stream();
25842 let mut b2 = __s_b2.launch_builder(&fc);
25843 b2.arg(&*part_o)
25844 .arg(&*part_m)
25845 .arg(&*part_l)
25846 .arg(o)
25847 .arg(&hd)
25848 .arg(&nh)
25849 .arg(&nsp);
25850 unsafe {
25851 b2.launch(cfg2)?;
25852 }
25853 Ok(())
25854 }
25855
25856 #[allow(clippy::too_many_arguments)] pub fn fa_decode_kvmod(
25858 &self,
25859 q: &CudaSlice<f32>,
25860 k: &cudarc::driver::CudaView<u8>,
25861 v: &cudarc::driver::CudaView<u8>,
25862 o: &mut CudaSlice<f32>,
25863 head_dim: usize,
25864 n_head: usize,
25865 n_head_kv: usize,
25866 t_kv: usize,
25867 scale: f32,
25868 k_tok_bytes: usize,
25869 v_tok_bytes: usize,
25870 g: bool,
25871 ) -> Result<(), Box<dyn std::error::Error>> {
25872 let q_view = q.as_view();
25873 let mut o_view = o.as_view_mut();
25874 self.fa_decode_kvmod_view(
25875 &q_view,
25876 k,
25877 v,
25878 &mut o_view,
25879 head_dim,
25880 n_head,
25881 n_head_kv,
25882 t_kv,
25883 scale,
25884 k_tok_bytes,
25885 v_tok_bytes,
25886 g,
25887 )
25888 }
25889
25890 #[allow(clippy::too_many_arguments)]
25895 #[allow(clippy::manual_div_ceil)] pub fn fa_decode_kvmod_view(
25897 &self,
25898 q: &cudarc::driver::CudaView<f32>,
25899 k: &cudarc::driver::CudaView<u8>,
25900 v: &cudarc::driver::CudaView<u8>,
25901 o: &mut cudarc::driver::CudaViewMut<f32>,
25902 head_dim: usize,
25903 n_head: usize,
25904 n_head_kv: usize,
25905 t_kv: usize,
25906 scale: f32,
25907 k_tok_bytes: usize,
25908 v_tok_bytes: usize,
25909 g: bool,
25910 ) -> Result<(), Box<dyn std::error::Error>> {
25911 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
25932 if g && head_dim == 256 && !fa_v4_at(t_kv) {
25936 fa_vec = false;
25937 }
25938 let sp = fa_split_keys(t_kv, n_head_kv);
25939 let n_splits = if fa_vec {
25940 ((t_kv + sp - 1) / sp).max(1)
25941 } else {
25942 ((t_kv + 255) / 256).max(1)
25943 };
25944 let o_len = n_head * n_splits * head_dim;
25945 let ml_len = n_head * n_splits;
25946 let mut part_guard = self.fa_part_pool.lock().unwrap();
25947 if part_guard
25948 .as_ref()
25949 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
25950 .unwrap_or(true)
25951 {
25952 let old = part_guard.take();
25963 let (co, cm) = old
25964 .as_ref()
25965 .map(|pp| (pp.0.len(), pp.1.len()))
25966 .unwrap_or((0, 0));
25967 if let Some(old) = old {
25968 self.fa_part_retired.lock().unwrap().push(old);
25969 }
25970 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
25971 eprintln!(
25972 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
25973 co, o_len, cm, ml_len
25974 );
25975 }
25976 *part_guard =
25977 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
25978 }
25979 let pg = part_guard.as_mut().unwrap();
25980 self.gpu
25981 .stream()
25982 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
25983 self.gpu
25984 .stream()
25985 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
25986 self.gpu
25987 .stream()
25988 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
25989 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
25990 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
25991 let (hd, nh, nhkv, tkvi, nsp) = (
25992 head_dim as i32,
25993 n_head as i32,
25994 n_head_kv as i32,
25995 t_kv as i32,
25996 n_splits as i32,
25997 );
25998 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
25999 let fa_vec = fa_vec && head_dim <= 512 && head_dim.is_multiple_of(32);
26003 let fa512_min = fa512_min_tkv();
26008 let deep = fa_vec
26011 && head_dim == 256
26012 && fa_v4_at(t_kv)
26013 && !g
26014 && fa_deep_at(t_kv)
26015 && !matches!(fa_v4_mode(), "noB3" | "stage");
26016 let (f, cfg) = if fa_vec && head_dim == 512 && t_kv >= fa512_min {
26017 let gqa = (n_head / n_head_kv).max(1) as u32;
26020 let fv = self.fa_func("fa_decode_vec_q_dpl16", head_dim);
26021 (
26022 fv,
26023 LaunchConfig {
26024 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26025 block_dim: (32, gqa, 1),
26026 shared_mem_bytes: 0,
26027 },
26028 )
26029 } else if fa_vec && head_dim <= 256 {
26030 let gqa = (n_head / n_head_kv).max(1) as u32;
26031 static SMEM_TKV: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
26042 let smem_tkv = *SMEM_TKV.get_or_init(|| {
26043 std::env::var("MEMRA_FA_SMEM_TKV")
26044 .ok()
26045 .and_then(|v| v.parse().ok())
26046 .unwrap_or_else(|| {
26047 FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
26048 })
26049 });
26050 if fa_v4_at(t_kv) && head_dim == 256 {
26051 let v4name = match fa_v4_mode() {
26055 "noB3" => "fa_decode_vec_q_v4_noB3", "stage" => "fa_decode_vec_q_v4_stage", _ if deep => "fa_decode_vec_q_v4_deep",
26058 _ => "fa_decode_vec_q_v4",
26059 };
26060 let fv = if g {
26061 self.func_g(v4name)
26062 } else {
26063 self.func(v4name)
26064 };
26065 let shmem = (if deep { 12160 } else { 11520 }
26068 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
26069 use cudarc::driver::sys::CUfunction_attribute_enum as A;
26070 fv.set_attribute(
26071 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26072 shmem as i32,
26073 )?;
26074 (
26075 fv,
26076 LaunchConfig {
26077 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26078 block_dim: (32, gqa, 1),
26079 shared_mem_bytes: shmem,
26080 },
26081 )
26082 } else if fa_v3_active(head_dim) {
26083 let fv = if g {
26086 self.func_g("fa_decode_vec_q_v3")
26087 } else {
26088 self.func("fa_decode_vec_q_v3")
26089 };
26090 let shmem = (32 * head_dim * 2) as u32; (
26092 fv,
26093 LaunchConfig {
26094 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26095 block_dim: (32, gqa, 1),
26096 shared_mem_bytes: shmem,
26097 },
26098 )
26099 } else if fa_v2_on() {
26100 let fv = if g {
26104 self.func_g("fa_decode_vec_q_v2")
26105 } else {
26106 self.func("fa_decode_vec_q_v2")
26107 };
26108 let shmem = (2 * 32 * head_dim * 2) as u32; (
26110 fv,
26111 LaunchConfig {
26112 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26113 block_dim: (32, gqa, 1),
26114 shared_mem_bytes: shmem,
26115 },
26116 )
26117 } else if smem_tkv > 0 && t_kv >= smem_tkv && !g && !(head_dim == 512 && Self::gkv_on())
26118 {
26119 let fv = if g {
26123 self.func_g("fa_decode_vec_q_smem")
26124 } else {
26125 self.func("fa_decode_vec_q_smem")
26126 };
26127 let shmem = (2 * 32 * head_dim * 2) as u32; use cudarc::driver::sys::CUfunction_attribute_enum as A;
26129 fv.set_attribute(
26130 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26131 shmem as i32,
26132 )?;
26133 (
26134 fv,
26135 LaunchConfig {
26136 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26137 block_dim: (32, gqa, 1),
26138 shared_mem_bytes: shmem,
26139 },
26140 )
26141 } else {
26142 let fv = if g {
26145 self.func_g("fa_decode_vec_q")
26146 } else {
26147 self.func("fa_decode_vec_q")
26148 };
26149 (
26150 fv,
26151 LaunchConfig {
26152 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
26153 block_dim: (32, gqa, 1),
26154 shared_mem_bytes: 0,
26155 },
26156 )
26157 }
26158 } else {
26159 return self.fa_decode_scalar_unified(
26162 q,
26163 k,
26164 v,
26165 o,
26166 head_dim,
26167 n_head,
26168 n_head_kv,
26169 t_kv,
26170 None,
26171 scale,
26172 n_splits,
26173 if fa_vec { sp } else { 256 },
26174 k_tok_bytes,
26175 v_tok_bytes,
26176 g,
26177 part_o,
26178 part_m,
26179 part_l,
26180 None,
26181 );
26182 };
26183 let __s_b = self.gpu.stream();
26184 let mut b = __s_b.launch_builder(&f);
26185 b.arg(q)
26186 .arg(k)
26187 .arg(v)
26188 .arg(&mut *part_o)
26189 .arg(&mut *part_m)
26190 .arg(&mut *part_l)
26191 .arg(&hd)
26192 .arg(&nh)
26193 .arg(&nhkv)
26194 .arg(&tkvi)
26195 .arg(&scale)
26196 .arg(&nsp)
26197 .arg(&ktb)
26198 .arg(&vtb);
26199 unsafe {
26200 b.launch(cfg)?;
26201 }
26202 let (fc, cfg2) = (
26205 if g {
26206 self.func_g("fa_decode_combine_f32")
26207 } else {
26208 self.fa_func("fa_decode_combine_f32", head_dim)
26209 },
26210 LaunchConfig {
26211 grid_dim: (n_head as u32, 1, 1),
26212 block_dim: (head_dim as u32, 1, 1),
26213 shared_mem_bytes: 0,
26214 },
26215 );
26216 let __s_b2 = self.gpu.stream();
26217 let mut b2 = __s_b2.launch_builder(&fc);
26218 b2.arg(&*part_o)
26219 .arg(&*part_m)
26220 .arg(&*part_l)
26221 .arg(o)
26222 .arg(&hd)
26223 .arg(&nh)
26224 .arg(&nsp);
26225 unsafe {
26226 b2.launch(cfg2)?;
26227 }
26228 Ok(())
26229 }
26230
26231 #[allow(clippy::too_many_arguments)]
26242 pub fn fa_decode_batch_seqs_v4(
26243 &self,
26244 q: &CudaSlice<f32>,
26245 kv_ptrs: &cudarc::driver::CudaView<u64>,
26246 pos_seq: &CudaSlice<i32>,
26247 o: &mut CudaSlice<f32>,
26248 head_dim: usize,
26249 n_head: usize,
26250 n_head_kv: usize,
26251 b_n: usize,
26252 t_kv_max: usize,
26253 scale: f32,
26254 split_keys: usize,
26255 k_tok_bytes: usize,
26256 v_tok_bytes: usize,
26257 ) -> Result<(), Box<dyn std::error::Error>> {
26258 debug_assert!(head_dim == 256, "seqs twin is v4-stamped (hd256 only)");
26259 #[allow(clippy::manual_div_ceil)]
26260 let n_splits_max = (t_kv_max + split_keys - 1) / split_keys;
26262 let o_len = b_n * n_head * n_splits_max * head_dim;
26263 let ml_len = b_n * n_head * n_splits_max;
26264 let mut part_guard = self.fa_part_pool.lock().unwrap();
26265 if part_guard
26266 .as_ref()
26267 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
26268 .unwrap_or(true)
26269 {
26270 let old = part_guard.take();
26281 let (co, cm) = old
26282 .as_ref()
26283 .map(|pp| (pp.0.len(), pp.1.len()))
26284 .unwrap_or((0, 0));
26285 if let Some(old) = old {
26286 self.fa_part_retired.lock().unwrap().push(old);
26287 }
26288 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
26289 eprintln!(
26290 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
26291 co, o_len, cm, ml_len
26292 );
26293 }
26294 *part_guard =
26295 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
26296 }
26297 let pg = part_guard.as_mut().unwrap();
26298 self.gpu
26299 .stream()
26300 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
26301 self.gpu
26302 .stream()
26303 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
26304 self.gpu
26305 .stream()
26306 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
26307 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
26308 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
26309 let (nspm, spk) = (n_splits_max as i32, split_keys as i32);
26310 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
26311 let gqa = (n_head / n_head_kv).max(1) as u32;
26312 let f = self.func("fa_decode_vec_q_seqs_v4");
26313 let shmem = (11520 + 32 * head_dim * 2) as u32;
26315 use cudarc::driver::sys::CUfunction_attribute_enum as A;
26316 f.set_attribute(
26317 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26318 shmem as i32,
26319 )?;
26320 let cfg = LaunchConfig {
26321 grid_dim: (n_head_kv as u32, n_splits_max as u32, b_n as u32),
26322 block_dim: (32, gqa, 1),
26323 shared_mem_bytes: shmem,
26324 };
26325 {
26326 let __s_b = self.gpu.stream();
26327 let mut b = __s_b.launch_builder(&f);
26328 b.arg(q)
26329 .arg(kv_ptrs)
26330 .arg(pos_seq)
26331 .arg(&mut *part_o)
26332 .arg(&mut *part_m)
26333 .arg(&mut *part_l)
26334 .arg(&hd)
26335 .arg(&nh)
26336 .arg(&nhkv)
26337 .arg(&scale)
26338 .arg(&nspm)
26339 .arg(&spk)
26340 .arg(&ktb)
26341 .arg(&vtb);
26342 unsafe {
26343 b.launch(cfg)?;
26344 }
26345 }
26346 let fc = self.func("fa_decode_combine_seqs");
26347 let cfg2 = LaunchConfig {
26348 grid_dim: (n_head as u32, b_n as u32, 1),
26349 block_dim: (head_dim as u32, 1, 1),
26350 shared_mem_bytes: 0,
26351 };
26352 let __s_b2 = self.gpu.stream();
26353 let mut b2 = __s_b2.launch_builder(&fc);
26354 b2.arg(&*part_o)
26355 .arg(&*part_m)
26356 .arg(&*part_l)
26357 .arg(o)
26358 .arg(&hd)
26359 .arg(&nh)
26360 .arg(pos_seq)
26361 .arg(&nspm)
26362 .arg(&spk);
26363 unsafe {
26364 b2.launch(cfg2)?;
26365 }
26366 Ok(())
26367 }
26368
26369 #[allow(clippy::too_many_arguments)]
26376 pub fn append_kv_quantized_seqs(
26377 &self,
26378 k_rows: &CudaSlice<f32>,
26379 v_rows: &CudaSlice<f32>,
26380 kv_ptrs: &cudarc::driver::CudaView<u64>,
26381 pos_seq: &CudaSlice<i32>,
26382 b_n: usize,
26383 kv_dim_k: usize,
26384 kv_dim_v: usize,
26385 k_tok_bytes: usize,
26386 v_tok_bytes: usize,
26387 ) -> Result<(), Box<dyn std::error::Error>> {
26388 let f = self.func("append_quantize_kv_q8_0_q5_1_seqs");
26389 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
26390 let cfg = LaunchConfig {
26391 grid_dim: (nblk, b_n as u32, 1),
26392 block_dim: (32, 1, 1),
26393 shared_mem_bytes: 0,
26394 };
26395 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
26396 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
26397 let __s_b = self.gpu.stream();
26398 let mut b = __s_b.launch_builder(&f);
26399 b.arg(k_rows)
26400 .arg(v_rows)
26401 .arg(kv_ptrs)
26402 .arg(pos_seq)
26403 .arg(&kdk)
26404 .arg(&kdv)
26405 .arg(&ktb)
26406 .arg(&vtb);
26407 unsafe {
26408 b.launch(cfg)?;
26409 }
26410 Ok(())
26411 }
26412
26413 pub fn fa_rows_eligible(&self, base_len: usize, head_dim: usize) -> bool {
26419 std::env::var("MEMRA_NO_FA_VEC").is_err()
26420 && std::env::var("MEMRA_FA_ROWS_OFF").is_err()
26421 && base_len + 1 >= fa_vec_min_tkv()
26422 && head_dim <= 256
26423 && head_dim.is_multiple_of(32)
26424 }
26425
26426 #[allow(clippy::too_many_arguments)]
26435 pub fn fa_decode_rows(
26436 &self,
26437 q: &CudaSlice<f32>,
26438 k: &cudarc::driver::CudaView<u8>,
26439 v: &cudarc::driver::CudaView<u8>,
26440 o: &mut CudaSlice<f32>,
26441 head_dim: usize,
26442 n_head: usize,
26443 n_head_kv: usize,
26444 base_len: usize,
26445 t: usize,
26446 scale: f32,
26447 k_tok_bytes: usize,
26448 v_tok_bytes: usize,
26449 base_dev: Option<(&CudaSlice<i32>, i32)>,
26453 kv_shared: bool,
26456 g: bool,
26460 mut q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
26463 ) -> Result<(), Box<dyn std::error::Error>> {
26464 debug_assert!(
26465 base_len + 1 >= fa_vec_min_tkv() && head_dim <= 512 && head_dim.is_multiple_of(32)
26466 );
26467 let t_kv_max = base_len + t; let mut sp = fa_split_keys(t_kv_max, n_head_kv); if head_dim == 512 {
26474 static SP512: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
26475 let v = *SP512.get_or_init(|| {
26478 std::env::var("MEMRA_FA_SP512")
26479 .ok()
26480 .and_then(|x| x.parse().ok())
26481 .unwrap_or(0)
26482 });
26483 sp = if v >= 8 {
26484 v
26485 } else {
26486 FA_SP512_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
26487 };
26488 }
26489 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
26490 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
26491 let gqa = (n_head / n_head_kv).max(1) as u32;
26492 let mut groups: Vec<(usize, usize, usize)> = Vec::new(); if head_dim == 512 || fa_split_keys(base_len + 1, n_head_kv) == sp {
26503 groups.push((0, t, sp));
26504 } else {
26505 let mut r0 = 0usize;
26506 while r0 < t {
26507 let sp_g = fa_split_keys(base_len + r0 + 1, n_head_kv);
26508 let mut r1 = r0 + 1;
26509 while r1 < t && fa_split_keys(base_len + r1 + 1, n_head_kv) == sp_g {
26510 r1 += 1;
26511 }
26512 groups.push((r0, r1 - r0, sp_g));
26513 r0 = r1;
26514 }
26515 }
26516 static SMEM_TKV_R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
26520 let smem_tkv = *SMEM_TKV_R.get_or_init(|| {
26521 std::env::var("MEMRA_FA_SMEM_TKV")
26522 .ok()
26523 .and_then(|v| v.parse().ok())
26524 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
26525 });
26526 let v4 = fa_v4_at(base_len + t) && head_dim == 256;
26527 let v3 = fa_v3_active(head_dim);
26528 let smem_rows =
26529 head_dim <= 256 && !v3 && !fa_v2_on() && smem_tkv > 0 && t_kv_max >= smem_tkv;
26530 let _ = kv_shared;
26535 let i2 = head_dim == 512 && std::env::var("MEMRA_FA_I2").as_deref() != Ok("0");
26538 static TB512: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
26552 let tb512 = head_dim == 512
26554 && sp <= 32
26555 && n_head / n_head_kv.max(1) <= 16
26556 && *TB512.get_or_init(|| std::env::var("MEMRA_FA_TB512").as_deref() != Ok("0"));
26557 let fname = if tb512 {
26558 "fa_decode_vec_q_rows_v4_512_tb"
26559 } else if i2 {
26560 "fa_decode_vec_q_rows_dpl16_i2"
26561 } else if head_dim == 512 {
26562 "fa_decode_vec_q_rows_dpl16"
26563 }
26564 else if v4 {
26566 "fa_decode_vec_q_rows_v4"
26567 } else if v3 {
26568 "fa_decode_vec_q_rows_v3"
26569 } else if fa_v2_on() {
26570 "fa_decode_vec_q_rows_v2"
26571 } else if smem_rows {
26572 "fa_decode_vec_q_rows_smem"
26573 } else {
26574 "fa_decode_vec_q_rows"
26575 };
26576 let f = if head_dim == 512 {
26577 self.fa_func(fname, head_dim)
26578 } else if g {
26579 self.func_g(if smem_rows {
26587 "fa_decode_vec_q_rows"
26588 } else {
26589 fname
26590 })
26591 } else {
26592 self.func(fname)
26593 };
26594 let shmem = if tb512 {
26595 let gk = Self::gkv_on();
26597 let sh =
26598 (8192 + 1024 + 32 * 512 + 32 * 64 + 32 * head_dim * if gk { 1 } else { 2 }) as u32;
26599 use cudarc::driver::sys::CUfunction_attribute_enum as A;
26600 f.set_attribute(
26601 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26602 sh as i32,
26603 )?;
26604 sh
26605 } else if v4 || v3 || smem_rows || fa_v2_on() {
26606 let sh = (if v4 {
26608 11520 + 32 * head_dim * if g { 1 } else { 2 }
26609 } else if v3 {
26610 32 * head_dim * 2
26611 } else {
26612 2 * 32 * head_dim * 2
26613 }) as u32;
26614 use cudarc::driver::sys::CUfunction_attribute_enum as A;
26615 f.set_attribute(
26616 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26617 sh as i32,
26618 )?;
26619 sh
26620 } else {
26621 0
26622 };
26623 for &(r0, t_g, sp_g) in &groups {
26627 let n_splits_g = (base_len + r0 + t_g).div_ceil(sp_g);
26628 let (nspm, spk) = (n_splits_g as i32, sp_g as i32);
26629 let base_i = (base_len + r0) as i32;
26630 let o_len = t_g * n_head * n_splits_g * head_dim;
26631 let ml_len = t_g * n_head * n_splits_g;
26632 let mut part_guard = self.fa_part_pool.lock().unwrap();
26633 if part_guard
26634 .as_ref()
26635 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
26636 .unwrap_or(true)
26637 {
26638 let old = part_guard.take();
26649 let (co, cm) = old
26650 .as_ref()
26651 .map(|pp| (pp.0.len(), pp.1.len()))
26652 .unwrap_or((0, 0));
26653 if let Some(old) = old {
26654 self.fa_part_retired.lock().unwrap().push(old);
26655 }
26656 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
26657 eprintln!(
26658 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
26659 co, o_len, cm, ml_len
26660 );
26661 }
26662 *part_guard =
26663 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
26664 }
26665 let pg = part_guard.as_mut().unwrap();
26666 self.gpu
26667 .stream()
26668 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
26669 self.gpu
26670 .stream()
26671 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
26672 self.gpu
26673 .stream()
26674 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
26675 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
26676 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
26677 let qv = self.view(q, t * n_head * head_dim);
26678 let q_g = qv.slice(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
26679 let cfg = LaunchConfig {
26680 grid_dim: (n_head_kv as u32, n_splits_g as u32, t_g as u32),
26681 block_dim: (32, gqa, 1),
26682 shared_mem_bytes: shmem,
26683 };
26684 {
26685 let __s_b = self.gpu.stream();
26686 let mut b = __s_b.launch_builder(&f);
26687 if tb512 {
26688 let (bd, plus) =
26690 base_dev.expect("hd512 rows twin requires a device base counter");
26691 let plus_g = plus + r0 as i32;
26692 let nr = t_g as i32;
26693 if Self::pdl_on() && Self::pdl_wb_on() {
26694 use cudarc::driver::{DevicePtr, DevicePtrMut};
26696 let s = &self.gpu.stream();
26697 let (pq, _b0) = q_g.device_ptr(s);
26698 let (pk, _b1) = k.device_ptr(s);
26699 let (pv, _b2) = v.device_ptr(s);
26700 let (po, _b3) = part_o.device_ptr_mut(s);
26701 let (pm, _b4) = part_m.device_ptr_mut(s);
26702 let (pl, _b5) = part_l.device_ptr_mut(s);
26703 let (pb, _b6) = bd.device_ptr(s);
26704 let mut ps = [
26705 &pq as *const _ as *mut std::ffi::c_void,
26706 &pk as *const _ as *mut _,
26707 &pv as *const _ as *mut _,
26708 &po as *const _ as *mut _,
26709 &pm as *const _ as *mut _,
26710 &pl as *const _ as *mut _,
26711 &hd as *const _ as *mut _,
26712 &nh as *const _ as *mut _,
26713 &nhkv as *const _ as *mut _,
26714 &pb as *const _ as *mut _,
26715 &plus_g as *const _ as *mut _,
26716 &scale as *const _ as *mut _,
26717 &nspm as *const _ as *mut _,
26718 &spk as *const _ as *mut _,
26719 &ktb as *const _ as *mut _,
26720 &vtb as *const _ as *mut _,
26721 &nr as *const _ as *mut _,
26722 ];
26723 unsafe {
26724 self.launch_pdl_flash(
26725 Self::gkv_on(),
26726 "fa_decode_vec_q_rows_v4_512_tb",
26727 (n_head_kv as u32, n_splits_g as u32, 1),
26728 (32, gqa, 1),
26729 shmem,
26730 &mut ps,
26731 )?;
26732 }
26733 } else {
26734 let cfg_tb = LaunchConfig {
26735 grid_dim: (n_head_kv as u32, n_splits_g as u32, 1),
26736 block_dim: (32, gqa, 1),
26737 shared_mem_bytes: shmem,
26738 };
26739 b.arg(&q_g)
26740 .arg(k)
26741 .arg(v)
26742 .arg(&mut *part_o)
26743 .arg(&mut *part_m)
26744 .arg(&mut *part_l)
26745 .arg(&hd)
26746 .arg(&nh)
26747 .arg(&nhkv)
26748 .arg(bd)
26749 .arg(&plus_g)
26750 .arg(&scale)
26751 .arg(&nspm)
26752 .arg(&spk)
26753 .arg(&ktb)
26754 .arg(&vtb)
26755 .arg(&nr);
26756 unsafe {
26757 b.launch(cfg_tb)?;
26758 }
26759 }
26760 } else if head_dim == 512 {
26761 let (bd, plus) =
26762 base_dev.expect("hd512 rows twin requires a device base counter");
26763 let plus_g = plus + r0 as i32;
26764 b.arg(&q_g)
26765 .arg(k)
26766 .arg(v)
26767 .arg(&mut *part_o)
26768 .arg(&mut *part_m)
26769 .arg(&mut *part_l)
26770 .arg(&hd)
26771 .arg(&nh)
26772 .arg(&nhkv)
26773 .arg(bd)
26774 .arg(&plus_g)
26775 .arg(&scale)
26776 .arg(&nspm)
26777 .arg(&spk)
26778 .arg(&ktb)
26779 .arg(&vtb);
26780 unsafe {
26781 b.launch(cfg)?;
26782 }
26783 } else {
26784 b.arg(&q_g)
26785 .arg(k)
26786 .arg(v)
26787 .arg(&mut *part_o)
26788 .arg(&mut *part_m)
26789 .arg(&mut *part_l)
26790 .arg(&hd)
26791 .arg(&nh)
26792 .arg(&nhkv)
26793 .arg(&base_i)
26794 .arg(&scale)
26795 .arg(&nspm)
26796 .arg(&spk)
26797 .arg(&ktb)
26798 .arg(&vtb);
26799 unsafe {
26800 b.launch(cfg)?;
26801 }
26802 }
26803 }
26804 let cfg2 = LaunchConfig {
26805 grid_dim: (n_head as u32, t_g as u32, 1),
26806 block_dim: (head_dim as u32, 1, 1),
26807 shared_mem_bytes: 0,
26808 };
26809 let mut o_g = o.slice_mut(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
26810 if head_dim == 512 {
26811 let (bd, plus) = base_dev.unwrap();
26814 let plus_g = plus + r0 as i32;
26815 if let Some((oq, od)) = q8_out.as_mut() {
26816 debug_assert!(t == 1, "rows q8 emit is a t=1 decode arm");
26818 if Self::pdl_on() && Self::pdl_wb_on() {
26819 use cudarc::driver::{DevicePtr, DevicePtrMut};
26821 let s = &self.gpu.stream();
26822 let (po, _g0) = part_o.device_ptr(s);
26823 let (pm, _g1) = part_m.device_ptr(s);
26824 let (pl, _g2) = part_l.device_ptr(s);
26825 let (pq, _g3) = oq.device_ptr_mut(s);
26826 let (pd, _g4) = od.device_ptr_mut(s);
26827 let (pb, _g5) = bd.device_ptr(s);
26828 let mut ps = [
26829 &po as *const _ as *mut std::ffi::c_void,
26830 &pm as *const _ as *mut _,
26831 &pl as *const _ as *mut _,
26832 &pq as *const _ as *mut _,
26833 &pd as *const _ as *mut _,
26834 &hd as *const _ as *mut _,
26835 &nh as *const _ as *mut _,
26836 &pb as *const _ as *mut _,
26837 &plus_g as *const _ as *mut _,
26838 &nspm as *const _ as *mut _,
26839 &spk as *const _ as *mut _,
26840 ];
26841 unsafe {
26842 self.launch_pdl_flash(
26843 Self::gkv_on(),
26844 "fa_decode_combine_rows_dc_q8_1",
26845 cfg2.grid_dim,
26846 cfg2.block_dim,
26847 0,
26848 &mut ps,
26849 )?;
26850 }
26851 continue;
26852 }
26853 let fc = self.fa_func("fa_decode_combine_rows_dc_q8_1", head_dim);
26854 let __s_b2 = self.gpu.stream();
26855 let mut b2 = __s_b2.launch_builder(&fc);
26856 b2.arg(&*part_o)
26857 .arg(&*part_m)
26858 .arg(&*part_l)
26859 .arg(&mut **oq)
26860 .arg(&mut **od)
26861 .arg(&hd)
26862 .arg(&nh)
26863 .arg(bd)
26864 .arg(&plus_g)
26865 .arg(&nspm)
26866 .arg(&spk);
26867 unsafe {
26868 b2.launch(cfg2)?;
26869 }
26870 continue;
26871 }
26872 let fc = self.fa_func("fa_decode_combine_rows_dc", head_dim);
26873 let __s_b2 = self.gpu.stream();
26874 let mut b2 = __s_b2.launch_builder(&fc);
26875 b2.arg(&*part_o)
26876 .arg(&*part_m)
26877 .arg(&*part_l)
26878 .arg(&mut o_g)
26879 .arg(&hd)
26880 .arg(&nh)
26881 .arg(bd)
26882 .arg(&plus_g)
26883 .arg(&nspm)
26884 .arg(&spk);
26885 unsafe {
26886 b2.launch(cfg2)?;
26887 }
26888 } else {
26889 assert!(
26892 q8_out.is_none(),
26893 "rows q8 emit requires the hd512 dc combine"
26894 );
26895 let fc = self.func("fa_decode_combine_rows");
26896 let __s_b2 = self.gpu.stream();
26897 let mut b2 = __s_b2.launch_builder(&fc);
26898 b2.arg(&*part_o)
26899 .arg(&*part_m)
26900 .arg(&*part_l)
26901 .arg(&mut o_g)
26902 .arg(&hd)
26903 .arg(&nh)
26904 .arg(&base_i)
26905 .arg(&nspm)
26906 .arg(&spk);
26907 unsafe {
26908 b2.launch(cfg2)?;
26909 }
26910 }
26911 }
26912 Ok(())
26913 }
26914
26915 #[allow(clippy::too_many_arguments)]
26919 pub fn fa_decode_rows_w(
26920 &self,
26921 q: &CudaSlice<f32>,
26922 k: &cudarc::driver::CudaView<u8>,
26923 v: &cudarc::driver::CudaView<u8>,
26924 o: &mut CudaSlice<f32>,
26925 head_dim: usize,
26926 n_head: usize,
26927 n_head_kv: usize,
26928 base_dev: &CudaSlice<i32>,
26929 base_plus: i32,
26930 t: usize,
26931 scale: f32,
26932 window: usize,
26933 k_tok_bytes: usize,
26934 v_tok_bytes: usize,
26935 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
26936 ) -> Result<(), Box<dyn std::error::Error>> {
26937 debug_assert!(head_dim == 256);
26942 let sp = {
26950 static SPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
26951 let v = *SPW.get_or_init(|| {
26952 std::env::var("MEMRA_FA_SPW")
26953 .ok()
26954 .and_then(|x| x.parse().ok())
26955 .unwrap_or(0)
26956 });
26957 if v >= 8 {
26958 v
26959 } else {
26960 FA_SPW_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
26961 }
26962 };
26963 #[allow(clippy::manual_div_ceil)]
26964 let n_splits_max = (window + sp - 1) / sp;
26966 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
26967 let (nspm, spk, wini) = (n_splits_max as i32, sp as i32, window as i32);
26968 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
26969 let gqa = (n_head / n_head_kv).max(1) as u32;
26970 let o_len = t * n_head * n_splits_max * head_dim;
26971 let ml_len = t * n_head * n_splits_max;
26972 let mut part_guard = self.fa_part_pool.lock().unwrap();
26973 if part_guard
26974 .as_ref()
26975 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
26976 .unwrap_or(true)
26977 {
26978 let old = part_guard.take();
26989 let (co, cm) = old
26990 .as_ref()
26991 .map(|pp| (pp.0.len(), pp.1.len()))
26992 .unwrap_or((0, 0));
26993 if let Some(old) = old {
26994 self.fa_part_retired.lock().unwrap().push(old);
26995 }
26996 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
26997 eprintln!(
26998 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
26999 co, o_len, cm, ml_len
27000 );
27001 }
27002 *part_guard =
27003 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
27004 }
27005 let pg = part_guard.as_mut().unwrap();
27006 self.gpu
27007 .stream()
27008 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
27009 self.gpu
27010 .stream()
27011 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
27012 self.gpu
27013 .stream()
27014 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
27015 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
27016 static SMEM_TKV_W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
27022 let smem_tkv = *SMEM_TKV_W.get_or_init(|| {
27023 std::env::var("MEMRA_FA_SMEM_TKV")
27024 .ok()
27025 .and_then(|v| v.parse().ok())
27026 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
27027 });
27028 use cudarc::driver::sys::CUfunction_attribute_enum as A;
27034 let wg = Self::wkv_on();
27039 let sp2 =
27042 gqa <= 4 && fa_v4_at(window) && std::env::var("MEMRA_FA_SPW2").as_deref() != Ok("0");
27043 if sp2 {
27044 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
27045 if Self::pdl_on() && Self::pdl_wb_on() {
27046 use cudarc::driver::{DevicePtr, DevicePtrMut};
27048 let s = &self.gpu.stream();
27049 let (pq, _b0) = q.device_ptr(s);
27050 let (pk, _b1) = k.device_ptr(s);
27051 let (pv, _b2) = v.device_ptr(s);
27052 let (po, _b3) = part_o.device_ptr_mut(s);
27053 let (pm, _b4) = part_m.device_ptr_mut(s);
27054 let (pl, _b5) = part_l.device_ptr_mut(s);
27055 let (pb, _b6) = base_dev.device_ptr(s);
27056 let mut ps = [
27057 &pq as *const _ as *mut std::ffi::c_void,
27058 &pk as *const _ as *mut _,
27059 &pv as *const _ as *mut _,
27060 &po as *const _ as *mut _,
27061 &pm as *const _ as *mut _,
27062 &pl as *const _ as *mut _,
27063 &hd as *const _ as *mut _,
27064 &nh as *const _ as *mut _,
27065 &nhkv as *const _ as *mut _,
27066 &pb as *const _ as *mut _,
27067 &base_plus as *const _ as *mut _,
27068 &scale as *const _ as *mut _,
27069 &nspm as *const _ as *mut _,
27070 &spk as *const _ as *mut _,
27071 &ktb as *const _ as *mut _,
27072 &vtb as *const _ as *mut _,
27073 &wini as *const _ as *mut _,
27074 ];
27075 unsafe {
27076 self.launch_pdl_flash(
27077 wg,
27078 "fa_decode_vec_q_rows_v4_w_sp",
27079 (n_head_kv as u32, n_splits_max as u32, t as u32),
27080 (32, gqa + 1, 1),
27081 sh,
27082 &mut ps,
27083 )?;
27084 }
27085 } else {
27086 let f = if wg {
27087 self.func_g("fa_decode_vec_q_rows_v4_w_sp")
27088 } else {
27089 self.func("fa_decode_vec_q_rows_v4_w_sp")
27090 };
27091 f.set_attribute(
27092 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27093 sh as i32,
27094 )?;
27095 let cfg = LaunchConfig {
27096 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
27097 block_dim: (32, gqa + 1, 1),
27098 shared_mem_bytes: sh,
27099 };
27100 let __s_b = self.gpu.stream();
27101 let mut b = __s_b.launch_builder(&f);
27102 b.arg(q)
27103 .arg(k)
27104 .arg(v)
27105 .arg(&mut *part_o)
27106 .arg(&mut *part_m)
27107 .arg(&mut *part_l)
27108 .arg(&hd)
27109 .arg(&nh)
27110 .arg(&nhkv)
27111 .arg(base_dev)
27112 .arg(&base_plus)
27113 .arg(&scale)
27114 .arg(&nspm)
27115 .arg(&spk)
27116 .arg(&ktb)
27117 .arg(&vtb)
27118 .arg(&wini);
27119 unsafe {
27120 b.launch(cfg)?;
27121 }
27122 }
27123 } else {
27124 if fa_v4_at(window) && Self::pdl_on() && Self::pdl_wb_on() {
27125 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
27127 use cudarc::driver::{DevicePtr, DevicePtrMut};
27128 let s = &self.gpu.stream();
27129 let (pq, _b0) = q.device_ptr(s);
27130 let (pk, _b1) = k.device_ptr(s);
27131 let (pv, _b2) = v.device_ptr(s);
27132 let (po, _b3) = part_o.device_ptr_mut(s);
27133 let (pm, _b4) = part_m.device_ptr_mut(s);
27134 let (pl, _b5) = part_l.device_ptr_mut(s);
27135 let (pb, _b6) = base_dev.device_ptr(s);
27136 let mut ps = [
27137 &pq as *const _ as *mut std::ffi::c_void,
27138 &pk as *const _ as *mut _,
27139 &pv as *const _ as *mut _,
27140 &po as *const _ as *mut _,
27141 &pm as *const _ as *mut _,
27142 &pl as *const _ as *mut _,
27143 &hd as *const _ as *mut _,
27144 &nh as *const _ as *mut _,
27145 &nhkv as *const _ as *mut _,
27146 &pb as *const _ as *mut _,
27147 &base_plus as *const _ as *mut _,
27148 &scale as *const _ as *mut _,
27149 &nspm as *const _ as *mut _,
27150 &spk as *const _ as *mut _,
27151 &ktb as *const _ as *mut _,
27152 &vtb as *const _ as *mut _,
27153 &wini as *const _ as *mut _,
27154 ];
27155 unsafe {
27156 self.launch_pdl_flash(
27157 wg,
27158 "fa_decode_vec_q_rows_v4_w",
27159 (n_head_kv as u32, n_splits_max as u32, t as u32),
27160 (32, gqa, 1),
27161 sh,
27162 &mut ps,
27163 )?;
27164 }
27165 } else {
27166 let pick = |name: &str| {
27167 if wg {
27168 self.func_g(name)
27169 } else {
27170 self.func(name)
27171 }
27172 };
27173 let (f, sh) = if fa_v4_at(window) {
27174 let f = pick("fa_decode_vec_q_rows_v4_w");
27175 (f, (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32)
27176 } else if smem_tkv > 0 && window >= smem_tkv {
27177 (
27180 pick("fa_decode_vec_q_rows_smem_w"),
27181 (2 * 32 * head_dim * 2) as u32,
27182 )
27183 } else {
27184 (pick("fa_decode_vec_q_rows_reg_w"), 0u32)
27185 };
27186 f.set_attribute(
27187 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27188 sh as i32,
27189 )?;
27190 let cfg = LaunchConfig {
27191 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
27192 block_dim: (32, gqa, 1),
27193 shared_mem_bytes: sh,
27194 };
27195 let __s_b = self.gpu.stream();
27196 let mut b = __s_b.launch_builder(&f);
27197 b.arg(q)
27198 .arg(k)
27199 .arg(v)
27200 .arg(&mut *part_o)
27201 .arg(&mut *part_m)
27202 .arg(&mut *part_l)
27203 .arg(&hd)
27204 .arg(&nh)
27205 .arg(&nhkv)
27206 .arg(base_dev)
27207 .arg(&base_plus)
27208 .arg(&scale)
27209 .arg(&nspm)
27210 .arg(&spk)
27211 .arg(&ktb)
27212 .arg(&vtb)
27213 .arg(&wini);
27214 unsafe {
27215 b.launch(cfg)?;
27216 }
27217 }
27218 }
27219 let cfg2 = LaunchConfig {
27220 grid_dim: (n_head as u32, t as u32, 1),
27221 block_dim: (head_dim as u32, 1, 1),
27222 shared_mem_bytes: 0,
27223 };
27224 if let Some((oq, od)) = q8_out {
27225 if Self::pdl_on() && Self::pdl_wb_on() {
27228 use cudarc::driver::{DevicePtr, DevicePtrMut};
27230 let s = &self.gpu.stream();
27231 let (po, _g0) = part_o.device_ptr(s);
27232 let (pm, _g1) = part_m.device_ptr(s);
27233 let (pl, _g2) = part_l.device_ptr(s);
27234 let (pq, _g3) = oq.device_ptr_mut(s);
27235 let (pd, _g4) = od.device_ptr_mut(s);
27236 let mut ps = [
27237 &po as *const _ as *mut std::ffi::c_void,
27238 &pm as *const _ as *mut _,
27239 &pl as *const _ as *mut _,
27240 &pq as *const _ as *mut _,
27241 &pd as *const _ as *mut _,
27242 &hd as *const _ as *mut _,
27243 &nh as *const _ as *mut _,
27244 &nspm as *const _ as *mut _,
27245 &spk as *const _ as *mut _,
27246 &wini as *const _ as *mut _,
27247 ];
27248 unsafe {
27249 self.launch_pdl_flash(
27250 wg,
27251 "fa_decode_combine_rows_w_q8_1",
27252 cfg2.grid_dim,
27253 cfg2.block_dim,
27254 0,
27255 &mut ps,
27256 )?;
27257 }
27258 return Ok(());
27259 }
27260 let fc = if wg {
27261 self.func_g("fa_decode_combine_rows_w_q8_1")
27262 } else {
27263 self.func("fa_decode_combine_rows_w_q8_1")
27264 };
27265 let __s_b2 = self.gpu.stream();
27266 let mut b2 = __s_b2.launch_builder(&fc);
27267 b2.arg(&*part_o)
27268 .arg(&*part_m)
27269 .arg(&*part_l)
27270 .arg(oq)
27271 .arg(od)
27272 .arg(&hd)
27273 .arg(&nh)
27274 .arg(&nspm)
27275 .arg(&spk)
27276 .arg(&wini);
27277 unsafe {
27278 b2.launch(cfg2)?;
27279 }
27280 return Ok(());
27281 }
27282 let fc = if wg {
27283 self.func_g("fa_decode_combine_rows_w")
27284 } else {
27285 self.func("fa_decode_combine_rows_w")
27286 };
27287 let __s_b2 = self.gpu.stream();
27288 let mut b2 = __s_b2.launch_builder(&fc);
27289 b2.arg(&*part_o)
27290 .arg(&*part_m)
27291 .arg(&*part_l)
27292 .arg(o)
27293 .arg(&hd)
27294 .arg(&nh)
27295 .arg(&nspm)
27296 .arg(&spk)
27297 .arg(&wini);
27298 unsafe {
27299 b2.launch(cfg2)?;
27300 }
27301 Ok(())
27302 }
27303
27304 #[allow(clippy::too_many_arguments)]
27310 pub fn fa_decode_rows_dc(
27311 &self,
27312 q: &CudaSlice<f32>,
27313 k: &cudarc::driver::CudaView<u8>,
27314 v: &cudarc::driver::CudaView<u8>,
27315 o: &mut CudaSlice<f32>,
27316 head_dim: usize,
27317 n_head: usize,
27318 n_head_kv: usize,
27319 base_dev: &CudaSlice<i32>,
27320 t_kv_upper: usize,
27321 t: usize,
27322 scale: f32,
27323 k_tok_bytes: usize,
27324 v_tok_bytes: usize,
27325 base_plus: i32,
27326 g: bool,
27327 ) -> Result<(), Box<dyn std::error::Error>> {
27328 let v4 = head_dim == 256 && fa_v4_at(t_kv_upper);
27329 assert!(
27330 v4 || fa_v3_active(head_dim),
27331 "stream fa rows requires the v3 or v4 lane"
27332 );
27333 assert!(v4 || base_plus == 0, "v3_dc kernel takes no plus arg");
27334 if v4 {
27335 let sp = fa_split_keys(t_kv_upper, n_head_kv);
27336 #[allow(clippy::manual_div_ceil)]
27337 let n_splits_max = (t_kv_upper + sp - 1) / sp;
27339 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
27340 let (nspm, spk) = (n_splits_max as i32, sp as i32);
27341 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
27342 let gqa = (n_head / n_head_kv).max(1) as u32;
27343 let o_len = t * n_head * n_splits_max * head_dim;
27344 let ml_len = t * n_head * n_splits_max;
27345 let mut part_guard = self.fa_part_pool.lock().unwrap();
27346 if part_guard
27347 .as_ref()
27348 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
27349 .unwrap_or(true)
27350 {
27351 let old = part_guard.take();
27362 let (co, cm) = old
27363 .as_ref()
27364 .map(|pp| (pp.0.len(), pp.1.len()))
27365 .unwrap_or((0, 0));
27366 if let Some(old) = old {
27367 self.fa_part_retired.lock().unwrap().push(old);
27368 }
27369 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
27370 eprintln!(
27371 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
27372 co, o_len, cm, ml_len
27373 );
27374 }
27375 *part_guard =
27376 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
27377 }
27378 let pg = part_guard.as_mut().unwrap();
27379 self.gpu
27380 .stream()
27381 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
27382 self.gpu
27383 .stream()
27384 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
27385 self.gpu
27386 .stream()
27387 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
27388 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
27389 let f = if g {
27390 self.func_g("fa_decode_vec_q_rows_v4_dc")
27391 } else {
27392 self.func("fa_decode_vec_q_rows_v4_dc")
27393 };
27394 let sh = (11520 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
27395 use cudarc::driver::sys::CUfunction_attribute_enum as A;
27396 f.set_attribute(
27397 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27398 sh as i32,
27399 )?;
27400 let cfg = LaunchConfig {
27401 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
27402 block_dim: (32, gqa, 1),
27403 shared_mem_bytes: sh,
27404 };
27405 let __s_b = self.gpu.stream();
27406 let mut b = __s_b.launch_builder(&f);
27407 b.arg(q)
27408 .arg(k)
27409 .arg(v)
27410 .arg(&mut *part_o)
27411 .arg(&mut *part_m)
27412 .arg(&mut *part_l)
27413 .arg(&hd)
27414 .arg(&nh)
27415 .arg(&nhkv)
27416 .arg(base_dev)
27417 .arg(&base_plus)
27418 .arg(&scale)
27419 .arg(&nspm)
27420 .arg(&spk)
27421 .arg(&ktb)
27422 .arg(&vtb);
27423 unsafe {
27424 b.launch(cfg)?;
27425 }
27426 let fc = self.func("fa_decode_combine_rows_dc");
27427 let cfg2 = LaunchConfig {
27428 grid_dim: (n_head as u32, t as u32, 1),
27429 block_dim: (head_dim as u32, 1, 1),
27430 shared_mem_bytes: 0,
27431 };
27432 let __s_b2 = self.gpu.stream();
27433 let mut b2 = __s_b2.launch_builder(&fc);
27434 b2.arg(&*part_o)
27435 .arg(&*part_m)
27436 .arg(&*part_l)
27437 .arg(o)
27438 .arg(&hd)
27439 .arg(&nh)
27440 .arg(base_dev)
27441 .arg(&base_plus)
27442 .arg(&nspm)
27443 .arg(&spk);
27444 unsafe {
27445 b2.launch(cfg2)?;
27446 }
27447 return Ok(());
27448 }
27449 let sp = fa_split_keys(t_kv_upper, n_head_kv);
27450 #[allow(clippy::manual_div_ceil)]
27451 let n_splits_max = (t_kv_upper + sp - 1) / sp;
27453 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
27454 let (nspm, spk) = (n_splits_max as i32, sp as i32);
27455 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
27456 let gqa = (n_head / n_head_kv).max(1) as u32;
27457 let o_len = t * n_head * n_splits_max * head_dim;
27458 let ml_len = t * n_head * n_splits_max;
27459 let mut part_guard = self.fa_part_pool.lock().unwrap();
27460 if part_guard
27461 .as_ref()
27462 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
27463 .unwrap_or(true)
27464 {
27465 let old = part_guard.take();
27476 let (co, cm) = old
27477 .as_ref()
27478 .map(|pp| (pp.0.len(), pp.1.len()))
27479 .unwrap_or((0, 0));
27480 if let Some(old) = old {
27481 self.fa_part_retired.lock().unwrap().push(old);
27482 }
27483 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
27484 eprintln!(
27485 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
27486 co, o_len, cm, ml_len
27487 );
27488 }
27489 *part_guard =
27490 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
27491 }
27492 let pg = part_guard.as_mut().unwrap();
27493 self.gpu
27494 .stream()
27495 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
27496 self.gpu
27497 .stream()
27498 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
27499 self.gpu
27500 .stream()
27501 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
27502 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
27503 let f = self.func("fa_decode_vec_q_rows_v3_dc");
27504 let sh = (32 * head_dim * 2) as u32;
27505 use cudarc::driver::sys::CUfunction_attribute_enum as A;
27506 f.set_attribute(
27507 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27508 sh as i32,
27509 )?;
27510 let cfg = LaunchConfig {
27511 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
27512 block_dim: (32, gqa, 1),
27513 shared_mem_bytes: sh,
27514 };
27515 let __s_b = self.gpu.stream();
27516 let mut b = __s_b.launch_builder(&f);
27517 b.arg(q)
27518 .arg(k)
27519 .arg(v)
27520 .arg(&mut *part_o)
27521 .arg(&mut *part_m)
27522 .arg(&mut *part_l)
27523 .arg(&hd)
27524 .arg(&nh)
27525 .arg(&nhkv)
27526 .arg(base_dev)
27527 .arg(&scale)
27528 .arg(&nspm)
27529 .arg(&spk)
27530 .arg(&ktb)
27531 .arg(&vtb);
27532 unsafe {
27533 b.launch(cfg)?;
27534 }
27535 let fc = self.func("fa_decode_combine_rows_dc");
27536 let cfg2 = LaunchConfig {
27537 grid_dim: (n_head as u32, t as u32, 1),
27538 block_dim: (head_dim as u32, 1, 1),
27539 shared_mem_bytes: 0,
27540 };
27541 let plus0 = 0i32;
27542 let __s_b2 = self.gpu.stream();
27543 let mut b2 = __s_b2.launch_builder(&fc);
27544 b2.arg(&*part_o)
27545 .arg(&*part_m)
27546 .arg(&*part_l)
27547 .arg(o)
27548 .arg(&hd)
27549 .arg(&nh)
27550 .arg(base_dev)
27551 .arg(&plus0)
27552 .arg(&nspm)
27553 .arg(&spk);
27554 unsafe {
27555 b2.launch(cfg2)?;
27556 }
27557 Ok(())
27558 }
27559
27560 #[allow(clippy::too_many_arguments)] pub fn fa_decode_dc(
27572 &self,
27573 q: &CudaSlice<f32>,
27574 k: &cudarc::driver::CudaView<u8>,
27575 v: &cudarc::driver::CudaView<u8>,
27576 o: &mut CudaSlice<f32>,
27577 head_dim: usize,
27578 n_head: usize,
27579 n_head_kv: usize,
27580 t_kv_dev: &CudaSlice<i32>,
27581 bucket_max: usize,
27582 scale: f32,
27583 k_tok_bytes: usize,
27584 v_tok_bytes: usize,
27585 g: bool,
27586 ) -> Result<(), Box<dyn std::error::Error>> {
27587 self.fa_decode_dc_q8(
27588 q,
27589 k,
27590 v,
27591 o,
27592 head_dim,
27593 n_head,
27594 n_head_kv,
27595 t_kv_dev,
27596 bucket_max,
27597 scale,
27598 k_tok_bytes,
27599 v_tok_bytes,
27600 g,
27601 None,
27602 )
27603 }
27604
27605 #[allow(clippy::too_many_arguments)]
27608 #[allow(clippy::manual_div_ceil)] pub fn fa_decode_dc_q8(
27610 &self,
27611 q: &CudaSlice<f32>,
27612 k: &cudarc::driver::CudaView<u8>,
27613 v: &cudarc::driver::CudaView<u8>,
27614 o: &mut CudaSlice<f32>,
27615 head_dim: usize,
27616 n_head: usize,
27617 n_head_kv: usize,
27618 t_kv_dev: &CudaSlice<i32>,
27619 bucket_max: usize,
27620 scale: f32,
27621 k_tok_bytes: usize,
27622 v_tok_bytes: usize,
27623 g: bool,
27624 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
27625 ) -> Result<(), Box<dyn std::error::Error>> {
27626 let mut fa_vec =
27634 std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
27635 if g && head_dim == 256 && !fa_v4_at(bucket_max) {
27636 fa_vec = false;
27637 } let sp = fa_split_keys(bucket_max, n_head_kv);
27639 let n_splits = if fa_vec {
27640 ((bucket_max + sp - 1) / sp).max(1)
27641 } else {
27642 ((bucket_max + 255) / 256).max(1)
27643 };
27644 let o_len = n_head * n_splits * head_dim;
27645 let ml_len = n_head * n_splits;
27646 let mut part_guard = self.fa_part_pool.lock().unwrap();
27647 if part_guard
27648 .as_ref()
27649 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
27650 .unwrap_or(true)
27651 {
27652 let old = part_guard.take();
27663 let (co, cm) = old
27664 .as_ref()
27665 .map(|pp| (pp.0.len(), pp.1.len()))
27666 .unwrap_or((0, 0));
27667 if let Some(old) = old {
27668 self.fa_part_retired.lock().unwrap().push(old);
27669 }
27670 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
27671 eprintln!(
27672 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
27673 co, o_len, cm, ml_len
27674 );
27675 }
27676 *part_guard =
27677 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
27678 }
27679 let pg = part_guard.as_mut().unwrap();
27680 self.gpu
27681 .stream()
27682 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
27683 self.gpu
27684 .stream()
27685 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
27686 self.gpu
27687 .stream()
27688 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
27689 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
27690 let (hd, nh, nhkv, nsp) = (
27691 head_dim as i32,
27692 n_head as i32,
27693 n_head_kv as i32,
27694 n_splits as i32,
27695 );
27696 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
27697 let fa_vec = fa_vec && head_dim <= 512 && head_dim.is_multiple_of(32);
27698 let deep = fa_vec
27701 && head_dim == 256
27702 && fa_v4_at(bucket_max)
27703 && !g
27704 && fa_deep_at(bucket_max)
27705 && !matches!(fa_v4_mode(), "noB3" | "stage");
27706 let (f, cfg) = if fa_vec
27707 && head_dim == 512
27708 && bucket_max >= {
27709 static FA512_MIN_DC: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
27710 *FA512_MIN_DC.get_or_init(|| {
27711 std::env::var("MEMRA_FA512_MIN")
27712 .ok()
27713 .and_then(|v| v.parse().ok())
27714 .unwrap_or(512)
27715 })
27716 } {
27717 let gqa = (n_head / n_head_kv).max(1) as u32;
27719 (
27720 self.fa_func("fa_decode_vec_q_dpl16_dc", head_dim),
27721 LaunchConfig {
27722 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27723 block_dim: (32, gqa, 1),
27724 shared_mem_bytes: 0,
27725 },
27726 )
27727 } else if fa_vec && head_dim == 512 {
27728 let q_view = q.as_view();
27731 let mut o_view = o.as_view_mut();
27732 return self.fa_decode_scalar_unified(
27733 &q_view,
27734 k,
27735 v,
27736 &mut o_view,
27737 head_dim,
27738 n_head,
27739 n_head_kv,
27740 0,
27741 Some(t_kv_dev),
27742 scale,
27743 n_splits,
27744 sp,
27745 k_tok_bytes,
27746 v_tok_bytes,
27747 g,
27748 &mut *part_o,
27749 &mut *part_m,
27750 &mut *part_l,
27751 q8_out,
27752 );
27753 } else if fa_vec && head_dim == 256 && fa_v4_at(bucket_max) {
27754 let gqa = (n_head / n_head_kv).max(1) as u32;
27757 let fv = if g {
27758 self.func_g("fa_decode_vec_q_v4_dc")
27759 } else if deep {
27760 self.func("fa_decode_vec_q_v4_deep_dc")
27761 } else {
27762 self.func("fa_decode_vec_q_v4_dc")
27763 };
27764 let shmem =
27765 (if deep { 12160 } else { 11520 } + 32 * head_dim * if g { 1 } else { 2 }) as u32;
27766 use cudarc::driver::sys::CUfunction_attribute_enum as A;
27767 fv.set_attribute(
27768 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
27769 shmem as i32,
27770 )?;
27771 (
27772 fv,
27773 LaunchConfig {
27774 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27775 block_dim: (32, gqa, 1),
27776 shared_mem_bytes: shmem,
27777 },
27778 )
27779 } else if fa_vec && fa_v3_active(head_dim) {
27780 let gqa = (n_head / n_head_kv).max(1) as u32;
27783 let fv = if g {
27784 self.func_g("fa_decode_vec_q_v3_dc")
27785 } else {
27786 self.func("fa_decode_vec_q_v3_dc")
27787 };
27788 let shmem = (32 * head_dim * 2) as u32; (
27790 fv,
27791 LaunchConfig {
27792 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27793 block_dim: (32, gqa, 1),
27794 shared_mem_bytes: shmem,
27795 },
27796 )
27797 } else if fa_vec && fa_v2_on() {
27798 let gqa = (n_head / n_head_kv).max(1) as u32;
27802 let fv = if g {
27803 self.func_g("fa_decode_vec_q_v2_dc")
27804 } else {
27805 self.func("fa_decode_vec_q_v2_dc")
27806 };
27807 let shmem = (2 * 32 * head_dim * 2) as u32; (
27809 fv,
27810 LaunchConfig {
27811 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27812 block_dim: (32, gqa, 1),
27813 shared_mem_bytes: shmem,
27814 },
27815 )
27816 } else if fa_vec {
27817 let gqa = (n_head / n_head_kv).max(1) as u32;
27818 let fv = if g {
27820 self.func_g("fa_decode_vec_q_dc")
27821 } else {
27822 self.func("fa_decode_vec_q_dc")
27823 };
27824 (
27825 fv,
27826 LaunchConfig {
27827 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
27828 block_dim: (32, gqa, 1),
27829 shared_mem_bytes: 0,
27830 },
27831 )
27832 } else {
27833 let q_view = q.as_view();
27834 let mut o_view = o.as_view_mut();
27835 return self.fa_decode_scalar_unified(
27836 &q_view,
27837 k,
27838 v,
27839 &mut o_view,
27840 head_dim,
27841 n_head,
27842 n_head_kv,
27843 0,
27844 Some(t_kv_dev),
27845 scale,
27846 n_splits,
27847 if fa_vec { sp } else { 256 },
27848 k_tok_bytes,
27849 v_tok_bytes,
27850 g,
27851 &mut *part_o,
27852 &mut *part_m,
27853 &mut *part_l,
27854 q8_out,
27855 );
27856 };
27857 let ski = sp as i32; let __s_b = self.gpu.stream();
27859 let mut b = __s_b.launch_builder(&f);
27860 b.arg(q)
27861 .arg(k)
27862 .arg(v)
27863 .arg(&mut *part_o)
27864 .arg(&mut *part_m)
27865 .arg(&mut *part_l)
27866 .arg(&hd)
27867 .arg(&nh)
27868 .arg(&nhkv)
27869 .arg(t_kv_dev)
27870 .arg(&scale)
27871 .arg(&nsp)
27872 .arg(&ski)
27873 .arg(&ktb)
27874 .arg(&vtb);
27875 unsafe {
27876 b.launch(cfg)?;
27877 }
27878 let cfg2 = LaunchConfig {
27879 grid_dim: (n_head as u32, 1, 1),
27880 block_dim: (head_dim as u32, 1, 1),
27881 shared_mem_bytes: 0,
27882 };
27883 if let Some((oq, od)) = q8_out {
27884 let fc = if g {
27885 self.func_g("fa_decode_combine_q8_1")
27886 } else {
27887 self.fa_func("fa_decode_combine_q8_1", head_dim)
27888 };
27889 let __s_b2 = self.gpu.stream();
27890 let mut b2 = __s_b2.launch_builder(&fc);
27891 b2.arg(&*part_o)
27892 .arg(&*part_m)
27893 .arg(&*part_l)
27894 .arg(oq)
27895 .arg(od)
27896 .arg(&hd)
27897 .arg(&nh)
27898 .arg(&nsp);
27899 unsafe {
27900 b2.launch(cfg2)?;
27901 }
27902 return Ok(());
27903 }
27904 let fc = if g {
27905 self.func_g("fa_decode_combine_f32")
27906 } else {
27907 self.fa_func("fa_decode_combine_f32", head_dim)
27908 };
27909 let __s_b2 = self.gpu.stream();
27910 let mut b2 = __s_b2.launch_builder(&fc);
27911 b2.arg(&*part_o)
27912 .arg(&*part_m)
27913 .arg(&*part_l)
27914 .arg(o)
27915 .arg(&hd)
27916 .arg(&nh)
27917 .arg(&nsp);
27918 unsafe {
27919 b2.launch(cfg2)?;
27920 }
27921 Ok(())
27922 }
27923
27924 #[allow(clippy::too_many_arguments)]
27928 pub fn append_kv_quantized_dcw(
27929 &self,
27930 k_row: &CudaSlice<f32>,
27931 v_row: &CudaSlice<f32>,
27932 kc: &mut CudaSlice<u8>,
27933 vc: &mut CudaSlice<u8>,
27934 len_dev: &CudaSlice<i32>,
27935 base_dev: Option<&CudaSlice<i32>>,
27936 kv_dim_k: usize,
27937 kv_dim_v: usize,
27938 k_tok_bytes: usize,
27939 v_tok_bytes: usize,
27940 ) -> Result<(), Box<dyn std::error::Error>> {
27941 let f = self.func("append_quantize_kv_q8_0_q5_1_dcw");
27942 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
27943 let cfg = LaunchConfig {
27944 grid_dim: (nblk, 1, 1),
27945 block_dim: (32, 1, 1),
27946 shared_mem_bytes: 0,
27947 };
27948 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
27949 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
27950 let null: u64 = 0;
27951 let __s_b = self.gpu.stream();
27952 let mut b = __s_b.launch_builder(&f);
27953 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(len_dev);
27954 match base_dev {
27955 Some(base) => {
27956 b.arg(base);
27957 }
27958 None => {
27959 b.arg(&null);
27960 }
27961 }
27962 b.arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
27963 unsafe {
27964 b.launch(cfg)?;
27965 }
27966 Ok(())
27967 }
27968
27969 pub fn inc_i32(&self, counter: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
27971 let f = self.func("inc_i32");
27972 let cfg = LaunchConfig {
27973 grid_dim: (1, 1, 1),
27974 block_dim: (1, 1, 1),
27975 shared_mem_bytes: 0,
27976 };
27977 let __s_b = self.gpu.stream();
27978 let mut b = __s_b.launch_builder(&f);
27979 b.arg(counter);
27980 unsafe {
27981 b.launch(cfg)?;
27982 }
27983 Ok(())
27984 }
27985
27986 #[allow(clippy::too_many_arguments)]
27995 #[allow(clippy::type_complexity)] fn fa_part_alloc(
28020 &self,
28021 o_len: usize,
28022 ml_len: usize,
28023 co: usize,
28024 cm: usize,
28025 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
28026 static GROWS: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
28027 let n = GROWS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
28028 if n < 64 {
28029 eprintln!(
28030 "[fa-pool] grow #{n} dev={} o_len {co} -> {o_len} ml_len {cm} -> {ml_len} (retired kept, zero={})",
28031 self.ctx().ordinal(),
28032 fa_part_zero_on()
28033 );
28034 }
28035 let mut po = self.alloc_uninit::<f32>(o_len)?;
28036 let mut pm = self.alloc_uninit::<f32>(ml_len)?;
28037 let mut pl = self.alloc_uninit::<f32>(ml_len)?;
28038 if fa_part_zero_on() {
28039 self.gpu.stream().memset_zeros(&mut po)?;
28040 self.gpu.stream().memset_zeros(&mut pm)?;
28041 self.gpu.stream().memset_zeros(&mut pl)?;
28042 }
28043 Ok((po, pm, pl))
28044 }
28045
28046 fn fa_part_pool_grow(
28047 &self,
28048 part_guard: &mut Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>,
28049 o_len: usize,
28050 ml_len: usize,
28051 ) -> Result<(), Box<dyn std::error::Error>> {
28052 if part_guard
28053 .as_ref()
28054 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
28055 .unwrap_or(true)
28056 {
28057 let old = part_guard.take();
28058 let (co, cm) = old
28059 .as_ref()
28060 .map(|pp| (pp.0.len(), pp.1.len()))
28061 .unwrap_or((0, 0));
28062 if let Some(old) = old {
28063 self.fa_part_retired.lock().unwrap().push(old);
28064 }
28065 *part_guard =
28078 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
28079 }
28080 Ok(())
28081 }
28082
28083 pub fn fa_dcw_pool_ensure(
28086 &self,
28087 head_dim: usize,
28088 n_head: usize,
28089 n_head_kv: usize,
28090 bucket_max: usize,
28091 ) -> Result<(), Box<dyn std::error::Error>> {
28092 let sp = fa_split_keys(bucket_max, n_head_kv);
28093 #[allow(clippy::manual_div_ceil)]
28094 let n_splits = ((bucket_max + sp - 1) / sp).max(1);
28096 let o_len = n_head * n_splits * head_dim;
28097 let ml_len = n_head * n_splits;
28098 let mut part_guard = self.fa_part_pool.lock().unwrap();
28099 self.fa_part_pool_grow(&mut part_guard, o_len, ml_len)
28100 }
28101
28102 #[allow(clippy::too_many_arguments)]
28110 pub fn fa_decode_dcw2(
28111 &self,
28112 q2: &CudaSlice<f32>,
28113 k_ring: &cudarc::driver::CudaView<u8>,
28114 v_ring: &cudarc::driver::CudaView<u8>,
28115 o2: &mut CudaSlice<f32>,
28116 head_dim: usize,
28117 n_head: usize,
28118 n_head_kv: usize,
28119 len_dev: &CudaSlice<i32>,
28120 base_dev: Option<&CudaSlice<i32>>,
28121 window: usize,
28122 bucket_max: usize,
28123 scale: f32,
28124 k_tok_bytes: usize,
28125 v_tok_bytes: usize,
28126 gate2: &CudaSlice<f32>,
28127 ) -> Result<(), Box<dyn std::error::Error>> {
28128 let fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
28129 if !fa_vec || head_dim > 256 || !head_dim.is_multiple_of(32) || !fa_v3_on() {
28130 return Err("fa_decode_dcw2 supports the default v3-vec class only".into());
28131 }
28132 let sp = fa_split_keys(bucket_max, n_head_kv);
28133 #[allow(clippy::manual_div_ceil)]
28134 let n_splits = ((bucket_max + sp - 1) / sp).max(1);
28136 let o_len = 2 * n_head * n_splits * head_dim;
28138 let ml_len = 2 * n_head * n_splits;
28139 let mut part_guard = self.fa_part_pool.lock().unwrap();
28140 Self::fa_part_pool_grow(self, &mut part_guard, o_len, ml_len)?;
28141 let pg = part_guard.as_mut().unwrap();
28142 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
28143 let (hd, nh, nhkv, nsp) = (
28144 head_dim as i32,
28145 n_head as i32,
28146 n_head_kv as i32,
28147 n_splits as i32,
28148 );
28149 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
28150 let (ski, win) = (sp as i32, window as i32);
28151 let gqa = (n_head / n_head_kv).max(1) as u32;
28152 let smem = (32 * head_dim * 2) as u32;
28153 let f = self.func("fa_decode_vec_q_v3_dcw2");
28154 let cfg = LaunchConfig {
28155 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
28156 block_dim: (32, gqa, 1),
28157 shared_mem_bytes: smem,
28158 };
28159 let null: u64 = 0;
28160 {
28161 let __s_b = self.gpu.stream();
28162 let mut b = __s_b.launch_builder(&f);
28163 b.arg(q2)
28164 .arg(k_ring)
28165 .arg(v_ring)
28166 .arg(&mut *part_o)
28167 .arg(&mut *part_m)
28168 .arg(&mut *part_l)
28169 .arg(&hd)
28170 .arg(&nh)
28171 .arg(&nhkv)
28172 .arg(len_dev);
28173 match base_dev {
28174 Some(base) => {
28175 b.arg(base);
28176 }
28177 None => {
28178 b.arg(&null);
28179 }
28180 }
28181 b.arg(&win)
28182 .arg(&scale)
28183 .arg(&nsp)
28184 .arg(&ski)
28185 .arg(&ktb)
28186 .arg(&vtb);
28187 unsafe {
28188 b.launch(cfg)?;
28189 }
28190 }
28191 let fc = {
28195 static CS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28196 if *CS.get_or_init(|| std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1")) {
28197 self.func("fa_decode_combine_gate_f32_s")
28198 } else {
28199 self.func("fa_decode_combine_gate_f32")
28200 }
28201 };
28202 let combine_shared = std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1");
28203 let nh2 = (2 * n_head) as i32;
28204 let cfg2 = LaunchConfig {
28205 grid_dim: ((2 * n_head) as u32, 1, 1),
28206 block_dim: (head_dim as u32, 1, 1),
28207 shared_mem_bytes: if combine_shared {
28208 (2 * n_splits * 4) as u32
28209 } else {
28210 0
28211 },
28212 };
28213 let __s_b2 = self.gpu.stream();
28214 let mut b2 = __s_b2.launch_builder(&fc);
28215 b2.arg(&*part_o)
28216 .arg(&*part_m)
28217 .arg(&*part_l)
28218 .arg(gate2)
28219 .arg(o2)
28220 .arg(&hd)
28221 .arg(&nh2)
28222 .arg(&nsp);
28223 unsafe {
28224 b2.launch(cfg2)?;
28225 }
28226 Ok(())
28227 }
28228
28229 #[allow(clippy::too_many_arguments)]
28238 pub fn fa_decode_dcw_rows(
28239 &self,
28240 q_rows: &CudaSlice<f32>,
28241 tab: &CudaSlice<u64>,
28242 o_rows: &mut CudaSlice<f32>,
28243 t: usize,
28244 head_dim: usize,
28245 n_head: usize,
28246 n_head_kv: usize,
28247 window: usize,
28248 max_ns: usize,
28249 scale: f32,
28250 k_tok_bytes: usize,
28251 v_tok_bytes: usize,
28252 gate_rows: &CudaSlice<f32>,
28253 ) -> Result<(), Box<dyn std::error::Error>> {
28254 if std::env::var("MEMRA_NO_FA_VEC").is_ok()
28255 || head_dim > 256
28256 || !head_dim.is_multiple_of(32)
28257 || !fa_v3_on()
28258 {
28259 return Err("fa_decode_dcw_rows supports the default v3-vec class only".into());
28260 }
28261 if fa_sm_count() < 128
28262 || std::env::var("MEMRA_FA_SPLIT").is_ok()
28263 || std::env::var("MEMRA_FA_SP_SHORT").is_ok()
28264 || std::env::var("MEMRA_FA_SP16").is_ok()
28265 {
28266 return Err(
28267 "fa_decode_dcw_rows embeds the big-rig split ladder; env split overrides \
28268 (or a <128-SM rig) keep the per-row path"
28269 .into(),
28270 );
28271 }
28272 if t == 0 || t > 32 || max_ns == 0 || tab.len() < t * 6 {
28273 return Err("fa_decode_dcw_rows geometry".into());
28274 }
28275 let o_len = t * n_head * max_ns * head_dim;
28276 let ml_len = t * n_head * max_ns;
28277 let mut part_guard = self.fa_part_pool.lock().unwrap();
28278 Self::fa_part_pool_grow(self, &mut part_guard, o_len, ml_len)?;
28279 let pg = part_guard.as_mut().unwrap();
28280 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
28281 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
28282 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
28283 let (win, mns) = (window as i32, max_ns as i32);
28284 let gqa = (n_head / n_head_kv).max(1) as u32;
28285 let smem = (32 * head_dim * 2) as u32;
28286 let f = self.func("fa_decode_vec_q_v3_dcw_rows");
28287 let cfg = LaunchConfig {
28288 grid_dim: (n_head_kv as u32, max_ns as u32, t as u32),
28289 block_dim: (32, gqa, 1),
28290 shared_mem_bytes: smem,
28291 };
28292 {
28293 let __s_b = self.gpu.stream();
28294 let mut b = __s_b.launch_builder(&f);
28295 b.arg(q_rows)
28296 .arg(tab)
28297 .arg(&mut *part_o)
28298 .arg(&mut *part_m)
28299 .arg(&mut *part_l)
28300 .arg(&hd)
28301 .arg(&nh)
28302 .arg(&nhkv)
28303 .arg(&win)
28304 .arg(&scale)
28305 .arg(&mns)
28306 .arg(&ktb)
28307 .arg(&vtb);
28308 unsafe {
28309 b.launch(cfg)?;
28310 }
28311 }
28312 let fc = {
28316 static CS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28317 if *CS.get_or_init(|| std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1")) {
28318 self.func("fa_decode_combine_gate_f32_s")
28319 } else {
28320 self.func("fa_decode_combine_gate_f32")
28321 }
28322 };
28323 let combine_shared = std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1");
28324 let nht = (t * n_head) as i32;
28325 let cfg2 = LaunchConfig {
28326 grid_dim: ((t * n_head) as u32, 1, 1),
28327 block_dim: (head_dim as u32, 1, 1),
28328 shared_mem_bytes: if combine_shared {
28329 (2 * max_ns * 4) as u32
28330 } else {
28331 0
28332 },
28333 };
28334 let __s_b2 = self.gpu.stream();
28335 let mut b2 = __s_b2.launch_builder(&fc);
28336 b2.arg(&*part_o)
28337 .arg(&*part_m)
28338 .arg(&*part_l)
28339 .arg(gate_rows)
28340 .arg(o_rows)
28341 .arg(&hd)
28342 .arg(&nht)
28343 .arg(&mns);
28344 unsafe {
28345 b2.launch(cfg2)?;
28346 }
28347 Ok(())
28348 }
28349
28350 #[allow(clippy::too_many_arguments)] pub fn fa_decode_dcw(
28352 &self,
28353 q: &CudaSlice<f32>,
28354 k_ring: &cudarc::driver::CudaView<u8>,
28355 v_ring: &cudarc::driver::CudaView<u8>,
28356 o: &mut CudaSlice<f32>,
28357 head_dim: usize,
28358 n_head: usize,
28359 n_head_kv: usize,
28360 len_dev: &CudaSlice<i32>,
28361 base_dev: Option<&CudaSlice<i32>>,
28362 window: usize,
28363 bucket_max: usize,
28364 scale: f32,
28365 k_tok_bytes: usize,
28366 v_tok_bytes: usize,
28367 fused_gate: Option<&CudaSlice<f32>>,
28371 ) -> Result<(), Box<dyn std::error::Error>> {
28372 let fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
28373 if !fa_vec || head_dim > 256 || !head_dim.is_multiple_of(32) || !fa_v3_on() {
28374 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"
28375 .into());
28376 }
28377 let sp = fa_split_keys(bucket_max, n_head_kv);
28378 #[allow(clippy::manual_div_ceil)]
28379 let n_splits = ((bucket_max + sp - 1) / sp).max(1);
28381 let o_len = n_head * n_splits * head_dim;
28382 let ml_len = n_head * n_splits;
28383 let mut part_guard = self.fa_part_pool.lock().unwrap();
28384 Self::fa_part_pool_grow(self, &mut part_guard, o_len, ml_len)?;
28385 let pg = part_guard.as_mut().unwrap();
28386 static MEMSET_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28391 let memset_on = *MEMSET_ON
28396 .get_or_init(|| std::env::var("MEMRA_FA_DCW_MEMSET").as_deref() != Ok("0"))
28397 || crate::tp::token_graph_building();
28398 if memset_on {
28399 self.gpu
28400 .stream()
28401 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
28402 self.gpu
28403 .stream()
28404 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
28405 self.gpu
28406 .stream()
28407 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
28408 }
28409 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
28410 let (hd, nh, nhkv, nsp) = (
28411 head_dim as i32,
28412 n_head as i32,
28413 n_head_kv as i32,
28414 n_splits as i32,
28415 );
28416 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
28417 let (ski, win) = (sp as i32, window as i32);
28418 let gqa = (n_head / n_head_kv).max(1) as u32;
28419 let smem = (32 * head_dim * 2) as u32; static U8: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28423 static HOIST: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
28424 let hoist = *HOIST.get_or_init(|| match std::env::var("MEMRA_FA_HOIST").as_deref() {
28425 Ok("2") => 2,
28426 Ok("1") => 1,
28427 _ => 0,
28428 });
28429 static FPROF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28434 let fprof = *FPROF.get_or_init(|| std::env::var("MEMRA_FA_PROF").as_deref() == Ok("1"));
28435 static PROF_BUF: std::sync::Mutex<Option<(usize, CudaSlice<u64>)>> =
28436 std::sync::Mutex::new(None);
28437 static HS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28441 let hs2 = *HS.get_or_init(|| std::env::var("MEMRA_FA_HSPLIT").as_deref() == Ok("2"))
28442 && (n_head / n_head_kv).is_multiple_of(2)
28443 && (n_head / n_head_kv) >= 2;
28444 let f = if fprof {
28445 self.func("fa_decode_vec_q_v3_dcw_prof")
28446 } else if hs2 {
28447 self.func("fa_decode_vec_q_v3_dcw_hs2")
28448 } else if hoist == 2 {
28449 self.func("fa_decode_vec_q_v3_dcw_hc")
28451 } else if hoist == 1 {
28452 self.func("fa_decode_vec_q_v3_dcw_h")
28454 } else if *U8.get_or_init(|| std::env::var("MEMRA_FA_UNROLL").as_deref() == Ok("8")) {
28455 self.func("fa_decode_vec_q_v3_dcw_u8")
28456 } else {
28457 self.func("fa_decode_vec_q_v3_dcw")
28458 };
28459 let cfg = LaunchConfig {
28460 grid_dim: if hs2 {
28461 ((2 * n_head_kv) as u32, n_splits as u32, 1)
28462 } else {
28463 (n_head_kv as u32, n_splits as u32, 1)
28464 },
28465 block_dim: if hs2 { (32, gqa / 2, 1) } else { (32, gqa, 1) },
28466 shared_mem_bytes: smem,
28467 };
28468 let null: u64 = 0;
28469 let __s_b = self.gpu.stream();
28470 let mut b = __s_b.launch_builder(&f);
28471 b.arg(q)
28472 .arg(k_ring)
28473 .arg(v_ring)
28474 .arg(&mut *part_o)
28475 .arg(&mut *part_m)
28476 .arg(&mut *part_l)
28477 .arg(&hd)
28478 .arg(&nh)
28479 .arg(&nhkv)
28480 .arg(len_dev);
28481 match base_dev {
28482 Some(base) => {
28483 b.arg(base);
28484 }
28485 None => {
28486 b.arg(&null);
28487 }
28488 }
28489 b.arg(&win)
28490 .arg(&scale)
28491 .arg(&nsp)
28492 .arg(&ski)
28493 .arg(&ktb)
28494 .arg(&vtb);
28495 if fprof {
28496 let mut guard = PROF_BUF.lock().map_err(|_| "fa prof buffer lock")?;
28497 if guard
28498 .as_ref()
28499 .is_none_or(|(d, _)| *d != self.ctx().ordinal())
28500 {
28501 *guard = Some((self.ctx().ordinal(), self.htod_u64(&[0u64; 8])?));
28502 }
28503 let (_, buf) = guard.as_mut().expect("armed above");
28504 b.arg(&*buf);
28505 unsafe {
28506 b.launch(cfg)?;
28507 }
28508 static CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
28509 let n = CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1;
28510 if n.is_multiple_of(430) {
28511 self.stream().synchronize()?;
28512 let h = self.dtoh_u64(buf)?;
28513 let phases = ["setup", "stageV", "b1_klo", "b2_soft", "sync", "b3_vacc"];
28514 let tot: u64 = h[..6].iter().sum();
28515 let mut line = format!("[fa-prof] calls={n} keys={} cycles={tot}", h[6]);
28516 for (i, name) in phases.iter().enumerate() {
28517 let pct = if tot > 0 {
28518 h[i] as f64 / tot as f64 * 100.0
28519 } else {
28520 0.0
28521 };
28522 line.push_str(&format!(" {name}={pct:.1}%"));
28523 }
28524 if h[6] > 0 {
28525 line.push_str(&format!(" cyc/key={:.0}", tot as f64 / h[6] as f64));
28526 }
28527 eprintln!("{line}");
28528 }
28529 } else {
28530 unsafe {
28531 b.launch(cfg)?;
28532 }
28533 }
28534 let mut combine_shared = false;
28535 let fc = if fused_gate.is_some() {
28536 static CS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28539 if *CS.get_or_init(|| std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1")) {
28540 combine_shared = true;
28541 self.func("fa_decode_combine_gate_f32_s")
28542 } else {
28543 self.func("fa_decode_combine_gate_f32")
28544 }
28545 } else {
28546 self.fa_func("fa_decode_combine_f32", head_dim)
28547 };
28548 let cfg2 = LaunchConfig {
28549 grid_dim: (n_head as u32, 1, 1),
28550 block_dim: (head_dim as u32, 1, 1),
28551 shared_mem_bytes: if combine_shared {
28552 (2 * n_splits * 4) as u32
28553 } else {
28554 0
28555 },
28556 };
28557 let __s_b2 = self.gpu.stream();
28558 let mut b2 = __s_b2.launch_builder(&fc);
28559 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l);
28560 if let Some(gate_row) = fused_gate {
28561 b2.arg(gate_row);
28562 }
28563 b2.arg(o).arg(&hd).arg(&nh).arg(&nsp);
28564 unsafe {
28565 b2.launch(cfg2)?;
28566 }
28567 Ok(())
28568 }
28569
28570 #[allow(clippy::manual_div_ceil)] pub fn fa_geom_eager(
28577 &self,
28578 t_kv: usize,
28579 head_dim: usize,
28580 n_head_kv: usize,
28581 g: bool,
28582 ) -> (bool, usize) {
28583 let fa_ok = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
28587 let vec512 = fa_ok && head_dim == 512 && t_kv >= fa512_min_tkv();
28593 let mut fa_vec = vec512 || (fa_ok && head_dim <= 256 && head_dim.is_multiple_of(32));
28594 if g && head_dim == 256 && !fa_v4_at(t_kv) {
28600 fa_vec = false;
28601 }
28602 let sp = fa_split_keys(t_kv, n_head_kv);
28603 let n_splits = if fa_vec {
28604 ((t_kv + sp - 1) / sp).max(1)
28605 } else {
28606 ((t_kv + 255) / 256).max(1)
28607 };
28608 (fa_vec, n_splits)
28609 }
28610
28611 pub fn fa_bucket_key(
28617 &self,
28618 t_kv: usize,
28619 head_dim: usize,
28620 n_head_kv: usize,
28621 g: bool,
28622 ) -> (bool, usize) {
28623 self.fa_geom_eager(t_kv, head_dim, n_head_kv, g)
28624 }
28625
28626 #[allow(clippy::type_complexity)] pub fn capture_graph_retained<F>(
28639 &self,
28640 step: F,
28641 ) -> Result<
28642 (
28643 cudarc::driver::CudaGraph,
28644 Vec<Box<dyn std::any::Any + Send>>,
28645 ),
28646 Box<dyn std::error::Error>,
28647 >
28648 where
28649 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
28650 {
28651 use cudarc::driver::sys::CUgraphInstantiate_flags;
28652 self.capture_graph_retained_flags(
28653 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
28654 step,
28655 )
28656 }
28657
28658 #[allow(clippy::type_complexity)] pub fn capture_graph_retained_flags<F>(
28664 &self,
28665 flags: cudarc::driver::sys::CUgraphInstantiate_flags,
28666 mut step: F,
28667 ) -> Result<
28668 (
28669 cudarc::driver::CudaGraph,
28670 Vec<Box<dyn std::any::Any + Send>>,
28671 ),
28672 Box<dyn std::error::Error>,
28673 >
28674 where
28675 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
28676 {
28677 use cudarc::driver::sys::CUstreamCaptureMode;
28678 self.capture_keep.lock().unwrap().clear();
28686 let was_tracking = self.gpu.ctx.is_event_tracking();
28687 if was_tracking {
28688 unsafe {
28689 self.gpu.ctx.disable_event_tracking();
28690 }
28691 }
28692 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
28693 self.capture_keep_on
28694 .store(true, std::sync::atomic::Ordering::Relaxed);
28695 let w = (|| {
28696 step(self)?;
28697 step(self)
28698 })();
28699 self.capture_keep_on
28700 .store(false, std::sync::atomic::Ordering::Relaxed);
28701 w?;
28702 self.gpu.stream().synchronize()?;
28703 self.gpu
28704 .stream()
28705 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
28706 let r = step(self);
28707 let g = self.gpu.stream().end_capture(flags);
28708 r?;
28709 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
28710 graph.upload()?;
28711 Ok(graph)
28712 };
28713 let result = run();
28714 self.capture_keep_on
28715 .store(false, std::sync::atomic::Ordering::Relaxed);
28716 if was_tracking {
28717 unsafe {
28718 self.gpu.ctx.enable_event_tracking();
28719 }
28720 }
28721 let keeper = std::mem::take(&mut *self.capture_keep.lock().unwrap());
28722 Ok((result?, keeper))
28723 }
28724
28725 #[allow(clippy::type_complexity)] pub fn capture_graph_retained_nowarm<F>(
28732 &self,
28733 mut step: F,
28734 ) -> Result<
28735 (
28736 cudarc::driver::CudaGraph,
28737 Vec<Box<dyn std::any::Any + Send>>,
28738 ),
28739 Box<dyn std::error::Error>,
28740 >
28741 where
28742 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
28743 {
28744 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
28745 let was_tracking = self.gpu.ctx.is_event_tracking();
28746 if was_tracking {
28747 unsafe {
28748 self.gpu.ctx.disable_event_tracking();
28749 }
28750 }
28751 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
28752 self.gpu.stream().synchronize()?;
28753 self.gpu
28754 .stream()
28755 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
28756 let r = step(self);
28757 let g = self.gpu.stream().end_capture(
28758 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
28759 );
28760 r?;
28761 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
28762 graph.upload()?;
28763 Ok(graph)
28764 };
28765 let result = run();
28766 if was_tracking {
28767 unsafe {
28768 self.gpu.ctx.enable_event_tracking();
28769 }
28770 }
28771 Ok((result?, Vec::new()))
28772 }
28773
28774 pub fn capture_graph<F>(
28775 &self,
28776 mut step: F,
28777 ) -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>>
28778 where
28779 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
28780 {
28781 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
28782 let was_tracking = self.gpu.ctx.is_event_tracking();
28790 if was_tracking {
28791 unsafe {
28792 self.gpu.ctx.disable_event_tracking();
28793 }
28794 }
28795 let iflag = {
28802 static F: std::sync::OnceLock<CUgraphInstantiate_flags> = std::sync::OnceLock::new();
28803 *F.get_or_init(|| match std::env::var("MEMRA_GRAPH_IFLAG").as_deref() {
28804 Ok("upload") => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_UPLOAD,
28807 Ok("priority") => {
28808 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY
28809 }
28810 _ => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
28811 })
28812 };
28813 let ct = {
28820 static T: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
28821 *T.get_or_init(|| std::env::var("MEMRA_GRAPH_CAPTIME").as_deref() == Ok("1"))
28822 };
28823 let warmups = {
28846 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
28847 *W.get_or_init(|| {
28848 std::env::var("MEMRA_GRAPH_WARMUPS")
28849 .ok()
28850 .and_then(|v| v.parse().ok())
28851 .filter(|n| *n >= 1)
28852 .unwrap_or(1)
28853 })
28854 };
28855 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
28856 let t_w = std::time::Instant::now();
28857 for _ in 0..warmups {
28859 step(self)?;
28860 }
28861 self.gpu.stream().synchronize()?;
28862 let ms_warm = t_w.elapsed().as_secs_f64() * 1e3;
28863 let t_c = std::time::Instant::now();
28865 self.gpu
28866 .stream()
28867 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
28868 let r = step(self);
28871 let ms_body = t_c.elapsed().as_secs_f64() * 1e3;
28872 let t_i = std::time::Instant::now();
28873 let g = self.gpu.stream().end_capture(iflag);
28874 let ms_inst = t_i.elapsed().as_secs_f64() * 1e3;
28875 r?;
28876 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
28877 let t_u = std::time::Instant::now();
28878 graph.upload()?;
28879 if ct {
28880 println!(
28881 "[graph-captime] warmup2x {ms_warm:.2} ms capture-body {ms_body:.2} ms \
28882 instantiate {ms_inst:.2} ms upload {:.2} ms",
28883 t_u.elapsed().as_secs_f64() * 1e3
28884 );
28885 }
28886 Ok(graph)
28887 };
28888 let result = run();
28889 if was_tracking {
28890 unsafe {
28891 self.gpu.ctx.enable_event_tracking();
28892 }
28893 }
28894 result
28895 }
28896
28897 #[allow(clippy::too_many_arguments)] pub fn gdn_scan_s128_view(
28900 &self,
28901 q: &CudaSlice<f32>,
28902 k: &CudaSlice<f32>,
28903 v: &CudaSlice<f32>,
28904 g: &CudaSlice<f32>,
28905 beta: &CudaSlice<f32>,
28906 state_in: &cudarc::driver::CudaView<f32>,
28907 state_out: &mut cudarc::driver::CudaViewMut<f32>,
28908 o: &mut CudaSlice<f32>,
28909 n_head: usize,
28910 t: usize,
28911 scale: f32,
28912 ) -> Result<(), Box<dyn std::error::Error>> {
28913 let f = self.func("gdn_scan_s128");
28914 const S_V: u32 = 128;
28915 const WARP: u32 = 32;
28916 const COLS: u32 = 4;
28917 let cfg = LaunchConfig {
28918 grid_dim: (n_head as u32, 1, S_V / COLS),
28919 block_dim: (WARP, COLS, 1),
28920 shared_mem_bytes: 0,
28921 };
28922 let (h, ti) = (n_head as i32, t as i32);
28923 let __s_b = self.gpu.stream();
28924 let mut b = __s_b.launch_builder(&f);
28925 b.arg(q)
28926 .arg(k)
28927 .arg(v)
28928 .arg(g)
28929 .arg(beta)
28930 .arg(state_in)
28931 .arg(state_out)
28932 .arg(o)
28933 .arg(&h)
28934 .arg(&ti)
28935 .arg(&scale);
28936 unsafe {
28937 b.launch(cfg)?;
28938 }
28939 Ok(())
28940 }
28941
28942 #[allow(clippy::too_many_arguments)]
28944 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_view(
28947 &self,
28948 x: &cudarc::driver::CudaView<f32>,
28949 w: &CudaSlice<f32>,
28950 y: &mut CudaSlice<f32>,
28951 conv_dim: usize,
28952 t: usize,
28953 d_conv: usize,
28954 silu: bool,
28955 ) -> Result<(), Box<dyn std::error::Error>> {
28956 let f = self.func("ssm_conv1d_silu_f32");
28957 let cfg = LaunchConfig {
28959 grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
28960 block_dim: (256, 1, 1),
28961 shared_mem_bytes: 0,
28962 };
28963 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
28964 let __s_b = self.gpu.stream();
28965 let mut b = __s_b.launch_builder(&f);
28966 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
28967 unsafe {
28968 b.launch(cfg)?;
28969 }
28970 Ok(())
28971 }
28972
28973 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_tm(
28981 &self,
28982 qkv_tm: &CudaSlice<f32>,
28983 w: &CudaSlice<f32>,
28984 y: &mut CudaSlice<f32>,
28985 conv_dim: usize,
28986 t: usize,
28987 d_conv: usize,
28988 ) -> Result<(), Box<dyn std::error::Error>> {
28989 let f = self.func("ssm_conv1d_tm_f32");
28990 let cfg = LaunchConfig {
28991 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
28992 block_dim: (256, 1, 1),
28993 shared_mem_bytes: 0,
28994 };
28995 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
28996 let __s_b = self.gpu.stream();
28997 let mut b = __s_b.launch_builder(&f);
28998 b.arg(qkv_tm).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
28999 unsafe {
29000 b.launch(cfg)?;
29001 }
29002 Ok(())
29003 }
29004
29005 #[allow(clippy::too_many_arguments)] pub fn ssm_conv1d_tm_state(
29014 &self,
29015 qkv_tm: &CudaSlice<f32>,
29016 conv_state: &mut CudaSlice<f32>,
29017 w: &CudaSlice<f32>,
29018 y: &mut CudaSlice<f32>,
29019 conv_dim: usize,
29020 t: usize,
29021 d_conv: usize,
29022 ) -> Result<(), Box<dyn std::error::Error>> {
29023 self.ssm_conv1d_tm_state_pad(qkv_tm, conv_state, w, y, conv_dim, t, d_conv, None)
29024 }
29025
29026 #[allow(clippy::too_many_arguments)]
29029 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_tm_state_pad(
29031 &self,
29032 qkv_tm: &CudaSlice<f32>,
29033 conv_state: &mut CudaSlice<f32>,
29034 w: &CudaSlice<f32>,
29035 y: &mut CudaSlice<f32>,
29036 conv_dim: usize,
29037 t: usize,
29038 d_conv: usize,
29039 pad_len: Option<&CudaSlice<i32>>,
29040 ) -> Result<(), Box<dyn std::error::Error>> {
29041 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
29042 let ring_old = if t < d_conv - 1 {
29046 Some(self.clone_dtod(conv_state)?)
29047 } else {
29048 None
29049 };
29050 {
29051 let f = self.func("ssm_conv1d_tm_state_f32");
29052 let cfg = LaunchConfig {
29053 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
29054 block_dim: (256, 1, 1),
29055 shared_mem_bytes: 0,
29056 };
29057 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29058 let __s_b = self.gpu.stream();
29059 let mut b = __s_b.launch_builder(&f);
29060 b.arg(qkv_tm)
29061 .arg(&*conv_state)
29062 .arg(w)
29063 .arg(y)
29064 .arg(&cd)
29065 .arg(&ti)
29066 .arg(&dc);
29067 unsafe {
29068 b.launch(cfg)?;
29069 }
29070 }
29071 match (ring_old, pad_len) {
29072 (None, Some(len_d)) => {
29073 let f = self.func("ssm_conv_ring_update_dev_f32");
29074 let n = conv_dim * (d_conv - 1);
29075 let cfg = LaunchConfig::for_num_elems(n as u32);
29076 let (cd, dc) = (conv_dim as i32, d_conv as i32);
29077 let __s_b = self.gpu.stream();
29078 let mut b = __s_b.launch_builder(&f);
29079 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
29080 unsafe {
29081 b.launch(cfg)?;
29082 }
29083 }
29084 (None, None) => {
29085 let f = self.func("ssm_conv_ring_update_f32");
29086 let n = conv_dim * (d_conv - 1);
29087 let cfg = LaunchConfig::for_num_elems(n as u32);
29088 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29089 let __s_b = self.gpu.stream();
29090 let mut b = __s_b.launch_builder(&f);
29091 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
29092 unsafe {
29093 b.launch(cfg)?;
29094 }
29095 }
29096 (Some(old), _) => {
29097 self.ssm_conv_ring_rebuild(qkv_tm, &old, conv_state, conv_dim, t, d_conv)?
29098 }
29099 }
29100 Ok(())
29101 }
29102
29103 #[allow(clippy::too_many_arguments)]
29105 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_tm_state_pad_v(
29108 &self,
29109 qkv_tm: &cudarc::driver::CudaView<f32>,
29110 conv_state: &mut CudaSlice<f32>,
29111 w: &CudaSlice<f32>,
29112 y: &mut CudaSlice<f32>,
29113 conv_dim: usize,
29114 t: usize,
29115 d_conv: usize,
29116 pad_len: Option<&CudaSlice<i32>>,
29117 ) -> Result<(), Box<dyn std::error::Error>> {
29118 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
29119 let ring_old = if t < d_conv - 1 {
29123 Some(self.clone_dtod(conv_state)?)
29124 } else {
29125 None
29126 };
29127 {
29128 let f = self.func("ssm_conv1d_tm_state_f32");
29129 let cfg = LaunchConfig {
29130 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
29131 block_dim: (256, 1, 1),
29132 shared_mem_bytes: 0,
29133 };
29134 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29135 let __s_b = self.gpu.stream();
29136 let mut b = __s_b.launch_builder(&f);
29137 b.arg(qkv_tm)
29138 .arg(&*conv_state)
29139 .arg(w)
29140 .arg(y)
29141 .arg(&cd)
29142 .arg(&ti)
29143 .arg(&dc);
29144 unsafe {
29145 b.launch(cfg)?;
29146 }
29147 }
29148 match (ring_old, pad_len) {
29149 (None, Some(len_d)) => {
29150 let f = self.func("ssm_conv_ring_update_dev_f32");
29151 let n = conv_dim * (d_conv - 1);
29152 let cfg = LaunchConfig::for_num_elems(n as u32);
29153 let (cd, dc) = (conv_dim as i32, d_conv as i32);
29154 let __s_b = self.gpu.stream();
29155 let mut b = __s_b.launch_builder(&f);
29156 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
29157 unsafe {
29158 b.launch(cfg)?;
29159 }
29160 }
29161 (None, None) => {
29162 let f = self.func("ssm_conv_ring_update_f32");
29163 let n = conv_dim * (d_conv - 1);
29164 let cfg = LaunchConfig::for_num_elems(n as u32);
29165 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29166 let __s_b = self.gpu.stream();
29167 let mut b = __s_b.launch_builder(&f);
29168 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
29169 unsafe {
29170 b.launch(cfg)?;
29171 }
29172 }
29173 (Some(_), _) => unreachable!(
29174 "ssm_conv1d_tm_state_pad_v: T < d_conv-1 has no view path (PRIME_MIN_T gates it)"
29175 ),
29176 }
29177 Ok(())
29178 }
29179
29180 pub fn ssm_conv_ring_rebuild(
29185 &self,
29186 qkv_tm: &CudaSlice<f32>,
29187 ring_old: &CudaSlice<f32>,
29188 conv_state: &mut CudaSlice<f32>,
29189 conv_dim: usize,
29190 tc: usize,
29191 d_conv: usize,
29192 ) -> Result<(), Box<dyn std::error::Error>> {
29193 let f = self.func("ssm_conv_ring_rebuild_f32");
29194 let n = conv_dim * (d_conv - 1);
29195 let cfg = LaunchConfig::for_num_elems(n as u32);
29196 let (cd, ti, dc) = (conv_dim as i32, tc 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)
29200 .arg(ring_old)
29201 .arg(conv_state)
29202 .arg(&cd)
29203 .arg(&ti)
29204 .arg(&dc);
29205 unsafe {
29206 b.launch(cfg)?;
29207 }
29208 Ok(())
29209 }
29210
29211 #[allow(clippy::too_many_arguments)]
29216 pub fn gdn_prep_decode(
29217 &self,
29218 conv_out: &CudaSlice<f32>,
29219 beta_raw: &CudaSlice<f32>,
29220 alpha: &CudaSlice<f32>,
29221 dt_bias: &CudaSlice<f32>,
29222 a: &CudaSlice<f32>,
29223 q_l2: &mut CudaSlice<f32>,
29224 k_l2: &mut CudaSlice<f32>,
29225 v_g: &mut CudaSlice<f32>,
29226 beta: &mut CudaSlice<f32>,
29227 g_log: &mut CudaSlice<f32>,
29228 d_state: usize,
29229 num_v: usize,
29230 num_k: usize,
29231 key_dim: usize,
29232 eps: f32,
29233 ) -> Result<(), Box<dyn std::error::Error>> {
29234 let f = self.func("gdn_prep_decode_f32");
29235 let cfg = LaunchConfig {
29236 grid_dim: (num_v as u32, 1, 1),
29237 block_dim: (32, 4, 1),
29238 shared_mem_bytes: 0,
29239 };
29240 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
29241 let __s_b = self.gpu.stream();
29242 let mut b = __s_b.launch_builder(&f);
29243 b.arg(conv_out)
29244 .arg(beta_raw)
29245 .arg(alpha)
29246 .arg(dt_bias)
29247 .arg(a)
29248 .arg(q_l2)
29249 .arg(k_l2)
29250 .arg(v_g)
29251 .arg(beta)
29252 .arg(g_log)
29253 .arg(&ds)
29254 .arg(&nv)
29255 .arg(&nk)
29256 .arg(&kd)
29257 .arg(&eps);
29258 unsafe {
29259 b.launch(cfg)?;
29260 }
29261 Ok(())
29262 }
29263
29264 #[allow(clippy::too_many_arguments)]
29268 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_gdn(
29270 &self,
29271 qkv_tm: &CudaSlice<f32>,
29272 w: &CudaSlice<f32>,
29273 q_g: &mut CudaSlice<f32>,
29274 k_g: &mut CudaSlice<f32>,
29275 v_g: &mut CudaSlice<f32>,
29276 conv_dim: usize,
29277 t: usize,
29278 d_conv: usize,
29279 d_state: usize,
29280 num_v: usize,
29281 num_k: usize,
29282 key_dim: usize,
29283 ) -> Result<(), Box<dyn std::error::Error>> {
29284 let f = self.func("ssm_conv1d_gdn_f32");
29285 let cfg = LaunchConfig {
29286 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
29287 block_dim: (256, 1, 1),
29288 shared_mem_bytes: 0,
29289 };
29290 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29291 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
29292 let __s_b = self.gpu.stream();
29293 let mut b = __s_b.launch_builder(&f);
29294 b.arg(qkv_tm)
29295 .arg(w)
29296 .arg(q_g)
29297 .arg(k_g)
29298 .arg(v_g)
29299 .arg(&cd)
29300 .arg(&ti)
29301 .arg(&dc)
29302 .arg(&ds)
29303 .arg(&nv)
29304 .arg(&nk)
29305 .arg(&kd);
29306 unsafe {
29307 b.launch(cfg)?;
29308 }
29309 Ok(())
29310 }
29311
29312 #[allow(clippy::too_many_arguments)]
29313 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d(
29316 &self,
29317 x: &CudaSlice<f32>,
29318 w: &CudaSlice<f32>,
29319 y: &mut CudaSlice<f32>,
29320 conv_dim: usize,
29321 t: usize,
29322 d_conv: usize,
29323 silu: bool,
29324 ) -> Result<(), Box<dyn std::error::Error>> {
29325 let f = self.func("ssm_conv1d_silu_f32");
29326 let cfg = LaunchConfig {
29327 grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
29328 block_dim: (256, 1, 1),
29329 shared_mem_bytes: 0,
29330 };
29331 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
29332 let __s_b = self.gpu.stream();
29333 let mut b = __s_b.launch_builder(&f);
29334 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
29335 unsafe {
29336 b.launch(cfg)?;
29337 }
29338 Ok(())
29339 }
29340
29341 #[allow(clippy::too_many_arguments)] pub fn gdn_scan_s128(
29345 &self,
29346 q: &CudaSlice<f32>,
29347 k: &CudaSlice<f32>,
29348 v: &CudaSlice<f32>,
29349 g: &CudaSlice<f32>,
29350 beta: &CudaSlice<f32>,
29351 state_in: &CudaSlice<f32>,
29352 state_out: &mut CudaSlice<f32>,
29353 o: &mut CudaSlice<f32>,
29354 n_head: usize,
29355 t: usize,
29356 scale: f32,
29357 ) -> Result<(), Box<dyn std::error::Error>> {
29358 let f = self.func("gdn_scan_s128");
29359 const S_V: u32 = 128;
29360 const WARP: u32 = 32;
29361 const COLS_PER_BLOCK: u32 = 4;
29362 let cfg = LaunchConfig {
29363 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
29364 block_dim: (WARP, COLS_PER_BLOCK, 1),
29365 shared_mem_bytes: 0,
29366 };
29367 let (h, ti) = (n_head as i32, t as i32);
29368 let __s_b = self.gpu.stream();
29369 let mut b = __s_b.launch_builder(&f);
29370 b.arg(q)
29371 .arg(k)
29372 .arg(v)
29373 .arg(g)
29374 .arg(beta)
29375 .arg(state_in)
29376 .arg(state_out)
29377 .arg(o)
29378 .arg(&h)
29379 .arg(&ti)
29380 .arg(&scale);
29381 unsafe {
29382 b.launch(cfg)?;
29383 }
29384 Ok(())
29385 }
29386
29387 #[allow(clippy::too_many_arguments)]
29392 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_fused_decode_b(
29394 &self,
29395 qkv_cols: &CudaSlice<f32>,
29396 conv_state_ptrs: &cudarc::driver::CudaView<u64>,
29397 w: &CudaSlice<f32>,
29398 conv_outs: &mut CudaSlice<f32>,
29399 conv_dim: usize,
29400 d_conv: usize,
29401 b_n: usize,
29402 ) -> Result<(), Box<dyn std::error::Error>> {
29403 let f = self.func("ssm_conv1d_fused_decode_b_f32");
29404 let cfg = LaunchConfig {
29405 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
29406 block_dim: (256, 1, 1),
29407 shared_mem_bytes: 0,
29408 };
29409 let (cd, dc) = (conv_dim as i32, d_conv as i32);
29410 let __s_b = self.gpu.stream();
29411 let mut b = __s_b.launch_builder(&f);
29412 b.arg(qkv_cols)
29413 .arg(conv_state_ptrs)
29414 .arg(w)
29415 .arg(conv_outs)
29416 .arg(&cd)
29417 .arg(&dc);
29418 unsafe {
29419 b.launch(cfg)?;
29420 }
29421 Ok(())
29422 }
29423
29424 #[allow(clippy::too_many_arguments)]
29425 pub fn gdn_prep_decode_b(
29426 &self,
29427 conv_outs: &CudaSlice<f32>,
29428 beta_raws: &CudaSlice<f32>,
29429 alphas: &CudaSlice<f32>,
29430 dt_bias: &CudaSlice<f32>,
29431 a: &CudaSlice<f32>,
29432 q_l2: &mut CudaSlice<f32>,
29433 k_l2: &mut CudaSlice<f32>,
29434 v_g: &mut CudaSlice<f32>,
29435 beta: &mut CudaSlice<f32>,
29436 g_log: &mut CudaSlice<f32>,
29437 d_state: usize,
29438 num_v: usize,
29439 num_k: usize,
29440 key_dim: usize,
29441 eps: f32,
29442 conv_dim: usize,
29443 b_n: usize,
29444 ) -> Result<(), Box<dyn std::error::Error>> {
29445 let f = self.func("gdn_prep_decode_b_f32");
29446 let cfg = LaunchConfig {
29447 grid_dim: (num_v as u32, 1, b_n as u32),
29448 block_dim: (32, 4, 1),
29449 shared_mem_bytes: 0,
29450 };
29451 let (ds, nv, nk, kd, cd) = (
29452 d_state as i32,
29453 num_v as i32,
29454 num_k as i32,
29455 key_dim as i32,
29456 conv_dim as i32,
29457 );
29458 let __s_b = self.gpu.stream();
29459 let mut b = __s_b.launch_builder(&f);
29460 b.arg(conv_outs)
29461 .arg(beta_raws)
29462 .arg(alphas)
29463 .arg(dt_bias)
29464 .arg(a)
29465 .arg(q_l2)
29466 .arg(k_l2)
29467 .arg(v_g)
29468 .arg(beta)
29469 .arg(g_log)
29470 .arg(&ds)
29471 .arg(&nv)
29472 .arg(&nk)
29473 .arg(&kd)
29474 .arg(&eps)
29475 .arg(&cd);
29476 unsafe {
29477 b.launch(cfg)?;
29478 }
29479 Ok(())
29480 }
29481
29482 #[allow(clippy::too_many_arguments)]
29483 pub fn gdn_scan_s128_batched(
29484 &self,
29485 q: &CudaSlice<f32>,
29486 k: &CudaSlice<f32>,
29487 v: &CudaSlice<f32>,
29488 g: &CudaSlice<f32>,
29489 beta: &CudaSlice<f32>,
29490 state_in_ptrs: &cudarc::driver::CudaView<u64>,
29491 state_out_ptrs: &cudarc::driver::CudaView<u64>,
29492 o: &mut CudaSlice<f32>,
29493 n_head: usize,
29494 b_n: usize,
29495 scale: f32,
29496 ) -> Result<(), Box<dyn std::error::Error>> {
29497 let f = self.func("gdn_scan_s128_b");
29498 const S_V: u32 = 128;
29499 const WARP: u32 = 32;
29500 const COLS_PER_BLOCK: u32 = 4;
29501 let cfg = LaunchConfig {
29502 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
29503 block_dim: (WARP, COLS_PER_BLOCK, 1),
29504 shared_mem_bytes: 0,
29505 };
29506 let h = n_head as i32;
29507 let __s_b = self.gpu.stream();
29508 let mut b = __s_b.launch_builder(&f);
29509 b.arg(q)
29510 .arg(k)
29511 .arg(v)
29512 .arg(g)
29513 .arg(beta)
29514 .arg(state_in_ptrs)
29515 .arg(state_out_ptrs)
29516 .arg(o)
29517 .arg(&h)
29518 .arg(&scale);
29519 unsafe {
29520 b.launch(cfg)?;
29521 }
29522 Ok(())
29523 }
29524
29525 #[allow(clippy::too_many_arguments)]
29531 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_fused_decode_b_view(
29533 &self,
29534 qkv_cols: &cudarc::driver::CudaView<f32>,
29535 conv_state_ptrs: &cudarc::driver::CudaView<u64>,
29536 w: &CudaSlice<f32>,
29537 conv_outs: &mut CudaSlice<f32>,
29538 conv_dim: usize,
29539 d_conv: usize,
29540 b_n: usize,
29541 ) -> Result<(), Box<dyn std::error::Error>> {
29542 let f = self.func("ssm_conv1d_fused_decode_b_f32");
29543 let cfg = LaunchConfig {
29544 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
29545 block_dim: (256, 1, 1),
29546 shared_mem_bytes: 0,
29547 };
29548 let (cd, dc) = (conv_dim as i32, d_conv as i32);
29549 let __s_b = self.gpu.stream();
29550 let mut b = __s_b.launch_builder(&f);
29551 b.arg(qkv_cols)
29552 .arg(conv_state_ptrs)
29553 .arg(w)
29554 .arg(conv_outs)
29555 .arg(&cd)
29556 .arg(&dc);
29557 unsafe {
29558 b.launch(cfg)?;
29559 }
29560 Ok(())
29561 }
29562
29563 #[allow(clippy::too_many_arguments)]
29564 pub fn gdn_prep_decode_b_view(
29565 &self,
29566 conv_outs: &CudaSlice<f32>,
29567 beta_raws: &cudarc::driver::CudaView<f32>,
29568 alphas: &cudarc::driver::CudaView<f32>,
29569 dt_bias: &CudaSlice<f32>,
29570 a: &CudaSlice<f32>,
29571 q_l2: &mut CudaSlice<f32>,
29572 k_l2: &mut CudaSlice<f32>,
29573 v_g: &mut CudaSlice<f32>,
29574 beta: &mut CudaSlice<f32>,
29575 g_log: &mut CudaSlice<f32>,
29576 d_state: usize,
29577 num_v: usize,
29578 num_k: usize,
29579 key_dim: usize,
29580 eps: f32,
29581 conv_dim: usize,
29582 b_n: usize,
29583 ) -> Result<(), Box<dyn std::error::Error>> {
29584 let f = self.func("gdn_prep_decode_b_f32");
29585 let cfg = LaunchConfig {
29586 grid_dim: (num_v as u32, 1, b_n as u32),
29587 block_dim: (32, 4, 1),
29588 shared_mem_bytes: 0,
29589 };
29590 let (ds, nv, nk, kd, cd) = (
29591 d_state as i32,
29592 num_v as i32,
29593 num_k as i32,
29594 key_dim as i32,
29595 conv_dim as i32,
29596 );
29597 let __s_b = self.gpu.stream();
29598 let mut b = __s_b.launch_builder(&f);
29599 b.arg(conv_outs)
29600 .arg(beta_raws)
29601 .arg(alphas)
29602 .arg(dt_bias)
29603 .arg(a)
29604 .arg(q_l2)
29605 .arg(k_l2)
29606 .arg(v_g)
29607 .arg(beta)
29608 .arg(g_log)
29609 .arg(&ds)
29610 .arg(&nv)
29611 .arg(&nk)
29612 .arg(&kd)
29613 .arg(&eps)
29614 .arg(&cd);
29615 unsafe {
29616 b.launch(cfg)?;
29617 }
29618 Ok(())
29619 }
29620
29621 #[allow(clippy::too_many_arguments)]
29622 pub fn gdn_scan_s128_batched_view(
29623 &self,
29624 q: &CudaSlice<f32>,
29625 k: &CudaSlice<f32>,
29626 v: &CudaSlice<f32>,
29627 g: &CudaSlice<f32>,
29628 beta: &CudaSlice<f32>,
29629 state_in_ptrs: &cudarc::driver::CudaView<u64>,
29630 state_out_ptrs: &cudarc::driver::CudaView<u64>,
29631 o: &mut cudarc::driver::CudaViewMut<f32>,
29632 n_head: usize,
29633 b_n: usize,
29634 scale: f32,
29635 ) -> Result<(), Box<dyn std::error::Error>> {
29636 let f = self.func("gdn_scan_s128_b");
29637 const S_V: u32 = 128;
29638 const WARP: u32 = 32;
29639 const COLS_PER_BLOCK: u32 = 4;
29640 let cfg = LaunchConfig {
29641 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
29642 block_dim: (WARP, COLS_PER_BLOCK, 1),
29643 shared_mem_bytes: 0,
29644 };
29645 let h = n_head as i32;
29646 let __s_b = self.gpu.stream();
29647 let mut b = __s_b.launch_builder(&f);
29648 b.arg(q)
29649 .arg(k)
29650 .arg(v)
29651 .arg(g)
29652 .arg(beta)
29653 .arg(state_in_ptrs)
29654 .arg(state_out_ptrs)
29655 .arg(o)
29656 .arg(&h)
29657 .arg(&scale);
29658 unsafe {
29659 b.launch(cfg)?;
29660 }
29661 Ok(())
29662 }
29663
29664 pub fn gdn_chunked_enabled() -> bool {
29673 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
29674 *E.get_or_init(|| {
29675 std::env::var("MEMRA_GDN_CHUNKED")
29676 .map(|v| v != "0")
29677 .unwrap_or(true)
29678 })
29679 }
29680
29681 pub fn gdn_chunk_size() -> usize {
29686 static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
29687 *C.get_or_init(|| {
29688 let c: usize = std::env::var("MEMRA_GDN_CHUNK")
29689 .ok()
29690 .and_then(|v| v.parse().ok())
29691 .unwrap_or(32);
29692 c.clamp(32, 128) / 32 * 32
29693 })
29694 }
29695
29696 #[allow(clippy::too_many_arguments)]
29701 #[allow(clippy::too_many_arguments, clippy::type_complexity)]
29704 #[allow(clippy::too_many_arguments)]
29705 pub fn gdn_chunk_k123(
29706 &self,
29707 q: &CudaSlice<f32>,
29708 k: &CudaSlice<f32>,
29709 v: &CudaSlice<f32>,
29710 g: &CudaSlice<f32>,
29711 beta: &CudaSlice<f32>,
29712 wb16: Option<&mut CudaSlice<u8>>,
29713 n_head: usize,
29714 t: usize,
29715 c: usize,
29716 hk: usize,
29717 k2w: Option<(&CudaSlice<u8>, &CudaSlice<u8>, &mut CudaSlice<u8>)>,
29718 ) -> Result<
29719 (
29720 CudaSlice<f32>,
29721 CudaSlice<f32>,
29722 CudaSlice<f32>,
29723 CudaSlice<f32>,
29724 ),
29725 Box<dyn std::error::Error>,
29726 > {
29727 const D: usize = 128;
29728 let h = n_head;
29729 #[allow(clippy::manual_div_ceil)]
29730 let nc = (t + c - 1) / c;
29732 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
29733 let mut gcum = self.uninit(t * h)?;
29734 let mut a = self.uninit(nc * h * c * c)?;
29735 let mut p = self.uninit(nc * h * c * c)?;
29736 let mut u = self.uninit(nc * h * c * D)?;
29737 let mut w = self.uninit(nc * h * c * D)?;
29738 {
29739 let f = self.func("gdn_chunk_cumgate_f32");
29741 let cfg = LaunchConfig {
29742 grid_dim: (nc as u32, h as u32, 1),
29743 block_dim: (32, 1, 1),
29744 shared_mem_bytes: 0,
29745 };
29746 let __s_b = self.gpu.stream();
29747 let mut b = __s_b.launch_builder(&f);
29748 b.arg(g).arg(&mut gcum).arg(&hi).arg(&ti).arg(&ci);
29749 unsafe {
29750 b.launch(cfg)?;
29751 }
29752 }
29753 if let Some((qb, kb, pb)) = k2w {
29754 assert!(c == 32, "gdn_k2_wgmma is a C==32 tile");
29757 let f = self.func("gdn_k2_wgmma");
29758 let cfg = LaunchConfig {
29759 grid_dim: (nc as u32, h as u32, 1),
29760 block_dim: (128, 1, 1),
29761 shared_mem_bytes: 0,
29762 };
29763 let hki = hk as i32;
29764 let __s_b = self.gpu.stream();
29765 let mut b = __s_b.launch_builder(&f);
29766 b.arg(qb)
29767 .arg(kb)
29768 .arg(&gcum)
29769 .arg(beta)
29770 .arg(&mut a)
29771 .arg(&mut *pb)
29772 .arg(&hi)
29773 .arg(&ti)
29774 .arg(&ci)
29775 .arg(&hki);
29776 unsafe {
29777 b.launch(cfg)?;
29778 }
29779 } else if c <= 64 && !portable_mma_gated() {
29780 let f = self.func("gdn_chunk_attn_f32");
29782 f.set_attribute(
29783 CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
29784 GDN_K2_DYNAMIC_SHARED_BYTES as i32,
29785 )?;
29786 #[allow(clippy::manual_div_ceil)]
29787 let jt = ((c + 31) / 32) as u32;
29789 let cfg = LaunchConfig {
29790 grid_dim: (nc as u32, h as u32, jt),
29791 block_dim: (256, 1, 1),
29792 shared_mem_bytes: GDN_K2_DYNAMIC_SHARED_BYTES,
29793 };
29794 let hki = hk as i32;
29795 let __s_b = self.gpu.stream();
29796 let mut b = __s_b.launch_builder(&f);
29797 b.arg(q)
29798 .arg(k)
29799 .arg(&gcum)
29800 .arg(beta)
29801 .arg(&mut a)
29802 .arg(&mut p)
29803 .arg(&hi)
29804 .arg(&ti)
29805 .arg(&ci)
29806 .arg(&hki);
29807 unsafe {
29808 b.launch(cfg)?;
29809 }
29810 } else {
29811 assert!(
29813 hk == h,
29814 "generic K2 is broadcast-only (de-broadcast rides C==32)"
29815 );
29816 let f = self.func("gdn_chunk_attn_g_f32");
29817 let cfg = LaunchConfig {
29818 grid_dim: (nc as u32, h as u32, 1),
29819 block_dim: (32, 8, 1),
29820 shared_mem_bytes: 0,
29821 };
29822 let __s_b = self.gpu.stream();
29823 let mut b = __s_b.launch_builder(&f);
29824 b.arg(q)
29825 .arg(k)
29826 .arg(&gcum)
29827 .arg(beta)
29828 .arg(&mut a)
29829 .arg(&mut p)
29830 .arg(&hi)
29831 .arg(&ti)
29832 .arg(&ci);
29833 unsafe {
29834 b.launch(cfg)?;
29835 }
29836 }
29837 {
29838 let cfg = LaunchConfig {
29840 grid_dim: (nc as u32, h as u32, 1),
29841 block_dim: (256, 1, 1),
29842 shared_mem_bytes: 0,
29843 };
29844 match c {
29845 32 | 64 => {
29846 let f = self.func(if c == 32 {
29847 "gdn_chunk_solve32_f32"
29848 } else {
29849 "gdn_chunk_solve64_f32"
29850 });
29851 let wb: u64 = match wb16 {
29853 Some(d) => self.addr_u8(d),
29854 None => 0,
29855 };
29856 let hki = hk as i32;
29857 let __s_b = self.gpu.stream();
29858 let mut b = __s_b.launch_builder(&f);
29859 b.arg(v)
29860 .arg(k)
29861 .arg(&a)
29862 .arg(&gcum)
29863 .arg(&mut u)
29864 .arg(&mut w)
29865 .arg(&wb)
29866 .arg(&hi)
29867 .arg(&ti)
29868 .arg(&hki);
29869 unsafe {
29870 b.launch(cfg)?;
29871 }
29872 }
29873 _ => {
29874 assert!(hk == h, "generic K3 is broadcast-only");
29875 let f = self.func("gdn_chunk_solve_f32");
29876 let __s_b = self.gpu.stream();
29877 let mut b = __s_b.launch_builder(&f);
29878 b.arg(v)
29879 .arg(k)
29880 .arg(&a)
29881 .arg(&gcum)
29882 .arg(&mut u)
29883 .arg(&mut w)
29884 .arg(&hi)
29885 .arg(&ti)
29886 .arg(&ci);
29887 unsafe {
29888 b.launch(cfg)?;
29889 }
29890 }
29891 }
29892 }
29893 Ok((gcum, p, u, w))
29894 }
29895
29896 pub fn gdn_db_on() -> bool {
29900 std::env::var("MEMRA_GDN_DB").as_deref() != Ok("0")
29901 }
29902
29903 pub fn gdn_mma_enabled(&self, c: usize) -> bool {
29912 !portable_mma_gated()
29913 && c == 32
29914 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
29915 Ok("1") => true,
29916 Ok("0") => false,
29917 _ => gdn_mma_default_on(),
29918 }
29919 }
29920
29921 pub fn gdn_wgmma_on(&self, c: usize) -> bool {
29928 cfg!(memra_hopper_mma)
29929 && self.gdn_mma_enabled(c)
29930 && std::env::var("MEMRA_GDN_WGMMA").as_deref() != Ok("0")
29931 }
29932
29933 #[allow(clippy::too_many_arguments)]
29938 #[allow(clippy::manual_div_ceil)] pub fn ssm_conv1d_gdn_state_pad(
29940 &self,
29941 qkv_tm: &cudarc::driver::CudaView<f32>,
29942 conv_state: &mut CudaSlice<f32>,
29943 w: &CudaSlice<f32>,
29944 q_g: &mut CudaSlice<f32>,
29945 k_g: &mut CudaSlice<f32>,
29946 v_g: &mut CudaSlice<f32>,
29947 conv_dim: usize,
29948 t: usize,
29949 d_conv: usize,
29950 d_state: usize,
29951 num_v: usize,
29952 num_k: usize,
29953 key_dim: usize,
29954 hk: usize,
29955 pad_len: Option<&CudaSlice<i32>>,
29956 ) -> Result<(), Box<dyn std::error::Error>> {
29957 assert!(
29958 t >= d_conv - 1,
29959 "fused state conv requires T >= pad (PRIME_MIN_T gates)"
29960 );
29961 {
29962 let f = self.func("ssm_conv1d_gdn_state_f32");
29963 let cfg = LaunchConfig {
29964 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
29965 block_dim: (256, 1, 1),
29966 shared_mem_bytes: 0,
29967 };
29968 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
29969 let (ds, nv, nk, kd, hki) = (
29970 d_state as i32,
29971 num_v as i32,
29972 num_k as i32,
29973 key_dim as i32,
29974 hk as i32,
29975 );
29976 let __s_b = self.gpu.stream();
29977 let mut b = __s_b.launch_builder(&f);
29978 b.arg(qkv_tm)
29979 .arg(&*conv_state)
29980 .arg(w)
29981 .arg(q_g)
29982 .arg(k_g)
29983 .arg(v_g)
29984 .arg(&cd)
29985 .arg(&ti)
29986 .arg(&dc)
29987 .arg(&ds)
29988 .arg(&nv)
29989 .arg(&nk)
29990 .arg(&kd)
29991 .arg(&hki);
29992 unsafe {
29993 b.launch(cfg)?;
29994 }
29995 }
29996 match pad_len {
29997 Some(len_d) => {
29998 let f = self.func("ssm_conv_ring_update_dev_f32");
29999 let n = conv_dim * (d_conv - 1);
30000 let cfg = LaunchConfig::for_num_elems(n as u32);
30001 let (cd, dc) = (conv_dim as i32, d_conv as i32);
30002 let __s_b = self.gpu.stream();
30003 let mut b = __s_b.launch_builder(&f);
30004 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
30005 unsafe {
30006 b.launch(cfg)?;
30007 }
30008 }
30009 None => {
30010 let f = self.func("ssm_conv_ring_update_f32");
30011 let n = conv_dim * (d_conv - 1);
30012 let cfg = LaunchConfig::for_num_elems(n as u32);
30013 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
30014 let __s_b = self.gpu.stream();
30015 let mut b = __s_b.launch_builder(&f);
30016 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
30017 unsafe {
30018 b.launch(cfg)?;
30019 }
30020 }
30021 }
30022 Ok(())
30023 }
30024
30025 pub fn gdn_chunk_alloc(
30029 &self,
30030 n_head: usize,
30031 t: usize,
30032 c: usize,
30033 hk: usize,
30034 ) -> Result<GdnChunkBufs, Box<dyn std::error::Error>> {
30035 const D: usize = 128;
30036 assert!(
30037 c == 32,
30038 "gdn_chunk_alloc: varlen chain is the C==32 mma pair"
30039 );
30040 let h = n_head;
30041 #[allow(clippy::manual_div_ceil)]
30042 let nc = (t + c - 1) / c;
30044 Ok(GdnChunkBufs {
30045 gcum: self.uninit(t * h)?,
30046 a: self.uninit(nc * h * c * c)?,
30047 p: self.uninit(nc * h * c * c)?,
30048 u: self.uninit(nc * h * c * D)?,
30049 w: self.uninit(nc * h * c * D)?,
30050 kb16: self.alloc_u8_uninit(t * hk * D * 2)?,
30051 wb16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
30052 y16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
30053 ssnap16: self.alloc_u8_uninit(nc * h * D * D * 2)?,
30054 qb16: self.alloc_u8_uninit(t * hk * D * 2)?,
30055 pb16: self.alloc_u8_uninit(nc * h * c * c * 2)?,
30056 o: self.uninit(D * h * t)?,
30057 t,
30058 nc,
30059 })
30060 }
30061
30062 pub fn f32_to_bf16_v(
30064 &self,
30065 x: &cudarc::driver::CudaView<f32>,
30066 dst: &mut CudaSlice<u8>,
30067 n: usize,
30068 ) -> Result<(), Box<dyn std::error::Error>> {
30069 let f = self.func("f32_to_bf16_bulk");
30070 let ni = n as i64;
30071 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
30072 let __s_b = self.gpu.stream();
30073 let mut b = __s_b.launch_builder(&f);
30074 b.arg(x).arg(dst).arg(&ni);
30075 unsafe {
30076 b.launch(cfg)?;
30077 }
30078 Ok(())
30079 }
30080
30081 pub fn f32_to_bf16_into(
30083 &self,
30084 x: &CudaSlice<f32>,
30085 dst: &mut CudaSlice<u8>,
30086 n: usize,
30087 ) -> Result<(), Box<dyn std::error::Error>> {
30088 let f = self.func("f32_to_bf16_bulk");
30089 let ni = n as i64;
30090 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
30091 let __s_b = self.gpu.stream();
30092 let mut b = __s_b.launch_builder(&f);
30093 b.arg(x).arg(dst).arg(&ni);
30094 unsafe {
30095 b.launch(cfg)?;
30096 }
30097 Ok(())
30098 }
30099
30100 pub fn gdn_chunk_k123_vl8(
30103 &self,
30104 seqs: &[GdnSeqVl],
30105 n_head: usize,
30106 hk: usize,
30107 wq: Option<&GdnWVl8>,
30108 ) -> Result<(), Box<dyn std::error::Error>> {
30109 let b = seqs.len();
30110 assert!((1..=8).contains(&b), "gdn_chunk_k123_vl8: 1..=8 sequences");
30111 let mut packed = [GdnSeqVl::default(); 8];
30112 packed[..b].copy_from_slice(seqs);
30113 let v = GdnVl8(packed);
30114 let (hi, ci) = (n_head as i32, 32i32);
30115 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
30116 {
30117 let f = self.func("gdn_chunk_cumgate_vl");
30118 let cfg = LaunchConfig {
30119 grid_dim: (max_nc, n_head as u32, b as u32),
30120 block_dim: (32, 1, 1),
30121 shared_mem_bytes: 0,
30122 };
30123 let __s_lb = self.gpu.stream();
30124 let mut lb = __s_lb.launch_builder(&f);
30125 lb.arg(&v).arg(&hi).arg(&ci);
30126 unsafe {
30127 lb.launch(cfg)?;
30128 }
30129 }
30130 let hki = hk as i32;
30131 if let Some(w) = wq {
30132 let f = self.func("gdn_k2_wgmma_vl");
30134 let cfg = LaunchConfig {
30135 grid_dim: (max_nc, n_head as u32, b as u32),
30136 block_dim: (128, 1, 1),
30137 shared_mem_bytes: 0,
30138 };
30139 let __s_lb = self.gpu.stream();
30140 let mut lb = __s_lb.launch_builder(&f);
30141 lb.arg(&v).arg(w).arg(&hi).arg(&ci).arg(&hki);
30142 unsafe {
30143 lb.launch(cfg)?;
30144 }
30145 } else {
30146 let f = self.func("gdn_chunk_attn_vl");
30147 f.set_attribute(
30148 CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
30149 GDN_K2_DYNAMIC_SHARED_BYTES as i32,
30150 )?;
30151 let cfg = LaunchConfig {
30152 grid_dim: (max_nc, n_head as u32, b as u32),
30153 block_dim: (256, 1, 1),
30154 shared_mem_bytes: GDN_K2_DYNAMIC_SHARED_BYTES,
30155 };
30156 let __s_lb = self.gpu.stream();
30157 let mut lb = __s_lb.launch_builder(&f);
30158 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
30159 unsafe {
30160 lb.launch(cfg)?;
30161 }
30162 }
30163 {
30164 let f = self.func("gdn_chunk_solve32_vl");
30165 let cfg = LaunchConfig {
30166 grid_dim: (max_nc, n_head as u32, b as u32),
30167 block_dim: (256, 1, 1),
30168 shared_mem_bytes: 0,
30169 };
30170 let __s_lb = self.gpu.stream();
30171 let mut lb = __s_lb.launch_builder(&f);
30172 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
30173 unsafe {
30174 lb.launch(cfg)?;
30175 }
30176 }
30177 Ok(())
30178 }
30179
30180 #[allow(clippy::too_many_arguments)]
30184 pub fn gdn_prep_vl8(
30185 &self,
30186 seqs: &[GdnPrepVl],
30187 conv_w: &CudaSlice<f32>,
30188 dt_bias: &CudaSlice<f32>,
30189 a: &CudaSlice<f32>,
30190 conv_dim: usize,
30191 d_conv: usize,
30192 d_state: usize,
30193 num_v: usize,
30194 num_k: usize,
30195 key_dim: usize,
30196 hk: usize,
30197 eps: f32,
30198 ) -> Result<(), Box<dyn std::error::Error>> {
30199 let b = seqs.len();
30200 assert!((1..=8).contains(&b));
30201 let mut packed = [GdnPrepVl::default(); 8];
30202 packed[..b].copy_from_slice(seqs);
30203 let v = GdnPrepVl8(packed);
30204 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
30205 let (cdi, dci) = (conv_dim as i32, d_conv as i32);
30206 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
30207 assert!(
30208 conv_fuse || hk == num_v,
30209 "de-broadcast requires the fused conv"
30210 );
30211 if conv_fuse {
30212 let f = self.func("ssm_conv1d_gdn_state_vl");
30213 let cfg = LaunchConfig {
30214 grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32),
30215 block_dim: (256, 1, 1),
30216 shared_mem_bytes: 0,
30217 };
30218 let (dsi, nvi, nki, kdi, hki) = (
30219 d_state as i32,
30220 num_v as i32,
30221 num_k as i32,
30222 key_dim as i32,
30223 hk as i32,
30224 );
30225 let __s_lb = self.gpu.stream();
30226 let mut lb = __s_lb.launch_builder(&f);
30227 lb.arg(&v)
30228 .arg(conv_w)
30229 .arg(&cdi)
30230 .arg(&dci)
30231 .arg(&dsi)
30232 .arg(&nvi)
30233 .arg(&nki)
30234 .arg(&kdi)
30235 .arg(&hki);
30236 unsafe {
30237 lb.launch(cfg)?;
30238 }
30239 } else {
30240 let f = self.func("ssm_conv1d_tm_state_vl");
30241 let cfg = LaunchConfig {
30242 grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32),
30243 block_dim: (256, 1, 1),
30244 shared_mem_bytes: 0,
30245 };
30246 let __s_lb = self.gpu.stream();
30247 let mut lb = __s_lb.launch_builder(&f);
30248 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci);
30249 unsafe {
30250 lb.launch(cfg)?;
30251 }
30252 }
30253 {
30254 let f = self.func("ssm_conv_ring_update_vl");
30255 let n = (conv_dim * (d_conv - 1)) as u32;
30256 let cfg = LaunchConfig {
30257 grid_dim: (n.div_ceil(256), 1, b as u32),
30258 block_dim: (256, 1, 1),
30259 shared_mem_bytes: 0,
30260 };
30261 let __s_lb = self.gpu.stream();
30262 let mut lb = __s_lb.launch_builder(&f);
30263 lb.arg(&v).arg(&cdi).arg(&dci);
30264 unsafe {
30265 lb.launch(cfg)?;
30266 }
30267 }
30268 if !conv_fuse {
30269 let f = self.func("qkv_to_gdn_repack_vl");
30270 let n = max_t * (num_v * d_state) as u32;
30271 let cfg = LaunchConfig {
30272 grid_dim: (n.div_ceil(256), 1, b as u32),
30273 block_dim: (256, 1, 1),
30274 shared_mem_bytes: 0,
30275 };
30276 let (dsi, nvi, nki, kdi) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
30277 let __s_lb = self.gpu.stream();
30278 let mut lb = __s_lb.launch_builder(&f);
30279 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi);
30280 unsafe {
30281 lb.launch(cfg)?;
30282 }
30283 }
30284 if Self::l2_v2_on(d_state) {
30285 let f = self.func("gdn_l2_v2_vl");
30286 let cfg = LaunchConfig {
30287 grid_dim: ((max_t * hk as u32).div_ceil(8), 2, b as u32),
30288 block_dim: (256, 1, 1),
30289 shared_mem_bytes: 0,
30290 };
30291 let (dsi, nvi) = (d_state as i32, hk as i32);
30292 let __s_lb = self.gpu.stream();
30293 let mut lb = __s_lb.launch_builder(&f);
30294 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
30295 unsafe {
30296 lb.launch(cfg)?;
30297 }
30298 } else {
30299 let f = self.func("gdn_l2_vl");
30300 let cfg = LaunchConfig {
30301 grid_dim: (max_t * hk as u32, 2, b as u32),
30302 block_dim: (256, 1, 1),
30303 shared_mem_bytes: 0,
30304 };
30305 let (dsi, nvi) = (d_state as i32, hk as i32);
30306 let __s_lb = self.gpu.stream();
30307 let mut lb = __s_lb.launch_builder(&f);
30308 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
30309 unsafe {
30310 lb.launch(cfg)?;
30311 }
30312 }
30313 {
30314 let f = self.func("gdn_gate_prep_vl");
30315 let n = max_t * num_v as u32;
30316 let cfg = LaunchConfig {
30317 grid_dim: (n.div_ceil(256), 1, b as u32),
30318 block_dim: (256, 1, 1),
30319 shared_mem_bytes: 0,
30320 };
30321 let nvi = num_v as i32;
30322 let __s_lb = self.gpu.stream();
30323 let mut lb = __s_lb.launch_builder(&f);
30324 lb.arg(&v).arg(dt_bias).arg(a).arg(&nvi);
30325 unsafe {
30326 lb.launch(cfg)?;
30327 }
30328 }
30329 Ok(())
30330 }
30331
30332 pub fn gdn_mirror_vl8(
30334 &self,
30335 seqs: &[GdnSeqVl],
30336 n_head: usize,
30337 which: i32,
30338 hk: usize,
30339 ) -> Result<(), Box<dyn std::error::Error>> {
30340 let b = seqs.len();
30341 assert!((1..=8).contains(&b));
30342 let mut packed = [GdnSeqVl::default(); 8];
30343 packed[..b].copy_from_slice(seqs);
30344 let v = GdnVl8(packed);
30345 let ept = (if which == 0 { hk } else { n_head } * 128) as i32;
30346 let max_n = seqs
30347 .iter()
30348 .map(|s| {
30349 if which == 0 {
30350 s.t as i64 * ept as i64
30351 } else {
30352 s.nc as i64 * ept as i64 * 32
30353 }
30354 })
30355 .max()
30356 .unwrap();
30357 let f = self.func("gdn_mirror_vl");
30358 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
30359 let cfg = LaunchConfig {
30360 grid_dim: (blocks, 1, b as u32),
30361 block_dim: (256, 1, 1),
30362 shared_mem_bytes: 0,
30363 };
30364 let __s_lb = self.gpu.stream();
30365 let mut lb = __s_lb.launch_builder(&f);
30366 lb.arg(&v).arg(&ept).arg(&which);
30367 unsafe {
30368 lb.launch(cfg)?;
30369 }
30370 Ok(())
30371 }
30372
30373 pub fn gdn_tail_vl8(
30375 &self,
30376 seqs: &[GdnPrepVl],
30377 norm_w: &CudaSlice<f32>,
30378 d_state: usize,
30379 num_v: usize,
30380 eps: f32,
30381 ) -> Result<(), Box<dyn std::error::Error>> {
30382 let b = seqs.len();
30383 assert!((1..=8).contains(&b));
30384 let mut packed = [GdnPrepVl::default(); 8];
30385 packed[..b].copy_from_slice(seqs);
30386 let v = GdnPrepVl8(packed);
30387 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
30388 let f = self.func("gated_rmsnorm_f16out_vl");
30389 let cfg = LaunchConfig {
30391 grid_dim: (max_t * num_v as u32, 1, b as u32),
30392 block_dim: (128, 1, 1),
30393 shared_mem_bytes: 0,
30394 };
30395 let (dsi, nvi) = (d_state as i32, num_v as i32);
30396 let __s_lb = self.gpu.stream();
30397 let mut lb = __s_lb.launch_builder(&f);
30398 lb.arg(&v).arg(norm_w).arg(&dsi).arg(&nvi).arg(&eps);
30399 unsafe {
30400 lb.launch(cfg)?;
30401 }
30402 Ok(())
30403 }
30404
30405 pub fn addr_f32(&self, x: &CudaSlice<f32>) -> u64 {
30408 use cudarc::driver::DevicePtr;
30409 let s = self.gpu.stream();
30410 let (p, _g) = x.device_ptr(&s);
30411 p
30412 }
30413 pub fn addr_f32_mut(&self, x: &mut CudaSlice<f32>) -> u64 {
30414 use cudarc::driver::DevicePtrMut;
30415 let s = self.gpu.stream();
30416 let (p, _g) = x.device_ptr_mut(&s);
30417 p
30418 }
30419 pub fn addr_f32v(&self, x: &cudarc::driver::CudaView<f32>) -> u64 {
30420 use cudarc::driver::DevicePtr;
30421 let s = self.gpu.stream();
30422 let (p, _g) = x.device_ptr(&s);
30423 p
30424 }
30425 pub fn addr_u8(&self, x: &CudaSlice<u8>) -> u64 {
30426 use cudarc::driver::DevicePtr;
30427 let s = self.gpu.stream();
30428 let (p, _g) = x.device_ptr(&s);
30429 p
30430 }
30431
30432 pub fn gdn_chunk_vl8(
30436 &self,
30437 seqs: &[GdnSeqVl],
30438 n_head: usize,
30439 scale: f32,
30440 hk: usize,
30441 wq: Option<&GdnWVl8>,
30442 ) -> Result<(), Box<dyn std::error::Error>> {
30443 const NSPLIT: u32 = 4;
30444 let b = seqs.len();
30445 assert!((1..=8).contains(&b), "gdn_chunk_vl8: 1..=8 sequences");
30446 let mut packed = [GdnSeqVl::default(); 8];
30447 packed[..b].copy_from_slice(seqs);
30448 let v = GdnVl8(packed);
30449 let (hi, ci) = (n_head as i32, 32i32);
30450 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
30451 let hki = hk as i32;
30452 if let Some(w) = wq {
30453 let f = self.func("gdn_k45_wgmma_vl");
30455 let cfg = LaunchConfig {
30456 grid_dim: (n_head as u32, NSPLIT, b as u32),
30457 block_dim: (256, 1, 1),
30458 shared_mem_bytes: 0,
30459 };
30460 let __s_lb = self.gpu.stream();
30461 let mut lb = __s_lb.launch_builder(&f);
30462 lb.arg(&v).arg(w).arg(&scale).arg(&hi).arg(&ci).arg(&hki);
30463 unsafe {
30464 lb.launch(cfg)?;
30465 }
30466 let _ = max_nc;
30467 return Ok(());
30468 }
30469 {
30470 let f = self.func("gdn_chunk_state_mma_vl");
30471 let cfg = LaunchConfig {
30472 grid_dim: (n_head as u32, NSPLIT, b as u32),
30473 block_dim: (256, 1, 1),
30474 shared_mem_bytes: 0,
30475 };
30476 let __s_lb = self.gpu.stream();
30477 let mut lb = __s_lb.launch_builder(&f);
30478 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
30479 unsafe {
30480 lb.launch(cfg)?;
30481 }
30482 }
30483 {
30484 let f = self.func("gdn_chunk_output_mma_vl");
30485 let cfg = LaunchConfig {
30486 grid_dim: (max_nc, n_head as u32, b as u32),
30487 block_dim: (256, 1, 1),
30488 shared_mem_bytes: 0,
30489 };
30490 let __s_lb = self.gpu.stream();
30491 let mut lb = __s_lb.launch_builder(&f);
30492 lb.arg(&v).arg(&hi).arg(&ci).arg(&scale).arg(&hki);
30493 unsafe {
30494 lb.launch(cfg)?;
30495 }
30496 }
30497 Ok(())
30498 }
30499 #[allow(clippy::too_many_arguments)] pub fn gdn_scan_chunked(
30501 &self,
30502 q: &CudaSlice<f32>,
30503 k: &CudaSlice<f32>,
30504 v: &CudaSlice<f32>,
30505 g: &CudaSlice<f32>,
30506 beta: &CudaSlice<f32>,
30507 kb16_pre: Option<&CudaSlice<u8>>,
30508 qb16_pre: Option<&CudaSlice<u8>>,
30509 state_in: &CudaSlice<f32>,
30510 state_out: &mut CudaSlice<f32>,
30511 o: &mut CudaSlice<f32>,
30512 n_head: usize,
30513 t: usize,
30514 scale: f32,
30515 c: usize,
30516 hk: usize,
30517 ) -> Result<(), Box<dyn std::error::Error>> {
30518 const D: usize = 128;
30519 const NSPLIT: u32 = 4;
30520 assert!(
30521 (1..=128).contains(&c),
30522 "gdn_scan_chunked: C must be in 1..=128"
30523 );
30524 let h = n_head;
30525 #[allow(clippy::manual_div_ceil)]
30526 let nc = (t + c - 1) / c;
30528 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
30529 let gdn_mma_pre = !portable_mma_gated()
30534 && c == 32
30535 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
30536 Ok("1") => true,
30537 Ok("0") => false,
30538 _ => gdn_mma_default_on(),
30539 };
30540 let mut wb16_pre: Option<CudaSlice<u8>> = if gdn_mma_pre {
30541 Some(self.alloc_u8_uninit(nc * h * c * D * 2)?)
30542 } else {
30543 None
30544 };
30545 let gdn_wgmma_pre = cfg!(memra_hopper_mma)
30550 && gdn_mma_pre
30551 && std::env::var("MEMRA_GDN_WGMMA").as_deref() != Ok("0");
30552 let nk = t * hk * D;
30553 let mut kb16_local: Option<CudaSlice<u8>> = None;
30554 if gdn_mma_pre && kb16_pre.is_none() {
30555 let mut kb = self.alloc_u8_uninit(nk * 2)?;
30556 let f = self.func("f32_to_bf16_bulk");
30557 let n2 = nk as i64;
30558 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
30559 let __s_b = self.gpu.stream();
30560 let mut b = __s_b.launch_builder(&f);
30561 b.arg(k).arg(&mut kb).arg(&n2);
30562 unsafe {
30563 b.launch(cfg2)?;
30564 }
30565 kb16_local = Some(kb);
30566 }
30567 let kb16_ref0: Option<&CudaSlice<u8>> = kb16_local.as_ref().or(kb16_pre);
30568 if let Some(kb) = kb16_pre {
30569 assert!(kb.len() >= nk * 2, "kb16_pre too small");
30570 }
30571 let mut qb16: Option<CudaSlice<u8>> = None;
30572 let mut pb16: Option<CudaSlice<u8>> = None;
30573 if gdn_wgmma_pre {
30574 if qb16_pre.is_none() {
30577 let mut qb = self.alloc_u8_uninit(nk * 2)?;
30578 let f = self.func("f32_to_bf16_bulk");
30579 let n2 = nk as i64;
30580 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
30581 let __s_b = self.gpu.stream();
30582 let mut b = __s_b.launch_builder(&f);
30583 b.arg(q).arg(&mut qb).arg(&n2);
30584 unsafe {
30585 b.launch(cfg2)?;
30586 }
30587 qb16 = Some(qb);
30588 } else if let Some(qb) = qb16_pre {
30589 assert!(qb.len() >= nk * 2, "qb16_pre too small");
30590 }
30591 pb16 = Some(self.alloc_u8_uninit(nc * h * c * c * 2)?);
30592 }
30593 let qb16_ref0: Option<&CudaSlice<u8>> = qb16.as_ref().or(qb16_pre);
30594 let k2w = if gdn_wgmma_pre {
30595 Some((
30596 *qb16_ref0.as_ref().unwrap(),
30597 *kb16_ref0.as_ref().unwrap(),
30598 pb16.as_mut().unwrap(),
30599 ))
30600 } else {
30601 None
30602 };
30603 let (gcum, p, u, w) =
30604 self.gdn_chunk_k123(q, k, v, g, beta, wb16_pre.as_mut(), n_head, t, c, hk, k2w)?;
30605 let _ = &w;
30606 let mut y = self.uninit(nc * h * c * D)?;
30607 let mut ssnap = self.uninit(nc * h * D * D)?; let gdn_mma = !portable_mma_gated()
30621 && c == 32
30622 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
30623 Ok("1") => true,
30624 Ok("0") => false,
30625 _ => gdn_mma_default_on(),
30626 };
30627 if gdn_mma {
30628 let wb16 = wb16_pre
30629 .take()
30630 .expect("mma path pre-allocates wb16 (K3 store fold)");
30631 let kb16_ref: &CudaSlice<u8> = kb16_ref0.expect("mma path pre-builds kb16 above K123");
30632 if gdn_wgmma_pre {
30644 let qb16 = qb16_ref0.unwrap();
30646 let pb16 = pb16.as_ref().unwrap();
30647 {
30648 let f = self.func("gdn_k45_wgmma");
30649 let cfg = LaunchConfig {
30650 grid_dim: (h as u32, 4, 1),
30651 block_dim: (256, 1, 1),
30652 shared_mem_bytes: 0,
30653 };
30654 let hki = hk as i32;
30655 let __s_b = self.gpu.stream();
30656 let mut b = __s_b.launch_builder(&f);
30657 b.arg(kb16_ref)
30658 .arg(&gcum)
30659 .arg(beta)
30660 .arg(&u)
30661 .arg(&wb16)
30662 .arg(qb16)
30663 .arg(pb16)
30664 .arg(o)
30665 .arg(&scale)
30666 .arg(state_in)
30667 .arg(&mut *state_out)
30668 .arg(&hi)
30669 .arg(&ti)
30670 .arg(&ci)
30671 .arg(&hki);
30672 unsafe {
30673 b.launch(cfg)?;
30674 }
30675 }
30676 return Ok(());
30677 }
30678 let mut y16 = self.alloc_u8_uninit(nc * h * c * D * 2)?;
30682 let mut ssnap16 = self.alloc_u8_uninit(nc * h * D * D * 2)?;
30683 {
30684 let f = self.func("gdn_chunk_state_mma");
30685 let cfg = LaunchConfig {
30686 grid_dim: (h as u32, NSPLIT, 1),
30687 block_dim: (256, 1, 1),
30688 shared_mem_bytes: 0,
30689 };
30690 let hki = hk as i32;
30691 let __s_b = self.gpu.stream();
30692 let mut b = __s_b.launch_builder(&f);
30693 b.arg(kb16_ref)
30694 .arg(&gcum)
30695 .arg(beta)
30696 .arg(&u)
30697 .arg(&wb16)
30698 .arg(&mut y16)
30699 .arg(&mut ssnap16)
30700 .arg(state_in)
30701 .arg(&mut *state_out)
30702 .arg(&hi)
30703 .arg(&ti)
30704 .arg(&ci)
30705 .arg(&hki);
30706 unsafe {
30707 b.launch(cfg)?;
30708 }
30709 }
30710 {
30711 let f = self.func("gdn_chunk_output_mma");
30713 #[allow(clippy::manual_div_ceil)]
30714 let jt = ((c + 31) / 32) as u32;
30716 let cfg = LaunchConfig {
30717 grid_dim: (nc as u32, h as u32, jt),
30718 block_dim: (256, 1, 1),
30719 shared_mem_bytes: 0,
30720 };
30721 let hki = hk as i32;
30722 let __s_b = self.gpu.stream();
30723 let mut b = __s_b.launch_builder(&f);
30724 b.arg(q)
30725 .arg(&gcum)
30726 .arg(&p)
30727 .arg(&y16)
30728 .arg(&ssnap16)
30729 .arg(o)
30730 .arg(&hi)
30731 .arg(&ti)
30732 .arg(&ci)
30733 .arg(&scale)
30734 .arg(&hki);
30735 unsafe {
30736 b.launch(cfg)?;
30737 }
30738 }
30739 return Ok(());
30740 }
30741 {
30742 let f = self.func("gdn_chunk_state_f32");
30744 let cfg = LaunchConfig {
30745 grid_dim: (h as u32, NSPLIT, 1),
30746 block_dim: (256, 1, 1),
30747 shared_mem_bytes: 0,
30748 };
30749 let __s_b = self.gpu.stream();
30750 let mut b = __s_b.launch_builder(&f);
30751 b.arg(k)
30752 .arg(&gcum)
30753 .arg(beta)
30754 .arg(&u)
30755 .arg(&w)
30756 .arg(&mut y)
30757 .arg(&mut ssnap)
30758 .arg(state_in)
30759 .arg(&mut *state_out)
30760 .arg(&hi)
30761 .arg(&ti)
30762 .arg(&ci);
30763 unsafe {
30764 b.launch(cfg)?;
30765 }
30766 }
30767 {
30768 let f = self.func("gdn_chunk_output_f32");
30770 #[allow(clippy::manual_div_ceil)]
30771 let jt = ((c + 31) / 32) as u32;
30773 let cfg = LaunchConfig {
30774 grid_dim: (nc as u32, h as u32, jt),
30775 block_dim: (256, 1, 1),
30776 shared_mem_bytes: 0,
30777 };
30778 let __s_b = self.gpu.stream();
30779 let mut b = __s_b.launch_builder(&f);
30780 b.arg(q)
30781 .arg(&gcum)
30782 .arg(&p)
30783 .arg(&y)
30784 .arg(&ssnap)
30785 .arg(o)
30786 .arg(&hi)
30787 .arg(&ti)
30788 .arg(&ci)
30789 .arg(&scale);
30790 unsafe {
30791 b.launch(cfg)?;
30792 }
30793 }
30794 Ok(())
30795 }
30796
30797 #[allow(clippy::too_many_arguments)]
30806 #[allow(clippy::too_many_arguments)]
30807 pub fn gdn_scan_prefill(
30808 &self,
30809 q: &CudaSlice<f32>,
30810 k: &CudaSlice<f32>,
30811 v: &CudaSlice<f32>,
30812 g: &CudaSlice<f32>,
30813 beta: &CudaSlice<f32>,
30814 kb16_pre: Option<&CudaSlice<u8>>,
30815 qb16_pre: Option<&CudaSlice<u8>>,
30816 state_in: &CudaSlice<f32>,
30817 state_out: &mut CudaSlice<f32>,
30818 o: &mut CudaSlice<f32>,
30819 n_head: usize,
30820 t: usize,
30821 scale: f32,
30822 hk: usize,
30823 ) -> Result<(), Box<dyn std::error::Error>> {
30824 if std::env::var("MEMRA_GDN_DIFF").is_ok() && t >= 16 {
30825 assert!(hk == n_head, "GDN_DIFF oracle is broadcast-only");
30826 return self.gdn_scan_diff(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale);
30827 }
30828 if Self::gdn_chunked_enabled() && t >= 16 {
30829 self.gdn_scan_chunked(
30830 q,
30831 k,
30832 v,
30833 g,
30834 beta,
30835 kb16_pre,
30836 qb16_pre,
30837 state_in,
30838 state_out,
30839 o,
30840 n_head,
30841 t,
30842 scale,
30843 Self::gdn_chunk_size(),
30844 hk,
30845 )
30846 } else {
30847 assert!(
30848 hk == n_head,
30849 "s128 scan is broadcast-only (prep guarantees by predicate)"
30850 );
30851 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)
30852 }
30853 }
30854
30855 #[allow(clippy::too_many_arguments)]
30857 fn gdn_scan_diff(
30858 &self,
30859 q: &CudaSlice<f32>,
30860 k: &CudaSlice<f32>,
30861 v: &CudaSlice<f32>,
30862 g: &CudaSlice<f32>,
30863 beta: &CudaSlice<f32>,
30864 state_in: &CudaSlice<f32>,
30865 state_out: &mut CudaSlice<f32>,
30866 o: &mut CudaSlice<f32>,
30867 n_head: usize,
30868 t: usize,
30869 scale: f32,
30870 ) -> Result<(), Box<dyn std::error::Error>> {
30871 static CALL: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
30872 let call = CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
30873 let mut o_c = self.uninit(o.len())?;
30874 let mut st_c = self.uninit(state_out.len())?;
30875 self.gdn_scan_chunked(
30876 q,
30877 k,
30878 v,
30879 g,
30880 beta,
30881 None,
30882 None,
30883 state_in,
30884 &mut st_c,
30885 &mut o_c,
30886 n_head,
30887 t,
30888 scale,
30889 Self::gdn_chunk_size(),
30890 n_head,
30891 )?;
30892 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)?;
30893 let (oh_s, oh_c) = (self.dtoh(o)?, self.dtoh(&o_c)?);
30894 let (sh_s, sh_c) = (self.dtoh(state_out)?, self.dtoh(&st_c)?);
30895 let stats = |a: &[f32], b: &[f32]| -> (f32, f32, f64) {
30896 let mut max_abs = 0f32;
30897 let mut max_rel = 0f32;
30898 let mut sum_rel = 0f64;
30899 for (x, y) in a.iter().zip(b) {
30900 let ad = (x - y).abs();
30901 let rel = ad / x.abs().max(y.abs()).max(1e-3);
30902 if ad > max_abs {
30903 max_abs = ad;
30904 }
30905 if rel > max_rel {
30906 max_rel = rel;
30907 }
30908 sum_rel += rel as f64;
30909 }
30910 (max_abs, max_rel, sum_rel / a.len() as f64)
30911 };
30912 let (o_ma, o_mr, o_mean) = stats(&oh_s, &oh_c);
30913 let (s_ma, s_mr, s_mean) = stats(&sh_s, &sh_c);
30914 println!(
30915 "[gdn-diff call {call:3} T={t} C={}] out: max_abs={o_ma:.3e} max_rel={o_mr:.3e} mean_rel={o_mean:.3e} | \
30916 state: max_abs={s_ma:.3e} max_rel={s_mr:.3e} mean_rel={s_mean:.3e}",
30917 Self::gdn_chunk_size()
30918 );
30919 Ok(())
30920 }
30921
30922 pub fn gdn_glog(
30924 &self,
30925 alpha: &CudaSlice<f32>,
30926 dt_bias: &CudaSlice<f32>,
30927 a: &CudaSlice<f32>,
30928 g_log: &mut CudaSlice<f32>,
30929 n_head: usize,
30930 t: usize,
30931 ) -> Result<(), Box<dyn std::error::Error>> {
30932 let f = self.func("gdn_glog_f32");
30933 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
30934 let (h, ti) = (n_head as i32, t as i32);
30935 let __s_b = self.gpu.stream();
30936 let mut b = __s_b.launch_builder(&f);
30937 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
30938 unsafe {
30939 b.launch(cfg)?;
30940 }
30941 Ok(())
30942 }
30943
30944 pub fn sigmoid_v(
30947 &self,
30948 x: &cudarc::driver::CudaView<f32>,
30949 y: &mut CudaSlice<f32>,
30950 n: usize,
30951 ) -> Result<(), Box<dyn std::error::Error>> {
30952 let f = self.func("sigmoid_f32");
30953 let cfg = LaunchConfig::for_num_elems(n as u32);
30954 let ni = n as i32;
30955 let __s_b = self.gpu.stream();
30956 let mut b = __s_b.launch_builder(&f);
30957 b.arg(x).arg(y).arg(&ni);
30958 unsafe {
30959 b.launch(cfg)?;
30960 }
30961 Ok(())
30962 }
30963
30964 pub fn gdn_glog_v(
30965 &self,
30966 alpha: &cudarc::driver::CudaView<f32>,
30967 dt_bias: &CudaSlice<f32>,
30968 a: &CudaSlice<f32>,
30969 g_log: &mut CudaSlice<f32>,
30970 n_head: usize,
30971 t: usize,
30972 ) -> Result<(), Box<dyn std::error::Error>> {
30973 let f = self.func("gdn_glog_f32");
30974 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
30975 let (h, ti) = (n_head as i32, t as i32);
30976 let __s_b = self.gpu.stream();
30977 let mut b = __s_b.launch_builder(&f);
30978 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
30979 unsafe {
30980 b.launch(cfg)?;
30981 }
30982 Ok(())
30983 }
30984
30985 pub fn sigmoid(
30986 &self,
30987 x: &CudaSlice<f32>,
30988 y: &mut CudaSlice<f32>,
30989 n: usize,
30990 ) -> Result<(), Box<dyn std::error::Error>> {
30991 let f = self.func("sigmoid_f32");
30992 let cfg = LaunchConfig::for_num_elems(n as u32);
30993 let ni = n as i32;
30994 let __s_b = self.gpu.stream();
30995 let mut b = __s_b.launch_builder(&f);
30996 b.arg(x).arg(y).arg(&ni);
30997 unsafe {
30998 b.launch(cfg)?;
30999 }
31000 Ok(())
31001 }
31002
31003 pub fn sig_mul_f16out(
31006 &self,
31007 a: &CudaSlice<f32>,
31008 g: &CudaSlice<f32>,
31009 dst: &mut CudaSlice<f32>,
31010 dst16: &mut CudaSlice<u8>,
31011 n: usize,
31012 ) -> Result<(), Box<dyn std::error::Error>> {
31013 let f = self.func("sig_mul_f16out_f32");
31014 let cfg = LaunchConfig::for_num_elems(n as u32);
31015 let ni = n as i32;
31016 let __s_b = self.gpu.stream();
31017 let mut b = __s_b.launch_builder(&f);
31018 b.arg(a).arg(g).arg(dst).arg(dst16).arg(&ni);
31019 unsafe {
31020 b.launch(cfg)?;
31021 }
31022 Ok(())
31023 }
31024
31025 #[allow(clippy::too_many_arguments)]
31034 pub fn attn_head_gate(
31035 &self,
31036 a: &CudaSlice<f32>,
31037 g: &CudaSlice<f32>,
31038 dst: &mut CudaSlice<f32>,
31039 dst16: Option<&mut CudaSlice<u8>>,
31040 head_dim: usize,
31041 n_head: usize,
31042 t: usize,
31043 ) -> Result<(), Box<dyn std::error::Error>> {
31044 let f = self.func("attn_head_gate_f32");
31045 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
31046 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
31047 let d16: u64 = match dst16 {
31049 Some(d) => self.addr_u8(d),
31050 None => 0,
31051 };
31052 let __s_b = self.gpu.stream();
31053 let mut b = __s_b.launch_builder(&f);
31054 b.arg(a)
31055 .arg(g)
31056 .arg(dst)
31057 .arg(&d16)
31058 .arg(&hd)
31059 .arg(&nh)
31060 .arg(&ti);
31061 unsafe {
31062 b.launch(cfg)?;
31063 }
31064 Ok(())
31065 }
31066
31067 #[allow(clippy::too_many_arguments)]
31076 pub fn swiglu_clamped_mul_scaled(
31077 &self,
31078 gate: &CudaSlice<f32>,
31079 up: &CudaSlice<f32>,
31080 gs: f32,
31081 us: f32,
31082 limit: f32,
31083 dst: &mut CudaSlice<f32>,
31084 n: usize,
31085 ) -> Result<(), Box<dyn std::error::Error>> {
31086 debug_assert!(
31087 limit > 1e-6,
31088 "swiglu_clamped needs a live limit; use silu_mul_scaled"
31089 );
31090 let f = self.func("swiglu_clamped_mul_scaled_f32");
31091 let cfg = LaunchConfig::for_num_elems(n as u32);
31092 let ni = n as i32;
31093 let __s_b = self.gpu.stream();
31094 let mut b = __s_b.launch_builder(&f);
31095 b.arg(gate)
31096 .arg(up)
31097 .arg(&gs)
31098 .arg(&us)
31099 .arg(&limit)
31100 .arg(dst)
31101 .arg(&ni);
31102 unsafe {
31103 b.launch(cfg)?;
31104 }
31105 Ok(())
31106 }
31107
31108 #[allow(clippy::too_many_arguments)]
31117 pub fn swiglu_preclamped_mul_scaled(
31118 &self,
31119 gate: &CudaSlice<f32>,
31120 up: &CudaSlice<f32>,
31121 gs: f32,
31122 us: f32,
31123 limit: f32,
31124 dst: &mut CudaSlice<f32>,
31125 n: usize,
31126 ) -> Result<(), Box<dyn std::error::Error>> {
31127 debug_assert!(
31128 limit > 1e-6,
31129 "swiglu_preclamped needs a live limit; use silu_mul_scaled"
31130 );
31131 let f = self.func("swiglu_preclamped_mul_scaled_f32");
31132 let cfg = LaunchConfig::for_num_elems(n as u32);
31133 let ni = n as i32;
31134 let __s_b = self.gpu.stream();
31135 let mut b = __s_b.launch_builder(&f);
31136 b.arg(gate)
31137 .arg(up)
31138 .arg(&gs)
31139 .arg(&us)
31140 .arg(&limit)
31141 .arg(dst)
31142 .arg(&ni);
31143 unsafe {
31144 b.launch(cfg)?;
31145 }
31146 Ok(())
31147 }
31148
31149 #[allow(clippy::too_many_arguments)] pub fn gated_rmsnorm(
31152 &self,
31153 o: &CudaSlice<f32>,
31154 w: &CudaSlice<f32>,
31155 z: &CudaSlice<f32>,
31156 dst: &mut CudaSlice<f32>,
31157 ncols: usize,
31158 nrows: usize,
31159 eps: f32,
31160 ) -> Result<(), Box<dyn std::error::Error>> {
31161 let f = self.func("gated_rmsnorm_f32");
31162 let cfg = LaunchConfig {
31163 grid_dim: (nrows as u32, 1, 1),
31164 block_dim: (128, 1, 1),
31165 shared_mem_bytes: 0,
31166 };
31167 let (nc, e) = (ncols as i32, eps);
31168 let __s_b = self.gpu.stream();
31169 let mut b = __s_b.launch_builder(&f);
31170 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
31171 unsafe {
31172 b.launch(cfg)?;
31173 }
31174 Ok(())
31175 }
31176
31177 #[allow(clippy::too_many_arguments)] pub fn gated_rmsnorm_f16out(
31181 &self,
31182 o: &CudaSlice<f32>,
31183 w: &CudaSlice<f32>,
31184 z: &CudaSlice<f32>,
31185 dst: &mut CudaSlice<f32>,
31186 dst16: &mut CudaSlice<u8>,
31187 ncols: usize,
31188 nrows: usize,
31189 eps: f32,
31190 ) -> Result<(), Box<dyn std::error::Error>> {
31191 let f = self.func("gated_rmsnorm_f16out_f32");
31192 let cfg = LaunchConfig {
31194 grid_dim: (nrows as u32, 1, 1),
31195 block_dim: (128, 1, 1),
31196 shared_mem_bytes: 0,
31197 };
31198 let (nc, e) = (ncols as i32, eps);
31199 let __s_b = self.gpu.stream();
31200 let mut b = __s_b.launch_builder(&f);
31201 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
31202 unsafe {
31203 b.launch(cfg)?;
31204 }
31205 Ok(())
31206 }
31207
31208 #[allow(clippy::too_many_arguments)]
31212 pub fn add_rms_norm_zq8(
31213 &self,
31214 a: &CudaSlice<f32>,
31215 b_in: &CudaSlice<f32>,
31216 w: &CudaSlice<f32>,
31217 res: &mut CudaSlice<f32>,
31218 z: &mut CudaSlice<f32>,
31219 ncols: usize,
31220 nrows: usize,
31221 eps: f32,
31222 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
31223 assert!(ncols.is_multiple_of(32));
31224 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
31225 let mut d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
31226 let f = self.func("add_rms_norm_zq8");
31227 let cfg = LaunchConfig {
31228 grid_dim: (nrows as u32, 1, 1),
31229 block_dim: (1024, 1, 1),
31230 shared_mem_bytes: 0,
31231 };
31232 let (nc, ep) = (ncols as i32, eps);
31233 let __s_b = self.gpu.stream();
31234 let mut b = __s_b.launch_builder(&f);
31235 b.arg(a)
31236 .arg(b_in)
31237 .arg(w)
31238 .arg(res)
31239 .arg(z)
31240 .arg(&mut q)
31241 .arg(&mut d)
31242 .arg(&nc)
31243 .arg(&ep);
31244 unsafe {
31245 b.launch(cfg)?;
31246 }
31247 Ok((q, d))
31248 }
31249
31250 #[allow(clippy::too_many_arguments)] pub fn gated_rmsnorm_zv(
31256 &self,
31257 o: &CudaSlice<f32>,
31258 w: &CudaSlice<f32>,
31259 z: &cudarc::driver::CudaView<f32>,
31260 dst: &mut CudaSlice<f32>,
31261 ncols: usize,
31262 nrows: usize,
31263 eps: f32,
31264 ) -> Result<(), Box<dyn std::error::Error>> {
31265 let f = self.func("gated_rmsnorm_f32");
31266 let cfg = LaunchConfig {
31267 grid_dim: (nrows as u32, 1, 1),
31268 block_dim: (128, 1, 1),
31269 shared_mem_bytes: 0,
31270 };
31271 let (nc, e) = (ncols as i32, eps);
31272 let __s_b = self.gpu.stream();
31273 let mut b = __s_b.launch_builder(&f);
31274 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
31275 unsafe {
31276 b.launch(cfg)?;
31277 }
31278 Ok(())
31279 }
31280
31281 #[allow(clippy::too_many_arguments)] pub fn gated_rmsnorm_f16out_zv(
31283 &self,
31284 o: &CudaSlice<f32>,
31285 w: &CudaSlice<f32>,
31286 z: &cudarc::driver::CudaView<f32>,
31287 dst: &mut CudaSlice<f32>,
31288 dst16: &mut CudaSlice<u8>,
31289 ncols: usize,
31290 nrows: usize,
31291 eps: f32,
31292 ) -> Result<(), Box<dyn std::error::Error>> {
31293 let f = self.func("gated_rmsnorm_f16out_f32");
31294 let cfg = LaunchConfig {
31296 grid_dim: (nrows as u32, 1, 1),
31297 block_dim: (128, 1, 1),
31298 shared_mem_bytes: 0,
31299 };
31300 let (nc, e) = (ncols as i32, eps);
31301 let __s_b = self.gpu.stream();
31302 let mut b = __s_b.launch_builder(&f);
31303 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
31304 unsafe {
31305 b.launch(cfg)?;
31306 }
31307 Ok(())
31308 }
31309
31310 pub fn gated_rmsnorm_q8_1(
31311 &self,
31312 o: &CudaSlice<f32>,
31313 w: &CudaSlice<f32>,
31314 z: &CudaSlice<f32>,
31315 ncols: usize,
31316 nrows: usize,
31317 eps: f32,
31318 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
31319 assert!(ncols.is_multiple_of(32));
31320 let f = self.func("gated_rmsnorm_q8_1");
31321 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
31322 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
31323 let cfg = LaunchConfig {
31324 grid_dim: (nrows as u32, 1, 1),
31325 block_dim: (128, 1, 1),
31326 shared_mem_bytes: 0,
31327 };
31328 let (nc, ep) = (ncols as i32, eps);
31329 let __s_b = self.gpu.stream();
31330 let mut b = __s_b.launch_builder(&f);
31331 b.arg(o)
31332 .arg(w)
31333 .arg(z)
31334 .arg(&mut out_q)
31335 .arg(&mut out_d)
31336 .arg(&nc)
31337 .arg(&ep);
31338 unsafe {
31339 b.launch(cfg)?;
31340 }
31341 Ok((out_q, out_d))
31342 }
31343
31344 pub fn transpose(
31346 &self,
31347 inp: &CudaSlice<f32>,
31348 rows: usize,
31349 cols: usize,
31350 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31351 let f = self.func("transpose_f32");
31352 let mut out = self.zeros(rows * cols)?;
31353 let cfg = LaunchConfig::for_num_elems((rows * cols) as u32);
31354 let (r, c) = (rows as i32, cols as i32);
31355 let __s_b = self.gpu.stream();
31356 let mut b = __s_b.launch_builder(&f);
31357 b.arg(inp).arg(&mut out).arg(&r).arg(&c);
31358 unsafe {
31359 b.launch(cfg)?;
31360 }
31361 Ok(out)
31362 }
31363
31364 pub fn repeat_heads(
31366 &self,
31367 inp: &CudaSlice<f32>,
31368 out: &mut CudaSlice<f32>,
31369 head_dim: usize,
31370 n_in: usize,
31371 n_out: usize,
31372 t: usize,
31373 ) -> Result<(), Box<dyn std::error::Error>> {
31374 let f = self.func("repeat_heads_f32");
31375 let cfg = LaunchConfig::for_num_elems((head_dim * n_out * t) as u32);
31376 let (hd, ni, no, ti) = (head_dim as i32, n_in as i32, n_out as i32, t as i32);
31377 let __s_b = self.gpu.stream();
31378 let mut b = __s_b.launch_builder(&f);
31379 b.arg(inp).arg(out).arg(&hd).arg(&ni).arg(&no).arg(&ti);
31380 unsafe {
31381 b.launch(cfg)?;
31382 }
31383 Ok(())
31384 }
31385
31386 pub fn q_gate_split(
31393 &self,
31394 qf: &CudaSlice<f32>,
31395 q_out: &mut CudaSlice<f32>,
31396 gate_out: &mut CudaSlice<f32>,
31397 head_dim: usize,
31398 n_head: usize,
31399 t: usize,
31400 ) -> Result<(), Box<dyn std::error::Error>> {
31401 memra_gguf::config::check_fused_q_gate_extent(qf.len(), head_dim, n_head, t)?;
31402 let out_need = head_dim * n_head * t;
31403 if q_out.len() < out_need || gate_out.len() < out_need {
31404 return Err(format!(
31405 "q_gate_split destinations too small: need {out_need} each, have q={} gate={}",
31406 q_out.len(),
31407 gate_out.len()
31408 )
31409 .into());
31410 }
31411 let f = self.func("q_gate_split_f32");
31412 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
31413 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
31414 let __s_b = self.gpu.stream();
31415 let mut b = __s_b.launch_builder(&f);
31416 b.arg(qf)
31417 .arg(q_out)
31418 .arg(gate_out)
31419 .arg(&hd)
31420 .arg(&nh)
31421 .arg(&ti);
31422 unsafe {
31423 b.launch(cfg)?;
31424 }
31425 Ok(())
31426 }
31427
31428 #[allow(clippy::too_many_arguments)] pub fn qkv_to_gdn_repack(
31433 &self,
31434 conv_out: &CudaSlice<f32>,
31435 q_g: &mut CudaSlice<f32>,
31436 k_g: &mut CudaSlice<f32>,
31437 v_g: &mut CudaSlice<f32>,
31438 d_state: usize,
31439 num_v: usize,
31440 num_k: usize,
31441 key_dim: usize,
31442 t: usize,
31443 ) -> Result<(), Box<dyn std::error::Error>> {
31444 let f = self.func("qkv_to_gdn_repack_f32");
31445 let cfg = LaunchConfig::for_num_elems((d_state * num_v * t) as u32);
31446 let (ds, nv, nk, kd, ti) = (
31447 d_state as i32,
31448 num_v as i32,
31449 num_k as i32,
31450 key_dim as i32,
31451 t as i32,
31452 );
31453 let __s_b = self.gpu.stream();
31454 let mut b = __s_b.launch_builder(&f);
31455 b.arg(conv_out)
31456 .arg(q_g)
31457 .arg(k_g)
31458 .arg(v_g)
31459 .arg(&ds)
31460 .arg(&nv)
31461 .arg(&nk)
31462 .arg(&kd)
31463 .arg(&ti);
31464 unsafe {
31465 b.launch(cfg)?;
31466 }
31467 Ok(())
31468 }
31469
31470 pub fn conv_left_pad(
31473 &self,
31474 src: &CudaSlice<f32>,
31475 dst: &mut CudaSlice<f32>,
31476 conv_dim: usize,
31477 t: usize,
31478 pad: usize,
31479 ) -> Result<(), Box<dyn std::error::Error>> {
31480 let f = self.func("conv_left_pad_f32");
31481 let cfg = LaunchConfig::for_num_elems((conv_dim * t) as u32);
31482 let (cd, ti, p) = (conv_dim as i32, t as i32, pad as i32);
31483 let __s_b = self.gpu.stream();
31484 let mut b = __s_b.launch_builder(&f);
31485 b.arg(src).arg(dst).arg(&cd).arg(&ti).arg(&p);
31486 unsafe {
31487 b.launch(cfg)?;
31488 }
31489 Ok(())
31490 }
31491
31492 pub fn conv_assemble_and_roll(
31496 &self,
31497 qkv_col: &CudaSlice<f32>,
31498 conv_state: &mut CudaSlice<f32>,
31499 conv_in: &mut CudaSlice<f32>,
31500 conv_dim: usize,
31501 pad: usize,
31502 ) -> Result<(), Box<dyn std::error::Error>> {
31503 let f = self.func("conv_assemble_and_roll_f32");
31504 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
31505 let (cd, p) = (conv_dim as i32, pad as i32);
31506 let __s_b = self.gpu.stream();
31507 let mut b = __s_b.launch_builder(&f);
31508 b.arg(qkv_col).arg(conv_state).arg(conv_in).arg(&cd).arg(&p);
31509 unsafe {
31510 b.launch(cfg)?;
31511 }
31512 Ok(())
31513 }
31514
31515 pub fn ssm_conv1d_fused_decode(
31521 &self,
31522 qkv_col: &CudaSlice<f32>,
31523 conv_state: &mut CudaSlice<f32>,
31524 w: &CudaSlice<f32>,
31525 conv_out: &mut CudaSlice<f32>,
31526 conv_dim: usize,
31527 d_conv: usize,
31528 ) -> Result<(), Box<dyn std::error::Error>> {
31529 let f = self.func("ssm_conv1d_fused_decode_f32");
31530 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
31531 let (cd, dc) = (conv_dim as i32, d_conv as i32);
31532 let __s_b = self.gpu.stream();
31533 let mut b = __s_b.launch_builder(&f);
31534 b.arg(qkv_col)
31535 .arg(conv_state)
31536 .arg(w)
31537 .arg(conv_out)
31538 .arg(&cd)
31539 .arg(&dc);
31540 unsafe {
31541 b.launch(cfg)?;
31542 }
31543 Ok(())
31544 }
31545
31546 pub fn slice_range(
31549 &self,
31550 src: &CudaSlice<f32>,
31551 start: usize,
31552 len: usize,
31553 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31554 let host = self.gpu.stream().clone_dtoh(src)?;
31555 self.gpu.stream().synchronize()?;
31556 self.htod(&host[start..start + len])
31557 }
31558}
31559
31560#[cfg(test)]
31561mod target_dispatch_tests {
31562 use super::legacy_quant_gemm_allowed;
31563
31564 #[test]
31565 fn legacy_quant_gemm_arch_policy_honors_the_escape_hatch() {
31566 assert!(legacy_quant_gemm_allowed(false, false, false));
31568 assert!(!legacy_quant_gemm_allowed(false, false, true));
31569 assert!(!legacy_quant_gemm_allowed(true, false, false));
31571 assert!(!legacy_quant_gemm_allowed(true, false, true));
31572 assert!(legacy_quant_gemm_allowed(true, true, false));
31574 assert!(!legacy_quant_gemm_allowed(true, true, true));
31575 }
31576
31577 #[cfg(all(memra_portable_cuda, not(memra_hopper_mma)))]
31578 #[test]
31579 fn portable_build_disables_legacy_quant_gemm_without_an_env_override() {
31580 assert!(!legacy_quant_gemm_allowed(
31581 cfg!(memra_portable_cuda),
31582 cfg!(memra_hopper_mma),
31583 false
31584 ));
31585 }
31586
31587 #[cfg(memra_hopper_mma)]
31588 #[test]
31589 fn hopper_mma_build_re_admits_legacy_quant_gemm() {
31590 assert!(legacy_quant_gemm_allowed(
31591 cfg!(memra_portable_cuda),
31592 cfg!(memra_hopper_mma),
31593 false
31594 ));
31595 assert!(super::portable_mma_gated() == false);
31596 }
31597}
31598
31599impl memra_kv::KvDev for Engine {
31602 fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31603 Engine::zeros(self, n)
31604 }
31605 fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31606 Engine::uninit(self, n)
31607 }
31608 fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
31609 Engine::alloc_u8(self, n)
31610 }
31611 fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
31612 Engine::htod_i32(self, v)
31613 }
31614 fn clone_dtod(
31615 &self,
31616 src: &CudaSlice<f32>,
31617 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
31618 Engine::clone_dtod(self, src)
31619 }
31620 fn copy_into(
31621 &self,
31622 dst: &mut CudaSlice<f32>,
31623 off: usize,
31624 src: &CudaSlice<f32>,
31625 len: usize,
31626 ) -> Result<(), Box<dyn std::error::Error>> {
31627 Engine::copy_into(self, dst, off, src, len)
31628 }
31629 fn copy_range_into(
31630 &self,
31631 dst: &mut CudaSlice<f32>,
31632 dst_off: usize,
31633 src: &CudaSlice<f32>,
31634 src_off: usize,
31635 len: usize,
31636 ) -> Result<(), Box<dyn std::error::Error>> {
31637 Engine::copy_range_into(self, dst, dst_off, src, src_off, len)
31638 }
31639 fn set_i32_one(
31640 &self,
31641 d: &mut CudaSlice<i32>,
31642 v: i32,
31643 ) -> Result<(), Box<dyn std::error::Error>> {
31644 Engine::set_i32_one(self, d, v)
31645 }
31646}
31647
31648#[cfg(test)]
31649mod fused_gate_bounds_tests {
31650 use super::*;
31651
31652 #[test]
31665 #[ignore = "requires a CUDA GPU"]
31666 fn q_gate_split_refuses_a_separate_gate_wq_instead_of_reading_past_it() {
31667 let e = Engine::new(0).unwrap();
31668 let (head_dim, n_head, t) = (8usize, 4usize, 2usize);
31669 let fused = 2 * head_dim * n_head * t;
31670 let out_n = head_dim * n_head * t;
31671
31672 let narrow = e.htod(&vec![1.0f32; out_n]).unwrap();
31674 let mut q = e.uninit(out_n).unwrap();
31675 let mut gate = e.uninit(out_n).unwrap();
31676 let err = e
31677 .q_gate_split(&narrow, &mut q, &mut gate, head_dim, n_head, t)
31678 .expect_err("half-width wq must be refused, not read past")
31679 .to_string();
31680 assert!(err.contains("NO fused gate"), "{err}");
31681 assert!(err.contains(&format!("{fused}")), "{err}");
31682
31683 let host: Vec<f32> = (0..fused).map(|i| i as f32).collect();
31686 let wide = e.htod(&host).unwrap();
31687 e.q_gate_split(&wide, &mut q, &mut gate, head_dim, n_head, t)
31688 .expect("full-width wq splits");
31689 let (qh, gh) = (e.dtoh(&q).unwrap(), e.dtoh(&gate).unwrap());
31690 for tok in 0..t {
31691 for hh in 0..n_head {
31692 for d in 0..head_dim {
31693 let base = tok * (n_head * 2 * head_dim) + hh * (2 * head_dim);
31694 let idx = tok * (n_head * head_dim) + hh * head_dim + d;
31695 assert_eq!(qh[idx], host[base + d], "q t{tok} h{hh} d{d}");
31696 assert_eq!(gh[idx], host[base + head_dim + d], "gate t{tok} h{hh} d{d}");
31697 }
31698 }
31699 }
31700
31701 let mut small = e.uninit(out_n - 1).unwrap();
31703 assert!(
31704 e.q_gate_split(&wide, &mut small, &mut gate, head_dim, n_head, t)
31705 .is_err()
31706 );
31707 }
31708}
31709
31710#[cfg(test)]
31714mod fused_rope_width_tests {
31715 use super::Engine;
31716
31717 #[test]
31720 fn full_width_is_accepted() {
31721 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope", 256, 256).is_ok());
31722 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_cat", 512, 512).is_ok());
31723 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_append", 128, 128).is_ok());
31724 }
31725
31726 #[test]
31739 fn gemma4_official_artifact_widths_pass() {
31740 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope", 512, 512).is_ok());
31741 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_append_dc", 256, 256).is_ok());
31742 }
31743
31744 #[test]
31747 fn partial_rotary_is_refused_with_the_geometry_named() {
31748 let err = Engine::full_width_rope_only("rms_norm_qkv_rope", 64, 256)
31750 .expect_err("partial rotary must refuse");
31751 let msg = err.to_string();
31752 assert!(msg.contains("PARTIAL ROTARY REFUSED"), "{msg}");
31753 assert!(msg.contains("n_rot 64"), "{msg}");
31754 assert!(msg.contains("head_dim 256"), "{msg}");
31755 assert!(
31756 msg.contains("64..256"),
31757 "names the band it would corrupt: {msg}"
31758 );
31759 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_append_dc", 64, 128).is_err());
31761 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope", 256, 128).is_err());
31763 }
31764}