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 model;
62pub mod sigrouter_contract;
63pub mod vision;
64pub mod vision_gemma;
65pub mod vision_pre;
66pub mod vision_step;
67pub mod cache {
70 pub use memra_kv::*;
71}
72pub mod decode;
73pub mod decode_batch;
74pub mod dflash;
75pub mod eagle;
76pub mod gemma_spec;
77pub mod graph_update;
78pub mod mla;
82pub mod moesd;
83pub mod parallel;
84pub mod plan_backend;
85pub mod pp;
86pub mod round_stream;
87pub mod spec;
88pub mod tp;
89pub use memra_sampling as sampler;
90
91pub fn moe_f16g_mode() -> u8 {
135 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
136 *M.get_or_init(|| match std::env::var("MEMRA_MOE_F16G").as_deref() {
137 Ok("0") => 0,
138 Ok("2") => 2,
139 Ok("3") => 3,
140 Ok(_) => 1,
141 Err(_) => 2,
144 })
145}
146pub fn moe_f16g_sk_params() -> (i32, i32) {
160 static P: std::sync::OnceLock<(i32, i32)> = std::sync::OnceLock::new();
161 *P.get_or_init(|| match std::env::var("MEMRA_F16G_SK").as_deref() {
162 Ok("0") => (-1, 0),
163 Ok("32") => (0, i32::MAX),
164 Ok("128") => (0, 1),
165 _ => {
166 let cross = std::env::var("MEMRA_F16G_SK_CROSS")
167 .ok()
168 .and_then(|v| v.parse().ok())
169 .unwrap_or(64);
170 (0, cross)
171 }
172 })
173}
174pub fn moe_f16g_direct_on(qtype: i32) -> bool {
185 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
186 let m = *M.get_or_init(|| match std::env::var("MEMRA_F16G_DIRECT").as_deref() {
187 Ok("0") => 0,
188 Ok("kq") => 1,
189 _ => 2,
190 });
191 match m {
192 0 => false,
193 1 => qtype == QT_Q4_K || qtype == QT_Q6_K,
194 _ => true,
195 }
196}
197pub fn moe_f16g_tail_on() -> bool {
206 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
207 *ON.get_or_init(|| std::env::var("MEMRA_F16G_TAIL").as_deref() != Ok("0"))
208}
209
210pub fn moe_f16g_gemma_on() -> bool {
217 static M: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
218 *M.get_or_init(|| !matches!(std::env::var("MEMRA_MOE_F16G").as_deref(), Ok("0") | Err(_)))
219}
220
221pub fn moe_fuse_actq_on() -> bool {
225 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
226 *ON.get_or_init(|| std::env::var("MEMRA_MOE_FUSE_ACTQ").as_deref() != Ok("0"))
227}
228
229pub fn router_prefill_exact_on() -> bool {
239 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
240 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_PREFILL_EXACT").as_deref() != Ok("0"))
241}
242
243pub fn router_kernel_on() -> bool {
244 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
245 *ON.get_or_init(|| {
246 let on = std::env::var("MEMRA_ROUTER_KERNEL").as_deref() != Ok("0");
247 if !on {
248 eprintln!("[memra] router kernel OFF (rollback: per-column cuBLAS gemv)");
249 }
250 on
251 })
252}
253
254pub const ROUTER_BATCH_MIN_T: usize = 8;
269pub fn router_batch_on() -> bool {
270 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
271 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_BATCH").as_deref() != Ok("0"))
272}
273mod cpu_experts;
274#[cfg(memra_cutlass)]
275pub mod cutlass_ffi;
276pub mod dsv4_ffi;
277pub mod dsv4_gpu;
278pub mod f16_ffi;
279pub mod fp8_ffi;
280pub mod mmq_ffi;
281pub mod moe_cache;
282pub mod prime_graph;
283pub mod spill;
284mod spill_pread;
285
286const FATBIN: &[u8] = include_bytes!(env!("MEMRA_ENGINE_FATBIN"));
293const HYBRID_FATBIN: &[u8] = include_bytes!(env!("MEMRA_HYBRID_FATBIN"));
294const QMATVEC_FATBIN: &[u8] = include_bytes!(env!("MEMRA_QMATVEC_FATBIN"));
295const FLASH_FATBIN: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN"));
296const GEMM_FATBIN: &[u8] = include_bytes!(env!("MEMRA_GEMM_FATBIN"));
297const ROUTER_FATBIN: &[u8] = include_bytes!(env!("MEMRA_ROUTER_FATBIN"));
298const SAMPLE_FATBIN: &[u8] = include_bytes!(env!("MEMRA_SAMPLE_FATBIN"));
300
301fn gemm_fatbin_bytes() -> std::borrow::Cow<'static, [u8]> {
307 assert!(
308 !(portable_mma_gated() && std::env::var_os("MEMRA_GEMM_FATBIN").is_some()),
309 "MEMRA_GEMM_FATBIN overrides are not allowed in the portable CUDA lane"
310 );
311 match std::env::var("MEMRA_GEMM_FATBIN") {
312 Ok(path) => std::borrow::Cow::Owned(
313 std::fs::read(&path).unwrap_or_else(|e| panic!("MEMRA_GEMM_FATBIN read {path}: {e}")),
314 ),
315 Err(_) => std::borrow::Cow::Borrowed(GEMM_FATBIN),
316 }
317}
318
319pub(crate) const fn portable_mma_gated() -> bool {
326 cfg!(memra_portable_cuda) && !cfg!(memra_hopper_mma)
327}
328
329#[track_caller]
346pub(crate) fn refuse_portable_force(var: &str, needs: &str) {
347 assert!(
348 !portable_mma_gated(),
349 "{var} forces a kernel path this build does not contain: it needs {needs}, and this is a \
350 portable-CUDA build (sm_89). Unset {var} — the default path serves this arch."
351 );
352}
353
354pub(crate) const fn gdn_mma_default_on() -> bool {
363 cfg!(memra_hopper_mma) || konst_eq(env!("MEMRA_BUILT_CUDA_ARCH"), "120a")
364}
365
366const fn konst_eq(a: &str, b: &str) -> bool {
368 let (a, b) = (a.as_bytes(), b.as_bytes());
369 if a.len() != b.len() {
370 return false;
371 }
372 let mut i = 0;
373 while i < a.len() {
374 if a[i] != b[i] {
375 return false;
376 }
377 i += 1;
378 }
379 true
380}
381
382const fn legacy_quant_gemm_allowed(portable_cuda: bool, hopper_mma: bool, no_gemm: bool) -> bool {
387 (!portable_cuda || hopper_mma) && !no_gemm
388}
389
390const FLASH_FATBIN_VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VQ4"));
398const FLASH_FATBIN_VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VF8"));
399const FLASH_FATBIN_KF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8"));
400const FLASH_FATBIN_KF8VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VQ4"));
401const FLASH_FATBIN_KF8VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VF8"));
402
403pub use memra_kv::{kv_blk_bytes, kv_cache_formats};
406
407fn flash_fatbin_bytes() -> &'static [u8] {
409 match kv_cache_formats() {
410 ("q8_0", "q5_1") => FLASH_FATBIN,
411 ("q8_0", "q4_0") => FLASH_FATBIN_VQ4,
412 ("q8_0", "fp8") => FLASH_FATBIN_VF8,
413 ("fp8", "q5_1") => FLASH_FATBIN_KF8,
414 ("fp8", "q4_0") => FLASH_FATBIN_KF8VQ4,
415 ("fp8", "fp8") => FLASH_FATBIN_KF8VF8,
416 other => unreachable!("kv_cache_formats returned {other:?}"),
417 }
418}
419
420fn k1_launch_override() -> Option<(u32, u32, u32)> {
427 static K1: std::sync::OnceLock<Option<(u32, u32, u32)>> = std::sync::OnceLock::new();
428 *K1.get_or_init(|| {
429 let v = std::env::var("MEMRA_GEMM_K1_LAUNCH").ok()?;
430 let p: Vec<u32> = v.split(',').filter_map(|s| s.trim().parse().ok()).collect();
431 match p.as_slice() {
432 [bm, bn, w] => Some((*bm, *bn, *w)),
433 _ => None,
434 }
435 })
436}
437
438pub(crate) fn wgmma_gemm_enabled() -> bool {
445 static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
446 *V.get_or_init(|| std::env::var("MEMRA_WGMMA").as_deref() == Ok("1"))
447}
448
449pub const FA_VEC_MIN_TKV: usize = 96;
464pub fn fa_vec_min_tkv() -> usize {
468 static V: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
469 *V.get_or_init(|| {
470 std::env::var("MEMRA_FA_VEC_MIN")
471 .ok()
472 .and_then(|v| v.parse().ok())
473 .unwrap_or_else(|| FA_VEC_MIN_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
474 })
475}
476
477pub fn fa_f16pv_on() -> bool {
488 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
489 *ON.get_or_init(|| {
490 std::env::var("MEMRA_FA_F16PV")
491 .map(|v| v != "0")
492 .unwrap_or_else(|_| std::env::var("MEMRA_DRAFT").is_err())
493 })
494}
495
496pub fn fa512_hp_on() -> bool {
500 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
501 *ON.get_or_init(|| std::env::var("MEMRA_FA512_HP").as_deref() != Ok("0"))
502}
503
504pub fn faw_hp_on() -> bool {
508 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
509 *ON.get_or_init(|| std::env::var("MEMRA_FAW_HP").as_deref() != Ok("0"))
510}
511
512pub fn fa512_wide_warps() -> usize {
516 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
517 *N.get_or_init(|| match std::env::var("MEMRA_FA512_W4").as_deref() {
518 Ok("1") => 4,
519 _ => 2,
520 })
521}
522
523pub fn fa512_min_tkv() -> usize {
526 static FA512_MIN: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
527 *FA512_MIN.get_or_init(|| {
528 std::env::var("MEMRA_FA512_MIN")
529 .ok()
530 .and_then(|v| v.parse().ok())
531 .unwrap_or(512)
532 })
533}
534pub static FA_VEC_MIN_DEFAULT: std::sync::atomic::AtomicUsize =
538 std::sync::atomic::AtomicUsize::new(FA_VEC_MIN_TKV);
539pub static FA_SPW_DEFAULT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(32);
543pub static FUSED_MR1_DEFAULT: std::sync::atomic::AtomicBool =
549 std::sync::atomic::AtomicBool::new(false);
550pub static ROUTER_W8_DEFAULT: std::sync::atomic::AtomicBool =
557 std::sync::atomic::AtomicBool::new(true);
558pub static FA_SP512_DEFAULT: std::sync::atomic::AtomicUsize =
559 std::sync::atomic::AtomicUsize::new(16);
560pub static RMS_BLOCK_DEFAULT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(256);
565pub static FA_SP_GEMMA: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
567pub static MMQ_SK_FORCE: std::sync::atomic::AtomicI8 = std::sync::atomic::AtomicI8::new(-1);
572pub use memra_kv::KV_FP8_FORCE;
575pub(crate) fn mmv_block() -> u32 {
580 static V: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
581 *V.get_or_init(|| {
582 std::env::var("MEMRA_MMV_BLOCK")
583 .ok()
584 .and_then(|v| v.parse().ok())
585 .filter(|&b: &u32| (64..=256).contains(&b) && b % 32 == 0)
586 .unwrap_or(128)
587 })
588}
589
590static STEP37_SERVING_DEFAULTS: std::sync::atomic::AtomicBool =
616 std::sync::atomic::AtomicBool::new(false);
617
618pub fn arm_step37_serving_defaults() {
619 STEP37_SERVING_DEFAULTS.store(true, std::sync::atomic::Ordering::Relaxed);
620 crate::cache::set_swa_ring_default(true);
621 eprintln!(
622 "[step37-defaults] serving doors armed ON for the SlidingGatedMoe program \
623 (per-flag =0 kills, =1 forces; owner flip 2026-08-27)"
624 );
625}
626
627pub(crate) fn step37_defaults_armed() -> bool {
628 STEP37_SERVING_DEFAULTS.load(std::sync::atomic::Ordering::Relaxed)
629}
630
631pub(crate) fn step37_door(cell: &'static std::sync::OnceLock<Option<bool>>, name: &str) -> bool {
635 match *cell.get_or_init(|| match std::env::var(name).ok().as_deref() {
636 Some("1") => Some(true),
637 Some("0") => Some(false),
638 _ => None,
639 }) {
640 Some(forced) => forced,
641 None => step37_defaults_armed(),
642 }
643}
644
645pub(crate) fn w8_hybrid_on() -> bool {
646 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
647 step37_door(&ENV, "MEMRA_W8_HYBRID")
648}
649
650pub(crate) fn step_tp_w8_on() -> bool {
651 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
652 step37_door(&ENV, "MEMRA_STEP_TP_W8")
653}
654
655pub(crate) fn w8_view_on() -> bool {
660 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
661 *ON.get_or_init(|| std::env::var("MEMRA_W8_VIEW").as_deref() == Ok("1"))
662}
663
664pub(crate) fn step_gemm_prime_on() -> bool {
678 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
679 step37_door(&ENV, "MEMRA_STEP_GEMM_PRIME")
680}
681
682pub(crate) fn q8t_wonce_on() -> bool {
683 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
684 step37_door(&ENV, "MEMRA_Q8T_WONCE")
685}
686
687pub(crate) fn sig_expf_dev_on() -> bool {
692 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
693 *ON.get_or_init(|| std::env::var("MEMRA_SIG_EXPF_DEV").as_deref() == Ok("1"))
694}
695
696pub(crate) fn topk_fast_on() -> bool {
697 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
698 *ON.get_or_init(|| std::env::var("MEMRA_TOPK_FAST").as_deref() == Ok("1"))
699}
700
701fn sigmoid_topk_kernel(sig_expf: bool, fast: bool, n_used: usize) -> &'static str {
706 match (sig_expf, fast && n_used <= 8) {
707 (true, true) => "moe_router_sigmoid_topk_f32_dexp_fast",
708 (true, false) => "moe_router_sigmoid_topk_f32_dexp",
709 (false, true) => "moe_router_sigmoid_topk_f32_fast",
710 (false, false) => "moe_router_sigmoid_topk_f32",
711 }
712}
713
714#[cfg(test)]
715mod sigmoid_topk_dispatch_tests {
716 #[test]
717 fn fast_kernel_refuses_wide_topk_and_composes_with_dexp() {
718 use super::sigmoid_topk_kernel;
719
720 assert_eq!(
721 sigmoid_topk_kernel(false, true, 8),
722 "moe_router_sigmoid_topk_f32_fast"
723 );
724 assert_eq!(
725 sigmoid_topk_kernel(true, true, 8),
726 "moe_router_sigmoid_topk_f32_dexp_fast"
727 );
728 assert_eq!(
729 sigmoid_topk_kernel(false, true, 9),
730 "moe_router_sigmoid_topk_f32"
731 );
732 assert_eq!(
733 sigmoid_topk_kernel(true, true, 9),
734 "moe_router_sigmoid_topk_f32_dexp"
735 );
736 }
737}
738
739pub(crate) fn rms_block() -> u32 {
740 static V: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
741 *V.get_or_init(|| {
742 std::env::var("MEMRA_RMS_BLOCK")
743 .ok()
744 .and_then(|v| v.parse().ok())
745 .unwrap_or_else(|| RMS_BLOCK_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
746 })
747}
748
749pub(crate) fn fa_split_keys(t_kv: usize, n_head_kv: usize) -> usize {
750 static S: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
751 if let Some(forced) = *S.get_or_init(|| {
752 std::env::var("MEMRA_FA_SPLIT")
753 .ok()
754 .and_then(|v| v.parse().ok())
755 .filter(|&s: &usize| s >= 8 && s % 8 == 0)
756 }) {
757 return forced;
758 }
759 if FA_SP_GEMMA.load(std::sync::atomic::Ordering::Relaxed)
777 && std::env::var("MEMRA_FA_SP16").as_deref() == Ok("1")
778 {
779 return if t_kv <= 8192 {
780 16
781 } else if t_kv <= 16384 {
782 64
783 } else {
784 128
785 };
786 }
787 let big_rig = fa_sm_count() >= 128;
788 if big_rig {
789 let _ = n_head_kv;
790 if t_kv <= 2048 {
791 static SHORT: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
800 if let Some(sp) = *SHORT.get_or_init(|| {
801 std::env::var("MEMRA_FA_SP_SHORT")
802 .ok()
803 .and_then(|v| v.parse().ok())
804 .filter(|&s: &usize| s >= 8 && s % 8 == 0)
805 }) {
806 return sp;
807 }
808 16
809 } else if t_kv <= 16384 {
810 64
811 } else {
812 128
813 }
814 } else if n_head_kv <= 4 {
815 if t_kv <= 512 {
836 8
837 } else if t_kv <= 16384 {
838 64
839 } else {
840 128
841 }
842 } else {
843 if t_kv <= 8192 {
844 32
845 } else if t_kv <= 16384 {
846 64
847 } else {
848 128
849 }
850 }
851}
852
853pub(crate) fn fa_sm_count() -> i32 {
856 static N: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
857 *N.get_or_init(|| {
858 cudarc::driver::result::init().ok();
859 cudarc::driver::result::device::get(0)
860 .and_then(|d| unsafe { cudarc::driver::result::device::get_attribute(
861 d, cudarc::driver::sys::CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT) })
862 .unwrap_or(82)
863 })
864}
865
866fn fa_hd_suffix(head_dim: usize) -> Result<&'static str, Box<dyn std::error::Error>> {
870 match head_dim {
871 256 => Ok(""),
872 128 => Ok("_hd128"),
873 d => Err(format!(
874 "fa_prefill: no kernel stamped for head_dim={d} (only 256/128); \
875 callers must gate to sdpa_naive"
876 )
877 .into()),
878 }
879}
880
881pub const QT_Q8_0: i32 = 0;
883pub const QT_Q4_K: i32 = 1;
884pub const QT_Q6_K: i32 = 2;
885pub const QT_Q5_K: i32 = 3;
886pub const QT_Q3_K: i32 = 4;
887pub const QT_IQ4_XS: i32 = 5;
888pub const QT_IQ3_S: i32 = 6;
889pub const QT_NVFP4: i32 = 7;
890pub const QT_NVFP4_V2: i32 = 107;
893pub const QT_F8_E4M3: i32 = 10;
899pub const QT_NVFP4_RP: i32 = 9;
902pub const QT_F32: i32 = 8;
904pub const QT_BF16: i32 = 11;
905pub const QT_Q4_0: i32 = 12; pub const QT_Q2_K: i32 = 13;
910pub const QT_F8_E4M3_BLK: i32 = 14;
926
927pub struct Engine {
929 pub gpu: memra_runtime::Gpu,
930 module: Arc<CudaModule>,
931 hybrid: Arc<CudaModule>,
932 qmatvec: Arc<CudaModule>,
933 flash: Arc<CudaModule>,
934 flash_g: std::sync::OnceLock<Arc<CudaModule>>,
938 gemm: Arc<CudaModule>,
939 router: Arc<CudaModule>,
940 sample: Arc<CudaModule>,
942 moe_cache: Mutex<Option<crate::moe_cache::MoeSlotCache>>,
946 w8_mirrors: Mutex<std::collections::HashMap<(u64, u32, u32), CudaSlice<u8>>>,
955 w8_act: Mutex<std::collections::HashMap<usize, (CudaSlice<i8>, CudaSlice<f32>)>>,
958 moe_cache_layout: Mutex<Option<Vec<usize>>>,
962 capture_keep_on: std::sync::atomic::AtomicBool,
968 verify_exact: std::sync::atomic::AtomicBool,
973 capture_keep: Mutex<Vec<Box<dyn std::any::Any + Send>>>,
974 pub copy_stream: Arc<CudaStream>,
976 #[cfg(memra_cutlass)]
983 cutlass_scratch: Mutex<Option<crate::cutlass_ffi::CutlassScratch>>,
984 fp8_scratch: Mutex<Option<crate::fp8_ffi::Fp8Scratch>>,
988 fa_vf16_scratch: Mutex<Option<CudaSlice<u8>>>,
991 fa_part_pool: Mutex<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
995 fa_part_retired: Mutex<Vec<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
999 fn_cache: Mutex<std::collections::HashMap<String, CudaFunction>>,
1001 f16_scratch: Mutex<Option<crate::f16_ffi::F16Scratch>>,
1002 argmax_partials: Mutex<Option<(CudaSlice<f32>, CudaSlice<i32>)>>,
1007 prime_deqw_ws: Mutex<Option<(CudaSlice<u8>, CudaSlice<u8>)>>,
1012 router_stage: Mutex<Option<PinnedStage>>,
1016}
1017
1018fn fa_v2_on() -> bool {
1028 std::env::var("MEMRA_FA_V2")
1034 .map(|v| v != "0")
1035 .unwrap_or(true)
1036}
1037
1038pub(crate) fn fa_part_zero_on() -> bool {
1048 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1049 *ON.get_or_init(|| std::env::var("MEMRA_FA_PART_ZERO").as_deref() == Ok("1"))
1050}
1051
1052pub(crate) fn fa_v3_on() -> bool {
1053 std::env::var("MEMRA_FA_V3")
1057 .map(|v| v != "0")
1058 .unwrap_or(true)
1059}
1060
1061fn fa_v4_mode() -> &'static str {
1066 static M: std::sync::OnceLock<String> = std::sync::OnceLock::new();
1067 M.get_or_init(|| std::env::var("MEMRA_FA_V4").unwrap_or_default())
1068}
1069fn fa_v4_on() -> bool {
1070 fa_v4_mode() != "0"
1071} pub static FA_SMEM_TKV_DEFAULT: std::sync::atomic::AtomicUsize =
1080 std::sync::atomic::AtomicUsize::new(1024);
1081pub static FA_V4_MAX_DEFAULT: std::sync::atomic::AtomicUsize =
1082 std::sync::atomic::AtomicUsize::new(usize::MAX);
1083pub fn fa_v4_at_pub(t_kv: usize) -> bool {
1084 fa_v4_at(t_kv)
1085}
1086fn fa_v4_at(t_kv: usize) -> bool {
1087 static M: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
1088 let mx = *M.get_or_init(|| {
1089 std::env::var("MEMRA_FA_V4_MAX")
1090 .ok()
1091 .and_then(|v| v.parse().ok())
1092 .unwrap_or_else(|| FA_V4_MAX_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
1093 });
1094 fa_v4_on() && t_kv < mx
1095}
1096pub const FA_DEEP_MIN_DEFAULT: usize = 0;
1110fn fa_deep_at(t_kv: usize) -> bool {
1111 if std::env::var("MEMRA_FA_DEEP").as_deref() == Ok("0") {
1112 return false;
1113 }
1114 let min = std::env::var("MEMRA_FA_DEEP_MIN")
1115 .ok()
1116 .and_then(|v| v.parse().ok())
1117 .unwrap_or(FA_DEEP_MIN_DEFAULT);
1118 t_kv >= min
1119}
1120pub fn fa_deep_at_pub(t_kv: usize) -> bool {
1122 fa_deep_at(t_kv)
1123}
1124
1125fn fa_v3_active(head_dim: usize) -> bool {
1126 fa_v3_on()
1129 && head_dim % 128 == 0
1130 && kv_cache_formats() == ("q8_0", "q5_1")
1131 && !Engine::kv_fp8_on()
1132}
1133
1134pub fn fa_seqs_eligible(t_kv: usize, head_dim: usize) -> bool {
1142 std::env::var("MEMRA_NO_FA_VEC").is_err()
1143 && t_kv >= fa_vec_min_tkv()
1144 && head_dim == 256
1145 && fa_v4_at(t_kv)
1146 && !matches!(fa_v4_mode(), "noB3" | "stage")
1147 && !Engine::kv_fp8_on()
1148}
1149pub fn fa_split_keys_pub(t_kv: usize, n_head_kv: usize) -> usize {
1151 fa_split_keys(t_kv, n_head_kv)
1152}
1153
1154struct PinnedStage {
1159 ptr: *mut u8,
1160 cap: usize,
1161}
1162unsafe impl Send for PinnedStage {}
1163impl PinnedStage {
1164 fn new(cap: usize) -> Result<Self, Box<dyn std::error::Error>> {
1165 let ptr = unsafe { cudarc::driver::result::malloc_host(cap, 0)? } as *mut u8;
1166 Ok(PinnedStage { ptr, cap })
1167 }
1168}
1169impl Drop for PinnedStage {
1170 fn drop(&mut self) {
1171 let _ = unsafe { cudarc::driver::result::free_host(self.ptr as _) };
1172 }
1173}
1174
1175pub const ARGMAX_NB: usize = 256;
1178
1179pub(crate) use memra_fa3_vl as fa3_vl_raw;
1181
1182unsafe extern "C" {
1183 fn memra_fa3_prefill(
1185 q16: *const core::ffi::c_void,
1186 k16: *const core::ffi::c_void,
1187 v16: *const core::ffi::c_void,
1188 o: *mut f32,
1189 t: i32,
1190 h: i32,
1191 hkv: i32,
1192 d: i32,
1193 scale: f32,
1194 stream: *mut core::ffi::c_void,
1195 ) -> i32;
1196 pub(crate) fn memra_fa3_vl(
1198 q16s: *const *const core::ffi::c_void,
1199 k16s: *const *const core::ffi::c_void,
1200 v16s: *const *const core::ffi::c_void,
1201 os: *const *mut f32,
1202 ts: *const i32,
1203 b: i32,
1204 h: i32,
1205 hkv: i32,
1206 d: i32,
1207 scale: f32,
1208 stream: *mut core::ffi::c_void,
1209 ) -> i32;
1210}
1211
1212#[repr(C)]
1217#[derive(Clone, Copy)]
1218pub struct WPtr8(pub [u64; 8]);
1219unsafe impl cudarc::driver::DeviceRepr for WPtr8 {}
1220
1221#[repr(C)]
1226#[derive(Clone, Copy, Default)]
1227pub struct GdnSeqVl {
1228 pub kb16: u64,
1229 pub gcum: u64,
1230 pub beta: u64,
1231 pub u: u64,
1232 pub wb16: u64,
1233 pub y: u64,
1234 pub ssnap: u64,
1235 pub state_in: u64,
1236 pub state_out: u64,
1237 pub q: u64,
1238 pub p: u64,
1239 pub o: u64,
1240 pub k: u64,
1241 pub v: u64,
1242 pub g: u64,
1243 pub a: u64,
1244 pub w: u64,
1245 pub t: i32,
1246 pub nc: i32,
1247}
1248unsafe impl cudarc::driver::DeviceRepr for GdnSeqVl {}
1249#[repr(C)]
1250#[derive(Clone, Copy)]
1251pub struct GdnVl8(pub [GdnSeqVl; 8]);
1252unsafe impl cudarc::driver::DeviceRepr for GdnVl8 {}
1253
1254#[repr(C)]
1257#[derive(Clone, Copy, Default)]
1258pub struct GdnWVl {
1259 pub qb16: u64,
1260 pub pb16: u64,
1261}
1262unsafe impl cudarc::driver::DeviceRepr for GdnWVl {}
1263#[repr(C)]
1264#[derive(Clone, Copy)]
1265pub struct GdnWVl8(pub [GdnWVl; 8]);
1266unsafe impl cudarc::driver::DeviceRepr for GdnWVl8 {}
1267
1268#[repr(C)]
1270#[derive(Clone, Copy, Default)]
1271pub struct GdnPrepVl {
1272 pub qkv: u64,
1273 pub conv_state: u64,
1274 pub conv_out: u64,
1275 pub q_g: u64,
1276 pub k_g: u64,
1277 pub v_g: u64,
1278 pub q_l2: u64,
1279 pub k_l2: u64,
1280 pub beta_raw: u64,
1281 pub alpha: u64,
1282 pub beta: u64,
1283 pub g_log: u64,
1284 pub o: u64,
1285 pub z: u64,
1286 pub gn: u64,
1287 pub gn16: u64,
1288 pub kb16: u64,
1289 pub qb16: u64,
1290 pub t: i32,
1291 pub pad: i32,
1292}
1293unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl {}
1294#[repr(C)]
1295#[derive(Clone, Copy)]
1296pub struct GdnPrepVl8(pub [GdnPrepVl; 8]);
1297unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl8 {}
1298
1299#[repr(C)]
1301#[derive(Clone, Copy, Default)]
1302pub struct FaSeqVl {
1303 pub q: u64,
1304 pub k16: u64,
1305 pub v16: u64,
1306 pub o: u64,
1307 pub kf: u64,
1308 pub vf: u64,
1309 pub t: i32,
1310 pub pad: i32,
1311}
1312unsafe impl cudarc::driver::DeviceRepr for FaSeqVl {}
1313#[repr(C)]
1314#[derive(Clone, Copy)]
1315pub struct FaVl8(pub [FaSeqVl; 8]);
1316unsafe impl cudarc::driver::DeviceRepr for FaVl8 {}
1317
1318#[repr(C)]
1320#[derive(Clone, Copy, Default)]
1321pub struct AttnPreVl {
1322 pub qf: u64,
1323 pub kf: u64,
1324 pub vf: u64,
1325 pub q: u64,
1326 pub gate: u64,
1327 pub qn: u64,
1328 pub kn: u64,
1329 pub kc: u64,
1330 pub vc: u64,
1331 pub t: i32,
1332 pub pad: i32,
1333}
1334unsafe impl cudarc::driver::DeviceRepr for AttnPreVl {}
1335#[repr(C)]
1336#[derive(Clone, Copy)]
1337pub struct AttnPreVl8(pub [AttnPreVl; 8]);
1338unsafe impl cudarc::driver::DeviceRepr for AttnPreVl8 {}
1339
1340pub struct GdnChunkBufs {
1343 pub gcum: CudaSlice<f32>,
1344 pub a: CudaSlice<f32>,
1345 pub p: CudaSlice<f32>,
1346 pub u: CudaSlice<f32>,
1347 pub w: CudaSlice<f32>,
1348 pub kb16: CudaSlice<u8>,
1349 pub wb16: CudaSlice<u8>,
1350 pub y16: CudaSlice<u8>,
1351 pub ssnap16: CudaSlice<u8>,
1352 pub qb16: CudaSlice<u8>,
1353 pub pb16: CudaSlice<u8>,
1354 pub o: CudaSlice<f32>,
1355 pub t: usize,
1356 pub nc: usize,
1357}
1358
1359#[repr(C)]
1361#[derive(Clone, Copy)]
1362pub struct F32x8(pub [f32; 8]);
1363unsafe impl cudarc::driver::DeviceRepr for F32x8 {}
1364
1365pub static PRIME_NANOS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
1369
1370#[must_use = "dropping immediately ends the exact scope"]
1375pub struct ExactScope<'a> {
1376 flag: &'a std::sync::atomic::AtomicBool,
1377 prev: bool,
1378}
1379
1380impl<'a> ExactScope<'a> {
1381 pub(crate) fn set(flag: &'a std::sync::atomic::AtomicBool, on: bool) -> Self {
1382 let prev = flag.load(std::sync::atomic::Ordering::Relaxed);
1383 flag.store(on, std::sync::atomic::Ordering::Relaxed);
1384 ExactScope { flag, prev }
1385 }
1386}
1387
1388impl Drop for ExactScope<'_> {
1389 fn drop(&mut self) {
1390 self.flag
1391 .store(self.prev, std::sync::atomic::Ordering::Relaxed);
1392 }
1393}
1394
1395#[cfg(test)]
1396mod exact_scope_tests {
1397 use std::sync::atomic::{AtomicBool, Ordering};
1398
1399 #[test]
1400 fn error_path_restores_verify_exact() {
1401 let flag = AtomicBool::new(false);
1406 let failing = |flag: &AtomicBool| -> Result<(), &'static str> {
1407 let _scope = super::ExactScope::set(flag, true);
1408 assert!(flag.load(Ordering::Relaxed), "scope arms the flag");
1409 Err("draft forward failed")? };
1411 assert!(failing(&flag).is_err());
1412 assert!(
1413 !flag.load(Ordering::Relaxed),
1414 "error propagation must restore the pre-scope value"
1415 );
1416 let flag = AtomicBool::new(true);
1418 {
1419 let _scope = super::ExactScope::set(&flag, true);
1420 }
1421 assert!(flag.load(Ordering::Relaxed));
1422 let flag = AtomicBool::new(false);
1424 let scope = super::ExactScope::set(&flag, true);
1425 drop(scope);
1426 assert!(!flag.load(Ordering::Relaxed));
1427 }
1428}
1429
1430impl Engine {
1431 pub fn new(ordinal: usize) -> Result<Self, Box<dyn std::error::Error>> {
1432 let gpu = memra_runtime::Gpu::new(ordinal)?;
1433 if std::env::var("MEMRA_ARCH_CHECK").as_deref() != Ok("0") {
1437 use cudarc::driver::sys::CUdevice_attribute_enum as A;
1438 let (maj, min) = cudarc::driver::result::device::get(ordinal as i32)
1439 .and_then(|d| unsafe {
1440 Ok((
1441 cudarc::driver::result::device::get_attribute(
1442 d,
1443 A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR,
1444 )?,
1445 cudarc::driver::result::device::get_attribute(
1446 d,
1447 A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR,
1448 )?,
1449 ))
1450 })
1451 .unwrap_or((0, 0));
1452 let built = env!("MEMRA_BUILT_CUDA_ARCH");
1453 let ok = matches!(
1454 (built, maj, min),
1455 ("120a", 12, 0) | ("120a", 12, 1) | ("100a", 10, 0) | ("90a", 9, 0) | ("89", 8, 9)
1456 );
1457 if !ok {
1458 return Err(format!(
1459 "memra was built for sm_{built} but device {ordinal} reports compute \
1460 capability {maj}.{min}. Rebuild on this machine (MEMRA_CUDA_ARCH \
1461 auto-detects the GPU) or set MEMRA_ARCH_CHECK=0 to bypass."
1462 )
1463 .into());
1464 }
1465 }
1466 unsafe {
1471 use cudarc::driver::sys;
1472 let dev: sys::CUdevice = ordinal as sys::CUdevice;
1473 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
1474 if sys::cuDeviceGetDefaultMemPool(&mut pool, dev) == sys::CUresult::CUDA_SUCCESS {
1475 let mut thresh: u64 = u64::MAX;
1476 let _ = sys::cuMemPoolSetAttribute(
1477 pool,
1478 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
1479 &mut thresh as *mut u64 as *mut core::ffi::c_void,
1480 );
1481 }
1482 }
1483 let module = gpu.ctx.load_module(Ptx::from_binary(FATBIN.to_vec()))?;
1484 let hybrid = gpu
1485 .ctx
1486 .load_module(Ptx::from_binary(HYBRID_FATBIN.to_vec()))?;
1487 let qmatvec = gpu
1488 .ctx
1489 .load_module(Ptx::from_binary(QMATVEC_FATBIN.to_vec()))?;
1490 let flash = gpu
1491 .ctx
1492 .load_module(Ptx::from_binary(flash_fatbin_bytes().to_vec()))?;
1493 let gemm = gpu
1494 .ctx
1495 .load_module(Ptx::from_binary(gemm_fatbin_bytes().into_owned()))?;
1496 let router = gpu
1497 .ctx
1498 .load_module(Ptx::from_binary(ROUTER_FATBIN.to_vec()))?;
1499 let sample = gpu
1500 .ctx
1501 .load_module(Ptx::from_binary(SAMPLE_FATBIN.to_vec()))?;
1502 let copy_stream = gpu.ctx.new_stream()?;
1503 if std::env::var("MEMRA_EVT")
1519 .map(|v| v == "1")
1520 .unwrap_or(false)
1521 {
1522 } else {
1524 unsafe {
1525 gpu.ctx.disable_event_tracking();
1526 }
1527 }
1528 Ok(Self {
1529 gpu,
1530 module,
1531 hybrid,
1532 qmatvec,
1533 flash,
1534 flash_g: std::sync::OnceLock::new(),
1535 gemm,
1536 router,
1537 sample,
1538 moe_cache: Mutex::new(None),
1539 w8_mirrors: Mutex::new(std::collections::HashMap::new()),
1540 w8_act: Mutex::new(std::collections::HashMap::new()),
1541 moe_cache_layout: Mutex::new(None),
1542 copy_stream,
1543 capture_keep_on: std::sync::atomic::AtomicBool::new(false),
1544 verify_exact: std::sync::atomic::AtomicBool::new(false),
1545 capture_keep: Mutex::new(Vec::new()),
1546 argmax_partials: Mutex::new(None),
1547 prime_deqw_ws: Mutex::new(None),
1548 router_stage: Mutex::new(None),
1549 fp8_scratch: Mutex::new(None),
1550 fa_vf16_scratch: Mutex::new(None),
1551 fa_part_pool: Mutex::new(None),
1552 fa_part_retired: Mutex::new(Vec::new()),
1553 fn_cache: Mutex::new(Default::default()),
1554 f16_scratch: Mutex::new(None),
1555 #[cfg(memra_cutlass)]
1556 cutlass_scratch: Mutex::new(None),
1557 })
1558 }
1559
1560 pub fn ctx(&self) -> &Arc<CudaContext> {
1561 &self.gpu.ctx
1562 }
1563
1564 pub fn pool_cached_bytes(&self) -> usize {
1582 let (reserved, used) = self.pool_reserved_used();
1583 reserved.saturating_sub(used)
1584 }
1585
1586 pub fn device_graph_mem_reserved(&self) -> usize {
1596 use cudarc::driver::sys as cus;
1597 let Ok(dev) = cudarc::driver::result::device::get(self.gpu.ctx.ordinal() as i32) else {
1598 return 0;
1599 };
1600 let mut bytes: u64 = 0;
1601 let rc = unsafe {
1602 cus::cuDeviceGetGraphMemAttribute(
1603 dev,
1604 cus::CUgraphMem_attribute::CU_GRAPH_MEM_ATTR_RESERVED_MEM_CURRENT,
1605 &mut bytes as *mut u64 as *mut std::ffi::c_void,
1606 )
1607 };
1608 if rc == cus::cudaError_enum::CUDA_SUCCESS {
1609 bytes as usize
1610 } else {
1611 0
1612 }
1613 }
1614
1615 pub fn pool_trim_to_zero(&self) -> usize {
1629 use cudarc::driver::sys;
1630 let (before, _) = self.pool_reserved_used();
1631 unsafe {
1632 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
1633 if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
1634 != sys::CUresult::CUDA_SUCCESS
1635 {
1636 return 0;
1637 }
1638 let _ = sys::cuMemPoolTrimTo(pool, 0);
1639 }
1640 let (after, _) = self.pool_reserved_used();
1641 before.saturating_sub(after)
1642 }
1643
1644 pub fn pool_reserved_used(&self) -> (usize, usize) {
1645 use cudarc::driver::sys;
1646 unsafe {
1647 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
1648 if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
1649 != sys::CUresult::CUDA_SUCCESS
1650 {
1651 return (0, 0);
1652 }
1653 let (mut reserved, mut used) = (0u64, 0u64);
1654 if sys::cuMemPoolGetAttribute(
1655 pool,
1656 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RESERVED_MEM_CURRENT,
1657 &mut reserved as *mut u64 as *mut core::ffi::c_void,
1658 ) != sys::CUresult::CUDA_SUCCESS
1659 {
1660 return (0, 0);
1661 }
1662 if sys::cuMemPoolGetAttribute(
1663 pool,
1664 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_USED_MEM_CURRENT,
1665 &mut used as *mut u64 as *mut core::ffi::c_void,
1666 ) != sys::CUresult::CUDA_SUCCESS
1667 {
1668 return (0, 0);
1669 }
1670 (reserved as usize, used as usize)
1671 }
1672 }
1673
1674 pub fn stream(&self) -> Arc<CudaStream> {
1677 self.gpu.stream()
1678 }
1679 pub fn gkv_on() -> bool {
1682 memra_kv::gkv_on()
1683 }
1684
1685 pub fn wkv_on() -> bool {
1697 memra_kv::wkv_on()
1698 }
1699
1700 pub fn kv_fp8_on() -> bool {
1706 memra_kv::kv_fp8_on()
1707 }
1708
1709 fn fa_func(&self, name: &str, head_dim: usize) -> CudaFunction {
1712 if head_dim == 512 && Self::gkv_on() {
1713 self.func_g(name)
1714 } else {
1715 self.func(name)
1716 }
1717 }
1718
1719 fn func_g(&self, name: &str) -> CudaFunction {
1723 let m = self.flash_g.get_or_init(|| {
1724 self.gpu
1725 .ctx
1726 .load_module(cudarc::nvrtc::Ptx::from_binary(
1727 FLASH_FATBIN_KF8VF8.to_vec(),
1728 ))
1729 .expect("load kf8vf8 flash fatbin (fp8-globals arm)")
1730 });
1731 let key = format!("g:{name}");
1732 if let Some(f) = self.fn_cache.lock().unwrap().get(&key) {
1733 return f.clone();
1734 }
1735 let f = match m.load_function(name) {
1736 Ok(f) => f,
1737 Err(_) => self.func(name),
1738 };
1739 self.fn_cache.lock().unwrap().insert(key, f.clone());
1740 f
1741 }
1742
1743 fn func(&self, name: &str) -> CudaFunction {
1744 if let Some(f) = self.fn_cache.lock().unwrap().get(name) {
1747 return f.clone();
1748 }
1749 let f = self
1750 .module
1751 .load_function(name)
1752 .or_else(|_| self.hybrid.load_function(name))
1753 .or_else(|_| self.qmatvec.load_function(name))
1754 .or_else(|_| self.flash.load_function(name))
1755 .or_else(|_| self.gemm.load_function(name))
1756 .or_else(|_| self.router.load_function(name))
1757 .or_else(|_| self.sample.load_function(name))
1758 .unwrap_or_else(|_| panic!("kernel {name} not in any fatbin"));
1759 self.fn_cache
1760 .lock()
1761 .unwrap()
1762 .insert(name.to_string(), f.clone());
1763 f
1764 }
1765
1766 pub fn scatter_trim_logits(
1769 &self,
1770 src: &CudaSlice<f32>,
1771 d2t: &CudaSlice<u32>,
1772 dst: &mut CudaSlice<f32>,
1773 d_vocab: usize,
1774 n_vocab: usize,
1775 ) -> Result<(), Box<dyn std::error::Error>> {
1776 let f1 = self.func("scatter_trim_logits_f32");
1777 let f2 = self.func("scatter_trim_logits_pass2_f32");
1778 let (dv, nv) = (d_vocab as i32, n_vocab as i32);
1779 let cfg1 = LaunchConfig {
1780 grid_dim: (256, 1, 1),
1781 block_dim: (256, 1, 1),
1782 shared_mem_bytes: 0,
1783 };
1784 let __s_b1 = self.gpu.stream();
1785 let mut b1 = __s_b1.launch_builder(&f1);
1786 b1.arg(src).arg(d2t).arg(&mut *dst).arg(&dv).arg(&nv);
1787 unsafe {
1788 b1.launch(cfg1)?;
1789 }
1790 let cfg2 = LaunchConfig {
1791 grid_dim: (d_vocab.div_ceil(256) as u32, 1, 1),
1792 block_dim: (256, 1, 1),
1793 shared_mem_bytes: 0,
1794 };
1795 let __s_b2 = self.gpu.stream();
1796 let mut b2 = __s_b2.launch_builder(&f2);
1797 b2.arg(src).arg(d2t).arg(&mut *dst).arg(&dv);
1798 unsafe {
1799 b2.launch(cfg2)?;
1800 }
1801 Ok(())
1802 }
1803
1804 #[allow(clippy::too_many_arguments)]
1810 pub fn filter_stats(
1811 &self,
1812 x: &CudaSlice<f32>,
1813 row_stride: usize,
1814 rows: &CudaSlice<i32>,
1815 out_th: &mut CudaSlice<f32>,
1816 out_z: &mut CudaSlice<f32>,
1817 out_max: &mut CudaSlice<f32>,
1818 n: usize,
1819 nrow: usize,
1820 temp: f32,
1821 top_k: i32,
1822 top_p: f32,
1823 min_p: f32,
1824 ) -> Result<(), Box<dyn std::error::Error>> {
1825 static COOP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1852 let coop_on =
1853 *COOP_ON.get_or_init(|| std::env::var("MEMRA_FILTER_COOP").as_deref() != Ok("0"));
1854 if coop_on && self.sm_count() >= 16 {
1855 let cap = self.sm_count() as usize / 16;
1856 let mut done = 0usize;
1857 while done < nrow {
1858 let chunk = cap.min(nrow - done);
1859 self.filter_stats_coop_chunk(
1860 x, row_stride, rows, done, out_th, out_z, out_max, n, chunk, temp, top_k,
1861 top_p, min_p,
1862 )?;
1863 done += chunk;
1864 }
1865 return Ok(());
1866 }
1867 self.filter_stats_plain_program(
1868 x, row_stride, rows, out_th, out_z, out_max, n, nrow, temp, top_k, top_p, min_p,
1869 )
1870 }
1871
1872 #[allow(clippy::too_many_arguments)]
1877 pub fn filter_stats_coop_chunk(
1878 &self,
1879 x: &CudaSlice<f32>,
1880 row_stride: usize,
1881 rows: &CudaSlice<i32>,
1882 row0: usize,
1883 out_th: &mut CudaSlice<f32>,
1884 out_z: &mut CudaSlice<f32>,
1885 out_max: &mut CudaSlice<f32>,
1886 n: usize,
1887 chunk: usize,
1888 temp: f32,
1889 top_k: i32,
1890 top_p: f32,
1891 min_p: f32,
1892 ) -> Result<(), Box<dyn std::error::Error>> {
1893 let (ni, nr, rs) = (n as i32, chunk as i32, row_stride as i64);
1894 let f = self.func("filter_stats_coop_f32");
1895 let mut ws = self.alloc_uninit::<f32>(chunk * (2 * 16 + 2))?;
1896 let cfg = LaunchConfig {
1897 grid_dim: (16, chunk as u32, 1),
1898 block_dim: (512, 1, 1),
1899 shared_mem_bytes: 0,
1900 };
1901 let rows_v = rows.slice(row0..row0 + chunk);
1902 let mut th_v = out_th.slice_mut(row0..row0 + chunk);
1903 let mut z_v = out_z.slice_mut(row0..row0 + chunk);
1904 let mut mx_v = out_max.slice_mut(row0..row0 + chunk);
1905 let __s_b = self.gpu.stream();
1906 let mut b = __s_b.launch_builder(&f);
1907 b.arg(x)
1908 .arg(&rs)
1909 .arg(&rows_v)
1910 .arg(&mut th_v)
1911 .arg(&mut z_v)
1912 .arg(&mut mx_v)
1913 .arg(&mut ws)
1914 .arg(&ni)
1915 .arg(&nr)
1916 .arg(&temp)
1917 .arg(&top_k)
1918 .arg(&top_p)
1919 .arg(&min_p);
1920 unsafe {
1921 b.launch_cooperative(cfg)?;
1922 }
1923 Ok(())
1924 }
1925
1926 #[allow(clippy::too_many_arguments)]
1930 pub fn filter_stats_plain_program(
1931 &self,
1932 x: &CudaSlice<f32>,
1933 row_stride: usize,
1934 rows: &CudaSlice<i32>,
1935 out_th: &mut CudaSlice<f32>,
1936 out_z: &mut CudaSlice<f32>,
1937 out_max: &mut CudaSlice<f32>,
1938 n: usize,
1939 nrow: usize,
1940 temp: f32,
1941 top_k: i32,
1942 top_p: f32,
1943 min_p: f32,
1944 ) -> Result<(), Box<dyn std::error::Error>> {
1945 let (ni, nr, rs) = (n as i32, nrow as i32, row_stride as i64);
1946 let f = self.func("filter_stats_f32");
1947 let cfg = LaunchConfig {
1948 grid_dim: (nrow as u32, 1, 1),
1949 block_dim: (1024, 1, 1),
1950 shared_mem_bytes: 0,
1951 };
1952 let __s_b = self.gpu.stream();
1953 let mut b = __s_b.launch_builder(&f);
1954 b.arg(x)
1955 .arg(&rs)
1956 .arg(rows)
1957 .arg(&mut *out_th)
1958 .arg(&mut *out_z)
1959 .arg(&mut *out_max)
1960 .arg(&ni)
1961 .arg(&nr)
1962 .arg(&temp)
1963 .arg(&top_k)
1964 .arg(&top_p)
1965 .arg(&min_p);
1966 unsafe {
1967 b.launch(cfg)?;
1968 }
1969 Ok(())
1970 }
1971
1972 #[allow(clippy::too_many_arguments)]
1974 pub fn softmax_gather_filtered(
1975 &self,
1976 x: &CudaSlice<f32>,
1977 row_stride: usize,
1978 ids: &CudaSlice<u32>,
1979 rows: &CudaSlice<i32>,
1980 th: &CudaSlice<f32>,
1981 z: &CudaSlice<f32>,
1982 out: &mut CudaSlice<f32>,
1983 n: usize,
1984 npair: usize,
1985 temp: f32,
1986 ) -> Result<(), Box<dyn std::error::Error>> {
1987 let f = self.func("softmax_gather_filtered_f32");
1988 let (ni, np, rs) = (n as i32, npair as i32, row_stride as i64);
1989 let cfg = LaunchConfig {
1990 grid_dim: (npair as u32, 1, 1),
1991 block_dim: (256, 1, 1),
1992 shared_mem_bytes: 0,
1993 };
1994 let __s_b = self.gpu.stream();
1995 let mut b = __s_b.launch_builder(&f);
1996 b.arg(x)
1997 .arg(&rs)
1998 .arg(ids)
1999 .arg(rows)
2000 .arg(th)
2001 .arg(z)
2002 .arg(&mut *out)
2003 .arg(&ni)
2004 .arg(&np)
2005 .arg(&temp);
2006 unsafe {
2007 b.launch(cfg)?;
2008 }
2009 Ok(())
2010 }
2011
2012 #[allow(clippy::too_many_arguments)]
2014 pub fn residual_sample_filtered(
2015 &self,
2016 p: &CudaSlice<f32>,
2017 q: Option<&CudaSlice<f32>>,
2018 n: usize,
2019 temp: f32,
2020 seed: u64,
2021 stream_pos: u32,
2022 p_stats: (f32, f32, f32),
2023 q_stats: (f32, f32, f32),
2024 out_tok: &mut CudaSlice<u32>,
2025 ) -> Result<(), Box<dyn std::error::Error>> {
2026 let f = self.func("residual_sample_filtered_f32");
2027 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2028 let has_q: i32 = q.is_some() as i32;
2029 let qbuf = q.unwrap_or(p);
2030 let (pm, pth, pz) = p_stats;
2031 let (qm, qth, qz) = q_stats;
2032 let cfg = LaunchConfig {
2033 grid_dim: (1, 1, 1),
2034 block_dim: (1024, 1, 1),
2035 shared_mem_bytes: 0,
2036 };
2037 let __s_b = self.gpu.stream();
2038 let mut b = __s_b.launch_builder(&f);
2039 b.arg(p)
2040 .arg(qbuf)
2041 .arg(&has_q)
2042 .arg(&ni)
2043 .arg(&temp)
2044 .arg(&slo)
2045 .arg(&shi)
2046 .arg(&stream_pos)
2047 .arg(&pm)
2048 .arg(&pth)
2049 .arg(&pz)
2050 .arg(&qm)
2051 .arg(&qth)
2052 .arg(&qz)
2053 .arg(&mut *out_tok);
2054 unsafe {
2055 b.launch(cfg)?;
2056 }
2057 Ok(())
2058 }
2059
2060 #[allow(clippy::too_many_arguments)]
2066 pub fn residual_sample_sparse_q(
2067 &self,
2068 p: &CudaSlice<f32>,
2069 cand_ids: &CudaSlice<u32>,
2070 q_probs: &CudaSlice<f32>,
2071 n_cand: usize,
2072 n: usize,
2073 temp: f32,
2074 seed: u64,
2075 stream_pos: u32,
2076 p_stats: (f32, f32, f32),
2077 out_tok: &mut CudaSlice<u32>,
2078 ) -> Result<(), Box<dyn std::error::Error>> {
2079 assert!(
2080 n_cand >= 1 && n_cand <= 32,
2081 "residual_sample_sparse_q supports 1..=32 candidates, got {n_cand}"
2082 );
2083 let f = self.func("residual_sample_sparse_q_f32");
2084 let (ni, nc) = (n as i32, n_cand as i32);
2085 let (slo, shi) = ((seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2086 let (pm, pth, pz) = p_stats;
2087 let cfg = LaunchConfig {
2088 grid_dim: (1, 1, 1),
2089 block_dim: (1024, 1, 1),
2090 shared_mem_bytes: 0,
2091 };
2092 let __s_b = self.gpu.stream();
2093 let mut b = __s_b.launch_builder(&f);
2094 b.arg(p)
2095 .arg(cand_ids)
2096 .arg(q_probs)
2097 .arg(&nc)
2098 .arg(&ni)
2099 .arg(&temp)
2100 .arg(&slo)
2101 .arg(&shi)
2102 .arg(&stream_pos)
2103 .arg(&pm)
2104 .arg(&pth)
2105 .arg(&pz)
2106 .arg(&mut *out_tok);
2107 unsafe {
2108 b.launch(cfg)?;
2109 }
2110 Ok(())
2111 }
2112
2113 #[allow(clippy::too_many_arguments)]
2115 pub fn gumbel_perturb_filtered(
2116 &self,
2117 x: &CudaSlice<f32>,
2118 y: &mut CudaSlice<f32>,
2119 n: usize,
2120 seed: u64,
2121 stream_pos: u32,
2122 temp: f32,
2123 row_max: f32,
2124 th: f32,
2125 ) -> Result<(), Box<dyn std::error::Error>> {
2126 let f = self.func("gumbel_perturb_filtered_f32");
2127 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2128 let cfg = LaunchConfig {
2129 grid_dim: (n.div_ceil(256) as u32, 1, 1),
2130 block_dim: (256, 1, 1),
2131 shared_mem_bytes: 0,
2132 };
2133 let __s_b = self.gpu.stream();
2134 let mut b = __s_b.launch_builder(&f);
2135 b.arg(x)
2136 .arg(&mut *y)
2137 .arg(&ni)
2138 .arg(&slo)
2139 .arg(&shi)
2140 .arg(&stream_pos)
2141 .arg(&temp)
2142 .arg(&row_max)
2143 .arg(&th);
2144 unsafe {
2145 b.launch(cfg)?;
2146 }
2147 Ok(())
2148 }
2149
2150 #[allow(clippy::too_many_arguments)]
2154 pub fn penalize_logits(
2155 &self,
2156 x: &mut CudaSlice<f32>,
2157 hist: &CudaSlice<u32>,
2158 n_hist: usize,
2159 rep: f32,
2160 freq: f32,
2161 present: f32,
2162 n: usize,
2163 ) -> Result<(), Box<dyn std::error::Error>> {
2164 if n_hist == 0 {
2165 return Ok(());
2166 }
2167 let f = self.func("penalize_logits_f32");
2168 let (nh, ni) = (n_hist as i32, n as i32);
2169 let cfg = LaunchConfig {
2170 grid_dim: (n_hist.div_ceil(128) as u32, 1, 1),
2171 block_dim: (128, 1, 1),
2172 shared_mem_bytes: 0,
2173 };
2174 let __s_b = self.gpu.stream();
2175 let mut b = __s_b.launch_builder(&f);
2176 b.arg(&mut *x)
2177 .arg(hist)
2178 .arg(&nh)
2179 .arg(&rep)
2180 .arg(&freq)
2181 .arg(&present)
2182 .arg(&ni);
2183 unsafe {
2184 b.launch(cfg)?;
2185 }
2186 Ok(())
2187 }
2188
2189 #[allow(clippy::too_many_arguments)]
2191 pub fn penalize_logits_rows(
2192 &self,
2193 x: &mut CudaSlice<f32>,
2194 hist: &CudaSlice<u32>,
2195 n_hist: usize,
2196 rep: f32,
2197 freq: f32,
2198 present: f32,
2199 n: usize,
2200 nrow: usize,
2201 ) -> Result<(), Box<dyn std::error::Error>> {
2202 if n_hist == 0 || nrow == 0 {
2203 return Ok(());
2204 }
2205 let f = self.func("penalize_logits_rows_f32");
2206 let (nh, ni, nr) = (n_hist as i32, n as i32, nrow as i32);
2207 let cfg = LaunchConfig {
2208 grid_dim: (n_hist.div_ceil(128) as u32, nrow as u32, 1),
2209 block_dim: (128, 1, 1),
2210 shared_mem_bytes: 0,
2211 };
2212 let __s_b = self.gpu.stream();
2213 let mut b = __s_b.launch_builder(&f);
2214 b.arg(&mut *x)
2215 .arg(hist)
2216 .arg(&nh)
2217 .arg(&rep)
2218 .arg(&freq)
2219 .arg(&present)
2220 .arg(&ni)
2221 .arg(&nr);
2222 unsafe {
2223 b.launch(cfg)?;
2224 }
2225 Ok(())
2226 }
2227
2228 #[allow(clippy::too_many_arguments)]
2234 pub fn penalize_logits_sparse_rows(
2235 &self,
2236 x: &mut CudaSlice<f32>,
2237 ids: &[u32],
2238 counts: &[u32],
2239 offsets: &[i32],
2240 rows: &[i32],
2241 reps: &[f32],
2242 freqs: &[f32],
2243 presents: &[f32],
2244 n: usize,
2245 ) -> Result<(), Box<dyn std::error::Error>> {
2246 let nrow = rows.len();
2247 if nrow == 0 {
2248 return Ok(());
2249 }
2250 let _ni = i32::try_from(n).map_err(|_| "sparse penalty logits width must fit CUDA i32")?;
2251 let _nr = i32::try_from(nrow).map_err(|_| "sparse penalty row count must fit CUDA i32")?;
2252 let entry_count =
2253 i32::try_from(ids.len()).map_err(|_| "sparse penalty entry count must fit CUDA i32")?;
2254 if ids.len() != counts.len()
2255 || offsets.len() != nrow + 1
2256 || reps.len() != nrow
2257 || freqs.len() != nrow
2258 || presents.len() != nrow
2259 || offsets.first().copied() != Some(0)
2260 || offsets.last().copied() != Some(entry_count)
2261 {
2262 return Err("sparse penalty row metadata shape mismatch".into());
2263 }
2264 if counts.contains(&0) {
2265 return Err("sparse penalty counts must be positive".into());
2266 }
2267 let mut max_len = 0usize;
2268 for pair in offsets.windows(2) {
2269 if pair[0] < 0 || pair[1] < pair[0] {
2270 return Err("sparse penalty offsets must be monotonic".into());
2271 }
2272 max_len = max_len.max((pair[1] - pair[0]) as usize);
2273 }
2274 if max_len == 0 {
2275 return Ok(());
2276 }
2277
2278 let mut seen = std::collections::HashSet::with_capacity(ids.len());
2279 for (r, &row) in rows.iter().enumerate() {
2280 if row < 0 || (row as usize + 1).saturating_mul(n) > x.len() {
2281 return Err("sparse penalty row index exceeds logits shape".into());
2282 }
2283 let begin = offsets[r] as usize;
2284 let end = offsets[r + 1] as usize;
2285 for &id in &ids[begin..end] {
2286 if id as usize >= n {
2287 return Err("sparse penalty token id exceeds logits row".into());
2288 }
2289 if !seen.insert((row, id)) {
2290 return Err("sparse penalty entries must be unique per logits row".into());
2291 }
2292 }
2293 }
2294
2295 unsafe {
2297 self.penalize_logits_sparse_rows_unchecked(
2298 x, ids, counts, offsets, rows, reps, freqs, presents, n,
2299 )
2300 }
2301 }
2302
2303 #[allow(clippy::too_many_arguments)]
2312 pub(crate) unsafe fn penalize_logits_sparse_rows_unchecked(
2313 &self,
2314 x: &mut CudaSlice<f32>,
2315 ids: &[u32],
2316 counts: &[u32],
2317 offsets: &[i32],
2318 rows: &[i32],
2319 reps: &[f32],
2320 freqs: &[f32],
2321 presents: &[f32],
2322 n: usize,
2323 ) -> Result<(), Box<dyn std::error::Error>> {
2324 let nrow = rows.len();
2325 if nrow == 0 {
2326 return Ok(());
2327 }
2328 let max_len = offsets
2329 .windows(2)
2330 .map(|pair| (pair[1] - pair[0]) as usize)
2331 .max()
2332 .unwrap_or(0);
2333 if max_len == 0 {
2334 return Ok(());
2335 }
2336 let ids_d = self.htod_u32_v(ids)?;
2337 let counts_d = self.htod_u32_v(counts)?;
2338 let offsets_d = self.htod_i32(offsets)?;
2339 let rows_d = self.htod_i32(rows)?;
2340 let reps_d = self.htod(reps)?;
2341 let freqs_d = self.htod(freqs)?;
2342 let presents_d = self.htod(presents)?;
2343 let f = self.func("penalize_logits_sparse_rows_f32");
2344 let ni = i32::try_from(n).map_err(|_| "sparse penalty logits width must fit CUDA i32")?;
2345 let nr = i32::try_from(nrow).map_err(|_| "sparse penalty row count must fit CUDA i32")?;
2346 let cfg = LaunchConfig {
2347 grid_dim: (max_len.div_ceil(128) as u32, nrow as u32, 1),
2348 block_dim: (128, 1, 1),
2349 shared_mem_bytes: 0,
2350 };
2351 let __s_b = self.gpu.stream();
2352 let mut b = __s_b.launch_builder(&f);
2353 b.arg(&mut *x)
2354 .arg(&ids_d)
2355 .arg(&counts_d)
2356 .arg(&offsets_d)
2357 .arg(&rows_d)
2358 .arg(&reps_d)
2359 .arg(&freqs_d)
2360 .arg(&presents_d)
2361 .arg(&ni)
2362 .arg(&nr);
2363 unsafe {
2364 b.launch(cfg)?;
2365 }
2366 Ok(())
2367 }
2368
2369 #[allow(clippy::too_many_arguments)]
2377 pub fn penalize_logits_rows_inc(
2378 &self,
2379 x: &mut CudaSlice<f32>,
2380 hist: &CudaSlice<u32>,
2381 n_hist0: usize,
2382 rep: f32,
2383 freq: f32,
2384 present: f32,
2385 n: usize,
2386 nrow: usize,
2387 win: usize,
2388 ) -> Result<(), Box<dyn std::error::Error>> {
2389 if nrow == 0 || win == 0 || (n_hist0 == 0 && nrow == 1) {
2390 return Ok(());
2391 }
2392 debug_assert!(
2393 hist.len() >= n_hist0 + nrow - 1,
2394 "rows-inc hist must carry n_hist0 + nrow - 1 ids"
2395 );
2396 let f = self.func("penalize_logits_rows_inc_f32");
2397 let max_len = win.min(n_hist0 + nrow - 1).max(1);
2398 let (nh, ni, nr, wi) = (n_hist0 as i32, n as i32, nrow as i32, win as i32);
2399 let cfg = LaunchConfig {
2400 grid_dim: (max_len.div_ceil(128) as u32, nrow as u32, 1),
2401 block_dim: (128, 1, 1),
2402 shared_mem_bytes: 0,
2403 };
2404 let __s_b = self.gpu.stream();
2405 let mut b = __s_b.launch_builder(&f);
2406 b.arg(&mut *x)
2407 .arg(hist)
2408 .arg(&nh)
2409 .arg(&rep)
2410 .arg(&freq)
2411 .arg(&present)
2412 .arg(&ni)
2413 .arg(&nr)
2414 .arg(&wi);
2415 unsafe {
2416 b.launch(cfg)?;
2417 }
2418 Ok(())
2419 }
2420
2421 pub fn wpf_level() -> u32 {
2429 static ON: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
2430 *ON.get_or_init(|| {
2431 std::env::var("MEMRA_WPF")
2432 .ok()
2433 .and_then(|v| v.parse().ok())
2434 .unwrap_or(1)
2435 })
2436 }
2437
2438 pub fn set_verify_exact(&self, on: bool) {
2454 self.verify_exact
2455 .store(on, std::sync::atomic::Ordering::Relaxed);
2456 }
2457 pub(crate) fn verify_exact_on(&self) -> bool {
2458 self.verify_exact.load(std::sync::atomic::Ordering::Relaxed)
2459 }
2460
2461 pub fn exact_scope(&self, on: bool) -> ExactScope<'_> {
2467 ExactScope::set(&self.verify_exact, on)
2468 }
2469
2470 pub fn qkv_append_on() -> bool {
2473 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2474 *ON.get_or_init(|| {
2475 std::env::var("MEMRA_QKV_APPEND")
2476 .map(|v| v != "0")
2477 .unwrap_or(true)
2478 })
2479 }
2480
2481 pub fn pdl_wb_on() -> bool {
2484 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2485 *ON.get_or_init(|| {
2486 std::env::var("MEMRA_PDL_WB")
2487 .map(|v| v != "0")
2488 .unwrap_or(true)
2489 })
2490 }
2491
2492 pub fn norm_ilp_on() -> bool {
2500 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2501 *ON.get_or_init(|| {
2502 std::env::var("MEMRA_NORM_ILP")
2503 .map(|v| v != "0")
2504 .unwrap_or(true)
2505 })
2506 }
2507
2508 pub fn tk_ffn_dual_on() -> bool {
2516 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2517 *ON.get_or_init(|| {
2518 std::env::var("MEMRA_TK_FFN_DUAL")
2519 .map(|v| v != "0")
2520 .unwrap_or(true)
2521 })
2522 }
2523
2524 pub fn pdl_mmvq_on() -> bool {
2528 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2529 *ON.get_or_init(|| {
2530 std::env::var("MEMRA_PDL_MMVQ")
2531 .map(|v| v != "0")
2532 .unwrap_or(true)
2533 })
2534 }
2535
2536 pub fn pdl_on() -> bool {
2537 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2538 *ON.get_or_init(|| std::env::var("MEMRA_PDL").map(|v| v != "0").unwrap_or(true))
2539 }
2540
2541 pub fn pdl_nvfp4q8_on() -> bool {
2547 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2548 *ON.get_or_init(|| {
2549 std::env::var("MEMRA_PDL_NVFP4")
2550 .map(|v| v != "0")
2551 .unwrap_or(true)
2552 })
2553 }
2554
2555 fn q40_mr1_on() -> bool {
2561 static Q40MR: std::sync::OnceLock<Option<u32>> = std::sync::OnceLock::new();
2562 match *Q40MR.get_or_init(|| {
2563 std::env::var("MEMRA_Q40_MR")
2564 .ok()
2565 .and_then(|v| v.parse().ok())
2566 }) {
2567 Some(v) => v == 1,
2568 None => crate::FUSED_MR1_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
2569 }
2570 }
2571
2572 fn pdl_func_flash(
2577 &self,
2578 g: bool,
2579 name: &'static str,
2580 ) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
2581 use cudarc::driver::sys as cu;
2582 static MODS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool), usize>>> =
2589 std::sync::Mutex::new(None);
2590 static FNS: std::sync::Mutex<
2591 Option<std::collections::HashMap<(usize, bool, &'static str), usize>>,
2592 > = std::sync::Mutex::new(None);
2593 let ctx_key = self.ctx().cu_ctx() as usize;
2594 if let Some(&f) = FNS
2595 .lock()
2596 .unwrap()
2597 .get_or_insert_with(Default::default)
2598 .get(&(ctx_key, g, name))
2599 {
2600 return Ok(f as cu::CUfunction);
2601 }
2602 let module = {
2603 let mut mods = MODS.lock().unwrap();
2604 let map = mods.get_or_insert_with(Default::default);
2605 match map.get(&(ctx_key, g)) {
2606 Some(&m) => m,
2607 None => {
2608 let m = self.pdl_load_module_in_ctx(if g {
2609 FLASH_FATBIN_KF8VF8
2610 } else {
2611 FLASH_FATBIN
2612 })?;
2613 map.insert((ctx_key, g), m);
2614 m
2615 }
2616 }
2617 };
2618 let cname = std::ffi::CString::new(name)?;
2619 let mut f: cu::CUfunction = std::ptr::null_mut();
2620 let r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
2621 if r != cu::CUresult::CUDA_SUCCESS {
2622 return Err(format!("pdl_func_flash {name} (g={g}): {r:?}").into());
2623 }
2624 FNS.lock()
2625 .unwrap()
2626 .get_or_insert_with(Default::default)
2627 .insert((ctx_key, g, name), f as usize);
2628 Ok(f)
2629 }
2630
2631 fn pdl_load_module_in_ctx(&self, bytes: &[u8]) -> Result<usize, Box<dyn std::error::Error>> {
2636 use cudarc::driver::sys as cu;
2637 let mut prev: cu::CUcontext = std::ptr::null_mut();
2638 unsafe {
2639 cu::cuCtxGetCurrent(&mut prev).result()?;
2640 }
2641 self.ctx().bind_to_thread()?;
2642 let mut m: cu::CUmodule = std::ptr::null_mut();
2643 let r = unsafe { cu::cuModuleLoadData(&mut m, bytes.as_ptr() as *const std::ffi::c_void) };
2644 let restore = if prev.is_null() {
2645 cu::CUresult::CUDA_SUCCESS
2646 } else {
2647 unsafe { cu::cuCtxSetCurrent(prev) }
2648 };
2649 if r != cu::CUresult::CUDA_SUCCESS {
2650 return Err(format!("pdl module load: {r:?}").into());
2651 }
2652 if restore != cu::CUresult::CUDA_SUCCESS {
2653 return Err(format!("pdl module load: ctx restore {restore:?}").into());
2654 }
2655 Ok(m as usize)
2656 }
2657
2658 pub fn raw_kernel_function(
2661 &self,
2662 name: &'static str,
2663 ) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
2664 self.pdl_func(name)
2665 }
2666
2667 fn pdl_func(
2668 &self,
2669 name: &'static str,
2670 ) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
2671 use cudarc::driver::sys as cu;
2672 static MODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
2675 std::sync::Mutex::new(None);
2676 static QMODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
2679 std::sync::Mutex::new(None);
2680 static FNS: std::sync::Mutex<
2681 Option<std::collections::HashMap<(usize, &'static str), usize>>,
2682 > = std::sync::Mutex::new(None);
2683 let ctx_key = self.ctx().cu_ctx() as usize;
2684 if let Some(&f) = FNS
2685 .lock()
2686 .unwrap()
2687 .get_or_insert_with(Default::default)
2688 .get(&(ctx_key, name))
2689 {
2690 return Ok(f as cu::CUfunction);
2691 }
2692 let module = {
2693 let mut mods = MODULES.lock().unwrap();
2694 let map = mods.get_or_insert_with(Default::default);
2695 match map.get(&ctx_key) {
2696 Some(&m) => m,
2697 None => {
2698 let m = self.pdl_load_module_in_ctx(FATBIN)?;
2699 map.insert(ctx_key, m);
2700 m
2701 }
2702 }
2703 };
2704 let cname = std::ffi::CString::new(name)?;
2705 let mut f: cu::CUfunction = std::ptr::null_mut();
2706 let mut r =
2707 unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
2708 if r == cu::CUresult::CUDA_ERROR_NOT_FOUND {
2709 let qmodule = {
2710 let mut mods = QMODULES.lock().unwrap();
2711 let map = mods.get_or_insert_with(Default::default);
2712 match map.get(&ctx_key) {
2713 Some(&m) => m,
2714 None => {
2715 let m = self.pdl_load_module_in_ctx(QMATVEC_FATBIN)?;
2716 map.insert(ctx_key, m);
2717 m
2718 }
2719 }
2720 };
2721 r = unsafe { cu::cuModuleGetFunction(&mut f, qmodule as cu::CUmodule, cname.as_ptr()) };
2722 }
2723 if r != cu::CUresult::CUDA_SUCCESS {
2724 return Err(format!("pdl_func {name}: {r:?}").into());
2725 }
2726 FNS.lock()
2727 .unwrap()
2728 .get_or_insert_with(Default::default)
2729 .insert((ctx_key, name), f as usize);
2730 Ok(f)
2731 }
2732
2733 unsafe fn launch_pdl_flash(
2745 &self,
2746 g: bool,
2747 name: &'static str,
2748 grid: (u32, u32, u32),
2749 block: (u32, u32, u32),
2750 smem: u32,
2751 params: &mut [*mut std::ffi::c_void],
2752 ) -> Result<(), Box<dyn std::error::Error>> {
2753 use cudarc::driver::sys as cu;
2754 let f = self.pdl_func_flash(g, name)?;
2755 if smem > 0 {
2756 let r =
2758 unsafe {
2759 cu::cuFuncSetAttribute(f,
2760 cu::CUfunction_attribute_enum::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
2761 smem as i32)
2762 };
2763 if r != cu::CUresult::CUDA_SUCCESS {
2764 return Err(format!("pdl smem attr {name}: {r:?}").into());
2765 }
2766 }
2767 let mut attr = cu::CUlaunchAttribute {
2768 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
2769 pad: [0; 4],
2770 value: cu::CUlaunchAttributeValue {
2771 programmaticStreamSerializationAllowed: 1,
2772 },
2773 };
2774 let cfg = cu::CUlaunchConfig {
2775 gridDimX: grid.0,
2776 gridDimY: grid.1,
2777 gridDimZ: grid.2,
2778 blockDimX: block.0,
2779 blockDimY: block.1,
2780 blockDimZ: block.2,
2781 sharedMemBytes: smem,
2782 hStream: self.gpu.stream().cu_stream(),
2783 attrs: &mut attr,
2784 numAttrs: 1,
2785 };
2786 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
2787 if r != cu::CUresult::CUDA_SUCCESS {
2788 return Err(format!("launch_pdl_flash {name}: {r:?}").into());
2789 }
2790 Ok(())
2791 }
2792
2793 unsafe fn launch_pdl(
2794 &self,
2795 name: &'static str,
2796 grid: (u32, u32, u32),
2797 block: (u32, u32, u32),
2798 params: &mut [*mut std::ffi::c_void],
2799 ) -> Result<(), Box<dyn std::error::Error>> {
2800 use cudarc::driver::sys as cu;
2801 let f = self.pdl_func(name)?;
2802 let mut attr = cu::CUlaunchAttribute {
2803 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
2804 pad: [0; 4],
2805 value: cu::CUlaunchAttributeValue {
2806 programmaticStreamSerializationAllowed: 1,
2807 },
2808 };
2809 let cfg = cu::CUlaunchConfig {
2810 gridDimX: grid.0,
2811 gridDimY: grid.1,
2812 gridDimZ: grid.2,
2813 blockDimX: block.0,
2814 blockDimY: block.1,
2815 blockDimZ: block.2,
2816 sharedMemBytes: 0,
2817 hStream: self.gpu.stream().cu_stream(),
2818 attrs: &mut attr,
2819 numAttrs: 1,
2820 };
2821 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
2822 if r != cu::CUresult::CUDA_SUCCESS {
2823 return Err(format!("launch_pdl {name}: {r:?}").into());
2824 }
2825 Ok(())
2826 }
2827
2828 pub fn prefetch_weight_l2(
2831 &self,
2832 w: &crate::model::GpuTensor,
2833 ) -> Result<(), Box<dyn std::error::Error>> {
2834 if let crate::model::GpuTensor::Quant { bytes, rp4, .. } = w {
2835 let p = rp4.as_ref().unwrap_or(bytes);
2836 self.prefetch_l2(p, p.len())?;
2837 }
2838 Ok(())
2839 }
2840
2841 pub fn gather_row_bf16(
2844 &self,
2845 table: &CudaSlice<u8>,
2846 tok: &CudaSlice<u32>,
2847 idx: usize,
2848 dst: &mut CudaSlice<f32>,
2849 ncols: usize,
2850 ) -> Result<(), Box<dyn std::error::Error>> {
2851 let f = self.func("gather_row_bf16_f32");
2852 let cfg = LaunchConfig {
2853 grid_dim: (ncols.div_ceil(256) as u32, 1, 1),
2854 block_dim: (256, 1, 1),
2855 shared_mem_bytes: 0,
2856 };
2857 let (nc, ix) = (ncols as i32, idx as i32);
2858 let __s_b = self.gpu.stream();
2859 let mut b = __s_b.launch_builder(&f);
2860 b.arg(table).arg(tok).arg(&ix).arg(dst).arg(&nc);
2861 unsafe {
2862 b.launch(cfg)?;
2863 }
2864 Ok(())
2865 }
2866
2867 #[allow(clippy::too_many_arguments)]
2873 pub fn dflash2_dynconv(
2874 &self,
2875 x: &CudaSlice<f32>,
2876 dyn_: &CudaSlice<f32>,
2877 base: &CudaSlice<f32>,
2878 out: &mut CudaSlice<f32>,
2879 rows: usize,
2880 hidden: usize,
2881 group_size: usize,
2882 ksize: usize,
2883 half: usize,
2884 ) -> Result<(), Box<dyn std::error::Error>> {
2885 assert_eq!(hidden % group_size, 0, "hidden % group_size != 0");
2886 let f = self.func("dflash2_dynconv_f32");
2887 let n = rows * hidden;
2888 let cfg = LaunchConfig {
2889 grid_dim: (n.div_ceil(256) as u32, 1, 1),
2890 block_dim: (256, 1, 1),
2891 shared_mem_bytes: 0,
2892 };
2893 let (ri, hi, gi, ki, hf) = (
2894 rows as i32,
2895 hidden as i32,
2896 group_size as i32,
2897 ksize as i32,
2898 half as i32,
2899 );
2900 let __s_b = self.gpu.stream();
2901 let mut b = __s_b.launch_builder(&f);
2902 b.arg(x)
2903 .arg(dyn_)
2904 .arg(base)
2905 .arg(out)
2906 .arg(&ri)
2907 .arg(&hi)
2908 .arg(&gi)
2909 .arg(&ki)
2910 .arg(&hf);
2911 unsafe {
2912 b.launch(cfg)?;
2913 }
2914 Ok(())
2915 }
2916
2917 pub fn topk_rows(
2921 &self,
2922 logits: &CudaSlice<f32>,
2923 n_rows: usize,
2924 n_cols: usize,
2925 k: usize,
2926 ) -> Result<(CudaSlice<f32>, CudaSlice<u32>), Box<dyn std::error::Error>> {
2927 assert!(k <= 32 && k >= 1, "topk_rows supports 1..=32, got {k}");
2928 assert!(k <= n_cols, "topk_rows: k {k} > n_cols {n_cols}");
2929 let f = self.func("topk_rows_f32");
2930 let nth = 256usize;
2931 let mut vals = self.uninit(n_rows * k)?;
2932 let mut idxs = self.gpu.stream().alloc_zeros::<u32>(n_rows * k)?;
2933 let cfg = LaunchConfig {
2934 grid_dim: (n_rows as u32, 1, 1),
2935 block_dim: (nth as u32, 1, 1),
2936 shared_mem_bytes: (nth * k * 8) as u32,
2937 };
2938 let (nr, nc, ki) = (n_rows as i32, n_cols as i32, k as i32);
2939 let __s_b = self.gpu.stream();
2940 let mut b = __s_b.launch_builder(&f);
2941 b.arg(logits)
2942 .arg(&nr)
2943 .arg(&nc)
2944 .arg(&ki)
2945 .arg(&mut vals)
2946 .arg(&mut idxs);
2947 unsafe {
2948 b.launch(cfg)?;
2949 }
2950 Ok((vals, idxs))
2951 }
2952
2953 pub fn add_row_inplace(
2955 &self,
2956 logits: &mut CudaSlice<f32>,
2957 bias: &CudaSlice<f32>,
2958 n: usize,
2959 row_off: usize,
2960 ) -> Result<(), Box<dyn std::error::Error>> {
2961 let f = self.func("add_row_inplace_f32");
2962 let cfg = LaunchConfig {
2963 grid_dim: (n.div_ceil(256) as u32, 1, 1),
2964 block_dim: (256, 1, 1),
2965 shared_mem_bytes: 0,
2966 };
2967 let (ni, off) = (n as i32, row_off as i64);
2968 let __s_b = self.gpu.stream();
2969 let mut b = __s_b.launch_builder(&f);
2970 b.arg(logits).arg(bias).arg(&ni).arg(&off);
2971 unsafe {
2972 b.launch(cfg)?;
2973 }
2974 Ok(())
2975 }
2976
2977 pub fn prefetch_l2(
2979 &self,
2980 p: &CudaSlice<u8>,
2981 n: usize,
2982 ) -> Result<(), Box<dyn std::error::Error>> {
2983 let f = self.func("prefetch_l2_bytes");
2984 let lines = n.div_ceil(128);
2985 let ni = n as i64;
2986 let cfg = LaunchConfig {
2987 grid_dim: (lines.div_ceil(256) as u32, 1, 1),
2988 block_dim: (256, 1, 1),
2989 shared_mem_bytes: 0,
2990 };
2991 let __s_b = self.gpu.stream();
2992 let mut b = __s_b.launch_builder(&f);
2993 b.arg(p).arg(&ni);
2994 unsafe {
2995 b.launch(cfg)?;
2996 }
2997 Ok(())
2998 }
2999
3000 pub fn router_gemv(
3003 &self,
3004 w: &CudaSlice<f32>,
3005 x: &CudaSlice<f32>,
3006 n_embd: usize,
3007 n_experts: usize,
3008 t: usize,
3009 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3010 let w8 = match std::env::var("MEMRA_ROUTER_V2").as_deref() {
3016 Ok("0") => false,
3017 Ok(_) => true,
3018 Err(_) => ROUTER_W8_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
3019 };
3020 let batch = w8 && t >= ROUTER_BATCH_MIN_T && router_batch_on();
3030 self.router_gemv_form(w, x, n_embd, n_experts, t, w8, batch)
3031 }
3032
3033 pub fn router_gemv_form(
3036 &self,
3037 w: &CudaSlice<f32>,
3038 x: &CudaSlice<f32>,
3039 n_embd: usize,
3040 n_experts: usize,
3041 t: usize,
3042 w8: bool,
3043 batch: bool,
3044 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3045 debug_assert!(!batch || w8, "batch twin exists for the w8 form only");
3046 let mut y = self.alloc_uninit::<f32>(t * n_experts)?;
3047 let f = if batch {
3048 self.func("router_gemv_f32_w8_batch")
3049 } else if w8 {
3050 self.func("router_gemv_f32_w8")
3051 } else {
3052 self.func("router_gemv_f32")
3053 };
3054 let (ne, nx, ti) = (n_embd as i32, n_experts as i32, t as i32);
3055 let cfg = if batch {
3056 LaunchConfig {
3057 grid_dim: (n_experts.div_ceil(8) as u32, t.div_ceil(8) as u32, 1),
3058 block_dim: (32, 8, 1),
3059 shared_mem_bytes: 0,
3060 }
3061 } else {
3062 LaunchConfig {
3063 grid_dim: (n_experts as u32, t as u32, 1),
3064 block_dim: (32, if w8 { 8 } else { 1 }, 1),
3065 shared_mem_bytes: 0,
3066 }
3067 };
3068 let __s_b = self.gpu.stream();
3069 let mut b = __s_b.launch_builder(&f);
3070 b.arg(w).arg(x).arg(&mut y).arg(&ne).arg(&nx).arg(&ti);
3071 unsafe {
3072 b.launch(cfg)?;
3073 }
3074 Ok(y)
3075 }
3076
3077 pub fn router_gemv_into(
3080 &self,
3081 w: &CudaSlice<f32>,
3082 x: &CudaSlice<f32>,
3083 y: &mut CudaSlice<f32>,
3084 n_embd: usize,
3085 n_experts: usize,
3086 t: usize,
3087 ) -> Result<(), Box<dyn std::error::Error>> {
3088 if y.len() < t * n_experts {
3089 return Err("router_gemv_into output too small".into());
3090 }
3091 let w8 = match std::env::var("MEMRA_ROUTER_V2").as_deref() {
3092 Ok("0") => false,
3093 Ok(_) => true,
3094 Err(_) => ROUTER_W8_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
3095 };
3096 let f = if w8 {
3097 self.func("router_gemv_f32_w8")
3098 } else {
3099 self.func("router_gemv_f32")
3100 };
3101 let (ne, nx, ti) = (n_embd as i32, n_experts as i32, t as i32);
3102 let cfg = LaunchConfig {
3103 grid_dim: (n_experts as u32, t as u32, 1),
3104 block_dim: (32, if w8 { 8 } else { 1 }, 1),
3105 shared_mem_bytes: 0,
3106 };
3107 let __s_b = self.gpu.stream();
3108 let mut b = __s_b.launch_builder(&f);
3109 b.arg(w).arg(x).arg(&mut *y).arg(&ne).arg(&nx).arg(&ti);
3110 unsafe {
3111 b.launch(cfg)?;
3112 }
3113 Ok(())
3114 }
3115
3116 pub fn rows_permute(
3118 &self,
3119 src: &CudaSlice<f32>,
3120 idx: &CudaSlice<i32>,
3121 nrows: usize,
3122 ncols: usize,
3123 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3124 let mut dst = self.alloc_uninit::<f32>(nrows * ncols)?;
3125 let f = self.func("rows_permute_f32");
3126 let (nc, nr) = (ncols as i32, nrows as i32);
3127 let cfg = LaunchConfig {
3128 grid_dim: (nrows as u32, 1, 1),
3129 block_dim: (256, 1, 1),
3130 shared_mem_bytes: 0,
3131 };
3132 let __s_b = self.gpu.stream();
3133 let mut b = __s_b.launch_builder(&f);
3134 b.arg(src).arg(idx).arg(&mut dst).arg(&nc).arg(&nr);
3135 unsafe {
3136 b.launch(cfg)?;
3137 }
3138 Ok(dst)
3139 }
3140
3141 pub fn sigmoid_dot_rows(
3146 &self,
3147 x: &CudaSlice<f32>,
3148 w: &CudaSlice<f32>,
3149 n_embd: usize,
3150 t: usize,
3151 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3152 static OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3155 if *OFF.get_or_init(|| std::env::var("MEMRA_SHEXP_DOT").as_deref() == Ok("0")) {
3156 let gs = self.linear(x, w, t, n_embd, 1)?;
3157 let mut g = self.uninit(t)?;
3158 self.sigmoid(&gs, &mut g, t)?;
3159 return Ok(g);
3160 }
3161 let mut g = self.alloc_uninit::<f32>(t)?;
3167 let f = self.func("sigmoid_dot_rows_f32");
3168 let (ne, ti) = (n_embd as i32, t as i32);
3169 let cfg = LaunchConfig {
3170 grid_dim: (t as u32, 1, 1),
3171 block_dim: (32, 8, 1),
3172 shared_mem_bytes: 0,
3173 };
3174 let __s_b = self.gpu.stream();
3175 let mut b = __s_b.launch_builder(&f);
3176 b.arg(x).arg(w).arg(&mut g).arg(&ne).arg(&ti);
3177 unsafe {
3178 b.launch(cfg)?;
3179 }
3180 Ok(g)
3181 }
3182
3183 pub fn sigmoid_dot_rows_into(
3185 &self,
3186 x: &CudaSlice<f32>,
3187 w: &CudaSlice<f32>,
3188 g: &mut CudaSlice<f32>,
3189 n_embd: usize,
3190 t: usize,
3191 ) -> Result<(), Box<dyn std::error::Error>> {
3192 if g.len() < t {
3193 return Err("sigmoid_dot_rows_into output too small".into());
3194 }
3195 let f = self.func("sigmoid_dot_rows_f32");
3196 let (ne, ti) = (n_embd as i32, t as i32);
3197 let cfg = LaunchConfig {
3198 grid_dim: (t as u32, 1, 1),
3199 block_dim: (32, 8, 1),
3200 shared_mem_bytes: 0,
3201 };
3202 let __s_b = self.gpu.stream();
3203 let mut b = __s_b.launch_builder(&f);
3204 b.arg(x).arg(w).arg(&mut *g).arg(&ne).arg(&ti);
3205 unsafe {
3206 b.launch(cfg)?;
3207 }
3208 Ok(())
3209 }
3210
3211 pub fn spec_rollback_stream(
3213 &self,
3214 len_ptrs: &CudaSlice<u64>,
3215 pos_start: &CudaSlice<i32>,
3216 acc: &CudaSlice<u32>,
3217 base: usize,
3218 n_rows: usize,
3219 ) -> Result<(), Box<dyn std::error::Error>> {
3220 let f = self.func("spec_rollback_stream");
3221 let (b, nr) = (base as i32, n_rows as i32);
3222 let cfg = LaunchConfig {
3223 grid_dim: (n_rows.div_ceil(64) as u32, 1, 1),
3224 block_dim: (64, 1, 1),
3225 shared_mem_bytes: 0,
3226 };
3227 let __s_bl = self.gpu.stream();
3228 let mut bl = __s_bl.launch_builder(&f);
3229 bl.arg(len_ptrs).arg(pos_start).arg(acc).arg(&b).arg(&nr);
3230 unsafe {
3231 bl.launch(cfg)?;
3232 }
3233 Ok(())
3234 }
3235
3236 pub fn plain_tok_ring(
3238 &self,
3239 vam: &CudaSlice<u32>,
3240 pos_start: &CudaSlice<i32>,
3241 base: usize,
3242 ring: &mut CudaSlice<u32>,
3243 ) -> Result<(), Box<dyn std::error::Error>> {
3244 let f = self.func("plain_tok_ring");
3245 let (b, cap) = (base as i32, ring.len() as i32);
3246 let cfg = LaunchConfig {
3247 grid_dim: (1, 1, 1),
3248 block_dim: (32, 1, 1),
3249 shared_mem_bytes: 0,
3250 };
3251 let __s_bl = self.gpu.stream();
3252 let mut bl = __s_bl.launch_builder(&f);
3253 bl.arg(vam).arg(pos_start).arg(&b).arg(&mut *ring).arg(&cap);
3254 unsafe {
3255 bl.launch(cfg)?;
3256 }
3257 Ok(())
3258 }
3259
3260 pub fn spec_ring_commit(
3262 &self,
3263 vtok: &CudaSlice<u32>,
3264 acc: &CudaSlice<u32>,
3265 brk: &CudaSlice<u32>,
3266 ring: &mut CudaSlice<u32>,
3267 pend: &mut CudaSlice<u32>,
3268 ) -> Result<(), Box<dyn std::error::Error>> {
3269 let f = self.func("spec_ring_commit");
3270 let cfg = LaunchConfig {
3271 grid_dim: (1, 1, 1),
3272 block_dim: (32, 1, 1),
3273 shared_mem_bytes: 0,
3274 };
3275 let __s_b = self.gpu.stream();
3276 let mut b = __s_b.launch_builder(&f);
3277 b.arg(vtok).arg(acc).arg(brk).arg(ring).arg(pend);
3278 unsafe {
3279 b.launch(cfg)?;
3280 }
3281 Ok(())
3282 }
3283 pub fn i32_copy_add(
3284 &self,
3285 src: &CudaSlice<i32>,
3286 dst: &mut CudaSlice<i32>,
3287 delta: i32,
3288 ) -> Result<(), Box<dyn std::error::Error>> {
3289 let f = self.func("i32_copy_add");
3290 let cfg = LaunchConfig {
3291 grid_dim: (1, 1, 1),
3292 block_dim: (32, 1, 1),
3293 shared_mem_bytes: 0,
3294 };
3295 let __s_b = self.gpu.stream();
3296 let mut b = __s_b.launch_builder(&f);
3297 b.arg(src).arg(dst).arg(&delta);
3298 unsafe {
3299 b.launch(cfg)?;
3300 }
3301 Ok(())
3302 }
3303 pub fn u32_copy(
3304 &self,
3305 src: &CudaSlice<u32>,
3306 dst: &mut CudaSlice<u32>,
3307 ) -> Result<(), Box<dyn std::error::Error>> {
3308 let f = self.func("u32_copy");
3309 let cfg = LaunchConfig {
3310 grid_dim: (1, 1, 1),
3311 block_dim: (32, 1, 1),
3312 shared_mem_bytes: 0,
3313 };
3314 let __s_b = self.gpu.stream();
3315 let mut b = __s_b.launch_builder(&f);
3316 b.arg(src).arg(dst);
3317 unsafe {
3318 b.launch(cfg)?;
3319 }
3320 Ok(())
3321 }
3322
3323 pub fn spec_adapt_k(
3327 &self,
3328 acc: &CudaSlice<u32>,
3329 brk: &mut CudaSlice<u32>,
3330 floor: usize,
3331 cap: usize,
3332 ) -> Result<(), Box<dyn std::error::Error>> {
3333 let f = self.func("spec_adapt_k");
3334 let (fl, cp) = (floor as i32, cap as i32);
3335 let cfg = LaunchConfig {
3336 grid_dim: (1, 1, 1),
3337 block_dim: (32, 1, 1),
3338 shared_mem_bytes: 0,
3339 };
3340 let __s_b = self.gpu.stream();
3341 let mut b = __s_b.launch_builder(&f);
3342 b.arg(acc).arg(brk).arg(&fl).arg(&cp);
3343 unsafe {
3344 b.launch(cfg)?;
3345 }
3346 Ok(())
3347 }
3348
3349 pub fn spec_accept_greedy_dc(
3351 &self,
3352 preds: &CudaSlice<u32>,
3353 vtok: &CudaSlice<u32>,
3354 last_pred: &CudaSlice<u32>,
3355 brk: &CudaSlice<u32>,
3356 out: &mut CudaSlice<u32>,
3357 ) -> Result<(), Box<dyn std::error::Error>> {
3358 let f = self.func("spec_accept_greedy_dc");
3359 let cfg = LaunchConfig {
3360 grid_dim: (1, 1, 1),
3361 block_dim: (32, 1, 1),
3362 shared_mem_bytes: 0,
3363 };
3364 let __s_b = self.gpu.stream();
3365 let mut b = __s_b.launch_builder(&f);
3366 b.arg(preds).arg(vtok).arg(last_pred).arg(brk).arg(out);
3367 unsafe {
3368 b.launch(cfg)?;
3369 }
3370 Ok(())
3371 }
3372
3373 pub fn pos_iota(
3375 &self,
3376 pos0: &CudaSlice<i32>,
3377 out: &mut CudaSlice<i32>,
3378 t: usize,
3379 ) -> Result<(), Box<dyn std::error::Error>> {
3380 let f = self.func("pos_iota_i32");
3381 let ti = t as i32;
3382 let cfg = LaunchConfig {
3383 grid_dim: (1, 1, 1),
3384 block_dim: (t.max(1) as u32, 1, 1),
3385 shared_mem_bytes: 0,
3386 };
3387 let __s_b = self.gpu.stream();
3388 let mut b = __s_b.launch_builder(&f);
3389 b.arg(pos0).arg(out).arg(&ti);
3390 unsafe {
3391 b.launch(cfg)?;
3392 }
3393 Ok(())
3394 }
3395 #[allow(clippy::too_many_arguments)]
3396 pub fn append_kv_quantized_rows_dc(
3397 &self,
3398 k_rows: &CudaSlice<f32>,
3399 v_rows: &CudaSlice<f32>,
3400 kc: &mut CudaSlice<u8>,
3401 vc: &mut CudaSlice<u8>,
3402 t0_dev: &CudaSlice<i32>,
3403 t: usize,
3404 kv_dim_k: usize,
3405 kv_dim_v: usize,
3406 k_tok_bytes: usize,
3407 v_tok_bytes: usize,
3408 g: bool,
3409 ) -> Result<(), Box<dyn std::error::Error>> {
3410 let f = if g {
3411 self.func_g("append_quantize_kv_q8_0_q5_1_rows_dc")
3412 } else {
3413 self.func("append_quantize_kv_q8_0_q5_1_rows_dc")
3414 };
3415 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
3416 let cfg = LaunchConfig {
3417 grid_dim: (nblk, t as u32, 1),
3418 block_dim: (32, 1, 1),
3419 shared_mem_bytes: 0,
3420 };
3421 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
3422 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
3423 let __s_b = self.gpu.stream();
3424 let mut b = __s_b.launch_builder(&f);
3425 b.arg(k_rows)
3426 .arg(v_rows)
3427 .arg(kc)
3428 .arg(vc)
3429 .arg(t0_dev)
3430 .arg(&kdk)
3431 .arg(&kdv)
3432 .arg(&ktb)
3433 .arg(&vtb);
3434 unsafe {
3435 b.launch(cfg)?;
3436 }
3437 Ok(())
3438 }
3439
3440 #[allow(clippy::too_many_arguments)]
3443 pub fn append_kv_quantized_row_dc_inc(
3444 &self,
3445 k_row: &CudaSlice<f32>,
3446 v_row: &CudaSlice<f32>,
3447 kc: &mut CudaSlice<u8>,
3448 vc: &mut CudaSlice<u8>,
3449 t0_dev: &mut CudaSlice<i32>,
3450 kv_dim_k: usize,
3451 kv_dim_v: usize,
3452 k_tok_bytes: usize,
3453 v_tok_bytes: usize,
3454 g: bool,
3455 ) -> Result<(), Box<dyn std::error::Error>> {
3456 let f = if g {
3457 self.func_g("append_quantize_kv_q8_0_q5_1_dc_inc")
3458 } else {
3459 self.func("append_quantize_kv_q8_0_q5_1_dc_inc")
3460 };
3461 let nthreads = ((kv_dim_k.max(kv_dim_v) / 32) * 32).min(1024) as u32;
3462 let cfg = LaunchConfig {
3463 grid_dim: (1, 1, 1),
3464 block_dim: (nthreads, 1, 1),
3465 shared_mem_bytes: 0,
3466 };
3467 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
3468 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
3469 let __s_b = self.gpu.stream();
3470 let mut b = __s_b.launch_builder(&f);
3471 b.arg(k_row)
3472 .arg(v_row)
3473 .arg(kc)
3474 .arg(vc)
3475 .arg(t0_dev)
3476 .arg(&kdk)
3477 .arg(&kdv)
3478 .arg(&ktb)
3479 .arg(&vtb);
3480 unsafe {
3481 b.launch(cfg)?;
3482 }
3483 Ok(())
3484 }
3485
3486 pub fn pack_tok_p(
3488 &self,
3489 tok: &CudaSlice<u32>,
3490 p: &CudaSlice<f32>,
3491 out: &mut CudaSlice<u32>,
3492 slot: usize,
3493 ) -> Result<(), Box<dyn std::error::Error>> {
3494 let f = self.func("pack_tok_p");
3495 let sl = slot as i32;
3496 let cfg = LaunchConfig {
3497 grid_dim: (1, 1, 1),
3498 block_dim: (32, 1, 1),
3499 shared_mem_bytes: 0,
3500 };
3501 let __s_b = self.gpu.stream();
3502 let mut b = __s_b.launch_builder(&f);
3503 b.arg(tok).arg(p).arg(out).arg(&sl);
3504 unsafe {
3505 b.launch(cfg)?;
3506 }
3507 Ok(())
3508 }
3509 pub fn tok_map_u32(
3510 &self,
3511 tok: &mut CudaSlice<u32>,
3512 map: &CudaSlice<u32>,
3513 ) -> Result<(), Box<dyn std::error::Error>> {
3514 let f = self.func("tok_map_u32");
3515 let cfg = LaunchConfig {
3516 grid_dim: (1, 1, 1),
3517 block_dim: (32, 1, 1),
3518 shared_mem_bytes: 0,
3519 };
3520 let __s_b = self.gpu.stream();
3521 let mut b = __s_b.launch_builder(&f);
3522 b.arg(tok).arg(map);
3523 unsafe {
3524 b.launch(cfg)?;
3525 }
3526 Ok(())
3527 }
3528
3529 #[allow(clippy::too_many_arguments)]
3531 pub fn spec_assemble_verify(
3532 &self,
3533 tokp: &CudaSlice<u32>,
3534 pend: &CudaSlice<u32>,
3535 d2t: Option<&CudaSlice<u32>>,
3536 vtok: &mut CudaSlice<u32>,
3537 brk: &mut CudaSlice<u32>,
3538 p_min: f32,
3539 k: usize,
3540 pmin0: bool,
3541 ) -> Result<(), Box<dyn std::error::Error>> {
3542 let f = self.func("spec_assemble_verify");
3543 let (ki, pm) = (k as i32, if pmin0 { 1i32 } else { 0i32 });
3544 let cfg = LaunchConfig {
3545 grid_dim: (1, 1, 1),
3546 block_dim: (32, 1, 1),
3547 shared_mem_bytes: 0,
3548 };
3549 let __s_b = self.gpu.stream();
3550 let mut b = __s_b.launch_builder(&f);
3551 match d2t {
3552 Some(m) => {
3553 b.arg(tokp)
3554 .arg(pend)
3555 .arg(m)
3556 .arg(vtok)
3557 .arg(brk)
3558 .arg(&p_min)
3559 .arg(&ki)
3560 .arg(&pm);
3561 unsafe {
3562 b.launch(cfg)?;
3563 }
3564 }
3565 None => {
3566 let null: u64 = 0;
3567 b.arg(tokp)
3568 .arg(pend)
3569 .arg(&null)
3570 .arg(vtok)
3571 .arg(brk)
3572 .arg(&p_min)
3573 .arg(&ki)
3574 .arg(&pm);
3575 unsafe {
3576 b.launch(cfg)?;
3577 }
3578 }
3579 }
3580 Ok(())
3581 }
3582
3583 #[allow(clippy::too_many_arguments)]
3585 pub fn ssm_conv_ring_rebuild_dc(
3586 &self,
3587 qkv_tm: &CudaSlice<f32>,
3588 ring_old: &CudaSlice<f32>,
3589 conv_state: &mut CudaSlice<f32>,
3590 conv_dim: usize,
3591 acc: &CudaSlice<u32>,
3592 base: usize,
3593 t_v: usize,
3594 d_conv: usize,
3595 ) -> Result<(), Box<dyn std::error::Error>> {
3596 let f = self.func("ssm_conv_ring_rebuild_f32_dc");
3597 let n = conv_dim * (d_conv - 1);
3598 let cfg = LaunchConfig::for_num_elems(n as u32);
3599 let (cd, b0, tv, dc) = (conv_dim as i32, base as i32, t_v as i32, d_conv as i32);
3600 let __s_b = self.gpu.stream();
3601 let mut b = __s_b.launch_builder(&f);
3602 b.arg(qkv_tm)
3603 .arg(ring_old)
3604 .arg(conv_state)
3605 .arg(&cd)
3606 .arg(acc)
3607 .arg(&b0)
3608 .arg(&tv)
3609 .arg(&dc);
3610 unsafe {
3611 b.launch(cfg)?;
3612 }
3613 Ok(())
3614 }
3615 #[allow(clippy::too_many_arguments)]
3616 pub fn gdn_scan_s128_dc(
3617 &self,
3618 q: &CudaSlice<f32>,
3619 k: &CudaSlice<f32>,
3620 v: &CudaSlice<f32>,
3621 g: &CudaSlice<f32>,
3622 beta: &CudaSlice<f32>,
3623 state_in: &CudaSlice<f32>,
3624 state_out: &mut CudaSlice<f32>,
3625 o: &mut CudaSlice<f32>,
3626 n_head: usize,
3627 acc: &CudaSlice<u32>,
3628 base: usize,
3629 t_v: usize,
3630 scale: f32,
3631 ) -> Result<(), Box<dyn std::error::Error>> {
3632 let f = self.func("gdn_scan_s128_dc");
3633 const S_V: u32 = 128;
3634 const WARP: u32 = 32;
3635 const COLS_PER_BLOCK: u32 = 4;
3636 let cfg = LaunchConfig {
3637 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
3638 block_dim: (WARP, COLS_PER_BLOCK, 1),
3639 shared_mem_bytes: 0,
3640 };
3641 let (h, b0, tv) = (n_head as i32, base as i32, t_v as i32);
3642 let __s_b = self.gpu.stream();
3643 let mut b = __s_b.launch_builder(&f);
3644 b.arg(q)
3645 .arg(k)
3646 .arg(v)
3647 .arg(g)
3648 .arg(beta)
3649 .arg(state_in)
3650 .arg(state_out)
3651 .arg(o)
3652 .arg(&h)
3653 .arg(acc)
3654 .arg(&b0)
3655 .arg(&tv)
3656 .arg(&scale);
3657 unsafe {
3658 b.launch(cfg)?;
3659 }
3660 Ok(())
3661 }
3662
3663 pub fn spec_rollback_kv(
3665 &self,
3666 len_ptrs: &CudaSlice<u64>,
3667 saved: &CudaSlice<i32>,
3668 acc: &CudaSlice<u32>,
3669 base: usize,
3670 n_layer: usize,
3671 ) -> Result<(), Box<dyn std::error::Error>> {
3672 let f = self.func("spec_rollback_kv");
3673 let (b, nl) = (base as i32, n_layer as i32);
3674 let cfg = LaunchConfig {
3675 grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
3676 block_dim: (64, 1, 1),
3677 shared_mem_bytes: 0,
3678 };
3679 let __s_bl = self.gpu.stream();
3680 let mut bl = __s_bl.launch_builder(&f);
3681 bl.arg(len_ptrs).arg(saved).arg(acc).arg(&b).arg(&nl);
3682 unsafe {
3683 bl.launch(cfg)?;
3684 }
3685 Ok(())
3686 }
3687
3688 pub fn spec_fork_valid(
3690 &self,
3691 acc: &CudaSlice<u32>,
3692 optimistic_pending: u32,
3693 valid: &mut CudaSlice<u32>,
3694 ) -> Result<(), Box<dyn std::error::Error>> {
3695 let f = self.func("spec_fork_valid");
3696 let cfg = LaunchConfig {
3697 grid_dim: (1, 1, 1),
3698 block_dim: (1, 1, 1),
3699 shared_mem_bytes: 0,
3700 };
3701 let __s_bl = self.gpu.stream();
3702 let mut bl = __s_bl.launch_builder(&f);
3703 bl.arg(acc).arg(&optimistic_pending).arg(valid);
3704 unsafe {
3705 bl.launch(cfg)?;
3706 }
3707 Ok(())
3708 }
3709
3710 pub fn spec_fork_reconcile_kv(
3712 &self,
3713 len_ptrs: &CudaSlice<u64>,
3714 saved: &CudaSlice<i32>,
3715 acc: &CudaSlice<u32>,
3716 valid: &CudaSlice<u32>,
3717 base: usize,
3718 n_layer: usize,
3719 ) -> Result<(), Box<dyn std::error::Error>> {
3720 let f = self.func("spec_fork_reconcile_kv");
3721 let (b, nl) = (base as i32, n_layer as i32);
3722 let cfg = LaunchConfig {
3723 grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
3724 block_dim: (64, 1, 1),
3725 shared_mem_bytes: 0,
3726 };
3727 let __s_bl = self.gpu.stream();
3728 let mut bl = __s_bl.launch_builder(&f);
3729 bl.arg(len_ptrs)
3730 .arg(saved)
3731 .arg(acc)
3732 .arg(valid)
3733 .arg(&b)
3734 .arg(&nl);
3735 unsafe {
3736 bl.launch(cfg)?;
3737 }
3738 Ok(())
3739 }
3740
3741 pub fn spec_fork_restore_f32(
3743 &self,
3744 snapshot: &CudaSlice<f32>,
3745 state: &mut CudaSlice<f32>,
3746 valid: &CudaSlice<u32>,
3747 ) -> Result<(), Box<dyn std::error::Error>> {
3748 assert_eq!(
3749 snapshot.len(),
3750 state.len(),
3751 "fork recurrent snapshot shape mismatch"
3752 );
3753 let f = self.func("spec_fork_restore_f32");
3754 let n = state.len() as i32;
3755 let blocks = state.len().div_ceil(256).min(65535).max(1) as u32;
3756 let cfg = LaunchConfig {
3757 grid_dim: (blocks, 1, 1),
3758 block_dim: (256, 1, 1),
3759 shared_mem_bytes: 0,
3760 };
3761 let __s_bl = self.gpu.stream();
3762 let mut bl = __s_bl.launch_builder(&f);
3763 bl.arg(snapshot).arg(state).arg(valid).arg(&n);
3764 unsafe {
3765 bl.launch(cfg)?;
3766 }
3767 Ok(())
3768 }
3769
3770 pub fn spec_seed_gather(
3773 &self,
3774 vx: &CudaSlice<f32>,
3775 fill_prev: &CudaSlice<f32>,
3776 acc: &CudaSlice<u32>,
3777 h_seed: &mut CudaSlice<f32>,
3778 base: usize,
3779 n_embd: usize,
3780 ) -> Result<(), Box<dyn std::error::Error>> {
3781 let f = self.func("spec_seed_gather");
3782 let (b, ne) = (base as i32, n_embd as i32);
3783 let cfg = LaunchConfig {
3784 grid_dim: (n_embd.div_ceil(256) as u32, 1, 1),
3785 block_dim: (256, 1, 1),
3786 shared_mem_bytes: 0,
3787 };
3788 let __s_bl = self.gpu.stream();
3789 let mut bl = __s_bl.launch_builder(&f);
3790 bl.arg(vx)
3791 .arg(fill_prev)
3792 .arg(acc)
3793 .arg(h_seed)
3794 .arg(&b)
3795 .arg(&ne);
3796 unsafe {
3797 bl.launch(cfg)?;
3798 }
3799 Ok(())
3800 }
3801
3802 pub fn spec_accept_greedy(
3804 &self,
3805 preds: &CudaSlice<u32>,
3806 draft: &CudaSlice<u32>,
3807 last_pred: u32,
3808 base: usize,
3809 k_round: usize,
3810 out: &mut CudaSlice<u32>,
3811 ) -> Result<(), Box<dyn std::error::Error>> {
3812 let f = self.func("spec_accept_greedy");
3813 let (b, k) = (base as i32, k_round as i32);
3814 let cfg = LaunchConfig {
3815 grid_dim: (1, 1, 1),
3816 block_dim: (32, 1, 1),
3817 shared_mem_bytes: 0,
3818 };
3819 let __s_bl = self.gpu.stream();
3820 let mut bl = __s_bl.launch_builder(&f);
3821 bl.arg(preds)
3822 .arg(draft)
3823 .arg(&last_pred)
3824 .arg(&b)
3825 .arg(&k)
3826 .arg(out);
3827 unsafe {
3828 bl.launch(cfg)?;
3829 }
3830 Ok(())
3831 }
3832
3833 pub fn gumbel_perturb(
3840 &self,
3841 x: &CudaSlice<f32>,
3842 y: &mut CudaSlice<f32>,
3843 n: usize,
3844 seed: u64,
3845 stream_pos: u32,
3846 temp: f32,
3847 ) -> Result<(), Box<dyn std::error::Error>> {
3848 let f = self.func("gumbel_perturb_f32");
3849 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
3850 let cfg = LaunchConfig {
3851 grid_dim: (n.div_ceil(256) as u32, 1, 1),
3852 block_dim: (256, 1, 1),
3853 shared_mem_bytes: 0,
3854 };
3855 let __s_b = self.gpu.stream();
3856 let mut b = __s_b.launch_builder(&f);
3857 b.arg(x)
3858 .arg(&mut *y)
3859 .arg(&ni)
3860 .arg(&slo)
3861 .arg(&shi)
3862 .arg(&stream_pos)
3863 .arg(&temp);
3864 unsafe {
3865 b.launch(cfg)?;
3866 }
3867 Ok(())
3868 }
3869
3870 pub fn mask_logits_col(
3878 &self,
3879 logits: &mut CudaSlice<f32>,
3880 mask: &CudaSlice<u32>,
3881 col: usize,
3882 n: usize,
3883 mask_words: usize,
3884 ) -> Result<(), Box<dyn std::error::Error>> {
3885 let f = self.func("mask_logits_f32");
3886 let (ci, ni, mw) = (col as i32, n as i32, mask_words as i32);
3887 let cfg = LaunchConfig {
3888 grid_dim: (n.div_ceil(256).min(1024) as u32, 1, 1),
3889 block_dim: (256, 1, 1),
3890 shared_mem_bytes: 0,
3891 };
3892 let __s_b = self.gpu.stream();
3893 let mut b = __s_b.launch_builder(&f);
3894 b.arg(&mut *logits).arg(mask).arg(&ci).arg(&ni).arg(&mw);
3895 unsafe {
3896 b.launch(cfg)?;
3897 }
3898 Ok(())
3899 }
3900
3901 pub fn gumbel_perturb_col(
3908 &self,
3909 x: &CudaSlice<f32>,
3910 col: usize,
3911 y: &mut CudaSlice<f32>,
3912 n: usize,
3913 seed: u64,
3914 stream_pos: u32,
3915 temp: f32,
3916 ) -> Result<(), Box<dyn std::error::Error>> {
3917 let f = self.func("gumbel_perturb_f32");
3918 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
3919 let col_view = x.slice(col * n..(col + 1) * n);
3920 let cfg = LaunchConfig {
3921 grid_dim: (n.div_ceil(256) as u32, 1, 1),
3922 block_dim: (256, 1, 1),
3923 shared_mem_bytes: 0,
3924 };
3925 let __s_b = self.gpu.stream();
3926 let mut b = __s_b.launch_builder(&f);
3927 b.arg(&col_view)
3928 .arg(&mut *y)
3929 .arg(&ni)
3930 .arg(&slo)
3931 .arg(&shi)
3932 .arg(&stream_pos)
3933 .arg(&temp);
3934 unsafe {
3935 b.launch(cfg)?;
3936 }
3937 Ok(())
3938 }
3939
3940 #[allow(clippy::too_many_arguments)]
3946 pub fn gumbel_perturb_filtered_col(
3947 &self,
3948 x: &CudaSlice<f32>,
3949 col: usize,
3950 y: &mut CudaSlice<f32>,
3951 n: usize,
3952 seed: u64,
3953 stream_pos: u32,
3954 temp: f32,
3955 stat_max: &CudaSlice<f32>,
3956 stat_th: &CudaSlice<f32>,
3957 stat_idx: usize,
3958 ) -> Result<(), Box<dyn std::error::Error>> {
3959 let f = self.func("gumbel_perturb_filtered_col_f32");
3960 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
3961 let (ci, si) = (col as i32, stat_idx as i32);
3962 let cfg = LaunchConfig {
3963 grid_dim: (n.div_ceil(256) as u32, 1, 1),
3964 block_dim: (256, 1, 1),
3965 shared_mem_bytes: 0,
3966 };
3967 let __s_b = self.gpu.stream();
3968 let mut b = __s_b.launch_builder(&f);
3969 b.arg(x)
3970 .arg(&ci)
3971 .arg(&mut *y)
3972 .arg(&ni)
3973 .arg(&slo)
3974 .arg(&shi)
3975 .arg(&stream_pos)
3976 .arg(&temp)
3977 .arg(stat_max)
3978 .arg(stat_th)
3979 .arg(&si);
3980 unsafe {
3981 b.launch(cfg)?;
3982 }
3983 Ok(())
3984 }
3985
3986 pub fn sctr_inc(&self, ctr: &mut CudaSlice<u32>) -> Result<(), Box<dyn std::error::Error>> {
3991 let f = self.func("memra_sctr_inc");
3992 let cfg = LaunchConfig {
3993 grid_dim: (1, 1, 1),
3994 block_dim: (1, 1, 1),
3995 shared_mem_bytes: 0,
3996 };
3997 let __s_b = self.gpu.stream();
3998 let mut b = __s_b.launch_builder(&f);
3999 b.arg(&mut *ctr);
4000 unsafe {
4001 b.launch(cfg)?;
4002 }
4003 Ok(())
4004 }
4005
4006 pub fn gumbel_perturb_ctr(
4011 &self,
4012 x: &CudaSlice<f32>,
4013 y: &mut CudaSlice<f32>,
4014 n: usize,
4015 seed: u64,
4016 ctr: &CudaSlice<u32>,
4017 temp: f32,
4018 ) -> Result<(), Box<dyn std::error::Error>> {
4019 let f = self.func("gumbel_perturb_ctr_f32");
4020 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
4021 let cfg = LaunchConfig {
4022 grid_dim: (n.div_ceil(256) as u32, 1, 1),
4023 block_dim: (256, 1, 1),
4024 shared_mem_bytes: 0,
4025 };
4026 let __s_b = self.gpu.stream();
4027 let mut b = __s_b.launch_builder(&f);
4028 b.arg(x)
4029 .arg(&mut *y)
4030 .arg(&ni)
4031 .arg(&slo)
4032 .arg(&shi)
4033 .arg(ctr)
4034 .arg(&temp);
4035 unsafe {
4036 b.launch(cfg)?;
4037 }
4038 Ok(())
4039 }
4040
4041 #[allow(clippy::too_many_arguments)]
4049 pub fn gumbel_perturb_filtered_ctr(
4050 &self,
4051 x: &CudaSlice<f32>,
4052 y: &mut CudaSlice<f32>,
4053 n: usize,
4054 seed: u64,
4055 ctr: &CudaSlice<u32>,
4056 temp: f32,
4057 stat_max: &CudaSlice<f32>,
4058 stat_th: &CudaSlice<f32>,
4059 ) -> Result<(), Box<dyn std::error::Error>> {
4060 let f = self.func("gumbel_perturb_filtered_ctr_f32");
4061 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
4062 let cfg = LaunchConfig {
4063 grid_dim: (n.div_ceil(256) as u32, 1, 1),
4064 block_dim: (256, 1, 1),
4065 shared_mem_bytes: 0,
4066 };
4067 let __s_b = self.gpu.stream();
4068 let mut b = __s_b.launch_builder(&f);
4069 b.arg(x)
4070 .arg(&mut *y)
4071 .arg(&ni)
4072 .arg(&slo)
4073 .arg(&shi)
4074 .arg(ctr)
4075 .arg(&temp)
4076 .arg(stat_max)
4077 .arg(stat_th);
4078 unsafe {
4079 b.launch(cfg)?;
4080 }
4081 Ok(())
4082 }
4083
4084 pub fn softmax_gather(
4088 &self,
4089 x: &CudaSlice<f32>,
4090 row_stride: usize,
4091 ids: &CudaSlice<u32>,
4092 rows: &CudaSlice<i32>,
4093 out: &mut CudaSlice<f32>,
4094 n: usize,
4095 npair: usize,
4096 temp: f32,
4097 ) -> Result<(), Box<dyn std::error::Error>> {
4098 let f = self.func("softmax_gather_f32");
4099 let (ni, rs) = (n as i32, row_stride as i64);
4100 let np = npair as i32;
4101 let cfg = LaunchConfig {
4102 grid_dim: (npair as u32, 1, 1),
4103 block_dim: (256, 1, 1),
4104 shared_mem_bytes: 0,
4105 };
4106 let __s_b = self.gpu.stream();
4107 let mut b = __s_b.launch_builder(&f);
4108 b.arg(x)
4109 .arg(&rs)
4110 .arg(ids)
4111 .arg(rows)
4112 .arg(&mut *out)
4113 .arg(&ni)
4114 .arg(&np)
4115 .arg(&temp);
4116 unsafe {
4117 b.launch(cfg)?;
4118 }
4119 Ok(())
4120 }
4121
4122 pub fn residual_sample(
4126 &self,
4127 p: &CudaSlice<f32>,
4128 q: Option<&CudaSlice<f32>>,
4129 n: usize,
4130 temp: f32,
4131 seed: u64,
4132 stream_pos: u32,
4133 out_tok: &mut CudaSlice<u32>,
4134 ) -> Result<(), Box<dyn std::error::Error>> {
4135 let f = self.func("residual_sample_f32");
4136 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
4137 let nth = 1024u32;
4138 let cfg = LaunchConfig {
4139 grid_dim: (1, 1, 1),
4140 block_dim: (nth, 1, 1),
4141 shared_mem_bytes: 0,
4142 };
4143 let has_q: i32 = q.is_some() as i32;
4144 let qbuf = q.unwrap_or(p); let __s_b = self.gpu.stream();
4146 let mut b = __s_b.launch_builder(&f);
4147 b.arg(p)
4148 .arg(qbuf)
4149 .arg(&has_q)
4150 .arg(&ni)
4151 .arg(&temp)
4152 .arg(&slo)
4153 .arg(&shi)
4154 .arg(&stream_pos)
4155 .arg(&mut *out_tok);
4156 unsafe {
4157 b.launch(cfg)?;
4158 }
4159 Ok(())
4160 }
4161
4162 pub fn with_moe_cache<R>(
4167 &self,
4168 max_block_bytes: usize,
4169 f: impl FnOnce(
4170 &mut crate::moe_cache::MoeSlotCache,
4171 &Engine,
4172 ) -> Result<R, Box<dyn std::error::Error>>,
4173 ) -> Result<R, Box<dyn std::error::Error>> {
4174 let mut guard = self.moe_cache.lock().unwrap();
4175 if guard.is_none() {
4176 *guard = Some(crate::moe_cache::MoeSlotCache::new(self, max_block_bytes)?);
4177 }
4178 let cache = guard.as_mut().unwrap();
4179 f(cache, self)
4180 }
4181
4182 pub fn freeze_moe_cache(&self) {
4185 if let Some(cache) = self.moe_cache.lock().unwrap().as_mut() {
4186 cache.freeze();
4187 }
4188 }
4189
4190 pub fn export_moe_residency(&self) -> Option<Vec<(u16, u8, u16)>> {
4193 self.moe_cache
4194 .lock()
4195 .unwrap()
4196 .as_ref()
4197 .map(crate::moe_cache::MoeSlotCache::export_residency)
4198 }
4199
4200 pub(crate) fn moe_cache_frozen(&self) -> bool {
4201 self.moe_cache
4202 .lock()
4203 .unwrap()
4204 .as_ref()
4205 .is_some_and(crate::moe_cache::MoeSlotCache::is_frozen)
4206 }
4207
4208 pub fn frozen_cpu_experts_prefer_tokenwise_prime(&self) -> bool {
4215 crate::cpu_experts::configured()
4216 && self.moe_cache_frozen()
4217 && std::env::var("MEMRA_CPU_EXPERT_BATCHED_PRIME").as_deref() != Ok("1")
4218 }
4219
4220 pub(crate) fn configure_moe_cache_layout(&self, block_bytes: Vec<usize>) {
4222 assert!(
4223 self.moe_cache.lock().unwrap().is_none(),
4224 "MoE cache layout configured after cache construction"
4225 );
4226 *self.moe_cache_layout.lock().unwrap() = Some(block_bytes);
4227 }
4228
4229 pub(crate) fn moe_cache_layout(&self) -> Option<Vec<usize>> {
4230 self.moe_cache_layout.lock().unwrap().clone()
4231 }
4232
4233 pub fn moe_cache_enabled() -> bool {
4235 std::env::var("MEMRA_MOE_CACHE").as_deref() != Ok("0")
4236 }
4237
4238 pub fn moe_cache_stats(&self) -> Option<(u64, u64, u64, usize)> {
4241 let guard = self.moe_cache.lock().unwrap();
4242 guard
4243 .as_ref()
4244 .map(|c| (c.hits, c.misses, c.staged_bytes, c.n_slots()))
4245 }
4246
4247 pub fn cpu_expert_stats(
4251 &self,
4252 ) -> Option<(u64, u64, u64, u64, u64, u64, u64, u64, u64, u64, u64)> {
4253 crate::cpu_experts::configured().then(crate::cpu_experts::stats)
4254 }
4255
4256 pub fn cpu_expert_predictor_stats(&self) -> (u64, u64) {
4259 crate::cpu_experts::predictor_stats()
4260 }
4261
4262 pub fn cpu_expert_exposed_wait_ns(&self) -> Option<u64> {
4263 crate::cpu_experts::configured().then(crate::cpu_experts::exposed_wait_ns)
4264 }
4265
4266 pub fn cpu_expert_gpu_residency_stats(&self) -> Option<(u64, u64, u64)> {
4269 crate::cpu_experts::configured().then(crate::cpu_experts::incomplete_gpu_residency_stats)
4270 }
4271
4272 pub fn moe_pread_stats(&self) -> Option<(u64, u64, u64, u64, u64, u64, u64)> {
4275 let guard = self.moe_cache.lock().unwrap();
4276 guard
4277 .as_ref()
4278 .and_then(|cache| cache.pread_stats())
4279 .map(|stats| {
4280 (
4281 stats.reads,
4282 stats.bytes,
4283 stats.read_errors,
4284 stats.short_reads,
4285 stats.fallbacks,
4286 stats.buffer_waits,
4287 stats.ring_full,
4288 )
4289 })
4290 }
4291
4292 pub fn spill_config_fallbacks(&self) -> u64 {
4294 crate::spill_pread::config_fallbacks()
4295 }
4296
4297 pub fn moe_cache_reset_counters(&self) {
4299 if let Some(c) = self.moe_cache.lock().unwrap().as_mut() {
4300 c.reset_counters();
4301 }
4302 }
4303
4304 pub fn htod_bytes(&self, v: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
4305 Ok(self.gpu.stream().clone_htod(v)?)
4306 }
4307
4308 pub fn htod_bytes_padded(
4312 &self,
4313 v: &[u8],
4314 pad: usize,
4315 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
4316 let mut d = self.alloc_u8_uninit(v.len() + pad)?;
4317 {
4318 let mut view = d.slice_mut(0..v.len());
4319 self.gpu.stream().memcpy_htod(v, &mut view)?;
4320 }
4321 Ok(d)
4322 }
4323
4324 pub fn copy_into(
4326 &self,
4327 dst: &mut CudaSlice<f32>,
4328 off: usize,
4329 src: &CudaSlice<f32>,
4330 len: usize,
4331 ) -> Result<(), Box<dyn std::error::Error>> {
4332 let mut view = dst.slice_mut(off..off + len);
4333 self.gpu
4334 .stream()
4335 .memcpy_dtod(&src.slice(0..len), &mut view)?;
4336 Ok(())
4337 }
4338
4339 pub fn copy_range_into(
4343 &self,
4344 dst: &mut CudaSlice<f32>,
4345 dst_off: usize,
4346 src: &CudaSlice<f32>,
4347 src_off: usize,
4348 len: usize,
4349 ) -> Result<(), Box<dyn std::error::Error>> {
4350 let mut view = dst.slice_mut(dst_off..dst_off + len);
4351 self.gpu
4352 .stream()
4353 .memcpy_dtod(&src.slice(src_off..src_off + len), &mut view)?;
4354 Ok(())
4355 }
4356
4357 pub fn copy_u8_into(
4360 &self,
4361 dst: &mut CudaSlice<u8>,
4362 off: usize,
4363 src: &CudaSlice<u8>,
4364 len: usize,
4365 ) -> Result<(), Box<dyn std::error::Error>> {
4366 let cap = dst.len();
4370 let mut view = dst.try_slice_mut(off..off + len).ok_or_else(|| {
4371 format!(
4372 "copy_u8_into dst range [{off},{}) exceeds capacity {cap}",
4373 off + len,
4374 )
4375 })?;
4376 self.gpu
4377 .stream()
4378 .memcpy_dtod(&src.slice(0..len), &mut view)?;
4379 Ok(())
4380 }
4381
4382 pub fn copy_u8_range_into(
4384 &self,
4385 dst: &mut CudaSlice<u8>,
4386 dst_off: usize,
4387 src: &CudaSlice<u8>,
4388 src_off: usize,
4389 len: usize,
4390 ) -> Result<(), Box<dyn std::error::Error>> {
4391 let cap = dst.len();
4394 let mut dst_view = dst.try_slice_mut(dst_off..dst_off + len).ok_or_else(|| {
4395 format!(
4396 "copy_u8_range_into dst range [{dst_off},{}) exceeds capacity {cap}",
4397 dst_off + len,
4398 )
4399 })?;
4400 self.gpu
4401 .stream()
4402 .memcpy_dtod(&src.slice(src_off..src_off + len), &mut dst_view)?;
4403 Ok(())
4404 }
4405
4406 #[track_caller]
4416 pub fn prepare_kv_append(
4417 &self,
4418 kv: &mut crate::cache::KvLayer,
4419 retain_from: usize,
4420 append_rows: usize,
4421 ) -> Result<usize, Box<dyn std::error::Error>> {
4422 let caller = std::panic::Location::caller();
4423 let base_before = kv.ring.as_ref().map(|r| r.base());
4424 let Some(plan) = kv
4425 .ring
4426 .as_ref()
4427 .map(|ring| ring.append_plan(kv.len, retain_from, append_rows))
4428 .transpose()
4429 .map_err(|err| -> Box<dyn std::error::Error> {
4430 format!(
4431 "{err} [append len={} retain_from={retain_from} append_rows={append_rows} base={base_before:?} called from {caller}]",
4432 kv.len
4433 )
4434 .into()
4435 })?
4436 else {
4437 return Ok(kv.len);
4438 };
4439 match plan {
4440 crate::cache::KvRingAppend::Contiguous { write_row } => Ok(write_row),
4441 crate::cache::KvRingAppend::Rebase {
4442 src_row,
4443 keep_rows,
4444 new_base,
4445 write_row,
4446 } => {
4447 if keep_rows > 0 {
4448 let k_len = keep_rows * kv.k_tok_bytes;
4449 let v_len = keep_rows * kv.v_tok_bytes;
4450 let mut k_tmp = self.alloc_u8_uninit(k_len)?;
4451 let mut v_tmp = self.alloc_u8_uninit(v_len)?;
4452 self.copy_u8_range_into(&mut k_tmp, 0, &kv.k, src_row * kv.k_tok_bytes, k_len)?;
4453 self.copy_u8_range_into(&mut v_tmp, 0, &kv.v, src_row * kv.v_tok_bytes, v_len)?;
4454 self.copy_u8_into(&mut kv.k, 0, &k_tmp, k_len)?;
4455 self.copy_u8_into(&mut kv.v, 0, &v_tmp, v_len)?;
4456 }
4457 if std::env::var("MEMRA_KV_REBASE_TRACE").as_deref() == Ok("1") {
4460 eprintln!(
4461 "[kv-rebase] new_base={new_base} keep_rows={keep_rows} len={} \
4462 retain_from={retain_from} called from {caller}",
4463 kv.len
4464 );
4465 }
4466 kv.ring.as_mut().unwrap().apply_rebase(new_base);
4467 if let Some(base_d) = kv.base_d.as_mut() {
4471 self.set_i32_one(base_d, new_base as i32)?;
4472 }
4473 Ok(write_row)
4474 }
4475 }
4476 }
4477
4478 pub fn htod_u8_into(
4481 &self,
4482 dst: &mut CudaSlice<u8>,
4483 off: usize,
4484 src: &[u8],
4485 ) -> Result<(), Box<dyn std::error::Error>> {
4486 let mut view = dst.slice_mut(off..off + src.len());
4487 self.gpu.stream().memcpy_htod(src, &mut view)?;
4488 Ok(())
4489 }
4490
4491 pub fn view<'a>(&self, b: &'a CudaSlice<f32>, len: usize) -> cudarc::driver::CudaView<'a, f32> {
4492 b.slice(0..len)
4493 }
4494
4495 pub fn view_u8_range<'a>(
4498 &self,
4499 b: &'a CudaSlice<u8>,
4500 start: usize,
4501 end: usize,
4502 ) -> cudarc::driver::CudaView<'a, u8> {
4503 b.slice(start..end)
4504 }
4505 pub fn view_u8<'a>(
4506 &self,
4507 b: &'a CudaSlice<u8>,
4508 len: usize,
4509 ) -> cudarc::driver::CudaView<'a, u8> {
4510 b.slice(0..len)
4511 }
4512
4513 pub fn append_kv_quantized(
4517 &self,
4518 k_row: &CudaSlice<f32>,
4519 v_row: &CudaSlice<f32>,
4520 kc: &mut CudaSlice<u8>,
4521 vc: &mut CudaSlice<u8>,
4522 t: usize,
4523 kv_dim_k: usize,
4524 kv_dim_v: usize,
4525 k_tok_bytes: usize,
4526 v_tok_bytes: usize,
4527 g: bool,
4528 ) -> Result<(), Box<dyn std::error::Error>> {
4529 let f = if g {
4530 self.func_g("append_quantize_kv_q8_0_q5_1")
4531 } else {
4532 self.func("append_quantize_kv_q8_0_q5_1")
4533 };
4534 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
4535 let cfg = LaunchConfig {
4536 grid_dim: (nblk, 1, 1),
4537 block_dim: (32, 1, 1),
4538 shared_mem_bytes: 0,
4539 };
4540 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
4541 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
4542 let __s_b = self.gpu.stream();
4543 let mut b = __s_b.launch_builder(&f);
4544 b.arg(k_row)
4545 .arg(v_row)
4546 .arg(kc)
4547 .arg(vc)
4548 .arg(&ti)
4549 .arg(&kdk)
4550 .arg(&kdv)
4551 .arg(&ktb)
4552 .arg(&vtb);
4553 unsafe {
4554 b.launch(cfg)?;
4555 }
4556 Ok(())
4557 }
4558
4559 pub fn append_kv_quantized_dc(
4563 &self,
4564 k_row: &CudaSlice<f32>,
4565 v_row: &CudaSlice<f32>,
4566 kc: &mut CudaSlice<u8>,
4567 vc: &mut CudaSlice<u8>,
4568 t_dev: &CudaSlice<i32>,
4569 kv_dim_k: usize,
4570 kv_dim_v: usize,
4571 k_tok_bytes: usize,
4572 v_tok_bytes: usize,
4573 g: bool,
4574 ) -> Result<(), Box<dyn std::error::Error>> {
4575 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
4576 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
4577 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
4578 if Self::pdl_on() && Self::pdl_wb_on() {
4580 use cudarc::driver::{DevicePtr, DevicePtrMut};
4581 let s = &self.gpu.stream();
4582 let (pk, _g0) = k_row.device_ptr(s);
4583 let (pv, _g1) = v_row.device_ptr(s);
4584 let (pkc, _g2) = kc.device_ptr_mut(s);
4585 let (pvc, _g3) = vc.device_ptr_mut(s);
4586 let (pt, _g4) = t_dev.device_ptr(s);
4587 let mut ps = [
4588 &pk as *const _ as *mut std::ffi::c_void,
4589 &pv as *const _ as *mut _,
4590 &pkc as *const _ as *mut _,
4591 &pvc as *const _ as *mut _,
4592 &pt as *const _ as *mut _,
4593 &kdk as *const _ as *mut _,
4594 &kdv as *const _ as *mut _,
4595 &ktb as *const _ as *mut _,
4596 &vtb as *const _ as *mut _,
4597 ];
4598 unsafe {
4599 self.launch_pdl_flash(
4600 g,
4601 "append_quantize_kv_q8_0_q5_1_dc",
4602 (nblk, 1, 1),
4603 (32, 1, 1),
4604 0,
4605 &mut ps,
4606 )?;
4607 }
4608 return Ok(());
4609 }
4610 let f = if g {
4611 self.func_g("append_quantize_kv_q8_0_q5_1_dc")
4612 } else {
4613 self.func("append_quantize_kv_q8_0_q5_1_dc")
4614 };
4615 let cfg = LaunchConfig {
4616 grid_dim: (nblk, 1, 1),
4617 block_dim: (32, 1, 1),
4618 shared_mem_bytes: 0,
4619 };
4620 let __s_b = self.gpu.stream();
4621 let mut b = __s_b.launch_builder(&f);
4622 b.arg(k_row)
4623 .arg(v_row)
4624 .arg(kc)
4625 .arg(vc)
4626 .arg(t_dev)
4627 .arg(&kdk)
4628 .arg(&kdv)
4629 .arg(&ktb)
4630 .arg(&vtb);
4631 unsafe {
4632 b.launch(cfg)?;
4633 }
4634 Ok(())
4635 }
4636
4637 #[allow(clippy::too_many_arguments)]
4644 pub fn append_kv_quantized_rows(
4645 &self,
4646 k_rows: &CudaSlice<f32>,
4647 v_rows: &CudaSlice<f32>,
4648 kc: &mut CudaSlice<u8>,
4649 vc: &mut CudaSlice<u8>,
4650 t0: usize,
4651 t: usize,
4652 kv_dim_k: usize,
4653 kv_dim_v: usize,
4654 k_tok_bytes: usize,
4655 v_tok_bytes: usize,
4656 g: bool,
4657 ) -> Result<(), Box<dyn std::error::Error>> {
4658 if std::env::var("MEMRA_PRIME_APPEND_LOOP").is_ok() {
4659 for i in 0..t {
4660 let k_row = k_rows.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
4661 let v_row = v_rows.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
4662 self.append_kv_quantized_view(
4663 &k_row,
4664 &v_row,
4665 kc,
4666 vc,
4667 t0 + i,
4668 kv_dim_k,
4669 kv_dim_v,
4670 k_tok_bytes,
4671 v_tok_bytes,
4672 g,
4673 )?;
4674 }
4675 return Ok(());
4676 }
4677 let f = if g {
4678 self.func_g("append_quantize_kv_q8_0_q5_1_rows")
4679 } else {
4680 self.func("append_quantize_kv_q8_0_q5_1_rows")
4681 };
4682 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
4683 let cfg = LaunchConfig {
4684 grid_dim: (nblk, t as u32, 1),
4685 block_dim: (32, 1, 1),
4686 shared_mem_bytes: 0,
4687 };
4688 let (t0i, kdk, kdv) = (t0 as i32, kv_dim_k as i32, kv_dim_v as i32);
4689 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
4690 let __s_b = self.gpu.stream();
4691 let mut b = __s_b.launch_builder(&f);
4692 b.arg(k_rows)
4693 .arg(v_rows)
4694 .arg(kc)
4695 .arg(vc)
4696 .arg(&t0i)
4697 .arg(&kdk)
4698 .arg(&kdv)
4699 .arg(&ktb)
4700 .arg(&vtb);
4701 unsafe {
4702 b.launch(cfg)?;
4703 }
4704 Ok(())
4705 }
4706
4707 pub fn inc_seqlen(&self, p: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
4711 let f = self.func("inc_i32");
4712 let cfg = LaunchConfig {
4713 grid_dim: (1, 1, 1),
4714 block_dim: (1, 1, 1),
4715 shared_mem_bytes: 0,
4716 };
4717 let __s_b = self.gpu.stream();
4718 let mut b = __s_b.launch_builder(&f);
4719 b.arg(p);
4720 unsafe {
4721 b.launch(cfg)?;
4722 }
4723 Ok(())
4724 }
4725
4726 pub fn append_kv_quantized_view(
4729 &self,
4730 k_row: &cudarc::driver::CudaView<f32>,
4731 v_row: &cudarc::driver::CudaView<f32>,
4732 kc: &mut CudaSlice<u8>,
4733 vc: &mut CudaSlice<u8>,
4734 t: usize,
4735 kv_dim_k: usize,
4736 kv_dim_v: usize,
4737 k_tok_bytes: usize,
4738 v_tok_bytes: usize,
4739 g: bool,
4740 ) -> Result<(), Box<dyn std::error::Error>> {
4741 let stream = self.gpu.stream();
4742 ensure_tensor_stream_device(k_row, &stream, "append_kv_quantized_view.k_row")?;
4743 ensure_tensor_stream_device(v_row, &stream, "append_kv_quantized_view.v_row")?;
4744 ensure_tensor_stream_device(kc, &stream, "append_kv_quantized_view.k_cache")?;
4745 ensure_tensor_stream_device(vc, &stream, "append_kv_quantized_view.v_cache")?;
4746 let f = if g {
4747 self.func_g("append_quantize_kv_q8_0_q5_1")
4748 } else {
4749 self.func("append_quantize_kv_q8_0_q5_1")
4750 };
4751 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
4752 let cfg = LaunchConfig {
4753 grid_dim: (nblk, 1, 1),
4754 block_dim: (32, 1, 1),
4755 shared_mem_bytes: 0,
4756 };
4757 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
4758 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
4759 let mut b = stream.launch_builder(&f);
4760 b.arg(k_row)
4761 .arg(v_row)
4762 .arg(kc)
4763 .arg(vc)
4764 .arg(&ti)
4765 .arg(&kdk)
4766 .arg(&kdv)
4767 .arg(&ktb)
4768 .arg(&vtb);
4769 unsafe {
4770 b.launch(cfg)?;
4771 }
4772 Ok(())
4773 }
4774
4775 pub fn copy_view_into(
4778 &self,
4779 dst: &mut CudaSlice<f32>,
4780 off: usize,
4781 src: &cudarc::driver::CudaView<f32>,
4782 len: usize,
4783 ) -> Result<(), Box<dyn std::error::Error>> {
4784 let mut view = dst.slice_mut(off..off + len);
4785 self.gpu
4786 .stream()
4787 .memcpy_dtod(&src.slice(0..len), &mut view)?;
4788 Ok(())
4789 }
4790
4791 pub fn clone_dtod(
4795 &self,
4796 src: &CudaSlice<f32>,
4797 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4798 let mut dst = self.gpu.stream().alloc_zeros::<f32>(src.len())?;
4799 self.gpu.stream().memcpy_dtod(src, &mut dst)?;
4800 Ok(dst)
4801 }
4802
4803 pub fn dtod_copy_view(
4806 &self,
4807 src: &cudarc::driver::CudaView<f32>,
4808 dst: &mut CudaSlice<f32>,
4809 ) -> Result<(), Box<dyn std::error::Error>> {
4810 self.gpu.stream().memcpy_dtod(src, dst)?;
4811 Ok(())
4812 }
4813
4814 pub fn dtod_copy_view_i8(
4816 &self,
4817 src: &cudarc::driver::CudaView<i8>,
4818 dst: &mut CudaSlice<i8>,
4819 ) -> Result<(), Box<dyn std::error::Error>> {
4820 self.gpu.stream().memcpy_dtod(src, dst)?;
4821 Ok(())
4822 }
4823
4824 pub fn dtod_copy_into(
4826 &self,
4827 src: &CudaSlice<f32>,
4828 dst: &mut CudaSlice<f32>,
4829 offset: usize,
4830 ) -> Result<(), Box<dyn std::error::Error>> {
4831 let n = src.len();
4832 let mut dv = dst.slice_mut(offset..offset + n);
4833 self.gpu.stream().memcpy_dtod(src, &mut dv)?;
4834 Ok(())
4835 }
4836
4837 pub fn copy_batch_uniform_f32(
4843 &self,
4844 table: &CudaSlice<u64>,
4845 n: usize,
4846 words: usize,
4847 ) -> Result<(), Box<dyn std::error::Error>> {
4848 if n == 0 || words == 0 {
4849 return Ok(());
4850 }
4851 debug_assert!(
4852 table.len() >= 2 * n,
4853 "pointer table must hold n srcs + n dsts"
4854 );
4855 let f = self.func("copy_batch_uniform_f32");
4856 let chunks = (words / 4).max(1).div_ceil(256).min(48) as u32;
4859 let (ni, wi) = (n as i32, words as i32);
4860 let cfg = LaunchConfig {
4861 grid_dim: (chunks, n as u32, 1),
4862 block_dim: (256, 1, 1),
4863 shared_mem_bytes: 0,
4864 };
4865 let __s = self.gpu.stream();
4866 let mut b = __s.launch_builder(&f);
4867 b.arg(table).arg(&ni).arg(&wi);
4868 unsafe {
4869 b.launch(cfg)?;
4870 }
4871 Ok(())
4872 }
4873
4874 pub fn htod_u64_into(
4877 &self,
4878 v: &[u64],
4879 dst: &mut CudaSlice<u64>,
4880 ) -> Result<(), Box<dyn std::error::Error>> {
4881 let mut view = dst.slice_mut(0..v.len());
4882 self.gpu.stream().memcpy_htod(v, &mut view)?;
4883 Ok(())
4884 }
4885
4886 pub fn copy_indirect_src_f32(
4891 &self,
4892 src_entry: &cudarc::driver::CudaView<u64>,
4893 dst: &mut CudaSlice<f32>,
4894 dst_off: usize,
4895 words: usize,
4896 ) -> Result<(), Box<dyn std::error::Error>> {
4897 let f = self.func("copy_indirect_src_f32");
4898 let chunks = (words / 4).max(1).div_ceil(256).min(48) as u32;
4899 let wi = words as i32;
4900 let cfg = LaunchConfig {
4901 grid_dim: (chunks, 1, 1),
4902 block_dim: (256, 1, 1),
4903 shared_mem_bytes: 0,
4904 };
4905 let mut dv = dst.slice_mut(dst_off..dst_off + words);
4906 let __s = self.gpu.stream();
4907 let mut b = __s.launch_builder(&f);
4908 b.arg(src_entry).arg(&mut dv).arg(&wi);
4909 unsafe {
4910 b.launch(cfg)?;
4911 }
4912 Ok(())
4913 }
4914
4915 pub fn uninit_i8(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
4917 self.alloc_uninit::<i8>(n)
4918 }
4919
4920 pub fn qmatvec(
4922 &self,
4923 w: &CudaSlice<u8>,
4924 x: &CudaSlice<f32>,
4925 m: usize,
4926 in_f: usize,
4927 out_f: usize,
4928 qtype: i32,
4929 row_bytes: usize,
4930 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4931 let f = self.func("qmatvec_f32");
4932 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
4934 grid_dim: (out_f as u32, m as u32, 1),
4935 block_dim: (256, 1, 1),
4936 shared_mem_bytes: 0,
4937 };
4938 let (inf, outf, mi, qt, rb) =
4939 (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
4940 let __s_b = self.gpu.stream();
4941 let mut b = __s_b.launch_builder(&f);
4942 b.arg(w)
4943 .arg(x)
4944 .arg(&mut y)
4945 .arg(&inf)
4946 .arg(&outf)
4947 .arg(&mi)
4948 .arg(&qt)
4949 .arg(&rb);
4950 unsafe {
4951 b.launch(cfg)?;
4952 }
4953 Ok(y)
4954 }
4955
4956 pub fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
4958 let s = self.gpu.stream().alloc_zeros::<u8>(n)?;
4959 self.keep_if_capturing(&s);
4960 Ok(s)
4961 }
4962
4963 pub fn alloc_u8_uninit(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
4967 let s = unsafe { self.gpu.stream().alloc::<u8>(n)? };
4968 self.keep_if_capturing(&s);
4969 Ok(s)
4970 }
4971
4972 pub fn memset_zeros_view(
4975 &self,
4976 dst: &mut cudarc::driver::CudaViewMut<f32>,
4977 ) -> Result<(), Box<dyn std::error::Error>> {
4978 self.gpu.stream().memset_zeros(dst)?;
4979 Ok(())
4980 }
4981
4982 pub fn stage_expert(
4988 &self,
4989 host_bytes: &[u8],
4990 scratch: &mut CudaSlice<u8>,
4991 off: usize,
4992 ) -> Result<(), Box<dyn std::error::Error>> {
4993 let mut dst = scratch.slice_mut(off..off + host_bytes.len()); self.gpu.stream().memcpy_htod(host_bytes, &mut dst)?; Ok(())
4996 }
4997
4998 pub fn moe_router_topk(
5004 &self,
5005 logits: &CudaSlice<f32>,
5006 t: usize,
5007 n_expert: usize,
5008 n_used: usize,
5009 ) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5010 let f = self.func("moe_router_topk_f32");
5011 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 {
5014 grid_dim: (t as u32, 1, 1),
5015 block_dim: (n_expert as u32, 1, 1),
5016 shared_mem_bytes: 0,
5017 };
5018 let (ne, nu) = (n_expert as i32, n_used as i32);
5019 let __s_b = self.gpu.stream();
5020 let mut b = __s_b.launch_builder(&f);
5021 b.arg(logits)
5022 .arg(&mut sel_idx)
5023 .arg(&mut sel_w)
5024 .arg(&ne)
5025 .arg(&nu);
5026 unsafe {
5027 b.launch(cfg)?;
5028 }
5029 Ok((sel_idx, sel_w))
5030 }
5031
5032 pub fn moe_router_topk_scaled(
5035 &self,
5036 logits: &CudaSlice<f32>,
5037 t: usize,
5038 n_expert: usize,
5039 n_used: usize,
5040 ex_scale: &CudaSlice<f32>,
5041 ) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5042 let f = self.func("moe_router_topk_scaled_f32");
5047 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
5048 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
5049 let cfg = LaunchConfig {
5050 grid_dim: (t as u32, 1, 1),
5051 block_dim: (n_expert as u32, 1, 1),
5052 shared_mem_bytes: 0,
5053 };
5054 let (ne, nu) = (n_expert as i32, n_used as i32);
5055 let __s_b = self.gpu.stream();
5056 let mut b = __s_b.launch_builder(&f);
5057 b.arg(logits)
5058 .arg(&mut sel_idx)
5059 .arg(&mut sel_w)
5060 .arg(&ne)
5061 .arg(&nu)
5062 .arg(ex_scale);
5063 unsafe {
5064 b.launch(cfg)?;
5065 }
5066 Ok((sel_idx, sel_w))
5067 }
5068
5069 pub fn moe_router_topk_host(
5077 &self,
5078 logits: &CudaSlice<f32>,
5079 t: usize,
5080 n_expert: usize,
5081 n_used: usize,
5082 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
5083 let f = self.func("moe_router_topk_f32");
5084 let n = t * n_used;
5085 let mut sel_idx = self.alloc_uninit::<i32>(n)?;
5086 let mut sel_w = self.alloc_uninit::<f32>(n)?;
5087 let cfg = LaunchConfig {
5088 grid_dim: (t as u32, 1, 1),
5089 block_dim: (n_expert as u32, 1, 1),
5090 shared_mem_bytes: 0,
5091 };
5092 let (ne, nu) = (n_expert as i32, n_used as i32);
5093 let __s_b = self.gpu.stream();
5094 let mut b = __s_b.launch_builder(&f);
5095 b.arg(logits)
5096 .arg(&mut sel_idx)
5097 .arg(&mut sel_w)
5098 .arg(&ne)
5099 .arg(&nu);
5100 unsafe {
5101 b.launch(cfg)?;
5102 }
5103 let bytes = n * 8;
5105 let mut guard = self.router_stage.lock().unwrap();
5106 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
5107 *guard = Some(PinnedStage::new(bytes.max(4096))?);
5108 }
5109 let stage = guard.as_mut().unwrap();
5110 let (si, sw) = unsafe {
5111 (
5112 std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
5113 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n),
5114 )
5115 };
5116 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()))
5120 }
5121
5122 #[allow(clippy::too_many_arguments)]
5126 pub fn moe_router_sigmoid_topk(
5127 &self,
5128 logits: &CudaSlice<f32>,
5129 t: usize,
5130 n_expert: usize,
5131 n_used: usize,
5132 active_count: usize,
5133 correction_bias: &CudaSlice<f32>,
5134 active: &CudaSlice<u8>,
5135 scaling_factor: f32,
5136 route_norm: bool,
5137 ) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5138 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
5139 if n_expert == 0 || n_expert > 1024 || n_used == 0 || n_used > n_expert {
5140 return Err(format!(
5141 "sigmoid router shape unsupported: n_expert={n_expert}, n_used={n_used}",
5142 )
5143 .into());
5144 }
5145 if logits.len() < t * n_expert
5146 || correction_bias.len() != n_expert
5147 || active.len() != n_expert
5148 {
5149 return Err(format!(
5150 "sigmoid router buffer mismatch: logits={} bias={} active={} expected logits>={} row={}",
5151 logits.len(), correction_bias.len(), active.len(), t * n_expert, n_expert,
5152 ).into());
5153 }
5154 let f = self.func(crate::sigmoid_topk_kernel(
5155 crate::sig_expf_dev_on(),
5156 crate::topk_fast_on(),
5157 n_used,
5158 ));
5159 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
5160 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
5161 let threads = n_expert.div_ceil(32) * 32;
5162 let cfg = LaunchConfig {
5163 grid_dim: (t as u32, 1, 1),
5164 block_dim: (threads as u32, 1, 1),
5165 shared_mem_bytes: 0,
5166 };
5167 let (ne, nu, rn) = (n_expert as i32, n_used as i32, i32::from(route_norm));
5168 let __s_b = self.gpu.stream();
5169 let mut b = __s_b.launch_builder(&f);
5170 b.arg(logits)
5171 .arg(correction_bias)
5172 .arg(active)
5173 .arg(&mut sel_idx)
5174 .arg(&mut sel_w)
5175 .arg(&ne)
5176 .arg(&nu)
5177 .arg(&scaling_factor)
5178 .arg(&rn);
5179 unsafe {
5180 b.launch(cfg)?;
5181 }
5182 Ok((sel_idx, sel_w))
5183 }
5184
5185 #[allow(clippy::too_many_arguments)]
5188 pub fn ring_flag_raw(&self, ptr: u64, value: u32) -> Result<(), Box<dyn std::error::Error>> {
5192 if ptr == 0 {
5193 return Err("ring_flag_raw: unarmed flag".into());
5194 }
5195 let f = self.func("memra_ring_flag");
5196 let cfg = LaunchConfig {
5197 grid_dim: (1, 1, 1),
5198 block_dim: (32, 1, 1),
5199 shared_mem_bytes: 0,
5200 };
5201 let __s_b = self.gpu.stream();
5202 let mut b = __s_b.launch_builder(&f);
5203 b.arg(&ptr).arg(&value);
5204 unsafe {
5205 b.launch(cfg)?;
5206 }
5207 Ok(())
5208 }
5209
5210 pub fn moe_sel_w_mirror(
5213 &self,
5214 sel_src: &CudaSlice<i32>,
5215 w_src: &CudaSlice<f32>,
5216 sel_dst: &mut CudaSlice<i32>,
5217 w_dst: &mut CudaSlice<f32>,
5218 n: usize,
5219 ) -> Result<(), Box<dyn std::error::Error>> {
5220 if n == 0
5221 || n > 32
5222 || sel_src.len() < n
5223 || w_src.len() < n
5224 || sel_dst.len() < n
5225 || w_dst.len() < n
5226 {
5227 return Err(format!("moe_sel_w_mirror geometry n={n}").into());
5228 }
5229 let f = self.func("moe_sel_w_mirror");
5230 let cfg = LaunchConfig {
5231 grid_dim: (1, 1, 1),
5232 block_dim: (32, 1, 1),
5233 shared_mem_bytes: 0,
5234 };
5235 let ni = n as i32;
5236 let __s_b = self.gpu.stream();
5237 let mut b = __s_b.launch_builder(&f);
5238 b.arg(sel_src).arg(w_src).arg(sel_dst).arg(w_dst).arg(&ni);
5239 unsafe {
5240 b.launch(cfg)?;
5241 }
5242 Ok(())
5243 }
5244
5245 pub fn moe_router_sigmoid_topk_into(
5246 &self,
5247 logits: &CudaSlice<f32>,
5248 t: usize,
5249 n_expert: usize,
5250 n_used: usize,
5251 active_count: usize,
5252 correction_bias: &CudaSlice<f32>,
5253 active: &CudaSlice<u8>,
5254 scaling_factor: f32,
5255 route_norm: bool,
5256 sel_idx: &mut CudaSlice<i32>,
5257 sel_w: &mut CudaSlice<f32>,
5258 ) -> Result<(), Box<dyn std::error::Error>> {
5259 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
5260 if n_expert == 0
5261 || n_expert > 1024
5262 || n_used == 0
5263 || n_used > 32 || n_used > n_expert
5265 || logits.len() < t * n_expert
5266 || correction_bias.len() != n_expert
5267 || active.len() != n_expert
5268 || sel_idx.len() < t * n_used
5269 || sel_w.len() < t * n_used
5270 {
5271 return Err("sigmoid router _into geometry mismatch".into());
5272 }
5273 let f = self.func(crate::sigmoid_topk_kernel(
5274 crate::sig_expf_dev_on(),
5275 crate::topk_fast_on(),
5276 n_used,
5277 ));
5278 let threads = n_expert.div_ceil(32) * 32;
5279 let cfg = LaunchConfig {
5280 grid_dim: (t as u32, 1, 1),
5281 block_dim: (threads as u32, 1, 1),
5282 shared_mem_bytes: 0,
5283 };
5284 let (ne, nu, rn) = (n_expert as i32, n_used as i32, i32::from(route_norm));
5285 let __s_b = self.gpu.stream();
5286 let mut b = __s_b.launch_builder(&f);
5287 b.arg(logits)
5288 .arg(correction_bias)
5289 .arg(active)
5290 .arg(&mut *sel_idx)
5291 .arg(&mut *sel_w)
5292 .arg(&ne)
5293 .arg(&nu)
5294 .arg(&scaling_factor)
5295 .arg(&rn);
5296 unsafe {
5297 b.launch(cfg)?;
5298 }
5299 Ok(())
5300 }
5301
5302 #[allow(clippy::too_many_arguments)]
5305 pub fn moe_router_sigmoid_topk_host(
5306 &self,
5307 logits: &CudaSlice<f32>,
5308 t: usize,
5309 n_expert: usize,
5310 n_used: usize,
5311 active_count: usize,
5312 correction_bias: &CudaSlice<f32>,
5313 active: &CudaSlice<u8>,
5314 scaling_factor: f32,
5315 route_norm: bool,
5316 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
5317 let (sel_idx, sel_w) = self.moe_router_sigmoid_topk(
5318 logits,
5319 t,
5320 n_expert,
5321 n_used,
5322 active_count,
5323 correction_bias,
5324 active,
5325 scaling_factor,
5326 route_norm,
5327 )?;
5328 let n = t * n_used;
5329 let bytes = n * 8;
5330 let mut guard = self.router_stage.lock().unwrap();
5331 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
5332 *guard = Some(PinnedStage::new(bytes.max(4096))?);
5333 }
5334 let stage = guard.as_mut().unwrap();
5335 let (si, sw) = unsafe {
5336 (
5337 std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
5338 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n),
5339 )
5340 };
5341 self.gpu.stream().memcpy_dtoh(&sel_idx, si)?;
5342 self.gpu.stream().memcpy_dtoh(&sel_w, sw)?;
5343 self.gpu.stream().synchronize()?;
5344 Ok((si.iter().map(|&i| i as u32).collect(), sw.to_vec()))
5345 }
5346
5347 pub fn stage_expert_async(
5351 &self,
5352 host_bytes: &[u8],
5353 scratch: &mut CudaSlice<u8>,
5354 off: usize,
5355 ) -> Result<cudarc::driver::CudaEvent, Box<dyn std::error::Error>> {
5356 let mut dst = scratch.slice_mut(off..off + host_bytes.len());
5357 self.copy_stream.memcpy_htod(host_bytes, &mut dst)?;
5358 Ok(self.copy_stream.record_event(None)?)
5359 }
5360
5361 pub fn compute_wait(
5363 &self,
5364 ev: &cudarc::driver::CudaEvent,
5365 ) -> Result<(), Box<dyn std::error::Error>> {
5366 self.gpu.stream().wait(ev)?;
5367 Ok(())
5368 }
5369
5370 pub fn qmatvec_view(
5375 &self,
5376 w: &CudaSlice<u8>,
5377 range: std::ops::Range<usize>,
5378 x: &cudarc::driver::CudaView<f32>,
5379 m: usize,
5380 in_f: usize,
5381 out_f: usize,
5382 qtype: i32,
5383 row_bytes: usize,
5384 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5385 let f = self.func("qmatvec_f32");
5386 let wv = w.slice(range); let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
5389 grid_dim: (out_f as u32, m as u32, 1),
5390 block_dim: (256, 1, 1),
5391 shared_mem_bytes: 0,
5392 };
5393 let (inf, outf, mi, qt, rb) =
5394 (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
5395 let __s_b = self.gpu.stream();
5396 let mut b = __s_b.launch_builder(&f);
5397 b.arg(&wv)
5398 .arg(x)
5399 .arg(&mut y)
5400 .arg(&inf)
5401 .arg(&outf)
5402 .arg(&mi)
5403 .arg(&qt)
5404 .arg(&rb);
5405 unsafe {
5406 b.launch(cfg)?;
5407 }
5408 Ok(y)
5409 }
5410
5411 #[allow(clippy::too_many_arguments)]
5418 pub fn moe_gate_up_silu8_q8(
5422 &self,
5423 gp: WPtr8,
5424 up: WPtr8,
5425 aq: &CudaSlice<i8>,
5426 ad: &CudaSlice<f32>,
5427 in_f: usize,
5428 n_ff: usize,
5429 n_used: usize,
5430 qt_g: i32,
5431 qt_u: i32,
5432 rb_g: usize,
5433 rb_u: usize,
5434 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5435 let f = self.func("moe_gate_up_silu8_q8");
5436 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
5437 let cfg = LaunchConfig {
5438 grid_dim: (n_ff as u32, n_used as u32, 1),
5439 block_dim: (32, 1, 1),
5440 shared_mem_bytes: 0,
5441 };
5442 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
5443 let __s_b = self.gpu.stream();
5444 let mut b = __s_b.launch_builder(&f);
5445 b.arg(&gp)
5446 .arg(&up)
5447 .arg(aq)
5448 .arg(ad)
5449 .arg(&mut act)
5450 .arg(&inf)
5451 .arg(&nff)
5452 .arg(&qt_g)
5453 .arg(&qt_u)
5454 .arg(&rbg)
5455 .arg(&rbu);
5456 unsafe {
5457 b.launch(cfg)?;
5458 }
5459 Ok(act)
5460 }
5461
5462 #[allow(clippy::too_many_arguments)]
5463 pub fn moe_down8_fma_q8(
5464 &self,
5465 dp: WPtr8,
5466 w: F32x8,
5467 aq2: &CudaSlice<i8>,
5468 ad2: &CudaSlice<f32>,
5469 dst: &mut cudarc::driver::CudaViewMut<f32>,
5470 in_f: usize,
5471 out_f: usize,
5472 n_used: usize,
5473 qt: i32,
5474 rb: usize,
5475 ) -> Result<(), Box<dyn std::error::Error>> {
5476 let f = self.func("moe_down8_fma_q8");
5477 let cfg = LaunchConfig {
5478 grid_dim: (out_f as u32, 1, 1),
5479 block_dim: (32, 1, 1),
5480 shared_mem_bytes: 0,
5481 };
5482 let (inf, outf, nu, rbi) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
5483 let __s_b = self.gpu.stream();
5484 let mut b = __s_b.launch_builder(&f);
5485 b.arg(&dp)
5486 .arg(&w)
5487 .arg(aq2)
5488 .arg(ad2)
5489 .arg(dst)
5490 .arg(&inf)
5491 .arg(&outf)
5492 .arg(&nu)
5493 .arg(&qt)
5494 .arg(&rbi);
5495 unsafe {
5496 b.launch(cfg)?;
5497 }
5498 Ok(())
5499 }
5500
5501 pub fn qmatvec_expert_q8(
5503 &self,
5504 w: &CudaSlice<u8>,
5505 range: std::ops::Range<usize>,
5506 aq: &CudaSlice<i8>,
5507 ad: &CudaSlice<f32>,
5508 m: usize,
5509 in_f: usize,
5510 out_f: usize,
5511 qtype: i32,
5512 row_bytes: usize,
5513 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5514 let f = self.func("qmatvec_expert_q8");
5515 let wv = w.slice(range);
5516 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
5517 const ROWS: u32 = 4; let cfg = LaunchConfig {
5519 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, m as u32, 1),
5520 block_dim: (32, ROWS, 1),
5521 shared_mem_bytes: 0,
5522 };
5523 let (inf, outf, mi, rbi) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
5524 let __s_b = self.gpu.stream();
5525 let mut b = __s_b.launch_builder(&f);
5526 b.arg(&wv)
5527 .arg(aq)
5528 .arg(ad)
5529 .arg(&mut y)
5530 .arg(&inf)
5531 .arg(&outf)
5532 .arg(&mi)
5533 .arg(&qtype)
5534 .arg(&rbi);
5535 unsafe {
5536 b.launch(cfg)?;
5537 }
5538 Ok(y)
5539 }
5540
5541 pub fn moe_gate_up_silu8(
5542 &self,
5543 gp: WPtr8,
5544 up: WPtr8,
5545 x: &cudarc::driver::CudaView<f32>,
5546 in_f: usize,
5547 n_ff: usize,
5548 n_used: usize,
5549 qt_g: i32,
5550 qt_u: i32,
5551 rb_g: usize,
5552 rb_u: usize,
5553 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5554 let f = self.func("moe_gate_up_silu8_f32");
5555 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig {
5557 grid_dim: (n_ff as u32, n_used as u32, 1),
5558 block_dim: (256, 1, 1),
5559 shared_mem_bytes: 0,
5560 };
5561 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
5562 let __s_b = self.gpu.stream();
5563 let mut b = __s_b.launch_builder(&f);
5564 b.arg(&gp)
5565 .arg(&up)
5566 .arg(x)
5567 .arg(&mut act)
5568 .arg(&inf)
5569 .arg(&nff)
5570 .arg(&qt_g)
5571 .arg(&qt_u)
5572 .arg(&rbg)
5573 .arg(&rbu);
5574 unsafe {
5575 b.launch(cfg)?;
5576 }
5577 Ok(act)
5578 }
5579
5580 #[allow(clippy::too_many_arguments)]
5586 pub fn moe_down8_fma_into(
5587 &self,
5588 dp: WPtr8,
5589 w: F32x8,
5590 act: &CudaSlice<f32>,
5591 dst: &mut cudarc::driver::CudaViewMut<f32>,
5592 in_f: usize,
5593 out_f: usize,
5594 n_used: usize,
5595 qt: i32,
5596 rb: usize,
5597 ) -> Result<(), Box<dyn std::error::Error>> {
5598 let f = self.func("moe_down8_fma_f32");
5599 let cfg = LaunchConfig {
5600 grid_dim: (out_f as u32, 1, 1),
5601 block_dim: (256, 1, 1),
5602 shared_mem_bytes: 0,
5603 };
5604 let (inf, outf, nu, rbv) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
5605 let __s_b = self.gpu.stream();
5606 let mut b = __s_b.launch_builder(&f);
5607 b.arg(&dp)
5608 .arg(&w)
5609 .arg(act)
5610 .arg(dst)
5611 .arg(&inf)
5612 .arg(&outf)
5613 .arg(&nu)
5614 .arg(&qt)
5615 .arg(&rbv);
5616 unsafe {
5617 b.launch(cfg)?;
5618 }
5619 Ok(())
5620 }
5621
5622 #[allow(clippy::too_many_arguments)]
5627 #[allow(clippy::too_many_arguments)]
5642 #[allow(clippy::too_many_arguments)]
5644 pub fn moe_pairs_matvec_q8(
5645 &self,
5646 table: &CudaSlice<u64>,
5647 proj: i32,
5648 pair_tok: &CudaSlice<i32>,
5649 pair_ex: &CudaSlice<i32>,
5650 aq: &CudaSlice<i8>,
5651 ad: &CudaSlice<f32>,
5652 in_f: usize,
5653 out_f: usize,
5654 n_expert: usize,
5655 n_pairs: usize,
5656 qtype: i32,
5657 row_bytes: usize,
5658 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5659 let f = self.func("moe_pairs_matvec_q8");
5660 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
5661 const ROWS: u32 = 4;
5662 let cfg = LaunchConfig {
5663 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_pairs as u32, 1),
5664 block_dim: (32, ROWS, 1),
5665 shared_mem_bytes: 0,
5666 };
5667 let (inf, outf, ne, np, rbi) = (
5668 in_f as i32,
5669 out_f as i32,
5670 n_expert as i32,
5671 n_pairs as i32,
5672 row_bytes as i64,
5673 );
5674 let __s_b = self.gpu.stream();
5675 let mut b = __s_b.launch_builder(&f);
5676 b.arg(table)
5677 .arg(&proj)
5678 .arg(pair_tok)
5679 .arg(pair_ex)
5680 .arg(aq)
5681 .arg(ad)
5682 .arg(&mut y)
5683 .arg(&inf)
5684 .arg(&outf)
5685 .arg(&ne)
5686 .arg(&np)
5687 .arg(&qtype)
5688 .arg(&rbi);
5689 unsafe {
5690 b.launch(cfg)?;
5691 }
5692 Ok(y)
5693 }
5694
5695 #[allow(clippy::too_many_arguments)]
5697 pub fn moe_pairs_matvec_q8_em(
5698 &self,
5699 table: &CudaSlice<u64>,
5700 proj: i32,
5701 ex_ids: &CudaSlice<i32>,
5702 ex_off: &CudaSlice<i32>,
5703 ex_pairs: &CudaSlice<i32>,
5704 pair_tok: &CudaSlice<i32>,
5705 aq: &CudaSlice<i8>,
5706 ad: &CudaSlice<f32>,
5707 in_f: usize,
5708 out_f: usize,
5709 n_expert: usize,
5710 n_active: usize,
5711 n_pairs: usize,
5712 qtype: i32,
5713 row_bytes: usize,
5714 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5715 let f = self.func("moe_pairs_matvec_q8_em");
5716 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
5717 const ROWS: u32 = 4;
5718 let cfg = LaunchConfig {
5719 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
5720 block_dim: (32, ROWS, 1),
5721 shared_mem_bytes: 0,
5722 };
5723 let (inf, outf, ne, na, rbi) = (
5724 in_f as i32,
5725 out_f as i32,
5726 n_expert as i32,
5727 n_active as i32,
5728 row_bytes as i64,
5729 );
5730 let __s_b = self.gpu.stream();
5731 let mut b = __s_b.launch_builder(&f);
5732 b.arg(table)
5733 .arg(&proj)
5734 .arg(ex_ids)
5735 .arg(ex_off)
5736 .arg(ex_pairs)
5737 .arg(pair_tok)
5738 .arg(aq)
5739 .arg(ad)
5740 .arg(&mut y)
5741 .arg(&inf)
5742 .arg(&outf)
5743 .arg(&ne)
5744 .arg(&na)
5745 .arg(&qtype)
5746 .arg(&rbi);
5747 unsafe {
5748 b.launch(cfg)?;
5749 }
5750 Ok(y)
5751 }
5752
5753 #[allow(clippy::too_many_arguments)]
5756 pub fn moe_pairs_matvec_q8_dec(
5757 &self,
5758 table: &CudaSlice<u64>,
5759 proj: i32,
5760 ex_ids: &CudaSlice<i32>,
5761 ex_off: &CudaSlice<i32>,
5762 ex_pairs: &CudaSlice<i32>,
5763 pair_tok: &CudaSlice<i32>,
5764 aq: &CudaSlice<i8>,
5765 ad: &CudaSlice<f32>,
5766 in_f: usize,
5767 out_f: usize,
5768 n_expert: usize,
5769 n_active: usize,
5770 n_pairs: usize,
5771 qtype: i32,
5772 row_bytes: usize,
5773 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5774 let f = self.func("moe_pairs_matvec_q8_dec");
5775 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
5776 const ROWS: u32 = 4;
5777 let cfg = LaunchConfig {
5778 grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
5779 block_dim: (32, ROWS, 1),
5780 shared_mem_bytes: 0,
5781 };
5782 let (inf, outf, ne, na, rbi) = (
5783 in_f as i32,
5784 out_f as i32,
5785 n_expert as i32,
5786 n_active as i32,
5787 row_bytes as i64,
5788 );
5789 let __s_b = self.gpu.stream();
5790 let mut b = __s_b.launch_builder(&f);
5791 b.arg(table)
5792 .arg(&proj)
5793 .arg(ex_ids)
5794 .arg(ex_off)
5795 .arg(ex_pairs)
5796 .arg(pair_tok)
5797 .arg(aq)
5798 .arg(ad)
5799 .arg(&mut y)
5800 .arg(&inf)
5801 .arg(&outf)
5802 .arg(&ne)
5803 .arg(&na)
5804 .arg(&qtype)
5805 .arg(&rbi);
5806 unsafe {
5807 b.launch(cfg)?;
5808 }
5809 Ok(y)
5810 }
5811
5812 pub fn moe_pairs_gelu_mul(
5813 &self,
5814 gate: &CudaSlice<f32>,
5815 up: &CudaSlice<f32>,
5816 n: usize,
5817 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5818 let f = self.func("moe_pairs_gelu_mul");
5819 let mut act = self.alloc_uninit::<f32>(n)?;
5820 let cfg = LaunchConfig::for_num_elems(n as u32);
5821 let nl = n as i64;
5822 let __s_b = self.gpu.stream();
5823 let mut b = __s_b.launch_builder(&f);
5824 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
5825 unsafe {
5826 b.launch(cfg)?;
5827 }
5828 Ok(act)
5829 }
5830
5831 pub fn moe_pairs_silu_mul(
5832 &self,
5833 gate: &CudaSlice<f32>,
5834 up: &CudaSlice<f32>,
5835 n: usize,
5836 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5837 let f = self.func("moe_pairs_silu_mul");
5838 let mut act = self.alloc_uninit::<f32>(n)?;
5839 let cfg = LaunchConfig::for_num_elems(n as u32);
5840 let nl = n as i64;
5841 let __s_b = self.gpu.stream();
5842 let mut b = __s_b.launch_builder(&f);
5843 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
5844 unsafe {
5845 b.launch(cfg)?;
5846 }
5847 Ok(act)
5848 }
5849
5850 #[allow(clippy::too_many_arguments)]
5851 pub fn moe_pairs_scatter(
5852 &self,
5853 y_down: &CudaSlice<f32>,
5854 pair_w: &CudaSlice<f32>,
5855 tok_pair_off: &CudaSlice<i32>,
5856 tok_pair_ids: &CudaSlice<i32>,
5857 moe_out: &mut CudaSlice<f32>,
5858 t: usize,
5859 n_embd: usize,
5860 ) -> Result<(), Box<dyn std::error::Error>> {
5861 let f = self.func("moe_pairs_scatter");
5862 let cfg = LaunchConfig {
5863 grid_dim: (((n_embd + 255) / 256) as u32, t as u32, 1),
5864 block_dim: (256, 1, 1),
5865 shared_mem_bytes: 0,
5866 };
5867 let ne = n_embd as i32;
5868 let __s_b = self.gpu.stream();
5869 let mut b = __s_b.launch_builder(&f);
5870 b.arg(y_down)
5871 .arg(pair_w)
5872 .arg(tok_pair_off)
5873 .arg(tok_pair_ids)
5874 .arg(moe_out)
5875 .arg(&ne);
5876 unsafe {
5877 b.launch(cfg)?;
5878 }
5879 Ok(())
5880 }
5881
5882 #[allow(clippy::too_many_arguments)]
5886 pub fn moe_gate_up_gelu8_dev_q8(
5887 &self,
5888 table: &CudaSlice<u64>,
5889 sel: &cudarc::driver::CudaView<i32>,
5890 aq: &CudaSlice<i8>,
5891 ad: &CudaSlice<f32>,
5892 in_f: usize,
5893 n_ff: usize,
5894 n_used: usize,
5895 n_expert: usize,
5896 qt_g: i32,
5897 qt_u: i32,
5898 rb_g: usize,
5899 rb_u: usize,
5900 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5901 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
5902 let (inf, nff, ne, rbg, rbu) = (
5903 in_f as i32,
5904 n_ff as i32,
5905 n_expert as i32,
5906 rb_g as i64,
5907 rb_u as i64,
5908 );
5909 let f = self.func("moe_gate_up_gelu8_dev_q8");
5910 let cfg = LaunchConfig {
5911 grid_dim: (n_ff as u32, n_used as u32, 1),
5912 block_dim: (32, 1, 1),
5913 shared_mem_bytes: 0,
5914 };
5915 let __s_b = self.gpu.stream();
5916 let mut b = __s_b.launch_builder(&f);
5917 b.arg(table)
5918 .arg(sel)
5919 .arg(aq)
5920 .arg(ad)
5921 .arg(&mut act)
5922 .arg(&inf)
5923 .arg(&nff)
5924 .arg(&ne)
5925 .arg(&qt_g)
5926 .arg(&qt_u)
5927 .arg(&rbg)
5928 .arg(&rbu);
5929 unsafe {
5930 b.launch(cfg)?;
5931 }
5932 Ok(act)
5933 }
5934
5935 #[allow(clippy::too_many_arguments)]
5937 pub fn moe_gate_up_gelu8_dev_q8_rows(
5938 &self,
5939 table: &CudaSlice<u64>,
5940 sel: &CudaSlice<i32>,
5941 aq: &CudaSlice<i8>,
5942 ad: &CudaSlice<f32>,
5943 t: usize,
5944 in_f: usize,
5945 n_ff: usize,
5946 n_used: usize,
5947 n_expert: usize,
5948 qt_g: i32,
5949 qt_u: i32,
5950 rb_g: usize,
5951 rb_u: usize,
5952 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5953 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
5954 let (inf, nff, ne, rbg, rbu, nu) = (
5955 in_f as i32,
5956 n_ff as i32,
5957 n_expert as i32,
5958 rb_g as i64,
5959 rb_u as i64,
5960 n_used as i32,
5961 );
5962 let f = self.func("moe_gate_up_gelu8_dev_q8_rows");
5963 let cfg = LaunchConfig {
5964 grid_dim: (n_ff as u32, n_used as u32, t as u32),
5965 block_dim: (32, 1, 1),
5966 shared_mem_bytes: 0,
5967 };
5968 let __s_b = self.gpu.stream();
5969 let mut b = __s_b.launch_builder(&f);
5970 b.arg(table)
5971 .arg(sel)
5972 .arg(aq)
5973 .arg(ad)
5974 .arg(&mut act)
5975 .arg(&inf)
5976 .arg(&nff)
5977 .arg(&ne)
5978 .arg(&qt_g)
5979 .arg(&qt_u)
5980 .arg(&rbg)
5981 .arg(&rbu)
5982 .arg(&nu);
5983 unsafe {
5984 b.launch(cfg)?;
5985 }
5986 Ok(act)
5987 }
5988
5989 #[allow(clippy::too_many_arguments)]
5991 pub fn moe_gate_up_gelu8_dev_q8_csr(
5992 &self,
5993 table: &CudaSlice<u64>,
5994 sel: &CudaSlice<i32>,
5995 aq: &CudaSlice<i8>,
5996 ad: &CudaSlice<f32>,
5997 n_pairs: usize,
5998 in_f: usize,
5999 n_ff: usize,
6000 n_used: usize,
6001 n_expert: usize,
6002 qt_g: i32,
6003 qt_u: i32,
6004 rb_g: usize,
6005 rb_u: usize,
6006 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6007 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
6008 let (inf, nff, ne, rbg, rbu, nu, npi) = (
6009 in_f as i32,
6010 n_ff as i32,
6011 n_expert as i32,
6012 rb_g as i64,
6013 rb_u as i64,
6014 n_used as i32,
6015 n_pairs as i32,
6016 );
6017 let f = self.func("moe_gate_up_gelu8_dev_q8_csr");
6018 let cfg = LaunchConfig {
6019 grid_dim: (n_ff as u32, n_pairs as u32, 1),
6020 block_dim: (32, 1, 1),
6021 shared_mem_bytes: 0,
6022 };
6023 let __s_b = self.gpu.stream();
6024 let mut b = __s_b.launch_builder(&f);
6025 b.arg(table)
6026 .arg(sel)
6027 .arg(aq)
6028 .arg(ad)
6029 .arg(&mut act)
6030 .arg(&inf)
6031 .arg(&nff)
6032 .arg(&ne)
6033 .arg(&qt_g)
6034 .arg(&qt_u)
6035 .arg(&rbg)
6036 .arg(&rbu)
6037 .arg(&nu)
6038 .arg(&npi);
6039 unsafe {
6040 b.launch(cfg)?;
6041 }
6042 Ok(act)
6043 }
6044
6045 #[allow(clippy::too_many_arguments)]
6047 pub fn moe_down8_fma_dev_q8_rows_g(
6048 &self,
6049 table: &CudaSlice<u64>,
6050 sel: &CudaSlice<i32>,
6051 w: &CudaSlice<f32>,
6052 aq2: &CudaSlice<i8>,
6053 ad2: &CudaSlice<f32>,
6054 dst: &mut CudaSlice<f32>,
6055 t: usize,
6056 in_f: usize,
6057 out_f: usize,
6058 n_used: usize,
6059 n_expert: usize,
6060 qt: i32,
6061 rb: usize,
6062 ) -> Result<(), Box<dyn std::error::Error>> {
6063 let (inf, outf, nu, ne, rbi) = (
6064 in_f as i32,
6065 out_f as i32,
6066 n_used as i32,
6067 n_expert as i32,
6068 rb as i64,
6069 );
6070 let step_b1_w8 = t == 1 && in_f == 1280 && out_f == 4096 && n_used == 8 && qt == QT_IQ4_XS;
6074 let f = self.func(if step_b1_w8 {
6075 "moe_down8_fma_dev_q8_rows_w8"
6076 } else {
6077 "moe_down8_fma_dev_q8_rows_g"
6078 });
6079 let cfg = LaunchConfig {
6080 grid_dim: (out_f as u32, 1, t as u32),
6081 block_dim: (32, if step_b1_w8 { 8 } else { 1 }, 1),
6082 shared_mem_bytes: 0,
6083 };
6084 let __s_b = self.gpu.stream();
6085 let mut b = __s_b.launch_builder(&f);
6086 b.arg(table)
6087 .arg(sel)
6088 .arg(w)
6089 .arg(aq2)
6090 .arg(ad2)
6091 .arg(dst)
6092 .arg(&inf)
6093 .arg(&outf)
6094 .arg(&nu)
6095 .arg(&ne)
6096 .arg(&qt)
6097 .arg(&rbi);
6098 unsafe {
6099 b.launch(cfg)?;
6100 }
6101 Ok(())
6102 }
6103
6104 pub fn rp_probe_q4(&self, m: usize) -> Result<(f64, f64), Box<dyn std::error::Error>> {
6108 let (out_f, in_f) = (2048usize, 2816usize);
6109 let nblk = in_f / 32;
6110 let mut seed = 0x9E3779B97F4A7C15u64;
6111 let mut rng = move || {
6112 seed = seed
6113 .wrapping_mul(6364136223846793005)
6114 .wrapping_add(1442695040888963407);
6115 (seed >> 33) as u8
6116 };
6117 let mut w = vec![0u8; out_f * nblk * 18];
6118 for b in w.iter_mut() {
6119 *b = rng();
6120 }
6121 for r in 0..out_f {
6122 for g in 0..nblk {
6123 let off = (r * nblk + g) * 18;
6124 w[off] = 0x00;
6125 w[off + 1] = 0x2C; }
6127 }
6128 let qplane = out_f * nblk * 16;
6129 let mut wrp = vec![0u8; w.len()];
6130 for r in 0..out_f {
6131 for g in 0..nblk {
6132 let src = &w[(r * nblk + g) * 18..(r * nblk + g) * 18 + 18];
6133 wrp[qplane + (r * nblk + g) * 2..qplane + (r * nblk + g) * 2 + 2]
6134 .copy_from_slice(&src[0..2]);
6135 wrp[(r * nblk + g) * 16..(r * nblk + g) * 16 + 16].copy_from_slice(&src[2..18]);
6136 }
6137 }
6138 let w_d = self.htod_bytes(&w)?;
6139 let wrp_d = self.htod_bytes(&wrp)?;
6140 let mut aq = vec![0i8; m * in_f];
6141 for v in aq.iter_mut() {
6142 *v = rng() as i8;
6143 }
6144 let aq_d = self.htod_i8(&aq)?;
6145 let ad_d = self.htod(&vec![0.03125f32; m * nblk])?;
6146 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
6147 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
6148 const RPB: u32 = 4;
6149 let cfg = LaunchConfig {
6150 grid_dim: ((out_f as u32).div_ceil(RPB), 1, 1),
6151 block_dim: (32, RPB, 1),
6152 shared_mem_bytes: 0,
6153 };
6154 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
6155 let (rb, qp) = ((nblk * 18) as i64, qplane as i64);
6156 let fb = self.func("qmatvec_q4_0_mmvq_b4");
6157 let fr = self.func("qmatvec_q4_0_mmvq_b4_rp");
6158 {
6159 let __s_b = self.gpu.stream();
6160 let mut b = __s_b.launch_builder(&fb);
6161 b.arg(&w_d)
6162 .arg(&aq_d)
6163 .arg(&ad_d)
6164 .arg(&mut y0)
6165 .arg(&inf)
6166 .arg(&outf)
6167 .arg(&mi)
6168 .arg(&rb);
6169 unsafe {
6170 b.launch(cfg)?;
6171 }
6172 let __s_b = self.gpu.stream();
6173 let mut b = __s_b.launch_builder(&fr);
6174 b.arg(&wrp_d)
6175 .arg(&aq_d)
6176 .arg(&ad_d)
6177 .arg(&mut y1)
6178 .arg(&inf)
6179 .arg(&outf)
6180 .arg(&mi)
6181 .arg(&qp);
6182 unsafe {
6183 b.launch(cfg)?;
6184 }
6185 }
6186 self.gpu.stream().synchronize()?;
6187 let (h0, h1) = (self.dtoh(&y0)?, self.dtoh(&y1)?);
6188 let nd = h0
6189 .iter()
6190 .zip(&h1)
6191 .filter(|(a, b)| a.to_bits() != b.to_bits())
6192 .count();
6193 if nd != 0 {
6194 return Err(format!("rp twin not bitwise: {nd}/{} diffs", h0.len()).into());
6195 }
6196 let mut time = |rp: bool| -> Result<f64, Box<dyn std::error::Error>> {
6197 self.gpu.stream().synchronize()?;
6198 let t0 = std::time::Instant::now();
6199 for _ in 0..500 {
6200 if rp {
6201 let __s_b = self.gpu.stream();
6202 let mut b = __s_b.launch_builder(&fr);
6203 b.arg(&wrp_d)
6204 .arg(&aq_d)
6205 .arg(&ad_d)
6206 .arg(&mut y1)
6207 .arg(&inf)
6208 .arg(&outf)
6209 .arg(&mi)
6210 .arg(&qp);
6211 unsafe {
6212 b.launch(cfg)?;
6213 }
6214 } else {
6215 let __s_b = self.gpu.stream();
6216 let mut b = __s_b.launch_builder(&fb);
6217 b.arg(&w_d)
6218 .arg(&aq_d)
6219 .arg(&ad_d)
6220 .arg(&mut y0)
6221 .arg(&inf)
6222 .arg(&outf)
6223 .arg(&mi)
6224 .arg(&rb);
6225 unsafe {
6226 b.launch(cfg)?;
6227 }
6228 }
6229 }
6230 self.gpu.stream().synchronize()?;
6231 Ok(t0.elapsed().as_secs_f64() * 1e6 / 500.0)
6232 };
6233 let _ = time(false)?;
6234 let _ = time(true)?; Ok((time(false)?, time(true)?))
6236 }
6237
6238 pub fn build_q4_rp4(
6243 &self,
6244 t: &mut crate::model::GpuTensor,
6245 ) -> Result<(), Box<dyn std::error::Error>> {
6246 use crate::model::GpuTensor;
6247 let GpuTensor::Quant {
6248 bytes,
6249 qtype,
6250 row_bytes,
6251 ne,
6252 rp4,
6253 ..
6254 } = t
6255 else {
6256 return Ok(());
6257 };
6258 if *qtype != QT_Q4_0 || rp4.is_some() || ne.len() != 2 {
6259 return Ok(());
6260 }
6261 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
6262 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 18 {
6263 return Ok(());
6264 }
6265 let nblk = in_f / 32;
6266 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 18)?;
6267 let f = self.func("q4_0_split_rp_build");
6268 let n = (out_f * nblk) as i32;
6269 let cfg = LaunchConfig {
6270 grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
6271 block_dim: (256, 1, 1),
6272 shared_mem_bytes: 0,
6273 };
6274 let (of, nb) = (out_f as i32, nblk as i32);
6275 let _ = n;
6276 let __s_b = self.gpu.stream();
6277 let mut b = __s_b.launch_builder(&f);
6278 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
6279 unsafe {
6280 b.launch(cfg)?;
6281 }
6282 *rp4 = Some(dst);
6283 Ok(())
6284 }
6285
6286 pub fn build_q8_rp4(
6291 &self,
6292 t: &mut crate::model::GpuTensor,
6293 ) -> Result<(), Box<dyn std::error::Error>> {
6294 use crate::model::GpuTensor;
6295 let GpuTensor::Quant {
6296 bytes,
6297 qtype,
6298 row_bytes,
6299 ne,
6300 rp4,
6301 ..
6302 } = t
6303 else {
6304 return Ok(());
6305 };
6306 if *qtype != QT_Q8_0 || rp4.is_some() || ne.len() != 2 {
6307 return Ok(());
6308 }
6309 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
6310 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 34 {
6311 return Ok(());
6312 }
6313 *rp4 = Some(self.build_q8_rp4_raw(bytes, in_f, out_f)?);
6314 Ok(())
6315 }
6316
6317 pub fn build_q8_rp4_raw(
6320 &self,
6321 bytes: &CudaSlice<u8>,
6322 in_f: usize,
6323 out_f: usize,
6324 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
6325 assert!(in_f % 32 == 0);
6326 let nblk = in_f / 32;
6327 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 34)?;
6328 let f = self.func("q8_0_split_rp_build");
6329 let cfg = LaunchConfig {
6330 grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
6331 block_dim: (256, 1, 1),
6332 shared_mem_bytes: 0,
6333 };
6334 let (of, nb) = (out_f as i32, nblk as i32);
6335 let __s_b = self.gpu.stream();
6336 let mut b = __s_b.launch_builder(&f);
6337 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
6338 unsafe {
6339 b.launch(cfg)?;
6340 }
6341 Ok(dst)
6342 }
6343
6344 pub fn build_q4k_rp4(
6352 &self,
6353 t: &mut crate::model::GpuTensor,
6354 ) -> Result<(), Box<dyn std::error::Error>> {
6355 use crate::model::GpuTensor;
6356 let GpuTensor::Quant {
6357 bytes,
6358 qtype,
6359 row_bytes,
6360 ne,
6361 rp4,
6362 ..
6363 } = t
6364 else {
6365 return Ok(());
6366 };
6367 if *qtype != QT_Q4_K || rp4.is_some() || ne.len() != 2 {
6368 return Ok(());
6369 }
6370 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
6371 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 144 {
6372 return Ok(());
6373 }
6374 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q4_K)?);
6375 Ok(())
6376 }
6377
6378 pub fn build_q6k_rp4(
6379 &self,
6380 t: &mut crate::model::GpuTensor,
6381 ) -> Result<(), Box<dyn std::error::Error>> {
6382 use crate::model::GpuTensor;
6383 let GpuTensor::Quant {
6384 bytes,
6385 qtype,
6386 row_bytes,
6387 ne,
6388 rp4,
6389 ..
6390 } = t
6391 else {
6392 return Ok(());
6393 };
6394 if *qtype != QT_Q6_K || rp4.is_some() || ne.len() != 2 {
6395 return Ok(());
6396 }
6397 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
6398 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 210 {
6399 return Ok(());
6400 }
6401 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q6_K)?);
6402 Ok(())
6403 }
6404
6405 pub fn build_kq_rp4_raw(
6407 &self,
6408 bytes: &CudaSlice<u8>,
6409 in_f: usize,
6410 out_f: usize,
6411 qtype: i32,
6412 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
6413 assert!(in_f % 256 == 0);
6414 let nsbk = in_f / 256;
6415 let (sb_bytes, kname) = match qtype {
6416 QT_Q4_K => (144usize, "q4_K_split_rp_build"),
6417 QT_Q6_K => (210usize, "q6_K_split_rp_build"),
6418 _ => return Err(format!("build_kq_rp4_raw: qtype {qtype} has no rp mirror").into()),
6419 };
6420 let mut dst = self.alloc_uninit::<u8>(out_f * nsbk * sb_bytes)?;
6421 let f = self.func(kname);
6422 let cfg = LaunchConfig {
6423 grid_dim: (((out_f * nsbk) as u32).div_ceil(256), 1, 1),
6424 block_dim: (256, 1, 1),
6425 shared_mem_bytes: 0,
6426 };
6427 let (of, nb) = (out_f as i32, nsbk as i32);
6428 let __s_b = self.gpu.stream();
6429 let mut b = __s_b.launch_builder(&f);
6430 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
6431 unsafe {
6432 b.launch(cfg)?;
6433 }
6434 Ok(dst)
6435 }
6436
6437 pub fn kqrp_enabled() -> bool {
6441 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6442 *ON.get_or_init(|| match std::env::var("MEMRA_KQRP").as_deref() {
6443 Ok("0") => false,
6444 Ok(_) => true,
6445 Err(_) => cfg!(memra_hopper_mma),
6446 })
6447 }
6448
6449 pub fn build_q4_rp_swap(
6455 &self,
6456 t: &mut crate::model::GpuTensor,
6457 ) -> Result<bool, Box<dyn std::error::Error>> {
6458 use crate::model::GpuTensor;
6459 if !matches!(t, GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0) {
6469 return Ok(false);
6470 }
6471 self.build_q4_rp4(t)?;
6472 self.gpu.stream().synchronize()?; let GpuTensor::Quant { bytes, rp4, rp, .. } = t else {
6474 return Ok(false);
6475 };
6476 match rp4.take() {
6477 Some(split) => {
6478 *bytes = split; *rp = true;
6480 Ok(true)
6481 }
6482 None => Ok(false),
6483 }
6484 }
6485
6486 pub fn q4rp_enabled() -> bool {
6488 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6489 *ON.get_or_init(|| {
6490 std::env::var("MEMRA_Q4RP")
6491 .map(|v| v != "0")
6492 .unwrap_or(true)
6493 })
6494 }
6495
6496 pub fn copy_rows_strided(
6499 &self,
6500 src: &CudaSlice<f32>,
6501 dst: &mut CudaSlice<f32>,
6502 row_elems: usize,
6503 n_rows: usize,
6504 src_stride: usize,
6505 src_off: usize,
6506 ) -> Result<(), Box<dyn std::error::Error>> {
6507 let f = self.func("copy_rows_strided_f32");
6508 let cfg = LaunchConfig {
6509 grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
6510 block_dim: (256, 1, 1),
6511 shared_mem_bytes: 0,
6512 };
6513 let (re, nr) = (row_elems as i32, n_rows as i32);
6514 let (st, off) = (src_stride as i64, src_off as i64);
6515 let __s_b = self.gpu.stream();
6516 let mut b = __s_b.launch_builder(&f);
6517 b.arg(src)
6518 .arg(&mut *dst)
6519 .arg(&re)
6520 .arg(&nr)
6521 .arg(&st)
6522 .arg(&off);
6523 unsafe {
6524 b.launch(cfg)?;
6525 }
6526 Ok(())
6527 }
6528
6529 pub fn place_rows_strided(
6535 &self,
6536 src: &CudaSlice<f32>,
6537 dst: &mut CudaSlice<f32>,
6538 row_elems: usize,
6539 n_rows: usize,
6540 dst_stride: usize,
6541 dst_off: usize,
6542 ) -> Result<(), Box<dyn std::error::Error>> {
6543 if row_elems == 0 || n_rows == 0 {
6544 return Err("strided row placement requires nonzero rows and row width".into());
6545 }
6546 let src_len = n_rows
6547 .checked_mul(row_elems)
6548 .ok_or("strided row placement source size overflow")?;
6549 let dst_len = n_rows
6550 .checked_sub(1)
6551 .and_then(|rows| rows.checked_mul(dst_stride))
6552 .and_then(|base| base.checked_add(dst_off))
6553 .and_then(|base| base.checked_add(row_elems))
6554 .ok_or("strided row placement destination size overflow")?;
6555 let row_end = dst_off
6556 .checked_add(row_elems)
6557 .ok_or("strided row placement row size overflow")?;
6558 if src.len() < src_len || dst.len() < dst_len || row_end > dst_stride {
6559 return Err(format!(
6560 "strided row placement geometry mismatch: src={} need_src={src_len} \
6561 dst={} need_dst={dst_len} row_elems={row_elems} rows={n_rows} \
6562 dst_stride={dst_stride} dst_off={dst_off}",
6563 src.len(),
6564 dst.len(),
6565 )
6566 .into());
6567 }
6568 if row_elems > i32::MAX as usize || n_rows > i32::MAX as usize {
6569 return Err("strided row placement exceeds CUDA kernel geometry".into());
6570 }
6571 let f = self.func("place_rows_strided_f32");
6572 let cfg = LaunchConfig {
6573 grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
6574 block_dim: (256, 1, 1),
6575 shared_mem_bytes: 0,
6576 };
6577 let (re, nr) = (row_elems as i32, n_rows as i32);
6578 let (st, off) = (dst_stride as i64, dst_off as i64);
6579 let __s_b = self.gpu.stream();
6580 let mut b = __s_b.launch_builder(&f);
6581 b.arg(src)
6582 .arg(&mut *dst)
6583 .arg(&re)
6584 .arg(&nr)
6585 .arg(&st)
6586 .arg(&off);
6587 unsafe {
6588 b.launch(cfg)?;
6589 }
6590 Ok(())
6591 }
6592
6593 pub fn u32_set_k(
6595 &self,
6596 dst: &mut CudaSlice<u32>,
6597 v: u32,
6598 idx: usize,
6599 ) -> Result<(), Box<dyn std::error::Error>> {
6600 let f = self.func("u32_set_k");
6601 let cfg = LaunchConfig {
6602 grid_dim: (1, 1, 1),
6603 block_dim: (1, 1, 1),
6604 shared_mem_bytes: 0,
6605 };
6606 let ii = idx as i32;
6607 let __s_b = self.gpu.stream();
6608 let mut b = __s_b.launch_builder(&f);
6609 b.arg(dst).arg(&v).arg(&ii);
6610 unsafe {
6611 b.launch(cfg)?;
6612 }
6613 Ok(())
6614 }
6615
6616 pub fn i32_add_k(
6618 &self,
6619 d: &mut CudaSlice<i32>,
6620 v: i32,
6621 ) -> Result<(), Box<dyn std::error::Error>> {
6622 let f = self.func("i32_add_k");
6623 let cfg = LaunchConfig {
6624 grid_dim: (1, 1, 1),
6625 block_dim: (32, 1, 1),
6626 shared_mem_bytes: 0,
6627 };
6628 let __s_b = self.gpu.stream();
6629 let mut b = __s_b.launch_builder(&f);
6630 b.arg(d).arg(&v);
6631 unsafe {
6632 b.launch(cfg)?;
6633 }
6634 Ok(())
6635 }
6636
6637 pub fn i32_iota_from(
6639 &self,
6640 ctr: &CudaSlice<i32>,
6641 dst: &mut CudaSlice<i32>,
6642 n: usize,
6643 ) -> Result<(), Box<dyn std::error::Error>> {
6644 let f = self.func("i32_iota_from");
6645 let cfg = LaunchConfig::for_num_elems(n as u32);
6646 let ni = n as i32;
6647 let __s_b = self.gpu.stream();
6648 let mut b = __s_b.launch_builder(&f);
6649 b.arg(ctr).arg(dst).arg(&ni);
6650 unsafe {
6651 b.launch(cfg)?;
6652 }
6653 Ok(())
6654 }
6655
6656 pub fn u32_map_k(
6658 &self,
6659 buf: &mut CudaSlice<u32>,
6660 map: &CudaSlice<u32>,
6661 idx: usize,
6662 ) -> Result<(), Box<dyn std::error::Error>> {
6663 let f = self.func("u32_map_k");
6664 let cfg = LaunchConfig {
6665 grid_dim: (1, 1, 1),
6666 block_dim: (1, 1, 1),
6667 shared_mem_bytes: 0,
6668 };
6669 let ii = idx as i32;
6670 let __s_b = self.gpu.stream();
6671 let mut b = __s_b.launch_builder(&f);
6672 b.arg(buf).arg(map).arg(&ii);
6673 unsafe {
6674 b.launch(cfg)?;
6675 }
6676 Ok(())
6677 }
6678
6679 #[allow(clippy::too_many_arguments)]
6681 pub fn u32_pack2(
6682 &self,
6683 a: &CudaSlice<u32>,
6684 off_a: usize,
6685 n1: usize,
6686 b_in: &CudaSlice<u32>,
6687 n2: usize,
6688 out: &mut CudaSlice<u32>,
6689 ) -> Result<(), Box<dyn std::error::Error>> {
6690 let f = self.func("u32_pack2");
6691 let cfg = LaunchConfig::for_num_elems((n1 + n2) as u32);
6692 let (oa, i1, i2) = (off_a as i32, n1 as i32, n2 as i32);
6693 let __s_b = self.gpu.stream();
6694 let mut b = __s_b.launch_builder(&f);
6695 b.arg(a).arg(&oa).arg(&i1).arg(b_in).arg(&i2).arg(out);
6696 unsafe {
6697 b.launch(cfg)?;
6698 }
6699 Ok(())
6700 }
6701
6702 pub fn moe_w_exscale(
6704 &self,
6705 w: &mut CudaSlice<f32>,
6706 sel: &CudaSlice<i32>,
6707 s: &CudaSlice<f32>,
6708 n: usize,
6709 ) -> Result<(), Box<dyn std::error::Error>> {
6710 let f = self.func("moe_w_exscale");
6711 let cfg = LaunchConfig::for_num_elems(n as u32);
6712 let ni = n as i32;
6713 let __s_b = self.gpu.stream();
6714 let mut b = __s_b.launch_builder(&f);
6715 b.arg(w).arg(sel).arg(s).arg(&ni);
6716 unsafe {
6717 b.launch(cfg)?;
6718 }
6719 Ok(())
6720 }
6721
6722 pub fn moe_w_scale_by_expert(
6725 &self,
6726 w: &mut CudaSlice<f32>,
6727 sel: &CudaSlice<i32>,
6728 macros: &CudaSlice<f32>,
6729 n_expert: usize,
6730 n: usize,
6731 ) -> Result<(), Box<dyn std::error::Error>> {
6732 let f = self.func("moe_w_scale_by_expert");
6733 let cfg = LaunchConfig {
6734 grid_dim: (n.div_ceil(64) as u32, 1, 1),
6735 block_dim: (64, 1, 1),
6736 shared_mem_bytes: 0,
6737 };
6738 let (ne, nn) = (n_expert as i32, n as i32);
6739 let __s_b = self.gpu.stream();
6740 let mut b = __s_b.launch_builder(&f);
6741 b.arg(w).arg(sel).arg(macros).arg(&ne).arg(&nn);
6742 unsafe {
6743 b.launch(cfg)?;
6744 }
6745 Ok(())
6746 }
6747
6748 pub fn moe_gate_up_silu8_dev_q8(
6749 &self,
6750 table: &CudaSlice<u64>,
6751 sel: &cudarc::driver::CudaView<i32>,
6752 aq: &CudaSlice<i8>,
6753 ad: &CudaSlice<f32>,
6754 in_f: usize,
6755 n_ff: usize,
6756 n_used: usize,
6757 n_expert: usize,
6758 qt_g: i32,
6759 qt_u: i32,
6760 rb_g: usize,
6761 rb_u: usize,
6762 macros: &CudaSlice<f32>,
6763 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6764 static GU: std::sync::OnceLock<(String, u32)> = std::sync::OnceLock::new();
6765 let (mode, wpb) = GU.get_or_init(|| {
6766 let mode = std::env::var("MEMRA_MOE_DEVQ8_GU").unwrap_or_default();
6767 let wpb = std::env::var("MEMRA_MOE_DEVQ8_WPB")
6768 .ok()
6769 .and_then(|v| v.parse().ok())
6770 .unwrap_or(4u32)
6771 .clamp(1, 16);
6772 (mode, wpb)
6773 });
6774 let (mode, wpb) = (mode.as_str(), *wpb);
6775 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
6776 let (inf, nff, ne, rbg, rbu) = (
6777 in_f as i32,
6778 n_ff as i32,
6779 n_expert as i32,
6780 rb_g as i64,
6781 rb_u as i64,
6782 );
6783 let (f, cfg) = match mode {
6784 "1" | "2" | "4" => {
6785 let rpw: u32 = mode.parse().unwrap();
6786 let f = self.func(match rpw {
6787 1 => "moe_gate_up_silu8_dev_q8_r1",
6788 2 => "moe_gate_up_silu8_dev_q8_r2",
6789 _ => "moe_gate_up_silu8_dev_q8_r4",
6790 });
6791 let rows_per_block = (rpw * wpb) as usize;
6792 let gx = n_ff.div_ceil(rows_per_block) as u32;
6793 (
6794 f,
6795 LaunchConfig {
6796 grid_dim: (gx, n_used as u32, 1),
6797 block_dim: (32, wpb, 1),
6798 shared_mem_bytes: 0,
6799 },
6800 )
6801 }
6802 "j8" if n_used <= 32 => (
6803 self.func("moe_gate_up_silu8_dev_q8_j8"),
6804 LaunchConfig {
6805 grid_dim: (n_ff as u32, 1, 1),
6806 block_dim: (32, n_used as u32, 1),
6807 shared_mem_bytes: 0,
6808 },
6809 ),
6810 "vsm2" => {
6812 let f = self.func("moe_gate_up_silu8_dev_q8_vsm2");
6813 let sh = (rb_g + rb_u) as u32;
6814 use cudarc::driver::sys::CUfunction_attribute_enum as A;
6815 f.set_attribute(
6816 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
6817 sh as i32,
6818 )?;
6819 (
6820 f,
6821 LaunchConfig {
6822 grid_dim: (n_ff as u32, n_used as u32, 1),
6823 block_dim: (32, 1, 1),
6824 shared_mem_bytes: sh,
6825 },
6826 )
6827 }
6828 "vsm" => {
6829 let f = self.func("moe_gate_up_silu8_dev_q8_vsm");
6830 let sh = (rb_g + rb_u) as u32;
6831 use cudarc::driver::sys::CUfunction_attribute_enum as A;
6832 f.set_attribute(
6833 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
6834 sh as i32,
6835 )?;
6836 (
6837 f,
6838 LaunchConfig {
6839 grid_dim: (n_ff as u32, n_used as u32, 1),
6840 block_dim: (32, 1, 1),
6841 shared_mem_bytes: sh,
6842 },
6843 )
6844 }
6845 "sg" => (
6846 self.func("moe_gate_up_silu8_dev_q8_sg"),
6847 LaunchConfig {
6848 grid_dim: (n_ff as u32, n_used as u32, 1),
6849 block_dim: (32, 1, 1),
6850 shared_mem_bytes: 0,
6851 },
6852 ),
6853 "j8sg" if n_used <= 32 => (
6854 self.func("moe_gate_up_silu8_dev_q8_j8sg"),
6855 LaunchConfig {
6856 grid_dim: (n_ff as u32, 1, 1),
6857 block_dim: (32, n_used as u32, 1),
6858 shared_mem_bytes: 0,
6859 },
6860 ),
6861 "u64" if in_f == 2048 => (
6862 self.func("moe_gate_up_silu8_dev_q8_u64"),
6863 LaunchConfig {
6864 grid_dim: (n_ff as u32, n_used as u32, 1),
6865 block_dim: (32, 1, 1),
6866 shared_mem_bytes: 0,
6867 },
6868 ),
6869 "gs4" if in_f == 2048 => (
6870 self.func("moe_gate_up_silu8_dev_q8_gs4"),
6871 LaunchConfig {
6872 grid_dim: (n_ff as u32, n_used as u32, 1),
6873 block_dim: (32, 4, 1),
6874 shared_mem_bytes: 0,
6875 },
6876 ),
6877 "v" | "" => (
6879 self.func("moe_gate_up_silu8_dev_q8_v"),
6880 LaunchConfig {
6881 grid_dim: (n_ff as u32, n_used as u32, 1),
6882 block_dim: (32, 1, 1),
6883 shared_mem_bytes: 0,
6884 },
6885 ),
6886 "s2" => (
6887 self.func("moe_gate_up_silu8_dev_q8_s2"),
6888 LaunchConfig {
6889 grid_dim: (n_ff as u32, n_used as u32, 1),
6890 block_dim: (32, 2, 1),
6891 shared_mem_bytes: 0,
6892 },
6893 ),
6894 "s2z" => {
6895 let rz = wpb.min(16); (
6897 self.func("moe_gate_up_silu8_dev_q8_s2z"),
6898 LaunchConfig {
6899 grid_dim: (n_ff.div_ceil(rz as usize) as u32, n_used as u32, 1),
6900 block_dim: (32, 2, rz),
6901 shared_mem_bytes: 0,
6902 },
6903 )
6904 }
6905 _ => (
6906 self.func("moe_gate_up_silu8_dev_q8"),
6907 LaunchConfig {
6908 grid_dim: (n_ff as u32, n_used as u32, 1),
6909 block_dim: (32, 1, 1),
6910 shared_mem_bytes: 0,
6911 },
6912 ),
6913 };
6914 let __s_b = self.gpu.stream();
6915 let mut b = __s_b.launch_builder(&f);
6916 b.arg(table)
6917 .arg(sel)
6918 .arg(aq)
6919 .arg(ad)
6920 .arg(&mut act)
6921 .arg(&inf)
6922 .arg(&nff)
6923 .arg(&ne)
6924 .arg(&qt_g)
6925 .arg(&qt_u)
6926 .arg(&rbg)
6927 .arg(&rbu)
6928 .arg(macros);
6929 unsafe {
6930 b.launch(cfg)?;
6931 }
6932 Ok(act)
6933 }
6934
6935 #[allow(clippy::too_many_arguments)]
6936 pub fn moe_down8_fma_dev_q8(
6937 &self,
6938 table: &CudaSlice<u64>,
6939 sel: &cudarc::driver::CudaView<i32>,
6940 w: &cudarc::driver::CudaView<f32>,
6941 aq2: &CudaSlice<i8>,
6942 ad2: &CudaSlice<f32>,
6943 dst: &mut cudarc::driver::CudaViewMut<f32>,
6944 in_f: usize,
6945 out_f: usize,
6946 n_used: usize,
6947 n_expert: usize,
6948 qt: i32,
6949 rb: usize,
6950 ) -> Result<(), Box<dyn std::error::Error>> {
6951 static DOWN: std::sync::OnceLock<String> = std::sync::OnceLock::new();
6952 let mode = DOWN.get_or_init(|| std::env::var("MEMRA_MOE_DEVQ8_DOWN").unwrap_or_default());
6953 let (inf, outf, nu, ne, rbi) = (
6954 in_f as i32,
6955 out_f as i32,
6956 n_used as i32,
6957 n_expert as i32,
6958 rb as i64,
6959 );
6960 let (f, cfg) = match mode.as_str() {
6963 m @ ("1" | "2" | "4") if n_used <= 8 => {
6964 let rpw: usize = m.parse().unwrap();
6965 let f = self.func(match rpw {
6966 1 => "moe_down8_fma_dev_q8_w8r1",
6967 2 => "moe_down8_fma_dev_q8_w8r2",
6968 _ => "moe_down8_fma_dev_q8_w8r4",
6969 });
6970 (
6971 f,
6972 LaunchConfig {
6973 grid_dim: (out_f.div_ceil(rpw) as u32, 1, 1),
6974 block_dim: (32, n_used as u32, 1),
6975 shared_mem_bytes: 0,
6976 },
6977 )
6978 }
6979 "h2" if in_f == 512 => (
6980 self.func("moe_down8_fma_dev_q8_h2"),
6981 LaunchConfig {
6982 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
6983 block_dim: (32, 1, 1),
6984 shared_mem_bytes: 0,
6985 },
6986 ),
6987 "" if in_f == 704 && n_used <= 8 => (
6990 self.func("moe_down8_fma_dev_q8_w8r2"),
6991 LaunchConfig {
6992 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
6993 block_dim: (32, n_used as u32, 1),
6994 shared_mem_bytes: 0,
6995 },
6996 ),
6997 "w8h2v" | "" if in_f == 512 && n_used <= 8 => (
7001 self.func("moe_down8_fma_dev_q8_w8h2v"),
7002 LaunchConfig {
7003 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
7004 block_dim: (32, n_used as u32, 1),
7005 shared_mem_bytes: 0,
7006 },
7007 ),
7008 "w8h2r2v" if in_f == 512 && n_used <= 8 => (
7009 self.func("moe_down8_fma_dev_q8_w8h2r2v"),
7010 LaunchConfig {
7011 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
7012 block_dim: (32, n_used as u32, 1),
7013 shared_mem_bytes: 0,
7014 },
7015 ),
7016 "w8h2r2" if in_f == 512 && n_used <= 8 => (
7017 self.func("moe_down8_fma_dev_q8_w8h2r2"),
7018 LaunchConfig {
7019 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
7020 block_dim: (32, n_used as u32, 1),
7021 shared_mem_bytes: 0,
7022 },
7023 ),
7024 "w8h2" if in_f == 512 && n_used <= 8 => (
7025 self.func("moe_down8_fma_dev_q8_w8h2"),
7026 LaunchConfig {
7027 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
7028 block_dim: (32, n_used as u32, 1),
7029 shared_mem_bytes: 0,
7030 },
7031 ),
7032 _ => (
7033 self.func("moe_down8_fma_dev_q8"),
7034 LaunchConfig {
7035 grid_dim: (out_f as u32, 1, 1),
7036 block_dim: (32, 1, 1),
7037 shared_mem_bytes: 0,
7038 },
7039 ),
7040 };
7041 let __s_b = self.gpu.stream();
7042 let mut b = __s_b.launch_builder(&f);
7043 b.arg(table)
7044 .arg(sel)
7045 .arg(w)
7046 .arg(aq2)
7047 .arg(ad2)
7048 .arg(dst)
7049 .arg(&inf)
7050 .arg(&outf)
7051 .arg(&nu)
7052 .arg(&ne)
7053 .arg(&qt)
7054 .arg(&rbi);
7055 unsafe {
7056 b.launch(cfg)?;
7057 }
7058 Ok(())
7059 }
7060
7061 #[allow(clippy::too_many_arguments)]
7068 pub fn moe_gate_up_silu8_dev_q8_rows(
7069 &self,
7070 table: &CudaSlice<u64>,
7071 sel: &CudaSlice<i32>,
7072 aq: &CudaSlice<i8>,
7073 ad: &CudaSlice<f32>,
7074 t: usize,
7075 in_f: usize,
7076 n_ff: usize,
7077 n_used: usize,
7078 n_expert: usize,
7079 qt_g: i32,
7080 qt_u: i32,
7081 rb_g: usize,
7082 rb_u: usize,
7083 macros: &CudaSlice<f32>,
7084 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7085 let f = self.func("moe_gate_up_silu8_dev_q8_v_rows");
7086 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
7087 let cfg = LaunchConfig {
7088 grid_dim: (n_ff as u32, n_used as u32, t as u32),
7089 block_dim: (32, 1, 1),
7090 shared_mem_bytes: 0,
7091 };
7092 let (inf, nff, ne, nu, rbg, rbu) = (
7093 in_f as i32,
7094 n_ff as i32,
7095 n_expert as i32,
7096 n_used as i32,
7097 rb_g as i64,
7098 rb_u as i64,
7099 );
7100 let __s_b = self.gpu.stream();
7101 let mut b = __s_b.launch_builder(&f);
7102 b.arg(table)
7103 .arg(sel)
7104 .arg(aq)
7105 .arg(ad)
7106 .arg(&mut act)
7107 .arg(&inf)
7108 .arg(&nff)
7109 .arg(&ne)
7110 .arg(&qt_g)
7111 .arg(&qt_u)
7112 .arg(&rbg)
7113 .arg(&rbu)
7114 .arg(&nu)
7115 .arg(macros);
7116 unsafe {
7117 b.launch(cfg)?;
7118 }
7119 Ok(act)
7120 }
7121
7122 #[allow(clippy::too_many_arguments)]
7127 pub fn moe_down8_fma_dev_q8_rows(
7128 &self,
7129 table: &CudaSlice<u64>,
7130 sel: &CudaSlice<i32>,
7131 w: &CudaSlice<f32>,
7132 aq2: &CudaSlice<i8>,
7133 ad2: &CudaSlice<f32>,
7134 dst: &mut CudaSlice<f32>,
7135 t: usize,
7136 in_f: usize,
7137 out_f: usize,
7138 n_used: usize,
7139 n_expert: usize,
7140 qt: i32,
7141 rb: usize,
7142 ) -> Result<(), Box<dyn std::error::Error>> {
7143 assert!(
7144 in_f == 512 && n_used <= 8,
7145 "down rows twin is w8h2v shape-gated"
7146 );
7147 let f = self.func("moe_down8_fma_dev_q8_w8h2v_rows");
7148 let cfg = LaunchConfig {
7149 grid_dim: (out_f.div_ceil(2) as u32, 1, t as u32),
7150 block_dim: (32, n_used as u32, 1),
7151 shared_mem_bytes: 0,
7152 };
7153 let (inf, outf, nu, ne, rbi) = (
7154 in_f as i32,
7155 out_f as i32,
7156 n_used as i32,
7157 n_expert as i32,
7158 rb as i64,
7159 );
7160 let __s_b = self.gpu.stream();
7161 let mut b = __s_b.launch_builder(&f);
7162 b.arg(table)
7163 .arg(sel)
7164 .arg(w)
7165 .arg(aq2)
7166 .arg(ad2)
7167 .arg(dst)
7168 .arg(&inf)
7169 .arg(&outf)
7170 .arg(&nu)
7171 .arg(&ne)
7172 .arg(&qt)
7173 .arg(&rbi);
7174 unsafe {
7175 b.launch(cfg)?;
7176 }
7177 Ok(())
7178 }
7179
7180 #[allow(clippy::too_many_arguments)]
7184 pub fn moe_gate_up_silu8_dev_q8_csr(
7185 &self,
7186 table: &CudaSlice<u64>,
7187 sel: &CudaSlice<i32>,
7188 aq: &CudaSlice<i8>,
7189 ad: &CudaSlice<f32>,
7190 n_pairs: usize,
7191 in_f: usize,
7192 n_ff: usize,
7193 n_used: usize,
7194 n_expert: usize,
7195 qt_g: i32,
7196 qt_u: i32,
7197 rb_g: usize,
7198 rb_u: usize,
7199 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7200 let f = if qt_g == crate::QT_NVFP4 {
7203 self.func("moe_gate_up_silu8_dev_q8_csr_nvfp4")
7204 } else {
7205 self.func("moe_gate_up_silu8_dev_q8_csr_iq4")
7206 };
7207 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
7208 let cfg = LaunchConfig {
7209 grid_dim: (n_ff as u32, n_pairs as u32, 1),
7210 block_dim: (32, 1, 1),
7211 shared_mem_bytes: 0,
7212 };
7213 let (inf, nff, ne, nu, npi, rbg, rbu) = (
7214 in_f as i32,
7215 n_ff as i32,
7216 n_expert as i32,
7217 n_used as i32,
7218 n_pairs as i32,
7219 rb_g as i64,
7220 rb_u as i64,
7221 );
7222 let __s_b = self.gpu.stream();
7223 let mut b = __s_b.launch_builder(&f);
7224 b.arg(table)
7225 .arg(sel)
7226 .arg(aq)
7227 .arg(ad)
7228 .arg(&mut act)
7229 .arg(&inf)
7230 .arg(&nff)
7231 .arg(&ne)
7232 .arg(&qt_g)
7233 .arg(&qt_u)
7234 .arg(&rbg)
7235 .arg(&rbu)
7236 .arg(&nu)
7237 .arg(&npi);
7238 unsafe {
7239 b.launch(cfg)?;
7240 }
7241 Ok(act)
7242 }
7243
7244 #[allow(clippy::too_many_arguments)]
7248 pub fn moe_down8_fma_dev_q8_variant(
7249 &self,
7250 variant: &str,
7251 table: &CudaSlice<u64>,
7252 sel: &cudarc::driver::CudaView<i32>,
7253 w: &cudarc::driver::CudaView<f32>,
7254 aq2: &CudaSlice<i8>,
7255 ad2: &CudaSlice<f32>,
7256 dst: &mut cudarc::driver::CudaViewMut<f32>,
7257 in_f: usize,
7258 out_f: usize,
7259 n_used: usize,
7260 n_expert: usize,
7261 qt: i32,
7262 rb: usize,
7263 ) -> Result<(), Box<dyn std::error::Error>> {
7264 let (inf, outf, nu, ne, rbi) = (
7265 in_f as i32,
7266 out_f as i32,
7267 n_used as i32,
7268 n_expert as i32,
7269 rb as i64,
7270 );
7271 let (f, cfg) = match variant {
7272 "w8h2" | "w8h2v" => (
7273 self.func(if variant == "w8h2" {
7274 "moe_down8_fma_dev_q8_w8h2"
7275 } else {
7276 "moe_down8_fma_dev_q8_w8h2v"
7277 }),
7278 LaunchConfig {
7279 grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
7280 block_dim: (32, n_used as u32, 1),
7281 shared_mem_bytes: 0,
7282 },
7283 ),
7284 "w8h2r2" | "w8h2r2v" => (
7285 self.func(if variant == "w8h2r2" {
7286 "moe_down8_fma_dev_q8_w8h2r2"
7287 } else {
7288 "moe_down8_fma_dev_q8_w8h2r2v"
7289 }),
7290 LaunchConfig {
7291 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
7292 block_dim: (32, n_used as u32, 1),
7293 shared_mem_bytes: 0,
7294 },
7295 ),
7296 _ => (
7297 self.func("moe_down8_fma_dev_q8"),
7298 LaunchConfig {
7299 grid_dim: (out_f as u32, 1, 1),
7300 block_dim: (32, 1, 1),
7301 shared_mem_bytes: 0,
7302 },
7303 ),
7304 };
7305 let __s_b = self.gpu.stream();
7306 let mut b = __s_b.launch_builder(&f);
7307 b.arg(table)
7308 .arg(sel)
7309 .arg(w)
7310 .arg(aq2)
7311 .arg(ad2)
7312 .arg(dst)
7313 .arg(&inf)
7314 .arg(&outf)
7315 .arg(&nu)
7316 .arg(&ne)
7317 .arg(&qt)
7318 .arg(&rbi);
7319 unsafe {
7320 b.launch(cfg)?;
7321 }
7322 Ok(())
7323 }
7324
7325 #[allow(clippy::too_many_arguments)]
7327 pub fn moe_gate_up_silu8_dev_q8_variant(
7328 &self,
7329 variant: &str,
7330 table: &CudaSlice<u64>,
7331 sel: &cudarc::driver::CudaView<i32>,
7332 aq: &CudaSlice<i8>,
7333 ad: &CudaSlice<f32>,
7334 in_f: usize,
7335 n_ff: usize,
7336 n_used: usize,
7337 n_expert: usize,
7338 qt_g: i32,
7339 qt_u: i32,
7340 rb_g: usize,
7341 rb_u: usize,
7342 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7343 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
7344 let (inf, nff, ne, rbg, rbu) = (
7345 in_f as i32,
7346 n_ff as i32,
7347 n_expert as i32,
7348 rb_g as i64,
7349 rb_u as i64,
7350 );
7351 let f = self.func(if variant == "v" {
7352 "moe_gate_up_silu8_dev_q8_v"
7353 } else {
7354 "moe_gate_up_silu8_dev_q8"
7355 });
7356 let cfg = LaunchConfig {
7357 grid_dim: (n_ff as u32, n_used as u32, 1),
7358 block_dim: (32, 1, 1),
7359 shared_mem_bytes: 0,
7360 };
7361 let __s_b = self.gpu.stream();
7362 let mut b = __s_b.launch_builder(&f);
7363 b.arg(table)
7364 .arg(sel)
7365 .arg(aq)
7366 .arg(ad)
7367 .arg(&mut act)
7368 .arg(&inf)
7369 .arg(&nff)
7370 .arg(&ne)
7371 .arg(&qt_g)
7372 .arg(&qt_u)
7373 .arg(&rbg)
7374 .arg(&rbu);
7375 unsafe {
7376 b.launch(cfg)?;
7377 }
7378 Ok(act)
7379 }
7380
7381 pub fn moe_gate_up_silu8_dev(
7382 &self,
7383 table: &CudaSlice<u64>,
7384 sel: &cudarc::driver::CudaView<i32>,
7385 x: &cudarc::driver::CudaView<f32>,
7386 in_f: usize,
7387 n_ff: usize,
7388 n_used: usize,
7389 n_expert: usize,
7390 qt_g: i32,
7391 qt_u: i32,
7392 rb_g: usize,
7393 rb_u: usize,
7394 macros: &CudaSlice<f32>,
7395 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7396 let f = self.func("moe_gate_up_silu8_dev");
7397 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig {
7399 grid_dim: (n_ff as u32, n_used as u32, 1),
7400 block_dim: (256, 1, 1),
7401 shared_mem_bytes: 0,
7402 };
7403 let (inf, nff, ne, rbg, rbu) = (
7404 in_f as i32,
7405 n_ff as i32,
7406 n_expert as i32,
7407 rb_g as i64,
7408 rb_u as i64,
7409 );
7410 let __s_b = self.gpu.stream();
7411 let mut b = __s_b.launch_builder(&f);
7412 b.arg(table)
7413 .arg(sel)
7414 .arg(x)
7415 .arg(&mut act)
7416 .arg(&inf)
7417 .arg(&nff)
7418 .arg(&ne)
7419 .arg(&qt_g)
7420 .arg(&qt_u)
7421 .arg(&rbg)
7422 .arg(&rbu)
7423 .arg(macros);
7424 unsafe {
7425 b.launch(cfg)?;
7426 }
7427 Ok(act)
7428 }
7429
7430 #[allow(clippy::too_many_arguments)]
7433 pub fn moe_down8_fma_dev(
7434 &self,
7435 table: &CudaSlice<u64>,
7436 sel: &cudarc::driver::CudaView<i32>,
7437 w: &cudarc::driver::CudaView<f32>,
7438 act: &CudaSlice<f32>,
7439 dst: &mut cudarc::driver::CudaViewMut<f32>,
7440 in_f: usize,
7441 out_f: usize,
7442 n_used: usize,
7443 n_expert: usize,
7444 qt: i32,
7445 rb: usize,
7446 ) -> Result<(), Box<dyn std::error::Error>> {
7447 let f = self.func("moe_down8_fma_dev");
7448 let cfg = LaunchConfig {
7449 grid_dim: (out_f as u32, 1, 1),
7450 block_dim: (256, 1, 1),
7451 shared_mem_bytes: 0,
7452 };
7453 let (inf, outf, nu, ne, rbv) = (
7454 in_f as i32,
7455 out_f as i32,
7456 n_used as i32,
7457 n_expert as i32,
7458 rb as i64,
7459 );
7460 let __s_b = self.gpu.stream();
7461 let mut b = __s_b.launch_builder(&f);
7462 b.arg(table)
7463 .arg(sel)
7464 .arg(w)
7465 .arg(act)
7466 .arg(dst)
7467 .arg(&inf)
7468 .arg(&outf)
7469 .arg(&nu)
7470 .arg(&ne)
7471 .arg(&qt)
7472 .arg(&rbv);
7473 unsafe {
7474 b.launch(cfg)?;
7475 }
7476 Ok(())
7477 }
7478
7479 pub fn axpy_into(
7481 &self,
7482 src: &CudaSlice<f32>,
7483 alpha: f32,
7484 dst: &mut cudarc::driver::CudaViewMut<f32>,
7485 n: usize,
7486 ) -> Result<(), Box<dyn std::error::Error>> {
7487 let f = self.func("axpy_f32");
7488 let cfg = LaunchConfig::for_num_elems(n as u32);
7489 let (a, ni) = (alpha, n as i32);
7490 let __s_b = self.gpu.stream();
7491 let mut b = __s_b.launch_builder(&f);
7492 b.arg(src).arg(dst).arg(&a).arg(&ni);
7493 unsafe {
7494 b.launch(cfg)?;
7495 }
7496 Ok(())
7497 }
7498
7499 pub fn axpy_host_into(
7501 &self,
7502 src: &cudarc::driver::CudaView<'_, f32>,
7503 alpha: f32,
7504 dst: &mut cudarc::driver::CudaViewMut<f32>,
7505 n: usize,
7506 ) -> Result<(), Box<dyn std::error::Error>> {
7507 let f = self.func("axpy_host_f32");
7508 let cfg = LaunchConfig::for_num_elems(n as u32);
7509 let (a, ni) = (alpha, n as i32);
7510 let __s_b = self.gpu.stream();
7511 let mut b = __s_b.launch_builder(&f);
7512 b.arg(src).arg(dst).arg(&a).arg(&ni);
7513 unsafe {
7514 b.launch(cfg)?;
7515 }
7516 Ok(())
7517 }
7518
7519 pub fn add_scaled_rows(
7521 &self,
7522 src: &CudaSlice<f32>,
7523 scale: &CudaSlice<f32>,
7524 dst: &mut CudaSlice<f32>,
7525 ncols: usize,
7526 nrows: usize,
7527 ) -> Result<(), Box<dyn std::error::Error>> {
7528 let f = self.func("add_scaled_rows_f32");
7529 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
7530 let (nc, nr) = (ncols as i32, nrows as i32);
7531 let __s_b = self.gpu.stream();
7532 let mut b = __s_b.launch_builder(&f);
7533 b.arg(src).arg(scale).arg(dst).arg(&nc).arg(&nr);
7534 unsafe {
7535 b.launch(cfg)?;
7536 }
7537 Ok(())
7538 }
7539
7540 pub fn scale_rows(
7543 &self,
7544 y: &mut CudaSlice<f32>,
7545 s: &CudaSlice<f32>,
7546 ncols: usize,
7547 nrows: usize,
7548 ) -> Result<(), Box<dyn std::error::Error>> {
7549 let f = self.func("scale_rows_f32");
7550 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
7551 let (nc, nr) = (ncols as i32, nrows as i32);
7552 let __s_b = self.gpu.stream();
7553 let mut b = __s_b.launch_builder(&f);
7554 b.arg(&mut *y).arg(s).arg(&nc).arg(&nr);
7555 unsafe {
7556 b.launch(cfg)?;
7557 }
7558 Ok(())
7559 }
7560
7561 #[allow(clippy::too_many_arguments)]
7565 pub fn moe_prime_join_scatter(
7566 &self,
7567 y0: &CudaSlice<f32>,
7568 y1: &CudaSlice<f32>,
7569 inv: &CudaSlice<i32>,
7570 w: &CudaSlice<f32>,
7571 out: &mut CudaSlice<f32>,
7572 ncols: usize,
7573 n_used: usize,
7574 t: usize,
7575 ) -> Result<(), Box<dyn std::error::Error>> {
7576 let f = self.func("moe_prime_join_scatter_f32");
7577 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
7578 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
7579 let __s_b = self.gpu.stream();
7580 let mut b = __s_b.launch_builder(&f);
7581 b.arg(y0)
7582 .arg(y1)
7583 .arg(inv)
7584 .arg(w)
7585 .arg(&mut *out)
7586 .arg(&nc)
7587 .arg(&nu)
7588 .arg(&ti);
7589 unsafe {
7590 b.launch(cfg)?;
7591 }
7592 Ok(())
7593 }
7594
7595 pub fn moe_pairs_weighted_scatter(
7598 &self,
7599 y: &CudaSlice<f32>,
7600 w: &CudaSlice<f32>,
7601 out: &mut CudaSlice<f32>,
7602 ncols: usize,
7603 n_used: usize,
7604 t: usize,
7605 ) -> Result<(), Box<dyn std::error::Error>> {
7606 let f = self.func("moe_pairs_weighted_scatter_f32");
7607 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
7608 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
7609 let __s_b = self.gpu.stream();
7610 let mut b = __s_b.launch_builder(&f);
7611 b.arg(y).arg(w).arg(&mut *out).arg(&nc).arg(&nu).arg(&ti);
7612 unsafe {
7613 b.launch(cfg)?;
7614 }
7615 Ok(())
7616 }
7617
7618 pub fn gather_rows(
7622 &self,
7623 src: &CudaSlice<f32>,
7624 idx: &CudaSlice<i32>,
7625 dst: &mut CudaSlice<f32>,
7626 ncols: usize,
7627 m_e: usize,
7628 ) -> Result<(), Box<dyn std::error::Error>> {
7629 let f = self.func("gather_rows_f32");
7630 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
7631 let (nc, me) = (ncols as i32, m_e as i32);
7632 let __s_b = self.gpu.stream();
7633 let mut b = __s_b.launch_builder(&f);
7634 b.arg(src).arg(idx).arg(dst).arg(&nc).arg(&me);
7635 unsafe {
7636 b.launch(cfg)?;
7637 }
7638 Ok(())
7639 }
7640
7641 pub fn scatter_slot(
7646 &self,
7647 src: &CudaSlice<f32>,
7648 tok_idx: &CudaSlice<i32>,
7649 slot_idx: &CudaSlice<i32>,
7650 weight: &CudaSlice<f32>,
7651 dst: &mut CudaSlice<f32>,
7652 wbuf: &mut CudaSlice<f32>,
7653 ncols: usize,
7654 n_used: usize,
7655 m_e: usize,
7656 ) -> Result<(), Box<dyn std::error::Error>> {
7657 let f = self.func("scatter_add_slot_f32");
7658 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
7659 let (nc, nu, me) = (ncols as i32, n_used as i32, m_e as i32);
7660 let __s_b = self.gpu.stream();
7661 let mut b = __s_b.launch_builder(&f);
7662 b.arg(src)
7663 .arg(tok_idx)
7664 .arg(slot_idx)
7665 .arg(weight)
7666 .arg(dst)
7667 .arg(wbuf)
7668 .arg(&nc)
7669 .arg(&nu)
7670 .arg(&me);
7671 unsafe {
7672 b.launch(cfg)?;
7673 }
7674 Ok(())
7675 }
7676
7677 pub fn reduce_slots(
7681 &self,
7682 slots: &CudaSlice<f32>,
7683 wbuf: &CudaSlice<f32>,
7684 dst: &mut CudaSlice<f32>,
7685 ncols: usize,
7686 n_used: usize,
7687 t: usize,
7688 ) -> Result<(), Box<dyn std::error::Error>> {
7689 let f = self.func("reduce_slots_f32");
7690 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
7691 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
7692 let __s_b = self.gpu.stream();
7693 let mut b = __s_b.launch_builder(&f);
7694 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
7695 unsafe {
7696 b.launch(cfg)?;
7697 }
7698 Ok(())
7699 }
7700
7701 pub fn reduce_slots_host(
7706 &self,
7707 slots: &CudaSlice<f32>,
7708 wbuf: &CudaSlice<f32>,
7709 dst: &mut CudaSlice<f32>,
7710 ncols: usize,
7711 n_used: usize,
7712 t: usize,
7713 ) -> Result<(), Box<dyn std::error::Error>> {
7714 let f = self.func("reduce_slots_host_f32");
7715 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
7716 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
7717 let __s_b = self.gpu.stream();
7718 let mut b = __s_b.launch_builder(&f);
7719 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
7720 unsafe {
7721 b.launch(cfg)?;
7722 }
7723 Ok(())
7724 }
7725
7726 pub fn quantize_q8_1_view(
7733 &self,
7734 x: &cudarc::driver::CudaView<f32>,
7735 m: usize,
7736 in_f: usize,
7737 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7738 let f = self.func("quantize_q8_1");
7739 let nblk = in_f / 32;
7740 let mut q = self.alloc_uninit::<i8>(m * in_f)?;
7741 let mut d = self.alloc_uninit::<f32>(m * nblk)?;
7742 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
7743 let (inf, mi) = (in_f as i32, m as i32);
7744 let __s_b = self.gpu.stream();
7745 let mut b = __s_b.launch_builder(&f);
7746 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
7747 unsafe {
7748 b.launch(cfg)?;
7749 }
7750 Ok((q, d))
7751 }
7752
7753 pub fn quantize_q8_1(
7754 &self,
7755 x: &CudaSlice<f32>,
7756 m: usize,
7757 in_f: usize,
7758 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7759 let nblk = in_f / 32;
7760 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);
7764 let (inf, mi) = (in_f as i32, m as i32);
7765 if Self::pdl_on() && Self::pdl_wb_on() {
7766 {
7767 use cudarc::driver::{DevicePtr, DevicePtrMut};
7768 let s = &self.gpu.stream();
7769 let (px, _g0) = x.device_ptr(s);
7770 let (pq, _g1) = q.device_ptr_mut(s);
7771 let (pd, _g2) = d.device_ptr_mut(s);
7772 let mut ps = [
7773 &px as *const _ as *mut std::ffi::c_void,
7774 &pq as *const _ as *mut _,
7775 &pd as *const _ as *mut _,
7776 &inf as *const _ as *mut _,
7777 &mi as *const _ as *mut _,
7778 ];
7779 unsafe {
7780 self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?;
7781 }
7782 }
7783 return Ok((q, d));
7784 }
7785 let f = self.func("quantize_q8_1");
7786 let __s_b = self.gpu.stream();
7787 let mut b = __s_b.launch_builder(&f);
7788 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
7789 unsafe {
7790 b.launch(cfg)?;
7791 }
7792 Ok((q, d))
7793 }
7794
7795 pub fn quantize_fp4_act(
7799 &self,
7800 x: &CudaSlice<f32>,
7801 m: usize,
7802 in_f: usize,
7803 ) -> Result<(CudaSlice<u32>, CudaSlice<u8>), Box<dyn std::error::Error>> {
7804 let f = self.func("quantize_fp4_act");
7805 let nb16 = in_f / 16;
7806 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);
7809 let (inf, mi) = (in_f as i32, m as i32);
7810 let __s_b = self.gpu.stream();
7811 let mut b = __s_b.launch_builder(&f);
7812 b.arg(x).arg(&mut aq4).arg(&mut ad4).arg(&inf).arg(&mi);
7813 unsafe {
7814 b.launch(cfg)?;
7815 }
7816 Ok((aq4, ad4))
7817 }
7818
7819 pub fn qmatvec_gemm_nvfp4_fp4(
7824 &self,
7825 bytes: &CudaSlice<u8>,
7826 x: &CudaSlice<f32>,
7827 m: usize,
7828 in_f: usize,
7829 out_f: usize,
7830 row_bytes: usize,
7831 scale: f32,
7832 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7833 assert!(
7834 in_f % 64 == 0,
7835 "FP4 GEMM requires in_f % 64 == 0, got {in_f}"
7836 );
7837 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
7838 let mut y = self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)?;
7839 if scale != 1.0 {
7840 self.scale_inplace(&mut y, scale, m * out_f)?;
7841 }
7842 Ok(y)
7843 }
7844
7845 fn fp4_gemm_launch(
7848 &self,
7849 bytes: &CudaSlice<u8>,
7850 aq4: &CudaSlice<u32>,
7851 ad4: &CudaSlice<u8>,
7852 m: usize,
7853 in_f: usize,
7854 out_f: usize,
7855 row_bytes: usize,
7856 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7857 let f = self.func("qmatvec_gemm_nvfp4_fp4");
7858 let mut y = self.alloc_uninit::<f32>(m * out_f)?; const BM: u32 = 64;
7860 const BN: u32 = 256;
7861 let cfg = LaunchConfig {
7862 grid_dim: ((out_f as u32 + BM - 1) / BM, (m as u32 + BN - 1) / BN, 1),
7863 block_dim: (32, 4, 1),
7864 shared_mem_bytes: 0,
7865 };
7866 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7867 let __s_b = self.gpu.stream();
7868 let mut b = __s_b.launch_builder(&f);
7869 b.arg(bytes)
7870 .arg(aq4)
7871 .arg(ad4)
7872 .arg(&mut y)
7873 .arg(&inf)
7874 .arg(&outf)
7875 .arg(&mi)
7876 .arg(&rb);
7877 unsafe {
7878 b.launch(cfg)?;
7879 }
7880 Ok(y)
7881 }
7882
7883 pub fn qmatvec_gemm_nvfp4_fp4_raw(
7885 &self,
7886 bytes: &CudaSlice<u8>,
7887 x: &CudaSlice<f32>,
7888 m: usize,
7889 in_f: usize,
7890 out_f: usize,
7891 row_bytes: usize,
7892 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7893 assert!(
7894 in_f % 64 == 0,
7895 "FP4 GEMM requires in_f % 64 == 0, got {in_f}"
7896 );
7897 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
7898 self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)
7899 }
7900
7901 pub fn qmatvec_q8_0_fast(
7903 &self,
7904 w: &CudaSlice<u8>,
7905 x: &CudaSlice<f32>,
7906 m: usize,
7907 in_f: usize,
7908 out_f: usize,
7909 row_bytes: usize,
7910 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7911 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7912 let f = self.func("qmatvec_q8_0_dp4a");
7913 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
7915 grid_dim: (out_f as u32, m as u32, 1),
7916 block_dim: (128, 1, 1),
7917 shared_mem_bytes: 0,
7918 };
7919 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7920 let __s_b = self.gpu.stream();
7921 let mut b = __s_b.launch_builder(&f);
7922 b.arg(w)
7923 .arg(&aq)
7924 .arg(&ad)
7925 .arg(&mut y)
7926 .arg(&inf)
7927 .arg(&outf)
7928 .arg(&mi)
7929 .arg(&rb);
7930 unsafe {
7931 b.launch(cfg)?;
7932 }
7933 Ok(y)
7934 }
7935
7936 #[allow(non_snake_case)] pub fn qmatvec_q4_K_fast(
7939 &self,
7940 w: &CudaSlice<u8>,
7941 x: &CudaSlice<f32>,
7942 m: usize,
7943 in_f: usize,
7944 out_f: usize,
7945 row_bytes: usize,
7946 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7947 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7948 let f = self.func("qmatvec_q4_K_dp4a");
7949 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
7951 grid_dim: (out_f as u32, m as u32, 1),
7952 block_dim: (128, 1, 1),
7953 shared_mem_bytes: 0,
7954 };
7955 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7956 let __s_b = self.gpu.stream();
7957 let mut b = __s_b.launch_builder(&f);
7958 b.arg(w)
7959 .arg(&aq)
7960 .arg(&ad)
7961 .arg(&mut y)
7962 .arg(&inf)
7963 .arg(&outf)
7964 .arg(&mi)
7965 .arg(&rb);
7966 unsafe {
7967 b.launch(cfg)?;
7968 }
7969 Ok(y)
7970 }
7971
7972 #[allow(non_snake_case)] pub fn qmatvec_q6_K_fast(
7975 &self,
7976 w: &CudaSlice<u8>,
7977 x: &CudaSlice<f32>,
7978 m: usize,
7979 in_f: usize,
7980 out_f: usize,
7981 row_bytes: usize,
7982 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7983 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7984 let f = self.func("qmatvec_q6_K_dp4a");
7985 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
7987 grid_dim: (out_f as u32, m as u32, 1),
7988 block_dim: (128, 1, 1),
7989 shared_mem_bytes: 0,
7990 };
7991 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7992 let __s_b = self.gpu.stream();
7993 let mut b = __s_b.launch_builder(&f);
7994 b.arg(w)
7995 .arg(&aq)
7996 .arg(&ad)
7997 .arg(&mut y)
7998 .arg(&inf)
7999 .arg(&outf)
8000 .arg(&mi)
8001 .arg(&rb);
8002 unsafe {
8003 b.launch(cfg)?;
8004 }
8005 Ok(y)
8006 }
8007
8008 #[allow(non_snake_case)] pub fn qmatvec_q5_K_fast(
8011 &self,
8012 w: &CudaSlice<u8>,
8013 x: &CudaSlice<f32>,
8014 m: usize,
8015 in_f: usize,
8016 out_f: usize,
8017 row_bytes: usize,
8018 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8019 self.qmatvec_dp4a_named(
8020 "qmatvec_q5_K_dp4a",
8021 &w.slice(0..w.len()),
8022 x,
8023 m,
8024 in_f,
8025 out_f,
8026 row_bytes,
8027 )
8028 }
8029 #[allow(non_snake_case)] pub fn qmatvec_q3_K_fast(
8032 &self,
8033 w: &CudaSlice<u8>,
8034 x: &CudaSlice<f32>,
8035 m: usize,
8036 in_f: usize,
8037 out_f: usize,
8038 row_bytes: usize,
8039 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8040 self.qmatvec_dp4a_named(
8041 "qmatvec_q3_K_dp4a",
8042 &w.slice(0..w.len()),
8043 x,
8044 m,
8045 in_f,
8046 out_f,
8047 row_bytes,
8048 )
8049 }
8050 pub fn qmatvec_nvfp4_fast_rp(
8052 &self,
8053 w: &CudaSlice<u8>,
8054 x: &CudaSlice<f32>,
8055 m: usize,
8056 in_f: usize,
8057 out_f: usize,
8058 row_bytes: usize,
8059 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8060 assert!(
8061 in_f % 64 == 0,
8062 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
8063 );
8064 self.qmatvec_dp4a_named(
8065 "qmatvec_nvfp4_dp4a_rp",
8066 &w.slice(0..w.len()),
8067 x,
8068 m,
8069 in_f,
8070 out_f,
8071 row_bytes,
8072 )
8073 }
8074 pub fn qmatvec_nvfp4_fast(
8076 &self,
8077 w: &cudarc::driver::CudaView<'_, u8>,
8078 x: &CudaSlice<f32>,
8079 m: usize,
8080 in_f: usize,
8081 out_f: usize,
8082 row_bytes: usize,
8083 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8084 assert!(
8087 in_f % 64 == 0,
8088 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
8089 );
8090 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a", w, x, m, in_f, out_f, row_bytes)
8091 }
8092 pub fn qmatvec_nvfp4_fast_v2(
8097 &self,
8098 w: &cudarc::driver::CudaView<'_, u8>,
8099 x: &CudaSlice<f32>,
8100 m: usize,
8101 in_f: usize,
8102 out_f: usize,
8103 row_bytes: usize,
8104 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8105 assert!(
8106 in_f % 64 == 0,
8107 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
8108 );
8109 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a_v2", w, x, m, in_f, out_f, row_bytes)
8110 }
8111 #[allow(non_snake_case)] pub fn qmatvec_iq4_XS_fast(
8114 &self,
8115 w: &CudaSlice<u8>,
8116 x: &CudaSlice<f32>,
8117 m: usize,
8118 in_f: usize,
8119 out_f: usize,
8120 row_bytes: usize,
8121 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8122 self.qmatvec_dp4a_named(
8123 "qmatvec_iq4_XS_dp4a",
8124 &w.slice(0..w.len()),
8125 x,
8126 m,
8127 in_f,
8128 out_f,
8129 row_bytes,
8130 )
8131 }
8132
8133 fn qmatvec_dp4a_named(
8135 &self,
8136 name: &str,
8137 w: &cudarc::driver::CudaView<'_, u8>,
8138 x: &CudaSlice<f32>,
8139 m: usize,
8140 in_f: usize,
8141 out_f: usize,
8142 row_bytes: usize,
8143 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8144 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
8145 let f = self.func(name);
8146 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
8148 grid_dim: (out_f as u32, m as u32, 1),
8149 block_dim: (128, 1, 1),
8150 shared_mem_bytes: 0,
8151 };
8152 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8153 let __s_b = self.gpu.stream();
8154 let mut b = __s_b.launch_builder(&f);
8155 b.arg(w)
8156 .arg(&aq)
8157 .arg(&ad)
8158 .arg(&mut y)
8159 .arg(&inf)
8160 .arg(&outf)
8161 .arg(&mi)
8162 .arg(&rb);
8163 unsafe {
8164 b.launch(cfg)?;
8165 }
8166 Ok(y)
8167 }
8168
8169 #[allow(clippy::too_many_arguments)]
8175 pub fn qmatvec_nvfp4_fast_prequant_into(
8176 &self,
8177 w: &CudaSlice<u8>,
8178 aq: &CudaSlice<i8>,
8179 ad: &CudaSlice<f32>,
8180 y: &mut CudaSlice<f32>,
8181 m: usize,
8182 in_f: usize,
8183 out_f: usize,
8184 row_bytes: usize,
8185 ) -> Result<(), Box<dyn std::error::Error>> {
8186 assert!(
8187 in_f % 64 == 0,
8188 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
8189 );
8190 if y.len() < m * out_f {
8191 return Err(format!(
8192 "NVFP4 prequant output {} is shorter than {m}x{out_f}",
8193 y.len()
8194 )
8195 .into());
8196 }
8197 let f = self.func("qmatvec_nvfp4_dp4a");
8198 let cfg = LaunchConfig {
8199 grid_dim: (out_f as u32, m as u32, 1),
8200 block_dim: (128, 1, 1),
8201 shared_mem_bytes: 0,
8202 };
8203 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8204 let __s_b = self.gpu.stream();
8205 let mut b = __s_b.launch_builder(&f);
8206 b.arg(w)
8207 .arg(aq)
8208 .arg(ad)
8209 .arg(y)
8210 .arg(&inf)
8211 .arg(&outf)
8212 .arg(&mi)
8213 .arg(&rb);
8214 unsafe {
8215 b.launch(cfg)?;
8216 }
8217 Ok(())
8218 }
8219
8220 #[allow(clippy::too_many_arguments)]
8223 pub fn matvec_f32_qkv_into(
8224 &self,
8225 wq: &CudaSlice<f32>,
8226 wk: &CudaSlice<f32>,
8227 wv: &CudaSlice<f32>,
8228 wg: &CudaSlice<f32>,
8229 x: &CudaSlice<f32>,
8230 yq: &mut CudaSlice<f32>,
8231 yk: &mut CudaSlice<f32>,
8232 yv: &mut CudaSlice<f32>,
8233 yg: &mut CudaSlice<f32>,
8234 in_f: usize,
8235 out_q: usize,
8236 out_kv: usize,
8237 out_g: usize,
8238 ) -> Result<(), Box<dyn std::error::Error>> {
8239 if in_f % 4 != 0
8240 || wq.len() != out_q * in_f
8241 || wk.len() != out_kv * in_f
8242 || wv.len() != out_kv * in_f
8243 || wg.len() < out_g * in_f
8244 || x.len() < in_f
8245 || yq.len() < out_q
8246 || yk.len() < out_kv
8247 || yv.len() < out_kv
8248 || (out_g > 0 && yg.len() < out_g)
8249 {
8250 return Err(format!(
8251 "fused QKV geometry in={in_f} out_q={out_q} out_kv={out_kv} out_g={out_g} \
8252 wq={} wk={} wv={} wg={}",
8253 wq.len(),
8254 wk.len(),
8255 wv.len(),
8256 wg.len()
8257 )
8258 .into());
8259 }
8260 let f = self.func("matvec_f32_qkv");
8261 let cfg = LaunchConfig {
8262 grid_dim: ((out_q + 2 * out_kv + out_g) as u32, 1, 1),
8263 block_dim: (128, 1, 1),
8264 shared_mem_bytes: 0,
8265 };
8266 let (inf, oq, okv, og) = (in_f as i32, out_q as i32, out_kv as i32, out_g as i32);
8267 let __s_b = self.gpu.stream();
8268 let mut b = __s_b.launch_builder(&f);
8269 b.arg(wq)
8270 .arg(wk)
8271 .arg(wv)
8272 .arg(wg)
8273 .arg(x)
8274 .arg(yq)
8275 .arg(yk)
8276 .arg(yv)
8277 .arg(yg)
8278 .arg(&inf)
8279 .arg(&oq)
8280 .arg(&okv)
8281 .arg(&og);
8282 unsafe {
8283 b.launch(cfg)?;
8284 }
8285 Ok(())
8286 }
8287
8288 #[allow(clippy::too_many_arguments)]
8291 pub fn qmatvec_nvfp4_sel_gu_ep_into(
8292 &self,
8293 gate_bank: &CudaSlice<u8>,
8294 up_bank: &CudaSlice<u8>,
8295 sel: &CudaSlice<i32>,
8296 aq: &CudaSlice<i8>,
8297 ad: &CudaSlice<f32>,
8298 yg: &mut CudaSlice<f32>,
8299 yu: &mut CudaSlice<f32>,
8300 n_sel: usize,
8301 in_f: usize,
8302 out_f: usize,
8303 row_bytes: usize,
8304 expert_stride: usize,
8305 owner: usize,
8306 ) -> Result<(), Box<dyn std::error::Error>> {
8307 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0");
8308 if yg.len() < n_sel * out_f || yu.len() < n_sel * out_f || sel.len() < n_sel {
8309 return Err("NVFP4 gu ep geometry".into());
8310 }
8311 let f = self.func("qmatvec_nvfp4_dp4a_sel_v2_gu_ep");
8312 let cfg = LaunchConfig {
8313 grid_dim: ((2 * out_f) as u32, n_sel as u32, 1),
8314 block_dim: (128, 1, 1),
8315 shared_mem_bytes: 0,
8316 };
8317 let (inf, outf, ns, own) = (in_f as i32, out_f as i32, n_sel as i32, owner as i32);
8318 let (rb, es) = (row_bytes as i64, expert_stride as i64);
8319 let (ars, adrs) = (0i64, 0i64);
8320 let __s_b = self.gpu.stream();
8321 let mut b = __s_b.launch_builder(&f);
8322 b.arg(gate_bank)
8323 .arg(up_bank)
8324 .arg(sel)
8325 .arg(aq)
8326 .arg(ad)
8327 .arg(yg)
8328 .arg(yu)
8329 .arg(&inf)
8330 .arg(&outf)
8331 .arg(&ns)
8332 .arg(&rb)
8333 .arg(&es)
8334 .arg(&ars)
8335 .arg(&adrs)
8336 .arg(&own);
8337 unsafe {
8338 b.launch(cfg)?;
8339 }
8340 Ok(())
8341 }
8342
8343 #[allow(clippy::too_many_arguments)]
8345 pub fn silu_mul_scaled_q8_1_sel_ep_into(
8346 &self,
8347 gate: &CudaSlice<f32>,
8348 up: &CudaSlice<f32>,
8349 gmac: &CudaSlice<f32>,
8350 umac: &CudaSlice<f32>,
8351 sel: &CudaSlice<i32>,
8352 limit: Option<f32>,
8353 out_q: &mut CudaSlice<i8>,
8354 out_d: &mut CudaSlice<f32>,
8355 n_per: usize,
8356 n_sel: usize,
8357 owner: usize,
8358 ) -> Result<(), Box<dyn std::error::Error>> {
8359 if n_per % 32 != 0 || out_q.len() < n_sel * n_per || out_d.len() < n_sel * n_per / 32 {
8360 return Err("NVFP4 silu ep geometry".into());
8361 }
8362 let f = self.func("silu_mul_scaled_q8_1_sel_ep");
8363 let warps = n_sel * n_per / 32;
8364 let cfg = LaunchConfig {
8365 grid_dim: ((warps as u32).div_ceil(4), 1, 1),
8366 block_dim: (128, 1, 1),
8367 shared_mem_bytes: 0,
8368 };
8369 let (np, ns, own) = (n_per as i32, n_sel as i32, owner as i32);
8370 let (lim, has) = match limit {
8371 Some(l) => (l, 1i32),
8372 None => (0.0f32, 0i32),
8373 };
8374 let __s_b = self.gpu.stream();
8375 let mut b = __s_b.launch_builder(&f);
8376 b.arg(gate)
8377 .arg(up)
8378 .arg(gmac)
8379 .arg(umac)
8380 .arg(sel)
8381 .arg(&lim)
8382 .arg(&has)
8383 .arg(out_q)
8384 .arg(out_d)
8385 .arg(&np)
8386 .arg(&ns)
8387 .arg(&own);
8388 unsafe {
8389 b.launch(cfg)?;
8390 }
8391 Ok(())
8392 }
8393
8394 #[allow(clippy::too_many_arguments)]
8396 pub fn qmatvec_nvfp4_sel_down8_ep_into(
8397 &self,
8398 bank: &CudaSlice<u8>,
8399 sel: &CudaSlice<i32>,
8400 aq: &CudaSlice<i8>,
8401 ad: &CudaSlice<f32>,
8402 route_w: &CudaSlice<f32>,
8403 md: &CudaSlice<f32>,
8404 dst: &mut CudaSlice<f32>,
8405 n_sel: usize,
8406 in_f: usize,
8407 out_f: usize,
8408 row_bytes: usize,
8409 expert_stride: usize,
8410 act_row_stride: usize,
8411 ad_row_stride: usize,
8412 owner: usize,
8413 ) -> Result<(), Box<dyn std::error::Error>> {
8414 if in_f % 64 != 0 || n_sel == 0 || n_sel > 8 || (in_f >> 5) > 64 || dst.len() < out_f {
8415 return Err("NVFP4 down8 ep geometry".into());
8416 }
8417 let f = self.func("qmatvec_nvfp4_dp4a_sel_v2_down8_ep");
8418 let cfg = LaunchConfig {
8419 grid_dim: (out_f as u32, 1, 1),
8420 block_dim: (32, n_sel as u32, 1),
8421 shared_mem_bytes: 0,
8422 };
8423 let (inf, outf, ns, own) = (in_f as i32, out_f as i32, n_sel as i32, owner as i32);
8424 let (rb, es) = (row_bytes as i64, expert_stride as i64);
8425 let (ars, adrs) = (act_row_stride as i64, ad_row_stride as i64);
8426 let __s_b = self.gpu.stream();
8427 let mut b = __s_b.launch_builder(&f);
8428 b.arg(bank)
8429 .arg(sel)
8430 .arg(aq)
8431 .arg(ad)
8432 .arg(route_w)
8433 .arg(md)
8434 .arg(dst)
8435 .arg(&inf)
8436 .arg(&outf)
8437 .arg(&ns)
8438 .arg(&rb)
8439 .arg(&es)
8440 .arg(&ars)
8441 .arg(&adrs)
8442 .arg(&own);
8443 unsafe {
8444 b.launch(cfg)?;
8445 }
8446 Ok(())
8447 }
8448
8449 #[allow(clippy::too_many_arguments)]
8455 pub fn qmatvec_nvfp4_sel_into(
8456 &self,
8457 bank: &CudaSlice<u8>,
8458 sel: &CudaSlice<i32>,
8459 aq: &CudaSlice<i8>,
8460 ad: &CudaSlice<f32>,
8461 y: &mut CudaSlice<f32>,
8462 n_sel: usize,
8463 in_f: usize,
8464 out_f: usize,
8465 row_bytes: usize,
8466 expert_stride: usize,
8467 act_row_stride: usize,
8468 ad_row_stride: usize,
8469 ) -> Result<(), Box<dyn std::error::Error>> {
8470 assert!(
8471 in_f % 64 == 0,
8472 "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
8473 );
8474 if y.len() < n_sel * out_f || sel.len() < n_sel {
8475 return Err(format!(
8476 "NVFP4 sel output {} / sel {} shorter than {n_sel}x{out_f}",
8477 y.len(),
8478 sel.len()
8479 )
8480 .into());
8481 }
8482 static MR: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
8489 let mode = *MR.get_or_init(|| {
8490 if std::env::var("MEMRA_SEL_STREAM").as_deref() == Ok("1") {
8491 2
8492 } else if std::env::var("MEMRA_SEL_MR").as_deref() == Ok("1") {
8493 1
8494 } else {
8495 0
8496 }
8497 });
8498 let mode = if mode == 2 && in_f > 4096 { 0 } else { mode };
8499 let f = match mode {
8500 2 => self.func("qmatvec_nvfp4_dp4a_sel_stream"),
8501 1 => self.func("qmatvec_nvfp4_dp4a_sel_mr4"),
8502 _ => self.func("qmatvec_nvfp4_dp4a_sel"),
8503 };
8504 let nsb = in_f >> 5;
8509 let fit_block: u32 = if mode == 0 && nsb <= 32 {
8510 32
8511 } else if mode == 1 {
8512 512
8513 } else {
8514 128
8515 };
8516 let cfg = LaunchConfig {
8517 grid_dim: (
8518 match mode {
8519 2 => (out_f as u32).div_ceil(16),
8520 1 => (out_f as u32).div_ceil(4),
8521 _ => out_f as u32,
8522 },
8523 n_sel as u32,
8524 1,
8525 ),
8526 block_dim: (fit_block, 1, 1),
8527 shared_mem_bytes: 0,
8528 };
8529 let (inf, outf, ns) = (in_f as i32, out_f as i32, n_sel as i32);
8530 let (rb, es, ars, adrs) = (
8531 row_bytes as i64,
8532 expert_stride as i64,
8533 act_row_stride as i64,
8534 ad_row_stride as i64,
8535 );
8536 let __s_b = self.gpu.stream();
8537 let mut b = __s_b.launch_builder(&f);
8538 b.arg(bank)
8539 .arg(sel)
8540 .arg(aq)
8541 .arg(ad)
8542 .arg(y)
8543 .arg(&inf)
8544 .arg(&outf)
8545 .arg(&ns)
8546 .arg(&rb)
8547 .arg(&es)
8548 .arg(&ars)
8549 .arg(&adrs);
8550 unsafe {
8551 b.launch(cfg)?;
8552 }
8553 Ok(())
8554 }
8555
8556 #[allow(clippy::too_many_arguments)]
8561 pub fn silu_mul_scaled_q8_1_sel_into(
8562 &self,
8563 gate: &CudaSlice<f32>,
8564 up: &CudaSlice<f32>,
8565 gmac: &CudaSlice<f32>,
8566 umac: &CudaSlice<f32>,
8567 sel: &CudaSlice<i32>,
8568 limit: Option<f32>,
8569 out_q: &mut CudaSlice<i8>,
8570 out_d: &mut CudaSlice<f32>,
8571 n_per: usize,
8572 n_sel: usize,
8573 ) -> Result<(), Box<dyn std::error::Error>> {
8574 let n = n_per * n_sel;
8575 if n_per % 32 != 0 || out_q.len() < n || out_d.len() < n / 32 {
8576 return Err(format!(
8577 "silu sel geometry n_per={n_per} n_sel={n_sel} q={} d={}",
8578 out_q.len(),
8579 out_d.len()
8580 )
8581 .into());
8582 }
8583 if let Some(limit) = limit {
8584 if limit <= 1e-6 {
8585 return Err(format!(
8586 "silu sel clamp limit {limit} is at or below the 1e-6 eps gate"
8587 )
8588 .into());
8589 }
8590 let f = self.func("silu_mul_scaled_q8_1_sel_clamp");
8591 let cfg = LaunchConfig::for_num_elems(n as u32);
8592 let (np, ns) = (n_per as i32, n_sel as i32);
8593 let __s_b = self.gpu.stream();
8594 let mut b = __s_b.launch_builder(&f);
8595 b.arg(gate)
8596 .arg(up)
8597 .arg(gmac)
8598 .arg(umac)
8599 .arg(sel)
8600 .arg(&limit)
8601 .arg(out_q)
8602 .arg(out_d)
8603 .arg(&np)
8604 .arg(&ns);
8605 unsafe {
8606 b.launch(cfg)?;
8607 }
8608 return Ok(());
8609 }
8610 let f = self.func("silu_mul_scaled_q8_1_sel");
8611 let cfg = LaunchConfig::for_num_elems(n as u32);
8612 let (np, ns) = (n_per as i32, n_sel as i32);
8613 let __s_b = self.gpu.stream();
8614 let mut b = __s_b.launch_builder(&f);
8615 b.arg(gate)
8616 .arg(up)
8617 .arg(gmac)
8618 .arg(umac)
8619 .arg(sel)
8620 .arg(out_q)
8621 .arg(out_d)
8622 .arg(&np)
8623 .arg(&ns);
8624 unsafe {
8625 b.launch(cfg)?;
8626 }
8627 Ok(())
8628 }
8629
8630 pub fn htod(&self, v: &[f32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8631 Ok(self.gpu.stream().clone_htod(v)?)
8632 }
8633 pub fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
8634 Ok(self.gpu.stream().clone_htod(v)?)
8635 }
8636 pub fn htod_i8(&self, v: &[i8]) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
8638 Ok(self.gpu.stream().clone_htod(v)?)
8639 }
8640 pub fn htod_u64(&self, v: &[u64]) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
8641 Ok(self.gpu.stream().clone_htod(v)?)
8642 }
8643 pub fn dtoh_view(
8645 &self,
8646 d: &cudarc::driver::CudaView<f32>,
8647 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
8648 let v = self.gpu.stream().clone_dtoh(d)?;
8649 self.gpu.stream().synchronize()?;
8650 Ok(v)
8651 }
8652 pub fn dtoh(&self, d: &CudaSlice<f32>) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
8653 let v = self.gpu.stream().clone_dtoh(d)?;
8654 self.gpu.stream().synchronize()?;
8655 Ok(v)
8656 }
8657 pub fn dtoh_pair(
8661 &self,
8662 a: &CudaSlice<f32>,
8663 b: &CudaSlice<f32>,
8664 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
8665 let av = self.gpu.stream().clone_dtoh(a)?;
8666 let bv = self.gpu.stream().clone_dtoh(b)?;
8667 self.gpu.stream().synchronize()?;
8668 Ok((av, bv))
8669 }
8670 pub fn dtoh_pair_views(
8673 &self,
8674 a: &cudarc::driver::CudaView<f32>,
8675 b: &cudarc::driver::CudaView<f32>,
8676 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
8677 let av = self.gpu.stream().clone_dtoh(a)?;
8678 let bv = self.gpu.stream().clone_dtoh(b)?;
8679 self.gpu.stream().synchronize()?;
8680 Ok((av, bv))
8681 }
8682 pub fn dtoh_i32(&self, d: &CudaSlice<i32>) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
8684 let v = self.gpu.stream().clone_dtoh(d)?;
8685 self.gpu.stream().synchronize()?;
8686 Ok(v)
8687 }
8688 pub fn dtoh_u8(&self, d: &CudaSlice<u8>) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
8690 let v = self.gpu.stream().clone_dtoh(d)?;
8691 self.gpu.stream().synchronize()?;
8692 Ok(v)
8693 }
8694 pub fn dtoh_u8_view(
8695 &self,
8696 d: &cudarc::driver::CudaView<u8>,
8697 ) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
8698 let v = self.gpu.stream().clone_dtoh(d)?;
8699 self.gpu.stream().synchronize()?;
8700 Ok(v)
8701 }
8702 pub fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8703 let s = self.gpu.stream().alloc_zeros::<f32>(n)?;
8704 self.keep_if_capturing(&s);
8705 Ok(s)
8706 }
8707
8708 pub fn prob_of_token_device(
8717 &self,
8718 logits: &CudaSlice<f32>,
8719 tok: &CudaSlice<u32>,
8720 n_vocab: usize,
8721 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8722 let nb = ARGMAX_NB;
8723 let mut part = self.alloc_uninit::<f32>(nb)?;
8724 let mut p = self.alloc_uninit::<f32>(1)?;
8725 let f1 = self.func("prob_of_token_partial_f32");
8726 let cfg1 = LaunchConfig {
8727 grid_dim: (nb as u32, 1, 1),
8728 block_dim: (256, 1, 1),
8729 shared_mem_bytes: 0,
8730 };
8731 let nv = n_vocab as i32;
8732 let __s_b1 = self.gpu.stream();
8733 let mut b1 = __s_b1.launch_builder(&f1);
8734 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
8735 unsafe {
8736 b1.launch(cfg1)?;
8737 }
8738 let f2 = self.func("prob_of_token_final_f32");
8739 let cfg2 = LaunchConfig {
8740 grid_dim: (1, 1, 1),
8741 block_dim: (256, 1, 1),
8742 shared_mem_bytes: 0,
8743 };
8744 let nbi = nb as i32;
8745 let __s_b2 = self.gpu.stream();
8746 let mut b2 = __s_b2.launch_builder(&f2);
8747 b2.arg(&part).arg(&mut p).arg(&nbi);
8748 unsafe {
8749 b2.launch(cfg2)?;
8750 }
8751 Ok(p)
8752 }
8753
8754 pub fn prob_of_token_device_col(
8761 &self,
8762 logits: &CudaSlice<f32>,
8763 tok_all: &CudaSlice<u32>,
8764 tok_idx: usize,
8765 p_out: &mut CudaSlice<f32>,
8766 p_idx: usize,
8767 n_vocab: usize,
8768 ) -> Result<(), Box<dyn std::error::Error>> {
8769 let tok_v = tok_all.slice(tok_idx..tok_idx + 1);
8770 let mut p_v = p_out.slice_mut(p_idx..p_idx + 1);
8771 let nb = ARGMAX_NB;
8772 let mut part = self.alloc_uninit::<f32>(nb)?;
8773 let f1 = self.func("prob_of_token_partial_f32");
8774 let cfg1 = LaunchConfig {
8775 grid_dim: (nb as u32, 1, 1),
8776 block_dim: (256, 1, 1),
8777 shared_mem_bytes: 0,
8778 };
8779 let nv = n_vocab as i32;
8780 let __s_b1 = self.gpu.stream();
8781 let mut b1 = __s_b1.launch_builder(&f1);
8782 b1.arg(logits).arg(&tok_v).arg(&mut part).arg(&nv);
8783 unsafe {
8784 b1.launch(cfg1)?;
8785 }
8786 let f2 = self.func("prob_of_token_final_f32");
8787 let cfg2 = LaunchConfig {
8788 grid_dim: (1, 1, 1),
8789 block_dim: (256, 1, 1),
8790 shared_mem_bytes: 0,
8791 };
8792 let nbi = nb as i32;
8793 let __s_b2 = self.gpu.stream();
8794 let mut b2 = __s_b2.launch_builder(&f2);
8795 b2.arg(&part).arg(&mut p_v).arg(&nbi);
8796 unsafe {
8797 b2.launch(cfg2)?;
8798 }
8799 Ok(())
8800 }
8801
8802 pub fn prob_of_token_device_into(
8803 &self,
8804 logits: &CudaSlice<f32>,
8805 tok: &CudaSlice<u32>,
8806 p_out: &mut CudaSlice<f32>,
8807 n_vocab: usize,
8808 ) -> Result<(), Box<dyn std::error::Error>> {
8809 let nb = ARGMAX_NB;
8810 let mut part = self.alloc_uninit::<f32>(nb)?;
8811 let f1 = self.func("prob_of_token_partial_f32");
8812 let cfg1 = LaunchConfig {
8813 grid_dim: (nb as u32, 1, 1),
8814 block_dim: (256, 1, 1),
8815 shared_mem_bytes: 0,
8816 };
8817 let nv = n_vocab as i32;
8818 let __s_b1 = self.gpu.stream();
8819 let mut b1 = __s_b1.launch_builder(&f1);
8820 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
8821 unsafe {
8822 b1.launch(cfg1)?;
8823 }
8824 let f2 = self.func("prob_of_token_final_f32");
8825 let cfg2 = LaunchConfig {
8826 grid_dim: (1, 1, 1),
8827 block_dim: (256, 1, 1),
8828 shared_mem_bytes: 0,
8829 };
8830 let nbi = nb as i32;
8831 let __s_b2 = self.gpu.stream();
8832 let mut b2 = __s_b2.launch_builder(&f2);
8833 b2.arg(&part).arg(p_out).arg(&nbi);
8834 unsafe {
8835 b2.launch(cfg2)?;
8836 }
8837 Ok(())
8838 }
8839
8840 pub fn u32_hist_append(
8843 &self,
8844 tok: &CudaSlice<u32>,
8845 hist: &mut CudaSlice<u32>,
8846 idx: &mut CudaSlice<i32>,
8847 ) -> Result<(), Box<dyn std::error::Error>> {
8848 let f = self.func("u32_hist_append");
8849 let cfg = LaunchConfig {
8850 grid_dim: (1, 1, 1),
8851 block_dim: (32, 1, 1),
8852 shared_mem_bytes: 0,
8853 };
8854 let __s_b = self.gpu.stream();
8855 let mut b = __s_b.launch_builder(&f);
8856 b.arg(tok).arg(&mut *hist).arg(&mut *idx);
8857 unsafe {
8858 b.launch(cfg)?;
8859 }
8860 Ok(())
8861 }
8862
8863 pub fn argmax_token_device(
8864 &self,
8865 logits: &CudaSlice<f32>,
8866 n_vocab: usize,
8867 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
8868 let mut tok = unsafe { self.gpu.stream().alloc::<u32>(1)? };
8869 self.argmax_token_device_into(logits, &mut tok, n_vocab)?;
8870 Ok(tok)
8871 }
8872 pub fn argmax_token_device_into(
8879 &self,
8880 logits: &CudaSlice<f32>,
8881 tok: &mut CudaSlice<u32>,
8882 n_vocab: usize,
8883 ) -> Result<(), Box<dyn std::error::Error>> {
8884 let nb = ARGMAX_NB;
8885 let f1 = self.func("argmax_partial_f32");
8886 let f2 = self.func("argmax_final_f32");
8887 let mut guard = self.argmax_partials.lock().unwrap();
8888 if guard.is_none() {
8889 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
8892 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
8893 *guard = Some((pv, pi));
8894 }
8895 let (part_v, part_i) = guard.as_mut().unwrap();
8896 let nv = n_vocab as i32;
8897 let nbi = nb as i32;
8898 let cfg1 = LaunchConfig {
8900 grid_dim: (nb as u32, 1, 1),
8901 block_dim: (256, 1, 1),
8902 shared_mem_bytes: 0,
8903 };
8904 let __s_b1 = self.gpu.stream();
8905 let mut b1 = __s_b1.launch_builder(&f1);
8906 b1.arg(logits).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
8907 unsafe {
8908 b1.launch(cfg1)?;
8909 }
8910 let cfg2 = LaunchConfig {
8912 grid_dim: (1, 1, 1),
8913 block_dim: (256, 1, 1),
8914 shared_mem_bytes: 0,
8915 };
8916 let __s_b2 = self.gpu.stream();
8917 let mut b2 = __s_b2.launch_builder(&f2);
8918 b2.arg(&*part_v).arg(&*part_i).arg(tok).arg(&nbi);
8919 unsafe {
8920 b2.launch(cfg2)?;
8921 }
8922 Ok(())
8923 }
8924 pub fn argmax_token_device_col(
8930 &self,
8931 logits: &CudaSlice<f32>,
8932 col: usize,
8933 n_vocab: usize,
8934 toks: &mut CudaSlice<u32>,
8935 out_idx: usize,
8936 ) -> Result<(), Box<dyn std::error::Error>> {
8937 let nb = ARGMAX_NB;
8938 let f1 = self.func("argmax_partial_f32");
8939 let f2 = self.func("argmax_final_f32");
8940 let mut guard = self.argmax_partials.lock().unwrap();
8941 if guard.is_none() {
8942 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
8943 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
8944 *guard = Some((pv, pi));
8945 }
8946 let (part_v, part_i) = guard.as_mut().unwrap();
8947 let col_view = logits.slice(col * n_vocab..(col + 1) * n_vocab);
8948 let nv = n_vocab as i32;
8949 let nbi = nb as i32;
8950 let cfg1 = LaunchConfig {
8951 grid_dim: (nb as u32, 1, 1),
8952 block_dim: (256, 1, 1),
8953 shared_mem_bytes: 0,
8954 };
8955 let __s_b1 = self.gpu.stream();
8956 let mut b1 = __s_b1.launch_builder(&f1);
8957 b1.arg(&col_view)
8958 .arg(&mut *part_v)
8959 .arg(&mut *part_i)
8960 .arg(&nv);
8961 unsafe {
8962 b1.launch(cfg1)?;
8963 }
8964 let mut tok_view = toks.slice_mut(out_idx..out_idx + 1);
8965 let cfg2 = LaunchConfig {
8966 grid_dim: (1, 1, 1),
8967 block_dim: (256, 1, 1),
8968 shared_mem_bytes: 0,
8969 };
8970 let __s_b2 = self.gpu.stream();
8971 let mut b2 = __s_b2.launch_builder(&f2);
8972 b2.arg(&*part_v).arg(&*part_i).arg(&mut tok_view).arg(&nbi);
8973 unsafe {
8974 b2.launch(cfg2)?;
8975 }
8976 Ok(())
8977 }
8978 pub fn htod_u32_v(&self, v: &[u32]) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
8980 Ok(self.gpu.stream().clone_htod(v)?)
8981 }
8982 pub fn dtoh_u64(&self, d: &CudaSlice<u64>) -> Result<Vec<u64>, Box<dyn std::error::Error>> {
8983 let v = self.gpu.stream().clone_dtoh(d)?;
8984 self.gpu.stream().synchronize()?;
8985 Ok(v)
8986 }
8987
8988 pub fn dtoh_u32(&self, d: &CudaSlice<u32>) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
8989 let v = self.gpu.stream().clone_dtoh(d)?;
8990 self.gpu.stream().synchronize()?;
8991 Ok(v)
8992 }
8993 pub fn htod_u32_into(
8997 &self,
8998 dst: &mut CudaSlice<u32>,
8999 src: &[u32],
9000 ) -> Result<(), Box<dyn std::error::Error>> {
9001 let mut view = dst.slice_mut(0..src.len());
9002 self.gpu.stream().memcpy_htod(src, &mut view)?;
9003 Ok(())
9004 }
9005
9006 pub fn htod_i32_into(
9009 &self,
9010 dst: &mut CudaSlice<i32>,
9011 src: &[i32],
9012 ) -> Result<(), Box<dyn std::error::Error>> {
9013 let mut view = dst.slice_mut(0..src.len());
9014 self.gpu.stream().memcpy_htod(src, &mut view)?;
9015 Ok(())
9016 }
9017
9018 pub fn alloc_u32_zeroed(&self, n: usize) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
9019 let s = self.gpu.stream().alloc_zeros::<u32>(n)?;
9020 self.keep_if_capturing(&s);
9021 Ok(s)
9022 }
9023 pub fn embed_gather_device_into(
9026 &self,
9027 embd: &CudaSlice<u8>,
9028 token_d: &CudaSlice<u32>,
9029 x_out: &mut CudaSlice<f32>,
9030 n_embd: usize,
9031 qtype: i32,
9032 row_bytes: usize,
9033 ) -> Result<(), Box<dyn std::error::Error>> {
9034 let f = self.func("embed_gather_u32");
9035 let cfg = LaunchConfig {
9036 grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
9037 block_dim: (256, 1, 1),
9038 shared_mem_bytes: 0,
9039 };
9040 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
9041 let __s_b = self.gpu.stream();
9042 let mut b = __s_b.launch_builder(&f);
9043 b.arg(embd)
9044 .arg(token_d)
9045 .arg(x_out)
9046 .arg(&ne)
9047 .arg(&qt)
9048 .arg(&rb);
9049 unsafe {
9050 b.launch(cfg)?;
9051 }
9052 Ok(())
9053 }
9054 pub fn dtoh_i32_one(&self, d: &CudaSlice<i32>) -> Result<i32, Box<dyn std::error::Error>> {
9056 let v = self.gpu.stream().clone_dtoh(d)?;
9057 self.gpu.stream().synchronize()?;
9058 Ok(v[0])
9059 }
9060 pub fn i32_set_k(
9067 &self,
9068 dst: &mut CudaSlice<i32>,
9069 v: i32,
9070 ) -> Result<(), Box<dyn std::error::Error>> {
9071 let f = self.func("i32_set_k");
9072 let cfg = LaunchConfig {
9073 grid_dim: (1, 1, 1),
9074 block_dim: (1, 1, 1),
9075 shared_mem_bytes: 0,
9076 };
9077 let idx = 0i32;
9078 let __s_b = self.gpu.stream();
9079 let mut b = __s_b.launch_builder(&f);
9080 b.arg(dst).arg(&v).arg(&idx);
9081 unsafe {
9082 b.launch(cfg)?;
9083 }
9084 Ok(())
9085 }
9086
9087 pub fn set_i32_one(
9088 &self,
9089 d: &mut CudaSlice<i32>,
9090 v: i32,
9091 ) -> Result<(), Box<dyn std::error::Error>> {
9092 self.gpu.stream().memcpy_htod(&[v], d)?;
9093 Ok(())
9094 }
9095 pub fn set_u32_one(
9098 &self,
9099 d: &mut CudaSlice<u32>,
9100 v: u32,
9101 ) -> Result<(), Box<dyn std::error::Error>> {
9102 self.gpu.stream().memcpy_htod(&[v], d)?;
9103 Ok(())
9104 }
9105 pub fn dtoh_u32_one(&self, d: &CudaSlice<u32>) -> Result<u32, Box<dyn std::error::Error>> {
9107 let v = self.gpu.stream().clone_dtoh(d)?;
9108 self.gpu.stream().synchronize()?;
9109 Ok(v[0])
9110 }
9111 pub fn upload_u8(&self, bytes: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9113 Ok(self.gpu.stream().clone_htod(bytes)?)
9114 }
9115 pub fn embed_gather_device(
9119 &self,
9120 embd: &CudaSlice<u8>,
9121 token_d: &CudaSlice<u32>,
9122 n_embd: usize,
9123 qtype: i32,
9124 row_bytes: usize,
9125 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9126 let f = self.func("embed_gather_u32");
9127 let mut x = self.alloc_uninit::<f32>(n_embd)?;
9128 let cfg = LaunchConfig {
9129 grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
9130 block_dim: (256, 1, 1),
9131 shared_mem_bytes: 0,
9132 };
9133 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
9134 let __s_b = self.gpu.stream();
9135 let mut b = __s_b.launch_builder(&f);
9136 b.arg(embd)
9137 .arg(token_d)
9138 .arg(&mut x)
9139 .arg(&ne)
9140 .arg(&qt)
9141 .arg(&rb);
9142 unsafe {
9143 b.launch(cfg)?;
9144 }
9145 Ok(x)
9146 }
9147
9148 pub fn embed_gather_device_t(
9152 &self,
9153 embd: &CudaSlice<u8>,
9154 tokens: &[u32],
9155 n_embd: usize,
9156 qtype: i32,
9157 row_bytes: usize,
9158 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9159 let t = tokens.len();
9160 let tok_d = self.gpu.stream().clone_htod(tokens)?;
9161 let f = self.func("embed_gather_u32_t");
9162 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
9163 let cfg = LaunchConfig {
9164 grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
9165 block_dim: (256, 1, 1),
9166 shared_mem_bytes: 0,
9167 };
9168 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
9169 let __s_b = self.gpu.stream();
9170 let mut b = __s_b.launch_builder(&f);
9171 b.arg(embd)
9172 .arg(&tok_d)
9173 .arg(&mut x)
9174 .arg(&ne)
9175 .arg(&qt)
9176 .arg(&rb)
9177 .arg(&ti);
9178 unsafe {
9179 b.launch(cfg)?;
9180 }
9181 Ok(x)
9182 }
9183
9184 pub fn embed_gather_device_tv(
9189 &self,
9190 embd: &CudaSlice<u8>,
9191 tok_v: &cudarc::driver::CudaView<u32>,
9192 t: usize,
9193 n_embd: usize,
9194 qtype: i32,
9195 row_bytes: usize,
9196 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9197 let f = self.func("embed_gather_u32_t");
9198 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
9199 let cfg = LaunchConfig {
9200 grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
9201 block_dim: (256, 1, 1),
9202 shared_mem_bytes: 0,
9203 };
9204 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
9205 let __s_b = self.gpu.stream();
9206 let mut b = __s_b.launch_builder(&f);
9207 b.arg(embd)
9208 .arg(tok_v)
9209 .arg(&mut x)
9210 .arg(&ne)
9211 .arg(&qt)
9212 .arg(&rb)
9213 .arg(&ti);
9214 unsafe {
9215 b.launch(cfg)?;
9216 }
9217 Ok(x)
9218 }
9219
9220 pub fn embed_gather_device_td(
9221 &self,
9222 embd: &CudaSlice<u8>,
9223 tok_d: &CudaSlice<u32>,
9224 t: usize,
9225 n_embd: usize,
9226 qtype: i32,
9227 row_bytes: usize,
9228 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9229 let f = self.func("embed_gather_u32_t");
9230 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
9231 let cfg = LaunchConfig {
9232 grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
9233 block_dim: (256, 1, 1),
9234 shared_mem_bytes: 0,
9235 };
9236 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
9237 let __s_b = self.gpu.stream();
9238 let mut b = __s_b.launch_builder(&f);
9239 b.arg(embd)
9240 .arg(tok_d)
9241 .arg(&mut x)
9242 .arg(&ne)
9243 .arg(&qt)
9244 .arg(&rb)
9245 .arg(&ti);
9246 unsafe {
9247 b.launch(cfg)?;
9248 }
9249 Ok(x)
9250 }
9251
9252 #[inline]
9258 fn keep_if_capturing<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, s: &CudaSlice<T>) {
9260 if self
9261 .capture_keep_on
9262 .load(std::sync::atomic::Ordering::Relaxed)
9263 {
9264 self.capture_keep.lock().unwrap().push(Box::new(s.clone()));
9265 }
9266 }
9267
9268 fn alloc_uninit<T: cudarc::driver::DeviceRepr + Send + 'static>(
9269 &self,
9270 n: usize,
9271 ) -> Result<CudaSlice<T>, Box<dyn std::error::Error>> {
9272 let mut s = unsafe { self.gpu.stream().alloc::<T>(n)? };
9273 {
9277 static Z: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9278 if *Z.get_or_init(|| std::env::var("MEMRA_DEBUG_ZERO_ALLOCS").as_deref() == Ok("1")) {
9279 use cudarc::driver::DevicePtrMut;
9281 let n_bytes = s.len() * std::mem::size_of::<T>();
9282 let stream = self.gpu.stream();
9283 let (p_, _g) = s.device_ptr_mut(&stream);
9284 unsafe {
9285 cudarc::driver::sys::cuMemsetD8Async(p_, 0, n_bytes, stream.cu_stream())
9286 .result()?;
9287 }
9288 }
9289 }
9290 self.keep_if_capturing(&s);
9291 Ok(s)
9292 }
9293
9294 pub fn uninit_q8_pair(
9299 &self,
9300 n: usize,
9301 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9302 Ok((
9303 self.alloc_uninit::<i8>(n)?,
9304 self.alloc_uninit::<f32>(n / 32)?,
9305 ))
9306 }
9307
9308 pub fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9309 self.alloc_uninit::<f32>(n)
9310 }
9311
9312 pub fn alloc_i8_uninit(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
9314 self.alloc_uninit::<i8>(n)
9315 }
9316
9317 #[allow(clippy::too_many_arguments)]
9321 pub fn rms_norm3(
9322 &self,
9323 x: &CudaSlice<f32>,
9324 w0: &CudaSlice<f32>,
9325 w1: &CudaSlice<f32>,
9326 w2: &CudaSlice<f32>,
9327 d0: &mut CudaSlice<f32>,
9328 d1: &mut CudaSlice<f32>,
9329 d2: &mut CudaSlice<f32>,
9330 ncols: usize,
9331 nrows: usize,
9332 eps: f32,
9333 ) -> Result<(), Box<dyn std::error::Error>> {
9334 let f = self.func("rms_norm3_f32");
9335 let cfg = LaunchConfig {
9336 grid_dim: (nrows as u32, 1, 1),
9337 block_dim: (rms_block(), 1, 1),
9338 shared_mem_bytes: 0,
9339 };
9340 let (nc, e) = (ncols as i32, eps);
9341 let __s_b = self.gpu.stream();
9342 let mut b = __s_b.launch_builder(&f);
9343 b.arg(x)
9344 .arg(w0)
9345 .arg(w1)
9346 .arg(w2)
9347 .arg(d0)
9348 .arg(d1)
9349 .arg(d2)
9350 .arg(&nc)
9351 .arg(&e);
9352 unsafe {
9353 b.launch(cfg)?;
9354 }
9355 Ok(())
9356 }
9357
9358 #[allow(clippy::too_many_arguments)]
9360 pub fn qkvnorm_w_on_prefill(rows: usize, ncols: usize) -> bool {
9363 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9364 *WARP_ON.get_or_init(|| {
9365 std::env::var("MEMRA_QKVNORM_W")
9366 .map(|v| v != "0")
9367 .unwrap_or(true)
9368 }) && ncols % 4 == 0
9369 && rows >= 64
9370 }
9371
9372 #[allow(clippy::too_many_arguments)]
9375 pub fn rms_norm_qkv_w4b(
9376 &self,
9377 q: &CudaSlice<f32>,
9378 k: &CudaSlice<f32>,
9379 v: &CudaSlice<f32>,
9380 wq: &CudaSlice<f32>,
9381 wk: &CudaSlice<f32>,
9382 wv: &CudaSlice<f32>,
9383 dq: &mut CudaSlice<f32>,
9384 dk: &mut CudaSlice<f32>,
9385 dv: &mut CudaSlice<f32>,
9386 dvb: &mut CudaSlice<u8>,
9387 ncols: usize,
9388 rq: usize,
9389 rk: usize,
9390 eps: f32,
9391 vf16: bool,
9392 ) -> Result<(), Box<dyn std::error::Error>> {
9393 assert!(ncols % 4 == 0 && rq + 2 * rk >= 64);
9394 let f = self.func("rms_norm_qkv_w4b_f32");
9395 let rows = (rq + 2 * rk) as u32;
9396 let cfg = LaunchConfig {
9397 grid_dim: (rows.div_ceil(8), 1, 1),
9398 block_dim: (256, 1, 1),
9399 shared_mem_bytes: 0,
9400 };
9401 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
9402 let vf = vf16 as i32;
9403 let __s_b = self.gpu.stream();
9404 let mut b = __s_b.launch_builder(&f);
9405 b.arg(q)
9406 .arg(k)
9407 .arg(v)
9408 .arg(wq)
9409 .arg(wk)
9410 .arg(wv)
9411 .arg(dq)
9412 .arg(dk)
9413 .arg(dv)
9414 .arg(&mut *dvb)
9415 .arg(&nc)
9416 .arg(&rqi)
9417 .arg(&rki)
9418 .arg(&rvi)
9419 .arg(&e)
9420 .arg(&vf);
9421 unsafe {
9422 b.launch(cfg)?;
9423 }
9424 Ok(())
9425 }
9426
9427 pub fn rms_norm_qkv(
9428 &self,
9429 q: &CudaSlice<f32>,
9430 k: &CudaSlice<f32>,
9431 v: &CudaSlice<f32>,
9432 wq: &CudaSlice<f32>,
9433 wk: &CudaSlice<f32>,
9434 wv: &CudaSlice<f32>,
9435 dq: &mut CudaSlice<f32>,
9436 dk: &mut CudaSlice<f32>,
9437 dv: &mut CudaSlice<f32>,
9438 ncols: usize,
9439 rq: usize,
9440 rk: usize,
9441 eps: f32,
9442 ) -> Result<(), Box<dyn std::error::Error>> {
9443 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9447 let warp_on = *WARP_ON.get_or_init(|| {
9448 std::env::var("MEMRA_QKVNORM_W")
9449 .map(|v| v != "0")
9450 .unwrap_or(true)
9451 });
9452 if warp_on && ncols % 4 == 0 && rq + 2 * rk >= 64 {
9455 let f = self.func("rms_norm_qkv_w4_f32");
9456 let rows = (rq + 2 * rk) as u32;
9457 let cfg = LaunchConfig {
9458 grid_dim: (rows.div_ceil(8), 1, 1),
9459 block_dim: (256, 1, 1),
9460 shared_mem_bytes: 0,
9461 };
9462 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
9463 let __s_b = self.gpu.stream();
9464 let mut b = __s_b.launch_builder(&f);
9465 b.arg(q)
9466 .arg(k)
9467 .arg(v)
9468 .arg(wq)
9469 .arg(wk)
9470 .arg(wv)
9471 .arg(dq)
9472 .arg(dk)
9473 .arg(dv)
9474 .arg(&nc)
9475 .arg(&rqi)
9476 .arg(&rki)
9477 .arg(&rvi)
9478 .arg(&e);
9479 unsafe {
9480 b.launch(cfg)?;
9481 }
9482 return Ok(());
9483 }
9484 let f = self.func("rms_norm_qkv_f32");
9485 let grid = (rq + 2 * rk) as u32;
9486 let cfg = LaunchConfig {
9487 grid_dim: (grid, 1, 1),
9488 block_dim: (rms_block(), 1, 1),
9489 shared_mem_bytes: 0,
9490 };
9491 let (nc, rqi, rki, e) = (ncols as i32, rq as i32, rk as i32, eps);
9492 let __s_b = self.gpu.stream();
9493 let mut b = __s_b.launch_builder(&f);
9494 b.arg(q)
9495 .arg(k)
9496 .arg(v)
9497 .arg(wq)
9498 .arg(wk)
9499 .arg(wv)
9500 .arg(dq)
9501 .arg(dk)
9502 .arg(dv)
9503 .arg(&nc)
9504 .arg(&rqi)
9505 .arg(&rki)
9506 .arg(&e);
9507 unsafe {
9508 b.launch(cfg)?;
9509 }
9510 Ok(())
9511 }
9512
9513 #[allow(clippy::too_many_arguments)]
9515 pub fn rms_norm2x(
9516 &self,
9517 a: &CudaSlice<f32>,
9518 bb: &CudaSlice<f32>,
9519 wa: &CudaSlice<f32>,
9520 wb: &CudaSlice<f32>,
9521 da: &mut CudaSlice<f32>,
9522 db: &mut CudaSlice<f32>,
9523 ncols: usize,
9524 nrows: usize,
9525 eps: f32,
9526 ) -> Result<(), Box<dyn std::error::Error>> {
9527 let f = self.func("rms_norm2x_f32");
9528 let cfg = LaunchConfig {
9529 grid_dim: (2 * nrows as u32, 1, 1),
9530 block_dim: (rms_block(), 1, 1),
9531 shared_mem_bytes: 0,
9532 };
9533 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
9534 let __s_b = self.gpu.stream();
9535 let mut b = __s_b.launch_builder(&f);
9536 b.arg(a)
9537 .arg(bb)
9538 .arg(wa)
9539 .arg(wb)
9540 .arg(da)
9541 .arg(db)
9542 .arg(&nc)
9543 .arg(&nr)
9544 .arg(&e);
9545 unsafe {
9546 b.launch(cfg)?;
9547 }
9548 Ok(())
9549 }
9550
9551 pub fn softcap(
9553 &self,
9554 y: &mut CudaSlice<f32>,
9555 cap: f32,
9556 n: usize,
9557 ) -> Result<(), Box<dyn std::error::Error>> {
9558 let f = self.func("softcap_f32");
9559 let cfg = LaunchConfig::for_num_elems(n as u32);
9560 let ni = n as i32;
9561 let __s_b = self.gpu.stream();
9562 let mut b = __s_b.launch_builder(&f);
9563 b.arg(y).arg(&cap).arg(&ni);
9564 unsafe {
9565 b.launch(cfg)?;
9566 }
9567 Ok(())
9568 }
9569
9570 pub fn mask_ids_rows(
9573 &self,
9574 y: &mut CudaSlice<f32>,
9575 ids: &CudaSlice<i32>,
9576 n_ids: usize,
9577 n_vocab: usize,
9578 t: usize,
9579 ) -> Result<(), Box<dyn std::error::Error>> {
9580 let f = self.func("mask_ids_rows_f32");
9581 let cfg = LaunchConfig::for_num_elems((n_ids * t) as u32);
9582 let (ni, nv, ti) = (n_ids as i32, n_vocab as i32, t as i32);
9583 let __s_b = self.gpu.stream();
9584 let mut b = __s_b.launch_builder(&f);
9585 b.arg(y).arg(ids).arg(&ni).arg(&nv).arg(&ti);
9586 unsafe {
9587 b.launch(cfg)?;
9588 }
9589 Ok(())
9590 }
9591
9592 #[allow(clippy::too_many_arguments)]
9594 pub fn add_scale_rms_norm(
9595 &self,
9596 a: &CudaSlice<f32>,
9597 b_in: &CudaSlice<f32>,
9598 c: f32,
9599 w: &CudaSlice<f32>,
9600 res: &mut CudaSlice<f32>,
9601 dst: &mut CudaSlice<f32>,
9602 ncols: usize,
9603 nrows: usize,
9604 eps: f32,
9605 ) -> Result<(), Box<dyn std::error::Error>> {
9606 let f = self.func("add_scale_rms_norm_f32");
9607 let cfg = LaunchConfig {
9608 grid_dim: (nrows as u32, 1, 1),
9609 block_dim: (rms_block(), 1, 1),
9610 shared_mem_bytes: 0,
9611 };
9612 let (nc, e2) = (ncols as i32, eps);
9613 let __s_b = self.gpu.stream();
9614 let mut b = __s_b.launch_builder(&f);
9615 b.arg(a)
9616 .arg(b_in)
9617 .arg(&c)
9618 .arg(w)
9619 .arg(res)
9620 .arg(dst)
9621 .arg(&nc)
9622 .arg(&e2);
9623 unsafe {
9624 b.launch(cfg)?;
9625 }
9626 Ok(())
9627 }
9628
9629 #[allow(clippy::too_many_arguments)]
9632 pub fn add_scale_rms_norm_q8_1(
9633 &self,
9634 a: &CudaSlice<f32>,
9635 b_in: &CudaSlice<f32>,
9636 c: f32,
9637 w: &CudaSlice<f32>,
9638 res: &mut CudaSlice<f32>,
9639 ncols: usize,
9640 nrows: usize,
9641 eps: f32,
9642 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9643 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
9644 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
9645 let (nc, e2) = (ncols as i32, eps);
9646 if Self::pdl_on() && Self::pdl_wb_on() {
9647 {
9648 use cudarc::driver::{DevicePtr, DevicePtrMut};
9649 let s = &self.gpu.stream();
9650 let (pa, _g0) = a.device_ptr(s);
9651 let (pb, _g1) = b_in.device_ptr(s);
9652 let (pw, _g2) = w.device_ptr(s);
9653 let (pr, _g3) = res.device_ptr_mut(s);
9654 let (pq, _g4) = out_q.device_ptr_mut(s);
9655 let (pd, _g5) = out_d.device_ptr_mut(s);
9656 let mut ps = [
9657 &pa as *const _ as *mut std::ffi::c_void,
9658 &pb as *const _ as *mut _,
9659 &c as *const _ as *mut _,
9660 &pw as *const _ as *mut _,
9661 &pr as *const _ as *mut _,
9662 &pq as *const _ as *mut _,
9663 &pd as *const _ as *mut _,
9664 &nc as *const _ as *mut _,
9665 &e2 as *const _ as *mut _,
9666 ];
9667 unsafe {
9668 self.launch_pdl(
9669 "add_scale_rms_norm_q8_1",
9670 (nrows as u32, 1, 1),
9671 (rms_block(), 1, 1),
9672 &mut ps,
9673 )?;
9674 }
9675 }
9676 return Ok((out_q, out_d));
9677 }
9678 let f = self.func("add_scale_rms_norm_q8_1");
9679 let cfg = LaunchConfig {
9680 grid_dim: (nrows as u32, 1, 1),
9681 block_dim: (rms_block(), 1, 1),
9682 shared_mem_bytes: 0,
9683 };
9684 let __s_b = self.gpu.stream();
9685 let mut b = __s_b.launch_builder(&f);
9686 b.arg(a)
9687 .arg(b_in)
9688 .arg(&c)
9689 .arg(w)
9690 .arg(res)
9691 .arg(&mut out_q)
9692 .arg(&mut out_d)
9693 .arg(&nc)
9694 .arg(&e2);
9695 unsafe {
9696 b.launch(cfg)?;
9697 }
9698 Ok((out_q, out_d))
9699 }
9700
9701 #[allow(clippy::too_many_arguments)]
9703 pub fn add_scale_rms_norm_q8_1_into(
9704 &self,
9705 a: &CudaSlice<f32>,
9706 b_in: &CudaSlice<f32>,
9707 c: f32,
9708 w: &CudaSlice<f32>,
9709 res: &mut CudaSlice<f32>,
9710 ncols: usize,
9711 nrows: usize,
9712 eps: f32,
9713 out_q: &mut CudaSlice<i8>,
9714 out_d: &mut CudaSlice<f32>,
9715 ) -> Result<(), Box<dyn std::error::Error>> {
9716 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
9717 let (nc, e2) = (ncols as i32, eps);
9718 if Self::pdl_on() && Self::pdl_wb_on() {
9719 use cudarc::driver::{DevicePtr, DevicePtrMut};
9720 let s = &self.gpu.stream();
9721 let (pa, _g0) = a.device_ptr(s);
9722 let (pb, _g1) = b_in.device_ptr(s);
9723 let (pw, _g2) = w.device_ptr(s);
9724 let (pr, _g3) = res.device_ptr_mut(s);
9725 let (pq, _g4) = out_q.device_ptr_mut(s);
9726 let (pd, _g5) = out_d.device_ptr_mut(s);
9727 let mut ps = [
9728 &pa as *const _ as *mut std::ffi::c_void,
9729 &pb as *const _ as *mut _,
9730 &c as *const _ as *mut _,
9731 &pw as *const _ as *mut _,
9732 &pr as *const _ as *mut _,
9733 &pq as *const _ as *mut _,
9734 &pd as *const _ as *mut _,
9735 &nc as *const _ as *mut _,
9736 &e2 as *const _ as *mut _,
9737 ];
9738 unsafe {
9739 self.launch_pdl(
9740 "add_scale_rms_norm_q8_1",
9741 (nrows as u32, 1, 1),
9742 (rms_block(), 1, 1),
9743 &mut ps,
9744 )?;
9745 }
9746 return Ok(());
9747 }
9748 let f = self.func("add_scale_rms_norm_q8_1");
9749 let cfg = LaunchConfig {
9750 grid_dim: (nrows as u32, 1, 1),
9751 block_dim: (rms_block(), 1, 1),
9752 shared_mem_bytes: 0,
9753 };
9754 let __s_b = self.gpu.stream();
9755 let mut b = __s_b.launch_builder(&f);
9756 b.arg(a)
9757 .arg(b_in)
9758 .arg(&c)
9759 .arg(w)
9760 .arg(res)
9761 .arg(&mut *out_q)
9762 .arg(&mut *out_d)
9763 .arg(&nc)
9764 .arg(&e2);
9765 unsafe {
9766 b.launch(cfg)?;
9767 }
9768 Ok(())
9769 }
9770
9771 #[allow(clippy::too_many_arguments)]
9774 pub fn rms_pre_add_scale_rms_norm_q8_1(
9775 &self,
9776 a: &CudaSlice<f32>,
9777 wa: &CudaSlice<f32>,
9778 b_in: &CudaSlice<f32>,
9779 c: f32,
9780 w: &CudaSlice<f32>,
9781 res: &mut CudaSlice<f32>,
9782 ncols: usize,
9783 nrows: usize,
9784 eps: f32,
9785 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9786 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
9787 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
9788 let (nc, e2) = (ncols as i32, eps);
9789 if Self::pdl_on() {
9790 {
9791 use cudarc::driver::{DevicePtr, DevicePtrMut};
9792 let s = &self.gpu.stream();
9793 let (pa, _g0) = a.device_ptr(s);
9794 let (pwa, _g1) = wa.device_ptr(s);
9795 let (pb, _g2) = b_in.device_ptr(s);
9796 let (pw, _g3) = w.device_ptr(s);
9797 let (pr, _g4) = res.device_ptr_mut(s);
9798 let (pq, _g5) = out_q.device_ptr_mut(s);
9799 let (pd, _g6) = out_d.device_ptr_mut(s);
9800 let mut ps = [
9801 &pa as *const _ as *mut std::ffi::c_void,
9802 &pwa as *const _ as *mut _,
9803 &pb as *const _ as *mut _,
9804 &c as *const _ as *mut _,
9805 &pw as *const _ as *mut _,
9806 &pr as *const _ as *mut _,
9807 &pq as *const _ as *mut _,
9808 &pd as *const _ as *mut _,
9809 &nc as *const _ as *mut _,
9810 &e2 as *const _ as *mut _,
9811 ];
9812 unsafe {
9813 self.launch_pdl(
9814 "rms_pre_add_scale_rms_norm_q8_1",
9815 (nrows as u32, 1, 1),
9816 (rms_block(), 1, 1),
9817 &mut ps,
9818 )?;
9819 }
9820 }
9821 return Ok((out_q, out_d));
9822 }
9823 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
9824 let cfg = LaunchConfig {
9825 grid_dim: (nrows as u32, 1, 1),
9826 block_dim: (rms_block(), 1, 1),
9827 shared_mem_bytes: 0,
9828 };
9829 let __s_b = self.gpu.stream();
9830 let mut b = __s_b.launch_builder(&f);
9831 b.arg(a)
9832 .arg(wa)
9833 .arg(b_in)
9834 .arg(&c)
9835 .arg(w)
9836 .arg(res)
9837 .arg(&mut out_q)
9838 .arg(&mut out_d)
9839 .arg(&nc)
9840 .arg(&e2);
9841 unsafe {
9842 b.launch(cfg)?;
9843 }
9844 Ok((out_q, out_d))
9845 }
9846
9847 pub fn gelu_tanh_mul_q8_1(
9850 &self,
9851 gate: &CudaSlice<f32>,
9852 up: &cudarc::driver::CudaView<f32>,
9853 act: &mut CudaSlice<f32>,
9854 ncols: usize,
9855 nrows: usize,
9856 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9857 debug_assert!(ncols % 128 == 0);
9858 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
9859 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
9860 let nc = ncols as i32;
9861 if Self::pdl_on() {
9862 {
9863 use cudarc::driver::{DevicePtr, DevicePtrMut};
9864 let s = &self.gpu.stream();
9865 let (pg, _g0) = gate.device_ptr(s);
9866 let (pu, _g1) = up.device_ptr(s);
9867 let (pact, _g2) = act.device_ptr_mut(s);
9868 let (pq, _g3) = out_q.device_ptr_mut(s);
9869 let (pd, _g4) = out_d.device_ptr_mut(s);
9870 let mut ps = [
9871 &pg as *const _ as *mut std::ffi::c_void,
9872 &pu as *const _ as *mut _,
9873 &pact as *const _ as *mut _,
9874 &pq as *const _ as *mut _,
9875 &pd as *const _ as *mut _,
9876 &nc as *const _ as *mut _,
9877 ];
9878 unsafe {
9879 self.launch_pdl(
9880 "gelu_tanh_mul_q8_1",
9881 (nrows as u32, 1, 1),
9882 (rms_block(), 1, 1),
9883 &mut ps,
9884 )?;
9885 }
9886 }
9887 return Ok((out_q, out_d));
9888 }
9889 let f = self.func("gelu_tanh_mul_q8_1");
9890 let cfg = LaunchConfig {
9891 grid_dim: (nrows as u32, 1, 1),
9892 block_dim: (rms_block(), 1, 1),
9893 shared_mem_bytes: 0,
9894 };
9895 let __s_b = self.gpu.stream();
9896 let mut b = __s_b.launch_builder(&f);
9897 b.arg(gate)
9898 .arg(up)
9899 .arg(act)
9900 .arg(&mut out_q)
9901 .arg(&mut out_d)
9902 .arg(&nc);
9903 unsafe {
9904 b.launch(cfg)?;
9905 }
9906 Ok((out_q, out_d))
9907 }
9908
9909 #[allow(clippy::too_many_arguments)]
9911 pub fn gelu_tanh_mul_q8_1_into(
9912 &self,
9913 gate: &CudaSlice<f32>,
9914 up: &cudarc::driver::CudaView<f32>,
9915 act: &mut CudaSlice<f32>,
9916 ncols: usize,
9917 nrows: usize,
9918 out_q: &mut CudaSlice<i8>,
9919 out_d: &mut CudaSlice<f32>,
9920 ) -> Result<(), Box<dyn std::error::Error>> {
9921 debug_assert!(ncols % 128 == 0);
9922 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
9923 let nc = ncols as i32;
9924 if Self::pdl_on() {
9925 use cudarc::driver::{DevicePtr, DevicePtrMut};
9926 let s = &self.gpu.stream();
9927 let (pg, _g0) = gate.device_ptr(s);
9928 let (pu, _g1) = up.device_ptr(s);
9929 let (pact, _g2) = act.device_ptr_mut(s);
9930 let (pq, _g3) = out_q.device_ptr_mut(s);
9931 let (pd, _g4) = out_d.device_ptr_mut(s);
9932 let mut ps = [
9933 &pg as *const _ as *mut std::ffi::c_void,
9934 &pu as *const _ as *mut _,
9935 &pact as *const _ as *mut _,
9936 &pq as *const _ as *mut _,
9937 &pd as *const _ as *mut _,
9938 &nc as *const _ as *mut _,
9939 ];
9940 unsafe {
9941 self.launch_pdl(
9942 "gelu_tanh_mul_q8_1",
9943 (nrows as u32, 1, 1),
9944 (rms_block(), 1, 1),
9945 &mut ps,
9946 )?;
9947 }
9948 return Ok(());
9949 }
9950 let f = self.func("gelu_tanh_mul_q8_1");
9951 let cfg = LaunchConfig {
9952 grid_dim: (nrows as u32, 1, 1),
9953 block_dim: (rms_block(), 1, 1),
9954 shared_mem_bytes: 0,
9955 };
9956 let __s_b = self.gpu.stream();
9957 let mut b = __s_b.launch_builder(&f);
9958 b.arg(gate)
9959 .arg(up)
9960 .arg(&mut *act)
9961 .arg(&mut *out_q)
9962 .arg(&mut *out_d)
9963 .arg(&nc);
9964 unsafe {
9965 b.launch(cfg)?;
9966 }
9967 Ok(())
9968 }
9969
9970 #[allow(clippy::too_many_arguments)]
9972 pub fn add_rms_norm3_q8z(
9973 &self,
9974 a: &CudaSlice<f32>,
9975 b_in: &CudaSlice<f32>,
9976 w0: &CudaSlice<f32>,
9977 w1: &CudaSlice<f32>,
9978 w2: &CudaSlice<f32>,
9979 res: &mut CudaSlice<f32>,
9980 out1: &mut CudaSlice<f32>,
9981 ncols: usize,
9982 nrows: usize,
9983 eps: f32,
9984 ) -> Result<
9985 (
9986 (CudaSlice<i8>, CudaSlice<f32>),
9987 (CudaSlice<i8>, CudaSlice<f32>),
9988 ),
9989 Box<dyn std::error::Error>,
9990 > {
9991 let mut q0 = self.alloc_uninit::<i8>(nrows * ncols)?;
9992 let mut d0 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
9993 let mut q2 = self.alloc_uninit::<i8>(nrows * ncols)?;
9994 let mut d2 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
9995 let f = self.func("add_rms_norm3_q8z_f32");
9996 let cfg = LaunchConfig {
9997 grid_dim: (nrows as u32, 1, 1),
9998 block_dim: (rms_block(), 1, 1),
9999 shared_mem_bytes: 0,
10000 };
10001 let (nc, e2) = (ncols as i32, eps);
10002 let __s_b = self.gpu.stream();
10003 let mut b = __s_b.launch_builder(&f);
10004 b.arg(a)
10005 .arg(b_in)
10006 .arg(w0)
10007 .arg(w1)
10008 .arg(w2)
10009 .arg(res)
10010 .arg(&mut q0)
10011 .arg(&mut d0)
10012 .arg(out1)
10013 .arg(&mut q2)
10014 .arg(&mut d2)
10015 .arg(&nc)
10016 .arg(&e2);
10017 unsafe {
10018 b.launch(cfg)?;
10019 }
10020 Ok(((q0, d0), (q2, d2)))
10021 }
10022
10023 #[allow(clippy::too_many_arguments)]
10025 pub fn add_rms_norm3(
10026 &self,
10027 a: &CudaSlice<f32>,
10028 b_in: &CudaSlice<f32>,
10029 w0: &CudaSlice<f32>,
10030 w1: &CudaSlice<f32>,
10031 w2: &CudaSlice<f32>,
10032 res: &mut CudaSlice<f32>,
10033 d0: &mut CudaSlice<f32>,
10034 d1: &mut CudaSlice<f32>,
10035 d2: &mut CudaSlice<f32>,
10036 ncols: usize,
10037 nrows: usize,
10038 eps: f32,
10039 ) -> Result<(), Box<dyn std::error::Error>> {
10040 let f = self.func("add_rms_norm3_f32");
10041 let cfg = LaunchConfig {
10042 grid_dim: (nrows as u32, 1, 1),
10043 block_dim: (rms_block(), 1, 1),
10044 shared_mem_bytes: 0,
10045 };
10046 let (nc, e2) = (ncols as i32, eps);
10047 let __s_b = self.gpu.stream();
10048 let mut b = __s_b.launch_builder(&f);
10049 b.arg(a)
10050 .arg(b_in)
10051 .arg(w0)
10052 .arg(w1)
10053 .arg(w2)
10054 .arg(res)
10055 .arg(d0)
10056 .arg(d1)
10057 .arg(d2)
10058 .arg(&nc)
10059 .arg(&e2);
10060 unsafe {
10061 b.launch(cfg)?;
10062 }
10063 Ok(())
10064 }
10065
10066 pub fn add_scale(
10068 &self,
10069 a: &CudaSlice<f32>,
10070 b_in: &CudaSlice<f32>,
10071 c: f32,
10072 dst: &mut CudaSlice<f32>,
10073 n: usize,
10074 ) -> Result<(), Box<dyn std::error::Error>> {
10075 let f = self.func("add_scale_f32");
10076 let cfg = LaunchConfig::for_num_elems(n as u32);
10077 let ni = n as i32;
10078 let __s_b = self.gpu.stream();
10079 let mut b = __s_b.launch_builder(&f);
10080 b.arg(a).arg(b_in).arg(&c).arg(dst).arg(&ni);
10081 unsafe {
10082 b.launch(cfg)?;
10083 }
10084 Ok(())
10085 }
10086
10087 pub fn layer_norm_bias(
10089 &self,
10090 x: &CudaSlice<f32>,
10091 w: &CudaSlice<f32>,
10092 b: &CudaSlice<f32>,
10093 dst: &mut CudaSlice<f32>,
10094 ncols: usize,
10095 nrows: usize,
10096 eps: f32,
10097 ) -> Result<(), Box<dyn std::error::Error>> {
10098 let f = self.func("layer_norm_bias_f32");
10099 let (nc, e) = (ncols as i32, eps);
10100 let cfg = LaunchConfig {
10101 grid_dim: (nrows as u32, 1, 1),
10102 block_dim: (256, 1, 1),
10103 shared_mem_bytes: 0,
10104 };
10105 let __s_b = self.gpu.stream();
10106 let mut lb = __s_b.launch_builder(&f);
10107 lb.arg(x).arg(w).arg(b).arg(&mut *dst).arg(&nc).arg(&e);
10108 unsafe {
10109 lb.launch(cfg)?;
10110 }
10111 Ok(())
10112 }
10113
10114 pub fn gelu_tanh(
10116 &self,
10117 x: &CudaSlice<f32>,
10118 dst: &mut CudaSlice<f32>,
10119 n: usize,
10120 ) -> Result<(), Box<dyn std::error::Error>> {
10121 let f = self.func("gelu_tanh_f32");
10122 let ni = n as i64;
10123 let cfg = LaunchConfig {
10124 grid_dim: (n.div_ceil(256) as u32, 1, 1),
10125 block_dim: (256, 1, 1),
10126 shared_mem_bytes: 0,
10127 };
10128 let __s_b = self.gpu.stream();
10129 let mut lb = __s_b.launch_builder(&f);
10130 lb.arg(x).arg(&mut *dst).arg(&ni);
10131 unsafe {
10132 lb.launch(cfg)?;
10133 }
10134 Ok(())
10135 }
10136
10137 pub fn row_softmax(
10139 &self,
10140 x: &mut CudaSlice<f32>,
10141 ncols: usize,
10142 nrows: usize,
10143 ) -> Result<(), Box<dyn std::error::Error>> {
10144 let f = self.func("row_softmax_f32");
10145 let nc = ncols as i32;
10146 let cfg = LaunchConfig {
10147 grid_dim: (nrows as u32, 1, 1),
10148 block_dim: (256, 1, 1),
10149 shared_mem_bytes: 0,
10150 };
10151 let __s_b = self.gpu.stream();
10152 let mut lb = __s_b.launch_builder(&f);
10153 lb.arg(&mut *x).arg(&nc);
10154 unsafe {
10155 lb.launch(cfg)?;
10156 }
10157 Ok(())
10158 }
10159
10160 pub fn rms_norm(
10161 &self,
10162 x: &CudaSlice<f32>,
10163 w: &CudaSlice<f32>,
10164 dst: &mut CudaSlice<f32>,
10165 ncols: usize,
10166 nrows: usize,
10167 eps: f32,
10168 ) -> Result<(), Box<dyn std::error::Error>> {
10169 let (nc, e) = (ncols as i32, eps);
10170 let kname = if Self::norm_ilp_on() {
10171 "rms_norm_f32_v2"
10172 } else {
10173 "rms_norm_f32"
10174 };
10175 if Self::pdl_on() && Self::pdl_wb_on() {
10176 use cudarc::driver::{DevicePtr, DevicePtrMut};
10177 let s = &self.gpu.stream();
10178 let (px, _g0) = x.device_ptr(s);
10179 let (pw, _g1) = w.device_ptr(s);
10180 let (pd, _g2) = dst.device_ptr_mut(s);
10181 let mut ps = [
10182 &px as *const _ as *mut std::ffi::c_void,
10183 &pw as *const _ as *mut _,
10184 &pd as *const _ as *mut _,
10185 &nc as *const _ as *mut _,
10186 &e as *const _ as *mut _,
10187 ];
10188 unsafe {
10189 self.launch_pdl(kname, (nrows as u32, 1, 1), (rms_block(), 1, 1), &mut ps)?;
10190 }
10191 return Ok(());
10192 }
10193 let f = self.func(kname);
10194 let cfg = LaunchConfig {
10195 grid_dim: (nrows as u32, 1, 1),
10196 block_dim: (rms_block(), 1, 1),
10197 shared_mem_bytes: 0,
10198 };
10199 let __s_b = self.gpu.stream();
10200 let mut b = __s_b.launch_builder(&f);
10201 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
10202 unsafe {
10203 b.launch(cfg)?;
10204 }
10205 Ok(())
10206 }
10207
10208 pub fn rms_norm_decode(
10216 &self,
10217 x: &CudaSlice<f32>,
10218 w: &CudaSlice<f32>,
10219 dst: &mut CudaSlice<f32>,
10220 ncols: usize,
10221 nrows: usize,
10222 eps: f32,
10223 ) -> Result<(), Box<dyn std::error::Error>> {
10224 let f = self.func(if Self::norm_ilp_on() {
10225 "rms_norm_f32_v2"
10226 } else {
10227 "rms_norm_f32"
10228 });
10229 let cfg = LaunchConfig {
10230 grid_dim: (nrows as u32, 1, 1),
10231 block_dim: (1024, 1, 1),
10232 shared_mem_bytes: 0,
10233 };
10234 let (nc, e) = (ncols as i32, eps);
10235 let __s_b = self.gpu.stream();
10236 let mut b = __s_b.launch_builder(&f);
10237 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
10238 unsafe {
10239 b.launch(cfg)?;
10240 }
10241 Ok(())
10242 }
10243
10244 pub fn rms_norm_q8_1(
10248 &self,
10249 x: &CudaSlice<f32>,
10250 w: &CudaSlice<f32>,
10251 ncols: usize,
10252 nrows: usize,
10253 eps: f32,
10254 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
10255 let nblk = ncols / 32;
10256 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
10257 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
10258 let (nc, e) = (ncols as i32, eps);
10259 if Self::pdl_on() {
10260 {
10261 use cudarc::driver::{DevicePtr, DevicePtrMut};
10262 let s = &self.gpu.stream();
10263 let (px, _g0) = x.device_ptr(s);
10264 let (pw, _g1) = w.device_ptr(s);
10265 let (pq, _g2) = q.device_ptr_mut(s);
10266 let (pd, _g3) = d.device_ptr_mut(s);
10267 let mut ps = [
10268 &px as *const _ as *mut std::ffi::c_void,
10269 &pw as *const _ as *mut _,
10270 &pq as *const _ as *mut _,
10271 &pd as *const _ as *mut _,
10272 &nc as *const _ as *mut _,
10273 &e as *const _ as *mut _,
10274 ];
10275 unsafe {
10276 self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1), &mut ps)?;
10277 }
10278 }
10279 return Ok((q, d));
10280 }
10281 let f = self.func("rms_norm_q8_1");
10282 let cfg = LaunchConfig {
10285 grid_dim: (nrows as u32, 1, 1),
10286 block_dim: (1024, 1, 1),
10287 shared_mem_bytes: 0,
10288 };
10289 let __s_b = self.gpu.stream();
10290 let mut b = __s_b.launch_builder(&f);
10291 b.arg(x).arg(w).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
10292 unsafe {
10293 b.launch(cfg)?;
10294 }
10295 Ok((q, d))
10296 }
10297
10298 pub fn rms_norm_q8_1_into(
10301 &self,
10302 x: &CudaSlice<f32>,
10303 w: &CudaSlice<f32>,
10304 ncols: usize,
10305 nrows: usize,
10306 eps: f32,
10307 q: &mut CudaSlice<i8>,
10308 d: &mut CudaSlice<f32>,
10309 ) -> Result<(), Box<dyn std::error::Error>> {
10310 let nblk = ncols / 32;
10311 debug_assert!(q.len() >= nrows * ncols && d.len() >= nrows * nblk);
10312 let (nc, e) = (ncols as i32, eps);
10313 if Self::pdl_on() {
10314 use cudarc::driver::{DevicePtr, DevicePtrMut};
10315 let s = &self.gpu.stream();
10316 let (px, _g0) = x.device_ptr(s);
10317 let (pw, _g1) = w.device_ptr(s);
10318 let (pq, _g2) = q.device_ptr_mut(s);
10319 let (pd, _g3) = d.device_ptr_mut(s);
10320 let mut ps = [
10321 &px as *const _ as *mut std::ffi::c_void,
10322 &pw as *const _ as *mut _,
10323 &pq as *const _ as *mut _,
10324 &pd as *const _ as *mut _,
10325 &nc as *const _ as *mut _,
10326 &e as *const _ as *mut _,
10327 ];
10328 unsafe {
10329 self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1), &mut ps)?;
10330 }
10331 return Ok(());
10332 }
10333 let f = self.func("rms_norm_q8_1");
10334 let cfg = LaunchConfig {
10335 grid_dim: (nrows as u32, 1, 1),
10336 block_dim: (1024, 1, 1),
10337 shared_mem_bytes: 0,
10338 };
10339 let __s_b = self.gpu.stream();
10340 let mut b = __s_b.launch_builder(&f);
10341 b.arg(x).arg(w).arg(&mut *q).arg(&mut *d).arg(&nc).arg(&e);
10342 unsafe {
10343 b.launch(cfg)?;
10344 }
10345 Ok(())
10346 }
10347
10348 pub fn quantize_q8_1_into(
10350 &self,
10351 x: &CudaSlice<f32>,
10352 m: usize,
10353 in_f: usize,
10354 q: &mut CudaSlice<i8>,
10355 d: &mut CudaSlice<f32>,
10356 ) -> Result<(), Box<dyn std::error::Error>> {
10357 let nblk = in_f / 32;
10358 debug_assert!(q.len() >= m * in_f && d.len() >= m * nblk);
10359 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
10360 let (inf, mi) = (in_f as i32, m as i32);
10361 if Self::pdl_on() && Self::pdl_wb_on() {
10362 use cudarc::driver::{DevicePtr, DevicePtrMut};
10363 let s = &self.gpu.stream();
10364 let (px, _g0) = x.device_ptr(s);
10365 let (pq, _g1) = q.device_ptr_mut(s);
10366 let (pd, _g2) = d.device_ptr_mut(s);
10367 let mut ps = [
10368 &px as *const _ as *mut std::ffi::c_void,
10369 &pq as *const _ as *mut _,
10370 &pd as *const _ as *mut _,
10371 &inf as *const _ as *mut _,
10372 &mi as *const _ as *mut _,
10373 ];
10374 unsafe {
10375 self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?;
10376 }
10377 return Ok(());
10378 }
10379 let f = self.func("quantize_q8_1");
10380 let __s_b = self.gpu.stream();
10381 let mut b = __s_b.launch_builder(&f);
10382 b.arg(x).arg(&mut *q).arg(&mut *d).arg(&inf).arg(&mi);
10383 unsafe {
10384 b.launch(cfg)?;
10385 }
10386 Ok(())
10387 }
10388
10389 pub fn add_rms_norm_q8_1(
10393 &self,
10394 a: &CudaSlice<f32>,
10395 b_in: &CudaSlice<f32>,
10396 w: &CudaSlice<f32>,
10397 res: &mut CudaSlice<f32>,
10398 ncols: usize,
10399 nrows: usize,
10400 eps: f32,
10401 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
10402 let nblk = ncols / 32;
10403 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
10404 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
10405 let f = self.func("add_rms_norm_q8_1");
10406 let cfg = LaunchConfig {
10408 grid_dim: (nrows as u32, 1, 1),
10409 block_dim: (1024, 1, 1),
10410 shared_mem_bytes: 0,
10411 };
10412 let (nc, e) = (ncols as i32, eps);
10413 let __s_bld = self.gpu.stream();
10414 let mut bld = __s_bld.launch_builder(&f);
10415 bld.arg(a)
10416 .arg(b_in)
10417 .arg(w)
10418 .arg(res)
10419 .arg(&mut q)
10420 .arg(&mut d)
10421 .arg(&nc)
10422 .arg(&e);
10423 unsafe {
10424 bld.launch(cfg)?;
10425 }
10426 Ok((q, d))
10427 }
10428
10429 #[allow(clippy::too_many_arguments)]
10435 pub fn join_add_rms_norm_raw(
10436 &self,
10437 a0_raw: u64,
10438 a1_raw: u64,
10439 x: &CudaSlice<f32>,
10440 w: &CudaSlice<f32>,
10441 res: &mut CudaSlice<f32>,
10442 dst: &mut CudaSlice<f32>,
10443 ncols: usize,
10444 eps: f32,
10445 ) -> Result<(), Box<dyn std::error::Error>> {
10446 if a0_raw == 0 || a1_raw == 0 || x.len() < ncols || res.len() < ncols || dst.len() < ncols {
10447 return Err("join_add_rms_norm geometry".into());
10448 }
10449 let f = self.func("join_add_rms_norm_f32");
10450 let cfg = LaunchConfig {
10451 grid_dim: (1, 1, 1),
10452 block_dim: (rms_block(), 1, 1),
10453 shared_mem_bytes: 0,
10454 };
10455 let (nc, e) = (ncols as i32, eps);
10456 let __s_b = self.gpu.stream();
10457 let mut b = __s_b.launch_builder(&f);
10458 b.arg(&a0_raw)
10459 .arg(&a1_raw)
10460 .arg(x)
10461 .arg(w)
10462 .arg(&mut *res)
10463 .arg(&mut *dst)
10464 .arg(&nc)
10465 .arg(&e);
10466 unsafe {
10467 b.launch(cfg)?;
10468 }
10469 Ok(())
10470 }
10471
10472 pub fn add_rms_norm(
10473 &self,
10474 a: &CudaSlice<f32>,
10475 b: &CudaSlice<f32>,
10476 w: &CudaSlice<f32>,
10477 res: &mut CudaSlice<f32>,
10478 dst: &mut CudaSlice<f32>,
10479 ncols: usize,
10480 nrows: usize,
10481 eps: f32,
10482 ) -> Result<(), Box<dyn std::error::Error>> {
10483 let (nc, e) = (ncols as i32, eps);
10484 let kname = if Self::norm_ilp_on() {
10485 "add_rms_norm_f32_v2"
10486 } else {
10487 "add_rms_norm_f32"
10488 };
10489 if Self::pdl_on() && Self::pdl_wb_on() {
10490 use cudarc::driver::{DevicePtr, DevicePtrMut};
10491 let s = &self.gpu.stream();
10492 let (pa, _g0) = a.device_ptr(s);
10493 let (pb, _g1) = b.device_ptr(s);
10494 let (pw, _g2) = w.device_ptr(s);
10495 let (pr, _g3) = res.device_ptr_mut(s);
10496 let (pd, _g4) = dst.device_ptr_mut(s);
10497 let mut ps = [
10498 &pa as *const _ as *mut std::ffi::c_void,
10499 &pb as *const _ as *mut _,
10500 &pw as *const _ as *mut _,
10501 &pr as *const _ as *mut _,
10502 &pd as *const _ as *mut _,
10503 &nc as *const _ as *mut _,
10504 &e as *const _ as *mut _,
10505 ];
10506 unsafe {
10507 self.launch_pdl(kname, (nrows as u32, 1, 1), (rms_block(), 1, 1), &mut ps)?;
10508 }
10509 return Ok(());
10510 }
10511 let f = self.func(kname);
10512 let cfg = LaunchConfig {
10513 grid_dim: (nrows as u32, 1, 1),
10514 block_dim: (rms_block(), 1, 1),
10515 shared_mem_bytes: 0,
10516 };
10517 let __s_b2 = self.gpu.stream();
10518 let mut b2 = __s_b2.launch_builder(&f);
10519 b2.arg(a)
10520 .arg(b)
10521 .arg(w)
10522 .arg(&mut *res)
10523 .arg(&mut *dst)
10524 .arg(&nc)
10525 .arg(&e);
10526 unsafe {
10527 b2.launch(cfg)?;
10528 }
10529 Ok(())
10530 }
10531
10532 #[allow(clippy::too_many_arguments)]
10535 pub fn rms_pre_add_rms_norm(
10536 &self,
10537 a: &CudaSlice<f32>,
10538 wa: &CudaSlice<f32>,
10539 b: &CudaSlice<f32>,
10540 w: &CudaSlice<f32>,
10541 res: &mut CudaSlice<f32>,
10542 dst: &mut CudaSlice<f32>,
10543 ncols: usize,
10544 nrows: usize,
10545 eps: f32,
10546 ) -> Result<(), Box<dyn std::error::Error>> {
10547 let f = self.func("rms_pre_add_rms_norm_f32");
10548 let cfg = LaunchConfig {
10549 grid_dim: (nrows as u32, 1, 1),
10550 block_dim: (rms_block(), 1, 1),
10551 shared_mem_bytes: 0,
10552 };
10553 let (nc, e) = (ncols as i32, eps);
10554 let __s_b2 = self.gpu.stream();
10555 let mut b2 = __s_b2.launch_builder(&f);
10556 b2.arg(a)
10557 .arg(wa)
10558 .arg(b)
10559 .arg(w)
10560 .arg(&mut *res)
10561 .arg(&mut *dst)
10562 .arg(&nc)
10563 .arg(&e);
10564 unsafe {
10565 b2.launch(cfg)?;
10566 }
10567 Ok(())
10568 }
10569
10570 #[allow(clippy::too_many_arguments)]
10572 pub fn rms_pre_add_rms_norm_q8z(
10573 &self,
10574 a: &CudaSlice<f32>,
10575 wa: &CudaSlice<f32>,
10576 b: &CudaSlice<f32>,
10577 w: &CudaSlice<f32>,
10578 res: &mut CudaSlice<f32>,
10579 dst: &mut CudaSlice<f32>,
10580 ncols: usize,
10581 nrows: usize,
10582 eps: f32,
10583 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
10584 debug_assert!(ncols % 128 == 0);
10585 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
10586 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
10587 let (nc, e) = (ncols as i32, eps);
10588 if Self::pdl_on() {
10589 {
10590 use cudarc::driver::{DevicePtr, DevicePtrMut};
10591 let s = &self.gpu.stream();
10592 let (pa, _g0) = a.device_ptr(s);
10593 let (pwa, _g1) = wa.device_ptr(s);
10594 let (pb, _g2) = b.device_ptr(s);
10595 let (pw, _g3) = w.device_ptr(s);
10596 let (pr, _g4) = res.device_ptr_mut(s);
10597 let (pdst, _g5) = dst.device_ptr_mut(s);
10598 let (pq, _g6) = out_q.device_ptr_mut(s);
10599 let (pd, _g7) = out_d.device_ptr_mut(s);
10600 let mut ps = [
10601 &pa as *const _ as *mut std::ffi::c_void,
10602 &pwa as *const _ as *mut _,
10603 &pb as *const _ as *mut _,
10604 &pw as *const _ as *mut _,
10605 &pr as *const _ as *mut _,
10606 &pdst as *const _ as *mut _,
10607 &pq as *const _ as *mut _,
10608 &pd as *const _ as *mut _,
10609 &nc as *const _ as *mut _,
10610 &e as *const _ as *mut _,
10611 ];
10612 unsafe {
10613 self.launch_pdl(
10614 "rms_pre_add_rms_norm_q8z_f32",
10615 (nrows as u32, 1, 1),
10616 (rms_block(), 1, 1),
10617 &mut ps,
10618 )?;
10619 }
10620 }
10621 return Ok((out_q, out_d));
10622 }
10623 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
10624 let cfg = LaunchConfig {
10625 grid_dim: (nrows as u32, 1, 1),
10626 block_dim: (rms_block(), 1, 1),
10627 shared_mem_bytes: 0,
10628 };
10629 let __s_b2 = self.gpu.stream();
10630 let mut b2 = __s_b2.launch_builder(&f);
10631 b2.arg(a)
10632 .arg(wa)
10633 .arg(b)
10634 .arg(w)
10635 .arg(&mut *res)
10636 .arg(&mut *dst)
10637 .arg(&mut out_q)
10638 .arg(&mut out_d)
10639 .arg(&nc)
10640 .arg(&e);
10641 unsafe {
10642 b2.launch(cfg)?;
10643 }
10644 Ok((out_q, out_d))
10645 }
10646
10647 #[allow(clippy::too_many_arguments)]
10651 pub fn rms_pre_add_rms_norm_q8z_into(
10652 &self,
10653 a: &CudaSlice<f32>,
10654 wa: &CudaSlice<f32>,
10655 b: &CudaSlice<f32>,
10656 w: &CudaSlice<f32>,
10657 res: &mut CudaSlice<f32>,
10658 dst: &mut CudaSlice<f32>,
10659 ncols: usize,
10660 nrows: usize,
10661 eps: f32,
10662 out_q: &mut CudaSlice<i8>,
10663 out_d: &mut CudaSlice<f32>,
10664 ) -> Result<(), Box<dyn std::error::Error>> {
10665 debug_assert!(ncols % 128 == 0);
10666 let (nc, e) = (ncols as i32, eps);
10667 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
10668 let cfg = LaunchConfig {
10669 grid_dim: (nrows as u32, 1, 1),
10670 block_dim: (rms_block(), 1, 1),
10671 shared_mem_bytes: 0,
10672 };
10673 let __s_b = self.gpu.stream();
10674 let mut b2 = __s_b.launch_builder(&f);
10675 b2.arg(a)
10676 .arg(wa)
10677 .arg(b)
10678 .arg(w)
10679 .arg(&mut *res)
10680 .arg(&mut *dst)
10681 .arg(&mut *out_q)
10682 .arg(&mut *out_d)
10683 .arg(&nc)
10684 .arg(&e);
10685 unsafe {
10686 b2.launch(cfg)?;
10687 }
10688 Ok(())
10689 }
10690
10691 #[allow(clippy::too_many_arguments)]
10694 pub fn rms_pre_add_scale_rms_norm_q8_1_into(
10695 &self,
10696 a: &CudaSlice<f32>,
10697 wa: &CudaSlice<f32>,
10698 b_in: &CudaSlice<f32>,
10699 c: f32,
10700 w: &CudaSlice<f32>,
10701 res: &mut CudaSlice<f32>,
10702 ncols: usize,
10703 nrows: usize,
10704 eps: f32,
10705 out_q: &mut CudaSlice<i8>,
10706 out_d: &mut CudaSlice<f32>,
10707 ) -> Result<(), Box<dyn std::error::Error>> {
10708 debug_assert!(ncols % 128 == 0);
10709 let (nc, e2) = (ncols as i32, eps);
10710 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
10711 let cfg = LaunchConfig {
10712 grid_dim: (nrows as u32, 1, 1),
10713 block_dim: (rms_block(), 1, 1),
10714 shared_mem_bytes: 0,
10715 };
10716 let __s_b = self.gpu.stream();
10717 let mut b2 = __s_b.launch_builder(&f);
10718 b2.arg(a)
10719 .arg(wa)
10720 .arg(b_in)
10721 .arg(&c)
10722 .arg(w)
10723 .arg(&mut *res)
10724 .arg(&mut *out_q)
10725 .arg(&mut *out_d)
10726 .arg(&nc)
10727 .arg(&e2);
10728 unsafe {
10729 b2.launch(cfg)?;
10730 }
10731 Ok(())
10732 }
10733
10734 pub fn g4_pnfold_on() -> bool {
10742 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10743 *ON.get_or_init(|| {
10744 std::env::var("MEMRA_G4_PNFOLD")
10745 .map(|v| v != "0")
10746 .unwrap_or(true)
10747 })
10748 }
10749
10750 pub fn build_q4_out_concat3(
10754 &self,
10755 w0: &crate::model::GpuTensor,
10756 w1: &crate::model::GpuTensor,
10757 w2: &crate::model::GpuTensor,
10758 ) -> Result<Option<crate::model::GpuTensor>, Box<dyn std::error::Error>> {
10759 use crate::model::GpuTensor;
10760 let part = |w: &GpuTensor| -> Option<(usize, usize)> {
10761 match w {
10762 GpuTensor::Quant {
10763 qtype,
10764 row_bytes,
10765 rp,
10766 ..
10767 } if *qtype == QT_Q4_0 && !*rp => Some((*row_bytes, w.out_features())),
10768 _ => None,
10769 }
10770 };
10771 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (part(w0), part(w1), part(w2))
10772 else {
10773 return Ok(None);
10774 };
10775 if rb0 != rb1
10776 || rb0 != rb2
10777 || w0.in_features() != w1.in_features()
10778 || w0.in_features() != w2.in_features()
10779 {
10780 return Ok(None);
10781 }
10782 fn bytes_of(w: &crate::model::GpuTensor) -> &CudaSlice<u8> {
10783 match w {
10784 crate::model::GpuTensor::Quant { bytes, .. } => bytes,
10785 _ => unreachable!(),
10786 }
10787 }
10788 let (b0, b1, b2) = (bytes_of(w0), bytes_of(w1), bytes_of(w2));
10789 let total = rb0 * (o0 + o1 + o2);
10790 let mut cat = self.alloc_u8(total)?;
10791 self.copy_u8_into(&mut cat, 0, b0, rb0 * o0)?;
10792 self.copy_u8_into(&mut cat, rb0 * o0, b1, rb1 * o1)?;
10793 self.copy_u8_into(&mut cat, rb0 * (o0 + o1), b2, rb2 * o2)?;
10794 Ok(Some(GpuTensor::Quant {
10795 bytes: cat,
10796 qtype: QT_Q4_0,
10797 row_bytes: rb0,
10798 ne: vec![w0.in_features() as u64, (o0 + o1 + o2) as u64],
10799 scale: 1.0,
10800 rp: false,
10801 #[cfg(memra_cutlass)]
10802 cutlass: None,
10803 fp8: None,
10804 blk: None,
10805 rp4: None,
10806 f16: None,
10807 }))
10808 }
10809
10810 fn full_width_rope_only(
10828 kernel: &str,
10829 n_rot: usize,
10830 head_dim: usize,
10831 ) -> Result<(), Box<dyn std::error::Error>> {
10832 if n_rot == head_dim {
10833 return Ok(());
10834 }
10835 Err(format!(
10836 "{kernel}: PARTIAL ROTARY REFUSED — n_rot {n_rot} != head_dim {head_dim}. This fused \
10837 rms_norm+qkv+rope kernel carries no n_dims parameter and rotates the full head \
10838 width (half = ncols/2), so it would rotate dims {n_rot}..{head_dim} that must pass \
10839 through unrotated. Use the split path (rms_norm_qkv + rope_neox/rope_neox2 with \
10840 n_dims={n_rot}), or add an n_dims early-return to the kernel and widen this guard."
10841 )
10842 .into())
10843 }
10844
10845 #[allow(clippy::too_many_arguments)]
10849 pub fn rms_norm_qkv_rope_cat(
10850 &self,
10851 qkv: &CudaSlice<f32>,
10852 wq: &CudaSlice<f32>,
10853 wk: &CudaSlice<f32>,
10854 wv: &CudaSlice<f32>,
10855 q: &mut CudaSlice<f32>,
10856 k: &mut CudaSlice<f32>,
10857 v: &mut CudaSlice<f32>,
10858 head_dim: usize,
10859 n_rot: usize,
10860 rq: usize,
10861 rk: usize,
10862 pos: &CudaSlice<i32>,
10863 nh_q: usize,
10864 nh_k: usize,
10865 base: f32,
10866 freq_scale: f32,
10867 ff: Option<&CudaSlice<f32>>,
10868 eps: f32,
10869 ) -> Result<(), Box<dyn std::error::Error>> {
10870 Self::full_width_rope_only("rms_norm_qkv_rope_cat", n_rot, head_dim)?;
10871 let rows = rq + rk + rk;
10872 let theta_scale = base.powf(-2.0 / head_dim as f32);
10873 let (nc, rqi, rki, nhq, nhk) = (
10874 head_dim as i32,
10875 rq as i32,
10876 rk as i32,
10877 nh_q as i32,
10878 nh_k as i32,
10879 );
10880 if Self::pdl_on() {
10881 use cudarc::driver::{DevicePtr, DevicePtrMut};
10882 let s = &self.gpu.stream();
10883 let (pqkv, _g0) = qkv.device_ptr(s);
10884 let (pwq, _g1) = wq.device_ptr(s);
10885 let (pwk, _g2) = wk.device_ptr(s);
10886 let (pwv, _g3) = wv.device_ptr(s);
10887 let (pq, _g4) = q.device_ptr_mut(s);
10888 let (pk, _g5) = k.device_ptr_mut(s);
10889 let (pv, _g6) = v.device_ptr_mut(s);
10890 let (ppos, _g7) = pos.device_ptr(s);
10891 let (pff, _g8) = match ff {
10892 Some(t) => {
10893 let (p, g) = t.device_ptr(s);
10894 (p, Some(g))
10895 }
10896 None => (0, None),
10897 };
10898 let mut ps = [
10899 &pqkv as *const _ as *mut std::ffi::c_void,
10900 &pwq as *const _ as *mut _,
10901 &pwk as *const _ as *mut _,
10902 &pwv as *const _ as *mut _,
10903 &pq as *const _ as *mut _,
10904 &pk as *const _ as *mut _,
10905 &pv as *const _ as *mut _,
10906 &nc as *const _ as *mut _,
10907 &rqi as *const _ as *mut _,
10908 &rki as *const _ as *mut _,
10909 &ppos as *const _ as *mut _,
10910 &nhq as *const _ as *mut _,
10911 &nhk as *const _ as *mut _,
10912 &theta_scale as *const _ as *mut _,
10913 &freq_scale as *const _ as *mut _,
10914 &pff as *const _ as *mut _,
10915 &eps as *const _ as *mut _,
10916 ];
10917 unsafe {
10918 self.launch_pdl(
10919 "rms_norm_qkv_rope_cat_f32",
10920 (rows as u32, 1, 1),
10921 (rms_block(), 1, 1),
10922 &mut ps,
10923 )?;
10924 }
10925 return Ok(());
10926 }
10927 let f = self.func("rms_norm_qkv_rope_cat_f32");
10928 let cfg = LaunchConfig {
10929 grid_dim: (rows as u32, 1, 1),
10930 block_dim: (rms_block(), 1, 1),
10931 shared_mem_bytes: 0,
10932 };
10933 let __s_b = self.gpu.stream();
10934 let mut b = __s_b.launch_builder(&f);
10935 match ff {
10936 Some(t) => {
10937 b.arg(qkv)
10938 .arg(wq)
10939 .arg(wk)
10940 .arg(wv)
10941 .arg(&mut *q)
10942 .arg(&mut *k)
10943 .arg(&mut *v)
10944 .arg(&nc)
10945 .arg(&rqi)
10946 .arg(&rki)
10947 .arg(pos)
10948 .arg(&nhq)
10949 .arg(&nhk)
10950 .arg(&theta_scale)
10951 .arg(&freq_scale)
10952 .arg(t)
10953 .arg(&eps);
10954 unsafe {
10955 b.launch(cfg)?;
10956 }
10957 }
10958 None => {
10959 let null: u64 = 0;
10960 b.arg(qkv)
10961 .arg(wq)
10962 .arg(wk)
10963 .arg(wv)
10964 .arg(&mut *q)
10965 .arg(&mut *k)
10966 .arg(&mut *v)
10967 .arg(&nc)
10968 .arg(&rqi)
10969 .arg(&rki)
10970 .arg(pos)
10971 .arg(&nhq)
10972 .arg(&nhk)
10973 .arg(&theta_scale)
10974 .arg(&freq_scale)
10975 .arg(&null)
10976 .arg(&eps);
10977 unsafe {
10978 b.launch(cfg)?;
10979 }
10980 }
10981 }
10982 Ok(())
10983 }
10984
10985 #[allow(clippy::too_many_arguments)]
10989 pub fn rms_norm_qkv_rope(
10990 &self,
10991 q0: &CudaSlice<f32>,
10992 k0: &CudaSlice<f32>,
10993 v0: &CudaSlice<f32>,
10994 wq: &CudaSlice<f32>,
10995 wk: &CudaSlice<f32>,
10996 wv: &CudaSlice<f32>,
10997 q: &mut CudaSlice<f32>,
10998 k: &mut CudaSlice<f32>,
10999 v: &mut CudaSlice<f32>,
11000 head_dim: usize,
11001 n_rot: usize,
11002 rq: usize,
11003 rk: usize,
11004 pos: &CudaSlice<i32>,
11005 nh_q: usize,
11006 nh_k: usize,
11007 base: f32,
11008 freq_scale: f32,
11009 ff: Option<&CudaSlice<f32>>,
11010 eps: f32,
11011 ) -> Result<(), Box<dyn std::error::Error>> {
11012 Self::full_width_rope_only("rms_norm_qkv_rope", n_rot, head_dim)?;
11013 let f = self.func("rms_norm_qkv_rope_f32");
11014 let rows = rq + rk + rk; let cfg = LaunchConfig {
11016 grid_dim: (rows as u32, 1, 1),
11017 block_dim: (rms_block(), 1, 1),
11018 shared_mem_bytes: 0,
11019 };
11020 let theta_scale = base.powf(-2.0 / head_dim as f32);
11021 let (nc, rqi, rki, nhq, nhk) = (
11022 head_dim as i32,
11023 rq as i32,
11024 rk as i32,
11025 nh_q as i32,
11026 nh_k as i32,
11027 );
11028 let __s_b = self.gpu.stream();
11029 let mut b = __s_b.launch_builder(&f);
11030 match ff {
11031 Some(t) => {
11032 b.arg(q0)
11033 .arg(k0)
11034 .arg(v0)
11035 .arg(wq)
11036 .arg(wk)
11037 .arg(wv)
11038 .arg(&mut *q)
11039 .arg(&mut *k)
11040 .arg(&mut *v)
11041 .arg(&nc)
11042 .arg(&rqi)
11043 .arg(&rki)
11044 .arg(pos)
11045 .arg(&nhq)
11046 .arg(&nhk)
11047 .arg(&theta_scale)
11048 .arg(&freq_scale)
11049 .arg(t)
11050 .arg(&eps);
11051 unsafe {
11052 b.launch(cfg)?;
11053 }
11054 }
11055 None => {
11056 let null: u64 = 0;
11057 b.arg(q0)
11058 .arg(k0)
11059 .arg(v0)
11060 .arg(wq)
11061 .arg(wk)
11062 .arg(wv)
11063 .arg(&mut *q)
11064 .arg(&mut *k)
11065 .arg(&mut *v)
11066 .arg(&nc)
11067 .arg(&rqi)
11068 .arg(&rki)
11069 .arg(pos)
11070 .arg(&nhq)
11071 .arg(&nhk)
11072 .arg(&theta_scale)
11073 .arg(&freq_scale)
11074 .arg(&null)
11075 .arg(&eps);
11076 unsafe {
11077 b.launch(cfg)?;
11078 }
11079 }
11080 }
11081 Ok(())
11082 }
11083
11084 #[allow(clippy::too_many_arguments)]
11090 pub fn rms_norm_qkv_rope_append_dc(
11091 &self,
11092 q0: &CudaSlice<f32>,
11093 k0: &CudaSlice<f32>,
11094 v0: &CudaSlice<f32>,
11095 wq: &CudaSlice<f32>,
11096 wk: &CudaSlice<f32>,
11097 wv: &CudaSlice<f32>,
11098 q: &mut CudaSlice<f32>,
11099 k: &mut CudaSlice<f32>,
11100 v: &mut CudaSlice<f32>,
11101 head_dim: usize,
11102 n_rot: usize,
11103 rq: usize,
11104 rk: usize,
11105 pos: &CudaSlice<i32>,
11106 nh_q: usize,
11107 nh_k: usize,
11108 base: f32,
11109 freq_scale: f32,
11110 ff: Option<&CudaSlice<f32>>,
11111 eps: f32,
11112 kc: &mut CudaSlice<u8>,
11113 vc: &mut CudaSlice<u8>,
11114 t_dev: &CudaSlice<i32>,
11115 k_tok_bytes: usize,
11116 v_tok_bytes: usize,
11117 g: bool,
11118 ) -> Result<(), Box<dyn std::error::Error>> {
11119 Self::full_width_rope_only("rms_norm_qkv_rope_append_dc", n_rot, head_dim)?;
11120 let rows = rq + rk + rk;
11121 let theta_scale = base.powf(-2.0 / head_dim as f32);
11122 let (nc, rqi, rki, nhq, nhk) = (
11123 head_dim as i32,
11124 rq as i32,
11125 rk as i32,
11126 nh_q as i32,
11127 nh_k as i32,
11128 );
11129 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
11130 if Self::pdl_on() && Self::pdl_wb_on() {
11131 use cudarc::driver::{DevicePtr, DevicePtrMut};
11132 let s = &self.gpu.stream();
11133 let (p0, _a0) = q0.device_ptr(s);
11134 let (p1, _a1) = k0.device_ptr(s);
11135 let (p2, _a2) = v0.device_ptr(s);
11136 let (pwq, _a3) = wq.device_ptr(s);
11137 let (pwk, _a4) = wk.device_ptr(s);
11138 let (pwv, _a5) = wv.device_ptr(s);
11139 let (pq, _a6) = q.device_ptr_mut(s);
11140 let (pk, _a7) = k.device_ptr_mut(s);
11141 let (pv, _a8) = v.device_ptr_mut(s);
11142 let (pp, _a9) = pos.device_ptr(s);
11143 let pff: u64 = match ff {
11144 Some(t) => {
11145 let (p, _gg) = t.device_ptr(s);
11146 p as u64
11147 }
11148 None => 0,
11149 };
11150 let (pkc, _a10) = kc.device_ptr_mut(s);
11151 let (pvc, _a11) = vc.device_ptr_mut(s);
11152 let (pt, _a12) = t_dev.device_ptr(s);
11153 let mut ps = [
11154 &p0 as *const _ as *mut std::ffi::c_void,
11155 &p1 as *const _ as *mut _,
11156 &p2 as *const _ as *mut _,
11157 &pwq as *const _ as *mut _,
11158 &pwk as *const _ as *mut _,
11159 &pwv as *const _ as *mut _,
11160 &pq as *const _ as *mut _,
11161 &pk as *const _ as *mut _,
11162 &pv as *const _ as *mut _,
11163 &nc as *const _ as *mut _,
11164 &rqi as *const _ as *mut _,
11165 &rki as *const _ as *mut _,
11166 &pp as *const _ as *mut _,
11167 &nhq as *const _ as *mut _,
11168 &nhk as *const _ as *mut _,
11169 &theta_scale as *const _ as *mut _,
11170 &freq_scale as *const _ as *mut _,
11171 &pff as *const _ as *mut _,
11172 &eps as *const _ as *mut _,
11173 &pkc as *const _ as *mut _,
11174 &pvc as *const _ as *mut _,
11175 &pt as *const _ as *mut _,
11176 &ktb as *const _ as *mut _,
11177 &vtb as *const _ as *mut _,
11178 ];
11179 unsafe {
11180 self.launch_pdl_flash(
11181 g,
11182 "rms_norm_qkv_rope_append_dc_f32",
11183 (rows as u32, 1, 1),
11184 (rms_block(), 1, 1),
11185 0,
11186 &mut ps,
11187 )?;
11188 }
11189 return Ok(());
11190 }
11191 let f = if g {
11192 self.func_g("rms_norm_qkv_rope_append_dc_f32")
11193 } else {
11194 self.func("rms_norm_qkv_rope_append_dc_f32")
11195 };
11196 let cfg = LaunchConfig {
11197 grid_dim: (rows as u32, 1, 1),
11198 block_dim: (rms_block(), 1, 1),
11199 shared_mem_bytes: 0,
11200 };
11201 let __s_b = self.gpu.stream();
11202 let mut b = __s_b.launch_builder(&f);
11203 match ff {
11204 Some(t) => {
11205 b.arg(q0)
11206 .arg(k0)
11207 .arg(v0)
11208 .arg(wq)
11209 .arg(wk)
11210 .arg(wv)
11211 .arg(&mut *q)
11212 .arg(&mut *k)
11213 .arg(&mut *v)
11214 .arg(&nc)
11215 .arg(&rqi)
11216 .arg(&rki)
11217 .arg(pos)
11218 .arg(&nhq)
11219 .arg(&nhk)
11220 .arg(&theta_scale)
11221 .arg(&freq_scale)
11222 .arg(t)
11223 .arg(&eps)
11224 .arg(&mut *kc)
11225 .arg(&mut *vc)
11226 .arg(t_dev)
11227 .arg(&ktb)
11228 .arg(&vtb);
11229 unsafe {
11230 b.launch(cfg)?;
11231 }
11232 }
11233 None => {
11234 let null: u64 = 0;
11235 b.arg(q0)
11236 .arg(k0)
11237 .arg(v0)
11238 .arg(wq)
11239 .arg(wk)
11240 .arg(wv)
11241 .arg(&mut *q)
11242 .arg(&mut *k)
11243 .arg(&mut *v)
11244 .arg(&nc)
11245 .arg(&rqi)
11246 .arg(&rki)
11247 .arg(pos)
11248 .arg(&nhq)
11249 .arg(&nhk)
11250 .arg(&theta_scale)
11251 .arg(&freq_scale)
11252 .arg(&null)
11253 .arg(&eps)
11254 .arg(&mut *kc)
11255 .arg(&mut *vc)
11256 .arg(t_dev)
11257 .arg(&ktb)
11258 .arg(&vtb);
11259 unsafe {
11260 b.launch(cfg)?;
11261 }
11262 }
11263 }
11264 Ok(())
11265 }
11266
11267 #[allow(clippy::too_many_arguments)]
11275 pub fn rms_norm_qkv_rope_append(
11276 &self,
11277 q0: &CudaSlice<f32>,
11278 k0: &CudaSlice<f32>,
11279 v0: &CudaSlice<f32>,
11280 wq: &CudaSlice<f32>,
11281 wk: &CudaSlice<f32>,
11282 wv: &CudaSlice<f32>,
11283 q: &mut CudaSlice<f32>,
11284 k: &mut CudaSlice<f32>,
11285 v: &mut CudaSlice<f32>,
11286 head_dim: usize,
11287 n_rot: usize,
11288 rq: usize,
11289 rk: usize,
11290 pos: &CudaSlice<i32>,
11291 nh_q: usize,
11292 nh_k: usize,
11293 base: f32,
11294 freq_scale: f32,
11295 ff: Option<&CudaSlice<f32>>,
11296 eps: f32,
11297 kc: &mut CudaSlice<u8>,
11298 vc: &mut CudaSlice<u8>,
11299 t: usize,
11300 k_tok_bytes: usize,
11301 v_tok_bytes: usize,
11302 g: bool,
11303 ) -> Result<(), Box<dyn std::error::Error>> {
11304 Self::full_width_rope_only("rms_norm_qkv_rope_append", n_rot, head_dim)?;
11305 let rows = rq + rk + rk;
11306 let theta_scale = base.powf(-2.0 / head_dim as f32);
11307 let (nc, rqi, rki, nhq, nhk) = (
11308 head_dim as i32,
11309 rq as i32,
11310 rk as i32,
11311 nh_q as i32,
11312 nh_k as i32,
11313 );
11314 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
11315 let ti = t as i32;
11316 if Self::pdl_on() && Self::pdl_wb_on() {
11317 use cudarc::driver::{DevicePtr, DevicePtrMut};
11318 let s = &self.gpu.stream();
11319 let (p0, _a0) = q0.device_ptr(s);
11320 let (p1, _a1) = k0.device_ptr(s);
11321 let (p2, _a2) = v0.device_ptr(s);
11322 let (pwq, _a3) = wq.device_ptr(s);
11323 let (pwk, _a4) = wk.device_ptr(s);
11324 let (pwv, _a5) = wv.device_ptr(s);
11325 let (pq, _a6) = q.device_ptr_mut(s);
11326 let (pk, _a7) = k.device_ptr_mut(s);
11327 let (pv, _a8) = v.device_ptr_mut(s);
11328 let (pp, _a9) = pos.device_ptr(s);
11329 let pff: u64 = match ff {
11330 Some(t) => {
11331 let (p, _gg) = t.device_ptr(s);
11332 p as u64
11333 }
11334 None => 0,
11335 };
11336 let (pkc, _a10) = kc.device_ptr_mut(s);
11337 let (pvc, _a11) = vc.device_ptr_mut(s);
11338 let mut ps = [
11339 &p0 as *const _ as *mut std::ffi::c_void,
11340 &p1 as *const _ as *mut _,
11341 &p2 as *const _ as *mut _,
11342 &pwq as *const _ as *mut _,
11343 &pwk as *const _ as *mut _,
11344 &pwv as *const _ as *mut _,
11345 &pq as *const _ as *mut _,
11346 &pk as *const _ as *mut _,
11347 &pv as *const _ as *mut _,
11348 &nc as *const _ as *mut _,
11349 &rqi as *const _ as *mut _,
11350 &rki as *const _ as *mut _,
11351 &pp as *const _ as *mut _,
11352 &nhq as *const _ as *mut _,
11353 &nhk as *const _ as *mut _,
11354 &theta_scale as *const _ as *mut _,
11355 &freq_scale as *const _ as *mut _,
11356 &pff as *const _ as *mut _,
11357 &eps as *const _ as *mut _,
11358 &pkc as *const _ as *mut _,
11359 &pvc as *const _ as *mut _,
11360 &ti as *const _ as *mut _,
11361 &ktb as *const _ as *mut _,
11362 &vtb as *const _ as *mut _,
11363 ];
11364 unsafe {
11365 self.launch_pdl_flash(
11366 g,
11367 "rms_norm_qkv_rope_append_f32",
11368 (rows as u32, 1, 1),
11369 (rms_block(), 1, 1),
11370 0,
11371 &mut ps,
11372 )?;
11373 }
11374 return Ok(());
11375 }
11376 let f = if g {
11377 self.func_g("rms_norm_qkv_rope_append_f32")
11378 } else {
11379 self.func("rms_norm_qkv_rope_append_f32")
11380 };
11381 let cfg = LaunchConfig {
11382 grid_dim: (rows as u32, 1, 1),
11383 block_dim: (rms_block(), 1, 1),
11384 shared_mem_bytes: 0,
11385 };
11386 let __s_b = self.gpu.stream();
11387 let mut b = __s_b.launch_builder(&f);
11388 let null: u64 = 0;
11389 b.arg(q0)
11390 .arg(k0)
11391 .arg(v0)
11392 .arg(wq)
11393 .arg(wk)
11394 .arg(wv)
11395 .arg(&mut *q)
11396 .arg(&mut *k)
11397 .arg(&mut *v)
11398 .arg(&nc)
11399 .arg(&rqi)
11400 .arg(&rki)
11401 .arg(pos)
11402 .arg(&nhq)
11403 .arg(&nhk)
11404 .arg(&theta_scale)
11405 .arg(&freq_scale);
11406 match ff {
11407 Some(t) => {
11408 b.arg(t);
11409 }
11410 None => {
11411 b.arg(&null);
11412 }
11413 }
11414 b.arg(&eps)
11415 .arg(&mut *kc)
11416 .arg(&mut *vc)
11417 .arg(&ti)
11418 .arg(&ktb)
11419 .arg(&vtb);
11420 unsafe {
11421 b.launch(cfg)?;
11422 }
11423 Ok(())
11424 }
11425
11426 pub fn add_q8_1(
11427 &self,
11428 a: &CudaSlice<f32>,
11429 b: &CudaSlice<f32>,
11430 res: &mut CudaSlice<f32>,
11431 ncols: usize,
11432 nrows: usize,
11433 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11434 debug_assert!(ncols % 128 == 0);
11435 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
11436 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
11437 let f = self.func("add_q8_1_f32");
11438 let cfg = LaunchConfig {
11439 grid_dim: (nrows as u32, 1, 1),
11440 block_dim: (rms_block(), 1, 1),
11441 shared_mem_bytes: 0,
11442 };
11443 let nc = ncols as i32;
11444 let __s_b2 = self.gpu.stream();
11445 let mut b2 = __s_b2.launch_builder(&f);
11446 b2.arg(a)
11447 .arg(b)
11448 .arg(&mut *res)
11449 .arg(&mut out_q)
11450 .arg(&mut out_d)
11451 .arg(&nc);
11452 unsafe {
11453 b2.launch(cfg)?;
11454 }
11455 Ok((out_q, out_d))
11456 }
11457
11458 pub fn rms_pre_add_q8_1(
11462 &self,
11463 a: &CudaSlice<f32>,
11464 wa: &CudaSlice<f32>,
11465 b: &CudaSlice<f32>,
11466 res: &mut CudaSlice<f32>,
11467 ncols: usize,
11468 nrows: usize,
11469 eps: f32,
11470 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11471 debug_assert!(ncols % 128 == 0);
11472 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
11473 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
11474 let f = self.func("rms_pre_add_q8_1_f32");
11475 let cfg = LaunchConfig {
11476 grid_dim: (nrows as u32, 1, 1),
11477 block_dim: (rms_block(), 1, 1),
11478 shared_mem_bytes: 0,
11479 };
11480 let (nc, ep) = (ncols as i32, eps);
11481 let __s_b2 = self.gpu.stream();
11482 let mut b2 = __s_b2.launch_builder(&f);
11483 b2.arg(a)
11484 .arg(wa)
11485 .arg(b)
11486 .arg(&mut *res)
11487 .arg(&mut out_q)
11488 .arg(&mut out_d)
11489 .arg(&nc)
11490 .arg(&ep);
11491 unsafe {
11492 b2.launch(cfg)?;
11493 }
11494 Ok((out_q, out_d))
11495 }
11496
11497 pub fn l2_v2_on(ncols: usize) -> bool {
11501 ncols == 128 && std::env::var("MEMRA_L2_V2").as_deref() != Ok("0")
11502 }
11503
11504 pub fn l2_norm_pp(
11505 &self,
11506 x: &CudaSlice<f32>,
11507 dst: &mut CudaSlice<f32>,
11508 dst16: Option<&mut CudaSlice<u8>>,
11509 ncols: usize,
11510 nrows: usize,
11511 eps: f32,
11512 ) -> Result<(), Box<dyn std::error::Error>> {
11513 if Self::l2_v2_on(ncols) {
11514 let f = self.func("l2_norm_pp_v2_f32");
11515 let rows_per_block = 8u32; let cfg = LaunchConfig {
11517 grid_dim: ((nrows as u32).div_ceil(rows_per_block), 1, 1),
11518 block_dim: (256, 1, 1),
11519 shared_mem_bytes: 0,
11520 };
11521 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
11522 let d16: u64 = match dst16 {
11524 Some(d) => self.addr_u8(d),
11525 None => 0,
11526 };
11527 let __s_b = self.gpu.stream();
11528 let mut b = __s_b.launch_builder(&f);
11529 b.arg(x).arg(dst).arg(&d16).arg(&nc).arg(&nr).arg(&e);
11530 unsafe {
11531 b.launch(cfg)?;
11532 }
11533 return Ok(());
11534 }
11535 self.l2_norm(x, dst, ncols, nrows, eps)
11536 }
11537
11538 pub fn l2_norm(
11539 &self,
11540 x: &CudaSlice<f32>,
11541 dst: &mut CudaSlice<f32>,
11542 ncols: usize,
11543 nrows: usize,
11544 eps: f32,
11545 ) -> Result<(), Box<dyn std::error::Error>> {
11546 let f = self.func("l2_norm_f32");
11547 let cfg = LaunchConfig {
11548 grid_dim: (nrows as u32, 1, 1),
11549 block_dim: (256, 1, 1),
11550 shared_mem_bytes: 0,
11551 };
11552 let (nc, e) = (ncols as i32, eps);
11553 let __s_b = self.gpu.stream();
11554 let mut b = __s_b.launch_builder(&f);
11555 b.arg(x).arg(dst).arg(&nc).arg(&e);
11556 unsafe {
11557 b.launch(cfg)?;
11558 }
11559 Ok(())
11560 }
11561
11562 pub fn l2_norm_decode(
11568 &self,
11569 x: &CudaSlice<f32>,
11570 dst: &mut CudaSlice<f32>,
11571 ncols: usize,
11572 nrows: usize,
11573 eps: f32,
11574 ) -> Result<(), Box<dyn std::error::Error>> {
11575 let f = self.func("l2_norm_f32");
11576 let cfg = LaunchConfig {
11577 grid_dim: (nrows as u32, 1, 1),
11578 block_dim: (32, 1, 1),
11579 shared_mem_bytes: 0,
11580 };
11581 let (nc, e) = (ncols as i32, eps);
11582 let __s_b = self.gpu.stream();
11583 let mut b = __s_b.launch_builder(&f);
11584 b.arg(x).arg(dst).arg(&nc).arg(&e);
11585 unsafe {
11586 b.launch(cfg)?;
11587 }
11588 Ok(())
11589 }
11590
11591 pub fn rope_neox(
11593 &self,
11594 x: &mut CudaSlice<f32>,
11595 pos: &CudaSlice<i32>,
11596 head_dim: usize,
11597 n_dims: usize,
11598 n_heads: usize,
11599 n_tokens: usize,
11600 freq_base: f32,
11601 freq_scale: f32,
11602 ) -> Result<(), Box<dyn std::error::Error>> {
11603 let f = self.func("rope_neox_f32");
11604 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
11605 let grid = (n_heads * n_tokens) as u32;
11606 let cfg = LaunchConfig {
11607 grid_dim: (grid, 1, 1),
11608 block_dim: ((head_dim / 2) as u32, 1, 1),
11609 shared_mem_bytes: 0,
11610 };
11611 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
11612 let __s_b = self.gpu.stream();
11613 let mut b = __s_b.launch_builder(&f);
11614 b.arg(x)
11615 .arg(pos)
11616 .arg(&hd)
11617 .arg(&nd)
11618 .arg(&nh)
11619 .arg(&theta_scale)
11620 .arg(&freq_scale);
11621 unsafe {
11622 b.launch(cfg)?;
11623 }
11624 Ok(())
11625 }
11626
11627 pub fn rope_neox_ff(
11629 &self,
11630 x: &mut CudaSlice<f32>,
11631 pos: &CudaSlice<i32>,
11632 head_dim: usize,
11633 n_dims: usize,
11634 n_heads: usize,
11635 n_tokens: usize,
11636 freq_base: f32,
11637 freq_scale: f32,
11638 ff: &CudaSlice<f32>,
11639 ) -> Result<(), Box<dyn std::error::Error>> {
11640 let f = self.func("rope_neox_ff_f32");
11641 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
11642 let grid = (n_heads * n_tokens) as u32;
11643 let cfg = LaunchConfig {
11644 grid_dim: (grid, 1, 1),
11645 block_dim: ((head_dim / 2) as u32, 1, 1),
11646 shared_mem_bytes: 0,
11647 };
11648 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
11649 let __s_b = self.gpu.stream();
11650 let mut b = __s_b.launch_builder(&f);
11651 b.arg(x)
11652 .arg(pos)
11653 .arg(&hd)
11654 .arg(&nd)
11655 .arg(&nh)
11656 .arg(&theta_scale)
11657 .arg(&freq_scale)
11658 .arg(ff);
11659 unsafe {
11660 b.launch(cfg)?;
11661 }
11662 Ok(())
11663 }
11664
11665 #[allow(clippy::too_many_arguments)]
11667 pub fn rope_neox2(
11668 &self,
11669 q: &mut CudaSlice<f32>,
11670 k: &mut CudaSlice<f32>,
11671 pos: &CudaSlice<i32>,
11672 head_dim: usize,
11673 n_dims: usize,
11674 nh_q: usize,
11675 nh_k: usize,
11676 n_tokens: usize,
11677 freq_base: f32,
11678 freq_scale: f32,
11679 ff: Option<&CudaSlice<f32>>,
11680 ) -> Result<(), Box<dyn std::error::Error>> {
11681 let f = self.func("rope_neox2_f32");
11682 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
11683 let grid = ((nh_q + nh_k) * n_tokens) as u32;
11684 let cfg = LaunchConfig {
11685 grid_dim: (grid, 1, 1),
11686 block_dim: ((head_dim / 2) as u32, 1, 1),
11687 shared_mem_bytes: 0,
11688 };
11689 let (hd, nd, nq, nk, nt) = (
11690 head_dim as i32,
11691 n_dims as i32,
11692 nh_q as i32,
11693 nh_k as i32,
11694 n_tokens as i32,
11695 );
11696 let __s_b = self.gpu.stream();
11697 let mut b = __s_b.launch_builder(&f);
11698 b.arg(q)
11699 .arg(k)
11700 .arg(pos)
11701 .arg(&hd)
11702 .arg(&nd)
11703 .arg(&nq)
11704 .arg(&nk)
11705 .arg(&nt)
11706 .arg(&theta_scale)
11707 .arg(&freq_scale);
11708 match ff {
11709 Some(ffv) => {
11710 b.arg(ffv);
11711 unsafe {
11712 b.launch(cfg)?;
11713 }
11714 }
11715 None => {
11716 let null: u64 = 0;
11717 b.arg(&null);
11718 unsafe {
11719 b.launch(cfg)?;
11720 }
11721 }
11722 }
11723 Ok(())
11724 }
11725
11726 pub fn gelu_tanh_mul(
11728 &self,
11729 gate: &CudaSlice<f32>,
11730 up: &CudaSlice<f32>,
11731 dst: &mut CudaSlice<f32>,
11732 n: usize,
11733 ) -> Result<(), Box<dyn std::error::Error>> {
11734 let f = self.func("gelu_tanh_mul_f32");
11735 let cfg = LaunchConfig::for_num_elems(n as u32);
11736 let ni = n as i32;
11737 let __s_b = self.gpu.stream();
11738 let mut b = __s_b.launch_builder(&f);
11739 b.arg(gate).arg(up).arg(dst).arg(&ni);
11740 unsafe {
11741 b.launch(cfg)?;
11742 }
11743 Ok(())
11744 }
11745
11746 pub fn silu_mul(
11747 &self,
11748 gate: &CudaSlice<f32>,
11749 up: &CudaSlice<f32>,
11750 dst: &mut CudaSlice<f32>,
11751 n: usize,
11752 ) -> Result<(), Box<dyn std::error::Error>> {
11753 let f = self.func("silu_mul_f32");
11754 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11756 let ni = n as i32;
11757 let __s_b = self.gpu.stream();
11758 let mut b = __s_b.launch_builder(&f);
11759 b.arg(gate).arg(up).arg(dst).arg(&ni);
11760 unsafe {
11761 b.launch(cfg)?;
11762 }
11763 Ok(())
11764 }
11765
11766 pub fn silu_mul_host_expf(
11768 &self,
11769 gate: &CudaSlice<f32>,
11770 up: &CudaSlice<f32>,
11771 dst: &mut CudaSlice<f32>,
11772 n: usize,
11773 ) -> Result<(), Box<dyn std::error::Error>> {
11774 let f = self.func("silu_mul_host_expf_f32");
11775 let cfg = LaunchConfig::for_num_elems(n as u32);
11776 let ni = n as i32;
11777 let __s_b = self.gpu.stream();
11778 let mut b = __s_b.launch_builder(&f);
11779 b.arg(gate).arg(up).arg(dst).arg(&ni);
11780 unsafe {
11781 b.launch(cfg)?;
11782 }
11783 Ok(())
11784 }
11785
11786 pub fn silu_clamped_mul_host_expf(
11788 &self,
11789 gate: &CudaSlice<f32>,
11790 up: &CudaSlice<f32>,
11791 limit: f32,
11792 dst: &mut CudaSlice<f32>,
11793 n: usize,
11794 ) -> Result<(), Box<dyn std::error::Error>> {
11795 if !limit.is_finite() || limit <= 0.0 {
11796 return Err(
11797 format!("Step routed-expert clamp limit must be positive, got {limit}").into(),
11798 );
11799 }
11800 let f = self.func("silu_clamped_mul_host_expf_f32");
11801 let cfg = LaunchConfig::for_num_elems(n as u32);
11802 let ni = n as i32;
11803 let __s_b = self.gpu.stream();
11804 let mut b = __s_b.launch_builder(&f);
11805 b.arg(gate).arg(up).arg(&limit).arg(dst).arg(&ni);
11806 unsafe {
11807 b.launch(cfg)?;
11808 }
11809 Ok(())
11810 }
11811
11812 pub fn silu_mul_f16out(
11815 &self,
11816 gate: &CudaSlice<f32>,
11817 up: &CudaSlice<f32>,
11818 dst: &mut CudaSlice<f32>,
11819 dst16: &mut CudaSlice<u8>,
11820 n: usize,
11821 ) -> Result<(), Box<dyn std::error::Error>> {
11822 let f = self.func("silu_mul_f16out_f32");
11823 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11824 let ni = n as i32;
11825 let __s_b = self.gpu.stream();
11826 let mut b = __s_b.launch_builder(&f);
11827 b.arg(gate).arg(up).arg(dst).arg(dst16).arg(&ni);
11828 unsafe {
11829 b.launch(cfg)?;
11830 }
11831 Ok(())
11832 }
11833
11834 pub fn silu_mul_scaled(
11841 &self,
11842 gate: &CudaSlice<f32>,
11843 up: &CudaSlice<f32>,
11844 gs: f32,
11845 us: f32,
11846 dst: &mut CudaSlice<f32>,
11847 n: usize,
11848 ) -> Result<(), Box<dyn std::error::Error>> {
11849 let f = self.func("silu_mul_scaled_f32");
11850 let cfg = LaunchConfig::for_num_elems(n as u32);
11851 let ni = n as i32;
11852 let (gsf, usf) = (gs, us);
11853 let __s_b = self.gpu.stream();
11854 let mut b = __s_b.launch_builder(&f);
11855 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(dst).arg(&ni);
11856 unsafe {
11857 b.launch(cfg)?;
11858 }
11859 Ok(())
11860 }
11861
11862 #[allow(clippy::too_many_arguments)]
11866 pub fn swigluoai_mul_scaled(
11867 &self,
11868 gate: &CudaSlice<f32>,
11869 up: &CudaSlice<f32>,
11870 gs: f32,
11871 us: f32,
11872 alpha: f32,
11873 limit: f32,
11874 dst: &mut CudaSlice<f32>,
11875 n: usize,
11876 ) -> Result<(), Box<dyn std::error::Error>> {
11877 let f = self.func("swigluoai_mul_scaled_f32");
11878 let cfg = LaunchConfig::for_num_elems(n as u32);
11879 let ni = n as i32;
11880 let __s_b = self.gpu.stream();
11881 let mut b = __s_b.launch_builder(&f);
11882 b.arg(gate)
11883 .arg(up)
11884 .arg(&gs)
11885 .arg(&us)
11886 .arg(&alpha)
11887 .arg(&limit)
11888 .arg(dst)
11889 .arg(&ni);
11890 unsafe {
11891 b.launch(cfg)?;
11892 }
11893 Ok(())
11894 }
11895
11896 pub fn silu_mul_scaled_q8_1(
11904 &self,
11905 gate: &CudaSlice<f32>,
11906 up: &CudaSlice<f32>,
11907 gs: f32,
11908 us: f32,
11909 n: usize,
11910 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11911 let f = self.func("silu_mul_scaled_q8_1");
11912 let nblk = n / 32;
11913 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);
11917 let (gsf, usf, ni) = (gs, us, n as i32);
11918 let __s_b = self.gpu.stream();
11919 let mut b = __s_b.launch_builder(&f);
11920 b.arg(gate)
11921 .arg(up)
11922 .arg(&gsf)
11923 .arg(&usf)
11924 .arg(&mut aq)
11925 .arg(&mut ad)
11926 .arg(&ni);
11927 unsafe {
11928 b.launch(cfg)?;
11929 }
11930 Ok((aq, ad))
11931 }
11932
11933 pub fn add(
11934 &self,
11935 a: &CudaSlice<f32>,
11936 b_in: &CudaSlice<f32>,
11937 dst: &mut CudaSlice<f32>,
11938 n: usize,
11939 ) -> Result<(), Box<dyn std::error::Error>> {
11940 let f = self.func("add_f32");
11941 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11943 let ni = n as i32;
11944 let __s_bld = self.gpu.stream();
11945 let mut bld = __s_bld.launch_builder(&f);
11946 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
11947 unsafe {
11948 bld.launch(cfg)?;
11949 }
11950 Ok(())
11951 }
11952
11953 pub fn mul(
11954 &self,
11955 a: &CudaSlice<f32>,
11956 b_in: &CudaSlice<f32>,
11957 dst: &mut CudaSlice<f32>,
11958 n: usize,
11959 ) -> Result<(), Box<dyn std::error::Error>> {
11960 let f = self.func("mul_f32");
11961 let cfg = LaunchConfig::for_num_elems(n as u32);
11962 let ni = n as i32;
11963 let __s_bld = self.gpu.stream();
11964 let mut bld = __s_bld.launch_builder(&f);
11965 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
11966 unsafe {
11967 bld.launch(cfg)?;
11968 }
11969 Ok(())
11970 }
11971
11972 pub fn matmul(
11975 &self,
11976 w: &crate::model::GpuTensor,
11977 x: &CudaSlice<f32>,
11978 m: usize,
11979 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11980 use crate::model::GpuTensor;
11981 let in_f = w.in_features();
11982 let out_f = w.out_features();
11983 #[allow(non_snake_case)]
11991 let GEMM_M_THRESHOLD = if self.verify_exact_on() {
11994 usize::MAX
11995 } else {
11996 16usize
11997 };
11998
11999 const GEMM_MIN_OUT_F: usize = 128; if m >= GEMM_M_THRESHOLD {
12024 if let Some(y) = self.try_fp8_gemm(w, x, m)? {
12025 return Ok(y);
12026 }
12027 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? {
12034 return Ok(y);
12035 }
12036 if let Some(y) = self.try_f16_gemm(w, x, m)? {
12039 return Ok(y);
12040 }
12041 }
12042 if let GpuTensor::Quant { qtype, .. } = w {
12057 if *qtype == QT_F8_E4M3_BLK {
12058 if m >= GEMM_M_THRESHOLD {
12059 if let Some(y) = self.try_e4m3_blk_prefill(w, x, m)? {
12060 return Ok(y);
12061 }
12062 }
12063 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
12064 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? {
12065 return Ok(y);
12066 }
12067 }
12068 }
12069 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.mmq_supports(w) {
12070 return self.qmatvec_mmq(w, x, m);
12071 }
12072 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.gemm_supports(w) {
12073 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
12074 return self.qmatvec_gemm(w, &aq, &ad, m);
12075 }
12076 if m >= GEMM_M_THRESHOLD {
12079 if let Some(y) = self.try_fp4_gemm(w, x, m, in_f, out_f)? {
12080 return Ok(y);
12081 }
12082 }
12083 let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
12087 if m == 1 && fast {
12092 if let GpuTensor::Quant {
12093 bytes,
12094 qtype,
12095 row_bytes,
12096 rp,
12097 rp4,
12098 scale,
12099 ..
12100 } = w
12101 {
12102 if self.mmvq_supports(*qtype) {
12103 let (bytes, rp) = match rp4 {
12107 Some(m4) => (m4, true),
12108 None => (bytes, *rp),
12109 };
12110 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
12111 return self.qmatvec_mmvq(
12112 bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, rp,
12113 );
12114 }
12115 }
12116 }
12117 if (2..=16).contains(&m)
12133 && fast
12134 && std::env::var("MEMRA_NO_BATCHED").is_err()
12135 && (m <= 4 || Self::b8_enabled())
12136 {
12137 let m_ok = m <= 8
12147 || matches!(w, GpuTensor::Quant { qtype, .. }
12148 if *qtype == QT_Q4_0 || *qtype == QT_Q6_K || *qtype == QT_F8_E4M3
12149 || *qtype == QT_NVFP4 || *qtype == QT_Q4_K || *qtype == QT_Q5_K || *qtype == QT_Q8_0);
12150 if m_ok {
12151 if let GpuTensor::Quant {
12152 bytes,
12153 qtype,
12154 row_bytes,
12155 rp,
12156 rp4,
12157 ..
12158 } = w
12159 {
12160 if self.batched_supports(*qtype) && self.mmvq_supports(*qtype) {
12161 let (bytes, rp) = match rp4 {
12162 Some(m4) => (m4, true),
12163 None => (bytes, *rp),
12164 };
12165 let mcols = Self::batched_mcols(m);
12166 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
12167 let mut y = self.qmatvec_mmvq_batched(
12168 bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, mcols, 1.0, rp,
12169 )?;
12170 if let GpuTensor::Quant { scale, .. } = w {
12171 if *scale != 1.0 {
12172 self.scale_inplace(&mut y, *scale, m * out_f)?;
12173 }
12174 }
12175 return Ok(y);
12176 }
12177 }
12178 }
12179 }
12180 if fast {
12186 if let GpuTensor::Quant {
12187 bytes,
12188 qtype,
12189 row_bytes,
12190 scale,
12191 ..
12192 } = w
12193 {
12194 if *qtype == QT_F8_E4M3 {
12195 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
12196 return self.qmatvec_mmvq(
12197 bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, false,
12198 );
12199 }
12200 }
12201 }
12202 let mut y = match w {
12203 GpuTensor::Quant {
12204 bytes,
12205 qtype,
12206 row_bytes,
12207 ..
12208 } if fast && *qtype == QT_Q8_0 => {
12209 self.qmatvec_q8_0_fast(bytes, x, m, in_f, out_f, *row_bytes)?
12210 }
12211 GpuTensor::Quant {
12212 bytes,
12213 qtype,
12214 row_bytes,
12215 ..
12216 } if fast && *qtype == QT_Q4_K => {
12217 self.qmatvec_q4_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
12218 }
12219 GpuTensor::Quant {
12220 bytes,
12221 qtype,
12222 row_bytes,
12223 ..
12224 } if fast && *qtype == QT_Q6_K => {
12225 self.qmatvec_q6_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
12226 }
12227 GpuTensor::Quant {
12228 bytes,
12229 qtype,
12230 row_bytes,
12231 ..
12232 } if fast && *qtype == QT_Q5_K => {
12233 self.qmatvec_q5_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
12234 }
12235 GpuTensor::Quant {
12236 bytes,
12237 qtype,
12238 row_bytes,
12239 ..
12240 } if fast && *qtype == QT_Q3_K => {
12241 self.qmatvec_q3_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
12242 }
12243 GpuTensor::Quant {
12244 bytes,
12245 qtype,
12246 row_bytes,
12247 rp,
12248 ..
12249 } if fast && *qtype == QT_NVFP4 => self.qmatvec_dp4a_named(
12250 if *rp {
12251 "qmatvec_nvfp4_dp4a_rp"
12252 } else {
12253 "qmatvec_nvfp4_dp4a"
12254 },
12255 &bytes.slice(0..bytes.len()),
12256 x,
12257 m,
12258 in_f,
12259 out_f,
12260 *row_bytes,
12261 )?,
12262 GpuTensor::Quant {
12266 bytes,
12267 qtype,
12268 row_bytes,
12269 ..
12270 } if fast && *qtype == QT_IQ4_XS && Self::iq_fast_enabled() => {
12271 self.qmatvec_iq4_XS_fast(bytes, x, m, in_f, out_f, *row_bytes)?
12272 }
12273 GpuTensor::Quant {
12278 bytes,
12279 qtype,
12280 row_bytes,
12281 rp,
12282 ..
12283 } =>
12284 {
12287 self.qmatvec(
12288 bytes,
12289 x,
12290 m,
12291 in_f,
12292 out_f,
12293 if *rp && *qtype == QT_NVFP4 {
12294 QT_NVFP4_RP
12295 } else {
12296 *qtype
12297 },
12298 *row_bytes,
12299 )?
12300 }
12301 GpuTensor::Float { data, .. } => self.linear(x, data, m, in_f, out_f)?,
12302 GpuTensor::FloatBf16 { data, .. } => {
12305 if (1..=32).contains(&m) && Self::bf16_mmv_on() && in_f % 8 == 0 {
12312 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
12313 self.matvec_bf16_rows_into(data, x, &mut y, in_f, out_f, m)?;
12314 y
12315 } else {
12316 self.linear_bf16_chunked(x, data, m, in_f, out_f, false, None)?
12317 }
12318 }
12319 };
12320 if let GpuTensor::Quant { scale, .. } = w {
12322 if *scale != 1.0 {
12323 self.scale_inplace(&mut y, *scale, m * out_f)?;
12324 }
12325 }
12326 Ok(y)
12327 }
12328
12329 pub fn stage_a_raw_needed() -> bool {
12339 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12340 *ON.get_or_init(|| std::env::var("MEMRA_FAST").as_deref() == Ok("0"))
12341 }
12342
12343 pub fn uses_q8_1_fast(&self, w: &crate::model::GpuTensor) -> bool {
12346 use crate::model::GpuTensor;
12347 if std::env::var("MEMRA_FAST").as_deref() == Ok("0") {
12348 return false;
12349 }
12350 match w {
12351 GpuTensor::Quant { qtype, .. } => {
12358 matches!(
12359 *qtype,
12360 QT_Q8_0
12361 | QT_Q4_K
12362 | QT_Q6_K
12363 | QT_Q5_K
12364 | QT_Q3_K
12365 | QT_NVFP4
12366 | QT_F8_E4M3
12367 | QT_F8_E4M3_BLK
12368 | QT_Q4_0
12369 ) || (*qtype == QT_IQ4_XS && Self::iq_fast_enabled())
12370 }
12371 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
12372 }
12373 }
12374
12375 pub fn matmul_pre(
12380 &self,
12381 w: &crate::model::GpuTensor,
12382 aq: &CudaSlice<i8>,
12383 ad: &CudaSlice<f32>,
12384 x_fallback: &CudaSlice<f32>,
12385 m: usize,
12386 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12387 use crate::model::GpuTensor;
12388 let x_raw_ok = x_fallback.len() >= m * w.in_features();
12394 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
12397 if let Some(y) = self.try_fp8_gemm(w, x_fallback, m)? {
12398 return Ok(y);
12399 }
12400 if let Some(y) = self.try_fp8_blk_mmq(w, x_fallback, m)? {
12403 return Ok(y);
12404 }
12405 if let Some(y) = self.try_f16_gemm(w, x_fallback, m)? {
12407 return Ok(y);
12408 }
12409 }
12410 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
12416 if let Some(y) = self.try_e4m3_blk_prefill(w, x_fallback, m)? {
12417 return Ok(y);
12418 }
12419 }
12420 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
12421 return Ok(y);
12422 }
12423 if m >= 16
12428 && w.out_features() >= 128
12429 && self.mmq_supports(w)
12430 && !self.verify_exact_on()
12431 && x_raw_ok
12432 {
12433 return self.qmatvec_mmq(w, x_fallback, m);
12434 }
12435 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
12438 if let Some(y) =
12439 self.try_fp4_gemm(w, x_fallback, m, w.in_features(), w.out_features())?
12440 {
12441 return Ok(y);
12442 }
12443 }
12444 if m >= 16 && self.gemm_supports(w) && !self.verify_exact_on() {
12447 return self.qmatvec_gemm(w, aq, ad, m);
12448 }
12449 if !self.uses_q8_1_fast(w) {
12468 if !x_raw_ok {
12469 return Err(format!(
12470 "matmul_pre: q8_1-fast is off for this weight but x_fallback holds {} f32 \
12471 (need m*in_f = {}*{} = {}). This call site pre-quantized its activation and \
12472 dropped the f32, so there is nothing to fall back to — pass the real f32 \
12473 activation (see Engine::rms_norm_decode, which is bit-identical to \
12474 rms_norm_q8_1's reduction) or keep the weight on the q8_1 path.",
12475 x_fallback.len(),
12476 m,
12477 w.in_features(),
12478 m * w.in_features()
12479 )
12480 .into());
12481 }
12482 return self.matmul(w, x_fallback, m);
12483 }
12484 let in_f = w.in_features();
12485 let out_f = w.out_features();
12486 let (bytes, qtype, row_bytes, scale, rp) = match w {
12487 GpuTensor::Quant {
12488 bytes,
12489 qtype,
12490 row_bytes,
12491 scale,
12492 rp,
12493 ..
12494 } => (bytes, *qtype, *row_bytes, *scale, *rp),
12495 _ => unreachable!("uses_q8_1_fast guaranteed Quant"),
12496 };
12497 let (mbytes, mrp) = match w {
12500 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
12501 _ => (bytes, rp),
12502 };
12503 if m == 1 && self.mmvq_supports(qtype) {
12507 return self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, mrp);
12508 }
12509 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
12522 && std::env::var("MEMRA_NO_BATCHED").is_err()
12523 && (m <= 4 || Self::b8_enabled())
12524 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_NVFP4
12528 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_F8_E4M3 || qtype == QT_Q8_0)
12529 {
12530 let mcols = Self::batched_mcols(m);
12531 return self.qmatvec_mmvq_batched(
12532 mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, mrp,
12533 );
12534 }
12535 if qtype == QT_F8_E4M3 || qtype == QT_Q4_0 {
12541 let (b2, r2) = if qtype == QT_Q4_0 {
12542 (mbytes, mrp)
12543 } else {
12544 (bytes, rp)
12545 };
12546 return self.qmatvec_mmvq(b2, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, r2);
12547 }
12548 let name = match qtype {
12549 QT_Q8_0 => "qmatvec_q8_0_dp4a",
12550 QT_Q4_K => "qmatvec_q4_K_dp4a",
12551 QT_Q6_K => "qmatvec_q6_K_dp4a",
12552 QT_Q5_K => "qmatvec_q5_K_dp4a",
12553 QT_Q3_K => "qmatvec_q3_K_dp4a",
12554 QT_NVFP4 => {
12555 if rp {
12556 "qmatvec_nvfp4_dp4a_rp"
12557 } else {
12558 "qmatvec_nvfp4_dp4a"
12559 }
12560 }
12561 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
12562 _ => unreachable!(),
12563 };
12564 let f = self.func(name);
12565 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
12567 grid_dim: (out_f as u32, m as u32, 1),
12568 block_dim: (128, 1, 1),
12569 shared_mem_bytes: 0,
12570 };
12571 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
12572 let __s_b = self.gpu.stream();
12573 let mut b = __s_b.launch_builder(&f);
12574 b.arg(bytes)
12575 .arg(aq)
12576 .arg(ad)
12577 .arg(&mut y)
12578 .arg(&inf)
12579 .arg(&outf)
12580 .arg(&mi)
12581 .arg(&rb);
12582 unsafe {
12583 b.launch(cfg)?;
12584 }
12585 if scale != 1.0 {
12586 self.scale_inplace(&mut y, scale, m * out_f)?;
12587 }
12588 Ok(y)
12589 }
12590
12591 pub fn matmul_decode_exact(
12599 &self,
12600 w: &crate::model::GpuTensor,
12601 x: &CudaSlice<f32>,
12602 m: usize,
12603 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12604 use crate::model::GpuTensor;
12605 if let GpuTensor::Float { data, .. } = w {
12613 return self.linear_decode_exact(x, data, m, w.in_features(), w.out_features());
12614 }
12615 if let GpuTensor::FloatBf16 { data, .. } = w {
12618 let (in_f, out_f) = (w.in_features(), w.out_features());
12619 if (1..=32).contains(&m) && Self::bf16_mmv_on() && in_f % 8 == 0 {
12622 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
12623 self.matvec_bf16_rows_into(data, x, &mut y, in_f, out_f, m)?;
12624 return Ok(y);
12625 }
12626 return self.linear_bf16_chunked(x, data, m, in_f, out_f, true, None);
12627 }
12628 if !self.uses_q8_1_fast(w) {
12629 return self.matmul(w, x, m);
12630 }
12631 let in_f = w.in_features();
12632 let out_f = w.out_features();
12633 let (bytes, qtype, row_bytes, scale, rp) = match w {
12634 GpuTensor::Quant {
12635 bytes,
12636 qtype,
12637 row_bytes,
12638 scale,
12639 rp,
12640 ..
12641 } => (bytes, *qtype, *row_bytes, *scale, *rp),
12642 _ => return self.matmul(w, x, m),
12643 };
12644 let (bytes, rp) = match w {
12647 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
12648 _ => (bytes, rp),
12649 };
12650 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
12651 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? {
12655 return Ok(y);
12656 }
12657 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
12666 && std::env::var("MEMRA_NO_BATCHED").is_err()
12667 && (m <= 4 || Self::b8_enabled())
12668 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
12671 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0)
12672 {
12673 let mcols = Self::batched_mcols(m);
12674 return self.qmatvec_mmvq_batched(
12675 bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp,
12676 );
12677 }
12678 if self.mmvq_supports(qtype) {
12679 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
12682 }
12683 self.matmul_pre(w, &aq, &ad, x, m)
12686 }
12687
12688 pub fn matmul_decode_exact_pre(
12698 &self,
12699 w: &crate::model::GpuTensor,
12700 aq: &CudaSlice<i8>,
12701 ad: &CudaSlice<f32>,
12702 m: usize,
12703 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12704 use crate::model::GpuTensor;
12705 debug_assert!(
12706 self.uses_q8_1_fast(w),
12707 "matmul_decode_exact_pre: caller must guarantee q8_1-fast"
12708 );
12709 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
12711 return Ok(y);
12712 }
12713 let in_f = w.in_features();
12714 let out_f = w.out_features();
12715 let (bytes, qtype, row_bytes, scale, rp) = match w {
12716 GpuTensor::Quant {
12717 bytes,
12718 qtype,
12719 row_bytes,
12720 scale,
12721 rp,
12722 ..
12723 } => (bytes, *qtype, *row_bytes, *scale, *rp),
12724 _ => {
12725 return Err(
12726 "matmul_decode_exact_pre: Quant tensor required (q8_1-fast contract)".into(),
12727 );
12728 }
12729 };
12730 let (bytes, rp) = match w {
12732 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
12733 _ => (bytes, rp),
12734 };
12735 if (2..=16).contains(&m)
12737 && self.batched_supports(qtype)
12738 && self.mmvq_supports(qtype)
12739 && std::env::var("MEMRA_NO_BATCHED").is_err()
12740 && (m <= 4 || Self::b8_enabled())
12741 && (m <= 8
12742 || qtype == QT_Q4_0
12743 || qtype == QT_Q6_K
12744 || qtype == QT_F8_E4M3
12745 || qtype == QT_NVFP4
12746 || qtype == QT_Q4_K
12747 || qtype == QT_Q5_K
12748 || qtype == QT_Q8_0)
12749 {
12750 let mcols = Self::batched_mcols(m);
12751 return self.qmatvec_mmvq_batched(
12752 bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp,
12753 );
12754 }
12755 if self.mmvq_supports(qtype) {
12756 return self.qmatvec_mmvq(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
12757 }
12758 let x0 = self.zeros(0)?;
12761 self.matmul_pre(w, aq, ad, &x0, m)
12762 }
12763
12764 pub fn matmul_decode_exact_dual_pre(
12773 &self,
12774 w0: &crate::model::GpuTensor,
12775 w1: &crate::model::GpuTensor,
12776 aq: &CudaSlice<i8>,
12777 ad: &CudaSlice<f32>,
12778 m: usize,
12779 ) -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>>
12780 {
12781 use crate::model::GpuTensor;
12782 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12783 let on = *ON.get_or_init(|| {
12784 std::env::var("MEMRA_SPEC_DUAL_T")
12785 .map(|v| v != "0")
12786 .unwrap_or(true)
12787 });
12788 if !on
12789 || !(2..=7).contains(&m)
12790 || std::env::var("MEMRA_NO_BATCHED").is_ok()
12791 || !self.uses_q8_1_fast(w0)
12792 || !self.uses_q8_1_fast(w1)
12793 {
12794 return Ok(None);
12795 }
12796 if !self.mmvq_supports(QT_NVFP4) {
12801 return Ok(None);
12802 }
12803 let (in_f, out_f) = (w0.in_features(), w0.out_features());
12804 if w1.in_features() != in_f || w1.out_features() != out_f {
12805 return Ok(None);
12806 }
12807 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
12808 (
12809 GpuTensor::Quant {
12810 bytes: b0,
12811 qtype: q0,
12812 row_bytes: rb0,
12813 scale: s0,
12814 rp: rp0,
12815 rp4: None,
12816 ..
12817 },
12818 GpuTensor::Quant {
12819 bytes: b1,
12820 qtype: q1,
12821 row_bytes: rb1,
12822 scale: s1,
12823 rp: rp1,
12824 rp4: None,
12825 ..
12826 },
12827 ) if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 => {
12828 (b0, b1, *rb0, *s0, *s1, *rp0)
12829 }
12830 _ => return Ok(None),
12831 };
12832 if m > 4 && !(rp && Self::b8_enabled() && std::env::var("MEMRA_B567").as_deref() != Ok("0"))
12835 {
12836 return Ok(None);
12837 }
12838 let (y0, y1) =
12839 self.qmatvec_batched_dual_raw(b0, b1, aq, ad, m, in_f, out_f, row_bytes, rp)?;
12840 Ok(Some(((y0, s0), (y1, s1))))
12841 }
12842
12843 pub fn matmul_decode_exact_group4_pre(
12855 &self,
12856 ws: [&crate::model::GpuTensor; 4],
12857 aq: &CudaSlice<i8>,
12858 ad: &CudaSlice<f32>,
12859 m: usize,
12860 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
12861 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12862 let on = *ON.get_or_init(|| {
12863 std::env::var("MEMRA_TK_GDN_GROUP")
12864 .map(|v| v != "0")
12865 .unwrap_or(true)
12866 });
12867 self.matmul_decode_exact_group_pre(&ws, aq, ad, m, on, "GDN group4")
12868 }
12869
12870 pub fn matmul_decode_exact_group3_pre(
12875 &self,
12876 ws: [&crate::model::GpuTensor; 3],
12877 aq: &CudaSlice<i8>,
12878 ad: &CudaSlice<f32>,
12879 m: usize,
12880 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
12881 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12882 let on = *ON.get_or_init(|| {
12883 std::env::var("MEMRA_TK_FA_GROUP")
12884 .map(|v| v != "0")
12885 .unwrap_or(true)
12886 });
12887 self.matmul_decode_exact_group_pre(&ws, aq, ad, m, on, "FA group3")
12888 }
12889
12890 fn matmul_decode_exact_group_pre(
12894 &self,
12895 ws: &[&crate::model::GpuTensor],
12896 aq: &CudaSlice<i8>,
12897 ad: &CudaSlice<f32>,
12898 m: usize,
12899 on: bool,
12900 tag: &'static str,
12901 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
12902 use crate::model::GpuTensor;
12903 if !on
12904 || !(2..=16).contains(&m)
12905 || std::env::var("MEMRA_NO_BATCHED").is_ok()
12906 || (m > 4 && !Self::b8_enabled())
12907 || !self.mmvq_supports(QT_NVFP4)
12908 || !self.batched_supports(QT_NVFP4)
12909 {
12910 return Ok(None);
12911 }
12912 let in_f = ws[0].in_features();
12913 let mut parts: Vec<(&CudaSlice<u8>, usize, f32)> = Vec::with_capacity(4);
12914 for w in ws {
12915 if !self.uses_q8_1_fast(w) || w.in_features() != in_f {
12916 return Ok(None);
12917 }
12918 match w {
12919 GpuTensor::Quant {
12920 bytes,
12921 qtype,
12922 scale,
12923 rp: true,
12924 rp4: None,
12925 ..
12926 } if *qtype == QT_NVFP4 && w.out_features() % 8 == 0 => {
12927 parts.push((bytes, w.out_features(), *scale));
12928 }
12929 _ => return Ok(None),
12930 }
12931 }
12932 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12934 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
12935 let mcols = if (5..=7).contains(&m) && b567 {
12936 m
12937 } else {
12938 Self::batched_mcols(m)
12939 };
12940 let kname: &'static str = match mcols {
12941 2 => "qmatvec_nvfp4_mmvq_group4_b2_rp",
12942 4 => "qmatvec_nvfp4_mmvq_group4_b4_rp",
12943 5 => "qmatvec_nvfp4_mmvq_group4_b5_rp",
12944 6 => "qmatvec_nvfp4_mmvq_group4_b6_rp",
12945 7 => "qmatvec_nvfp4_mmvq_group4_b7_rp",
12946 8 => "qmatvec_nvfp4_mmvq_group4_b8_rp",
12947 16 => "qmatvec_nvfp4_mmvq_group4_b16_rp",
12948 _ => return Ok(None),
12949 };
12950 if std::env::var("MEMRA_DEBUG").is_ok() {
12953 use std::sync::Mutex;
12954 static SEEN: Mutex<Vec<&'static str>> = Mutex::new(Vec::new());
12955 let mut seen = SEEN.lock().unwrap();
12956 if !seen.contains(&tag) {
12957 seen.push(tag);
12958 eprintln!("[memra] {tag} batched ENGAGED (m={m})");
12959 }
12960 }
12961 const ROWS_PER_BLOCK: u32 = 4; let rows_per_block = ROWS_PER_BLOCK * 2; let total: usize = parts.iter().map(|p| p.1).sum();
12964 let three = parts.len() == 3;
12965 let mut y0 = self.alloc_uninit::<f32>(m * parts[0].1)?;
12966 let mut y1 = self.alloc_uninit::<f32>(m * parts[1].1)?;
12967 let mut y2 = self.alloc_uninit::<f32>(m * parts[2].1)?;
12968 let mut y3 = self.alloc_uninit::<f32>(if three { 1 } else { m * parts[3].1 })?;
12971 let cfg = LaunchConfig {
12972 grid_dim: ((total as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
12973 block_dim: (32, ROWS_PER_BLOCK, 1),
12974 shared_mem_bytes: 0,
12975 };
12976 let (inf, mi) = (in_f as i32, m as i32);
12977 let (n0, n1, n2) = (parts[0].1 as i32, parts[1].1 as i32, parts[2].1 as i32);
12978 let n3 = if three { 0i32 } else { parts[3].1 as i32 };
12979 let (s0, s1, s2) = (parts[0].2, parts[1].2, parts[2].2);
12980 let s3 = if three { 1.0f32 } else { parts[3].2 };
12981 let w3 = if three { parts[0].0 } else { parts[3].0 };
12982 let f = self.func(kname);
12983 let __s_b = self.gpu.stream();
12984 let mut b = __s_b.launch_builder(&f);
12985 b.arg(parts[0].0)
12986 .arg(parts[1].0)
12987 .arg(parts[2].0)
12988 .arg(w3)
12989 .arg(aq)
12990 .arg(ad)
12991 .arg(&mut y0)
12992 .arg(&mut y1)
12993 .arg(&mut y2)
12994 .arg(&mut y3)
12995 .arg(&inf)
12996 .arg(&n0)
12997 .arg(&n1)
12998 .arg(&n2)
12999 .arg(&n3)
13000 .arg(&mi)
13001 .arg(&s0)
13002 .arg(&s1)
13003 .arg(&s2)
13004 .arg(&s3);
13005 unsafe {
13006 b.launch(cfg)?;
13007 }
13008 Ok(Some(if three {
13009 vec![y0, y1, y2]
13010 } else {
13011 vec![y0, y1, y2, y3]
13012 }))
13013 }
13014
13015 pub fn matmul_decode_exact_dual(
13031 &self,
13032 w0: &crate::model::GpuTensor,
13033 w1: &crate::model::GpuTensor,
13034 x: &CudaSlice<f32>,
13035 m: usize,
13036 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
13037 use crate::model::GpuTensor;
13038 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13039 let on = *ON.get_or_init(|| {
13040 std::env::var("MEMRA_SPEC_DUAL_T")
13041 .map(|v| v != "0")
13042 .unwrap_or(true)
13043 });
13044 if !on
13045 || !(2..=4).contains(&m)
13046 || std::env::var("MEMRA_NO_BATCHED").is_ok()
13047 || !self.uses_q8_1_fast(w0)
13048 || !self.uses_q8_1_fast(w1)
13049 {
13050 return Ok(None);
13051 }
13052 if !self.mmvq_supports(QT_NVFP4) {
13057 return Ok(None);
13058 }
13059 let (in_f, out_f) = (w0.in_features(), w0.out_features());
13060 if w1.in_features() != in_f || w1.out_features() != out_f {
13061 return Ok(None);
13062 }
13063 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
13064 (
13065 GpuTensor::Quant {
13066 bytes: b0,
13067 qtype: q0,
13068 row_bytes: rb0,
13069 scale: s0,
13070 rp: rp0,
13071 rp4: None,
13072 ..
13073 },
13074 GpuTensor::Quant {
13075 bytes: b1,
13076 qtype: q1,
13077 row_bytes: rb1,
13078 scale: s1,
13079 rp: rp1,
13080 rp4: None,
13081 ..
13082 },
13083 ) if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 => {
13084 (b0, b1, *rb0, *s0, *s1, *rp0)
13085 }
13086 _ => return Ok(None),
13087 };
13088 if std::env::var("MEMRA_DEBUG").is_ok() {
13091 static ONCE: std::sync::Once = std::sync::Once::new();
13092 ONCE.call_once(|| eprintln!("[memra] dual gate+up batched ENGAGED (m={m} rp={rp})"));
13093 }
13094 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
13095 let (y0, y1) =
13096 self.qmatvec_batched_dual_raw(b0, b1, &aq, &ad, m, in_f, out_f, row_bytes, rp)?;
13097 let mut y0 = y0;
13098 let mut y1 = y1;
13099 if s0 != 1.0 {
13100 self.scale_inplace(&mut y0, s0, m * out_f)?;
13101 }
13102 if s1 != 1.0 {
13103 self.scale_inplace(&mut y1, s1, m * out_f)?;
13104 }
13105 Ok(Some((y0, y1)))
13106 }
13107
13108 #[allow(clippy::too_many_arguments)]
13113 pub fn qmatvec_batched_dual_raw(
13114 &self,
13115 b0: &CudaSlice<u8>,
13116 b1: &CudaSlice<u8>,
13117 aq: &CudaSlice<i8>,
13118 ad: &CudaSlice<f32>,
13119 m: usize,
13120 in_f: usize,
13121 out_f: usize,
13122 row_bytes: usize,
13123 rp: bool,
13124 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13125 const ROWS_PER_BLOCK: u32 = 4;
13126 let mcols = Self::batched_mcols(m);
13127 let tiny_rp1 = rp
13130 && mcols == 4
13131 && out_f <= 128
13132 && std::env::var("MEMRA_NVFP4_AUX_DUAL").as_deref() != Ok("0");
13133 let (name, rows_per_block) = if tiny_rp1 {
13134 ("qmatvec_nvfp4_mmvq_dual_b4_rp", ROWS_PER_BLOCK)
13135 } else {
13136 match (mcols, rp, m) {
13137 (2, false, _) => ("qmatvec_nvfp4_mmvq_dual_b2", ROWS_PER_BLOCK),
13138 (4, false, _) => ("qmatvec_nvfp4_mmvq_dual_b4_r2", ROWS_PER_BLOCK * 2),
13139 (2, true, _) => ("qmatvec_nvfp4_mmvq_dual_b2_rp", ROWS_PER_BLOCK),
13140 (4, true, _) => ("qmatvec_nvfp4_mmvq_dual_b4_rpr2", ROWS_PER_BLOCK * 2),
13141 (8, true, 5) => ("qmatvec_nvfp4_mmvq_dual_b5_rpr2", ROWS_PER_BLOCK * 2),
13142 (8, true, 6) => ("qmatvec_nvfp4_mmvq_dual_b6_rpr2", ROWS_PER_BLOCK * 2),
13143 (8, true, 7) => ("qmatvec_nvfp4_mmvq_dual_b7_rpr2", ROWS_PER_BLOCK * 2),
13144 _ => {
13145 return Err(
13146 format!("qmatvec_batched_dual_raw: no dual kernel for m {m}").into(),
13147 );
13148 }
13149 }
13150 };
13151 let f = self.func(name);
13152 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
13153 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
13154 let cfg = LaunchConfig {
13155 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
13156 block_dim: (32, ROWS_PER_BLOCK, 1),
13157 shared_mem_bytes: 0,
13158 };
13159 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
13160 let __s_b = self.gpu.stream();
13161 let mut b = __s_b.launch_builder(&f);
13162 b.arg(b0)
13163 .arg(b1)
13164 .arg(aq)
13165 .arg(ad)
13166 .arg(&mut y0)
13167 .arg(&mut y1)
13168 .arg(&inf)
13169 .arg(&outf)
13170 .arg(&mi)
13171 .arg(&rb);
13172 unsafe {
13173 b.launch(cfg)?;
13174 }
13175 Ok((y0, y1))
13176 }
13177
13178 pub fn matmul_pre_dual_noscale(
13190 &self,
13191 w0: &crate::model::GpuTensor,
13192 w1: &crate::model::GpuTensor,
13193 aq: &CudaSlice<i8>,
13194 ad: &CudaSlice<f32>,
13195 m: usize,
13196 ) -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>>
13197 {
13198 use crate::model::GpuTensor;
13199 if m != 1 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
13200 return Ok(None);
13201 }
13202 if !self.mmvq_supports(QT_NVFP4) {
13212 return Ok(None);
13213 }
13214 let (in_f, out_f) = (w0.in_features(), w0.out_features());
13215 if w1.in_features() != in_f || w1.out_features() != out_f {
13216 return Ok(None);
13217 }
13218 let no_mirror =
13231 |w: &crate::model::GpuTensor| !matches!(w, GpuTensor::Quant { rp4: Some(_), .. });
13232 if self.q8_ffn_fuse2_on()
13233 && no_mirror(w0)
13234 && no_mirror(w1)
13235 && let Some([p0, p1]) = self.q8_fused_params(&[w0, w1])
13236 {
13237 let (y0, y1) = self.q8_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2)?;
13238 return Ok(Some(((y0, 1.0), (y1, 1.0))));
13239 }
13240 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
13250 let (y0, y1) =
13251 self.e4m3_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2, 1.0, 1.0)?;
13252 return Ok(Some(((y0, p0.3), (y1, p1.3))));
13253 }
13254 let (b0, q0, rb0, s0, rp0) = match w0 {
13255 GpuTensor::Quant {
13256 bytes,
13257 qtype,
13258 row_bytes,
13259 scale,
13260 rp,
13261 ..
13262 } => (bytes, *qtype, *row_bytes, *scale, *rp),
13263 _ => return Ok(None),
13264 };
13265 let (b1, q1, rb1, s1, rp1) = match w1 {
13266 GpuTensor::Quant {
13267 bytes,
13268 qtype,
13269 row_bytes,
13270 scale,
13271 rp,
13272 ..
13273 } => (bytes, *qtype, *row_bytes, *scale, *rp),
13274 _ => return Ok(None),
13275 };
13276 if q0 != QT_NVFP4 || q1 != QT_NVFP4 || rb0 != rb1 || rp0 != rp1 {
13277 return Ok(None);
13278 }
13279 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
13281 let rows_per_block = ROWS_PER_BLOCK * RPW;
13282 let f = self.func(if rp0 {
13283 "qmatvec_nvfp4_mmvq_dual_mr2_rp"
13284 } else {
13285 "qmatvec_nvfp4_mmvq_dual_mr2"
13286 });
13287 let mut y0 = self.alloc_uninit::<f32>(out_f)?;
13288 let mut y1 = self.alloc_uninit::<f32>(out_f)?;
13289 let cfg = LaunchConfig {
13290 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
13291 block_dim: (32, ROWS_PER_BLOCK, 1),
13292 shared_mem_bytes: 0,
13293 };
13294 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, rb0 as i64);
13295 let one = 1.0f32;
13298 let __s_b = self.gpu.stream();
13299 let mut b = __s_b.launch_builder(&f);
13300 b.arg(b0)
13301 .arg(b1)
13302 .arg(aq)
13303 .arg(ad)
13304 .arg(&mut y0)
13305 .arg(&mut y1)
13306 .arg(&inf)
13307 .arg(&outf)
13308 .arg(&mi)
13309 .arg(&rb)
13310 .arg(&one)
13311 .arg(&one);
13312 unsafe {
13313 b.launch(cfg)?;
13314 }
13315 Ok(Some(((y0, s0), (y1, s1))))
13316 }
13317
13318 #[allow(clippy::too_many_arguments)]
13326 pub fn matmul_nvfp4_fused3(
13327 &self,
13328 w0: &crate::model::GpuTensor,
13329 w1: &crate::model::GpuTensor,
13330 w2: &crate::model::GpuTensor,
13331 aq: &CudaSlice<i8>,
13332 ad: &CudaSlice<f32>,
13333 m: usize,
13334 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
13335 {
13336 use crate::model::GpuTensor;
13337 if !self.mmvq_supports(QT_NVFP4)
13344 || !self.uses_q8_1_fast(w0)
13345 || !self.uses_q8_1_fast(w1)
13346 || !self.uses_q8_1_fast(w2)
13347 {
13348 return Ok(None);
13349 }
13350 if (9..=16).contains(&m) {
13353 return Ok(
13354 match self.matmul_decode_exact_group3_pre([w0, w1, w2], aq, ad, m)? {
13355 Some(mut ys) => {
13356 let y2 = ys.pop().unwrap();
13357 let y1 = ys.pop().unwrap();
13358 let y0 = ys.pop().unwrap();
13359 Some((y0, y1, y2))
13360 }
13361 None => None,
13362 },
13363 );
13364 }
13365 if !(1..=8).contains(&m) {
13366 return Ok(None);
13367 }
13368 if m > 1 {
13369 let in_f = w0.in_features();
13370 if std::env::var("MEMRA_NVFP4_FUSED3B").as_deref() == Ok("0")
13371 || !self.batched_supports(QT_NVFP4)
13372 || std::env::var("MEMRA_NO_BATCHED").is_ok()
13373 || (m > 4 && !Self::b8_enabled())
13374 || in_f % 512 != 0
13375 || in_f / 64 > 272
13376 {
13377 return Ok(None);
13378 }
13379 }
13380 let unpack = |w: &crate::model::GpuTensor| match w {
13381 GpuTensor::Quant {
13382 bytes,
13383 qtype,
13384 scale,
13385 rp,
13386 ..
13387 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
13388 _ => None,
13389 };
13390 let (Some(p0), Some(p1), Some(p2)) = (unpack(w0), unpack(w1), unpack(w2)) else {
13391 return Ok(None);
13392 };
13393 let in_f = w0.in_features();
13394 if w1.in_features() != in_f || w2.in_features() != in_f {
13395 return Ok(None);
13396 }
13397 let (o0, o1, o2) = (w0.out_features(), w1.out_features(), w2.out_features());
13398 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
13400 let rows_pb = ROWS_PER_BLOCK * RPW;
13401 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
13402 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
13403 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
13404 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
13405 let (inf, oi0, oi1, oi2, mi) = (in_f as i32, o0 as i32, o1 as i32, o2 as i32, m as i32);
13406 let (b0, b1, b2) = unsafe { (&*p0.0, &*p1.0, &*p2.0) };
13409 if m > 1 {
13410 if p0.1 != 1.0 || p1.1 != 1.0 || p2.1 != 1.0 {
13412 return Ok(None);
13413 }
13414 let f = self.func("qmatvec_nvfp4_mmvq_fused3_b8_rpsc");
13415 let cfg = LaunchConfig {
13416 grid_dim: (nb(o0) + nb(o1) + nb(o2), 1, 1),
13417 block_dim: (32, ROWS_PER_BLOCK, 1),
13418 shared_mem_bytes: 0,
13419 };
13420 let __s_b = self.gpu.stream();
13421 let mut b = __s_b.launch_builder(&f);
13422 b.arg(b0)
13423 .arg(b1)
13424 .arg(b2)
13425 .arg(aq)
13426 .arg(ad)
13427 .arg(&mut y0)
13428 .arg(&mut y1)
13429 .arg(&mut y2)
13430 .arg(&inf)
13431 .arg(&oi0)
13432 .arg(&oi1)
13433 .arg(&oi2)
13434 .arg(&mi);
13435 unsafe {
13436 b.launch(cfg)?;
13437 }
13438 return Ok(Some((y0, y1, y2)));
13439 }
13440 let f = self.func("qmatvec_nvfp4_mmvq_fused3_rp");
13441 let cfg = LaunchConfig {
13442 grid_dim: (nb(o0) + nb(o1) + nb(o2), m as u32, 1),
13443 block_dim: (32, ROWS_PER_BLOCK, 1),
13444 shared_mem_bytes: 0,
13445 };
13446 let __s_b = self.gpu.stream();
13447 let mut b = __s_b.launch_builder(&f);
13448 b.arg(b0)
13449 .arg(b1)
13450 .arg(b2)
13451 .arg(aq)
13452 .arg(ad)
13453 .arg(&mut y0)
13454 .arg(&mut y1)
13455 .arg(&mut y2)
13456 .arg(&inf)
13457 .arg(&oi0)
13458 .arg(&oi1)
13459 .arg(&oi2)
13460 .arg(&mi)
13461 .arg(&p0.1)
13462 .arg(&p1.1)
13463 .arg(&p2.1);
13464 unsafe {
13465 b.launch(cfg)?;
13466 }
13467 Ok(Some((y0, y1, y2)))
13468 }
13469
13470 pub fn matmul_nvfp4_fused2(
13479 &self,
13480 w0: &crate::model::GpuTensor,
13481 w1: &crate::model::GpuTensor,
13482 aq: &CudaSlice<i8>,
13483 ad: &CudaSlice<f32>,
13484 m: usize,
13485 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
13486 use crate::model::GpuTensor;
13487 static FUSED2_OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13488 let off =
13489 *FUSED2_OFF.get_or_init(|| std::env::var("MEMRA_NVFP4_FUSED2").as_deref() == Ok("0"));
13490 if off
13493 || m != 1
13494 || !self.mmvq_supports(QT_NVFP4)
13495 || !self.uses_q8_1_fast(w0)
13496 || !self.uses_q8_1_fast(w1)
13497 {
13498 return Ok(None);
13499 }
13500 let unpack = |w: &crate::model::GpuTensor| match w {
13501 GpuTensor::Quant {
13502 bytes,
13503 qtype,
13504 scale,
13505 rp,
13506 ..
13507 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
13508 _ => None,
13509 };
13510 let (Some(p0), Some(p1)) = (unpack(w0), unpack(w1)) else {
13511 return Ok(None);
13512 };
13513 let in_f = w0.in_features();
13514 if w1.in_features() != in_f {
13515 return Ok(None);
13516 }
13517 let (o0, o1) = (w0.out_features(), w1.out_features());
13518 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
13520 let rows_pb = ROWS_PER_BLOCK * RPW;
13521 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
13522 let f = self.func("qmatvec_nvfp4_mmvq_fused2_rp");
13523 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
13524 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
13525 let cfg = LaunchConfig {
13526 grid_dim: (nb(o0) + nb(o1), m as u32, 1),
13527 block_dim: (32, ROWS_PER_BLOCK, 1),
13528 shared_mem_bytes: 0,
13529 };
13530 let (inf, oi0, oi1, mi) = (in_f as i32, o0 as i32, o1 as i32, m as i32);
13531 let (b0, b1) = unsafe { (&*p0.0, &*p1.0) };
13534 if Self::pdl_on() && Self::pdl_mmvq_on() && Self::pdl_nvfp4q8_on() {
13537 {
13538 use cudarc::driver::{DevicePtr, DevicePtrMut};
13539 let s = &self.gpu.stream();
13540 let (pw0, _g0) = b0.device_ptr(s);
13541 let (pw1, _g1) = b1.device_ptr(s);
13542 let (paq, _g2) = aq.device_ptr(s);
13543 let (pad, _g3) = ad.device_ptr(s);
13544 let (py0, _g4) = y0.device_ptr_mut(s);
13545 let (py1, _g5) = y1.device_ptr_mut(s);
13546 let (s0, s1) = (p0.1, p1.1);
13547 let mut ps = [
13548 &pw0 as *const _ as *mut std::ffi::c_void,
13549 &pw1 as *const _ as *mut _,
13550 &paq as *const _ as *mut _,
13551 &pad as *const _ as *mut _,
13552 &py0 as *const _ as *mut _,
13553 &py1 as *const _ as *mut _,
13554 &inf as *const _ as *mut _,
13555 &oi0 as *const _ as *mut _,
13556 &oi1 as *const _ as *mut _,
13557 &mi as *const _ as *mut _,
13558 &s0 as *const _ as *mut _,
13559 &s1 as *const _ as *mut _,
13560 ];
13561 unsafe {
13562 self.launch_pdl(
13563 "qmatvec_nvfp4_mmvq_fused2_rp",
13564 cfg.grid_dim,
13565 cfg.block_dim,
13566 &mut ps,
13567 )?;
13568 }
13569 }
13570 return Ok(Some((y0, y1)));
13571 }
13572 let __s_b = self.gpu.stream();
13573 let mut b = __s_b.launch_builder(&f);
13574 b.arg(b0)
13575 .arg(b1)
13576 .arg(aq)
13577 .arg(ad)
13578 .arg(&mut y0)
13579 .arg(&mut y1)
13580 .arg(&inf)
13581 .arg(&oi0)
13582 .arg(&oi1)
13583 .arg(&mi)
13584 .arg(&p0.1)
13585 .arg(&p1.1);
13586 unsafe {
13587 b.launch(cfg)?;
13588 }
13589 Ok(Some((y0, y1)))
13590 }
13591
13592 pub fn matmul_nvfp4_fused2_into(
13597 &self,
13598 w0: &crate::model::GpuTensor,
13599 w1: &crate::model::GpuTensor,
13600 aq: &CudaSlice<i8>,
13601 ad: &CudaSlice<f32>,
13602 y0: &mut CudaSlice<f32>,
13603 y1: &mut CudaSlice<f32>,
13604 ) -> Result<bool, Box<dyn std::error::Error>> {
13605 use crate::model::GpuTensor;
13606 static FUSED2_OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13607 let off =
13608 *FUSED2_OFF.get_or_init(|| std::env::var("MEMRA_NVFP4_FUSED2").as_deref() == Ok("0"));
13609 if off
13610 || !self.mmvq_supports(QT_NVFP4)
13611 || !self.uses_q8_1_fast(w0)
13612 || !self.uses_q8_1_fast(w1)
13613 {
13614 return Ok(false);
13615 }
13616 let unpack = |w: &crate::model::GpuTensor| match w {
13617 GpuTensor::Quant {
13618 bytes,
13619 qtype,
13620 scale,
13621 rp,
13622 ..
13623 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
13624 _ => None,
13625 };
13626 let (Some(p0), Some(p1)) = (unpack(w0), unpack(w1)) else {
13627 return Ok(false);
13628 };
13629 let in_f = w0.in_features();
13630 if w1.in_features() != in_f {
13631 return Ok(false);
13632 }
13633 let (o0, o1) = (w0.out_features(), w1.out_features());
13634 if y0.len() < o0 || y1.len() < o1 {
13635 return Ok(false);
13636 }
13637 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
13639 let rows_pb = ROWS_PER_BLOCK * RPW;
13640 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
13641 let f = self.func("qmatvec_nvfp4_mmvq_fused2_rp");
13642 let cfg = LaunchConfig {
13643 grid_dim: (nb(o0) + nb(o1), 1, 1),
13644 block_dim: (32, ROWS_PER_BLOCK, 1),
13645 shared_mem_bytes: 0,
13646 };
13647 let (inf, oi0, oi1, mi) = (in_f as i32, o0 as i32, o1 as i32, 1i32);
13648 let (b0, b1) = unsafe { (&*p0.0, &*p1.0) };
13651 let __s_b = self.gpu.stream();
13652 let mut b = __s_b.launch_builder(&f);
13653 b.arg(b0)
13654 .arg(b1)
13655 .arg(aq)
13656 .arg(ad)
13657 .arg(&mut *y0)
13658 .arg(&mut *y1)
13659 .arg(&inf)
13660 .arg(&oi0)
13661 .arg(&oi1)
13662 .arg(&mi)
13663 .arg(&p0.1)
13664 .arg(&p1.1);
13665 unsafe {
13666 b.launch(cfg)?;
13667 }
13668 Ok(true)
13669 }
13670
13671 #[allow(clippy::type_complexity)]
13676 pub fn matmul_nvfp4_fused4(
13677 &self,
13678 w0: &crate::model::GpuTensor,
13679 w1: &crate::model::GpuTensor,
13680 w2: &crate::model::GpuTensor,
13681 w3: &crate::model::GpuTensor,
13682 aq: &CudaSlice<i8>,
13683 ad: &CudaSlice<f32>,
13684 m: usize,
13685 ) -> Result<
13686 Option<(
13687 CudaSlice<f32>,
13688 CudaSlice<f32>,
13689 CudaSlice<f32>,
13690 CudaSlice<f32>,
13691 )>,
13692 Box<dyn std::error::Error>,
13693 > {
13694 use crate::model::GpuTensor;
13695 if std::env::var("MEMRA_NVFP4_FUSED4").as_deref() == Ok("0")
13702 || !self.mmvq_supports(QT_NVFP4)
13703 || !self.uses_q8_1_fast(w0)
13704 || !self.uses_q8_1_fast(w1)
13705 || !self.uses_q8_1_fast(w2)
13706 || !self.uses_q8_1_fast(w3)
13707 {
13708 return Ok(None);
13709 }
13710 if (9..=16).contains(&m) {
13715 return Ok(
13716 match self.matmul_decode_exact_group4_pre([w0, w1, w2, w3], aq, ad, m)? {
13717 Some(mut ys) => {
13718 let y3 = ys.pop().unwrap();
13719 let y2 = ys.pop().unwrap();
13720 let y1 = ys.pop().unwrap();
13721 let y0 = ys.pop().unwrap();
13722 Some((y0, y1, y2, y3))
13723 }
13724 None => None,
13725 },
13726 );
13727 }
13728 if !(1..=8).contains(&m) {
13729 return Ok(None);
13730 }
13731 if m > 1 {
13732 let in_f = w0.in_features();
13735 if !self.batched_supports(QT_NVFP4)
13736 || std::env::var("MEMRA_NO_BATCHED").is_ok()
13737 || (m > 4 && !Self::b8_enabled())
13738 || in_f % 512 != 0
13739 || in_f / 64 > 272
13740 {
13741 return Ok(None);
13742 }
13743 }
13744 let unpack = |w: &crate::model::GpuTensor| match w {
13745 GpuTensor::Quant {
13746 bytes,
13747 qtype,
13748 scale,
13749 rp,
13750 ..
13751 } if *qtype == QT_NVFP4 && *rp => Some((bytes as *const CudaSlice<u8>, *scale)),
13752 _ => None,
13753 };
13754 let (Some(p0), Some(p1), Some(p2), Some(p3)) =
13755 (unpack(w0), unpack(w1), unpack(w2), unpack(w3))
13756 else {
13757 return Ok(None);
13758 };
13759 let in_f = w0.in_features();
13760 if w1.in_features() != in_f || w2.in_features() != in_f || w3.in_features() != in_f {
13761 return Ok(None);
13762 }
13763 let (o0, o1, o2, o3) = (
13764 w0.out_features(),
13765 w1.out_features(),
13766 w2.out_features(),
13767 w3.out_features(),
13768 );
13769 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
13771 let rows_pb = ROWS_PER_BLOCK * RPW;
13772 let nb = |o: usize| (o as u32).div_ceil(rows_pb);
13773 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
13774 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
13775 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
13776 let mut y3 = self.alloc_uninit::<f32>(m * o3)?;
13777 let (inf, oi0, oi1, oi2, oi3, mi) = (
13778 in_f as i32,
13779 o0 as i32,
13780 o1 as i32,
13781 o2 as i32,
13782 o3 as i32,
13783 m as i32,
13784 );
13785 let (b0, b1, b2, b3) = unsafe { (&*p0.0, &*p1.0, &*p2.0, &*p3.0) };
13788 if m > 1 {
13789 if p0.1 != 1.0 || p1.1 != 1.0 || p2.1 != 1.0 || p3.1 != 1.0 {
13792 return Ok(None);
13793 }
13794 let f = self.func("qmatvec_nvfp4_mmvq_fused4_b8_rpsc");
13795 let cfg = LaunchConfig {
13796 grid_dim: (nb(o0) + nb(o1) + nb(o2) + nb(o3), 1, 1),
13797 block_dim: (32, ROWS_PER_BLOCK, 1),
13798 shared_mem_bytes: 0,
13799 };
13800 let __s_b = self.gpu.stream();
13801 let mut b = __s_b.launch_builder(&f);
13802 b.arg(b0)
13803 .arg(b1)
13804 .arg(b2)
13805 .arg(b3)
13806 .arg(aq)
13807 .arg(ad)
13808 .arg(&mut y0)
13809 .arg(&mut y1)
13810 .arg(&mut y2)
13811 .arg(&mut y3)
13812 .arg(&inf)
13813 .arg(&oi0)
13814 .arg(&oi1)
13815 .arg(&oi2)
13816 .arg(&oi3)
13817 .arg(&mi);
13818 unsafe {
13819 b.launch(cfg)?;
13820 }
13821 return Ok(Some((y0, y1, y2, y3)));
13822 }
13823 let f = self.func("qmatvec_nvfp4_mmvq_fused4_rp");
13824 let cfg = LaunchConfig {
13825 grid_dim: (nb(o0) + nb(o1) + nb(o2) + nb(o3), m as u32, 1),
13826 block_dim: (32, ROWS_PER_BLOCK, 1),
13827 shared_mem_bytes: 0,
13828 };
13829 let __s_b = self.gpu.stream();
13830 let mut b = __s_b.launch_builder(&f);
13831 b.arg(b0)
13832 .arg(b1)
13833 .arg(b2)
13834 .arg(b3)
13835 .arg(aq)
13836 .arg(ad)
13837 .arg(&mut y0)
13838 .arg(&mut y1)
13839 .arg(&mut y2)
13840 .arg(&mut y3)
13841 .arg(&inf)
13842 .arg(&oi0)
13843 .arg(&oi1)
13844 .arg(&oi2)
13845 .arg(&oi3)
13846 .arg(&mi)
13847 .arg(&p0.1)
13848 .arg(&p1.1)
13849 .arg(&p2.1)
13850 .arg(&p3.1);
13851 unsafe {
13852 b.launch(cfg)?;
13853 }
13854 Ok(Some((y0, y1, y2, y3)))
13855 }
13856
13857 pub fn matmul_q8_fused2(
13865 &self,
13866 w0: &crate::model::GpuTensor,
13867 w1: &crate::model::GpuTensor,
13868 aq: &CudaSlice<i8>,
13869 ad: &CudaSlice<f32>,
13870 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
13871 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
13877 return Ok(Some(self.e4m3_fused2_core(
13878 p0.0,
13879 p1.0,
13880 aq,
13881 ad,
13882 w0.in_features(),
13883 p0.1,
13884 p1.1,
13885 p0.2,
13886 p0.3,
13887 p1.3,
13888 )?));
13889 }
13890 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
13891 return Ok(None);
13892 };
13893 Ok(Some(self.q8_fused2_core(
13894 p0.0,
13895 p1.0,
13896 aq,
13897 ad,
13898 w0.in_features(),
13899 p0.1,
13900 p1.1,
13901 p0.2,
13902 )?))
13903 }
13904
13905 #[allow(clippy::too_many_arguments)]
13906 fn q8_fused2_core(
13907 &self,
13908 b0: &CudaSlice<u8>,
13909 b1: &CudaSlice<u8>,
13910 aq: &CudaSlice<i8>,
13911 ad: &CudaSlice<f32>,
13912 in_f: usize,
13913 out0: usize,
13914 out1: usize,
13915 row_bytes: usize,
13916 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13917 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
13919 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
13920 let f = self.func("qmatvec_q8_0_mmvq_fused2");
13921 let mut y0 = self.alloc_uninit::<f32>(out0)?;
13922 let mut y1 = self.alloc_uninit::<f32>(out1)?;
13923 let cfg = LaunchConfig {
13924 grid_dim: (nb0 + nb1, 1, 1),
13925 block_dim: (32, ROWS_PER_BLOCK, 1),
13926 shared_mem_bytes: 0,
13927 };
13928 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
13929 let __s_b = self.gpu.stream();
13930 let mut b = __s_b.launch_builder(&f);
13931 b.arg(b0)
13932 .arg(b1)
13933 .arg(aq)
13934 .arg(ad)
13935 .arg(&mut y0)
13936 .arg(&mut y1)
13937 .arg(&inf)
13938 .arg(&o0)
13939 .arg(&o1)
13940 .arg(&rbl);
13941 unsafe {
13942 b.launch(cfg)?;
13943 }
13944 Ok((y0, y1))
13945 }
13946
13947 pub fn matmul_q8_fused2_x(
13953 &self,
13954 w0: &crate::model::GpuTensor,
13955 w1: &crate::model::GpuTensor,
13956 x: &CudaSlice<f32>,
13957 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
13958 if !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
13959 return Ok(None);
13960 }
13961 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
13962 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
13963 return Ok(Some(self.e4m3_fused2_core(
13964 p0.0,
13965 p1.0,
13966 &aq,
13967 &ad,
13968 w0.in_features(),
13969 p0.1,
13970 p1.1,
13971 p0.2,
13972 p0.3,
13973 p1.3,
13974 )?));
13975 }
13976 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
13977 return Ok(None);
13978 };
13979 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
13980 Ok(Some(self.q8_fused2_core(
13981 p0.0,
13982 p1.0,
13983 &aq,
13984 &ad,
13985 w0.in_features(),
13986 p0.1,
13987 p1.1,
13988 p0.2,
13989 )?))
13990 }
13991
13992 #[allow(clippy::too_many_arguments)]
13995 pub fn qmatvec_q8_fused2_raw(
13996 &self,
13997 b0: &CudaSlice<u8>,
13998 b1: &CudaSlice<u8>,
13999 x: &CudaSlice<f32>,
14000 in_f: usize,
14001 out0: usize,
14002 out1: usize,
14003 row_bytes: usize,
14004 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14005 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
14006 self.q8_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes)
14007 }
14008
14009 pub fn matmul_q4_fused3(
14015 &self,
14016 w0: &crate::model::GpuTensor,
14017 w1: &crate::model::GpuTensor,
14018 w2: &crate::model::GpuTensor,
14019 aq: &CudaSlice<i8>,
14020 ad: &CudaSlice<f32>,
14021 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
14022 {
14023 use crate::model::GpuTensor;
14024 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
14025 match w {
14026 GpuTensor::Quant {
14027 qtype, row_bytes, ..
14028 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
14029 _ => None,
14030 }
14031 };
14032 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2)) else {
14033 return Ok(None);
14034 };
14035 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
14036 return Ok(None);
14037 }
14038 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
14042 match w {
14043 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
14044 Some(m) => (m, true),
14045 None => (bytes, *rp),
14046 },
14047 _ => unreachable!(),
14048 }
14049 }
14050 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
14051 if rp0 != rp1 || rp1 != rp2 {
14052 return Ok(None);
14053 }
14054 let rp = rp0;
14055 let rpb: u32 = 4;
14056 let mr1 = rp && Self::q40_mr1_on();
14060 let nb = |o: usize| {
14061 if mr1 {
14062 (o as u32).div_ceil(rpb)
14063 } else {
14064 (o as u32).div_ceil(2).div_ceil(rpb)
14065 }
14066 };
14067 let grid = nb(o0) + nb(o1) + nb(o2);
14068 let mut y0 = self.alloc_uninit::<f32>(o0)?;
14069 let mut y1 = self.alloc_uninit::<f32>(o1)?;
14070 let mut y2 = self.alloc_uninit::<f32>(o2)?;
14071 let f = self.func(if mr1 {
14072 "qmatvec_q4_0_mmvq_fused3_mr1_rp"
14073 } else if rp {
14074 "qmatvec_q4_0_mmvq_fused3_rp"
14075 } else {
14076 "qmatvec_q4_0_mmvq_fused3"
14077 });
14078 let cfg = LaunchConfig {
14079 grid_dim: (grid, 1, 1),
14080 block_dim: (32, rpb, 1),
14081 shared_mem_bytes: 0,
14082 };
14083 let inf = w0.in_features() as i32;
14084 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
14085 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
14086 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
14089 {
14090 use cudarc::driver::{DevicePtr, DevicePtrMut};
14091 let s = &self.gpu.stream();
14092 let (p0, _g0) = b0.device_ptr(s);
14093 let (p1, _g1) = b1.device_ptr(s);
14094 let (p2, _g2) = b2.device_ptr(s);
14095 let (paq, _g3) = aq.device_ptr(s);
14096 let (pad, _g4) = ad.device_ptr(s);
14097 let (py0, _g5) = y0.device_ptr_mut(s);
14098 let (py1, _g6) = y1.device_ptr_mut(s);
14099 let (py2, _g7) = y2.device_ptr_mut(s);
14100 let mut ps = [
14101 &p0 as *const _ as *mut std::ffi::c_void,
14102 &p1 as *const _ as *mut _,
14103 &p2 as *const _ as *mut _,
14104 &paq as *const _ as *mut _,
14105 &pad as *const _ as *mut _,
14106 &py0 as *const _ as *mut _,
14107 &py1 as *const _ as *mut _,
14108 &py2 as *const _ as *mut _,
14109 &inf as *const _ as *mut _,
14110 &oo0 as *const _ as *mut _,
14111 &oo1 as *const _ as *mut _,
14112 &oo2 as *const _ as *mut _,
14113 &r0 as *const _ as *mut _,
14114 &r1 as *const _ as *mut _,
14115 &r2 as *const _ as *mut _,
14116 ];
14117 unsafe {
14118 self.launch_pdl(
14119 "qmatvec_q4_0_mmvq_fused3_mr1_rp",
14120 (grid, 1, 1),
14121 (32, rpb, 1),
14122 &mut ps,
14123 )?;
14124 }
14125 }
14126 return Ok(Some((y0, y1, y2)));
14127 }
14128 let __s_b = self.gpu.stream();
14129 let mut b = __s_b.launch_builder(&f);
14130 b.arg(b0)
14131 .arg(b1)
14132 .arg(b2)
14133 .arg(aq)
14134 .arg(ad)
14135 .arg(&mut y0)
14136 .arg(&mut y1)
14137 .arg(&mut y2)
14138 .arg(&inf)
14139 .arg(&oo0)
14140 .arg(&oo1)
14141 .arg(&oo2)
14142 .arg(&r0)
14143 .arg(&r1)
14144 .arg(&r2);
14145 unsafe {
14146 b.launch(cfg)?;
14147 }
14148 Ok(Some((y0, y1, y2)))
14149 }
14150
14151 #[allow(clippy::too_many_arguments)]
14154 pub fn matmul_q4_fused3_into(
14155 &self,
14156 w0: &crate::model::GpuTensor,
14157 w1: &crate::model::GpuTensor,
14158 w2: &crate::model::GpuTensor,
14159 aq: &CudaSlice<i8>,
14160 ad: &CudaSlice<f32>,
14161 y0: &mut CudaSlice<f32>,
14162 y1: &mut CudaSlice<f32>,
14163 y2: &mut CudaSlice<f32>,
14164 ) -> Result<bool, Box<dyn std::error::Error>> {
14165 use crate::model::GpuTensor;
14166 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
14167 match w {
14168 GpuTensor::Quant {
14169 qtype, row_bytes, ..
14170 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
14171 _ => None,
14172 }
14173 };
14174 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2)) else {
14175 return Ok(false);
14176 };
14177 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
14178 return Ok(false);
14179 }
14180 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
14181 match w {
14182 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
14183 Some(m) => (m, true),
14184 None => (bytes, *rp),
14185 },
14186 _ => unreachable!(),
14187 }
14188 }
14189 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
14190 if rp0 != rp1 || rp1 != rp2 {
14191 return Ok(false);
14192 }
14193 let rp = rp0;
14194 let rpb: u32 = 4;
14195 let mr1 = rp && Self::q40_mr1_on();
14196 let nb = |o: usize| {
14197 if mr1 {
14198 (o as u32).div_ceil(rpb)
14199 } else {
14200 (o as u32).div_ceil(2).div_ceil(rpb)
14201 }
14202 };
14203 let grid = nb(o0) + nb(o1) + nb(o2);
14204 debug_assert!(y0.len() >= o0 && y1.len() >= o1 && y2.len() >= o2);
14205 let f = self.func(if mr1 {
14206 "qmatvec_q4_0_mmvq_fused3_mr1_rp"
14207 } else if rp {
14208 "qmatvec_q4_0_mmvq_fused3_rp"
14209 } else {
14210 "qmatvec_q4_0_mmvq_fused3"
14211 });
14212 let cfg = LaunchConfig {
14213 grid_dim: (grid, 1, 1),
14214 block_dim: (32, rpb, 1),
14215 shared_mem_bytes: 0,
14216 };
14217 let inf = w0.in_features() as i32;
14218 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
14219 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
14220 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
14222 use cudarc::driver::{DevicePtr, DevicePtrMut};
14223 let s = &self.gpu.stream();
14224 let (p0, _g0) = b0.device_ptr(s);
14225 let (p1, _g1) = b1.device_ptr(s);
14226 let (p2, _g2) = b2.device_ptr(s);
14227 let (paq, _g3) = aq.device_ptr(s);
14228 let (pad, _g4) = ad.device_ptr(s);
14229 let (py0, _g5) = y0.device_ptr_mut(s);
14230 let (py1, _g6) = y1.device_ptr_mut(s);
14231 let (py2, _g7) = y2.device_ptr_mut(s);
14232 let mut ps = [
14233 &p0 as *const _ as *mut std::ffi::c_void,
14234 &p1 as *const _ as *mut _,
14235 &p2 as *const _ as *mut _,
14236 &paq as *const _ as *mut _,
14237 &pad as *const _ as *mut _,
14238 &py0 as *const _ as *mut _,
14239 &py1 as *const _ as *mut _,
14240 &py2 as *const _ as *mut _,
14241 &inf as *const _ as *mut _,
14242 &oo0 as *const _ as *mut _,
14243 &oo1 as *const _ as *mut _,
14244 &oo2 as *const _ as *mut _,
14245 &r0 as *const _ as *mut _,
14246 &r1 as *const _ as *mut _,
14247 &r2 as *const _ as *mut _,
14248 ];
14249 unsafe {
14250 self.launch_pdl(
14251 "qmatvec_q4_0_mmvq_fused3_mr1_rp",
14252 (grid, 1, 1),
14253 (32, rpb, 1),
14254 &mut ps,
14255 )?;
14256 }
14257 return Ok(true);
14258 }
14259 let __s_b = self.gpu.stream();
14260 let mut b = __s_b.launch_builder(&f);
14261 b.arg(b0)
14262 .arg(b1)
14263 .arg(b2)
14264 .arg(aq)
14265 .arg(ad)
14266 .arg(&mut *y0)
14267 .arg(&mut *y1)
14268 .arg(&mut *y2)
14269 .arg(&inf)
14270 .arg(&oo0)
14271 .arg(&oo1)
14272 .arg(&oo2)
14273 .arg(&r0)
14274 .arg(&r1)
14275 .arg(&r2);
14276 unsafe {
14277 b.launch(cfg)?;
14278 }
14279 Ok(true)
14280 }
14281
14282 pub fn matmul_q4_fused2(
14284 &self,
14285 w0: &crate::model::GpuTensor,
14286 w1: &crate::model::GpuTensor,
14287 aq: &CudaSlice<i8>,
14288 ad: &CudaSlice<f32>,
14289 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
14290 use crate::model::GpuTensor;
14291 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
14292 match w {
14293 GpuTensor::Quant {
14294 qtype, row_bytes, ..
14295 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
14296 _ => None,
14297 }
14298 };
14299 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else {
14300 return Ok(None);
14301 };
14302 if w0.in_features() != w1.in_features() {
14303 return Ok(None);
14304 }
14305 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
14307 match w {
14308 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
14309 Some(m) => (m, true),
14310 None => (bytes, *rp),
14311 },
14312 _ => unreachable!(),
14313 }
14314 }
14315 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
14316 if rp0 != rp1 {
14317 return Ok(None);
14318 }
14319 let rp = rp0;
14320 let rpb: u32 = 4;
14321 let mr1 = rp && Self::q40_mr1_on();
14323 let nb = |o: usize| {
14324 if mr1 {
14325 (o as u32).div_ceil(rpb)
14326 } else {
14327 (o as u32).div_ceil(2).div_ceil(rpb)
14328 }
14329 };
14330 let grid = nb(o0) + nb(o1);
14331 let mut y0 = self.alloc_uninit::<f32>(o0)?;
14332 let mut y1 = self.alloc_uninit::<f32>(o1)?;
14333 let f = self.func(if mr1 {
14334 "qmatvec_q4_0_mmvq_fused2_mr1_rp"
14335 } else if rp {
14336 "qmatvec_q4_0_mmvq_fused2_rp"
14337 } else {
14338 "qmatvec_q4_0_mmvq_fused2"
14339 });
14340 let cfg = LaunchConfig {
14341 grid_dim: (grid, 1, 1),
14342 block_dim: (32, rpb, 1),
14343 shared_mem_bytes: 0,
14344 };
14345 let inf = w0.in_features() as i32;
14346 let (oo0, oo1) = (o0 as i32, o1 as i32);
14347 let (r0, r1) = (rb0 as i64, rb1 as i64);
14348 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
14350 {
14351 use cudarc::driver::{DevicePtr, DevicePtrMut};
14352 let s = &self.gpu.stream();
14353 let (p0, _g0) = b0.device_ptr(s);
14354 let (p1, _g1) = b1.device_ptr(s);
14355 let (paq, _g2) = aq.device_ptr(s);
14356 let (pad, _g3) = ad.device_ptr(s);
14357 let (py0, _g4) = y0.device_ptr_mut(s);
14358 let (py1, _g5) = y1.device_ptr_mut(s);
14359 let mut ps = [
14360 &p0 as *const _ as *mut std::ffi::c_void,
14361 &p1 as *const _ as *mut _,
14362 &paq as *const _ as *mut _,
14363 &pad as *const _ as *mut _,
14364 &py0 as *const _ as *mut _,
14365 &py1 as *const _ as *mut _,
14366 &inf as *const _ as *mut _,
14367 &oo0 as *const _ as *mut _,
14368 &oo1 as *const _ as *mut _,
14369 &r0 as *const _ as *mut _,
14370 &r1 as *const _ as *mut _,
14371 ];
14372 unsafe {
14373 self.launch_pdl(
14374 "qmatvec_q4_0_mmvq_fused2_mr1_rp",
14375 (grid, 1, 1),
14376 (32, rpb, 1),
14377 &mut ps,
14378 )?;
14379 }
14380 }
14381 return Ok(Some((y0, y1)));
14382 }
14383 let __s_b = self.gpu.stream();
14384 let mut b = __s_b.launch_builder(&f);
14385 b.arg(b0)
14386 .arg(b1)
14387 .arg(aq)
14388 .arg(ad)
14389 .arg(&mut y0)
14390 .arg(&mut y1)
14391 .arg(&inf)
14392 .arg(&oo0)
14393 .arg(&oo1)
14394 .arg(&r0)
14395 .arg(&r1);
14396 unsafe {
14397 b.launch(cfg)?;
14398 }
14399 Ok(Some((y0, y1)))
14400 }
14401
14402 pub fn matmul_q4_fused2_into(
14404 &self,
14405 w0: &crate::model::GpuTensor,
14406 w1: &crate::model::GpuTensor,
14407 aq: &CudaSlice<i8>,
14408 ad: &CudaSlice<f32>,
14409 y0: &mut CudaSlice<f32>,
14410 y1: &mut CudaSlice<f32>,
14411 ) -> Result<bool, Box<dyn std::error::Error>> {
14412 use crate::model::GpuTensor;
14413 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
14414 match w {
14415 GpuTensor::Quant {
14416 qtype, row_bytes, ..
14417 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
14418 _ => None,
14419 }
14420 };
14421 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else {
14422 return Ok(false);
14423 };
14424 if w0.in_features() != w1.in_features() {
14425 return Ok(false);
14426 }
14427 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
14428 match w {
14429 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
14430 Some(m) => (m, true),
14431 None => (bytes, *rp),
14432 },
14433 _ => unreachable!(),
14434 }
14435 }
14436 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
14437 if rp0 != rp1 {
14438 return Ok(false);
14439 }
14440 let rp = rp0;
14441 let rpb: u32 = 4;
14442 let mr1 = rp && Self::q40_mr1_on();
14443 let nb = |o: usize| {
14444 if mr1 {
14445 (o as u32).div_ceil(rpb)
14446 } else {
14447 (o as u32).div_ceil(2).div_ceil(rpb)
14448 }
14449 };
14450 let grid = nb(o0) + nb(o1);
14451 debug_assert!(y0.len() >= o0 && y1.len() >= o1);
14452 let f = self.func(if mr1 {
14453 "qmatvec_q4_0_mmvq_fused2_mr1_rp"
14454 } else if rp {
14455 "qmatvec_q4_0_mmvq_fused2_rp"
14456 } else {
14457 "qmatvec_q4_0_mmvq_fused2"
14458 });
14459 let cfg = LaunchConfig {
14460 grid_dim: (grid, 1, 1),
14461 block_dim: (32, rpb, 1),
14462 shared_mem_bytes: 0,
14463 };
14464 let inf = w0.in_features() as i32;
14465 let (oo0, oo1) = (o0 as i32, o1 as i32);
14466 let (r0, r1) = (rb0 as i64, rb1 as i64);
14467 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
14469 use cudarc::driver::{DevicePtr, DevicePtrMut};
14470 let s = &self.gpu.stream();
14471 let (p0, _g0) = b0.device_ptr(s);
14472 let (p1, _g1) = b1.device_ptr(s);
14473 let (paq, _g2) = aq.device_ptr(s);
14474 let (pad, _g3) = ad.device_ptr(s);
14475 let (py0, _g4) = y0.device_ptr_mut(s);
14476 let (py1, _g5) = y1.device_ptr_mut(s);
14477 let mut ps = [
14478 &p0 as *const _ as *mut std::ffi::c_void,
14479 &p1 as *const _ as *mut _,
14480 &paq as *const _ as *mut _,
14481 &pad as *const _ as *mut _,
14482 &py0 as *const _ as *mut _,
14483 &py1 as *const _ as *mut _,
14484 &inf as *const _ as *mut _,
14485 &oo0 as *const _ as *mut _,
14486 &oo1 as *const _ as *mut _,
14487 &r0 as *const _ as *mut _,
14488 &r1 as *const _ as *mut _,
14489 ];
14490 unsafe {
14491 self.launch_pdl(
14492 "qmatvec_q4_0_mmvq_fused2_mr1_rp",
14493 (grid, 1, 1),
14494 (32, rpb, 1),
14495 &mut ps,
14496 )?;
14497 }
14498 return Ok(true);
14499 }
14500 let __s_b = self.gpu.stream();
14501 let mut b = __s_b.launch_builder(&f);
14502 b.arg(b0)
14503 .arg(b1)
14504 .arg(aq)
14505 .arg(ad)
14506 .arg(&mut *y0)
14507 .arg(&mut *y1)
14508 .arg(&inf)
14509 .arg(&oo0)
14510 .arg(&oo1)
14511 .arg(&r0)
14512 .arg(&r1);
14513 unsafe {
14514 b.launch(cfg)?;
14515 }
14516 Ok(true)
14517 }
14518
14519 pub fn matmul_q4_fused2_batched(
14524 &self,
14525 w0: &crate::model::GpuTensor,
14526 w1: &crate::model::GpuTensor,
14527 aq: &CudaSlice<i8>,
14528 ad: &CudaSlice<f32>,
14529 m: usize,
14530 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
14531 use crate::model::GpuTensor;
14532 if m < 2 || m > 8 {
14533 return Ok(None);
14534 }
14535 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
14536 match w {
14537 GpuTensor::Quant {
14538 qtype, row_bytes, ..
14539 } if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
14540 _ => None,
14541 }
14542 };
14543 let (Some((rb0, o0)), Some((_rb1, o1))) = (q4(w0), q4(w1)) else {
14544 return Ok(None);
14545 };
14546 if w0.in_features() != w1.in_features() {
14547 return Ok(None);
14548 }
14549 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
14550 match w {
14551 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
14552 Some(mr) => (mr, true),
14553 None => (bytes, *rp),
14554 },
14555 _ => unreachable!(),
14556 }
14557 }
14558 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
14559 if !rp0 || !rp1 {
14560 return Ok(None);
14561 }
14562 let mcols = Self::batched_mcols(m);
14563 let rpb: u32 = 4;
14564 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
14565 let grid = nb(o0) + nb(o1);
14566 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
14567 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
14568 let f = self.func(match mcols {
14569 2 => "qmatvec_q4_0_mmvq_b2_f2_rp",
14570 4 => "qmatvec_q4_0_mmvq_b4_f2_rp",
14571 _ => "qmatvec_q4_0_mmvq_b8_f2_rp",
14572 });
14573 let cfg = LaunchConfig {
14574 grid_dim: (grid, 1, 1),
14575 block_dim: (32, rpb, 1),
14576 shared_mem_bytes: 0,
14577 };
14578 let inf = w0.in_features() as i32;
14579 let (oo0, oo1, mi) = (o0 as i32, o1 as i32, m as i32);
14580 let rb = rb0 as i64;
14581 let __s_b = self.gpu.stream();
14582 let mut b = __s_b.launch_builder(&f);
14583 b.arg(b0)
14584 .arg(b1)
14585 .arg(aq)
14586 .arg(ad)
14587 .arg(&mut y0)
14588 .arg(&mut y1)
14589 .arg(&inf)
14590 .arg(&oo0)
14591 .arg(&oo1)
14592 .arg(&mi)
14593 .arg(&rb);
14594 unsafe {
14595 b.launch(cfg)?;
14596 }
14597 Ok(Some((y0, y1)))
14598 }
14599
14600 #[allow(clippy::too_many_arguments)]
14603 pub fn matmul_q4_fused3_batched(
14604 &self,
14605 w0: &crate::model::GpuTensor,
14606 w1: &crate::model::GpuTensor,
14607 w2: &crate::model::GpuTensor,
14608 aq: &CudaSlice<i8>,
14609 ad: &CudaSlice<f32>,
14610 m: usize,
14611 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
14612 {
14613 use crate::model::GpuTensor;
14614 if m < 2 || m > 8 {
14615 return Ok(None);
14616 }
14617 let q4 = |w: &GpuTensor| -> Option<usize> {
14618 match w {
14619 GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0 => Some(w.out_features()),
14620 _ => None,
14621 }
14622 };
14623 let (Some(o0), Some(o1), Some(o2)) = (q4(w0), q4(w1), q4(w2)) else {
14624 return Ok(None);
14625 };
14626 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
14627 return Ok(None);
14628 }
14629 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
14630 match w {
14631 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
14632 Some(mr) => (mr, true),
14633 None => (bytes, *rp),
14634 },
14635 _ => unreachable!(),
14636 }
14637 }
14638 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
14639 if !rp0 || !rp1 || !rp2 {
14640 return Ok(None);
14641 }
14642 let mcols = Self::batched_mcols(m);
14643 let rpb: u32 = 4;
14644 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
14645 let grid = nb(o0) + nb(o1) + nb(o2);
14646 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
14647 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
14648 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
14649 let f = self.func(match mcols {
14650 2 => "qmatvec_q4_0_mmvq_b2_f3_rp",
14651 4 => "qmatvec_q4_0_mmvq_b4_f3_rp",
14652 _ => "qmatvec_q4_0_mmvq_b8_f3_rp",
14653 });
14654 let cfg = LaunchConfig {
14655 grid_dim: (grid, 1, 1),
14656 block_dim: (32, rpb, 1),
14657 shared_mem_bytes: 0,
14658 };
14659 let inf = w0.in_features() as i32;
14660 let (oo0, oo1, oo2, mi) = (o0 as i32, o1 as i32, o2 as i32, m as i32);
14661 let rb = 0i64;
14662 let __s_b = self.gpu.stream();
14663 let mut b = __s_b.launch_builder(&f);
14664 b.arg(b0)
14665 .arg(b1)
14666 .arg(b2)
14667 .arg(aq)
14668 .arg(ad)
14669 .arg(&mut y0)
14670 .arg(&mut y1)
14671 .arg(&mut y2)
14672 .arg(&inf)
14673 .arg(&oo0)
14674 .arg(&oo1)
14675 .arg(&oo2)
14676 .arg(&mi)
14677 .arg(&rb);
14678 unsafe {
14679 b.launch(cfg)?;
14680 }
14681 Ok(Some((y0, y1, y2)))
14682 }
14683
14684 pub fn matmul_q8_fused3(
14685 &self,
14686 w0: &crate::model::GpuTensor,
14687 w1: &crate::model::GpuTensor,
14688 w2: &crate::model::GpuTensor,
14689 aq: &CudaSlice<i8>,
14690 ad: &CudaSlice<f32>,
14691 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
14692 {
14693 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
14696 return Ok(Some(self.e4m3_fused3_core(
14697 p0.0,
14698 p1.0,
14699 p2.0,
14700 aq,
14701 ad,
14702 w0.in_features(),
14703 p0.1,
14704 p1.1,
14705 p2.1,
14706 p0.2,
14707 p0.3,
14708 p1.3,
14709 p2.3,
14710 )?));
14711 }
14712 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else {
14713 return Ok(None);
14714 };
14715 Ok(Some(self.q8_fused3_core(
14716 p0.0,
14717 p1.0,
14718 p2.0,
14719 aq,
14720 ad,
14721 w0.in_features(),
14722 p0.1,
14723 p1.1,
14724 p2.1,
14725 p0.2,
14726 )?))
14727 }
14728
14729 #[allow(clippy::too_many_arguments)]
14730 fn q8_fused3_core(
14731 &self,
14732 b0: &CudaSlice<u8>,
14733 b1: &CudaSlice<u8>,
14734 b2: &CudaSlice<u8>,
14735 aq: &CudaSlice<i8>,
14736 ad: &CudaSlice<f32>,
14737 in_f: usize,
14738 out0: usize,
14739 out1: usize,
14740 out2: usize,
14741 row_bytes: usize,
14742 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14743 const ROWS_PER_BLOCK: u32 = 4;
14744 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
14745 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
14746 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
14747 let f = self.func("qmatvec_q8_0_mmvq_fused3");
14748 let mut y0 = self.alloc_uninit::<f32>(out0)?;
14749 let mut y1 = self.alloc_uninit::<f32>(out1)?;
14750 let mut y2 = self.alloc_uninit::<f32>(out2)?;
14751 let cfg = LaunchConfig {
14752 grid_dim: (nb0 + nb1 + nb2, 1, 1),
14753 block_dim: (32, ROWS_PER_BLOCK, 1),
14754 shared_mem_bytes: 0,
14755 };
14756 let (inf, o0, o1, o2, rbl) = (
14757 in_f as i32,
14758 out0 as i32,
14759 out1 as i32,
14760 out2 as i32,
14761 row_bytes as i64,
14762 );
14763 let __s_b = self.gpu.stream();
14764 let mut b = __s_b.launch_builder(&f);
14765 b.arg(b0)
14766 .arg(b1)
14767 .arg(b2)
14768 .arg(aq)
14769 .arg(ad)
14770 .arg(&mut y0)
14771 .arg(&mut y1)
14772 .arg(&mut y2)
14773 .arg(&inf)
14774 .arg(&o0)
14775 .arg(&o1)
14776 .arg(&o2)
14777 .arg(&rbl);
14778 unsafe {
14779 b.launch(cfg)?;
14780 }
14781 Ok((y0, y1, y2))
14782 }
14783
14784 #[allow(clippy::too_many_arguments)]
14786 pub fn qmatvec_q8_fused3_raw(
14787 &self,
14788 b0: &CudaSlice<u8>,
14789 b1: &CudaSlice<u8>,
14790 b2: &CudaSlice<u8>,
14791 x: &CudaSlice<f32>,
14792 in_f: usize,
14793 out0: usize,
14794 out1: usize,
14795 out2: usize,
14796 row_bytes: usize,
14797 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14798 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
14799 self.q8_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes)
14800 }
14801
14802 pub fn matmul_q8_fused2_t(
14813 &self,
14814 w0: &crate::model::GpuTensor,
14815 w1: &crate::model::GpuTensor,
14816 aq: &CudaSlice<i8>,
14817 ad: &CudaSlice<f32>,
14818 m: usize,
14819 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
14820 if !(2..=8).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() {
14824 return Ok(None);
14825 }
14826 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
14829 if m > 4 && !Self::b8_enabled() {
14830 return Ok(None);
14831 }
14832 return Ok(Some(self.e4m3_fused2_t_core(
14833 p0.0,
14834 p1.0,
14835 aq,
14836 ad,
14837 m,
14838 w0.in_features(),
14839 p0.1,
14840 p1.1,
14841 p0.2,
14842 p0.3,
14843 p1.3,
14844 )?));
14845 }
14846 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
14847 return Ok(None);
14848 };
14849 Ok(Some(self.q8_fused2_t_core(
14850 p0.0,
14851 p1.0,
14852 aq,
14853 ad,
14854 m,
14855 w0.in_features(),
14856 p0.1,
14857 p1.1,
14858 p0.2,
14859 )?))
14860 }
14861
14862 #[allow(clippy::too_many_arguments)]
14863 fn q8_fused2_t_core(
14864 &self,
14865 b0: &CudaSlice<u8>,
14866 b1: &CudaSlice<u8>,
14867 aq: &CudaSlice<i8>,
14868 ad: &CudaSlice<f32>,
14869 m: usize,
14870 in_f: usize,
14871 out0: usize,
14872 out1: usize,
14873 row_bytes: usize,
14874 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14875 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
14877 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
14878 let f = self.func(match Self::batched_mcols(m) {
14879 2 => "qmatvec_q8_0_mmvq_fused2_b2",
14880 4 => "qmatvec_q8_0_mmvq_fused2_b4",
14881 _ => "qmatvec_q8_0_mmvq_fused2_b8",
14883 });
14884 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
14885 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
14886 let cfg = LaunchConfig {
14887 grid_dim: (nb0 + nb1, 1, 1),
14888 block_dim: (32, ROWS_PER_BLOCK, 1),
14889 shared_mem_bytes: 0,
14890 };
14891 let (inf, o0, o1, mi, rbl) = (
14892 in_f as i32,
14893 out0 as i32,
14894 out1 as i32,
14895 m as i32,
14896 row_bytes as i64,
14897 );
14898 let __s_b = self.gpu.stream();
14899 let mut b = __s_b.launch_builder(&f);
14900 b.arg(b0)
14901 .arg(b1)
14902 .arg(aq)
14903 .arg(ad)
14904 .arg(&mut y0)
14905 .arg(&mut y1)
14906 .arg(&inf)
14907 .arg(&o0)
14908 .arg(&o1)
14909 .arg(&mi)
14910 .arg(&rbl);
14911 unsafe {
14912 b.launch(cfg)?;
14913 }
14914 Ok((y0, y1))
14915 }
14916
14917 #[allow(clippy::too_many_arguments)]
14920 pub fn qmatvec_q8_fused2_t_raw(
14921 &self,
14922 b0: &CudaSlice<u8>,
14923 b1: &CudaSlice<u8>,
14924 x: &CudaSlice<f32>,
14925 m: usize,
14926 in_f: usize,
14927 out0: usize,
14928 out1: usize,
14929 row_bytes: usize,
14930 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14931 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
14932 self.q8_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes)
14933 }
14934
14935 #[allow(clippy::too_many_arguments)]
14938 pub fn matmul_q8_fused3_t(
14939 &self,
14940 w0: &crate::model::GpuTensor,
14941 w1: &crate::model::GpuTensor,
14942 w2: &crate::model::GpuTensor,
14943 aq: &CudaSlice<i8>,
14944 ad: &CudaSlice<f32>,
14945 m: usize,
14946 ) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
14947 {
14948 if !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() {
14949 return Ok(None);
14950 }
14951 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
14952 return Ok(Some(self.e4m3_fused3_t_core(
14953 p0.0,
14954 p1.0,
14955 p2.0,
14956 aq,
14957 ad,
14958 m,
14959 w0.in_features(),
14960 p0.1,
14961 p1.1,
14962 p2.1,
14963 p0.2,
14964 p0.3,
14965 p1.3,
14966 p2.3,
14967 )?));
14968 }
14969 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else {
14970 return Ok(None);
14971 };
14972 Ok(Some(self.q8_fused3_t_core(
14973 p0.0,
14974 p1.0,
14975 p2.0,
14976 aq,
14977 ad,
14978 m,
14979 w0.in_features(),
14980 p0.1,
14981 p1.1,
14982 p2.1,
14983 p0.2,
14984 )?))
14985 }
14986
14987 #[allow(clippy::too_many_arguments)]
14988 fn q8_fused3_t_core(
14989 &self,
14990 b0: &CudaSlice<u8>,
14991 b1: &CudaSlice<u8>,
14992 b2: &CudaSlice<u8>,
14993 aq: &CudaSlice<i8>,
14994 ad: &CudaSlice<f32>,
14995 m: usize,
14996 in_f: usize,
14997 out0: usize,
14998 out1: usize,
14999 out2: usize,
15000 row_bytes: usize,
15001 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15002 const ROWS_PER_BLOCK: u32 = 4;
15003 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
15004 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
15005 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
15006 let f = self.func(if Self::batched_mcols(m) == 2 {
15007 "qmatvec_q8_0_mmvq_fused3_b2"
15008 } else {
15009 "qmatvec_q8_0_mmvq_fused3_b4"
15010 });
15011 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
15012 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
15013 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
15014 let cfg = LaunchConfig {
15015 grid_dim: (nb0 + nb1 + nb2, 1, 1),
15016 block_dim: (32, ROWS_PER_BLOCK, 1),
15017 shared_mem_bytes: 0,
15018 };
15019 let (inf, o0, o1, o2, mi, rbl) = (
15020 in_f as i32,
15021 out0 as i32,
15022 out1 as i32,
15023 out2 as i32,
15024 m as i32,
15025 row_bytes as i64,
15026 );
15027 let __s_b = self.gpu.stream();
15028 let mut b = __s_b.launch_builder(&f);
15029 b.arg(b0)
15030 .arg(b1)
15031 .arg(b2)
15032 .arg(aq)
15033 .arg(ad)
15034 .arg(&mut y0)
15035 .arg(&mut y1)
15036 .arg(&mut y2)
15037 .arg(&inf)
15038 .arg(&o0)
15039 .arg(&o1)
15040 .arg(&o2)
15041 .arg(&mi)
15042 .arg(&rbl);
15043 unsafe {
15044 b.launch(cfg)?;
15045 }
15046 Ok((y0, y1, y2))
15047 }
15048
15049 #[allow(clippy::too_many_arguments)]
15051 pub fn qmatvec_q8_fused3_t_raw(
15052 &self,
15053 b0: &CudaSlice<u8>,
15054 b1: &CudaSlice<u8>,
15055 b2: &CudaSlice<u8>,
15056 x: &CudaSlice<f32>,
15057 m: usize,
15058 in_f: usize,
15059 out0: usize,
15060 out1: usize,
15061 out2: usize,
15062 row_bytes: usize,
15063 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15064 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15065 self.q8_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes)
15066 }
15067
15068 pub fn q8_ffn_fuse2_on(&self) -> bool {
15072 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15073 *ON.get_or_init(|| std::env::var("MEMRA_Q8_FFN_FUSE2").as_deref() != Ok("0"))
15074 }
15075
15076 #[allow(clippy::type_complexity)]
15082 fn q8_fused_params<'w, const N: usize>(
15083 &self,
15084 ws: &[&'w crate::model::GpuTensor; N],
15085 ) -> Option<[(&'w CudaSlice<u8>, usize, usize); N]> {
15086 use crate::model::GpuTensor;
15087 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") {
15088 return None;
15089 }
15090 if std::env::var("MEMRA_Q8_DUAL").is_ok_and(|v| v == "0") {
15091 return None;
15092 }
15093 let in_f = ws[0].in_features();
15094 let mut out: [Option<(&CudaSlice<u8>, usize, usize)>; N] = [None; N];
15095 for (i, w) in ws.iter().enumerate() {
15096 match w {
15097 GpuTensor::Quant {
15098 bytes,
15099 qtype,
15100 row_bytes,
15101 scale,
15102 ..
15103 } if *qtype == QT_Q8_0 && *scale == 1.0 && w.in_features() == in_f => {
15104 out[i] = Some((bytes, w.out_features(), *row_bytes))
15105 }
15106 _ => return None,
15107 }
15108 }
15109 Some(out.map(|o| o.unwrap()))
15110 }
15111
15112 pub fn e4m3_dual_on(&self) -> bool {
15115 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15116 *ON.get_or_init(|| std::env::var("MEMRA_E4M3_DUAL").as_deref() != Ok("0"))
15117 }
15118
15119 #[allow(clippy::type_complexity)]
15131 fn e4m3_fused_params<'w, const N: usize>(
15132 &self,
15133 ws: &[&'w crate::model::GpuTensor; N],
15134 ) -> Option<[(&'w CudaSlice<u8>, usize, usize, f32); N]> {
15135 use crate::model::GpuTensor;
15136 if !self.e4m3_dual_on() {
15137 return None;
15138 }
15139 let in_f = ws[0].in_features();
15140 let mut out: [Option<(&CudaSlice<u8>, usize, usize, f32)>; N] = [None; N];
15141 for (i, w) in ws.iter().enumerate() {
15142 match w {
15143 GpuTensor::Quant {
15144 bytes,
15145 qtype,
15146 row_bytes,
15147 scale,
15148 rp,
15149 rp4,
15150 ..
15151 } if *qtype == QT_F8_E4M3
15152 && w.in_features() == in_f
15153 && *row_bytes == in_f
15154 && !*rp
15155 && rp4.is_none() =>
15156 {
15157 out[i] = Some((bytes, w.out_features(), *row_bytes, *scale))
15158 }
15159 _ => return None,
15160 }
15161 }
15162 Some(out.map(|o| o.unwrap()))
15163 }
15164
15165 #[allow(clippy::too_many_arguments)]
15169 fn e4m3_fused2_core(
15170 &self,
15171 b0: &CudaSlice<u8>,
15172 b1: &CudaSlice<u8>,
15173 aq: &CudaSlice<i8>,
15174 ad: &CudaSlice<f32>,
15175 in_f: usize,
15176 out0: usize,
15177 out1: usize,
15178 row_bytes: usize,
15179 ws0: f32,
15180 ws1: f32,
15181 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15182 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
15184 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
15185 let f = self.func("qmatvec_e4m3_mmvq_fused2");
15186 let mut y0 = self.alloc_uninit::<f32>(out0)?;
15187 let mut y1 = self.alloc_uninit::<f32>(out1)?;
15188 let cfg = LaunchConfig {
15189 grid_dim: (nb0 + nb1, 1, 1),
15190 block_dim: (32, ROWS_PER_BLOCK, 1),
15191 shared_mem_bytes: 0,
15192 };
15193 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
15194 let __s_b = self.gpu.stream();
15195 let mut b = __s_b.launch_builder(&f);
15196 b.arg(b0)
15197 .arg(b1)
15198 .arg(aq)
15199 .arg(ad)
15200 .arg(&mut y0)
15201 .arg(&mut y1)
15202 .arg(&inf)
15203 .arg(&o0)
15204 .arg(&o1)
15205 .arg(&rbl)
15206 .arg(&ws0)
15207 .arg(&ws1);
15208 unsafe {
15209 b.launch(cfg)?;
15210 }
15211 Ok((y0, y1))
15212 }
15213
15214 #[allow(clippy::too_many_arguments)]
15216 fn e4m3_fused3_core(
15217 &self,
15218 b0: &CudaSlice<u8>,
15219 b1: &CudaSlice<u8>,
15220 b2: &CudaSlice<u8>,
15221 aq: &CudaSlice<i8>,
15222 ad: &CudaSlice<f32>,
15223 in_f: usize,
15224 out0: usize,
15225 out1: usize,
15226 out2: usize,
15227 row_bytes: usize,
15228 ws0: f32,
15229 ws1: f32,
15230 ws2: f32,
15231 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15232 const ROWS_PER_BLOCK: u32 = 4;
15233 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
15234 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
15235 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
15236 let f = self.func("qmatvec_e4m3_mmvq_fused3");
15237 let mut y0 = self.alloc_uninit::<f32>(out0)?;
15238 let mut y1 = self.alloc_uninit::<f32>(out1)?;
15239 let mut y2 = self.alloc_uninit::<f32>(out2)?;
15240 let cfg = LaunchConfig {
15241 grid_dim: (nb0 + nb1 + nb2, 1, 1),
15242 block_dim: (32, ROWS_PER_BLOCK, 1),
15243 shared_mem_bytes: 0,
15244 };
15245 let (inf, o0, o1, o2, rbl) = (
15246 in_f as i32,
15247 out0 as i32,
15248 out1 as i32,
15249 out2 as i32,
15250 row_bytes as i64,
15251 );
15252 let __s_b = self.gpu.stream();
15253 let mut b = __s_b.launch_builder(&f);
15254 b.arg(b0)
15255 .arg(b1)
15256 .arg(b2)
15257 .arg(aq)
15258 .arg(ad)
15259 .arg(&mut y0)
15260 .arg(&mut y1)
15261 .arg(&mut y2)
15262 .arg(&inf)
15263 .arg(&o0)
15264 .arg(&o1)
15265 .arg(&o2)
15266 .arg(&rbl)
15267 .arg(&ws0)
15268 .arg(&ws1)
15269 .arg(&ws2);
15270 unsafe {
15271 b.launch(cfg)?;
15272 }
15273 Ok((y0, y1, y2))
15274 }
15275
15276 #[allow(clippy::too_many_arguments)]
15280 fn e4m3_fused2_t_core(
15281 &self,
15282 b0: &CudaSlice<u8>,
15283 b1: &CudaSlice<u8>,
15284 aq: &CudaSlice<i8>,
15285 ad: &CudaSlice<f32>,
15286 m: usize,
15287 in_f: usize,
15288 out0: usize,
15289 out1: usize,
15290 row_bytes: usize,
15291 ws0: f32,
15292 ws1: f32,
15293 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15294 const ROWS_PER_BLOCK: u32 = 4;
15295 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
15296 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
15297 let f = self.func(match Self::batched_mcols(m) {
15298 2 => "qmatvec_e4m3_mmvq_fused2_b2",
15299 4 => "qmatvec_e4m3_mmvq_fused2_b4",
15300 _ => "qmatvec_e4m3_mmvq_fused2_b8",
15301 });
15302 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
15303 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
15304 let cfg = LaunchConfig {
15305 grid_dim: (nb0 + nb1, 1, 1),
15306 block_dim: (32, ROWS_PER_BLOCK, 1),
15307 shared_mem_bytes: 0,
15308 };
15309 let (inf, o0, o1, mi, rbl) = (
15310 in_f as i32,
15311 out0 as i32,
15312 out1 as i32,
15313 m as i32,
15314 row_bytes as i64,
15315 );
15316 let __s_b = self.gpu.stream();
15317 let mut b = __s_b.launch_builder(&f);
15318 b.arg(b0)
15319 .arg(b1)
15320 .arg(aq)
15321 .arg(ad)
15322 .arg(&mut y0)
15323 .arg(&mut y1)
15324 .arg(&inf)
15325 .arg(&o0)
15326 .arg(&o1)
15327 .arg(&mi)
15328 .arg(&rbl);
15329 unsafe {
15330 b.launch(cfg)?;
15331 }
15332 if ws0 != 1.0 {
15333 self.scale_inplace(&mut y0, ws0, m * out0)?;
15334 }
15335 if ws1 != 1.0 {
15336 self.scale_inplace(&mut y1, ws1, m * out1)?;
15337 }
15338 Ok((y0, y1))
15339 }
15340
15341 #[allow(clippy::too_many_arguments)]
15343 fn e4m3_fused3_t_core(
15344 &self,
15345 b0: &CudaSlice<u8>,
15346 b1: &CudaSlice<u8>,
15347 b2: &CudaSlice<u8>,
15348 aq: &CudaSlice<i8>,
15349 ad: &CudaSlice<f32>,
15350 m: usize,
15351 in_f: usize,
15352 out0: usize,
15353 out1: usize,
15354 out2: usize,
15355 row_bytes: usize,
15356 ws0: f32,
15357 ws1: f32,
15358 ws2: f32,
15359 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15360 const ROWS_PER_BLOCK: u32 = 4;
15361 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
15362 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
15363 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
15364 let f = self.func(if Self::batched_mcols(m) == 2 {
15365 "qmatvec_e4m3_mmvq_fused3_b2"
15366 } else {
15367 "qmatvec_e4m3_mmvq_fused3_b4"
15368 });
15369 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
15370 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
15371 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
15372 let cfg = LaunchConfig {
15373 grid_dim: (nb0 + nb1 + nb2, 1, 1),
15374 block_dim: (32, ROWS_PER_BLOCK, 1),
15375 shared_mem_bytes: 0,
15376 };
15377 let (inf, o0, o1, o2, mi, rbl) = (
15378 in_f as i32,
15379 out0 as i32,
15380 out1 as i32,
15381 out2 as i32,
15382 m as i32,
15383 row_bytes as i64,
15384 );
15385 let __s_b = self.gpu.stream();
15386 let mut b = __s_b.launch_builder(&f);
15387 b.arg(b0)
15388 .arg(b1)
15389 .arg(b2)
15390 .arg(aq)
15391 .arg(ad)
15392 .arg(&mut y0)
15393 .arg(&mut y1)
15394 .arg(&mut y2)
15395 .arg(&inf)
15396 .arg(&o0)
15397 .arg(&o1)
15398 .arg(&o2)
15399 .arg(&mi)
15400 .arg(&rbl);
15401 unsafe {
15402 b.launch(cfg)?;
15403 }
15404 if ws0 != 1.0 {
15405 self.scale_inplace(&mut y0, ws0, m * out0)?;
15406 }
15407 if ws1 != 1.0 {
15408 self.scale_inplace(&mut y1, ws1, m * out1)?;
15409 }
15410 if ws2 != 1.0 {
15411 self.scale_inplace(&mut y2, ws2, m * out2)?;
15412 }
15413 Ok((y0, y1, y2))
15414 }
15415
15416 pub fn qmatvec_e4m3_blk_mmvq(
15426 &self,
15427 bytes: &CudaSlice<u8>,
15428 aq: &CudaSlice<i8>,
15429 ad: &CudaSlice<f32>,
15430 scales: &CudaSlice<f32>,
15431 m: usize,
15432 in_f: usize,
15433 out_f: usize,
15434 row_bytes: usize,
15435 scale_cols: usize,
15436 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15437 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_e4m3_blk_mmvq_into(
15439 bytes, aq, ad, scales, m, in_f, out_f, row_bytes, scale_cols, &mut y,
15440 )?;
15441 Ok(y)
15442 }
15443
15444 #[allow(clippy::too_many_arguments)]
15446 pub fn qmatvec_e4m3_blk_mmvq_into(
15447 &self,
15448 bytes: &CudaSlice<u8>,
15449 aq: &CudaSlice<i8>,
15450 ad: &CudaSlice<f32>,
15451 scales: &CudaSlice<f32>,
15452 m: usize,
15453 in_f: usize,
15454 out_f: usize,
15455 row_bytes: usize,
15456 scale_cols: usize,
15457 y: &mut CudaSlice<f32>,
15458 ) -> Result<(), Box<dyn std::error::Error>> {
15459 const ROWS_PER_BLOCK: u32 = 4; let f = self.func("qmatvec_e4m3_blk_mmvq");
15461 let cfg = LaunchConfig {
15462 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), m as u32, 1),
15463 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
15466 let (inf, outf, mi, rb, sc) = (
15467 in_f as i32,
15468 out_f as i32,
15469 m as i32,
15470 row_bytes as i64,
15471 scale_cols as i32,
15472 );
15473 let __s_b = self.gpu.stream();
15474 let mut b = __s_b.launch_builder(&f);
15475 b.arg(bytes)
15476 .arg(aq)
15477 .arg(ad)
15478 .arg(scales)
15479 .arg(&mut *y)
15480 .arg(&inf)
15481 .arg(&outf)
15482 .arg(&mi)
15483 .arg(&rb)
15484 .arg(&sc);
15485 unsafe {
15486 b.launch(cfg)?;
15487 }
15488 Ok(())
15489 }
15490
15491 #[allow(clippy::too_many_arguments)]
15497 pub fn qmatvec_e4m3_blk_mmvq_batched(
15498 &self,
15499 bytes: &CudaSlice<u8>,
15500 aq: &CudaSlice<i8>,
15501 ad: &CudaSlice<f32>,
15502 scales: &CudaSlice<f32>,
15503 m: usize,
15504 in_f: usize,
15505 out_f: usize,
15506 row_bytes: usize,
15507 scale_cols: usize,
15508 mcols: usize,
15509 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15510 const ROWS_PER_BLOCK: u32 = 4; debug_assert!(mcols >= m, "blk batched: mcols {mcols} < m {m}");
15512 let name = match mcols {
15513 2 => "qmatvec_e4m3_blk_mmvq_b2",
15514 4 => "qmatvec_e4m3_blk_mmvq_b4",
15515 8 => "qmatvec_e4m3_blk_mmvq_b8",
15516 16 => "qmatvec_e4m3_blk_mmvq_b16",
15517 _ => {
15518 return Err(
15519 format!("qmatvec_e4m3_blk_mmvq_batched: no kernel for mcols {mcols}").into(),
15520 );
15521 }
15522 };
15523 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
15524 let f = self.func(name);
15525 let cfg = LaunchConfig {
15526 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
15527 block_dim: (32, ROWS_PER_BLOCK, 1),
15528 shared_mem_bytes: 0,
15529 };
15530 let (inf, outf, mi, rb, sc) = (
15531 in_f as i32,
15532 out_f as i32,
15533 m as i32,
15534 row_bytes as i64,
15535 scale_cols as i32,
15536 );
15537 let __s_b = self.gpu.stream();
15538 let mut b = __s_b.launch_builder(&f);
15539 b.arg(bytes)
15540 .arg(aq)
15541 .arg(ad)
15542 .arg(scales)
15543 .arg(&mut y)
15544 .arg(&inf)
15545 .arg(&outf)
15546 .arg(&mi)
15547 .arg(&rb)
15548 .arg(&sc);
15549 unsafe {
15550 b.launch(cfg)?;
15551 }
15552 Ok(y)
15553 }
15554
15555 #[allow(clippy::too_many_arguments)]
15558 pub fn qmatvec_e4m3_blk_batched_raw(
15559 &self,
15560 bytes: &CudaSlice<u8>,
15561 x: &CudaSlice<f32>,
15562 scales: &CudaSlice<f32>,
15563 m: usize,
15564 in_f: usize,
15565 out_f: usize,
15566 row_bytes: usize,
15567 scale_cols: usize,
15568 mcols: usize,
15569 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15570 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15571 self.qmatvec_e4m3_blk_mmvq_batched(
15572 bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols, mcols,
15573 )
15574 }
15575
15576 #[allow(clippy::too_many_arguments)]
15579 pub fn qmatvec_e4m3_blk_mmvq_raw(
15580 &self,
15581 bytes: &CudaSlice<u8>,
15582 x: &CudaSlice<f32>,
15583 scales: &CudaSlice<f32>,
15584 m: usize,
15585 in_f: usize,
15586 out_f: usize,
15587 row_bytes: usize,
15588 scale_cols: usize,
15589 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15590 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15591 self.qmatvec_e4m3_blk_mmvq(
15592 bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols,
15593 )
15594 }
15595
15596 #[allow(clippy::too_many_arguments)]
15599 pub fn qmatvec_e4m3_fused2_raw(
15600 &self,
15601 b0: &CudaSlice<u8>,
15602 b1: &CudaSlice<u8>,
15603 x: &CudaSlice<f32>,
15604 in_f: usize,
15605 out0: usize,
15606 out1: usize,
15607 row_bytes: usize,
15608 ws0: f32,
15609 ws1: f32,
15610 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15611 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
15612 self.e4m3_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes, ws0, ws1)
15613 }
15614
15615 #[allow(clippy::too_many_arguments)]
15616 pub fn qmatvec_e4m3_fused3_raw(
15617 &self,
15618 b0: &CudaSlice<u8>,
15619 b1: &CudaSlice<u8>,
15620 b2: &CudaSlice<u8>,
15621 x: &CudaSlice<f32>,
15622 in_f: usize,
15623 out0: usize,
15624 out1: usize,
15625 out2: usize,
15626 row_bytes: usize,
15627 ws0: f32,
15628 ws1: f32,
15629 ws2: f32,
15630 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15631 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
15632 self.e4m3_fused3_core(
15633 b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes, ws0, ws1, ws2,
15634 )
15635 }
15636
15637 #[allow(clippy::too_many_arguments)]
15638 pub fn qmatvec_e4m3_fused2_t_raw(
15639 &self,
15640 b0: &CudaSlice<u8>,
15641 b1: &CudaSlice<u8>,
15642 x: &CudaSlice<f32>,
15643 m: usize,
15644 in_f: usize,
15645 out0: usize,
15646 out1: usize,
15647 row_bytes: usize,
15648 ws0: f32,
15649 ws1: f32,
15650 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15651 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15652 self.e4m3_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes, ws0, ws1)
15653 }
15654
15655 #[allow(clippy::too_many_arguments)]
15656 pub fn qmatvec_e4m3_fused3_t_raw(
15657 &self,
15658 b0: &CudaSlice<u8>,
15659 b1: &CudaSlice<u8>,
15660 b2: &CudaSlice<u8>,
15661 x: &CudaSlice<f32>,
15662 m: usize,
15663 in_f: usize,
15664 out0: usize,
15665 out1: usize,
15666 out2: usize,
15667 row_bytes: usize,
15668 ws0: f32,
15669 ws1: f32,
15670 ws2: f32,
15671 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15672 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
15673 self.e4m3_fused3_t_core(
15674 b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes, ws0, ws1, ws2,
15675 )
15676 }
15677
15678 fn try_e4m3_blk_pre(
15689 &self,
15690 w: &crate::model::GpuTensor,
15691 aq: &CudaSlice<i8>,
15692 ad: &CudaSlice<f32>,
15693 m: usize,
15694 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
15695 use crate::model::GpuTensor;
15696 if let GpuTensor::Quant {
15697 bytes,
15698 qtype,
15699 row_bytes,
15700 blk: Some(g),
15701 ..
15702 } = w
15703 {
15704 if *qtype == QT_F8_E4M3_BLK {
15705 if (2..=16).contains(&m)
15711 && std::env::var("MEMRA_NO_BATCHED").is_err()
15712 && (m <= 4 || Self::b8_enabled())
15713 {
15714 let mcols = Self::batched_mcols(m);
15715 return Ok(Some(self.qmatvec_e4m3_blk_mmvq_batched(
15716 bytes,
15717 aq,
15718 ad,
15719 &g.scales,
15720 m,
15721 w.in_features(),
15722 w.out_features(),
15723 *row_bytes,
15724 g.cols,
15725 mcols,
15726 )?));
15727 }
15728 return Ok(Some(self.qmatvec_e4m3_blk_mmvq(
15729 bytes,
15730 aq,
15731 ad,
15732 &g.scales,
15733 m,
15734 w.in_features(),
15735 w.out_features(),
15736 *row_bytes,
15737 g.cols,
15738 )?));
15739 }
15740 }
15741 Ok(None)
15742 }
15743
15744 fn try_e4m3_blk_prefill(
15791 &self,
15792 w: &crate::model::GpuTensor,
15793 x: &CudaSlice<f32>,
15794 m: usize,
15795 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
15796 use crate::model::GpuTensor;
15797 let GpuTensor::Quant {
15798 bytes,
15799 qtype,
15800 blk: Some(g),
15801 ..
15802 } = w
15803 else {
15804 return Ok(None);
15805 };
15806 if *qtype != QT_F8_E4M3_BLK {
15807 return Ok(None);
15808 }
15809 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? {
15814 return Ok(Some(y));
15815 }
15816 let (in_f, out_f) = (w.in_features(), w.out_features());
15817 let slab = self.fp8_blk_dequant_q8_0_dev(bytes, &g.scales, out_f, in_f)?;
15818 let tmp = GpuTensor::Quant {
15819 bytes: slab,
15820 qtype: QT_Q8_0,
15821 row_bytes: in_f / 32 * 34,
15822 ne: vec![in_f as u64, out_f as u64],
15823 scale: 1.0,
15824 rp: false,
15825 #[cfg(memra_cutlass)]
15826 cutlass: None,
15827 fp8: None,
15828 blk: None,
15829 f16: None,
15830 rp4: None,
15831 };
15832 Ok(Some(self.matmul(&tmp, x, m)?))
15834 }
15835
15836 pub fn matmul_pre_noscale(
15837 &self,
15838 w: &crate::model::GpuTensor,
15839 aq: &CudaSlice<i8>,
15840 ad: &CudaSlice<f32>,
15841 m: usize,
15842 ) -> Result<Option<(CudaSlice<f32>, f32)>, Box<dyn std::error::Error>> {
15843 use crate::model::GpuTensor;
15844 if m == 1 {
15848 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
15849 return Ok(Some((y, 1.0)));
15850 }
15851 }
15852 if m != 1 || !self.uses_q8_1_fast(w) {
15854 return Ok(None);
15855 }
15856 let in_f = w.in_features();
15857 let out_f = w.out_features();
15858 let (bytes, qtype, row_bytes, scale, rp) = match w {
15859 GpuTensor::Quant {
15860 bytes,
15861 qtype,
15862 row_bytes,
15863 scale,
15864 rp,
15865 ..
15866 } => (bytes, *qtype, *row_bytes, *scale, *rp),
15867 _ => return Ok(None),
15868 };
15869 if self.mmvq_supports(qtype) {
15871 let (mbytes, mrp) = match w {
15873 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
15874 _ => (bytes, rp),
15875 };
15876 let y = self.qmatvec_mmvq(
15877 mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, 1.0, mrp,
15878 )?;
15879 return Ok(Some((y, scale)));
15880 }
15881 let name = match qtype {
15883 QT_Q8_0 => "qmatvec_q8_0_dp4a",
15884 QT_Q4_K => "qmatvec_q4_K_dp4a",
15885 QT_Q6_K => "qmatvec_q6_K_dp4a",
15886 QT_Q5_K => "qmatvec_q5_K_dp4a",
15887 QT_Q3_K => "qmatvec_q3_K_dp4a",
15888 QT_NVFP4 => {
15889 if rp {
15890 "qmatvec_nvfp4_dp4a_rp"
15891 } else {
15892 "qmatvec_nvfp4_dp4a"
15893 }
15894 }
15895 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
15896 _ => return Ok(None),
15897 };
15898 let f = self.func(name);
15899 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
15900 let cfg = LaunchConfig {
15901 grid_dim: (out_f as u32, m as u32, 1),
15902 block_dim: (128, 1, 1),
15903 shared_mem_bytes: 0,
15904 };
15905 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
15906 let __s_b = self.gpu.stream();
15907 let mut b = __s_b.launch_builder(&f);
15908 b.arg(bytes)
15909 .arg(aq)
15910 .arg(ad)
15911 .arg(&mut y)
15912 .arg(&inf)
15913 .arg(&outf)
15914 .arg(&mi)
15915 .arg(&rb);
15916 unsafe {
15917 b.launch(cfg)?;
15918 }
15919 Ok(Some((y, scale)))
15920 }
15921
15922 pub fn mmvq_supports(&self, qtype: i32) -> bool {
15925 if qtype == QT_F8_E4M3 {
15930 return true;
15931 }
15932 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") {
15933 return false;
15934 }
15935 matches!(
15936 qtype,
15937 QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_Q4_0
15938 )
15939 }
15940
15941 pub fn qmatvec_mmvq(
15946 &self,
15947 bytes: &CudaSlice<u8>,
15948 aq: &CudaSlice<i8>,
15949 ad: &CudaSlice<f32>,
15950 m: usize,
15951 in_f: usize,
15952 out_f: usize,
15953 qtype: i32,
15954 row_bytes: usize,
15955 scale: f32,
15956 rp: bool,
15957 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15958 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_mmvq_into(
15960 bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp, &mut y,
15961 )?;
15962 Ok(y)
15963 }
15964
15965 #[allow(clippy::too_many_arguments)]
15967 pub fn qmatvec_mmvq_into(
15968 &self,
15969 bytes: &CudaSlice<u8>,
15970 aq: &CudaSlice<i8>,
15971 ad: &CudaSlice<f32>,
15972 m: usize,
15973 in_f: usize,
15974 out_f: usize,
15975 qtype: i32,
15976 row_bytes: usize,
15977 scale: f32,
15978 rp: bool,
15979 y: &mut CudaSlice<f32>,
15980 ) -> Result<(), Box<dyn std::error::Error>> {
15981 debug_assert!(y.len() >= m * out_f);
15982 const ROWS_PER_BLOCK: u32 = 4; if qtype == QT_Q8_0
15988 && rp
15989 && m == 1
15990 && out_f >= 64
15991 && (out_f as u32).div_ceil(ROWS_PER_BLOCK) < 4 * self.sm_count() as u32
15992 && {
15993 static G2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15994 *G2.get_or_init(|| std::env::var("MEMRA_Q80_G2").as_deref() != Ok("0"))
15995 }
15996 {
15997 let f = self.func("qmatvec_q8_0_mmvq_rp_g2");
15998 let cfg = LaunchConfig {
15999 grid_dim: ((out_f as u32).div_ceil(2), 1, 1),
16000 block_dim: (32, 2, 1),
16001 shared_mem_bytes: 0,
16002 };
16003 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, row_bytes as i64);
16004 let __s_b = self.gpu.stream();
16005 let mut b = __s_b.launch_builder(&f);
16006 b.arg(bytes)
16007 .arg(aq)
16008 .arg(ad)
16009 .arg(&mut *y)
16010 .arg(&inf)
16011 .arg(&outf)
16012 .arg(&mi)
16013 .arg(&rb);
16014 unsafe {
16015 b.launch(cfg)?;
16016 }
16017 if scale != 1.0 {
16018 self.scale_inplace(y, scale, out_f)?;
16019 }
16020 return Ok(());
16021 }
16022 let mut mr: u32 = if m == 1 && (qtype == QT_NVFP4 || qtype == QT_Q5_K) {
16031 2
16032 } else {
16033 1
16034 };
16035 if m == 1 && qtype == QT_Q4_0 {
16040 static Q40MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
16041 mr = *Q40MR.get_or_init(|| {
16044 std::env::var("MEMRA_Q40_MR")
16045 .ok()
16046 .and_then(|v| v.parse().ok())
16047 .unwrap_or(1)
16048 });
16049 }
16050 let q5_mode = std::env::var("MEMRA_Q5K_ISSUE").ok();
16061 let q5_force = q5_mode.as_deref() == Some("2");
16062 let q5_il = qtype == QT_Q5_K
16065 && m == 1
16066 && (q5_force || q5_mode.as_deref().map(|v| v != "0").unwrap_or(true));
16067 if q5_il && !q5_force && out_f > 65536 {
16068 mr = 1;
16069 }
16070 if qtype == QT_Q4_0 && rp && mr != 1 {
16073 mr = 2;
16074 }
16075 if qtype == QT_Q8_0 && rp {
16079 static Q80MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
16080 mr = *Q80MR.get_or_init(|| {
16081 std::env::var("MEMRA_Q80_MR")
16082 .ok()
16083 .and_then(|v| v.parse().ok())
16084 .unwrap_or(1)
16085 });
16086 }
16087 let name = match (qtype, mr, rp) {
16088 (QT_NVFP4, 2, false) => "qmatvec_nvfp4_mmvq_mr2",
16089 (QT_NVFP4, 2, true) => "qmatvec_nvfp4_mmvq_mr2_rp",
16090 (QT_NVFP4, _, true) => "qmatvec_nvfp4_mmvq_rp",
16091 (QT_Q4_0, 1, true) => "qmatvec_q4_0_mmvq_rp",
16092 (QT_Q4_0, _, true) => "qmatvec_q4_0_mmvq_mr2_rp",
16093 (QT_Q5_K, 2, _) => {
16094 if q5_il {
16095 "qmatvec_q5_K_mmvq_mr2_il"
16096 } else {
16097 "qmatvec_q5_K_mmvq_mr2"
16098 }
16099 }
16100 (QT_Q8_0, 2, true) => "qmatvec_q8_0_mmvq_mr2_rp",
16101 (QT_Q8_0, _, true)
16106 if in_f % 1024 == 0 && {
16107 static CA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16108 *CA.get_or_init(|| std::env::var("MEMRA_Q80_CA").as_deref() == Ok("1"))
16109 } =>
16110 {
16111 "qmatvec_q8_0_mmvq_rpca"
16112 }
16113 (QT_Q8_0, _, true) => "qmatvec_q8_0_mmvq_rp",
16114 (QT_Q8_0, _, _) => "qmatvec_q8_0_mmvq",
16115 (QT_Q4_K, _, true) => "qmatvec_q4_K_mmvq_rp",
16119 (QT_Q6_K, _, true) => "qmatvec_q6_K_mmvq_rp",
16120 (QT_Q4_K, _, _) => "qmatvec_q4_K_mmvq",
16121 (QT_Q4_0, 2, false) => "qmatvec_q4_0_mmvq_mr2",
16122 (QT_Q4_0, _, false) => "qmatvec_q4_0_mmvq",
16123 (QT_Q5_K, _, _) => {
16124 if q5_il {
16125 "qmatvec_q5_K_mmvq_il"
16126 } else {
16127 "qmatvec_q5_K_mmvq"
16128 }
16129 }
16130 (QT_Q6_K, _, _) => "qmatvec_q6_K_mmvq",
16131 (QT_NVFP4, _, false) => "qmatvec_nvfp4_mmvq",
16132 (QT_F8_E4M3, _, _) => "qmatvec_e4m3_mmvq",
16133 _ => panic!("qmatvec_mmvq: qtype {qtype} has no MMVQ kernel"),
16134 };
16135 let f = self.func(name);
16136 let rows_per_block = ROWS_PER_BLOCK * mr;
16138 let cfg = LaunchConfig {
16139 grid_dim: (
16140 (out_f as u32 + rows_per_block - 1) / rows_per_block,
16141 m as u32,
16142 1,
16143 ),
16144 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
16147 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
16148 let __s_b = self.gpu.stream();
16149 let mut b = __s_b.launch_builder(&f);
16150 if qtype == QT_NVFP4 || qtype == QT_F8_E4M3 {
16155 if Self::pdl_on()
16158 && Self::pdl_mmvq_on()
16159 && Self::pdl_nvfp4q8_on()
16160 && name == "qmatvec_nvfp4_mmvq_mr2_rp"
16161 {
16162 use cudarc::driver::{DevicePtr, DevicePtrMut};
16163 let s = &self.gpu.stream();
16164 let (pw, _g0) = bytes.device_ptr(s);
16165 let (paq, _g1) = aq.device_ptr(s);
16166 let (pad, _g2) = ad.device_ptr(s);
16167 let (py, _g3) = y.device_ptr_mut(s);
16168 let mut ps = [
16169 &pw as *const _ as *mut std::ffi::c_void,
16170 &paq as *const _ as *mut _,
16171 &pad as *const _ as *mut _,
16172 &py as *const _ as *mut _,
16173 &inf as *const _ as *mut _,
16174 &outf as *const _ as *mut _,
16175 &mi as *const _ as *mut _,
16176 &rb as *const _ as *mut _,
16177 &scale as *const _ as *mut _,
16178 ];
16179 unsafe {
16180 self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?;
16181 }
16182 return Ok(());
16183 }
16184 b.arg(bytes)
16185 .arg(aq)
16186 .arg(ad)
16187 .arg(&mut *y)
16188 .arg(&inf)
16189 .arg(&outf)
16190 .arg(&mi)
16191 .arg(&rb)
16192 .arg(&scale);
16193 unsafe {
16194 b.launch(cfg)?;
16195 }
16196 } else if Self::pdl_on()
16197 && Self::pdl_mmvq_on()
16198 && (matches!(
16199 name,
16200 "qmatvec_q4_0_mmvq_rp" | "qmatvec_q6_K_mmvq" | "qmatvec_q6_K_mmvq_rp"
16201 ) || (Self::pdl_nvfp4q8_on()
16202 && matches!(name, "qmatvec_q8_0_mmvq_rp" | "qmatvec_q8_0_mmvq_mr2_rp")))
16203 {
16204 {
16208 use cudarc::driver::{DevicePtr, DevicePtrMut};
16209 let s = &self.gpu.stream();
16210 let (pw, _g0) = bytes.device_ptr(s);
16211 let (paq, _g1) = aq.device_ptr(s);
16212 let (pad, _g2) = ad.device_ptr(s);
16213 let (py, _g3) = y.device_ptr_mut(s);
16214 let mut ps = [
16215 &pw as *const _ as *mut std::ffi::c_void,
16216 &paq as *const _ as *mut _,
16217 &pad as *const _ as *mut _,
16218 &py as *const _ as *mut _,
16219 &inf as *const _ as *mut _,
16220 &outf as *const _ as *mut _,
16221 &mi as *const _ as *mut _,
16222 &rb as *const _ as *mut _,
16223 ];
16224 unsafe {
16225 self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?;
16226 }
16227 }
16228 if scale != 1.0 {
16229 self.scale_inplace(y, scale, m * out_f)?;
16230 }
16231 } else {
16232 b.arg(bytes)
16233 .arg(aq)
16234 .arg(ad)
16235 .arg(&mut *y)
16236 .arg(&inf)
16237 .arg(&outf)
16238 .arg(&mi)
16239 .arg(&rb);
16240 unsafe {
16241 b.launch(cfg)?;
16242 }
16243 if scale != 1.0 {
16244 self.scale_inplace(y, scale, m * out_f)?;
16245 }
16246 }
16247 Ok(())
16248 }
16249
16250 pub fn qmatvec_mmvq_raw(
16254 &self,
16255 bytes: &CudaSlice<u8>,
16256 x: &CudaSlice<f32>,
16257 m: usize,
16258 in_f: usize,
16259 out_f: usize,
16260 qtype: i32,
16261 row_bytes: usize,
16262 rp: bool,
16263 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16264 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
16265 self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, 1.0, rp)
16266 }
16267
16268 pub fn batched_supports(&self, qtype: i32) -> bool {
16272 matches!(
16273 qtype,
16274 QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_F8_E4M3 | QT_Q4_0
16275 )
16276 }
16277
16278 pub fn iq_fast_enabled() -> bool {
16286 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16287 *ON.get_or_init(|| {
16288 std::env::var("MEMRA_IQ_FAST")
16289 .map(|v| v != "0")
16290 .unwrap_or(true)
16291 })
16292 }
16293
16294 pub fn b8_enabled() -> bool {
16297 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16298 *ON.get_or_init(|| std::env::var("MEMRA_B8").map(|v| v != "0").unwrap_or(true))
16299 }
16300
16301 pub fn batched_mcols(m: usize) -> usize {
16303 if m == 2 {
16304 2
16305 } else if m <= 4 {
16306 4
16307 } else if m <= 8 {
16308 8
16309 } else {
16310 16
16311 }
16312 }
16313
16314 fn batched_kernel_name(qtype: i32, mcols: usize) -> Option<&'static str> {
16319 Some(match (qtype, mcols) {
16320 (QT_Q8_0, 2) => "qmatvec_q8_0_mmvq_b2",
16321 (QT_Q8_0, 4) => "qmatvec_q8_0_mmvq_b4",
16322 (QT_Q8_0, 8) => "qmatvec_q8_0_mmvq_b8",
16323 (QT_Q8_0, 16) => "qmatvec_q8_0_mmvq_b16",
16329 (QT_Q4_K, 2) => "qmatvec_q4_K_mmvq_b2",
16330 (QT_Q4_K, 4) => "qmatvec_q4_K_mmvq_b4",
16331 (QT_Q4_K, 8) => "qmatvec_q4_K_mmvq_b8",
16332 (QT_Q4_K, 16) => "qmatvec_q4_K_mmvq_b16",
16335 (QT_Q5_K, 2) => "qmatvec_q5_K_mmvq_b2",
16336 (QT_Q5_K, 4) => "qmatvec_q5_K_mmvq_b4",
16337 (QT_Q5_K, 8) => "qmatvec_q5_K_mmvq_b8",
16338 (QT_Q5_K, 16) => "qmatvec_q5_K_mmvq_b16",
16341 (QT_Q6_K, 2) => "qmatvec_q6_K_mmvq_b2",
16342 (QT_Q6_K, 4) => "qmatvec_q6_K_mmvq_b4",
16343 (QT_Q6_K, 8) => "qmatvec_q6_K_mmvq_b8",
16344 (QT_Q6_K, 16) => "qmatvec_q6_K_mmvq_b16",
16345 (QT_NVFP4, 2) => "qmatvec_nvfp4_mmvq_b2",
16346 (QT_NVFP4, 4) => "qmatvec_nvfp4_mmvq_b4",
16347 (QT_NVFP4, 8) => "qmatvec_nvfp4_mmvq_b8",
16348 (QT_NVFP4, 16) => "qmatvec_nvfp4_mmvq_b16",
16352 (QT_F8_E4M3, 2) => "qmatvec_e4m3_mmvq_b2",
16353 (QT_F8_E4M3, 4) => "qmatvec_e4m3_mmvq_b4",
16354 (QT_F8_E4M3, 8) => "qmatvec_e4m3_mmvq_b8",
16355 (QT_F8_E4M3, 16) => "qmatvec_e4m3_mmvq_b16",
16359 (QT_Q4_0, 2) => "qmatvec_q4_0_mmvq_b2",
16360 (QT_Q4_0, 4) => "qmatvec_q4_0_mmvq_b4",
16361 (QT_Q4_0, 8) => "qmatvec_q4_0_mmvq_b8",
16362 (QT_Q4_0, 16) => "qmatvec_q4_0_mmvq_b16",
16363 _ => return None,
16364 })
16365 }
16366
16367 pub fn sm_count(&self) -> i32 {
16402 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
16403 *SMS.get_or_init(|| {
16404 use cudarc::driver::sys::CUdevice_attribute_enum as A;
16405 self.gpu
16406 .ctx
16407 .attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT)
16408 .unwrap_or(82)
16409 })
16410 }
16411
16412 pub fn batched_variant(
16413 &self,
16414 _m: usize,
16415 in_f: usize,
16416 out_f: usize,
16417 qtype: i32,
16418 row_bytes: usize,
16419 mcols: usize,
16420 rp: bool,
16421 ) -> &'static str {
16422 if qtype == QT_Q8_0 {
16427 return if rp { "rp" } else { "base" };
16428 }
16429 static BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
16430 let bv = *BV.get_or_init(|| match std::env::var("MEMRA_MMVQ_BV").as_deref() {
16431 Ok("base") => "base",
16432 Ok("pf") => "pf",
16433 Ok("r2") => "r2",
16434 Ok("r2w8") => "r2w8",
16435 Ok("pfr2") => "pfr2",
16436 Ok("ca") => "ca",
16437 Ok("car2") => "car2",
16438 Ok("rp") => "rp",
16441 Ok("rpr2") => "rpr2",
16442 Ok("rpr2w8") => "rpr2w8",
16443 Ok("rpca") => "rpca",
16446 Ok("rpcar2") => "rpcar2",
16447 Ok("rpsc") => "rpsc",
16454 Ok("rpms") => "rpms",
16455 Ok("rpmsc") => "rpmsc",
16456 Ok("rpks") => "rpks",
16457 Ok("rpksc") => "rpksc",
16458 _ => "auto",
16459 });
16460 let ca_ok = qtype == QT_NVFP4 && (row_bytes % 16 == 0) && (in_f % 1024 == 0);
16464 static KS_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16469 let ks_on = *KS_ON.get_or_init(|| std::env::var("MEMRA_KS").as_deref() != Ok("0"));
16470 let sc_ok = ks_on && qtype == QT_NVFP4 && (in_f % 256 == 0) && (in_f / 64 <= 272);
16471 let ks_ok = ks_on && qtype == QT_NVFP4 && (in_f % 512 == 0) && (in_f / 64 <= 272);
16472 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
16473 let sms = *SMS.get_or_init(|| {
16474 use cudarc::driver::sys::CUdevice_attribute_enum as A;
16475 self.gpu
16476 .ctx
16477 .attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT)
16478 .unwrap_or(82)
16479 });
16480 let kq_r2 = matches!(qtype, QT_Q4_K | QT_Q5_K | QT_Q6_K);
16500 static KQBV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
16503 let kq_bv = *KQBV.get_or_init(|| match std::env::var("MEMRA_KQ_BV").as_deref() {
16504 Ok("base") => "base",
16505 Ok("r2") => "r2",
16506 Ok("r2w8") => "r2w8",
16507 _ => "auto",
16508 });
16509 let variant: &'static str = if qtype == QT_Q4_0 {
16510 static Q40BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
16514 let q40 = *Q40BV.get_or_init(|| match std::env::var("MEMRA_Q40_BV").as_deref() {
16515 Ok("base") => "base",
16521 Ok("r2") => "r2",
16522 Ok("ms") => "ms",
16523 Ok("sm") => "sm",
16524 Ok("la") => "la",
16525 _ => "auto",
16526 });
16527 let v = if q40 != "auto" {
16528 q40
16529 } else if (out_f as u32).div_ceil(8) >= 4 * sms as u32 {
16530 "r2"
16531 } else {
16532 "base"
16533 };
16534 if rp {
16539 match v {
16540 "ms" => "r2ms_rp",
16541 "sm" => "r2sm_rp",
16542 "la" => "r2la_rp",
16543 "r2" => "r2_rp",
16544 _ => "rp",
16545 }
16546 } else if matches!(v, "ms" | "sm" | "la") {
16547 "r2"
16548 } else {
16549 v
16550 }
16551 } else if qtype != QT_NVFP4 && !kq_r2 {
16552 "base"
16553 } else if kq_r2 && rp {
16554 "rp"
16558 } else if kq_r2 {
16559 if kq_bv != "auto" {
16562 if kq_bv == "r2w8" && mcols != 4 {
16563 "r2"
16564 } else {
16565 kq_bv
16566 }
16567 } else if bv != "auto" {
16568 match bv {
16569 "r2" | "pfr2" | "rpr2" | "car2" => "r2",
16570 "r2w8" | "rpr2w8" => {
16571 if mcols != 4 {
16572 "r2"
16573 } else {
16574 "r2w8"
16575 }
16576 }
16577 _ => "base", }
16579 } else {
16580 let blocks = (out_f + 7) / 8;
16581 let waves = blocks as f64 / (7 * sms as usize) as f64;
16582 let filled = blocks >= 4 * sms as usize;
16583 let use_r2 = if qtype == QT_Q4_K {
16584 filled
16585 } else {
16586 waves >= 2.0
16587 };
16588 if use_r2 { "r2" } else { "base" }
16589 }
16590 } else if bv != "auto" {
16591 let v = if bv == "r2w8" && mcols == 2 {
16596 "r2"
16597 } else if bv == "ca" && (!ca_ok || mcols == 8) {
16598 "pf"
16599 } else if bv == "car2" && (!ca_ok || mcols == 8) {
16600 "r2"
16601 } else if bv == "pfr2" && mcols == 8 {
16602 "r2"
16603 } else if (bv == "rpr2w8" || bv == "rpr2") && mcols == 2 {
16604 "rpr2"
16605 }
16606 else if (bv == "rpca" || bv == "rpcar2") && (!ca_ok || mcols == 8) {
16608 if mcols == 8 { "rpr2w8" } else { "rpr2" }
16609 } else if bv == "rpcar2" && mcols == 2 {
16610 "rpca"
16611 }
16612 else if (bv == "rpsc" || bv == "rpmsc") && !sc_ok {
16615 "rpr2"
16616 } else if (bv == "rpks" || bv == "rpksc") && !ks_ok {
16617 "rpr2"
16618 } else {
16619 bv
16620 };
16621 if rp {
16622 match v {
16623 "base" | "pf" | "ca" | "rp" => "rp",
16624 "r2" | "pfr2" | "car2" | "rpr2" => "rpr2",
16625 "r2w8" | "rpr2w8" => {
16626 if mcols == 2 {
16627 "rpr2"
16628 } else {
16629 "rpr2w8"
16630 }
16631 }
16632 other => other, }
16634 } else {
16635 v
16636 }
16637 } else if mcols == 8 {
16638 if rp {
16649 if sc_ok { "rpsc" } else { "rpr2w8" }
16650 } else {
16651 "r2w8"
16652 }
16653 } else if mcols >= 4 {
16654 let blocks = (out_f + 7) / 8;
16658 let r7 = 7 * sms as usize;
16659 let r8 = 8 * sms as usize;
16660 let waves = blocks as f64 / r7 as f64;
16661 let filled = blocks >= 4 * sms as usize;
16662 if filled && blocks.div_ceil(r8) < blocks.div_ceil(r7) {
16666 if rp { "rpr2w8" } else { "r2w8" }
16670 } else if waves >= 2.0 || (waves <= 1.0 && filled) {
16671 if rp { "rpr2" } else { "r2" }
16674 } else {
16675 if rp { "rp" } else { "pf" }
16679 }
16680 } else if in_f >= 6144 {
16681 if rp { "rpr2" } else { "r2" }
16685 } else if rp {
16686 let waves = ((out_f + 7) / 8) as f64 / (7 * sms as usize) as f64;
16691 if sc_ok && waves >= 0.9 && waves <= 1.1 {
16692 "rpsc"
16693 } else {
16694 "rp"
16695 }
16696 } else {
16697 "base"
16698 };
16699 variant
16700 }
16701
16702 pub fn qmatvec_mmvq_batched(
16703 &self,
16704 bytes: &CudaSlice<u8>,
16705 aq: &CudaSlice<i8>,
16706 ad: &CudaSlice<f32>,
16707 m: usize,
16708 in_f: usize,
16709 out_f: usize,
16710 qtype: i32,
16711 row_bytes: usize,
16712 mcols: usize,
16713 scale: f32,
16714 rp: bool,
16715 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16716 const ROWS_PER_BLOCK: u32 = 4;
16717 let forced: Option<&'static str> = {
16722 static V: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
16723 V.get_or_init(|| std::env::var("MEMRA_BVAR").ok())
16724 .as_deref()
16725 .map(|s| Box::leak(s.to_string().into_boxed_str()) as &'static str)
16726 };
16727 let variant = match forced {
16728 Some(v) if !rp || v.contains("rp") => v,
16729 _ => self.batched_variant(m, in_f, out_f, qtype, row_bytes, mcols, rp),
16730 };
16731 let base_name = Self::batched_kernel_name(qtype, mcols).ok_or_else(|| {
16732 format!("qmatvec_mmvq_batched: no kernel for qtype {qtype} mcols {mcols}")
16733 })?;
16734 let variant = if mcols == 16 {
16738 if rp { "rp" } else { "base" }
16739 } else {
16740 variant
16741 };
16742 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16749 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
16750 if b567
16751 && qtype == QT_NVFP4
16752 && rp
16753 && mcols == 8
16754 && (5..=7).contains(&m)
16755 && matches!(variant, "rpsc" | "rpr2w8")
16756 {
16757 let f = self.func(&format!("qmatvec_nvfp4_mmvq_b{m}_{variant}"));
16758 let rows_per_block = ROWS_PER_BLOCK * 2; let mut y = self.alloc_uninit::<f32>(m * out_f)?;
16760 let cfg = LaunchConfig {
16761 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
16762 block_dim: (32, ROWS_PER_BLOCK, 1),
16763 shared_mem_bytes: 0,
16764 };
16765 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
16766 let __s_b = self.gpu.stream();
16767 let mut b = __s_b.launch_builder(&f);
16768 b.arg(bytes)
16769 .arg(aq)
16770 .arg(ad)
16771 .arg(&mut y)
16772 .arg(&inf)
16773 .arg(&outf)
16774 .arg(&mi)
16775 .arg(&rb);
16776 unsafe {
16777 b.launch(cfg)?;
16778 }
16779 if scale != 1.0 {
16780 self.scale_inplace(&mut y, scale, m * out_f)?;
16781 }
16782 return Ok(y);
16783 }
16784 let (name, rows_per_block): (std::borrow::Cow<'static, str>, u32) = match variant {
16785 "base" => (base_name.into(), ROWS_PER_BLOCK),
16786 "pf" => (format!("{base_name}_pf").into(), ROWS_PER_BLOCK),
16787 "ca" => (format!("{base_name}_ca").into(), ROWS_PER_BLOCK),
16788 "rp" => (format!("{base_name}_rp").into(), ROWS_PER_BLOCK),
16789 "rpca" => (format!("{base_name}_rpca").into(), ROWS_PER_BLOCK), "rpks" => (format!("{base_name}_rpks").into(), ROWS_PER_BLOCK),
16793 "rpksc" => (format!("{base_name}_rpksc").into(), ROWS_PER_BLOCK),
16794 "rpms" => (format!("{base_name}_rpms").into(), ROWS_PER_BLOCK),
16795 "rpmsc" => (format!("{base_name}_rpmsc").into(), ROWS_PER_BLOCK),
16796 "r2ms_rp" => (format!("{base_name}_r2ms_rp").into(), ROWS_PER_BLOCK),
16797 "r2sm_rp" => (format!("{base_name}_r2sm_rp").into(), ROWS_PER_BLOCK * 2),
16798 "r2la_rp" => (format!("{base_name}_r2la_rp").into(), ROWS_PER_BLOCK * 2),
16799 v => (format!("{base_name}_{v}").into(), ROWS_PER_BLOCK * 2), };
16801 debug_assert!(
16802 !rp || name.contains("_rp"),
16803 "rp weight dispatched to a GGUF-layout kernel"
16804 );
16805 let f = self.func(&name);
16806 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
16807 let smem = if name.contains("_r2sm_rp") {
16809 (mcols * 32 * 9 * 4 + mcols * 32 * 4) as u32
16810 } else {
16811 0
16812 };
16813 let cfg = LaunchConfig {
16814 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
16815 block_dim: (32, ROWS_PER_BLOCK, 1),
16816 shared_mem_bytes: smem,
16817 };
16818 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
16819 let __s_b = self.gpu.stream();
16820 let mut b = __s_b.launch_builder(&f);
16821 b.arg(bytes)
16822 .arg(aq)
16823 .arg(ad)
16824 .arg(&mut y)
16825 .arg(&inf)
16826 .arg(&outf)
16827 .arg(&mi)
16828 .arg(&rb);
16829 unsafe {
16830 b.launch(cfg)?;
16831 }
16832 if scale != 1.0 {
16833 self.scale_inplace(&mut y, scale, m * out_f)?;
16834 }
16835 Ok(y)
16836 }
16837
16838 pub fn qmatvec_batched_raw(
16842 &self,
16843 bytes: &CudaSlice<u8>,
16844 x: &CudaSlice<f32>,
16845 m: usize,
16846 in_f: usize,
16847 out_f: usize,
16848 qtype: i32,
16849 row_bytes: usize,
16850 mcols: usize,
16851 rp: bool,
16852 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16853 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
16854 self.qmatvec_mmvq_batched(
16855 bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, 1.0, rp,
16856 )
16857 }
16858
16859 pub fn qmatvec_nvfp4_batched_raw(
16861 &self,
16862 bytes: &CudaSlice<u8>,
16863 x: &CudaSlice<f32>,
16864 m: usize,
16865 in_f: usize,
16866 out_f: usize,
16867 row_bytes: usize,
16868 mcols: usize,
16869 rp: bool,
16870 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16871 self.qmatvec_batched_raw(bytes, x, m, in_f, out_f, QT_NVFP4, row_bytes, mcols, rp)
16872 }
16873
16874 fn try_fp4_gemm(
16878 &self,
16879 w: &crate::model::GpuTensor,
16880 x: &CudaSlice<f32>,
16881 m: usize,
16882 in_f: usize,
16883 out_f: usize,
16884 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
16885 use crate::model::GpuTensor;
16886 if cfg!(memra_portable_cuda) {
16887 return Ok(None);
16888 }
16889 if std::env::var("MEMRA_FP4").is_ok() {
16892 refuse_portable_force("MEMRA_FP4", "the sm_120a mxf4 block-scale MMA");
16893 }
16894 if std::env::var("MEMRA_FP4").is_err() {
16895 return Ok(None);
16896 }
16897 #[cfg(memra_cutlass)]
16906 if m >= 128 && std::env::var("MEMRA_FP4_CUTLASS").is_ok() {
16907 if let GpuTensor::Quant {
16908 bytes,
16909 qtype,
16910 scale,
16911 row_bytes,
16912 cutlass,
16913 ..
16914 } = w
16915 {
16916 if *qtype == QT_NVFP4 && in_f % 64 == 0 {
16917 if let Some(cw) = cutlass {
16918 let y = self.cutlass_fp4_gemm(
16920 &cw.b_packed,
16921 &cw.sfb_swizzled,
16922 x,
16923 *scale,
16924 m,
16925 out_f,
16926 in_f,
16927 )?;
16928 return Ok(Some(y));
16929 } else if std::env::var("MEMRA_FP4_CUTLASS_OTF").is_ok() {
16930 let (b_packed, sfb_sw) =
16935 self.build_cutlass_weight(bytes, out_f, in_f, *row_bytes)?;
16936 let y =
16937 self.cutlass_fp4_gemm(&b_packed, &sfb_sw, x, *scale, m, out_f, in_f)?;
16938 return Ok(Some(y));
16939 }
16940 }
16941 }
16942 }
16943 if let GpuTensor::Quant {
16944 bytes,
16945 qtype,
16946 row_bytes,
16947 scale,
16948 rp,
16949 ..
16950 } = w
16951 {
16952 if *qtype == QT_NVFP4 && in_f % 64 == 0 && !*rp {
16955 let y =
16956 self.qmatvec_gemm_nvfp4_fp4(bytes, x, m, in_f, out_f, *row_bytes, *scale)?;
16957 return Ok(Some(y));
16958 }
16959 }
16960 Ok(None)
16961 }
16962
16963 pub fn rms_norm_f16out(
16967 &self,
16968 x: &CudaSlice<f32>,
16969 w: &CudaSlice<f32>,
16970 dst: &mut CudaSlice<f32>,
16971 dst16: &mut CudaSlice<u8>,
16972 ncols: usize,
16973 nrows: usize,
16974 eps: f32,
16975 ) -> Result<(), Box<dyn std::error::Error>> {
16976 let f = self.func("rms_norm_f16out_f32");
16977 let cfg = LaunchConfig {
16978 grid_dim: (nrows as u32, 1, 1),
16979 block_dim: (rms_block(), 1, 1),
16980 shared_mem_bytes: 0,
16981 };
16982 let (nc, e) = (ncols as i32, eps);
16983 let __s_b = self.gpu.stream();
16984 let mut b = __s_b.launch_builder(&f);
16985 b.arg(x).arg(w).arg(dst).arg(dst16).arg(&nc).arg(&e);
16986 unsafe {
16987 b.launch(cfg)?;
16988 }
16989 Ok(())
16990 }
16991
16992 #[allow(clippy::too_many_arguments)]
16995 pub fn add_rms_norm_f16out(
16996 &self,
16997 a: &CudaSlice<f32>,
16998 b: &CudaSlice<f32>,
16999 w: &CudaSlice<f32>,
17000 res: &mut CudaSlice<f32>,
17001 dst: &mut CudaSlice<f32>,
17002 dst16: &mut CudaSlice<u8>,
17003 ncols: usize,
17004 nrows: usize,
17005 eps: f32,
17006 ) -> Result<(), Box<dyn std::error::Error>> {
17007 let f = self.func("add_rms_norm_f16out_f32");
17008 let cfg = LaunchConfig {
17009 grid_dim: (nrows as u32, 1, 1),
17010 block_dim: (rms_block(), 1, 1),
17011 shared_mem_bytes: 0,
17012 };
17013 let (nc, e) = (ncols as i32, eps);
17014 let __s_lb = self.gpu.stream();
17015 let mut lb = __s_lb.launch_builder(&f);
17016 lb.arg(a)
17017 .arg(b)
17018 .arg(w)
17019 .arg(res)
17020 .arg(dst)
17021 .arg(dst16)
17022 .arg(&nc)
17023 .arg(&e);
17024 unsafe {
17025 lb.launch(cfg)?;
17026 }
17027 Ok(())
17028 }
17029
17030 pub fn matmul_group_xh(
17033 &self,
17034 ws: &[&crate::model::GpuTensor],
17035 x: &CudaSlice<f32>,
17036 xh: &CudaSlice<u8>,
17037 m: usize,
17038 ) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
17039 let mut out = Vec::with_capacity(ws.len());
17040 let in_f = ws[0].in_features();
17041 for w in ws {
17042 if w.in_features() == in_f && m >= 16 && !self.verify_exact_on() {
17043 if let Some(y) = self.try_f16_gemm_pre(w, xh, m)? {
17044 out.push(y);
17045 continue;
17046 }
17047 }
17048 out.push(self.matmul(w, x, m)?);
17049 }
17050 Ok(out)
17051 }
17052
17053 pub fn gdn_pad_mask(
17056 &self,
17057 beta: &mut CudaSlice<f32>,
17058 g_log: &mut CudaSlice<f32>,
17059 len_d: &CudaSlice<i32>,
17060 h: usize,
17061 t: usize,
17062 ) -> Result<(), Box<dyn std::error::Error>> {
17063 let f = self.func("gdn_pad_mask_f32");
17064 let cfg = LaunchConfig::for_num_elems((t * h) as u32);
17065 let (hi, ti) = (h as i32, t as i32);
17066 let __s_b = self.gpu.stream();
17067 let mut b = __s_b.launch_builder(&f);
17068 b.arg(beta).arg(g_log).arg(len_d).arg(&hi).arg(&ti);
17069 unsafe {
17070 b.launch(cfg)?;
17071 }
17072 Ok(())
17073 }
17074
17075 pub fn row_gather_dev(
17078 &self,
17079 src: &CudaSlice<f32>,
17080 dst: &mut CudaSlice<f32>,
17081 len_d: &CudaSlice<i32>,
17082 ncols: usize,
17083 ) -> Result<(), Box<dyn std::error::Error>> {
17084 let f = self.func("row_gather_dev_f32");
17085 let cfg = LaunchConfig::for_num_elems(ncols as u32);
17086 let nc = ncols as i32;
17087 let __s_b = self.gpu.stream();
17088 let mut b = __s_b.launch_builder(&f);
17089 b.arg(src).arg(dst).arg(len_d).arg(&nc);
17090 unsafe {
17091 b.launch(cfg)?;
17092 }
17093 Ok(())
17094 }
17095
17096 pub fn matmul_group(
17103 &self,
17104 ws: &[&crate::model::GpuTensor],
17105 x: &CudaSlice<f32>,
17106 m: usize,
17107 ) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
17108 use crate::model::GpuTensor;
17109 let mut out = Vec::with_capacity(ws.len());
17110 let any_mirror = ws
17111 .iter()
17112 .any(|w| matches!(w, GpuTensor::Quant { f16: Some(_), .. }));
17113 if m >= 16 && any_mirror && !self.verify_exact_on() {
17114 let in_f = ws[0].in_features();
17115 let xh = self.f16_act(x, m * in_f, in_f)?;
17116 for w in ws {
17117 if w.in_features() == in_f {
17118 if let Some(y) = self.try_f16_gemm_pre(w, &xh, m)? {
17119 out.push(y);
17120 continue;
17121 }
17122 }
17123 out.push(self.matmul(w, x, m)?);
17124 }
17125 return Ok(out);
17126 }
17127 for w in ws {
17128 out.push(self.matmul(w, x, m)?);
17129 }
17130 Ok(out)
17131 }
17132
17133 pub fn matmul_group_multi(
17140 &self,
17141 ws: &[&crate::model::GpuTensor],
17142 xs: &[&CudaSlice<f32>],
17143 ms: &[usize],
17144 ) -> Result<Vec<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
17145 assert_eq!(xs.len(), ms.len());
17146 let in_f = ws[0].in_features();
17147 let total: usize = ms.iter().sum();
17148 let mut xcat = self.uninit(total * in_f)?;
17149 let mut off = 0usize;
17150 for (x, &m) in xs.iter().zip(ms) {
17151 self.copy_into(&mut xcat, off * in_f, x, m * in_f)?;
17152 off += m;
17153 }
17154 let ys = self.matmul_group(ws, &xcat, total)?;
17155 let mut out: Vec<Vec<CudaSlice<f32>>> = (0..xs.len()).map(|_| Vec::new()).collect();
17156 for (w, y) in ws.iter().zip(ys) {
17157 let out_f = w.out_features();
17158 let mut off = 0usize;
17159 for (s, &m) in ms.iter().enumerate() {
17160 let mut ys_s = self.uninit(m * out_f)?;
17161 let src = y.slice(off * out_f..(off + m) * out_f);
17162 self.gpu.stream().memcpy_dtod(&src, &mut ys_s)?;
17163 out[s].push(ys_s);
17164 off += m;
17165 }
17166 }
17167 Ok(out)
17168 }
17169
17170 pub fn gemm_supports(&self, w: &crate::model::GpuTensor) -> bool {
17180 use crate::model::GpuTensor;
17181 if !legacy_quant_gemm_allowed(
17182 cfg!(memra_portable_cuda),
17183 cfg!(memra_hopper_mma),
17184 std::env::var_os("MEMRA_NO_GEMM").is_some(),
17185 ) {
17186 return false;
17187 }
17188 match w {
17189 GpuTensor::Quant { qtype, .. } => {
17190 matches!(*qtype, QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q4_0)
17191 || (*qtype == QT_NVFP4 && w.in_features() % 64 == 0)
17192 }
17193 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
17194 }
17195 }
17196
17197 pub fn qmatvec_gemm(
17204 &self,
17205 w: &crate::model::GpuTensor,
17206 aq: &CudaSlice<i8>,
17207 ad: &CudaSlice<f32>,
17208 m: usize,
17209 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17210 use crate::model::GpuTensor;
17211 let in_f = w.in_features();
17212 let out_f = w.out_features();
17213 let (bytes, qtype, row_bytes, scale, rp) = match w {
17214 GpuTensor::Quant {
17215 bytes,
17216 qtype,
17217 row_bytes,
17218 scale,
17219 rp,
17220 ..
17221 } => (bytes, *qtype, *row_bytes, *scale, *rp),
17222 _ => unreachable!("gemm_supports guaranteed Quant"),
17223 };
17224 if cfg!(memra_hopper_mma) && qtype == QT_Q8_0 && out_f % 64 == 0 && wgmma_gemm_enabled() {
17230 if let GpuTensor::Quant { rp4: Some(m4), .. } = w {
17231 let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, aq, ad, m, in_f, out_f)?;
17232 if scale != 1.0 {
17233 self.scale_inplace(&mut y, scale, m * out_f)?;
17234 }
17235 return Ok(y);
17236 }
17237 }
17238 let name = match qtype {
17239 QT_Q8_0 => "qmatvec_gemm_q8_0",
17240 QT_Q4_K => "qmatvec_gemm_q4_K",
17241 QT_Q4_0 => {
17242 if rp {
17243 "qmatvec_gemm_q4_0_rp"
17244 } else {
17245 "qmatvec_gemm_q4_0"
17246 }
17247 }
17248 QT_Q5_K => "qmatvec_gemm_q5_K",
17249 QT_Q6_K => "qmatvec_gemm_q6_K",
17250 QT_NVFP4 => {
17251 if rp {
17252 "qmatvec_gemm_nvfp4_rp"
17253 } else {
17254 "qmatvec_gemm_nvfp4"
17255 }
17256 }
17257 _ => unreachable!(),
17258 };
17259 let f = self.func(name);
17260 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);
17265 let k1_tile = if is_k1 {
17267 k1_launch_override().unwrap_or((128, 128, 8))
17268 } else {
17269 (128, 128, 8)
17270 };
17271 let (bm, bn): (u32, u32) = if is_k1 {
17272 (k1_tile.0, k1_tile.1)
17273 } else {
17274 (64, 256)
17275 };
17276 let warps: u32 = if is_k1 {
17277 k1_tile.2
17278 } else {
17279 match qtype {
17280 QT_NVFP4 => 8,
17281 _ => 4,
17282 }
17283 };
17284 let cfg = LaunchConfig {
17285 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
17286 block_dim: (32, warps, 1),
17287 shared_mem_bytes: 0,
17288 };
17289 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
17290 let __s_b = self.gpu.stream();
17291 let mut b = __s_b.launch_builder(&f);
17292 b.arg(bytes)
17293 .arg(aq)
17294 .arg(ad)
17295 .arg(&mut y)
17296 .arg(&inf)
17297 .arg(&outf)
17298 .arg(&mi)
17299 .arg(&rb);
17300 unsafe {
17301 b.launch(cfg)?;
17302 }
17303 if scale != 1.0 {
17304 self.scale_inplace(&mut y, scale, m * out_f)?;
17305 }
17306 Ok(y)
17307 }
17308
17309 pub fn qmatvec_gemm_raw(
17314 &self,
17315 bytes: &CudaSlice<u8>,
17316 x: &CudaSlice<f32>,
17317 m: usize,
17318 in_f: usize,
17319 out_f: usize,
17320 qtype: i32,
17321 row_bytes: usize,
17322 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17323 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
17324 let name = match qtype {
17325 QT_Q8_0 => "qmatvec_gemm_q8_0",
17326 QT_Q4_K => "qmatvec_gemm_q4_K",
17327 QT_Q4_0 => "qmatvec_gemm_q4_0",
17328 QT_Q5_K => "qmatvec_gemm_q5_K",
17329 QT_Q6_K => "qmatvec_gemm_q6_K",
17330 QT_NVFP4 => "qmatvec_gemm_nvfp4",
17331 QT_NVFP4_RP => "qmatvec_gemm_nvfp4_rp",
17332 _ => panic!("qmatvec_gemm_raw: qtype {qtype} has no GEMM kernel"),
17333 };
17334 let f = self.func(name);
17335 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);
17339 let k1_tile = if is_k1 {
17341 k1_launch_override().unwrap_or((128, 128, 8))
17342 } else {
17343 (128, 128, 8)
17344 };
17345 let (bm, bn): (u32, u32) = if is_k1 {
17346 (k1_tile.0, k1_tile.1)
17347 } else {
17348 (64, 256)
17349 };
17350 let warps: u32 = if is_k1 {
17351 k1_tile.2
17352 } else {
17353 match qtype {
17354 QT_NVFP4 | QT_NVFP4_RP => 8,
17355 _ => 4,
17356 }
17357 };
17358 let cfg = LaunchConfig {
17359 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
17360 block_dim: (32, warps, 1),
17361 shared_mem_bytes: 0,
17362 };
17363 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
17364 let __s_b = self.gpu.stream();
17365 let mut b = __s_b.launch_builder(&f);
17366 b.arg(bytes)
17367 .arg(&aq)
17368 .arg(&ad)
17369 .arg(&mut y)
17370 .arg(&inf)
17371 .arg(&outf)
17372 .arg(&mi)
17373 .arg(&rb);
17374 unsafe {
17375 b.launch(cfg)?;
17376 }
17377 Ok(y)
17378 }
17379
17380 pub fn qmatvec_gemm_q8_0_wgmma_raw(
17387 &self,
17388 rp4: &CudaSlice<u8>,
17389 aq: &CudaSlice<i8>,
17390 ad: &CudaSlice<f32>,
17391 m: usize,
17392 in_f: usize,
17393 out_f: usize,
17394 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17395 assert!(
17396 out_f % 64 == 0 && in_f % 32 == 0,
17397 "wgmma GEMM needs out_f%64==0, in_f%32==0"
17398 );
17399 let f = self.func("qmatvec_gemm_q8_0_wgmma");
17400 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
17402 grid_dim: ((out_f / 64) as u32, (m as u32).div_ceil(64), 1),
17403 block_dim: (128, 1, 1),
17404 shared_mem_bytes: 0,
17405 };
17406 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
17407 let __s_b = self.gpu.stream();
17408 let mut b = __s_b.launch_builder(&f);
17409 b.arg(rp4)
17410 .arg(aq)
17411 .arg(ad)
17412 .arg(&mut y)
17413 .arg(&inf)
17414 .arg(&outf)
17415 .arg(&mi);
17416 unsafe {
17417 b.launch(cfg)?;
17418 }
17419 Ok(y)
17420 }
17421
17422 pub fn scale_inplace(
17424 &self,
17425 y: &mut CudaSlice<f32>,
17426 s: f32,
17427 n: usize,
17428 ) -> Result<(), Box<dyn std::error::Error>> {
17429 let f = self.func("scale_f32");
17430 let cfg = LaunchConfig::for_num_elems(n as u32);
17431 let (sf, ni) = (s, n as i32);
17432 let __s_b = self.gpu.stream();
17433 let mut b = __s_b.launch_builder(&f);
17434 b.arg(y).arg(&sf).arg(&ni);
17435 unsafe {
17436 b.launch(cfg)?;
17437 }
17438 Ok(())
17439 }
17440
17441 pub fn bf16_to_f32(
17446 &self,
17447 data: &cudarc::driver::CudaView<'_, u8>,
17448 n: usize,
17449 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17450 let mut out = self.alloc_uninit::<f32>(n)?;
17451 let f = self.func("bf16_to_f32");
17452 let cfg = LaunchConfig::for_num_elems(n as u32);
17453 let ni = n as i32;
17454 let __s_b = self.gpu.stream();
17455 let mut b = __s_b.launch_builder(&f);
17456 b.arg(data).arg(&mut out).arg(&ni);
17457 unsafe {
17458 b.launch(cfg)?;
17459 }
17460 Ok(out)
17461 }
17462
17463 fn linear_bf16_chunked(
17470 &self,
17471 x: &CudaSlice<f32>,
17472 data: &CudaSlice<u8>,
17473 m: usize,
17474 in_f: usize,
17475 out_f: usize,
17476 exact: bool,
17477 canonical_chunk_rows: Option<usize>,
17478 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17479 static EXP_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17483 static EXP_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17484 static EXP_WBYTES: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17485 let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
17486 let started = timing.then(std::time::Instant::now);
17487 let result =
17488 self.linear_bf16_chunked_inner(x, data, m, in_f, out_f, exact, canonical_chunk_rows);
17489 if let Some(started) = started {
17490 use std::sync::atomic::Ordering;
17491 self.stream().synchronize()?;
17492 let ns = EXP_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
17493 + started.elapsed().as_nanos() as u64;
17494 let wb = EXP_WBYTES.fetch_add((in_f * out_f * 2) as u64, Ordering::Relaxed)
17495 + (in_f * out_f * 2) as u64;
17496 let calls = EXP_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
17497 if calls % 1024 == 0 {
17498 eprintln!(
17499 "[bf16-expand-timing] calls={calls} total_ms={:.1} avg_us={:.1} \
17500 weight_gb={:.2}",
17501 ns as f64 / 1.0e6,
17502 ns as f64 / calls as f64 / 1.0e3,
17503 wb as f64 / 1.0e9,
17504 );
17505 }
17506 }
17507 result
17508 }
17509
17510 pub(crate) fn bf16_mmv_on() -> bool {
17515 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17516 *ON.get_or_init(|| std::env::var("MEMRA_BF16_MMV").as_deref() == Ok("1"))
17517 }
17518
17519 fn matvec_bf16(
17522 &self,
17523 data: &CudaSlice<u8>,
17524 x: &CudaSlice<f32>,
17525 in_f: usize,
17526 out_f: usize,
17527 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17528 if data.len() != in_f * out_f * 2 || x.len() < in_f || in_f % 8 != 0 {
17529 return Err(format!(
17530 "matvec_bf16 geometry bytes={} x={} in={in_f} out={out_f}",
17531 data.len(),
17532 x.len()
17533 )
17534 .into());
17535 }
17536 let mut y = self.alloc_uninit::<f32>(out_f)?;
17537 let f = self.func("matvec_bf16_f32acc");
17538 let cfg = LaunchConfig {
17539 grid_dim: (out_f as u32, 1, 1),
17540 block_dim: (mmv_block(), 1, 1),
17541 shared_mem_bytes: 0,
17542 };
17543 let ini = in_f as i32;
17544 let __s_bld = self.gpu.stream();
17545 let mut bld = __s_bld.launch_builder(&f);
17546 bld.arg(data).arg(x).arg(&mut y).arg(&ini);
17547 unsafe {
17548 bld.launch(cfg)?;
17549 }
17550 Ok(y)
17551 }
17552
17553 #[allow(clippy::too_many_arguments)]
17557 #[allow(clippy::too_many_arguments)]
17562 #[allow(clippy::too_many_arguments)]
17567 pub fn qk_norm_rope_append_inc_dcw_rows(
17568 &self,
17569 q_raw_t: &CudaSlice<f32>,
17570 k_raw_t: &CudaSlice<f32>,
17571 v_raw_t: &CudaSlice<f32>,
17572 qw: &CudaSlice<f32>,
17573 kw: &CudaSlice<f32>,
17574 q_out_t: &mut CudaSlice<f32>,
17575 k_out_t: &mut CudaSlice<f32>,
17576 tab: &CudaSlice<u64>,
17577 pos_t: &CudaSlice<i32>,
17578 same_session: bool,
17579 t: usize,
17580 kv_dim_k: usize,
17581 kv_dim_v: usize,
17582 k_tok_bytes: usize,
17583 v_tok_bytes: usize,
17584 head_dim: usize,
17585 n_dims: usize,
17586 nh_q: usize,
17587 nh_k: usize,
17588 eps: f32,
17589 freq_base: f32,
17590 freq_scale: f32,
17591 ff: Option<&CudaSlice<f32>>,
17592 ) -> Result<(), Box<dyn std::error::Error>> {
17593 if head_dim != 128
17594 || kv_dim_v != kv_dim_k
17595 || kv_dim_k != nh_k * head_dim
17596 || t == 0
17597 || t > 32
17598 || tab.len() < t * 6
17599 || pos_t.len() < t
17600 || q_raw_t.len() < t * nh_q * head_dim
17601 || k_raw_t.len() < t * nh_k * head_dim
17602 || v_raw_t.len() < t * kv_dim_v
17603 || q_out_t.len() < t * nh_q * head_dim
17604 || k_out_t.len() < t * nh_k * head_dim
17605 {
17606 return Err(format!(
17607 "qk_norm_rope_append_inc_rows geometry head_dim={head_dim} t={t} \
17608 nh_q={nh_q} nh_k={nh_k}"
17609 )
17610 .into());
17611 }
17612 let f = self.func("qk_norm_rope_append_inc_dcw_rows");
17613 let same_t: i32 = if same_session { t as i32 } else { 0 };
17614 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
17615 let cfg = LaunchConfig {
17616 grid_dim: ((nh_q + nh_k) as u32, 1, t as u32),
17617 block_dim: (128, 1, 1),
17618 shared_mem_bytes: 0,
17619 };
17620 let (kvk, kvv) = (kv_dim_k as i32, kv_dim_v as i32);
17621 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
17622 let (hd, nd, nq, nk) = (head_dim as i32, n_dims as i32, nh_q as i32, nh_k as i32);
17623 let null: u64 = 0;
17624 let __s_b = self.gpu.stream();
17625 let mut b = __s_b.launch_builder(&f);
17626 b.arg(q_raw_t)
17627 .arg(k_raw_t)
17628 .arg(v_raw_t)
17629 .arg(qw)
17630 .arg(kw)
17631 .arg(q_out_t)
17632 .arg(k_out_t)
17633 .arg(tab)
17634 .arg(pos_t)
17635 .arg(&same_t)
17636 .arg(&kvk)
17637 .arg(&kvv)
17638 .arg(&ktb)
17639 .arg(&vtb)
17640 .arg(&hd)
17641 .arg(&nd)
17642 .arg(&nq)
17643 .arg(&nk)
17644 .arg(&eps)
17645 .arg(&theta_scale)
17646 .arg(&freq_scale);
17647 match ff {
17648 Some(freqs) => {
17649 b.arg(freqs);
17650 }
17651 None => {
17652 b.arg(&null);
17653 }
17654 }
17655 unsafe {
17656 b.launch(cfg)?;
17657 }
17658 Ok(())
17659 }
17660
17661 pub fn qk_norm_rope_append_inc_dcw(
17662 &self,
17663 q_raw: &CudaSlice<f32>,
17664 k_raw: &CudaSlice<f32>,
17665 v_raw: &CudaSlice<f32>,
17666 qw: &CudaSlice<f32>,
17667 kw: &CudaSlice<f32>,
17668 q_out: &mut CudaSlice<f32>,
17669 k_out: &mut CudaSlice<f32>,
17670 pos: &CudaSlice<i32>,
17671 k_plane: &mut CudaSlice<u8>,
17672 v_plane: &mut CudaSlice<u8>,
17673 len_dev: &CudaSlice<i32>,
17676 base_dev: Option<&CudaSlice<i32>>,
17677 done_ctr: &mut CudaSlice<u32>,
17678 kv_dim_k: usize,
17679 kv_dim_v: usize,
17680 k_tok_bytes: usize,
17681 v_tok_bytes: usize,
17682 head_dim: usize,
17683 n_dims: usize,
17684 nh_q: usize,
17685 nh_k: usize,
17686 eps: f32,
17687 freq_base: f32,
17688 freq_scale: f32,
17689 ff: Option<&CudaSlice<f32>>,
17690 ) -> Result<(), Box<dyn std::error::Error>> {
17691 if head_dim != 128
17692 || kv_dim_v != kv_dim_k
17693 || kv_dim_k != nh_k * head_dim
17694 || q_raw.len() < nh_q * head_dim
17695 || k_raw.len() < nh_k * head_dim
17696 || v_raw.len() < kv_dim_v
17697 || q_out.len() < nh_q * head_dim
17698 || k_out.len() < nh_k * head_dim
17699 || pos.is_empty()
17700 || done_ctr.is_empty()
17701 {
17702 return Err(format!(
17703 "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}"
17704 )
17705 .into());
17706 }
17707 let f = self.func("qk_norm_rope_append_inc_dcw");
17708 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
17709 let cfg = LaunchConfig {
17710 grid_dim: ((nh_q + nh_k) as u32, 1, 1),
17711 block_dim: (128, 1, 1),
17712 shared_mem_bytes: 0,
17713 };
17714 let (kvk, kvv) = (kv_dim_k as i32, kv_dim_v as i32);
17715 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
17716 let (hd, nd, nq) = (head_dim as i32, n_dims as i32, nh_q as i32);
17717 let null: u64 = 0;
17718 let __s_b = self.gpu.stream();
17719 let mut b = __s_b.launch_builder(&f);
17720 b.arg(q_raw)
17721 .arg(k_raw)
17722 .arg(v_raw)
17723 .arg(qw)
17724 .arg(kw)
17725 .arg(q_out)
17726 .arg(k_out)
17727 .arg(pos)
17728 .arg(&mut *k_plane)
17729 .arg(&mut *v_plane)
17730 .arg(len_dev);
17731 match base_dev {
17732 Some(base) => {
17733 b.arg(base);
17734 }
17735 None => {
17736 b.arg(&null);
17737 }
17738 }
17739 b.arg(&mut *done_ctr)
17740 .arg(&kvk)
17741 .arg(&kvv)
17742 .arg(&ktb)
17743 .arg(&vtb)
17744 .arg(&hd)
17745 .arg(&nd)
17746 .arg(&nq)
17747 .arg(&eps)
17748 .arg(&theta_scale)
17749 .arg(&freq_scale);
17750 match ff {
17751 Some(freqs) => {
17752 b.arg(freqs);
17753 }
17754 None => {
17755 b.arg(&null);
17756 }
17757 }
17758 unsafe {
17759 b.launch(cfg)?;
17760 }
17761 Ok(())
17762 }
17763
17764 pub fn qk_norm_rope_into(
17765 &self,
17766 q_raw: &CudaSlice<f32>,
17767 k_raw: &CudaSlice<f32>,
17768 qw: &CudaSlice<f32>,
17769 kw: &CudaSlice<f32>,
17770 q_out: &mut CudaSlice<f32>,
17771 k_out: &mut CudaSlice<f32>,
17772 pos: &CudaSlice<i32>,
17773 head_dim: usize,
17774 n_dims: usize,
17775 nh_q: usize,
17776 nh_k: usize,
17777 eps: f32,
17778 freq_base: f32,
17779 freq_scale: f32,
17780 ff: Option<&CudaSlice<f32>>,
17781 ) -> Result<(), Box<dyn std::error::Error>> {
17782 if head_dim > 512
17783 || q_raw.len() < nh_q * head_dim
17784 || k_raw.len() < nh_k * head_dim
17785 || q_out.len() < nh_q * head_dim
17786 || k_out.len() < nh_k * head_dim
17787 || qw.len() < head_dim
17788 || kw.len() < head_dim
17789 || pos.is_empty()
17790 {
17791 return Err(format!(
17792 "qk_norm_rope geometry head_dim={head_dim} nh_q={nh_q} nh_k={nh_k}"
17793 )
17794 .into());
17795 }
17796 let f = self.func("qk_norm_rope_f32");
17797 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
17798 let cfg = LaunchConfig {
17799 grid_dim: ((nh_q + nh_k) as u32, 1, 1),
17800 block_dim: (128, 1, 1),
17801 shared_mem_bytes: 0,
17802 };
17803 let (hd, nd, nq) = (head_dim as i32, n_dims as i32, nh_q as i32);
17804 let __s_b = self.gpu.stream();
17805 let mut b = __s_b.launch_builder(&f);
17806 b.arg(q_raw)
17807 .arg(k_raw)
17808 .arg(qw)
17809 .arg(kw)
17810 .arg(q_out)
17811 .arg(k_out)
17812 .arg(pos)
17813 .arg(&hd)
17814 .arg(&nd)
17815 .arg(&nq)
17816 .arg(&eps)
17817 .arg(&theta_scale)
17818 .arg(&freq_scale);
17819 match ff {
17820 Some(ffv) => {
17821 b.arg(ffv);
17822 unsafe {
17823 b.launch(cfg)?;
17824 }
17825 }
17826 None => {
17827 let null: u64 = 0;
17828 b.arg(&null);
17829 unsafe {
17830 b.launch(cfg)?;
17831 }
17832 }
17833 }
17834 Ok(())
17835 }
17836
17837 #[allow(clippy::too_many_arguments)]
17840 pub fn matvec_f32_b4_into(
17841 &self,
17842 w: [&CudaSlice<f32>; 4],
17843 x: &CudaSlice<f32>,
17844 y: &mut CudaSlice<f32>,
17845 block_cols: usize,
17846 out_f: usize,
17847 ) -> Result<(), Box<dyn std::error::Error>> {
17848 if block_cols % 4 != 0
17849 || x.len() < 4 * block_cols
17850 || y.len() < out_f
17851 || w.iter().any(|w| w.len() != out_f * block_cols)
17852 {
17853 return Err(format!(
17854 "matvec_f32_b4 geometry block_cols={block_cols} out={out_f} x={}",
17855 x.len()
17856 )
17857 .into());
17858 }
17859 let f = self.func("matvec_f32_b4");
17860 let cfg = LaunchConfig {
17861 grid_dim: (out_f as u32, 1, 1),
17862 block_dim: (128, 1, 1),
17863 shared_mem_bytes: 0,
17864 };
17865 let (bc, of) = (block_cols as i32, out_f as i32);
17866 let __s_b = self.gpu.stream();
17867 let mut b = __s_b.launch_builder(&f);
17868 b.arg(w[0])
17869 .arg(w[1])
17870 .arg(w[2])
17871 .arg(w[3])
17872 .arg(x)
17873 .arg(y)
17874 .arg(&bc)
17875 .arg(&of);
17876 unsafe {
17877 b.launch(cfg)?;
17878 }
17879 Ok(())
17880 }
17881
17882 pub fn axpy_rows_seq_into(
17885 &self,
17886 x: &CudaSlice<f32>,
17887 w: &CudaSlice<f32>,
17888 y: &mut CudaSlice<f32>,
17889 width: usize,
17890 n_rows: usize,
17891 ) -> Result<(), Box<dyn std::error::Error>> {
17892 if x.len() < n_rows * width || w.len() < n_rows || y.len() < width {
17893 return Err(format!(
17894 "axpy_rows_seq geometry x={} w={} y={} width={width} rows={n_rows}",
17895 x.len(),
17896 w.len(),
17897 y.len()
17898 )
17899 .into());
17900 }
17901 let f = self.func("axpy_rows_seq_f32");
17902 let cfg = LaunchConfig::for_num_elems(width as u32);
17903 let (wi, nr) = (width as i32, n_rows as i32);
17904 let __s_b = self.gpu.stream();
17905 let mut b = __s_b.launch_builder(&f);
17906 b.arg(x).arg(w).arg(y).arg(&wi).arg(&nr);
17907 unsafe {
17908 b.launch(cfg)?;
17909 }
17910 Ok(())
17911 }
17912
17913 #[allow(clippy::too_many_arguments)]
17917 pub fn axpy_rows_seq_md_off_into(
17918 &self,
17919 x: &CudaSlice<f32>,
17920 w_route: &CudaSlice<f32>,
17921 md: &CudaSlice<f32>,
17922 sel: &CudaSlice<i32>,
17923 y: &mut CudaSlice<f32>,
17924 width: usize,
17925 n_rows: usize,
17926 row0: usize,
17927 ) -> Result<(), Box<dyn std::error::Error>> {
17928 if x.len() < (row0 + n_rows) * width
17929 || w_route.len() < row0 + n_rows
17930 || sel.len() < row0 + n_rows
17931 || y.len() < width
17932 {
17933 return Err(format!(
17934 "axpy_rows_seq_md_off geometry x={} w={} sel={} y={} width={width} \
17935 rows={n_rows} row0={row0}",
17936 x.len(),
17937 w_route.len(),
17938 sel.len(),
17939 y.len()
17940 )
17941 .into());
17942 }
17943 let f = self.func("axpy_rows_seq_md_off_f32");
17944 let cfg = LaunchConfig::for_num_elems(width as u32);
17945 let (wi, nr, r0) = (width as i32, n_rows as i32, row0 as i32);
17946 let __s_b = self.gpu.stream();
17947 let mut b = __s_b.launch_builder(&f);
17948 b.arg(x)
17949 .arg(w_route)
17950 .arg(md)
17951 .arg(sel)
17952 .arg(y)
17953 .arg(&wi)
17954 .arg(&nr)
17955 .arg(&r0);
17956 unsafe {
17957 b.launch(cfg)?;
17958 }
17959 Ok(())
17960 }
17961
17962 #[allow(clippy::too_many_arguments)]
17965 pub fn axpy_rows_seq_md_into(
17966 &self,
17967 x: &CudaSlice<f32>,
17968 w_route: &CudaSlice<f32>,
17969 md: &CudaSlice<f32>,
17970 sel: &CudaSlice<i32>,
17971 y: &mut CudaSlice<f32>,
17972 width: usize,
17973 n_rows: usize,
17974 ) -> Result<(), Box<dyn std::error::Error>> {
17975 if x.len() < n_rows * width
17976 || w_route.len() < n_rows
17977 || sel.len() < n_rows
17978 || y.len() < width
17979 {
17980 return Err(format!(
17981 "axpy_rows_seq_md geometry x={} w={} sel={} y={} width={width} rows={n_rows}",
17982 x.len(),
17983 w_route.len(),
17984 sel.len(),
17985 y.len()
17986 )
17987 .into());
17988 }
17989 let f = self.func("axpy_rows_seq_md_f32");
17990 let cfg = LaunchConfig::for_num_elems(width as u32);
17991 let (wi, nr) = (width as i32, n_rows as i32);
17992 let __s_b = self.gpu.stream();
17993 let mut b = __s_b.launch_builder(&f);
17994 b.arg(x)
17995 .arg(w_route)
17996 .arg(md)
17997 .arg(sel)
17998 .arg(y)
17999 .arg(&wi)
18000 .arg(&nr);
18001 unsafe {
18002 b.launch(cfg)?;
18003 }
18004 Ok(())
18005 }
18006
18007 #[allow(clippy::too_many_arguments)]
18009 #[allow(clippy::too_many_arguments)]
18013 pub fn matvec_bf16_qkvg_tcol_into(
18014 &self,
18015 wq: &CudaSlice<u8>,
18016 wk: &CudaSlice<u8>,
18017 wv: &CudaSlice<u8>,
18018 wg: &CudaSlice<u8>,
18019 x_t: &CudaSlice<f32>,
18020 yq: &mut CudaSlice<f32>,
18021 yk: &mut CudaSlice<f32>,
18022 yv: &mut CudaSlice<f32>,
18023 yg: &mut CudaSlice<f32>,
18024 in_f: usize,
18025 out_q: usize,
18026 out_kv: usize,
18027 out_g: usize,
18028 t: usize,
18029 ) -> Result<(), Box<dyn std::error::Error>> {
18030 if t == 0
18031 || t > 8
18032 || in_f % 8 != 0
18033 || x_t.len() < t * in_f
18034 || yq.len() < t * out_q
18035 || yk.len() < t * out_kv
18036 || yv.len() < t * out_kv
18037 || (out_g > 0 && yg.len() < t * out_g)
18038 {
18039 return Err("matvec_bf16_qkvg_tcol geometry".into());
18040 }
18041 let grid = out_q + 2 * out_kv + out_g;
18042 let cfg = LaunchConfig {
18043 grid_dim: (grid as u32, 1, 1),
18044 block_dim: (mmv_block(), 1, 1),
18045 shared_mem_bytes: 0,
18046 };
18047 let (ini, oq, okv, og, ti) = (
18048 in_f as i32,
18049 out_q as i32,
18050 out_kv as i32,
18051 out_g as i32,
18052 t as i32,
18053 );
18054 let __s_b = self.gpu.stream();
18055 let f = self.func("matvec_bf16_qkvg_tcol");
18061 let mut b = __s_b.launch_builder(&f);
18062 b.arg(wq)
18063 .arg(wk)
18064 .arg(wv)
18065 .arg(wg)
18066 .arg(x_t)
18067 .arg(yq)
18068 .arg(yk)
18069 .arg(yv)
18070 .arg(yg)
18071 .arg(&ini)
18072 .arg(&oq)
18073 .arg(&okv)
18074 .arg(&og)
18075 .arg(&ti);
18076 unsafe {
18077 b.launch(cfg)?;
18078 }
18079 Ok(())
18080 }
18081
18082 pub fn matvec_bf16_qkvg_into(
18083 &self,
18084 wq: &CudaSlice<u8>,
18085 wk: &CudaSlice<u8>,
18086 wv: &CudaSlice<u8>,
18087 wg: &CudaSlice<u8>,
18088 x: &CudaSlice<f32>,
18089 yq: &mut CudaSlice<f32>,
18090 yk: &mut CudaSlice<f32>,
18091 yv: &mut CudaSlice<f32>,
18092 yg: &mut CudaSlice<f32>,
18093 in_f: usize,
18094 out_q: usize,
18095 out_kv: usize,
18096 out_g: usize,
18097 ) -> Result<(), Box<dyn std::error::Error>> {
18098 if in_f % 8 != 0
18099 || wq.len() != out_q * in_f * 2
18100 || wk.len() != out_kv * in_f * 2
18101 || wv.len() != out_kv * in_f * 2
18102 || wg.len() < out_g * in_f * 2
18103 || x.len() < in_f
18104 || yq.len() < out_q
18105 || yk.len() < out_kv
18106 || yv.len() < out_kv
18107 || (out_g > 0 && yg.len() < out_g)
18108 {
18109 return Err(format!(
18110 "fused bf16 QKV geometry in={in_f} out_q={out_q} out_kv={out_kv} out_g={out_g}"
18111 )
18112 .into());
18113 }
18114 let f = self.func("matvec_bf16_qkvg");
18115 let cfg = LaunchConfig {
18116 grid_dim: ((out_q + 2 * out_kv + out_g) as u32, 1, 1),
18117 block_dim: (mmv_block(), 1, 1),
18118 shared_mem_bytes: 0,
18119 };
18120 let (inf, oq, okv, og) = (in_f as i32, out_q as i32, out_kv as i32, out_g as i32);
18121 let __s_b = self.gpu.stream();
18122 let mut b = __s_b.launch_builder(&f);
18123 b.arg(wq)
18124 .arg(wk)
18125 .arg(wv)
18126 .arg(wg)
18127 .arg(x)
18128 .arg(yq)
18129 .arg(yk)
18130 .arg(yv)
18131 .arg(yg)
18132 .arg(&inf)
18133 .arg(&oq)
18134 .arg(&okv)
18135 .arg(&og);
18136 unsafe {
18137 b.launch(cfg)?;
18138 }
18139 Ok(())
18140 }
18141
18142 pub fn matvec_bf16_b4_into(
18144 &self,
18145 w: [&CudaSlice<u8>; 4],
18146 x: &CudaSlice<f32>,
18147 y: &mut CudaSlice<f32>,
18148 block_cols: usize,
18149 out_f: usize,
18150 ) -> Result<(), Box<dyn std::error::Error>> {
18151 if block_cols % 8 != 0
18152 || x.len() < 4 * block_cols
18153 || y.len() < out_f
18154 || w.iter().any(|w| w.len() != out_f * block_cols * 2)
18155 {
18156 return Err(format!(
18157 "bf16 b4 geometry block_cols={block_cols} out={out_f} x={}",
18158 x.len()
18159 )
18160 .into());
18161 }
18162 static B4_X2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18165 let x2 = *B4_X2.get_or_init(|| std::env::var("MEMRA_B4_X2").as_deref() == Ok("1"));
18166 let f = self.func(if x2 {
18167 "matvec_bf16_b4_x2"
18168 } else {
18169 "matvec_bf16_b4"
18170 });
18171 let grid = if x2 { out_f.div_ceil(2) } else { out_f };
18172 let cfg = LaunchConfig {
18173 grid_dim: (grid as u32, 1, 1),
18174 block_dim: (mmv_block(), 1, 1),
18175 shared_mem_bytes: 0,
18176 };
18177 let (bc, of) = (block_cols as i32, out_f as i32);
18178 let __s_b = self.gpu.stream();
18179 let mut b = __s_b.launch_builder(&f);
18180 b.arg(w[0])
18181 .arg(w[1])
18182 .arg(w[2])
18183 .arg(w[3])
18184 .arg(x)
18185 .arg(y)
18186 .arg(&bc)
18187 .arg(&of);
18188 unsafe {
18189 b.launch(cfg)?;
18190 }
18191 Ok(())
18192 }
18193
18194 pub fn matvec_bf16_b4_tcol_into(
18200 &self,
18201 w: [&CudaSlice<u8>; 4],
18202 x_t: &CudaSlice<f32>,
18203 y_t: &mut CudaSlice<f32>,
18204 block_cols: usize,
18205 out_f: usize,
18206 t: usize,
18207 ) -> Result<(), Box<dyn std::error::Error>> {
18208 if block_cols % 8 != 0
18209 || t == 0
18210 || t > 8
18211 || x_t.len() < t * 4 * block_cols
18212 || y_t.len() < t * out_f
18213 || w.iter().any(|w| w.len() != out_f * block_cols * 2)
18214 {
18215 return Err(format!(
18216 "bf16 b4 tcol geometry block_cols={block_cols} out={out_f} t={t} x={}",
18217 x_t.len()
18218 )
18219 .into());
18220 }
18221 if std::env::var("MEMRA_B4_X2").as_deref() == Ok("1") {
18222 return Err(
18223 "b4 tcol verify is qualified against the plain b4 kernel only \
18224 (MEMRA_B4_X2=1 is a different t=1 program)"
18225 .into(),
18226 );
18227 }
18228 let cfg = LaunchConfig {
18232 grid_dim: (out_f as u32, 1, 1),
18233 block_dim: (mmv_block(), 1, 1),
18234 shared_mem_bytes: 0,
18235 };
18236 let (bc, of, ti) = (block_cols as i32, out_f as i32, t as i32);
18237 let __s_b = self.gpu.stream();
18238 let f = self.func("matvec_bf16_b4_tcol");
18239 let mut b = __s_b.launch_builder(&f);
18240 b.arg(w[0])
18241 .arg(w[1])
18242 .arg(w[2])
18243 .arg(w[3])
18244 .arg(x_t)
18245 .arg(y_t)
18246 .arg(&bc)
18247 .arg(&of)
18248 .arg(&ti);
18249 unsafe {
18250 b.launch(cfg)?;
18251 }
18252 Ok(())
18253 }
18254
18255 pub fn q8_0_row_bytes(in_f: usize) -> usize {
18258 in_f / 32 * 34
18259 }
18260
18261 pub fn encode_q8_0_from_bf16(
18265 &self,
18266 w_bf16: &CudaSlice<u8>,
18267 out: &mut CudaSlice<u8>,
18268 in_f: usize,
18269 out_f: usize,
18270 ) -> Result<(), Box<dyn std::error::Error>> {
18271 if in_f % 32 != 0
18272 || w_bf16.len() < in_f * out_f * 2
18273 || out.len() < out_f * Self::q8_0_row_bytes(in_f)
18274 {
18275 return Err(format!(
18276 "encode_q8_0_from_bf16 geometry in={in_f} out={out_f} src={} dst={}",
18277 w_bf16.len(),
18278 out.len()
18279 )
18280 .into());
18281 }
18282 let f = self.func("encode_q8_0_rows_from_bf16");
18283 const PAIRS_PER_BLOCK: u32 = 4;
18286 let pairs = (out_f * (in_f / 32)) as u64;
18287 let cfg = LaunchConfig {
18288 grid_dim: ((pairs.div_ceil(PAIRS_PER_BLOCK as u64)) as u32, 1, 1),
18289 block_dim: (32, PAIRS_PER_BLOCK, 1),
18290 shared_mem_bytes: 0,
18291 };
18292 let (ini, outi) = (in_f as i32, out_f as i32);
18293 let __s_b = self.gpu.stream();
18294 let mut b = __s_b.launch_builder(&f);
18295 b.arg(w_bf16).arg(out).arg(&ini).arg(&outi);
18296 unsafe {
18297 b.launch(cfg)?;
18298 }
18299 Ok(())
18300 }
18301
18302 pub fn encode_q8_0_from_bf16_view(
18306 &self,
18307 w_bf16: &cudarc::driver::CudaView<'_, u8>,
18308 out: &mut CudaSlice<u8>,
18309 in_f: usize,
18310 out_f: usize,
18311 ) -> Result<(), Box<dyn std::error::Error>> {
18312 if in_f % 32 != 0
18313 || w_bf16.len() < in_f * out_f * 2
18314 || out.len() < out_f * Self::q8_0_row_bytes(in_f)
18315 {
18316 return Err(format!(
18317 "encode_q8_0_from_bf16_view geometry in={in_f} out={out_f} src={} dst={}",
18318 w_bf16.len(),
18319 out.len()
18320 )
18321 .into());
18322 }
18323 let f = self.func("encode_q8_0_rows_from_bf16");
18324 const PAIRS_PER_BLOCK: u32 = 4;
18325 let pairs = (out_f * (in_f / 32)) as u64;
18326 let cfg = LaunchConfig {
18327 grid_dim: ((pairs.div_ceil(PAIRS_PER_BLOCK as u64)) as u32, 1, 1),
18328 block_dim: (32, PAIRS_PER_BLOCK, 1),
18329 shared_mem_bytes: 0,
18330 };
18331 let (ini, outi) = (in_f as i32, out_f as i32);
18332 let __s_b = self.gpu.stream();
18333 let mut b = __s_b.launch_builder(&f);
18334 b.arg(w_bf16).arg(out).arg(&ini).arg(&outi);
18335 unsafe {
18336 b.launch(cfg)?;
18337 }
18338 Ok(())
18339 }
18340
18341 #[allow(clippy::too_many_arguments)]
18346 pub fn qmatvec_q8_0_qkv_rp_into(
18347 &self,
18348 wq: &CudaSlice<u8>,
18349 wk: &CudaSlice<u8>,
18350 wv: &CudaSlice<u8>,
18351 aq: &CudaSlice<i8>,
18352 ad: &CudaSlice<f32>,
18353 yq: &mut CudaSlice<f32>,
18354 yk: &mut CudaSlice<f32>,
18355 yv: &mut CudaSlice<f32>,
18356 in_f: usize,
18357 out_q: usize,
18358 out_kv: usize,
18359 ) -> Result<(), Box<dyn std::error::Error>> {
18360 const ROWS_PER_BLOCK: u32 = 4; let rows = out_q + 2 * out_kv;
18362 let nblk = in_f / 32;
18363 if in_f % 32 != 0
18364 || aq.len() < in_f
18365 || ad.len() < nblk
18366 || yq.len() < out_q
18367 || yk.len() < out_kv
18368 || yv.len() < out_kv
18369 || wq.len() < out_q * nblk * 34
18370 || wk.len() < out_kv * nblk * 34
18371 || wv.len() < out_kv * nblk * 34
18372 {
18373 return Err(
18374 format!("q8_0 qkv rp geometry in={in_f} out_q={out_q} out_kv={out_kv}").into(),
18375 );
18376 }
18377 let f = self.func("qmatvec_q8_0_qkv_rp");
18378 let cfg = LaunchConfig {
18379 grid_dim: ((rows as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
18380 block_dim: (32, ROWS_PER_BLOCK, 1),
18381 shared_mem_bytes: 0,
18382 };
18383 let (ini, oq, okv) = (in_f as i32, out_q as i32, out_kv as i32);
18384 let __s_b = self.gpu.stream();
18385 let mut b = __s_b.launch_builder(&f);
18386 b.arg(wq)
18387 .arg(wk)
18388 .arg(wv)
18389 .arg(aq)
18390 .arg(ad)
18391 .arg(yq)
18392 .arg(yk)
18393 .arg(yv)
18394 .arg(&ini)
18395 .arg(&oq)
18396 .arg(&okv);
18397 unsafe {
18398 b.launch(cfg)?;
18399 }
18400 Ok(())
18401 }
18402
18403 #[allow(clippy::too_many_arguments)]
18407 pub fn qmatvec_q8_0_b4_rp_into(
18408 &self,
18409 w: [&CudaSlice<u8>; 4],
18410 aq: &CudaSlice<i8>,
18411 ad: &CudaSlice<f32>,
18412 y: &mut CudaSlice<f32>,
18413 block_cols: usize,
18414 out_f: usize,
18415 ) -> Result<(), Box<dyn std::error::Error>> {
18416 const ROWS_PER_BLOCK: u32 = 4; let nblk = block_cols / 32;
18418 if block_cols % 32 != 0
18419 || aq.len() < 4 * block_cols
18420 || ad.len() < 4 * nblk
18421 || y.len() < out_f
18422 || w.iter().any(|p| p.len() < out_f * nblk * 34)
18423 {
18424 return Err(format!("q8_0 b4 rp geometry block_cols={block_cols} out={out_f}").into());
18425 }
18426 let f = self.func("qmatvec_q8_0_b4_rp");
18427 let cfg = LaunchConfig {
18428 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
18429 block_dim: (32, ROWS_PER_BLOCK, 1),
18430 shared_mem_bytes: 0,
18431 };
18432 let (bc, of) = (block_cols as i32, out_f as i32);
18433 let __s_b = self.gpu.stream();
18434 let mut b = __s_b.launch_builder(&f);
18435 b.arg(w[0])
18436 .arg(w[1])
18437 .arg(w[2])
18438 .arg(w[3])
18439 .arg(aq)
18440 .arg(ad)
18441 .arg(y)
18442 .arg(&bc)
18443 .arg(&of);
18444 unsafe {
18445 b.launch(cfg)?;
18446 }
18447 Ok(())
18448 }
18449
18450 fn matvec_bf16_via_q8_mirror_t(
18453 &self,
18454 data: &CudaSlice<u8>,
18455 x: &CudaSlice<f32>,
18456 y: &mut CudaSlice<f32>,
18457 in_f: usize,
18458 out_f: usize,
18459 t: usize,
18460 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
18461 use cudarc::driver::DevicePtr;
18462 let key = {
18463 let s = self.gpu.stream();
18464 let (p, _g) = data.device_ptr(&s);
18465 (p as u64, in_f as u32, out_f as u32)
18466 };
18467 {
18468 let mut mirrors = self
18469 .w8_mirrors
18470 .lock()
18471 .map_err(|_| "w8 mirror map is poisoned")?;
18472 if !mirrors.contains_key(&key) {
18473 let mut interleaved = self.alloc_u8_uninit(out_f * Self::q8_0_row_bytes(in_f))?;
18474 self.encode_q8_0_from_bf16(data, &mut interleaved, in_f, out_f)?;
18475 let planar = self.build_q8_rp4_raw(&interleaved, in_f, out_f)?;
18476 mirrors.insert(key, planar);
18477 }
18478 }
18479 let nblk = in_f / 32;
18480 let akey = in_f * 64 + t.min(32);
18482 {
18483 let mut act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
18484 if !act.contains_key(&akey) {
18485 let aq = self.alloc_i8_uninit(32 * in_f)?;
18486 let ad = self.alloc_uninit::<f32>(32 * nblk)?;
18487 act.insert(akey, (aq, ad));
18488 }
18489 let (aq, ad) = act.get_mut(&akey).expect("just inserted");
18490 self.quantize_q8_1_into(x, t, in_f, aq, ad)?;
18491 }
18492 let mirrors = self
18493 .w8_mirrors
18494 .lock()
18495 .map_err(|_| "w8 mirror map is poisoned")?;
18496 let act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
18497 let mirror = mirrors.get(&key).expect("built above");
18498 let (aq, ad) = act.get(&akey).expect("built above");
18499 const ROWS_PER_BLOCK: u32 = 4;
18500 let (ini, of) = (in_f as i32, out_f as i32);
18501 if q8t_wonce_on() && t <= 32 {
18505 let f = self.func(if t <= 8 {
18506 "qmatvec_q8_0_rows_tw"
18507 } else {
18508 "qmatvec_q8_0_rows_tw32"
18509 });
18510 let cfg = LaunchConfig {
18511 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
18512 block_dim: (32, ROWS_PER_BLOCK, 1),
18513 shared_mem_bytes: 0,
18514 };
18515 let ti = t as i32;
18516 let __s_b = self.gpu.stream();
18517 let mut b = __s_b.launch_builder(&f);
18518 b.arg(mirror)
18519 .arg(aq)
18520 .arg(ad)
18521 .arg(&mut *y)
18522 .arg(&ini)
18523 .arg(&of)
18524 .arg(&ti);
18525 unsafe {
18526 b.launch(cfg)?;
18527 }
18528 return Ok(Some(()));
18529 }
18530 let f = self.func("qmatvec_q8_0_rows_t");
18531 let cfg = LaunchConfig {
18532 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), t as u32, 1),
18533 block_dim: (32, ROWS_PER_BLOCK, 1),
18534 shared_mem_bytes: 0,
18535 };
18536 let __s_b = self.gpu.stream();
18537 let mut b = __s_b.launch_builder(&f);
18538 b.arg(mirror)
18539 .arg(aq)
18540 .arg(ad)
18541 .arg(&mut *y)
18542 .arg(&ini)
18543 .arg(&of);
18544 unsafe {
18545 b.launch(cfg)?;
18546 }
18547 Ok(Some(()))
18548 }
18549
18550 fn matvec_bf16_via_q8_mirror(
18553 &self,
18554 data: &CudaSlice<u8>,
18555 x: &CudaSlice<f32>,
18556 y: &mut CudaSlice<f32>,
18557 in_f: usize,
18558 out_f: usize,
18559 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
18560 use cudarc::driver::DevicePtr;
18561 let key = {
18562 let s = self.gpu.stream();
18563 let (p, _g) = data.device_ptr(&s);
18564 (p as u64, in_f as u32, out_f as u32)
18565 };
18566 {
18567 let mut mirrors = self
18568 .w8_mirrors
18569 .lock()
18570 .map_err(|_| "w8 mirror map is poisoned")?;
18571 if !mirrors.contains_key(&key) {
18572 let mut interleaved = self.alloc_u8_uninit(out_f * Self::q8_0_row_bytes(in_f))?;
18573 self.encode_q8_0_from_bf16(data, &mut interleaved, in_f, out_f)?;
18574 let planar = self.build_q8_rp4_raw(&interleaved, in_f, out_f)?;
18575 mirrors.insert(key, planar);
18576 if std::env::var("MEMRA_W8_TRACE").as_deref() == Ok("1") {
18582 eprintln!(
18583 "[w8-mirror] built in_f={in_f} out_f={out_f} mirrors={}",
18584 mirrors.len()
18585 );
18586 }
18587 }
18588 }
18589 let nblk = in_f / 32;
18590 {
18591 let mut act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
18592 if !act.contains_key(&in_f) {
18593 let aq = self.alloc_uninit::<i8>(in_f)?;
18594 let ad = self.alloc_uninit::<f32>(nblk)?;
18595 act.insert(in_f, (aq, ad));
18596 }
18597 let (aq, ad) = act.get_mut(&in_f).expect("just inserted");
18598 self.quantize_q8_1_into(x, 1, in_f, aq, ad)?;
18599 }
18600 let mirrors = self
18601 .w8_mirrors
18602 .lock()
18603 .map_err(|_| "w8 mirror map is poisoned")?;
18604 let act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
18605 let mirror = mirrors.get(&key).expect("built above");
18606 let (aq, ad) = act.get(&in_f).expect("built above");
18607 self.qmatvec_mmvq_into(
18608 mirror,
18609 aq,
18610 ad,
18611 1,
18612 in_f,
18613 out_f,
18614 QT_Q8_0,
18615 Self::q8_0_row_bytes(in_f),
18616 1.0,
18617 true,
18618 y,
18619 )?;
18620 Ok(Some(()))
18621 }
18622
18623 #[allow(clippy::too_many_arguments)]
18628 pub fn qmatvec_q8_0_qkv_rp_t_into(
18629 &self,
18630 wq: &CudaSlice<u8>,
18631 wk: &CudaSlice<u8>,
18632 wv: &CudaSlice<u8>,
18633 aq: &CudaSlice<i8>,
18634 ad: &CudaSlice<f32>,
18635 yq: &mut CudaSlice<f32>,
18636 yk: &mut CudaSlice<f32>,
18637 yv: &mut CudaSlice<f32>,
18638 in_f: usize,
18639 out_q: usize,
18640 out_kv: usize,
18641 t: usize,
18642 ) -> Result<(), Box<dyn std::error::Error>> {
18643 const ROWS_PER_BLOCK: u32 = 4;
18644 let rows = out_q + 2 * out_kv;
18645 let nblk = in_f / 32;
18646 if in_f % 32 != 0
18647 || t == 0
18648 || aq.len() < t * in_f
18649 || ad.len() < t * nblk
18650 || yq.len() < t * out_q
18651 || yk.len() < t * out_kv
18652 || yv.len() < t * out_kv
18653 {
18654 return Err(format!("q8_0 qkv rp_t geometry in={in_f} t={t}").into());
18655 }
18656 let (ini, oq, okv) = (in_f as i32, out_q as i32, out_kv as i32);
18657 if q8t_wonce_on() && t <= 32 {
18661 let f = self.func(if t <= 8 {
18662 "qmatvec_q8_0_qkv_rp_tw"
18663 } else {
18664 "qmatvec_q8_0_qkv_rp_tw32"
18665 });
18666 let cfg = LaunchConfig {
18667 grid_dim: ((rows as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
18668 block_dim: (32, ROWS_PER_BLOCK, 1),
18669 shared_mem_bytes: 0,
18670 };
18671 let ti = t as i32;
18672 let __s_b = self.gpu.stream();
18673 let mut b = __s_b.launch_builder(&f);
18674 b.arg(wq)
18675 .arg(wk)
18676 .arg(wv)
18677 .arg(aq)
18678 .arg(ad)
18679 .arg(yq)
18680 .arg(yk)
18681 .arg(yv)
18682 .arg(&ini)
18683 .arg(&oq)
18684 .arg(&okv)
18685 .arg(&ti);
18686 unsafe {
18687 b.launch(cfg)?;
18688 }
18689 return Ok(());
18690 }
18691 let f = self.func("qmatvec_q8_0_qkv_rp_t");
18692 let cfg = LaunchConfig {
18693 grid_dim: ((rows as u32).div_ceil(ROWS_PER_BLOCK), t as u32, 1),
18694 block_dim: (32, ROWS_PER_BLOCK, 1),
18695 shared_mem_bytes: 0,
18696 };
18697 let __s_b = self.gpu.stream();
18698 let mut b = __s_b.launch_builder(&f);
18699 b.arg(wq)
18700 .arg(wk)
18701 .arg(wv)
18702 .arg(aq)
18703 .arg(ad)
18704 .arg(yq)
18705 .arg(yk)
18706 .arg(yv)
18707 .arg(&ini)
18708 .arg(&oq)
18709 .arg(&okv);
18710 unsafe {
18711 b.launch(cfg)?;
18712 }
18713 Ok(())
18714 }
18715
18716 #[allow(clippy::too_many_arguments)]
18719 pub fn qmatvec_q8_0_b4_rp_t_into(
18720 &self,
18721 w: [&CudaSlice<u8>; 4],
18722 aq: &CudaSlice<i8>,
18723 ad: &CudaSlice<f32>,
18724 y: &mut CudaSlice<f32>,
18725 block_cols: usize,
18726 out_f: usize,
18727 t: usize,
18728 ) -> Result<(), Box<dyn std::error::Error>> {
18729 const ROWS_PER_BLOCK: u32 = 4;
18730 let nblk = block_cols / 32;
18731 if block_cols % 32 != 0
18732 || t == 0
18733 || aq.len() < t * 4 * block_cols
18734 || ad.len() < t * 4 * nblk
18735 || y.len() < t * out_f
18736 {
18737 return Err(format!("q8_0 b4 rp_t geometry cols={block_cols} t={t}").into());
18738 }
18739 let (bc, of) = (block_cols as i32, out_f as i32);
18740 if q8t_wonce_on() && t <= 32 {
18743 let f = self.func(if t <= 8 {
18744 "qmatvec_q8_0_b4_rp_tw"
18745 } else {
18746 "qmatvec_q8_0_b4_rp_tw32"
18747 });
18748 let cfg = LaunchConfig {
18749 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
18750 block_dim: (32, ROWS_PER_BLOCK, 1),
18751 shared_mem_bytes: 0,
18752 };
18753 let ti = t as i32;
18754 let __s_b = self.gpu.stream();
18755 let mut b = __s_b.launch_builder(&f);
18756 b.arg(w[0])
18757 .arg(w[1])
18758 .arg(w[2])
18759 .arg(w[3])
18760 .arg(aq)
18761 .arg(ad)
18762 .arg(y)
18763 .arg(&bc)
18764 .arg(&of)
18765 .arg(&ti);
18766 unsafe {
18767 b.launch(cfg)?;
18768 }
18769 return Ok(());
18770 }
18771 let f = self.func("qmatvec_q8_0_b4_rp_t");
18772 let cfg = LaunchConfig {
18773 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), t as u32, 1),
18774 block_dim: (32, ROWS_PER_BLOCK, 1),
18775 shared_mem_bytes: 0,
18776 };
18777 let __s_b = self.gpu.stream();
18778 let mut b = __s_b.launch_builder(&f);
18779 b.arg(w[0])
18780 .arg(w[1])
18781 .arg(w[2])
18782 .arg(w[3])
18783 .arg(aq)
18784 .arg(ad)
18785 .arg(y)
18786 .arg(&bc)
18787 .arg(&of);
18788 unsafe {
18789 b.launch(cfg)?;
18790 }
18791 Ok(())
18792 }
18793
18794 fn matvec_bf16_view_via_q8_mirror(
18804 &self,
18805 data: &cudarc::driver::CudaView<'_, u8>,
18806 x: &CudaSlice<f32>,
18807 y: &mut CudaSlice<f32>,
18808 in_f: usize,
18809 out_f: usize,
18810 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
18811 use cudarc::driver::DevicePtr;
18812 let key = {
18813 let s = self.gpu.stream();
18814 let (p, _g) = data.device_ptr(&s);
18815 (p as u64, in_f as u32, out_f as u32)
18816 };
18817 {
18818 let mut mirrors = self
18819 .w8_mirrors
18820 .lock()
18821 .map_err(|_| "w8 mirror map is poisoned")?;
18822 if !mirrors.contains_key(&key) {
18823 let mut interleaved = self.alloc_u8_uninit(out_f * Self::q8_0_row_bytes(in_f))?;
18824 self.encode_q8_0_from_bf16_view(data, &mut interleaved, in_f, out_f)?;
18825 let planar = self.build_q8_rp4_raw(&interleaved, in_f, out_f)?;
18826 mirrors.insert(key, planar);
18827 eprintln!("[w8-view] mirror built in_f={in_f} out_f={out_f}");
18831 }
18832 }
18833 let nblk = in_f / 32;
18834 {
18835 let mut act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
18836 if !act.contains_key(&in_f) {
18837 let aq = self.alloc_uninit::<i8>(in_f)?;
18838 let ad = self.alloc_uninit::<f32>(nblk)?;
18839 act.insert(in_f, (aq, ad));
18840 }
18841 let (aq, ad) = act.get_mut(&in_f).expect("just inserted");
18842 self.quantize_q8_1_into(x, 1, in_f, aq, ad)?;
18843 }
18844 let mirrors = self
18845 .w8_mirrors
18846 .lock()
18847 .map_err(|_| "w8 mirror map is poisoned")?;
18848 let act = self.w8_act.lock().map_err(|_| "w8 act map is poisoned")?;
18849 let mirror = mirrors.get(&key).expect("built above");
18850 let (aq, ad) = act.get(&in_f).expect("built above");
18851 self.qmatvec_mmvq_into(
18852 mirror,
18853 aq,
18854 ad,
18855 1,
18856 in_f,
18857 out_f,
18858 QT_Q8_0,
18859 Self::q8_0_row_bytes(in_f),
18860 1.0,
18861 true,
18862 y,
18863 )?;
18864 Ok(Some(()))
18865 }
18866
18867 pub fn matvec_bf16_into(
18868 &self,
18869 data: &CudaSlice<u8>,
18870 x: &CudaSlice<f32>,
18871 y: &mut CudaSlice<f32>,
18872 in_f: usize,
18873 out_f: usize,
18874 ) -> Result<(), Box<dyn std::error::Error>> {
18875 if data.len() != in_f * out_f * 2 || x.len() < in_f || in_f % 8 != 0 || y.len() < out_f {
18876 return Err(format!(
18877 "matvec_bf16_into geometry bytes={} x={} y={} in={in_f} out={out_f}",
18878 data.len(),
18879 x.len(),
18880 y.len()
18881 )
18882 .into());
18883 }
18884 if step_tp_w8_on() && w8_hybrid_on() && in_f % 32 == 0 && out_f >= 64 {
18891 if let Some(()) = self.matvec_bf16_via_q8_mirror(data, x, y, in_f, out_f)? {
18892 return Ok(());
18893 }
18894 }
18895 static X4: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18899 let x4 = *X4.get_or_init(|| std::env::var("MEMRA_DOWN_X4").as_deref() == Ok("1"))
18900 && in_f <= 2048;
18901 if x4 {
18902 let f = self.func("matvec_bf16_f32acc_x4");
18903 let cfg = LaunchConfig {
18904 grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
18905 block_dim: (mmv_block(), 1, 1),
18906 shared_mem_bytes: 0,
18907 };
18908 let (ini, outi) = (in_f as i32, out_f as i32);
18909 let __s_b = self.gpu.stream();
18910 let mut b = __s_b.launch_builder(&f);
18911 b.arg(data).arg(x).arg(y).arg(&ini).arg(&outi);
18912 unsafe {
18913 b.launch(cfg)?;
18914 }
18915 return Ok(());
18916 }
18917 let f = self.func("matvec_bf16_f32acc");
18918 let cfg = LaunchConfig {
18919 grid_dim: (out_f as u32, 1, 1),
18920 block_dim: (mmv_block(), 1, 1),
18921 shared_mem_bytes: 0,
18922 };
18923 let ini = in_f as i32;
18924 let __s_b = self.gpu.stream();
18925 let mut b = __s_b.launch_builder(&f);
18926 b.arg(data).arg(x).arg(y).arg(&ini);
18927 unsafe {
18928 b.launch(cfg)?;
18929 }
18930 Ok(())
18931 }
18932
18933 pub fn matvec_bf16_view_into(
18936 &self,
18937 data: &cudarc::driver::CudaView<'_, u8>,
18938 x: &CudaSlice<f32>,
18939 y: &mut CudaSlice<f32>,
18940 in_f: usize,
18941 out_f: usize,
18942 ) -> Result<(), Box<dyn std::error::Error>> {
18943 if data.len() != in_f * out_f * 2 || x.len() < in_f || in_f % 8 != 0 || y.len() < out_f {
18944 return Err(format!(
18945 "matvec_bf16_view_into geometry bytes={} x={} y={} in={in_f} out={out_f}",
18946 data.len(),
18947 x.len(),
18948 y.len()
18949 )
18950 .into());
18951 }
18952 if w8_view_on() && step_tp_w8_on() && w8_hybrid_on() && in_f % 32 == 0 && out_f >= 64 {
18953 if let Some(()) = self.matvec_bf16_view_via_q8_mirror(data, x, y, in_f, out_f)? {
18954 return Ok(());
18955 }
18956 }
18957 let f = self.func("matvec_bf16_f32acc");
18958 let cfg = LaunchConfig {
18959 grid_dim: (out_f as u32, 1, 1),
18960 block_dim: (mmv_block(), 1, 1),
18961 shared_mem_bytes: 0,
18962 };
18963 let ini = in_f as i32;
18964 let __s_b = self.gpu.stream();
18965 let mut b = __s_b.launch_builder(&f);
18966 b.arg(data).arg(x).arg(y).arg(&ini);
18967 unsafe {
18968 b.launch(cfg)?;
18969 }
18970 Ok(())
18971 }
18972
18973 pub fn matvec_bf16_raw_out(
18976 &self,
18977 w: &CudaSlice<u8>,
18978 x: &CudaSlice<f32>,
18979 y_raw: u64,
18980 in_f: usize,
18981 out_f: usize,
18982 ) -> Result<(), Box<dyn std::error::Error>> {
18983 if w.len() != in_f * out_f * 2 || x.len() < in_f || in_f % 8 != 0 || y_raw == 0 {
18984 return Err("matvec_bf16_raw_out geometry".into());
18985 }
18986 let f = self.func("matvec_bf16_f32acc");
18987 let cfg = LaunchConfig {
18988 grid_dim: (out_f as u32, 1, 1),
18989 block_dim: (mmv_block(), 1, 1),
18990 shared_mem_bytes: 0,
18991 };
18992 let ini = in_f as i32;
18993 let __s_b = self.gpu.stream();
18994 let mut b = __s_b.launch_builder(&f);
18995 b.arg(w).arg(x).arg(&y_raw).arg(&ini);
18996 unsafe {
18997 b.launch(cfg)?;
18998 }
18999 Ok(())
19000 }
19001
19002 pub fn add3_raw(
19006 &self,
19007 a: &CudaSlice<f32>,
19008 b: &CudaSlice<f32>,
19009 sh_raw: u64,
19010 scale_raw: u64,
19011 dst: &mut CudaSlice<f32>,
19012 n: usize,
19013 ) -> Result<(), Box<dyn std::error::Error>> {
19014 if a.len() < n || b.len() < n || dst.len() < n || sh_raw == 0 || scale_raw == 0 {
19015 return Err("add3_raw geometry".into());
19016 }
19017 let f = self.func("add3_f32");
19018 let cfg = LaunchConfig {
19019 grid_dim: ((n as u32).div_ceil(256), 1, 1),
19020 block_dim: (256, 1, 1),
19021 shared_mem_bytes: 0,
19022 };
19023 let ni = n as i32;
19024 let __s_b = self.gpu.stream();
19025 let mut bld = __s_b.launch_builder(&f);
19026 bld.arg(a)
19027 .arg(b)
19028 .arg(&sh_raw)
19029 .arg(&scale_raw)
19030 .arg(dst)
19031 .arg(&ni);
19032 unsafe {
19033 bld.launch(cfg)?;
19034 }
19035 Ok(())
19036 }
19037
19038 pub fn matvec_bf16_down_addscale_into(
19041 &self,
19042 w: &CudaSlice<u8>,
19043 x: &CudaSlice<f32>,
19044 scale: &CudaSlice<f32>,
19045 dst: &mut CudaSlice<f32>,
19046 in_f: usize,
19047 out_f: usize,
19048 ) -> Result<(), Box<dyn std::error::Error>> {
19049 if w.len() != in_f * out_f * 2
19050 || x.len() < in_f
19051 || in_f % 8 != 0
19052 || dst.len() < out_f
19053 || scale.is_empty()
19054 {
19055 return Err("matvec_bf16_down_addscale geometry".into());
19056 }
19057 let f = self.func("matvec_bf16_down_addscale");
19058 let cfg = LaunchConfig {
19059 grid_dim: (out_f as u32, 1, 1),
19060 block_dim: (mmv_block(), 1, 1),
19061 shared_mem_bytes: 0,
19062 };
19063 let ini = in_f as i32;
19064 let __s_b = self.gpu.stream();
19065 let mut b = __s_b.launch_builder(&f);
19066 b.arg(w).arg(x).arg(scale).arg(dst).arg(&ini);
19067 unsafe {
19068 b.launch(cfg)?;
19069 }
19070 Ok(())
19071 }
19072
19073 #[allow(clippy::too_many_arguments)]
19077 pub fn matvec_bf16_dual_silu_rows_into(
19078 &self,
19079 wg: &CudaSlice<u8>,
19080 wu: &CudaSlice<u8>,
19081 x: &CudaSlice<f32>,
19082 act: &mut CudaSlice<f32>,
19083 in_f: usize,
19084 out_f: usize,
19085 limit: Option<f32>,
19086 t: usize,
19087 ) -> Result<(), Box<dyn std::error::Error>> {
19088 if x.len() < t * in_f || act.len() < t * out_f || t == 0 || t > 32 {
19089 return Err("matvec_bf16_dual_silu_rows geometry".into());
19090 }
19091 let f = self.func("matvec_bf16_dual_silu_rows");
19092 let cfg = LaunchConfig {
19093 grid_dim: (out_f as u32, t as u32, 1),
19094 block_dim: (mmv_block(), 1, 1),
19095 shared_mem_bytes: 0,
19096 };
19097 let (ini, outi) = (in_f as i32, out_f as i32);
19098 let lim = limit.unwrap_or(0.0);
19099 let __s_b = self.gpu.stream();
19100 let mut b = __s_b.launch_builder(&f);
19101 b.arg(wg)
19102 .arg(wu)
19103 .arg(x)
19104 .arg(&mut *act)
19105 .arg(&ini)
19106 .arg(&outi)
19107 .arg(&lim);
19108 unsafe {
19109 b.launch(cfg)?;
19110 }
19111 Ok(())
19112 }
19113
19114 pub fn matvec_bf16_rows_into(
19116 &self,
19117 w: &CudaSlice<u8>,
19118 x: &CudaSlice<f32>,
19119 y: &mut CudaSlice<f32>,
19120 in_f: usize,
19121 out_f: usize,
19122 t: usize,
19123 ) -> Result<(), Box<dyn std::error::Error>> {
19124 if x.len() < t * in_f || y.len() < t * out_f || t == 0 || t > 32 || in_f % 8 != 0 {
19125 return Err("matvec_bf16_rows geometry".into());
19126 }
19127 if t >= 2 && t <= 32 && step_tp_w8_on() && w8_hybrid_on() && in_f % 32 == 0 && out_f >= 64 {
19132 if let Some(()) = self.matvec_bf16_via_q8_mirror_t(w, x, y, in_f, out_f, t)? {
19133 return Ok(());
19134 }
19135 }
19136 if t == 1 && step_tp_w8_on() && w8_hybrid_on() && in_f % 32 == 0 && out_f >= 64 {
19142 if let Some(()) = self.matvec_bf16_via_q8_mirror(w, x, y, in_f, out_f)? {
19143 return Ok(());
19144 }
19145 }
19146 let f = self.func("matvec_bf16_f32acc_x4_rows");
19147 let cfg = LaunchConfig {
19148 grid_dim: (out_f.div_ceil(4) as u32, t as u32, 1),
19149 block_dim: (mmv_block(), 1, 1),
19150 shared_mem_bytes: 0,
19151 };
19152 let (ini, outi) = (in_f as i32, out_f as i32);
19153 let __s_b = self.gpu.stream();
19154 let mut b = __s_b.launch_builder(&f);
19155 b.arg(w).arg(x).arg(&mut *y).arg(&ini).arg(&outi);
19156 unsafe {
19157 b.launch(cfg)?;
19158 }
19159 Ok(())
19160 }
19161
19162 pub fn matvec_bf16_dual_silu_into(
19163 &self,
19164 wg: &CudaSlice<u8>,
19165 wu: &CudaSlice<u8>,
19166 x: &CudaSlice<f32>,
19167 act: &mut CudaSlice<f32>,
19168 in_f: usize,
19169 out_f: usize,
19170 limit: Option<f32>,
19171 ) -> Result<(), Box<dyn std::error::Error>> {
19172 if wg.len() != in_f * out_f * 2
19173 || wu.len() != in_f * out_f * 2
19174 || x.len() < in_f
19175 || in_f % 8 != 0
19176 || act.len() < out_f
19177 {
19178 return Err("matvec_bf16_dual_silu geometry".into());
19179 }
19180 let f = self.func("matvec_bf16_dual_silu");
19181 let cfg = LaunchConfig {
19182 grid_dim: (out_f as u32, 1, 1),
19183 block_dim: (mmv_block(), 1, 1),
19184 shared_mem_bytes: 0,
19185 };
19186 let (ini, outi) = (in_f as i32, out_f as i32);
19187 let lim = limit.unwrap_or(0.0);
19188 let __s_b = self.gpu.stream();
19189 let mut b = __s_b.launch_builder(&f);
19190 b.arg(wg)
19191 .arg(wu)
19192 .arg(x)
19193 .arg(act)
19194 .arg(&ini)
19195 .arg(&outi)
19196 .arg(&lim);
19197 unsafe {
19198 b.launch(cfg)?;
19199 }
19200 Ok(())
19201 }
19202
19203 #[allow(clippy::too_many_arguments)]
19206 pub fn matvec_bf16_dual_view_into(
19207 &self,
19208 wg: &cudarc::driver::CudaView<'_, u8>,
19209 wu: &cudarc::driver::CudaView<'_, u8>,
19210 x: &CudaSlice<f32>,
19211 yg: &mut CudaSlice<f32>,
19212 yu: &mut CudaSlice<f32>,
19213 in_f: usize,
19214 out_f: usize,
19215 ) -> Result<(), Box<dyn std::error::Error>> {
19216 if wg.len() != in_f * out_f * 2
19217 || wu.len() != in_f * out_f * 2
19218 || x.len() < in_f
19219 || in_f % 8 != 0
19220 || yg.len() < out_f
19221 || yu.len() < out_f
19222 {
19223 return Err(format!(
19224 "matvec_bf16_dual_view_into geometry wg={} wu={} x={} in={in_f} out={out_f}",
19225 wg.len(),
19226 wu.len(),
19227 x.len()
19228 )
19229 .into());
19230 }
19231 let f = self.func("matvec_bf16_dual");
19232 let cfg = LaunchConfig {
19233 grid_dim: ((2 * out_f) as u32, 1, 1),
19234 block_dim: (mmv_block(), 1, 1),
19235 shared_mem_bytes: 0,
19236 };
19237 let (ini, outi) = (in_f as i32, out_f as i32);
19238 let __s_b = self.gpu.stream();
19239 let mut b = __s_b.launch_builder(&f);
19240 b.arg(wg)
19241 .arg(wu)
19242 .arg(x)
19243 .arg(yg)
19244 .arg(yu)
19245 .arg(&ini)
19246 .arg(&outi);
19247 unsafe {
19248 b.launch(cfg)?;
19249 }
19250 Ok(())
19251 }
19252
19253 #[allow(clippy::too_many_arguments)]
19255 pub fn matvec_bf16_dual_into(
19256 &self,
19257 wg: &CudaSlice<u8>,
19258 wu: &CudaSlice<u8>,
19259 x: &CudaSlice<f32>,
19260 yg: &mut CudaSlice<f32>,
19261 yu: &mut CudaSlice<f32>,
19262 in_f: usize,
19263 out_f: usize,
19264 ) -> Result<(), Box<dyn std::error::Error>> {
19265 if wg.len() != in_f * out_f * 2
19266 || wu.len() != in_f * out_f * 2
19267 || x.len() < in_f
19268 || in_f % 8 != 0
19269 || yg.len() < out_f
19270 || yu.len() < out_f
19271 {
19272 return Err(format!(
19273 "matvec_bf16_dual_into geometry wg={} wu={} x={} in={in_f} out={out_f}",
19274 wg.len(),
19275 wu.len(),
19276 x.len()
19277 )
19278 .into());
19279 }
19280 let f = self.func("matvec_bf16_dual");
19281 let cfg = LaunchConfig {
19282 grid_dim: ((2 * out_f) as u32, 1, 1),
19283 block_dim: (mmv_block(), 1, 1),
19284 shared_mem_bytes: 0,
19285 };
19286 let (ini, outi) = (in_f as i32, out_f as i32);
19287 let __s_b = self.gpu.stream();
19288 let mut b = __s_b.launch_builder(&f);
19289 b.arg(wg)
19290 .arg(wu)
19291 .arg(x)
19292 .arg(yg)
19293 .arg(yu)
19294 .arg(&ini)
19295 .arg(&outi);
19296 unsafe {
19297 b.launch(cfg)?;
19298 }
19299 Ok(())
19300 }
19301
19302 pub(crate) fn matvec_bf16_dual(
19305 &self,
19306 wg: &CudaSlice<u8>,
19307 wu: &CudaSlice<u8>,
19308 x: &CudaSlice<f32>,
19309 in_f: usize,
19310 out_f: usize,
19311 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
19312 if wg.len() != in_f * out_f * 2
19313 || wu.len() != in_f * out_f * 2
19314 || x.len() < in_f
19315 || in_f % 8 != 0
19316 {
19317 return Err(format!(
19318 "matvec_bf16_dual geometry wg={} wu={} x={} in={in_f} out={out_f}",
19319 wg.len(),
19320 wu.len(),
19321 x.len()
19322 )
19323 .into());
19324 }
19325 let mut yg = self.alloc_uninit::<f32>(out_f)?;
19326 let mut yu = self.alloc_uninit::<f32>(out_f)?;
19327 let f = self.func("matvec_bf16_dual");
19328 let cfg = LaunchConfig {
19329 grid_dim: ((2 * out_f) as u32, 1, 1),
19330 block_dim: (mmv_block(), 1, 1),
19331 shared_mem_bytes: 0,
19332 };
19333 let (ini, outi) = (in_f as i32, out_f as i32);
19334 let __s_b = self.gpu.stream();
19335 let mut b = __s_b.launch_builder(&f);
19336 b.arg(wg)
19337 .arg(wu)
19338 .arg(x)
19339 .arg(&mut yg)
19340 .arg(&mut yu)
19341 .arg(&ini)
19342 .arg(&outi);
19343 unsafe {
19344 b.launch(cfg)?;
19345 }
19346 Ok((yg, yu))
19347 }
19348
19349 #[allow(clippy::too_many_arguments)]
19350 fn linear_bf16_chunked_inner(
19351 &self,
19352 x: &CudaSlice<f32>,
19353 data: &CudaSlice<u8>,
19354 m: usize,
19355 in_f: usize,
19356 out_f: usize,
19357 exact: bool,
19358 canonical_chunk_rows: Option<usize>,
19359 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19360 const CHUNK_BYTES: usize = 256 << 20;
19361 if m == 1
19364 && !exact
19365 && canonical_chunk_rows.is_none()
19366 && in_f % 8 == 0
19367 && Self::bf16_mmv_on()
19368 {
19369 return self.matvec_bf16(data, x, in_f, out_f);
19370 }
19371 if m >= 16
19376 && !exact
19377 && canonical_chunk_rows.is_none()
19378 && data.len() == in_f * out_f * 2
19379 && crate::f16_ffi::pp_bf16_enabled()
19380 {
19381 if let Some(y) = self.bf16_tc_gemm(data, x, m, in_f, out_f)? {
19384 return Ok(y);
19385 }
19386 }
19387 let row_bytes = in_f
19388 .checked_mul(std::mem::size_of::<f32>())
19389 .ok_or("BF16 chunk row byte count overflow")?;
19390 if row_bytes == 0 || out_f == 0 {
19391 return Err("BF16 chunk dimensions must be nonzero".into());
19392 }
19393 let max_chunk_rows = (CHUNK_BYTES / row_bytes).max(1).min(out_f);
19394 let chunk_rows = match canonical_chunk_rows {
19395 Some(rows) if rows == 0 => {
19396 return Err("canonical BF16 chunk rows must be nonzero".into());
19397 }
19398 Some(rows) if rows > max_chunk_rows => {
19399 return Err(format!(
19400 "canonical BF16 chunk rows {rows} exceed the {max_chunk_rows}-row scratch limit"
19401 )
19402 .into());
19403 }
19404 Some(rows) if out_f % rows != 0 => {
19405 return Err(format!(
19406 "BF16 output width {out_f} is not divisible by canonical {rows}-row chunks"
19407 )
19408 .into());
19409 }
19410 Some(rows) => rows,
19411 None => max_chunk_rows,
19412 };
19413 if chunk_rows >= out_f {
19414 let wf32 = self.bf16_to_f32(&data.slice(0..in_f * out_f * 2), in_f * out_f)?;
19415 return if exact {
19416 self.linear_decode_exact(x, &wf32, m, in_f, out_f)
19417 } else {
19418 self.linear(x, &wf32, m, in_f, out_f)
19419 };
19420 }
19421 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
19422 let mut r0 = 0usize;
19423 while r0 < out_f {
19424 let rows = chunk_rows.min(out_f - r0);
19425 let wslice = data.slice(r0 * in_f * 2..(r0 + rows) * in_f * 2);
19426 let wf32 = self.bf16_to_f32(&wslice, in_f * rows)?;
19427 let yc = if exact {
19428 self.linear_decode_exact(x, &wf32, m, in_f, rows)?
19429 } else {
19430 self.linear(x, &wf32, m, in_f, rows)?
19431 };
19432 for mi in 0..m {
19434 let src = yc.slice(mi * rows..(mi + 1) * rows);
19435 let mut dst = y.slice_mut(mi * out_f + r0..mi * out_f + r0 + rows);
19436 self.gpu.stream().memcpy_dtod(&src, &mut dst)?;
19437 }
19438 r0 += rows;
19439 }
19440 Ok(y)
19441 }
19442
19443 pub fn linear_bf16_resident(
19447 &self,
19448 x: &CudaSlice<f32>,
19449 data: &CudaSlice<u8>,
19450 m: usize,
19451 in_f: usize,
19452 out_f: usize,
19453 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19454 if data.len() != in_f * out_f * 2 {
19455 return Err(format!("resident BF16 bytes {} != {out_f}x{in_f}x2", data.len()).into());
19456 }
19457 self.linear_bf16_chunked(x, data, m, in_f, out_f, false, None)
19458 }
19459
19460 pub fn linear_bf16_resident_canonical_rows(
19466 &self,
19467 x: &CudaSlice<f32>,
19468 data: &CudaSlice<u8>,
19469 m: usize,
19470 in_f: usize,
19471 out_f: usize,
19472 canonical_chunk_rows: usize,
19473 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19474 if data.len() != in_f * out_f * 2 {
19475 return Err(format!("resident BF16 bytes {} != {out_f}x{in_f}x2", data.len()).into());
19476 }
19477 self.linear_bf16_chunked(x, data, m, in_f, out_f, false, Some(canonical_chunk_rows))
19478 }
19479
19480 pub fn linear_f32_resident_canonical_rows(
19485 &self,
19486 x: &CudaSlice<f32>,
19487 data: &CudaSlice<f32>,
19488 m: usize,
19489 in_f: usize,
19490 out_f: usize,
19491 canonical_chunk_rows: usize,
19492 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19493 self.linear_f32_resident_canonical_rows_inner(
19494 x,
19495 data,
19496 m,
19497 in_f,
19498 out_f,
19499 canonical_chunk_rows,
19500 false,
19501 )
19502 }
19503
19504 pub fn linear_f32_resident_canonical_rows_strided(
19510 &self,
19511 x: &CudaSlice<f32>,
19512 data: &CudaSlice<f32>,
19513 m: usize,
19514 in_f: usize,
19515 out_f: usize,
19516 canonical_chunk_rows: usize,
19517 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19518 self.linear_f32_resident_canonical_rows_inner(
19519 x,
19520 data,
19521 m,
19522 in_f,
19523 out_f,
19524 canonical_chunk_rows,
19525 true,
19526 )
19527 }
19528
19529 fn linear_f32_resident_canonical_rows_inner(
19530 &self,
19531 x: &CudaSlice<f32>,
19532 data: &CudaSlice<f32>,
19533 m: usize,
19534 in_f: usize,
19535 out_f: usize,
19536 canonical_chunk_rows: usize,
19537 strided_output: bool,
19538 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19539 if data.len() != in_f * out_f {
19540 return Err(format!("resident F32 values {} != {out_f}x{in_f}", data.len()).into());
19541 }
19542 if canonical_chunk_rows == 0
19543 || canonical_chunk_rows > out_f
19544 || out_f % canonical_chunk_rows != 0
19545 {
19546 return Err(format!(
19547 "invalid canonical F32 chunk rows {canonical_chunk_rows} for output width {out_f}"
19548 )
19549 .into());
19550 }
19551 if canonical_chunk_rows == out_f {
19552 return self.linear(x, data, m, in_f, out_f);
19553 }
19554
19555 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
19556 let input = x.slice(0..x.len());
19557 for r0 in (0..out_f).step_by(canonical_chunk_rows) {
19558 let weights = data.slice(r0 * in_f..(r0 + canonical_chunk_rows) * in_f);
19559 if m == 1 {
19560 let mut destination = y.slice_mut(r0..r0 + canonical_chunk_rows);
19561 self.linear_device_into(
19562 &input,
19563 &weights,
19564 &mut destination,
19565 1,
19566 in_f,
19567 canonical_chunk_rows,
19568 )?;
19569 continue;
19570 }
19571 let chunk = self.linear_device(&input, &weights, m, in_f, canonical_chunk_rows)?;
19572 if strided_output {
19573 self.place_rows_strided(&chunk, &mut y, canonical_chunk_rows, m, out_f, r0)?;
19574 } else {
19575 for token in 0..m {
19576 let source = chunk
19577 .slice(token * canonical_chunk_rows..(token + 1) * canonical_chunk_rows);
19578 let mut destination =
19579 y.slice_mut(token * out_f + r0..token * out_f + r0 + canonical_chunk_rows);
19580 self.gpu.stream().memcpy_dtod(&source, &mut destination)?;
19581 }
19582 }
19583 }
19584 Ok(y)
19585 }
19586
19587 pub fn linear_f32_resident_canonical_rows_t1_into(
19592 &self,
19593 x: &CudaSlice<f32>,
19594 data: &CudaSlice<f32>,
19595 y: &mut CudaSlice<f32>,
19596 in_f: usize,
19597 out_f: usize,
19598 canonical_chunk_rows: usize,
19599 ) -> Result<(), Box<dyn std::error::Error>> {
19600 if data.len() != in_f * out_f {
19601 return Err(format!("resident F32 values {} != {out_f}x{in_f}", data.len()).into());
19602 }
19603 if y.len() != out_f || x.len() != in_f {
19604 return Err(format!(
19605 "resident F32 t1 shapes x={} y={} != in {in_f} out {out_f}",
19606 x.len(),
19607 y.len()
19608 )
19609 .into());
19610 }
19611 if canonical_chunk_rows == 0
19612 || canonical_chunk_rows > out_f
19613 || out_f % canonical_chunk_rows != 0
19614 {
19615 return Err(format!(
19616 "invalid canonical F32 chunk rows {canonical_chunk_rows} for output width {out_f}"
19617 )
19618 .into());
19619 }
19620 let input = x.slice(0..x.len());
19621 for r0 in (0..out_f).step_by(canonical_chunk_rows) {
19622 let weights = data.slice(r0 * in_f..(r0 + canonical_chunk_rows) * in_f);
19623 let mut destination = y.slice_mut(r0..r0 + canonical_chunk_rows);
19624 self.linear_device_into(
19625 &input,
19626 &weights,
19627 &mut destination,
19628 1,
19629 in_f,
19630 canonical_chunk_rows,
19631 )?;
19632 }
19633 Ok(())
19634 }
19635
19636 pub fn linear_t1_into(
19639 &self,
19640 x: &cudarc::driver::CudaView<'_, f32>,
19641 w: &cudarc::driver::CudaView<'_, f32>,
19642 y: &mut cudarc::driver::CudaViewMut<'_, f32>,
19643 in_f: usize,
19644 out_f: usize,
19645 ) -> Result<(), Box<dyn std::error::Error>> {
19646 self.linear_device_into(x, w, y, 1, in_f, out_f)
19647 }
19648
19649 pub fn linear_decode_exact(
19656 &self,
19657 x: &CudaSlice<f32>,
19658 w: &CudaSlice<f32>,
19659 m_tokens: usize,
19660 in_f: usize,
19661 out_f: usize,
19662 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19663 if m_tokens == 1 {
19664 return self.linear(x, w, 1, in_f, out_f);
19665 }
19666 let xv = self.view(x, m_tokens * in_f);
19667 let mut y = self.alloc_uninit::<f32>(m_tokens * out_f)?;
19668 for t in 0..m_tokens {
19669 let row = xv.slice(t * in_f..(t + 1) * in_f);
19670 let mut xr = self.alloc_uninit::<f32>(in_f)?;
19671 self.copy_view_into(&mut xr, 0, &row, in_f)?;
19672 let yr = self.linear(&xr, w, 1, in_f, out_f)?;
19673 self.copy_into(&mut y, t * out_f, &yr, out_f)?;
19674 }
19675 Ok(y)
19676 }
19677
19678 pub fn linear(
19679 &self,
19680 x: &CudaSlice<f32>,
19681 w: &CudaSlice<f32>,
19682 m_tokens: usize,
19683 in_f: usize,
19684 out_f: usize,
19685 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19686 self.linear_device(x, w, m_tokens, in_f, out_f)
19687 }
19688
19689 fn linear_device<I>(
19690 &self,
19691 x: &I,
19692 w: &I,
19693 m_tokens: usize,
19694 in_f: usize,
19695 out_f: usize,
19696 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>>
19697 where
19698 I: cudarc::driver::DevicePtr<f32>,
19699 {
19700 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)?;
19702 Ok(c)
19703 }
19704
19705 fn linear_device_into<I, O>(
19706 &self,
19707 x: &I,
19708 w: &I,
19709 c: &mut O,
19710 m_tokens: usize,
19711 in_f: usize,
19712 out_f: usize,
19713 ) -> Result<(), Box<dyn std::error::Error>>
19714 where
19715 I: cudarc::driver::DevicePtr<f32>,
19716 O: cudarc::driver::DevicePtrMut<f32>,
19717 {
19718 use cudarc::cublaslt::{Matmul, MatmulConfig};
19719 let cfg = MatmulConfig {
19720 transa: true,
19721 transb: false,
19722 transc: false,
19723 m: out_f as u64,
19724 n: m_tokens as u64,
19725 k: in_f as u64,
19726 alpha: 1.0,
19727 lda: in_f as i64,
19728 ldb: in_f as i64,
19729 beta: 0.0,
19730 ldc: out_f as i64,
19731 stride_a: None,
19732 stride_b: None,
19733 stride_c: None,
19734 stride_bias: None,
19735 batch_size: None,
19736 };
19737 let blas = self.gpu.blas();
19738 unsafe {
19739 blas.matmul(cfg, w, x, c, None, None)?;
19740 }
19741 Ok(())
19742 }
19743
19744 pub fn sdpa_naive(
19753 &self,
19754 q: &CudaSlice<f32>,
19755 k: &CudaSlice<f32>,
19756 v: &CudaSlice<f32>,
19757 o: &mut CudaSlice<f32>,
19758 head_dim: usize,
19759 n_head: usize,
19760 n_head_kv: usize,
19761 t: usize,
19762 t_kv: usize,
19763 scale: f32,
19764 causal: bool,
19765 ) -> Result<(), Box<dyn std::error::Error>> {
19766 if t_kv * 4 > SDPA_NAIVE_SMEM_MAX {
19767 return self.sdpa_naive_gmem(
19768 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
19769 );
19770 }
19771 let f = self.func("sdpa_naive_f32");
19772 let cfg = LaunchConfig {
19773 grid_dim: (n_head as u32, t as u32, 1),
19774 block_dim: (128, 1, 1),
19775 shared_mem_bytes: (t_kv * 4) as u32,
19776 };
19777 let (hd, nh, nhkv, ti, tkvi, cz) = (
19778 head_dim as i32,
19779 n_head as i32,
19780 n_head_kv as i32,
19781 t as i32,
19782 t_kv as i32,
19783 causal as i32,
19784 );
19785 let __s_b = self.gpu.stream();
19786 let mut b = __s_b.launch_builder(&f);
19787 b.arg(q)
19788 .arg(k)
19789 .arg(v)
19790 .arg(o)
19791 .arg(&hd)
19792 .arg(&nh)
19793 .arg(&nhkv)
19794 .arg(&ti)
19795 .arg(&tkvi)
19796 .arg(&scale)
19797 .arg(&cz);
19798 unsafe {
19799 b.launch(cfg)?;
19800 }
19801 Ok(())
19802 }
19803
19804 #[allow(clippy::too_many_arguments)]
19813 pub fn sdpa_naive_gmem(
19814 &self,
19815 q: &CudaSlice<f32>,
19816 k: &CudaSlice<f32>,
19817 v: &CudaSlice<f32>,
19818 o: &mut CudaSlice<f32>,
19819 head_dim: usize,
19820 n_head: usize,
19821 n_head_kv: usize,
19822 t: usize,
19823 t_kv: usize,
19824 scale: f32,
19825 causal: bool,
19826 ) -> Result<(), Box<dyn std::error::Error>> {
19827 let ws_len = n_head
19828 .checked_mul(t)
19829 .and_then(|x| x.checked_mul(t_kv))
19830 .ok_or("sdpa_naive_gmem: scores workspace size overflow")?;
19831 let ws_bytes = ws_len
19832 .checked_mul(std::mem::size_of::<f32>())
19833 .ok_or("sdpa_naive_gmem: scores workspace byte count overflow")?;
19834 if ws_bytes > SDPA_NAIVE_GMEM_WS_MAX {
19835 return Err(format!(
19836 "sdpa_naive_gmem: scores workspace {ws_bytes} bytes (heads {n_head} x T {t} x \
19837 T_kv {t_kv}) exceeds the {SDPA_NAIVE_GMEM_WS_MAX}-byte guard — this shape \
19838 needs a tiled/flash kernel, not the naive oracle"
19839 )
19840 .into());
19841 }
19842 let mut scores = self.uninit(ws_len)?;
19843 let f = self.func("sdpa_naive_gmem_f32");
19844 let cfg = LaunchConfig {
19845 grid_dim: (n_head as u32, t as u32, 1),
19846 block_dim: (128, 1, 1),
19847 shared_mem_bytes: 0,
19848 };
19849 let (hd, nh, nhkv, ti, tkvi, cz) = (
19850 head_dim as i32,
19851 n_head as i32,
19852 n_head_kv as i32,
19853 t as i32,
19854 t_kv as i32,
19855 causal as i32,
19856 );
19857 let __s_b = self.gpu.stream();
19858 let mut b = __s_b.launch_builder(&f);
19859 b.arg(q)
19860 .arg(k)
19861 .arg(v)
19862 .arg(o)
19863 .arg(&mut scores)
19864 .arg(&hd)
19865 .arg(&nh)
19866 .arg(&nhkv)
19867 .arg(&ti)
19868 .arg(&tkvi)
19869 .arg(&scale)
19870 .arg(&cz);
19871 unsafe {
19872 b.launch(cfg)?;
19873 }
19874 Ok(())
19875 }
19876
19877 #[allow(clippy::too_many_arguments)]
19882 pub fn sdpa_naive_island(
19883 &self,
19884 q: &CudaSlice<f32>,
19885 k: &CudaSlice<f32>,
19886 v: &CudaSlice<f32>,
19887 o: &mut CudaSlice<f32>,
19888 span_id: &CudaSlice<i32>,
19889 head_dim: usize,
19890 n_head: usize,
19891 n_head_kv: usize,
19892 t: usize,
19893 t_kv: usize,
19894 scale: f32,
19895 window: usize,
19896 ) -> Result<(), Box<dyn std::error::Error>> {
19897 let f = self.func("sdpa_naive_island_f32");
19898 let cfg = LaunchConfig {
19899 grid_dim: (n_head as u32, t as u32, 1),
19900 block_dim: (128, 1, 1),
19901 shared_mem_bytes: (t_kv * 4) as u32,
19902 };
19903 let (hd, nh, nhkv, ti, tkvi, wi) = (
19904 head_dim as i32,
19905 n_head as i32,
19906 n_head_kv as i32,
19907 t as i32,
19908 t_kv as i32,
19909 window as i32,
19910 );
19911 let __s_b = self.gpu.stream();
19912 let mut b = __s_b.launch_builder(&f);
19913 b.arg(q)
19914 .arg(k)
19915 .arg(v)
19916 .arg(o)
19917 .arg(span_id)
19918 .arg(&hd)
19919 .arg(&nh)
19920 .arg(&nhkv)
19921 .arg(&ti)
19922 .arg(&tkvi)
19923 .arg(&scale)
19924 .arg(&wi);
19925 unsafe {
19926 b.launch(cfg)?;
19927 }
19928 Ok(())
19929 }
19930
19931 #[allow(clippy::too_many_arguments)]
19933 pub fn sdpa_naive_w(
19934 &self,
19935 q: &CudaSlice<f32>,
19936 k: &CudaSlice<f32>,
19937 v: &CudaSlice<f32>,
19938 o: &mut CudaSlice<f32>,
19939 head_dim: usize,
19940 n_head: usize,
19941 n_head_kv: usize,
19942 t: usize,
19943 t_kv: usize,
19944 scale: f32,
19945 causal: bool,
19946 window: usize,
19947 ) -> Result<(), Box<dyn std::error::Error>> {
19948 let f = self.func("sdpa_naive_w_f32");
19949 let cfg = LaunchConfig {
19950 grid_dim: (n_head as u32, t as u32, 1),
19951 block_dim: (128, 1, 1),
19952 shared_mem_bytes: (t_kv * 4) as u32,
19953 };
19954 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
19955 head_dim as i32,
19956 n_head as i32,
19957 n_head_kv as i32,
19958 t as i32,
19959 t_kv as i32,
19960 causal as i32,
19961 window as i32,
19962 );
19963 let __s_b = self.gpu.stream();
19964 let mut b = __s_b.launch_builder(&f);
19965 b.arg(q)
19966 .arg(k)
19967 .arg(v)
19968 .arg(o)
19969 .arg(&hd)
19970 .arg(&nh)
19971 .arg(&nhkv)
19972 .arg(&ti)
19973 .arg(&tkvi)
19974 .arg(&scale)
19975 .arg(&cz)
19976 .arg(&wi);
19977 unsafe {
19978 b.launch(cfg)?;
19979 }
19980 Ok(())
19981 }
19982
19983 #[allow(clippy::too_many_arguments)]
19993 pub fn sdpa_naive_w_lo(
19994 &self,
19995 q: &CudaSlice<f32>,
19996 k: &CudaSlice<f32>,
19997 v: &CudaSlice<f32>,
19998 o: &mut CudaSlice<f32>,
19999 head_dim: usize,
20000 n_head: usize,
20001 n_head_kv: usize,
20002 t: usize,
20003 t_kv: usize,
20004 scale: f32,
20005 causal: bool,
20006 window: usize,
20007 ) -> Result<(), Box<dyn std::error::Error>> {
20008 let kv_lo = if window > 0 {
20009 (t_kv - t + 1).saturating_sub(window)
20010 } else {
20011 0
20012 };
20013 let smem = (t_kv - kv_lo) * 4;
20014 if smem > 48 * 1024 {
20015 return Err(format!(
20016 "sdpa_naive_w_lo: window {window} + T {t} rows need {smem} bytes of dynamic \
20017 shared memory (> 48KB launch bound) — this kernel clips the OLD side only; \
20018 a window this wide needs the multi-pass long-ctx kernel"
20019 )
20020 .into());
20021 }
20022 let f = self.func("sdpa_naive_w_lo_f32");
20023 let cfg = LaunchConfig {
20024 grid_dim: (n_head as u32, t as u32, 1),
20025 block_dim: (128, 1, 1),
20026 shared_mem_bytes: smem as u32,
20027 };
20028 let (hd, nh, nhkv, ti, tkvi, cz, wi, lo) = (
20029 head_dim as i32,
20030 n_head as i32,
20031 n_head_kv as i32,
20032 t as i32,
20033 t_kv as i32,
20034 causal as i32,
20035 window as i32,
20036 kv_lo as i32,
20037 );
20038 let __s_b = self.gpu.stream();
20039 let mut b = __s_b.launch_builder(&f);
20040 b.arg(q)
20041 .arg(k)
20042 .arg(v)
20043 .arg(o)
20044 .arg(&hd)
20045 .arg(&nh)
20046 .arg(&nhkv)
20047 .arg(&ti)
20048 .arg(&tkvi)
20049 .arg(&scale)
20050 .arg(&cz)
20051 .arg(&wi)
20052 .arg(&lo);
20053 unsafe {
20054 b.launch(cfg)?;
20055 }
20056 Ok(())
20057 }
20058
20059 pub fn sdpa_naive_view(
20061 &self,
20062 q: &CudaSlice<f32>,
20063 k: &cudarc::driver::CudaView<f32>,
20064 v: &cudarc::driver::CudaView<f32>,
20065 o: &mut CudaSlice<f32>,
20066 head_dim: usize,
20067 n_head: usize,
20068 n_head_kv: usize,
20069 t: usize,
20070 t_kv: usize,
20071 scale: f32,
20072 causal: bool,
20073 ) -> Result<(), Box<dyn std::error::Error>> {
20074 let f = self.func("sdpa_naive_f32");
20075 let cfg = LaunchConfig {
20076 grid_dim: (n_head as u32, t as u32, 1),
20077 block_dim: (128, 1, 1),
20078 shared_mem_bytes: (t_kv * 4) as u32,
20079 };
20080 let (hd, nh, nhkv, ti, tkvi, cz) = (
20081 head_dim as i32,
20082 n_head as i32,
20083 n_head_kv as i32,
20084 t as i32,
20085 t_kv as i32,
20086 causal as i32,
20087 );
20088 let __s_b = self.gpu.stream();
20089 let mut b = __s_b.launch_builder(&f);
20090 b.arg(q)
20091 .arg(k)
20092 .arg(v)
20093 .arg(o)
20094 .arg(&hd)
20095 .arg(&nh)
20096 .arg(&nhkv)
20097 .arg(&ti)
20098 .arg(&tkvi)
20099 .arg(&scale)
20100 .arg(&cz);
20101 unsafe {
20102 b.launch(cfg)?;
20103 }
20104 Ok(())
20105 }
20106
20107 #[allow(clippy::too_many_arguments)]
20115 pub fn fa_dequant_kv_view_f32(
20116 &self,
20117 k: &cudarc::driver::CudaView<u8>,
20118 v: &cudarc::driver::CudaView<u8>,
20119 kf: &mut CudaSlice<f32>,
20120 vf: &mut CudaSlice<f32>,
20121 kv_dim_k: usize,
20122 kv_dim_v: usize,
20123 t_kv: usize,
20124 k_tok_bytes: usize,
20125 v_tok_bytes: usize,
20126 g: bool,
20127 ) -> Result<(), Box<dyn std::error::Error>> {
20128 let f = if g {
20129 self.func_g("fa_dequant_kv_ws_f32")
20130 } else {
20131 self.func("fa_dequant_kv_ws_f32")
20132 };
20133 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
20134 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
20135 let cfg = LaunchConfig {
20136 grid_dim: (nblk.max(1), 1, 1),
20137 block_dim: (256, 1, 1),
20138 shared_mem_bytes: 0,
20139 };
20140 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
20141 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
20142 let __s_b = self.gpu.stream();
20143 let mut b = __s_b.launch_builder(&f);
20144 b.arg(k)
20145 .arg(v)
20146 .arg(&mut *kf)
20147 .arg(&mut *vf)
20148 .arg(&kdk)
20149 .arg(&kdv)
20150 .arg(&tkvi)
20151 .arg(&ktb)
20152 .arg(&vtb);
20153 unsafe {
20154 b.launch(cfg)?;
20155 }
20156 Ok(())
20157 }
20158
20159 #[allow(clippy::too_many_arguments)]
20160 pub fn sdpa_naive_quantized_view(
20161 &self,
20162 q: &CudaSlice<f32>,
20163 k: &cudarc::driver::CudaView<u8>,
20164 v: &cudarc::driver::CudaView<u8>,
20165 o: &mut CudaSlice<f32>,
20166 head_dim: usize,
20167 n_head: usize,
20168 n_head_kv: usize,
20169 t: usize,
20170 t_kv: usize,
20171 scale: f32,
20172 causal: bool,
20173 k_tok_bytes: usize,
20174 v_tok_bytes: usize,
20175 ) -> Result<(), Box<dyn std::error::Error>> {
20176 let kv_dim = n_head_kv * head_dim;
20177 let mut kf = self.uninit(t_kv * kv_dim)?;
20178 let mut vf = self.uninit(t_kv * kv_dim)?;
20179 let f = self.func("fa_dequant_kv_ws_f32");
20180 let total = (2 * t_kv * kv_dim) as u64;
20181 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
20182 let cfg = LaunchConfig {
20183 grid_dim: (nblk.max(1), 1, 1),
20184 block_dim: (256, 1, 1),
20185 shared_mem_bytes: 0,
20186 };
20187 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
20188 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
20189 let __s_b = self.gpu.stream();
20190 let mut b = __s_b.launch_builder(&f);
20191 b.arg(k)
20192 .arg(v)
20193 .arg(&mut kf)
20194 .arg(&mut vf)
20195 .arg(&kv_dim_i)
20196 .arg(&kv_dim_i)
20197 .arg(&t_kv_i)
20198 .arg(&k_tok_bytes_i)
20199 .arg(&v_tok_bytes_i);
20200 unsafe { b.launch(cfg)? };
20201 self.sdpa_naive(
20202 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
20203 )
20204 }
20205
20206 #[allow(clippy::too_many_arguments)]
20218 pub fn sdpa_naive_w_quantized_view(
20219 &self,
20220 q: &CudaSlice<f32>,
20221 k: &cudarc::driver::CudaView<u8>,
20222 v: &cudarc::driver::CudaView<u8>,
20223 o: &mut CudaSlice<f32>,
20224 head_dim: usize,
20225 n_head: usize,
20226 n_head_kv: usize,
20227 t: usize,
20228 t_kv: usize,
20229 scale: f32,
20230 causal: bool,
20231 window: usize,
20232 k_tok_bytes: usize,
20233 v_tok_bytes: usize,
20234 ) -> Result<(), Box<dyn std::error::Error>> {
20235 let kv_dim = n_head_kv * head_dim;
20236 let mut kf = self.uninit(t_kv * kv_dim)?;
20237 let mut vf = self.uninit(t_kv * kv_dim)?;
20238 let f = self.func("fa_dequant_kv_ws_f32");
20239 let total = (2 * t_kv * kv_dim) as u64;
20240 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
20241 let cfg = LaunchConfig {
20242 grid_dim: (nblk.max(1), 1, 1),
20243 block_dim: (256, 1, 1),
20244 shared_mem_bytes: 0,
20245 };
20246 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
20247 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
20248 let __s_b = self.gpu.stream();
20249 let mut b = __s_b.launch_builder(&f);
20250 b.arg(k)
20251 .arg(v)
20252 .arg(&mut kf)
20253 .arg(&mut vf)
20254 .arg(&kv_dim_i)
20255 .arg(&kv_dim_i)
20256 .arg(&t_kv_i)
20257 .arg(&k_tok_bytes_i)
20258 .arg(&v_tok_bytes_i);
20259 unsafe { b.launch(cfg)? };
20260 self.sdpa_naive_w(
20261 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
20262 )
20263 }
20264
20265 pub fn fa_prefill(
20269 &self,
20270 q: &CudaSlice<f32>,
20271 k: &CudaSlice<f32>,
20272 v: &CudaSlice<f32>,
20273 o: &mut CudaSlice<f32>,
20274 head_dim: usize,
20275 n_head: usize,
20276 n_head_kv: usize,
20277 t: usize,
20278 t_kv: usize,
20279 scale: f32,
20280 causal: bool,
20281 ) -> Result<(), Box<dyn std::error::Error>> {
20282 if portable_mma_gated() {
20283 return self.sdpa_naive(
20284 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
20285 );
20286 }
20287 let fa3_on = head_dim == 256
20295 && causal
20296 && t == t_kv
20297 && match std::env::var("MEMRA_FA3").as_deref() {
20298 Ok("0") => false,
20299 Ok("1") => {
20303 refuse_portable_force("MEMRA_FA3=1", "the sm_90a fa3/bf16 kernels");
20304 true
20305 }
20306 _ => cfg!(memra_hopper_mma),
20307 };
20308 if fa3_on {
20309 let n = t * n_head * head_dim;
20310 let nkv = t * n_head_kv * head_dim;
20311 let mut q16 = self.alloc_u8_uninit(n * 2)?;
20312 let mut k16 = self.alloc_u8_uninit(nkv * 2)?;
20313 let mut v16 = self.alloc_u8_uninit(nkv * 2)?;
20314 self.f32_to_bf16_into(q, &mut q16, n)?;
20315 self.f32_to_bf16_into(k, &mut k16, nkv)?;
20316 self.f32_to_bf16_into(v, &mut v16, nkv)?;
20317 let rc = {
20318 use cudarc::driver::{DevicePtr, DevicePtrMut};
20319 let stream = self.gpu.stream();
20320 let (qp, _g1) = q16.device_ptr(&stream);
20321 let (kp, _g2) = k16.device_ptr(&stream);
20322 let (vp, _g3) = v16.device_ptr(&stream);
20323 let (op, _g4) = o.device_ptr_mut(&stream);
20324 unsafe {
20325 memra_fa3_prefill(
20326 qp as *const core::ffi::c_void,
20327 kp as *const core::ffi::c_void,
20328 vp as *const core::ffi::c_void,
20329 op as *mut f32,
20330 t as i32,
20331 n_head as i32,
20332 n_head_kv as i32,
20333 head_dim as i32,
20334 scale,
20335 stream.cu_stream() as *mut core::ffi::c_void,
20336 )
20337 }
20338 };
20339 if rc != 0 {
20340 return Err(format!("memra_fa3_prefill rc={rc}").into());
20341 }
20342 return Ok(());
20343 }
20344 static FA_P1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20349 let fa_p1 = *FA_P1.get_or_init(|| std::env::var("MEMRA_FA_P1").as_deref() == Ok("1"));
20350 if fa_p1 && head_dim == 256 && !std::env::var("MEMRA_FA_FLOOR").is_ok() {
20351 const BLOCK_Q: usize = 64;
20352 const BKX: usize = 32;
20353 let f = self.func("fa_prefill_bf16_p1");
20354 let shmem = (2 * (2 * BKX * head_dim + BLOCK_Q * BKX)
20355 + 4 * (BLOCK_Q * BKX + 2 * BLOCK_Q)) as u32;
20356 use cudarc::driver::sys::CUfunction_attribute_enum as A;
20357 f.set_attribute(
20358 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
20359 shmem as i32,
20360 )?;
20361 let cfg = LaunchConfig {
20362 grid_dim: (
20363 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
20364 n_head as u32,
20365 1,
20366 ),
20367 block_dim: (32, 4, 1),
20368 shared_mem_bytes: shmem,
20369 };
20370 let (hd, nh, nhkv, ti, tkvi, cz) = (
20371 head_dim as i32,
20372 n_head as i32,
20373 n_head_kv as i32,
20374 t as i32,
20375 t_kv as i32,
20376 causal as i32,
20377 );
20378 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
20379 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
20380 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
20381 let __s_b = self.gpu.stream();
20382 let mut b = __s_b.launch_builder(&f);
20383 b.arg(&qb)
20384 .arg(&kb)
20385 .arg(&vb)
20386 .arg(o)
20387 .arg(&hd)
20388 .arg(&nh)
20389 .arg(&nhkv)
20390 .arg(&ti)
20391 .arg(&tkvi)
20392 .arg(&scale)
20393 .arg(&cz);
20394 unsafe {
20395 b.launch(cfg)?;
20396 }
20397 return Ok(());
20398 }
20399 const BK: usize = 32;
20405 let w2 = std::env::var("MEMRA_FA_PP_W2").as_deref() == Ok("1");
20408 let (block_q, warps, w2_sfx): (usize, u32, &str) =
20409 if w2 { (32, 2, "_w2") } else { (64, 4, "") };
20410 let hd_sfx = fa_hd_suffix(head_dim)?;
20414 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
20415 let bf16kv = !floor && !w2 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0");
20420 let (kb16, vb16) = if bf16kv {
20421 let n = t_kv * n_head_kv * head_dim;
20422 let mut kb = self.alloc_u8_uninit(n * 2)?;
20423 let mut vb = self.alloc_u8_uninit(n * 2)?;
20424 let fcv = self.func("f32_to_bf16_bulk");
20425 let ni = n as i64;
20426 let cfgc = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
20427 let __s_b = self.gpu.stream();
20428 let mut b = __s_b.launch_builder(&fcv);
20429 b.arg(k).arg(&mut kb).arg(&ni);
20430 unsafe {
20431 b.launch(cfgc)?;
20432 }
20433 let __s_b = self.gpu.stream();
20434 let mut b = __s_b.launch_builder(&fcv);
20435 b.arg(v).arg(&mut vb).arg(&ni);
20436 unsafe {
20437 b.launch(cfgc)?;
20438 }
20439 (Some(kb), Some(vb))
20440 } else {
20441 (None, None)
20442 };
20443 let f = self.func(&if bf16kv {
20444 format!("fa_prefill_bf16kv_pp{hd_sfx}")
20445 } else {
20446 format!(
20447 "fa_prefill_f32{}{}{hd_sfx}",
20448 if floor { "" } else { "_pp" },
20449 if floor { "" } else { w2_sfx }
20450 )
20451 });
20452 let kv_stages = if bf16kv { 2 } else { 1 };
20455 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
20456 + 4 * (block_q * BK + 2 * block_q)) as u32;
20457 use cudarc::driver::sys::CUfunction_attribute_enum as A;
20458 f.set_attribute(
20459 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
20460 shmem as i32,
20461 )?;
20462 let cfg = LaunchConfig {
20463 grid_dim: (
20464 (t as u32 + block_q as u32 - 1) / block_q as u32,
20465 n_head as u32,
20466 1,
20467 ),
20468 block_dim: (32, warps, 1),
20469 shared_mem_bytes: shmem,
20470 };
20471 let (hd, nh, nhkv, ti, tkvi, cz) = (
20472 head_dim as i32,
20473 n_head as i32,
20474 n_head_kv as i32,
20475 t as i32,
20476 t_kv as i32,
20477 causal as i32,
20478 );
20479 let __s_b = self.gpu.stream();
20480 let mut b = __s_b.launch_builder(&f);
20481 b.arg(q);
20482 match (&kb16, &vb16) {
20483 (Some(kb), Some(vb)) => {
20484 b.arg(kb).arg(vb);
20485 }
20486 _ => {
20487 b.arg(k).arg(v);
20488 }
20489 }
20490 b.arg(o)
20491 .arg(&hd)
20492 .arg(&nh)
20493 .arg(&nhkv)
20494 .arg(&ti)
20495 .arg(&tkvi)
20496 .arg(&scale)
20497 .arg(&cz);
20498 unsafe {
20499 b.launch(cfg)?;
20500 }
20501 Ok(())
20502 }
20503
20504 #[allow(clippy::too_many_arguments)]
20508 pub fn fa_prefill_w(
20509 &self,
20510 q: &CudaSlice<f32>,
20511 k: &CudaSlice<f32>,
20512 v: &CudaSlice<f32>,
20513 o: &mut CudaSlice<f32>,
20514 head_dim: usize,
20515 n_head: usize,
20516 n_head_kv: usize,
20517 t: usize,
20518 t_kv: usize,
20519 scale: f32,
20520 causal: bool,
20521 window: usize,
20522 ) -> Result<(), Box<dyn std::error::Error>> {
20523 if portable_mma_gated() {
20526 return self.sdpa_naive_w(
20527 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
20528 );
20529 }
20530 static FAW_F32: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20534 let faw_f32 =
20535 *FAW_F32.get_or_init(|| std::env::var("MEMRA_FAW_STAGE").as_deref() == Ok("f32"));
20536 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
20537 self.fa_prefill_w_arm(
20538 q,
20539 k,
20540 v,
20541 o,
20542 head_dim,
20543 n_head,
20544 n_head_kv,
20545 t,
20546 t_kv,
20547 scale,
20548 causal,
20549 window,
20550 floor || faw_f32,
20551 floor,
20552 )
20553 }
20554
20555 #[allow(clippy::too_many_arguments)]
20558 pub fn fa_prefill_w_pre(
20559 &self,
20560 qb: &CudaSlice<u8>,
20561 kb: &CudaSlice<u8>,
20562 vb: &CudaSlice<u8>,
20563 o: &mut CudaSlice<f32>,
20564 head_dim: usize,
20565 n_head: usize,
20566 n_head_kv: usize,
20567 t: usize,
20568 t_kv: usize,
20569 scale: f32,
20570 causal: bool,
20571 window: usize,
20572 v_f16: bool,
20573 ) -> Result<(), Box<dyn std::error::Error>> {
20574 const BLOCK_Q: usize = 64;
20575 const BK: usize = 32;
20576 debug_assert_eq!(head_dim, 256);
20577 let hp = fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
20578 debug_assert!(!v_f16 || hp, "f16 V emitted but the SWA hp arm is off");
20579 if hp {
20580 const BLOCK_QH: usize = 32;
20581 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
20584 let vh: &CudaSlice<u8> = if v_f16 {
20585 vb
20586 } else {
20587 let n = t_kv * n_head_kv * head_dim;
20588 if vguard.as_ref().map(|b| b.len() < n * 2).unwrap_or(true) {
20589 *vguard = Some(self.alloc_uninit::<u8>(n * 2)?);
20590 }
20591 self.bf16_to_f16_into(vb, n, vguard.as_mut().unwrap())?;
20592 vguard.as_ref().unwrap()
20593 };
20594 let f = self.func("fa_prefill_w_bf16_p1h2");
20595 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK) + 4 * (2 * BLOCK_QH)) as u32;
20596 use cudarc::driver::sys::CUfunction_attribute_enum as A;
20597 f.set_attribute(
20598 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
20599 shmem as i32,
20600 )?;
20601 let cfg = LaunchConfig {
20602 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
20603 block_dim: (32, 4, 1),
20604 shared_mem_bytes: shmem,
20605 };
20606 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
20607 head_dim as i32,
20608 n_head as i32,
20609 n_head_kv as i32,
20610 t as i32,
20611 t_kv as i32,
20612 causal as i32,
20613 window as i32,
20614 );
20615 let __s_b = self.gpu.stream();
20616 let mut b = __s_b.launch_builder(&f);
20617 b.arg(qb)
20618 .arg(kb)
20619 .arg(vh)
20620 .arg(o)
20621 .arg(&hd)
20622 .arg(&nh)
20623 .arg(&nhkv)
20624 .arg(&ti)
20625 .arg(&tkvi)
20626 .arg(&scale)
20627 .arg(&cz)
20628 .arg(&wi);
20629 unsafe {
20630 b.launch(cfg)?;
20631 }
20632 return Ok(());
20633 }
20634 let f = self.func("fa_prefill_w_bf16_p1");
20635 let shmem =
20636 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
20637 use cudarc::driver::sys::CUfunction_attribute_enum as A;
20638 f.set_attribute(
20639 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
20640 shmem as i32,
20641 )?;
20642 let cfg = LaunchConfig {
20643 grid_dim: (
20644 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
20645 n_head as u32,
20646 1,
20647 ),
20648 block_dim: (32, 4, 1),
20649 shared_mem_bytes: shmem,
20650 };
20651 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
20652 head_dim as i32,
20653 n_head as i32,
20654 n_head_kv as i32,
20655 t as i32,
20656 t_kv as i32,
20657 causal as i32,
20658 window as i32,
20659 );
20660 let __s_b = self.gpu.stream();
20661 let mut b = __s_b.launch_builder(&f);
20662 b.arg(qb)
20663 .arg(kb)
20664 .arg(vb)
20665 .arg(o)
20666 .arg(&hd)
20667 .arg(&nh)
20668 .arg(&nhkv)
20669 .arg(&ti)
20670 .arg(&tkvi)
20671 .arg(&scale)
20672 .arg(&cz)
20673 .arg(&wi);
20674 unsafe {
20675 b.launch(cfg)?;
20676 }
20677 Ok(())
20678 }
20679
20680 #[allow(clippy::too_many_arguments)]
20682 pub fn fa_prefill_w_arm(
20683 &self,
20684 q: &CudaSlice<f32>,
20685 k: &CudaSlice<f32>,
20686 v: &CudaSlice<f32>,
20687 o: &mut CudaSlice<f32>,
20688 head_dim: usize,
20689 n_head: usize,
20690 n_head_kv: usize,
20691 t: usize,
20692 t_kv: usize,
20693 scale: f32,
20694 causal: bool,
20695 window: usize,
20696 f32_stage: bool,
20697 floor: bool,
20698 ) -> Result<(), Box<dyn std::error::Error>> {
20699 const BLOCK_Q: usize = 64;
20700 const BK: usize = 32;
20701 debug_assert_eq!(head_dim, 256, "fa_prefill_w is stamped hd256 only");
20702 static P1_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20706 let p1 = !floor
20707 && !f32_stage
20708 && *P1_ON.get_or_init(|| {
20709 std::env::var("MEMRA_FAW_P1")
20710 .map(|v| v != "0")
20711 .unwrap_or(true)
20712 });
20713 let hp =
20714 p1 && fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
20715 if hp {
20716 const BLOCK_QH: usize = 32;
20717 let f = self.func("fa_prefill_w_bf16_p1h2");
20718 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK) + 4 * (2 * BLOCK_QH)) as u32;
20719 use cudarc::driver::sys::CUfunction_attribute_enum as A;
20720 f.set_attribute(
20721 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
20722 shmem as i32,
20723 )?;
20724 let cfg = LaunchConfig {
20725 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
20726 block_dim: (32, 4, 1),
20727 shared_mem_bytes: shmem,
20728 };
20729 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
20730 head_dim as i32,
20731 n_head as i32,
20732 n_head_kv as i32,
20733 t as i32,
20734 t_kv as i32,
20735 causal as i32,
20736 window as i32,
20737 );
20738 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
20739 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
20740 let vh = self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?;
20741 let __s_b = self.gpu.stream();
20742 let mut b = __s_b.launch_builder(&f);
20743 b.arg(&qb)
20744 .arg(&kb)
20745 .arg(&vh)
20746 .arg(o)
20747 .arg(&hd)
20748 .arg(&nh)
20749 .arg(&nhkv)
20750 .arg(&ti)
20751 .arg(&tkvi)
20752 .arg(&scale)
20753 .arg(&cz)
20754 .arg(&wi);
20755 unsafe {
20756 b.launch(cfg)?;
20757 }
20758 return Ok(());
20759 }
20760 if p1 {
20761 let f = self.func("fa_prefill_w_bf16_p1");
20762 let shmem =
20763 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
20764 use cudarc::driver::sys::CUfunction_attribute_enum as A;
20765 f.set_attribute(
20766 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
20767 shmem as i32,
20768 )?;
20769 let cfg = LaunchConfig {
20770 grid_dim: (
20771 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
20772 n_head as u32,
20773 1,
20774 ),
20775 block_dim: (32, 4, 1),
20776 shared_mem_bytes: shmem,
20777 };
20778 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
20779 head_dim as i32,
20780 n_head as i32,
20781 n_head_kv as i32,
20782 t as i32,
20783 t_kv as i32,
20784 causal as i32,
20785 window as i32,
20786 );
20787 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
20788 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
20789 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
20790 let __s_b = self.gpu.stream();
20791 let mut b = __s_b.launch_builder(&f);
20792 b.arg(&qb)
20793 .arg(&kb)
20794 .arg(&vb)
20795 .arg(o)
20796 .arg(&hd)
20797 .arg(&nh)
20798 .arg(&nhkv)
20799 .arg(&ti)
20800 .arg(&tkvi)
20801 .arg(&scale)
20802 .arg(&cz)
20803 .arg(&wi);
20804 unsafe {
20805 b.launch(cfg)?;
20806 }
20807 return Ok(());
20808 }
20809 static G4_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20812 let g4 = !floor
20813 && !f32_stage
20814 && n_head_kv == 1
20815 && n_head % 4 == 0
20816 && *G4_ON.get_or_init(|| {
20817 std::env::var("MEMRA_FAW_G4")
20818 .map(|v| v != "0")
20819 .unwrap_or(true)
20820 });
20821 if g4 {
20822 const SP_M: usize = 16;
20823 static O2_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20826 let o2 = *O2_ON.get_or_init(|| {
20827 std::env::var("MEMRA_FAW_O2")
20828 .map(|v| v != "0")
20829 .unwrap_or(true)
20830 });
20831 let f = self.func(if o2 {
20832 "fa_prefill_w_bf16_g4o2"
20833 } else {
20834 "fa_prefill_w_bf16_g4"
20835 });
20836 let shmem = if o2 {
20837 (2 * (4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M)) as u32
20838 } else {
20839 (2 * (2 * BK * head_dim + 4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M))
20840 as u32
20841 };
20842 use cudarc::driver::sys::CUfunction_attribute_enum as A;
20843 f.set_attribute(
20844 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
20845 shmem as i32,
20846 )?;
20847 let cfg = LaunchConfig {
20848 grid_dim: ((t as u32).div_ceil(SP_M as u32), (n_head / 4) as u32, 1),
20849 block_dim: (32, 4, 1),
20850 shared_mem_bytes: shmem,
20851 };
20852 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
20853 head_dim as i32,
20854 n_head as i32,
20855 n_head_kv as i32,
20856 t as i32,
20857 t_kv as i32,
20858 causal as i32,
20859 window as i32,
20860 );
20861 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
20862 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
20863 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
20864 let __s_b = self.gpu.stream();
20865 let mut b = __s_b.launch_builder(&f);
20866 b.arg(&qb)
20867 .arg(&kb)
20868 .arg(&vb)
20869 .arg(o)
20870 .arg(&hd)
20871 .arg(&nh)
20872 .arg(&nhkv)
20873 .arg(&ti)
20874 .arg(&tkvi)
20875 .arg(&scale)
20876 .arg(&cz)
20877 .arg(&wi);
20878 unsafe {
20879 b.launch(cfg)?;
20880 }
20881 return Ok(());
20882 }
20883 let f = self.func(if floor {
20884 "fa_prefill_w_f32"
20885 } else if f32_stage {
20886 "fa_prefill_w_f32_pp"
20887 } else {
20888 "fa_prefill_w_bf16_pp"
20889 });
20890 let shmem =
20891 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
20892 use cudarc::driver::sys::CUfunction_attribute_enum as A;
20893 f.set_attribute(
20894 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
20895 shmem as i32,
20896 )?;
20897 let cfg = LaunchConfig {
20898 grid_dim: (
20899 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
20900 n_head as u32,
20901 1,
20902 ),
20903 block_dim: (32, 4, 1),
20904 shared_mem_bytes: shmem,
20905 };
20906 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
20907 head_dim as i32,
20908 n_head as i32,
20909 n_head_kv as i32,
20910 t as i32,
20911 t_kv as i32,
20912 causal as i32,
20913 window as i32,
20914 );
20915 if f32_stage {
20916 let __s_b = self.gpu.stream();
20917 let mut b = __s_b.launch_builder(&f);
20918 b.arg(q)
20919 .arg(k)
20920 .arg(v)
20921 .arg(o)
20922 .arg(&hd)
20923 .arg(&nh)
20924 .arg(&nhkv)
20925 .arg(&ti)
20926 .arg(&tkvi)
20927 .arg(&scale)
20928 .arg(&cz)
20929 .arg(&wi);
20930 unsafe {
20931 b.launch(cfg)?;
20932 }
20933 } else {
20934 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
20935 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
20936 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
20937 let __s_b = self.gpu.stream();
20938 let mut b = __s_b.launch_builder(&f);
20939 b.arg(&qb)
20940 .arg(&kb)
20941 .arg(&vb)
20942 .arg(o)
20943 .arg(&hd)
20944 .arg(&nh)
20945 .arg(&nhkv)
20946 .arg(&ti)
20947 .arg(&tkvi)
20948 .arg(&scale)
20949 .arg(&cz)
20950 .arg(&wi);
20951 unsafe {
20952 b.launch(cfg)?;
20953 }
20954 }
20955 Ok(())
20956 }
20957
20958 #[allow(clippy::too_many_arguments)]
20962 pub fn fa_prefill_hd512(
20963 &self,
20964 q: &CudaSlice<f32>,
20965 k: &CudaSlice<f32>,
20966 v: &CudaSlice<f32>,
20967 o: &mut CudaSlice<f32>,
20968 head_dim: usize,
20969 n_head: usize,
20970 n_head_kv: usize,
20971 t: usize,
20972 t_kv: usize,
20973 scale: f32,
20974 causal: bool,
20975 ) -> Result<(), Box<dyn std::error::Error>> {
20976 if portable_mma_gated() {
20978 return self.sdpa_naive(
20979 q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
20980 );
20981 }
20982 static F32_STAGE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20988 let f32_stage =
20989 *F32_STAGE.get_or_init(|| std::env::var("MEMRA_FA512_STAGE").as_deref() == Ok("f32"));
20990 static SP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20994 let sp = !f32_stage
20995 && *SP_ON.get_or_init(|| {
20996 std::env::var("MEMRA_FA512_SP")
20997 .map(|v| v != "0")
20998 .unwrap_or(true)
20999 });
21000 self.fa_prefill_hd512_arm(
21001 q,
21002 k,
21003 v,
21004 o,
21005 head_dim,
21006 n_head,
21007 n_head_kv,
21008 t,
21009 t_kv,
21010 scale,
21011 causal,
21012 f32_stage,
21013 sp,
21014 sp && fa_f16pv_on(),
21015 )
21016 }
21017
21018 #[allow(clippy::too_many_arguments)]
21020 pub fn fa_prefill_hd512_pre(
21021 &self,
21022 qb: &CudaSlice<u8>,
21023 kb: &CudaSlice<u8>,
21024 vb: &CudaSlice<u8>,
21025 o: &mut CudaSlice<f32>,
21026 head_dim: usize,
21027 n_head: usize,
21028 n_head_kv: usize,
21029 t: usize,
21030 t_kv: usize,
21031 scale: f32,
21032 causal: bool,
21033 v_f16: bool,
21034 ) -> Result<(), Box<dyn std::error::Error>> {
21035 debug_assert_eq!(head_dim, 512);
21036 const SP_M: usize = 16;
21037 const BKS: usize = 32;
21038 let f16pv = fa_f16pv_on();
21042 let nw = if f16pv { fa512_wide_warps() } else { 2 };
21043 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
21044 debug_assert!(!v_f16 || f16pv, "f16 V emitted without the door on");
21045 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
21046 let vref: &CudaSlice<u8> = if f16pv && !v_f16 {
21047 let n = t_kv * n_head_kv * head_dim;
21049 let need = n * 2;
21050 if vguard.as_ref().map(|b| b.len() < need).unwrap_or(true) {
21051 *vguard = Some(self.alloc_uninit::<u8>(need)?);
21052 }
21053 let dst = vguard.as_mut().unwrap();
21054 self.bf16_to_f16_into(vb, n, dst)?;
21055 vguard.as_ref().unwrap()
21056 } else {
21057 vb
21058 };
21059 let f = self.func(if hp {
21060 "fa_prefill_bf16_hd512_sp16h2"
21061 } else {
21062 match (f16pv, nw) {
21063 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
21064 (true, _) => "fa_prefill_bf16_hd512_sp16",
21065 _ => "fa_prefill_bf16_hd512_sp",
21066 }
21067 });
21068 let (nwarp, npart) = if hp {
21069 (4usize, 4usize)
21070 } else if nw > 2 {
21071 (nw, nw)
21072 } else {
21073 (2, 1)
21074 };
21075 let shmem = if hp {
21077 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS) + 4 * (2 * npart * SP_M * BKS + 2 * SP_M))
21078 as u32
21079 } else {
21080 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
21081 + 4 * (npart * SP_M * BKS + SP_M)) as u32
21082 };
21083 use cudarc::driver::sys::CUfunction_attribute_enum as A;
21084 f.set_attribute(
21085 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
21086 shmem as i32,
21087 )?;
21088 let grid_y = if hp {
21089 (n_head / 2) as u32
21090 } else {
21091 n_head as u32
21092 };
21093 let cfg = LaunchConfig {
21094 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
21095 block_dim: (32, nwarp as u32, 1),
21096 shared_mem_bytes: shmem,
21097 };
21098 let (hd, nh, nhkv, ti, tkvi, cz) = (
21099 head_dim as i32,
21100 n_head as i32,
21101 n_head_kv as i32,
21102 t as i32,
21103 t_kv as i32,
21104 causal as i32,
21105 );
21106 let __s_b = self.gpu.stream();
21107 let mut b = __s_b.launch_builder(&f);
21108 b.arg(qb)
21109 .arg(kb)
21110 .arg(vref)
21111 .arg(o)
21112 .arg(&hd)
21113 .arg(&nh)
21114 .arg(&nhkv)
21115 .arg(&ti)
21116 .arg(&tkvi)
21117 .arg(&scale)
21118 .arg(&cz);
21119 unsafe {
21120 b.launch(cfg)?;
21121 }
21122 Ok(())
21123 }
21124
21125 #[allow(clippy::too_many_arguments)]
21128 pub fn fa_prefill_hd512_arm(
21129 &self,
21130 q: &CudaSlice<f32>,
21131 k: &CudaSlice<f32>,
21132 v: &CudaSlice<f32>,
21133 o: &mut CudaSlice<f32>,
21134 head_dim: usize,
21135 n_head: usize,
21136 n_head_kv: usize,
21137 t: usize,
21138 t_kv: usize,
21139 scale: f32,
21140 causal: bool,
21141 f32_stage: bool,
21142 sp: bool,
21143 f16pv: bool,
21144 ) -> Result<(), Box<dyn std::error::Error>> {
21145 debug_assert_eq!(head_dim, 512, "fa_prefill_hd512 is hd512 only");
21146 if sp && !f32_stage {
21147 const SP_M: usize = 16;
21151 const BKS: usize = 32;
21152 let nw = if f16pv { fa512_wide_warps() } else { 2 };
21153 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
21154 let f = self.func(if hp {
21155 "fa_prefill_bf16_hd512_sp16h2"
21156 } else {
21157 match (f16pv, nw) {
21158 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
21159 (true, _) => "fa_prefill_bf16_hd512_sp16",
21160 _ => "fa_prefill_bf16_hd512_sp",
21161 }
21162 });
21163 let (nwarp, npart) = if hp {
21164 (4usize, 4usize)
21165 } else if nw > 2 {
21166 (nw, nw)
21167 } else {
21168 (2, 1)
21169 };
21170 let shmem = if hp {
21171 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
21172 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
21173 } else {
21174 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
21175 + 4 * (npart * SP_M * BKS + SP_M)) as u32
21176 };
21177 use cudarc::driver::sys::CUfunction_attribute_enum as A;
21178 f.set_attribute(
21179 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
21180 shmem as i32,
21181 )?;
21182 let grid_y = if hp {
21183 (n_head / 2) as u32
21184 } else {
21185 n_head as u32
21186 };
21187 let cfg = LaunchConfig {
21188 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
21189 block_dim: (32, nwarp as u32, 1),
21190 shared_mem_bytes: shmem,
21191 };
21192 let (hd, nh, nhkv, ti, tkvi, cz) = (
21193 head_dim as i32,
21194 n_head as i32,
21195 n_head_kv as i32,
21196 t as i32,
21197 t_kv as i32,
21198 causal as i32,
21199 );
21200 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
21201 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
21202 let vb = if f16pv {
21203 self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?
21204 } else {
21205 self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?
21206 };
21207 let __s_b = self.gpu.stream();
21208 let mut b = __s_b.launch_builder(&f);
21209 b.arg(&qb)
21210 .arg(&kb)
21211 .arg(&vb)
21212 .arg(o)
21213 .arg(&hd)
21214 .arg(&nh)
21215 .arg(&nhkv)
21216 .arg(&ti)
21217 .arg(&tkvi)
21218 .arg(&scale)
21219 .arg(&cz);
21220 unsafe {
21221 b.launch(cfg)?;
21222 }
21223 return Ok(());
21224 }
21225 const BLOCK_Q: usize = 32;
21226 const BK: usize = 32;
21227 const HALF: usize = 256;
21228 let f = self.func(if f32_stage {
21229 "fa_prefill_f32_hd512"
21230 } else {
21231 "fa_prefill_bf16_hd512"
21232 });
21233 let shmem = (2 * (BLOCK_Q * head_dim + BK * head_dim + BK * HALF + BLOCK_Q * BK)
21235 + 4 * BLOCK_Q) as u32;
21236 use cudarc::driver::sys::CUfunction_attribute_enum as A;
21237 f.set_attribute(
21238 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
21239 shmem as i32,
21240 )?;
21241 let cfg = LaunchConfig {
21242 grid_dim: (
21243 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
21244 n_head as u32,
21245 2,
21246 ),
21247 block_dim: (32, 2, 1),
21248 shared_mem_bytes: shmem,
21249 };
21250 let (hd, nh, nhkv, ti, tkvi, cz) = (
21251 head_dim as i32,
21252 n_head as i32,
21253 n_head_kv as i32,
21254 t as i32,
21255 t_kv as i32,
21256 causal as i32,
21257 );
21258 if f32_stage {
21259 let __s_b = self.gpu.stream();
21260 let mut b = __s_b.launch_builder(&f);
21261 b.arg(q)
21262 .arg(k)
21263 .arg(v)
21264 .arg(o)
21265 .arg(&hd)
21266 .arg(&nh)
21267 .arg(&nhkv)
21268 .arg(&ti)
21269 .arg(&tkvi)
21270 .arg(&scale)
21271 .arg(&cz);
21272 unsafe {
21273 b.launch(cfg)?;
21274 }
21275 } else {
21276 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
21277 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
21278 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
21279 let __s_b = self.gpu.stream();
21280 let mut b = __s_b.launch_builder(&f);
21281 b.arg(&qb)
21282 .arg(&kb)
21283 .arg(&vb)
21284 .arg(o)
21285 .arg(&hd)
21286 .arg(&nh)
21287 .arg(&nhkv)
21288 .arg(&ti)
21289 .arg(&tkvi)
21290 .arg(&scale)
21291 .arg(&cz);
21292 unsafe {
21293 b.launch(cfg)?;
21294 }
21295 }
21296 Ok(())
21297 }
21298
21299 #[allow(clippy::too_many_arguments)]
21303 pub fn rope_neox2_bf16e(
21304 &self,
21305 q: &mut CudaSlice<f32>,
21306 k: &mut CudaSlice<f32>,
21307 qb: &mut CudaSlice<u8>,
21308 kb: &mut CudaSlice<u8>,
21309 pos: &CudaSlice<i32>,
21310 head_dim: usize,
21311 n_dims: usize,
21312 nh_q: usize,
21313 nh_k: usize,
21314 n_tokens: usize,
21315 base: f32,
21316 freq_scale: f32,
21317 ff: Option<&CudaSlice<f32>>,
21318 ) -> Result<(), Box<dyn std::error::Error>> {
21319 let f = self.func("rope_neox2_bf16e_f32");
21320 let rows = ((nh_q + nh_k) * n_tokens) as u32;
21321 let cfg = LaunchConfig {
21322 grid_dim: (rows, 1, 1),
21323 block_dim: ((head_dim / 2) as u32, 1, 1),
21324 shared_mem_bytes: 0,
21325 };
21326 let theta_scale = base.powf(-2.0 / n_dims as f32);
21327 let (hd, nd, nhq, nhk, nt) = (
21328 head_dim as i32,
21329 n_dims as i32,
21330 nh_q as i32,
21331 nh_k as i32,
21332 n_tokens as i32,
21333 );
21334 let __s_b = self.gpu.stream();
21335 let mut b = __s_b.launch_builder(&f);
21336 match ff {
21337 Some(t) => {
21338 b.arg(&mut *q)
21339 .arg(&mut *k)
21340 .arg(&mut *qb)
21341 .arg(&mut *kb)
21342 .arg(pos)
21343 .arg(&hd)
21344 .arg(&nd)
21345 .arg(&nhq)
21346 .arg(&nhk)
21347 .arg(&nt)
21348 .arg(&theta_scale)
21349 .arg(&freq_scale)
21350 .arg(t);
21351 unsafe {
21352 b.launch(cfg)?;
21353 }
21354 }
21355 None => {
21356 let null: u64 = 0;
21357 b.arg(&mut *q)
21358 .arg(&mut *k)
21359 .arg(&mut *qb)
21360 .arg(&mut *kb)
21361 .arg(pos)
21362 .arg(&hd)
21363 .arg(&nd)
21364 .arg(&nhq)
21365 .arg(&nhk)
21366 .arg(&nt)
21367 .arg(&theta_scale)
21368 .arg(&freq_scale)
21369 .arg(&null);
21370 unsafe {
21371 b.launch(cfg)?;
21372 }
21373 }
21374 }
21375 Ok(())
21376 }
21377
21378 pub fn f32_to_bf16(
21381 &self,
21382 x: &CudaSlice<f32>,
21383 n: usize,
21384 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
21385 assert!(n % 4 == 0, "f32_to_bf16 requires n % 4 == 0, got {n}");
21386 let mut y = self.alloc_uninit::<u8>(n * 2)?;
21387 let f = self.func("f32_to_bf16_flat");
21388 let n_i = n as i64;
21389 let cfg = LaunchConfig {
21390 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
21391 block_dim: (256, 1, 1),
21392 shared_mem_bytes: 0,
21393 };
21394 let __s_b = self.gpu.stream();
21395 let mut b = __s_b.launch_builder(&f);
21396 b.arg(x).arg(&mut y).arg(&n_i);
21397 unsafe {
21398 b.launch(cfg)?;
21399 }
21400 Ok(y)
21401 }
21402
21403 pub fn f32_to_f16(
21404 &self,
21405 x: &CudaSlice<f32>,
21406 n: usize,
21407 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
21408 assert!(n % 4 == 0, "f32_to_f16 requires n % 4 == 0, got {n}");
21409 let mut y = self.alloc_uninit::<u8>(n * 2)?;
21410 let f = self.func("f32_to_f16_flat");
21411 let n_i = n as i64;
21412 let cfg = LaunchConfig {
21413 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
21414 block_dim: (256, 1, 1),
21415 shared_mem_bytes: 0,
21416 };
21417 let __s_b = self.gpu.stream();
21418 let mut b = __s_b.launch_builder(&f);
21419 b.arg(x).arg(&mut y).arg(&n_i);
21420 unsafe {
21421 b.launch(cfg)?;
21422 }
21423 Ok(y)
21424 }
21425
21426 pub fn bf16_to_f16(
21428 &self,
21429 xb: &CudaSlice<u8>,
21430 n: usize,
21431 ) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
21432 let mut y = self.alloc_uninit::<u8>(n * 2)?;
21433 self.bf16_to_f16_into(xb, n, &mut y)?;
21434 Ok(y)
21435 }
21436
21437 pub fn bf16_to_f16_into(
21439 &self,
21440 xb: &CudaSlice<u8>,
21441 n: usize,
21442 y: &mut CudaSlice<u8>,
21443 ) -> Result<(), Box<dyn std::error::Error>> {
21444 assert!(n % 2 == 0, "bf16_to_f16 requires n % 2 == 0, got {n}");
21445 assert!(y.len() >= n * 2);
21446 let f = self.func("bf16_to_f16_flat");
21447 let n2 = (n / 2) as i64;
21448 let cfg = LaunchConfig {
21449 grid_dim: (((n / 2) as u32).div_ceil(256), 1, 1),
21450 block_dim: (256, 1, 1),
21451 shared_mem_bytes: 0,
21452 };
21453 let __s_b = self.gpu.stream();
21454 let mut b = __s_b.launch_builder(&f);
21455 b.arg(xb).arg(y).arg(&n2);
21456 unsafe {
21457 b.launch(cfg)?;
21458 }
21459 Ok(())
21460 }
21461
21462 #[allow(clippy::too_many_arguments)]
21467 pub fn fa_prefill_vl8(
21468 &self,
21469 seqs: &[FaSeqVl],
21470 head_dim: usize,
21471 n_head: usize,
21472 n_head_kv: usize,
21473 scale: f32,
21474 ) -> Result<(), Box<dyn std::error::Error>> {
21475 const BK: usize = 32;
21476 let b = seqs.len();
21477 assert!(b >= 1 && b <= 8);
21478 let mut packed = [FaSeqVl::default(); 8];
21479 packed[..b].copy_from_slice(seqs);
21480 let v = FaVl8(packed);
21481 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
21482 let ept = (n_head_kv * head_dim) as i32;
21483 {
21484 let f = self.func("fa_mirror_vl");
21485 let max_n = (max_t as i64) * ept as i64;
21486 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
21487 for which in 0..2i32 {
21488 let cfg = LaunchConfig {
21489 grid_dim: (blocks, 1, b as u32),
21490 block_dim: (256, 1, 1),
21491 shared_mem_bytes: 0,
21492 };
21493 let __s_lb = self.gpu.stream();
21494 let mut lb = __s_lb.launch_builder(&f);
21495 lb.arg(&v).arg(&ept).arg(&which);
21496 unsafe {
21497 lb.launch(cfg)?;
21498 }
21499 }
21500 }
21501 let hd_sfx = fa_hd_suffix(head_dim)?;
21502 let f = self.func(&format!("fa_prefill_bf16kv_vl{hd_sfx}"));
21503 let block_q = 64usize;
21504 let kv_stages = 2usize;
21505 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
21506 + 4 * (block_q * BK + 2 * block_q)) as u32;
21507 use cudarc::driver::sys::CUfunction_attribute_enum as A;
21508 f.set_attribute(
21509 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
21510 shmem as i32,
21511 )?;
21512 let cfg = LaunchConfig {
21513 grid_dim: (max_t.div_ceil(block_q as u32), n_head as u32, b as u32),
21514 block_dim: (32, 4, 1),
21515 shared_mem_bytes: shmem,
21516 };
21517 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
21518 let __s_lb = self.gpu.stream();
21519 let mut lb = __s_lb.launch_builder(&f);
21520 lb.arg(&v).arg(&hd).arg(&nh).arg(&nhkv).arg(&scale);
21521 unsafe {
21522 lb.launch(cfg)?;
21523 }
21524 Ok(())
21525 }
21526
21527 #[allow(clippy::too_many_arguments)]
21531 pub fn attn_pre_vl8(
21532 &self,
21533 seqs: &[AttnPreVl],
21534 wq: &CudaSlice<f32>,
21535 wk: &CudaSlice<f32>,
21536 head_dim: usize,
21537 rope_dims: usize,
21538 n_head: usize,
21539 n_head_kv: usize,
21540 eps: f32,
21541 freq_base: f32,
21542 freq_scale: f32,
21543 kv_dim_k: usize,
21544 kv_dim_v: usize,
21545 k_tok_bytes: usize,
21546 v_tok_bytes: usize,
21547 ) -> Result<(), Box<dyn std::error::Error>> {
21548 let b = seqs.len();
21549 assert!(b >= 1 && b <= 8);
21550 let mut packed = [AttnPreVl::default(); 8];
21551 packed[..b].copy_from_slice(seqs);
21552 let v = AttnPreVl8(packed);
21553 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
21554 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
21555 {
21556 let f = self.func("q_gate_split_vl");
21557 let n = max_t * (n_head * head_dim) as u32;
21558 let cfg = LaunchConfig {
21559 grid_dim: (n.div_ceil(256), 1, b as u32),
21560 block_dim: (256, 1, 1),
21561 shared_mem_bytes: 0,
21562 };
21563 let __s_lb = self.gpu.stream();
21564 let mut lb = __s_lb.launch_builder(&f);
21565 lb.arg(&v).arg(&hd).arg(&nh);
21566 unsafe {
21567 lb.launch(cfg)?;
21568 }
21569 }
21570 {
21571 let f = self.func("attn_rms_vl");
21572 let cfg = LaunchConfig {
21573 grid_dim: (max_t * n_head as u32, 2, b as u32),
21574 block_dim: (rms_block(), 1, 1),
21575 shared_mem_bytes: 0,
21576 };
21577 let __s_lb = self.gpu.stream();
21578 let mut lb = __s_lb.launch_builder(&f);
21579 lb.arg(&v)
21580 .arg(wq)
21581 .arg(wk)
21582 .arg(&hd)
21583 .arg(&nh)
21584 .arg(&nhkv)
21585 .arg(&eps);
21586 unsafe {
21587 lb.launch(cfg)?;
21588 }
21589 }
21590 {
21591 let f = self.func("attn_rope_vl");
21592 let theta_scale = freq_base.powf(-2.0 / rope_dims as f32);
21593 let nd = rope_dims as i32;
21594 let cfg = LaunchConfig {
21595 grid_dim: (max_t * n_head as u32, 2, b as u32),
21596 block_dim: ((head_dim / 2) as u32, 1, 1),
21597 shared_mem_bytes: 0,
21598 };
21599 let __s_lb = self.gpu.stream();
21600 let mut lb = __s_lb.launch_builder(&f);
21601 lb.arg(&v)
21602 .arg(&hd)
21603 .arg(&nd)
21604 .arg(&nh)
21605 .arg(&nhkv)
21606 .arg(&theta_scale)
21607 .arg(&freq_scale);
21608 unsafe {
21609 lb.launch(cfg)?;
21610 }
21611 }
21612 {
21613 let f = self.func("append_kv_vl");
21614 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
21615 let cfg = LaunchConfig {
21616 grid_dim: (nblk, max_t, b as u32),
21617 block_dim: (32, 1, 1),
21618 shared_mem_bytes: 0,
21619 };
21620 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
21621 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
21622 let __s_lb = self.gpu.stream();
21623 let mut lb = __s_lb.launch_builder(&f);
21624 lb.arg(&v).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
21625 unsafe {
21626 lb.launch(cfg)?;
21627 }
21628 }
21629 Ok(())
21630 }
21631
21632 pub fn fa_prefill_view(
21637 &self,
21638 q: &CudaSlice<f32>,
21639 k: &cudarc::driver::CudaView<u8>,
21640 v: &cudarc::driver::CudaView<u8>,
21641 o: &mut CudaSlice<f32>,
21642 head_dim: usize,
21643 n_head: usize,
21644 n_head_kv: usize,
21645 t: usize,
21646 t_kv: usize,
21647 scale: f32,
21648 causal: bool,
21649 k_tok_bytes: usize,
21650 v_tok_bytes: usize,
21651 g: bool,
21652 ) -> Result<(), Box<dyn std::error::Error>> {
21653 if portable_mma_gated() {
21654 return self.sdpa_naive_quantized_view(
21655 q,
21656 k,
21657 v,
21658 o,
21659 head_dim,
21660 n_head,
21661 n_head_kv,
21662 t,
21663 t_kv,
21664 scale,
21665 causal,
21666 k_tok_bytes,
21667 v_tok_bytes,
21668 );
21669 }
21670 const BLOCK_Q: usize = 64;
21671 const BK: usize = 32;
21672 let name = format!("fa_prefill_q{}", fa_hd_suffix(head_dim)?);
21675 let f = if g {
21676 self.func_g(&name)
21677 } else {
21678 self.func(&name)
21679 };
21680 let shmem =
21681 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
21682 use cudarc::driver::sys::CUfunction_attribute_enum as A;
21683 f.set_attribute(
21684 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
21685 shmem as i32,
21686 )?;
21687 let cfg = LaunchConfig {
21688 grid_dim: (
21689 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
21690 n_head as u32,
21691 1,
21692 ),
21693 block_dim: (32, 4, 1),
21694 shared_mem_bytes: shmem,
21695 };
21696 let (hd, nh, nhkv, ti, tkvi, cz) = (
21697 head_dim as i32,
21698 n_head as i32,
21699 n_head_kv as i32,
21700 t as i32,
21701 t_kv as i32,
21702 causal as i32,
21703 );
21704 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
21705 let __s_b = self.gpu.stream();
21706 let mut b = __s_b.launch_builder(&f);
21707 b.arg(q)
21708 .arg(k)
21709 .arg(v)
21710 .arg(o)
21711 .arg(&hd)
21712 .arg(&nh)
21713 .arg(&nhkv)
21714 .arg(&ti)
21715 .arg(&tkvi)
21716 .arg(&scale)
21717 .arg(&cz)
21718 .arg(&ktb)
21719 .arg(&vtb);
21720 unsafe {
21721 b.launch(cfg)?;
21722 }
21723 Ok(())
21724 }
21725
21726 #[allow(clippy::too_many_arguments)]
21736 pub fn fa_prefill_view_ws(
21737 &self,
21738 q: &CudaSlice<f32>,
21739 k: &cudarc::driver::CudaView<u8>,
21740 v: &cudarc::driver::CudaView<u8>,
21741 o: &mut CudaSlice<f32>,
21742 head_dim: usize,
21743 n_head: usize,
21744 n_head_kv: usize,
21745 t: usize,
21746 t_kv: usize,
21747 scale: f32,
21748 causal: bool,
21749 k_tok_bytes: usize,
21750 v_tok_bytes: usize,
21751 g: bool,
21752 ) -> Result<(), Box<dyn std::error::Error>> {
21753 if portable_mma_gated() {
21754 return self.sdpa_naive_quantized_view(
21755 q,
21756 k,
21757 v,
21758 o,
21759 head_dim,
21760 n_head,
21761 n_head_kv,
21762 t,
21763 t_kv,
21764 scale,
21765 causal,
21766 k_tok_bytes,
21767 v_tok_bytes,
21768 );
21769 }
21770 const BLOCK_Q: usize = 64;
21771 const BK: usize = 32;
21772 let kv_dim_k = n_head_kv * head_dim;
21773 let kv_dim_v = n_head_kv * head_dim;
21774 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
21776 let mut guard = self.prime_deqw_ws.lock().unwrap();
21778 let need_grow = match guard.as_ref() {
21779 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
21780 None => true,
21781 };
21782 if need_grow {
21783 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
21784 let (ck, cv) = guard
21785 .as_ref()
21786 .map(|(a, b)| (a.len(), b.len()))
21787 .unwrap_or((0, 0));
21788 *guard = Some((
21789 self.alloc_u8(grow(ck, k_ws_bytes))?,
21790 self.alloc_u8(grow(cv, v_ws_bytes))?,
21791 ));
21792 }
21793 let (kw, vw) = guard.as_mut().unwrap();
21794 {
21796 let f = if g {
21798 self.func_g("fa_dequant_kv_ws_bf16")
21799 } else {
21800 self.func("fa_dequant_kv_ws_bf16")
21801 };
21802 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
21803 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
21804 let cfg = LaunchConfig {
21805 grid_dim: (nblk.max(1), 1, 1),
21806 block_dim: (256, 1, 1),
21807 shared_mem_bytes: 0,
21808 };
21809 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
21810 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
21811 let __s_b = self.gpu.stream();
21812 let mut b = __s_b.launch_builder(&f);
21813 b.arg(k)
21814 .arg(v)
21815 .arg(&mut *kw)
21816 .arg(&mut *vw)
21817 .arg(&kdk)
21818 .arg(&kdv)
21819 .arg(&tkvi)
21820 .arg(&ktb)
21821 .arg(&vtb);
21822 unsafe {
21823 b.launch(cfg)?;
21824 }
21825 }
21826 let db = std::env::var("MEMRA_PRIME_DEQW_DB")
21834 .map(|v| v != "0")
21835 .unwrap_or(true);
21836 {
21837 let hd_sfx = fa_hd_suffix(head_dim)?;
21838 let f = self.func(&format!(
21839 "fa_prefill_qw{}{hd_sfx}",
21840 if db { "_db" } else { "" }
21841 ));
21842 let shmem = if db {
21843 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
21845 } else {
21846 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
21847 };
21848 use cudarc::driver::sys::CUfunction_attribute_enum as A;
21849 f.set_attribute(
21850 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
21851 shmem as i32,
21852 )?;
21853 let cfg = LaunchConfig {
21854 grid_dim: (
21855 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
21856 n_head as u32,
21857 1,
21858 ),
21859 block_dim: (32, 4, 1),
21860 shared_mem_bytes: shmem,
21861 };
21862 let (hd, nh, nhkv, ti, tkvi, cz) = (
21863 head_dim as i32,
21864 n_head as i32,
21865 n_head_kv as i32,
21866 t as i32,
21867 t_kv as i32,
21868 causal as i32,
21869 );
21870 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
21871 let __s_b = self.gpu.stream();
21872 let mut b = __s_b.launch_builder(&f);
21873 b.arg(q)
21874 .arg(&*kw)
21875 .arg(&*vw)
21876 .arg(o)
21877 .arg(&hd)
21878 .arg(&nh)
21879 .arg(&nhkv)
21880 .arg(&ti)
21881 .arg(&tkvi)
21882 .arg(&scale)
21883 .arg(&cz)
21884 .arg(&kdk)
21885 .arg(&kdv);
21886 unsafe {
21887 b.launch(cfg)?;
21888 }
21889 }
21890 Ok(())
21891 }
21892
21893 #[allow(clippy::too_many_arguments)]
21909 pub fn fa_prefill_view_ws_w_hd128(
21910 &self,
21911 q: &CudaSlice<f32>,
21912 k: &cudarc::driver::CudaView<u8>,
21913 v: &cudarc::driver::CudaView<u8>,
21914 o: &mut CudaSlice<f32>,
21915 head_dim: usize,
21916 n_head: usize,
21917 n_head_kv: usize,
21918 t: usize,
21919 t_kv: usize,
21920 scale: f32,
21921 causal: bool,
21922 window: usize,
21923 k_tok_bytes: usize,
21924 v_tok_bytes: usize,
21925 ) -> Result<(), Box<dyn std::error::Error>> {
21926 assert_eq!(
21927 head_dim, 128,
21928 "fa_prefill_view_ws_w_hd128: only the hd128 twin is stamped"
21929 );
21930 if portable_mma_gated() {
21931 return self.sdpa_naive_w_quantized_view(
21932 q,
21933 k,
21934 v,
21935 o,
21936 head_dim,
21937 n_head,
21938 n_head_kv,
21939 t,
21940 t_kv,
21941 scale,
21942 causal,
21943 window,
21944 k_tok_bytes,
21945 v_tok_bytes,
21946 );
21947 }
21948 const BLOCK_Q: usize = 64;
21949 const BK: usize = 32;
21950 let kv_dim_k = n_head_kv * head_dim;
21951 let kv_dim_v = n_head_kv * head_dim;
21952 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
21954 let mut guard = self.prime_deqw_ws.lock().unwrap();
21955 let need_grow = match guard.as_ref() {
21956 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
21957 None => true,
21958 };
21959 if need_grow {
21960 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
21961 let (ck, cv) = guard
21962 .as_ref()
21963 .map(|(a, b)| (a.len(), b.len()))
21964 .unwrap_or((0, 0));
21965 *guard = Some((
21966 self.alloc_u8(grow(ck, k_ws_bytes))?,
21967 self.alloc_u8(grow(cv, v_ws_bytes))?,
21968 ));
21969 }
21970 let (kw, vw) = guard.as_mut().unwrap();
21971 {
21974 let f = self.func("fa_dequant_kv_ws_bf16");
21975 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
21976 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
21977 let cfg = LaunchConfig {
21978 grid_dim: (nblk.max(1), 1, 1),
21979 block_dim: (256, 1, 1),
21980 shared_mem_bytes: 0,
21981 };
21982 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
21983 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
21984 let __s_b = self.gpu.stream();
21985 let mut b = __s_b.launch_builder(&f);
21986 b.arg(k)
21987 .arg(v)
21988 .arg(&mut *kw)
21989 .arg(&mut *vw)
21990 .arg(&kdk)
21991 .arg(&kdv)
21992 .arg(&tkvi)
21993 .arg(&ktb)
21994 .arg(&vtb);
21995 unsafe {
21996 b.launch(cfg)?;
21997 }
21998 }
21999 let db = std::env::var("MEMRA_PRIME_DEQW_DB")
22001 .map(|v| v != "0")
22002 .unwrap_or(true);
22003 {
22004 let f = self.func(if db {
22005 "fa_prefill_qw_db_w_hd128"
22006 } else {
22007 "fa_prefill_qw_w_hd128"
22008 });
22009 let shmem = if db {
22010 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
22011 } else {
22012 (2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
22013 };
22014 use cudarc::driver::sys::CUfunction_attribute_enum as A;
22015 f.set_attribute(
22016 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
22017 shmem as i32,
22018 )?;
22019 let cfg = LaunchConfig {
22020 grid_dim: (
22021 (t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
22022 n_head as u32,
22023 1,
22024 ),
22025 block_dim: (32, 4, 1),
22026 shared_mem_bytes: shmem,
22027 };
22028 let (hd, nh, nhkv, ti, tkvi, cz) = (
22029 head_dim as i32,
22030 n_head as i32,
22031 n_head_kv as i32,
22032 t as i32,
22033 t_kv as i32,
22034 causal as i32,
22035 );
22036 let (kdk, kdv, wnd) = (kv_dim_k as i32, kv_dim_v as i32, window as i32);
22037 let __s_b = self.gpu.stream();
22038 let mut b = __s_b.launch_builder(&f);
22039 b.arg(q)
22040 .arg(&*kw)
22041 .arg(&*vw)
22042 .arg(o)
22043 .arg(&hd)
22044 .arg(&nh)
22045 .arg(&nhkv)
22046 .arg(&ti)
22047 .arg(&tkvi)
22048 .arg(&scale)
22049 .arg(&cz)
22050 .arg(&kdk)
22051 .arg(&kdv)
22052 .arg(&wnd);
22053 unsafe {
22054 b.launch(cfg)?;
22055 }
22056 }
22057 Ok(())
22058 }
22059
22060 pub fn fa_decode(
22064 &self,
22065 q: &CudaSlice<f32>,
22066 k: &cudarc::driver::CudaView<u8>,
22067 v: &cudarc::driver::CudaView<u8>,
22068 o: &mut CudaSlice<f32>,
22069 head_dim: usize,
22070 n_head: usize,
22071 n_head_kv: usize,
22072 t_kv: usize,
22073 scale: f32,
22074 k_tok_bytes: usize,
22075 v_tok_bytes: usize,
22076 ) -> Result<(), Box<dyn std::error::Error>> {
22077 self.fa_decode_kvmod(
22078 q,
22079 k,
22080 v,
22081 o,
22082 head_dim,
22083 n_head,
22084 n_head_kv,
22085 t_kv,
22086 scale,
22087 k_tok_bytes,
22088 v_tok_bytes,
22089 false,
22090 )
22091 }
22092
22093 #[allow(clippy::too_many_arguments)]
22097 #[allow(clippy::too_many_arguments)]
22101 #[allow(clippy::too_many_arguments)]
22102 fn fa_decode_scalar_unified(
22103 &self,
22104 q: &cudarc::driver::CudaView<f32>,
22105 k: &cudarc::driver::CudaView<u8>,
22106 v: &cudarc::driver::CudaView<u8>,
22107 o: &mut cudarc::driver::CudaViewMut<f32>,
22108 head_dim: usize,
22109 n_head: usize,
22110 n_head_kv: usize,
22111 t_kv_host: usize,
22112 t_kv_dev: Option<&CudaSlice<i32>>,
22113 scale: f32,
22114 n_splits: usize,
22115 split_keys: usize,
22116 k_tok_bytes: usize,
22117 v_tok_bytes: usize,
22118 g: bool,
22119 part_o: &mut CudaSlice<f32>,
22120 part_m: &mut CudaSlice<f32>,
22121 part_l: &mut CudaSlice<f32>,
22122 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
22123 ) -> Result<(), Box<dyn std::error::Error>> {
22124 let f = if g {
22125 self.func_g("fa_decode_f32")
22126 } else {
22127 self.fa_func("fa_decode_f32", head_dim)
22128 };
22129 let cfg = LaunchConfig {
22130 grid_dim: (n_head as u32, n_splits as u32, 1),
22131 block_dim: (head_dim as u32, 1, 1),
22132 shared_mem_bytes: (4 * (head_dim + 32)) as u32,
22133 };
22134 let (hd, nh, nhkv, nsp) = (
22135 head_dim as i32,
22136 n_head as i32,
22137 n_head_kv as i32,
22138 n_splits as i32,
22139 );
22140 let (ktb, vtb, tkvi, ski) = (
22141 k_tok_bytes as i64,
22142 v_tok_bytes as i64,
22143 t_kv_host as i32,
22144 split_keys as i32,
22145 );
22146 let __s_b = self.gpu.stream();
22147 let mut b = __s_b.launch_builder(&f);
22148 match t_kv_dev {
22149 Some(d) => {
22150 b.arg(q)
22151 .arg(k)
22152 .arg(v)
22153 .arg(&mut *part_o)
22154 .arg(&mut *part_m)
22155 .arg(&mut *part_l)
22156 .arg(&hd)
22157 .arg(&nh)
22158 .arg(&nhkv)
22159 .arg(&tkvi)
22160 .arg(d)
22161 .arg(&scale)
22162 .arg(&nsp)
22163 .arg(&ski)
22164 .arg(&ktb)
22165 .arg(&vtb);
22166 unsafe {
22167 b.launch(cfg)?;
22168 }
22169 }
22170 None => {
22171 let null: u64 = 0;
22172 b.arg(q)
22173 .arg(k)
22174 .arg(v)
22175 .arg(&mut *part_o)
22176 .arg(&mut *part_m)
22177 .arg(&mut *part_l)
22178 .arg(&hd)
22179 .arg(&nh)
22180 .arg(&nhkv)
22181 .arg(&tkvi)
22182 .arg(&null)
22183 .arg(&scale)
22184 .arg(&nsp)
22185 .arg(&ski)
22186 .arg(&ktb)
22187 .arg(&vtb);
22188 unsafe {
22189 b.launch(cfg)?;
22190 }
22191 }
22192 }
22193 let cfg2 = LaunchConfig {
22194 grid_dim: (n_head as u32, 1, 1),
22195 block_dim: (head_dim as u32, 1, 1),
22196 shared_mem_bytes: 0,
22197 };
22198 if let Some((oq, od)) = q8_out {
22199 let fc = if g {
22201 self.func_g("fa_decode_combine_q8_1")
22202 } else {
22203 self.fa_func("fa_decode_combine_q8_1", head_dim)
22204 };
22205 let __s_b2 = self.gpu.stream();
22206 let mut b2 = __s_b2.launch_builder(&fc);
22207 b2.arg(&*part_o)
22208 .arg(&*part_m)
22209 .arg(&*part_l)
22210 .arg(oq)
22211 .arg(od)
22212 .arg(&hd)
22213 .arg(&nh)
22214 .arg(&nsp);
22215 unsafe {
22216 b2.launch(cfg2)?;
22217 }
22218 return Ok(());
22219 }
22220 let fc = if g {
22221 self.func_g("fa_decode_combine_f32")
22222 } else {
22223 self.fa_func("fa_decode_combine_f32", head_dim)
22224 };
22225 let __s_b2 = self.gpu.stream();
22226 let mut b2 = __s_b2.launch_builder(&fc);
22227 b2.arg(&*part_o)
22228 .arg(&*part_m)
22229 .arg(&*part_l)
22230 .arg(o)
22231 .arg(&hd)
22232 .arg(&nh)
22233 .arg(&nsp);
22234 unsafe {
22235 b2.launch(cfg2)?;
22236 }
22237 Ok(())
22238 }
22239
22240 pub fn fa_decode_kvmod(
22241 &self,
22242 q: &CudaSlice<f32>,
22243 k: &cudarc::driver::CudaView<u8>,
22244 v: &cudarc::driver::CudaView<u8>,
22245 o: &mut CudaSlice<f32>,
22246 head_dim: usize,
22247 n_head: usize,
22248 n_head_kv: usize,
22249 t_kv: usize,
22250 scale: f32,
22251 k_tok_bytes: usize,
22252 v_tok_bytes: usize,
22253 g: bool,
22254 ) -> Result<(), Box<dyn std::error::Error>> {
22255 let q_view = q.as_view();
22256 let mut o_view = o.as_view_mut();
22257 self.fa_decode_kvmod_view(
22258 &q_view,
22259 k,
22260 v,
22261 &mut o_view,
22262 head_dim,
22263 n_head,
22264 n_head_kv,
22265 t_kv,
22266 scale,
22267 k_tok_bytes,
22268 v_tok_bytes,
22269 g,
22270 )
22271 }
22272
22273 #[allow(clippy::too_many_arguments)]
22278 pub fn fa_decode_kvmod_view(
22279 &self,
22280 q: &cudarc::driver::CudaView<f32>,
22281 k: &cudarc::driver::CudaView<u8>,
22282 v: &cudarc::driver::CudaView<u8>,
22283 o: &mut cudarc::driver::CudaViewMut<f32>,
22284 head_dim: usize,
22285 n_head: usize,
22286 n_head_kv: usize,
22287 t_kv: usize,
22288 scale: f32,
22289 k_tok_bytes: usize,
22290 v_tok_bytes: usize,
22291 g: bool,
22292 ) -> Result<(), Box<dyn std::error::Error>> {
22293 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
22314 if g && head_dim == 256 && !fa_v4_at(t_kv) {
22318 fa_vec = false;
22319 }
22320 let sp = fa_split_keys(t_kv, n_head_kv);
22321 let n_splits = if fa_vec {
22322 ((t_kv + sp - 1) / sp).max(1)
22323 } else {
22324 ((t_kv + 255) / 256).max(1)
22325 };
22326 let o_len = n_head * n_splits * head_dim;
22327 let ml_len = n_head * n_splits;
22328 let mut part_guard = self.fa_part_pool.lock().unwrap();
22329 if part_guard
22330 .as_ref()
22331 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
22332 .unwrap_or(true)
22333 {
22334 let old = part_guard.take();
22345 let (co, cm) = old
22346 .as_ref()
22347 .map(|pp| (pp.0.len(), pp.1.len()))
22348 .unwrap_or((0, 0));
22349 if let Some(old) = old {
22350 self.fa_part_retired.lock().unwrap().push(old);
22351 }
22352 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
22353 eprintln!(
22354 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
22355 co, o_len, cm, ml_len
22356 );
22357 }
22358 *part_guard =
22359 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
22360 }
22361 let pg = part_guard.as_mut().unwrap();
22362 self.gpu
22363 .stream()
22364 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
22365 self.gpu
22366 .stream()
22367 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
22368 self.gpu
22369 .stream()
22370 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
22371 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
22372 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
22373 let (hd, nh, nhkv, tkvi, nsp) = (
22374 head_dim as i32,
22375 n_head as i32,
22376 n_head_kv as i32,
22377 t_kv as i32,
22378 n_splits as i32,
22379 );
22380 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
22381 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
22385 let fa512_min = fa512_min_tkv();
22390 let deep = fa_vec
22393 && head_dim == 256
22394 && fa_v4_at(t_kv)
22395 && !g
22396 && fa_deep_at(t_kv)
22397 && !matches!(fa_v4_mode(), "noB3" | "stage");
22398 let (f, cfg) = if fa_vec && head_dim == 512 && t_kv >= fa512_min {
22399 let gqa = (n_head / n_head_kv).max(1) as u32;
22402 let fv = self.fa_func("fa_decode_vec_q_dpl16", head_dim);
22403 (
22404 fv,
22405 LaunchConfig {
22406 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
22407 block_dim: (32, gqa, 1),
22408 shared_mem_bytes: 0,
22409 },
22410 )
22411 } else if fa_vec && head_dim <= 256 {
22412 let gqa = (n_head / n_head_kv).max(1) as u32;
22413 static SMEM_TKV: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
22424 let smem_tkv = *SMEM_TKV.get_or_init(|| {
22425 std::env::var("MEMRA_FA_SMEM_TKV")
22426 .ok()
22427 .and_then(|v| v.parse().ok())
22428 .unwrap_or_else(|| {
22429 FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
22430 })
22431 });
22432 if fa_v4_at(t_kv) && head_dim == 256 {
22433 let v4name = match fa_v4_mode() {
22437 "noB3" => "fa_decode_vec_q_v4_noB3", "stage" => "fa_decode_vec_q_v4_stage", _ if deep => "fa_decode_vec_q_v4_deep",
22440 _ => "fa_decode_vec_q_v4",
22441 };
22442 let fv = if g {
22443 self.func_g(v4name)
22444 } else {
22445 self.func(v4name)
22446 };
22447 let shmem = (if deep { 12160 } else { 11520 }
22450 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
22451 use cudarc::driver::sys::CUfunction_attribute_enum as A;
22452 fv.set_attribute(
22453 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
22454 shmem as i32,
22455 )?;
22456 (
22457 fv,
22458 LaunchConfig {
22459 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
22460 block_dim: (32, gqa, 1),
22461 shared_mem_bytes: shmem,
22462 },
22463 )
22464 } else if fa_v3_active(head_dim) {
22465 let fv = if g {
22468 self.func_g("fa_decode_vec_q_v3")
22469 } else {
22470 self.func("fa_decode_vec_q_v3")
22471 };
22472 let shmem = (32 * head_dim * 2) as u32; (
22474 fv,
22475 LaunchConfig {
22476 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
22477 block_dim: (32, gqa, 1),
22478 shared_mem_bytes: shmem,
22479 },
22480 )
22481 } else if fa_v2_on() {
22482 let fv = if g {
22486 self.func_g("fa_decode_vec_q_v2")
22487 } else {
22488 self.func("fa_decode_vec_q_v2")
22489 };
22490 let shmem = (2 * 32 * head_dim * 2) as u32; (
22492 fv,
22493 LaunchConfig {
22494 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
22495 block_dim: (32, gqa, 1),
22496 shared_mem_bytes: shmem,
22497 },
22498 )
22499 } else if smem_tkv > 0 && t_kv >= smem_tkv && !g && !(head_dim == 512 && Self::gkv_on())
22500 {
22501 let fv = if g {
22505 self.func_g("fa_decode_vec_q_smem")
22506 } else {
22507 self.func("fa_decode_vec_q_smem")
22508 };
22509 let shmem = (2 * 32 * head_dim * 2) as u32; use cudarc::driver::sys::CUfunction_attribute_enum as A;
22511 fv.set_attribute(
22512 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
22513 shmem as i32,
22514 )?;
22515 (
22516 fv,
22517 LaunchConfig {
22518 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
22519 block_dim: (32, gqa, 1),
22520 shared_mem_bytes: shmem,
22521 },
22522 )
22523 } else {
22524 let fv = if g {
22527 self.func_g("fa_decode_vec_q")
22528 } else {
22529 self.func("fa_decode_vec_q")
22530 };
22531 (
22532 fv,
22533 LaunchConfig {
22534 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
22535 block_dim: (32, gqa, 1),
22536 shared_mem_bytes: 0,
22537 },
22538 )
22539 }
22540 } else {
22541 return self.fa_decode_scalar_unified(
22544 q,
22545 k,
22546 v,
22547 o,
22548 head_dim,
22549 n_head,
22550 n_head_kv,
22551 t_kv,
22552 None,
22553 scale,
22554 n_splits,
22555 if fa_vec { sp } else { 256 },
22556 k_tok_bytes,
22557 v_tok_bytes,
22558 g,
22559 part_o,
22560 part_m,
22561 part_l,
22562 None,
22563 );
22564 };
22565 let __s_b = self.gpu.stream();
22566 let mut b = __s_b.launch_builder(&f);
22567 b.arg(q)
22568 .arg(k)
22569 .arg(v)
22570 .arg(&mut *part_o)
22571 .arg(&mut *part_m)
22572 .arg(&mut *part_l)
22573 .arg(&hd)
22574 .arg(&nh)
22575 .arg(&nhkv)
22576 .arg(&tkvi)
22577 .arg(&scale)
22578 .arg(&nsp)
22579 .arg(&ktb)
22580 .arg(&vtb);
22581 unsafe {
22582 b.launch(cfg)?;
22583 }
22584 let (fc, cfg2) = (
22587 if g {
22588 self.func_g("fa_decode_combine_f32")
22589 } else {
22590 self.fa_func("fa_decode_combine_f32", head_dim)
22591 },
22592 LaunchConfig {
22593 grid_dim: (n_head as u32, 1, 1),
22594 block_dim: (head_dim as u32, 1, 1),
22595 shared_mem_bytes: 0,
22596 },
22597 );
22598 let __s_b2 = self.gpu.stream();
22599 let mut b2 = __s_b2.launch_builder(&fc);
22600 b2.arg(&*part_o)
22601 .arg(&*part_m)
22602 .arg(&*part_l)
22603 .arg(o)
22604 .arg(&hd)
22605 .arg(&nh)
22606 .arg(&nsp);
22607 unsafe {
22608 b2.launch(cfg2)?;
22609 }
22610 Ok(())
22611 }
22612
22613 #[allow(clippy::too_many_arguments)]
22624 pub fn fa_decode_batch_seqs_v4(
22625 &self,
22626 q: &CudaSlice<f32>,
22627 kv_ptrs: &cudarc::driver::CudaView<u64>,
22628 pos_seq: &CudaSlice<i32>,
22629 o: &mut CudaSlice<f32>,
22630 head_dim: usize,
22631 n_head: usize,
22632 n_head_kv: usize,
22633 b_n: usize,
22634 t_kv_max: usize,
22635 scale: f32,
22636 split_keys: usize,
22637 k_tok_bytes: usize,
22638 v_tok_bytes: usize,
22639 ) -> Result<(), Box<dyn std::error::Error>> {
22640 debug_assert!(head_dim == 256, "seqs twin is v4-stamped (hd256 only)");
22641 let n_splits_max = (t_kv_max + split_keys - 1) / split_keys;
22642 let o_len = b_n * n_head * n_splits_max * head_dim;
22643 let ml_len = b_n * n_head * n_splits_max;
22644 let mut part_guard = self.fa_part_pool.lock().unwrap();
22645 if part_guard
22646 .as_ref()
22647 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
22648 .unwrap_or(true)
22649 {
22650 let old = part_guard.take();
22661 let (co, cm) = old
22662 .as_ref()
22663 .map(|pp| (pp.0.len(), pp.1.len()))
22664 .unwrap_or((0, 0));
22665 if let Some(old) = old {
22666 self.fa_part_retired.lock().unwrap().push(old);
22667 }
22668 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
22669 eprintln!(
22670 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
22671 co, o_len, cm, ml_len
22672 );
22673 }
22674 *part_guard =
22675 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
22676 }
22677 let pg = part_guard.as_mut().unwrap();
22678 self.gpu
22679 .stream()
22680 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
22681 self.gpu
22682 .stream()
22683 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
22684 self.gpu
22685 .stream()
22686 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
22687 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
22688 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
22689 let (nspm, spk) = (n_splits_max as i32, split_keys as i32);
22690 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
22691 let gqa = (n_head / n_head_kv).max(1) as u32;
22692 let f = self.func("fa_decode_vec_q_seqs_v4");
22693 let shmem = (11520 + 32 * head_dim * 2) as u32;
22695 use cudarc::driver::sys::CUfunction_attribute_enum as A;
22696 f.set_attribute(
22697 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
22698 shmem as i32,
22699 )?;
22700 let cfg = LaunchConfig {
22701 grid_dim: (n_head_kv as u32, n_splits_max as u32, b_n as u32),
22702 block_dim: (32, gqa, 1),
22703 shared_mem_bytes: shmem,
22704 };
22705 {
22706 let __s_b = self.gpu.stream();
22707 let mut b = __s_b.launch_builder(&f);
22708 b.arg(q)
22709 .arg(kv_ptrs)
22710 .arg(pos_seq)
22711 .arg(&mut *part_o)
22712 .arg(&mut *part_m)
22713 .arg(&mut *part_l)
22714 .arg(&hd)
22715 .arg(&nh)
22716 .arg(&nhkv)
22717 .arg(&scale)
22718 .arg(&nspm)
22719 .arg(&spk)
22720 .arg(&ktb)
22721 .arg(&vtb);
22722 unsafe {
22723 b.launch(cfg)?;
22724 }
22725 }
22726 let fc = self.func("fa_decode_combine_seqs");
22727 let cfg2 = LaunchConfig {
22728 grid_dim: (n_head as u32, b_n as u32, 1),
22729 block_dim: (head_dim as u32, 1, 1),
22730 shared_mem_bytes: 0,
22731 };
22732 let __s_b2 = self.gpu.stream();
22733 let mut b2 = __s_b2.launch_builder(&fc);
22734 b2.arg(&*part_o)
22735 .arg(&*part_m)
22736 .arg(&*part_l)
22737 .arg(o)
22738 .arg(&hd)
22739 .arg(&nh)
22740 .arg(pos_seq)
22741 .arg(&nspm)
22742 .arg(&spk);
22743 unsafe {
22744 b2.launch(cfg2)?;
22745 }
22746 Ok(())
22747 }
22748
22749 #[allow(clippy::too_many_arguments)]
22756 pub fn append_kv_quantized_seqs(
22757 &self,
22758 k_rows: &CudaSlice<f32>,
22759 v_rows: &CudaSlice<f32>,
22760 kv_ptrs: &cudarc::driver::CudaView<u64>,
22761 pos_seq: &CudaSlice<i32>,
22762 b_n: usize,
22763 kv_dim_k: usize,
22764 kv_dim_v: usize,
22765 k_tok_bytes: usize,
22766 v_tok_bytes: usize,
22767 ) -> Result<(), Box<dyn std::error::Error>> {
22768 let f = self.func("append_quantize_kv_q8_0_q5_1_seqs");
22769 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
22770 let cfg = LaunchConfig {
22771 grid_dim: (nblk, b_n as u32, 1),
22772 block_dim: (32, 1, 1),
22773 shared_mem_bytes: 0,
22774 };
22775 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
22776 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
22777 let __s_b = self.gpu.stream();
22778 let mut b = __s_b.launch_builder(&f);
22779 b.arg(k_rows)
22780 .arg(v_rows)
22781 .arg(kv_ptrs)
22782 .arg(pos_seq)
22783 .arg(&kdk)
22784 .arg(&kdv)
22785 .arg(&ktb)
22786 .arg(&vtb);
22787 unsafe {
22788 b.launch(cfg)?;
22789 }
22790 Ok(())
22791 }
22792
22793 pub fn fa_rows_eligible(&self, base_len: usize, head_dim: usize) -> bool {
22799 std::env::var("MEMRA_NO_FA_VEC").is_err()
22800 && std::env::var("MEMRA_FA_ROWS_OFF").is_err()
22801 && base_len + 1 >= fa_vec_min_tkv()
22802 && head_dim <= 256
22803 && head_dim % 32 == 0
22804 }
22805
22806 #[allow(clippy::too_many_arguments)]
22815 pub fn fa_decode_rows(
22816 &self,
22817 q: &CudaSlice<f32>,
22818 k: &cudarc::driver::CudaView<u8>,
22819 v: &cudarc::driver::CudaView<u8>,
22820 o: &mut CudaSlice<f32>,
22821 head_dim: usize,
22822 n_head: usize,
22823 n_head_kv: usize,
22824 base_len: usize,
22825 t: usize,
22826 scale: f32,
22827 k_tok_bytes: usize,
22828 v_tok_bytes: usize,
22829 base_dev: Option<(&CudaSlice<i32>, i32)>,
22833 kv_shared: bool,
22836 g: bool,
22840 mut q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
22843 ) -> Result<(), Box<dyn std::error::Error>> {
22844 debug_assert!(base_len + 1 >= fa_vec_min_tkv() && head_dim <= 512 && head_dim % 32 == 0);
22845 let t_kv_max = base_len + t; let mut sp = fa_split_keys(t_kv_max, n_head_kv); if head_dim == 512 {
22852 static SP512: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
22853 let v = *SP512.get_or_init(|| {
22856 std::env::var("MEMRA_FA_SP512")
22857 .ok()
22858 .and_then(|x| x.parse().ok())
22859 .unwrap_or(0)
22860 });
22861 sp = if v >= 8 {
22862 v
22863 } else {
22864 FA_SP512_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
22865 };
22866 }
22867 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
22868 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
22869 let gqa = (n_head / n_head_kv).max(1) as u32;
22870 let mut groups: Vec<(usize, usize, usize)> = Vec::new(); if head_dim == 512 || fa_split_keys(base_len + 1, n_head_kv) == sp {
22881 groups.push((0, t, sp));
22882 } else {
22883 let mut r0 = 0usize;
22884 while r0 < t {
22885 let sp_g = fa_split_keys(base_len + r0 + 1, n_head_kv);
22886 let mut r1 = r0 + 1;
22887 while r1 < t && fa_split_keys(base_len + r1 + 1, n_head_kv) == sp_g {
22888 r1 += 1;
22889 }
22890 groups.push((r0, r1 - r0, sp_g));
22891 r0 = r1;
22892 }
22893 }
22894 static SMEM_TKV_R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
22898 let smem_tkv = *SMEM_TKV_R.get_or_init(|| {
22899 std::env::var("MEMRA_FA_SMEM_TKV")
22900 .ok()
22901 .and_then(|v| v.parse().ok())
22902 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
22903 });
22904 let v4 = fa_v4_at(base_len + t) && head_dim == 256;
22905 let v3 = fa_v3_active(head_dim);
22906 let smem_rows =
22907 head_dim <= 256 && !v3 && !fa_v2_on() && smem_tkv > 0 && t_kv_max >= smem_tkv;
22908 let _ = kv_shared;
22913 let i2 = head_dim == 512 && std::env::var("MEMRA_FA_I2").as_deref() != Ok("0");
22916 static TB512: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
22930 let tb512 = head_dim == 512
22932 && sp <= 32
22933 && n_head / n_head_kv.max(1) <= 16
22934 && *TB512.get_or_init(|| std::env::var("MEMRA_FA_TB512").as_deref() != Ok("0"));
22935 let fname = if tb512 {
22936 "fa_decode_vec_q_rows_v4_512_tb"
22937 } else if i2 {
22938 "fa_decode_vec_q_rows_dpl16_i2"
22939 } else if head_dim == 512 {
22940 "fa_decode_vec_q_rows_dpl16"
22941 }
22942 else if v4 {
22944 "fa_decode_vec_q_rows_v4"
22945 } else if v3 {
22946 "fa_decode_vec_q_rows_v3"
22947 } else if fa_v2_on() {
22948 "fa_decode_vec_q_rows_v2"
22949 } else if smem_rows {
22950 "fa_decode_vec_q_rows_smem"
22951 } else {
22952 "fa_decode_vec_q_rows"
22953 };
22954 let f = if head_dim == 512 {
22955 self.fa_func(fname, head_dim)
22956 } else if g {
22957 self.func_g(if smem_rows {
22965 "fa_decode_vec_q_rows"
22966 } else {
22967 fname
22968 })
22969 } else {
22970 self.func(fname)
22971 };
22972 let shmem = if tb512 {
22973 let gk = Self::gkv_on();
22975 let sh =
22976 (8192 + 1024 + 32 * 512 + 32 * 64 + 32 * head_dim * if gk { 1 } else { 2 }) as u32;
22977 use cudarc::driver::sys::CUfunction_attribute_enum as A;
22978 f.set_attribute(
22979 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
22980 sh as i32,
22981 )?;
22982 sh
22983 } else if v4 || v3 || smem_rows || fa_v2_on() {
22984 let sh = (if v4 {
22986 11520 + 32 * head_dim * if g { 1 } else { 2 }
22987 } else if v3 {
22988 32 * head_dim * 2
22989 } else {
22990 2 * 32 * head_dim * 2
22991 }) as u32;
22992 use cudarc::driver::sys::CUfunction_attribute_enum as A;
22993 f.set_attribute(
22994 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
22995 sh as i32,
22996 )?;
22997 sh
22998 } else {
22999 0
23000 };
23001 for &(r0, t_g, sp_g) in &groups {
23005 let n_splits_g = (base_len + r0 + t_g).div_ceil(sp_g);
23006 let (nspm, spk) = (n_splits_g as i32, sp_g as i32);
23007 let base_i = (base_len + r0) as i32;
23008 let o_len = t_g * n_head * n_splits_g * head_dim;
23009 let ml_len = t_g * n_head * n_splits_g;
23010 let mut part_guard = self.fa_part_pool.lock().unwrap();
23011 if part_guard
23012 .as_ref()
23013 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
23014 .unwrap_or(true)
23015 {
23016 let old = part_guard.take();
23027 let (co, cm) = old
23028 .as_ref()
23029 .map(|pp| (pp.0.len(), pp.1.len()))
23030 .unwrap_or((0, 0));
23031 if let Some(old) = old {
23032 self.fa_part_retired.lock().unwrap().push(old);
23033 }
23034 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
23035 eprintln!(
23036 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
23037 co, o_len, cm, ml_len
23038 );
23039 }
23040 *part_guard =
23041 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
23042 }
23043 let pg = part_guard.as_mut().unwrap();
23044 self.gpu
23045 .stream()
23046 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
23047 self.gpu
23048 .stream()
23049 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
23050 self.gpu
23051 .stream()
23052 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
23053 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
23054 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
23055 let qv = self.view(q, t * n_head * head_dim);
23056 let q_g = qv.slice(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
23057 let cfg = LaunchConfig {
23058 grid_dim: (n_head_kv as u32, n_splits_g as u32, t_g as u32),
23059 block_dim: (32, gqa, 1),
23060 shared_mem_bytes: shmem,
23061 };
23062 {
23063 let __s_b = self.gpu.stream();
23064 let mut b = __s_b.launch_builder(&f);
23065 if tb512 {
23066 let (bd, plus) =
23068 base_dev.expect("hd512 rows twin requires a device base counter");
23069 let plus_g = plus + r0 as i32;
23070 let nr = t_g as i32;
23071 if Self::pdl_on() && Self::pdl_wb_on() {
23072 use cudarc::driver::{DevicePtr, DevicePtrMut};
23074 let s = &self.gpu.stream();
23075 let (pq, _b0) = q_g.device_ptr(s);
23076 let (pk, _b1) = k.device_ptr(s);
23077 let (pv, _b2) = v.device_ptr(s);
23078 let (po, _b3) = part_o.device_ptr_mut(s);
23079 let (pm, _b4) = part_m.device_ptr_mut(s);
23080 let (pl, _b5) = part_l.device_ptr_mut(s);
23081 let (pb, _b6) = bd.device_ptr(s);
23082 let mut ps = [
23083 &pq as *const _ as *mut std::ffi::c_void,
23084 &pk as *const _ as *mut _,
23085 &pv as *const _ as *mut _,
23086 &po as *const _ as *mut _,
23087 &pm as *const _ as *mut _,
23088 &pl as *const _ as *mut _,
23089 &hd as *const _ as *mut _,
23090 &nh as *const _ as *mut _,
23091 &nhkv as *const _ as *mut _,
23092 &pb as *const _ as *mut _,
23093 &plus_g as *const _ as *mut _,
23094 &scale as *const _ as *mut _,
23095 &nspm as *const _ as *mut _,
23096 &spk as *const _ as *mut _,
23097 &ktb as *const _ as *mut _,
23098 &vtb as *const _ as *mut _,
23099 &nr as *const _ as *mut _,
23100 ];
23101 unsafe {
23102 self.launch_pdl_flash(
23103 Self::gkv_on(),
23104 "fa_decode_vec_q_rows_v4_512_tb",
23105 (n_head_kv as u32, n_splits_g as u32, 1),
23106 (32, gqa, 1),
23107 shmem,
23108 &mut ps,
23109 )?;
23110 }
23111 } else {
23112 let cfg_tb = LaunchConfig {
23113 grid_dim: (n_head_kv as u32, n_splits_g as u32, 1),
23114 block_dim: (32, gqa, 1),
23115 shared_mem_bytes: shmem,
23116 };
23117 b.arg(&q_g)
23118 .arg(k)
23119 .arg(v)
23120 .arg(&mut *part_o)
23121 .arg(&mut *part_m)
23122 .arg(&mut *part_l)
23123 .arg(&hd)
23124 .arg(&nh)
23125 .arg(&nhkv)
23126 .arg(bd)
23127 .arg(&plus_g)
23128 .arg(&scale)
23129 .arg(&nspm)
23130 .arg(&spk)
23131 .arg(&ktb)
23132 .arg(&vtb)
23133 .arg(&nr);
23134 unsafe {
23135 b.launch(cfg_tb)?;
23136 }
23137 }
23138 } else if head_dim == 512 {
23139 let (bd, plus) =
23140 base_dev.expect("hd512 rows twin requires a device base counter");
23141 let plus_g = plus + r0 as i32;
23142 b.arg(&q_g)
23143 .arg(k)
23144 .arg(v)
23145 .arg(&mut *part_o)
23146 .arg(&mut *part_m)
23147 .arg(&mut *part_l)
23148 .arg(&hd)
23149 .arg(&nh)
23150 .arg(&nhkv)
23151 .arg(bd)
23152 .arg(&plus_g)
23153 .arg(&scale)
23154 .arg(&nspm)
23155 .arg(&spk)
23156 .arg(&ktb)
23157 .arg(&vtb);
23158 unsafe {
23159 b.launch(cfg)?;
23160 }
23161 } else {
23162 b.arg(&q_g)
23163 .arg(k)
23164 .arg(v)
23165 .arg(&mut *part_o)
23166 .arg(&mut *part_m)
23167 .arg(&mut *part_l)
23168 .arg(&hd)
23169 .arg(&nh)
23170 .arg(&nhkv)
23171 .arg(&base_i)
23172 .arg(&scale)
23173 .arg(&nspm)
23174 .arg(&spk)
23175 .arg(&ktb)
23176 .arg(&vtb);
23177 unsafe {
23178 b.launch(cfg)?;
23179 }
23180 }
23181 }
23182 let cfg2 = LaunchConfig {
23183 grid_dim: (n_head as u32, t_g as u32, 1),
23184 block_dim: (head_dim as u32, 1, 1),
23185 shared_mem_bytes: 0,
23186 };
23187 let mut o_g = o.slice_mut(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
23188 if head_dim == 512 {
23189 let (bd, plus) = base_dev.unwrap();
23192 let plus_g = plus + r0 as i32;
23193 if let Some((oq, od)) = q8_out.as_mut() {
23194 debug_assert!(t == 1, "rows q8 emit is a t=1 decode arm");
23196 if Self::pdl_on() && Self::pdl_wb_on() {
23197 use cudarc::driver::{DevicePtr, DevicePtrMut};
23199 let s = &self.gpu.stream();
23200 let (po, _g0) = part_o.device_ptr(s);
23201 let (pm, _g1) = part_m.device_ptr(s);
23202 let (pl, _g2) = part_l.device_ptr(s);
23203 let (pq, _g3) = oq.device_ptr_mut(s);
23204 let (pd, _g4) = od.device_ptr_mut(s);
23205 let (pb, _g5) = bd.device_ptr(s);
23206 let mut ps = [
23207 &po as *const _ as *mut std::ffi::c_void,
23208 &pm as *const _ as *mut _,
23209 &pl as *const _ as *mut _,
23210 &pq as *const _ as *mut _,
23211 &pd as *const _ as *mut _,
23212 &hd as *const _ as *mut _,
23213 &nh as *const _ as *mut _,
23214 &pb as *const _ as *mut _,
23215 &plus_g as *const _ as *mut _,
23216 &nspm as *const _ as *mut _,
23217 &spk as *const _ as *mut _,
23218 ];
23219 unsafe {
23220 self.launch_pdl_flash(
23221 Self::gkv_on(),
23222 "fa_decode_combine_rows_dc_q8_1",
23223 cfg2.grid_dim,
23224 cfg2.block_dim,
23225 0,
23226 &mut ps,
23227 )?;
23228 }
23229 continue;
23230 }
23231 let fc = self.fa_func("fa_decode_combine_rows_dc_q8_1", head_dim);
23232 let __s_b2 = self.gpu.stream();
23233 let mut b2 = __s_b2.launch_builder(&fc);
23234 b2.arg(&*part_o)
23235 .arg(&*part_m)
23236 .arg(&*part_l)
23237 .arg(&mut **oq)
23238 .arg(&mut **od)
23239 .arg(&hd)
23240 .arg(&nh)
23241 .arg(bd)
23242 .arg(&plus_g)
23243 .arg(&nspm)
23244 .arg(&spk);
23245 unsafe {
23246 b2.launch(cfg2)?;
23247 }
23248 continue;
23249 }
23250 let fc = self.fa_func("fa_decode_combine_rows_dc", head_dim);
23251 let __s_b2 = self.gpu.stream();
23252 let mut b2 = __s_b2.launch_builder(&fc);
23253 b2.arg(&*part_o)
23254 .arg(&*part_m)
23255 .arg(&*part_l)
23256 .arg(&mut o_g)
23257 .arg(&hd)
23258 .arg(&nh)
23259 .arg(bd)
23260 .arg(&plus_g)
23261 .arg(&nspm)
23262 .arg(&spk);
23263 unsafe {
23264 b2.launch(cfg2)?;
23265 }
23266 } else {
23267 assert!(
23270 q8_out.is_none(),
23271 "rows q8 emit requires the hd512 dc combine"
23272 );
23273 let fc = self.func("fa_decode_combine_rows");
23274 let __s_b2 = self.gpu.stream();
23275 let mut b2 = __s_b2.launch_builder(&fc);
23276 b2.arg(&*part_o)
23277 .arg(&*part_m)
23278 .arg(&*part_l)
23279 .arg(&mut o_g)
23280 .arg(&hd)
23281 .arg(&nh)
23282 .arg(&base_i)
23283 .arg(&nspm)
23284 .arg(&spk);
23285 unsafe {
23286 b2.launch(cfg2)?;
23287 }
23288 }
23289 }
23290 Ok(())
23291 }
23292
23293 #[allow(clippy::too_many_arguments)]
23297 pub fn fa_decode_rows_w(
23298 &self,
23299 q: &CudaSlice<f32>,
23300 k: &cudarc::driver::CudaView<u8>,
23301 v: &cudarc::driver::CudaView<u8>,
23302 o: &mut CudaSlice<f32>,
23303 head_dim: usize,
23304 n_head: usize,
23305 n_head_kv: usize,
23306 base_dev: &CudaSlice<i32>,
23307 base_plus: i32,
23308 t: usize,
23309 scale: f32,
23310 window: usize,
23311 k_tok_bytes: usize,
23312 v_tok_bytes: usize,
23313 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
23314 ) -> Result<(), Box<dyn std::error::Error>> {
23315 debug_assert!(head_dim == 256);
23320 let sp = {
23328 static SPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
23329 let v = *SPW.get_or_init(|| {
23330 std::env::var("MEMRA_FA_SPW")
23331 .ok()
23332 .and_then(|x| x.parse().ok())
23333 .unwrap_or(0)
23334 });
23335 if v >= 8 {
23336 v
23337 } else {
23338 FA_SPW_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
23339 }
23340 };
23341 let n_splits_max = (window + sp - 1) / sp;
23342 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
23343 let (nspm, spk, wini) = (n_splits_max as i32, sp as i32, window as i32);
23344 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
23345 let gqa = (n_head / n_head_kv).max(1) as u32;
23346 let o_len = t * n_head * n_splits_max * head_dim;
23347 let ml_len = t * n_head * n_splits_max;
23348 let mut part_guard = self.fa_part_pool.lock().unwrap();
23349 if part_guard
23350 .as_ref()
23351 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
23352 .unwrap_or(true)
23353 {
23354 let old = part_guard.take();
23365 let (co, cm) = old
23366 .as_ref()
23367 .map(|pp| (pp.0.len(), pp.1.len()))
23368 .unwrap_or((0, 0));
23369 if let Some(old) = old {
23370 self.fa_part_retired.lock().unwrap().push(old);
23371 }
23372 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
23373 eprintln!(
23374 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
23375 co, o_len, cm, ml_len
23376 );
23377 }
23378 *part_guard =
23379 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
23380 }
23381 let pg = part_guard.as_mut().unwrap();
23382 self.gpu
23383 .stream()
23384 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
23385 self.gpu
23386 .stream()
23387 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
23388 self.gpu
23389 .stream()
23390 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
23391 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
23392 static SMEM_TKV_W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
23398 let smem_tkv = *SMEM_TKV_W.get_or_init(|| {
23399 std::env::var("MEMRA_FA_SMEM_TKV")
23400 .ok()
23401 .and_then(|v| v.parse().ok())
23402 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
23403 });
23404 use cudarc::driver::sys::CUfunction_attribute_enum as A;
23410 let wg = Self::wkv_on();
23415 let sp2 =
23418 gqa <= 4 && fa_v4_at(window) && std::env::var("MEMRA_FA_SPW2").as_deref() != Ok("0");
23419 if sp2 {
23420 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
23421 if Self::pdl_on() && Self::pdl_wb_on() {
23422 use cudarc::driver::{DevicePtr, DevicePtrMut};
23424 let s = &self.gpu.stream();
23425 let (pq, _b0) = q.device_ptr(s);
23426 let (pk, _b1) = k.device_ptr(s);
23427 let (pv, _b2) = v.device_ptr(s);
23428 let (po, _b3) = part_o.device_ptr_mut(s);
23429 let (pm, _b4) = part_m.device_ptr_mut(s);
23430 let (pl, _b5) = part_l.device_ptr_mut(s);
23431 let (pb, _b6) = base_dev.device_ptr(s);
23432 let mut ps = [
23433 &pq as *const _ as *mut std::ffi::c_void,
23434 &pk as *const _ as *mut _,
23435 &pv as *const _ as *mut _,
23436 &po as *const _ as *mut _,
23437 &pm as *const _ as *mut _,
23438 &pl as *const _ as *mut _,
23439 &hd as *const _ as *mut _,
23440 &nh as *const _ as *mut _,
23441 &nhkv as *const _ as *mut _,
23442 &pb as *const _ as *mut _,
23443 &base_plus as *const _ as *mut _,
23444 &scale as *const _ as *mut _,
23445 &nspm as *const _ as *mut _,
23446 &spk as *const _ as *mut _,
23447 &ktb as *const _ as *mut _,
23448 &vtb as *const _ as *mut _,
23449 &wini as *const _ as *mut _,
23450 ];
23451 unsafe {
23452 self.launch_pdl_flash(
23453 wg,
23454 "fa_decode_vec_q_rows_v4_w_sp",
23455 (n_head_kv as u32, n_splits_max as u32, t as u32),
23456 (32, gqa + 1, 1),
23457 sh,
23458 &mut ps,
23459 )?;
23460 }
23461 } else {
23462 let f = if wg {
23463 self.func_g("fa_decode_vec_q_rows_v4_w_sp")
23464 } else {
23465 self.func("fa_decode_vec_q_rows_v4_w_sp")
23466 };
23467 f.set_attribute(
23468 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
23469 sh as i32,
23470 )?;
23471 let cfg = LaunchConfig {
23472 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
23473 block_dim: (32, gqa + 1, 1),
23474 shared_mem_bytes: sh,
23475 };
23476 let __s_b = self.gpu.stream();
23477 let mut b = __s_b.launch_builder(&f);
23478 b.arg(q)
23479 .arg(k)
23480 .arg(v)
23481 .arg(&mut *part_o)
23482 .arg(&mut *part_m)
23483 .arg(&mut *part_l)
23484 .arg(&hd)
23485 .arg(&nh)
23486 .arg(&nhkv)
23487 .arg(base_dev)
23488 .arg(&base_plus)
23489 .arg(&scale)
23490 .arg(&nspm)
23491 .arg(&spk)
23492 .arg(&ktb)
23493 .arg(&vtb)
23494 .arg(&wini);
23495 unsafe {
23496 b.launch(cfg)?;
23497 }
23498 }
23499 } else {
23500 if fa_v4_at(window) && Self::pdl_on() && Self::pdl_wb_on() {
23501 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
23503 use cudarc::driver::{DevicePtr, DevicePtrMut};
23504 let s = &self.gpu.stream();
23505 let (pq, _b0) = q.device_ptr(s);
23506 let (pk, _b1) = k.device_ptr(s);
23507 let (pv, _b2) = v.device_ptr(s);
23508 let (po, _b3) = part_o.device_ptr_mut(s);
23509 let (pm, _b4) = part_m.device_ptr_mut(s);
23510 let (pl, _b5) = part_l.device_ptr_mut(s);
23511 let (pb, _b6) = base_dev.device_ptr(s);
23512 let mut ps = [
23513 &pq as *const _ as *mut std::ffi::c_void,
23514 &pk as *const _ as *mut _,
23515 &pv as *const _ as *mut _,
23516 &po as *const _ as *mut _,
23517 &pm as *const _ as *mut _,
23518 &pl as *const _ as *mut _,
23519 &hd as *const _ as *mut _,
23520 &nh as *const _ as *mut _,
23521 &nhkv as *const _ as *mut _,
23522 &pb as *const _ as *mut _,
23523 &base_plus as *const _ as *mut _,
23524 &scale as *const _ as *mut _,
23525 &nspm as *const _ as *mut _,
23526 &spk as *const _ as *mut _,
23527 &ktb as *const _ as *mut _,
23528 &vtb as *const _ as *mut _,
23529 &wini as *const _ as *mut _,
23530 ];
23531 unsafe {
23532 self.launch_pdl_flash(
23533 wg,
23534 "fa_decode_vec_q_rows_v4_w",
23535 (n_head_kv as u32, n_splits_max as u32, t as u32),
23536 (32, gqa, 1),
23537 sh,
23538 &mut ps,
23539 )?;
23540 }
23541 } else {
23542 let pick = |name: &str| {
23543 if wg {
23544 self.func_g(name)
23545 } else {
23546 self.func(name)
23547 }
23548 };
23549 let (f, sh) = if fa_v4_at(window) {
23550 let f = pick("fa_decode_vec_q_rows_v4_w");
23551 (f, (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32)
23552 } else if smem_tkv > 0 && window >= smem_tkv {
23553 (
23556 pick("fa_decode_vec_q_rows_smem_w"),
23557 (2 * 32 * head_dim * 2) as u32,
23558 )
23559 } else {
23560 (pick("fa_decode_vec_q_rows_reg_w"), 0u32)
23561 };
23562 f.set_attribute(
23563 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
23564 sh as i32,
23565 )?;
23566 let cfg = LaunchConfig {
23567 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
23568 block_dim: (32, gqa, 1),
23569 shared_mem_bytes: sh,
23570 };
23571 let __s_b = self.gpu.stream();
23572 let mut b = __s_b.launch_builder(&f);
23573 b.arg(q)
23574 .arg(k)
23575 .arg(v)
23576 .arg(&mut *part_o)
23577 .arg(&mut *part_m)
23578 .arg(&mut *part_l)
23579 .arg(&hd)
23580 .arg(&nh)
23581 .arg(&nhkv)
23582 .arg(base_dev)
23583 .arg(&base_plus)
23584 .arg(&scale)
23585 .arg(&nspm)
23586 .arg(&spk)
23587 .arg(&ktb)
23588 .arg(&vtb)
23589 .arg(&wini);
23590 unsafe {
23591 b.launch(cfg)?;
23592 }
23593 }
23594 }
23595 let cfg2 = LaunchConfig {
23596 grid_dim: (n_head as u32, t as u32, 1),
23597 block_dim: (head_dim as u32, 1, 1),
23598 shared_mem_bytes: 0,
23599 };
23600 if let Some((oq, od)) = q8_out {
23601 if Self::pdl_on() && Self::pdl_wb_on() {
23604 use cudarc::driver::{DevicePtr, DevicePtrMut};
23606 let s = &self.gpu.stream();
23607 let (po, _g0) = part_o.device_ptr(s);
23608 let (pm, _g1) = part_m.device_ptr(s);
23609 let (pl, _g2) = part_l.device_ptr(s);
23610 let (pq, _g3) = oq.device_ptr_mut(s);
23611 let (pd, _g4) = od.device_ptr_mut(s);
23612 let mut ps = [
23613 &po as *const _ as *mut std::ffi::c_void,
23614 &pm as *const _ as *mut _,
23615 &pl as *const _ as *mut _,
23616 &pq as *const _ as *mut _,
23617 &pd as *const _ as *mut _,
23618 &hd as *const _ as *mut _,
23619 &nh as *const _ as *mut _,
23620 &nspm as *const _ as *mut _,
23621 &spk as *const _ as *mut _,
23622 &wini as *const _ as *mut _,
23623 ];
23624 unsafe {
23625 self.launch_pdl_flash(
23626 wg,
23627 "fa_decode_combine_rows_w_q8_1",
23628 cfg2.grid_dim,
23629 cfg2.block_dim,
23630 0,
23631 &mut ps,
23632 )?;
23633 }
23634 return Ok(());
23635 }
23636 let fc = if wg {
23637 self.func_g("fa_decode_combine_rows_w_q8_1")
23638 } else {
23639 self.func("fa_decode_combine_rows_w_q8_1")
23640 };
23641 let __s_b2 = self.gpu.stream();
23642 let mut b2 = __s_b2.launch_builder(&fc);
23643 b2.arg(&*part_o)
23644 .arg(&*part_m)
23645 .arg(&*part_l)
23646 .arg(oq)
23647 .arg(od)
23648 .arg(&hd)
23649 .arg(&nh)
23650 .arg(&nspm)
23651 .arg(&spk)
23652 .arg(&wini);
23653 unsafe {
23654 b2.launch(cfg2)?;
23655 }
23656 return Ok(());
23657 }
23658 let fc = if wg {
23659 self.func_g("fa_decode_combine_rows_w")
23660 } else {
23661 self.func("fa_decode_combine_rows_w")
23662 };
23663 let __s_b2 = self.gpu.stream();
23664 let mut b2 = __s_b2.launch_builder(&fc);
23665 b2.arg(&*part_o)
23666 .arg(&*part_m)
23667 .arg(&*part_l)
23668 .arg(o)
23669 .arg(&hd)
23670 .arg(&nh)
23671 .arg(&nspm)
23672 .arg(&spk)
23673 .arg(&wini);
23674 unsafe {
23675 b2.launch(cfg2)?;
23676 }
23677 Ok(())
23678 }
23679
23680 #[allow(clippy::too_many_arguments)]
23686 pub fn fa_decode_rows_dc(
23687 &self,
23688 q: &CudaSlice<f32>,
23689 k: &cudarc::driver::CudaView<u8>,
23690 v: &cudarc::driver::CudaView<u8>,
23691 o: &mut CudaSlice<f32>,
23692 head_dim: usize,
23693 n_head: usize,
23694 n_head_kv: usize,
23695 base_dev: &CudaSlice<i32>,
23696 t_kv_upper: usize,
23697 t: usize,
23698 scale: f32,
23699 k_tok_bytes: usize,
23700 v_tok_bytes: usize,
23701 base_plus: i32,
23702 g: bool,
23703 ) -> Result<(), Box<dyn std::error::Error>> {
23704 let v4 = head_dim == 256 && fa_v4_at(t_kv_upper);
23705 assert!(
23706 v4 || fa_v3_active(head_dim),
23707 "stream fa rows requires the v3 or v4 lane"
23708 );
23709 assert!(v4 || base_plus == 0, "v3_dc kernel takes no plus arg");
23710 if v4 {
23711 let sp = fa_split_keys(t_kv_upper, n_head_kv);
23712 let n_splits_max = (t_kv_upper + sp - 1) / sp;
23713 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
23714 let (nspm, spk) = (n_splits_max as i32, sp as i32);
23715 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
23716 let gqa = (n_head / n_head_kv).max(1) as u32;
23717 let o_len = t * n_head * n_splits_max * head_dim;
23718 let ml_len = t * n_head * n_splits_max;
23719 let mut part_guard = self.fa_part_pool.lock().unwrap();
23720 if part_guard
23721 .as_ref()
23722 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
23723 .unwrap_or(true)
23724 {
23725 let old = part_guard.take();
23736 let (co, cm) = old
23737 .as_ref()
23738 .map(|pp| (pp.0.len(), pp.1.len()))
23739 .unwrap_or((0, 0));
23740 if let Some(old) = old {
23741 self.fa_part_retired.lock().unwrap().push(old);
23742 }
23743 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
23744 eprintln!(
23745 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
23746 co, o_len, cm, ml_len
23747 );
23748 }
23749 *part_guard =
23750 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
23751 }
23752 let pg = part_guard.as_mut().unwrap();
23753 self.gpu
23754 .stream()
23755 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
23756 self.gpu
23757 .stream()
23758 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
23759 self.gpu
23760 .stream()
23761 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
23762 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
23763 let f = if g {
23764 self.func_g("fa_decode_vec_q_rows_v4_dc")
23765 } else {
23766 self.func("fa_decode_vec_q_rows_v4_dc")
23767 };
23768 let sh = (11520 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
23769 use cudarc::driver::sys::CUfunction_attribute_enum as A;
23770 f.set_attribute(
23771 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
23772 sh as i32,
23773 )?;
23774 let cfg = LaunchConfig {
23775 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
23776 block_dim: (32, gqa, 1),
23777 shared_mem_bytes: sh,
23778 };
23779 let __s_b = self.gpu.stream();
23780 let mut b = __s_b.launch_builder(&f);
23781 b.arg(q)
23782 .arg(k)
23783 .arg(v)
23784 .arg(&mut *part_o)
23785 .arg(&mut *part_m)
23786 .arg(&mut *part_l)
23787 .arg(&hd)
23788 .arg(&nh)
23789 .arg(&nhkv)
23790 .arg(base_dev)
23791 .arg(&base_plus)
23792 .arg(&scale)
23793 .arg(&nspm)
23794 .arg(&spk)
23795 .arg(&ktb)
23796 .arg(&vtb);
23797 unsafe {
23798 b.launch(cfg)?;
23799 }
23800 let fc = self.func("fa_decode_combine_rows_dc");
23801 let cfg2 = LaunchConfig {
23802 grid_dim: (n_head as u32, t as u32, 1),
23803 block_dim: (head_dim as u32, 1, 1),
23804 shared_mem_bytes: 0,
23805 };
23806 let __s_b2 = self.gpu.stream();
23807 let mut b2 = __s_b2.launch_builder(&fc);
23808 b2.arg(&*part_o)
23809 .arg(&*part_m)
23810 .arg(&*part_l)
23811 .arg(o)
23812 .arg(&hd)
23813 .arg(&nh)
23814 .arg(base_dev)
23815 .arg(&base_plus)
23816 .arg(&nspm)
23817 .arg(&spk);
23818 unsafe {
23819 b2.launch(cfg2)?;
23820 }
23821 return Ok(());
23822 }
23823 let sp = fa_split_keys(t_kv_upper, n_head_kv);
23824 let n_splits_max = (t_kv_upper + sp - 1) / sp;
23825 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
23826 let (nspm, spk) = (n_splits_max as i32, sp as i32);
23827 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
23828 let gqa = (n_head / n_head_kv).max(1) as u32;
23829 let o_len = t * n_head * n_splits_max * head_dim;
23830 let ml_len = t * n_head * n_splits_max;
23831 let mut part_guard = self.fa_part_pool.lock().unwrap();
23832 if part_guard
23833 .as_ref()
23834 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
23835 .unwrap_or(true)
23836 {
23837 let old = part_guard.take();
23848 let (co, cm) = old
23849 .as_ref()
23850 .map(|pp| (pp.0.len(), pp.1.len()))
23851 .unwrap_or((0, 0));
23852 if let Some(old) = old {
23853 self.fa_part_retired.lock().unwrap().push(old);
23854 }
23855 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
23856 eprintln!(
23857 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
23858 co, o_len, cm, ml_len
23859 );
23860 }
23861 *part_guard =
23862 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
23863 }
23864 let pg = part_guard.as_mut().unwrap();
23865 self.gpu
23866 .stream()
23867 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
23868 self.gpu
23869 .stream()
23870 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
23871 self.gpu
23872 .stream()
23873 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
23874 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
23875 let f = self.func("fa_decode_vec_q_rows_v3_dc");
23876 let sh = (32 * head_dim * 2) as u32;
23877 use cudarc::driver::sys::CUfunction_attribute_enum as A;
23878 f.set_attribute(
23879 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
23880 sh as i32,
23881 )?;
23882 let cfg = LaunchConfig {
23883 grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
23884 block_dim: (32, gqa, 1),
23885 shared_mem_bytes: sh,
23886 };
23887 let __s_b = self.gpu.stream();
23888 let mut b = __s_b.launch_builder(&f);
23889 b.arg(q)
23890 .arg(k)
23891 .arg(v)
23892 .arg(&mut *part_o)
23893 .arg(&mut *part_m)
23894 .arg(&mut *part_l)
23895 .arg(&hd)
23896 .arg(&nh)
23897 .arg(&nhkv)
23898 .arg(base_dev)
23899 .arg(&scale)
23900 .arg(&nspm)
23901 .arg(&spk)
23902 .arg(&ktb)
23903 .arg(&vtb);
23904 unsafe {
23905 b.launch(cfg)?;
23906 }
23907 let fc = self.func("fa_decode_combine_rows_dc");
23908 let cfg2 = LaunchConfig {
23909 grid_dim: (n_head as u32, t as u32, 1),
23910 block_dim: (head_dim as u32, 1, 1),
23911 shared_mem_bytes: 0,
23912 };
23913 let plus0 = 0i32;
23914 let __s_b2 = self.gpu.stream();
23915 let mut b2 = __s_b2.launch_builder(&fc);
23916 b2.arg(&*part_o)
23917 .arg(&*part_m)
23918 .arg(&*part_l)
23919 .arg(o)
23920 .arg(&hd)
23921 .arg(&nh)
23922 .arg(base_dev)
23923 .arg(&plus0)
23924 .arg(&nspm)
23925 .arg(&spk);
23926 unsafe {
23927 b2.launch(cfg2)?;
23928 }
23929 Ok(())
23930 }
23931
23932 pub fn fa_decode_dc(
23943 &self,
23944 q: &CudaSlice<f32>,
23945 k: &cudarc::driver::CudaView<u8>,
23946 v: &cudarc::driver::CudaView<u8>,
23947 o: &mut CudaSlice<f32>,
23948 head_dim: usize,
23949 n_head: usize,
23950 n_head_kv: usize,
23951 t_kv_dev: &CudaSlice<i32>,
23952 bucket_max: usize,
23953 scale: f32,
23954 k_tok_bytes: usize,
23955 v_tok_bytes: usize,
23956 g: bool,
23957 ) -> Result<(), Box<dyn std::error::Error>> {
23958 self.fa_decode_dc_q8(
23959 q,
23960 k,
23961 v,
23962 o,
23963 head_dim,
23964 n_head,
23965 n_head_kv,
23966 t_kv_dev,
23967 bucket_max,
23968 scale,
23969 k_tok_bytes,
23970 v_tok_bytes,
23971 g,
23972 None,
23973 )
23974 }
23975
23976 #[allow(clippy::too_many_arguments)]
23979 pub fn fa_decode_dc_q8(
23980 &self,
23981 q: &CudaSlice<f32>,
23982 k: &cudarc::driver::CudaView<u8>,
23983 v: &cudarc::driver::CudaView<u8>,
23984 o: &mut CudaSlice<f32>,
23985 head_dim: usize,
23986 n_head: usize,
23987 n_head_kv: usize,
23988 t_kv_dev: &CudaSlice<i32>,
23989 bucket_max: usize,
23990 scale: f32,
23991 k_tok_bytes: usize,
23992 v_tok_bytes: usize,
23993 g: bool,
23994 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
23995 ) -> Result<(), Box<dyn std::error::Error>> {
23996 let mut fa_vec =
24004 std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
24005 if g && head_dim == 256 && !fa_v4_at(bucket_max) {
24006 fa_vec = false;
24007 } let sp = fa_split_keys(bucket_max, n_head_kv);
24009 let n_splits = if fa_vec {
24010 ((bucket_max + sp - 1) / sp).max(1)
24011 } else {
24012 ((bucket_max + 255) / 256).max(1)
24013 };
24014 let o_len = n_head * n_splits * head_dim;
24015 let ml_len = n_head * n_splits;
24016 let mut part_guard = self.fa_part_pool.lock().unwrap();
24017 if part_guard
24018 .as_ref()
24019 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
24020 .unwrap_or(true)
24021 {
24022 let old = part_guard.take();
24033 let (co, cm) = old
24034 .as_ref()
24035 .map(|pp| (pp.0.len(), pp.1.len()))
24036 .unwrap_or((0, 0));
24037 if let Some(old) = old {
24038 self.fa_part_retired.lock().unwrap().push(old);
24039 }
24040 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
24041 eprintln!(
24042 "[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
24043 co, o_len, cm, ml_len
24044 );
24045 }
24046 *part_guard =
24047 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
24048 }
24049 let pg = part_guard.as_mut().unwrap();
24050 self.gpu
24051 .stream()
24052 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
24053 self.gpu
24054 .stream()
24055 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
24056 self.gpu
24057 .stream()
24058 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
24059 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
24060 let (hd, nh, nhkv, nsp) = (
24061 head_dim as i32,
24062 n_head as i32,
24063 n_head_kv as i32,
24064 n_splits as i32,
24065 );
24066 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
24067 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
24068 let deep = fa_vec
24071 && head_dim == 256
24072 && fa_v4_at(bucket_max)
24073 && !g
24074 && fa_deep_at(bucket_max)
24075 && !matches!(fa_v4_mode(), "noB3" | "stage");
24076 let (f, cfg) = if fa_vec
24077 && head_dim == 512
24078 && bucket_max >= {
24079 static FA512_MIN_DC: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
24080 *FA512_MIN_DC.get_or_init(|| {
24081 std::env::var("MEMRA_FA512_MIN")
24082 .ok()
24083 .and_then(|v| v.parse().ok())
24084 .unwrap_or(512)
24085 })
24086 } {
24087 let gqa = (n_head / n_head_kv).max(1) as u32;
24089 (
24090 self.fa_func("fa_decode_vec_q_dpl16_dc", head_dim),
24091 LaunchConfig {
24092 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
24093 block_dim: (32, gqa, 1),
24094 shared_mem_bytes: 0,
24095 },
24096 )
24097 } else if fa_vec && head_dim == 512 {
24098 let q_view = q.as_view();
24101 let mut o_view = o.as_view_mut();
24102 return self.fa_decode_scalar_unified(
24103 &q_view,
24104 k,
24105 v,
24106 &mut o_view,
24107 head_dim,
24108 n_head,
24109 n_head_kv,
24110 0,
24111 Some(t_kv_dev),
24112 scale,
24113 n_splits,
24114 sp,
24115 k_tok_bytes,
24116 v_tok_bytes,
24117 g,
24118 &mut *part_o,
24119 &mut *part_m,
24120 &mut *part_l,
24121 q8_out,
24122 );
24123 } else if fa_vec && head_dim == 256 && fa_v4_at(bucket_max) {
24124 let gqa = (n_head / n_head_kv).max(1) as u32;
24127 let fv = if g {
24128 self.func_g("fa_decode_vec_q_v4_dc")
24129 } else if deep {
24130 self.func("fa_decode_vec_q_v4_deep_dc")
24131 } else {
24132 self.func("fa_decode_vec_q_v4_dc")
24133 };
24134 let shmem =
24135 (if deep { 12160 } else { 11520 } + 32 * head_dim * if g { 1 } else { 2 }) as u32;
24136 use cudarc::driver::sys::CUfunction_attribute_enum as A;
24137 fv.set_attribute(
24138 A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
24139 shmem as i32,
24140 )?;
24141 (
24142 fv,
24143 LaunchConfig {
24144 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
24145 block_dim: (32, gqa, 1),
24146 shared_mem_bytes: shmem,
24147 },
24148 )
24149 } else if fa_vec && fa_v3_active(head_dim) {
24150 let gqa = (n_head / n_head_kv).max(1) as u32;
24153 let fv = if g {
24154 self.func_g("fa_decode_vec_q_v3_dc")
24155 } else {
24156 self.func("fa_decode_vec_q_v3_dc")
24157 };
24158 let shmem = (32 * head_dim * 2) as u32; (
24160 fv,
24161 LaunchConfig {
24162 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
24163 block_dim: (32, gqa, 1),
24164 shared_mem_bytes: shmem,
24165 },
24166 )
24167 } else if fa_vec && fa_v2_on() {
24168 let gqa = (n_head / n_head_kv).max(1) as u32;
24172 let fv = if g {
24173 self.func_g("fa_decode_vec_q_v2_dc")
24174 } else {
24175 self.func("fa_decode_vec_q_v2_dc")
24176 };
24177 let shmem = (2 * 32 * head_dim * 2) as u32; (
24179 fv,
24180 LaunchConfig {
24181 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
24182 block_dim: (32, gqa, 1),
24183 shared_mem_bytes: shmem,
24184 },
24185 )
24186 } else if fa_vec {
24187 let gqa = (n_head / n_head_kv).max(1) as u32;
24188 let fv = if g {
24190 self.func_g("fa_decode_vec_q_dc")
24191 } else {
24192 self.func("fa_decode_vec_q_dc")
24193 };
24194 (
24195 fv,
24196 LaunchConfig {
24197 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
24198 block_dim: (32, gqa, 1),
24199 shared_mem_bytes: 0,
24200 },
24201 )
24202 } else {
24203 let q_view = q.as_view();
24204 let mut o_view = o.as_view_mut();
24205 return self.fa_decode_scalar_unified(
24206 &q_view,
24207 k,
24208 v,
24209 &mut o_view,
24210 head_dim,
24211 n_head,
24212 n_head_kv,
24213 0,
24214 Some(t_kv_dev),
24215 scale,
24216 n_splits,
24217 if fa_vec { sp } else { 256 },
24218 k_tok_bytes,
24219 v_tok_bytes,
24220 g,
24221 &mut *part_o,
24222 &mut *part_m,
24223 &mut *part_l,
24224 q8_out,
24225 );
24226 };
24227 let ski = sp as i32; let __s_b = self.gpu.stream();
24229 let mut b = __s_b.launch_builder(&f);
24230 b.arg(q)
24231 .arg(k)
24232 .arg(v)
24233 .arg(&mut *part_o)
24234 .arg(&mut *part_m)
24235 .arg(&mut *part_l)
24236 .arg(&hd)
24237 .arg(&nh)
24238 .arg(&nhkv)
24239 .arg(t_kv_dev)
24240 .arg(&scale)
24241 .arg(&nsp)
24242 .arg(&ski)
24243 .arg(&ktb)
24244 .arg(&vtb);
24245 unsafe {
24246 b.launch(cfg)?;
24247 }
24248 let cfg2 = LaunchConfig {
24249 grid_dim: (n_head as u32, 1, 1),
24250 block_dim: (head_dim as u32, 1, 1),
24251 shared_mem_bytes: 0,
24252 };
24253 if let Some((oq, od)) = q8_out {
24254 let fc = if g {
24255 self.func_g("fa_decode_combine_q8_1")
24256 } else {
24257 self.fa_func("fa_decode_combine_q8_1", head_dim)
24258 };
24259 let __s_b2 = self.gpu.stream();
24260 let mut b2 = __s_b2.launch_builder(&fc);
24261 b2.arg(&*part_o)
24262 .arg(&*part_m)
24263 .arg(&*part_l)
24264 .arg(oq)
24265 .arg(od)
24266 .arg(&hd)
24267 .arg(&nh)
24268 .arg(&nsp);
24269 unsafe {
24270 b2.launch(cfg2)?;
24271 }
24272 return Ok(());
24273 }
24274 let fc = if g {
24275 self.func_g("fa_decode_combine_f32")
24276 } else {
24277 self.fa_func("fa_decode_combine_f32", head_dim)
24278 };
24279 let __s_b2 = self.gpu.stream();
24280 let mut b2 = __s_b2.launch_builder(&fc);
24281 b2.arg(&*part_o)
24282 .arg(&*part_m)
24283 .arg(&*part_l)
24284 .arg(o)
24285 .arg(&hd)
24286 .arg(&nh)
24287 .arg(&nsp);
24288 unsafe {
24289 b2.launch(cfg2)?;
24290 }
24291 Ok(())
24292 }
24293
24294 #[allow(clippy::too_many_arguments)]
24298 pub fn append_kv_quantized_dcw(
24299 &self,
24300 k_row: &CudaSlice<f32>,
24301 v_row: &CudaSlice<f32>,
24302 kc: &mut CudaSlice<u8>,
24303 vc: &mut CudaSlice<u8>,
24304 len_dev: &CudaSlice<i32>,
24305 base_dev: Option<&CudaSlice<i32>>,
24306 kv_dim_k: usize,
24307 kv_dim_v: usize,
24308 k_tok_bytes: usize,
24309 v_tok_bytes: usize,
24310 ) -> Result<(), Box<dyn std::error::Error>> {
24311 let f = self.func("append_quantize_kv_q8_0_q5_1_dcw");
24312 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
24313 let cfg = LaunchConfig {
24314 grid_dim: (nblk, 1, 1),
24315 block_dim: (32, 1, 1),
24316 shared_mem_bytes: 0,
24317 };
24318 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
24319 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
24320 let null: u64 = 0;
24321 let __s_b = self.gpu.stream();
24322 let mut b = __s_b.launch_builder(&f);
24323 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(len_dev);
24324 match base_dev {
24325 Some(base) => {
24326 b.arg(base);
24327 }
24328 None => {
24329 b.arg(&null);
24330 }
24331 }
24332 b.arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
24333 unsafe {
24334 b.launch(cfg)?;
24335 }
24336 Ok(())
24337 }
24338
24339 pub fn inc_i32(&self, counter: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
24341 let f = self.func("inc_i32");
24342 let cfg = LaunchConfig {
24343 grid_dim: (1, 1, 1),
24344 block_dim: (1, 1, 1),
24345 shared_mem_bytes: 0,
24346 };
24347 let __s_b = self.gpu.stream();
24348 let mut b = __s_b.launch_builder(&f);
24349 b.arg(counter);
24350 unsafe {
24351 b.launch(cfg)?;
24352 }
24353 Ok(())
24354 }
24355
24356 #[allow(clippy::too_many_arguments)]
24365 fn fa_part_alloc(
24389 &self,
24390 o_len: usize,
24391 ml_len: usize,
24392 co: usize,
24393 cm: usize,
24394 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
24395 static GROWS: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
24396 let n = GROWS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
24397 if n < 64 {
24398 eprintln!(
24399 "[fa-pool] grow #{n} dev={} o_len {co} -> {o_len} ml_len {cm} -> {ml_len} (retired kept, zero={})",
24400 self.ctx().ordinal(),
24401 fa_part_zero_on()
24402 );
24403 }
24404 let mut po = self.alloc_uninit::<f32>(o_len)?;
24405 let mut pm = self.alloc_uninit::<f32>(ml_len)?;
24406 let mut pl = self.alloc_uninit::<f32>(ml_len)?;
24407 if fa_part_zero_on() {
24408 self.gpu.stream().memset_zeros(&mut po)?;
24409 self.gpu.stream().memset_zeros(&mut pm)?;
24410 self.gpu.stream().memset_zeros(&mut pl)?;
24411 }
24412 Ok((po, pm, pl))
24413 }
24414
24415 fn fa_part_pool_grow(
24416 &self,
24417 part_guard: &mut Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>,
24418 o_len: usize,
24419 ml_len: usize,
24420 ) -> Result<(), Box<dyn std::error::Error>> {
24421 if part_guard
24422 .as_ref()
24423 .map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
24424 .unwrap_or(true)
24425 {
24426 let old = part_guard.take();
24427 let (co, cm) = old
24428 .as_ref()
24429 .map(|pp| (pp.0.len(), pp.1.len()))
24430 .unwrap_or((0, 0));
24431 if let Some(old) = old {
24432 self.fa_part_retired.lock().unwrap().push(old);
24433 }
24434 *part_guard =
24447 Some(self.fa_part_alloc(o_len.max(2 * co), ml_len.max(2 * cm), co, cm)?);
24448 }
24449 Ok(())
24450 }
24451
24452 pub fn fa_dcw_pool_ensure(
24455 &self,
24456 head_dim: usize,
24457 n_head: usize,
24458 n_head_kv: usize,
24459 bucket_max: usize,
24460 ) -> Result<(), Box<dyn std::error::Error>> {
24461 let sp = fa_split_keys(bucket_max, n_head_kv);
24462 let n_splits = ((bucket_max + sp - 1) / sp).max(1);
24463 let o_len = n_head * n_splits * head_dim;
24464 let ml_len = n_head * n_splits;
24465 let mut part_guard = self.fa_part_pool.lock().unwrap();
24466 self.fa_part_pool_grow(&mut part_guard, o_len, ml_len)
24467 }
24468
24469 #[allow(clippy::too_many_arguments)]
24477 pub fn fa_decode_dcw2(
24478 &self,
24479 q2: &CudaSlice<f32>,
24480 k_ring: &cudarc::driver::CudaView<u8>,
24481 v_ring: &cudarc::driver::CudaView<u8>,
24482 o2: &mut CudaSlice<f32>,
24483 head_dim: usize,
24484 n_head: usize,
24485 n_head_kv: usize,
24486 len_dev: &CudaSlice<i32>,
24487 base_dev: Option<&CudaSlice<i32>>,
24488 window: usize,
24489 bucket_max: usize,
24490 scale: f32,
24491 k_tok_bytes: usize,
24492 v_tok_bytes: usize,
24493 gate2: &CudaSlice<f32>,
24494 ) -> Result<(), Box<dyn std::error::Error>> {
24495 let fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
24496 if !fa_vec || head_dim > 256 || head_dim % 32 != 0 || !fa_v3_on() {
24497 return Err("fa_decode_dcw2 supports the default v3-vec class only".into());
24498 }
24499 let sp = fa_split_keys(bucket_max, n_head_kv);
24500 let n_splits = ((bucket_max + sp - 1) / sp).max(1);
24501 let o_len = 2 * n_head * n_splits * head_dim;
24503 let ml_len = 2 * n_head * n_splits;
24504 let mut part_guard = self.fa_part_pool.lock().unwrap();
24505 Self::fa_part_pool_grow(self, &mut part_guard, o_len, ml_len)?;
24506 let pg = part_guard.as_mut().unwrap();
24507 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
24508 let (hd, nh, nhkv, nsp) = (
24509 head_dim as i32,
24510 n_head as i32,
24511 n_head_kv as i32,
24512 n_splits as i32,
24513 );
24514 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
24515 let (ski, win) = (sp as i32, window as i32);
24516 let gqa = (n_head / n_head_kv).max(1) as u32;
24517 let smem = (32 * head_dim * 2) as u32;
24518 let f = self.func("fa_decode_vec_q_v3_dcw2");
24519 let cfg = LaunchConfig {
24520 grid_dim: (n_head_kv as u32, n_splits as u32, 1),
24521 block_dim: (32, gqa, 1),
24522 shared_mem_bytes: smem,
24523 };
24524 let null: u64 = 0;
24525 {
24526 let __s_b = self.gpu.stream();
24527 let mut b = __s_b.launch_builder(&f);
24528 b.arg(q2)
24529 .arg(k_ring)
24530 .arg(v_ring)
24531 .arg(&mut *part_o)
24532 .arg(&mut *part_m)
24533 .arg(&mut *part_l)
24534 .arg(&hd)
24535 .arg(&nh)
24536 .arg(&nhkv)
24537 .arg(len_dev);
24538 match base_dev {
24539 Some(base) => {
24540 b.arg(base);
24541 }
24542 None => {
24543 b.arg(&null);
24544 }
24545 }
24546 b.arg(&win)
24547 .arg(&scale)
24548 .arg(&nsp)
24549 .arg(&ski)
24550 .arg(&ktb)
24551 .arg(&vtb);
24552 unsafe {
24553 b.launch(cfg)?;
24554 }
24555 }
24556 let fc = {
24560 static CS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24561 if *CS.get_or_init(|| std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1")) {
24562 self.func("fa_decode_combine_gate_f32_s")
24563 } else {
24564 self.func("fa_decode_combine_gate_f32")
24565 }
24566 };
24567 let combine_shared = std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1");
24568 let nh2 = (2 * n_head) as i32;
24569 let cfg2 = LaunchConfig {
24570 grid_dim: ((2 * n_head) as u32, 1, 1),
24571 block_dim: (head_dim as u32, 1, 1),
24572 shared_mem_bytes: if combine_shared {
24573 (2 * n_splits * 4) as u32
24574 } else {
24575 0
24576 },
24577 };
24578 let __s_b2 = self.gpu.stream();
24579 let mut b2 = __s_b2.launch_builder(&fc);
24580 b2.arg(&*part_o)
24581 .arg(&*part_m)
24582 .arg(&*part_l)
24583 .arg(gate2)
24584 .arg(o2)
24585 .arg(&hd)
24586 .arg(&nh2)
24587 .arg(&nsp);
24588 unsafe {
24589 b2.launch(cfg2)?;
24590 }
24591 Ok(())
24592 }
24593
24594 #[allow(clippy::too_many_arguments)]
24603 pub fn fa_decode_dcw_rows(
24604 &self,
24605 q_rows: &CudaSlice<f32>,
24606 tab: &CudaSlice<u64>,
24607 o_rows: &mut CudaSlice<f32>,
24608 t: usize,
24609 head_dim: usize,
24610 n_head: usize,
24611 n_head_kv: usize,
24612 window: usize,
24613 max_ns: usize,
24614 scale: f32,
24615 k_tok_bytes: usize,
24616 v_tok_bytes: usize,
24617 gate_rows: &CudaSlice<f32>,
24618 ) -> Result<(), Box<dyn std::error::Error>> {
24619 if std::env::var("MEMRA_NO_FA_VEC").is_ok()
24620 || head_dim > 256
24621 || head_dim % 32 != 0
24622 || !fa_v3_on()
24623 {
24624 return Err("fa_decode_dcw_rows supports the default v3-vec class only".into());
24625 }
24626 if fa_sm_count() < 128
24627 || std::env::var("MEMRA_FA_SPLIT").is_ok()
24628 || std::env::var("MEMRA_FA_SP_SHORT").is_ok()
24629 || std::env::var("MEMRA_FA_SP16").is_ok()
24630 {
24631 return Err(
24632 "fa_decode_dcw_rows embeds the big-rig split ladder; env split overrides \
24633 (or a <128-SM rig) keep the per-row path"
24634 .into(),
24635 );
24636 }
24637 if t == 0 || t > 32 || max_ns == 0 || tab.len() < t * 6 {
24638 return Err("fa_decode_dcw_rows geometry".into());
24639 }
24640 let o_len = t * n_head * max_ns * head_dim;
24641 let ml_len = t * n_head * max_ns;
24642 let mut part_guard = self.fa_part_pool.lock().unwrap();
24643 Self::fa_part_pool_grow(self, &mut part_guard, o_len, ml_len)?;
24644 let pg = part_guard.as_mut().unwrap();
24645 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
24646 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
24647 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
24648 let (win, mns) = (window as i32, max_ns as i32);
24649 let gqa = (n_head / n_head_kv).max(1) as u32;
24650 let smem = (32 * head_dim * 2) as u32;
24651 let f = self.func("fa_decode_vec_q_v3_dcw_rows");
24652 let cfg = LaunchConfig {
24653 grid_dim: (n_head_kv as u32, max_ns as u32, t as u32),
24654 block_dim: (32, gqa, 1),
24655 shared_mem_bytes: smem,
24656 };
24657 {
24658 let __s_b = self.gpu.stream();
24659 let mut b = __s_b.launch_builder(&f);
24660 b.arg(q_rows)
24661 .arg(tab)
24662 .arg(&mut *part_o)
24663 .arg(&mut *part_m)
24664 .arg(&mut *part_l)
24665 .arg(&hd)
24666 .arg(&nh)
24667 .arg(&nhkv)
24668 .arg(&win)
24669 .arg(&scale)
24670 .arg(&mns)
24671 .arg(&ktb)
24672 .arg(&vtb);
24673 unsafe {
24674 b.launch(cfg)?;
24675 }
24676 }
24677 let fc = {
24681 static CS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24682 if *CS.get_or_init(|| std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1")) {
24683 self.func("fa_decode_combine_gate_f32_s")
24684 } else {
24685 self.func("fa_decode_combine_gate_f32")
24686 }
24687 };
24688 let combine_shared = std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1");
24689 let nht = (t * n_head) as i32;
24690 let cfg2 = LaunchConfig {
24691 grid_dim: ((t * n_head) as u32, 1, 1),
24692 block_dim: (head_dim as u32, 1, 1),
24693 shared_mem_bytes: if combine_shared {
24694 (2 * max_ns * 4) as u32
24695 } else {
24696 0
24697 },
24698 };
24699 let __s_b2 = self.gpu.stream();
24700 let mut b2 = __s_b2.launch_builder(&fc);
24701 b2.arg(&*part_o)
24702 .arg(&*part_m)
24703 .arg(&*part_l)
24704 .arg(gate_rows)
24705 .arg(o_rows)
24706 .arg(&hd)
24707 .arg(&nht)
24708 .arg(&mns);
24709 unsafe {
24710 b2.launch(cfg2)?;
24711 }
24712 Ok(())
24713 }
24714
24715 pub fn fa_decode_dcw(
24716 &self,
24717 q: &CudaSlice<f32>,
24718 k_ring: &cudarc::driver::CudaView<u8>,
24719 v_ring: &cudarc::driver::CudaView<u8>,
24720 o: &mut CudaSlice<f32>,
24721 head_dim: usize,
24722 n_head: usize,
24723 n_head_kv: usize,
24724 len_dev: &CudaSlice<i32>,
24725 base_dev: Option<&CudaSlice<i32>>,
24726 window: usize,
24727 bucket_max: usize,
24728 scale: f32,
24729 k_tok_bytes: usize,
24730 v_tok_bytes: usize,
24731 fused_gate: Option<&CudaSlice<f32>>,
24735 ) -> Result<(), Box<dyn std::error::Error>> {
24736 let fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
24737 if !fa_vec || head_dim > 256 || head_dim % 32 != 0 || !fa_v3_on() {
24738 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"
24739 .into());
24740 }
24741 let sp = fa_split_keys(bucket_max, n_head_kv);
24742 let n_splits = ((bucket_max + sp - 1) / sp).max(1);
24743 let o_len = n_head * n_splits * head_dim;
24744 let ml_len = n_head * n_splits;
24745 let mut part_guard = self.fa_part_pool.lock().unwrap();
24746 Self::fa_part_pool_grow(self, &mut part_guard, o_len, ml_len)?;
24747 let pg = part_guard.as_mut().unwrap();
24748 static MEMSET_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24753 let memset_on = *MEMSET_ON
24758 .get_or_init(|| std::env::var("MEMRA_FA_DCW_MEMSET").as_deref() != Ok("0"))
24759 || crate::tp::token_graph_building();
24760 if memset_on {
24761 self.gpu
24762 .stream()
24763 .memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
24764 self.gpu
24765 .stream()
24766 .memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
24767 self.gpu
24768 .stream()
24769 .memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
24770 }
24771 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
24772 let (hd, nh, nhkv, nsp) = (
24773 head_dim as i32,
24774 n_head as i32,
24775 n_head_kv as i32,
24776 n_splits as i32,
24777 );
24778 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
24779 let (ski, win) = (sp as i32, window as i32);
24780 let gqa = (n_head / n_head_kv).max(1) as u32;
24781 let smem = (32 * head_dim * 2) as u32; static U8: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24785 static HOIST: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
24786 let hoist = *HOIST.get_or_init(|| match std::env::var("MEMRA_FA_HOIST").as_deref() {
24787 Ok("2") => 2,
24788 Ok("1") => 1,
24789 _ => 0,
24790 });
24791 static FPROF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24796 let fprof = *FPROF.get_or_init(|| std::env::var("MEMRA_FA_PROF").as_deref() == Ok("1"));
24797 static PROF_BUF: std::sync::Mutex<Option<(usize, CudaSlice<u64>)>> =
24798 std::sync::Mutex::new(None);
24799 static HS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24803 let hs2 = *HS.get_or_init(|| std::env::var("MEMRA_FA_HSPLIT").as_deref() == Ok("2"))
24804 && (n_head / n_head_kv) % 2 == 0
24805 && (n_head / n_head_kv) >= 2;
24806 let f = if fprof {
24807 self.func("fa_decode_vec_q_v3_dcw_prof")
24808 } else if hs2 {
24809 self.func("fa_decode_vec_q_v3_dcw_hs2")
24810 } else if hoist == 2 {
24811 self.func("fa_decode_vec_q_v3_dcw_hc")
24813 } else if hoist == 1 {
24814 self.func("fa_decode_vec_q_v3_dcw_h")
24816 } else if *U8.get_or_init(|| std::env::var("MEMRA_FA_UNROLL").as_deref() == Ok("8")) {
24817 self.func("fa_decode_vec_q_v3_dcw_u8")
24818 } else {
24819 self.func("fa_decode_vec_q_v3_dcw")
24820 };
24821 let cfg = LaunchConfig {
24822 grid_dim: if hs2 {
24823 ((2 * n_head_kv) as u32, n_splits as u32, 1)
24824 } else {
24825 (n_head_kv as u32, n_splits as u32, 1)
24826 },
24827 block_dim: if hs2 { (32, gqa / 2, 1) } else { (32, gqa, 1) },
24828 shared_mem_bytes: smem,
24829 };
24830 let null: u64 = 0;
24831 let __s_b = self.gpu.stream();
24832 let mut b = __s_b.launch_builder(&f);
24833 b.arg(q)
24834 .arg(k_ring)
24835 .arg(v_ring)
24836 .arg(&mut *part_o)
24837 .arg(&mut *part_m)
24838 .arg(&mut *part_l)
24839 .arg(&hd)
24840 .arg(&nh)
24841 .arg(&nhkv)
24842 .arg(len_dev);
24843 match base_dev {
24844 Some(base) => {
24845 b.arg(base);
24846 }
24847 None => {
24848 b.arg(&null);
24849 }
24850 }
24851 b.arg(&win)
24852 .arg(&scale)
24853 .arg(&nsp)
24854 .arg(&ski)
24855 .arg(&ktb)
24856 .arg(&vtb);
24857 if fprof {
24858 let mut guard = PROF_BUF.lock().map_err(|_| "fa prof buffer lock")?;
24859 if guard
24860 .as_ref()
24861 .is_none_or(|(d, _)| *d != self.ctx().ordinal())
24862 {
24863 *guard = Some((self.ctx().ordinal(), self.htod_u64(&vec![0u64; 8])?));
24864 }
24865 let (_, buf) = guard.as_mut().expect("armed above");
24866 b.arg(&*buf);
24867 unsafe {
24868 b.launch(cfg)?;
24869 }
24870 static CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
24871 let n = CALLS.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1;
24872 if n % 430 == 0 {
24873 self.stream().synchronize()?;
24874 let h = self.dtoh_u64(buf)?;
24875 let phases = ["setup", "stageV", "b1_klo", "b2_soft", "sync", "b3_vacc"];
24876 let tot: u64 = h[..6].iter().sum();
24877 let mut line = format!("[fa-prof] calls={n} keys={} cycles={tot}", h[6]);
24878 for (i, name) in phases.iter().enumerate() {
24879 let pct = if tot > 0 {
24880 h[i] as f64 / tot as f64 * 100.0
24881 } else {
24882 0.0
24883 };
24884 line.push_str(&format!(" {name}={pct:.1}%"));
24885 }
24886 if h[6] > 0 {
24887 line.push_str(&format!(" cyc/key={:.0}", tot as f64 / h[6] as f64));
24888 }
24889 eprintln!("{line}");
24890 }
24891 } else {
24892 unsafe {
24893 b.launch(cfg)?;
24894 }
24895 }
24896 let mut combine_shared = false;
24897 let fc = if fused_gate.is_some() {
24898 static CS: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24901 if *CS.get_or_init(|| std::env::var("MEMRA_FA_COMBINE_S").as_deref() == Ok("1")) {
24902 combine_shared = true;
24903 self.func("fa_decode_combine_gate_f32_s")
24904 } else {
24905 self.func("fa_decode_combine_gate_f32")
24906 }
24907 } else {
24908 self.fa_func("fa_decode_combine_f32", head_dim)
24909 };
24910 let cfg2 = LaunchConfig {
24911 grid_dim: (n_head as u32, 1, 1),
24912 block_dim: (head_dim as u32, 1, 1),
24913 shared_mem_bytes: if combine_shared {
24914 (2 * n_splits * 4) as u32
24915 } else {
24916 0
24917 },
24918 };
24919 let __s_b2 = self.gpu.stream();
24920 let mut b2 = __s_b2.launch_builder(&fc);
24921 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l);
24922 if let Some(gate_row) = fused_gate {
24923 b2.arg(gate_row);
24924 }
24925 b2.arg(o).arg(&hd).arg(&nh).arg(&nsp);
24926 unsafe {
24927 b2.launch(cfg2)?;
24928 }
24929 Ok(())
24930 }
24931
24932 pub fn fa_geom_eager(
24938 &self,
24939 t_kv: usize,
24940 head_dim: usize,
24941 n_head_kv: usize,
24942 g: bool,
24943 ) -> (bool, usize) {
24944 let fa_ok = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
24948 let vec512 = fa_ok && head_dim == 512 && t_kv >= fa512_min_tkv();
24954 let mut fa_vec = vec512 || (fa_ok && head_dim <= 256 && head_dim % 32 == 0);
24955 if g && head_dim == 256 && !fa_v4_at(t_kv) {
24961 fa_vec = false;
24962 }
24963 let sp = fa_split_keys(t_kv, n_head_kv);
24964 let n_splits = if fa_vec {
24965 ((t_kv + sp - 1) / sp).max(1)
24966 } else {
24967 ((t_kv + 255) / 256).max(1)
24968 };
24969 (fa_vec, n_splits)
24970 }
24971
24972 pub fn fa_bucket_key(
24978 &self,
24979 t_kv: usize,
24980 head_dim: usize,
24981 n_head_kv: usize,
24982 g: bool,
24983 ) -> (bool, usize) {
24984 self.fa_geom_eager(t_kv, head_dim, n_head_kv, g)
24985 }
24986
24987 pub fn capture_graph_retained<F>(
24999 &self,
25000 step: F,
25001 ) -> Result<
25002 (
25003 cudarc::driver::CudaGraph,
25004 Vec<Box<dyn std::any::Any + Send>>,
25005 ),
25006 Box<dyn std::error::Error>,
25007 >
25008 where
25009 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
25010 {
25011 use cudarc::driver::sys::CUgraphInstantiate_flags;
25012 self.capture_graph_retained_flags(
25013 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
25014 step,
25015 )
25016 }
25017
25018 pub fn capture_graph_retained_flags<F>(
25023 &self,
25024 flags: cudarc::driver::sys::CUgraphInstantiate_flags,
25025 mut step: F,
25026 ) -> Result<
25027 (
25028 cudarc::driver::CudaGraph,
25029 Vec<Box<dyn std::any::Any + Send>>,
25030 ),
25031 Box<dyn std::error::Error>,
25032 >
25033 where
25034 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
25035 {
25036 use cudarc::driver::sys::CUstreamCaptureMode;
25037 self.capture_keep.lock().unwrap().clear();
25045 let was_tracking = self.gpu.ctx.is_event_tracking();
25046 if was_tracking {
25047 unsafe {
25048 self.gpu.ctx.disable_event_tracking();
25049 }
25050 }
25051 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
25052 self.capture_keep_on
25053 .store(true, std::sync::atomic::Ordering::Relaxed);
25054 let w = (|| {
25055 step(self)?;
25056 step(self)
25057 })();
25058 self.capture_keep_on
25059 .store(false, std::sync::atomic::Ordering::Relaxed);
25060 w?;
25061 self.gpu.stream().synchronize()?;
25062 self.gpu
25063 .stream()
25064 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
25065 let r = step(self);
25066 let g = self.gpu.stream().end_capture(flags);
25067 r?;
25068 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
25069 graph.upload()?;
25070 Ok(graph)
25071 };
25072 let result = run();
25073 self.capture_keep_on
25074 .store(false, std::sync::atomic::Ordering::Relaxed);
25075 if was_tracking {
25076 unsafe {
25077 self.gpu.ctx.enable_event_tracking();
25078 }
25079 }
25080 let keeper = std::mem::take(&mut *self.capture_keep.lock().unwrap());
25081 Ok((result?, keeper))
25082 }
25083
25084 pub fn capture_graph_retained_nowarm<F>(
25090 &self,
25091 mut step: F,
25092 ) -> Result<
25093 (
25094 cudarc::driver::CudaGraph,
25095 Vec<Box<dyn std::any::Any + Send>>,
25096 ),
25097 Box<dyn std::error::Error>,
25098 >
25099 where
25100 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
25101 {
25102 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
25103 let was_tracking = self.gpu.ctx.is_event_tracking();
25104 if was_tracking {
25105 unsafe {
25106 self.gpu.ctx.disable_event_tracking();
25107 }
25108 }
25109 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
25110 self.gpu.stream().synchronize()?;
25111 self.gpu
25112 .stream()
25113 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
25114 let r = step(self);
25115 let g = self.gpu.stream().end_capture(
25116 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
25117 );
25118 r?;
25119 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
25120 graph.upload()?;
25121 Ok(graph)
25122 };
25123 let result = run();
25124 if was_tracking {
25125 unsafe {
25126 self.gpu.ctx.enable_event_tracking();
25127 }
25128 }
25129 Ok((result?, Vec::new()))
25130 }
25131
25132 pub fn capture_graph<F>(
25133 &self,
25134 mut step: F,
25135 ) -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>>
25136 where
25137 F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
25138 {
25139 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
25140 let was_tracking = self.gpu.ctx.is_event_tracking();
25148 if was_tracking {
25149 unsafe {
25150 self.gpu.ctx.disable_event_tracking();
25151 }
25152 }
25153 let iflag = {
25160 static F: std::sync::OnceLock<CUgraphInstantiate_flags> = std::sync::OnceLock::new();
25161 *F.get_or_init(|| match std::env::var("MEMRA_GRAPH_IFLAG").as_deref() {
25162 Ok("upload") => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_UPLOAD,
25165 Ok("priority") => {
25166 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY
25167 }
25168 _ => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
25169 })
25170 };
25171 let ct = {
25178 static T: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
25179 *T.get_or_init(|| std::env::var("MEMRA_GRAPH_CAPTIME").as_deref() == Ok("1"))
25180 };
25181 let warmups = {
25204 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
25205 *W.get_or_init(|| {
25206 std::env::var("MEMRA_GRAPH_WARMUPS")
25207 .ok()
25208 .and_then(|v| v.parse().ok())
25209 .filter(|n| *n >= 1)
25210 .unwrap_or(1)
25211 })
25212 };
25213 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
25214 let t_w = std::time::Instant::now();
25215 for _ in 0..warmups {
25217 step(self)?;
25218 }
25219 self.gpu.stream().synchronize()?;
25220 let ms_warm = t_w.elapsed().as_secs_f64() * 1e3;
25221 let t_c = std::time::Instant::now();
25223 self.gpu
25224 .stream()
25225 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
25226 let r = step(self);
25229 let ms_body = t_c.elapsed().as_secs_f64() * 1e3;
25230 let t_i = std::time::Instant::now();
25231 let g = self.gpu.stream().end_capture(iflag);
25232 let ms_inst = t_i.elapsed().as_secs_f64() * 1e3;
25233 r?;
25234 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
25235 let t_u = std::time::Instant::now();
25236 graph.upload()?;
25237 if ct {
25238 println!(
25239 "[graph-captime] warmup2x {ms_warm:.2} ms capture-body {ms_body:.2} ms \
25240 instantiate {ms_inst:.2} ms upload {:.2} ms",
25241 t_u.elapsed().as_secs_f64() * 1e3
25242 );
25243 }
25244 Ok(graph)
25245 };
25246 let result = run();
25247 if was_tracking {
25248 unsafe {
25249 self.gpu.ctx.enable_event_tracking();
25250 }
25251 }
25252 result
25253 }
25254
25255 pub fn gdn_scan_s128_view(
25257 &self,
25258 q: &CudaSlice<f32>,
25259 k: &CudaSlice<f32>,
25260 v: &CudaSlice<f32>,
25261 g: &CudaSlice<f32>,
25262 beta: &CudaSlice<f32>,
25263 state_in: &cudarc::driver::CudaView<f32>,
25264 state_out: &mut cudarc::driver::CudaViewMut<f32>,
25265 o: &mut CudaSlice<f32>,
25266 n_head: usize,
25267 t: usize,
25268 scale: f32,
25269 ) -> Result<(), Box<dyn std::error::Error>> {
25270 let f = self.func("gdn_scan_s128");
25271 const S_V: u32 = 128;
25272 const WARP: u32 = 32;
25273 const COLS: u32 = 4;
25274 let cfg = LaunchConfig {
25275 grid_dim: (n_head as u32, 1, S_V / COLS),
25276 block_dim: (WARP, COLS, 1),
25277 shared_mem_bytes: 0,
25278 };
25279 let (h, ti) = (n_head as i32, t as i32);
25280 let __s_b = self.gpu.stream();
25281 let mut b = __s_b.launch_builder(&f);
25282 b.arg(q)
25283 .arg(k)
25284 .arg(v)
25285 .arg(g)
25286 .arg(beta)
25287 .arg(state_in)
25288 .arg(state_out)
25289 .arg(o)
25290 .arg(&h)
25291 .arg(&ti)
25292 .arg(&scale);
25293 unsafe {
25294 b.launch(cfg)?;
25295 }
25296 Ok(())
25297 }
25298
25299 pub fn ssm_conv1d_view(
25301 &self,
25302 x: &cudarc::driver::CudaView<f32>,
25303 w: &CudaSlice<f32>,
25304 y: &mut CudaSlice<f32>,
25305 conv_dim: usize,
25306 t: usize,
25307 d_conv: usize,
25308 silu: bool,
25309 ) -> Result<(), Box<dyn std::error::Error>> {
25310 let f = self.func("ssm_conv1d_silu_f32");
25311 let cfg = LaunchConfig {
25313 grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
25314 block_dim: (256, 1, 1),
25315 shared_mem_bytes: 0,
25316 };
25317 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
25318 let __s_b = self.gpu.stream();
25319 let mut b = __s_b.launch_builder(&f);
25320 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
25321 unsafe {
25322 b.launch(cfg)?;
25323 }
25324 Ok(())
25325 }
25326
25327 pub fn ssm_conv1d_tm(
25334 &self,
25335 qkv_tm: &CudaSlice<f32>,
25336 w: &CudaSlice<f32>,
25337 y: &mut CudaSlice<f32>,
25338 conv_dim: usize,
25339 t: usize,
25340 d_conv: usize,
25341 ) -> Result<(), Box<dyn std::error::Error>> {
25342 let f = self.func("ssm_conv1d_tm_f32");
25343 let cfg = LaunchConfig {
25344 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
25345 block_dim: (256, 1, 1),
25346 shared_mem_bytes: 0,
25347 };
25348 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
25349 let __s_b = self.gpu.stream();
25350 let mut b = __s_b.launch_builder(&f);
25351 b.arg(qkv_tm).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
25352 unsafe {
25353 b.launch(cfg)?;
25354 }
25355 Ok(())
25356 }
25357
25358 pub fn ssm_conv1d_tm_state(
25366 &self,
25367 qkv_tm: &CudaSlice<f32>,
25368 conv_state: &mut CudaSlice<f32>,
25369 w: &CudaSlice<f32>,
25370 y: &mut CudaSlice<f32>,
25371 conv_dim: usize,
25372 t: usize,
25373 d_conv: usize,
25374 ) -> Result<(), Box<dyn std::error::Error>> {
25375 self.ssm_conv1d_tm_state_pad(qkv_tm, conv_state, w, y, conv_dim, t, d_conv, None)
25376 }
25377
25378 #[allow(clippy::too_many_arguments)]
25381 pub fn ssm_conv1d_tm_state_pad(
25382 &self,
25383 qkv_tm: &CudaSlice<f32>,
25384 conv_state: &mut CudaSlice<f32>,
25385 w: &CudaSlice<f32>,
25386 y: &mut CudaSlice<f32>,
25387 conv_dim: usize,
25388 t: usize,
25389 d_conv: usize,
25390 pad_len: Option<&CudaSlice<i32>>,
25391 ) -> Result<(), Box<dyn std::error::Error>> {
25392 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
25393 let ring_old = if t < d_conv - 1 {
25397 Some(self.clone_dtod(conv_state)?)
25398 } else {
25399 None
25400 };
25401 {
25402 let f = self.func("ssm_conv1d_tm_state_f32");
25403 let cfg = LaunchConfig {
25404 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
25405 block_dim: (256, 1, 1),
25406 shared_mem_bytes: 0,
25407 };
25408 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
25409 let __s_b = self.gpu.stream();
25410 let mut b = __s_b.launch_builder(&f);
25411 b.arg(qkv_tm)
25412 .arg(&*conv_state)
25413 .arg(w)
25414 .arg(y)
25415 .arg(&cd)
25416 .arg(&ti)
25417 .arg(&dc);
25418 unsafe {
25419 b.launch(cfg)?;
25420 }
25421 }
25422 match (ring_old, pad_len) {
25423 (None, Some(len_d)) => {
25424 let f = self.func("ssm_conv_ring_update_dev_f32");
25425 let n = conv_dim * (d_conv - 1);
25426 let cfg = LaunchConfig::for_num_elems(n as u32);
25427 let (cd, dc) = (conv_dim as i32, d_conv as i32);
25428 let __s_b = self.gpu.stream();
25429 let mut b = __s_b.launch_builder(&f);
25430 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
25431 unsafe {
25432 b.launch(cfg)?;
25433 }
25434 }
25435 (None, None) => {
25436 let f = self.func("ssm_conv_ring_update_f32");
25437 let n = conv_dim * (d_conv - 1);
25438 let cfg = LaunchConfig::for_num_elems(n as u32);
25439 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
25440 let __s_b = self.gpu.stream();
25441 let mut b = __s_b.launch_builder(&f);
25442 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
25443 unsafe {
25444 b.launch(cfg)?;
25445 }
25446 }
25447 (Some(old), _) => {
25448 self.ssm_conv_ring_rebuild(qkv_tm, &old, conv_state, conv_dim, t, d_conv)?
25449 }
25450 }
25451 Ok(())
25452 }
25453
25454 pub fn ssm_conv1d_tm_state_pad_v(
25456 &self,
25457 qkv_tm: &cudarc::driver::CudaView<f32>,
25458 conv_state: &mut CudaSlice<f32>,
25459 w: &CudaSlice<f32>,
25460 y: &mut CudaSlice<f32>,
25461 conv_dim: usize,
25462 t: usize,
25463 d_conv: usize,
25464 pad_len: Option<&CudaSlice<i32>>,
25465 ) -> Result<(), Box<dyn std::error::Error>> {
25466 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
25467 let ring_old = if t < d_conv - 1 {
25471 Some(self.clone_dtod(conv_state)?)
25472 } else {
25473 None
25474 };
25475 {
25476 let f = self.func("ssm_conv1d_tm_state_f32");
25477 let cfg = LaunchConfig {
25478 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
25479 block_dim: (256, 1, 1),
25480 shared_mem_bytes: 0,
25481 };
25482 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
25483 let __s_b = self.gpu.stream();
25484 let mut b = __s_b.launch_builder(&f);
25485 b.arg(qkv_tm)
25486 .arg(&*conv_state)
25487 .arg(w)
25488 .arg(y)
25489 .arg(&cd)
25490 .arg(&ti)
25491 .arg(&dc);
25492 unsafe {
25493 b.launch(cfg)?;
25494 }
25495 }
25496 match (ring_old, pad_len) {
25497 (None, Some(len_d)) => {
25498 let f = self.func("ssm_conv_ring_update_dev_f32");
25499 let n = conv_dim * (d_conv - 1);
25500 let cfg = LaunchConfig::for_num_elems(n as u32);
25501 let (cd, dc) = (conv_dim as i32, d_conv as i32);
25502 let __s_b = self.gpu.stream();
25503 let mut b = __s_b.launch_builder(&f);
25504 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
25505 unsafe {
25506 b.launch(cfg)?;
25507 }
25508 }
25509 (None, None) => {
25510 let f = self.func("ssm_conv_ring_update_f32");
25511 let n = conv_dim * (d_conv - 1);
25512 let cfg = LaunchConfig::for_num_elems(n as u32);
25513 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
25514 let __s_b = self.gpu.stream();
25515 let mut b = __s_b.launch_builder(&f);
25516 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
25517 unsafe {
25518 b.launch(cfg)?;
25519 }
25520 }
25521 (Some(_), _) => unreachable!(
25522 "ssm_conv1d_tm_state_pad_v: T < d_conv-1 has no view path (PRIME_MIN_T gates it)"
25523 ),
25524 }
25525 Ok(())
25526 }
25527
25528 pub fn ssm_conv_ring_rebuild(
25533 &self,
25534 qkv_tm: &CudaSlice<f32>,
25535 ring_old: &CudaSlice<f32>,
25536 conv_state: &mut CudaSlice<f32>,
25537 conv_dim: usize,
25538 tc: usize,
25539 d_conv: usize,
25540 ) -> Result<(), Box<dyn std::error::Error>> {
25541 let f = self.func("ssm_conv_ring_rebuild_f32");
25542 let n = conv_dim * (d_conv - 1);
25543 let cfg = LaunchConfig::for_num_elems(n as u32);
25544 let (cd, ti, dc) = (conv_dim as i32, tc as i32, d_conv as i32);
25545 let __s_b = self.gpu.stream();
25546 let mut b = __s_b.launch_builder(&f);
25547 b.arg(qkv_tm)
25548 .arg(ring_old)
25549 .arg(conv_state)
25550 .arg(&cd)
25551 .arg(&ti)
25552 .arg(&dc);
25553 unsafe {
25554 b.launch(cfg)?;
25555 }
25556 Ok(())
25557 }
25558
25559 #[allow(clippy::too_many_arguments)]
25564 pub fn gdn_prep_decode(
25565 &self,
25566 conv_out: &CudaSlice<f32>,
25567 beta_raw: &CudaSlice<f32>,
25568 alpha: &CudaSlice<f32>,
25569 dt_bias: &CudaSlice<f32>,
25570 a: &CudaSlice<f32>,
25571 q_l2: &mut CudaSlice<f32>,
25572 k_l2: &mut CudaSlice<f32>,
25573 v_g: &mut CudaSlice<f32>,
25574 beta: &mut CudaSlice<f32>,
25575 g_log: &mut CudaSlice<f32>,
25576 d_state: usize,
25577 num_v: usize,
25578 num_k: usize,
25579 key_dim: usize,
25580 eps: f32,
25581 ) -> Result<(), Box<dyn std::error::Error>> {
25582 let f = self.func("gdn_prep_decode_f32");
25583 let cfg = LaunchConfig {
25584 grid_dim: (num_v as u32, 1, 1),
25585 block_dim: (32, 4, 1),
25586 shared_mem_bytes: 0,
25587 };
25588 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
25589 let __s_b = self.gpu.stream();
25590 let mut b = __s_b.launch_builder(&f);
25591 b.arg(conv_out)
25592 .arg(beta_raw)
25593 .arg(alpha)
25594 .arg(dt_bias)
25595 .arg(a)
25596 .arg(q_l2)
25597 .arg(k_l2)
25598 .arg(v_g)
25599 .arg(beta)
25600 .arg(g_log)
25601 .arg(&ds)
25602 .arg(&nv)
25603 .arg(&nk)
25604 .arg(&kd)
25605 .arg(&eps);
25606 unsafe {
25607 b.launch(cfg)?;
25608 }
25609 Ok(())
25610 }
25611
25612 #[allow(clippy::too_many_arguments)]
25616 pub fn ssm_conv1d_gdn(
25617 &self,
25618 qkv_tm: &CudaSlice<f32>,
25619 w: &CudaSlice<f32>,
25620 q_g: &mut CudaSlice<f32>,
25621 k_g: &mut CudaSlice<f32>,
25622 v_g: &mut CudaSlice<f32>,
25623 conv_dim: usize,
25624 t: usize,
25625 d_conv: usize,
25626 d_state: usize,
25627 num_v: usize,
25628 num_k: usize,
25629 key_dim: usize,
25630 ) -> Result<(), Box<dyn std::error::Error>> {
25631 let f = self.func("ssm_conv1d_gdn_f32");
25632 let cfg = LaunchConfig {
25633 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
25634 block_dim: (256, 1, 1),
25635 shared_mem_bytes: 0,
25636 };
25637 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
25638 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
25639 let __s_b = self.gpu.stream();
25640 let mut b = __s_b.launch_builder(&f);
25641 b.arg(qkv_tm)
25642 .arg(w)
25643 .arg(q_g)
25644 .arg(k_g)
25645 .arg(v_g)
25646 .arg(&cd)
25647 .arg(&ti)
25648 .arg(&dc)
25649 .arg(&ds)
25650 .arg(&nv)
25651 .arg(&nk)
25652 .arg(&kd);
25653 unsafe {
25654 b.launch(cfg)?;
25655 }
25656 Ok(())
25657 }
25658
25659 pub fn ssm_conv1d(
25660 &self,
25661 x: &CudaSlice<f32>,
25662 w: &CudaSlice<f32>,
25663 y: &mut CudaSlice<f32>,
25664 conv_dim: usize,
25665 t: usize,
25666 d_conv: usize,
25667 silu: bool,
25668 ) -> Result<(), Box<dyn std::error::Error>> {
25669 let f = self.func("ssm_conv1d_silu_f32");
25670 let cfg = LaunchConfig {
25671 grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
25672 block_dim: (256, 1, 1),
25673 shared_mem_bytes: 0,
25674 };
25675 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
25676 let __s_b = self.gpu.stream();
25677 let mut b = __s_b.launch_builder(&f);
25678 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
25679 unsafe {
25680 b.launch(cfg)?;
25681 }
25682 Ok(())
25683 }
25684
25685 pub fn gdn_scan_s128(
25688 &self,
25689 q: &CudaSlice<f32>,
25690 k: &CudaSlice<f32>,
25691 v: &CudaSlice<f32>,
25692 g: &CudaSlice<f32>,
25693 beta: &CudaSlice<f32>,
25694 state_in: &CudaSlice<f32>,
25695 state_out: &mut CudaSlice<f32>,
25696 o: &mut CudaSlice<f32>,
25697 n_head: usize,
25698 t: usize,
25699 scale: f32,
25700 ) -> Result<(), Box<dyn std::error::Error>> {
25701 let f = self.func("gdn_scan_s128");
25702 const S_V: u32 = 128;
25703 const WARP: u32 = 32;
25704 const COLS_PER_BLOCK: u32 = 4;
25705 let cfg = LaunchConfig {
25706 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
25707 block_dim: (WARP, COLS_PER_BLOCK, 1),
25708 shared_mem_bytes: 0,
25709 };
25710 let (h, ti) = (n_head as i32, t as i32);
25711 let __s_b = self.gpu.stream();
25712 let mut b = __s_b.launch_builder(&f);
25713 b.arg(q)
25714 .arg(k)
25715 .arg(v)
25716 .arg(g)
25717 .arg(beta)
25718 .arg(state_in)
25719 .arg(state_out)
25720 .arg(o)
25721 .arg(&h)
25722 .arg(&ti)
25723 .arg(&scale);
25724 unsafe {
25725 b.launch(cfg)?;
25726 }
25727 Ok(())
25728 }
25729
25730 #[allow(clippy::too_many_arguments)]
25735 pub fn ssm_conv1d_fused_decode_b(
25736 &self,
25737 qkv_cols: &CudaSlice<f32>,
25738 conv_state_ptrs: &cudarc::driver::CudaView<u64>,
25739 w: &CudaSlice<f32>,
25740 conv_outs: &mut CudaSlice<f32>,
25741 conv_dim: usize,
25742 d_conv: usize,
25743 b_n: usize,
25744 ) -> Result<(), Box<dyn std::error::Error>> {
25745 let f = self.func("ssm_conv1d_fused_decode_b_f32");
25746 let cfg = LaunchConfig {
25747 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
25748 block_dim: (256, 1, 1),
25749 shared_mem_bytes: 0,
25750 };
25751 let (cd, dc) = (conv_dim as i32, d_conv as i32);
25752 let __s_b = self.gpu.stream();
25753 let mut b = __s_b.launch_builder(&f);
25754 b.arg(qkv_cols)
25755 .arg(conv_state_ptrs)
25756 .arg(w)
25757 .arg(conv_outs)
25758 .arg(&cd)
25759 .arg(&dc);
25760 unsafe {
25761 b.launch(cfg)?;
25762 }
25763 Ok(())
25764 }
25765
25766 #[allow(clippy::too_many_arguments)]
25767 pub fn gdn_prep_decode_b(
25768 &self,
25769 conv_outs: &CudaSlice<f32>,
25770 beta_raws: &CudaSlice<f32>,
25771 alphas: &CudaSlice<f32>,
25772 dt_bias: &CudaSlice<f32>,
25773 a: &CudaSlice<f32>,
25774 q_l2: &mut CudaSlice<f32>,
25775 k_l2: &mut CudaSlice<f32>,
25776 v_g: &mut CudaSlice<f32>,
25777 beta: &mut CudaSlice<f32>,
25778 g_log: &mut CudaSlice<f32>,
25779 d_state: usize,
25780 num_v: usize,
25781 num_k: usize,
25782 key_dim: usize,
25783 eps: f32,
25784 conv_dim: usize,
25785 b_n: usize,
25786 ) -> Result<(), Box<dyn std::error::Error>> {
25787 let f = self.func("gdn_prep_decode_b_f32");
25788 let cfg = LaunchConfig {
25789 grid_dim: (num_v as u32, 1, b_n as u32),
25790 block_dim: (32, 4, 1),
25791 shared_mem_bytes: 0,
25792 };
25793 let (ds, nv, nk, kd, cd) = (
25794 d_state as i32,
25795 num_v as i32,
25796 num_k as i32,
25797 key_dim as i32,
25798 conv_dim as i32,
25799 );
25800 let __s_b = self.gpu.stream();
25801 let mut b = __s_b.launch_builder(&f);
25802 b.arg(conv_outs)
25803 .arg(beta_raws)
25804 .arg(alphas)
25805 .arg(dt_bias)
25806 .arg(a)
25807 .arg(q_l2)
25808 .arg(k_l2)
25809 .arg(v_g)
25810 .arg(beta)
25811 .arg(g_log)
25812 .arg(&ds)
25813 .arg(&nv)
25814 .arg(&nk)
25815 .arg(&kd)
25816 .arg(&eps)
25817 .arg(&cd);
25818 unsafe {
25819 b.launch(cfg)?;
25820 }
25821 Ok(())
25822 }
25823
25824 #[allow(clippy::too_many_arguments)]
25825 pub fn gdn_scan_s128_batched(
25826 &self,
25827 q: &CudaSlice<f32>,
25828 k: &CudaSlice<f32>,
25829 v: &CudaSlice<f32>,
25830 g: &CudaSlice<f32>,
25831 beta: &CudaSlice<f32>,
25832 state_in_ptrs: &cudarc::driver::CudaView<u64>,
25833 state_out_ptrs: &cudarc::driver::CudaView<u64>,
25834 o: &mut CudaSlice<f32>,
25835 n_head: usize,
25836 b_n: usize,
25837 scale: f32,
25838 ) -> Result<(), Box<dyn std::error::Error>> {
25839 let f = self.func("gdn_scan_s128_b");
25840 const S_V: u32 = 128;
25841 const WARP: u32 = 32;
25842 const COLS_PER_BLOCK: u32 = 4;
25843 let cfg = LaunchConfig {
25844 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
25845 block_dim: (WARP, COLS_PER_BLOCK, 1),
25846 shared_mem_bytes: 0,
25847 };
25848 let h = n_head as i32;
25849 let __s_b = self.gpu.stream();
25850 let mut b = __s_b.launch_builder(&f);
25851 b.arg(q)
25852 .arg(k)
25853 .arg(v)
25854 .arg(g)
25855 .arg(beta)
25856 .arg(state_in_ptrs)
25857 .arg(state_out_ptrs)
25858 .arg(o)
25859 .arg(&h)
25860 .arg(&scale);
25861 unsafe {
25862 b.launch(cfg)?;
25863 }
25864 Ok(())
25865 }
25866
25867 #[allow(clippy::too_many_arguments)]
25873 pub fn ssm_conv1d_fused_decode_b_view(
25874 &self,
25875 qkv_cols: &cudarc::driver::CudaView<f32>,
25876 conv_state_ptrs: &cudarc::driver::CudaView<u64>,
25877 w: &CudaSlice<f32>,
25878 conv_outs: &mut CudaSlice<f32>,
25879 conv_dim: usize,
25880 d_conv: usize,
25881 b_n: usize,
25882 ) -> Result<(), Box<dyn std::error::Error>> {
25883 let f = self.func("ssm_conv1d_fused_decode_b_f32");
25884 let cfg = LaunchConfig {
25885 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
25886 block_dim: (256, 1, 1),
25887 shared_mem_bytes: 0,
25888 };
25889 let (cd, dc) = (conv_dim as i32, d_conv as i32);
25890 let __s_b = self.gpu.stream();
25891 let mut b = __s_b.launch_builder(&f);
25892 b.arg(qkv_cols)
25893 .arg(conv_state_ptrs)
25894 .arg(w)
25895 .arg(conv_outs)
25896 .arg(&cd)
25897 .arg(&dc);
25898 unsafe {
25899 b.launch(cfg)?;
25900 }
25901 Ok(())
25902 }
25903
25904 #[allow(clippy::too_many_arguments)]
25905 pub fn gdn_prep_decode_b_view(
25906 &self,
25907 conv_outs: &CudaSlice<f32>,
25908 beta_raws: &cudarc::driver::CudaView<f32>,
25909 alphas: &cudarc::driver::CudaView<f32>,
25910 dt_bias: &CudaSlice<f32>,
25911 a: &CudaSlice<f32>,
25912 q_l2: &mut CudaSlice<f32>,
25913 k_l2: &mut CudaSlice<f32>,
25914 v_g: &mut CudaSlice<f32>,
25915 beta: &mut CudaSlice<f32>,
25916 g_log: &mut CudaSlice<f32>,
25917 d_state: usize,
25918 num_v: usize,
25919 num_k: usize,
25920 key_dim: usize,
25921 eps: f32,
25922 conv_dim: usize,
25923 b_n: usize,
25924 ) -> Result<(), Box<dyn std::error::Error>> {
25925 let f = self.func("gdn_prep_decode_b_f32");
25926 let cfg = LaunchConfig {
25927 grid_dim: (num_v as u32, 1, b_n as u32),
25928 block_dim: (32, 4, 1),
25929 shared_mem_bytes: 0,
25930 };
25931 let (ds, nv, nk, kd, cd) = (
25932 d_state as i32,
25933 num_v as i32,
25934 num_k as i32,
25935 key_dim as i32,
25936 conv_dim as i32,
25937 );
25938 let __s_b = self.gpu.stream();
25939 let mut b = __s_b.launch_builder(&f);
25940 b.arg(conv_outs)
25941 .arg(beta_raws)
25942 .arg(alphas)
25943 .arg(dt_bias)
25944 .arg(a)
25945 .arg(q_l2)
25946 .arg(k_l2)
25947 .arg(v_g)
25948 .arg(beta)
25949 .arg(g_log)
25950 .arg(&ds)
25951 .arg(&nv)
25952 .arg(&nk)
25953 .arg(&kd)
25954 .arg(&eps)
25955 .arg(&cd);
25956 unsafe {
25957 b.launch(cfg)?;
25958 }
25959 Ok(())
25960 }
25961
25962 #[allow(clippy::too_many_arguments)]
25963 pub fn gdn_scan_s128_batched_view(
25964 &self,
25965 q: &CudaSlice<f32>,
25966 k: &CudaSlice<f32>,
25967 v: &CudaSlice<f32>,
25968 g: &CudaSlice<f32>,
25969 beta: &CudaSlice<f32>,
25970 state_in_ptrs: &cudarc::driver::CudaView<u64>,
25971 state_out_ptrs: &cudarc::driver::CudaView<u64>,
25972 o: &mut cudarc::driver::CudaViewMut<f32>,
25973 n_head: usize,
25974 b_n: usize,
25975 scale: f32,
25976 ) -> Result<(), Box<dyn std::error::Error>> {
25977 let f = self.func("gdn_scan_s128_b");
25978 const S_V: u32 = 128;
25979 const WARP: u32 = 32;
25980 const COLS_PER_BLOCK: u32 = 4;
25981 let cfg = LaunchConfig {
25982 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
25983 block_dim: (WARP, COLS_PER_BLOCK, 1),
25984 shared_mem_bytes: 0,
25985 };
25986 let h = n_head as i32;
25987 let __s_b = self.gpu.stream();
25988 let mut b = __s_b.launch_builder(&f);
25989 b.arg(q)
25990 .arg(k)
25991 .arg(v)
25992 .arg(g)
25993 .arg(beta)
25994 .arg(state_in_ptrs)
25995 .arg(state_out_ptrs)
25996 .arg(o)
25997 .arg(&h)
25998 .arg(&scale);
25999 unsafe {
26000 b.launch(cfg)?;
26001 }
26002 Ok(())
26003 }
26004
26005 pub fn gdn_chunked_enabled() -> bool {
26014 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
26015 *E.get_or_init(|| {
26016 std::env::var("MEMRA_GDN_CHUNKED")
26017 .map(|v| v != "0")
26018 .unwrap_or(true)
26019 })
26020 }
26021
26022 pub fn gdn_chunk_size() -> usize {
26027 static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
26028 *C.get_or_init(|| {
26029 let c: usize = std::env::var("MEMRA_GDN_CHUNK")
26030 .ok()
26031 .and_then(|v| v.parse().ok())
26032 .unwrap_or(32);
26033 c.clamp(32, 128) / 32 * 32
26034 })
26035 }
26036
26037 #[allow(clippy::too_many_arguments)]
26042 #[allow(clippy::too_many_arguments, clippy::type_complexity)]
26045 #[allow(clippy::too_many_arguments)]
26046 pub fn gdn_chunk_k123(
26047 &self,
26048 q: &CudaSlice<f32>,
26049 k: &CudaSlice<f32>,
26050 v: &CudaSlice<f32>,
26051 g: &CudaSlice<f32>,
26052 beta: &CudaSlice<f32>,
26053 wb16: Option<&mut CudaSlice<u8>>,
26054 n_head: usize,
26055 t: usize,
26056 c: usize,
26057 hk: usize,
26058 k2w: Option<(&CudaSlice<u8>, &CudaSlice<u8>, &mut CudaSlice<u8>)>,
26059 ) -> Result<
26060 (
26061 CudaSlice<f32>,
26062 CudaSlice<f32>,
26063 CudaSlice<f32>,
26064 CudaSlice<f32>,
26065 ),
26066 Box<dyn std::error::Error>,
26067 > {
26068 const D: usize = 128;
26069 let h = n_head;
26070 let nc = (t + c - 1) / c;
26071 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
26072 let mut gcum = self.uninit(t * h)?;
26073 let mut a = self.uninit(nc * h * c * c)?;
26074 let mut p = self.uninit(nc * h * c * c)?;
26075 let mut u = self.uninit(nc * h * c * D)?;
26076 let mut w = self.uninit(nc * h * c * D)?;
26077 {
26078 let f = self.func("gdn_chunk_cumgate_f32");
26080 let cfg = LaunchConfig {
26081 grid_dim: (nc as u32, h as u32, 1),
26082 block_dim: (32, 1, 1),
26083 shared_mem_bytes: 0,
26084 };
26085 let __s_b = self.gpu.stream();
26086 let mut b = __s_b.launch_builder(&f);
26087 b.arg(g).arg(&mut gcum).arg(&hi).arg(&ti).arg(&ci);
26088 unsafe {
26089 b.launch(cfg)?;
26090 }
26091 }
26092 if let Some((qb, kb, pb)) = k2w {
26093 assert!(c == 32, "gdn_k2_wgmma is a C==32 tile");
26096 let f = self.func("gdn_k2_wgmma");
26097 let cfg = LaunchConfig {
26098 grid_dim: (nc as u32, h as u32, 1),
26099 block_dim: (128, 1, 1),
26100 shared_mem_bytes: 0,
26101 };
26102 let hki = hk as i32;
26103 let __s_b = self.gpu.stream();
26104 let mut b = __s_b.launch_builder(&f);
26105 b.arg(qb)
26106 .arg(kb)
26107 .arg(&gcum)
26108 .arg(beta)
26109 .arg(&mut a)
26110 .arg(&mut *pb)
26111 .arg(&hi)
26112 .arg(&ti)
26113 .arg(&ci)
26114 .arg(&hki);
26115 unsafe {
26116 b.launch(cfg)?;
26117 }
26118 } else if c <= 64 && !portable_mma_gated() {
26119 let f = self.func("gdn_chunk_attn_f32");
26121 f.set_attribute(
26122 CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26123 GDN_K2_DYNAMIC_SHARED_BYTES as i32,
26124 )?;
26125 let jt = ((c + 31) / 32) as u32;
26126 let cfg = LaunchConfig {
26127 grid_dim: (nc as u32, h as u32, jt),
26128 block_dim: (256, 1, 1),
26129 shared_mem_bytes: GDN_K2_DYNAMIC_SHARED_BYTES,
26130 };
26131 let hki = hk as i32;
26132 let __s_b = self.gpu.stream();
26133 let mut b = __s_b.launch_builder(&f);
26134 b.arg(q)
26135 .arg(k)
26136 .arg(&gcum)
26137 .arg(beta)
26138 .arg(&mut a)
26139 .arg(&mut p)
26140 .arg(&hi)
26141 .arg(&ti)
26142 .arg(&ci)
26143 .arg(&hki);
26144 unsafe {
26145 b.launch(cfg)?;
26146 }
26147 } else {
26148 assert!(
26150 hk == h,
26151 "generic K2 is broadcast-only (de-broadcast rides C==32)"
26152 );
26153 let f = self.func("gdn_chunk_attn_g_f32");
26154 let cfg = LaunchConfig {
26155 grid_dim: (nc as u32, h as u32, 1),
26156 block_dim: (32, 8, 1),
26157 shared_mem_bytes: 0,
26158 };
26159 let __s_b = self.gpu.stream();
26160 let mut b = __s_b.launch_builder(&f);
26161 b.arg(q)
26162 .arg(k)
26163 .arg(&gcum)
26164 .arg(beta)
26165 .arg(&mut a)
26166 .arg(&mut p)
26167 .arg(&hi)
26168 .arg(&ti)
26169 .arg(&ci);
26170 unsafe {
26171 b.launch(cfg)?;
26172 }
26173 }
26174 {
26175 let cfg = LaunchConfig {
26177 grid_dim: (nc as u32, h as u32, 1),
26178 block_dim: (256, 1, 1),
26179 shared_mem_bytes: 0,
26180 };
26181 match c {
26182 32 | 64 => {
26183 let f = self.func(if c == 32 {
26184 "gdn_chunk_solve32_f32"
26185 } else {
26186 "gdn_chunk_solve64_f32"
26187 });
26188 let wb: u64 = match wb16 {
26190 Some(d) => self.addr_u8(d),
26191 None => 0,
26192 };
26193 let hki = hk as i32;
26194 let __s_b = self.gpu.stream();
26195 let mut b = __s_b.launch_builder(&f);
26196 b.arg(v)
26197 .arg(k)
26198 .arg(&a)
26199 .arg(&gcum)
26200 .arg(&mut u)
26201 .arg(&mut w)
26202 .arg(&wb)
26203 .arg(&hi)
26204 .arg(&ti)
26205 .arg(&hki);
26206 unsafe {
26207 b.launch(cfg)?;
26208 }
26209 }
26210 _ => {
26211 assert!(hk == h, "generic K3 is broadcast-only");
26212 let f = self.func("gdn_chunk_solve_f32");
26213 let __s_b = self.gpu.stream();
26214 let mut b = __s_b.launch_builder(&f);
26215 b.arg(v)
26216 .arg(k)
26217 .arg(&a)
26218 .arg(&gcum)
26219 .arg(&mut u)
26220 .arg(&mut w)
26221 .arg(&hi)
26222 .arg(&ti)
26223 .arg(&ci);
26224 unsafe {
26225 b.launch(cfg)?;
26226 }
26227 }
26228 }
26229 }
26230 Ok((gcum, p, u, w))
26231 }
26232
26233 pub fn gdn_db_on() -> bool {
26237 std::env::var("MEMRA_GDN_DB").as_deref() != Ok("0")
26238 }
26239
26240 pub fn gdn_mma_enabled(&self, c: usize) -> bool {
26249 !portable_mma_gated()
26250 && c == 32
26251 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
26252 Ok("1") => true,
26253 Ok("0") => false,
26254 _ => gdn_mma_default_on(),
26255 }
26256 }
26257
26258 pub fn gdn_wgmma_on(&self, c: usize) -> bool {
26265 cfg!(memra_hopper_mma)
26266 && self.gdn_mma_enabled(c)
26267 && std::env::var("MEMRA_GDN_WGMMA").as_deref() != Ok("0")
26268 }
26269
26270 #[allow(clippy::too_many_arguments)]
26275 pub fn ssm_conv1d_gdn_state_pad(
26276 &self,
26277 qkv_tm: &cudarc::driver::CudaView<f32>,
26278 conv_state: &mut CudaSlice<f32>,
26279 w: &CudaSlice<f32>,
26280 q_g: &mut CudaSlice<f32>,
26281 k_g: &mut CudaSlice<f32>,
26282 v_g: &mut CudaSlice<f32>,
26283 conv_dim: usize,
26284 t: usize,
26285 d_conv: usize,
26286 d_state: usize,
26287 num_v: usize,
26288 num_k: usize,
26289 key_dim: usize,
26290 hk: usize,
26291 pad_len: Option<&CudaSlice<i32>>,
26292 ) -> Result<(), Box<dyn std::error::Error>> {
26293 assert!(
26294 t >= d_conv - 1,
26295 "fused state conv requires T >= pad (PRIME_MIN_T gates)"
26296 );
26297 {
26298 let f = self.func("ssm_conv1d_gdn_state_f32");
26299 let cfg = LaunchConfig {
26300 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
26301 block_dim: (256, 1, 1),
26302 shared_mem_bytes: 0,
26303 };
26304 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
26305 let (ds, nv, nk, kd, hki) = (
26306 d_state as i32,
26307 num_v as i32,
26308 num_k as i32,
26309 key_dim as i32,
26310 hk as i32,
26311 );
26312 let __s_b = self.gpu.stream();
26313 let mut b = __s_b.launch_builder(&f);
26314 b.arg(qkv_tm)
26315 .arg(&*conv_state)
26316 .arg(w)
26317 .arg(q_g)
26318 .arg(k_g)
26319 .arg(v_g)
26320 .arg(&cd)
26321 .arg(&ti)
26322 .arg(&dc)
26323 .arg(&ds)
26324 .arg(&nv)
26325 .arg(&nk)
26326 .arg(&kd)
26327 .arg(&hki);
26328 unsafe {
26329 b.launch(cfg)?;
26330 }
26331 }
26332 match pad_len {
26333 Some(len_d) => {
26334 let f = self.func("ssm_conv_ring_update_dev_f32");
26335 let n = conv_dim * (d_conv - 1);
26336 let cfg = LaunchConfig::for_num_elems(n as u32);
26337 let (cd, dc) = (conv_dim as i32, d_conv as i32);
26338 let __s_b = self.gpu.stream();
26339 let mut b = __s_b.launch_builder(&f);
26340 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
26341 unsafe {
26342 b.launch(cfg)?;
26343 }
26344 }
26345 None => {
26346 let f = self.func("ssm_conv_ring_update_f32");
26347 let n = conv_dim * (d_conv - 1);
26348 let cfg = LaunchConfig::for_num_elems(n as u32);
26349 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
26350 let __s_b = self.gpu.stream();
26351 let mut b = __s_b.launch_builder(&f);
26352 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
26353 unsafe {
26354 b.launch(cfg)?;
26355 }
26356 }
26357 }
26358 Ok(())
26359 }
26360
26361 pub fn gdn_chunk_alloc(
26365 &self,
26366 n_head: usize,
26367 t: usize,
26368 c: usize,
26369 hk: usize,
26370 ) -> Result<GdnChunkBufs, Box<dyn std::error::Error>> {
26371 const D: usize = 128;
26372 assert!(
26373 c == 32,
26374 "gdn_chunk_alloc: varlen chain is the C==32 mma pair"
26375 );
26376 let h = n_head;
26377 let nc = (t + c - 1) / c;
26378 Ok(GdnChunkBufs {
26379 gcum: self.uninit(t * h)?,
26380 a: self.uninit(nc * h * c * c)?,
26381 p: self.uninit(nc * h * c * c)?,
26382 u: self.uninit(nc * h * c * D)?,
26383 w: self.uninit(nc * h * c * D)?,
26384 kb16: self.alloc_u8_uninit(t * hk * D * 2)?,
26385 wb16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
26386 y16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
26387 ssnap16: self.alloc_u8_uninit(nc * h * D * D * 2)?,
26388 qb16: self.alloc_u8_uninit(t * hk * D * 2)?,
26389 pb16: self.alloc_u8_uninit(nc * h * c * c * 2)?,
26390 o: self.uninit(D * h * t)?,
26391 t,
26392 nc,
26393 })
26394 }
26395
26396 pub fn f32_to_bf16_v(
26398 &self,
26399 x: &cudarc::driver::CudaView<f32>,
26400 dst: &mut CudaSlice<u8>,
26401 n: usize,
26402 ) -> Result<(), Box<dyn std::error::Error>> {
26403 let f = self.func("f32_to_bf16_bulk");
26404 let ni = n as i64;
26405 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
26406 let __s_b = self.gpu.stream();
26407 let mut b = __s_b.launch_builder(&f);
26408 b.arg(x).arg(dst).arg(&ni);
26409 unsafe {
26410 b.launch(cfg)?;
26411 }
26412 Ok(())
26413 }
26414
26415 pub fn f32_to_bf16_into(
26417 &self,
26418 x: &CudaSlice<f32>,
26419 dst: &mut CudaSlice<u8>,
26420 n: usize,
26421 ) -> Result<(), Box<dyn std::error::Error>> {
26422 let f = self.func("f32_to_bf16_bulk");
26423 let ni = n as i64;
26424 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
26425 let __s_b = self.gpu.stream();
26426 let mut b = __s_b.launch_builder(&f);
26427 b.arg(x).arg(dst).arg(&ni);
26428 unsafe {
26429 b.launch(cfg)?;
26430 }
26431 Ok(())
26432 }
26433
26434 pub fn gdn_chunk_k123_vl8(
26437 &self,
26438 seqs: &[GdnSeqVl],
26439 n_head: usize,
26440 hk: usize,
26441 wq: Option<&GdnWVl8>,
26442 ) -> Result<(), Box<dyn std::error::Error>> {
26443 let b = seqs.len();
26444 assert!(b >= 1 && b <= 8, "gdn_chunk_k123_vl8: 1..=8 sequences");
26445 let mut packed = [GdnSeqVl::default(); 8];
26446 packed[..b].copy_from_slice(seqs);
26447 let v = GdnVl8(packed);
26448 let (hi, ci) = (n_head as i32, 32i32);
26449 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
26450 {
26451 let f = self.func("gdn_chunk_cumgate_vl");
26452 let cfg = LaunchConfig {
26453 grid_dim: (max_nc, n_head as u32, b as u32),
26454 block_dim: (32, 1, 1),
26455 shared_mem_bytes: 0,
26456 };
26457 let __s_lb = self.gpu.stream();
26458 let mut lb = __s_lb.launch_builder(&f);
26459 lb.arg(&v).arg(&hi).arg(&ci);
26460 unsafe {
26461 lb.launch(cfg)?;
26462 }
26463 }
26464 let hki = hk as i32;
26465 if let Some(w) = wq {
26466 let f = self.func("gdn_k2_wgmma_vl");
26468 let cfg = LaunchConfig {
26469 grid_dim: (max_nc, n_head as u32, b as u32),
26470 block_dim: (128, 1, 1),
26471 shared_mem_bytes: 0,
26472 };
26473 let __s_lb = self.gpu.stream();
26474 let mut lb = __s_lb.launch_builder(&f);
26475 lb.arg(&v).arg(w).arg(&hi).arg(&ci).arg(&hki);
26476 unsafe {
26477 lb.launch(cfg)?;
26478 }
26479 } else {
26480 let f = self.func("gdn_chunk_attn_vl");
26481 f.set_attribute(
26482 CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
26483 GDN_K2_DYNAMIC_SHARED_BYTES as i32,
26484 )?;
26485 let cfg = LaunchConfig {
26486 grid_dim: (max_nc, n_head as u32, b as u32),
26487 block_dim: (256, 1, 1),
26488 shared_mem_bytes: GDN_K2_DYNAMIC_SHARED_BYTES,
26489 };
26490 let __s_lb = self.gpu.stream();
26491 let mut lb = __s_lb.launch_builder(&f);
26492 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
26493 unsafe {
26494 lb.launch(cfg)?;
26495 }
26496 }
26497 {
26498 let f = self.func("gdn_chunk_solve32_vl");
26499 let cfg = LaunchConfig {
26500 grid_dim: (max_nc, n_head as u32, b as u32),
26501 block_dim: (256, 1, 1),
26502 shared_mem_bytes: 0,
26503 };
26504 let __s_lb = self.gpu.stream();
26505 let mut lb = __s_lb.launch_builder(&f);
26506 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
26507 unsafe {
26508 lb.launch(cfg)?;
26509 }
26510 }
26511 Ok(())
26512 }
26513
26514 #[allow(clippy::too_many_arguments)]
26518 pub fn gdn_prep_vl8(
26519 &self,
26520 seqs: &[GdnPrepVl],
26521 conv_w: &CudaSlice<f32>,
26522 dt_bias: &CudaSlice<f32>,
26523 a: &CudaSlice<f32>,
26524 conv_dim: usize,
26525 d_conv: usize,
26526 d_state: usize,
26527 num_v: usize,
26528 num_k: usize,
26529 key_dim: usize,
26530 hk: usize,
26531 eps: f32,
26532 ) -> Result<(), Box<dyn std::error::Error>> {
26533 let b = seqs.len();
26534 assert!(b >= 1 && b <= 8);
26535 let mut packed = [GdnPrepVl::default(); 8];
26536 packed[..b].copy_from_slice(seqs);
26537 let v = GdnPrepVl8(packed);
26538 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
26539 let (cdi, dci) = (conv_dim as i32, d_conv as i32);
26540 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
26541 assert!(
26542 conv_fuse || hk == num_v,
26543 "de-broadcast requires the fused conv"
26544 );
26545 if conv_fuse {
26546 let f = self.func("ssm_conv1d_gdn_state_vl");
26547 let cfg = LaunchConfig {
26548 grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32),
26549 block_dim: (256, 1, 1),
26550 shared_mem_bytes: 0,
26551 };
26552 let (dsi, nvi, nki, kdi, hki) = (
26553 d_state as i32,
26554 num_v as i32,
26555 num_k as i32,
26556 key_dim as i32,
26557 hk as i32,
26558 );
26559 let __s_lb = self.gpu.stream();
26560 let mut lb = __s_lb.launch_builder(&f);
26561 lb.arg(&v)
26562 .arg(conv_w)
26563 .arg(&cdi)
26564 .arg(&dci)
26565 .arg(&dsi)
26566 .arg(&nvi)
26567 .arg(&nki)
26568 .arg(&kdi)
26569 .arg(&hki);
26570 unsafe {
26571 lb.launch(cfg)?;
26572 }
26573 } else {
26574 let f = self.func("ssm_conv1d_tm_state_vl");
26575 let cfg = LaunchConfig {
26576 grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32),
26577 block_dim: (256, 1, 1),
26578 shared_mem_bytes: 0,
26579 };
26580 let __s_lb = self.gpu.stream();
26581 let mut lb = __s_lb.launch_builder(&f);
26582 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci);
26583 unsafe {
26584 lb.launch(cfg)?;
26585 }
26586 }
26587 {
26588 let f = self.func("ssm_conv_ring_update_vl");
26589 let n = (conv_dim * (d_conv - 1)) as u32;
26590 let cfg = LaunchConfig {
26591 grid_dim: (n.div_ceil(256), 1, b as u32),
26592 block_dim: (256, 1, 1),
26593 shared_mem_bytes: 0,
26594 };
26595 let __s_lb = self.gpu.stream();
26596 let mut lb = __s_lb.launch_builder(&f);
26597 lb.arg(&v).arg(&cdi).arg(&dci);
26598 unsafe {
26599 lb.launch(cfg)?;
26600 }
26601 }
26602 if !conv_fuse {
26603 let f = self.func("qkv_to_gdn_repack_vl");
26604 let n = max_t * (num_v * d_state) as u32;
26605 let cfg = LaunchConfig {
26606 grid_dim: (n.div_ceil(256), 1, b as u32),
26607 block_dim: (256, 1, 1),
26608 shared_mem_bytes: 0,
26609 };
26610 let (dsi, nvi, nki, kdi) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
26611 let __s_lb = self.gpu.stream();
26612 let mut lb = __s_lb.launch_builder(&f);
26613 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi);
26614 unsafe {
26615 lb.launch(cfg)?;
26616 }
26617 }
26618 if Self::l2_v2_on(d_state) {
26619 let f = self.func("gdn_l2_v2_vl");
26620 let cfg = LaunchConfig {
26621 grid_dim: ((max_t * hk as u32).div_ceil(8), 2, b as u32),
26622 block_dim: (256, 1, 1),
26623 shared_mem_bytes: 0,
26624 };
26625 let (dsi, nvi) = (d_state as i32, hk as i32);
26626 let __s_lb = self.gpu.stream();
26627 let mut lb = __s_lb.launch_builder(&f);
26628 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
26629 unsafe {
26630 lb.launch(cfg)?;
26631 }
26632 } else {
26633 let f = self.func("gdn_l2_vl");
26634 let cfg = LaunchConfig {
26635 grid_dim: (max_t * hk as u32, 2, b as u32),
26636 block_dim: (256, 1, 1),
26637 shared_mem_bytes: 0,
26638 };
26639 let (dsi, nvi) = (d_state as i32, hk as i32);
26640 let __s_lb = self.gpu.stream();
26641 let mut lb = __s_lb.launch_builder(&f);
26642 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
26643 unsafe {
26644 lb.launch(cfg)?;
26645 }
26646 }
26647 {
26648 let f = self.func("gdn_gate_prep_vl");
26649 let n = max_t * num_v as u32;
26650 let cfg = LaunchConfig {
26651 grid_dim: (n.div_ceil(256), 1, b as u32),
26652 block_dim: (256, 1, 1),
26653 shared_mem_bytes: 0,
26654 };
26655 let nvi = num_v as i32;
26656 let __s_lb = self.gpu.stream();
26657 let mut lb = __s_lb.launch_builder(&f);
26658 lb.arg(&v).arg(dt_bias).arg(a).arg(&nvi);
26659 unsafe {
26660 lb.launch(cfg)?;
26661 }
26662 }
26663 Ok(())
26664 }
26665
26666 pub fn gdn_mirror_vl8(
26668 &self,
26669 seqs: &[GdnSeqVl],
26670 n_head: usize,
26671 which: i32,
26672 hk: usize,
26673 ) -> Result<(), Box<dyn std::error::Error>> {
26674 let b = seqs.len();
26675 assert!(b >= 1 && b <= 8);
26676 let mut packed = [GdnSeqVl::default(); 8];
26677 packed[..b].copy_from_slice(seqs);
26678 let v = GdnVl8(packed);
26679 let ept = (if which == 0 { hk } else { n_head } * 128) as i32;
26680 let max_n = seqs
26681 .iter()
26682 .map(|s| {
26683 if which == 0 {
26684 s.t as i64 * ept as i64
26685 } else {
26686 s.nc as i64 * ept as i64 * 32
26687 }
26688 })
26689 .max()
26690 .unwrap();
26691 let f = self.func("gdn_mirror_vl");
26692 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
26693 let cfg = LaunchConfig {
26694 grid_dim: (blocks, 1, b as u32),
26695 block_dim: (256, 1, 1),
26696 shared_mem_bytes: 0,
26697 };
26698 let __s_lb = self.gpu.stream();
26699 let mut lb = __s_lb.launch_builder(&f);
26700 lb.arg(&v).arg(&ept).arg(&which);
26701 unsafe {
26702 lb.launch(cfg)?;
26703 }
26704 Ok(())
26705 }
26706
26707 pub fn gdn_tail_vl8(
26709 &self,
26710 seqs: &[GdnPrepVl],
26711 norm_w: &CudaSlice<f32>,
26712 d_state: usize,
26713 num_v: usize,
26714 eps: f32,
26715 ) -> Result<(), Box<dyn std::error::Error>> {
26716 let b = seqs.len();
26717 assert!(b >= 1 && b <= 8);
26718 let mut packed = [GdnPrepVl::default(); 8];
26719 packed[..b].copy_from_slice(seqs);
26720 let v = GdnPrepVl8(packed);
26721 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
26722 let f = self.func("gated_rmsnorm_f16out_vl");
26723 let cfg = LaunchConfig {
26725 grid_dim: (max_t * num_v as u32, 1, b as u32),
26726 block_dim: (128, 1, 1),
26727 shared_mem_bytes: 0,
26728 };
26729 let (dsi, nvi) = (d_state as i32, num_v as i32);
26730 let __s_lb = self.gpu.stream();
26731 let mut lb = __s_lb.launch_builder(&f);
26732 lb.arg(&v).arg(norm_w).arg(&dsi).arg(&nvi).arg(&eps);
26733 unsafe {
26734 lb.launch(cfg)?;
26735 }
26736 Ok(())
26737 }
26738
26739 pub fn addr_f32(&self, x: &CudaSlice<f32>) -> u64 {
26742 use cudarc::driver::DevicePtr;
26743 let s = self.gpu.stream();
26744 let (p, _g) = x.device_ptr(&s);
26745 p as u64
26746 }
26747 pub fn addr_f32_mut(&self, x: &mut CudaSlice<f32>) -> u64 {
26748 use cudarc::driver::DevicePtrMut;
26749 let s = self.gpu.stream();
26750 let (p, _g) = x.device_ptr_mut(&s);
26751 p as u64
26752 }
26753 pub fn addr_f32v(&self, x: &cudarc::driver::CudaView<f32>) -> u64 {
26754 use cudarc::driver::DevicePtr;
26755 let s = self.gpu.stream();
26756 let (p, _g) = x.device_ptr(&s);
26757 p as u64
26758 }
26759 pub fn addr_u8(&self, x: &CudaSlice<u8>) -> u64 {
26760 use cudarc::driver::DevicePtr;
26761 let s = self.gpu.stream();
26762 let (p, _g) = x.device_ptr(&s);
26763 p as u64
26764 }
26765
26766 pub fn gdn_chunk_vl8(
26770 &self,
26771 seqs: &[GdnSeqVl],
26772 n_head: usize,
26773 scale: f32,
26774 hk: usize,
26775 wq: Option<&GdnWVl8>,
26776 ) -> Result<(), Box<dyn std::error::Error>> {
26777 const NSPLIT: u32 = 4;
26778 let b = seqs.len();
26779 assert!(b >= 1 && b <= 8, "gdn_chunk_vl8: 1..=8 sequences");
26780 let mut packed = [GdnSeqVl::default(); 8];
26781 packed[..b].copy_from_slice(seqs);
26782 let v = GdnVl8(packed);
26783 let (hi, ci) = (n_head as i32, 32i32);
26784 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
26785 let hki = hk as i32;
26786 if let Some(w) = wq {
26787 let f = self.func("gdn_k45_wgmma_vl");
26789 let cfg = LaunchConfig {
26790 grid_dim: (n_head as u32, NSPLIT, b as u32),
26791 block_dim: (256, 1, 1),
26792 shared_mem_bytes: 0,
26793 };
26794 let __s_lb = self.gpu.stream();
26795 let mut lb = __s_lb.launch_builder(&f);
26796 lb.arg(&v).arg(w).arg(&scale).arg(&hi).arg(&ci).arg(&hki);
26797 unsafe {
26798 lb.launch(cfg)?;
26799 }
26800 let _ = max_nc;
26801 return Ok(());
26802 }
26803 {
26804 let f = self.func("gdn_chunk_state_mma_vl");
26805 let cfg = LaunchConfig {
26806 grid_dim: (n_head as u32, NSPLIT, b as u32),
26807 block_dim: (256, 1, 1),
26808 shared_mem_bytes: 0,
26809 };
26810 let __s_lb = self.gpu.stream();
26811 let mut lb = __s_lb.launch_builder(&f);
26812 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
26813 unsafe {
26814 lb.launch(cfg)?;
26815 }
26816 }
26817 {
26818 let f = self.func("gdn_chunk_output_mma_vl");
26819 let cfg = LaunchConfig {
26820 grid_dim: (max_nc, n_head as u32, b as u32),
26821 block_dim: (256, 1, 1),
26822 shared_mem_bytes: 0,
26823 };
26824 let __s_lb = self.gpu.stream();
26825 let mut lb = __s_lb.launch_builder(&f);
26826 lb.arg(&v).arg(&hi).arg(&ci).arg(&scale).arg(&hki);
26827 unsafe {
26828 lb.launch(cfg)?;
26829 }
26830 }
26831 Ok(())
26832 }
26833 pub fn gdn_scan_chunked(
26834 &self,
26835 q: &CudaSlice<f32>,
26836 k: &CudaSlice<f32>,
26837 v: &CudaSlice<f32>,
26838 g: &CudaSlice<f32>,
26839 beta: &CudaSlice<f32>,
26840 kb16_pre: Option<&CudaSlice<u8>>,
26841 qb16_pre: Option<&CudaSlice<u8>>,
26842 state_in: &CudaSlice<f32>,
26843 state_out: &mut CudaSlice<f32>,
26844 o: &mut CudaSlice<f32>,
26845 n_head: usize,
26846 t: usize,
26847 scale: f32,
26848 c: usize,
26849 hk: usize,
26850 ) -> Result<(), Box<dyn std::error::Error>> {
26851 const D: usize = 128;
26852 const NSPLIT: u32 = 4;
26853 assert!(c >= 1 && c <= 128, "gdn_scan_chunked: C must be in 1..=128");
26854 let h = n_head;
26855 let nc = (t + c - 1) / c;
26856 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
26857 let gdn_mma_pre = !portable_mma_gated()
26862 && c == 32
26863 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
26864 Ok("1") => true,
26865 Ok("0") => false,
26866 _ => gdn_mma_default_on(),
26867 };
26868 let mut wb16_pre: Option<CudaSlice<u8>> = if gdn_mma_pre {
26869 Some(self.alloc_u8_uninit(nc * h * c * D * 2)?)
26870 } else {
26871 None
26872 };
26873 let gdn_wgmma_pre = cfg!(memra_hopper_mma)
26878 && gdn_mma_pre
26879 && std::env::var("MEMRA_GDN_WGMMA").as_deref() != Ok("0");
26880 let nk = t * hk * D;
26881 let mut kb16_local: Option<CudaSlice<u8>> = None;
26882 if gdn_mma_pre && kb16_pre.is_none() {
26883 let mut kb = self.alloc_u8_uninit(nk * 2)?;
26884 let f = self.func("f32_to_bf16_bulk");
26885 let n2 = nk as i64;
26886 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
26887 let __s_b = self.gpu.stream();
26888 let mut b = __s_b.launch_builder(&f);
26889 b.arg(k).arg(&mut kb).arg(&n2);
26890 unsafe {
26891 b.launch(cfg2)?;
26892 }
26893 kb16_local = Some(kb);
26894 }
26895 let kb16_ref0: Option<&CudaSlice<u8>> = kb16_local.as_ref().or(kb16_pre);
26896 if let Some(kb) = kb16_pre {
26897 assert!(kb.len() >= nk * 2, "kb16_pre too small");
26898 }
26899 let mut qb16: Option<CudaSlice<u8>> = None;
26900 let mut pb16: Option<CudaSlice<u8>> = None;
26901 if gdn_wgmma_pre {
26902 if qb16_pre.is_none() {
26905 let mut qb = self.alloc_u8_uninit(nk * 2)?;
26906 let f = self.func("f32_to_bf16_bulk");
26907 let n2 = nk as i64;
26908 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
26909 let __s_b = self.gpu.stream();
26910 let mut b = __s_b.launch_builder(&f);
26911 b.arg(q).arg(&mut qb).arg(&n2);
26912 unsafe {
26913 b.launch(cfg2)?;
26914 }
26915 qb16 = Some(qb);
26916 } else if let Some(qb) = qb16_pre {
26917 assert!(qb.len() >= nk * 2, "qb16_pre too small");
26918 }
26919 pb16 = Some(self.alloc_u8_uninit(nc * h * c * c * 2)?);
26920 }
26921 let qb16_ref0: Option<&CudaSlice<u8>> = qb16.as_ref().or(qb16_pre);
26922 let k2w = if gdn_wgmma_pre {
26923 Some((
26924 *qb16_ref0.as_ref().unwrap(),
26925 *kb16_ref0.as_ref().unwrap(),
26926 pb16.as_mut().unwrap(),
26927 ))
26928 } else {
26929 None
26930 };
26931 let (gcum, p, u, w) =
26932 self.gdn_chunk_k123(q, k, v, g, beta, wb16_pre.as_mut(), n_head, t, c, hk, k2w)?;
26933 let _ = &w;
26934 let mut y = self.uninit(nc * h * c * D)?;
26935 let mut ssnap = self.uninit(nc * h * D * D)?; let gdn_mma = !portable_mma_gated()
26949 && c == 32
26950 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
26951 Ok("1") => true,
26952 Ok("0") => false,
26953 _ => gdn_mma_default_on(),
26954 };
26955 if gdn_mma {
26956 let wb16 = wb16_pre
26957 .take()
26958 .expect("mma path pre-allocates wb16 (K3 store fold)");
26959 let kb16_ref: &CudaSlice<u8> = kb16_ref0.expect("mma path pre-builds kb16 above K123");
26960 if gdn_wgmma_pre {
26972 let qb16 = qb16_ref0.unwrap();
26974 let pb16 = pb16.as_ref().unwrap();
26975 {
26976 let f = self.func("gdn_k45_wgmma");
26977 let cfg = LaunchConfig {
26978 grid_dim: (h as u32, 4, 1),
26979 block_dim: (256, 1, 1),
26980 shared_mem_bytes: 0,
26981 };
26982 let hki = hk as i32;
26983 let __s_b = self.gpu.stream();
26984 let mut b = __s_b.launch_builder(&f);
26985 b.arg(kb16_ref)
26986 .arg(&gcum)
26987 .arg(beta)
26988 .arg(&u)
26989 .arg(&wb16)
26990 .arg(qb16)
26991 .arg(pb16)
26992 .arg(o)
26993 .arg(&scale)
26994 .arg(state_in)
26995 .arg(&mut *state_out)
26996 .arg(&hi)
26997 .arg(&ti)
26998 .arg(&ci)
26999 .arg(&hki);
27000 unsafe {
27001 b.launch(cfg)?;
27002 }
27003 }
27004 return Ok(());
27005 }
27006 let mut y16 = self.alloc_u8_uninit(nc * h * c * D * 2)?;
27010 let mut ssnap16 = self.alloc_u8_uninit(nc * h * D * D * 2)?;
27011 {
27012 let f = self.func("gdn_chunk_state_mma");
27013 let cfg = LaunchConfig {
27014 grid_dim: (h as u32, NSPLIT, 1),
27015 block_dim: (256, 1, 1),
27016 shared_mem_bytes: 0,
27017 };
27018 let hki = hk as i32;
27019 let __s_b = self.gpu.stream();
27020 let mut b = __s_b.launch_builder(&f);
27021 b.arg(kb16_ref)
27022 .arg(&gcum)
27023 .arg(beta)
27024 .arg(&u)
27025 .arg(&wb16)
27026 .arg(&mut y16)
27027 .arg(&mut ssnap16)
27028 .arg(state_in)
27029 .arg(&mut *state_out)
27030 .arg(&hi)
27031 .arg(&ti)
27032 .arg(&ci)
27033 .arg(&hki);
27034 unsafe {
27035 b.launch(cfg)?;
27036 }
27037 }
27038 {
27039 let f = self.func("gdn_chunk_output_mma");
27041 let jt = ((c + 31) / 32) as u32;
27042 let cfg = LaunchConfig {
27043 grid_dim: (nc as u32, h as u32, jt),
27044 block_dim: (256, 1, 1),
27045 shared_mem_bytes: 0,
27046 };
27047 let hki = hk as i32;
27048 let __s_b = self.gpu.stream();
27049 let mut b = __s_b.launch_builder(&f);
27050 b.arg(q)
27051 .arg(&gcum)
27052 .arg(&p)
27053 .arg(&y16)
27054 .arg(&ssnap16)
27055 .arg(o)
27056 .arg(&hi)
27057 .arg(&ti)
27058 .arg(&ci)
27059 .arg(&scale)
27060 .arg(&hki);
27061 unsafe {
27062 b.launch(cfg)?;
27063 }
27064 }
27065 return Ok(());
27066 }
27067 {
27068 let f = self.func("gdn_chunk_state_f32");
27070 let cfg = LaunchConfig {
27071 grid_dim: (h as u32, NSPLIT, 1),
27072 block_dim: (256, 1, 1),
27073 shared_mem_bytes: 0,
27074 };
27075 let __s_b = self.gpu.stream();
27076 let mut b = __s_b.launch_builder(&f);
27077 b.arg(k)
27078 .arg(&gcum)
27079 .arg(beta)
27080 .arg(&u)
27081 .arg(&w)
27082 .arg(&mut y)
27083 .arg(&mut ssnap)
27084 .arg(state_in)
27085 .arg(&mut *state_out)
27086 .arg(&hi)
27087 .arg(&ti)
27088 .arg(&ci);
27089 unsafe {
27090 b.launch(cfg)?;
27091 }
27092 }
27093 {
27094 let f = self.func("gdn_chunk_output_f32");
27096 let jt = ((c + 31) / 32) as u32;
27097 let cfg = LaunchConfig {
27098 grid_dim: (nc as u32, h as u32, jt),
27099 block_dim: (256, 1, 1),
27100 shared_mem_bytes: 0,
27101 };
27102 let __s_b = self.gpu.stream();
27103 let mut b = __s_b.launch_builder(&f);
27104 b.arg(q)
27105 .arg(&gcum)
27106 .arg(&p)
27107 .arg(&y)
27108 .arg(&ssnap)
27109 .arg(o)
27110 .arg(&hi)
27111 .arg(&ti)
27112 .arg(&ci)
27113 .arg(&scale);
27114 unsafe {
27115 b.launch(cfg)?;
27116 }
27117 }
27118 Ok(())
27119 }
27120
27121 #[allow(clippy::too_many_arguments)]
27130 #[allow(clippy::too_many_arguments)]
27131 pub fn gdn_scan_prefill(
27132 &self,
27133 q: &CudaSlice<f32>,
27134 k: &CudaSlice<f32>,
27135 v: &CudaSlice<f32>,
27136 g: &CudaSlice<f32>,
27137 beta: &CudaSlice<f32>,
27138 kb16_pre: Option<&CudaSlice<u8>>,
27139 qb16_pre: Option<&CudaSlice<u8>>,
27140 state_in: &CudaSlice<f32>,
27141 state_out: &mut CudaSlice<f32>,
27142 o: &mut CudaSlice<f32>,
27143 n_head: usize,
27144 t: usize,
27145 scale: f32,
27146 hk: usize,
27147 ) -> Result<(), Box<dyn std::error::Error>> {
27148 if std::env::var("MEMRA_GDN_DIFF").is_ok() && t >= 16 {
27149 assert!(hk == n_head, "GDN_DIFF oracle is broadcast-only");
27150 return self.gdn_scan_diff(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale);
27151 }
27152 if Self::gdn_chunked_enabled() && t >= 16 {
27153 self.gdn_scan_chunked(
27154 q,
27155 k,
27156 v,
27157 g,
27158 beta,
27159 kb16_pre,
27160 qb16_pre,
27161 state_in,
27162 state_out,
27163 o,
27164 n_head,
27165 t,
27166 scale,
27167 Self::gdn_chunk_size(),
27168 hk,
27169 )
27170 } else {
27171 assert!(
27172 hk == n_head,
27173 "s128 scan is broadcast-only (prep guarantees by predicate)"
27174 );
27175 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)
27176 }
27177 }
27178
27179 #[allow(clippy::too_many_arguments)]
27181 fn gdn_scan_diff(
27182 &self,
27183 q: &CudaSlice<f32>,
27184 k: &CudaSlice<f32>,
27185 v: &CudaSlice<f32>,
27186 g: &CudaSlice<f32>,
27187 beta: &CudaSlice<f32>,
27188 state_in: &CudaSlice<f32>,
27189 state_out: &mut CudaSlice<f32>,
27190 o: &mut CudaSlice<f32>,
27191 n_head: usize,
27192 t: usize,
27193 scale: f32,
27194 ) -> Result<(), Box<dyn std::error::Error>> {
27195 static CALL: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
27196 let call = CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
27197 let mut o_c = self.uninit(o.len())?;
27198 let mut st_c = self.uninit(state_out.len())?;
27199 self.gdn_scan_chunked(
27200 q,
27201 k,
27202 v,
27203 g,
27204 beta,
27205 None,
27206 None,
27207 state_in,
27208 &mut st_c,
27209 &mut o_c,
27210 n_head,
27211 t,
27212 scale,
27213 Self::gdn_chunk_size(),
27214 n_head,
27215 )?;
27216 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)?;
27217 let (oh_s, oh_c) = (self.dtoh(o)?, self.dtoh(&o_c)?);
27218 let (sh_s, sh_c) = (self.dtoh(state_out)?, self.dtoh(&st_c)?);
27219 let stats = |a: &[f32], b: &[f32]| -> (f32, f32, f64) {
27220 let mut max_abs = 0f32;
27221 let mut max_rel = 0f32;
27222 let mut sum_rel = 0f64;
27223 for (x, y) in a.iter().zip(b) {
27224 let ad = (x - y).abs();
27225 let rel = ad / x.abs().max(y.abs()).max(1e-3);
27226 if ad > max_abs {
27227 max_abs = ad;
27228 }
27229 if rel > max_rel {
27230 max_rel = rel;
27231 }
27232 sum_rel += rel as f64;
27233 }
27234 (max_abs, max_rel, sum_rel / a.len() as f64)
27235 };
27236 let (o_ma, o_mr, o_mean) = stats(&oh_s, &oh_c);
27237 let (s_ma, s_mr, s_mean) = stats(&sh_s, &sh_c);
27238 println!(
27239 "[gdn-diff call {call:3} T={t} C={}] out: max_abs={o_ma:.3e} max_rel={o_mr:.3e} mean_rel={o_mean:.3e} | \
27240 state: max_abs={s_ma:.3e} max_rel={s_mr:.3e} mean_rel={s_mean:.3e}",
27241 Self::gdn_chunk_size()
27242 );
27243 Ok(())
27244 }
27245
27246 pub fn gdn_glog(
27248 &self,
27249 alpha: &CudaSlice<f32>,
27250 dt_bias: &CudaSlice<f32>,
27251 a: &CudaSlice<f32>,
27252 g_log: &mut CudaSlice<f32>,
27253 n_head: usize,
27254 t: usize,
27255 ) -> Result<(), Box<dyn std::error::Error>> {
27256 let f = self.func("gdn_glog_f32");
27257 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
27258 let (h, ti) = (n_head as i32, t as i32);
27259 let __s_b = self.gpu.stream();
27260 let mut b = __s_b.launch_builder(&f);
27261 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
27262 unsafe {
27263 b.launch(cfg)?;
27264 }
27265 Ok(())
27266 }
27267
27268 pub fn sigmoid_v(
27271 &self,
27272 x: &cudarc::driver::CudaView<f32>,
27273 y: &mut CudaSlice<f32>,
27274 n: usize,
27275 ) -> Result<(), Box<dyn std::error::Error>> {
27276 let f = self.func("sigmoid_f32");
27277 let cfg = LaunchConfig::for_num_elems(n as u32);
27278 let ni = n as i32;
27279 let __s_b = self.gpu.stream();
27280 let mut b = __s_b.launch_builder(&f);
27281 b.arg(x).arg(y).arg(&ni);
27282 unsafe {
27283 b.launch(cfg)?;
27284 }
27285 Ok(())
27286 }
27287
27288 pub fn gdn_glog_v(
27289 &self,
27290 alpha: &cudarc::driver::CudaView<f32>,
27291 dt_bias: &CudaSlice<f32>,
27292 a: &CudaSlice<f32>,
27293 g_log: &mut CudaSlice<f32>,
27294 n_head: usize,
27295 t: usize,
27296 ) -> Result<(), Box<dyn std::error::Error>> {
27297 let f = self.func("gdn_glog_f32");
27298 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
27299 let (h, ti) = (n_head as i32, t as i32);
27300 let __s_b = self.gpu.stream();
27301 let mut b = __s_b.launch_builder(&f);
27302 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
27303 unsafe {
27304 b.launch(cfg)?;
27305 }
27306 Ok(())
27307 }
27308
27309 pub fn sigmoid(
27310 &self,
27311 x: &CudaSlice<f32>,
27312 y: &mut CudaSlice<f32>,
27313 n: usize,
27314 ) -> Result<(), Box<dyn std::error::Error>> {
27315 let f = self.func("sigmoid_f32");
27316 let cfg = LaunchConfig::for_num_elems(n as u32);
27317 let ni = n as i32;
27318 let __s_b = self.gpu.stream();
27319 let mut b = __s_b.launch_builder(&f);
27320 b.arg(x).arg(y).arg(&ni);
27321 unsafe {
27322 b.launch(cfg)?;
27323 }
27324 Ok(())
27325 }
27326
27327 pub fn sig_mul_f16out(
27330 &self,
27331 a: &CudaSlice<f32>,
27332 g: &CudaSlice<f32>,
27333 dst: &mut CudaSlice<f32>,
27334 dst16: &mut CudaSlice<u8>,
27335 n: usize,
27336 ) -> Result<(), Box<dyn std::error::Error>> {
27337 let f = self.func("sig_mul_f16out_f32");
27338 let cfg = LaunchConfig::for_num_elems(n as u32);
27339 let ni = n as i32;
27340 let __s_b = self.gpu.stream();
27341 let mut b = __s_b.launch_builder(&f);
27342 b.arg(a).arg(g).arg(dst).arg(dst16).arg(&ni);
27343 unsafe {
27344 b.launch(cfg)?;
27345 }
27346 Ok(())
27347 }
27348
27349 #[allow(clippy::too_many_arguments)]
27358 pub fn attn_head_gate(
27359 &self,
27360 a: &CudaSlice<f32>,
27361 g: &CudaSlice<f32>,
27362 dst: &mut CudaSlice<f32>,
27363 dst16: Option<&mut CudaSlice<u8>>,
27364 head_dim: usize,
27365 n_head: usize,
27366 t: usize,
27367 ) -> Result<(), Box<dyn std::error::Error>> {
27368 let f = self.func("attn_head_gate_f32");
27369 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
27370 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
27371 let d16: u64 = match dst16 {
27373 Some(d) => self.addr_u8(d),
27374 None => 0,
27375 };
27376 let __s_b = self.gpu.stream();
27377 let mut b = __s_b.launch_builder(&f);
27378 b.arg(a)
27379 .arg(g)
27380 .arg(dst)
27381 .arg(&d16)
27382 .arg(&hd)
27383 .arg(&nh)
27384 .arg(&ti);
27385 unsafe {
27386 b.launch(cfg)?;
27387 }
27388 Ok(())
27389 }
27390
27391 #[allow(clippy::too_many_arguments)]
27400 pub fn swiglu_clamped_mul_scaled(
27401 &self,
27402 gate: &CudaSlice<f32>,
27403 up: &CudaSlice<f32>,
27404 gs: f32,
27405 us: f32,
27406 limit: f32,
27407 dst: &mut CudaSlice<f32>,
27408 n: usize,
27409 ) -> Result<(), Box<dyn std::error::Error>> {
27410 debug_assert!(
27411 limit > 1e-6,
27412 "swiglu_clamped needs a live limit; use silu_mul_scaled"
27413 );
27414 let f = self.func("swiglu_clamped_mul_scaled_f32");
27415 let cfg = LaunchConfig::for_num_elems(n as u32);
27416 let ni = n as i32;
27417 let __s_b = self.gpu.stream();
27418 let mut b = __s_b.launch_builder(&f);
27419 b.arg(gate)
27420 .arg(up)
27421 .arg(&gs)
27422 .arg(&us)
27423 .arg(&limit)
27424 .arg(dst)
27425 .arg(&ni);
27426 unsafe {
27427 b.launch(cfg)?;
27428 }
27429 Ok(())
27430 }
27431
27432 pub fn gated_rmsnorm(
27434 &self,
27435 o: &CudaSlice<f32>,
27436 w: &CudaSlice<f32>,
27437 z: &CudaSlice<f32>,
27438 dst: &mut CudaSlice<f32>,
27439 ncols: usize,
27440 nrows: usize,
27441 eps: f32,
27442 ) -> Result<(), Box<dyn std::error::Error>> {
27443 let f = self.func("gated_rmsnorm_f32");
27444 let cfg = LaunchConfig {
27445 grid_dim: (nrows as u32, 1, 1),
27446 block_dim: (128, 1, 1),
27447 shared_mem_bytes: 0,
27448 };
27449 let (nc, e) = (ncols as i32, eps);
27450 let __s_b = self.gpu.stream();
27451 let mut b = __s_b.launch_builder(&f);
27452 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
27453 unsafe {
27454 b.launch(cfg)?;
27455 }
27456 Ok(())
27457 }
27458
27459 pub fn gated_rmsnorm_f16out(
27462 &self,
27463 o: &CudaSlice<f32>,
27464 w: &CudaSlice<f32>,
27465 z: &CudaSlice<f32>,
27466 dst: &mut CudaSlice<f32>,
27467 dst16: &mut CudaSlice<u8>,
27468 ncols: usize,
27469 nrows: usize,
27470 eps: f32,
27471 ) -> Result<(), Box<dyn std::error::Error>> {
27472 let f = self.func("gated_rmsnorm_f16out_f32");
27473 let cfg = LaunchConfig {
27475 grid_dim: (nrows as u32, 1, 1),
27476 block_dim: (128, 1, 1),
27477 shared_mem_bytes: 0,
27478 };
27479 let (nc, e) = (ncols as i32, eps);
27480 let __s_b = self.gpu.stream();
27481 let mut b = __s_b.launch_builder(&f);
27482 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
27483 unsafe {
27484 b.launch(cfg)?;
27485 }
27486 Ok(())
27487 }
27488
27489 #[allow(clippy::too_many_arguments)]
27493 pub fn add_rms_norm_zq8(
27494 &self,
27495 a: &CudaSlice<f32>,
27496 b_in: &CudaSlice<f32>,
27497 w: &CudaSlice<f32>,
27498 res: &mut CudaSlice<f32>,
27499 z: &mut CudaSlice<f32>,
27500 ncols: usize,
27501 nrows: usize,
27502 eps: f32,
27503 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
27504 assert!(ncols % 32 == 0);
27505 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
27506 let mut d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
27507 let f = self.func("add_rms_norm_zq8");
27508 let cfg = LaunchConfig {
27509 grid_dim: (nrows as u32, 1, 1),
27510 block_dim: (1024, 1, 1),
27511 shared_mem_bytes: 0,
27512 };
27513 let (nc, ep) = (ncols as i32, eps);
27514 let __s_b = self.gpu.stream();
27515 let mut b = __s_b.launch_builder(&f);
27516 b.arg(a)
27517 .arg(b_in)
27518 .arg(w)
27519 .arg(res)
27520 .arg(z)
27521 .arg(&mut q)
27522 .arg(&mut d)
27523 .arg(&nc)
27524 .arg(&ep);
27525 unsafe {
27526 b.launch(cfg)?;
27527 }
27528 Ok((q, d))
27529 }
27530
27531 pub fn gated_rmsnorm_zv(
27536 &self,
27537 o: &CudaSlice<f32>,
27538 w: &CudaSlice<f32>,
27539 z: &cudarc::driver::CudaView<f32>,
27540 dst: &mut CudaSlice<f32>,
27541 ncols: usize,
27542 nrows: usize,
27543 eps: f32,
27544 ) -> Result<(), Box<dyn std::error::Error>> {
27545 let f = self.func("gated_rmsnorm_f32");
27546 let cfg = LaunchConfig {
27547 grid_dim: (nrows as u32, 1, 1),
27548 block_dim: (128, 1, 1),
27549 shared_mem_bytes: 0,
27550 };
27551 let (nc, e) = (ncols as i32, eps);
27552 let __s_b = self.gpu.stream();
27553 let mut b = __s_b.launch_builder(&f);
27554 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
27555 unsafe {
27556 b.launch(cfg)?;
27557 }
27558 Ok(())
27559 }
27560
27561 pub fn gated_rmsnorm_f16out_zv(
27562 &self,
27563 o: &CudaSlice<f32>,
27564 w: &CudaSlice<f32>,
27565 z: &cudarc::driver::CudaView<f32>,
27566 dst: &mut CudaSlice<f32>,
27567 dst16: &mut CudaSlice<u8>,
27568 ncols: usize,
27569 nrows: usize,
27570 eps: f32,
27571 ) -> Result<(), Box<dyn std::error::Error>> {
27572 let f = self.func("gated_rmsnorm_f16out_f32");
27573 let cfg = LaunchConfig {
27575 grid_dim: (nrows as u32, 1, 1),
27576 block_dim: (128, 1, 1),
27577 shared_mem_bytes: 0,
27578 };
27579 let (nc, e) = (ncols as i32, eps);
27580 let __s_b = self.gpu.stream();
27581 let mut b = __s_b.launch_builder(&f);
27582 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
27583 unsafe {
27584 b.launch(cfg)?;
27585 }
27586 Ok(())
27587 }
27588
27589 pub fn gated_rmsnorm_q8_1(
27590 &self,
27591 o: &CudaSlice<f32>,
27592 w: &CudaSlice<f32>,
27593 z: &CudaSlice<f32>,
27594 ncols: usize,
27595 nrows: usize,
27596 eps: f32,
27597 ) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
27598 assert!(ncols % 32 == 0);
27599 let f = self.func("gated_rmsnorm_q8_1");
27600 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
27601 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
27602 let cfg = LaunchConfig {
27603 grid_dim: (nrows as u32, 1, 1),
27604 block_dim: (128, 1, 1),
27605 shared_mem_bytes: 0,
27606 };
27607 let (nc, ep) = (ncols as i32, eps);
27608 let __s_b = self.gpu.stream();
27609 let mut b = __s_b.launch_builder(&f);
27610 b.arg(o)
27611 .arg(w)
27612 .arg(z)
27613 .arg(&mut out_q)
27614 .arg(&mut out_d)
27615 .arg(&nc)
27616 .arg(&ep);
27617 unsafe {
27618 b.launch(cfg)?;
27619 }
27620 Ok((out_q, out_d))
27621 }
27622
27623 pub fn transpose(
27625 &self,
27626 inp: &CudaSlice<f32>,
27627 rows: usize,
27628 cols: usize,
27629 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
27630 let f = self.func("transpose_f32");
27631 let mut out = self.zeros(rows * cols)?;
27632 let cfg = LaunchConfig::for_num_elems((rows * cols) as u32);
27633 let (r, c) = (rows as i32, cols as i32);
27634 let __s_b = self.gpu.stream();
27635 let mut b = __s_b.launch_builder(&f);
27636 b.arg(inp).arg(&mut out).arg(&r).arg(&c);
27637 unsafe {
27638 b.launch(cfg)?;
27639 }
27640 Ok(out)
27641 }
27642
27643 pub fn repeat_heads(
27645 &self,
27646 inp: &CudaSlice<f32>,
27647 out: &mut CudaSlice<f32>,
27648 head_dim: usize,
27649 n_in: usize,
27650 n_out: usize,
27651 t: usize,
27652 ) -> Result<(), Box<dyn std::error::Error>> {
27653 let f = self.func("repeat_heads_f32");
27654 let cfg = LaunchConfig::for_num_elems((head_dim * n_out * t) as u32);
27655 let (hd, ni, no, ti) = (head_dim as i32, n_in as i32, n_out as i32, t as i32);
27656 let __s_b = self.gpu.stream();
27657 let mut b = __s_b.launch_builder(&f);
27658 b.arg(inp).arg(out).arg(&hd).arg(&ni).arg(&no).arg(&ti);
27659 unsafe {
27660 b.launch(cfg)?;
27661 }
27662 Ok(())
27663 }
27664
27665 pub fn q_gate_split(
27672 &self,
27673 qf: &CudaSlice<f32>,
27674 q_out: &mut CudaSlice<f32>,
27675 gate_out: &mut CudaSlice<f32>,
27676 head_dim: usize,
27677 n_head: usize,
27678 t: usize,
27679 ) -> Result<(), Box<dyn std::error::Error>> {
27680 memra_gguf::config::check_fused_q_gate_extent(qf.len(), head_dim, n_head, t)?;
27681 let out_need = head_dim * n_head * t;
27682 if q_out.len() < out_need || gate_out.len() < out_need {
27683 return Err(format!(
27684 "q_gate_split destinations too small: need {out_need} each, have q={} gate={}",
27685 q_out.len(),
27686 gate_out.len()
27687 )
27688 .into());
27689 }
27690 let f = self.func("q_gate_split_f32");
27691 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
27692 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
27693 let __s_b = self.gpu.stream();
27694 let mut b = __s_b.launch_builder(&f);
27695 b.arg(qf)
27696 .arg(q_out)
27697 .arg(gate_out)
27698 .arg(&hd)
27699 .arg(&nh)
27700 .arg(&ti);
27701 unsafe {
27702 b.launch(cfg)?;
27703 }
27704 Ok(())
27705 }
27706
27707 pub fn qkv_to_gdn_repack(
27711 &self,
27712 conv_out: &CudaSlice<f32>,
27713 q_g: &mut CudaSlice<f32>,
27714 k_g: &mut CudaSlice<f32>,
27715 v_g: &mut CudaSlice<f32>,
27716 d_state: usize,
27717 num_v: usize,
27718 num_k: usize,
27719 key_dim: usize,
27720 t: usize,
27721 ) -> Result<(), Box<dyn std::error::Error>> {
27722 let f = self.func("qkv_to_gdn_repack_f32");
27723 let cfg = LaunchConfig::for_num_elems((d_state * num_v * t) as u32);
27724 let (ds, nv, nk, kd, ti) = (
27725 d_state as i32,
27726 num_v as i32,
27727 num_k as i32,
27728 key_dim as i32,
27729 t as i32,
27730 );
27731 let __s_b = self.gpu.stream();
27732 let mut b = __s_b.launch_builder(&f);
27733 b.arg(conv_out)
27734 .arg(q_g)
27735 .arg(k_g)
27736 .arg(v_g)
27737 .arg(&ds)
27738 .arg(&nv)
27739 .arg(&nk)
27740 .arg(&kd)
27741 .arg(&ti);
27742 unsafe {
27743 b.launch(cfg)?;
27744 }
27745 Ok(())
27746 }
27747
27748 pub fn conv_left_pad(
27751 &self,
27752 src: &CudaSlice<f32>,
27753 dst: &mut CudaSlice<f32>,
27754 conv_dim: usize,
27755 t: usize,
27756 pad: usize,
27757 ) -> Result<(), Box<dyn std::error::Error>> {
27758 let f = self.func("conv_left_pad_f32");
27759 let cfg = LaunchConfig::for_num_elems((conv_dim * t) as u32);
27760 let (cd, ti, p) = (conv_dim as i32, t as i32, pad as i32);
27761 let __s_b = self.gpu.stream();
27762 let mut b = __s_b.launch_builder(&f);
27763 b.arg(src).arg(dst).arg(&cd).arg(&ti).arg(&p);
27764 unsafe {
27765 b.launch(cfg)?;
27766 }
27767 Ok(())
27768 }
27769
27770 pub fn conv_assemble_and_roll(
27774 &self,
27775 qkv_col: &CudaSlice<f32>,
27776 conv_state: &mut CudaSlice<f32>,
27777 conv_in: &mut CudaSlice<f32>,
27778 conv_dim: usize,
27779 pad: usize,
27780 ) -> Result<(), Box<dyn std::error::Error>> {
27781 let f = self.func("conv_assemble_and_roll_f32");
27782 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
27783 let (cd, p) = (conv_dim as i32, pad as i32);
27784 let __s_b = self.gpu.stream();
27785 let mut b = __s_b.launch_builder(&f);
27786 b.arg(qkv_col).arg(conv_state).arg(conv_in).arg(&cd).arg(&p);
27787 unsafe {
27788 b.launch(cfg)?;
27789 }
27790 Ok(())
27791 }
27792
27793 pub fn ssm_conv1d_fused_decode(
27799 &self,
27800 qkv_col: &CudaSlice<f32>,
27801 conv_state: &mut CudaSlice<f32>,
27802 w: &CudaSlice<f32>,
27803 conv_out: &mut CudaSlice<f32>,
27804 conv_dim: usize,
27805 d_conv: usize,
27806 ) -> Result<(), Box<dyn std::error::Error>> {
27807 let f = self.func("ssm_conv1d_fused_decode_f32");
27808 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
27809 let (cd, dc) = (conv_dim as i32, d_conv as i32);
27810 let __s_b = self.gpu.stream();
27811 let mut b = __s_b.launch_builder(&f);
27812 b.arg(qkv_col)
27813 .arg(conv_state)
27814 .arg(w)
27815 .arg(conv_out)
27816 .arg(&cd)
27817 .arg(&dc);
27818 unsafe {
27819 b.launch(cfg)?;
27820 }
27821 Ok(())
27822 }
27823
27824 pub fn slice_range(
27827 &self,
27828 src: &CudaSlice<f32>,
27829 start: usize,
27830 len: usize,
27831 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
27832 let host = self.gpu.stream().clone_dtoh(src)?;
27833 self.gpu.stream().synchronize()?;
27834 Ok(self.htod(&host[start..start + len])?)
27835 }
27836}
27837
27838#[cfg(test)]
27839mod target_dispatch_tests {
27840 use super::legacy_quant_gemm_allowed;
27841
27842 #[test]
27843 fn legacy_quant_gemm_arch_policy_honors_the_escape_hatch() {
27844 assert!(legacy_quant_gemm_allowed(false, false, false));
27846 assert!(!legacy_quant_gemm_allowed(false, false, true));
27847 assert!(!legacy_quant_gemm_allowed(true, false, false));
27849 assert!(!legacy_quant_gemm_allowed(true, false, true));
27850 assert!(legacy_quant_gemm_allowed(true, true, false));
27852 assert!(!legacy_quant_gemm_allowed(true, true, true));
27853 }
27854
27855 #[cfg(all(memra_portable_cuda, not(memra_hopper_mma)))]
27856 #[test]
27857 fn portable_build_disables_legacy_quant_gemm_without_an_env_override() {
27858 assert!(!legacy_quant_gemm_allowed(
27859 cfg!(memra_portable_cuda),
27860 cfg!(memra_hopper_mma),
27861 false
27862 ));
27863 }
27864
27865 #[cfg(memra_hopper_mma)]
27866 #[test]
27867 fn hopper_mma_build_re_admits_legacy_quant_gemm() {
27868 assert!(legacy_quant_gemm_allowed(
27869 cfg!(memra_portable_cuda),
27870 cfg!(memra_hopper_mma),
27871 false
27872 ));
27873 assert!(super::portable_mma_gated() == false);
27874 }
27875}
27876
27877impl memra_kv::KvDev for Engine {
27880 fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
27881 Engine::zeros(self, n)
27882 }
27883 fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
27884 Engine::uninit(self, n)
27885 }
27886 fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
27887 Engine::alloc_u8(self, n)
27888 }
27889 fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
27890 Engine::htod_i32(self, v)
27891 }
27892 fn clone_dtod(
27893 &self,
27894 src: &CudaSlice<f32>,
27895 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
27896 Engine::clone_dtod(self, src)
27897 }
27898 fn copy_into(
27899 &self,
27900 dst: &mut CudaSlice<f32>,
27901 off: usize,
27902 src: &CudaSlice<f32>,
27903 len: usize,
27904 ) -> Result<(), Box<dyn std::error::Error>> {
27905 Engine::copy_into(self, dst, off, src, len)
27906 }
27907 fn set_i32_one(
27908 &self,
27909 d: &mut CudaSlice<i32>,
27910 v: i32,
27911 ) -> Result<(), Box<dyn std::error::Error>> {
27912 Engine::set_i32_one(self, d, v)
27913 }
27914}
27915
27916#[cfg(test)]
27917mod fused_gate_bounds_tests {
27918 use super::*;
27919
27920 #[test]
27933 #[ignore = "requires a CUDA GPU"]
27934 fn q_gate_split_refuses_a_separate_gate_wq_instead_of_reading_past_it() {
27935 let e = Engine::new(0).unwrap();
27936 let (head_dim, n_head, t) = (8usize, 4usize, 2usize);
27937 let fused = 2 * head_dim * n_head * t;
27938 let out_n = head_dim * n_head * t;
27939
27940 let narrow = e.htod(&vec![1.0f32; out_n]).unwrap();
27942 let mut q = e.uninit(out_n).unwrap();
27943 let mut gate = e.uninit(out_n).unwrap();
27944 let err = e
27945 .q_gate_split(&narrow, &mut q, &mut gate, head_dim, n_head, t)
27946 .expect_err("half-width wq must be refused, not read past")
27947 .to_string();
27948 assert!(err.contains("NO fused gate"), "{err}");
27949 assert!(err.contains(&format!("{fused}")), "{err}");
27950
27951 let host: Vec<f32> = (0..fused).map(|i| i as f32).collect();
27954 let wide = e.htod(&host).unwrap();
27955 e.q_gate_split(&wide, &mut q, &mut gate, head_dim, n_head, t)
27956 .expect("full-width wq splits");
27957 let (qh, gh) = (e.dtoh(&q).unwrap(), e.dtoh(&gate).unwrap());
27958 for tok in 0..t {
27959 for hh in 0..n_head {
27960 for d in 0..head_dim {
27961 let base = tok * (n_head * 2 * head_dim) + hh * (2 * head_dim);
27962 let idx = tok * (n_head * head_dim) + hh * head_dim + d;
27963 assert_eq!(qh[idx], host[base + d], "q t{tok} h{hh} d{d}");
27964 assert_eq!(gh[idx], host[base + head_dim + d], "gate t{tok} h{hh} d{d}");
27965 }
27966 }
27967 }
27968
27969 let mut small = e.uninit(out_n - 1).unwrap();
27971 assert!(
27972 e.q_gate_split(&wide, &mut small, &mut gate, head_dim, n_head, t)
27973 .is_err()
27974 );
27975 }
27976}
27977
27978#[cfg(test)]
27982mod fused_rope_width_tests {
27983 use super::Engine;
27984
27985 #[test]
27988 fn full_width_is_accepted() {
27989 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope", 256, 256).is_ok());
27990 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_cat", 512, 512).is_ok());
27991 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_append", 128, 128).is_ok());
27992 }
27993
27994 #[test]
28007 fn gemma4_official_artifact_widths_pass() {
28008 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope", 512, 512).is_ok());
28009 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_append_dc", 256, 256).is_ok());
28010 }
28011
28012 #[test]
28015 fn partial_rotary_is_refused_with_the_geometry_named() {
28016 let err = Engine::full_width_rope_only("rms_norm_qkv_rope", 64, 256)
28018 .expect_err("partial rotary must refuse");
28019 let msg = err.to_string();
28020 assert!(msg.contains("PARTIAL ROTARY REFUSED"), "{msg}");
28021 assert!(msg.contains("n_rot 64"), "{msg}");
28022 assert!(msg.contains("head_dim 256"), "{msg}");
28023 assert!(
28024 msg.contains("64..256"),
28025 "names the band it would corrupt: {msg}"
28026 );
28027 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope_append_dc", 64, 128).is_err());
28029 assert!(Engine::full_width_rope_only("rms_norm_qkv_rope", 256, 128).is_err());
28031 }
28032}