1use std::sync::{Arc, Mutex};
4use cudarc::driver::{CudaContext, CudaStream, CudaModule, CudaFunction, CudaSlice, LaunchConfig, PushKernelArg};
5use cudarc::nvrtc::Ptx;
6
7#[cfg(debug_assertions)]
8pub(crate) fn debug_assert_tensor_stream_device<T>(
9 tensor: &CudaSlice<T>,
10 stream: &CudaStream,
11 site: &str,
12) {
13 let tensor_dev = tensor.ordinal();
14 let stream_dev = stream.context().ordinal();
15 assert_eq!(
16 tensor_dev, stream_dev,
17 "PP cross-device tensor read at {site}: tensor on dev{tensor_dev}, stream on dev{stream_dev}"
18 );
19}
20
21pub use memra_gguf;
22pub use memra_runtime;
23
24pub mod model;
25pub mod forward;
26pub mod hybrid;
27pub mod hybrid_forward;
28pub mod sigrouter_contract;
29pub mod cache {
32 pub use memra_kv::*;
33}
34pub mod decode;
35pub mod decode_batch;
36pub mod mla;
40pub mod pp;
41pub mod spec;
42pub mod gemma_spec;
43pub mod round_stream;
44pub mod graph_update;
45pub mod dflash;
46pub mod eagle;
47pub use memra_sampling as sampler;
48
49pub fn moe_f16g_mode() -> u8 {
93 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
94 *M.get_or_init(|| match std::env::var("MEMRA_MOE_F16G").as_deref() {
95 Ok("0") => 0,
96 Ok("2") => 2,
97 Ok("3") => 3,
98 Ok(_) => 1,
99 Err(_) => 2,
102 })
103}
104pub fn moe_f16g_sk_params() -> (i32, i32) {
118 static P: std::sync::OnceLock<(i32, i32)> = std::sync::OnceLock::new();
119 *P.get_or_init(|| match std::env::var("MEMRA_F16G_SK").as_deref() {
120 Ok("0") => (-1, 0),
121 Ok("32") => (0, i32::MAX),
122 Ok("128") => (0, 1),
123 _ => {
124 let cross = std::env::var("MEMRA_F16G_SK_CROSS").ok()
125 .and_then(|v| v.parse().ok()).unwrap_or(64);
126 (0, cross)
127 }
128 })
129}
130pub fn moe_f16g_direct_on(qtype: i32) -> bool {
141 static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
142 let m = *M.get_or_init(|| match std::env::var("MEMRA_F16G_DIRECT").as_deref() {
143 Ok("0") => 0,
144 Ok("kq") => 1,
145 _ => 2,
146 });
147 match m {
148 0 => false,
149 1 => qtype == QT_Q4_K || qtype == QT_Q6_K,
150 _ => true,
151 }
152}
153pub fn moe_f16g_tail_on() -> bool {
162 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
163 *ON.get_or_init(|| std::env::var("MEMRA_F16G_TAIL").as_deref() != Ok("0"))
164}
165
166pub fn moe_f16g_gemma_on() -> bool {
173 static M: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
174 *M.get_or_init(|| !matches!(std::env::var("MEMRA_MOE_F16G").as_deref(), Ok("0") | Err(_)))
175}
176
177pub fn moe_fuse_actq_on() -> bool {
181 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
182 *ON.get_or_init(|| std::env::var("MEMRA_MOE_FUSE_ACTQ").as_deref() != Ok("0"))
183}
184
185pub fn router_prefill_exact_on() -> bool {
195 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
196 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_PREFILL_EXACT").as_deref() != Ok("0"))
197}
198
199pub fn router_kernel_on() -> bool {
200 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
201 *ON.get_or_init(|| {
202 let on = std::env::var("MEMRA_ROUTER_KERNEL").as_deref() != Ok("0");
203 if !on { eprintln!("[memra] router kernel OFF (rollback: per-column cuBLAS gemv)"); }
204 on
205 })
206}
207
208pub const ROUTER_BATCH_MIN_T: usize = 8;
223pub fn router_batch_on() -> bool {
224 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
225 *ON.get_or_init(|| std::env::var("MEMRA_ROUTER_BATCH").as_deref() != Ok("0"))
226}
227mod cpu_experts;
228pub mod moe_cache;
229pub mod spill;
230mod spill_pread;
231#[cfg(memra_cutlass)]
232pub mod cutlass_ffi;
233pub mod mmq_ffi;
234pub mod f16_ffi;
235pub mod prime_graph;
236pub mod fp8_ffi;
237
238const FATBIN: &[u8] = include_bytes!(env!("MEMRA_ENGINE_FATBIN"));
245const HYBRID_FATBIN: &[u8] = include_bytes!(env!("MEMRA_HYBRID_FATBIN"));
246const QMATVEC_FATBIN: &[u8] = include_bytes!(env!("MEMRA_QMATVEC_FATBIN"));
247const FLASH_FATBIN: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN"));
248const GEMM_FATBIN: &[u8] = include_bytes!(env!("MEMRA_GEMM_FATBIN"));
249const ROUTER_FATBIN: &[u8] = include_bytes!(env!("MEMRA_ROUTER_FATBIN"));
250const SAMPLE_FATBIN: &[u8] = include_bytes!(env!("MEMRA_SAMPLE_FATBIN"));
252
253fn gemm_fatbin_bytes() -> std::borrow::Cow<'static, [u8]> {
259 assert!(!(portable_mma_gated() && std::env::var_os("MEMRA_GEMM_FATBIN").is_some()),
260 "MEMRA_GEMM_FATBIN overrides are not allowed in the portable CUDA lane");
261 match std::env::var("MEMRA_GEMM_FATBIN") {
262 Ok(path) => std::borrow::Cow::Owned(
263 std::fs::read(&path).unwrap_or_else(|e| panic!("MEMRA_GEMM_FATBIN read {path}: {e}"))),
264 Err(_) => std::borrow::Cow::Borrowed(GEMM_FATBIN),
265 }
266}
267
268pub(crate) const fn portable_mma_gated() -> bool {
275 cfg!(memra_portable_cuda) && !cfg!(memra_hopper_mma)
276}
277
278const fn legacy_quant_gemm_allowed(portable_cuda: bool, hopper_mma: bool, no_gemm: bool) -> bool {
283 (!portable_cuda || hopper_mma) && !no_gemm
284}
285
286const FLASH_FATBIN_VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VQ4"));
294const FLASH_FATBIN_VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VF8"));
295const FLASH_FATBIN_KF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8"));
296const FLASH_FATBIN_KF8VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VQ4"));
297const FLASH_FATBIN_KF8VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VF8"));
298
299pub use memra_kv::{kv_blk_bytes, kv_cache_formats};
302
303fn flash_fatbin_bytes() -> &'static [u8] {
305 match kv_cache_formats() {
306 ("q8_0", "q5_1") => FLASH_FATBIN,
307 ("q8_0", "q4_0") => FLASH_FATBIN_VQ4,
308 ("q8_0", "fp8") => FLASH_FATBIN_VF8,
309 ("fp8", "q5_1") => FLASH_FATBIN_KF8,
310 ("fp8", "q4_0") => FLASH_FATBIN_KF8VQ4,
311 ("fp8", "fp8") => FLASH_FATBIN_KF8VF8,
312 other => unreachable!("kv_cache_formats returned {other:?}"),
313 }
314}
315
316fn k1_launch_override() -> Option<(u32, u32, u32)> {
323 static K1: std::sync::OnceLock<Option<(u32, u32, u32)>> = std::sync::OnceLock::new();
324 *K1.get_or_init(|| {
325 let v = std::env::var("MEMRA_GEMM_K1_LAUNCH").ok()?;
326 let p: Vec<u32> = v.split(',').filter_map(|s| s.trim().parse().ok()).collect();
327 match p.as_slice() { [bm, bn, w] => Some((*bm, *bn, *w)), _ => None }
328 })
329}
330
331pub(crate) fn wgmma_gemm_enabled() -> bool {
338 static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
339 *V.get_or_init(|| std::env::var("MEMRA_WGMMA").as_deref() == Ok("1"))
340}
341
342pub const FA_VEC_MIN_TKV: usize = 96;
357pub fn fa_vec_min_tkv() -> usize {
361 static V: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
362 *V.get_or_init(|| std::env::var("MEMRA_FA_VEC_MIN").ok()
363 .and_then(|v| v.parse().ok())
364 .unwrap_or_else(|| FA_VEC_MIN_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)))
365}
366
367pub fn fa_f16pv_on() -> bool {
378 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
379 *ON.get_or_init(|| std::env::var("MEMRA_FA_F16PV").map(|v| v != "0")
380 .unwrap_or_else(|_| std::env::var("MEMRA_DRAFT").is_err()))
381}
382
383pub fn fa512_hp_on() -> bool {
387 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
388 *ON.get_or_init(|| std::env::var("MEMRA_FA512_HP").as_deref() != Ok("0"))
389}
390
391pub fn faw_hp_on() -> bool {
395 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
396 *ON.get_or_init(|| std::env::var("MEMRA_FAW_HP").as_deref() != Ok("0"))
397}
398
399pub fn fa512_wide_warps() -> usize {
403 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
404 *N.get_or_init(|| match std::env::var("MEMRA_FA512_W4").as_deref() {
405 Ok("1") => 4, _ => 2,
406 })
407}
408
409pub fn fa512_min_tkv() -> usize {
412 static FA512_MIN: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
413 *FA512_MIN.get_or_init(|| std::env::var("MEMRA_FA512_MIN").ok()
414 .and_then(|v| v.parse().ok()).unwrap_or(512))
415}
416pub static FA_VEC_MIN_DEFAULT: std::sync::atomic::AtomicUsize =
420 std::sync::atomic::AtomicUsize::new(FA_VEC_MIN_TKV);
421pub static FA_SPW_DEFAULT: std::sync::atomic::AtomicUsize =
425 std::sync::atomic::AtomicUsize::new(32);
426pub static FUSED_MR1_DEFAULT: std::sync::atomic::AtomicBool =
432 std::sync::atomic::AtomicBool::new(false);
433pub static ROUTER_W8_DEFAULT: std::sync::atomic::AtomicBool =
440 std::sync::atomic::AtomicBool::new(true);
441pub static FA_SP512_DEFAULT: std::sync::atomic::AtomicUsize =
442 std::sync::atomic::AtomicUsize::new(16);
443pub static RMS_BLOCK_DEFAULT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(256);
448pub static FA_SP_GEMMA: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
450pub static MMQ_SK_FORCE: std::sync::atomic::AtomicI8 = std::sync::atomic::AtomicI8::new(-1);
455pub use memra_kv::KV_FP8_FORCE;
458pub(crate) fn rms_block() -> u32 {
459 static V: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
460 *V.get_or_init(|| std::env::var("MEMRA_RMS_BLOCK").ok()
461 .and_then(|v| v.parse().ok())
462 .unwrap_or_else(|| RMS_BLOCK_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)))
463}
464
465pub(crate) fn fa_split_keys(t_kv: usize, n_head_kv: usize) -> usize {
466 static S: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
467 if let Some(forced) = *S.get_or_init(|| {
468 std::env::var("MEMRA_FA_SPLIT").ok().and_then(|v| v.parse().ok())
469 .filter(|&s: &usize| s >= 8 && s % 8 == 0)
470 }) { return forced; }
471 if FA_SP_GEMMA.load(std::sync::atomic::Ordering::Relaxed)
489 && std::env::var("MEMRA_FA_SP16").as_deref() == Ok("1") {
490 return if t_kv <= 8192 { 16 } else if t_kv <= 16384 { 64 } else { 128 };
491 }
492 let big_rig = fa_sm_count() >= 128;
493 if big_rig {
494 let _ = n_head_kv;
495 if t_kv <= 2048 { 16 } else if t_kv <= 16384 { 64 } else { 128 }
496 } else if n_head_kv <= 4 {
497 if t_kv <= 512 { 8 } else if t_kv <= 16384 { 64 } else { 128 }
518 } else {
519 if t_kv <= 8192 { 32 } else if t_kv <= 16384 { 64 } else { 128 }
520 }
521}
522
523fn fa_sm_count() -> i32 {
526 static N: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
527 *N.get_or_init(|| {
528 cudarc::driver::result::init().ok();
529 cudarc::driver::result::device::get(0)
530 .and_then(|d| unsafe { cudarc::driver::result::device::get_attribute(
531 d, cudarc::driver::sys::CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT) })
532 .unwrap_or(82)
533 })
534}
535
536fn fa_hd_suffix(head_dim: usize) -> Result<&'static str, Box<dyn std::error::Error>> {
540 match head_dim {
541 256 => Ok(""),
542 128 => Ok("_hd128"),
543 d => Err(format!("fa_prefill: no kernel stamped for head_dim={d} (only 256/128); \
544 callers must gate to sdpa_naive").into()),
545 }
546}
547
548pub const QT_Q8_0: i32 = 0;
550pub const QT_Q4_K: i32 = 1;
551pub const QT_Q6_K: i32 = 2;
552pub const QT_Q5_K: i32 = 3;
553pub const QT_Q3_K: i32 = 4;
554pub const QT_IQ4_XS: i32 = 5;
555pub const QT_IQ3_S: i32 = 6;
556pub const QT_NVFP4: i32 = 7;
557pub const QT_F8_E4M3: i32 = 10;
563pub const QT_NVFP4_RP: i32 = 9;
566pub const QT_F32: i32 = 8;
568pub const QT_BF16: i32 = 11;
569pub const QT_Q4_0: i32 = 12; pub const QT_Q2_K: i32 = 13;
574pub const QT_F8_E4M3_BLK: i32 = 14;
590
591pub struct Engine {
593 pub gpu: memra_runtime::Gpu,
594 module: Arc<CudaModule>,
595 hybrid: Arc<CudaModule>,
596 qmatvec: Arc<CudaModule>,
597 flash: Arc<CudaModule>,
598 flash_g: std::sync::OnceLock<Arc<CudaModule>>,
602 gemm: Arc<CudaModule>,
603 router: Arc<CudaModule>,
604 sample: Arc<CudaModule>,
606 moe_cache: Mutex<Option<crate::moe_cache::MoeSlotCache>>,
610 moe_cache_layout: Mutex<Option<Vec<usize>>>,
614 capture_keep_on: std::sync::atomic::AtomicBool,
620 verify_exact: std::sync::atomic::AtomicBool,
625 capture_keep: Mutex<Vec<Box<dyn std::any::Any + Send>>>,
626 pub copy_stream: Arc<CudaStream>,
628 #[cfg(memra_cutlass)]
635 cutlass_scratch: Mutex<Option<crate::cutlass_ffi::CutlassScratch>>,
636 fp8_scratch: Mutex<Option<crate::fp8_ffi::Fp8Scratch>>,
640 fa_vf16_scratch: Mutex<Option<CudaSlice<u8>>>,
643 fa_part_pool: Mutex<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
647 fa_part_retired: Mutex<Vec<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
651 fn_cache: Mutex<std::collections::HashMap<String, CudaFunction>>,
653 f16_scratch: Mutex<Option<crate::f16_ffi::F16Scratch>>,
654 argmax_partials: Mutex<Option<(CudaSlice<f32>, CudaSlice<i32>)>>,
659 prime_deqw_ws: Mutex<Option<(CudaSlice<u8>, CudaSlice<u8>)>>,
664 router_stage: Mutex<Option<PinnedStage>>,
668}
669
670fn fa_v2_on() -> bool {
680 std::env::var("MEMRA_FA_V2").map(|v| v != "0").unwrap_or(true)
686}
687
688fn fa_v3_on() -> bool {
696 std::env::var("MEMRA_FA_V3").map(|v| v != "0").unwrap_or(true)
700}
701
702fn fa_v4_mode() -> &'static str {
707 static M: std::sync::OnceLock<String> = std::sync::OnceLock::new();
708 M.get_or_init(|| std::env::var("MEMRA_FA_V4").unwrap_or_default())
709}
710fn fa_v4_on() -> bool { fa_v4_mode() != "0" } pub static FA_SMEM_TKV_DEFAULT: std::sync::atomic::AtomicUsize =
719 std::sync::atomic::AtomicUsize::new(1024);
720pub static FA_V4_MAX_DEFAULT: std::sync::atomic::AtomicUsize =
721 std::sync::atomic::AtomicUsize::new(usize::MAX);
722pub fn fa_v4_at_pub(t_kv: usize) -> bool { fa_v4_at(t_kv) }
723fn fa_v4_at(t_kv: usize) -> bool {
724 static M: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
725 let mx = *M.get_or_init(|| std::env::var("MEMRA_FA_V4_MAX").ok()
726 .and_then(|v| v.parse().ok())
727 .unwrap_or_else(|| FA_V4_MAX_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)));
728 fa_v4_on() && t_kv < mx
729}
730pub const FA_DEEP_MIN_DEFAULT: usize = 0;
744fn fa_deep_at(t_kv: usize) -> bool {
745 if std::env::var("MEMRA_FA_DEEP").as_deref() == Ok("0") { return false; }
746 let min = std::env::var("MEMRA_FA_DEEP_MIN").ok().and_then(|v| v.parse().ok())
747 .unwrap_or(FA_DEEP_MIN_DEFAULT);
748 t_kv >= min
749}
750pub fn fa_deep_at_pub(t_kv: usize) -> bool { fa_deep_at(t_kv) }
752
753fn fa_v3_active(head_dim: usize) -> bool {
754 fa_v3_on() && head_dim % 128 == 0 && kv_cache_formats() == ("q8_0", "q5_1")
757 && !Engine::kv_fp8_on()
758}
759
760pub fn fa_seqs_eligible(t_kv: usize, head_dim: usize) -> bool {
768 std::env::var("MEMRA_NO_FA_VEC").is_err()
769 && t_kv >= fa_vec_min_tkv()
770 && head_dim == 256
771 && fa_v4_at(t_kv)
772 && !matches!(fa_v4_mode(), "noB3" | "stage")
773 && !Engine::kv_fp8_on()
774}
775pub fn fa_split_keys_pub(t_kv: usize, n_head_kv: usize) -> usize { fa_split_keys(t_kv, n_head_kv) }
777
778struct PinnedStage {
783 ptr: *mut u8,
784 cap: usize,
785}
786unsafe impl Send for PinnedStage {}
787impl PinnedStage {
788 fn new(cap: usize) -> Result<Self, Box<dyn std::error::Error>> {
789 let ptr = unsafe { cudarc::driver::result::malloc_host(cap, 0)? } as *mut u8;
790 Ok(PinnedStage { ptr, cap })
791 }
792}
793impl Drop for PinnedStage {
794 fn drop(&mut self) {
795 let _ = unsafe { cudarc::driver::result::free_host(self.ptr as _) };
796 }
797}
798
799pub const ARGMAX_NB: usize = 256;
802
803pub(crate) use memra_fa3_vl as fa3_vl_raw;
805
806unsafe extern "C" {
807 fn memra_fa3_prefill(q16: *const core::ffi::c_void, k16: *const core::ffi::c_void,
809 v16: *const core::ffi::c_void, o: *mut f32,
810 t: i32, h: i32, hkv: i32, d: i32, scale: f32,
811 stream: *mut core::ffi::c_void) -> i32;
812 pub(crate) fn memra_fa3_vl(q16s: *const *const core::ffi::c_void, k16s: *const *const core::ffi::c_void,
814 v16s: *const *const core::ffi::c_void, os: *const *mut f32,
815 ts: *const i32, b: i32, h: i32, hkv: i32, d: i32, scale: f32,
816 stream: *mut core::ffi::c_void) -> i32;
817}
818
819#[repr(C)]
824#[derive(Clone, Copy)]
825pub struct WPtr8(pub [u64; 8]);
826unsafe impl cudarc::driver::DeviceRepr for WPtr8 {}
827
828#[repr(C)]
833#[derive(Clone, Copy, Default)]
834pub struct GdnSeqVl {
835 pub kb16: u64, pub gcum: u64, pub beta: u64, pub u: u64, pub wb16: u64,
836 pub y: u64, pub ssnap: u64, pub state_in: u64, pub state_out: u64,
837 pub q: u64, pub p: u64, pub o: u64,
838 pub k: u64, pub v: u64, pub g: u64, pub a: u64, pub w: u64,
839 pub t: i32, pub nc: i32,
840}
841unsafe impl cudarc::driver::DeviceRepr for GdnSeqVl {}
842#[repr(C)]
843#[derive(Clone, Copy)]
844pub struct GdnVl8(pub [GdnSeqVl; 8]);
845unsafe impl cudarc::driver::DeviceRepr for GdnVl8 {}
846
847#[repr(C)]
850#[derive(Clone, Copy, Default)]
851pub struct GdnWVl { pub qb16: u64, pub pb16: u64 }
852unsafe impl cudarc::driver::DeviceRepr for GdnWVl {}
853#[repr(C)]
854#[derive(Clone, Copy)]
855pub struct GdnWVl8(pub [GdnWVl; 8]);
856unsafe impl cudarc::driver::DeviceRepr for GdnWVl8 {}
857
858#[repr(C)]
860#[derive(Clone, Copy, Default)]
861pub struct GdnPrepVl {
862 pub qkv: u64, pub conv_state: u64, pub conv_out: u64,
863 pub q_g: u64, pub k_g: u64, pub v_g: u64,
864 pub q_l2: u64, pub k_l2: u64,
865 pub beta_raw: u64, pub alpha: u64, pub beta: u64, pub g_log: u64,
866 pub o: u64, pub z: u64, pub gn: u64, pub gn16: u64,
867 pub kb16: u64,
868 pub qb16: u64,
869 pub t: i32, pub pad: i32,
870}
871unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl {}
872#[repr(C)]
873#[derive(Clone, Copy)]
874pub struct GdnPrepVl8(pub [GdnPrepVl; 8]);
875unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl8 {}
876
877#[repr(C)]
879#[derive(Clone, Copy, Default)]
880pub struct FaSeqVl {
881 pub q: u64, pub k16: u64, pub v16: u64, pub o: u64, pub kf: u64, pub vf: u64,
882 pub t: i32, pub pad: i32,
883}
884unsafe impl cudarc::driver::DeviceRepr for FaSeqVl {}
885#[repr(C)]
886#[derive(Clone, Copy)]
887pub struct FaVl8(pub [FaSeqVl; 8]);
888unsafe impl cudarc::driver::DeviceRepr for FaVl8 {}
889
890#[repr(C)]
892#[derive(Clone, Copy, Default)]
893pub struct AttnPreVl {
894 pub qf: u64, pub kf: u64, pub vf: u64,
895 pub q: u64, pub gate: u64, pub qn: u64, pub kn: u64,
896 pub kc: u64, pub vc: u64,
897 pub t: i32, pub pad: i32,
898}
899unsafe impl cudarc::driver::DeviceRepr for AttnPreVl {}
900#[repr(C)]
901#[derive(Clone, Copy)]
902pub struct AttnPreVl8(pub [AttnPreVl; 8]);
903unsafe impl cudarc::driver::DeviceRepr for AttnPreVl8 {}
904
905pub struct GdnChunkBufs {
908 pub gcum: CudaSlice<f32>,
909 pub a: CudaSlice<f32>,
910 pub p: CudaSlice<f32>,
911 pub u: CudaSlice<f32>,
912 pub w: CudaSlice<f32>,
913 pub kb16: CudaSlice<u8>,
914 pub wb16: CudaSlice<u8>,
915 pub y16: CudaSlice<u8>,
916 pub ssnap16: CudaSlice<u8>,
917 pub qb16: CudaSlice<u8>,
918 pub pb16: CudaSlice<u8>,
919 pub o: CudaSlice<f32>,
920 pub t: usize,
921 pub nc: usize,
922}
923
924#[repr(C)]
926#[derive(Clone, Copy)]
927pub struct F32x8(pub [f32; 8]);
928unsafe impl cudarc::driver::DeviceRepr for F32x8 {}
929
930pub static PRIME_NANOS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
934
935impl Engine {
936 pub fn new(ordinal: usize) -> Result<Self, Box<dyn std::error::Error>> {
937 let gpu = memra_runtime::Gpu::new(ordinal)?;
938 if std::env::var("MEMRA_ARCH_CHECK").as_deref() != Ok("0") {
942 use cudarc::driver::sys::CUdevice_attribute_enum as A;
943 let (maj, min) = cudarc::driver::result::device::get(ordinal as i32)
944 .and_then(|d| unsafe { Ok((
945 cudarc::driver::result::device::get_attribute(d, A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR)?,
946 cudarc::driver::result::device::get_attribute(d, A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR)?)) })
947 .unwrap_or((0, 0));
948 let built = env!("MEMRA_BUILT_CUDA_ARCH");
949 let ok = matches!((built, maj, min),
950 ("120a", 12, 0) | ("120a", 12, 1) | ("100a", 10, 0) | ("90a", 9, 0) | ("89", 8, 9));
951 if !ok {
952 return Err(format!(
953 "memra was built for sm_{built} but device {ordinal} reports compute \
954 capability {maj}.{min}. Rebuild on this machine (MEMRA_CUDA_ARCH \
955 auto-detects the GPU) or set MEMRA_ARCH_CHECK=0 to bypass.").into());
956 }
957 }
958 unsafe {
963 use cudarc::driver::sys;
964 let dev: sys::CUdevice = ordinal as sys::CUdevice;
965 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
966 if sys::cuDeviceGetDefaultMemPool(&mut pool, dev) == sys::CUresult::CUDA_SUCCESS {
967 let mut thresh: u64 = u64::MAX;
968 let _ = sys::cuMemPoolSetAttribute(
969 pool,
970 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
971 &mut thresh as *mut u64 as *mut core::ffi::c_void,
972 );
973 }
974 }
975 let module = gpu.ctx.load_module(Ptx::from_binary(FATBIN.to_vec()))?;
976 let hybrid = gpu.ctx.load_module(Ptx::from_binary(HYBRID_FATBIN.to_vec()))?;
977 let qmatvec = gpu.ctx.load_module(Ptx::from_binary(QMATVEC_FATBIN.to_vec()))?;
978 let flash = gpu.ctx.load_module(Ptx::from_binary(flash_fatbin_bytes().to_vec()))?;
979 let gemm = gpu.ctx.load_module(Ptx::from_binary(gemm_fatbin_bytes().into_owned()))?;
980 let router = gpu.ctx.load_module(Ptx::from_binary(ROUTER_FATBIN.to_vec()))?;
981 let sample = gpu.ctx.load_module(Ptx::from_binary(SAMPLE_FATBIN.to_vec()))?;
982 let copy_stream = gpu.ctx.new_stream()?;
983 if std::env::var("MEMRA_EVT").map(|v| v == "1").unwrap_or(false) {
999 } else {
1001 unsafe { gpu.ctx.disable_event_tracking(); }
1002 }
1003 Ok(Self { gpu, module, hybrid, qmatvec, flash, flash_g: std::sync::OnceLock::new(), gemm, router, sample,
1004 moe_cache: Mutex::new(None),
1005 moe_cache_layout: Mutex::new(None),
1006 copy_stream,
1007 capture_keep_on: std::sync::atomic::AtomicBool::new(false),
1008 verify_exact: std::sync::atomic::AtomicBool::new(false),
1009 capture_keep: Mutex::new(Vec::new()),
1010 argmax_partials: Mutex::new(None),
1011 prime_deqw_ws: Mutex::new(None),
1012 router_stage: Mutex::new(None),
1013 fp8_scratch: Mutex::new(None),
1014 fa_vf16_scratch: Mutex::new(None),
1015 fa_part_pool: Mutex::new(None),
1016 fa_part_retired: Mutex::new(Vec::new()),
1017 fn_cache: Mutex::new(Default::default()),
1018 f16_scratch: Mutex::new(None),
1019 #[cfg(memra_cutlass)]
1020 cutlass_scratch: Mutex::new(None) })
1021 }
1022
1023 pub fn ctx(&self) -> &Arc<CudaContext> { &self.gpu.ctx }
1024
1025 pub fn pool_cached_bytes(&self) -> usize {
1043 let (reserved, used) = self.pool_reserved_used();
1044 reserved.saturating_sub(used)
1045 }
1046
1047 pub fn pool_reserved_used(&self) -> (usize, usize) {
1054 use cudarc::driver::sys;
1055 unsafe {
1056 let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
1057 if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
1058 != sys::CUresult::CUDA_SUCCESS
1059 {
1060 return (0, 0);
1061 }
1062 let (mut reserved, mut used) = (0u64, 0u64);
1063 if sys::cuMemPoolGetAttribute(
1064 pool,
1065 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RESERVED_MEM_CURRENT,
1066 &mut reserved as *mut u64 as *mut core::ffi::c_void,
1067 ) != sys::CUresult::CUDA_SUCCESS {
1068 return (0, 0);
1069 }
1070 if sys::cuMemPoolGetAttribute(
1071 pool,
1072 sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_USED_MEM_CURRENT,
1073 &mut used as *mut u64 as *mut core::ffi::c_void,
1074 ) != sys::CUresult::CUDA_SUCCESS {
1075 return (0, 0);
1076 }
1077 (reserved as usize, used as usize)
1078 }
1079 }
1080
1081 pub fn stream(&self) -> Arc<CudaStream> { self.gpu.stream() }
1084 pub fn gkv_on() -> bool {
1087 memra_kv::gkv_on()
1088 }
1089
1090 pub fn wkv_on() -> bool {
1102 memra_kv::wkv_on()
1103 }
1104
1105 pub fn kv_fp8_on() -> bool {
1111 memra_kv::kv_fp8_on()
1112 }
1113
1114 fn fa_func(&self, name: &str, head_dim: usize) -> CudaFunction {
1117 if head_dim == 512 && Self::gkv_on() { self.func_g(name) } else { self.func(name) }
1118 }
1119
1120 fn func_g(&self, name: &str) -> CudaFunction {
1124 let m = self.flash_g.get_or_init(|| {
1125 self.gpu.ctx.load_module(cudarc::nvrtc::Ptx::from_binary(FLASH_FATBIN_KF8VF8.to_vec()))
1126 .expect("load kf8vf8 flash fatbin (fp8-globals arm)")
1127 });
1128 let key = format!("g:{name}");
1129 if let Some(f) = self.fn_cache.lock().unwrap().get(&key) { return f.clone(); }
1130 let f = match m.load_function(name) {
1131 Ok(f) => f,
1132 Err(_) => self.func(name),
1133 };
1134 self.fn_cache.lock().unwrap().insert(key, f.clone());
1135 f
1136 }
1137
1138 fn func(&self, name: &str) -> CudaFunction {
1139 if let Some(f) = self.fn_cache.lock().unwrap().get(name) { return f.clone(); }
1142 let f = self.module.load_function(name)
1143 .or_else(|_| self.hybrid.load_function(name))
1144 .or_else(|_| self.qmatvec.load_function(name))
1145 .or_else(|_| self.flash.load_function(name))
1146 .or_else(|_| self.gemm.load_function(name))
1147 .or_else(|_| self.router.load_function(name))
1148 .or_else(|_| self.sample.load_function(name))
1149 .unwrap_or_else(|_| panic!("kernel {name} not in any fatbin"));
1150 self.fn_cache.lock().unwrap().insert(name.to_string(), f.clone());
1151 f
1152 }
1153
1154 pub fn scatter_trim_logits(&self, src: &CudaSlice<f32>, d2t: &CudaSlice<u32>,
1157 dst: &mut CudaSlice<f32>, d_vocab: usize, n_vocab: usize)
1158 -> Result<(), Box<dyn std::error::Error>> {
1159 let f1 = self.func("scatter_trim_logits_f32");
1160 let f2 = self.func("scatter_trim_logits_pass2_f32");
1161 let (dv, nv) = (d_vocab as i32, n_vocab as i32);
1162 let cfg1 = LaunchConfig { grid_dim: (256, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1163 let __s_b1 = self.gpu.stream();
1164 let mut b1 = __s_b1.launch_builder(&f1);
1165 b1.arg(src).arg(d2t).arg(&mut *dst).arg(&dv).arg(&nv);
1166 unsafe { b1.launch(cfg1)?; }
1167 let cfg2 = LaunchConfig { grid_dim: (d_vocab.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1168 let __s_b2 = self.gpu.stream();
1169 let mut b2 = __s_b2.launch_builder(&f2);
1170 b2.arg(src).arg(d2t).arg(&mut *dst).arg(&dv);
1171 unsafe { b2.launch(cfg2)?; }
1172 Ok(())
1173 }
1174
1175 #[allow(clippy::too_many_arguments)]
1181 pub fn filter_stats(&self, x: &CudaSlice<f32>, row_stride: usize, rows: &CudaSlice<i32>,
1182 out_th: &mut CudaSlice<f32>, out_z: &mut CudaSlice<f32>,
1183 out_max: &mut CudaSlice<f32>, n: usize, nrow: usize,
1184 temp: f32, top_k: i32, top_p: f32, min_p: f32)
1185 -> Result<(), Box<dyn std::error::Error>> {
1186 let f = self.func("filter_stats_f32");
1187 let (ni, nr, rs) = (n as i32, nrow as i32, row_stride as i64);
1188 let cfg = LaunchConfig { grid_dim: (nrow as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
1189 let __s_b = self.gpu.stream();
1190 let mut b = __s_b.launch_builder(&f);
1191 b.arg(x).arg(&rs).arg(rows).arg(&mut *out_th).arg(&mut *out_z).arg(&mut *out_max)
1192 .arg(&ni).arg(&nr).arg(&temp).arg(&top_k).arg(&top_p).arg(&min_p);
1193 unsafe { b.launch(cfg)?; }
1194 Ok(())
1195 }
1196
1197 #[allow(clippy::too_many_arguments)]
1199 pub fn softmax_gather_filtered(&self, x: &CudaSlice<f32>, row_stride: usize,
1200 ids: &CudaSlice<u32>, rows: &CudaSlice<i32>,
1201 th: &CudaSlice<f32>, z: &CudaSlice<f32>,
1202 out: &mut CudaSlice<f32>, n: usize, npair: usize, temp: f32)
1203 -> Result<(), Box<dyn std::error::Error>> {
1204 let f = self.func("softmax_gather_filtered_f32");
1205 let (ni, np, rs) = (n as i32, npair as i32, row_stride as i64);
1206 let cfg = LaunchConfig { grid_dim: (npair as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1207 let __s_b = self.gpu.stream();
1208 let mut b = __s_b.launch_builder(&f);
1209 b.arg(x).arg(&rs).arg(ids).arg(rows).arg(th).arg(z).arg(&mut *out).arg(&ni).arg(&np).arg(&temp);
1210 unsafe { b.launch(cfg)?; }
1211 Ok(())
1212 }
1213
1214 #[allow(clippy::too_many_arguments)]
1216 pub fn residual_sample_filtered(&self, p: &CudaSlice<f32>, q: Option<&CudaSlice<f32>>, n: usize,
1217 temp: f32, seed: u64, stream_pos: u32,
1218 p_stats: (f32, f32, f32), q_stats: (f32, f32, f32),
1219 out_tok: &mut CudaSlice<u32>)
1220 -> Result<(), Box<dyn std::error::Error>> {
1221 let f = self.func("residual_sample_filtered_f32");
1222 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1223 let has_q: i32 = q.is_some() as i32;
1224 let qbuf = q.unwrap_or(p);
1225 let (pm, pth, pz) = p_stats; let (qm, qth, qz) = q_stats;
1226 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
1227 let __s_b = self.gpu.stream();
1228 let mut b = __s_b.launch_builder(&f);
1229 b.arg(p).arg(qbuf).arg(&has_q).arg(&ni).arg(&temp).arg(&slo).arg(&shi).arg(&stream_pos)
1230 .arg(&pm).arg(&pth).arg(&pz).arg(&qm).arg(&qth).arg(&qz).arg(&mut *out_tok);
1231 unsafe { b.launch(cfg)?; }
1232 Ok(())
1233 }
1234
1235 #[allow(clippy::too_many_arguments)]
1237 pub fn gumbel_perturb_filtered(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
1238 seed: u64, stream_pos: u32, temp: f32, row_max: f32, th: f32)
1239 -> Result<(), Box<dyn std::error::Error>> {
1240 let f = self.func("gumbel_perturb_filtered_f32");
1241 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
1242 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1243 let __s_b = self.gpu.stream();
1244 let mut b = __s_b.launch_builder(&f);
1245 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp).arg(&row_max).arg(&th);
1246 unsafe { b.launch(cfg)?; }
1247 Ok(())
1248 }
1249
1250 #[allow(clippy::too_many_arguments)]
1254 pub fn penalize_logits(&self, x: &mut CudaSlice<f32>, hist: &CudaSlice<u32>, n_hist: usize,
1255 rep: f32, freq: f32, present: f32, n: usize)
1256 -> Result<(), Box<dyn std::error::Error>> {
1257 if n_hist == 0 { return Ok(()); }
1258 let f = self.func("penalize_logits_f32");
1259 let (nh, ni) = (n_hist as i32, n as i32);
1260 let cfg = LaunchConfig { grid_dim: (n_hist.div_ceil(128) as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
1261 let __s_b = self.gpu.stream();
1262 let mut b = __s_b.launch_builder(&f);
1263 b.arg(&mut *x).arg(hist).arg(&nh).arg(&rep).arg(&freq).arg(&present).arg(&ni);
1264 unsafe { b.launch(cfg)?; }
1265 Ok(())
1266 }
1267
1268 #[allow(clippy::too_many_arguments)]
1270 pub fn penalize_logits_rows(&self, x: &mut CudaSlice<f32>, hist: &CudaSlice<u32>, n_hist: usize,
1271 rep: f32, freq: f32, present: f32, n: usize, nrow: usize)
1272 -> Result<(), Box<dyn std::error::Error>> {
1273 if n_hist == 0 || nrow == 0 { return Ok(()); }
1274 let f = self.func("penalize_logits_rows_f32");
1275 let (nh, ni, nr) = (n_hist as i32, n as i32, nrow as i32);
1276 let cfg = LaunchConfig { grid_dim: (n_hist.div_ceil(128) as u32, nrow as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
1277 let __s_b = self.gpu.stream();
1278 let mut b = __s_b.launch_builder(&f);
1279 b.arg(&mut *x).arg(hist).arg(&nh).arg(&rep).arg(&freq).arg(&present).arg(&ni).arg(&nr);
1280 unsafe { b.launch(cfg)?; }
1281 Ok(())
1282 }
1283
1284 pub fn wpf_level() -> u32 {
1292 static ON: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
1293 *ON.get_or_init(|| std::env::var("MEMRA_WPF").ok()
1294 .and_then(|v| v.parse().ok()).unwrap_or(1))
1295 }
1296
1297 pub fn set_verify_exact(&self, on: bool) {
1309 self.verify_exact.store(on, std::sync::atomic::Ordering::Relaxed);
1310 }
1311 pub(crate) fn verify_exact_on(&self) -> bool {
1312 self.verify_exact.load(std::sync::atomic::Ordering::Relaxed)
1313 }
1314
1315 pub fn qkv_append_on() -> bool {
1318 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1319 *ON.get_or_init(|| std::env::var("MEMRA_QKV_APPEND").map(|v| v != "0").unwrap_or(true))
1320 }
1321
1322 pub fn pdl_wb_on() -> bool {
1325 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1326 *ON.get_or_init(|| std::env::var("MEMRA_PDL_WB").map(|v| v != "0").unwrap_or(true))
1327 }
1328
1329 pub fn pdl_mmvq_on() -> bool {
1333 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1334 *ON.get_or_init(|| std::env::var("MEMRA_PDL_MMVQ").map(|v| v != "0").unwrap_or(true))
1335 }
1336
1337 pub fn pdl_on() -> bool {
1338 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1339 *ON.get_or_init(|| std::env::var("MEMRA_PDL").map(|v| v != "0").unwrap_or(true))
1340 }
1341
1342 fn q40_mr1_on() -> bool {
1348 static Q40MR: std::sync::OnceLock<Option<u32>> = std::sync::OnceLock::new();
1349 match *Q40MR.get_or_init(|| std::env::var("MEMRA_Q40_MR").ok()
1350 .and_then(|v| v.parse().ok())) {
1351 Some(v) => v == 1,
1352 None => crate::FUSED_MR1_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
1353 }
1354 }
1355
1356 fn pdl_func_flash(&self, g: bool, name: &'static str)
1361 -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
1362 use cudarc::driver::sys as cu;
1363 static MODS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool), usize>>> =
1370 std::sync::Mutex::new(None);
1371 static FNS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool, &'static str), usize>>> =
1372 std::sync::Mutex::new(None);
1373 let ctx_key = self.ctx().cu_ctx() as usize;
1374 if let Some(&f) = FNS.lock().unwrap().get_or_insert_with(Default::default)
1375 .get(&(ctx_key, g, name)) { return Ok(f as cu::CUfunction); }
1376 let module = {
1377 let mut mods = MODS.lock().unwrap();
1378 let map = mods.get_or_insert_with(Default::default);
1379 match map.get(&(ctx_key, g)) {
1380 Some(&m) => m,
1381 None => {
1382 let m = self.pdl_load_module_in_ctx(
1383 if g { FLASH_FATBIN_KF8VF8 } else { FLASH_FATBIN })?;
1384 map.insert((ctx_key, g), m);
1385 m
1386 }
1387 }
1388 };
1389 let cname = std::ffi::CString::new(name)?;
1390 let mut f: cu::CUfunction = std::ptr::null_mut();
1391 let r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
1392 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("pdl_func_flash {name} (g={g}): {r:?}").into()); }
1393 FNS.lock().unwrap().get_or_insert_with(Default::default)
1394 .insert((ctx_key, g, name), f as usize);
1395 Ok(f)
1396 }
1397
1398 fn pdl_load_module_in_ctx(&self, bytes: &[u8]) -> Result<usize, Box<dyn std::error::Error>> {
1403 use cudarc::driver::sys as cu;
1404 let mut prev: cu::CUcontext = std::ptr::null_mut();
1405 unsafe { cu::cuCtxGetCurrent(&mut prev).result()?; }
1406 self.ctx().bind_to_thread()?;
1407 let mut m: cu::CUmodule = std::ptr::null_mut();
1408 let r = unsafe { cu::cuModuleLoadData(&mut m, bytes.as_ptr() as *const std::ffi::c_void) };
1409 let restore = if prev.is_null() { cu::CUresult::CUDA_SUCCESS }
1410 else { unsafe { cu::cuCtxSetCurrent(prev) } };
1411 if r != cu::CUresult::CUDA_SUCCESS {
1412 return Err(format!("pdl module load: {r:?}").into());
1413 }
1414 if restore != cu::CUresult::CUDA_SUCCESS {
1415 return Err(format!("pdl module load: ctx restore {restore:?}").into());
1416 }
1417 Ok(m as usize)
1418 }
1419
1420 fn pdl_func(&self, name: &'static str) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
1421 use cudarc::driver::sys as cu;
1422 static MODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
1425 std::sync::Mutex::new(None);
1426 static QMODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
1429 std::sync::Mutex::new(None);
1430 static FNS: std::sync::Mutex<Option<std::collections::HashMap<(usize, &'static str), usize>>> =
1431 std::sync::Mutex::new(None);
1432 let ctx_key = self.ctx().cu_ctx() as usize;
1433 if let Some(&f) = FNS.lock().unwrap().get_or_insert_with(Default::default)
1434 .get(&(ctx_key, name)) { return Ok(f as cu::CUfunction); }
1435 let module = {
1436 let mut mods = MODULES.lock().unwrap();
1437 let map = mods.get_or_insert_with(Default::default);
1438 match map.get(&ctx_key) {
1439 Some(&m) => m,
1440 None => {
1441 let m = self.pdl_load_module_in_ctx(FATBIN)?;
1442 map.insert(ctx_key, m);
1443 m
1444 }
1445 }
1446 };
1447 let cname = std::ffi::CString::new(name)?;
1448 let mut f: cu::CUfunction = std::ptr::null_mut();
1449 let mut r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
1450 if r == cu::CUresult::CUDA_ERROR_NOT_FOUND {
1451 let qmodule = {
1452 let mut mods = QMODULES.lock().unwrap();
1453 let map = mods.get_or_insert_with(Default::default);
1454 match map.get(&ctx_key) {
1455 Some(&m) => m,
1456 None => {
1457 let m = self.pdl_load_module_in_ctx(QMATVEC_FATBIN)?;
1458 map.insert(ctx_key, m);
1459 m
1460 }
1461 }
1462 };
1463 r = unsafe { cu::cuModuleGetFunction(&mut f, qmodule as cu::CUmodule, cname.as_ptr()) };
1464 }
1465 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("pdl_func {name}: {r:?}").into()); }
1466 FNS.lock().unwrap().get_or_insert_with(Default::default)
1467 .insert((ctx_key, name), f as usize);
1468 Ok(f)
1469 }
1470
1471 unsafe fn launch_pdl_flash(&self, g: bool, name: &'static str, grid: (u32, u32, u32),
1483 block: (u32, u32, u32), smem: u32,
1484 params: &mut [*mut std::ffi::c_void])
1485 -> Result<(), Box<dyn std::error::Error>> {
1486 use cudarc::driver::sys as cu;
1487 let f = self.pdl_func_flash(g, name)?;
1488 if smem > 0 {
1489 let r = unsafe { cu::cuFuncSetAttribute(f,
1491 cu::CUfunction_attribute_enum::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
1492 smem as i32) };
1493 if r != cu::CUresult::CUDA_SUCCESS {
1494 return Err(format!("pdl smem attr {name}: {r:?}").into());
1495 }
1496 }
1497 let mut attr = cu::CUlaunchAttribute {
1498 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
1499 pad: [0; 4],
1500 value: cu::CUlaunchAttributeValue { programmaticStreamSerializationAllowed: 1 },
1501 };
1502 let cfg = cu::CUlaunchConfig {
1503 gridDimX: grid.0, gridDimY: grid.1, gridDimZ: grid.2,
1504 blockDimX: block.0, blockDimY: block.1, blockDimZ: block.2,
1505 sharedMemBytes: smem, hStream: self.gpu.stream().cu_stream(),
1506 attrs: &mut attr, numAttrs: 1,
1507 };
1508 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
1509 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("launch_pdl_flash {name}: {r:?}").into()); }
1510 Ok(())
1511 }
1512
1513 unsafe fn launch_pdl(&self, name: &'static str, grid: (u32, u32, u32), block: (u32, u32, u32),
1514 params: &mut [*mut std::ffi::c_void])
1515 -> Result<(), Box<dyn std::error::Error>> {
1516 use cudarc::driver::sys as cu;
1517 let f = self.pdl_func(name)?;
1518 let mut attr = cu::CUlaunchAttribute {
1519 id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
1520 pad: [0; 4],
1521 value: cu::CUlaunchAttributeValue { programmaticStreamSerializationAllowed: 1 },
1522 };
1523 let cfg = cu::CUlaunchConfig {
1524 gridDimX: grid.0, gridDimY: grid.1, gridDimZ: grid.2,
1525 blockDimX: block.0, blockDimY: block.1, blockDimZ: block.2,
1526 sharedMemBytes: 0, hStream: self.gpu.stream().cu_stream(),
1527 attrs: &mut attr, numAttrs: 1,
1528 };
1529 let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
1530 if r != cu::CUresult::CUDA_SUCCESS { return Err(format!("launch_pdl {name}: {r:?}").into()); }
1531 Ok(())
1532 }
1533
1534 pub fn prefetch_weight_l2(&self, w: &crate::model::GpuTensor)
1537 -> Result<(), Box<dyn std::error::Error>> {
1538 if let crate::model::GpuTensor::Quant { bytes, rp4, .. } = w {
1539 let p = rp4.as_ref().unwrap_or(bytes);
1540 self.prefetch_l2(p, p.len())?;
1541 }
1542 Ok(())
1543 }
1544
1545 pub fn gather_row_bf16(&self, table: &CudaSlice<u8>, tok: &CudaSlice<u32>, idx: usize,
1548 dst: &mut CudaSlice<f32>, ncols: usize)
1549 -> Result<(), Box<dyn std::error::Error>> {
1550 let f = self.func("gather_row_bf16_f32");
1551 let cfg = LaunchConfig { grid_dim: (ncols.div_ceil(256) as u32, 1, 1),
1552 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1553 let (nc, ix) = (ncols as i32, idx as i32);
1554 let __s_b = self.gpu.stream();
1555 let mut b = __s_b.launch_builder(&f);
1556 b.arg(table).arg(tok).arg(&ix).arg(dst).arg(&nc);
1557 unsafe { b.launch(cfg)?; }
1558 Ok(())
1559 }
1560
1561 pub fn add_row_inplace(&self, logits: &mut CudaSlice<f32>, bias: &CudaSlice<f32>,
1563 n: usize, row_off: usize)
1564 -> Result<(), Box<dyn std::error::Error>> {
1565 let f = self.func("add_row_inplace_f32");
1566 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1),
1567 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1568 let (ni, off) = (n as i32, row_off as i64);
1569 let __s_b = self.gpu.stream();
1570 let mut b = __s_b.launch_builder(&f);
1571 b.arg(logits).arg(bias).arg(&ni).arg(&off);
1572 unsafe { b.launch(cfg)?; }
1573 Ok(())
1574 }
1575
1576 pub fn prefetch_l2(&self, p: &CudaSlice<u8>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
1578 let f = self.func("prefetch_l2_bytes");
1579 let lines = n.div_ceil(128);
1580 let ni = n as i64;
1581 let cfg = LaunchConfig { grid_dim: (lines.div_ceil(256) as u32, 1, 1),
1582 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1583 let __s_b = self.gpu.stream();
1584 let mut b = __s_b.launch_builder(&f);
1585 b.arg(p).arg(&ni);
1586 unsafe { b.launch(cfg)?; }
1587 Ok(())
1588 }
1589
1590 pub fn router_gemv(&self, w: &CudaSlice<f32>, x: &CudaSlice<f32>, n_embd: usize,
1593 n_experts: usize, t: usize)
1594 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1595 let w8 = match std::env::var("MEMRA_ROUTER_V2").as_deref() {
1601 Ok("0") => false,
1602 Ok(_) => true,
1603 Err(_) => ROUTER_W8_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
1604 };
1605 let batch = w8 && t >= ROUTER_BATCH_MIN_T && router_batch_on();
1615 self.router_gemv_form(w, x, n_embd, n_experts, t, w8, batch)
1616 }
1617
1618 pub fn router_gemv_form(&self, w: &CudaSlice<f32>, x: &CudaSlice<f32>, n_embd: usize,
1621 n_experts: usize, t: usize, w8: bool, batch: bool)
1622 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1623 debug_assert!(!batch || w8, "batch twin exists for the w8 form only");
1624 let mut y = self.alloc_uninit::<f32>(t * n_experts)?;
1625 let f = if batch { self.func("router_gemv_f32_w8_batch") }
1626 else if w8 { self.func("router_gemv_f32_w8") }
1627 else { self.func("router_gemv_f32") };
1628 let (ne, nx, ti) = (n_embd as i32, n_experts as i32, t as i32);
1629 let cfg = if batch {
1630 LaunchConfig { grid_dim: (n_experts.div_ceil(8) as u32, t.div_ceil(8) as u32, 1),
1631 block_dim: (32, 8, 1), shared_mem_bytes: 0 }
1632 } else {
1633 LaunchConfig { grid_dim: (n_experts as u32, t as u32, 1),
1634 block_dim: (32, if w8 { 8 } else { 1 }, 1), shared_mem_bytes: 0 }
1635 };
1636 let __s_b = self.gpu.stream();
1637 let mut b = __s_b.launch_builder(&f);
1638 b.arg(w).arg(x).arg(&mut y).arg(&ne).arg(&nx).arg(&ti);
1639 unsafe { b.launch(cfg)?; }
1640 Ok(y)
1641 }
1642
1643 pub fn rows_permute(&self, src: &CudaSlice<f32>, idx: &CudaSlice<i32>, nrows: usize,
1645 ncols: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1646 let mut dst = self.alloc_uninit::<f32>(nrows * ncols)?;
1647 let f = self.func("rows_permute_f32");
1648 let (nc, nr) = (ncols as i32, nrows as i32);
1649 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (256, 1, 1),
1650 shared_mem_bytes: 0 };
1651 let __s_b = self.gpu.stream();
1652 let mut b = __s_b.launch_builder(&f);
1653 b.arg(src).arg(idx).arg(&mut dst).arg(&nc).arg(&nr);
1654 unsafe { b.launch(cfg)?; }
1655 Ok(dst)
1656 }
1657
1658 pub fn sigmoid_dot_rows(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, n_embd: usize,
1663 t: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1664 static OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1667 if *OFF.get_or_init(|| std::env::var("MEMRA_SHEXP_DOT").as_deref() == Ok("0")) {
1668 let gs = self.linear(x, w, t, n_embd, 1)?;
1669 let mut g = self.uninit(t)?;
1670 self.sigmoid(&gs, &mut g, t)?;
1671 return Ok(g);
1672 }
1673 let mut g = self.alloc_uninit::<f32>(t)?;
1679 let f = self.func("sigmoid_dot_rows_f32");
1680 let (ne, ti) = (n_embd as i32, t as i32);
1681 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (32, 8, 1),
1682 shared_mem_bytes: 0 };
1683 let __s_b = self.gpu.stream();
1684 let mut b = __s_b.launch_builder(&f);
1685 b.arg(x).arg(w).arg(&mut g).arg(&ne).arg(&ti);
1686 unsafe { b.launch(cfg)?; }
1687 Ok(g)
1688 }
1689
1690 pub fn spec_rollback_stream(&self, len_ptrs: &CudaSlice<u64>, pos_start: &CudaSlice<i32>,
1692 acc: &CudaSlice<u32>, base: usize, n_rows: usize)
1693 -> Result<(), Box<dyn std::error::Error>> {
1694 let f = self.func("spec_rollback_stream");
1695 let (b, nr) = (base as i32, n_rows as i32);
1696 let cfg = LaunchConfig { grid_dim: (n_rows.div_ceil(64) as u32, 1, 1),
1697 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1698 let __s_bl = self.gpu.stream();
1699 let mut bl = __s_bl.launch_builder(&f);
1700 bl.arg(len_ptrs).arg(pos_start).arg(acc).arg(&b).arg(&nr);
1701 unsafe { bl.launch(cfg)?; }
1702 Ok(())
1703 }
1704
1705 pub fn plain_tok_ring(&self, vam: &CudaSlice<u32>, pos_start: &CudaSlice<i32>,
1707 base: usize, ring: &mut CudaSlice<u32>)
1708 -> Result<(), Box<dyn std::error::Error>> {
1709 let f = self.func("plain_tok_ring");
1710 let (b, cap) = (base as i32, ring.len() as i32);
1711 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1712 let __s_bl = self.gpu.stream();
1713 let mut bl = __s_bl.launch_builder(&f);
1714 bl.arg(vam).arg(pos_start).arg(&b).arg(&mut *ring).arg(&cap);
1715 unsafe { bl.launch(cfg)?; }
1716 Ok(())
1717 }
1718
1719 pub fn spec_ring_commit(&self, vtok: &CudaSlice<u32>, acc: &CudaSlice<u32>,
1721 brk: &CudaSlice<u32>, ring: &mut CudaSlice<u32>,
1722 pend: &mut CudaSlice<u32>)
1723 -> Result<(), Box<dyn std::error::Error>> {
1724 let f = self.func("spec_ring_commit");
1725 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1726 let __s_b = self.gpu.stream();
1727 let mut b = __s_b.launch_builder(&f);
1728 b.arg(vtok).arg(acc).arg(brk).arg(ring).arg(pend);
1729 unsafe { b.launch(cfg)?; }
1730 Ok(())
1731 }
1732 pub fn i32_copy_add(&self, src: &CudaSlice<i32>, dst: &mut CudaSlice<i32>, delta: i32)
1733 -> Result<(), Box<dyn std::error::Error>> {
1734 let f = self.func("i32_copy_add");
1735 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1736 let __s_b = self.gpu.stream();
1737 let mut b = __s_b.launch_builder(&f);
1738 b.arg(src).arg(dst).arg(&delta);
1739 unsafe { b.launch(cfg)?; }
1740 Ok(())
1741 }
1742 pub fn u32_copy(&self, src: &CudaSlice<u32>, dst: &mut CudaSlice<u32>)
1743 -> Result<(), Box<dyn std::error::Error>> {
1744 let f = self.func("u32_copy");
1745 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1746 let __s_b = self.gpu.stream();
1747 let mut b = __s_b.launch_builder(&f);
1748 b.arg(src).arg(dst);
1749 unsafe { b.launch(cfg)?; }
1750 Ok(())
1751 }
1752
1753 pub fn spec_adapt_k(&self, acc: &CudaSlice<u32>, brk: &mut CudaSlice<u32>,
1757 floor: usize, cap: usize)
1758 -> Result<(), Box<dyn std::error::Error>> {
1759 let f = self.func("spec_adapt_k");
1760 let (fl, cp) = (floor as i32, cap as i32);
1761 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1762 let __s_b = self.gpu.stream();
1763 let mut b = __s_b.launch_builder(&f);
1764 b.arg(acc).arg(brk).arg(&fl).arg(&cp);
1765 unsafe { b.launch(cfg)?; }
1766 Ok(())
1767 }
1768
1769 pub fn spec_accept_greedy_dc(&self, preds: &CudaSlice<u32>, vtok: &CudaSlice<u32>,
1771 last_pred: &CudaSlice<u32>, brk: &CudaSlice<u32>,
1772 out: &mut CudaSlice<u32>)
1773 -> Result<(), Box<dyn std::error::Error>> {
1774 let f = self.func("spec_accept_greedy_dc");
1775 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1776 let __s_b = self.gpu.stream();
1777 let mut b = __s_b.launch_builder(&f);
1778 b.arg(preds).arg(vtok).arg(last_pred).arg(brk).arg(out);
1779 unsafe { b.launch(cfg)?; }
1780 Ok(())
1781 }
1782
1783 pub fn pos_iota(&self, pos0: &CudaSlice<i32>, out: &mut CudaSlice<i32>, t: usize)
1785 -> Result<(), Box<dyn std::error::Error>> {
1786 let f = self.func("pos_iota_i32");
1787 let ti = t as i32;
1788 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (t.max(1) as u32, 1, 1),
1789 shared_mem_bytes: 0 };
1790 let __s_b = self.gpu.stream();
1791 let mut b = __s_b.launch_builder(&f);
1792 b.arg(pos0).arg(out).arg(&ti);
1793 unsafe { b.launch(cfg)?; }
1794 Ok(())
1795 }
1796 #[allow(clippy::too_many_arguments)]
1797 pub fn append_kv_quantized_rows_dc(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
1798 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
1799 t0_dev: &CudaSlice<i32>, t: usize,
1800 kv_dim_k: usize, kv_dim_v: usize,
1801 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
1802 -> Result<(), Box<dyn std::error::Error>> {
1803 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_rows_dc") }
1804 else { self.func("append_quantize_kv_q8_0_q5_1_rows_dc") };
1805 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
1806 let cfg = LaunchConfig { grid_dim: (nblk, t as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1807 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
1808 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
1809 let __s_b = self.gpu.stream();
1810 let mut b = __s_b.launch_builder(&f);
1811 b.arg(k_rows).arg(v_rows).arg(kc).arg(vc).arg(t0_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
1812 unsafe { b.launch(cfg)?; }
1813 Ok(())
1814 }
1815
1816 #[allow(clippy::too_many_arguments)]
1819 pub fn append_kv_quantized_row_dc_inc(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
1820 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
1821 t0_dev: &mut CudaSlice<i32>,
1822 kv_dim_k: usize, kv_dim_v: usize,
1823 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
1824 -> Result<(), Box<dyn std::error::Error>> {
1825 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_dc_inc") }
1826 else { self.func("append_quantize_kv_q8_0_q5_1_dc_inc") };
1827 let nthreads = ((kv_dim_k.max(kv_dim_v) / 32) * 32).min(1024) as u32;
1828 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (nthreads, 1, 1),
1829 shared_mem_bytes: 0 };
1830 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
1831 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
1832 let __s_b = self.gpu.stream();
1833 let mut b = __s_b.launch_builder(&f);
1834 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(t0_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
1835 unsafe { b.launch(cfg)?; }
1836 Ok(())
1837 }
1838
1839 pub fn pack_tok_p(&self, tok: &CudaSlice<u32>, p: &CudaSlice<f32>, out: &mut CudaSlice<u32>,
1841 slot: usize) -> Result<(), Box<dyn std::error::Error>> {
1842 let f = self.func("pack_tok_p");
1843 let sl = slot as i32;
1844 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1845 let __s_b = self.gpu.stream();
1846 let mut b = __s_b.launch_builder(&f);
1847 b.arg(tok).arg(p).arg(out).arg(&sl);
1848 unsafe { b.launch(cfg)?; }
1849 Ok(())
1850 }
1851 pub fn tok_map_u32(&self, tok: &mut CudaSlice<u32>, map: &CudaSlice<u32>)
1852 -> Result<(), Box<dyn std::error::Error>> {
1853 let f = self.func("tok_map_u32");
1854 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1855 let __s_b = self.gpu.stream();
1856 let mut b = __s_b.launch_builder(&f);
1857 b.arg(tok).arg(map);
1858 unsafe { b.launch(cfg)?; }
1859 Ok(())
1860 }
1861
1862 #[allow(clippy::too_many_arguments)]
1864 pub fn spec_assemble_verify(&self, tokp: &CudaSlice<u32>, pend: &CudaSlice<u32>,
1865 d2t: Option<&CudaSlice<u32>>, vtok: &mut CudaSlice<u32>,
1866 brk: &mut CudaSlice<u32>, p_min: f32, k: usize, pmin0: bool)
1867 -> Result<(), Box<dyn std::error::Error>> {
1868 let f = self.func("spec_assemble_verify");
1869 let (ki, pm) = (k as i32, if pmin0 { 1i32 } else { 0i32 });
1870 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
1871 let __s_b = self.gpu.stream();
1872 let mut b = __s_b.launch_builder(&f);
1873 match d2t {
1874 Some(m) => { b.arg(tokp).arg(pend).arg(m).arg(vtok).arg(brk).arg(&p_min).arg(&ki).arg(&pm);
1875 unsafe { b.launch(cfg)?; } }
1876 None => { let null: u64 = 0;
1877 b.arg(tokp).arg(pend).arg(&null).arg(vtok).arg(brk).arg(&p_min).arg(&ki).arg(&pm);
1878 unsafe { b.launch(cfg)?; } }
1879 }
1880 Ok(())
1881 }
1882
1883 #[allow(clippy::too_many_arguments)]
1885 pub fn ssm_conv_ring_rebuild_dc(&self, qkv_tm: &CudaSlice<f32>, ring_old: &CudaSlice<f32>,
1886 conv_state: &mut CudaSlice<f32>, conv_dim: usize,
1887 acc: &CudaSlice<u32>, base: usize, t_v: usize, d_conv: usize)
1888 -> Result<(), Box<dyn std::error::Error>> {
1889 let f = self.func("ssm_conv_ring_rebuild_f32_dc");
1890 let n = conv_dim * (d_conv - 1);
1891 let cfg = LaunchConfig::for_num_elems(n as u32);
1892 let (cd, b0, tv, dc) = (conv_dim as i32, base as i32, t_v as i32, d_conv as i32);
1893 let __s_b = self.gpu.stream();
1894 let mut b = __s_b.launch_builder(&f);
1895 b.arg(qkv_tm).arg(ring_old).arg(conv_state).arg(&cd).arg(acc).arg(&b0).arg(&tv).arg(&dc);
1896 unsafe { b.launch(cfg)?; }
1897 Ok(())
1898 }
1899 #[allow(clippy::too_many_arguments)]
1900 pub fn gdn_scan_s128_dc(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
1901 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
1902 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
1903 n_head: usize, acc: &CudaSlice<u32>, base: usize, t_v: usize,
1904 scale: f32)
1905 -> Result<(), Box<dyn std::error::Error>> {
1906 let f = self.func("gdn_scan_s128_dc");
1907 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
1908 let cfg = LaunchConfig {
1909 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
1910 block_dim: (WARP, COLS_PER_BLOCK, 1),
1911 shared_mem_bytes: 0,
1912 };
1913 let (h, b0, tv) = (n_head as i32, base as i32, t_v as i32);
1914 let __s_b = self.gpu.stream();
1915 let mut b = __s_b.launch_builder(&f);
1916 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in).arg(state_out).arg(o)
1917 .arg(&h).arg(acc).arg(&b0).arg(&tv).arg(&scale);
1918 unsafe { b.launch(cfg)?; }
1919 Ok(())
1920 }
1921
1922 pub fn spec_rollback_kv(&self, len_ptrs: &CudaSlice<u64>, saved: &CudaSlice<i32>,
1924 acc: &CudaSlice<u32>, base: usize, n_layer: usize)
1925 -> Result<(), Box<dyn std::error::Error>> {
1926 let f = self.func("spec_rollback_kv");
1927 let (b, nl) = (base as i32, n_layer as i32);
1928 let cfg = LaunchConfig { grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
1929 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1930 let __s_bl = self.gpu.stream();
1931 let mut bl = __s_bl.launch_builder(&f);
1932 bl.arg(len_ptrs).arg(saved).arg(acc).arg(&b).arg(&nl);
1933 unsafe { bl.launch(cfg)?; }
1934 Ok(())
1935 }
1936
1937 pub fn spec_fork_valid(&self, acc: &CudaSlice<u32>, optimistic_pending: u32,
1939 valid: &mut CudaSlice<u32>)
1940 -> Result<(), Box<dyn std::error::Error>> {
1941 let f = self.func("spec_fork_valid");
1942 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1),
1943 shared_mem_bytes: 0 };
1944 let __s_bl = self.gpu.stream();
1945 let mut bl = __s_bl.launch_builder(&f);
1946 bl.arg(acc).arg(&optimistic_pending).arg(valid);
1947 unsafe { bl.launch(cfg)?; }
1948 Ok(())
1949 }
1950
1951 pub fn spec_fork_reconcile_kv(&self, len_ptrs: &CudaSlice<u64>, saved: &CudaSlice<i32>,
1953 acc: &CudaSlice<u32>, valid: &CudaSlice<u32>, base: usize,
1954 n_layer: usize)
1955 -> Result<(), Box<dyn std::error::Error>> {
1956 let f = self.func("spec_fork_reconcile_kv");
1957 let (b, nl) = (base as i32, n_layer as i32);
1958 let cfg = LaunchConfig { grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
1959 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
1960 let __s_bl = self.gpu.stream();
1961 let mut bl = __s_bl.launch_builder(&f);
1962 bl.arg(len_ptrs).arg(saved).arg(acc).arg(valid).arg(&b).arg(&nl);
1963 unsafe { bl.launch(cfg)?; }
1964 Ok(())
1965 }
1966
1967 pub fn spec_fork_restore_f32(&self, snapshot: &CudaSlice<f32>, state: &mut CudaSlice<f32>,
1969 valid: &CudaSlice<u32>)
1970 -> Result<(), Box<dyn std::error::Error>> {
1971 assert_eq!(snapshot.len(), state.len(), "fork recurrent snapshot shape mismatch");
1972 let f = self.func("spec_fork_restore_f32");
1973 let n = state.len() as i32;
1974 let blocks = state.len().div_ceil(256).min(65535).max(1) as u32;
1975 let cfg = LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1),
1976 shared_mem_bytes: 0 };
1977 let __s_bl = self.gpu.stream();
1978 let mut bl = __s_bl.launch_builder(&f);
1979 bl.arg(snapshot).arg(state).arg(valid).arg(&n);
1980 unsafe { bl.launch(cfg)?; }
1981 Ok(())
1982 }
1983
1984 pub fn spec_seed_gather(&self, vx: &CudaSlice<f32>, fill_prev: &CudaSlice<f32>,
1987 acc: &CudaSlice<u32>, h_seed: &mut CudaSlice<f32>,
1988 base: usize, n_embd: usize)
1989 -> Result<(), Box<dyn std::error::Error>> {
1990 let f = self.func("spec_seed_gather");
1991 let (b, ne) = (base as i32, n_embd as i32);
1992 let cfg = LaunchConfig { grid_dim: (n_embd.div_ceil(256) as u32, 1, 1),
1993 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
1994 let __s_bl = self.gpu.stream();
1995 let mut bl = __s_bl.launch_builder(&f);
1996 bl.arg(vx).arg(fill_prev).arg(acc).arg(h_seed).arg(&b).arg(&ne);
1997 unsafe { bl.launch(cfg)?; }
1998 Ok(())
1999 }
2000
2001
2002 pub fn spec_accept_greedy(&self, preds: &CudaSlice<u32>, draft: &CudaSlice<u32>,
2004 last_pred: u32, base: usize, k_round: usize,
2005 out: &mut CudaSlice<u32>)
2006 -> Result<(), Box<dyn std::error::Error>> {
2007 let f = self.func("spec_accept_greedy");
2008 let (b, k) = (base as i32, k_round as i32);
2009 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2010 let __s_bl = self.gpu.stream();
2011 let mut bl = __s_bl.launch_builder(&f);
2012 bl.arg(preds).arg(draft).arg(&last_pred).arg(&b).arg(&k).arg(out);
2013 unsafe { bl.launch(cfg)?; }
2014 Ok(())
2015 }
2016
2017 pub fn gumbel_perturb(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
2024 seed: u64, stream_pos: u32, temp: f32)
2025 -> Result<(), Box<dyn std::error::Error>> {
2026 let f = self.func("gumbel_perturb_f32");
2027 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2028 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2029 let __s_b = self.gpu.stream();
2030 let mut b = __s_b.launch_builder(&f);
2031 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
2032 unsafe { b.launch(cfg)?; }
2033 Ok(())
2034 }
2035
2036 pub fn mask_logits_col(&self, logits: &mut CudaSlice<f32>, mask: &CudaSlice<u32>,
2044 col: usize, n: usize, mask_words: usize)
2045 -> Result<(), Box<dyn std::error::Error>> {
2046 let f = self.func("mask_logits_f32");
2047 let (ci, ni, mw) = (col as i32, n as i32, mask_words as i32);
2048 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256).min(1024) as u32, 1, 1),
2049 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2050 let __s_b = self.gpu.stream();
2051 let mut b = __s_b.launch_builder(&f);
2052 b.arg(&mut *logits).arg(mask).arg(&ci).arg(&ni).arg(&mw);
2053 unsafe { b.launch(cfg)?; }
2054 Ok(())
2055 }
2056
2057 pub fn gumbel_perturb_col(&self, x: &CudaSlice<f32>, col: usize, y: &mut CudaSlice<f32>,
2064 n: usize, seed: u64, stream_pos: u32, temp: f32)
2065 -> Result<(), Box<dyn std::error::Error>> {
2066 let f = self.func("gumbel_perturb_f32");
2067 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2068 let col_view = x.slice(col * n..(col + 1) * n);
2069 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2070 let __s_b = self.gpu.stream();
2071 let mut b = __s_b.launch_builder(&f);
2072 b.arg(&col_view).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(&stream_pos).arg(&temp);
2073 unsafe { b.launch(cfg)?; }
2074 Ok(())
2075 }
2076
2077 pub fn sctr_inc(&self, ctr: &mut CudaSlice<u32>) -> Result<(), Box<dyn std::error::Error>> {
2082 let f = self.func("memra_sctr_inc");
2083 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
2084 let __s_b = self.gpu.stream();
2085 let mut b = __s_b.launch_builder(&f);
2086 b.arg(&mut *ctr);
2087 unsafe { b.launch(cfg)?; }
2088 Ok(())
2089 }
2090
2091 pub fn gumbel_perturb_ctr(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize,
2096 seed: u64, ctr: &CudaSlice<u32>, temp: f32)
2097 -> Result<(), Box<dyn std::error::Error>> {
2098 let f = self.func("gumbel_perturb_ctr_f32");
2099 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2100 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256) as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2101 let __s_b = self.gpu.stream();
2102 let mut b = __s_b.launch_builder(&f);
2103 b.arg(x).arg(&mut *y).arg(&ni).arg(&slo).arg(&shi).arg(ctr).arg(&temp);
2104 unsafe { b.launch(cfg)?; }
2105 Ok(())
2106 }
2107
2108 pub fn softmax_gather(&self, x: &CudaSlice<f32>, row_stride: usize,
2112 ids: &CudaSlice<u32>, rows: &CudaSlice<i32>,
2113 out: &mut CudaSlice<f32>, n: usize, npair: usize, temp: f32)
2114 -> Result<(), Box<dyn std::error::Error>> {
2115 let f = self.func("softmax_gather_f32");
2116 let (ni, rs) = (n as i32, row_stride as i64);
2117 let np = npair as i32;
2118 let cfg = LaunchConfig { grid_dim: (npair as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2119 let __s_b = self.gpu.stream();
2120 let mut b = __s_b.launch_builder(&f);
2121 b.arg(x).arg(&rs).arg(ids).arg(rows).arg(&mut *out).arg(&ni).arg(&np).arg(&temp);
2122 unsafe { b.launch(cfg)?; }
2123 Ok(())
2124 }
2125
2126 pub fn residual_sample(&self, p: &CudaSlice<f32>, q: Option<&CudaSlice<f32>>, n: usize,
2130 temp: f32, seed: u64, stream_pos: u32,
2131 out_tok: &mut CudaSlice<u32>)
2132 -> Result<(), Box<dyn std::error::Error>> {
2133 let f = self.func("residual_sample_f32");
2134 let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
2135 let nth = 1024u32;
2136 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (nth, 1, 1), shared_mem_bytes: 0 };
2137 let has_q: i32 = q.is_some() as i32;
2138 let qbuf = q.unwrap_or(p); let __s_b = self.gpu.stream();
2140 let mut b = __s_b.launch_builder(&f);
2141 b.arg(p).arg(qbuf).arg(&has_q).arg(&ni).arg(&temp).arg(&slo).arg(&shi).arg(&stream_pos)
2142 .arg(&mut *out_tok);
2143 unsafe { b.launch(cfg)?; }
2144 Ok(())
2145 }
2146
2147 pub fn with_moe_cache<R>(&self, max_block_bytes: usize,
2152 f: impl FnOnce(&mut crate::moe_cache::MoeSlotCache, &Engine) -> Result<R, Box<dyn std::error::Error>>)
2153 -> Result<R, Box<dyn std::error::Error>> {
2154 let mut guard = self.moe_cache.lock().unwrap();
2155 if guard.is_none() {
2156 *guard = Some(crate::moe_cache::MoeSlotCache::new(self, max_block_bytes)?);
2157 }
2158 let cache = guard.as_mut().unwrap();
2159 f(cache, self)
2160 }
2161
2162 pub fn freeze_moe_cache(&self) {
2165 if let Some(cache) = self.moe_cache.lock().unwrap().as_mut() {
2166 cache.freeze();
2167 }
2168 }
2169
2170 pub fn export_moe_residency(&self) -> Option<Vec<(u16, u8, u16)>> {
2173 self.moe_cache
2174 .lock()
2175 .unwrap()
2176 .as_ref()
2177 .map(crate::moe_cache::MoeSlotCache::export_residency)
2178 }
2179
2180 pub(crate) fn moe_cache_frozen(&self) -> bool {
2181 self.moe_cache
2182 .lock()
2183 .unwrap()
2184 .as_ref()
2185 .is_some_and(crate::moe_cache::MoeSlotCache::is_frozen)
2186 }
2187
2188 pub fn frozen_cpu_experts_prefer_tokenwise_prime(&self) -> bool {
2195 crate::cpu_experts::configured()
2196 && self.moe_cache_frozen()
2197 && std::env::var("MEMRA_CPU_EXPERT_BATCHED_PRIME").as_deref() != Ok("1")
2198 }
2199
2200 pub(crate) fn configure_moe_cache_layout(&self, block_bytes: Vec<usize>) {
2202 assert!(
2203 self.moe_cache.lock().unwrap().is_none(),
2204 "MoE cache layout configured after cache construction"
2205 );
2206 *self.moe_cache_layout.lock().unwrap() = Some(block_bytes);
2207 }
2208
2209 pub(crate) fn moe_cache_layout(&self) -> Option<Vec<usize>> {
2210 self.moe_cache_layout.lock().unwrap().clone()
2211 }
2212
2213 pub fn moe_cache_enabled() -> bool {
2215 std::env::var("MEMRA_MOE_CACHE").as_deref() != Ok("0")
2216 }
2217
2218 pub fn moe_cache_stats(&self) -> Option<(u64, u64, u64, usize)> {
2221 let guard = self.moe_cache.lock().unwrap();
2222 guard.as_ref() .map(|c| (c.hits, c.misses, c.staged_bytes, c.n_slots()))
2223 }
2224
2225 pub fn cpu_expert_stats(
2229 &self,
2230 ) -> Option<(u64, u64, u64, u64, u64, u64, u64, u64, u64, u64, u64)> {
2231 crate::cpu_experts::configured().then(crate::cpu_experts::stats)
2232 }
2233
2234 pub fn cpu_expert_predictor_stats(&self) -> (u64, u64) {
2237 crate::cpu_experts::predictor_stats()
2238 }
2239
2240 pub fn cpu_expert_exposed_wait_ns(&self) -> Option<u64> {
2241 crate::cpu_experts::configured().then(crate::cpu_experts::exposed_wait_ns)
2242 }
2243
2244 pub fn cpu_expert_gpu_residency_stats(&self) -> Option<(u64, u64, u64)> {
2247 crate::cpu_experts::configured().then(crate::cpu_experts::incomplete_gpu_residency_stats)
2248 }
2249
2250 pub fn moe_pread_stats(&self) -> Option<(u64, u64, u64, u64, u64, u64, u64)> {
2253
2254 let guard = self.moe_cache.lock().unwrap();
2255 guard.as_ref().and_then(|cache| cache.pread_stats()).map(|stats| (
2256 stats.reads,
2257 stats.bytes,
2258 stats.read_errors,
2259 stats.short_reads,
2260 stats.fallbacks,
2261 stats.buffer_waits,
2262 stats.ring_full,
2263 ))
2264 }
2265
2266 pub fn moe_cache_reset_counters(&self) {
2268 if let Some(c) = self.moe_cache.lock().unwrap().as_mut() { c.reset_counters(); }
2269 }
2270
2271 pub fn htod_bytes(&self, v: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2272 Ok(self.gpu.stream().clone_htod(v)?)
2273 }
2274
2275 pub fn htod_bytes_padded(&self, v: &[u8], pad: usize)
2279 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2280 let mut d = self.alloc_u8_uninit(v.len() + pad)?;
2281 {
2282 let mut view = d.slice_mut(0..v.len());
2283 self.gpu.stream().memcpy_htod(v, &mut view)?;
2284 }
2285 Ok(d)
2286 }
2287
2288 pub fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
2290 -> Result<(), Box<dyn std::error::Error>> {
2291 let mut view = dst.slice_mut(off..off + len);
2292 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2293 Ok(())
2294 }
2295
2296 pub fn copy_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &CudaSlice<u8>, len: usize)
2299 -> Result<(), Box<dyn std::error::Error>> {
2300 let mut view = dst.slice_mut(off..off + len);
2301 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2302 Ok(())
2303 }
2304
2305 pub fn copy_u8_range_into(
2307 &self,
2308 dst: &mut CudaSlice<u8>,
2309 dst_off: usize,
2310 src: &CudaSlice<u8>,
2311 src_off: usize,
2312 len: usize,
2313 ) -> Result<(), Box<dyn std::error::Error>> {
2314 let mut dst_view = dst.slice_mut(dst_off..dst_off + len);
2315 self.gpu
2316 .stream()
2317 .memcpy_dtod(&src.slice(src_off..src_off + len), &mut dst_view)?;
2318 Ok(())
2319 }
2320
2321 pub fn prepare_kv_append(
2325 &self,
2326 kv: &mut crate::cache::KvLayer,
2327 retain_from: usize,
2328 append_rows: usize,
2329 ) -> Result<usize, Box<dyn std::error::Error>> {
2330 let Some(plan) = kv
2331 .ring
2332 .as_ref()
2333 .map(|ring| ring.append_plan(kv.len, retain_from, append_rows))
2334 .transpose()?
2335 else {
2336 return Ok(kv.len);
2337 };
2338 match plan {
2339 crate::cache::KvRingAppend::Contiguous { write_row } => Ok(write_row),
2340 crate::cache::KvRingAppend::Rebase {
2341 src_row,
2342 keep_rows,
2343 new_base,
2344 write_row,
2345 } => {
2346 if keep_rows > 0 {
2347 let k_len = keep_rows * kv.k_tok_bytes;
2348 let v_len = keep_rows * kv.v_tok_bytes;
2349 let mut k_tmp = self.alloc_u8_uninit(k_len)?;
2350 let mut v_tmp = self.alloc_u8_uninit(v_len)?;
2351 self.copy_u8_range_into(
2352 &mut k_tmp,
2353 0,
2354 &kv.k,
2355 src_row * kv.k_tok_bytes,
2356 k_len,
2357 )?;
2358 self.copy_u8_range_into(
2359 &mut v_tmp,
2360 0,
2361 &kv.v,
2362 src_row * kv.v_tok_bytes,
2363 v_len,
2364 )?;
2365 self.copy_u8_into(&mut kv.k, 0, &k_tmp, k_len)?;
2366 self.copy_u8_into(&mut kv.v, 0, &v_tmp, v_len)?;
2367 }
2368 kv.ring.as_mut().unwrap().apply_rebase(new_base);
2369 Ok(write_row)
2370 }
2371 }
2372 }
2373
2374 pub fn htod_u8_into(&self, dst: &mut CudaSlice<u8>, off: usize, src: &[u8])
2377 -> Result<(), Box<dyn std::error::Error>> {
2378 let mut view = dst.slice_mut(off..off + src.len());
2379 self.gpu.stream().memcpy_htod(src, &mut view)?;
2380 Ok(())
2381 }
2382
2383 pub fn view<'a>(&self, b: &'a CudaSlice<f32>, len: usize) -> cudarc::driver::CudaView<'a, f32> {
2384 b.slice(0..len)
2385 }
2386
2387 pub fn view_u8_range<'a>(&self, b: &'a CudaSlice<u8>, start: usize, end: usize)
2390 -> cudarc::driver::CudaView<'a, u8> {
2391 b.slice(start..end)
2392 }
2393 pub fn view_u8<'a>(&self, b: &'a CudaSlice<u8>, len: usize) -> cudarc::driver::CudaView<'a, u8> {
2394 b.slice(0..len)
2395 }
2396
2397 pub fn append_kv_quantized(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2401 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2402 kv_dim_k: usize, kv_dim_v: usize,
2403 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2404 -> Result<(), Box<dyn std::error::Error>> {
2405 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") } else { self.func("append_quantize_kv_q8_0_q5_1") };
2406 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2407 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2408 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2409 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2410 let __s_b = self.gpu.stream();
2411 let mut b = __s_b.launch_builder(&f);
2412 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2413 unsafe { b.launch(cfg)?; }
2414 Ok(())
2415 }
2416
2417 pub fn append_kv_quantized_dc(&self, k_row: &CudaSlice<f32>, v_row: &CudaSlice<f32>,
2421 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t_dev: &CudaSlice<i32>,
2422 kv_dim_k: usize, kv_dim_v: usize,
2423 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2424 -> Result<(), Box<dyn std::error::Error>> {
2425 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2426 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
2427 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2428 if Self::pdl_on() && Self::pdl_wb_on() {
2430 use cudarc::driver::{DevicePtr, DevicePtrMut};
2431 let s = &self.gpu.stream();
2432 let (pk, _g0) = k_row.device_ptr(s); let (pv, _g1) = v_row.device_ptr(s);
2433 let (pkc, _g2) = kc.device_ptr_mut(s); let (pvc, _g3) = vc.device_ptr_mut(s);
2434 let (pt, _g4) = t_dev.device_ptr(s);
2435 let mut ps = [
2436 &pk as *const _ as *mut std::ffi::c_void, &pv as *const _ as *mut _,
2437 &pkc as *const _ as *mut _, &pvc as *const _ as *mut _,
2438 &pt as *const _ as *mut _, &kdk as *const _ as *mut _,
2439 &kdv as *const _ as *mut _, &ktb as *const _ as *mut _,
2440 &vtb as *const _ as *mut _,
2441 ];
2442 unsafe { self.launch_pdl_flash(g, "append_quantize_kv_q8_0_q5_1_dc",
2443 (nblk, 1, 1), (32, 1, 1), 0, &mut ps)?; }
2444 return Ok(());
2445 }
2446 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_dc") } else { self.func("append_quantize_kv_q8_0_q5_1_dc") };
2447 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2448 let __s_b = self.gpu.stream();
2449 let mut b = __s_b.launch_builder(&f);
2450 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(t_dev).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2451 unsafe { b.launch(cfg)?; }
2452 Ok(())
2453 }
2454
2455 #[allow(clippy::too_many_arguments)]
2462 pub fn append_kv_quantized_rows(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
2463 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
2464 t0: usize, t: usize, kv_dim_k: usize, kv_dim_v: usize,
2465 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2466 -> Result<(), Box<dyn std::error::Error>> {
2467 if std::env::var("MEMRA_PRIME_APPEND_LOOP").is_ok() {
2468 for i in 0..t {
2469 let k_row = k_rows.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
2470 let v_row = v_rows.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
2471 self.append_kv_quantized_view(&k_row, &v_row, kc, vc, t0 + i,
2472 kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes, g)?;
2473 }
2474 return Ok(());
2475 }
2476 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1_rows") } else { self.func("append_quantize_kv_q8_0_q5_1_rows") };
2477 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2478 let cfg = LaunchConfig { grid_dim: (nblk, t as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2479 let (t0i, kdk, kdv) = (t0 as i32, kv_dim_k as i32, kv_dim_v as i32);
2480 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2481 let __s_b = self.gpu.stream();
2482 let mut b = __s_b.launch_builder(&f);
2483 b.arg(k_rows).arg(v_rows).arg(kc).arg(vc).arg(&t0i).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2484 unsafe { b.launch(cfg)?; }
2485 Ok(())
2486 }
2487
2488 pub fn inc_seqlen(&self, p: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
2492 let f = self.func("inc_i32");
2493 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
2494 let __s_b = self.gpu.stream();
2495 let mut b = __s_b.launch_builder(&f);
2496 b.arg(p);
2497 unsafe { b.launch(cfg)?; }
2498 Ok(())
2499 }
2500
2501 pub fn append_kv_quantized_view(&self, k_row: &cudarc::driver::CudaView<f32>,
2504 v_row: &cudarc::driver::CudaView<f32>,
2505 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>, t: usize,
2506 kv_dim_k: usize, kv_dim_v: usize,
2507 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
2508 -> Result<(), Box<dyn std::error::Error>> {
2509 let f = if g { self.func_g("append_quantize_kv_q8_0_q5_1") }
2510 else { self.func("append_quantize_kv_q8_0_q5_1") };
2511 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
2512 let cfg = LaunchConfig { grid_dim: (nblk, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2513 let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
2514 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
2515 let __s_b = self.gpu.stream();
2516 let mut b = __s_b.launch_builder(&f);
2517 b.arg(k_row).arg(v_row).arg(kc).arg(vc).arg(&ti).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
2518 unsafe { b.launch(cfg)?; }
2519 Ok(())
2520 }
2521
2522 pub fn copy_view_into(&self, dst: &mut CudaSlice<f32>, off: usize,
2525 src: &cudarc::driver::CudaView<f32>, len: usize)
2526 -> Result<(), Box<dyn std::error::Error>> {
2527 let mut view = dst.slice_mut(off..off + len);
2528 self.gpu.stream().memcpy_dtod(&src.slice(0..len), &mut view)?;
2529 Ok(())
2530 }
2531
2532 pub fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2536 let mut dst = self.gpu.stream().alloc_zeros::<f32>(src.len())?;
2537 self.gpu.stream().memcpy_dtod(src, &mut dst)?;
2538 Ok(dst)
2539 }
2540
2541 pub fn dtod_copy_view(&self, src: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<f32>)
2544 -> Result<(), Box<dyn std::error::Error>> {
2545 self.gpu.stream().memcpy_dtod(src, dst)?;
2546 Ok(())
2547 }
2548
2549 pub fn dtod_copy_view_i8(&self, src: &cudarc::driver::CudaView<i8>, dst: &mut CudaSlice<i8>)
2551 -> Result<(), Box<dyn std::error::Error>> {
2552 self.gpu.stream().memcpy_dtod(src, dst)?;
2553 Ok(())
2554 }
2555
2556 pub fn dtod_copy_into(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, offset: usize)
2558 -> Result<(), Box<dyn std::error::Error>> {
2559 let n = src.len();
2560 let mut dv = dst.slice_mut(offset..offset + n);
2561 self.gpu.stream().memcpy_dtod(src, &mut dv)?;
2562 Ok(())
2563 }
2564
2565 pub fn uninit_i8(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
2567 self.alloc_uninit::<i8>(n)
2568 }
2569
2570 pub fn qmatvec(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
2572 qtype: i32, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2573 let f = self.func("qmatvec_f32");
2574 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2576 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2577 let __s_b = self.gpu.stream();
2578 let mut b = __s_b.launch_builder(&f);
2579 b.arg(w).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2580 unsafe { b.launch(cfg)?; }
2581 Ok(y)
2582 }
2583
2584 pub fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2586 let s = self.gpu.stream().alloc_zeros::<u8>(n)?;
2587 self.keep_if_capturing(&s);
2588 Ok(s)
2589 }
2590
2591 pub fn alloc_u8_uninit(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
2595 let s = unsafe { self.gpu.stream().alloc::<u8>(n)? };
2596 self.keep_if_capturing(&s);
2597 Ok(s)
2598 }
2599
2600 pub fn memset_zeros_view(&self, dst: &mut cudarc::driver::CudaViewMut<f32>)
2603 -> Result<(), Box<dyn std::error::Error>> {
2604 self.gpu.stream().memset_zeros(dst)?;
2605 Ok(())
2606 }
2607
2608 pub fn stage_expert(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2614 -> Result<(), Box<dyn std::error::Error>> {
2615 let mut dst = scratch.slice_mut(off..off + host_bytes.len()); self.gpu.stream().memcpy_htod(host_bytes, &mut dst)?; Ok(())
2618 }
2619
2620 pub fn moe_router_topk(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2626 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2627 let f = self.func("moe_router_topk_f32");
2628 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?; let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?; let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2631 shared_mem_bytes: 0 };
2632 let (ne, nu) = (n_expert as i32, n_used as i32);
2633 let __s_b = self.gpu.stream();
2634 let mut b = __s_b.launch_builder(&f);
2635 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2636 unsafe { b.launch(cfg)?; }
2637 Ok((sel_idx, sel_w))
2638 }
2639
2640 pub fn moe_router_topk_scaled(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize,
2643 n_used: usize, ex_scale: &CudaSlice<f32>)
2644 -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2645 let f = self.func("moe_router_topk_scaled_f32");
2650 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
2651 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
2652 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2653 shared_mem_bytes: 0 };
2654 let (ne, nu) = (n_expert as i32, n_used as i32);
2655 let __s_b = self.gpu.stream();
2656 let mut b = __s_b.launch_builder(&f);
2657 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu).arg(ex_scale);
2658 unsafe { b.launch(cfg)?; }
2659 Ok((sel_idx, sel_w))
2660 }
2661
2662 pub fn moe_router_topk_host(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2670 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2671 let f = self.func("moe_router_topk_f32");
2672 let n = t * n_used;
2673 let mut sel_idx = self.alloc_uninit::<i32>(n)?;
2674 let mut sel_w = self.alloc_uninit::<f32>(n)?;
2675 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (n_expert as u32, 1, 1),
2676 shared_mem_bytes: 0 };
2677 let (ne, nu) = (n_expert as i32, n_used as i32);
2678 let __s_b = self.gpu.stream();
2679 let mut b = __s_b.launch_builder(&f);
2680 b.arg(logits).arg(&mut sel_idx).arg(&mut sel_w).arg(&ne).arg(&nu);
2681 unsafe { b.launch(cfg)?; }
2682 let bytes = n * 8;
2684 let mut guard = self.router_stage.lock().unwrap();
2685 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
2686 *guard = Some(PinnedStage::new(bytes.max(4096))?);
2687 }
2688 let stage = guard.as_mut().unwrap();
2689 let (si, sw) = unsafe {
2690 (std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
2691 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n))
2692 };
2693 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()))
2697 }
2698
2699 #[allow(clippy::too_many_arguments)]
2703 pub fn moe_router_sigmoid_topk(&self, logits: &CudaSlice<f32>, t: usize, n_expert: usize,
2704 n_used: usize, active_count: usize,
2705 correction_bias: &CudaSlice<f32>,
2706 active: &CudaSlice<u8>, scaling_factor: f32, route_norm: bool)
2707 -> Result<(CudaSlice<i32>, CudaSlice<f32>),
2708 Box<dyn std::error::Error>> {
2709 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
2710 if n_expert == 0 || n_expert > 1024 || n_used == 0 || n_used > n_expert {
2711 return Err(format!(
2712 "sigmoid router shape unsupported: n_expert={n_expert}, n_used={n_used}",
2713 ).into());
2714 }
2715 if logits.len() < t * n_expert || correction_bias.len() != n_expert
2716 || active.len() != n_expert {
2717 return Err(format!(
2718 "sigmoid router buffer mismatch: logits={} bias={} active={} expected logits>={} row={}",
2719 logits.len(), correction_bias.len(), active.len(), t * n_expert, n_expert,
2720 ).into());
2721 }
2722 let f = self.func("moe_router_sigmoid_topk_f32");
2723 let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
2724 let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
2725 let threads = n_expert.div_ceil(32) * 32;
2726 let cfg = LaunchConfig { grid_dim: (t as u32, 1, 1), block_dim: (threads as u32, 1, 1),
2727 shared_mem_bytes: 0 };
2728 let (ne, nu, rn) = (n_expert as i32, n_used as i32, i32::from(route_norm));
2729 let __s_b = self.gpu.stream();
2730 let mut b = __s_b.launch_builder(&f);
2731 b.arg(logits).arg(correction_bias).arg(active).arg(&mut sel_idx).arg(&mut sel_w)
2732 .arg(&ne).arg(&nu).arg(&scaling_factor).arg(&rn);
2733 unsafe { b.launch(cfg)?; }
2734 Ok((sel_idx, sel_w))
2735 }
2736
2737 #[allow(clippy::too_many_arguments)]
2740 pub fn moe_router_sigmoid_topk_host(
2741 &self,
2742 logits: &CudaSlice<f32>,
2743 t: usize,
2744 n_expert: usize,
2745 n_used: usize,
2746 active_count: usize,
2747 correction_bias: &CudaSlice<f32>,
2748 active: &CudaSlice<u8>,
2749 scaling_factor: f32,
2750 route_norm: bool,
2751 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2752 let (sel_idx, sel_w) = self.moe_router_sigmoid_topk(
2753 logits, t, n_expert, n_used, active_count, correction_bias, active, scaling_factor,
2754 route_norm,
2755 )?;
2756 let n = t * n_used;
2757 let bytes = n * 8;
2758 let mut guard = self.router_stage.lock().unwrap();
2759 if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
2760 *guard = Some(PinnedStage::new(bytes.max(4096))?);
2761 }
2762 let stage = guard.as_mut().unwrap();
2763 let (si, sw) = unsafe {
2764 (std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
2765 std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n))
2766 };
2767 self.gpu.stream().memcpy_dtoh(&sel_idx, si)?;
2768 self.gpu.stream().memcpy_dtoh(&sel_w, sw)?;
2769 self.gpu.stream().synchronize()?;
2770 Ok((si.iter().map(|&i| i as u32).collect(), sw.to_vec()))
2771 }
2772
2773 pub fn stage_expert_async(&self, host_bytes: &[u8], scratch: &mut CudaSlice<u8>, off: usize)
2777 -> Result<cudarc::driver::CudaEvent, Box<dyn std::error::Error>> {
2778 let mut dst = scratch.slice_mut(off..off + host_bytes.len());
2779 self.copy_stream.memcpy_htod(host_bytes, &mut dst)?;
2780 Ok(self.copy_stream.record_event(None)?)
2781 }
2782
2783 pub fn compute_wait(&self, ev: &cudarc::driver::CudaEvent) -> Result<(), Box<dyn std::error::Error>> {
2785 self.gpu.stream().wait(ev)?;
2786 Ok(())
2787 }
2788
2789 pub fn qmatvec_view(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2794 x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize, out_f: usize,
2795 qtype: i32, row_bytes: usize)
2796 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2797 let f = self.func("qmatvec_f32");
2798 let wv = w.slice(range); let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2801 let (inf, outf, mi, qt, rb) = (in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
2802 let __s_b = self.gpu.stream();
2803 let mut b = __s_b.launch_builder(&f);
2804 b.arg(&wv).arg(x).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qt).arg(&rb);
2805 unsafe { b.launch(cfg)?; }
2806 Ok(y)
2807 }
2808
2809 #[allow(clippy::too_many_arguments)]
2816 pub fn moe_gate_up_silu8_q8(&self, gp: WPtr8, up: WPtr8,
2820 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2821 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2822 rb_g: usize, rb_u: usize)
2823 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2824 let f = self.func("moe_gate_up_silu8_q8");
2825 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
2826 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2827 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2828 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2829 let __s_b = self.gpu.stream();
2830 let mut b = __s_b.launch_builder(&f);
2831 b.arg(&gp).arg(&up).arg(aq).arg(ad).arg(&mut act)
2832 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2833 unsafe { b.launch(cfg)?; }
2834 Ok(act)
2835 }
2836
2837 #[allow(clippy::too_many_arguments)]
2838 pub fn moe_down8_fma_q8(&self, dp: WPtr8, w: F32x8,
2839 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
2840 dst: &mut cudarc::driver::CudaViewMut<f32>,
2841 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2842 -> Result<(), Box<dyn std::error::Error>> {
2843 let f = self.func("moe_down8_fma_q8");
2844 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2845 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
2846 let (inf, outf, nu, rbi) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2847 let __s_b = self.gpu.stream();
2848 let mut b = __s_b.launch_builder(&f);
2849 b.arg(&dp).arg(&w).arg(aq2).arg(ad2).arg(dst)
2850 .arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbi);
2851 unsafe { b.launch(cfg)?; }
2852 Ok(())
2853 }
2854
2855 pub fn qmatvec_expert_q8(&self, w: &CudaSlice<u8>, range: std::ops::Range<usize>,
2857 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
2858 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize)
2859 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2860 let f = self.func("qmatvec_expert_q8");
2861 let wv = w.slice(range);
2862 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
2863 const ROWS: u32 = 4; let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, m as u32, 1),
2865 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2866 let (inf, outf, mi, rbi) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
2867 let __s_b = self.gpu.stream();
2868 let mut b = __s_b.launch_builder(&f);
2869 b.arg(&wv).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&qtype).arg(&rbi);
2870 unsafe { b.launch(cfg)?; }
2871 Ok(y)
2872 }
2873
2874 pub fn moe_gate_up_silu8(&self, gp: WPtr8, up: WPtr8, x: &cudarc::driver::CudaView<f32>,
2875 in_f: usize, n_ff: usize, n_used: usize, qt_g: i32, qt_u: i32,
2876 rb_g: usize, rb_u: usize)
2877 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2878 let f = self.func("moe_gate_up_silu8_f32");
2879 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
2881 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2882 let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
2883 let __s_b = self.gpu.stream();
2884 let mut b = __s_b.launch_builder(&f);
2885 b.arg(&gp).arg(&up).arg(x).arg(&mut act)
2886 .arg(&inf).arg(&nff).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
2887 unsafe { b.launch(cfg)?; }
2888 Ok(act)
2889 }
2890
2891 #[allow(clippy::too_many_arguments)]
2897 pub fn moe_down8_fma_into(&self, dp: WPtr8, w: F32x8, act: &CudaSlice<f32>,
2898 dst: &mut cudarc::driver::CudaViewMut<f32>,
2899 in_f: usize, out_f: usize, n_used: usize, qt: i32, rb: usize)
2900 -> Result<(), Box<dyn std::error::Error>> {
2901 let f = self.func("moe_down8_fma_f32");
2902 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
2903 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
2904 let (inf, outf, nu, rbv) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
2905 let __s_b = self.gpu.stream();
2906 let mut b = __s_b.launch_builder(&f);
2907 b.arg(&dp).arg(&w).arg(act).arg(dst).arg(&inf).arg(&outf).arg(&nu).arg(&qt).arg(&rbv);
2908 unsafe { b.launch(cfg)?; }
2909 Ok(())
2910 }
2911
2912 #[allow(clippy::too_many_arguments)]
2917 #[allow(clippy::too_many_arguments)]
2932 #[allow(clippy::too_many_arguments)]
2934 pub fn moe_pairs_matvec_q8(&self, table: &CudaSlice<u64>, proj: i32,
2935 pair_tok: &CudaSlice<i32>, pair_ex: &CudaSlice<i32>,
2936 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2937 in_f: usize, out_f: usize, n_expert: usize, n_pairs: usize,
2938 qtype: i32, row_bytes: usize)
2939 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2940 let f = self.func("moe_pairs_matvec_q8");
2941 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2942 const ROWS: u32 = 4;
2943 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_pairs as u32, 1),
2944 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2945 let (inf, outf, ne, np, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2946 n_pairs as i32, row_bytes as i64);
2947 let __s_b = self.gpu.stream();
2948 let mut b = __s_b.launch_builder(&f);
2949 b.arg(table).arg(&proj).arg(pair_tok).arg(pair_ex).arg(aq).arg(ad).arg(&mut y)
2950 .arg(&inf).arg(&outf).arg(&ne).arg(&np).arg(&qtype).arg(&rbi);
2951 unsafe { b.launch(cfg)?; }
2952 Ok(y)
2953 }
2954
2955 #[allow(clippy::too_many_arguments)]
2957 pub fn moe_pairs_matvec_q8_em(&self, table: &CudaSlice<u64>, proj: i32,
2958 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2959 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2960 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2961 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2962 n_pairs: usize, qtype: i32, row_bytes: usize)
2963 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2964 let f = self.func("moe_pairs_matvec_q8_em");
2965 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2966 const ROWS: u32 = 4;
2967 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
2968 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2969 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2970 n_active as i32, row_bytes as i64);
2971 let __s_b = self.gpu.stream();
2972 let mut b = __s_b.launch_builder(&f);
2973 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
2974 .arg(aq).arg(ad).arg(&mut y)
2975 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
2976 unsafe { b.launch(cfg)?; }
2977 Ok(y)
2978 }
2979
2980 #[allow(clippy::too_many_arguments)]
2983 pub fn moe_pairs_matvec_q8_dec(&self, table: &CudaSlice<u64>, proj: i32,
2984 ex_ids: &CudaSlice<i32>, ex_off: &CudaSlice<i32>,
2985 ex_pairs: &CudaSlice<i32>, pair_tok: &CudaSlice<i32>,
2986 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
2987 in_f: usize, out_f: usize, n_expert: usize, n_active: usize,
2988 n_pairs: usize, qtype: i32, row_bytes: usize)
2989 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2990 let f = self.func("moe_pairs_matvec_q8_dec");
2991 let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
2992 const ROWS: u32 = 4;
2993 let cfg = LaunchConfig { grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
2994 block_dim: (32, ROWS, 1), shared_mem_bytes: 0 };
2995 let (inf, outf, ne, na, rbi) = (in_f as i32, out_f as i32, n_expert as i32,
2996 n_active as i32, row_bytes as i64);
2997 let __s_b = self.gpu.stream();
2998 let mut b = __s_b.launch_builder(&f);
2999 b.arg(table).arg(&proj).arg(ex_ids).arg(ex_off).arg(ex_pairs).arg(pair_tok)
3000 .arg(aq).arg(ad).arg(&mut y)
3001 .arg(&inf).arg(&outf).arg(&ne).arg(&na).arg(&qtype).arg(&rbi);
3002 unsafe { b.launch(cfg)?; }
3003 Ok(y)
3004 }
3005
3006 pub fn moe_pairs_gelu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
3007 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3008 let f = self.func("moe_pairs_gelu_mul");
3009 let mut act = self.alloc_uninit::<f32>(n)?;
3010 let cfg = LaunchConfig::for_num_elems(n as u32);
3011 let nl = n as i64;
3012 let __s_b = self.gpu.stream();
3013 let mut b = __s_b.launch_builder(&f);
3014 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
3015 unsafe { b.launch(cfg)?; }
3016 Ok(act)
3017 }
3018
3019 pub fn moe_pairs_silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, n: usize)
3020 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3021 let f = self.func("moe_pairs_silu_mul");
3022 let mut act = self.alloc_uninit::<f32>(n)?;
3023 let cfg = LaunchConfig::for_num_elems(n as u32);
3024 let nl = n as i64;
3025 let __s_b = self.gpu.stream();
3026 let mut b = __s_b.launch_builder(&f);
3027 b.arg(gate).arg(up).arg(&mut act).arg(&nl);
3028 unsafe { b.launch(cfg)?; }
3029 Ok(act)
3030 }
3031
3032 #[allow(clippy::too_many_arguments)]
3033 pub fn moe_pairs_scatter(&self, y_down: &CudaSlice<f32>, pair_w: &CudaSlice<f32>,
3034 tok_pair_off: &CudaSlice<i32>, tok_pair_ids: &CudaSlice<i32>,
3035 moe_out: &mut CudaSlice<f32>, t: usize, n_embd: usize)
3036 -> Result<(), Box<dyn std::error::Error>> {
3037 let f = self.func("moe_pairs_scatter");
3038 let cfg = LaunchConfig { grid_dim: (((n_embd + 255) / 256) as u32, t as u32, 1),
3039 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3040 let ne = n_embd as i32;
3041 let __s_b = self.gpu.stream();
3042 let mut b = __s_b.launch_builder(&f);
3043 b.arg(y_down).arg(pair_w).arg(tok_pair_off).arg(tok_pair_ids).arg(moe_out).arg(&ne);
3044 unsafe { b.launch(cfg)?; }
3045 Ok(())
3046 }
3047
3048 #[allow(clippy::too_many_arguments)]
3052 pub fn moe_gate_up_gelu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3053 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3054 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3055 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3056 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3057 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3058 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3059 rb_g as i64, rb_u as i64);
3060 let f = self.func("moe_gate_up_gelu8_dev_q8");
3061 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3062 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3063 let __s_b = self.gpu.stream();
3064 let mut b = __s_b.launch_builder(&f);
3065 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3066 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
3067 unsafe { b.launch(cfg)?; }
3068 Ok(act)
3069 }
3070
3071 #[allow(clippy::too_many_arguments)]
3073 pub fn moe_gate_up_gelu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3074 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
3075 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3076 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3077 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3078 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
3079 let (inf, nff, ne, rbg, rbu, nu) = (in_f as i32, n_ff as i32, n_expert as i32,
3080 rb_g as i64, rb_u as i64, n_used as i32);
3081 let f = self.func("moe_gate_up_gelu8_dev_q8_rows");
3082 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
3083 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3084 let __s_b = self.gpu.stream();
3085 let mut b = __s_b.launch_builder(&f);
3086 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3087 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu);
3088 unsafe { b.launch(cfg)?; }
3089 Ok(act)
3090 }
3091
3092 #[allow(clippy::too_many_arguments)]
3094 pub fn moe_gate_up_gelu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3095 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, n_pairs: usize,
3096 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3097 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3098 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3099 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
3100 let (inf, nff, ne, rbg, rbu, nu, npi) = (in_f as i32, n_ff as i32, n_expert as i32,
3101 rb_g as i64, rb_u as i64, n_used as i32,
3102 n_pairs as i32);
3103 let f = self.func("moe_gate_up_gelu8_dev_q8_csr");
3104 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
3105 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3106 let __s_b = self.gpu.stream();
3107 let mut b = __s_b.launch_builder(&f);
3108 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3109 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
3110 unsafe { b.launch(cfg)?; }
3111 Ok(act)
3112 }
3113
3114 #[allow(clippy::too_many_arguments)]
3116 pub fn moe_down8_fma_dev_q8_rows_g(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3117 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3118 dst: &mut CudaSlice<f32>, t: usize,
3119 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3120 qt: i32, rb: usize)
3121 -> Result<(), Box<dyn std::error::Error>> {
3122 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3123 n_expert as i32, rb as i64);
3124 let f = self.func("moe_down8_fma_dev_q8_rows_g");
3125 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, t as u32),
3126 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3127 let __s_b = self.gpu.stream();
3128 let mut b = __s_b.launch_builder(&f);
3129 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3130 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3131 unsafe { b.launch(cfg)?; }
3132 Ok(())
3133 }
3134
3135 pub fn rp_probe_q4(&self, m: usize) -> Result<(f64, f64), Box<dyn std::error::Error>> {
3139 let (out_f, in_f) = (2048usize, 2816usize);
3140 let nblk = in_f / 32;
3141 let mut seed = 0x9E3779B97F4A7C15u64;
3142 let mut rng = move || { seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407); (seed >> 33) as u8 };
3143 let mut w = vec![0u8; out_f * nblk * 18];
3144 for b in w.iter_mut() { *b = rng(); }
3145 for r in 0..out_f {
3146 for g in 0..nblk {
3147 let off = (r * nblk + g) * 18;
3148 w[off] = 0x00; w[off + 1] = 0x2C; }
3150 }
3151 let qplane = out_f * nblk * 16;
3152 let mut wrp = vec![0u8; w.len()];
3153 for r in 0..out_f {
3154 for g in 0..nblk {
3155 let src = &w[(r * nblk + g) * 18..(r * nblk + g) * 18 + 18];
3156 wrp[qplane + (r * nblk + g) * 2..qplane + (r * nblk + g) * 2 + 2]
3157 .copy_from_slice(&src[0..2]);
3158 wrp[(r * nblk + g) * 16..(r * nblk + g) * 16 + 16].copy_from_slice(&src[2..18]);
3159 }
3160 }
3161 let w_d = self.htod_bytes(&w)?;
3162 let wrp_d = self.htod_bytes(&wrp)?;
3163 let mut aq = vec![0i8; m * in_f];
3164 for v in aq.iter_mut() { *v = rng() as i8; }
3165 let aq_d = self.htod_i8(&aq)?;
3166 let ad_d = self.htod(&vec![0.03125f32; m * nblk])?;
3167 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
3168 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
3169 const RPB: u32 = 4;
3170 let cfg = LaunchConfig { grid_dim: ((out_f as u32).div_ceil(RPB), 1, 1),
3171 block_dim: (32, RPB, 1), shared_mem_bytes: 0 };
3172 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
3173 let (rb, qp) = ((nblk * 18) as i64, qplane as i64);
3174 let fb = self.func("qmatvec_q4_0_mmvq_b4");
3175 let fr = self.func("qmatvec_q4_0_mmvq_b4_rp");
3176 {
3177 let __s_b = self.gpu.stream();
3178 let mut b = __s_b.launch_builder(&fb);
3179 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3180 unsafe { b.launch(cfg)?; }
3181 let __s_b = self.gpu.stream();
3182 let mut b = __s_b.launch_builder(&fr);
3183 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1).arg(&inf).arg(&outf).arg(&mi).arg(&qp);
3184 unsafe { b.launch(cfg)?; }
3185 }
3186 self.gpu.stream().synchronize()?;
3187 let (h0, h1) = (self.dtoh(&y0)?, self.dtoh(&y1)?);
3188 let nd = h0.iter().zip(&h1).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
3189 if nd != 0 { return Err(format!("rp twin not bitwise: {nd}/{} diffs", h0.len()).into()); }
3190 let mut time = |rp: bool| -> Result<f64, Box<dyn std::error::Error>> {
3191 self.gpu.stream().synchronize()?;
3192 let t0 = std::time::Instant::now();
3193 for _ in 0..500 {
3194 if rp {
3195 let __s_b = self.gpu.stream();
3196 let mut b = __s_b.launch_builder(&fr);
3197 b.arg(&wrp_d).arg(&aq_d).arg(&ad_d).arg(&mut y1)
3198 .arg(&inf).arg(&outf).arg(&mi).arg(&qp);
3199 unsafe { b.launch(cfg)?; }
3200 } else {
3201 let __s_b = self.gpu.stream();
3202 let mut b = __s_b.launch_builder(&fb);
3203 b.arg(&w_d).arg(&aq_d).arg(&ad_d).arg(&mut y0)
3204 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3205 unsafe { b.launch(cfg)?; }
3206 }
3207 }
3208 self.gpu.stream().synchronize()?;
3209 Ok(t0.elapsed().as_secs_f64() * 1e6 / 500.0)
3210 };
3211 let _ = time(false)?; let _ = time(true)?; Ok((time(false)?, time(true)?))
3213 }
3214
3215 pub fn build_q4_rp4(&self, t: &mut crate::model::GpuTensor)
3220 -> Result<(), Box<dyn std::error::Error>> {
3221 use crate::model::GpuTensor;
3222 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3223 if *qtype != QT_Q4_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3224 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3225 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 18 { return Ok(()); }
3226 let nblk = in_f / 32;
3227 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 18)?;
3228 let f = self.func("q4_0_split_rp_build");
3229 let n = (out_f * nblk) as i32;
3230 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3231 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3232 let (of, nb) = (out_f as i32, nblk as i32);
3233 let _ = n;
3234 let __s_b = self.gpu.stream();
3235 let mut b = __s_b.launch_builder(&f);
3236 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3237 unsafe { b.launch(cfg)?; }
3238 *rp4 = Some(dst);
3239 Ok(())
3240 }
3241
3242 pub fn build_q8_rp4(&self, t: &mut crate::model::GpuTensor)
3247 -> Result<(), Box<dyn std::error::Error>> {
3248 use crate::model::GpuTensor;
3249 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3250 if *qtype != QT_Q8_0 || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3251 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3252 if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 34 { return Ok(()); }
3253 *rp4 = Some(self.build_q8_rp4_raw(bytes, in_f, out_f)?);
3254 Ok(())
3255 }
3256
3257 pub fn build_q8_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize)
3260 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3261 assert!(in_f % 32 == 0);
3262 let nblk = in_f / 32;
3263 let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 34)?;
3264 let f = self.func("q8_0_split_rp_build");
3265 let cfg = LaunchConfig { grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
3266 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3267 let (of, nb) = (out_f as i32, nblk as i32);
3268 let __s_b = self.gpu.stream();
3269 let mut b = __s_b.launch_builder(&f);
3270 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3271 unsafe { b.launch(cfg)?; }
3272 Ok(dst)
3273 }
3274
3275 pub fn build_q4k_rp4(&self, t: &mut crate::model::GpuTensor)
3283 -> Result<(), Box<dyn std::error::Error>> {
3284 use crate::model::GpuTensor;
3285 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3286 if *qtype != QT_Q4_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3287 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3288 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 144 { return Ok(()); }
3289 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q4_K)?);
3290 Ok(())
3291 }
3292
3293 pub fn build_q6k_rp4(&self, t: &mut crate::model::GpuTensor)
3294 -> Result<(), Box<dyn std::error::Error>> {
3295 use crate::model::GpuTensor;
3296 let GpuTensor::Quant { bytes, qtype, row_bytes, ne, rp4, .. } = t else { return Ok(()) };
3297 if *qtype != QT_Q6_K || rp4.is_some() || ne.len() != 2 { return Ok(()); }
3298 let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
3299 if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 210 { return Ok(()); }
3300 *rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q6_K)?);
3301 Ok(())
3302 }
3303
3304 pub fn build_kq_rp4_raw(&self, bytes: &CudaSlice<u8>, in_f: usize, out_f: usize, qtype: i32)
3306 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
3307 assert!(in_f % 256 == 0);
3308 let nsbk = in_f / 256;
3309 let (sb_bytes, kname) = match qtype {
3310 QT_Q4_K => (144usize, "q4_K_split_rp_build"),
3311 QT_Q6_K => (210usize, "q6_K_split_rp_build"),
3312 _ => return Err(format!("build_kq_rp4_raw: qtype {qtype} has no rp mirror").into()),
3313 };
3314 let mut dst = self.alloc_uninit::<u8>(out_f * nsbk * sb_bytes)?;
3315 let f = self.func(kname);
3316 let cfg = LaunchConfig { grid_dim: (((out_f * nsbk) as u32).div_ceil(256), 1, 1),
3317 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3318 let (of, nb) = (out_f as i32, nsbk as i32);
3319 let __s_b = self.gpu.stream();
3320 let mut b = __s_b.launch_builder(&f);
3321 b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
3322 unsafe { b.launch(cfg)?; }
3323 Ok(dst)
3324 }
3325
3326 pub fn kqrp_enabled() -> bool {
3330 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3331 *ON.get_or_init(|| match std::env::var("MEMRA_KQRP").as_deref() {
3332 Ok("0") => false,
3333 Ok(_) => true,
3334 Err(_) => cfg!(memra_hopper_mma),
3335 })
3336 }
3337
3338 pub fn build_q4_rp_swap(&self, t: &mut crate::model::GpuTensor)
3344 -> Result<bool, Box<dyn std::error::Error>> {
3345 self.build_q4_rp4(t)?;
3346 self.gpu.stream().synchronize()?; use crate::model::GpuTensor;
3348 let GpuTensor::Quant { bytes, rp4, rp, .. } = t else { return Ok(false) };
3349 match rp4.take() {
3350 Some(split) => {
3351 *bytes = split; *rp = true;
3353 Ok(true)
3354 }
3355 None => Ok(false),
3356 }
3357 }
3358
3359 pub fn q4rp_enabled() -> bool {
3361 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3362 *ON.get_or_init(|| std::env::var("MEMRA_Q4RP").map(|v| v != "0").unwrap_or(true))
3363 }
3364
3365 pub fn copy_rows_strided(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
3368 row_elems: usize, n_rows: usize, src_stride: usize, src_off: usize)
3369 -> Result<(), Box<dyn std::error::Error>> {
3370 let f = self.func("copy_rows_strided_f32");
3371 let cfg = LaunchConfig { grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
3372 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3373 let (re, nr) = (row_elems as i32, n_rows as i32);
3374 let (st, off) = (src_stride as i64, src_off as i64);
3375 let __s_b = self.gpu.stream();
3376 let mut b = __s_b.launch_builder(&f);
3377 b.arg(src).arg(&mut *dst).arg(&re).arg(&nr).arg(&st).arg(&off);
3378 unsafe { b.launch(cfg)?; }
3379 Ok(())
3380 }
3381
3382 pub fn u32_set_k(&self, dst: &mut CudaSlice<u32>, v: u32, idx: usize)
3384 -> Result<(), Box<dyn std::error::Error>> {
3385 let f = self.func("u32_set_k");
3386 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3387 let ii = idx as i32;
3388 let __s_b = self.gpu.stream();
3389 let mut b = __s_b.launch_builder(&f);
3390 b.arg(dst).arg(&v).arg(&ii);
3391 unsafe { b.launch(cfg)?; }
3392 Ok(())
3393 }
3394
3395 pub fn i32_add_k(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
3397 let f = self.func("i32_add_k");
3398 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3399 let __s_b = self.gpu.stream();
3400 let mut b = __s_b.launch_builder(&f);
3401 b.arg(d).arg(&v);
3402 unsafe { b.launch(cfg)?; }
3403 Ok(())
3404 }
3405
3406 pub fn i32_iota_from(&self, ctr: &CudaSlice<i32>, dst: &mut CudaSlice<i32>, n: usize)
3408 -> Result<(), Box<dyn std::error::Error>> {
3409 let f = self.func("i32_iota_from");
3410 let cfg = LaunchConfig::for_num_elems(n as u32);
3411 let ni = n as i32;
3412 let __s_b = self.gpu.stream();
3413 let mut b = __s_b.launch_builder(&f);
3414 b.arg(ctr).arg(dst).arg(&ni);
3415 unsafe { b.launch(cfg)?; }
3416 Ok(())
3417 }
3418
3419 pub fn u32_map_k(&self, buf: &mut CudaSlice<u32>, map: &CudaSlice<u32>, idx: usize)
3421 -> Result<(), Box<dyn std::error::Error>> {
3422 let f = self.func("u32_map_k");
3423 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
3424 let ii = idx as i32;
3425 let __s_b = self.gpu.stream();
3426 let mut b = __s_b.launch_builder(&f);
3427 b.arg(buf).arg(map).arg(&ii);
3428 unsafe { b.launch(cfg)?; }
3429 Ok(())
3430 }
3431
3432 #[allow(clippy::too_many_arguments)]
3434 pub fn u32_pack2(&self, a: &CudaSlice<u32>, off_a: usize, n1: usize,
3435 b_in: &CudaSlice<u32>, n2: usize, out: &mut CudaSlice<u32>)
3436 -> Result<(), Box<dyn std::error::Error>> {
3437 let f = self.func("u32_pack2");
3438 let cfg = LaunchConfig::for_num_elems((n1 + n2) as u32);
3439 let (oa, i1, i2) = (off_a as i32, n1 as i32, n2 as i32);
3440 let __s_b = self.gpu.stream();
3441 let mut b = __s_b.launch_builder(&f);
3442 b.arg(a).arg(&oa).arg(&i1).arg(b_in).arg(&i2).arg(out);
3443 unsafe { b.launch(cfg)?; }
3444 Ok(())
3445 }
3446
3447 pub fn moe_w_exscale(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3449 s: &CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
3450 let f = self.func("moe_w_exscale");
3451 let cfg = LaunchConfig::for_num_elems(n as u32);
3452 let ni = n as i32;
3453 let __s_b = self.gpu.stream();
3454 let mut b = __s_b.launch_builder(&f);
3455 b.arg(w).arg(sel).arg(s).arg(&ni);
3456 unsafe { b.launch(cfg)?; }
3457 Ok(())
3458 }
3459
3460 pub fn moe_w_scale_by_expert(&self, w: &mut CudaSlice<f32>, sel: &CudaSlice<i32>,
3463 macros: &CudaSlice<f32>, n_expert: usize, n: usize)
3464 -> Result<(), Box<dyn std::error::Error>> {
3465 let f = self.func("moe_w_scale_by_expert");
3466 let cfg = LaunchConfig { grid_dim: (n.div_ceil(64) as u32, 1, 1),
3467 block_dim: (64, 1, 1), shared_mem_bytes: 0 };
3468 let (ne, nn) = (n_expert as i32, n as i32);
3469 let __s_b = self.gpu.stream();
3470 let mut b = __s_b.launch_builder(&f);
3471 b.arg(w).arg(sel).arg(macros).arg(&ne).arg(&nn);
3472 unsafe { b.launch(cfg)?; }
3473 Ok(())
3474 }
3475
3476 pub fn moe_gate_up_silu8_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3477 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3478 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3479 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3480 macros: &CudaSlice<f32>)
3481 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3482 static GU: std::sync::OnceLock<(String, u32)> = std::sync::OnceLock::new();
3483 let (mode, wpb) = GU.get_or_init(|| {
3484 let mode = std::env::var("MEMRA_MOE_DEVQ8_GU").unwrap_or_default();
3485 let wpb = std::env::var("MEMRA_MOE_DEVQ8_WPB").ok()
3486 .and_then(|v| v.parse().ok()).unwrap_or(4u32).clamp(1, 16);
3487 (mode, wpb)
3488 });
3489 let (mode, wpb) = (mode.as_str(), *wpb);
3490 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3491 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3492 rb_g as i64, rb_u as i64);
3493 let (f, cfg) = match mode {
3494 "1" | "2" | "4" => {
3495 let rpw: u32 = mode.parse().unwrap();
3496 let f = self.func(match rpw { 1 => "moe_gate_up_silu8_dev_q8_r1",
3497 2 => "moe_gate_up_silu8_dev_q8_r2",
3498 _ => "moe_gate_up_silu8_dev_q8_r4" });
3499 let rows_per_block = (rpw * wpb) as usize;
3500 let gx = n_ff.div_ceil(rows_per_block) as u32;
3501 (f, LaunchConfig { grid_dim: (gx, n_used as u32, 1),
3502 block_dim: (32, wpb, 1), shared_mem_bytes: 0 })
3503 }
3504 "j8" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8"),
3505 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3506 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3507 "vsm2" => {
3509 let f = self.func("moe_gate_up_silu8_dev_q8_vsm2");
3510 let sh = (rb_g + rb_u) as u32;
3511 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3512 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3513 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3514 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3515 }
3516 "vsm" => {
3517 let f = self.func("moe_gate_up_silu8_dev_q8_vsm");
3518 let sh = (rb_g + rb_u) as u32;
3519 use cudarc::driver::sys::CUfunction_attribute_enum as A;
3520 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
3521 (f, LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3522 block_dim: (32, 1, 1), shared_mem_bytes: sh })
3523 }
3524 "sg" => (self.func("moe_gate_up_silu8_dev_q8_sg"),
3525 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3526 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3527 "j8sg" if n_used <= 32 => (self.func("moe_gate_up_silu8_dev_q8_j8sg"),
3528 LaunchConfig { grid_dim: (n_ff as u32, 1, 1),
3529 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3530 "u64" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_u64"),
3531 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3532 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3533 "gs4" if in_f == 2048 => (self.func("moe_gate_up_silu8_dev_q8_gs4"),
3534 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3535 block_dim: (32, 4, 1), shared_mem_bytes: 0 }),
3536 "v" | "" => (self.func("moe_gate_up_silu8_dev_q8_v"),
3538 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3539 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3540 "s2" => (self.func("moe_gate_up_silu8_dev_q8_s2"),
3541 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3542 block_dim: (32, 2, 1), shared_mem_bytes: 0 }),
3543 "s2z" => {
3544 let rz = wpb.min(16); (self.func("moe_gate_up_silu8_dev_q8_s2z"),
3546 LaunchConfig { grid_dim: (n_ff.div_ceil(rz as usize) as u32, n_used as u32, 1),
3547 block_dim: (32, 2, rz), shared_mem_bytes: 0 })
3548 }
3549 _ => (self.func("moe_gate_up_silu8_dev_q8"),
3550 LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3551 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3552 };
3553 let __s_b = self.gpu.stream();
3554 let mut b = __s_b.launch_builder(&f);
3555 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3556 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3557 unsafe { b.launch(cfg)?; }
3558 Ok(act)
3559 }
3560
3561 #[allow(clippy::too_many_arguments)]
3562 pub fn moe_down8_fma_dev_q8(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3563 w: &cudarc::driver::CudaView<f32>,
3564 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3565 dst: &mut cudarc::driver::CudaViewMut<f32>,
3566 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3567 qt: i32, rb: usize)
3568 -> Result<(), Box<dyn std::error::Error>> {
3569 static DOWN: std::sync::OnceLock<String> = std::sync::OnceLock::new();
3570 let mode = DOWN.get_or_init(|| std::env::var("MEMRA_MOE_DEVQ8_DOWN").unwrap_or_default());
3571 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3572 n_expert as i32, rb as i64);
3573 let (f, cfg) = match mode.as_str() {
3576 m @ ("1" | "2" | "4") if n_used <= 8 => {
3577 let rpw: usize = m.parse().unwrap();
3578 let f = self.func(match rpw { 1 => "moe_down8_fma_dev_q8_w8r1",
3579 2 => "moe_down8_fma_dev_q8_w8r2",
3580 _ => "moe_down8_fma_dev_q8_w8r4" });
3581 (f, LaunchConfig { grid_dim: (out_f.div_ceil(rpw) as u32, 1, 1),
3582 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3583 }
3584 "h2" if in_f == 512 => (self.func("moe_down8_fma_dev_q8_h2"),
3585 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3586 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3587 "" if in_f == 704 && n_used <= 8 =>
3590 (self.func("moe_down8_fma_dev_q8_w8r2"),
3591 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3592 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3593 "w8h2v" | "" if in_f == 512 && n_used <= 8 =>
3597 (self.func("moe_down8_fma_dev_q8_w8h2v"),
3598 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3599 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3600 "w8h2r2v" if in_f == 512 && n_used <= 8 =>
3601 (self.func("moe_down8_fma_dev_q8_w8h2r2v"),
3602 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3603 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3604 "w8h2r2" if in_f == 512 && n_used <= 8 =>
3605 (self.func("moe_down8_fma_dev_q8_w8h2r2"),
3606 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3607 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3608 "w8h2" if in_f == 512 && n_used <= 8 =>
3609 (self.func("moe_down8_fma_dev_q8_w8h2"),
3610 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3611 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 }),
3612 _ => (self.func("moe_down8_fma_dev_q8"),
3613 LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3614 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3615 };
3616 let __s_b = self.gpu.stream();
3617 let mut b = __s_b.launch_builder(&f);
3618 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3619 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3620 unsafe { b.launch(cfg)?; }
3621 Ok(())
3622 }
3623
3624 #[allow(clippy::too_many_arguments)]
3631 pub fn moe_gate_up_silu8_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3632 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, t: usize,
3633 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3634 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3635 macros: &CudaSlice<f32>)
3636 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3637 let f = self.func("moe_gate_up_silu8_dev_q8_v_rows");
3638 let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
3639 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, t as u32),
3640 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3641 let (inf, nff, ne, nu, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3642 n_used as i32, rb_g as i64, rb_u as i64);
3643 let __s_b = self.gpu.stream();
3644 let mut b = __s_b.launch_builder(&f);
3645 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3646 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(macros);
3647 unsafe { b.launch(cfg)?; }
3648 Ok(act)
3649 }
3650
3651 #[allow(clippy::too_many_arguments)]
3656 pub fn moe_down8_fma_dev_q8_rows(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3657 w: &CudaSlice<f32>, aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3658 dst: &mut CudaSlice<f32>, t: usize,
3659 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3660 qt: i32, rb: usize)
3661 -> Result<(), Box<dyn std::error::Error>> {
3662 assert!(in_f == 512 && n_used <= 8, "down rows twin is w8h2v shape-gated");
3663 let f = self.func("moe_down8_fma_dev_q8_w8h2v_rows");
3664 let cfg = LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, t as u32),
3665 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 };
3666 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3667 n_expert as i32, rb as i64);
3668 let __s_b = self.gpu.stream();
3669 let mut b = __s_b.launch_builder(&f);
3670 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3671 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3672 unsafe { b.launch(cfg)?; }
3673 Ok(())
3674 }
3675
3676 #[allow(clippy::too_many_arguments)]
3680 pub fn moe_gate_up_silu8_dev_q8_csr(&self, table: &CudaSlice<u64>, sel: &CudaSlice<i32>,
3681 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3682 n_pairs: usize, in_f: usize, n_ff: usize, n_used: usize,
3683 n_expert: usize, qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize)
3684 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3685 let f = self.func("moe_gate_up_silu8_dev_q8_csr_iq4");
3686 let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
3687 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_pairs as u32, 1),
3688 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3689 let (inf, nff, ne, nu, npi, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3690 n_used as i32, n_pairs as i32, rb_g as i64, rb_u as i64);
3691 let __s_b = self.gpu.stream();
3692 let mut b = __s_b.launch_builder(&f);
3693 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3694 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(&nu).arg(&npi);
3695 unsafe { b.launch(cfg)?; }
3696 Ok(act)
3697 }
3698
3699
3700 #[allow(clippy::too_many_arguments)]
3704 pub fn moe_down8_fma_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3705 sel: &cudarc::driver::CudaView<i32>,
3706 w: &cudarc::driver::CudaView<f32>,
3707 aq2: &CudaSlice<i8>, ad2: &CudaSlice<f32>,
3708 dst: &mut cudarc::driver::CudaViewMut<f32>,
3709 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3710 qt: i32, rb: usize)
3711 -> Result<(), Box<dyn std::error::Error>> {
3712 let (inf, outf, nu, ne, rbi) = (in_f as i32, out_f as i32, n_used as i32,
3713 n_expert as i32, rb as i64);
3714 let (f, cfg) = match variant {
3715 "w8h2" | "w8h2v" => {
3716 (self.func(if variant == "w8h2" { "moe_down8_fma_dev_q8_w8h2" }
3717 else { "moe_down8_fma_dev_q8_w8h2v" }),
3718 LaunchConfig { grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
3719 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3720 }
3721 "w8h2r2" | "w8h2r2v" => {
3722 (self.func(if variant == "w8h2r2" { "moe_down8_fma_dev_q8_w8h2r2" }
3723 else { "moe_down8_fma_dev_q8_w8h2r2v" }),
3724 LaunchConfig { grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
3725 block_dim: (32, n_used as u32, 1), shared_mem_bytes: 0 })
3726 }
3727 _ => (self.func("moe_down8_fma_dev_q8"),
3728 LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3729 block_dim: (32, 1, 1), shared_mem_bytes: 0 }),
3730 };
3731 let __s_b = self.gpu.stream();
3732 let mut b = __s_b.launch_builder(&f);
3733 b.arg(table).arg(sel).arg(w).arg(aq2).arg(ad2).arg(dst)
3734 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbi);
3735 unsafe { b.launch(cfg)?; }
3736 Ok(())
3737 }
3738
3739 #[allow(clippy::too_many_arguments)]
3741 pub fn moe_gate_up_silu8_dev_q8_variant(&self, variant: &str, table: &CudaSlice<u64>,
3742 sel: &cudarc::driver::CudaView<i32>,
3743 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
3744 in_f: usize, n_ff: usize, n_used: usize,
3745 n_expert: usize, qt_g: i32, qt_u: i32,
3746 rb_g: usize, rb_u: usize)
3747 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3748 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
3749 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3750 rb_g as i64, rb_u as i64);
3751 let f = self.func(if variant == "v" { "moe_gate_up_silu8_dev_q8_v" }
3752 else { "moe_gate_up_silu8_dev_q8" });
3753 let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3754 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
3755 let __s_b = self.gpu.stream();
3756 let mut b = __s_b.launch_builder(&f);
3757 b.arg(table).arg(sel).arg(aq).arg(ad).arg(&mut act)
3758 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu);
3759 unsafe { b.launch(cfg)?; }
3760 Ok(act)
3761 }
3762
3763 pub fn moe_gate_up_silu8_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3764 x: &cudarc::driver::CudaView<f32>,
3765 in_f: usize, n_ff: usize, n_used: usize, n_expert: usize,
3766 qt_g: i32, qt_u: i32, rb_g: usize, rb_u: usize,
3767 macros: &CudaSlice<f32>)
3768 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3769 let f = self.func("moe_gate_up_silu8_dev");
3770 let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?; let cfg = LaunchConfig { grid_dim: (n_ff as u32, n_used as u32, 1),
3772 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3773 let (inf, nff, ne, rbg, rbu) = (in_f as i32, n_ff as i32, n_expert as i32,
3774 rb_g as i64, rb_u as i64);
3775 let __s_b = self.gpu.stream();
3776 let mut b = __s_b.launch_builder(&f);
3777 b.arg(table).arg(sel).arg(x).arg(&mut act)
3778 .arg(&inf).arg(&nff).arg(&ne).arg(&qt_g).arg(&qt_u).arg(&rbg).arg(&rbu).arg(macros);
3779 unsafe { b.launch(cfg)?; }
3780 Ok(act)
3781 }
3782
3783 #[allow(clippy::too_many_arguments)]
3786 pub fn moe_down8_fma_dev(&self, table: &CudaSlice<u64>, sel: &cudarc::driver::CudaView<i32>,
3787 w: &cudarc::driver::CudaView<f32>, act: &CudaSlice<f32>,
3788 dst: &mut cudarc::driver::CudaViewMut<f32>,
3789 in_f: usize, out_f: usize, n_used: usize, n_expert: usize,
3790 qt: i32, rb: usize)
3791 -> Result<(), Box<dyn std::error::Error>> {
3792 let f = self.func("moe_down8_fma_dev");
3793 let cfg = LaunchConfig { grid_dim: (out_f as u32, 1, 1),
3794 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
3795 let (inf, outf, nu, ne, rbv) = (in_f as i32, out_f as i32, n_used as i32,
3796 n_expert as i32, rb as i64);
3797 let __s_b = self.gpu.stream();
3798 let mut b = __s_b.launch_builder(&f);
3799 b.arg(table).arg(sel).arg(w).arg(act).arg(dst)
3800 .arg(&inf).arg(&outf).arg(&nu).arg(&ne).arg(&qt).arg(&rbv);
3801 unsafe { b.launch(cfg)?; }
3802 Ok(())
3803 }
3804
3805 pub fn axpy_into(&self, src: &CudaSlice<f32>, alpha: f32,
3807 dst: &mut cudarc::driver::CudaViewMut<f32>, n: usize)
3808 -> Result<(), Box<dyn std::error::Error>> {
3809 let f = self.func("axpy_f32");
3810 let cfg = LaunchConfig::for_num_elems(n as u32);
3811 let (a, ni) = (alpha, n as i32);
3812 let __s_b = self.gpu.stream();
3813 let mut b = __s_b.launch_builder(&f);
3814 b.arg(src).arg(dst).arg(&a).arg(&ni);
3815 unsafe { b.launch(cfg)?; }
3816 Ok(())
3817 }
3818
3819 pub fn add_scaled_rows(&self, src: &CudaSlice<f32>, scale: &CudaSlice<f32>,
3821 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
3822 -> Result<(), Box<dyn std::error::Error>> {
3823 let f = self.func("add_scaled_rows_f32");
3824 let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
3825 let (nc, nr) = (ncols as i32, nrows as i32);
3826 let __s_b = self.gpu.stream();
3827 let mut b = __s_b.launch_builder(&f);
3828 b.arg(src).arg(scale).arg(dst).arg(&nc).arg(&nr);
3829 unsafe { b.launch(cfg)?; }
3830 Ok(())
3831 }
3832
3833 pub fn gather_rows(&self, src: &CudaSlice<f32>, idx: &CudaSlice<i32>,
3837 dst: &mut CudaSlice<f32>, ncols: usize, m_e: usize)
3838 -> Result<(), Box<dyn std::error::Error>> {
3839 let f = self.func("gather_rows_f32");
3840 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3841 let (nc, me) = (ncols as i32, m_e as i32);
3842 let __s_b = self.gpu.stream();
3843 let mut b = __s_b.launch_builder(&f);
3844 b.arg(src).arg(idx).arg(dst).arg(&nc).arg(&me);
3845 unsafe { b.launch(cfg)?; }
3846 Ok(())
3847 }
3848
3849 pub fn scatter_slot(&self, src: &CudaSlice<f32>, tok_idx: &CudaSlice<i32>,
3854 slot_idx: &CudaSlice<i32>, weight: &CudaSlice<f32>,
3855 dst: &mut CudaSlice<f32>, wbuf: &mut CudaSlice<f32>,
3856 ncols: usize, n_used: usize, m_e: usize)
3857 -> Result<(), Box<dyn std::error::Error>> {
3858 let f = self.func("scatter_add_slot_f32");
3859 let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
3860 let (nc, nu, me) = (ncols as i32, n_used as i32, m_e as i32);
3861 let __s_b = self.gpu.stream();
3862 let mut b = __s_b.launch_builder(&f);
3863 b.arg(src).arg(tok_idx).arg(slot_idx).arg(weight).arg(dst).arg(wbuf).arg(&nc).arg(&nu).arg(&me);
3864 unsafe { b.launch(cfg)?; }
3865 Ok(())
3866 }
3867
3868 pub fn reduce_slots(&self, slots: &CudaSlice<f32>, wbuf: &CudaSlice<f32>,
3872 dst: &mut CudaSlice<f32>, ncols: usize, n_used: usize, t: usize)
3873 -> Result<(), Box<dyn std::error::Error>> {
3874 let f = self.func("reduce_slots_f32");
3875 let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
3876 let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
3877 let __s_b = self.gpu.stream();
3878 let mut b = __s_b.launch_builder(&f);
3879 b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
3880 unsafe { b.launch(cfg)?; }
3881 Ok(())
3882 }
3883
3884 pub fn quantize_q8_1_view(&self, x: &cudarc::driver::CudaView<f32>, m: usize, in_f: usize)
3891 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3892 let f = self.func("quantize_q8_1");
3893 let nblk = in_f / 32;
3894 let mut q = self.alloc_uninit::<i8>(m * in_f)?;
3895 let mut d = self.alloc_uninit::<f32>(m * nblk)?;
3896 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
3897 let (inf, mi) = (in_f as i32, m as i32);
3898 let __s_b = self.gpu.stream();
3899 let mut b = __s_b.launch_builder(&f);
3900 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3901 unsafe { b.launch(cfg)?; }
3902 Ok((q, d))
3903 }
3904
3905 pub fn quantize_q8_1(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3906 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3907 let nblk = in_f / 32;
3908 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);
3912 let (inf, mi) = (in_f as i32, m as i32);
3913 if Self::pdl_on() && Self::pdl_wb_on() {
3914 {
3915 use cudarc::driver::{DevicePtr, DevicePtrMut};
3916 let s = &self.gpu.stream();
3917 let (px, _g0) = x.device_ptr(s);
3918 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
3919 let mut ps = [
3920 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
3921 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
3922 &mi as *const _ as *mut _,
3923 ];
3924 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
3925 }
3926 return Ok((q, d));
3927 }
3928 let f = self.func("quantize_q8_1");
3929 let __s_b = self.gpu.stream();
3930 let mut b = __s_b.launch_builder(&f);
3931 b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
3932 unsafe { b.launch(cfg)?; }
3933 Ok((q, d))
3934 }
3935
3936 pub fn quantize_fp4_act(&self, x: &CudaSlice<f32>, m: usize, in_f: usize)
3940 -> Result<(CudaSlice<u32>, CudaSlice<u8>), Box<dyn std::error::Error>> {
3941 let f = self.func("quantize_fp4_act");
3942 let nb16 = in_f / 16;
3943 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);
3946 let (inf, mi) = (in_f as i32, m as i32);
3947 let __s_b = self.gpu.stream();
3948 let mut b = __s_b.launch_builder(&f);
3949 b.arg(x).arg(&mut aq4).arg(&mut ad4).arg(&inf).arg(&mi);
3950 unsafe { b.launch(cfg)?; }
3951 Ok((aq4, ad4))
3952 }
3953
3954 pub fn qmatvec_gemm_nvfp4_fp4(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3959 in_f: usize, out_f: usize, row_bytes: usize, scale: f32)
3960 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3961 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
3962 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
3963 let mut y = self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)?;
3964 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
3965 Ok(y)
3966 }
3967
3968 fn fp4_gemm_launch(&self, bytes: &CudaSlice<u8>, aq4: &CudaSlice<u32>, ad4: &CudaSlice<u8>,
3971 m: usize, in_f: usize, out_f: usize, row_bytes: usize)
3972 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3973 let f = self.func("qmatvec_gemm_nvfp4_fp4");
3974 let mut y = self.alloc_uninit::<f32>(m * out_f)?; const BM: u32 = 64; const BN: u32 = 256;
3976 let cfg = LaunchConfig {
3977 grid_dim: ((out_f as u32 + BM - 1) / BM, (m as u32 + BN - 1) / BN, 1),
3978 block_dim: (32, 4, 1), shared_mem_bytes: 0,
3979 };
3980 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
3981 let __s_b = self.gpu.stream();
3982 let mut b = __s_b.launch_builder(&f);
3983 b.arg(bytes).arg(aq4).arg(ad4).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
3984 unsafe { b.launch(cfg)?; }
3985 Ok(y)
3986 }
3987
3988 pub fn qmatvec_gemm_nvfp4_fp4_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
3990 in_f: usize, out_f: usize, row_bytes: usize)
3991 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3992 assert!(in_f % 64 == 0, "FP4 GEMM requires in_f % 64 == 0, got {in_f}");
3993 let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
3994 self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)
3995 }
3996
3997 pub fn qmatvec_q8_0_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
3999 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4000 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
4001 let f = self.func("qmatvec_q8_0_dp4a");
4002 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
4004 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
4005 let __s_b = self.gpu.stream();
4006 let mut b = __s_b.launch_builder(&f);
4007 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
4008 unsafe { b.launch(cfg)?; }
4009 Ok(y)
4010 }
4011
4012 #[allow(non_snake_case)] pub fn qmatvec_q4_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4015 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4016 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
4017 let f = self.func("qmatvec_q4_K_dp4a");
4018 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
4020 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
4021 let __s_b = self.gpu.stream();
4022 let mut b = __s_b.launch_builder(&f);
4023 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
4024 unsafe { b.launch(cfg)?; }
4025 Ok(y)
4026 }
4027
4028 #[allow(non_snake_case)] pub fn qmatvec_q6_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4031 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4032 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
4033 let f = self.func("qmatvec_q6_K_dp4a");
4034 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
4036 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
4037 let __s_b = self.gpu.stream();
4038 let mut b = __s_b.launch_builder(&f);
4039 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
4040 unsafe { b.launch(cfg)?; }
4041 Ok(y)
4042 }
4043
4044 #[allow(non_snake_case)] pub fn qmatvec_q5_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4047 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4048 self.qmatvec_dp4a_named("qmatvec_q5_K_dp4a", w, x, m, in_f, out_f, row_bytes)
4049 }
4050 #[allow(non_snake_case)] pub fn qmatvec_q3_K_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4053 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4054 self.qmatvec_dp4a_named("qmatvec_q3_K_dp4a", w, x, m, in_f, out_f, row_bytes)
4055 }
4056 pub fn qmatvec_nvfp4_fast_rp(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4058 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4059 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
4060 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a_rp", w, x, m, in_f, out_f, row_bytes)
4061 }
4062 pub fn qmatvec_nvfp4_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4064 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4065 assert!(in_f % 64 == 0, "NVFP4 dp4a requires in_f % 64 == 0, got {in_f}");
4068 self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a", w, x, m, in_f, out_f, row_bytes)
4069 }
4070 #[allow(non_snake_case)] pub fn qmatvec_iq4_XS_fast(&self, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
4073 out_f: usize, row_bytes: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4074 self.qmatvec_dp4a_named("qmatvec_iq4_XS_dp4a", w, x, m, in_f, out_f, row_bytes)
4075 }
4076
4077 fn qmatvec_dp4a_named(&self, name: &str, w: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
4079 in_f: usize, out_f: usize, row_bytes: usize)
4080 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4081 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
4082 let f = self.func(name);
4083 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
4085 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
4086 let __s_b = self.gpu.stream();
4087 let mut b = __s_b.launch_builder(&f);
4088 b.arg(w).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
4089 unsafe { b.launch(cfg)?; }
4090 Ok(y)
4091 }
4092
4093 pub fn htod(&self, v: &[f32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4094 Ok(self.gpu.stream().clone_htod(v)?)
4095 }
4096 pub fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
4097 Ok(self.gpu.stream().clone_htod(v)?)
4098 }
4099 pub fn htod_i8(&self, v: &[i8]) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
4101 Ok(self.gpu.stream().clone_htod(v)?)
4102 }
4103 pub fn htod_u64(&self, v: &[u64]) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
4104 Ok(self.gpu.stream().clone_htod(v)?)
4105 }
4106 pub fn dtoh_view(&self, d: &cudarc::driver::CudaView<f32>)
4108 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
4109 let v = self.gpu.stream().clone_dtoh(d)?;
4110 self.gpu.stream().synchronize()?;
4111 Ok(v)
4112 }
4113 pub fn dtoh(&self, d: &CudaSlice<f32>) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
4114 let v = self.gpu.stream().clone_dtoh(d)?;
4115 self.gpu.stream().synchronize()?;
4116 Ok(v)
4117 }
4118 pub fn dtoh_pair(
4122 &self,
4123 a: &CudaSlice<f32>,
4124 b: &CudaSlice<f32>,
4125 ) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
4126 let av = self.gpu.stream().clone_dtoh(a)?;
4127 let bv = self.gpu.stream().clone_dtoh(b)?;
4128 self.gpu.stream().synchronize()?;
4129 Ok((av, bv))
4130 }
4131 pub fn dtoh_i32(&self, d: &CudaSlice<i32>) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
4133 let v = self.gpu.stream().clone_dtoh(d)?;
4134 self.gpu.stream().synchronize()?;
4135 Ok(v)
4136 }
4137 pub fn dtoh_u8(&self, d: &CudaSlice<u8>) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
4139 let v = self.gpu.stream().clone_dtoh(d)?;
4140 self.gpu.stream().synchronize()?;
4141 Ok(v)
4142 }
4143 pub fn dtoh_u8_view(&self, d: &cudarc::driver::CudaView<u8>)
4144 -> Result<Vec<u8>, Box<dyn std::error::Error>> {
4145 let v = self.gpu.stream().clone_dtoh(d)?;
4146 self.gpu.stream().synchronize()?;
4147 Ok(v)
4148 }
4149 pub fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4150 let s = self.gpu.stream().alloc_zeros::<f32>(n)?;
4151 self.keep_if_capturing(&s);
4152 Ok(s)
4153 }
4154
4155 pub fn prob_of_token_device(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>, n_vocab: usize)
4164 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4165 let nb = ARGMAX_NB;
4166 let mut part = self.alloc_uninit::<f32>(nb)?;
4167 let mut p = self.alloc_uninit::<f32>(1)?;
4168 let f1 = self.func("prob_of_token_partial_f32");
4169 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4170 let nv = n_vocab as i32;
4171 let __s_b1 = self.gpu.stream();
4172 let mut b1 = __s_b1.launch_builder(&f1);
4173 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
4174 unsafe { b1.launch(cfg1)?; }
4175 let f2 = self.func("prob_of_token_final_f32");
4176 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4177 let nbi = nb as i32;
4178 let __s_b2 = self.gpu.stream();
4179 let mut b2 = __s_b2.launch_builder(&f2);
4180 b2.arg(&part).arg(&mut p).arg(&nbi);
4181 unsafe { b2.launch(cfg2)?; }
4182 Ok(p)
4183 }
4184
4185 pub fn prob_of_token_device_col(&self, logits: &CudaSlice<f32>,
4192 tok_all: &CudaSlice<u32>, tok_idx: usize,
4193 p_out: &mut CudaSlice<f32>, p_idx: usize, n_vocab: usize)
4194 -> Result<(), Box<dyn std::error::Error>> {
4195 let tok_v = tok_all.slice(tok_idx..tok_idx + 1);
4196 let mut p_v = p_out.slice_mut(p_idx..p_idx + 1);
4197 let nb = ARGMAX_NB;
4198 let mut part = self.alloc_uninit::<f32>(nb)?;
4199 let f1 = self.func("prob_of_token_partial_f32");
4200 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4201 let nv = n_vocab as i32;
4202 let __s_b1 = self.gpu.stream();
4203 let mut b1 = __s_b1.launch_builder(&f1);
4204 b1.arg(logits).arg(&tok_v).arg(&mut part).arg(&nv);
4205 unsafe { b1.launch(cfg1)?; }
4206 let f2 = self.func("prob_of_token_final_f32");
4207 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4208 let nbi = nb as i32;
4209 let __s_b2 = self.gpu.stream();
4210 let mut b2 = __s_b2.launch_builder(&f2);
4211 b2.arg(&part).arg(&mut p_v).arg(&nbi);
4212 unsafe { b2.launch(cfg2)?; }
4213 Ok(())
4214 }
4215
4216 pub fn prob_of_token_device_into(&self, logits: &CudaSlice<f32>, tok: &CudaSlice<u32>,
4217 p_out: &mut CudaSlice<f32>, n_vocab: usize)
4218 -> Result<(), Box<dyn std::error::Error>> {
4219 let nb = ARGMAX_NB;
4220 let mut part = self.alloc_uninit::<f32>(nb)?;
4221 let f1 = self.func("prob_of_token_partial_f32");
4222 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4223 let nv = n_vocab as i32;
4224 let __s_b1 = self.gpu.stream();
4225 let mut b1 = __s_b1.launch_builder(&f1);
4226 b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
4227 unsafe { b1.launch(cfg1)?; }
4228 let f2 = self.func("prob_of_token_final_f32");
4229 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4230 let nbi = nb as i32;
4231 let __s_b2 = self.gpu.stream();
4232 let mut b2 = __s_b2.launch_builder(&f2);
4233 b2.arg(&part).arg(p_out).arg(&nbi);
4234 unsafe { b2.launch(cfg2)?; }
4235 Ok(())
4236 }
4237
4238 pub fn argmax_token_device(&self, logits: &CudaSlice<f32>, n_vocab: usize)
4239 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4240 let mut tok = unsafe { self.gpu.stream().alloc::<u32>(1)? };
4241 self.argmax_token_device_into(logits, &mut tok, n_vocab)?;
4242 Ok(tok)
4243 }
4244 pub fn argmax_token_device_into(&self, logits: &CudaSlice<f32>, tok: &mut CudaSlice<u32>,
4251 n_vocab: usize) -> Result<(), Box<dyn std::error::Error>> {
4252 let nb = ARGMAX_NB;
4253 let f1 = self.func("argmax_partial_f32");
4254 let f2 = self.func("argmax_final_f32");
4255 let mut guard = self.argmax_partials.lock().unwrap();
4256 if guard.is_none() {
4257 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4260 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4261 *guard = Some((pv, pi));
4262 }
4263 let (part_v, part_i) = guard.as_mut().unwrap();
4264 let nv = n_vocab as i32;
4265 let nbi = nb as i32;
4266 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4268 let __s_b1 = self.gpu.stream();
4269 let mut b1 = __s_b1.launch_builder(&f1);
4270 b1.arg(logits).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4271 unsafe { b1.launch(cfg1)?; }
4272 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4274 let __s_b2 = self.gpu.stream();
4275 let mut b2 = __s_b2.launch_builder(&f2);
4276 b2.arg(&*part_v).arg(&*part_i).arg(tok).arg(&nbi);
4277 unsafe { b2.launch(cfg2)?; }
4278 Ok(())
4279 }
4280 pub fn argmax_token_device_col(&self, logits: &CudaSlice<f32>, col: usize, n_vocab: usize,
4286 toks: &mut CudaSlice<u32>, out_idx: usize)
4287 -> Result<(), Box<dyn std::error::Error>> {
4288 let nb = ARGMAX_NB;
4289 let f1 = self.func("argmax_partial_f32");
4290 let f2 = self.func("argmax_final_f32");
4291 let mut guard = self.argmax_partials.lock().unwrap();
4292 if guard.is_none() {
4293 let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
4294 let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
4295 *guard = Some((pv, pi));
4296 }
4297 let (part_v, part_i) = guard.as_mut().unwrap();
4298 let col_view = logits.slice(col * n_vocab..(col + 1) * n_vocab);
4299 let nv = n_vocab as i32;
4300 let nbi = nb as i32;
4301 let cfg1 = LaunchConfig { grid_dim: (nb as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4302 let __s_b1 = self.gpu.stream();
4303 let mut b1 = __s_b1.launch_builder(&f1);
4304 b1.arg(&col_view).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
4305 unsafe { b1.launch(cfg1)?; }
4306 let mut tok_view = toks.slice_mut(out_idx..out_idx + 1);
4307 let cfg2 = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4308 let __s_b2 = self.gpu.stream();
4309 let mut b2 = __s_b2.launch_builder(&f2);
4310 b2.arg(&*part_v).arg(&*part_i).arg(&mut tok_view).arg(&nbi);
4311 unsafe { b2.launch(cfg2)?; }
4312 Ok(())
4313 }
4314 pub fn htod_u32_v(&self, v: &[u32]) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4316 Ok(self.gpu.stream().clone_htod(v)?)
4317 }
4318 pub fn dtoh_u32(&self, d: &CudaSlice<u32>) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
4319 let v = self.gpu.stream().clone_dtoh(d)?;
4320 self.gpu.stream().synchronize()?;
4321 Ok(v)
4322 }
4323 pub fn htod_u32_into(&self, dst: &mut CudaSlice<u32>, src: &[u32])
4327 -> Result<(), Box<dyn std::error::Error>> {
4328 let mut view = dst.slice_mut(0..src.len());
4329 self.gpu.stream().memcpy_htod(src, &mut view)?;
4330 Ok(())
4331 }
4332
4333 pub fn htod_i32_into(&self, dst: &mut CudaSlice<i32>, src: &[i32])
4336 -> Result<(), Box<dyn std::error::Error>> {
4337 let mut view = dst.slice_mut(0..src.len());
4338 self.gpu.stream().memcpy_htod(src, &mut view)?;
4339 Ok(())
4340 }
4341
4342 pub fn alloc_u32_zeroed(&self, n: usize) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
4343 let s = self.gpu.stream().alloc_zeros::<u32>(n)?;
4344 self.keep_if_capturing(&s);
4345 Ok(s)
4346 }
4347 pub fn embed_gather_device_into(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4350 x_out: &mut CudaSlice<f32>, n_embd: usize, qtype: i32,
4351 row_bytes: usize) -> Result<(), Box<dyn std::error::Error>> {
4352 let f = self.func("embed_gather_u32");
4353 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
4354 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4355 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
4356 let __s_b = self.gpu.stream();
4357 let mut b = __s_b.launch_builder(&f);
4358 b.arg(embd).arg(token_d).arg(x_out).arg(&ne).arg(&qt).arg(&rb);
4359 unsafe { b.launch(cfg)?; }
4360 Ok(())
4361 }
4362 pub fn dtoh_i32_one(&self, d: &CudaSlice<i32>) -> Result<i32, Box<dyn std::error::Error>> {
4364 let v = self.gpu.stream().clone_dtoh(d)?;
4365 self.gpu.stream().synchronize()?;
4366 Ok(v[0])
4367 }
4368 pub fn i32_set_k(&self, dst: &mut CudaSlice<i32>, v: i32)
4375 -> Result<(), Box<dyn std::error::Error>> {
4376 let f = self.func("i32_set_k");
4377 let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 };
4378 let idx = 0i32;
4379 let __s_b = self.gpu.stream();
4380 let mut b = __s_b.launch_builder(&f);
4381 b.arg(dst).arg(&v).arg(&idx);
4382 unsafe { b.launch(cfg)?; }
4383 Ok(())
4384 }
4385
4386 pub fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
4387 self.gpu.stream().memcpy_htod(&[v], d)?;
4388 Ok(())
4389 }
4390 pub fn set_u32_one(&self, d: &mut CudaSlice<u32>, v: u32) -> Result<(), Box<dyn std::error::Error>> {
4393 self.gpu.stream().memcpy_htod(&[v], d)?;
4394 Ok(())
4395 }
4396 pub fn dtoh_u32_one(&self, d: &CudaSlice<u32>) -> Result<u32, Box<dyn std::error::Error>> {
4398 let v = self.gpu.stream().clone_dtoh(d)?;
4399 self.gpu.stream().synchronize()?;
4400 Ok(v[0])
4401 }
4402 pub fn upload_u8(&self, bytes: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
4404 Ok(self.gpu.stream().clone_htod(bytes)?)
4405 }
4406 pub fn embed_gather_device(&self, embd: &CudaSlice<u8>, token_d: &CudaSlice<u32>,
4410 n_embd: usize, qtype: i32, row_bytes: usize)
4411 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4412 let f = self.func("embed_gather_u32");
4413 let mut x = self.alloc_uninit::<f32>(n_embd)?;
4414 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
4415 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4416 let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
4417 let __s_b = self.gpu.stream();
4418 let mut b = __s_b.launch_builder(&f);
4419 b.arg(embd).arg(token_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb);
4420 unsafe { b.launch(cfg)?; }
4421 Ok(x)
4422 }
4423
4424
4425 pub fn embed_gather_device_t(&self, embd: &CudaSlice<u8>, tokens: &[u32],
4429 n_embd: usize, qtype: i32, row_bytes: usize)
4430 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4431 let t = tokens.len();
4432 let tok_d = self.gpu.stream().clone_htod(tokens)?;
4433 let f = self.func("embed_gather_u32_t");
4434 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4435 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4436 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4437 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4438 let __s_b = self.gpu.stream();
4439 let mut b = __s_b.launch_builder(&f);
4440 b.arg(embd).arg(&tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4441 unsafe { b.launch(cfg)?; }
4442 Ok(x)
4443 }
4444
4445 pub fn embed_gather_device_tv(&self, embd: &CudaSlice<u8>, tok_v: &cudarc::driver::CudaView<u32>,
4450 t: usize, n_embd: usize, qtype: i32, row_bytes: usize)
4451 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4452 let f = self.func("embed_gather_u32_t");
4453 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4454 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4455 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4456 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4457 let __s_b = self.gpu.stream();
4458 let mut b = __s_b.launch_builder(&f);
4459 b.arg(embd).arg(tok_v).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4460 unsafe { b.launch(cfg)?; }
4461 Ok(x)
4462 }
4463
4464 pub fn embed_gather_device_td(&self, embd: &CudaSlice<u8>, tok_d: &CudaSlice<u32>, t: usize,
4465 n_embd: usize, qtype: i32, row_bytes: usize)
4466 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4467 let f = self.func("embed_gather_u32_t");
4468 let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
4469 let cfg = LaunchConfig { grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
4470 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
4471 let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
4472 let __s_b = self.gpu.stream();
4473 let mut b = __s_b.launch_builder(&f);
4474 b.arg(embd).arg(tok_d).arg(&mut x).arg(&ne).arg(&qt).arg(&rb).arg(&ti);
4475 unsafe { b.launch(cfg)?; }
4476 Ok(x)
4477 }
4478
4479 #[inline]
4485 fn keep_if_capturing<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, s: &CudaSlice<T>) {
4487 if self.capture_keep_on.load(std::sync::atomic::Ordering::Relaxed) {
4488 self.capture_keep.lock().unwrap().push(Box::new(s.clone()));
4489 }
4490 }
4491
4492 fn alloc_uninit<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, n: usize)
4493 -> Result<CudaSlice<T>, Box<dyn std::error::Error>> {
4494 let mut s = unsafe { self.gpu.stream().alloc::<T>(n)? };
4495 {
4499 static Z: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4500 if *Z.get_or_init(|| std::env::var("MEMRA_DEBUG_ZERO_ALLOCS").as_deref() == Ok("1")) {
4501 use cudarc::driver::DevicePtrMut;
4503 let n_bytes = s.len() * std::mem::size_of::<T>();
4504 let stream = self.gpu.stream();
4505 let (p_, _g) = s.device_ptr_mut(&stream);
4506 unsafe {
4507 cudarc::driver::sys::cuMemsetD8Async(p_, 0, n_bytes, stream.cu_stream())
4508 .result()?;
4509 }
4510 }
4511 }
4512 self.keep_if_capturing(&s);
4513 Ok(s)
4514 }
4515
4516 pub fn uninit_q8_pair(&self, n: usize)
4521 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4522 Ok((self.alloc_uninit::<i8>(n)?, self.alloc_uninit::<f32>(n / 32)?))
4523 }
4524
4525 pub fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4526 self.alloc_uninit::<f32>(n)
4527 }
4528
4529 pub fn alloc_i8_uninit(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
4531 self.alloc_uninit::<i8>(n)
4532 }
4533
4534 #[allow(clippy::too_many_arguments)]
4538 pub fn rms_norm3(&self, x: &CudaSlice<f32>, w0: &CudaSlice<f32>, w1: &CudaSlice<f32>,
4539 w2: &CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4540 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4541 -> Result<(), Box<dyn std::error::Error>> {
4542 let f = self.func("rms_norm3_f32");
4543 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4544 let (nc, e) = (ncols as i32, eps);
4545 let __s_b = self.gpu.stream();
4546 let mut b = __s_b.launch_builder(&f);
4547 b.arg(x).arg(w0).arg(w1).arg(w2).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e);
4548 unsafe { b.launch(cfg)?; }
4549 Ok(())
4550 }
4551
4552 #[allow(clippy::too_many_arguments)]
4554 pub fn qkvnorm_w_on_prefill(rows: usize, ncols: usize) -> bool {
4557 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4558 *WARP_ON.get_or_init(|| {
4559 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4560 }) && ncols % 4 == 0 && rows >= 64
4561 }
4562
4563 #[allow(clippy::too_many_arguments)]
4566 pub fn rms_norm_qkv_w4b(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4567 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4568 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4569 dvb: &mut CudaSlice<u8>,
4570 ncols: usize, rq: usize, rk: usize, eps: f32, vf16: bool)
4571 -> Result<(), Box<dyn std::error::Error>> {
4572 assert!(ncols % 4 == 0 && rq + 2 * rk >= 64);
4573 let f = self.func("rms_norm_qkv_w4b_f32");
4574 let rows = (rq + 2 * rk) as u32;
4575 let cfg = LaunchConfig {
4576 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4577 };
4578 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4579 let vf = vf16 as i32;
4580 let __s_b = self.gpu.stream();
4581 let mut b = __s_b.launch_builder(&f);
4582 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv).arg(&mut *dvb)
4583 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e).arg(&vf);
4584 unsafe { b.launch(cfg)?; }
4585 Ok(())
4586 }
4587
4588 pub fn rms_norm_qkv(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
4589 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
4590 dq: &mut CudaSlice<f32>, dk: &mut CudaSlice<f32>, dv: &mut CudaSlice<f32>,
4591 ncols: usize, rq: usize, rk: usize, eps: f32)
4592 -> Result<(), Box<dyn std::error::Error>> {
4593 static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4597 let warp_on = *WARP_ON.get_or_init(|| {
4598 std::env::var("MEMRA_QKVNORM_W").map(|v| v != "0").unwrap_or(true)
4599 });
4600 if warp_on && ncols % 4 == 0 && rq + 2 * rk >= 64 {
4603 let f = self.func("rms_norm_qkv_w4_f32");
4604 let rows = (rq + 2 * rk) as u32;
4605 let cfg = LaunchConfig {
4606 grid_dim: (rows.div_ceil(8), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0,
4607 };
4608 let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
4609 let __s_b = self.gpu.stream();
4610 let mut b = __s_b.launch_builder(&f);
4611 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4612 .arg(&nc).arg(&rqi).arg(&rki).arg(&rvi).arg(&e);
4613 unsafe { b.launch(cfg)?; }
4614 return Ok(());
4615 }
4616 let f = self.func("rms_norm_qkv_f32");
4617 let grid = (rq + 2 * rk) as u32;
4618 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4619 let (nc, rqi, rki, e) = (ncols as i32, rq as i32, rk as i32, eps);
4620 let __s_b = self.gpu.stream();
4621 let mut b = __s_b.launch_builder(&f);
4622 b.arg(q).arg(k).arg(v).arg(wq).arg(wk).arg(wv).arg(dq).arg(dk).arg(dv)
4623 .arg(&nc).arg(&rqi).arg(&rki).arg(&e);
4624 unsafe { b.launch(cfg)?; }
4625 Ok(())
4626 }
4627
4628 #[allow(clippy::too_many_arguments)]
4630 pub fn rms_norm2x(&self, a: &CudaSlice<f32>, bb: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4631 wb: &CudaSlice<f32>, da: &mut CudaSlice<f32>, db: &mut CudaSlice<f32>,
4632 ncols: usize, nrows: usize, eps: f32)
4633 -> Result<(), Box<dyn std::error::Error>> {
4634 let f = self.func("rms_norm2x_f32");
4635 let cfg = LaunchConfig { grid_dim: (2 * nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4636 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
4637 let __s_b = self.gpu.stream();
4638 let mut b = __s_b.launch_builder(&f);
4639 b.arg(a).arg(bb).arg(wa).arg(wb).arg(da).arg(db).arg(&nc).arg(&nr).arg(&e);
4640 unsafe { b.launch(cfg)?; }
4641 Ok(())
4642 }
4643
4644 pub fn softcap(&self, y: &mut CudaSlice<f32>, cap: f32, n: usize)
4646 -> Result<(), Box<dyn std::error::Error>> {
4647 let f = self.func("softcap_f32");
4648 let cfg = LaunchConfig::for_num_elems(n as u32);
4649 let ni = n as i32;
4650 let __s_b = self.gpu.stream();
4651 let mut b = __s_b.launch_builder(&f);
4652 b.arg(y).arg(&cap).arg(&ni);
4653 unsafe { b.launch(cfg)?; }
4654 Ok(())
4655 }
4656
4657 pub fn mask_ids_rows(&self, y: &mut CudaSlice<f32>, ids: &CudaSlice<i32>, n_ids: usize,
4660 n_vocab: usize, t: usize)
4661 -> Result<(), Box<dyn std::error::Error>> {
4662 let f = self.func("mask_ids_rows_f32");
4663 let cfg = LaunchConfig::for_num_elems((n_ids * t) as u32);
4664 let (ni, nv, ti) = (n_ids as i32, n_vocab as i32, t as i32);
4665 let __s_b = self.gpu.stream();
4666 let mut b = __s_b.launch_builder(&f);
4667 b.arg(y).arg(ids).arg(&ni).arg(&nv).arg(&ti);
4668 unsafe { b.launch(cfg)?; }
4669 Ok(())
4670 }
4671
4672 #[allow(clippy::too_many_arguments)]
4674 pub fn add_scale_rms_norm(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4675 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4676 ncols: usize, nrows: usize, eps: f32)
4677 -> Result<(), Box<dyn std::error::Error>> {
4678 let f = self.func("add_scale_rms_norm_f32");
4679 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4680 let (nc, e2) = (ncols as i32, eps);
4681 let __s_b = self.gpu.stream();
4682 let mut b = __s_b.launch_builder(&f);
4683 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(dst).arg(&nc).arg(&e2);
4684 unsafe { b.launch(cfg)?; }
4685 Ok(())
4686 }
4687
4688 #[allow(clippy::too_many_arguments)]
4691 pub fn add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4692 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4693 ncols: usize, nrows: usize, eps: f32)
4694 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4695 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4696 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4697 let (nc, e2) = (ncols as i32, eps);
4698 if Self::pdl_on() && Self::pdl_wb_on() {
4699 {
4700 use cudarc::driver::{DevicePtr, DevicePtrMut};
4701 let s = &self.gpu.stream();
4702 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4703 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4704 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4705 let mut ps = [
4706 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4707 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4708 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4709 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4710 &e2 as *const _ as *mut _,
4711 ];
4712 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4713 (rms_block(), 1, 1), &mut ps)?; }
4714 }
4715 return Ok((out_q, out_d));
4716 }
4717 let f = self.func("add_scale_rms_norm_q8_1");
4718 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4719 let __s_b = self.gpu.stream();
4720 let mut b = __s_b.launch_builder(&f);
4721 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e2);
4722 unsafe { b.launch(cfg)?; }
4723 Ok((out_q, out_d))
4724 }
4725
4726 #[allow(clippy::too_many_arguments)]
4728 pub fn add_scale_rms_norm_q8_1_into(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4729 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4730 ncols: usize, nrows: usize, eps: f32,
4731 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4732 -> Result<(), Box<dyn std::error::Error>> {
4733 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4734 let (nc, e2) = (ncols as i32, eps);
4735 if Self::pdl_on() && Self::pdl_wb_on() {
4736 use cudarc::driver::{DevicePtr, DevicePtrMut};
4737 let s = &self.gpu.stream();
4738 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b_in.device_ptr(s);
4739 let (pw, _g2) = w.device_ptr(s); let (pr, _g3) = res.device_ptr_mut(s);
4740 let (pq, _g4) = out_q.device_ptr_mut(s); let (pd, _g5) = out_d.device_ptr_mut(s);
4741 let mut ps = [
4742 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
4743 &c as *const _ as *mut _, &pw as *const _ as *mut _,
4744 &pr as *const _ as *mut _, &pq as *const _ as *mut _,
4745 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4746 &e2 as *const _ as *mut _,
4747 ];
4748 unsafe { self.launch_pdl("add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4749 (rms_block(), 1, 1), &mut ps)?; }
4750 return Ok(());
4751 }
4752 let f = self.func("add_scale_rms_norm_q8_1");
4753 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4754 let __s_b = self.gpu.stream();
4755 let mut b = __s_b.launch_builder(&f);
4756 b.arg(a).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut *out_q).arg(&mut *out_d).arg(&nc).arg(&e2);
4757 unsafe { b.launch(cfg)?; }
4758 Ok(())
4759 }
4760
4761 #[allow(clippy::too_many_arguments)]
4764 pub fn rms_pre_add_scale_rms_norm_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
4765 b_in: &CudaSlice<f32>, c: f32,
4766 w: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
4767 ncols: usize, nrows: usize, eps: f32)
4768 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4769 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4770 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4771 let (nc, e2) = (ncols as i32, eps);
4772 if Self::pdl_on() {
4773 {
4774 use cudarc::driver::{DevicePtr, DevicePtrMut};
4775 let s = &self.gpu.stream();
4776 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
4777 let (pb, _g2) = b_in.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
4778 let (pr, _g4) = res.device_ptr_mut(s);
4779 let (pq, _g5) = out_q.device_ptr_mut(s); let (pd, _g6) = out_d.device_ptr_mut(s);
4780 let mut ps = [
4781 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
4782 &pb as *const _ as *mut _, &c as *const _ as *mut _,
4783 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
4784 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4785 &nc as *const _ as *mut _, &e2 as *const _ as *mut _,
4786 ];
4787 unsafe { self.launch_pdl("rms_pre_add_scale_rms_norm_q8_1", (nrows as u32, 1, 1),
4788 (rms_block(), 1, 1), &mut ps)?; }
4789 }
4790 return Ok((out_q, out_d));
4791 }
4792 let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
4793 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4794 let __s_b = self.gpu.stream();
4795 let mut b = __s_b.launch_builder(&f);
4796 b.arg(a).arg(wa).arg(b_in).arg(&c).arg(w).arg(res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e2);
4797 unsafe { b.launch(cfg)?; }
4798 Ok((out_q, out_d))
4799 }
4800
4801 pub fn gelu_tanh_mul_q8_1(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4804 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize)
4805 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4806 debug_assert!(ncols % 128 == 0);
4807 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
4808 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4809 let nc = ncols as i32;
4810 if Self::pdl_on() {
4811 {
4812 use cudarc::driver::{DevicePtr, DevicePtrMut};
4813 let s = &self.gpu.stream();
4814 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4815 let (pact, _g2) = act.device_ptr_mut(s);
4816 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4817 let mut ps = [
4818 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4819 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4820 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4821 ];
4822 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4823 (rms_block(), 1, 1), &mut ps)?; }
4824 }
4825 return Ok((out_q, out_d));
4826 }
4827 let f = self.func("gelu_tanh_mul_q8_1");
4828 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4829 let __s_b = self.gpu.stream();
4830 let mut b = __s_b.launch_builder(&f);
4831 b.arg(gate).arg(up).arg(act).arg(&mut out_q).arg(&mut out_d).arg(&nc);
4832 unsafe { b.launch(cfg)?; }
4833 Ok((out_q, out_d))
4834 }
4835
4836 #[allow(clippy::too_many_arguments)]
4838 pub fn gelu_tanh_mul_q8_1_into(&self, gate: &CudaSlice<f32>, up: &cudarc::driver::CudaView<f32>,
4839 act: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
4840 out_q: &mut CudaSlice<i8>, out_d: &mut CudaSlice<f32>)
4841 -> Result<(), Box<dyn std::error::Error>> {
4842 debug_assert!(ncols % 128 == 0);
4843 debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
4844 let nc = ncols as i32;
4845 if Self::pdl_on() {
4846 use cudarc::driver::{DevicePtr, DevicePtrMut};
4847 let s = &self.gpu.stream();
4848 let (pg, _g0) = gate.device_ptr(s); let (pu, _g1) = up.device_ptr(s);
4849 let (pact, _g2) = act.device_ptr_mut(s);
4850 let (pq, _g3) = out_q.device_ptr_mut(s); let (pd, _g4) = out_d.device_ptr_mut(s);
4851 let mut ps = [
4852 &pg as *const _ as *mut std::ffi::c_void, &pu as *const _ as *mut _,
4853 &pact as *const _ as *mut _, &pq as *const _ as *mut _,
4854 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4855 ];
4856 unsafe { self.launch_pdl("gelu_tanh_mul_q8_1", (nrows as u32, 1, 1),
4857 (rms_block(), 1, 1), &mut ps)?; }
4858 return Ok(());
4859 }
4860 let f = self.func("gelu_tanh_mul_q8_1");
4861 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4862 let __s_b = self.gpu.stream();
4863 let mut b = __s_b.launch_builder(&f);
4864 b.arg(gate).arg(up).arg(&mut *act).arg(&mut *out_q).arg(&mut *out_d).arg(&nc);
4865 unsafe { b.launch(cfg)?; }
4866 Ok(())
4867 }
4868
4869 #[allow(clippy::too_many_arguments)]
4871 pub fn add_rms_norm3_q8z(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4872 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4873 res: &mut CudaSlice<f32>, out1: &mut CudaSlice<f32>,
4874 ncols: usize, nrows: usize, eps: f32)
4875 -> Result<((CudaSlice<i8>, CudaSlice<f32>), (CudaSlice<i8>, CudaSlice<f32>)), Box<dyn std::error::Error>> {
4876 let mut q0 = self.alloc_uninit::<i8>(nrows * ncols)?;
4877 let mut d0 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4878 let mut q2 = self.alloc_uninit::<i8>(nrows * ncols)?;
4879 let mut d2 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
4880 let f = self.func("add_rms_norm3_q8z_f32");
4881 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4882 let (nc, e2) = (ncols as i32, eps);
4883 let __s_b = self.gpu.stream();
4884 let mut b = __s_b.launch_builder(&f);
4885 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res)
4886 .arg(&mut q0).arg(&mut d0).arg(out1).arg(&mut q2).arg(&mut d2).arg(&nc).arg(&e2);
4887 unsafe { b.launch(cfg)?; }
4888 Ok(((q0, d0), (q2, d2)))
4889 }
4890
4891 #[allow(clippy::too_many_arguments)]
4893 pub fn add_rms_norm3(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>,
4894 w0: &CudaSlice<f32>, w1: &CudaSlice<f32>, w2: &CudaSlice<f32>,
4895 res: &mut CudaSlice<f32>, d0: &mut CudaSlice<f32>, d1: &mut CudaSlice<f32>,
4896 d2: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
4897 -> Result<(), Box<dyn std::error::Error>> {
4898 let f = self.func("add_rms_norm3_f32");
4899 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4900 let (nc, e2) = (ncols as i32, eps);
4901 let __s_b = self.gpu.stream();
4902 let mut b = __s_b.launch_builder(&f);
4903 b.arg(a).arg(b_in).arg(w0).arg(w1).arg(w2).arg(res).arg(d0).arg(d1).arg(d2).arg(&nc).arg(&e2);
4904 unsafe { b.launch(cfg)?; }
4905 Ok(())
4906 }
4907
4908 pub fn add_scale(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, c: f32,
4910 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
4911 let f = self.func("add_scale_f32");
4912 let cfg = LaunchConfig::for_num_elems(n as u32);
4913 let ni = n as i32;
4914 let __s_b = self.gpu.stream();
4915 let mut b = __s_b.launch_builder(&f);
4916 b.arg(a).arg(b_in).arg(&c).arg(dst).arg(&ni);
4917 unsafe { b.launch(cfg)?; }
4918 Ok(())
4919 }
4920
4921 pub fn rms_norm(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4922 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4923 let (nc, e) = (ncols as i32, eps);
4924 if Self::pdl_on() && Self::pdl_wb_on() {
4925 use cudarc::driver::{DevicePtr, DevicePtrMut};
4926 let s = &self.gpu.stream();
4927 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4928 let (pd, _g2) = dst.device_ptr_mut(s);
4929 let mut ps = [
4930 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4931 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
4932 &e as *const _ as *mut _,
4933 ];
4934 unsafe { self.launch_pdl("rms_norm_f32", (nrows as u32, 1, 1),
4935 (rms_block(), 1, 1), &mut ps)?; }
4936 return Ok(());
4937 }
4938 let f = self.func("rms_norm_f32");
4939 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
4940 let __s_b = self.gpu.stream();
4941 let mut b = __s_b.launch_builder(&f);
4942 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4943 unsafe { b.launch(cfg)?; }
4944 Ok(())
4945 }
4946
4947 pub fn rms_norm_decode(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
4955 ncols: usize, nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
4956 let f = self.func("rms_norm_f32");
4957 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4958 let (nc, e) = (ncols as i32, eps);
4959 let __s_b = self.gpu.stream();
4960 let mut b = __s_b.launch_builder(&f);
4961 b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
4962 unsafe { b.launch(cfg)?; }
4963 Ok(())
4964 }
4965
4966 pub fn rms_norm_q8_1(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize, nrows: usize,
4970 eps: f32) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4971 let nblk = ncols / 32;
4972 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
4973 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
4974 let (nc, e) = (ncols as i32, eps);
4975 if Self::pdl_on() {
4976 {
4977 use cudarc::driver::{DevicePtr, DevicePtrMut};
4978 let s = &self.gpu.stream();
4979 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
4980 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
4981 let mut ps = [
4982 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
4983 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
4984 &nc as *const _ as *mut _, &e as *const _ as *mut _,
4985 ];
4986 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
4987 &mut ps)?; }
4988 }
4989 return Ok((q, d));
4990 }
4991 let f = self.func("rms_norm_q8_1");
4992 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
4995 let __s_b = self.gpu.stream();
4996 let mut b = __s_b.launch_builder(&f);
4997 b.arg(x).arg(w).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
4998 unsafe { b.launch(cfg)?; }
4999 Ok((q, d))
5000 }
5001
5002 pub fn rms_norm_q8_1_into(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, ncols: usize,
5005 nrows: usize, eps: f32,
5006 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
5007 -> Result<(), Box<dyn std::error::Error>> {
5008 let nblk = ncols / 32;
5009 debug_assert!(q.len() >= nrows * ncols && d.len() >= nrows * nblk);
5010 let (nc, e) = (ncols as i32, eps);
5011 if Self::pdl_on() {
5012 use cudarc::driver::{DevicePtr, DevicePtrMut};
5013 let s = &self.gpu.stream();
5014 let (px, _g0) = x.device_ptr(s); let (pw, _g1) = w.device_ptr(s);
5015 let (pq, _g2) = q.device_ptr_mut(s); let (pd, _g3) = d.device_ptr_mut(s);
5016 let mut ps = [
5017 &px as *const _ as *mut std::ffi::c_void, &pw as *const _ as *mut _,
5018 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
5019 &nc as *const _ as *mut _, &e as *const _ as *mut _,
5020 ];
5021 unsafe { self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1),
5022 &mut ps)?; }
5023 return Ok(());
5024 }
5025 let f = self.func("rms_norm_q8_1");
5026 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
5027 let __s_b = self.gpu.stream();
5028 let mut b = __s_b.launch_builder(&f);
5029 b.arg(x).arg(w).arg(&mut *q).arg(&mut *d).arg(&nc).arg(&e);
5030 unsafe { b.launch(cfg)?; }
5031 Ok(())
5032 }
5033
5034 pub fn quantize_q8_1_into(&self, x: &CudaSlice<f32>, m: usize, in_f: usize,
5036 q: &mut CudaSlice<i8>, d: &mut CudaSlice<f32>)
5037 -> Result<(), Box<dyn std::error::Error>> {
5038 let nblk = in_f / 32;
5039 debug_assert!(q.len() >= m * in_f && d.len() >= m * nblk);
5040 let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
5041 let (inf, mi) = (in_f as i32, m as i32);
5042 if Self::pdl_on() && Self::pdl_wb_on() {
5043 use cudarc::driver::{DevicePtr, DevicePtrMut};
5044 let s = &self.gpu.stream();
5045 let (px, _g0) = x.device_ptr(s);
5046 let (pq, _g1) = q.device_ptr_mut(s); let (pd, _g2) = d.device_ptr_mut(s);
5047 let mut ps = [
5048 &px as *const _ as *mut std::ffi::c_void, &pq as *const _ as *mut _,
5049 &pd as *const _ as *mut _, &inf as *const _ as *mut _,
5050 &mi as *const _ as *mut _,
5051 ];
5052 unsafe { self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?; }
5053 return Ok(());
5054 }
5055 let f = self.func("quantize_q8_1");
5056 let __s_b = self.gpu.stream();
5057 let mut b = __s_b.launch_builder(&f);
5058 b.arg(x).arg(&mut *q).arg(&mut *d).arg(&inf).arg(&mi);
5059 unsafe { b.launch(cfg)?; }
5060 Ok(())
5061 }
5062
5063 pub fn add_rms_norm_q8_1(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
5067 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
5068 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5069 let nblk = ncols / 32;
5070 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
5071 let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
5072 let f = self.func("add_rms_norm_q8_1");
5073 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
5075 let (nc, e) = (ncols as i32, eps);
5076 let __s_bld = self.gpu.stream();
5077 let mut bld = __s_bld.launch_builder(&f);
5078 bld.arg(a).arg(b_in).arg(w).arg(res).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
5079 unsafe { bld.launch(cfg)?; }
5080 Ok((q, d))
5081 }
5082
5083 pub fn add_rms_norm(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
5087 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
5088 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5089 let (nc, e) = (ncols as i32, eps);
5090 if Self::pdl_on() && Self::pdl_wb_on() {
5091 use cudarc::driver::{DevicePtr, DevicePtrMut};
5092 let s = &self.gpu.stream();
5093 let (pa, _g0) = a.device_ptr(s); let (pb, _g1) = b.device_ptr(s);
5094 let (pw, _g2) = w.device_ptr(s);
5095 let (pr, _g3) = res.device_ptr_mut(s); let (pd, _g4) = dst.device_ptr_mut(s);
5096 let mut ps = [
5097 &pa as *const _ as *mut std::ffi::c_void, &pb as *const _ as *mut _,
5098 &pw as *const _ as *mut _, &pr as *const _ as *mut _,
5099 &pd as *const _ as *mut _, &nc as *const _ as *mut _,
5100 &e as *const _ as *mut _,
5101 ];
5102 unsafe { self.launch_pdl("add_rms_norm_f32", (nrows as u32, 1, 1),
5103 (rms_block(), 1, 1), &mut ps)?; }
5104 return Ok(());
5105 }
5106 let f = self.func("add_rms_norm_f32");
5107 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5108 let __s_b2 = self.gpu.stream();
5109 let mut b2 = __s_b2.launch_builder(&f);
5110 b2.arg(a).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
5111 unsafe { b2.launch(cfg)?; }
5112 Ok(())
5113 }
5114
5115 #[allow(clippy::too_many_arguments)]
5118 pub fn rms_pre_add_rms_norm(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
5119 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
5120 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5121 ncols: usize, nrows: usize, eps: f32)
5122 -> Result<(), Box<dyn std::error::Error>> {
5123 let f = self.func("rms_pre_add_rms_norm_f32");
5124 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5125 let (nc, e) = (ncols as i32, eps);
5126 let __s_b2 = self.gpu.stream();
5127 let mut b2 = __s_b2.launch_builder(&f);
5128 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst).arg(&nc).arg(&e);
5129 unsafe { b2.launch(cfg)?; }
5130 Ok(())
5131 }
5132
5133 #[allow(clippy::too_many_arguments)]
5135 pub fn rms_pre_add_rms_norm_q8z(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>,
5136 b: &CudaSlice<f32>, w: &CudaSlice<f32>,
5137 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5138 ncols: usize, nrows: usize, eps: f32)
5139 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5140 debug_assert!(ncols % 128 == 0);
5141 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5142 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5143 let (nc, e) = (ncols as i32, eps);
5144 if Self::pdl_on() {
5145 {
5146 use cudarc::driver::{DevicePtr, DevicePtrMut};
5147 let s = &self.gpu.stream();
5148 let (pa, _g0) = a.device_ptr(s); let (pwa, _g1) = wa.device_ptr(s);
5149 let (pb, _g2) = b.device_ptr(s); let (pw, _g3) = w.device_ptr(s);
5150 let (pr, _g4) = res.device_ptr_mut(s); let (pdst, _g5) = dst.device_ptr_mut(s);
5151 let (pq, _g6) = out_q.device_ptr_mut(s); let (pd, _g7) = out_d.device_ptr_mut(s);
5152 let mut ps = [
5153 &pa as *const _ as *mut std::ffi::c_void, &pwa as *const _ as *mut _,
5154 &pb as *const _ as *mut _, &pw as *const _ as *mut _,
5155 &pr as *const _ as *mut _, &pdst as *const _ as *mut _,
5156 &pq as *const _ as *mut _, &pd as *const _ as *mut _,
5157 &nc as *const _ as *mut _, &e as *const _ as *mut _,
5158 ];
5159 unsafe { self.launch_pdl("rms_pre_add_rms_norm_q8z_f32", (nrows as u32, 1, 1),
5160 (rms_block(), 1, 1), &mut ps)?; }
5161 }
5162 return Ok((out_q, out_d));
5163 }
5164 let f = self.func("rms_pre_add_rms_norm_q8z_f32");
5165 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5166 let __s_b2 = self.gpu.stream();
5167 let mut b2 = __s_b2.launch_builder(&f);
5168 b2.arg(a).arg(wa).arg(b).arg(w).arg(&mut *res).arg(&mut *dst)
5169 .arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&e);
5170 unsafe { b2.launch(cfg)?; }
5171 Ok((out_q, out_d))
5172 }
5173
5174 pub fn build_q4_out_concat3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
5178 w2: &crate::model::GpuTensor)
5179 -> Result<Option<crate::model::GpuTensor>, Box<dyn std::error::Error>> {
5180 use crate::model::GpuTensor;
5181 let part = |w: &GpuTensor| -> Option<(usize, usize)> {
5182 match w {
5183 GpuTensor::Quant { qtype, row_bytes, rp, .. }
5184 if *qtype == QT_Q4_0 && !*rp => Some((*row_bytes, w.out_features())),
5185 _ => None,
5186 }
5187 };
5188 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (part(w0), part(w1), part(w2))
5189 else { return Ok(None) };
5190 if rb0 != rb1 || rb0 != rb2
5191 || w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
5192 return Ok(None);
5193 }
5194 fn bytes_of(w: &crate::model::GpuTensor) -> &CudaSlice<u8> {
5195 match w { crate::model::GpuTensor::Quant { bytes, .. } => bytes, _ => unreachable!() }
5196 }
5197 let (b0, b1, b2) = (bytes_of(w0), bytes_of(w1), bytes_of(w2));
5198 let total = rb0 * (o0 + o1 + o2);
5199 let mut cat = self.alloc_u8(total)?;
5200 self.copy_u8_into(&mut cat, 0, b0, rb0 * o0)?;
5201 self.copy_u8_into(&mut cat, rb0 * o0, b1, rb1 * o1)?;
5202 self.copy_u8_into(&mut cat, rb0 * (o0 + o1), b2, rb2 * o2)?;
5203 Ok(Some(GpuTensor::Quant {
5204 bytes: cat, qtype: QT_Q4_0, row_bytes: rb0,
5205 ne: vec![w0.in_features() as u64, (o0 + o1 + o2) as u64], scale: 1.0, rp: false,
5206 #[cfg(memra_cutlass)]
5207 cutlass: None,
5208 fp8: None, blk: None, rp4: None, f16: None,
5209 }))
5210 }
5211
5212 #[allow(clippy::too_many_arguments)]
5214 pub fn rms_norm_qkv_rope_cat(&self, qkv: &CudaSlice<f32>,
5215 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5216 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5217 head_dim: usize, rq: usize, rk: usize,
5218 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5219 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32)
5220 -> Result<(), Box<dyn std::error::Error>> {
5221 let rows = rq + rk + rk;
5222 let theta_scale = base.powf(-2.0 / head_dim as f32);
5223 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5224 if Self::pdl_on() {
5225 use cudarc::driver::{DevicePtr, DevicePtrMut};
5226 let s = &self.gpu.stream();
5227 let (pqkv, _g0) = qkv.device_ptr(s);
5228 let (pwq, _g1) = wq.device_ptr(s); let (pwk, _g2) = wk.device_ptr(s);
5229 let (pwv, _g3) = wv.device_ptr(s);
5230 let (pq, _g4) = q.device_ptr_mut(s); let (pk, _g5) = k.device_ptr_mut(s);
5231 let (pv, _g6) = v.device_ptr_mut(s);
5232 let (ppos, _g7) = pos.device_ptr(s);
5233 let (pff, _g8) = match ff {
5234 Some(t) => { let (p, g) = t.device_ptr(s); (p, Some(g)) }
5235 None => (0, None),
5236 };
5237 let mut ps = [
5238 &pqkv as *const _ as *mut std::ffi::c_void,
5239 &pwq as *const _ as *mut _, &pwk as *const _ as *mut _,
5240 &pwv as *const _ as *mut _,
5241 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5242 &pv as *const _ as *mut _,
5243 &nc as *const _ as *mut _, &rqi as *const _ as *mut _,
5244 &rki as *const _ as *mut _, &ppos as *const _ as *mut _,
5245 &nhq as *const _ as *mut _, &nhk as *const _ as *mut _,
5246 &theta_scale as *const _ as *mut _, &freq_scale as *const _ as *mut _,
5247 &pff as *const _ as *mut _, &eps as *const _ as *mut _,
5248 ];
5249 unsafe { self.launch_pdl("rms_norm_qkv_rope_cat_f32", (rows as u32, 1, 1),
5250 (rms_block(), 1, 1), &mut ps)?; }
5251 return Ok(());
5252 }
5253 let f = self.func("rms_norm_qkv_rope_cat_f32");
5254 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5255 let __s_b = self.gpu.stream();
5256 let mut b = __s_b.launch_builder(&f);
5257 match ff {
5258 Some(t) => { b.arg(qkv).arg(wq).arg(wk).arg(wv)
5259 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5260 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5261 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5262 unsafe { b.launch(cfg)?; } }
5263 None => { let null: u64 = 0;
5264 b.arg(qkv).arg(wq).arg(wk).arg(wv)
5265 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5266 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5267 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5268 unsafe { b.launch(cfg)?; } }
5269 }
5270 Ok(())
5271 }
5272
5273 #[allow(clippy::too_many_arguments)]
5275 pub fn rms_norm_qkv_rope(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>, v0: &CudaSlice<f32>,
5276 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5277 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5278 head_dim: usize, rq: usize, rk: usize,
5279 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5280 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32)
5281 -> Result<(), Box<dyn std::error::Error>> {
5282 let f = self.func("rms_norm_qkv_rope_f32");
5283 let rows = rq + rk + rk; let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5285 let theta_scale = base.powf(-2.0 / head_dim as f32);
5286 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5287 let __s_b = self.gpu.stream();
5288 let mut b = __s_b.launch_builder(&f);
5289 match ff {
5290 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5291 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5292 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5293 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps);
5294 unsafe { b.launch(cfg)?; } }
5295 None => { let null: u64 = 0;
5296 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5297 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5298 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5299 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps);
5300 unsafe { b.launch(cfg)?; } }
5301 }
5302 Ok(())
5303 }
5304
5305 #[allow(clippy::too_many_arguments)]
5309 pub fn rms_norm_qkv_rope_append_dc(&self, q0: &CudaSlice<f32>, k0: &CudaSlice<f32>,
5310 v0: &CudaSlice<f32>,
5311 wq: &CudaSlice<f32>, wk: &CudaSlice<f32>, wv: &CudaSlice<f32>,
5312 q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>, v: &mut CudaSlice<f32>,
5313 head_dim: usize, rq: usize, rk: usize,
5314 pos: &CudaSlice<i32>, nh_q: usize, nh_k: usize,
5315 base: f32, freq_scale: f32, ff: Option<&CudaSlice<f32>>, eps: f32,
5316 kc: &mut CudaSlice<u8>, vc: &mut CudaSlice<u8>,
5317 t_dev: &CudaSlice<i32>, k_tok_bytes: usize, v_tok_bytes: usize,
5318 g: bool)
5319 -> Result<(), Box<dyn std::error::Error>> {
5320 let rows = rq + rk + rk;
5321 let theta_scale = base.powf(-2.0 / head_dim as f32);
5322 let (nc, rqi, rki, nhq, nhk) = (head_dim as i32, rq as i32, rk as i32, nh_q as i32, nh_k as i32);
5323 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
5324 if Self::pdl_on() && Self::pdl_wb_on() {
5325 use cudarc::driver::{DevicePtr, DevicePtrMut};
5326 let s = &self.gpu.stream();
5327 let (p0, _a0) = q0.device_ptr(s); let (p1, _a1) = k0.device_ptr(s);
5328 let (p2, _a2) = v0.device_ptr(s);
5329 let (pwq, _a3) = wq.device_ptr(s); let (pwk, _a4) = wk.device_ptr(s);
5330 let (pwv, _a5) = wv.device_ptr(s);
5331 let (pq, _a6) = q.device_ptr_mut(s); let (pk, _a7) = k.device_ptr_mut(s);
5332 let (pv, _a8) = v.device_ptr_mut(s);
5333 let (pp, _a9) = pos.device_ptr(s);
5334 let pff: u64 = match ff { Some(t) => { let (p, _gg) = t.device_ptr(s); p as u64 }
5335 None => 0 };
5336 let (pkc, _a10) = kc.device_ptr_mut(s); let (pvc, _a11) = vc.device_ptr_mut(s);
5337 let (pt, _a12) = t_dev.device_ptr(s);
5338 let mut ps = [
5339 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
5340 &p2 as *const _ as *mut _, &pwq as *const _ as *mut _,
5341 &pwk as *const _ as *mut _, &pwv as *const _ as *mut _,
5342 &pq as *const _ as *mut _, &pk as *const _ as *mut _,
5343 &pv as *const _ as *mut _, &nc as *const _ as *mut _,
5344 &rqi as *const _ as *mut _, &rki as *const _ as *mut _,
5345 &pp as *const _ as *mut _, &nhq as *const _ as *mut _,
5346 &nhk as *const _ as *mut _, &theta_scale as *const _ as *mut _,
5347 &freq_scale as *const _ as *mut _, &pff as *const _ as *mut _,
5348 &eps as *const _ as *mut _, &pkc as *const _ as *mut _,
5349 &pvc as *const _ as *mut _, &pt as *const _ as *mut _,
5350 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
5351 ];
5352 unsafe { self.launch_pdl_flash(g, "rms_norm_qkv_rope_append_dc_f32",
5353 (rows as u32, 1, 1), (rms_block(), 1, 1), 0, &mut ps)?; }
5354 return Ok(());
5355 }
5356 let f = if g { self.func_g("rms_norm_qkv_rope_append_dc_f32") }
5357 else { self.func("rms_norm_qkv_rope_append_dc_f32") };
5358 let cfg = LaunchConfig { grid_dim: (rows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5359 let __s_b = self.gpu.stream();
5360 let mut b = __s_b.launch_builder(&f);
5361 match ff {
5362 Some(t) => { b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5363 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5364 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5365 .arg(&theta_scale).arg(&freq_scale).arg(t).arg(&eps)
5366 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5367 unsafe { b.launch(cfg)?; } }
5368 None => { let null: u64 = 0;
5369 b.arg(q0).arg(k0).arg(v0).arg(wq).arg(wk).arg(wv)
5370 .arg(&mut *q).arg(&mut *k).arg(&mut *v)
5371 .arg(&nc).arg(&rqi).arg(&rki).arg(pos).arg(&nhq).arg(&nhk)
5372 .arg(&theta_scale).arg(&freq_scale).arg(&null).arg(&eps)
5373 .arg(&mut *kc).arg(&mut *vc).arg(t_dev).arg(&ktb).arg(&vtb);
5374 unsafe { b.launch(cfg)?; } }
5375 }
5376 Ok(())
5377 }
5378
5379 pub fn add_q8_1(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, res: &mut CudaSlice<f32>,
5381 ncols: usize, nrows: usize)
5382 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5383 debug_assert!(ncols % 128 == 0);
5384 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5385 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5386 let f = self.func("add_q8_1_f32");
5387 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
5388 let nc = ncols as i32;
5389 let __s_b2 = self.gpu.stream();
5390 let mut b2 = __s_b2.launch_builder(&f);
5391 b2.arg(a).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc);
5392 unsafe { b2.launch(cfg)?; }
5393 Ok((out_q, out_d))
5394 }
5395
5396 pub fn rms_pre_add_q8_1(&self, a: &CudaSlice<f32>, wa: &CudaSlice<f32>, b: &CudaSlice<f32>,
5400 res: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
5401 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5402 debug_assert!(ncols % 128 == 0);
5403 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
5404 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
5405 let f = self.func("rms_pre_add_q8_1_f32");
5406 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1),
5407 shared_mem_bytes: 0 };
5408 let (nc, ep) = (ncols as i32, eps);
5409 let __s_b2 = self.gpu.stream();
5410 let mut b2 = __s_b2.launch_builder(&f);
5411 b2.arg(a).arg(wa).arg(b).arg(&mut *res).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
5412 unsafe { b2.launch(cfg)?; }
5413 Ok((out_q, out_d))
5414 }
5415
5416 pub fn l2_v2_on(ncols: usize) -> bool {
5420 ncols == 128 && std::env::var("MEMRA_L2_V2").as_deref() != Ok("0")
5421 }
5422
5423 pub fn l2_norm_pp(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
5424 dst16: Option<&mut CudaSlice<u8>>, ncols: usize, nrows: usize,
5425 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5426 if Self::l2_v2_on(ncols) {
5427 let f = self.func("l2_norm_pp_v2_f32");
5428 let rows_per_block = 8u32; let cfg = LaunchConfig { grid_dim: ((nrows as u32).div_ceil(rows_per_block), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
5430 let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
5431 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
5433 let __s_b = self.gpu.stream();
5434 let mut b = __s_b.launch_builder(&f);
5435 b.arg(x).arg(dst).arg(&d16).arg(&nc).arg(&nr).arg(&e);
5436 unsafe { b.launch(cfg)?; }
5437 return Ok(());
5438 }
5439 self.l2_norm(x, dst, ncols, nrows, eps)
5440 }
5441
5442 pub fn l2_norm(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize,
5443 eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5444 let f = self.func("l2_norm_f32");
5445 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
5446 let (nc, e) = (ncols as i32, eps);
5447 let __s_b = self.gpu.stream();
5448 let mut b = __s_b.launch_builder(&f);
5449 b.arg(x).arg(dst).arg(&nc).arg(&e);
5450 unsafe { b.launch(cfg)?; }
5451 Ok(())
5452 }
5453
5454 pub fn l2_norm_decode(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, ncols: usize,
5460 nrows: usize, eps: f32) -> Result<(), Box<dyn std::error::Error>> {
5461 let f = self.func("l2_norm_f32");
5462 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
5463 let (nc, e) = (ncols as i32, eps);
5464 let __s_b = self.gpu.stream();
5465 let mut b = __s_b.launch_builder(&f);
5466 b.arg(x).arg(dst).arg(&nc).arg(&e);
5467 unsafe { b.launch(cfg)?; }
5468 Ok(())
5469 }
5470
5471 pub fn rope_neox(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5473 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32, freq_scale: f32)
5474 -> Result<(), Box<dyn std::error::Error>> {
5475 let f = self.func("rope_neox_f32");
5476 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5477 let grid = (n_heads * n_tokens) as u32;
5478 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5479 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5480 let __s_b = self.gpu.stream();
5481 let mut b = __s_b.launch_builder(&f);
5482 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale);
5483 unsafe { b.launch(cfg)?; }
5484 Ok(())
5485 }
5486
5487 pub fn rope_neox_ff(&self, x: &mut CudaSlice<f32>, pos: &CudaSlice<i32>, head_dim: usize,
5489 n_dims: usize, n_heads: usize, n_tokens: usize, freq_base: f32,
5490 freq_scale: f32, ff: &CudaSlice<f32>)
5491 -> Result<(), Box<dyn std::error::Error>> {
5492 let f = self.func("rope_neox_ff_f32");
5493 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5494 let grid = (n_heads * n_tokens) as u32;
5495 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5496 let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
5497 let __s_b = self.gpu.stream();
5498 let mut b = __s_b.launch_builder(&f);
5499 b.arg(x).arg(pos).arg(&hd).arg(&nd).arg(&nh).arg(&theta_scale).arg(&freq_scale).arg(ff);
5500 unsafe { b.launch(cfg)?; }
5501 Ok(())
5502 }
5503
5504 #[allow(clippy::too_many_arguments)]
5506 pub fn rope_neox2(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
5507 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
5508 nh_q: usize, nh_k: usize, n_tokens: usize, freq_base: f32,
5509 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
5510 -> Result<(), Box<dyn std::error::Error>> {
5511 let f = self.func("rope_neox2_f32");
5512 let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
5513 let grid = ((nh_q + nh_k) * n_tokens) as u32;
5514 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
5515 let (hd, nd, nq, nk, nt) = (head_dim as i32, n_dims as i32, nh_q as i32, nh_k as i32, n_tokens as i32);
5516 let __s_b = self.gpu.stream();
5517 let mut b = __s_b.launch_builder(&f);
5518 b.arg(q).arg(k).arg(pos).arg(&hd).arg(&nd).arg(&nq).arg(&nk).arg(&nt)
5519 .arg(&theta_scale).arg(&freq_scale);
5520 match ff {
5521 Some(ffv) => { b.arg(ffv); unsafe { b.launch(cfg)?; } }
5522 None => {
5523 let null: u64 = 0;
5524 b.arg(&null);
5525 unsafe { b.launch(cfg)?; }
5526 }
5527 }
5528 Ok(())
5529 }
5530
5531 pub fn gelu_tanh_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5533 -> Result<(), Box<dyn std::error::Error>> {
5534 let f = self.func("gelu_tanh_mul_f32");
5535 let cfg = LaunchConfig::for_num_elems(n as u32);
5536 let ni = n as i32;
5537 let __s_b = self.gpu.stream();
5538 let mut b = __s_b.launch_builder(&f);
5539 b.arg(gate).arg(up).arg(dst).arg(&ni);
5540 unsafe { b.launch(cfg)?; }
5541 Ok(())
5542 }
5543
5544 pub fn silu_mul(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5545 -> Result<(), Box<dyn std::error::Error>> {
5546 let f = self.func("silu_mul_f32");
5547 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5549 let ni = n as i32;
5550 let __s_b = self.gpu.stream();
5551 let mut b = __s_b.launch_builder(&f);
5552 b.arg(gate).arg(up).arg(dst).arg(&ni);
5553 unsafe { b.launch(cfg)?; }
5554 Ok(())
5555 }
5556
5557 pub fn silu_mul_f16out(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
5560 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
5561 -> Result<(), Box<dyn std::error::Error>> {
5562 let f = self.func("silu_mul_f16out_f32");
5563 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5564 let ni = n as i32;
5565 let __s_b = self.gpu.stream();
5566 let mut b = __s_b.launch_builder(&f);
5567 b.arg(gate).arg(up).arg(dst).arg(dst16).arg(&ni);
5568 unsafe { b.launch(cfg)?; }
5569 Ok(())
5570 }
5571
5572 pub fn silu_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5579 dst: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
5580 let f = self.func("silu_mul_scaled_f32");
5581 let cfg = LaunchConfig::for_num_elems(n as u32);
5582 let ni = n as i32;
5583 let (gsf, usf) = (gs, us);
5584 let __s_b = self.gpu.stream();
5585 let mut b = __s_b.launch_builder(&f);
5586 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(dst).arg(&ni);
5587 unsafe { b.launch(cfg)?; }
5588 Ok(())
5589 }
5590
5591 #[allow(clippy::too_many_arguments)]
5595 pub fn swigluoai_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5596 alpha: f32, limit: f32, dst: &mut CudaSlice<f32>, n: usize)
5597 -> Result<(), Box<dyn std::error::Error>> {
5598 let f = self.func("swigluoai_mul_scaled_f32");
5599 let cfg = LaunchConfig::for_num_elems(n as u32);
5600 let ni = n as i32;
5601 let __s_b = self.gpu.stream();
5602 let mut b = __s_b.launch_builder(&f);
5603 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&alpha).arg(&limit).arg(dst).arg(&ni);
5604 unsafe { b.launch(cfg)?; }
5605 Ok(())
5606 }
5607
5608 pub fn silu_mul_scaled_q8_1(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>, gs: f32, us: f32,
5616 n: usize)
5617 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5618 let f = self.func("silu_mul_scaled_q8_1");
5619 let nblk = n / 32;
5620 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);
5624 let (gsf, usf, ni) = (gs, us, n as i32);
5625 let __s_b = self.gpu.stream();
5626 let mut b = __s_b.launch_builder(&f);
5627 b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(&mut aq).arg(&mut ad).arg(&ni);
5628 unsafe { b.launch(cfg)?; }
5629 Ok((aq, ad))
5630 }
5631
5632 pub fn add(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5633 -> Result<(), Box<dyn std::error::Error>> {
5634 let f = self.func("add_f32");
5635 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
5637 let ni = n as i32;
5638 let __s_bld = self.gpu.stream();
5639 let mut bld = __s_bld.launch_builder(&f);
5640 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5641 unsafe { bld.launch(cfg)?; }
5642 Ok(())
5643 }
5644
5645 pub fn mul(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, dst: &mut CudaSlice<f32>, n: usize)
5646 -> Result<(), Box<dyn std::error::Error>> {
5647 let f = self.func("mul_f32");
5648 let cfg = LaunchConfig::for_num_elems(n as u32);
5649 let ni = n as i32;
5650 let __s_bld = self.gpu.stream();
5651 let mut bld = __s_bld.launch_builder(&f);
5652 bld.arg(a).arg(b_in).arg(dst).arg(&ni);
5653 unsafe { bld.launch(cfg)?; }
5654 Ok(())
5655 }
5656
5657 pub fn matmul(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
5660 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5661 use crate::model::GpuTensor;
5662 let in_f = w.in_features();
5663 let out_f = w.out_features();
5664 #[allow(non_snake_case)]
5672 let GEMM_M_THRESHOLD = if self.verify_exact_on() { usize::MAX } else { 16usize };
5675
5676 const GEMM_MIN_OUT_F: usize = 128; if m >= GEMM_M_THRESHOLD {
5701 if let Some(y) = self.try_fp8_gemm(w, x, m)? { return Ok(y); }
5702 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(y); }
5709 if let Some(y) = self.try_f16_gemm(w, x, m)? { return Ok(y); }
5712 }
5713 if let GpuTensor::Quant { qtype, .. } = w {
5728 if *qtype == QT_F8_E4M3_BLK {
5729 if m >= GEMM_M_THRESHOLD {
5730 if let Some(y) = self.try_e4m3_blk_prefill(w, x, m)? { return Ok(y); }
5731 }
5732 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5733 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
5734 }
5735 }
5736 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.mmq_supports(w) {
5737 return self.qmatvec_mmq(w, x, m);
5738 }
5739 if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.gemm_supports(w) {
5740 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5741 return self.qmatvec_gemm(w, &aq, &ad, m);
5742 }
5743 if m >= GEMM_M_THRESHOLD {
5746 if let Some(y) = self.try_fp4_gemm(w, x, m, in_f, out_f)? { return Ok(y); }
5747 }
5748 let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
5752 if m == 1 && fast {
5757 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, scale, .. } = w {
5758 if self.mmvq_supports(*qtype) {
5759 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5763 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5764 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, rp);
5765 }
5766 }
5767 }
5768 if (2..=16).contains(&m) && fast && std::env::var("MEMRA_NO_BATCHED").is_err()
5784 && (m <= 4 || Self::b8_enabled()) {
5785 let m_ok = m <= 8 || matches!(w, GpuTensor::Quant { qtype, .. }
5795 if *qtype == QT_Q4_0 || *qtype == QT_Q6_K || *qtype == QT_F8_E4M3
5796 || *qtype == QT_NVFP4 || *qtype == QT_Q4_K || *qtype == QT_Q5_K || *qtype == QT_Q8_0);
5797 if m_ok {
5798 if let GpuTensor::Quant { bytes, qtype, row_bytes, rp, rp4, .. } = w {
5799 if self.batched_supports(*qtype) && self.mmvq_supports(*qtype) {
5800 let (bytes, rp) = match rp4 { Some(m4) => (m4, true), None => (bytes, *rp) };
5801 let mcols = Self::batched_mcols(m);
5802 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5803 let mut y = self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, mcols, 1.0, rp)?;
5804 if let GpuTensor::Quant { scale, .. } = w {
5805 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5806 }
5807 return Ok(y);
5808 }
5809 }
5810 }
5811 }
5812 if fast {
5818 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. } = w {
5819 if *qtype == QT_F8_E4M3 {
5820 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
5821 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes,
5822 *scale, false);
5823 }
5824 }
5825 }
5826 let mut y = match w {
5827 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q8_0 =>
5828 self.qmatvec_q8_0_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5829 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q4_K =>
5830 self.qmatvec_q4_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5831 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q6_K =>
5832 self.qmatvec_q6_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5833 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q5_K =>
5834 self.qmatvec_q5_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5835 GpuTensor::Quant { bytes, qtype, row_bytes, .. } if fast && *qtype == QT_Q3_K =>
5836 self.qmatvec_q3_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5837 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } if fast && *qtype == QT_NVFP4 =>
5838 self.qmatvec_dp4a_named(
5839 if *rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
5840 bytes, x, m, in_f, out_f, *row_bytes)?,
5841 GpuTensor::Quant { bytes, qtype, row_bytes, .. }
5845 if fast && *qtype == QT_IQ4_XS && Self::iq_fast_enabled() =>
5846 self.qmatvec_iq4_XS_fast(bytes, x, m, in_f, out_f, *row_bytes)?,
5847 GpuTensor::Quant { bytes, qtype, row_bytes, rp, .. } =>
5852 self.qmatvec(bytes, x, m, in_f, out_f,
5855 if *rp && *qtype == QT_NVFP4 { QT_NVFP4_RP } else { *qtype },
5856 *row_bytes)?,
5857 GpuTensor::Float { data, .. } => self.linear(x, data, m, in_f, out_f)?,
5858 GpuTensor::FloatBf16 { data, .. } =>
5861 self.linear_bf16_chunked(x, data, m, in_f, out_f, false)?,
5862 };
5863 if let GpuTensor::Quant { scale, .. } = w {
5865 if *scale != 1.0 { self.scale_inplace(&mut y, *scale, m * out_f)?; }
5866 }
5867 Ok(y)
5868 }
5869
5870 pub fn uses_q8_1_fast(&self, w: &crate::model::GpuTensor) -> bool {
5873 use crate::model::GpuTensor;
5874 if std::env::var("MEMRA_FAST").as_deref() == Ok("0") { return false; }
5875 match w {
5876 GpuTensor::Quant { qtype, .. } => matches!(*qtype,
5883 QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q3_K | QT_NVFP4 | QT_F8_E4M3
5884 | QT_F8_E4M3_BLK | QT_Q4_0)
5885 || (*qtype == QT_IQ4_XS && Self::iq_fast_enabled()),
5886 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
5887 }
5888 }
5889
5890 pub fn matmul_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
5895 x_fallback: &CudaSlice<f32>, m: usize)
5896 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5897 use crate::model::GpuTensor;
5898 let x_raw_ok = x_fallback.len() >= m * w.in_features();
5904 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5907 if let Some(y) = self.try_fp8_gemm(w, x_fallback, m)? { return Ok(y); }
5908 if let Some(y) = self.try_fp8_blk_mmq(w, x_fallback, m)? { return Ok(y); }
5911 if let Some(y) = self.try_f16_gemm(w, x_fallback, m)? { return Ok(y); }
5913 }
5914 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5920 if let Some(y) = self.try_e4m3_blk_prefill(w, x_fallback, m)? { return Ok(y); }
5921 }
5922 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
5923 if m >= 16 && w.out_features() >= 128 && self.mmq_supports(w) && !self.verify_exact_on()
5928 && x_raw_ok {
5929 return self.qmatvec_mmq(w, x_fallback, m);
5930 }
5931 if m >= 16 && x_raw_ok && !self.verify_exact_on() {
5934 if let Some(y) = self.try_fp4_gemm(w, x_fallback, m, w.in_features(), w.out_features())? {
5935 return Ok(y);
5936 }
5937 }
5938 if m >= 16 && self.gemm_supports(w) && !self.verify_exact_on() {
5941 return self.qmatvec_gemm(w, aq, ad, m);
5942 }
5943 if !self.uses_q8_1_fast(w) { return self.matmul(w, x_fallback, m); }
5944 let in_f = w.in_features();
5945 let out_f = w.out_features();
5946 let (bytes, qtype, row_bytes, scale, rp) = match w {
5947 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
5948 _ => unreachable!("uses_q8_1_fast guaranteed Quant"),
5949 };
5950 let (mbytes, mrp) = match w {
5953 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5954 _ => (bytes, rp),
5955 };
5956 if m == 1 && self.mmvq_supports(qtype) {
5960 return self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, mrp);
5961 }
5962 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
5975 && std::env::var("MEMRA_NO_BATCHED").is_err()
5976 && (m <= 4 || Self::b8_enabled())
5977 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_NVFP4
5981 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_F8_E4M3 || qtype == QT_Q8_0) {
5982 let mcols = Self::batched_mcols(m);
5983 return self.qmatvec_mmvq_batched(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, mrp);
5984 }
5985 if qtype == QT_F8_E4M3 || qtype == QT_Q4_0 {
5991 let (b2, r2) = if qtype == QT_Q4_0 { (mbytes, mrp) } else { (bytes, rp) };
5992 return self.qmatvec_mmvq(b2, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, r2);
5993 }
5994 let name = match qtype {
5995 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
5996 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
5997 QT_Q3_K => "qmatvec_q3_K_dp4a",
5998 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
5999 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
6000 _ => unreachable!(),
6001 };
6002 let f = self.func(name);
6003 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
6005 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
6006 let __s_b = self.gpu.stream();
6007 let mut b = __s_b.launch_builder(&f);
6008 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
6009 unsafe { b.launch(cfg)?; }
6010 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
6011 Ok(y)
6012 }
6013
6014 pub fn matmul_decode_exact(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
6022 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6023 use crate::model::GpuTensor;
6024 if let GpuTensor::Float { data, .. } = w {
6032 return self.linear_decode_exact(x, data, m, w.in_features(), w.out_features());
6033 }
6034 if let GpuTensor::FloatBf16 { data, .. } = w {
6037 let (in_f, out_f) = (w.in_features(), w.out_features());
6038 return self.linear_bf16_chunked(x, data, m, in_f, out_f, true);
6039 }
6040 if !self.uses_q8_1_fast(w) { return self.matmul(w, x, m); }
6041 let in_f = w.in_features();
6042 let out_f = w.out_features();
6043 let (bytes, qtype, row_bytes, scale, rp) = match w {
6044 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6045 _ => return self.matmul(w, x, m),
6046 };
6047 let (bytes, rp) = match w {
6050 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
6051 _ => (bytes, rp),
6052 };
6053 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6054 if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? { return Ok(y); }
6058 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
6067 && std::env::var("MEMRA_NO_BATCHED").is_err()
6068 && (m <= 4 || Self::b8_enabled())
6069 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
6072 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
6073 let mcols = Self::batched_mcols(m);
6074 return self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
6075 }
6076 if self.mmvq_supports(qtype) {
6077 return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
6080 }
6081 self.matmul_pre(w, &aq, &ad, x, m)
6084 }
6085
6086 pub fn matmul_decode_exact_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
6096 ad: &CudaSlice<f32>, m: usize)
6097 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6098 use crate::model::GpuTensor;
6099 debug_assert!(self.uses_q8_1_fast(w),
6100 "matmul_decode_exact_pre: caller must guarantee q8_1-fast");
6101 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(y); }
6103 let in_f = w.in_features();
6104 let out_f = w.out_features();
6105 let (bytes, qtype, row_bytes, scale, rp) = match w {
6106 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
6107 (bytes, *qtype, *row_bytes, *scale, *rp),
6108 _ => return Err("matmul_decode_exact_pre: Quant tensor required (q8_1-fast contract)".into()),
6109 };
6110 let (bytes, rp) = match w {
6112 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
6113 _ => (bytes, rp),
6114 };
6115 if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
6117 && std::env::var("MEMRA_NO_BATCHED").is_err()
6118 && (m <= 4 || Self::b8_enabled())
6119 && (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
6120 || qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0) {
6121 let mcols = Self::batched_mcols(m);
6122 return self.qmatvec_mmvq_batched(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp);
6123 }
6124 if self.mmvq_supports(qtype) {
6125 return self.qmatvec_mmvq(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
6126 }
6127 let x0 = self.zeros(0)?;
6130 self.matmul_pre(w, aq, ad, &x0, m)
6131 }
6132
6133 pub fn matmul_decode_exact_dual_pre(&self, w0: &crate::model::GpuTensor,
6142 w1: &crate::model::GpuTensor,
6143 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6144 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
6145 use crate::model::GpuTensor;
6146 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6147 let on = *ON.get_or_init(|| {
6148 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
6149 });
6150 if !on || !(2..=7).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
6151 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
6152 return Ok(None);
6153 }
6154 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6159 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6160 if w1.in_features() != in_f || w1.out_features() != out_f {
6161 return Ok(None);
6162 }
6163 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
6164 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
6165 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
6166 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
6167 (b0, b1, *rb0, *s0, *s1, *rp0),
6168 _ => return Ok(None),
6169 };
6170 if m > 4 && !(rp && Self::b8_enabled()
6173 && std::env::var("MEMRA_B567").as_deref() != Ok("0")) {
6174 return Ok(None);
6175 }
6176 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, aq, ad, m, in_f, out_f, row_bytes, rp)?;
6177 Ok(Some(((y0, s0), (y1, s1))))
6178 }
6179
6180 pub fn matmul_decode_exact_dual(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6196 x: &CudaSlice<f32>, m: usize)
6197 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6198 use crate::model::GpuTensor;
6199 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6200 let on = *ON.get_or_init(|| {
6201 std::env::var("MEMRA_SPEC_DUAL_T").map(|v| v != "0").unwrap_or(true)
6202 });
6203 if !on || !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok()
6204 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
6205 return Ok(None);
6206 }
6207 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6212 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6213 if w1.in_features() != in_f || w1.out_features() != out_f {
6214 return Ok(None);
6215 }
6216 let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
6217 (GpuTensor::Quant { bytes: b0, qtype: q0, row_bytes: rb0, scale: s0, rp: rp0, rp4: None, .. },
6218 GpuTensor::Quant { bytes: b1, qtype: q1, row_bytes: rb1, scale: s1, rp: rp1, rp4: None, .. })
6219 if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 =>
6220 (b0, b1, *rb0, *s0, *s1, *rp0),
6221 _ => return Ok(None),
6222 };
6223 if std::env::var("MEMRA_DEBUG").is_ok() {
6226 static ONCE: std::sync::Once = std::sync::Once::new();
6227 ONCE.call_once(|| eprintln!("[memra] dual gate+up batched ENGAGED (m={m} rp={rp})"));
6228 }
6229 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6230 let (y0, y1) = self.qmatvec_batched_dual_raw(b0, b1, &aq, &ad, m, in_f, out_f, row_bytes, rp)?;
6231 let mut y0 = y0;
6232 let mut y1 = y1;
6233 if s0 != 1.0 { self.scale_inplace(&mut y0, s0, m * out_f)?; }
6234 if s1 != 1.0 { self.scale_inplace(&mut y1, s1, m * out_f)?; }
6235 Ok(Some((y0, y1)))
6236 }
6237
6238 #[allow(clippy::too_many_arguments)]
6243 pub fn qmatvec_batched_dual_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6244 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6245 m: usize, in_f: usize, out_f: usize, row_bytes: usize, rp: bool)
6246 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6247 const ROWS_PER_BLOCK: u32 = 4;
6248 let mcols = Self::batched_mcols(m);
6249 let tiny_rp1 = rp && mcols == 4 && out_f <= 128
6252 && std::env::var("MEMRA_NVFP4_AUX_DUAL").as_deref() != Ok("0");
6253 let (name, rows_per_block) = if tiny_rp1 {
6254 ("qmatvec_nvfp4_mmvq_dual_b4_rp", ROWS_PER_BLOCK)
6255 } else { match (mcols, rp, m) {
6256 (2, false, _) => ("qmatvec_nvfp4_mmvq_dual_b2", ROWS_PER_BLOCK),
6257 (4, false, _) => ("qmatvec_nvfp4_mmvq_dual_b4_r2", ROWS_PER_BLOCK * 2),
6258 (2, true, _) => ("qmatvec_nvfp4_mmvq_dual_b2_rp", ROWS_PER_BLOCK),
6259 (4, true, _) => ("qmatvec_nvfp4_mmvq_dual_b4_rpr2", ROWS_PER_BLOCK * 2),
6260 (8, true, 5) => ("qmatvec_nvfp4_mmvq_dual_b5_rpr2", ROWS_PER_BLOCK * 2),
6261 (8, true, 6) => ("qmatvec_nvfp4_mmvq_dual_b6_rpr2", ROWS_PER_BLOCK * 2),
6262 (8, true, 7) => ("qmatvec_nvfp4_mmvq_dual_b7_rpr2", ROWS_PER_BLOCK * 2),
6263 _ => return Err(format!("qmatvec_batched_dual_raw: no dual kernel for m {m}").into()),
6264 }};
6265 let f = self.func(name);
6266 let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
6267 let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
6268 let cfg = LaunchConfig {
6269 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6270 block_dim: (32, ROWS_PER_BLOCK, 1),
6271 shared_mem_bytes: 0,
6272 };
6273 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
6274 let __s_b = self.gpu.stream();
6275 let mut b = __s_b.launch_builder(&f);
6276 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6277 .arg(&inf).arg(&outf).arg(&mi).arg(&rb);
6278 unsafe { b.launch(cfg)?; }
6279 Ok((y0, y1))
6280 }
6281
6282 pub fn matmul_pre_dual_noscale(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6294 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6295 -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>> {
6296 use crate::model::GpuTensor;
6297 if m != 1 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6298 if !self.mmvq_supports(QT_NVFP4) { return Ok(None); }
6308 let (in_f, out_f) = (w0.in_features(), w0.out_features());
6309 if w1.in_features() != in_f || w1.out_features() != out_f { return Ok(None); }
6310 let no_mirror = |w: &crate::model::GpuTensor| {
6323 !matches!(w, GpuTensor::Quant { rp4: Some(_), .. })
6324 };
6325 if self.q8_ffn_fuse2_on()
6326 && no_mirror(w0) && no_mirror(w1)
6327 && let Some([p0, p1]) = self.q8_fused_params(&[w0, w1])
6328 {
6329 let (y0, y1) = self.q8_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2)?;
6330 return Ok(Some(((y0, 1.0), (y1, 1.0))));
6331 }
6332 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6342 let (y0, y1) = self.e4m3_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2,
6343 1.0, 1.0)?;
6344 return Ok(Some(((y0, p0.3), (y1, p1.3))));
6345 }
6346 let (b0, q0, rb0, s0, rp0) = match w0 {
6347 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6348 _ => return Ok(None),
6349 };
6350 let (b1, q1, rb1, s1, rp1) = match w1 {
6351 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
6352 _ => return Ok(None),
6353 };
6354 if q0 != QT_NVFP4 || q1 != QT_NVFP4 || rb0 != rb1 || rp0 != rp1 { return Ok(None); }
6355 const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
6357 let rows_per_block = ROWS_PER_BLOCK * RPW;
6358 let f = self.func(if rp0 { "qmatvec_nvfp4_mmvq_dual_mr2_rp" } else { "qmatvec_nvfp4_mmvq_dual_mr2" });
6359 let mut y0 = self.alloc_uninit::<f32>(out_f)?;
6360 let mut y1 = self.alloc_uninit::<f32>(out_f)?;
6361 let cfg = LaunchConfig {
6362 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
6363 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0,
6364 };
6365 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, rb0 as i64);
6366 let one = 1.0f32;
6369 let __s_b = self.gpu.stream();
6370 let mut b = __s_b.launch_builder(&f);
6371 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6372 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&one).arg(&one);
6373 unsafe { b.launch(cfg)?; }
6374 Ok(Some(((y0, s0), (y1, s1))))
6375 }
6376
6377 pub fn matmul_q8_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6385 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6386 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6387 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6393 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(),
6394 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6395 }
6396 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6397 Ok(Some(self.q8_fused2_core(p0.0, p1.0, aq, ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6398 }
6399
6400 #[allow(clippy::too_many_arguments)]
6401 fn q8_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6402 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6403 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6404 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6405 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6407 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6408 let f = self.func("qmatvec_q8_0_mmvq_fused2");
6409 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6410 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6411 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6412 shared_mem_bytes: 0 };
6413 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
6414 let __s_b = self.gpu.stream();
6415 let mut b = __s_b.launch_builder(&f);
6416 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6417 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl);
6418 unsafe { b.launch(cfg)?; }
6419 Ok((y0, y1))
6420 }
6421
6422 pub fn matmul_q8_fused2_x(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6428 x: &CudaSlice<f32>)
6429 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6430 if !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) { return Ok(None); }
6431 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6432 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6433 return Ok(Some(self.e4m3_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(),
6434 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6435 }
6436 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6437 let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
6438 Ok(Some(self.q8_fused2_core(p0.0, p1.0, &aq, &ad, w0.in_features(), p0.1, p1.1, p0.2)?))
6439 }
6440
6441 #[allow(clippy::too_many_arguments)]
6444 pub fn qmatvec_q8_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
6445 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6446 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6447 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6448 self.q8_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes)
6449 }
6450
6451 pub fn matmul_q4_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6457 w2: &crate::model::GpuTensor,
6458 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6459 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6460 use crate::model::GpuTensor;
6461 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6462 match w {
6463 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6464 Some((*row_bytes, w.out_features())),
6465 _ => None,
6466 }
6467 };
6468 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6469 else { return Ok(None) };
6470 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6471 return Ok(None);
6472 }
6473 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6477 match w {
6478 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6479 Some(m) => (m, true),
6480 None => (bytes, *rp),
6481 },
6482 _ => unreachable!(),
6483 }
6484 }
6485 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6486 if rp0 != rp1 || rp1 != rp2 { return Ok(None); }
6487 let rp = rp0;
6488 let rpb: u32 = 4;
6489 let mr1 = rp && Self::q40_mr1_on();
6493 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6494 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6495 let grid = nb(o0) + nb(o1) + nb(o2);
6496 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6497 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6498 let mut y2 = self.alloc_uninit::<f32>(o2)?;
6499 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6500 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6501 else { "qmatvec_q4_0_mmvq_fused3" });
6502 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6503 let inf = w0.in_features() as i32;
6504 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6505 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6506 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6509 {
6510 use cudarc::driver::{DevicePtr, DevicePtrMut};
6511 let s = &self.gpu.stream();
6512 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6513 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6514 let (pad, _g4) = ad.device_ptr(s);
6515 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6516 let (py2, _g7) = y2.device_ptr_mut(s);
6517 let mut ps = [
6518 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6519 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6520 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6521 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6522 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6523 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6524 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6525 &r2 as *const _ as *mut _,
6526 ];
6527 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6528 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6529 }
6530 return Ok(Some((y0, y1, y2)));
6531 }
6532 let __s_b = self.gpu.stream();
6533 let mut b = __s_b.launch_builder(&f);
6534 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6535 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6536 unsafe { b.launch(cfg)?; }
6537 Ok(Some((y0, y1, y2)))
6538 }
6539
6540 #[allow(clippy::too_many_arguments)]
6543 pub fn matmul_q4_fused3_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6544 w2: &crate::model::GpuTensor,
6545 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6546 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>,
6547 y2: &mut CudaSlice<f32>)
6548 -> Result<bool, Box<dyn std::error::Error>> {
6549 use crate::model::GpuTensor;
6550 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6551 match w {
6552 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6553 Some((*row_bytes, w.out_features())),
6554 _ => None,
6555 }
6556 };
6557 let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2))
6558 else { return Ok(false) };
6559 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6560 return Ok(false);
6561 }
6562 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6563 match w {
6564 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6565 Some(m) => (m, true),
6566 None => (bytes, *rp),
6567 },
6568 _ => unreachable!(),
6569 }
6570 }
6571 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6572 if rp0 != rp1 || rp1 != rp2 { return Ok(false); }
6573 let rp = rp0;
6574 let rpb: u32 = 4;
6575 let mr1 = rp && Self::q40_mr1_on();
6576 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6577 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6578 let grid = nb(o0) + nb(o1) + nb(o2);
6579 debug_assert!(y0.len() >= o0 && y1.len() >= o1 && y2.len() >= o2);
6580 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused3_mr1_rp" }
6581 else if rp { "qmatvec_q4_0_mmvq_fused3_rp" }
6582 else { "qmatvec_q4_0_mmvq_fused3" });
6583 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6584 let inf = w0.in_features() as i32;
6585 let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
6586 let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
6587 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6589 use cudarc::driver::{DevicePtr, DevicePtrMut};
6590 let s = &self.gpu.stream();
6591 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6592 let (p2, _g2) = b2.device_ptr(s); let (paq, _g3) = aq.device_ptr(s);
6593 let (pad, _g4) = ad.device_ptr(s);
6594 let (py0, _g5) = y0.device_ptr_mut(s); let (py1, _g6) = y1.device_ptr_mut(s);
6595 let (py2, _g7) = y2.device_ptr_mut(s);
6596 let mut ps = [
6597 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6598 &p2 as *const _ as *mut _, &paq as *const _ as *mut _,
6599 &pad as *const _ as *mut _, &py0 as *const _ as *mut _,
6600 &py1 as *const _ as *mut _, &py2 as *const _ as *mut _,
6601 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6602 &oo1 as *const _ as *mut _, &oo2 as *const _ as *mut _,
6603 &r0 as *const _ as *mut _, &r1 as *const _ as *mut _,
6604 &r2 as *const _ as *mut _,
6605 ];
6606 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused3_mr1_rp",
6607 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6608 return Ok(true);
6609 }
6610 let __s_b = self.gpu.stream();
6611 let mut b = __s_b.launch_builder(&f);
6612 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1).arg(&mut *y2)
6613 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&r0).arg(&r1).arg(&r2);
6614 unsafe { b.launch(cfg)?; }
6615 Ok(true)
6616 }
6617
6618 pub fn matmul_q4_fused2(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6620 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6621 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6622 use crate::model::GpuTensor;
6623 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6624 match w {
6625 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6626 Some((*row_bytes, w.out_features())),
6627 _ => None,
6628 }
6629 };
6630 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6631 if w0.in_features() != w1.in_features() { return Ok(None); }
6632 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6634 match w {
6635 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6636 Some(m) => (m, true),
6637 None => (bytes, *rp),
6638 },
6639 _ => unreachable!(),
6640 }
6641 }
6642 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6643 if rp0 != rp1 { return Ok(None); }
6644 let rp = rp0;
6645 let rpb: u32 = 4;
6646 let mr1 = rp && Self::q40_mr1_on();
6648 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6649 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6650 let grid = nb(o0) + nb(o1);
6651 let mut y0 = self.alloc_uninit::<f32>(o0)?;
6652 let mut y1 = self.alloc_uninit::<f32>(o1)?;
6653 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6654 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6655 else { "qmatvec_q4_0_mmvq_fused2" });
6656 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6657 let inf = w0.in_features() as i32;
6658 let (oo0, oo1) = (o0 as i32, o1 as i32);
6659 let (r0, r1) = (rb0 as i64, rb1 as i64);
6660 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6662 {
6663 use cudarc::driver::{DevicePtr, DevicePtrMut};
6664 let s = &self.gpu.stream();
6665 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6666 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6667 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6668 let mut ps = [
6669 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6670 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6671 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6672 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6673 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6674 &r1 as *const _ as *mut _,
6675 ];
6676 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6677 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6678 }
6679 return Ok(Some((y0, y1)));
6680 }
6681 let __s_b = self.gpu.stream();
6682 let mut b = __s_b.launch_builder(&f);
6683 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6684 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6685 unsafe { b.launch(cfg)?; }
6686 Ok(Some((y0, y1)))
6687 }
6688
6689 pub fn matmul_q4_fused2_into(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6691 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6692 y0: &mut CudaSlice<f32>, y1: &mut CudaSlice<f32>)
6693 -> Result<bool, Box<dyn std::error::Error>> {
6694 use crate::model::GpuTensor;
6695 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6696 match w {
6697 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6698 Some((*row_bytes, w.out_features())),
6699 _ => None,
6700 }
6701 };
6702 let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(false) };
6703 if w0.in_features() != w1.in_features() { return Ok(false); }
6704 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6705 match w {
6706 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6707 Some(m) => (m, true),
6708 None => (bytes, *rp),
6709 },
6710 _ => unreachable!(),
6711 }
6712 }
6713 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6714 if rp0 != rp1 { return Ok(false); }
6715 let rp = rp0;
6716 let rpb: u32 = 4;
6717 let mr1 = rp && Self::q40_mr1_on();
6718 let nb = |o: usize| if mr1 { (o as u32).div_ceil(rpb) }
6719 else { (o as u32).div_ceil(2).div_ceil(rpb) };
6720 let grid = nb(o0) + nb(o1);
6721 debug_assert!(y0.len() >= o0 && y1.len() >= o1);
6722 let f = self.func(if mr1 { "qmatvec_q4_0_mmvq_fused2_mr1_rp" }
6723 else if rp { "qmatvec_q4_0_mmvq_fused2_rp" }
6724 else { "qmatvec_q4_0_mmvq_fused2" });
6725 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1), shared_mem_bytes: 0 };
6726 let inf = w0.in_features() as i32;
6727 let (oo0, oo1) = (o0 as i32, o1 as i32);
6728 let (r0, r1) = (rb0 as i64, rb1 as i64);
6729 if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
6731 use cudarc::driver::{DevicePtr, DevicePtrMut};
6732 let s = &self.gpu.stream();
6733 let (p0, _g0) = b0.device_ptr(s); let (p1, _g1) = b1.device_ptr(s);
6734 let (paq, _g2) = aq.device_ptr(s); let (pad, _g3) = ad.device_ptr(s);
6735 let (py0, _g4) = y0.device_ptr_mut(s); let (py1, _g5) = y1.device_ptr_mut(s);
6736 let mut ps = [
6737 &p0 as *const _ as *mut std::ffi::c_void, &p1 as *const _ as *mut _,
6738 &paq as *const _ as *mut _, &pad as *const _ as *mut _,
6739 &py0 as *const _ as *mut _, &py1 as *const _ as *mut _,
6740 &inf as *const _ as *mut _, &oo0 as *const _ as *mut _,
6741 &oo1 as *const _ as *mut _, &r0 as *const _ as *mut _,
6742 &r1 as *const _ as *mut _,
6743 ];
6744 unsafe { self.launch_pdl("qmatvec_q4_0_mmvq_fused2_mr1_rp",
6745 (grid, 1, 1), (32, rpb, 1), &mut ps)?; }
6746 return Ok(true);
6747 }
6748 let __s_b = self.gpu.stream();
6749 let mut b = __s_b.launch_builder(&f);
6750 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut *y0).arg(&mut *y1)
6751 .arg(&inf).arg(&oo0).arg(&oo1).arg(&r0).arg(&r1);
6752 unsafe { b.launch(cfg)?; }
6753 Ok(true)
6754 }
6755
6756 pub fn matmul_q4_fused2_batched(&self, w0: &crate::model::GpuTensor,
6761 w1: &crate::model::GpuTensor,
6762 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6763 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6764 use crate::model::GpuTensor;
6765 if m < 2 || m > 8 { return Ok(None); }
6766 let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
6767 match w {
6768 GpuTensor::Quant { qtype, row_bytes, .. } if *qtype == QT_Q4_0 =>
6769 Some((*row_bytes, w.out_features())),
6770 _ => None,
6771 }
6772 };
6773 let (Some((rb0, o0)), Some((_rb1, o1))) = (q4(w0), q4(w1)) else { return Ok(None) };
6774 if w0.in_features() != w1.in_features() { return Ok(None); }
6775 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6776 match w {
6777 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6778 Some(mr) => (mr, true),
6779 None => (bytes, *rp),
6780 },
6781 _ => unreachable!(),
6782 }
6783 }
6784 let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
6785 if !rp0 || !rp1 { return Ok(None); }
6786 let mcols = Self::batched_mcols(m);
6787 let rpb: u32 = 4;
6788 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6789 let grid = nb(o0) + nb(o1);
6790 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6791 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6792 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f2_rp",
6793 4 => "qmatvec_q4_0_mmvq_b4_f2_rp",
6794 _ => "qmatvec_q4_0_mmvq_b8_f2_rp" });
6795 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6796 shared_mem_bytes: 0 };
6797 let inf = w0.in_features() as i32;
6798 let (oo0, oo1, mi) = (o0 as i32, o1 as i32, m as i32);
6799 let rb = rb0 as i64;
6800 let __s_b = self.gpu.stream();
6801 let mut b = __s_b.launch_builder(&f);
6802 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6803 .arg(&inf).arg(&oo0).arg(&oo1).arg(&mi).arg(&rb);
6804 unsafe { b.launch(cfg)?; }
6805 Ok(Some((y0, y1)))
6806 }
6807
6808 #[allow(clippy::too_many_arguments)]
6811 pub fn matmul_q4_fused3_batched(&self, w0: &crate::model::GpuTensor,
6812 w1: &crate::model::GpuTensor, w2: &crate::model::GpuTensor,
6813 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6814 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6815 use crate::model::GpuTensor;
6816 if m < 2 || m > 8 { return Ok(None); }
6817 let q4 = |w: &GpuTensor| -> Option<usize> {
6818 match w {
6819 GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0 => Some(w.out_features()),
6820 _ => None,
6821 }
6822 };
6823 let (Some(o0), Some(o1), Some(o2)) = (q4(w0), q4(w1), q4(w2)) else { return Ok(None) };
6824 if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
6825 return Ok(None);
6826 }
6827 fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
6828 match w {
6829 GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
6830 Some(mr) => (mr, true),
6831 None => (bytes, *rp),
6832 },
6833 _ => unreachable!(),
6834 }
6835 }
6836 let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
6837 if !rp0 || !rp1 || !rp2 { return Ok(None); }
6838 let mcols = Self::batched_mcols(m);
6839 let rpb: u32 = 4;
6840 let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
6841 let grid = nb(o0) + nb(o1) + nb(o2);
6842 let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
6843 let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
6844 let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
6845 let f = self.func(match mcols { 2 => "qmatvec_q4_0_mmvq_b2_f3_rp",
6846 4 => "qmatvec_q4_0_mmvq_b4_f3_rp",
6847 _ => "qmatvec_q4_0_mmvq_b8_f3_rp" });
6848 let cfg = LaunchConfig { grid_dim: (grid, 1, 1), block_dim: (32, rpb, 1),
6849 shared_mem_bytes: 0 };
6850 let inf = w0.in_features() as i32;
6851 let (oo0, oo1, oo2, mi) = (o0 as i32, o1 as i32, o2 as i32, m as i32);
6852 let rb = 0i64;
6853 let __s_b = self.gpu.stream();
6854 let mut b = __s_b.launch_builder(&f);
6855 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6856 .arg(&inf).arg(&oo0).arg(&oo1).arg(&oo2).arg(&mi).arg(&rb);
6857 unsafe { b.launch(cfg)?; }
6858 Ok(Some((y0, y1, y2)))
6859 }
6860
6861 pub fn matmul_q8_fused3(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6862 w2: &crate::model::GpuTensor,
6863 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
6864 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6865 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
6868 return Ok(Some(self.e4m3_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6869 p0.1, p1.1, p2.1, p0.2,
6870 p0.3, p1.3, p2.3)?));
6871 }
6872 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
6873 Ok(Some(self.q8_fused3_core(p0.0, p1.0, p2.0, aq, ad, w0.in_features(),
6874 p0.1, p1.1, p2.1, p0.2)?))
6875 }
6876
6877 #[allow(clippy::too_many_arguments)]
6878 fn q8_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6879 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
6880 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
6881 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6882 const ROWS_PER_BLOCK: u32 = 4;
6883 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6884 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6885 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
6886 let f = self.func("qmatvec_q8_0_mmvq_fused3");
6887 let mut y0 = self.alloc_uninit::<f32>(out0)?;
6888 let mut y1 = self.alloc_uninit::<f32>(out1)?;
6889 let mut y2 = self.alloc_uninit::<f32>(out2)?;
6890 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6891 shared_mem_bytes: 0 };
6892 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32, row_bytes as i64);
6893 let __s_b = self.gpu.stream();
6894 let mut b = __s_b.launch_builder(&f);
6895 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
6896 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl);
6897 unsafe { b.launch(cfg)?; }
6898 Ok((y0, y1, y2))
6899 }
6900
6901 #[allow(clippy::too_many_arguments)]
6903 pub fn qmatvec_q8_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6904 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
6905 out2: usize, row_bytes: usize)
6906 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6907 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
6908 self.q8_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes)
6909 }
6910
6911 pub fn matmul_q8_fused2_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6922 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6923 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6924 if !(2..=8).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
6928 if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
6931 if m > 4 && !Self::b8_enabled() { return Ok(None); }
6932 return Ok(Some(self.e4m3_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(),
6933 p0.1, p1.1, p0.2, p0.3, p1.3)?));
6934 }
6935 let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else { return Ok(None) };
6936 Ok(Some(self.q8_fused2_t_core(p0.0, p1.0, aq, ad, m, w0.in_features(), p0.1, p1.1, p0.2)?))
6937 }
6938
6939 #[allow(clippy::too_many_arguments)]
6940 fn q8_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6941 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6942 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6943 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6944 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
6946 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
6947 let f = self.func(match Self::batched_mcols(m) {
6948 2 => "qmatvec_q8_0_mmvq_fused2_b2",
6949 4 => "qmatvec_q8_0_mmvq_fused2_b4",
6950 _ => "qmatvec_q8_0_mmvq_fused2_b8",
6952 });
6953 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
6954 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
6955 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
6956 shared_mem_bytes: 0 };
6957 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32, row_bytes as i64);
6958 let __s_b = self.gpu.stream();
6959 let mut b = __s_b.launch_builder(&f);
6960 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
6961 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
6962 unsafe { b.launch(cfg)?; }
6963 Ok((y0, y1))
6964 }
6965
6966 #[allow(clippy::too_many_arguments)]
6969 pub fn qmatvec_q8_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
6970 x: &CudaSlice<f32>, m: usize,
6971 in_f: usize, out0: usize, out1: usize, row_bytes: usize)
6972 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6973 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
6974 self.q8_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes)
6975 }
6976
6977 #[allow(clippy::too_many_arguments)]
6980 pub fn matmul_q8_fused3_t(&self, w0: &crate::model::GpuTensor, w1: &crate::model::GpuTensor,
6981 w2: &crate::model::GpuTensor,
6982 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize)
6983 -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6984 if !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() { return Ok(None); }
6985 if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
6986 return Ok(Some(self.e4m3_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
6987 p0.1, p1.1, p2.1, p0.2,
6988 p0.3, p1.3, p2.3)?));
6989 }
6990 let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else { return Ok(None) };
6991 Ok(Some(self.q8_fused3_t_core(p0.0, p1.0, p2.0, aq, ad, m, w0.in_features(),
6992 p0.1, p1.1, p2.1, p0.2)?))
6993 }
6994
6995 #[allow(clippy::too_many_arguments)]
6996 fn q8_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
6997 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
6998 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize)
6999 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7000 const ROWS_PER_BLOCK: u32 = 4;
7001 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7002 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7003 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
7004 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_q8_0_mmvq_fused3_b2" }
7005 else { "qmatvec_q8_0_mmvq_fused3_b4" });
7006 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7007 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7008 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
7009 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7010 shared_mem_bytes: 0 };
7011 let (inf, o0, o1, o2, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
7012 m as i32, row_bytes as i64);
7013 let __s_b = self.gpu.stream();
7014 let mut b = __s_b.launch_builder(&f);
7015 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
7016 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&mi).arg(&rbl);
7017 unsafe { b.launch(cfg)?; }
7018 Ok((y0, y1, y2))
7019 }
7020
7021 #[allow(clippy::too_many_arguments)]
7023 pub fn qmatvec_q8_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7024 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
7025 out1: usize, out2: usize, row_bytes: usize)
7026 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7027 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7028 self.q8_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes)
7029 }
7030
7031 pub fn q8_ffn_fuse2_on(&self) -> bool {
7035 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7036 *ON.get_or_init(|| std::env::var("MEMRA_Q8_FFN_FUSE2").as_deref() != Ok("0"))
7037 }
7038
7039 #[allow(clippy::type_complexity)]
7045 fn q8_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
7046 -> Option<[(&'w CudaSlice<u8>, usize, usize); N]> {
7047 use crate::model::GpuTensor;
7048 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return None; }
7049 if std::env::var("MEMRA_Q8_DUAL").is_ok_and(|v| v == "0") { return None; }
7050 let in_f = ws[0].in_features();
7051 let mut out: [Option<(&CudaSlice<u8>, usize, usize)>; N] = [None; N];
7052 for (i, w) in ws.iter().enumerate() {
7053 match w {
7054 GpuTensor::Quant { bytes, qtype, row_bytes, scale, .. }
7055 if *qtype == QT_Q8_0 && *scale == 1.0 && w.in_features() == in_f =>
7056 out[i] = Some((bytes, w.out_features(), *row_bytes)),
7057 _ => return None,
7058 }
7059 }
7060 Some(out.map(|o| o.unwrap()))
7061 }
7062
7063 pub fn e4m3_dual_on(&self) -> bool {
7066 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7067 *ON.get_or_init(|| std::env::var("MEMRA_E4M3_DUAL").as_deref() != Ok("0"))
7068 }
7069
7070 #[allow(clippy::type_complexity)]
7082 fn e4m3_fused_params<'w, const N: usize>(&self, ws: &[&'w crate::model::GpuTensor; N])
7083 -> Option<[(&'w CudaSlice<u8>, usize, usize, f32); N]> {
7084 use crate::model::GpuTensor;
7085 if !self.e4m3_dual_on() { return None; }
7086 let in_f = ws[0].in_features();
7087 let mut out: [Option<(&CudaSlice<u8>, usize, usize, f32)>; N] = [None; N];
7088 for (i, w) in ws.iter().enumerate() {
7089 match w {
7090 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, rp4, .. }
7091 if *qtype == QT_F8_E4M3 && w.in_features() == in_f
7092 && *row_bytes == in_f && !*rp && rp4.is_none() =>
7093 out[i] = Some((bytes, w.out_features(), *row_bytes, *scale)),
7094 _ => return None,
7095 }
7096 }
7097 Some(out.map(|o| o.unwrap()))
7098 }
7099
7100 #[allow(clippy::too_many_arguments)]
7104 fn e4m3_fused2_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7105 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7106 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7107 ws0: f32, ws1: f32)
7108 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7109 const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7111 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7112 let f = self.func("qmatvec_e4m3_mmvq_fused2");
7113 let mut y0 = self.alloc_uninit::<f32>(out0)?;
7114 let mut y1 = self.alloc_uninit::<f32>(out1)?;
7115 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7116 shared_mem_bytes: 0 };
7117 let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
7118 let __s_b = self.gpu.stream();
7119 let mut b = __s_b.launch_builder(&f);
7120 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
7121 .arg(&inf).arg(&o0).arg(&o1).arg(&rbl).arg(&ws0).arg(&ws1);
7122 unsafe { b.launch(cfg)?; }
7123 Ok((y0, y1))
7124 }
7125
7126 #[allow(clippy::too_many_arguments)]
7128 fn e4m3_fused3_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7129 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7130 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
7131 ws0: f32, ws1: f32, ws2: f32)
7132 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7133 const ROWS_PER_BLOCK: u32 = 4;
7134 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7135 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7136 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
7137 let f = self.func("qmatvec_e4m3_mmvq_fused3");
7138 let mut y0 = self.alloc_uninit::<f32>(out0)?;
7139 let mut y1 = self.alloc_uninit::<f32>(out1)?;
7140 let mut y2 = self.alloc_uninit::<f32>(out2)?;
7141 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7142 shared_mem_bytes: 0 };
7143 let (inf, o0, o1, o2, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
7144 row_bytes as i64);
7145 let __s_b = self.gpu.stream();
7146 let mut b = __s_b.launch_builder(&f);
7147 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
7148 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&rbl).arg(&ws0).arg(&ws1).arg(&ws2);
7149 unsafe { b.launch(cfg)?; }
7150 Ok((y0, y1, y2))
7151 }
7152
7153 #[allow(clippy::too_many_arguments)]
7157 fn e4m3_fused2_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7158 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
7159 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7160 ws0: f32, ws1: f32)
7161 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7162 const ROWS_PER_BLOCK: u32 = 4;
7163 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7164 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7165 let f = self.func(match Self::batched_mcols(m) {
7166 2 => "qmatvec_e4m3_mmvq_fused2_b2",
7167 4 => "qmatvec_e4m3_mmvq_fused2_b4",
7168 _ => "qmatvec_e4m3_mmvq_fused2_b8",
7169 });
7170 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7171 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7172 let cfg = LaunchConfig { grid_dim: (nb0 + nb1, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7173 shared_mem_bytes: 0 };
7174 let (inf, o0, o1, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, m as i32,
7175 row_bytes as i64);
7176 let __s_b = self.gpu.stream();
7177 let mut b = __s_b.launch_builder(&f);
7178 b.arg(b0).arg(b1).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1)
7179 .arg(&inf).arg(&o0).arg(&o1).arg(&mi).arg(&rbl);
7180 unsafe { b.launch(cfg)?; }
7181 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
7182 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
7183 Ok((y0, y1))
7184 }
7185
7186 #[allow(clippy::too_many_arguments)]
7188 fn e4m3_fused3_t_core(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7189 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, m: usize,
7190 in_f: usize, out0: usize, out1: usize, out2: usize, row_bytes: usize,
7191 ws0: f32, ws1: f32, ws2: f32)
7192 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7193 const ROWS_PER_BLOCK: u32 = 4;
7194 let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
7195 let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
7196 let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
7197 let f = self.func(if Self::batched_mcols(m) == 2 { "qmatvec_e4m3_mmvq_fused3_b2" }
7198 else { "qmatvec_e4m3_mmvq_fused3_b4" });
7199 let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
7200 let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
7201 let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
7202 let cfg = LaunchConfig { grid_dim: (nb0 + nb1 + nb2, 1, 1), block_dim: (32, ROWS_PER_BLOCK, 1),
7203 shared_mem_bytes: 0 };
7204 let (inf, o0, o1, o2, mi, rbl) = (in_f as i32, out0 as i32, out1 as i32, out2 as i32,
7205 m as i32, row_bytes as i64);
7206 let __s_b = self.gpu.stream();
7207 let mut b = __s_b.launch_builder(&f);
7208 b.arg(b0).arg(b1).arg(b2).arg(aq).arg(ad).arg(&mut y0).arg(&mut y1).arg(&mut y2)
7209 .arg(&inf).arg(&o0).arg(&o1).arg(&o2).arg(&mi).arg(&rbl);
7210 unsafe { b.launch(cfg)?; }
7211 if ws0 != 1.0 { self.scale_inplace(&mut y0, ws0, m * out0)?; }
7212 if ws1 != 1.0 { self.scale_inplace(&mut y1, ws1, m * out1)?; }
7213 if ws2 != 1.0 { self.scale_inplace(&mut y2, ws2, m * out2)?; }
7214 Ok((y0, y1, y2))
7215 }
7216
7217 pub fn qmatvec_e4m3_blk_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7227 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7228 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7229 scale_cols: usize)
7230 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7231 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_e4m3_blk_mmvq_into(bytes, aq, ad, scales, m, in_f, out_f, row_bytes,
7233 scale_cols, &mut y)?;
7234 Ok(y)
7235 }
7236
7237 #[allow(clippy::too_many_arguments)]
7239 pub fn qmatvec_e4m3_blk_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7240 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7241 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7242 scale_cols: usize, y: &mut CudaSlice<f32>)
7243 -> Result<(), Box<dyn std::error::Error>> {
7244 const ROWS_PER_BLOCK: u32 = 4; let f = self.func("qmatvec_e4m3_blk_mmvq");
7246 let cfg = LaunchConfig {
7247 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), m as u32, 1),
7248 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7251 let (inf, outf, mi, rb, sc) =
7252 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7253 let __s_b = self.gpu.stream();
7254 let mut b = __s_b.launch_builder(&f);
7255 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut *y)
7256 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7257 unsafe { b.launch(cfg)?; }
7258 Ok(())
7259 }
7260
7261 #[allow(clippy::too_many_arguments)]
7267 pub fn qmatvec_e4m3_blk_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>,
7268 ad: &CudaSlice<f32>, scales: &CudaSlice<f32>,
7269 m: usize, in_f: usize, out_f: usize, row_bytes: usize,
7270 scale_cols: usize, mcols: usize)
7271 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7272 const ROWS_PER_BLOCK: u32 = 4; debug_assert!(mcols >= m, "blk batched: mcols {mcols} < m {m}");
7274 let name = match mcols {
7275 2 => "qmatvec_e4m3_blk_mmvq_b2",
7276 4 => "qmatvec_e4m3_blk_mmvq_b4",
7277 8 => "qmatvec_e4m3_blk_mmvq_b8",
7278 16 => "qmatvec_e4m3_blk_mmvq_b16",
7279 _ => return Err(format!("qmatvec_e4m3_blk_mmvq_batched: no kernel for mcols {mcols}").into()),
7280 };
7281 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7282 let f = self.func(name);
7283 let cfg = LaunchConfig {
7284 grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
7285 block_dim: (32, ROWS_PER_BLOCK, 1),
7286 shared_mem_bytes: 0,
7287 };
7288 let (inf, outf, mi, rb, sc) =
7289 (in_f as i32, out_f as i32, m as i32, row_bytes as i64, scale_cols as i32);
7290 let __s_b = self.gpu.stream();
7291 let mut b = __s_b.launch_builder(&f);
7292 b.arg(bytes).arg(aq).arg(ad).arg(scales).arg(&mut y)
7293 .arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&sc);
7294 unsafe { b.launch(cfg)?; }
7295 Ok(y)
7296 }
7297
7298 #[allow(clippy::too_many_arguments)]
7301 pub fn qmatvec_e4m3_blk_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7302 scales: &CudaSlice<f32>, m: usize, in_f: usize,
7303 out_f: usize, row_bytes: usize, scale_cols: usize,
7304 mcols: usize)
7305 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7306 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7307 self.qmatvec_e4m3_blk_mmvq_batched(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes,
7308 scale_cols, mcols)
7309 }
7310
7311 #[allow(clippy::too_many_arguments)]
7314 pub fn qmatvec_e4m3_blk_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>,
7315 scales: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize,
7316 row_bytes: usize, scale_cols: usize)
7317 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7318 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7319 self.qmatvec_e4m3_blk_mmvq(bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols)
7320 }
7321
7322 #[allow(clippy::too_many_arguments)]
7325 pub fn qmatvec_e4m3_fused2_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, x: &CudaSlice<f32>,
7326 in_f: usize, out0: usize, out1: usize, row_bytes: usize,
7327 ws0: f32, ws1: f32)
7328 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7329 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7330 self.e4m3_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes, ws0, ws1)
7331 }
7332
7333 #[allow(clippy::too_many_arguments)]
7334 pub fn qmatvec_e4m3_fused3_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>, b2: &CudaSlice<u8>,
7335 x: &CudaSlice<f32>, in_f: usize, out0: usize, out1: usize,
7336 out2: usize, row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7337 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7338 let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
7339 self.e4m3_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes,
7340 ws0, ws1, ws2)
7341 }
7342
7343 #[allow(clippy::too_many_arguments)]
7344 pub fn qmatvec_e4m3_fused2_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7345 x: &CudaSlice<f32>, m: usize, in_f: usize, out0: usize,
7346 out1: usize, row_bytes: usize, ws0: f32, ws1: f32)
7347 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7348 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7349 self.e4m3_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes, ws0, ws1)
7350 }
7351
7352 #[allow(clippy::too_many_arguments)]
7353 pub fn qmatvec_e4m3_fused3_t_raw(&self, b0: &CudaSlice<u8>, b1: &CudaSlice<u8>,
7354 b2: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
7355 in_f: usize, out0: usize, out1: usize, out2: usize,
7356 row_bytes: usize, ws0: f32, ws1: f32, ws2: f32)
7357 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7358 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7359 self.e4m3_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes,
7360 ws0, ws1, ws2)
7361 }
7362
7363 fn try_e4m3_blk_pre(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>,
7374 ad: &CudaSlice<f32>, m: usize)
7375 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7376 use crate::model::GpuTensor;
7377 if let GpuTensor::Quant { bytes, qtype, row_bytes, blk: Some(g), .. } = w {
7378 if *qtype == QT_F8_E4M3_BLK {
7379 if (2..=16).contains(&m) && std::env::var("MEMRA_NO_BATCHED").is_err()
7385 && (m <= 4 || Self::b8_enabled()) {
7386 let mcols = Self::batched_mcols(m);
7387 return Ok(Some(self.qmatvec_e4m3_blk_mmvq_batched(
7388 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7389 *row_bytes, g.cols, mcols)?));
7390 }
7391 return Ok(Some(self.qmatvec_e4m3_blk_mmvq(
7392 bytes, aq, ad, &g.scales, m, w.in_features(), w.out_features(),
7393 *row_bytes, g.cols)?));
7394 }
7395 }
7396 Ok(None)
7397 }
7398
7399 fn try_e4m3_blk_prefill(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize)
7446 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7447 use crate::model::GpuTensor;
7448 let GpuTensor::Quant { bytes, qtype, blk: Some(g), .. } = w else { return Ok(None) };
7449 if *qtype != QT_F8_E4M3_BLK { return Ok(None) }
7450 if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? { return Ok(Some(y)); }
7455 let (in_f, out_f) = (w.in_features(), w.out_features());
7456 let slab = self.fp8_blk_dequant_q8_0_dev(bytes, &g.scales, out_f, in_f)?;
7457 let tmp = GpuTensor::Quant {
7458 bytes: slab,
7459 qtype: QT_Q8_0,
7460 row_bytes: in_f / 32 * 34,
7461 ne: vec![in_f as u64, out_f as u64],
7462 scale: 1.0,
7463 rp: false,
7464 #[cfg(memra_cutlass)]
7465 cutlass: None,
7466 fp8: None, blk: None, f16: None, rp4: None,
7467 };
7468 Ok(Some(self.matmul(&tmp, x, m)?))
7470 }
7471
7472 pub fn matmul_pre_noscale(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7473 m: usize) -> Result<Option<(CudaSlice<f32>, f32)>, Box<dyn std::error::Error>> {
7474 use crate::model::GpuTensor;
7475 if m == 1 {
7479 if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? { return Ok(Some((y, 1.0))); }
7480 }
7481 if m != 1 || !self.uses_q8_1_fast(w) { return Ok(None); }
7483 let in_f = w.in_features();
7484 let out_f = w.out_features();
7485 let (bytes, qtype, row_bytes, scale, rp) = match w {
7486 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
7487 _ => return Ok(None),
7488 };
7489 if self.mmvq_supports(qtype) {
7491 let (mbytes, mrp) = match w {
7493 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
7494 _ => (bytes, rp),
7495 };
7496 let y = self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, 1.0, mrp)?;
7497 return Ok(Some((y, scale)));
7498 }
7499 let name = match qtype {
7501 QT_Q8_0 => "qmatvec_q8_0_dp4a", QT_Q4_K => "qmatvec_q4_K_dp4a",
7502 QT_Q6_K => "qmatvec_q6_K_dp4a", QT_Q5_K => "qmatvec_q5_K_dp4a",
7503 QT_Q3_K => "qmatvec_q3_K_dp4a",
7504 QT_NVFP4 => if rp { "qmatvec_nvfp4_dp4a_rp" } else { "qmatvec_nvfp4_dp4a" },
7505 QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
7506 _ => return Ok(None),
7507 };
7508 let f = self.func(name);
7509 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
7510 let cfg = LaunchConfig { grid_dim: (out_f as u32, m as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
7511 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7512 let __s_b = self.gpu.stream();
7513 let mut b = __s_b.launch_builder(&f);
7514 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7515 unsafe { b.launch(cfg)?; }
7516 Ok(Some((y, scale)))
7517 }
7518
7519 pub fn mmvq_supports(&self, qtype: i32) -> bool {
7522 if qtype == QT_F8_E4M3 { return true; }
7527 if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") { return false; }
7528 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_Q4_0)
7529 }
7530
7531 pub fn qmatvec_mmvq(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7536 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7537 rp: bool)
7538 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7539 let mut y = self.alloc_uninit::<f32>(m * out_f)?; self.qmatvec_mmvq_into(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp, &mut y)?;
7541 Ok(y)
7542 }
7543
7544 #[allow(clippy::too_many_arguments)]
7546 pub fn qmatvec_mmvq_into(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
7547 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, scale: f32,
7548 rp: bool, y: &mut CudaSlice<f32>)
7549 -> Result<(), Box<dyn std::error::Error>> {
7550 debug_assert!(y.len() >= m * out_f);
7551 const ROWS_PER_BLOCK: u32 = 4; if qtype == QT_Q8_0 && rp && m == 1 && out_f >= 64
7557 && (out_f as u32).div_ceil(ROWS_PER_BLOCK) < 4 * self.sm_count() as u32
7558 && {
7559 static G2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7560 *G2.get_or_init(|| std::env::var("MEMRA_Q80_G2").as_deref() != Ok("0"))
7561 }
7562 {
7563 let f = self.func("qmatvec_q8_0_mmvq_rp_g2");
7564 let cfg = LaunchConfig {
7565 grid_dim: ((out_f as u32).div_ceil(2), 1, 1),
7566 block_dim: (32, 2, 1),
7567 shared_mem_bytes: 0,
7568 };
7569 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, row_bytes as i64);
7570 let __s_b = self.gpu.stream();
7571 let mut b = __s_b.launch_builder(&f);
7572 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7573 unsafe { b.launch(cfg)?; }
7574 if scale != 1.0 { self.scale_inplace(y, scale, out_f)?; }
7575 return Ok(());
7576 }
7577 let mut mr: u32 = if m == 1 && (qtype == QT_NVFP4 || qtype == QT_Q5_K) { 2 } else { 1 };
7586 if m == 1 && qtype == QT_Q4_0 {
7591 static Q40MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7592 mr = *Q40MR.get_or_init(|| std::env::var("MEMRA_Q40_MR").ok()
7595 .and_then(|v| v.parse().ok()).unwrap_or(1));
7596 }
7597 let q5_mode = std::env::var("MEMRA_Q5K_ISSUE").ok();
7608 let q5_force = q5_mode.as_deref() == Some("2");
7609 let q5_il = qtype == QT_Q5_K && m == 1
7612 && (q5_force || q5_mode.as_deref().map(|v| v != "0").unwrap_or(true));
7613 if q5_il && !q5_force && out_f > 65536 { mr = 1; }
7614 if qtype == QT_Q4_0 && rp && mr != 1 { mr = 2; }
7617 if qtype == QT_Q8_0 && rp {
7621 static Q80MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
7622 mr = *Q80MR.get_or_init(|| std::env::var("MEMRA_Q80_MR").ok()
7623 .and_then(|v| v.parse().ok()).unwrap_or(1));
7624 }
7625 let name = match (qtype, mr, rp) {
7626 (QT_NVFP4, 2, false) => "qmatvec_nvfp4_mmvq_mr2",
7627 (QT_NVFP4, 2, true) => "qmatvec_nvfp4_mmvq_mr2_rp",
7628 (QT_NVFP4, _, true) => "qmatvec_nvfp4_mmvq_rp",
7629 (QT_Q4_0, 1, true) => "qmatvec_q4_0_mmvq_rp",
7630 (QT_Q4_0, _, true) => "qmatvec_q4_0_mmvq_mr2_rp",
7631 (QT_Q5_K, 2, _) => if q5_il { "qmatvec_q5_K_mmvq_mr2_il" } else { "qmatvec_q5_K_mmvq_mr2" },
7632 (QT_Q8_0, 2, true) => "qmatvec_q8_0_mmvq_mr2_rp",
7633 (QT_Q8_0, _, true) if in_f % 1024 == 0 && {
7638 static CA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7639 *CA.get_or_init(|| std::env::var("MEMRA_Q80_CA").as_deref() == Ok("1"))
7640 } => "qmatvec_q8_0_mmvq_rpca",
7641 (QT_Q8_0, _, true) => "qmatvec_q8_0_mmvq_rp",
7642 (QT_Q8_0, _, _) => "qmatvec_q8_0_mmvq",
7643 (QT_Q4_K, _, true) => "qmatvec_q4_K_mmvq_rp",
7647 (QT_Q6_K, _, true) => "qmatvec_q6_K_mmvq_rp",
7648 (QT_Q4_K, _, _) => "qmatvec_q4_K_mmvq",
7649 (QT_Q4_0, 2, false) => "qmatvec_q4_0_mmvq_mr2",
7650 (QT_Q4_0, _, false) => "qmatvec_q4_0_mmvq",
7651 (QT_Q5_K, _, _) => if q5_il { "qmatvec_q5_K_mmvq_il" } else { "qmatvec_q5_K_mmvq" },
7652 (QT_Q6_K, _, _) => "qmatvec_q6_K_mmvq",
7653 (QT_NVFP4, _, false) => "qmatvec_nvfp4_mmvq",
7654 (QT_F8_E4M3, _, _) => "qmatvec_e4m3_mmvq",
7655 _ => panic!("qmatvec_mmvq: qtype {qtype} has no MMVQ kernel"),
7656 };
7657 let f = self.func(name);
7658 let rows_per_block = ROWS_PER_BLOCK * mr;
7660 let cfg = LaunchConfig {
7661 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, m as u32, 1),
7662 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
7665 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
7666 let __s_b = self.gpu.stream();
7667 let mut b = __s_b.launch_builder(&f);
7668 if qtype == QT_NVFP4 || qtype == QT_F8_E4M3 {
7673 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb).arg(&scale);
7674 unsafe { b.launch(cfg)?; }
7675 } else if Self::pdl_on() && Self::pdl_mmvq_on()
7676 && matches!(name, "qmatvec_q4_0_mmvq_rp" | "qmatvec_q6_K_mmvq"
7677 | "qmatvec_q6_K_mmvq_rp") {
7678 {
7682 use cudarc::driver::{DevicePtr, DevicePtrMut};
7683 let s = &self.gpu.stream();
7684 let (pw, _g0) = bytes.device_ptr(s); let (paq, _g1) = aq.device_ptr(s);
7685 let (pad, _g2) = ad.device_ptr(s); let (py, _g3) = y.device_ptr_mut(s);
7686 let mut ps = [
7687 &pw as *const _ as *mut std::ffi::c_void, &paq as *const _ as *mut _,
7688 &pad as *const _ as *mut _, &py as *const _ as *mut _,
7689 &inf as *const _ as *mut _, &outf as *const _ as *mut _,
7690 &mi as *const _ as *mut _, &rb as *const _ as *mut _,
7691 ];
7692 unsafe { self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?; }
7693 }
7694 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7695 } else {
7696 b.arg(bytes).arg(aq).arg(ad).arg(&mut *y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
7697 unsafe { b.launch(cfg)?; }
7698 if scale != 1.0 { self.scale_inplace(y, scale, m * out_f)?; }
7699 }
7700 Ok(())
7701 }
7702
7703 pub fn qmatvec_mmvq_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
7707 out_f: usize, qtype: i32, row_bytes: usize, rp: bool)
7708 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7709 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
7710 self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, 1.0, rp)
7711 }
7712
7713 pub fn batched_supports(&self, qtype: i32) -> bool {
7717 matches!(qtype, QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_F8_E4M3 | QT_Q4_0)
7718 }
7719
7720 pub fn iq_fast_enabled() -> bool {
7728 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7729 *ON.get_or_init(|| std::env::var("MEMRA_IQ_FAST").map(|v| v != "0").unwrap_or(true))
7730 }
7731
7732 pub fn b8_enabled() -> bool {
7735 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7736 *ON.get_or_init(|| std::env::var("MEMRA_B8").map(|v| v != "0").unwrap_or(true))
7737 }
7738
7739 pub fn batched_mcols(m: usize) -> usize {
7741 if m == 2 { 2 } else if m <= 4 { 4 } else if m <= 8 { 8 } else { 16 }
7742 }
7743
7744 fn batched_kernel_name(qtype: i32, mcols: usize) -> Option<&'static str> {
7749 Some(match (qtype, mcols) {
7750 (QT_Q8_0, 2) => "qmatvec_q8_0_mmvq_b2", (QT_Q8_0, 4) => "qmatvec_q8_0_mmvq_b4",
7751 (QT_Q8_0, 8) => "qmatvec_q8_0_mmvq_b8",
7752 (QT_Q8_0, 16) => "qmatvec_q8_0_mmvq_b16",
7758 (QT_Q4_K, 2) => "qmatvec_q4_K_mmvq_b2", (QT_Q4_K, 4) => "qmatvec_q4_K_mmvq_b4",
7759 (QT_Q4_K, 8) => "qmatvec_q4_K_mmvq_b8",
7760 (QT_Q4_K, 16) => "qmatvec_q4_K_mmvq_b16",
7763 (QT_Q5_K, 2) => "qmatvec_q5_K_mmvq_b2", (QT_Q5_K, 4) => "qmatvec_q5_K_mmvq_b4",
7764 (QT_Q5_K, 8) => "qmatvec_q5_K_mmvq_b8",
7765 (QT_Q5_K, 16) => "qmatvec_q5_K_mmvq_b16",
7768 (QT_Q6_K, 2) => "qmatvec_q6_K_mmvq_b2", (QT_Q6_K, 4) => "qmatvec_q6_K_mmvq_b4",
7769 (QT_Q6_K, 8) => "qmatvec_q6_K_mmvq_b8", (QT_Q6_K, 16) => "qmatvec_q6_K_mmvq_b16",
7770 (QT_NVFP4, 2) => "qmatvec_nvfp4_mmvq_b2", (QT_NVFP4, 4) => "qmatvec_nvfp4_mmvq_b4",
7771 (QT_NVFP4, 8) => "qmatvec_nvfp4_mmvq_b8",
7772 (QT_NVFP4, 16) => "qmatvec_nvfp4_mmvq_b16",
7776 (QT_F8_E4M3, 2) => "qmatvec_e4m3_mmvq_b2", (QT_F8_E4M3, 4) => "qmatvec_e4m3_mmvq_b4",
7777 (QT_F8_E4M3, 8) => "qmatvec_e4m3_mmvq_b8",
7778 (QT_F8_E4M3, 16) => "qmatvec_e4m3_mmvq_b16",
7782 (QT_Q4_0, 2) => "qmatvec_q4_0_mmvq_b2", (QT_Q4_0, 4) => "qmatvec_q4_0_mmvq_b4",
7783 (QT_Q4_0, 8) => "qmatvec_q4_0_mmvq_b8", (QT_Q4_0, 16) => "qmatvec_q4_0_mmvq_b16",
7784 _ => return None,
7785 })
7786 }
7787
7788 pub fn sm_count(&self) -> i32 {
7823 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7824 *SMS.get_or_init(|| {
7825 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7826 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7827 })
7828 }
7829
7830 pub fn batched_variant(&self, _m: usize, in_f: usize, out_f: usize, qtype: i32,
7831 row_bytes: usize, mcols: usize, rp: bool) -> &'static str {
7832 if qtype == QT_Q8_0 {
7837 return if rp { "rp" } else { "base" };
7838 }
7839 static BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7840 let bv = *BV.get_or_init(|| match std::env::var("MEMRA_MMVQ_BV").as_deref() {
7841 Ok("base") => "base", Ok("pf") => "pf", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7842 Ok("pfr2") => "pfr2", Ok("ca") => "ca", Ok("car2") => "car2",
7843 Ok("rp") => "rp", Ok("rpr2") => "rpr2", Ok("rpr2w8") => "rpr2w8",
7846 Ok("rpca") => "rpca", Ok("rpcar2") => "rpcar2",
7849 Ok("rpsc") => "rpsc", Ok("rpms") => "rpms", Ok("rpmsc") => "rpmsc",
7856 Ok("rpks") => "rpks", Ok("rpksc") => "rpksc",
7857 _ => "auto",
7858 });
7859 let ca_ok = qtype == QT_NVFP4 && (row_bytes % 16 == 0) && (in_f % 1024 == 0);
7863 static KS_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7868 let ks_on = *KS_ON.get_or_init(|| std::env::var("MEMRA_KS").as_deref() != Ok("0"));
7869 let sc_ok = ks_on && qtype == QT_NVFP4 && (in_f % 256 == 0) && (in_f / 64 <= 272);
7870 let ks_ok = ks_on && qtype == QT_NVFP4 && (in_f % 512 == 0) && (in_f / 64 <= 272);
7871 static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
7872 let sms = *SMS.get_or_init(|| {
7873 use cudarc::driver::sys::CUdevice_attribute_enum as A;
7874 self.gpu.ctx.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT).unwrap_or(82)
7875 });
7876 let kq_r2 = matches!(qtype, QT_Q4_K | QT_Q5_K | QT_Q6_K);
7896 static KQBV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7899 let kq_bv = *KQBV.get_or_init(|| match std::env::var("MEMRA_KQ_BV").as_deref() {
7900 Ok("base") => "base", Ok("r2") => "r2", Ok("r2w8") => "r2w8",
7901 _ => "auto",
7902 });
7903 let variant: &'static str = if qtype == QT_Q4_0 {
7904 static Q40BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
7908 let q40 = *Q40BV.get_or_init(|| match std::env::var("MEMRA_Q40_BV").as_deref() {
7909 Ok("base") => "base", Ok("r2") => "r2", Ok("ms") => "ms", Ok("sm") => "sm",
7915 Ok("la") => "la", _ => "auto",
7916 });
7917 let v = if q40 != "auto" { q40 }
7918 else if (out_f as u32).div_ceil(8) >= 4 * sms as u32 { "r2" } else { "base" };
7919 if rp { match v { "ms" => "r2ms_rp", "sm" => "r2sm_rp", "la" => "r2la_rp",
7924 "r2" => "r2_rp", _ => "rp" } }
7925 else if matches!(v, "ms" | "sm" | "la") { "r2" } else { v }
7926 } else if qtype != QT_NVFP4 && !kq_r2 {
7927 "base"
7928 } else if kq_r2 && rp {
7929 "rp"
7933 } else if kq_r2 {
7934 if kq_bv != "auto" {
7937 if kq_bv == "r2w8" && mcols != 4 { "r2" } else { kq_bv }
7938 } else if bv != "auto" {
7939 match bv {
7940 "r2" | "pfr2" | "rpr2" | "car2" => "r2",
7941 "r2w8" | "rpr2w8" => if mcols != 4 { "r2" } else { "r2w8" },
7942 _ => "base", }
7944 } else {
7945 let blocks = (out_f + 7) / 8;
7946 let waves = blocks as f64 / (7 * sms as usize) as f64;
7947 let filled = blocks >= 4 * sms as usize;
7948 let use_r2 = if qtype == QT_Q4_K { filled } else { waves >= 2.0 };
7949 if use_r2 { "r2" } else { "base" }
7950 }
7951 } else if bv != "auto" {
7952 let v = if bv == "r2w8" && mcols == 2 { "r2" }
7957 else if bv == "ca" && (!ca_ok || mcols == 8) { "pf" }
7958 else if bv == "car2" && (!ca_ok || mcols == 8) { "r2" }
7959 else if bv == "pfr2" && mcols == 8 { "r2" }
7960 else if (bv == "rpr2w8" || bv == "rpr2") && mcols == 2 { "rpr2" }
7961 else if (bv == "rpca" || bv == "rpcar2") && (!ca_ok || mcols == 8) {
7963 if mcols == 8 { "rpr2w8" } else { "rpr2" }
7964 }
7965 else if bv == "rpcar2" && mcols == 2 { "rpca" }
7966 else if (bv == "rpsc" || bv == "rpmsc") && !sc_ok { "rpr2" }
7969 else if (bv == "rpks" || bv == "rpksc") && !ks_ok { "rpr2" }
7970 else { bv };
7971 if rp {
7972 match v {
7973 "base" | "pf" | "ca" | "rp" => "rp",
7974 "r2" | "pfr2" | "car2" | "rpr2" => "rpr2",
7975 "r2w8" | "rpr2w8" => if mcols == 2 { "rpr2" } else { "rpr2w8" },
7976 other => other, }
7978 } else { v }
7979 } else if mcols == 8 {
7980 if rp { if sc_ok { "rpsc" } else { "rpr2w8" } } else { "r2w8" }
7991 } else if mcols >= 4 {
7992 let blocks = (out_f + 7) / 8;
7996 let r7 = 7 * sms as usize;
7997 let r8 = 8 * sms as usize;
7998 let waves = blocks as f64 / r7 as f64;
7999 let filled = blocks >= 4 * sms as usize;
8000 if filled && blocks.div_ceil(r8) < blocks.div_ceil(r7) {
8004 if rp { "rpr2w8" } else { "r2w8" }
8008 } else if waves >= 2.0 || (waves <= 1.0 && filled) {
8009 if rp { "rpr2" } else { "r2" }
8012 } else {
8013 if rp { "rp" } else { "pf" }
8017 }
8018 } else if in_f >= 6144 {
8019 if rp { "rpr2" } else { "r2" }
8023 }
8024 else if rp {
8025 let waves = ((out_f + 7) / 8) as f64 / (7 * sms as usize) as f64;
8030 if sc_ok && waves >= 0.9 && waves <= 1.1 { "rpsc" } else { "rp" }
8031 } else { "base" };
8032 variant
8033 }
8034
8035 pub fn qmatvec_mmvq_batched(&self, bytes: &CudaSlice<u8>, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
8036 m: usize, in_f: usize, out_f: usize, qtype: i32, row_bytes: usize,
8037 mcols: usize, scale: f32, rp: bool)
8038 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8039 const ROWS_PER_BLOCK: u32 = 4;
8040 let forced: Option<&'static str> = {
8045 static V: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
8046 V.get_or_init(|| std::env::var("MEMRA_BVAR").ok())
8047 .as_deref()
8048 .map(|s| Box::leak(s.to_string().into_boxed_str()) as &'static str)
8049 };
8050 let variant = match forced {
8051 Some(v) if !rp || v.contains("rp") => v,
8052 _ => self.batched_variant(m, in_f, out_f, qtype, row_bytes, mcols, rp),
8053 };
8054 let base_name = Self::batched_kernel_name(qtype, mcols)
8055 .ok_or_else(|| format!("qmatvec_mmvq_batched: no kernel for qtype {qtype} mcols {mcols}"))?;
8056 let variant = if mcols == 16 { if rp { "rp" } else { "base" } } else { variant };
8060 static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8067 let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
8068 if b567 && qtype == QT_NVFP4 && rp && mcols == 8 && (5..=7).contains(&m)
8069 && matches!(variant, "rpsc" | "rpr2w8") {
8070 let f = self.func(&format!("qmatvec_nvfp4_mmvq_b{m}_{variant}"));
8071 let rows_per_block = ROWS_PER_BLOCK * 2; let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8073 let cfg = LaunchConfig {
8074 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
8075 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0 };
8076 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8077 let __s_b = self.gpu.stream();
8078 let mut b = __s_b.launch_builder(&f);
8079 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8080 unsafe { b.launch(cfg)?; }
8081 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8082 return Ok(y);
8083 }
8084 let (name, rows_per_block): (std::borrow::Cow<'static, str>, u32) = match variant {
8085 "base" => (base_name.into(), ROWS_PER_BLOCK),
8086 "pf" => (format!("{base_name}_pf").into(), ROWS_PER_BLOCK),
8087 "ca" => (format!("{base_name}_ca").into(), ROWS_PER_BLOCK),
8088 "rp" => (format!("{base_name}_rp").into(), ROWS_PER_BLOCK),
8089 "rpca" => (format!("{base_name}_rpca").into(), ROWS_PER_BLOCK), "rpks" => (format!("{base_name}_rpks").into(), ROWS_PER_BLOCK),
8093 "rpksc" => (format!("{base_name}_rpksc").into(), ROWS_PER_BLOCK),
8094 "rpms" => (format!("{base_name}_rpms").into(), ROWS_PER_BLOCK),
8095 "rpmsc" => (format!("{base_name}_rpmsc").into(), ROWS_PER_BLOCK),
8096 "r2ms_rp" => (format!("{base_name}_r2ms_rp").into(), ROWS_PER_BLOCK),
8097 "r2sm_rp" => (format!("{base_name}_r2sm_rp").into(), ROWS_PER_BLOCK * 2),
8098 "r2la_rp" => (format!("{base_name}_r2la_rp").into(), ROWS_PER_BLOCK * 2),
8099 v => (format!("{base_name}_{v}").into(), ROWS_PER_BLOCK * 2), };
8101 debug_assert!(!rp || name.contains("_rp"), "rp weight dispatched to a GGUF-layout kernel");
8102 let f = self.func(&name);
8103 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8104 let smem = if name.contains("_r2sm_rp") { (mcols * 32 * 9 * 4 + mcols * 32 * 4) as u32 }
8106 else { 0 };
8107 let cfg = LaunchConfig {
8108 grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
8109 block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: smem };
8110 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8111 let __s_b = self.gpu.stream();
8112 let mut b = __s_b.launch_builder(&f);
8113 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8114 unsafe { b.launch(cfg)?; }
8115 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8116 Ok(y)
8117 }
8118
8119 pub fn qmatvec_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
8123 in_f: usize, out_f: usize, qtype: i32, row_bytes: usize, mcols: usize,
8124 rp: bool)
8125 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8126 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
8127 self.qmatvec_mmvq_batched(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, 1.0, rp)
8128 }
8129
8130 pub fn qmatvec_nvfp4_batched_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize,
8132 in_f: usize, out_f: usize, row_bytes: usize, mcols: usize,
8133 rp: bool)
8134 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8135 self.qmatvec_batched_raw(bytes, x, m, in_f, out_f, QT_NVFP4, row_bytes, mcols, rp)
8136 }
8137
8138 fn try_fp4_gemm(&self, w: &crate::model::GpuTensor, x: &CudaSlice<f32>, m: usize,
8142 in_f: usize, out_f: usize)
8143 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8144 use crate::model::GpuTensor;
8145 if cfg!(memra_portable_cuda) { return Ok(None); }
8146 if std::env::var("MEMRA_FP4").is_err() { return Ok(None); }
8147 #[cfg(memra_cutlass)]
8156 if m >= 128 && std::env::var("MEMRA_FP4_CUTLASS").is_ok() {
8157 if let GpuTensor::Quant { bytes, qtype, scale, row_bytes, cutlass, .. } = w {
8158 if *qtype == QT_NVFP4 && in_f % 64 == 0 {
8159 if let Some(cw) = cutlass {
8160 let y = self.cutlass_fp4_gemm(&cw.b_packed, &cw.sfb_swizzled, x, *scale,
8162 m, out_f, in_f)?;
8163 return Ok(Some(y));
8164 } else if std::env::var("MEMRA_FP4_CUTLASS_OTF").is_ok() {
8165 let (b_packed, sfb_sw) = self.build_cutlass_weight(bytes, out_f, in_f, *row_bytes)?;
8170 let y = self.cutlass_fp4_gemm(&b_packed, &sfb_sw, x, *scale, m, out_f, in_f)?;
8171 return Ok(Some(y));
8172 }
8173 }
8174 }
8175 }
8176 if let GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } = w {
8177 if *qtype == QT_NVFP4 && in_f % 64 == 0 && !*rp {
8180 let y = self.qmatvec_gemm_nvfp4_fp4(bytes, x, m, in_f, out_f, *row_bytes, *scale)?;
8181 return Ok(Some(y));
8182 }
8183 }
8184 Ok(None)
8185 }
8186
8187 pub fn rms_norm_f16out(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>,
8191 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
8192 ncols: usize, nrows: usize, eps: f32)
8193 -> Result<(), Box<dyn std::error::Error>> {
8194 let f = self.func("rms_norm_f16out_f32");
8195 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
8196 let (nc, e) = (ncols as i32, eps);
8197 let __s_b = self.gpu.stream();
8198 let mut b = __s_b.launch_builder(&f);
8199 b.arg(x).arg(w).arg(dst).arg(dst16).arg(&nc).arg(&e);
8200 unsafe { b.launch(cfg)?; }
8201 Ok(())
8202 }
8203
8204 #[allow(clippy::too_many_arguments)]
8207 pub fn add_rms_norm_f16out(&self, a: &CudaSlice<f32>, b: &CudaSlice<f32>, w: &CudaSlice<f32>,
8208 res: &mut CudaSlice<f32>, dst: &mut CudaSlice<f32>,
8209 dst16: &mut CudaSlice<u8>, ncols: usize, nrows: usize, eps: f32)
8210 -> Result<(), Box<dyn std::error::Error>> {
8211 let f = self.func("add_rms_norm_f16out_f32");
8212 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
8213 let (nc, e) = (ncols as i32, eps);
8214 let __s_lb = self.gpu.stream();
8215 let mut lb = __s_lb.launch_builder(&f);
8216 lb.arg(a).arg(b).arg(w).arg(res).arg(dst).arg(dst16).arg(&nc).arg(&e);
8217 unsafe { lb.launch(cfg)?; }
8218 Ok(())
8219 }
8220
8221 pub fn matmul_group_xh(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>,
8224 xh: &CudaSlice<u8>, m: usize)
8225 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8226 let mut out = Vec::with_capacity(ws.len());
8227 let in_f = ws[0].in_features();
8228 for w in ws {
8229 if w.in_features() == in_f && m >= 16 && !self.verify_exact_on() {
8230 if let Some(y) = self.try_f16_gemm_pre(w, xh, m)? {
8231 out.push(y);
8232 continue;
8233 }
8234 }
8235 out.push(self.matmul(w, x, m)?);
8236 }
8237 Ok(out)
8238 }
8239
8240 pub fn gdn_pad_mask(&self, beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
8243 len_d: &CudaSlice<i32>, h: usize, t: usize)
8244 -> Result<(), Box<dyn std::error::Error>> {
8245 let f = self.func("gdn_pad_mask_f32");
8246 let cfg = LaunchConfig::for_num_elems((t * h) as u32);
8247 let (hi, ti) = (h as i32, t as i32);
8248 let __s_b = self.gpu.stream();
8249 let mut b = __s_b.launch_builder(&f);
8250 b.arg(beta).arg(g_log).arg(len_d).arg(&hi).arg(&ti);
8251 unsafe { b.launch(cfg)?; }
8252 Ok(())
8253 }
8254
8255 pub fn row_gather_dev(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
8258 len_d: &CudaSlice<i32>, ncols: usize)
8259 -> Result<(), Box<dyn std::error::Error>> {
8260 let f = self.func("row_gather_dev_f32");
8261 let cfg = LaunchConfig::for_num_elems(ncols as u32);
8262 let nc = ncols as i32;
8263 let __s_b = self.gpu.stream();
8264 let mut b = __s_b.launch_builder(&f);
8265 b.arg(src).arg(dst).arg(len_d).arg(&nc);
8266 unsafe { b.launch(cfg)?; }
8267 Ok(())
8268 }
8269
8270 pub fn matmul_group(&self, ws: &[&crate::model::GpuTensor], x: &CudaSlice<f32>, m: usize)
8277 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8278 use crate::model::GpuTensor;
8279 let mut out = Vec::with_capacity(ws.len());
8280 let any_mirror = ws.iter().any(|w| matches!(w, GpuTensor::Quant { f16: Some(_), .. }));
8281 if m >= 16 && any_mirror && !self.verify_exact_on() {
8282 let in_f = ws[0].in_features();
8283 let xh = self.f16_act(x, m * in_f, in_f)?;
8284 for w in ws {
8285 if w.in_features() == in_f {
8286 if let Some(y) = self.try_f16_gemm_pre(w, &xh, m)? {
8287 out.push(y);
8288 continue;
8289 }
8290 }
8291 out.push(self.matmul(w, x, m)?);
8292 }
8293 return Ok(out);
8294 }
8295 for w in ws {
8296 out.push(self.matmul(w, x, m)?);
8297 }
8298 Ok(out)
8299 }
8300
8301 pub fn matmul_group_multi(&self, ws: &[&crate::model::GpuTensor],
8308 xs: &[&CudaSlice<f32>], ms: &[usize])
8309 -> Result<Vec<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
8310 assert_eq!(xs.len(), ms.len());
8311 let in_f = ws[0].in_features();
8312 let total: usize = ms.iter().sum();
8313 let mut xcat = self.uninit(total * in_f)?;
8314 let mut off = 0usize;
8315 for (x, &m) in xs.iter().zip(ms) {
8316 self.copy_into(&mut xcat, off * in_f, x, m * in_f)?;
8317 off += m;
8318 }
8319 let ys = self.matmul_group(ws, &xcat, total)?;
8320 let mut out: Vec<Vec<CudaSlice<f32>>> = (0..xs.len()).map(|_| Vec::new()).collect();
8321 for (w, y) in ws.iter().zip(ys) {
8322 let out_f = w.out_features();
8323 let mut off = 0usize;
8324 for (s, &m) in ms.iter().enumerate() {
8325 let mut ys_s = self.uninit(m * out_f)?;
8326 let src = y.slice(off * out_f..(off + m) * out_f);
8327 self.gpu.stream().memcpy_dtod(&src, &mut ys_s)?;
8328 out[s].push(ys_s);
8329 off += m;
8330 }
8331 }
8332 Ok(out)
8333 }
8334
8335 pub fn gemm_supports(&self, w: &crate::model::GpuTensor) -> bool {
8345 use crate::model::GpuTensor;
8346 if !legacy_quant_gemm_allowed(
8347 cfg!(memra_portable_cuda),
8348 cfg!(memra_hopper_mma),
8349 std::env::var_os("MEMRA_NO_GEMM").is_some(),
8350 ) {
8351 return false;
8352 }
8353 match w {
8354 GpuTensor::Quant { qtype, .. } =>
8355 matches!(*qtype, QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q4_0)
8356 || (*qtype == QT_NVFP4 && w.in_features() % 64 == 0),
8357 GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
8358 }
8359 }
8360
8361 pub fn qmatvec_gemm(&self, w: &crate::model::GpuTensor, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>,
8368 m: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8369 use crate::model::GpuTensor;
8370 let in_f = w.in_features();
8371 let out_f = w.out_features();
8372 let (bytes, qtype, row_bytes, scale, rp) = match w {
8373 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } => (bytes, *qtype, *row_bytes, *scale, *rp),
8374 _ => unreachable!("gemm_supports guaranteed Quant"),
8375 };
8376 if cfg!(memra_hopper_mma) && qtype == QT_Q8_0 && out_f % 64 == 0 && wgmma_gemm_enabled() {
8382 if let GpuTensor::Quant { rp4: Some(m4), .. } = w {
8383 let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, aq, ad, m, in_f, out_f)?;
8384 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8385 return Ok(y);
8386 }
8387 }
8388 let name = match qtype {
8389 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8390 QT_Q4_0 => if rp { "qmatvec_gemm_q4_0_rp" } else { "qmatvec_gemm_q4_0" },
8391 QT_Q5_K => "qmatvec_gemm_q5_K",
8392 QT_Q6_K => "qmatvec_gemm_q6_K",
8393 QT_NVFP4 => if rp { "qmatvec_gemm_nvfp4_rp" } else { "qmatvec_gemm_nvfp4" },
8394 _ => unreachable!(),
8395 };
8396 let f = self.func(name);
8397 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);
8402 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8404 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8405 let warps: u32 = if is_k1 { k1_tile.2 } else {
8406 match qtype { QT_NVFP4 => 8, _ => 4 }
8407 };
8408 let cfg = LaunchConfig {
8409 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8410 block_dim: (32, warps, 1),
8411 shared_mem_bytes: 0,
8412 };
8413 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8414 let __s_b = self.gpu.stream();
8415 let mut b = __s_b.launch_builder(&f);
8416 b.arg(bytes).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8417 unsafe { b.launch(cfg)?; }
8418 if scale != 1.0 { self.scale_inplace(&mut y, scale, m * out_f)?; }
8419 Ok(y)
8420 }
8421
8422 pub fn qmatvec_gemm_raw(&self, bytes: &CudaSlice<u8>, x: &CudaSlice<f32>, m: usize, in_f: usize,
8427 out_f: usize, qtype: i32, row_bytes: usize)
8428 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8429 let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
8430 let name = match qtype {
8431 QT_Q8_0 => "qmatvec_gemm_q8_0", QT_Q4_K => "qmatvec_gemm_q4_K",
8432 QT_Q4_0 => "qmatvec_gemm_q4_0",
8433 QT_Q5_K => "qmatvec_gemm_q5_K",
8434 QT_Q6_K => "qmatvec_gemm_q6_K", QT_NVFP4 => "qmatvec_gemm_nvfp4",
8435 QT_NVFP4_RP => "qmatvec_gemm_nvfp4_rp",
8436 _ => panic!("qmatvec_gemm_raw: qtype {qtype} has no GEMM kernel"),
8437 };
8438 let f = self.func(name);
8439 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);
8443 let k1_tile = if is_k1 { k1_launch_override().unwrap_or((128, 128, 8)) } else { (128, 128, 8) };
8445 let (bm, bn): (u32, u32) = if is_k1 { (k1_tile.0, k1_tile.1) } else { (64, 256) };
8446 let warps: u32 = if is_k1 { k1_tile.2 } else {
8447 match qtype { QT_NVFP4 | QT_NVFP4_RP => 8, _ => 4 }
8448 };
8449 let cfg = LaunchConfig {
8450 grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
8451 block_dim: (32, warps, 1), shared_mem_bytes: 0,
8452 };
8453 let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
8454 let __s_b = self.gpu.stream();
8455 let mut b = __s_b.launch_builder(&f);
8456 b.arg(bytes).arg(&aq).arg(&ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi).arg(&rb);
8457 unsafe { b.launch(cfg)?; }
8458 Ok(y)
8459 }
8460
8461 pub fn qmatvec_gemm_q8_0_wgmma_raw(&self, rp4: &CudaSlice<u8>, aq: &CudaSlice<i8>,
8468 ad: &CudaSlice<f32>, m: usize, in_f: usize, out_f: usize)
8469 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8470 assert!(out_f % 64 == 0 && in_f % 32 == 0, "wgmma GEMM needs out_f%64==0, in_f%32==0");
8471 let f = self.func("qmatvec_gemm_q8_0_wgmma");
8472 let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
8474 grid_dim: ((out_f / 64) as u32, (m as u32).div_ceil(64), 1),
8475 block_dim: (128, 1, 1), shared_mem_bytes: 0,
8476 };
8477 let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
8478 let __s_b = self.gpu.stream();
8479 let mut b = __s_b.launch_builder(&f);
8480 b.arg(rp4).arg(aq).arg(ad).arg(&mut y).arg(&inf).arg(&outf).arg(&mi);
8481 unsafe { b.launch(cfg)?; }
8482 Ok(y)
8483 }
8484
8485 pub fn scale_inplace(&self, y: &mut CudaSlice<f32>, s: f32, n: usize)
8487 -> Result<(), Box<dyn std::error::Error>> {
8488 let f = self.func("scale_f32");
8489 let cfg = LaunchConfig::for_num_elems(n as u32);
8490 let (sf, ni) = (s, n as i32);
8491 let __s_b = self.gpu.stream();
8492 let mut b = __s_b.launch_builder(&f);
8493 b.arg(y).arg(&sf).arg(&ni);
8494 unsafe { b.launch(cfg)?; }
8495 Ok(())
8496 }
8497
8498 pub fn bf16_to_f32(&self, data: &cudarc::driver::CudaView<'_, u8>, n: usize)
8503 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8504 let mut out = self.alloc_uninit::<f32>(n)?;
8505 let f = self.func("bf16_to_f32");
8506 let cfg = LaunchConfig::for_num_elems(n as u32);
8507 let ni = n as i32;
8508 let __s_b = self.gpu.stream();
8509 let mut b = __s_b.launch_builder(&f);
8510 b.arg(data).arg(&mut out).arg(&ni);
8511 unsafe { b.launch(cfg)?; }
8512 Ok(out)
8513 }
8514
8515 fn linear_bf16_chunked(&self, x: &CudaSlice<f32>, data: &CudaSlice<u8>, m: usize,
8522 in_f: usize, out_f: usize, exact: bool)
8523 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8524 const CHUNK_BYTES: usize = 256 << 20;
8525 let chunk_rows = (CHUNK_BYTES / (in_f * 4)).max(1).min(out_f);
8526 if chunk_rows >= out_f {
8527 let wf32 = self.bf16_to_f32(&data.slice(0..in_f * out_f * 2), in_f * out_f)?;
8528 return if exact { self.linear_decode_exact(x, &wf32, m, in_f, out_f) }
8529 else { self.linear(x, &wf32, m, in_f, out_f) };
8530 }
8531 let mut y = self.alloc_uninit::<f32>(m * out_f)?;
8532 let mut r0 = 0usize;
8533 while r0 < out_f {
8534 let rows = chunk_rows.min(out_f - r0);
8535 let wslice = data.slice(r0 * in_f * 2..(r0 + rows) * in_f * 2);
8536 let wf32 = self.bf16_to_f32(&wslice, in_f * rows)?;
8537 let yc = if exact { self.linear_decode_exact(x, &wf32, m, in_f, rows)? }
8538 else { self.linear(x, &wf32, m, in_f, rows)? };
8539 for mi in 0..m {
8541 let src = yc.slice(mi * rows..(mi + 1) * rows);
8542 let mut dst = y.slice_mut(mi * out_f + r0..mi * out_f + r0 + rows);
8543 self.gpu.stream().memcpy_dtod(&src, &mut dst)?;
8544 }
8545 r0 += rows;
8546 }
8547 Ok(y)
8548 }
8549
8550 pub fn linear_decode_exact(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize,
8557 in_f: usize, out_f: usize)
8558 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8559 if m_tokens == 1 { return self.linear(x, w, 1, in_f, out_f); }
8560 let xv = self.view(x, m_tokens * in_f);
8561 let mut y = self.alloc_uninit::<f32>(m_tokens * out_f)?;
8562 for t in 0..m_tokens {
8563 let row = xv.slice(t * in_f..(t + 1) * in_f);
8564 let mut xr = self.alloc_uninit::<f32>(in_f)?;
8565 self.copy_view_into(&mut xr, 0, &row, in_f)?;
8566 let yr = self.linear(&xr, w, 1, in_f, out_f)?;
8567 self.copy_into(&mut y, t * out_f, &yr, out_f)?;
8568 }
8569 Ok(y)
8570 }
8571
8572 pub fn linear(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, m_tokens: usize, in_f: usize, out_f: usize)
8573 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8574 use cudarc::cublaslt::{Matmul, MatmulConfig};
8575 let mut c = self.alloc_uninit::<f32>(m_tokens * out_f)?; let cfg = MatmulConfig {
8577 transa: true, transb: false, transc: false,
8578 m: out_f as u64, n: m_tokens as u64, k: in_f as u64,
8579 alpha: 1.0, lda: in_f as i64, ldb: in_f as i64, beta: 0.0, ldc: out_f as i64,
8580 stride_a: None, stride_b: None, stride_c: None, stride_bias: None, batch_size: None,
8581 };
8582 unsafe { self.gpu.blas.matmul(cfg, w, x, &mut c, None, None)?; }
8583 Ok(c)
8584 }
8585
8586 pub fn sdpa_naive(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8588 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8589 t: usize, t_kv: usize, scale: f32, causal: bool)
8590 -> Result<(), Box<dyn std::error::Error>> {
8591 let f = self.func("sdpa_naive_f32");
8592 let cfg = LaunchConfig {
8593 grid_dim: (n_head as u32, t as u32, 1),
8594 block_dim: (128, 1, 1),
8595 shared_mem_bytes: (t_kv * 4) as u32,
8596 };
8597 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8598 let __s_b = self.gpu.stream();
8599 let mut b = __s_b.launch_builder(&f);
8600 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8601 unsafe { b.launch(cfg)?; }
8602 Ok(())
8603 }
8604
8605 #[allow(clippy::too_many_arguments)]
8607 pub fn sdpa_naive_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8608 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8609 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8610 -> Result<(), Box<dyn std::error::Error>> {
8611 let f = self.func("sdpa_naive_w_f32");
8612 let cfg = LaunchConfig {
8613 grid_dim: (n_head as u32, t as u32, 1),
8614 block_dim: (128, 1, 1),
8615 shared_mem_bytes: (t_kv * 4) as u32,
8616 };
8617 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
8618 t as i32, t_kv as i32, causal as i32, window as i32);
8619 let __s_b = self.gpu.stream();
8620 let mut b = __s_b.launch_builder(&f);
8621 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
8622 .arg(&scale).arg(&cz).arg(&wi);
8623 unsafe { b.launch(cfg)?; }
8624 Ok(())
8625 }
8626
8627 pub fn sdpa_naive_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<f32>,
8629 v: &cudarc::driver::CudaView<f32>, o: &mut CudaSlice<f32>,
8630 head_dim: usize, n_head: usize, n_head_kv: usize, t: usize, t_kv: usize,
8631 scale: f32, causal: bool) -> Result<(), Box<dyn std::error::Error>> {
8632 let f = self.func("sdpa_naive_f32");
8633 let cfg = LaunchConfig {
8634 grid_dim: (n_head as u32, t as u32, 1), block_dim: (128, 1, 1),
8635 shared_mem_bytes: (t_kv * 4) as u32,
8636 };
8637 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8638 let __s_b = self.gpu.stream();
8639 let mut b = __s_b.launch_builder(&f);
8640 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8641 unsafe { b.launch(cfg)?; }
8642 Ok(())
8643 }
8644
8645 #[allow(clippy::too_many_arguments)]
8653 pub fn fa_dequant_kv_view_f32(&self, k: &cudarc::driver::CudaView<u8>,
8654 v: &cudarc::driver::CudaView<u8>,
8655 kf: &mut CudaSlice<f32>, vf: &mut CudaSlice<f32>,
8656 kv_dim_k: usize, kv_dim_v: usize, t_kv: usize,
8657 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
8658 -> Result<(), Box<dyn std::error::Error>> {
8659 let f = if g { self.func_g("fa_dequant_kv_ws_f32") } else { self.func("fa_dequant_kv_ws_f32") };
8660 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
8661 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8662 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1),
8663 shared_mem_bytes: 0 };
8664 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
8665 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
8666 let __s_b = self.gpu.stream();
8667 let mut b = __s_b.launch_builder(&f);
8668 b.arg(k).arg(v).arg(&mut *kf).arg(&mut *vf).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
8669 unsafe { b.launch(cfg)?; }
8670 Ok(())
8671 }
8672
8673 #[allow(clippy::too_many_arguments)]
8674 pub fn sdpa_naive_quantized_view(
8675 &self,
8676 q: &CudaSlice<f32>,
8677 k: &cudarc::driver::CudaView<u8>,
8678 v: &cudarc::driver::CudaView<u8>,
8679 o: &mut CudaSlice<f32>,
8680 head_dim: usize,
8681 n_head: usize,
8682 n_head_kv: usize,
8683 t: usize,
8684 t_kv: usize,
8685 scale: f32,
8686 causal: bool,
8687 k_tok_bytes: usize,
8688 v_tok_bytes: usize,
8689 ) -> Result<(), Box<dyn std::error::Error>> {
8690 let kv_dim = n_head_kv * head_dim;
8691 let mut kf = self.uninit(t_kv * kv_dim)?;
8692 let mut vf = self.uninit(t_kv * kv_dim)?;
8693 let f = self.func("fa_dequant_kv_ws_f32");
8694 let total = (2 * t_kv * kv_dim) as u64;
8695 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8696 let cfg = LaunchConfig {
8697 grid_dim: (nblk.max(1), 1, 1),
8698 block_dim: (256, 1, 1),
8699 shared_mem_bytes: 0,
8700 };
8701 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8702 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8703 let __s_b = self.gpu.stream();
8704 let mut b = __s_b.launch_builder(&f);
8705 b.arg(k)
8706 .arg(v)
8707 .arg(&mut kf)
8708 .arg(&mut vf)
8709 .arg(&kv_dim_i)
8710 .arg(&kv_dim_i)
8711 .arg(&t_kv_i)
8712 .arg(&k_tok_bytes_i)
8713 .arg(&v_tok_bytes_i);
8714 unsafe { b.launch(cfg)? };
8715 self.sdpa_naive(
8716 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8717 )
8718 }
8719
8720 #[allow(clippy::too_many_arguments)]
8732 pub fn sdpa_naive_w_quantized_view(
8733 &self,
8734 q: &CudaSlice<f32>,
8735 k: &cudarc::driver::CudaView<u8>,
8736 v: &cudarc::driver::CudaView<u8>,
8737 o: &mut CudaSlice<f32>,
8738 head_dim: usize,
8739 n_head: usize,
8740 n_head_kv: usize,
8741 t: usize,
8742 t_kv: usize,
8743 scale: f32,
8744 causal: bool,
8745 window: usize,
8746 k_tok_bytes: usize,
8747 v_tok_bytes: usize,
8748 ) -> Result<(), Box<dyn std::error::Error>> {
8749 let kv_dim = n_head_kv * head_dim;
8750 let mut kf = self.uninit(t_kv * kv_dim)?;
8751 let mut vf = self.uninit(t_kv * kv_dim)?;
8752 let f = self.func("fa_dequant_kv_ws_f32");
8753 let total = (2 * t_kv * kv_dim) as u64;
8754 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
8755 let cfg = LaunchConfig {
8756 grid_dim: (nblk.max(1), 1, 1),
8757 block_dim: (256, 1, 1),
8758 shared_mem_bytes: 0,
8759 };
8760 let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
8761 let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
8762 let __s_b = self.gpu.stream();
8763 let mut b = __s_b.launch_builder(&f);
8764 b.arg(k)
8765 .arg(v)
8766 .arg(&mut kf)
8767 .arg(&mut vf)
8768 .arg(&kv_dim_i)
8769 .arg(&kv_dim_i)
8770 .arg(&t_kv_i)
8771 .arg(&k_tok_bytes_i)
8772 .arg(&v_tok_bytes_i);
8773 unsafe { b.launch(cfg)? };
8774 self.sdpa_naive_w(
8775 q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
8776 )
8777 }
8778
8779 pub fn fa_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8783 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8784 t: usize, t_kv: usize, scale: f32, causal: bool)
8785 -> Result<(), Box<dyn std::error::Error>> {
8786 if portable_mma_gated() {
8787 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
8788 t, t_kv, scale, causal);
8789 }
8790 let fa3_on = head_dim == 256 && causal && t == t_kv
8798 && match std::env::var("MEMRA_FA3").as_deref() {
8799 Ok("0") => false,
8800 Ok("1") => true,
8801 _ => cfg!(memra_hopper_mma),
8802 };
8803 if fa3_on {
8804 let n = t * n_head * head_dim;
8805 let nkv = t * n_head_kv * head_dim;
8806 let mut q16 = self.alloc_u8_uninit(n * 2)?;
8807 let mut k16 = self.alloc_u8_uninit(nkv * 2)?;
8808 let mut v16 = self.alloc_u8_uninit(nkv * 2)?;
8809 self.f32_to_bf16_into(q, &mut q16, n)?;
8810 self.f32_to_bf16_into(k, &mut k16, nkv)?;
8811 self.f32_to_bf16_into(v, &mut v16, nkv)?;
8812 let rc = {
8813 use cudarc::driver::{DevicePtr, DevicePtrMut};
8814 let stream = self.gpu.stream();
8815 let (qp, _g1) = q16.device_ptr(&stream);
8816 let (kp, _g2) = k16.device_ptr(&stream);
8817 let (vp, _g3) = v16.device_ptr(&stream);
8818 let (op, _g4) = o.device_ptr_mut(&stream);
8819 unsafe {
8820 memra_fa3_prefill(qp as *const core::ffi::c_void,
8821 kp as *const core::ffi::c_void,
8822 vp as *const core::ffi::c_void,
8823 op as *mut f32,
8824 t as i32, n_head as i32, n_head_kv as i32,
8825 head_dim as i32, scale,
8826 stream.cu_stream() as *mut core::ffi::c_void)
8827 }
8828 };
8829 if rc != 0 {
8830 return Err(format!("memra_fa3_prefill rc={rc}").into());
8831 }
8832 return Ok(());
8833 }
8834 static FA_P1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8839 let fa_p1 = *FA_P1.get_or_init(|| std::env::var("MEMRA_FA_P1").as_deref() == Ok("1"));
8840 if fa_p1 && head_dim == 256 && !std::env::var("MEMRA_FA_FLOOR").is_ok() {
8841 const BLOCK_Q: usize = 64; const BKX: usize = 32;
8842 let f = self.func("fa_prefill_bf16_p1");
8843 let shmem = (2 * (2 * BKX * head_dim + BLOCK_Q * BKX)
8844 + 4 * (BLOCK_Q * BKX + 2 * BLOCK_Q)) as u32;
8845 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8846 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8847 let cfg = LaunchConfig {
8848 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
8849 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8850 };
8851 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32,
8852 n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8853 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
8854 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
8855 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
8856 let __s_b = self.gpu.stream();
8857 let mut b = __s_b.launch_builder(&f);
8858 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti)
8859 .arg(&tkvi).arg(&scale).arg(&cz);
8860 unsafe { b.launch(cfg)?; }
8861 return Ok(());
8862 }
8863 const BK: usize = 32;
8869 let w2 = std::env::var("MEMRA_FA_PP_W2").as_deref() == Ok("1");
8872 let (block_q, warps, w2_sfx): (usize, u32, &str) =
8873 if w2 { (32, 2, "_w2") } else { (64, 4, "") };
8874 let hd_sfx = fa_hd_suffix(head_dim)?;
8878 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8879 let bf16kv = !floor && !w2
8884 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0");
8885 let (kb16, vb16) = if bf16kv {
8886 let n = t_kv * n_head_kv * head_dim;
8887 let mut kb = self.alloc_u8_uninit(n * 2)?;
8888 let mut vb = self.alloc_u8_uninit(n * 2)?;
8889 let fcv = self.func("f32_to_bf16_bulk");
8890 let ni = n as i64;
8891 let cfgc = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
8892 let __s_b = self.gpu.stream();
8893 let mut b = __s_b.launch_builder(&fcv);
8894 b.arg(k).arg(&mut kb).arg(&ni);
8895 unsafe { b.launch(cfgc)?; }
8896 let __s_b = self.gpu.stream();
8897 let mut b = __s_b.launch_builder(&fcv);
8898 b.arg(v).arg(&mut vb).arg(&ni);
8899 unsafe { b.launch(cfgc)?; }
8900 (Some(kb), Some(vb))
8901 } else {
8902 (None, None)
8903 };
8904 let f = self.func(&if bf16kv {
8905 format!("fa_prefill_bf16kv_pp{hd_sfx}")
8906 } else {
8907 format!("fa_prefill_f32{}{}{hd_sfx}",
8908 if floor { "" } else { "_pp" },
8909 if floor { "" } else { w2_sfx })
8910 });
8911 let kv_stages = if bf16kv { 2 } else { 1 };
8914 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
8915 + 4 * (block_q * BK + 2 * block_q)) as u32;
8916 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8917 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8918 let cfg = LaunchConfig {
8919 grid_dim: ((t as u32 + block_q as u32 - 1) / block_q as u32, n_head as u32, 1),
8920 block_dim: (32, warps, 1), shared_mem_bytes: shmem,
8921 };
8922 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
8923 let __s_b = self.gpu.stream();
8924 let mut b = __s_b.launch_builder(&f);
8925 b.arg(q);
8926 match (&kb16, &vb16) {
8927 (Some(kb), Some(vb)) => { b.arg(kb).arg(vb); }
8928 _ => { b.arg(k).arg(v); }
8929 }
8930 b.arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz);
8931 unsafe { b.launch(cfg)?; }
8932 Ok(())
8933 }
8934
8935 #[allow(clippy::too_many_arguments)]
8939 pub fn fa_prefill_w(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
8940 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, n_head_kv: usize,
8941 t: usize, t_kv: usize, scale: f32, causal: bool, window: usize)
8942 -> Result<(), Box<dyn std::error::Error>> {
8943 if portable_mma_gated() {
8946 return self.sdpa_naive_w(q, k, v, o, head_dim, n_head, n_head_kv,
8947 t, t_kv, scale, causal, window);
8948 }
8949 static FAW_F32: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8953 let faw_f32 = *FAW_F32.get_or_init(|| {
8954 std::env::var("MEMRA_FAW_STAGE").as_deref() == Ok("f32")
8955 });
8956 let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
8957 self.fa_prefill_w_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
8958 window, floor || faw_f32, floor)
8959 }
8960
8961 #[allow(clippy::too_many_arguments)]
8964 pub fn fa_prefill_w_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
8965 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
8966 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
8967 window: usize, v_f16: bool)
8968 -> Result<(), Box<dyn std::error::Error>> {
8969 const BLOCK_Q: usize = 64; const BK: usize = 32;
8970 debug_assert_eq!(head_dim, 256);
8971 let hp = fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
8972 && (n_head / n_head_kv) % 2 == 0;
8973 debug_assert!(!v_f16 || hp, "f16 V emitted but the SWA hp arm is off");
8974 if hp {
8975 const BLOCK_QH: usize = 32;
8976 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
8979 let vh: &CudaSlice<u8> = if v_f16 { vb } else {
8980 let n = t_kv * n_head_kv * head_dim;
8981 if vguard.as_ref().map(|b| b.len() < n * 2).unwrap_or(true) {
8982 *vguard = Some(self.alloc_uninit::<u8>(n * 2)?);
8983 }
8984 self.bf16_to_f16_into(vb, n, vguard.as_mut().unwrap())?;
8985 vguard.as_ref().unwrap()
8986 };
8987 let f = self.func("fa_prefill_w_bf16_p1h2");
8988 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
8989 + 4 * (2 * BLOCK_QH)) as u32;
8990 use cudarc::driver::sys::CUfunction_attribute_enum as A;
8991 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
8992 let cfg = LaunchConfig {
8993 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
8994 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
8995 };
8996 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
8997 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
8998 let __s_b = self.gpu.stream();
8999 let mut b = __s_b.launch_builder(&f);
9000 b.arg(qb).arg(kb).arg(vh).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9001 .arg(&scale).arg(&cz).arg(&wi);
9002 unsafe { b.launch(cfg)?; }
9003 return Ok(());
9004 }
9005 let f = self.func("fa_prefill_w_bf16_p1");
9006 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9007 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9008 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9009 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9010 let cfg = LaunchConfig {
9011 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9012 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9013 };
9014 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9015 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9016 let __s_b = self.gpu.stream();
9017 let mut b = __s_b.launch_builder(&f);
9018 b.arg(qb).arg(kb).arg(vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9019 .arg(&scale).arg(&cz).arg(&wi);
9020 unsafe { b.launch(cfg)?; }
9021 Ok(())
9022 }
9023
9024 #[allow(clippy::too_many_arguments)]
9026 pub fn fa_prefill_w_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9027 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9028 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9029 window: usize, f32_stage: bool, floor: bool)
9030 -> Result<(), Box<dyn std::error::Error>> {
9031 const BLOCK_Q: usize = 64; const BK: usize = 32;
9032 debug_assert_eq!(head_dim, 256, "fa_prefill_w is stamped hd256 only");
9033 static P1_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9037 let p1 = !floor && !f32_stage
9038 && *P1_ON.get_or_init(|| {
9039 std::env::var("MEMRA_FAW_P1").map(|v| v != "0").unwrap_or(true)
9040 });
9041 let hp = p1 && fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0
9042 && (n_head / n_head_kv) % 2 == 0;
9043 if hp {
9044 const BLOCK_QH: usize = 32;
9045 let f = self.func("fa_prefill_w_bf16_p1h2");
9046 let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK)
9047 + 4 * (2 * BLOCK_QH)) as u32;
9048 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9049 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9050 let cfg = LaunchConfig {
9051 grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
9052 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9053 };
9054 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9055 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9056 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9057 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9058 let vh = self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?;
9059 let __s_b = self.gpu.stream();
9060 let mut b = __s_b.launch_builder(&f);
9061 b.arg(&qb).arg(&kb).arg(&vh).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9062 .arg(&scale).arg(&cz).arg(&wi);
9063 unsafe { b.launch(cfg)?; }
9064 return Ok(());
9065 }
9066 if p1 {
9067 let f = self.func("fa_prefill_w_bf16_p1");
9068 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9069 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9070 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9071 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9072 let cfg = LaunchConfig {
9073 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9074 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9075 };
9076 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9077 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9078 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9079 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9080 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9081 let __s_b = self.gpu.stream();
9082 let mut b = __s_b.launch_builder(&f);
9083 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9084 .arg(&scale).arg(&cz).arg(&wi);
9085 unsafe { b.launch(cfg)?; }
9086 return Ok(());
9087 }
9088 static G4_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9091 let g4 = !floor && !f32_stage && n_head_kv == 1 && n_head % 4 == 0
9092 && *G4_ON.get_or_init(|| {
9093 std::env::var("MEMRA_FAW_G4").map(|v| v != "0").unwrap_or(true)
9094 });
9095 if g4 {
9096 const SP_M: usize = 16;
9097 static O2_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9100 let o2 = *O2_ON.get_or_init(|| {
9101 std::env::var("MEMRA_FAW_O2").map(|v| v != "0").unwrap_or(true)
9102 });
9103 let f = self.func(if o2 { "fa_prefill_w_bf16_g4o2" } else { "fa_prefill_w_bf16_g4" });
9104 let shmem = if o2 {
9105 (2 * (4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M)) as u32
9106 } else {
9107 (2 * (2 * BK * head_dim + 4 * SP_M * head_dim + 4 * SP_M * BK)
9108 + 4 * (4 * SP_M)) as u32
9109 };
9110 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9111 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9112 let cfg = LaunchConfig {
9113 grid_dim: ((t as u32).div_ceil(SP_M as u32), (n_head / 4) as u32, 1),
9114 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9115 };
9116 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32,
9117 n_head_kv as i32, t as i32, t_kv as i32, causal as i32, window as i32);
9118 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9119 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9120 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9121 let __s_b = self.gpu.stream();
9122 let mut b = __s_b.launch_builder(&f);
9123 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9124 .arg(&scale).arg(&cz).arg(&wi);
9125 unsafe { b.launch(cfg)?; }
9126 return Ok(());
9127 }
9128 let f = self.func(if floor { "fa_prefill_w_f32" }
9129 else if f32_stage { "fa_prefill_w_f32_pp" }
9130 else { "fa_prefill_w_bf16_pp" });
9131 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9132 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9133 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9134 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9135 let cfg = LaunchConfig {
9136 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9137 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9138 };
9139 let (hd, nh, nhkv, ti, tkvi, cz, wi) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9140 t as i32, t_kv as i32, causal as i32, window as i32);
9141 if f32_stage {
9142 let __s_b = self.gpu.stream();
9143 let mut b = __s_b.launch_builder(&f);
9144 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9145 .arg(&scale).arg(&cz).arg(&wi);
9146 unsafe { b.launch(cfg)?; }
9147 } else {
9148 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9149 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9150 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9151 let __s_b = self.gpu.stream();
9152 let mut b = __s_b.launch_builder(&f);
9153 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9154 .arg(&scale).arg(&cz).arg(&wi);
9155 unsafe { b.launch(cfg)?; }
9156 }
9157 Ok(())
9158 }
9159
9160 #[allow(clippy::too_many_arguments)]
9164 pub fn fa_prefill_hd512(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9165 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9166 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool)
9167 -> Result<(), Box<dyn std::error::Error>> {
9168 if portable_mma_gated() {
9170 return self.sdpa_naive(q, k, v, o, head_dim, n_head, n_head_kv,
9171 t, t_kv, scale, causal);
9172 }
9173 static F32_STAGE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9179 let f32_stage = *F32_STAGE.get_or_init(|| {
9180 std::env::var("MEMRA_FA512_STAGE").as_deref() == Ok("f32")
9181 });
9182 static SP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9186 let sp = !f32_stage
9187 && *SP_ON.get_or_init(|| {
9188 std::env::var("MEMRA_FA512_SP").map(|v| v != "0").unwrap_or(true)
9189 });
9190 self.fa_prefill_hd512_arm(q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale,
9191 causal, f32_stage, sp, sp && fa_f16pv_on())
9192 }
9193
9194 #[allow(clippy::too_many_arguments)]
9196 pub fn fa_prefill_hd512_pre(&self, qb: &CudaSlice<u8>, kb: &CudaSlice<u8>, vb: &CudaSlice<u8>,
9197 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9198 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9199 v_f16: bool)
9200 -> Result<(), Box<dyn std::error::Error>> {
9201 debug_assert_eq!(head_dim, 512);
9202 const SP_M: usize = 16; const BKS: usize = 32;
9203 let f16pv = fa_f16pv_on();
9207 let nw = if f16pv { fa512_wide_warps() } else { 2 };
9208 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
9209 debug_assert!(!v_f16 || f16pv, "f16 V emitted without the door on");
9210 let mut vguard = self.fa_vf16_scratch.lock().unwrap();
9211 let vref: &CudaSlice<u8> = if f16pv && !v_f16 {
9212 let n = t_kv * n_head_kv * head_dim;
9214 let need = n * 2;
9215 if vguard.as_ref().map(|b| b.len() < need).unwrap_or(true) {
9216 *vguard = Some(self.alloc_uninit::<u8>(need)?);
9217 }
9218 let dst = vguard.as_mut().unwrap();
9219 self.bf16_to_f16_into(vb, n, dst)?;
9220 vguard.as_ref().unwrap()
9221 } else { vb };
9222 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
9223 else { match (f16pv, nw) {
9224 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9225 (true, _) => "fa_prefill_bf16_hd512_sp16",
9226 _ => "fa_prefill_bf16_hd512_sp",
9227 } });
9228 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9229 let shmem = if hp {
9231 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9232 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9233 } else {
9234 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9235 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9236 };
9237 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9238 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9239 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9240 let cfg = LaunchConfig {
9241 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9242 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9243 };
9244 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9245 t as i32, t_kv as i32, causal as i32);
9246 let __s_b = self.gpu.stream();
9247 let mut b = __s_b.launch_builder(&f);
9248 b.arg(qb).arg(kb).arg(vref).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9249 .arg(&scale).arg(&cz);
9250 unsafe { b.launch(cfg)?; }
9251 Ok(())
9252 }
9253
9254 #[allow(clippy::too_many_arguments)]
9257 pub fn fa_prefill_hd512_arm(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
9258 o: &mut CudaSlice<f32>, head_dim: usize, n_head: usize,
9259 n_head_kv: usize, t: usize, t_kv: usize, scale: f32, causal: bool,
9260 f32_stage: bool, sp: bool, f16pv: bool)
9261 -> Result<(), Box<dyn std::error::Error>> {
9262 debug_assert_eq!(head_dim, 512, "fa_prefill_hd512 is hd512 only");
9263 if sp && !f32_stage {
9264 const SP_M: usize = 16; const BKS: usize = 32;
9268 let nw = if f16pv { fa512_wide_warps() } else { 2 };
9269 let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
9270 let f = self.func(if hp { "fa_prefill_bf16_hd512_sp16h2" }
9271 else { match (f16pv, nw) {
9272 (true, 4) => "fa_prefill_bf16_hd512_sp16w4",
9273 (true, _) => "fa_prefill_bf16_hd512_sp16",
9274 _ => "fa_prefill_bf16_hd512_sp",
9275 } });
9276 let (nwarp, npart) = if hp { (4usize, 4usize) } else if nw > 2 { (nw, nw) } else { (2, 1) };
9277 let shmem = if hp {
9278 (2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
9279 + 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
9280 } else {
9281 (2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
9282 + 4 * (npart * SP_M * BKS + SP_M)) as u32
9283 };
9284 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9285 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9286 let grid_y = if hp { (n_head / 2) as u32 } else { n_head as u32 };
9287 let cfg = LaunchConfig {
9288 grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
9289 block_dim: (32, nwarp as u32, 1), shared_mem_bytes: shmem,
9290 };
9291 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9292 t as i32, t_kv as i32, causal as i32);
9293 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9294 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9295 let vb = if f16pv { self.f32_to_f16(v, t_kv * n_head_kv * head_dim)? }
9296 else { self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)? };
9297 let __s_b = self.gpu.stream();
9298 let mut b = __s_b.launch_builder(&f);
9299 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9300 .arg(&scale).arg(&cz);
9301 unsafe { b.launch(cfg)?; }
9302 return Ok(());
9303 }
9304 const BLOCK_Q: usize = 32; const BK: usize = 32; const HALF: usize = 256;
9305 let f = self.func(if f32_stage { "fa_prefill_f32_hd512" } else { "fa_prefill_bf16_hd512" });
9306 let shmem = (2 * (BLOCK_Q * head_dim + BK * head_dim + BK * HALF + BLOCK_Q * BK)
9308 + 4 * BLOCK_Q) as u32;
9309 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9310 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9311 let cfg = LaunchConfig {
9312 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 2),
9313 block_dim: (32, 2, 1), shared_mem_bytes: shmem,
9314 };
9315 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32,
9316 t as i32, t_kv as i32, causal as i32);
9317 if f32_stage {
9318 let __s_b = self.gpu.stream();
9319 let mut b = __s_b.launch_builder(&f);
9320 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9321 .arg(&scale).arg(&cz);
9322 unsafe { b.launch(cfg)?; }
9323 } else {
9324 let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
9325 let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
9326 let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
9327 let __s_b = self.gpu.stream();
9328 let mut b = __s_b.launch_builder(&f);
9329 b.arg(&qb).arg(&kb).arg(&vb).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi)
9330 .arg(&scale).arg(&cz);
9331 unsafe { b.launch(cfg)?; }
9332 }
9333 Ok(())
9334 }
9335
9336 #[allow(clippy::too_many_arguments)]
9340 pub fn rope_neox2_bf16e(&self, q: &mut CudaSlice<f32>, k: &mut CudaSlice<f32>,
9341 qb: &mut CudaSlice<u8>, kb: &mut CudaSlice<u8>,
9342 pos: &CudaSlice<i32>, head_dim: usize, n_dims: usize,
9343 nh_q: usize, nh_k: usize, n_tokens: usize, base: f32,
9344 freq_scale: f32, ff: Option<&CudaSlice<f32>>)
9345 -> Result<(), Box<dyn std::error::Error>> {
9346 let f = self.func("rope_neox2_bf16e_f32");
9347 let rows = ((nh_q + nh_k) * n_tokens) as u32;
9348 let cfg = LaunchConfig { grid_dim: (rows, 1, 1),
9349 block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
9350 let theta_scale = base.powf(-2.0 / n_dims as f32);
9351 let (hd, nd, nhq, nhk, nt) = (head_dim as i32, n_dims as i32, nh_q as i32,
9352 nh_k as i32, n_tokens as i32);
9353 let __s_b = self.gpu.stream();
9354 let mut b = __s_b.launch_builder(&f);
9355 match ff {
9356 Some(t) => { b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9357 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9358 .arg(&theta_scale).arg(&freq_scale).arg(t);
9359 unsafe { b.launch(cfg)?; } }
9360 None => { let null: u64 = 0;
9361 b.arg(&mut *q).arg(&mut *k).arg(&mut *qb).arg(&mut *kb).arg(pos)
9362 .arg(&hd).arg(&nd).arg(&nhq).arg(&nhk).arg(&nt)
9363 .arg(&theta_scale).arg(&freq_scale).arg(&null);
9364 unsafe { b.launch(cfg)?; } }
9365 }
9366 Ok(())
9367 }
9368
9369 pub fn f32_to_bf16(&self, x: &CudaSlice<f32>, n: usize)
9372 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9373 assert!(n % 4 == 0, "f32_to_bf16 requires n % 4 == 0, got {n}");
9374 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9375 let f = self.func("f32_to_bf16_flat");
9376 let n_i = n as i64;
9377 let cfg = LaunchConfig {
9378 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9379 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9380 };
9381 let __s_b = self.gpu.stream();
9382 let mut b = __s_b.launch_builder(&f);
9383 b.arg(x).arg(&mut y).arg(&n_i);
9384 unsafe { b.launch(cfg)?; }
9385 Ok(y)
9386 }
9387
9388 pub fn f32_to_f16(&self, x: &CudaSlice<f32>, n: usize)
9389 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9390 assert!(n % 4 == 0, "f32_to_f16 requires n % 4 == 0, got {n}");
9391 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9392 let f = self.func("f32_to_f16_flat");
9393 let n_i = n as i64;
9394 let cfg = LaunchConfig {
9395 grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
9396 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9397 };
9398 let __s_b = self.gpu.stream();
9399 let mut b = __s_b.launch_builder(&f);
9400 b.arg(x).arg(&mut y).arg(&n_i);
9401 unsafe { b.launch(cfg)?; }
9402 Ok(y)
9403 }
9404
9405 pub fn bf16_to_f16(&self, xb: &CudaSlice<u8>, n: usize)
9407 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9408 let mut y = self.alloc_uninit::<u8>(n * 2)?;
9409 self.bf16_to_f16_into(xb, n, &mut y)?;
9410 Ok(y)
9411 }
9412
9413 pub fn bf16_to_f16_into(&self, xb: &CudaSlice<u8>, n: usize, y: &mut CudaSlice<u8>)
9415 -> Result<(), Box<dyn std::error::Error>> {
9416 assert!(n % 2 == 0, "bf16_to_f16 requires n % 2 == 0, got {n}");
9417 assert!(y.len() >= n * 2);
9418 let f = self.func("bf16_to_f16_flat");
9419 let n2 = (n / 2) as i64;
9420 let cfg = LaunchConfig {
9421 grid_dim: (((n / 2) as u32).div_ceil(256), 1, 1),
9422 block_dim: (256, 1, 1), shared_mem_bytes: 0,
9423 };
9424 let __s_b = self.gpu.stream();
9425 let mut b = __s_b.launch_builder(&f);
9426 b.arg(xb).arg(y).arg(&n2);
9427 unsafe { b.launch(cfg)?; }
9428 Ok(())
9429 }
9430
9431 #[allow(clippy::too_many_arguments)]
9436 pub fn fa_prefill_vl8(&self, seqs: &[FaSeqVl], head_dim: usize, n_head: usize,
9437 n_head_kv: usize, scale: f32)
9438 -> Result<(), Box<dyn std::error::Error>> {
9439 const BK: usize = 32;
9440 let b = seqs.len();
9441 assert!(b >= 1 && b <= 8);
9442 let mut packed = [FaSeqVl::default(); 8];
9443 packed[..b].copy_from_slice(seqs);
9444 let v = FaVl8(packed);
9445 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9446 let ept = (n_head_kv * head_dim) as i32;
9447 {
9448 let f = self.func("fa_mirror_vl");
9449 let max_n = (max_t as i64) * ept as i64;
9450 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
9451 for which in 0..2i32 {
9452 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9453 let __s_lb = self.gpu.stream();
9454 let mut lb = __s_lb.launch_builder(&f);
9455 lb.arg(&v).arg(&ept).arg(&which);
9456 unsafe { lb.launch(cfg)?; }
9457 }
9458 }
9459 let hd_sfx = fa_hd_suffix(head_dim)?;
9460 let f = self.func(&format!("fa_prefill_bf16kv_vl{hd_sfx}"));
9461 let block_q = 64usize;
9462 let kv_stages = 2usize;
9463 let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
9464 + 4 * (block_q * BK + 2 * block_q)) as u32;
9465 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9466 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9467 let cfg = LaunchConfig {
9468 grid_dim: (max_t.div_ceil(block_q as u32), n_head as u32, b as u32),
9469 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9470 };
9471 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9472 let __s_lb = self.gpu.stream();
9473 let mut lb = __s_lb.launch_builder(&f);
9474 lb.arg(&v).arg(&hd).arg(&nh).arg(&nhkv).arg(&scale);
9475 unsafe { lb.launch(cfg)?; }
9476 Ok(())
9477 }
9478
9479 #[allow(clippy::too_many_arguments)]
9483 pub fn attn_pre_vl8(&self, seqs: &[AttnPreVl], wq: &CudaSlice<f32>, wk: &CudaSlice<f32>,
9484 head_dim: usize, rope_dims: usize, n_head: usize, n_head_kv: usize,
9485 eps: f32, freq_base: f32, freq_scale: f32,
9486 kv_dim_k: usize, kv_dim_v: usize,
9487 k_tok_bytes: usize, v_tok_bytes: usize)
9488 -> Result<(), Box<dyn std::error::Error>> {
9489 let b = seqs.len();
9490 assert!(b >= 1 && b <= 8);
9491 let mut packed = [AttnPreVl::default(); 8];
9492 packed[..b].copy_from_slice(seqs);
9493 let v = AttnPreVl8(packed);
9494 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
9495 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
9496 {
9497 let f = self.func("q_gate_split_vl");
9498 let n = max_t * (n_head * head_dim) as u32;
9499 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9500 let __s_lb = self.gpu.stream();
9501 let mut lb = __s_lb.launch_builder(&f);
9502 lb.arg(&v).arg(&hd).arg(&nh);
9503 unsafe { lb.launch(cfg)?; }
9504 }
9505 {
9506 let f = self.func("attn_rms_vl");
9507 let cfg = LaunchConfig { grid_dim: (max_t * n_head as u32, 2, b as u32), block_dim: (rms_block(), 1, 1), shared_mem_bytes: 0 };
9508 let __s_lb = self.gpu.stream();
9509 let mut lb = __s_lb.launch_builder(&f);
9510 lb.arg(&v).arg(wq).arg(wk).arg(&hd).arg(&nh).arg(&nhkv).arg(&eps);
9511 unsafe { lb.launch(cfg)?; }
9512 }
9513 {
9514 let f = self.func("attn_rope_vl");
9515 let theta_scale = freq_base.powf(-2.0 / rope_dims as f32);
9516 let nd = rope_dims as i32;
9517 let cfg = LaunchConfig { grid_dim: (max_t * n_head as u32, 2, b as u32), block_dim: ((head_dim / 2) as u32, 1, 1), shared_mem_bytes: 0 };
9518 let __s_lb = self.gpu.stream();
9519 let mut lb = __s_lb.launch_builder(&f);
9520 lb.arg(&v).arg(&hd).arg(&nd).arg(&nh).arg(&nhkv).arg(&theta_scale).arg(&freq_scale);
9521 unsafe { lb.launch(cfg)?; }
9522 }
9523 {
9524 let f = self.func("append_kv_vl");
9525 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
9526 let cfg = LaunchConfig { grid_dim: (nblk, max_t, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
9527 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9528 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9529 let __s_lb = self.gpu.stream();
9530 let mut lb = __s_lb.launch_builder(&f);
9531 lb.arg(&v).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
9532 unsafe { lb.launch(cfg)?; }
9533 }
9534 Ok(())
9535 }
9536
9537 pub fn fa_prefill_view(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9542 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9543 head_dim: usize, n_head: usize, n_head_kv: usize,
9544 t: usize, t_kv: usize, scale: f32, causal: bool,
9545 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9546 -> Result<(), Box<dyn std::error::Error>> {
9547 if portable_mma_gated() {
9548 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9549 t, t_kv, scale, causal,
9550 k_tok_bytes, v_tok_bytes);
9551 }
9552 const BLOCK_Q: usize = 64; const BK: usize = 32;
9553 let name = format!("fa_prefill_q{}", fa_hd_suffix(head_dim)?);
9556 let f = if g { self.func_g(&name) } else { self.func(&name) };
9557 let shmem = (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9558 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
9559 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9560 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9561 let cfg = LaunchConfig {
9562 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9563 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9564 };
9565 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
9566 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9567 let __s_b = self.gpu.stream();
9568 let mut b = __s_b.launch_builder(&f);
9569 b.arg(q).arg(k).arg(v).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9570 .arg(&ktb).arg(&vtb);
9571 unsafe { b.launch(cfg)?; }
9572 Ok(())
9573 }
9574
9575 #[allow(clippy::too_many_arguments)]
9585 pub fn fa_prefill_view_ws(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9586 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9587 head_dim: usize, n_head: usize, n_head_kv: usize,
9588 t: usize, t_kv: usize, scale: f32, causal: bool,
9589 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9590 -> Result<(), Box<dyn std::error::Error>> {
9591 if portable_mma_gated() {
9592 return self.sdpa_naive_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9593 t, t_kv, scale, causal,
9594 k_tok_bytes, v_tok_bytes);
9595 }
9596 const BLOCK_Q: usize = 64; const BK: usize = 32;
9597 let kv_dim_k = n_head_kv * head_dim;
9598 let kv_dim_v = n_head_kv * head_dim;
9599 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9601 let mut guard = self.prime_deqw_ws.lock().unwrap();
9603 let need_grow = match guard.as_ref() {
9604 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9605 None => true,
9606 };
9607 if need_grow {
9608 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9609 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9610 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9611 }
9612 let (kw, vw) = guard.as_mut().unwrap();
9613 {
9615 let f = if g { self.func_g("fa_dequant_kv_ws_bf16") } else { self.func("fa_dequant_kv_ws_bf16") };
9617 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9618 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9619 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9620 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9621 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9622 let __s_b = self.gpu.stream();
9623 let mut b = __s_b.launch_builder(&f);
9624 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9625 unsafe { b.launch(cfg)?; }
9626 }
9627 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9635 {
9636 let hd_sfx = fa_hd_suffix(head_dim)?;
9637 let f = self.func(&format!("fa_prefill_qw{}{hd_sfx}", if db { "_db" } else { "" }));
9638 let shmem = if db {
9639 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9641 } else {
9642 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9643 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9644 };
9645 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9646 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9647 let cfg = LaunchConfig {
9648 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9649 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9650 };
9651 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
9652 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
9653 let __s_b = self.gpu.stream();
9654 let mut b = __s_b.launch_builder(&f);
9655 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9656 .arg(&kdk).arg(&kdv);
9657 unsafe { b.launch(cfg)?; }
9658 }
9659 Ok(())
9660 }
9661
9662 #[allow(clippy::too_many_arguments)]
9678 pub fn fa_prefill_view_ws_w_hd128(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9679 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9680 head_dim: usize, n_head: usize, n_head_kv: usize,
9681 t: usize, t_kv: usize, scale: f32, causal: bool,
9682 window: usize, k_tok_bytes: usize, v_tok_bytes: usize)
9683 -> Result<(), Box<dyn std::error::Error>> {
9684 assert_eq!(head_dim, 128, "fa_prefill_view_ws_w_hd128: only the hd128 twin is stamped");
9685 if portable_mma_gated() {
9686 return self.sdpa_naive_w_quantized_view(q, k, v, o, head_dim, n_head, n_head_kv,
9687 t, t_kv, scale, causal, window,
9688 k_tok_bytes, v_tok_bytes);
9689 }
9690 const BLOCK_Q: usize = 64; const BK: usize = 32;
9691 let kv_dim_k = n_head_kv * head_dim;
9692 let kv_dim_v = n_head_kv * head_dim;
9693 let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
9695 let mut guard = self.prime_deqw_ws.lock().unwrap();
9696 let need_grow = match guard.as_ref() {
9697 Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
9698 None => true,
9699 };
9700 if need_grow {
9701 let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
9702 let (ck, cv) = guard.as_ref().map(|(a, b)| (a.len(), b.len())).unwrap_or((0, 0));
9703 *guard = Some((self.alloc_u8(grow(ck, k_ws_bytes))?, self.alloc_u8(grow(cv, v_ws_bytes))?));
9704 }
9705 let (kw, vw) = guard.as_mut().unwrap();
9706 {
9709 let f = self.func("fa_dequant_kv_ws_bf16");
9710 let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
9711 let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
9712 let cfg = LaunchConfig { grid_dim: (nblk.max(1), 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
9713 let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
9714 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9715 let __s_b = self.gpu.stream();
9716 let mut b = __s_b.launch_builder(&f);
9717 b.arg(k).arg(v).arg(&mut *kw).arg(&mut *vw).arg(&kdk).arg(&kdv).arg(&tkvi).arg(&ktb).arg(&vtb);
9718 unsafe { b.launch(cfg)?; }
9719 }
9720 let db = std::env::var("MEMRA_PRIME_DEQW_DB").map(|v| v != "0").unwrap_or(true);
9722 {
9723 let f = self.func(if db { "fa_prefill_qw_db_w_hd128" } else { "fa_prefill_qw_w_hd128" });
9724 let shmem = if db {
9725 (2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
9726 } else {
9727 (2 * (2 * BK * head_dim + BLOCK_Q * BK)
9728 + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
9729 };
9730 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9731 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9732 let cfg = LaunchConfig {
9733 grid_dim: ((t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32, n_head as u32, 1),
9734 block_dim: (32, 4, 1), shared_mem_bytes: shmem,
9735 };
9736 let (hd, nh, nhkv, ti, tkvi, cz) = (head_dim as i32, n_head as i32, n_head_kv as i32, t as i32, t_kv as i32, causal as i32);
9737 let (kdk, kdv, wnd) = (kv_dim_k as i32, kv_dim_v as i32, window as i32);
9738 let __s_b = self.gpu.stream();
9739 let mut b = __s_b.launch_builder(&f);
9740 b.arg(q).arg(&*kw).arg(&*vw).arg(o).arg(&hd).arg(&nh).arg(&nhkv).arg(&ti).arg(&tkvi).arg(&scale).arg(&cz)
9741 .arg(&kdk).arg(&kdv).arg(&wnd);
9742 unsafe { b.launch(cfg)?; }
9743 }
9744 Ok(())
9745 }
9746
9747 pub fn fa_decode(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9751 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9752 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9753 k_tok_bytes: usize, v_tok_bytes: usize)
9754 -> Result<(), Box<dyn std::error::Error>> {
9755 self.fa_decode_kvmod(q, k, v, o, head_dim, n_head, n_head_kv, t_kv, scale,
9756 k_tok_bytes, v_tok_bytes, false)
9757 }
9758
9759 #[allow(clippy::too_many_arguments)]
9763 #[allow(clippy::too_many_arguments)]
9767 #[allow(clippy::too_many_arguments)]
9768 fn fa_decode_scalar_unified(&self, q: &cudarc::driver::CudaView<f32>,
9769 k: &cudarc::driver::CudaView<u8>,
9770 v: &cudarc::driver::CudaView<u8>,
9771 o: &mut cudarc::driver::CudaViewMut<f32>,
9772 head_dim: usize, n_head: usize, n_head_kv: usize,
9773 t_kv_host: usize, t_kv_dev: Option<&CudaSlice<i32>>,
9774 scale: f32, n_splits: usize, split_keys: usize,
9775 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
9776 part_o: &mut CudaSlice<f32>, part_m: &mut CudaSlice<f32>,
9777 part_l: &mut CudaSlice<f32>,
9778 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
9779 -> Result<(), Box<dyn std::error::Error>> {
9780 let f = if g { self.func_g("fa_decode_f32") } else { self.fa_func("fa_decode_f32", head_dim) };
9781 let cfg = LaunchConfig { grid_dim: (n_head as u32, n_splits as u32, 1),
9782 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: (4 * (head_dim + 32)) as u32 };
9783 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
9784 let (ktb, vtb, tkvi, ski) = (k_tok_bytes as i64, v_tok_bytes as i64, t_kv_host as i32,
9785 split_keys as i32);
9786 let __s_b = self.gpu.stream();
9787 let mut b = __s_b.launch_builder(&f);
9788 match t_kv_dev {
9789 Some(d) => { b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9790 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(d).arg(&scale).arg(&nsp)
9791 .arg(&ski).arg(&ktb).arg(&vtb);
9792 unsafe { b.launch(cfg)?; } }
9793 None => { let null: u64 = 0;
9794 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
9795 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&null).arg(&scale).arg(&nsp)
9796 .arg(&ski).arg(&ktb).arg(&vtb);
9797 unsafe { b.launch(cfg)?; } }
9798 }
9799 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1),
9800 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
9801 if let Some((oq, od)) = q8_out {
9802 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
9804 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
9805 let __s_b2 = self.gpu.stream();
9806 let mut b2 = __s_b2.launch_builder(&fc);
9807 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
9808 unsafe { b2.launch(cfg2)?; }
9809 return Ok(());
9810 }
9811 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
9812 let __s_b2 = self.gpu.stream();
9813 let mut b2 = __s_b2.launch_builder(&fc);
9814 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
9815 unsafe { b2.launch(cfg2)?; }
9816 Ok(())
9817 }
9818
9819 pub fn fa_decode_kvmod(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
9820 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
9821 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9822 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9823 -> Result<(), Box<dyn std::error::Error>> {
9824 let q_view = q.as_view();
9825 let mut o_view = o.as_view_mut();
9826 self.fa_decode_kvmod_view(&q_view, k, v, &mut o_view, head_dim, n_head, n_head_kv,
9827 t_kv, scale, k_tok_bytes, v_tok_bytes, g)
9828 }
9829
9830 #[allow(clippy::too_many_arguments)]
9835 pub fn fa_decode_kvmod_view(&self, q: &cudarc::driver::CudaView<f32>,
9836 k: &cudarc::driver::CudaView<u8>, v: &cudarc::driver::CudaView<u8>,
9837 o: &mut cudarc::driver::CudaViewMut<f32>,
9838 head_dim: usize, n_head: usize, n_head_kv: usize, t_kv: usize, scale: f32,
9839 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
9840 -> Result<(), Box<dyn std::error::Error>> {
9841 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
9862 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
9866 let sp = fa_split_keys(t_kv, n_head_kv);
9867 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
9868 let o_len = n_head * n_splits * head_dim;
9869 let ml_len = n_head * n_splits;
9870 let mut part_guard = self.fa_part_pool.lock().unwrap();
9871 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
9872 let old = part_guard.take();
9883 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
9884 if let Some(old) = old {
9885 self.fa_part_retired.lock().unwrap().push(old);
9886 }
9887 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
9888 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
9889 }
9890 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
9891 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
9892 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
9893 }
9894 let pg = part_guard.as_mut().unwrap();
9895 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
9896 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
9897 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
9898 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
9899 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
9900 let (hd, nh, nhkv, tkvi, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, t_kv as i32, n_splits as i32);
9901 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
9902 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
9906 let fa512_min = fa512_min_tkv();
9911 let deep = fa_vec && head_dim == 256 && fa_v4_at(t_kv) && !g
9914 && fa_deep_at(t_kv) && !matches!(fa_v4_mode(), "noB3" | "stage");
9915 let (f, cfg) = if fa_vec && head_dim == 512 && t_kv >= fa512_min {
9916 let gqa = (n_head / n_head_kv).max(1) as u32;
9919 let fv = self.fa_func("fa_decode_vec_q_dpl16", head_dim);
9920 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9921 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
9922 } else if fa_vec && head_dim <= 256 {
9923 let gqa = (n_head / n_head_kv).max(1) as u32;
9924 static SMEM_TKV: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9935 let smem_tkv = *SMEM_TKV.get_or_init(|| {
9936 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
9937 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
9938 });
9939 if fa_v4_at(t_kv) && head_dim == 256 {
9940 let v4name = match fa_v4_mode() {
9944 "noB3" => "fa_decode_vec_q_v4_noB3", "stage" => "fa_decode_vec_q_v4_stage", _ if deep => "fa_decode_vec_q_v4_deep",
9947 _ => "fa_decode_vec_q_v4",
9948 };
9949 let fv = if g { self.func_g(v4name) } else { self.func(v4name) };
9950 let shmem = (if deep { 12160 } else { 11520 }
9953 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
9954 use cudarc::driver::sys::CUfunction_attribute_enum as A;
9955 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9956 (fv,
9957 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9958 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9959 } else if fa_v3_active(head_dim) {
9960 let fv = if g { self.func_g("fa_decode_vec_q_v3") } else { self.func("fa_decode_vec_q_v3") };
9963 let shmem = (32 * head_dim * 2) as u32; (fv,
9965 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9966 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9967 } else if fa_v2_on() {
9968 let fv = if g { self.func_g("fa_decode_vec_q_v2") } else { self.func("fa_decode_vec_q_v2") };
9972 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
9974 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9975 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9976 } else if smem_tkv > 0 && t_kv >= smem_tkv && !g
9977 && !(head_dim == 512 && Self::gkv_on()) {
9978 let fv = if g { self.func_g("fa_decode_vec_q_smem") } else { self.func("fa_decode_vec_q_smem") };
9982 let shmem = (2 * 32 * head_dim * 2) as u32; use cudarc::driver::sys::CUfunction_attribute_enum as A;
9984 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
9985 (fv,
9986 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9987 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
9988 } else {
9989 let fv = if g { self.func_g("fa_decode_vec_q") } else { self.func("fa_decode_vec_q") };
9992 (fv,
9993 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
9994 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
9995 }
9996 } else {
9997 return self.fa_decode_scalar_unified(q, k, v, o, head_dim, n_head, n_head_kv,
10000 t_kv, None, scale, n_splits,
10001 if fa_vec { sp } else { 256 },
10002 k_tok_bytes, v_tok_bytes, g,
10003 part_o, part_m, part_l, None);
10004 };
10005 let __s_b = self.gpu.stream();
10006 let mut b = __s_b.launch_builder(&f);
10007 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10008 .arg(&hd).arg(&nh).arg(&nhkv).arg(&tkvi).arg(&scale).arg(&nsp).arg(&ktb).arg(&vtb);
10009 unsafe { b.launch(cfg)?; }
10010 let (fc, cfg2) = (if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) },
10013 LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 });
10014 let __s_b2 = self.gpu.stream();
10015 let mut b2 = __s_b2.launch_builder(&fc);
10016 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
10017 unsafe { b2.launch(cfg2)?; }
10018 Ok(())
10019 }
10020
10021 #[allow(clippy::too_many_arguments)]
10032 pub fn fa_decode_batch_seqs_v4(&self, q: &CudaSlice<f32>,
10033 kv_ptrs: &cudarc::driver::CudaView<u64>,
10034 pos_seq: &CudaSlice<i32>, o: &mut CudaSlice<f32>,
10035 head_dim: usize, n_head: usize, n_head_kv: usize,
10036 b_n: usize, t_kv_max: usize, scale: f32,
10037 split_keys: usize, k_tok_bytes: usize, v_tok_bytes: usize)
10038 -> Result<(), Box<dyn std::error::Error>> {
10039 debug_assert!(head_dim == 256, "seqs twin is v4-stamped (hd256 only)");
10040 let n_splits_max = (t_kv_max + split_keys - 1) / split_keys;
10041 let o_len = b_n * n_head * n_splits_max * head_dim;
10042 let ml_len = b_n * n_head * n_splits_max;
10043 let mut part_guard = self.fa_part_pool.lock().unwrap();
10044 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10045 let old = part_guard.take();
10056 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10057 if let Some(old) = old {
10058 self.fa_part_retired.lock().unwrap().push(old);
10059 }
10060 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10061 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10062 }
10063 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10064 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10065 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10066 }
10067 let pg = part_guard.as_mut().unwrap();
10068 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10069 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10070 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10071 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10072 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10073 let (nspm, spk) = (n_splits_max as i32, split_keys as i32);
10074 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10075 let gqa = (n_head / n_head_kv).max(1) as u32;
10076 let f = self.func("fa_decode_vec_q_seqs_v4");
10077 let shmem = (11520 + 32 * head_dim * 2) as u32;
10079 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10080 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
10081 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, b_n as u32),
10082 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10083 {
10084 let __s_b = self.gpu.stream();
10085 let mut b = __s_b.launch_builder(&f);
10086 b.arg(q).arg(kv_ptrs).arg(pos_seq).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10087 .arg(&hd).arg(&nh).arg(&nhkv).arg(&scale).arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
10088 unsafe { b.launch(cfg)?; }
10089 }
10090 let fc = self.func("fa_decode_combine_seqs");
10091 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, b_n as u32, 1),
10092 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10093 let __s_b2 = self.gpu.stream();
10094 let mut b2 = __s_b2.launch_builder(&fc);
10095 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10096 .arg(pos_seq).arg(&nspm).arg(&spk);
10097 unsafe { b2.launch(cfg2)?; }
10098 Ok(())
10099 }
10100
10101 #[allow(clippy::too_many_arguments)]
10108 pub fn append_kv_quantized_seqs(&self, k_rows: &CudaSlice<f32>, v_rows: &CudaSlice<f32>,
10109 kv_ptrs: &cudarc::driver::CudaView<u64>,
10110 pos_seq: &CudaSlice<i32>, b_n: usize,
10111 kv_dim_k: usize, kv_dim_v: usize,
10112 k_tok_bytes: usize, v_tok_bytes: usize)
10113 -> Result<(), Box<dyn std::error::Error>> {
10114 let f = self.func("append_quantize_kv_q8_0_q5_1_seqs");
10115 let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
10116 let cfg = LaunchConfig { grid_dim: (nblk, b_n as u32, 1),
10117 block_dim: (32, 1, 1), shared_mem_bytes: 0 };
10118 let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
10119 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10120 let __s_b = self.gpu.stream();
10121 let mut b = __s_b.launch_builder(&f);
10122 b.arg(k_rows).arg(v_rows).arg(kv_ptrs).arg(pos_seq)
10123 .arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
10124 unsafe { b.launch(cfg)?; }
10125 Ok(())
10126 }
10127
10128 pub fn fa_rows_eligible(&self, base_len: usize, head_dim: usize) -> bool {
10134 std::env::var("MEMRA_NO_FA_VEC").is_err()
10135 && std::env::var("MEMRA_FA_ROWS_OFF").is_err()
10136 && base_len + 1 >= fa_vec_min_tkv()
10137 && head_dim <= 256 && head_dim % 32 == 0
10138 }
10139
10140 #[allow(clippy::too_many_arguments)]
10149 pub fn fa_decode_rows(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10150 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10151 head_dim: usize, n_head: usize, n_head_kv: usize,
10152 base_len: usize, t: usize, scale: f32,
10153 k_tok_bytes: usize, v_tok_bytes: usize,
10154 base_dev: Option<(&CudaSlice<i32>, i32)>,
10158 kv_shared: bool,
10161 g: bool,
10165 mut q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10168 -> Result<(), Box<dyn std::error::Error>> {
10169 debug_assert!(base_len + 1 >= fa_vec_min_tkv() && head_dim <= 512 && head_dim % 32 == 0);
10170 let t_kv_max = base_len + t; let mut sp = fa_split_keys(t_kv_max, n_head_kv); if head_dim == 512 {
10177 static SP512: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10178 let v = *SP512.get_or_init(|| std::env::var("MEMRA_FA_SP512").ok()
10181 .and_then(|x| x.parse().ok()).unwrap_or(0));
10182 sp = if v >= 8 { v } else { FA_SP512_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) };
10183 }
10184 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10185 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10186 let gqa = (n_head / n_head_kv).max(1) as u32;
10187 let mut groups: Vec<(usize, usize, usize)> = Vec::new(); if head_dim == 512 || fa_split_keys(base_len + 1, n_head_kv) == sp {
10198 groups.push((0, t, sp));
10199 } else {
10200 let mut r0 = 0usize;
10201 while r0 < t {
10202 let sp_g = fa_split_keys(base_len + r0 + 1, n_head_kv);
10203 let mut r1 = r0 + 1;
10204 while r1 < t && fa_split_keys(base_len + r1 + 1, n_head_kv) == sp_g { r1 += 1; }
10205 groups.push((r0, r1 - r0, sp_g));
10206 r0 = r1;
10207 }
10208 }
10209 static SMEM_TKV_R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10213 let smem_tkv = *SMEM_TKV_R.get_or_init(|| {
10214 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
10215 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
10216 });
10217 let v4 = fa_v4_at(base_len + t) && head_dim == 256;
10218 let v3 = fa_v3_active(head_dim);
10219 let smem_rows = head_dim <= 256 && !v3 && !fa_v2_on() && smem_tkv > 0 && t_kv_max >= smem_tkv;
10220 let _ = kv_shared;
10225 let i2 = head_dim == 512 && std::env::var("MEMRA_FA_I2").as_deref() != Ok("0");
10228 static TB512: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10242 let tb512 = head_dim == 512 && sp <= 32 && n_head / n_head_kv.max(1) <= 16
10244 && *TB512.get_or_init(|| std::env::var("MEMRA_FA_TB512").as_deref() != Ok("0"));
10245 let fname = if tb512 { "fa_decode_vec_q_rows_v4_512_tb" }
10246 else if i2 { "fa_decode_vec_q_rows_dpl16_i2" }
10247 else if head_dim == 512 { "fa_decode_vec_q_rows_dpl16" } else if v4 { "fa_decode_vec_q_rows_v4" }
10249 else if v3 { "fa_decode_vec_q_rows_v3" }
10250 else if fa_v2_on() { "fa_decode_vec_q_rows_v2" }
10251 else if smem_rows { "fa_decode_vec_q_rows_smem" }
10252 else { "fa_decode_vec_q_rows" };
10253 let f = if head_dim == 512 { self.fa_func(fname, head_dim) }
10254 else if g {
10255 self.func_g(if smem_rows { "fa_decode_vec_q_rows" } else { fname })
10263 }
10264 else { self.func(fname) };
10265 let shmem = if tb512 {
10266 let gk = Self::gkv_on();
10268 let sh = (8192 + 1024 + 32 * 512 + 32 * 64
10269 + 32 * head_dim * if gk { 1 } else { 2 }) as u32;
10270 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10271 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10272 sh
10273 } else if v4 || v3 || smem_rows || fa_v2_on() {
10274 let sh = (if v4 { 11520 + 32 * head_dim * if g { 1 } else { 2 } }
10276 else if v3 { 32 * head_dim * 2 } else { 2 * 32 * head_dim * 2 }) as u32;
10277 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10278 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10279 sh
10280 } else { 0 };
10281 for &(r0, t_g, sp_g) in &groups {
10285 let n_splits_g = (base_len + r0 + t_g).div_ceil(sp_g);
10286 let (nspm, spk) = (n_splits_g as i32, sp_g as i32);
10287 let base_i = (base_len + r0) as i32;
10288 let o_len = t_g * n_head * n_splits_g * head_dim;
10289 let ml_len = t_g * n_head * n_splits_g;
10290 let mut part_guard = self.fa_part_pool.lock().unwrap();
10291 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10292 let old = part_guard.take();
10303 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10304 if let Some(old) = old {
10305 self.fa_part_retired.lock().unwrap().push(old);
10306 }
10307 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10308 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10309 }
10310 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10311 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10312 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10313 }
10314 let pg = part_guard.as_mut().unwrap();
10315 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10316 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10317 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10318 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10319 let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
10320 let qv = self.view(q, t * n_head * head_dim);
10321 let q_g = qv.slice(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10322 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_g as u32, t_g as u32),
10323 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10324 {
10325 let __s_b = self.gpu.stream();
10326 let mut b = __s_b.launch_builder(&f);
10327 if tb512 {
10328 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10330 let plus_g = plus + r0 as i32;
10331 let nr = t_g as i32;
10332 if Self::pdl_on() && Self::pdl_wb_on() {
10333 use cudarc::driver::{DevicePtr, DevicePtrMut};
10335 let s = &self.gpu.stream();
10336 let (pq, _b0) = q_g.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10337 let (pv, _b2) = v.device_ptr(s);
10338 let (po, _b3) = part_o.device_ptr_mut(s);
10339 let (pm, _b4) = part_m.device_ptr_mut(s);
10340 let (pl, _b5) = part_l.device_ptr_mut(s);
10341 let (pb, _b6) = bd.device_ptr(s);
10342 let mut ps = [
10343 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10344 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10345 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10346 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10347 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10348 &plus_g as *const _ as *mut _, &scale as *const _ as *mut _,
10349 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10350 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10351 &nr as *const _ as *mut _,
10352 ];
10353 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10354 "fa_decode_vec_q_rows_v4_512_tb",
10355 (n_head_kv as u32, n_splits_g as u32, 1), (32, gqa, 1),
10356 shmem, &mut ps)?; }
10357 } else {
10358 let cfg_tb = LaunchConfig {
10359 grid_dim: (n_head_kv as u32, n_splits_g as u32, 1),
10360 block_dim: (32, gqa, 1), shared_mem_bytes: shmem };
10361 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10362 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10363 .arg(&ktb).arg(&vtb).arg(&nr);
10364 unsafe { b.launch(cfg_tb)?; }
10365 }
10366 } else if head_dim == 512 {
10367 let (bd, plus) = base_dev.expect("hd512 rows twin requires a device base counter");
10368 let plus_g = plus + r0 as i32;
10369 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10370 .arg(&hd).arg(&nh).arg(&nhkv).arg(bd).arg(&plus_g).arg(&scale).arg(&nspm).arg(&spk)
10371 .arg(&ktb).arg(&vtb);
10372 unsafe { b.launch(cfg)?; }
10373 } else {
10374 b.arg(&q_g).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10375 .arg(&hd).arg(&nh).arg(&nhkv).arg(&base_i).arg(&scale).arg(&nspm).arg(&spk)
10376 .arg(&ktb).arg(&vtb);
10377 unsafe { b.launch(cfg)?; }
10378 }
10379 }
10380 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t_g as u32, 1),
10381 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10382 let mut o_g = o.slice_mut(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
10383 if head_dim == 512 {
10384 let (bd, plus) = base_dev.unwrap();
10387 let plus_g = plus + r0 as i32;
10388 if let Some((oq, od)) = q8_out.as_mut() {
10389 debug_assert!(t == 1, "rows q8 emit is a t=1 decode arm");
10391 if Self::pdl_on() && Self::pdl_wb_on() {
10392 use cudarc::driver::{DevicePtr, DevicePtrMut};
10394 let s = &self.gpu.stream();
10395 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10396 let (pl, _g2) = part_l.device_ptr(s);
10397 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10398 let (pb, _g5) = bd.device_ptr(s);
10399 let mut ps = [
10400 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10401 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10402 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10403 &nh as *const _ as *mut _, &pb as *const _ as *mut _,
10404 &plus_g as *const _ as *mut _, &nspm as *const _ as *mut _,
10405 &spk as *const _ as *mut _,
10406 ];
10407 unsafe { self.launch_pdl_flash(Self::gkv_on(),
10408 "fa_decode_combine_rows_dc_q8_1",
10409 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10410 continue;
10411 }
10412 let fc = self.fa_func("fa_decode_combine_rows_dc_q8_1", head_dim);
10413 let __s_b2 = self.gpu.stream();
10414 let mut b2 = __s_b2.launch_builder(&fc);
10415 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut **oq).arg(&mut **od)
10416 .arg(&hd).arg(&nh).arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10417 unsafe { b2.launch(cfg2)?; }
10418 continue;
10419 }
10420 let fc = self.fa_func("fa_decode_combine_rows_dc", head_dim);
10421 let __s_b2 = self.gpu.stream();
10422 let mut b2 = __s_b2.launch_builder(&fc);
10423 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10424 .arg(bd).arg(&plus_g).arg(&nspm).arg(&spk);
10425 unsafe { b2.launch(cfg2)?; }
10426 } else {
10427 assert!(q8_out.is_none(), "rows q8 emit requires the hd512 dc combine");
10430 let fc = self.func("fa_decode_combine_rows");
10431 let __s_b2 = self.gpu.stream();
10432 let mut b2 = __s_b2.launch_builder(&fc);
10433 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(&mut o_g).arg(&hd).arg(&nh)
10434 .arg(&base_i).arg(&nspm).arg(&spk);
10435 unsafe { b2.launch(cfg2)?; }
10436 }
10437 }
10438 Ok(())
10439 }
10440
10441 #[allow(clippy::too_many_arguments)]
10445 pub fn fa_decode_rows_w(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10446 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10447 head_dim: usize, n_head: usize, n_head_kv: usize,
10448 base_dev: &CudaSlice<i32>, base_plus: i32, t: usize, scale: f32,
10449 window: usize, k_tok_bytes: usize, v_tok_bytes: usize,
10450 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10451 -> Result<(), Box<dyn std::error::Error>> {
10452 debug_assert!(head_dim == 256);
10457 let sp = {
10465 static SPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10466 let v = *SPW.get_or_init(|| std::env::var("MEMRA_FA_SPW").ok()
10467 .and_then(|x| x.parse().ok()).unwrap_or(0));
10468 if v >= 8 { v } else { FA_SPW_DEFAULT.load(std::sync::atomic::Ordering::Relaxed) }
10469 };
10470 let n_splits_max = (window + sp - 1) / sp;
10471 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10472 let (nspm, spk, wini) = (n_splits_max as i32, sp as i32, window as i32);
10473 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10474 let gqa = (n_head / n_head_kv).max(1) as u32;
10475 let o_len = t * n_head * n_splits_max * head_dim;
10476 let ml_len = t * n_head * n_splits_max;
10477 let mut part_guard = self.fa_part_pool.lock().unwrap();
10478 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10479 let old = part_guard.take();
10490 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10491 if let Some(old) = old {
10492 self.fa_part_retired.lock().unwrap().push(old);
10493 }
10494 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10495 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10496 }
10497 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10498 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10499 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10500 }
10501 let pg = part_guard.as_mut().unwrap();
10502 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10503 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10504 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10505 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10506 static SMEM_TKV_W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10512 let smem_tkv = *SMEM_TKV_W.get_or_init(|| {
10513 std::env::var("MEMRA_FA_SMEM_TKV").ok().and_then(|v| v.parse().ok())
10514 .unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
10515 });
10516 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10522 let wg = Self::wkv_on();
10527 let sp2 = gqa <= 4 && fa_v4_at(window)
10530 && std::env::var("MEMRA_FA_SPW2").as_deref() != Ok("0");
10531 if sp2 {
10532 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10533 if Self::pdl_on() && Self::pdl_wb_on() {
10534 use cudarc::driver::{DevicePtr, DevicePtrMut};
10536 let s = &self.gpu.stream();
10537 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10538 let (pv, _b2) = v.device_ptr(s);
10539 let (po, _b3) = part_o.device_ptr_mut(s);
10540 let (pm, _b4) = part_m.device_ptr_mut(s);
10541 let (pl, _b5) = part_l.device_ptr_mut(s);
10542 let (pb, _b6) = base_dev.device_ptr(s);
10543 let mut ps = [
10544 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10545 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10546 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10547 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10548 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10549 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10550 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10551 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10552 &wini as *const _ as *mut _,
10553 ];
10554 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w_sp",
10555 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa + 1, 1),
10556 sh, &mut ps)?; }
10557 } else {
10558 let f = if wg { self.func_g("fa_decode_vec_q_rows_v4_w_sp") }
10559 else { self.func("fa_decode_vec_q_rows_v4_w_sp") };
10560 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10561 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10562 block_dim: (32, gqa + 1, 1), shared_mem_bytes: sh };
10563 let __s_b = self.gpu.stream();
10564 let mut b = __s_b.launch_builder(&f);
10565 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10566 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10567 .arg(&ktb).arg(&vtb).arg(&wini);
10568 unsafe { b.launch(cfg)?; }
10569 }
10570 } else {
10571 if fa_v4_at(window) && Self::pdl_on() && Self::pdl_wb_on() {
10572 let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
10574 use cudarc::driver::{DevicePtr, DevicePtrMut};
10575 let s = &self.gpu.stream();
10576 let (pq, _b0) = q.device_ptr(s); let (pk, _b1) = k.device_ptr(s);
10577 let (pv, _b2) = v.device_ptr(s);
10578 let (po, _b3) = part_o.device_ptr_mut(s);
10579 let (pm, _b4) = part_m.device_ptr_mut(s);
10580 let (pl, _b5) = part_l.device_ptr_mut(s);
10581 let (pb, _b6) = base_dev.device_ptr(s);
10582 let mut ps = [
10583 &pq as *const _ as *mut std::ffi::c_void, &pk as *const _ as *mut _,
10584 &pv as *const _ as *mut _, &po as *const _ as *mut _,
10585 &pm as *const _ as *mut _, &pl as *const _ as *mut _,
10586 &hd as *const _ as *mut _, &nh as *const _ as *mut _,
10587 &nhkv as *const _ as *mut _, &pb as *const _ as *mut _,
10588 &base_plus as *const _ as *mut _, &scale as *const _ as *mut _,
10589 &nspm as *const _ as *mut _, &spk as *const _ as *mut _,
10590 &ktb as *const _ as *mut _, &vtb as *const _ as *mut _,
10591 &wini as *const _ as *mut _,
10592 ];
10593 unsafe { self.launch_pdl_flash(wg, "fa_decode_vec_q_rows_v4_w",
10594 (n_head_kv as u32, n_splits_max as u32, t as u32), (32, gqa, 1),
10595 sh, &mut ps)?; }
10596 } else {
10597 let pick = |name: &str| if wg { self.func_g(name) } else { self.func(name) };
10598 let (f, sh) = if fa_v4_at(window) {
10599 let f = pick("fa_decode_vec_q_rows_v4_w");
10600 (f, (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32)
10601 } else if smem_tkv > 0 && window >= smem_tkv {
10602 (pick("fa_decode_vec_q_rows_smem_w"), (2 * 32 * head_dim * 2) as u32)
10605 } else {
10606 (pick("fa_decode_vec_q_rows_reg_w"), 0u32)
10607 };
10608 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10609 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10610 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10611 let __s_b = self.gpu.stream();
10612 let mut b = __s_b.launch_builder(&f);
10613 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10614 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale).arg(&nspm).arg(&spk)
10615 .arg(&ktb).arg(&vtb).arg(&wini);
10616 unsafe { b.launch(cfg)?; }
10617 }
10618 }
10619 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10620 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10621 if let Some((oq, od)) = q8_out {
10622 if Self::pdl_on() && Self::pdl_wb_on() {
10625 use cudarc::driver::{DevicePtr, DevicePtrMut};
10627 let s = &self.gpu.stream();
10628 let (po, _g0) = part_o.device_ptr(s); let (pm, _g1) = part_m.device_ptr(s);
10629 let (pl, _g2) = part_l.device_ptr(s);
10630 let (pq, _g3) = oq.device_ptr_mut(s); let (pd, _g4) = od.device_ptr_mut(s);
10631 let mut ps = [
10632 &po as *const _ as *mut std::ffi::c_void, &pm as *const _ as *mut _,
10633 &pl as *const _ as *mut _, &pq as *const _ as *mut _,
10634 &pd as *const _ as *mut _, &hd as *const _ as *mut _,
10635 &nh as *const _ as *mut _, &nspm as *const _ as *mut _,
10636 &spk as *const _ as *mut _, &wini as *const _ as *mut _,
10637 ];
10638 unsafe { self.launch_pdl_flash(wg, "fa_decode_combine_rows_w_q8_1",
10639 cfg2.grid_dim, cfg2.block_dim, 0, &mut ps)?; }
10640 return Ok(());
10641 }
10642 let fc = if wg { self.func_g("fa_decode_combine_rows_w_q8_1") }
10643 else { self.func("fa_decode_combine_rows_w_q8_1") };
10644 let __s_b2 = self.gpu.stream();
10645 let mut b2 = __s_b2.launch_builder(&fc);
10646 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh)
10647 .arg(&nspm).arg(&spk).arg(&wini);
10648 unsafe { b2.launch(cfg2)?; }
10649 return Ok(());
10650 }
10651 let fc = if wg { self.func_g("fa_decode_combine_rows_w") }
10652 else { self.func("fa_decode_combine_rows_w") };
10653 let __s_b2 = self.gpu.stream();
10654 let mut b2 = __s_b2.launch_builder(&fc);
10655 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10656 .arg(&nspm).arg(&spk).arg(&wini);
10657 unsafe { b2.launch(cfg2)?; }
10658 Ok(())
10659 }
10660
10661 #[allow(clippy::too_many_arguments)]
10667 pub fn fa_decode_rows_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10668 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10669 head_dim: usize, n_head: usize, n_head_kv: usize,
10670 base_dev: &CudaSlice<i32>, t_kv_upper: usize, t: usize, scale: f32,
10671 k_tok_bytes: usize, v_tok_bytes: usize, base_plus: i32, g: bool)
10672 -> Result<(), Box<dyn std::error::Error>> {
10673 let v4 = head_dim == 256 && fa_v4_at(t_kv_upper);
10674 assert!(v4 || fa_v3_active(head_dim), "stream fa rows requires the v3 or v4 lane");
10675 assert!(v4 || base_plus == 0, "v3_dc kernel takes no plus arg");
10676 if v4 {
10677 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10678 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10679 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10680 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10681 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10682 let gqa = (n_head / n_head_kv).max(1) as u32;
10683 let o_len = t * n_head * n_splits_max * head_dim;
10684 let ml_len = t * n_head * n_splits_max;
10685 let mut part_guard = self.fa_part_pool.lock().unwrap();
10686 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10687 let old = part_guard.take();
10698 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10699 if let Some(old) = old {
10700 self.fa_part_retired.lock().unwrap().push(old);
10701 }
10702 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10703 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10704 }
10705 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10706 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10707 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10708 }
10709 let pg = part_guard.as_mut().unwrap();
10710 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10711 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10712 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10713 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10714 let f = if g { self.func_g("fa_decode_vec_q_rows_v4_dc") }
10715 else { self.func("fa_decode_vec_q_rows_v4_dc") };
10716 let sh = (11520 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10717 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10718 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10719 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10720 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10721 let __s_b = self.gpu.stream();
10722 let mut b = __s_b.launch_builder(&f);
10723 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10724 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&base_plus).arg(&scale)
10725 .arg(&nspm).arg(&spk).arg(&ktb).arg(&vtb);
10726 unsafe { b.launch(cfg)?; }
10727 let fc = self.func("fa_decode_combine_rows_dc");
10728 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10729 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10730 let __s_b2 = self.gpu.stream();
10731 let mut b2 = __s_b2.launch_builder(&fc);
10732 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10733 .arg(base_dev).arg(&base_plus).arg(&nspm).arg(&spk);
10734 unsafe { b2.launch(cfg2)?; }
10735 return Ok(());
10736 }
10737 let sp = fa_split_keys(t_kv_upper, n_head_kv);
10738 let n_splits_max = (t_kv_upper + sp - 1) / sp;
10739 let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
10740 let (nspm, spk) = (n_splits_max as i32, sp as i32);
10741 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10742 let gqa = (n_head / n_head_kv).max(1) as u32;
10743 let o_len = t * n_head * n_splits_max * head_dim;
10744 let ml_len = t * n_head * n_splits_max;
10745 let mut part_guard = self.fa_part_pool.lock().unwrap();
10746 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10747 let old = part_guard.take();
10758 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10759 if let Some(old) = old {
10760 self.fa_part_retired.lock().unwrap().push(old);
10761 }
10762 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10763 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10764 }
10765 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10766 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10767 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10768 }
10769 let pg = part_guard.as_mut().unwrap();
10770 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10771 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10772 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10773 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10774 let f = self.func("fa_decode_vec_q_rows_v3_dc");
10775 let sh = (32 * head_dim * 2) as u32;
10776 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10777 f.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, sh as i32)?;
10778 let cfg = LaunchConfig { grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
10779 block_dim: (32, gqa, 1), shared_mem_bytes: sh };
10780 let __s_b = self.gpu.stream();
10781 let mut b = __s_b.launch_builder(&f);
10782 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10783 .arg(&hd).arg(&nh).arg(&nhkv).arg(base_dev).arg(&scale).arg(&nspm).arg(&spk)
10784 .arg(&ktb).arg(&vtb);
10785 unsafe { b.launch(cfg)?; }
10786 let fc = self.func("fa_decode_combine_rows_dc");
10787 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, t as u32, 1),
10788 block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10789 let plus0 = 0i32;
10790 let __s_b2 = self.gpu.stream();
10791 let mut b2 = __s_b2.launch_builder(&fc);
10792 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh)
10793 .arg(base_dev).arg(&plus0).arg(&nspm).arg(&spk);
10794 unsafe { b2.launch(cfg2)?; }
10795 Ok(())
10796 }
10797
10798 pub fn fa_decode_dc(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10809 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10810 head_dim: usize, n_head: usize, n_head_kv: usize,
10811 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10812 k_tok_bytes: usize, v_tok_bytes: usize, g: bool)
10813 -> Result<(), Box<dyn std::error::Error>> {
10814 self.fa_decode_dc_q8(q, k, v, o, head_dim, n_head, n_head_kv, t_kv_dev, bucket_max,
10815 scale, k_tok_bytes, v_tok_bytes, g, None)
10816 }
10817
10818 #[allow(clippy::too_many_arguments)]
10821 pub fn fa_decode_dc_q8(&self, q: &CudaSlice<f32>, k: &cudarc::driver::CudaView<u8>,
10822 v: &cudarc::driver::CudaView<u8>, o: &mut CudaSlice<f32>,
10823 head_dim: usize, n_head: usize, n_head_kv: usize,
10824 t_kv_dev: &CudaSlice<i32>, bucket_max: usize, scale: f32,
10825 k_tok_bytes: usize, v_tok_bytes: usize, g: bool,
10826 q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>)
10827 -> Result<(), Box<dyn std::error::Error>> {
10828 let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
10836 if g && head_dim == 256 && !fa_v4_at(bucket_max) { fa_vec = false; } let sp = fa_split_keys(bucket_max, n_head_kv);
10838 let n_splits = if fa_vec { ((bucket_max + sp - 1) / sp).max(1) } else { ((bucket_max + 255) / 256).max(1) };
10839 let o_len = n_head * n_splits * head_dim;
10840 let ml_len = n_head * n_splits;
10841 let mut part_guard = self.fa_part_pool.lock().unwrap();
10842 if part_guard.as_ref().map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len).unwrap_or(true) {
10843 let old = part_guard.take();
10854 let (co, cm) = old.as_ref().map(|pp| (pp.0.len(), pp.1.len())).unwrap_or((0, 0));
10855 if let Some(old) = old {
10856 self.fa_part_retired.lock().unwrap().push(old);
10857 }
10858 if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
10859 eprintln!("[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)", co, o_len, cm, ml_len);
10860 }
10861 *part_guard = Some((self.alloc_uninit::<f32>(o_len.max(2 * co))?,
10862 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
10863 self.alloc_uninit::<f32>(ml_len.max(2 * cm))?));
10864 }
10865 let pg = part_guard.as_mut().unwrap();
10866 self.gpu.stream().memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
10867 self.gpu.stream().memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
10868 self.gpu.stream().memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
10869 let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
10870 let (hd, nh, nhkv, nsp) = (head_dim as i32, n_head as i32, n_head_kv as i32, n_splits as i32);
10871 let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
10872 let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
10873 let deep = fa_vec && head_dim == 256 && fa_v4_at(bucket_max) && !g
10876 && fa_deep_at(bucket_max) && !matches!(fa_v4_mode(), "noB3" | "stage");
10877 let (f, cfg) = if fa_vec && head_dim == 512 && bucket_max >= {
10878 static FA512_MIN_DC: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10879 *FA512_MIN_DC.get_or_init(|| std::env::var("MEMRA_FA512_MIN").ok()
10880 .and_then(|v| v.parse().ok()).unwrap_or(512))
10881 } {
10882 let gqa = (n_head / n_head_kv).max(1) as u32;
10884 (self.fa_func("fa_decode_vec_q_dpl16_dc", head_dim),
10885 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10886 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10887 } else if fa_vec && head_dim == 512 {
10888 let q_view = q.as_view();
10891 let mut o_view = o.as_view_mut();
10892 return self.fa_decode_scalar_unified(&q_view, k, v, &mut o_view,
10893 head_dim, n_head, n_head_kv,
10894 0, Some(t_kv_dev), scale, n_splits, sp,
10895 k_tok_bytes, v_tok_bytes, g,
10896 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10897 } else if fa_vec && head_dim == 256 && fa_v4_at(bucket_max) {
10898 let gqa = (n_head / n_head_kv).max(1) as u32;
10901 let fv = if g { self.func_g("fa_decode_vec_q_v4_dc") }
10902 else if deep { self.func("fa_decode_vec_q_v4_deep_dc") }
10903 else { self.func("fa_decode_vec_q_v4_dc") };
10904 let shmem = (if deep { 12160 } else { 11520 }
10905 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
10906 use cudarc::driver::sys::CUfunction_attribute_enum as A;
10907 fv.set_attribute(A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES, shmem as i32)?;
10908 (fv, LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10909 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10910 } else if fa_vec && fa_v3_active(head_dim) {
10911 let gqa = (n_head / n_head_kv).max(1) as u32;
10914 let fv = if g { self.func_g("fa_decode_vec_q_v3_dc") } else { self.func("fa_decode_vec_q_v3_dc") };
10915 let shmem = (32 * head_dim * 2) as u32; (fv,
10917 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10918 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10919 } else if fa_vec && fa_v2_on() {
10920 let gqa = (n_head / n_head_kv).max(1) as u32;
10924 let fv = if g { self.func_g("fa_decode_vec_q_v2_dc") } else { self.func("fa_decode_vec_q_v2_dc") };
10925 let shmem = (2 * 32 * head_dim * 2) as u32; (fv,
10927 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10928 block_dim: (32, gqa, 1), shared_mem_bytes: shmem })
10929 } else if fa_vec {
10930 let gqa = (n_head / n_head_kv).max(1) as u32;
10931 let fv = if g { self.func_g("fa_decode_vec_q_dc") } else { self.func("fa_decode_vec_q_dc") };
10933 (fv,
10934 LaunchConfig { grid_dim: (n_head_kv as u32, n_splits as u32, 1),
10935 block_dim: (32, gqa, 1), shared_mem_bytes: 0 })
10936 } else {
10937 let q_view = q.as_view();
10938 let mut o_view = o.as_view_mut();
10939 return self.fa_decode_scalar_unified(&q_view, k, v, &mut o_view,
10940 head_dim, n_head, n_head_kv,
10941 0, Some(t_kv_dev), scale, n_splits,
10942 if fa_vec { sp } else { 256 },
10943 k_tok_bytes, v_tok_bytes, g,
10944 &mut *part_o, &mut *part_m, &mut *part_l, q8_out);
10945 };
10946 let ski = sp as i32; let __s_b = self.gpu.stream();
10948 let mut b = __s_b.launch_builder(&f);
10949 b.arg(q).arg(k).arg(v).arg(&mut *part_o).arg(&mut *part_m).arg(&mut *part_l)
10950 .arg(&hd).arg(&nh).arg(&nhkv).arg(t_kv_dev).arg(&scale).arg(&nsp).arg(&ski)
10951 .arg(&ktb).arg(&vtb);
10952 unsafe { b.launch(cfg)?; }
10953 let cfg2 = LaunchConfig { grid_dim: (n_head as u32, 1, 1), block_dim: (head_dim as u32, 1, 1), shared_mem_bytes: 0 };
10954 if let Some((oq, od)) = q8_out {
10955 let fc = if g { self.func_g("fa_decode_combine_q8_1") }
10956 else { self.fa_func("fa_decode_combine_q8_1", head_dim) };
10957 let __s_b2 = self.gpu.stream();
10958 let mut b2 = __s_b2.launch_builder(&fc);
10959 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(oq).arg(od).arg(&hd).arg(&nh).arg(&nsp);
10960 unsafe { b2.launch(cfg2)?; }
10961 return Ok(());
10962 }
10963 let fc = if g { self.func_g("fa_decode_combine_f32") } else { self.fa_func("fa_decode_combine_f32", head_dim) };
10964 let __s_b2 = self.gpu.stream();
10965 let mut b2 = __s_b2.launch_builder(&fc);
10966 b2.arg(&*part_o).arg(&*part_m).arg(&*part_l).arg(o).arg(&hd).arg(&nh).arg(&nsp);
10967 unsafe { b2.launch(cfg2)?; }
10968 Ok(())
10969 }
10970
10971 pub fn fa_geom_eager(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
10977 let fa_ok = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
10981 let vec512 = fa_ok && head_dim == 512 && t_kv >= fa512_min_tkv();
10987 let mut fa_vec = vec512 || (fa_ok && head_dim <= 256 && head_dim % 32 == 0);
10988 if g && head_dim == 256 && !fa_v4_at(t_kv) { fa_vec = false; }
10994 let sp = fa_split_keys(t_kv, n_head_kv);
10995 let n_splits = if fa_vec { ((t_kv + sp - 1) / sp).max(1) } else { ((t_kv + 255) / 256).max(1) };
10996 (fa_vec, n_splits)
10997 }
10998
10999 pub fn fa_bucket_key(&self, t_kv: usize, head_dim: usize, n_head_kv: usize, g: bool) -> (bool, usize) {
11005 self.fa_geom_eager(t_kv, head_dim, n_head_kv, g)
11006 }
11007
11008 pub fn capture_graph_retained<F>(&self, step: F)
11020 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
11021 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
11022 {
11023 use cudarc::driver::sys::CUgraphInstantiate_flags;
11024 self.capture_graph_retained_flags(
11025 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH, step)
11026 }
11027
11028 pub fn capture_graph_retained_flags<F>(&self,
11033 flags: cudarc::driver::sys::CUgraphInstantiate_flags, mut step: F)
11034 -> Result<(cudarc::driver::CudaGraph, Vec<Box<dyn std::any::Any + Send>>), Box<dyn std::error::Error>>
11035 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
11036 {
11037 use cudarc::driver::sys::CUstreamCaptureMode;
11038 self.capture_keep.lock().unwrap().clear();
11046 let was_tracking = self.gpu.ctx.is_event_tracking();
11047 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
11048 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
11049 self.capture_keep_on.store(true, std::sync::atomic::Ordering::Relaxed);
11050 let w = (|| { step(self)?; step(self) })();
11051 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
11052 w?;
11053 self.gpu.stream().synchronize()?;
11054 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
11055 let r = step(self);
11056 let g = self.gpu.stream().end_capture(flags);
11057 r?;
11058 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
11059 graph.upload()?;
11060 Ok(graph)
11061 };
11062 let result = run();
11063 self.capture_keep_on.store(false, std::sync::atomic::Ordering::Relaxed);
11064 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
11065 let keeper = std::mem::take(&mut *self.capture_keep.lock().unwrap());
11066 Ok((result?, keeper))
11067 }
11068
11069 pub fn capture_graph<F>(&self, mut step: F) -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>>
11070 where F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>
11071 {
11072 use cudarc::driver::sys::{CUstreamCaptureMode, CUgraphInstantiate_flags};
11073 let was_tracking = self.gpu.ctx.is_event_tracking();
11081 if was_tracking { unsafe { self.gpu.ctx.disable_event_tracking(); } }
11082 let iflag = {
11089 static F: std::sync::OnceLock<CUgraphInstantiate_flags> = std::sync::OnceLock::new();
11090 *F.get_or_init(|| match std::env::var("MEMRA_GRAPH_IFLAG").as_deref() {
11091 Ok("upload") => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_UPLOAD,
11094 Ok("priority") =>
11095 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
11096 _ => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
11097 })
11098 };
11099 let ct = {
11106 static T: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11107 *T.get_or_init(|| std::env::var("MEMRA_GRAPH_CAPTIME").as_deref() == Ok("1"))
11108 };
11109 let warmups = {
11132 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
11133 *W.get_or_init(|| std::env::var("MEMRA_GRAPH_WARMUPS").ok()
11134 .and_then(|v| v.parse().ok()).filter(|n| *n >= 1).unwrap_or(1))
11135 };
11136 let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
11137 let t_w = std::time::Instant::now();
11138 for _ in 0..warmups { step(self)?; }
11140 self.gpu.stream().synchronize()?;
11141 let ms_warm = t_w.elapsed().as_secs_f64() * 1e3;
11142 let t_c = std::time::Instant::now();
11144 self.gpu.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
11145 let r = step(self);
11148 let ms_body = t_c.elapsed().as_secs_f64() * 1e3;
11149 let t_i = std::time::Instant::now();
11150 let g = self.gpu.stream().end_capture(iflag);
11151 let ms_inst = t_i.elapsed().as_secs_f64() * 1e3;
11152 r?;
11153 let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
11154 let t_u = std::time::Instant::now();
11155 graph.upload()?;
11156 if ct {
11157 println!("[graph-captime] warmup2x {ms_warm:.2} ms capture-body {ms_body:.2} ms \
11158 instantiate {ms_inst:.2} ms upload {:.2} ms",
11159 t_u.elapsed().as_secs_f64() * 1e3);
11160 }
11161 Ok(graph)
11162 };
11163 let result = run();
11164 if was_tracking { unsafe { self.gpu.ctx.enable_event_tracking(); } }
11165 result
11166 }
11167
11168 pub fn gdn_scan_s128_view(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11170 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
11171 state_in: &cudarc::driver::CudaView<f32>,
11172 state_out: &mut cudarc::driver::CudaViewMut<f32>,
11173 o: &mut CudaSlice<f32>, n_head: usize, t: usize, scale: f32)
11174 -> Result<(), Box<dyn std::error::Error>> {
11175 let f = self.func("gdn_scan_s128");
11176 const S_V: u32 = 128; const WARP: u32 = 32; const COLS: u32 = 4;
11177 let cfg = LaunchConfig { grid_dim: (n_head as u32, 1, S_V / COLS), block_dim: (WARP, COLS, 1), shared_mem_bytes: 0 };
11178 let (h, ti) = (n_head as i32, t as i32);
11179 let __s_b = self.gpu.stream();
11180 let mut b = __s_b.launch_builder(&f);
11181 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in).arg(state_out).arg(o).arg(&h).arg(&ti).arg(&scale);
11182 unsafe { b.launch(cfg)?; }
11183 Ok(())
11184 }
11185
11186 pub fn ssm_conv1d_view(&self, x: &cudarc::driver::CudaView<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11188 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
11189 -> Result<(), Box<dyn std::error::Error>> {
11190 let f = self.func("ssm_conv1d_silu_f32");
11191 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
11193 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11194 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
11195 let __s_b = self.gpu.stream();
11196 let mut b = __s_b.launch_builder(&f);
11197 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
11198 unsafe { b.launch(cfg)?; }
11199 Ok(())
11200 }
11201
11202 pub fn ssm_conv1d_tm(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11209 conv_dim: usize, t: usize, d_conv: usize)
11210 -> Result<(), Box<dyn std::error::Error>> {
11211 let f = self.func("ssm_conv1d_tm_f32");
11212 let cfg = LaunchConfig {
11213 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11214 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11215 };
11216 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11217 let __s_b = self.gpu.stream();
11218 let mut b = __s_b.launch_builder(&f);
11219 b.arg(qkv_tm).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11220 unsafe { b.launch(cfg)?; }
11221 Ok(())
11222 }
11223
11224 pub fn ssm_conv1d_tm_state(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
11232 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11233 conv_dim: usize, t: usize, d_conv: usize)
11234 -> Result<(), Box<dyn std::error::Error>> {
11235 self.ssm_conv1d_tm_state_pad(qkv_tm, conv_state, w, y, conv_dim, t, d_conv, None)
11236 }
11237
11238 #[allow(clippy::too_many_arguments)]
11241 pub fn ssm_conv1d_tm_state_pad(&self, qkv_tm: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
11242 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11243 conv_dim: usize, t: usize, d_conv: usize,
11244 pad_len: Option<&CudaSlice<i32>>)
11245 -> Result<(), Box<dyn std::error::Error>> {
11246 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
11247 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11251 {
11252 let f = self.func("ssm_conv1d_tm_state_f32");
11253 let cfg = LaunchConfig {
11254 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11255 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11256 };
11257 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11258 let __s_b = self.gpu.stream();
11259 let mut b = __s_b.launch_builder(&f);
11260 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11261 unsafe { b.launch(cfg)?; }
11262 }
11263 match (ring_old, pad_len) {
11264 (None, Some(len_d)) => {
11265 let f = self.func("ssm_conv_ring_update_dev_f32");
11266 let n = conv_dim * (d_conv - 1);
11267 let cfg = LaunchConfig::for_num_elems(n as u32);
11268 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11269 let __s_b = self.gpu.stream();
11270 let mut b = __s_b.launch_builder(&f);
11271 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11272 unsafe { b.launch(cfg)?; }
11273 }
11274 (None, None) => {
11275 let f = self.func("ssm_conv_ring_update_f32");
11276 let n = conv_dim * (d_conv - 1);
11277 let cfg = LaunchConfig::for_num_elems(n as u32);
11278 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11279 let __s_b = self.gpu.stream();
11280 let mut b = __s_b.launch_builder(&f);
11281 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11282 unsafe { b.launch(cfg)?; }
11283 }
11284 (Some(old), _) => self.ssm_conv_ring_rebuild(qkv_tm, &old, conv_state, conv_dim, t, d_conv)?,
11285 }
11286 Ok(())
11287 }
11288
11289 pub fn ssm_conv1d_tm_state_pad_v(&self, qkv_tm: &cudarc::driver::CudaView<f32>, conv_state: &mut CudaSlice<f32>,
11291 w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11292 conv_dim: usize, t: usize, d_conv: usize,
11293 pad_len: Option<&CudaSlice<i32>>)
11294 -> Result<(), Box<dyn std::error::Error>> {
11295 assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
11296 let ring_old = if t < d_conv - 1 { Some(self.clone_dtod(conv_state)?) } else { None };
11300 {
11301 let f = self.func("ssm_conv1d_tm_state_f32");
11302 let cfg = LaunchConfig {
11303 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11304 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11305 };
11306 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11307 let __s_b = self.gpu.stream();
11308 let mut b = __s_b.launch_builder(&f);
11309 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
11310 unsafe { b.launch(cfg)?; }
11311 }
11312 match (ring_old, pad_len) {
11313 (None, Some(len_d)) => {
11314 let f = self.func("ssm_conv_ring_update_dev_f32");
11315 let n = conv_dim * (d_conv - 1);
11316 let cfg = LaunchConfig::for_num_elems(n as u32);
11317 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11318 let __s_b = self.gpu.stream();
11319 let mut b = __s_b.launch_builder(&f);
11320 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11321 unsafe { b.launch(cfg)?; }
11322 }
11323 (None, None) => {
11324 let f = self.func("ssm_conv_ring_update_f32");
11325 let n = conv_dim * (d_conv - 1);
11326 let cfg = LaunchConfig::for_num_elems(n as u32);
11327 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11328 let __s_b = self.gpu.stream();
11329 let mut b = __s_b.launch_builder(&f);
11330 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11331 unsafe { b.launch(cfg)?; }
11332 }
11333 (Some(_), _) => unreachable!(
11334 "ssm_conv1d_tm_state_pad_v: T < d_conv-1 has no view path (PRIME_MIN_T gates it)"),
11335 }
11336 Ok(())
11337 }
11338
11339 pub fn ssm_conv_ring_rebuild(&self, qkv_tm: &CudaSlice<f32>, ring_old: &CudaSlice<f32>,
11344 conv_state: &mut CudaSlice<f32>,
11345 conv_dim: usize, tc: usize, d_conv: usize)
11346 -> Result<(), Box<dyn std::error::Error>> {
11347 let f = self.func("ssm_conv_ring_rebuild_f32");
11348 let n = conv_dim * (d_conv - 1);
11349 let cfg = LaunchConfig::for_num_elems(n as u32);
11350 let (cd, ti, dc) = (conv_dim as i32, tc as i32, d_conv as i32);
11351 let __s_b = self.gpu.stream();
11352 let mut b = __s_b.launch_builder(&f);
11353 b.arg(qkv_tm).arg(ring_old).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11354 unsafe { b.launch(cfg)?; }
11355 Ok(())
11356 }
11357
11358 #[allow(clippy::too_many_arguments)]
11363 pub fn gdn_prep_decode(&self, conv_out: &CudaSlice<f32>, beta_raw: &CudaSlice<f32>,
11364 alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11365 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11366 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11367 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32)
11368 -> Result<(), Box<dyn std::error::Error>> {
11369 let f = self.func("gdn_prep_decode_f32");
11370 let cfg = LaunchConfig { grid_dim: (num_v as u32, 1, 1), block_dim: (32, 4, 1), shared_mem_bytes: 0 };
11371 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11372 let __s_b = self.gpu.stream();
11373 let mut b = __s_b.launch_builder(&f);
11374 b.arg(conv_out).arg(beta_raw).arg(alpha).arg(dt_bias).arg(a)
11375 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11376 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps);
11377 unsafe { b.launch(cfg)?; }
11378 Ok(())
11379 }
11380
11381 #[allow(clippy::too_many_arguments)]
11385 pub fn ssm_conv1d_gdn(&self, qkv_tm: &CudaSlice<f32>, w: &CudaSlice<f32>,
11386 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11387 conv_dim: usize, t: usize, d_conv: usize,
11388 d_state: usize, num_v: usize, num_k: usize, key_dim: usize)
11389 -> Result<(), Box<dyn std::error::Error>> {
11390 let f = self.func("ssm_conv1d_gdn_f32");
11391 let cfg = LaunchConfig {
11392 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11393 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11394 };
11395 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11396 let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11397 let __s_b = self.gpu.stream();
11398 let mut b = __s_b.launch_builder(&f);
11399 b.arg(qkv_tm).arg(w).arg(q_g).arg(k_g).arg(v_g)
11400 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd);
11401 unsafe { b.launch(cfg)?; }
11402 Ok(())
11403 }
11404
11405 pub fn ssm_conv1d(&self, x: &CudaSlice<f32>, w: &CudaSlice<f32>, y: &mut CudaSlice<f32>,
11406 conv_dim: usize, t: usize, d_conv: usize, silu: bool)
11407 -> Result<(), Box<dyn std::error::Error>> {
11408 let f = self.func("ssm_conv1d_silu_f32");
11409 let cfg = LaunchConfig { grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
11410 block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11411 let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
11412 let __s_b = self.gpu.stream();
11413 let mut b = __s_b.launch_builder(&f);
11414 b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
11415 unsafe { b.launch(cfg)?; }
11416 Ok(())
11417 }
11418
11419 pub fn gdn_scan_s128(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11422 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
11423 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
11424 n_head: usize, t: usize, scale: f32)
11425 -> Result<(), Box<dyn std::error::Error>> {
11426 let f = self.func("gdn_scan_s128");
11427 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11428 let cfg = LaunchConfig {
11429 grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
11430 block_dim: (WARP, COLS_PER_BLOCK, 1),
11431 shared_mem_bytes: 0,
11432 };
11433 let (h, ti) = (n_head as i32, t as i32);
11434 let __s_b = self.gpu.stream();
11435 let mut b = __s_b.launch_builder(&f);
11436 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in).arg(state_out).arg(o).arg(&h).arg(&ti).arg(&scale);
11437 unsafe { b.launch(cfg)?; }
11438 Ok(())
11439 }
11440
11441 #[allow(clippy::too_many_arguments)]
11446 pub fn ssm_conv1d_fused_decode_b(
11447 &self, qkv_cols: &CudaSlice<f32>, conv_state_ptrs: &cudarc::driver::CudaView<u64>,
11448 w: &CudaSlice<f32>, conv_outs: &mut CudaSlice<f32>, conv_dim: usize, d_conv: usize,
11449 b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11450 let f = self.func("ssm_conv1d_fused_decode_b_f32");
11451 let cfg = LaunchConfig {
11452 grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
11453 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11454 };
11455 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11456 let __s_b = self.gpu.stream();
11457 let mut b = __s_b.launch_builder(&f);
11458 b.arg(qkv_cols).arg(conv_state_ptrs).arg(w).arg(conv_outs).arg(&cd).arg(&dc);
11459 unsafe { b.launch(cfg)?; }
11460 Ok(())
11461 }
11462
11463 #[allow(clippy::too_many_arguments)]
11464 pub fn gdn_prep_decode_b(
11465 &self, conv_outs: &CudaSlice<f32>, beta_raws: &CudaSlice<f32>, alphas: &CudaSlice<f32>,
11466 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11467 q_l2: &mut CudaSlice<f32>, k_l2: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
11468 beta: &mut CudaSlice<f32>, g_log: &mut CudaSlice<f32>,
11469 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, eps: f32,
11470 conv_dim: usize, b_n: usize) -> Result<(), Box<dyn std::error::Error>> {
11471 let f = self.func("gdn_prep_decode_b_f32");
11472 let cfg = LaunchConfig {
11473 grid_dim: (num_v as u32, 1, b_n as u32),
11474 block_dim: (32, 4, 1), shared_mem_bytes: 0,
11475 };
11476 let (ds, nv, nk, kd, cd) =
11477 (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, conv_dim as i32);
11478 let __s_b = self.gpu.stream();
11479 let mut b = __s_b.launch_builder(&f);
11480 b.arg(conv_outs).arg(beta_raws).arg(alphas).arg(dt_bias).arg(a)
11481 .arg(q_l2).arg(k_l2).arg(v_g).arg(beta).arg(g_log)
11482 .arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&eps).arg(&cd);
11483 unsafe { b.launch(cfg)?; }
11484 Ok(())
11485 }
11486
11487 #[allow(clippy::too_many_arguments)]
11488 pub fn gdn_scan_s128_batched(
11489 &self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11490 g: &CudaSlice<f32>, beta: &CudaSlice<f32>,
11491 state_in_ptrs: &cudarc::driver::CudaView<u64>,
11492 state_out_ptrs: &cudarc::driver::CudaView<u64>,
11493 o: &mut CudaSlice<f32>, n_head: usize, b_n: usize, scale: f32)
11494 -> Result<(), Box<dyn std::error::Error>> {
11495 let f = self.func("gdn_scan_s128_b");
11496 const S_V: u32 = 128; const WARP: u32 = 32; const COLS_PER_BLOCK: u32 = 4;
11497 let cfg = LaunchConfig {
11498 grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
11499 block_dim: (WARP, COLS_PER_BLOCK, 1), shared_mem_bytes: 0,
11500 };
11501 let h = n_head as i32;
11502 let __s_b = self.gpu.stream();
11503 let mut b = __s_b.launch_builder(&f);
11504 b.arg(q).arg(k).arg(v).arg(g).arg(beta).arg(state_in_ptrs).arg(state_out_ptrs)
11505 .arg(o).arg(&h).arg(&scale);
11506 unsafe { b.launch(cfg)?; }
11507 Ok(())
11508 }
11509
11510 pub fn gdn_chunked_enabled() -> bool {
11519 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11520 *E.get_or_init(|| std::env::var("MEMRA_GDN_CHUNKED").map(|v| v != "0").unwrap_or(true))
11521 }
11522
11523 pub fn gdn_chunk_size() -> usize {
11528 static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
11529 *C.get_or_init(|| {
11530 let c: usize = std::env::var("MEMRA_GDN_CHUNK").ok()
11531 .and_then(|v| v.parse().ok()).unwrap_or(32);
11532 c.clamp(32, 128) / 32 * 32
11533 })
11534 }
11535
11536 #[allow(clippy::too_many_arguments)]
11541 #[allow(clippy::too_many_arguments, clippy::type_complexity)]
11544 #[allow(clippy::too_many_arguments)]
11545 pub fn gdn_chunk_k123(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
11546 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, wb16: Option<&mut CudaSlice<u8>>,
11547 n_head: usize, t: usize, c: usize, hk: usize,
11548 k2w: Option<(&CudaSlice<u8>, &CudaSlice<u8>, &mut CudaSlice<u8>)>)
11549 -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11550 const D: usize = 128;
11551 let h = n_head;
11552 let nc = (t + c - 1) / c;
11553 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
11554 let mut gcum = self.uninit(t * h)?;
11555 let mut a = self.uninit(nc * h * c * c)?;
11556 let mut p = self.uninit(nc * h * c * c)?;
11557 let mut u = self.uninit(nc * h * c * D)?;
11558 let mut w = self.uninit(nc * h * c * D)?;
11559 { let f = self.func("gdn_chunk_cumgate_f32");
11561 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11562 let __s_b = self.gpu.stream();
11563 let mut b = __s_b.launch_builder(&f);
11564 b.arg(g).arg(&mut gcum).arg(&hi).arg(&ti).arg(&ci);
11565 unsafe { b.launch(cfg)?; }
11566 }
11567 if let Some((qb, kb, pb)) = k2w {
11568 assert!(c == 32, "gdn_k2_wgmma is a C==32 tile");
11571 let f = self.func("gdn_k2_wgmma");
11572 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11573 let hki = hk as i32;
11574 let __s_b = self.gpu.stream();
11575 let mut b = __s_b.launch_builder(&f);
11576 b.arg(qb).arg(kb).arg(&gcum).arg(beta).arg(&mut a).arg(&mut *pb).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11577 unsafe { b.launch(cfg)?; }
11578 } else if c <= 64 && !portable_mma_gated() { let f = self.func("gdn_chunk_attn_f32");
11580 let jt = ((c + 31) / 32) as u32;
11581 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11582 let hki = hk as i32;
11583 let __s_b = self.gpu.stream();
11584 let mut b = __s_b.launch_builder(&f);
11585 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
11586 unsafe { b.launch(cfg)?; }
11587 } else { assert!(hk == h, "generic K2 is broadcast-only (de-broadcast rides C==32)");
11589 let f = self.func("gdn_chunk_attn_g_f32");
11590 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (32, 8, 1), shared_mem_bytes: 0 };
11591 let __s_b = self.gpu.stream();
11592 let mut b = __s_b.launch_builder(&f);
11593 b.arg(q).arg(k).arg(&gcum).arg(beta).arg(&mut a).arg(&mut p).arg(&hi).arg(&ti).arg(&ci);
11594 unsafe { b.launch(cfg)?; }
11595 }
11596 { let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11598 match c {
11599 32 | 64 => {
11600 let f = self.func(if c == 32 { "gdn_chunk_solve32_f32" } else { "gdn_chunk_solve64_f32" });
11601 let wb: u64 = match wb16 { Some(d) => self.addr_u8(d), None => 0 };
11603 let hki = hk as i32;
11604 let __s_b = self.gpu.stream();
11605 let mut b = __s_b.launch_builder(&f);
11606 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&wb).arg(&hi).arg(&ti).arg(&hki);
11607 unsafe { b.launch(cfg)?; }
11608 }
11609 _ => {
11610 assert!(hk == h, "generic K3 is broadcast-only");
11611 let f = self.func("gdn_chunk_solve_f32");
11612 let __s_b = self.gpu.stream();
11613 let mut b = __s_b.launch_builder(&f);
11614 b.arg(v).arg(k).arg(&a).arg(&gcum).arg(&mut u).arg(&mut w).arg(&hi).arg(&ti).arg(&ci);
11615 unsafe { b.launch(cfg)?; }
11616 }
11617 }
11618 }
11619 Ok((gcum, p, u, w))
11620 }
11621
11622 pub fn gdn_db_on() -> bool {
11626 std::env::var("MEMRA_GDN_DB").as_deref() != Ok("0")
11627 }
11628
11629 pub fn gdn_mma_enabled(&self, c: usize) -> bool {
11632 !portable_mma_gated() && c == 32
11633 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
11634 Ok("1") => true,
11635 Ok("0") => false,
11636 _ => cfg!(memra_hopper_mma),
11637 }
11638 }
11639
11640 pub fn gdn_wgmma_on(&self, c: usize) -> bool {
11643 self.gdn_mma_enabled(c)
11644 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
11645 Ok("0") => false,
11646 Ok("1") => true,
11647 _ => cfg!(memra_hopper_mma),
11648 }
11649 }
11650
11651 #[allow(clippy::too_many_arguments)]
11656 pub fn ssm_conv1d_gdn_state_pad(&self, qkv_tm: &cudarc::driver::CudaView<f32>,
11657 conv_state: &mut CudaSlice<f32>, w: &CudaSlice<f32>,
11658 q_g: &mut CudaSlice<f32>, k_g: &mut CudaSlice<f32>,
11659 v_g: &mut CudaSlice<f32>,
11660 conv_dim: usize, t: usize, d_conv: usize,
11661 d_state: usize, num_v: usize, num_k: usize, key_dim: usize,
11662 hk: usize,
11663 pad_len: Option<&CudaSlice<i32>>)
11664 -> Result<(), Box<dyn std::error::Error>> {
11665 assert!(t >= d_conv - 1, "fused state conv requires T >= pad (PRIME_MIN_T gates)");
11666 {
11667 let f = self.func("ssm_conv1d_gdn_state_f32");
11668 let cfg = LaunchConfig {
11669 grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
11670 block_dim: (256, 1, 1), shared_mem_bytes: 0,
11671 };
11672 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11673 let (ds, nv, nk, kd, hki) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, hk as i32);
11674 let __s_b = self.gpu.stream();
11675 let mut b = __s_b.launch_builder(&f);
11676 b.arg(qkv_tm).arg(&*conv_state).arg(w).arg(q_g).arg(k_g).arg(v_g)
11677 .arg(&cd).arg(&ti).arg(&dc).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&hki);
11678 unsafe { b.launch(cfg)?; }
11679 }
11680 match pad_len {
11681 Some(len_d) => {
11682 let f = self.func("ssm_conv_ring_update_dev_f32");
11683 let n = conv_dim * (d_conv - 1);
11684 let cfg = LaunchConfig::for_num_elems(n as u32);
11685 let (cd, dc) = (conv_dim as i32, d_conv as i32);
11686 let __s_b = self.gpu.stream();
11687 let mut b = __s_b.launch_builder(&f);
11688 b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
11689 unsafe { b.launch(cfg)?; }
11690 }
11691 None => {
11692 let f = self.func("ssm_conv_ring_update_f32");
11693 let n = conv_dim * (d_conv - 1);
11694 let cfg = LaunchConfig::for_num_elems(n as u32);
11695 let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
11696 let __s_b = self.gpu.stream();
11697 let mut b = __s_b.launch_builder(&f);
11698 b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
11699 unsafe { b.launch(cfg)?; }
11700 }
11701 }
11702 Ok(())
11703 }
11704
11705 pub fn gdn_chunk_alloc(&self, n_head: usize, t: usize, c: usize, hk: usize)
11709 -> Result<GdnChunkBufs, Box<dyn std::error::Error>> {
11710 const D: usize = 128;
11711 assert!(c == 32, "gdn_chunk_alloc: varlen chain is the C==32 mma pair");
11712 let h = n_head;
11713 let nc = (t + c - 1) / c;
11714 Ok(GdnChunkBufs {
11715 gcum: self.uninit(t * h)?,
11716 a: self.uninit(nc * h * c * c)?,
11717 p: self.uninit(nc * h * c * c)?,
11718 u: self.uninit(nc * h * c * D)?,
11719 w: self.uninit(nc * h * c * D)?,
11720 kb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11721 wb16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11722 y16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
11723 ssnap16: self.alloc_u8_uninit(nc * h * D * D * 2)?,
11724 qb16: self.alloc_u8_uninit(t * hk * D * 2)?,
11725 pb16: self.alloc_u8_uninit(nc * h * c * c * 2)?,
11726 o: self.uninit(D * h * t)?,
11727 t, nc,
11728 })
11729 }
11730
11731 pub fn f32_to_bf16_v(&self, x: &cudarc::driver::CudaView<f32>, dst: &mut CudaSlice<u8>, n: usize)
11733 -> Result<(), Box<dyn std::error::Error>> {
11734 let f = self.func("f32_to_bf16_bulk");
11735 let ni = n as i64;
11736 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11737 let __s_b = self.gpu.stream();
11738 let mut b = __s_b.launch_builder(&f);
11739 b.arg(x).arg(dst).arg(&ni);
11740 unsafe { b.launch(cfg)?; }
11741 Ok(())
11742 }
11743
11744 pub fn f32_to_bf16_into(&self, x: &CudaSlice<f32>, dst: &mut CudaSlice<u8>, n: usize)
11746 -> Result<(), Box<dyn std::error::Error>> {
11747 let f = self.func("f32_to_bf16_bulk");
11748 let ni = n as i64;
11749 let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
11750 let __s_b = self.gpu.stream();
11751 let mut b = __s_b.launch_builder(&f);
11752 b.arg(x).arg(dst).arg(&ni);
11753 unsafe { b.launch(cfg)?; }
11754 Ok(())
11755 }
11756
11757 pub fn gdn_chunk_k123_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, hk: usize,
11760 wq: Option<&GdnWVl8>)
11761 -> Result<(), Box<dyn std::error::Error>> {
11762 let b = seqs.len();
11763 assert!(b >= 1 && b <= 8, "gdn_chunk_k123_vl8: 1..=8 sequences");
11764 let mut packed = [GdnSeqVl::default(); 8];
11765 packed[..b].copy_from_slice(seqs);
11766 let v = GdnVl8(packed);
11767 let (hi, ci) = (n_head as i32, 32i32);
11768 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11769 {
11770 let f = self.func("gdn_chunk_cumgate_vl");
11771 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (32, 1, 1), shared_mem_bytes: 0 };
11772 let __s_lb = self.gpu.stream();
11773 let mut lb = __s_lb.launch_builder(&f);
11774 lb.arg(&v).arg(&hi).arg(&ci);
11775 unsafe { lb.launch(cfg)?; }
11776 }
11777 let hki = hk as i32;
11778 if let Some(w) = wq { let f = self.func("gdn_k2_wgmma_vl");
11780 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11781 let __s_lb = self.gpu.stream();
11782 let mut lb = __s_lb.launch_builder(&f);
11783 lb.arg(&v).arg(w).arg(&hi).arg(&ci).arg(&hki);
11784 unsafe { lb.launch(cfg)?; }
11785 } else {
11786 let f = self.func("gdn_chunk_attn_vl");
11787 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11788 let __s_lb = self.gpu.stream();
11789 let mut lb = __s_lb.launch_builder(&f);
11790 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11791 unsafe { lb.launch(cfg)?; }
11792 }
11793 {
11794 let f = self.func("gdn_chunk_solve32_vl");
11795 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11796 let __s_lb = self.gpu.stream();
11797 let mut lb = __s_lb.launch_builder(&f);
11798 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11799 unsafe { lb.launch(cfg)?; }
11800 }
11801 Ok(())
11802 }
11803
11804 #[allow(clippy::too_many_arguments)]
11808 pub fn gdn_prep_vl8(&self, seqs: &[GdnPrepVl], conv_w: &CudaSlice<f32>,
11809 dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
11810 conv_dim: usize, d_conv: usize, d_state: usize,
11811 num_v: usize, num_k: usize, key_dim: usize, hk: usize, eps: f32)
11812 -> Result<(), Box<dyn std::error::Error>> {
11813 let b = seqs.len();
11814 assert!(b >= 1 && b <= 8);
11815 let mut packed = [GdnPrepVl::default(); 8];
11816 packed[..b].copy_from_slice(seqs);
11817 let v = GdnPrepVl8(packed);
11818 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11819 let (cdi, dci) = (conv_dim as i32, d_conv as i32);
11820 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
11821 assert!(conv_fuse || hk == num_v, "de-broadcast requires the fused conv");
11822 if conv_fuse {
11823 let f = self.func("ssm_conv1d_gdn_state_vl");
11824 let cfg = LaunchConfig { grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11825 let (dsi, nvi, nki, kdi, hki) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, hk as i32);
11826 let __s_lb = self.gpu.stream();
11827 let mut lb = __s_lb.launch_builder(&f);
11828 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi).arg(&hki);
11829 unsafe { lb.launch(cfg)?; }
11830 } else {
11831 let f = self.func("ssm_conv1d_tm_state_vl");
11832 let cfg = LaunchConfig { grid_dim: ((conv_dim as u32).div_ceil(256), max_t, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11833 let __s_lb = self.gpu.stream();
11834 let mut lb = __s_lb.launch_builder(&f);
11835 lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci);
11836 unsafe { lb.launch(cfg)?; }
11837 }
11838 {
11839 let f = self.func("ssm_conv_ring_update_vl");
11840 let n = (conv_dim * (d_conv - 1)) as u32;
11841 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11842 let __s_lb = self.gpu.stream();
11843 let mut lb = __s_lb.launch_builder(&f);
11844 lb.arg(&v).arg(&cdi).arg(&dci);
11845 unsafe { lb.launch(cfg)?; }
11846 }
11847 if !conv_fuse {
11848 let f = self.func("qkv_to_gdn_repack_vl");
11849 let n = max_t * (num_v * d_state) as u32;
11850 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11851 let (dsi, nvi, nki, kdi) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
11852 let __s_lb = self.gpu.stream();
11853 let mut lb = __s_lb.launch_builder(&f);
11854 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi);
11855 unsafe { lb.launch(cfg)?; }
11856 }
11857 if Self::l2_v2_on(d_state) {
11858 let f = self.func("gdn_l2_v2_vl");
11859 let cfg = LaunchConfig { grid_dim: ((max_t * hk as u32).div_ceil(8), 2, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11860 let (dsi, nvi) = (d_state as i32, hk as i32);
11861 let __s_lb = self.gpu.stream();
11862 let mut lb = __s_lb.launch_builder(&f);
11863 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11864 unsafe { lb.launch(cfg)?; }
11865 } else {
11866 let f = self.func("gdn_l2_vl");
11867 let cfg = LaunchConfig { grid_dim: (max_t * hk as u32, 2, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11868 let (dsi, nvi) = (d_state as i32, hk as i32);
11869 let __s_lb = self.gpu.stream();
11870 let mut lb = __s_lb.launch_builder(&f);
11871 lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
11872 unsafe { lb.launch(cfg)?; }
11873 }
11874 {
11875 let f = self.func("gdn_gate_prep_vl");
11876 let n = max_t * num_v as u32;
11877 let cfg = LaunchConfig { grid_dim: (n.div_ceil(256), 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11878 let nvi = num_v as i32;
11879 let __s_lb = self.gpu.stream();
11880 let mut lb = __s_lb.launch_builder(&f);
11881 lb.arg(&v).arg(dt_bias).arg(a).arg(&nvi);
11882 unsafe { lb.launch(cfg)?; }
11883 }
11884 Ok(())
11885 }
11886
11887 pub fn gdn_mirror_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, which: i32, hk: usize)
11889 -> Result<(), Box<dyn std::error::Error>> {
11890 let b = seqs.len();
11891 assert!(b >= 1 && b <= 8);
11892 let mut packed = [GdnSeqVl::default(); 8];
11893 packed[..b].copy_from_slice(seqs);
11894 let v = GdnVl8(packed);
11895 let ept = (if which == 0 { hk } else { n_head } * 128) as i32;
11896 let max_n = seqs.iter().map(|s| if which == 0 { s.t as i64 * ept as i64 }
11897 else { s.nc as i64 * ept as i64 * 32 }).max().unwrap();
11898 let f = self.func("gdn_mirror_vl");
11899 let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
11900 let cfg = LaunchConfig { grid_dim: (blocks, 1, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11901 let __s_lb = self.gpu.stream();
11902 let mut lb = __s_lb.launch_builder(&f);
11903 lb.arg(&v).arg(&ept).arg(&which);
11904 unsafe { lb.launch(cfg)?; }
11905 Ok(())
11906 }
11907
11908 pub fn gdn_tail_vl8(&self, seqs: &[GdnPrepVl], norm_w: &CudaSlice<f32>,
11910 d_state: usize, num_v: usize, eps: f32)
11911 -> Result<(), Box<dyn std::error::Error>> {
11912 let b = seqs.len();
11913 assert!(b >= 1 && b <= 8);
11914 let mut packed = [GdnPrepVl::default(); 8];
11915 packed[..b].copy_from_slice(seqs);
11916 let v = GdnPrepVl8(packed);
11917 let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
11918 let f = self.func("gated_rmsnorm_f16out_vl");
11919 let cfg = LaunchConfig { grid_dim: (max_t * num_v as u32, 1, b as u32), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
11921 let (dsi, nvi) = (d_state as i32, num_v as i32);
11922 let __s_lb = self.gpu.stream();
11923 let mut lb = __s_lb.launch_builder(&f);
11924 lb.arg(&v).arg(norm_w).arg(&dsi).arg(&nvi).arg(&eps);
11925 unsafe { lb.launch(cfg)?; }
11926 Ok(())
11927 }
11928
11929 pub fn addr_f32(&self, x: &CudaSlice<f32>) -> u64 {
11932 use cudarc::driver::DevicePtr;
11933 let s = self.gpu.stream();
11934 let (p, _g) = x.device_ptr(&s);
11935 p as u64
11936 }
11937 pub fn addr_f32_mut(&self, x: &mut CudaSlice<f32>) -> u64 {
11938 use cudarc::driver::DevicePtrMut;
11939 let s = self.gpu.stream();
11940 let (p, _g) = x.device_ptr_mut(&s);
11941 p as u64
11942 }
11943 pub fn addr_f32v(&self, x: &cudarc::driver::CudaView<f32>) -> u64 {
11944 use cudarc::driver::DevicePtr;
11945 let s = self.gpu.stream();
11946 let (p, _g) = x.device_ptr(&s);
11947 p as u64
11948 }
11949 pub fn addr_u8(&self, x: &CudaSlice<u8>) -> u64 {
11950 use cudarc::driver::DevicePtr;
11951 let s = self.gpu.stream();
11952 let (p, _g) = x.device_ptr(&s);
11953 p as u64
11954 }
11955
11956 pub fn gdn_chunk_vl8(&self, seqs: &[GdnSeqVl], n_head: usize, scale: f32, hk: usize,
11960 wq: Option<&GdnWVl8>)
11961 -> Result<(), Box<dyn std::error::Error>> {
11962 const NSPLIT: u32 = 4;
11963 let b = seqs.len();
11964 assert!(b >= 1 && b <= 8, "gdn_chunk_vl8: 1..=8 sequences");
11965 let mut packed = [GdnSeqVl::default(); 8];
11966 packed[..b].copy_from_slice(seqs);
11967 let v = GdnVl8(packed);
11968 let (hi, ci) = (n_head as i32, 32i32);
11969 let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
11970 let hki = hk as i32;
11971 if let Some(w) = wq {
11972 let f = self.func("gdn_k45_wgmma_vl");
11974 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11975 let __s_lb = self.gpu.stream();
11976 let mut lb = __s_lb.launch_builder(&f);
11977 lb.arg(&v).arg(w).arg(&scale).arg(&hi).arg(&ci).arg(&hki);
11978 unsafe { lb.launch(cfg)?; }
11979 let _ = max_nc;
11980 return Ok(());
11981 }
11982 {
11983 let f = self.func("gdn_chunk_state_mma_vl");
11984 let cfg = LaunchConfig { grid_dim: (n_head as u32, NSPLIT, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11985 let __s_lb = self.gpu.stream();
11986 let mut lb = __s_lb.launch_builder(&f);
11987 lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
11988 unsafe { lb.launch(cfg)?; }
11989 }
11990 {
11991 let f = self.func("gdn_chunk_output_mma_vl");
11992 let cfg = LaunchConfig { grid_dim: (max_nc, n_head as u32, b as u32), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
11993 let __s_lb = self.gpu.stream();
11994 let mut lb = __s_lb.launch_builder(&f);
11995 lb.arg(&v).arg(&hi).arg(&ci).arg(&scale).arg(&hki);
11996 unsafe { lb.launch(cfg)?; }
11997 }
11998 Ok(())
11999 }
12000 pub fn gdn_scan_chunked(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12001 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
12002 qb16_pre: Option<&CudaSlice<u8>>,
12003 state_in: &CudaSlice<f32>,
12004 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12005 n_head: usize, t: usize, scale: f32, c: usize, hk: usize)
12006 -> Result<(), Box<dyn std::error::Error>> {
12007 const D: usize = 128;
12008 const NSPLIT: u32 = 4;
12009 assert!(c >= 1 && c <= 128, "gdn_scan_chunked: C must be in 1..=128");
12010 let h = n_head;
12011 let nc = (t + c - 1) / c;
12012 let (hi, ti, ci) = (h as i32, t as i32, c as i32);
12013 let gdn_mma_pre = !portable_mma_gated() && c == 32
12017 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
12018 Ok("1") => true,
12019 Ok("0") => false,
12020 _ => cfg!(memra_hopper_mma),
12021 };
12022 let mut wb16_pre: Option<CudaSlice<u8>> = if gdn_mma_pre {
12023 Some(self.alloc_u8_uninit(nc * h * c * D * 2)?)
12024 } else { None };
12025 let gdn_wgmma_pre = gdn_mma_pre
12029 && match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
12030 Ok("0") => false,
12031 Ok("1") => true,
12032 _ => cfg!(memra_hopper_mma),
12033 };
12034 let nk = t * hk * D;
12035 let mut kb16_local: Option<CudaSlice<u8>> = None;
12036 if gdn_mma_pre && kb16_pre.is_none() {
12037 let mut kb = self.alloc_u8_uninit(nk * 2)?;
12038 let f = self.func("f32_to_bf16_bulk");
12039 let n2 = nk as i64;
12040 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
12041 let __s_b = self.gpu.stream();
12042 let mut b = __s_b.launch_builder(&f);
12043 b.arg(k).arg(&mut kb).arg(&n2);
12044 unsafe { b.launch(cfg2)?; }
12045 kb16_local = Some(kb);
12046 }
12047 let kb16_ref0: Option<&CudaSlice<u8>> = kb16_local.as_ref().or(kb16_pre);
12048 if let Some(kb) = kb16_pre { assert!(kb.len() >= nk * 2, "kb16_pre too small"); }
12049 let mut qb16: Option<CudaSlice<u8>> = None;
12050 let mut pb16: Option<CudaSlice<u8>> = None;
12051 if gdn_wgmma_pre {
12052 if qb16_pre.is_none() {
12055 let mut qb = self.alloc_u8_uninit(nk * 2)?;
12056 let f = self.func("f32_to_bf16_bulk");
12057 let n2 = nk as i64;
12058 let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
12059 let __s_b = self.gpu.stream();
12060 let mut b = __s_b.launch_builder(&f);
12061 b.arg(q).arg(&mut qb).arg(&n2);
12062 unsafe { b.launch(cfg2)?; }
12063 qb16 = Some(qb);
12064 } else if let Some(qb) = qb16_pre {
12065 assert!(qb.len() >= nk * 2, "qb16_pre too small");
12066 }
12067 pb16 = Some(self.alloc_u8_uninit(nc * h * c * c * 2)?);
12068 }
12069 let qb16_ref0: Option<&CudaSlice<u8>> = qb16.as_ref().or(qb16_pre);
12070 let k2w = if gdn_wgmma_pre {
12071 Some((*qb16_ref0.as_ref().unwrap(),
12072 *kb16_ref0.as_ref().unwrap(),
12073 pb16.as_mut().unwrap()))
12074 } else { None };
12075 let (gcum, p, u, w) = self.gdn_chunk_k123(q, k, v, g, beta, wb16_pre.as_mut(), n_head, t, c, hk, k2w)?;
12076 let _ = &w;
12077 let mut y = self.uninit(nc * h * c * D)?;
12078 let mut ssnap = self.uninit(nc * h * D * D)?; let gdn_mma = !portable_mma_gated() && c == 32
12090 && match std::env::var("MEMRA_GDN_MMA").as_deref() {
12091 Ok("1") => true,
12092 Ok("0") => false,
12093 _ => cfg!(memra_hopper_mma),
12094 };
12095 if gdn_mma {
12096 let wb16 = wb16_pre.take().expect("mma path pre-allocates wb16 (K3 store fold)");
12097 let kb16_ref: &CudaSlice<u8> = kb16_ref0.expect("mma path pre-builds kb16 above K123");
12098 if gdn_wgmma_pre {
12110 let qb16 = qb16_ref0.unwrap();
12112 let pb16 = pb16.as_ref().unwrap();
12113 {
12114 let f = self.func("gdn_k45_wgmma");
12115 let cfg = LaunchConfig { grid_dim: (h as u32, 4, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12116 let hki = hk as i32;
12117 let __s_b = self.gpu.stream();
12118 let mut b = __s_b.launch_builder(&f);
12119 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(qb16).arg(pb16)
12120 .arg(o).arg(&scale).arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
12121 unsafe { b.launch(cfg)?; }
12122 }
12123 return Ok(());
12124 }
12125 let mut y16 = self.alloc_u8_uninit(nc * h * c * D * 2)?;
12129 let mut ssnap16 = self.alloc_u8_uninit(nc * h * D * D * 2)?;
12130 {
12131 let f = self.func("gdn_chunk_state_mma");
12132 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12133 let hki = hk as i32;
12134 let __s_b = self.gpu.stream();
12135 let mut b = __s_b.launch_builder(&f);
12136 b.arg(kb16_ref).arg(&gcum).arg(beta).arg(&u).arg(&wb16).arg(&mut y16).arg(&mut ssnap16)
12137 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci).arg(&hki);
12138 unsafe { b.launch(cfg)?; }
12139 }
12140 { let f = self.func("gdn_chunk_output_mma");
12142 let jt = ((c + 31) / 32) as u32;
12143 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12144 let hki = hk as i32;
12145 let __s_b = self.gpu.stream();
12146 let mut b = __s_b.launch_builder(&f);
12147 b.arg(q).arg(&gcum).arg(&p).arg(&y16).arg(&ssnap16).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale).arg(&hki);
12148 unsafe { b.launch(cfg)?; }
12149 }
12150 return Ok(());
12151 }
12152 { let f = self.func("gdn_chunk_state_f32");
12154 let cfg = LaunchConfig { grid_dim: (h as u32, NSPLIT, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12155 let __s_b = self.gpu.stream();
12156 let mut b = __s_b.launch_builder(&f);
12157 b.arg(k).arg(&gcum).arg(beta).arg(&u).arg(&w).arg(&mut y).arg(&mut ssnap)
12158 .arg(state_in).arg(&mut *state_out).arg(&hi).arg(&ti).arg(&ci);
12159 unsafe { b.launch(cfg)?; }
12160 }
12161 { let f = self.func("gdn_chunk_output_f32");
12163 let jt = ((c + 31) / 32) as u32;
12164 let cfg = LaunchConfig { grid_dim: (nc as u32, h as u32, jt), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
12165 let __s_b = self.gpu.stream();
12166 let mut b = __s_b.launch_builder(&f);
12167 b.arg(q).arg(&gcum).arg(&p).arg(&y).arg(&ssnap).arg(o).arg(&hi).arg(&ti).arg(&ci).arg(&scale);
12168 unsafe { b.launch(cfg)?; }
12169 }
12170 Ok(())
12171 }
12172
12173 #[allow(clippy::too_many_arguments)]
12182 #[allow(clippy::too_many_arguments)]
12183 pub fn gdn_scan_prefill(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12184 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, kb16_pre: Option<&CudaSlice<u8>>,
12185 qb16_pre: Option<&CudaSlice<u8>>,
12186 state_in: &CudaSlice<f32>,
12187 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12188 n_head: usize, t: usize, scale: f32, hk: usize)
12189 -> Result<(), Box<dyn std::error::Error>> {
12190 if std::env::var("MEMRA_GDN_DIFF").is_ok() && t >= 16 {
12191 assert!(hk == n_head, "GDN_DIFF oracle is broadcast-only");
12192 return self.gdn_scan_diff(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale);
12193 }
12194 if Self::gdn_chunked_enabled() && t >= 16 {
12195 self.gdn_scan_chunked(q, k, v, g, beta, kb16_pre, qb16_pre, state_in, state_out, o, n_head, t, scale,
12196 Self::gdn_chunk_size(), hk)
12197 } else {
12198 assert!(hk == n_head, "s128 scan is broadcast-only (prep guarantees by predicate)");
12199 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)
12200 }
12201 }
12202
12203 #[allow(clippy::too_many_arguments)]
12205 fn gdn_scan_diff(&self, q: &CudaSlice<f32>, k: &CudaSlice<f32>, v: &CudaSlice<f32>,
12206 g: &CudaSlice<f32>, beta: &CudaSlice<f32>, state_in: &CudaSlice<f32>,
12207 state_out: &mut CudaSlice<f32>, o: &mut CudaSlice<f32>,
12208 n_head: usize, t: usize, scale: f32)
12209 -> Result<(), Box<dyn std::error::Error>> {
12210 static CALL: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
12211 let call = CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12212 let mut o_c = self.uninit(o.len())?;
12213 let mut st_c = self.uninit(state_out.len())?;
12214 self.gdn_scan_chunked(q, k, v, g, beta, None, None, state_in, &mut st_c, &mut o_c,
12215 n_head, t, scale, Self::gdn_chunk_size(), n_head)?;
12216 self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)?;
12217 let (oh_s, oh_c) = (self.dtoh(o)?, self.dtoh(&o_c)?);
12218 let (sh_s, sh_c) = (self.dtoh(state_out)?, self.dtoh(&st_c)?);
12219 let stats = |a: &[f32], b: &[f32]| -> (f32, f32, f64) {
12220 let mut max_abs = 0f32; let mut max_rel = 0f32; let mut sum_rel = 0f64;
12221 for (x, y) in a.iter().zip(b) {
12222 let ad = (x - y).abs();
12223 let rel = ad / x.abs().max(y.abs()).max(1e-3);
12224 if ad > max_abs { max_abs = ad; }
12225 if rel > max_rel { max_rel = rel; }
12226 sum_rel += rel as f64;
12227 }
12228 (max_abs, max_rel, sum_rel / a.len() as f64)
12229 };
12230 let (o_ma, o_mr, o_mean) = stats(&oh_s, &oh_c);
12231 let (s_ma, s_mr, s_mean) = stats(&sh_s, &sh_c);
12232 println!("[gdn-diff call {call:3} T={t} C={}] out: max_abs={o_ma:.3e} max_rel={o_mr:.3e} mean_rel={o_mean:.3e} | \
12233 state: max_abs={s_ma:.3e} max_rel={s_mr:.3e} mean_rel={s_mean:.3e}",
12234 Self::gdn_chunk_size());
12235 Ok(())
12236 }
12237
12238 pub fn gdn_glog(&self, alpha: &CudaSlice<f32>, dt_bias: &CudaSlice<f32>, a: &CudaSlice<f32>,
12240 g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
12241 -> Result<(), Box<dyn std::error::Error>> {
12242 let f = self.func("gdn_glog_f32");
12243 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
12244 let (h, ti) = (n_head as i32, t as i32);
12245 let __s_b = self.gpu.stream();
12246 let mut b = __s_b.launch_builder(&f);
12247 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
12248 unsafe { b.launch(cfg)?; }
12249 Ok(())
12250 }
12251
12252 pub fn sigmoid_v(&self, x: &cudarc::driver::CudaView<f32>, y: &mut CudaSlice<f32>, n: usize)
12255 -> Result<(), Box<dyn std::error::Error>> {
12256 let f = self.func("sigmoid_f32");
12257 let cfg = LaunchConfig::for_num_elems(n as u32);
12258 let ni = n as i32;
12259 let __s_b = self.gpu.stream();
12260 let mut b = __s_b.launch_builder(&f);
12261 b.arg(x).arg(y).arg(&ni);
12262 unsafe { b.launch(cfg)?; }
12263 Ok(())
12264 }
12265
12266 pub fn gdn_glog_v(&self, alpha: &cudarc::driver::CudaView<f32>, dt_bias: &CudaSlice<f32>,
12267 a: &CudaSlice<f32>, g_log: &mut CudaSlice<f32>, n_head: usize, t: usize)
12268 -> Result<(), Box<dyn std::error::Error>> {
12269 let f = self.func("gdn_glog_f32");
12270 let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
12271 let (h, ti) = (n_head as i32, t as i32);
12272 let __s_b = self.gpu.stream();
12273 let mut b = __s_b.launch_builder(&f);
12274 b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
12275 unsafe { b.launch(cfg)?; }
12276 Ok(())
12277 }
12278
12279 pub fn sigmoid(&self, x: &CudaSlice<f32>, y: &mut CudaSlice<f32>, n: usize)
12280 -> Result<(), Box<dyn std::error::Error>> {
12281 let f = self.func("sigmoid_f32");
12282 let cfg = LaunchConfig::for_num_elems(n as u32);
12283 let ni = n as i32;
12284 let __s_b = self.gpu.stream();
12285 let mut b = __s_b.launch_builder(&f);
12286 b.arg(x).arg(y).arg(&ni);
12287 unsafe { b.launch(cfg)?; }
12288 Ok(())
12289 }
12290
12291 pub fn sig_mul_f16out(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12294 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>, n: usize)
12295 -> Result<(), Box<dyn std::error::Error>> {
12296 let f = self.func("sig_mul_f16out_f32");
12297 let cfg = LaunchConfig::for_num_elems(n as u32);
12298 let ni = n as i32;
12299 let __s_b = self.gpu.stream();
12300 let mut b = __s_b.launch_builder(&f);
12301 b.arg(a).arg(g).arg(dst).arg(dst16).arg(&ni);
12302 unsafe { b.launch(cfg)?; }
12303 Ok(())
12304 }
12305
12306 #[allow(clippy::too_many_arguments)]
12315 pub fn attn_head_gate(&self, a: &CudaSlice<f32>, g: &CudaSlice<f32>,
12316 dst: &mut CudaSlice<f32>, dst16: Option<&mut CudaSlice<u8>>,
12317 head_dim: usize, n_head: usize, t: usize)
12318 -> Result<(), Box<dyn std::error::Error>> {
12319 let f = self.func("attn_head_gate_f32");
12320 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12321 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12322 let d16: u64 = match dst16 { Some(d) => self.addr_u8(d), None => 0 };
12324 let __s_b = self.gpu.stream();
12325 let mut b = __s_b.launch_builder(&f);
12326 b.arg(a).arg(g).arg(dst).arg(&d16).arg(&hd).arg(&nh).arg(&ti);
12327 unsafe { b.launch(cfg)?; }
12328 Ok(())
12329 }
12330
12331 #[allow(clippy::too_many_arguments)]
12340 pub fn swiglu_clamped_mul_scaled(&self, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
12341 gs: f32, us: f32, limit: f32,
12342 dst: &mut CudaSlice<f32>, n: usize)
12343 -> Result<(), Box<dyn std::error::Error>> {
12344 debug_assert!(limit > 1e-6, "swiglu_clamped needs a live limit; use silu_mul_scaled");
12345 let f = self.func("swiglu_clamped_mul_scaled_f32");
12346 let cfg = LaunchConfig::for_num_elems(n as u32);
12347 let ni = n as i32;
12348 let __s_b = self.gpu.stream();
12349 let mut b = __s_b.launch_builder(&f);
12350 b.arg(gate).arg(up).arg(&gs).arg(&us).arg(&limit).arg(dst).arg(&ni);
12351 unsafe { b.launch(cfg)?; }
12352 Ok(())
12353 }
12354
12355 pub fn gated_rmsnorm(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12357 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12358 -> Result<(), Box<dyn std::error::Error>> {
12359 let f = self.func("gated_rmsnorm_f32");
12360 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12361 let (nc, e) = (ncols as i32, eps);
12362 let __s_b = self.gpu.stream();
12363 let mut b = __s_b.launch_builder(&f);
12364 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12365 unsafe { b.launch(cfg)?; }
12366 Ok(())
12367 }
12368
12369 pub fn gated_rmsnorm_f16out(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12372 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12373 ncols: usize, nrows: usize, eps: f32)
12374 -> Result<(), Box<dyn std::error::Error>> {
12375 let f = self.func("gated_rmsnorm_f16out_f32");
12376 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12378 let (nc, e) = (ncols as i32, eps);
12379 let __s_b = self.gpu.stream();
12380 let mut b = __s_b.launch_builder(&f);
12381 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12382 unsafe { b.launch(cfg)?; }
12383 Ok(())
12384 }
12385
12386 #[allow(clippy::too_many_arguments)]
12390 pub fn add_rms_norm_zq8(&self, a: &CudaSlice<f32>, b_in: &CudaSlice<f32>, w: &CudaSlice<f32>,
12391 res: &mut CudaSlice<f32>, z: &mut CudaSlice<f32>,
12392 ncols: usize, nrows: usize, eps: f32)
12393 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12394 assert!(ncols % 32 == 0);
12395 let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
12396 let mut d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12397 let f = self.func("add_rms_norm_zq8");
12398 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (1024, 1, 1), shared_mem_bytes: 0 };
12399 let (nc, ep) = (ncols as i32, eps);
12400 let __s_b = self.gpu.stream();
12401 let mut b = __s_b.launch_builder(&f);
12402 b.arg(a).arg(b_in).arg(w).arg(res).arg(z).arg(&mut q).arg(&mut d).arg(&nc).arg(&ep);
12403 unsafe { b.launch(cfg)?; }
12404 Ok((q, d))
12405 }
12406
12407 pub fn gated_rmsnorm_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12412 z: &cudarc::driver::CudaView<f32>,
12413 dst: &mut CudaSlice<f32>, ncols: usize, nrows: usize, eps: f32)
12414 -> Result<(), Box<dyn std::error::Error>> {
12415 let f = self.func("gated_rmsnorm_f32");
12416 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12417 let (nc, e) = (ncols as i32, eps);
12418 let __s_b = self.gpu.stream();
12419 let mut b = __s_b.launch_builder(&f);
12420 b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
12421 unsafe { b.launch(cfg)?; }
12422 Ok(())
12423 }
12424
12425 pub fn gated_rmsnorm_f16out_zv(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>,
12426 z: &cudarc::driver::CudaView<f32>,
12427 dst: &mut CudaSlice<f32>, dst16: &mut CudaSlice<u8>,
12428 ncols: usize, nrows: usize, eps: f32)
12429 -> Result<(), Box<dyn std::error::Error>> {
12430 let f = self.func("gated_rmsnorm_f16out_f32");
12431 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12433 let (nc, e) = (ncols as i32, eps);
12434 let __s_b = self.gpu.stream();
12435 let mut b = __s_b.launch_builder(&f);
12436 b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
12437 unsafe { b.launch(cfg)?; }
12438 Ok(())
12439 }
12440
12441 pub fn gated_rmsnorm_q8_1(&self, o: &CudaSlice<f32>, w: &CudaSlice<f32>, z: &CudaSlice<f32>,
12442 ncols: usize, nrows: usize, eps: f32)
12443 -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12444 assert!(ncols % 32 == 0);
12445 let f = self.func("gated_rmsnorm_q8_1");
12446 let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
12447 let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
12448 let cfg = LaunchConfig { grid_dim: (nrows as u32, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0 };
12449 let (nc, ep) = (ncols as i32, eps);
12450 let __s_b = self.gpu.stream();
12451 let mut b = __s_b.launch_builder(&f);
12452 b.arg(o).arg(w).arg(z).arg(&mut out_q).arg(&mut out_d).arg(&nc).arg(&ep);
12453 unsafe { b.launch(cfg)?; }
12454 Ok((out_q, out_d))
12455 }
12456
12457 pub fn transpose(&self, inp: &CudaSlice<f32>, rows: usize, cols: usize)
12459 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12460 let f = self.func("transpose_f32");
12461 let mut out = self.zeros(rows * cols)?;
12462 let cfg = LaunchConfig::for_num_elems((rows * cols) as u32);
12463 let (r, c) = (rows as i32, cols as i32);
12464 let __s_b = self.gpu.stream();
12465 let mut b = __s_b.launch_builder(&f);
12466 b.arg(inp).arg(&mut out).arg(&r).arg(&c);
12467 unsafe { b.launch(cfg)?; }
12468 Ok(out)
12469 }
12470
12471 pub fn repeat_heads(&self, inp: &CudaSlice<f32>, out: &mut CudaSlice<f32>,
12473 head_dim: usize, n_in: usize, n_out: usize, t: usize)
12474 -> Result<(), Box<dyn std::error::Error>> {
12475 let f = self.func("repeat_heads_f32");
12476 let cfg = LaunchConfig::for_num_elems((head_dim * n_out * t) as u32);
12477 let (hd, ni, no, ti) = (head_dim as i32, n_in as i32, n_out as i32, t as i32);
12478 let __s_b = self.gpu.stream();
12479 let mut b = __s_b.launch_builder(&f);
12480 b.arg(inp).arg(out).arg(&hd).arg(&ni).arg(&no).arg(&ti);
12481 unsafe { b.launch(cfg)?; }
12482 Ok(())
12483 }
12484
12485 pub fn q_gate_split(&self, qf: &CudaSlice<f32>, q_out: &mut CudaSlice<f32>,
12488 gate_out: &mut CudaSlice<f32>, head_dim: usize, n_head: usize, t: usize)
12489 -> Result<(), Box<dyn std::error::Error>> {
12490 let f = self.func("q_gate_split_f32");
12491 let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
12492 let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
12493 let __s_b = self.gpu.stream();
12494 let mut b = __s_b.launch_builder(&f);
12495 b.arg(qf).arg(q_out).arg(gate_out).arg(&hd).arg(&nh).arg(&ti);
12496 unsafe { b.launch(cfg)?; }
12497 Ok(())
12498 }
12499
12500 pub fn qkv_to_gdn_repack(&self, conv_out: &CudaSlice<f32>, q_g: &mut CudaSlice<f32>,
12504 k_g: &mut CudaSlice<f32>, v_g: &mut CudaSlice<f32>,
12505 d_state: usize, num_v: usize, num_k: usize, key_dim: usize, t: usize)
12506 -> Result<(), Box<dyn std::error::Error>> {
12507 let f = self.func("qkv_to_gdn_repack_f32");
12508 let cfg = LaunchConfig::for_num_elems((d_state * num_v * t) as u32);
12509 let (ds, nv, nk, kd, ti) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32, t as i32);
12510 let __s_b = self.gpu.stream();
12511 let mut b = __s_b.launch_builder(&f);
12512 b.arg(conv_out).arg(q_g).arg(k_g).arg(v_g).arg(&ds).arg(&nv).arg(&nk).arg(&kd).arg(&ti);
12513 unsafe { b.launch(cfg)?; }
12514 Ok(())
12515 }
12516
12517 pub fn conv_left_pad(&self, src: &CudaSlice<f32>, dst: &mut CudaSlice<f32>,
12520 conv_dim: usize, t: usize, pad: usize)
12521 -> Result<(), Box<dyn std::error::Error>> {
12522 let f = self.func("conv_left_pad_f32");
12523 let cfg = LaunchConfig::for_num_elems((conv_dim * t) as u32);
12524 let (cd, ti, p) = (conv_dim as i32, t as i32, pad as i32);
12525 let __s_b = self.gpu.stream();
12526 let mut b = __s_b.launch_builder(&f);
12527 b.arg(src).arg(dst).arg(&cd).arg(&ti).arg(&p);
12528 unsafe { b.launch(cfg)?; }
12529 Ok(())
12530 }
12531
12532 pub fn conv_assemble_and_roll(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12536 conv_in: &mut CudaSlice<f32>, conv_dim: usize, pad: usize)
12537 -> Result<(), Box<dyn std::error::Error>> {
12538 let f = self.func("conv_assemble_and_roll_f32");
12539 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12540 let (cd, p) = (conv_dim as i32, pad as i32);
12541 let __s_b = self.gpu.stream();
12542 let mut b = __s_b.launch_builder(&f);
12543 b.arg(qkv_col).arg(conv_state).arg(conv_in).arg(&cd).arg(&p);
12544 unsafe { b.launch(cfg)?; }
12545 Ok(())
12546 }
12547
12548 pub fn ssm_conv1d_fused_decode(&self, qkv_col: &CudaSlice<f32>, conv_state: &mut CudaSlice<f32>,
12554 w: &CudaSlice<f32>, conv_out: &mut CudaSlice<f32>,
12555 conv_dim: usize, d_conv: usize)
12556 -> Result<(), Box<dyn std::error::Error>> {
12557 let f = self.func("ssm_conv1d_fused_decode_f32");
12558 let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
12559 let (cd, dc) = (conv_dim as i32, d_conv as i32);
12560 let __s_b = self.gpu.stream();
12561 let mut b = __s_b.launch_builder(&f);
12562 b.arg(qkv_col).arg(conv_state).arg(w).arg(conv_out).arg(&cd).arg(&dc);
12563 unsafe { b.launch(cfg)?; }
12564 Ok(())
12565 }
12566
12567 pub fn slice_range(&self, src: &CudaSlice<f32>, start: usize, len: usize)
12570 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12571 let host = self.gpu.stream().clone_dtoh(src)?;
12572 self.gpu.stream().synchronize()?;
12573 Ok(self.htod(&host[start..start + len])?)
12574 }
12575}
12576
12577#[cfg(test)]
12578mod target_dispatch_tests {
12579 use super::legacy_quant_gemm_allowed;
12580
12581 #[test]
12582 fn legacy_quant_gemm_arch_policy_honors_the_escape_hatch() {
12583 assert!(legacy_quant_gemm_allowed(false, false, false));
12585 assert!(!legacy_quant_gemm_allowed(false, false, true));
12586 assert!(!legacy_quant_gemm_allowed(true, false, false));
12588 assert!(!legacy_quant_gemm_allowed(true, false, true));
12589 assert!(legacy_quant_gemm_allowed(true, true, false));
12591 assert!(!legacy_quant_gemm_allowed(true, true, true));
12592 }
12593
12594 #[cfg(all(memra_portable_cuda, not(memra_hopper_mma)))]
12595 #[test]
12596 fn portable_build_disables_legacy_quant_gemm_without_an_env_override() {
12597 assert!(!legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12598 }
12599
12600 #[cfg(memra_hopper_mma)]
12601 #[test]
12602 fn hopper_mma_build_re_admits_legacy_quant_gemm() {
12603 assert!(legacy_quant_gemm_allowed(cfg!(memra_portable_cuda), cfg!(memra_hopper_mma), false));
12604 assert!(super::portable_mma_gated() == false);
12605 }
12606}
12607
12608impl memra_kv::KvDev for Engine {
12611 fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12612 Engine::zeros(self, n)
12613 }
12614 fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12615 Engine::uninit(self, n)
12616 }
12617 fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
12618 Engine::alloc_u8(self, n)
12619 }
12620 fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
12621 Engine::htod_i32(self, v)
12622 }
12623 fn clone_dtod(&self, src: &CudaSlice<f32>) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12624 Engine::clone_dtod(self, src)
12625 }
12626 fn copy_into(&self, dst: &mut CudaSlice<f32>, off: usize, src: &CudaSlice<f32>, len: usize)
12627 -> Result<(), Box<dyn std::error::Error>> {
12628 Engine::copy_into(self, dst, off, src, len)
12629 }
12630 fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32) -> Result<(), Box<dyn std::error::Error>> {
12631 Engine::set_i32_one(self, d, v)
12632 }
12633}