use cudarc::driver::{
CudaContext, CudaFunction, CudaModule, CudaSlice, CudaStream, LaunchConfig, PushKernelArg,
};
use cudarc::nvrtc::Ptx;
use std::sync::{Arc, Mutex};
#[cfg(debug_assertions)]
pub(crate) fn debug_assert_tensor_stream_device<T>(
tensor: &CudaSlice<T>,
stream: &CudaStream,
site: &str,
) {
let tensor_dev = tensor.ordinal();
let stream_dev = stream.context().ordinal();
assert_eq!(
tensor_dev, stream_dev,
"PP cross-device tensor read at {site}: tensor on dev{tensor_dev}, stream on dev{stream_dev}"
);
}
pub use memra_gguf;
pub use memra_runtime;
pub mod forward;
pub mod hybrid;
pub mod hybrid_forward;
pub mod model;
pub mod sigrouter_contract;
pub mod cache {
pub use memra_kv::*;
}
pub mod decode;
pub mod decode_batch;
pub mod dflash;
pub mod eagle;
pub mod gemma_spec;
pub mod graph_update;
pub mod mla;
pub mod moesd;
pub mod pp;
pub mod round_stream;
pub mod spec;
pub use memra_sampling as sampler;
pub fn moe_f16g_mode() -> u8 {
static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
*M.get_or_init(|| match std::env::var("MEMRA_MOE_F16G").as_deref() {
Ok("0") => 0,
Ok("2") => 2,
Ok("3") => 3,
Ok(_) => 1,
Err(_) => 2,
})
}
pub fn moe_f16g_sk_params() -> (i32, i32) {
static P: std::sync::OnceLock<(i32, i32)> = std::sync::OnceLock::new();
*P.get_or_init(|| match std::env::var("MEMRA_F16G_SK").as_deref() {
Ok("0") => (-1, 0),
Ok("32") => (0, i32::MAX),
Ok("128") => (0, 1),
_ => {
let cross = std::env::var("MEMRA_F16G_SK_CROSS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(64);
(0, cross)
}
})
}
pub fn moe_f16g_direct_on(qtype: i32) -> bool {
static M: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
let m = *M.get_or_init(|| match std::env::var("MEMRA_F16G_DIRECT").as_deref() {
Ok("0") => 0,
Ok("kq") => 1,
_ => 2,
});
match m {
0 => false,
1 => qtype == QT_Q4_K || qtype == QT_Q6_K,
_ => true,
}
}
pub fn moe_f16g_tail_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_F16G_TAIL").as_deref() != Ok("0"))
}
pub fn moe_f16g_gemma_on() -> bool {
static M: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*M.get_or_init(|| !matches!(std::env::var("MEMRA_MOE_F16G").as_deref(), Ok("0") | Err(_)))
}
pub fn moe_fuse_actq_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_MOE_FUSE_ACTQ").as_deref() != Ok("0"))
}
pub fn router_prefill_exact_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_ROUTER_PREFILL_EXACT").as_deref() != Ok("0"))
}
pub fn router_kernel_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
let on = std::env::var("MEMRA_ROUTER_KERNEL").as_deref() != Ok("0");
if !on {
eprintln!("[memra] router kernel OFF (rollback: per-column cuBLAS gemv)");
}
on
})
}
pub const ROUTER_BATCH_MIN_T: usize = 8;
pub fn router_batch_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_ROUTER_BATCH").as_deref() != Ok("0"))
}
mod cpu_experts;
#[cfg(memra_cutlass)]
pub mod cutlass_ffi;
pub mod f16_ffi;
pub mod fp8_ffi;
pub mod mmq_ffi;
pub mod moe_cache;
pub mod prime_graph;
pub mod spill;
mod spill_pread;
const FATBIN: &[u8] = include_bytes!(env!("MEMRA_ENGINE_FATBIN"));
const HYBRID_FATBIN: &[u8] = include_bytes!(env!("MEMRA_HYBRID_FATBIN"));
const QMATVEC_FATBIN: &[u8] = include_bytes!(env!("MEMRA_QMATVEC_FATBIN"));
const FLASH_FATBIN: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN"));
const GEMM_FATBIN: &[u8] = include_bytes!(env!("MEMRA_GEMM_FATBIN"));
const ROUTER_FATBIN: &[u8] = include_bytes!(env!("MEMRA_ROUTER_FATBIN"));
const SAMPLE_FATBIN: &[u8] = include_bytes!(env!("MEMRA_SAMPLE_FATBIN"));
fn gemm_fatbin_bytes() -> std::borrow::Cow<'static, [u8]> {
assert!(
!(portable_mma_gated() && std::env::var_os("MEMRA_GEMM_FATBIN").is_some()),
"MEMRA_GEMM_FATBIN overrides are not allowed in the portable CUDA lane"
);
match std::env::var("MEMRA_GEMM_FATBIN") {
Ok(path) => std::borrow::Cow::Owned(
std::fs::read(&path).unwrap_or_else(|e| panic!("MEMRA_GEMM_FATBIN read {path}: {e}")),
),
Err(_) => std::borrow::Cow::Borrowed(GEMM_FATBIN),
}
}
pub(crate) const fn portable_mma_gated() -> bool {
cfg!(memra_portable_cuda) && !cfg!(memra_hopper_mma)
}
const fn legacy_quant_gemm_allowed(portable_cuda: bool, hopper_mma: bool, no_gemm: bool) -> bool {
(!portable_cuda || hopper_mma) && !no_gemm
}
const FLASH_FATBIN_VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VQ4"));
const FLASH_FATBIN_VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_VF8"));
const FLASH_FATBIN_KF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8"));
const FLASH_FATBIN_KF8VQ4: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VQ4"));
const FLASH_FATBIN_KF8VF8: &[u8] = include_bytes!(env!("MEMRA_FLASH_FATBIN_KF8VF8"));
pub use memra_kv::{kv_blk_bytes, kv_cache_formats};
fn flash_fatbin_bytes() -> &'static [u8] {
match kv_cache_formats() {
("q8_0", "q5_1") => FLASH_FATBIN,
("q8_0", "q4_0") => FLASH_FATBIN_VQ4,
("q8_0", "fp8") => FLASH_FATBIN_VF8,
("fp8", "q5_1") => FLASH_FATBIN_KF8,
("fp8", "q4_0") => FLASH_FATBIN_KF8VQ4,
("fp8", "fp8") => FLASH_FATBIN_KF8VF8,
other => unreachable!("kv_cache_formats returned {other:?}"),
}
}
fn k1_launch_override() -> Option<(u32, u32, u32)> {
static K1: std::sync::OnceLock<Option<(u32, u32, u32)>> = std::sync::OnceLock::new();
*K1.get_or_init(|| {
let v = std::env::var("MEMRA_GEMM_K1_LAUNCH").ok()?;
let p: Vec<u32> = v.split(',').filter_map(|s| s.trim().parse().ok()).collect();
match p.as_slice() {
[bm, bn, w] => Some((*bm, *bn, *w)),
_ => None,
}
})
}
pub(crate) fn wgmma_gemm_enabled() -> bool {
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*V.get_or_init(|| std::env::var("MEMRA_WGMMA").as_deref() == Ok("1"))
}
pub const FA_VEC_MIN_TKV: usize = 96;
pub fn fa_vec_min_tkv() -> usize {
static V: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*V.get_or_init(|| {
std::env::var("MEMRA_FA_VEC_MIN")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or_else(|| FA_VEC_MIN_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
})
}
pub fn fa_f16pv_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_FA_F16PV")
.map(|v| v != "0")
.unwrap_or_else(|_| std::env::var("MEMRA_DRAFT").is_err())
})
}
pub fn fa512_hp_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_FA512_HP").as_deref() != Ok("0"))
}
pub fn faw_hp_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_FAW_HP").as_deref() != Ok("0"))
}
pub fn fa512_wide_warps() -> usize {
static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*N.get_or_init(|| match std::env::var("MEMRA_FA512_W4").as_deref() {
Ok("1") => 4,
_ => 2,
})
}
pub fn fa512_min_tkv() -> usize {
static FA512_MIN: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*FA512_MIN.get_or_init(|| {
std::env::var("MEMRA_FA512_MIN")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(512)
})
}
pub static FA_VEC_MIN_DEFAULT: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(FA_VEC_MIN_TKV);
pub static FA_SPW_DEFAULT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(32);
pub static FUSED_MR1_DEFAULT: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
pub static ROUTER_W8_DEFAULT: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(true);
pub static FA_SP512_DEFAULT: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(16);
pub static RMS_BLOCK_DEFAULT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(256);
pub static FA_SP_GEMMA: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
pub static MMQ_SK_FORCE: std::sync::atomic::AtomicI8 = std::sync::atomic::AtomicI8::new(-1);
pub use memra_kv::KV_FP8_FORCE;
pub(crate) fn rms_block() -> u32 {
static V: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
*V.get_or_init(|| {
std::env::var("MEMRA_RMS_BLOCK")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or_else(|| RMS_BLOCK_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
})
}
pub(crate) fn fa_split_keys(t_kv: usize, n_head_kv: usize) -> usize {
static S: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
if let Some(forced) = *S.get_or_init(|| {
std::env::var("MEMRA_FA_SPLIT")
.ok()
.and_then(|v| v.parse().ok())
.filter(|&s: &usize| s >= 8 && s % 8 == 0)
}) {
return forced;
}
if FA_SP_GEMMA.load(std::sync::atomic::Ordering::Relaxed)
&& std::env::var("MEMRA_FA_SP16").as_deref() == Ok("1")
{
return if t_kv <= 8192 {
16
} else if t_kv <= 16384 {
64
} else {
128
};
}
let big_rig = fa_sm_count() >= 128;
if big_rig {
let _ = n_head_kv;
if t_kv <= 2048 {
16
} else if t_kv <= 16384 {
64
} else {
128
}
} else if n_head_kv <= 4 {
if t_kv <= 512 {
8
} else if t_kv <= 16384 {
64
} else {
128
}
} else {
if t_kv <= 8192 {
32
} else if t_kv <= 16384 {
64
} else {
128
}
}
}
fn fa_sm_count() -> i32 {
static N: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
*N.get_or_init(|| {
cudarc::driver::result::init().ok();
cudarc::driver::result::device::get(0)
.and_then(|d| unsafe { cudarc::driver::result::device::get_attribute(
d, cudarc::driver::sys::CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT) })
.unwrap_or(82)
})
}
fn fa_hd_suffix(head_dim: usize) -> Result<&'static str, Box<dyn std::error::Error>> {
match head_dim {
256 => Ok(""),
128 => Ok("_hd128"),
d => Err(format!(
"fa_prefill: no kernel stamped for head_dim={d} (only 256/128); \
callers must gate to sdpa_naive"
)
.into()),
}
}
pub const QT_Q8_0: i32 = 0;
pub const QT_Q4_K: i32 = 1;
pub const QT_Q6_K: i32 = 2;
pub const QT_Q5_K: i32 = 3;
pub const QT_Q3_K: i32 = 4;
pub const QT_IQ4_XS: i32 = 5;
pub const QT_IQ3_S: i32 = 6;
pub const QT_NVFP4: i32 = 7;
pub const QT_F8_E4M3: i32 = 10;
pub const QT_NVFP4_RP: i32 = 9;
pub const QT_F32: i32 = 8;
pub const QT_BF16: i32 = 11;
pub const QT_Q4_0: i32 = 12; pub const QT_Q2_K: i32 = 13;
pub const QT_F8_E4M3_BLK: i32 = 14;
pub struct Engine {
pub gpu: memra_runtime::Gpu,
module: Arc<CudaModule>,
hybrid: Arc<CudaModule>,
qmatvec: Arc<CudaModule>,
flash: Arc<CudaModule>,
flash_g: std::sync::OnceLock<Arc<CudaModule>>,
gemm: Arc<CudaModule>,
router: Arc<CudaModule>,
sample: Arc<CudaModule>,
moe_cache: Mutex<Option<crate::moe_cache::MoeSlotCache>>,
moe_cache_layout: Mutex<Option<Vec<usize>>>,
capture_keep_on: std::sync::atomic::AtomicBool,
verify_exact: std::sync::atomic::AtomicBool,
capture_keep: Mutex<Vec<Box<dyn std::any::Any + Send>>>,
pub copy_stream: Arc<CudaStream>,
#[cfg(memra_cutlass)]
cutlass_scratch: Mutex<Option<crate::cutlass_ffi::CutlassScratch>>,
fp8_scratch: Mutex<Option<crate::fp8_ffi::Fp8Scratch>>,
fa_vf16_scratch: Mutex<Option<CudaSlice<u8>>>,
fa_part_pool: Mutex<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
fa_part_retired: Mutex<Vec<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>>,
fn_cache: Mutex<std::collections::HashMap<String, CudaFunction>>,
f16_scratch: Mutex<Option<crate::f16_ffi::F16Scratch>>,
argmax_partials: Mutex<Option<(CudaSlice<f32>, CudaSlice<i32>)>>,
prime_deqw_ws: Mutex<Option<(CudaSlice<u8>, CudaSlice<u8>)>>,
router_stage: Mutex<Option<PinnedStage>>,
}
fn fa_v2_on() -> bool {
std::env::var("MEMRA_FA_V2")
.map(|v| v != "0")
.unwrap_or(true)
}
fn fa_v3_on() -> bool {
std::env::var("MEMRA_FA_V3")
.map(|v| v != "0")
.unwrap_or(true)
}
fn fa_v4_mode() -> &'static str {
static M: std::sync::OnceLock<String> = std::sync::OnceLock::new();
M.get_or_init(|| std::env::var("MEMRA_FA_V4").unwrap_or_default())
}
fn fa_v4_on() -> bool {
fa_v4_mode() != "0"
} pub static FA_SMEM_TKV_DEFAULT: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(1024);
pub static FA_V4_MAX_DEFAULT: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(usize::MAX);
pub fn fa_v4_at_pub(t_kv: usize) -> bool {
fa_v4_at(t_kv)
}
fn fa_v4_at(t_kv: usize) -> bool {
static M: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
let mx = *M.get_or_init(|| {
std::env::var("MEMRA_FA_V4_MAX")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or_else(|| FA_V4_MAX_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
});
fa_v4_on() && t_kv < mx
}
pub const FA_DEEP_MIN_DEFAULT: usize = 0;
fn fa_deep_at(t_kv: usize) -> bool {
if std::env::var("MEMRA_FA_DEEP").as_deref() == Ok("0") {
return false;
}
let min = std::env::var("MEMRA_FA_DEEP_MIN")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(FA_DEEP_MIN_DEFAULT);
t_kv >= min
}
pub fn fa_deep_at_pub(t_kv: usize) -> bool {
fa_deep_at(t_kv)
}
fn fa_v3_active(head_dim: usize) -> bool {
fa_v3_on()
&& head_dim % 128 == 0
&& kv_cache_formats() == ("q8_0", "q5_1")
&& !Engine::kv_fp8_on()
}
pub fn fa_seqs_eligible(t_kv: usize, head_dim: usize) -> bool {
std::env::var("MEMRA_NO_FA_VEC").is_err()
&& t_kv >= fa_vec_min_tkv()
&& head_dim == 256
&& fa_v4_at(t_kv)
&& !matches!(fa_v4_mode(), "noB3" | "stage")
&& !Engine::kv_fp8_on()
}
pub fn fa_split_keys_pub(t_kv: usize, n_head_kv: usize) -> usize {
fa_split_keys(t_kv, n_head_kv)
}
struct PinnedStage {
ptr: *mut u8,
cap: usize,
}
unsafe impl Send for PinnedStage {}
impl PinnedStage {
fn new(cap: usize) -> Result<Self, Box<dyn std::error::Error>> {
let ptr = unsafe { cudarc::driver::result::malloc_host(cap, 0)? } as *mut u8;
Ok(PinnedStage { ptr, cap })
}
}
impl Drop for PinnedStage {
fn drop(&mut self) {
let _ = unsafe { cudarc::driver::result::free_host(self.ptr as _) };
}
}
pub const ARGMAX_NB: usize = 256;
pub(crate) use memra_fa3_vl as fa3_vl_raw;
unsafe extern "C" {
fn memra_fa3_prefill(
q16: *const core::ffi::c_void,
k16: *const core::ffi::c_void,
v16: *const core::ffi::c_void,
o: *mut f32,
t: i32,
h: i32,
hkv: i32,
d: i32,
scale: f32,
stream: *mut core::ffi::c_void,
) -> i32;
pub(crate) fn memra_fa3_vl(
q16s: *const *const core::ffi::c_void,
k16s: *const *const core::ffi::c_void,
v16s: *const *const core::ffi::c_void,
os: *const *mut f32,
ts: *const i32,
b: i32,
h: i32,
hkv: i32,
d: i32,
scale: f32,
stream: *mut core::ffi::c_void,
) -> i32;
}
#[repr(C)]
#[derive(Clone, Copy)]
pub struct WPtr8(pub [u64; 8]);
unsafe impl cudarc::driver::DeviceRepr for WPtr8 {}
#[repr(C)]
#[derive(Clone, Copy, Default)]
pub struct GdnSeqVl {
pub kb16: u64,
pub gcum: u64,
pub beta: u64,
pub u: u64,
pub wb16: u64,
pub y: u64,
pub ssnap: u64,
pub state_in: u64,
pub state_out: u64,
pub q: u64,
pub p: u64,
pub o: u64,
pub k: u64,
pub v: u64,
pub g: u64,
pub a: u64,
pub w: u64,
pub t: i32,
pub nc: i32,
}
unsafe impl cudarc::driver::DeviceRepr for GdnSeqVl {}
#[repr(C)]
#[derive(Clone, Copy)]
pub struct GdnVl8(pub [GdnSeqVl; 8]);
unsafe impl cudarc::driver::DeviceRepr for GdnVl8 {}
#[repr(C)]
#[derive(Clone, Copy, Default)]
pub struct GdnWVl {
pub qb16: u64,
pub pb16: u64,
}
unsafe impl cudarc::driver::DeviceRepr for GdnWVl {}
#[repr(C)]
#[derive(Clone, Copy)]
pub struct GdnWVl8(pub [GdnWVl; 8]);
unsafe impl cudarc::driver::DeviceRepr for GdnWVl8 {}
#[repr(C)]
#[derive(Clone, Copy, Default)]
pub struct GdnPrepVl {
pub qkv: u64,
pub conv_state: u64,
pub conv_out: u64,
pub q_g: u64,
pub k_g: u64,
pub v_g: u64,
pub q_l2: u64,
pub k_l2: u64,
pub beta_raw: u64,
pub alpha: u64,
pub beta: u64,
pub g_log: u64,
pub o: u64,
pub z: u64,
pub gn: u64,
pub gn16: u64,
pub kb16: u64,
pub qb16: u64,
pub t: i32,
pub pad: i32,
}
unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl {}
#[repr(C)]
#[derive(Clone, Copy)]
pub struct GdnPrepVl8(pub [GdnPrepVl; 8]);
unsafe impl cudarc::driver::DeviceRepr for GdnPrepVl8 {}
#[repr(C)]
#[derive(Clone, Copy, Default)]
pub struct FaSeqVl {
pub q: u64,
pub k16: u64,
pub v16: u64,
pub o: u64,
pub kf: u64,
pub vf: u64,
pub t: i32,
pub pad: i32,
}
unsafe impl cudarc::driver::DeviceRepr for FaSeqVl {}
#[repr(C)]
#[derive(Clone, Copy)]
pub struct FaVl8(pub [FaSeqVl; 8]);
unsafe impl cudarc::driver::DeviceRepr for FaVl8 {}
#[repr(C)]
#[derive(Clone, Copy, Default)]
pub struct AttnPreVl {
pub qf: u64,
pub kf: u64,
pub vf: u64,
pub q: u64,
pub gate: u64,
pub qn: u64,
pub kn: u64,
pub kc: u64,
pub vc: u64,
pub t: i32,
pub pad: i32,
}
unsafe impl cudarc::driver::DeviceRepr for AttnPreVl {}
#[repr(C)]
#[derive(Clone, Copy)]
pub struct AttnPreVl8(pub [AttnPreVl; 8]);
unsafe impl cudarc::driver::DeviceRepr for AttnPreVl8 {}
pub struct GdnChunkBufs {
pub gcum: CudaSlice<f32>,
pub a: CudaSlice<f32>,
pub p: CudaSlice<f32>,
pub u: CudaSlice<f32>,
pub w: CudaSlice<f32>,
pub kb16: CudaSlice<u8>,
pub wb16: CudaSlice<u8>,
pub y16: CudaSlice<u8>,
pub ssnap16: CudaSlice<u8>,
pub qb16: CudaSlice<u8>,
pub pb16: CudaSlice<u8>,
pub o: CudaSlice<f32>,
pub t: usize,
pub nc: usize,
}
#[repr(C)]
#[derive(Clone, Copy)]
pub struct F32x8(pub [f32; 8]);
unsafe impl cudarc::driver::DeviceRepr for F32x8 {}
pub static PRIME_NANOS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
impl Engine {
pub fn new(ordinal: usize) -> Result<Self, Box<dyn std::error::Error>> {
let gpu = memra_runtime::Gpu::new(ordinal)?;
if std::env::var("MEMRA_ARCH_CHECK").as_deref() != Ok("0") {
use cudarc::driver::sys::CUdevice_attribute_enum as A;
let (maj, min) = cudarc::driver::result::device::get(ordinal as i32)
.and_then(|d| unsafe {
Ok((
cudarc::driver::result::device::get_attribute(
d,
A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR,
)?,
cudarc::driver::result::device::get_attribute(
d,
A::CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR,
)?,
))
})
.unwrap_or((0, 0));
let built = env!("MEMRA_BUILT_CUDA_ARCH");
let ok = matches!(
(built, maj, min),
("120a", 12, 0) | ("120a", 12, 1) | ("100a", 10, 0) | ("90a", 9, 0) | ("89", 8, 9)
);
if !ok {
return Err(format!(
"memra was built for sm_{built} but device {ordinal} reports compute \
capability {maj}.{min}. Rebuild on this machine (MEMRA_CUDA_ARCH \
auto-detects the GPU) or set MEMRA_ARCH_CHECK=0 to bypass."
)
.into());
}
}
unsafe {
use cudarc::driver::sys;
let dev: sys::CUdevice = ordinal as sys::CUdevice;
let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
if sys::cuDeviceGetDefaultMemPool(&mut pool, dev) == sys::CUresult::CUDA_SUCCESS {
let mut thresh: u64 = u64::MAX;
let _ = sys::cuMemPoolSetAttribute(
pool,
sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
&mut thresh as *mut u64 as *mut core::ffi::c_void,
);
}
}
let module = gpu.ctx.load_module(Ptx::from_binary(FATBIN.to_vec()))?;
let hybrid = gpu
.ctx
.load_module(Ptx::from_binary(HYBRID_FATBIN.to_vec()))?;
let qmatvec = gpu
.ctx
.load_module(Ptx::from_binary(QMATVEC_FATBIN.to_vec()))?;
let flash = gpu
.ctx
.load_module(Ptx::from_binary(flash_fatbin_bytes().to_vec()))?;
let gemm = gpu
.ctx
.load_module(Ptx::from_binary(gemm_fatbin_bytes().into_owned()))?;
let router = gpu
.ctx
.load_module(Ptx::from_binary(ROUTER_FATBIN.to_vec()))?;
let sample = gpu
.ctx
.load_module(Ptx::from_binary(SAMPLE_FATBIN.to_vec()))?;
let copy_stream = gpu.ctx.new_stream()?;
if std::env::var("MEMRA_EVT")
.map(|v| v == "1")
.unwrap_or(false)
{
} else {
unsafe {
gpu.ctx.disable_event_tracking();
}
}
Ok(Self {
gpu,
module,
hybrid,
qmatvec,
flash,
flash_g: std::sync::OnceLock::new(),
gemm,
router,
sample,
moe_cache: Mutex::new(None),
moe_cache_layout: Mutex::new(None),
copy_stream,
capture_keep_on: std::sync::atomic::AtomicBool::new(false),
verify_exact: std::sync::atomic::AtomicBool::new(false),
capture_keep: Mutex::new(Vec::new()),
argmax_partials: Mutex::new(None),
prime_deqw_ws: Mutex::new(None),
router_stage: Mutex::new(None),
fp8_scratch: Mutex::new(None),
fa_vf16_scratch: Mutex::new(None),
fa_part_pool: Mutex::new(None),
fa_part_retired: Mutex::new(Vec::new()),
fn_cache: Mutex::new(Default::default()),
f16_scratch: Mutex::new(None),
#[cfg(memra_cutlass)]
cutlass_scratch: Mutex::new(None),
})
}
pub fn ctx(&self) -> &Arc<CudaContext> {
&self.gpu.ctx
}
pub fn pool_cached_bytes(&self) -> usize {
let (reserved, used) = self.pool_reserved_used();
reserved.saturating_sub(used)
}
pub fn pool_reserved_used(&self) -> (usize, usize) {
use cudarc::driver::sys;
unsafe {
let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
if sys::cuDeviceGetDefaultMemPool(&mut pool, self.gpu.ctx.ordinal() as sys::CUdevice)
!= sys::CUresult::CUDA_SUCCESS
{
return (0, 0);
}
let (mut reserved, mut used) = (0u64, 0u64);
if sys::cuMemPoolGetAttribute(
pool,
sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_RESERVED_MEM_CURRENT,
&mut reserved as *mut u64 as *mut core::ffi::c_void,
) != sys::CUresult::CUDA_SUCCESS
{
return (0, 0);
}
if sys::cuMemPoolGetAttribute(
pool,
sys::CUmemPool_attribute_enum::CU_MEMPOOL_ATTR_USED_MEM_CURRENT,
&mut used as *mut u64 as *mut core::ffi::c_void,
) != sys::CUresult::CUDA_SUCCESS
{
return (0, 0);
}
(reserved as usize, used as usize)
}
}
pub fn stream(&self) -> Arc<CudaStream> {
self.gpu.stream()
}
pub fn gkv_on() -> bool {
memra_kv::gkv_on()
}
pub fn wkv_on() -> bool {
memra_kv::wkv_on()
}
pub fn kv_fp8_on() -> bool {
memra_kv::kv_fp8_on()
}
fn fa_func(&self, name: &str, head_dim: usize) -> CudaFunction {
if head_dim == 512 && Self::gkv_on() {
self.func_g(name)
} else {
self.func(name)
}
}
fn func_g(&self, name: &str) -> CudaFunction {
let m = self.flash_g.get_or_init(|| {
self.gpu
.ctx
.load_module(cudarc::nvrtc::Ptx::from_binary(
FLASH_FATBIN_KF8VF8.to_vec(),
))
.expect("load kf8vf8 flash fatbin (fp8-globals arm)")
});
let key = format!("g:{name}");
if let Some(f) = self.fn_cache.lock().unwrap().get(&key) {
return f.clone();
}
let f = match m.load_function(name) {
Ok(f) => f,
Err(_) => self.func(name),
};
self.fn_cache.lock().unwrap().insert(key, f.clone());
f
}
fn func(&self, name: &str) -> CudaFunction {
if let Some(f) = self.fn_cache.lock().unwrap().get(name) {
return f.clone();
}
let f = self
.module
.load_function(name)
.or_else(|_| self.hybrid.load_function(name))
.or_else(|_| self.qmatvec.load_function(name))
.or_else(|_| self.flash.load_function(name))
.or_else(|_| self.gemm.load_function(name))
.or_else(|_| self.router.load_function(name))
.or_else(|_| self.sample.load_function(name))
.unwrap_or_else(|_| panic!("kernel {name} not in any fatbin"));
self.fn_cache
.lock()
.unwrap()
.insert(name.to_string(), f.clone());
f
}
pub fn scatter_trim_logits(
&self,
src: &CudaSlice<f32>,
d2t: &CudaSlice<u32>,
dst: &mut CudaSlice<f32>,
d_vocab: usize,
n_vocab: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f1 = self.func("scatter_trim_logits_f32");
let f2 = self.func("scatter_trim_logits_pass2_f32");
let (dv, nv) = (d_vocab as i32, n_vocab as i32);
let cfg1 = LaunchConfig {
grid_dim: (256, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b1 = self.gpu.stream();
let mut b1 = __s_b1.launch_builder(&f1);
b1.arg(src).arg(d2t).arg(&mut *dst).arg(&dv).arg(&nv);
unsafe {
b1.launch(cfg1)?;
}
let cfg2 = LaunchConfig {
grid_dim: (d_vocab.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f2);
b2.arg(src).arg(d2t).arg(&mut *dst).arg(&dv);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn filter_stats(
&self,
x: &CudaSlice<f32>,
row_stride: usize,
rows: &CudaSlice<i32>,
out_th: &mut CudaSlice<f32>,
out_z: &mut CudaSlice<f32>,
out_max: &mut CudaSlice<f32>,
n: usize,
nrow: usize,
temp: f32,
top_k: i32,
top_p: f32,
min_p: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("filter_stats_f32");
let (ni, nr, rs) = (n as i32, nrow as i32, row_stride as i64);
let cfg = LaunchConfig {
grid_dim: (nrow as u32, 1, 1),
block_dim: (1024, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(&rs)
.arg(rows)
.arg(&mut *out_th)
.arg(&mut *out_z)
.arg(&mut *out_max)
.arg(&ni)
.arg(&nr)
.arg(&temp)
.arg(&top_k)
.arg(&top_p)
.arg(&min_p);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn softmax_gather_filtered(
&self,
x: &CudaSlice<f32>,
row_stride: usize,
ids: &CudaSlice<u32>,
rows: &CudaSlice<i32>,
th: &CudaSlice<f32>,
z: &CudaSlice<f32>,
out: &mut CudaSlice<f32>,
n: usize,
npair: usize,
temp: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("softmax_gather_filtered_f32");
let (ni, np, rs) = (n as i32, npair as i32, row_stride as i64);
let cfg = LaunchConfig {
grid_dim: (npair as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(&rs)
.arg(ids)
.arg(rows)
.arg(th)
.arg(z)
.arg(&mut *out)
.arg(&ni)
.arg(&np)
.arg(&temp);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn residual_sample_filtered(
&self,
p: &CudaSlice<f32>,
q: Option<&CudaSlice<f32>>,
n: usize,
temp: f32,
seed: u64,
stream_pos: u32,
p_stats: (f32, f32, f32),
q_stats: (f32, f32, f32),
out_tok: &mut CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("residual_sample_filtered_f32");
let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
let has_q: i32 = q.is_some() as i32;
let qbuf = q.unwrap_or(p);
let (pm, pth, pz) = p_stats;
let (qm, qth, qz) = q_stats;
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1024, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(p)
.arg(qbuf)
.arg(&has_q)
.arg(&ni)
.arg(&temp)
.arg(&slo)
.arg(&shi)
.arg(&stream_pos)
.arg(&pm)
.arg(&pth)
.arg(&pz)
.arg(&qm)
.arg(&qth)
.arg(&qz)
.arg(&mut *out_tok);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gumbel_perturb_filtered(
&self,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
n: usize,
seed: u64,
stream_pos: u32,
temp: f32,
row_max: f32,
th: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gumbel_perturb_filtered_f32");
let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(&mut *y)
.arg(&ni)
.arg(&slo)
.arg(&shi)
.arg(&stream_pos)
.arg(&temp)
.arg(&row_max)
.arg(&th);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn penalize_logits(
&self,
x: &mut CudaSlice<f32>,
hist: &CudaSlice<u32>,
n_hist: usize,
rep: f32,
freq: f32,
present: f32,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
if n_hist == 0 {
return Ok(());
}
let f = self.func("penalize_logits_f32");
let (nh, ni) = (n_hist as i32, n as i32);
let cfg = LaunchConfig {
grid_dim: (n_hist.div_ceil(128) as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&mut *x)
.arg(hist)
.arg(&nh)
.arg(&rep)
.arg(&freq)
.arg(&present)
.arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn penalize_logits_rows(
&self,
x: &mut CudaSlice<f32>,
hist: &CudaSlice<u32>,
n_hist: usize,
rep: f32,
freq: f32,
present: f32,
n: usize,
nrow: usize,
) -> Result<(), Box<dyn std::error::Error>> {
if n_hist == 0 || nrow == 0 {
return Ok(());
}
let f = self.func("penalize_logits_rows_f32");
let (nh, ni, nr) = (n_hist as i32, n as i32, nrow as i32);
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,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&mut *x)
.arg(hist)
.arg(&nh)
.arg(&rep)
.arg(&freq)
.arg(&present)
.arg(&ni)
.arg(&nr);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn wpf_level() -> u32 {
static ON: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_WPF")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(1)
})
}
pub fn set_verify_exact(&self, on: bool) {
self.verify_exact
.store(on, std::sync::atomic::Ordering::Relaxed);
}
pub(crate) fn verify_exact_on(&self) -> bool {
self.verify_exact.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn qkv_append_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_QKV_APPEND")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub fn pdl_wb_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_PDL_WB")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub fn pdl_mmvq_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_PDL_MMVQ")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub fn pdl_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_PDL").map(|v| v != "0").unwrap_or(true))
}
fn q40_mr1_on() -> bool {
static Q40MR: std::sync::OnceLock<Option<u32>> = std::sync::OnceLock::new();
match *Q40MR.get_or_init(|| {
std::env::var("MEMRA_Q40_MR")
.ok()
.and_then(|v| v.parse().ok())
}) {
Some(v) => v == 1,
None => crate::FUSED_MR1_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
}
}
fn pdl_func_flash(
&self,
g: bool,
name: &'static str,
) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
use cudarc::driver::sys as cu;
static MODS: std::sync::Mutex<Option<std::collections::HashMap<(usize, bool), usize>>> =
std::sync::Mutex::new(None);
static FNS: std::sync::Mutex<
Option<std::collections::HashMap<(usize, bool, &'static str), usize>>,
> = std::sync::Mutex::new(None);
let ctx_key = self.ctx().cu_ctx() as usize;
if let Some(&f) = FNS
.lock()
.unwrap()
.get_or_insert_with(Default::default)
.get(&(ctx_key, g, name))
{
return Ok(f as cu::CUfunction);
}
let module = {
let mut mods = MODS.lock().unwrap();
let map = mods.get_or_insert_with(Default::default);
match map.get(&(ctx_key, g)) {
Some(&m) => m,
None => {
let m = self.pdl_load_module_in_ctx(if g {
FLASH_FATBIN_KF8VF8
} else {
FLASH_FATBIN
})?;
map.insert((ctx_key, g), m);
m
}
}
};
let cname = std::ffi::CString::new(name)?;
let mut f: cu::CUfunction = std::ptr::null_mut();
let r = unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
if r != cu::CUresult::CUDA_SUCCESS {
return Err(format!("pdl_func_flash {name} (g={g}): {r:?}").into());
}
FNS.lock()
.unwrap()
.get_or_insert_with(Default::default)
.insert((ctx_key, g, name), f as usize);
Ok(f)
}
fn pdl_load_module_in_ctx(&self, bytes: &[u8]) -> Result<usize, Box<dyn std::error::Error>> {
use cudarc::driver::sys as cu;
let mut prev: cu::CUcontext = std::ptr::null_mut();
unsafe {
cu::cuCtxGetCurrent(&mut prev).result()?;
}
self.ctx().bind_to_thread()?;
let mut m: cu::CUmodule = std::ptr::null_mut();
let r = unsafe { cu::cuModuleLoadData(&mut m, bytes.as_ptr() as *const std::ffi::c_void) };
let restore = if prev.is_null() {
cu::CUresult::CUDA_SUCCESS
} else {
unsafe { cu::cuCtxSetCurrent(prev) }
};
if r != cu::CUresult::CUDA_SUCCESS {
return Err(format!("pdl module load: {r:?}").into());
}
if restore != cu::CUresult::CUDA_SUCCESS {
return Err(format!("pdl module load: ctx restore {restore:?}").into());
}
Ok(m as usize)
}
fn pdl_func(
&self,
name: &'static str,
) -> Result<cudarc::driver::sys::CUfunction, Box<dyn std::error::Error>> {
use cudarc::driver::sys as cu;
static MODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
std::sync::Mutex::new(None);
static QMODULES: std::sync::Mutex<Option<std::collections::HashMap<usize, usize>>> =
std::sync::Mutex::new(None);
static FNS: std::sync::Mutex<
Option<std::collections::HashMap<(usize, &'static str), usize>>,
> = std::sync::Mutex::new(None);
let ctx_key = self.ctx().cu_ctx() as usize;
if let Some(&f) = FNS
.lock()
.unwrap()
.get_or_insert_with(Default::default)
.get(&(ctx_key, name))
{
return Ok(f as cu::CUfunction);
}
let module = {
let mut mods = MODULES.lock().unwrap();
let map = mods.get_or_insert_with(Default::default);
match map.get(&ctx_key) {
Some(&m) => m,
None => {
let m = self.pdl_load_module_in_ctx(FATBIN)?;
map.insert(ctx_key, m);
m
}
}
};
let cname = std::ffi::CString::new(name)?;
let mut f: cu::CUfunction = std::ptr::null_mut();
let mut r =
unsafe { cu::cuModuleGetFunction(&mut f, module as cu::CUmodule, cname.as_ptr()) };
if r == cu::CUresult::CUDA_ERROR_NOT_FOUND {
let qmodule = {
let mut mods = QMODULES.lock().unwrap();
let map = mods.get_or_insert_with(Default::default);
match map.get(&ctx_key) {
Some(&m) => m,
None => {
let m = self.pdl_load_module_in_ctx(QMATVEC_FATBIN)?;
map.insert(ctx_key, m);
m
}
}
};
r = unsafe { cu::cuModuleGetFunction(&mut f, qmodule as cu::CUmodule, cname.as_ptr()) };
}
if r != cu::CUresult::CUDA_SUCCESS {
return Err(format!("pdl_func {name}: {r:?}").into());
}
FNS.lock()
.unwrap()
.get_or_insert_with(Default::default)
.insert((ctx_key, name), f as usize);
Ok(f)
}
unsafe fn launch_pdl_flash(
&self,
g: bool,
name: &'static str,
grid: (u32, u32, u32),
block: (u32, u32, u32),
smem: u32,
params: &mut [*mut std::ffi::c_void],
) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::sys as cu;
let f = self.pdl_func_flash(g, name)?;
if smem > 0 {
let r =
unsafe {
cu::cuFuncSetAttribute(f,
cu::CUfunction_attribute_enum::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
smem as i32)
};
if r != cu::CUresult::CUDA_SUCCESS {
return Err(format!("pdl smem attr {name}: {r:?}").into());
}
}
let mut attr = cu::CUlaunchAttribute {
id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
pad: [0; 4],
value: cu::CUlaunchAttributeValue {
programmaticStreamSerializationAllowed: 1,
},
};
let cfg = cu::CUlaunchConfig {
gridDimX: grid.0,
gridDimY: grid.1,
gridDimZ: grid.2,
blockDimX: block.0,
blockDimY: block.1,
blockDimZ: block.2,
sharedMemBytes: smem,
hStream: self.gpu.stream().cu_stream(),
attrs: &mut attr,
numAttrs: 1,
};
let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
if r != cu::CUresult::CUDA_SUCCESS {
return Err(format!("launch_pdl_flash {name}: {r:?}").into());
}
Ok(())
}
unsafe fn launch_pdl(
&self,
name: &'static str,
grid: (u32, u32, u32),
block: (u32, u32, u32),
params: &mut [*mut std::ffi::c_void],
) -> Result<(), Box<dyn std::error::Error>> {
use cudarc::driver::sys as cu;
let f = self.pdl_func(name)?;
let mut attr = cu::CUlaunchAttribute {
id: cu::CUlaunchAttributeID::CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION,
pad: [0; 4],
value: cu::CUlaunchAttributeValue {
programmaticStreamSerializationAllowed: 1,
},
};
let cfg = cu::CUlaunchConfig {
gridDimX: grid.0,
gridDimY: grid.1,
gridDimZ: grid.2,
blockDimX: block.0,
blockDimY: block.1,
blockDimZ: block.2,
sharedMemBytes: 0,
hStream: self.gpu.stream().cu_stream(),
attrs: &mut attr,
numAttrs: 1,
};
let r = unsafe { cu::cuLaunchKernelEx(&cfg, f, params.as_mut_ptr(), std::ptr::null_mut()) };
if r != cu::CUresult::CUDA_SUCCESS {
return Err(format!("launch_pdl {name}: {r:?}").into());
}
Ok(())
}
pub fn prefetch_weight_l2(
&self,
w: &crate::model::GpuTensor,
) -> Result<(), Box<dyn std::error::Error>> {
if let crate::model::GpuTensor::Quant { bytes, rp4, .. } = w {
let p = rp4.as_ref().unwrap_or(bytes);
self.prefetch_l2(p, p.len())?;
}
Ok(())
}
pub fn gather_row_bf16(
&self,
table: &CudaSlice<u8>,
tok: &CudaSlice<u32>,
idx: usize,
dst: &mut CudaSlice<f32>,
ncols: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gather_row_bf16_f32");
let cfg = LaunchConfig {
grid_dim: (ncols.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (nc, ix) = (ncols as i32, idx as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table).arg(tok).arg(&ix).arg(dst).arg(&nc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn add_row_inplace(
&self,
logits: &mut CudaSlice<f32>,
bias: &CudaSlice<f32>,
n: usize,
row_off: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("add_row_inplace_f32");
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (ni, off) = (n as i32, row_off as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(logits).arg(bias).arg(&ni).arg(&off);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn prefetch_l2(
&self,
p: &CudaSlice<u8>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("prefetch_l2_bytes");
let lines = n.div_ceil(128);
let ni = n as i64;
let cfg = LaunchConfig {
grid_dim: (lines.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(p).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn router_gemv(
&self,
w: &CudaSlice<f32>,
x: &CudaSlice<f32>,
n_embd: usize,
n_experts: usize,
t: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let w8 = match std::env::var("MEMRA_ROUTER_V2").as_deref() {
Ok("0") => false,
Ok(_) => true,
Err(_) => ROUTER_W8_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
};
let batch = w8 && t >= ROUTER_BATCH_MIN_T && router_batch_on();
self.router_gemv_form(w, x, n_embd, n_experts, t, w8, batch)
}
pub fn router_gemv_form(
&self,
w: &CudaSlice<f32>,
x: &CudaSlice<f32>,
n_embd: usize,
n_experts: usize,
t: usize,
w8: bool,
batch: bool,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
debug_assert!(!batch || w8, "batch twin exists for the w8 form only");
let mut y = self.alloc_uninit::<f32>(t * n_experts)?;
let f = if batch {
self.func("router_gemv_f32_w8_batch")
} else if w8 {
self.func("router_gemv_f32_w8")
} else {
self.func("router_gemv_f32")
};
let (ne, nx, ti) = (n_embd as i32, n_experts as i32, t as i32);
let cfg = if batch {
LaunchConfig {
grid_dim: (n_experts.div_ceil(8) as u32, t.div_ceil(8) as u32, 1),
block_dim: (32, 8, 1),
shared_mem_bytes: 0,
}
} else {
LaunchConfig {
grid_dim: (n_experts as u32, t as u32, 1),
block_dim: (32, if w8 { 8 } else { 1 }, 1),
shared_mem_bytes: 0,
}
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(w).arg(x).arg(&mut y).arg(&ne).arg(&nx).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn rows_permute(
&self,
src: &CudaSlice<f32>,
idx: &CudaSlice<i32>,
nrows: usize,
ncols: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let mut dst = self.alloc_uninit::<f32>(nrows * ncols)?;
let f = self.func("rows_permute_f32");
let (nc, nr) = (ncols as i32, nrows as i32);
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src).arg(idx).arg(&mut dst).arg(&nc).arg(&nr);
unsafe {
b.launch(cfg)?;
}
Ok(dst)
}
pub fn sigmoid_dot_rows(
&self,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
n_embd: usize,
t: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
static OFF: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
if *OFF.get_or_init(|| std::env::var("MEMRA_SHEXP_DOT").as_deref() == Ok("0")) {
let gs = self.linear(x, w, t, n_embd, 1)?;
let mut g = self.uninit(t)?;
self.sigmoid(&gs, &mut g, t)?;
return Ok(g);
}
let mut g = self.alloc_uninit::<f32>(t)?;
let f = self.func("sigmoid_dot_rows_f32");
let (ne, ti) = (n_embd as i32, t as i32);
let cfg = LaunchConfig {
grid_dim: (t as u32, 1, 1),
block_dim: (32, 8, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(w).arg(&mut g).arg(&ne).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(g)
}
pub fn spec_rollback_stream(
&self,
len_ptrs: &CudaSlice<u64>,
pos_start: &CudaSlice<i32>,
acc: &CudaSlice<u32>,
base: usize,
n_rows: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_rollback_stream");
let (b, nr) = (base as i32, n_rows as i32);
let cfg = LaunchConfig {
grid_dim: (n_rows.div_ceil(64) as u32, 1, 1),
block_dim: (64, 1, 1),
shared_mem_bytes: 0,
};
let __s_bl = self.gpu.stream();
let mut bl = __s_bl.launch_builder(&f);
bl.arg(len_ptrs).arg(pos_start).arg(acc).arg(&b).arg(&nr);
unsafe {
bl.launch(cfg)?;
}
Ok(())
}
pub fn plain_tok_ring(
&self,
vam: &CudaSlice<u32>,
pos_start: &CudaSlice<i32>,
base: usize,
ring: &mut CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("plain_tok_ring");
let (b, cap) = (base as i32, ring.len() as i32);
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_bl = self.gpu.stream();
let mut bl = __s_bl.launch_builder(&f);
bl.arg(vam).arg(pos_start).arg(&b).arg(&mut *ring).arg(&cap);
unsafe {
bl.launch(cfg)?;
}
Ok(())
}
pub fn spec_ring_commit(
&self,
vtok: &CudaSlice<u32>,
acc: &CudaSlice<u32>,
brk: &CudaSlice<u32>,
ring: &mut CudaSlice<u32>,
pend: &mut CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_ring_commit");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(vtok).arg(acc).arg(brk).arg(ring).arg(pend);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn i32_copy_add(
&self,
src: &CudaSlice<i32>,
dst: &mut CudaSlice<i32>,
delta: i32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("i32_copy_add");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src).arg(dst).arg(&delta);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn u32_copy(
&self,
src: &CudaSlice<u32>,
dst: &mut CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("u32_copy");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src).arg(dst);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn spec_adapt_k(
&self,
acc: &CudaSlice<u32>,
brk: &mut CudaSlice<u32>,
floor: usize,
cap: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_adapt_k");
let (fl, cp) = (floor as i32, cap as i32);
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(acc).arg(brk).arg(&fl).arg(&cp);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn spec_accept_greedy_dc(
&self,
preds: &CudaSlice<u32>,
vtok: &CudaSlice<u32>,
last_pred: &CudaSlice<u32>,
brk: &CudaSlice<u32>,
out: &mut CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_accept_greedy_dc");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(preds).arg(vtok).arg(last_pred).arg(brk).arg(out);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn pos_iota(
&self,
pos0: &CudaSlice<i32>,
out: &mut CudaSlice<i32>,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("pos_iota_i32");
let ti = t as i32;
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (t.max(1) as u32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(pos0).arg(out).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn append_kv_quantized_rows_dc(
&self,
k_rows: &CudaSlice<f32>,
v_rows: &CudaSlice<f32>,
kc: &mut CudaSlice<u8>,
vc: &mut CudaSlice<u8>,
t0_dev: &CudaSlice<i32>,
t: usize,
kv_dim_k: usize,
kv_dim_v: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = if g {
self.func_g("append_quantize_kv_q8_0_q5_1_rows_dc")
} else {
self.func("append_quantize_kv_q8_0_q5_1_rows_dc")
};
let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk, t as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k_rows)
.arg(v_rows)
.arg(kc)
.arg(vc)
.arg(t0_dev)
.arg(&kdk)
.arg(&kdv)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn append_kv_quantized_row_dc_inc(
&self,
k_row: &CudaSlice<f32>,
v_row: &CudaSlice<f32>,
kc: &mut CudaSlice<u8>,
vc: &mut CudaSlice<u8>,
t0_dev: &mut CudaSlice<i32>,
kv_dim_k: usize,
kv_dim_v: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = if g {
self.func_g("append_quantize_kv_q8_0_q5_1_dc_inc")
} else {
self.func("append_quantize_kv_q8_0_q5_1_dc_inc")
};
let nthreads = ((kv_dim_k.max(kv_dim_v) / 32) * 32).min(1024) as u32;
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (nthreads, 1, 1),
shared_mem_bytes: 0,
};
let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k_row)
.arg(v_row)
.arg(kc)
.arg(vc)
.arg(t0_dev)
.arg(&kdk)
.arg(&kdv)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn pack_tok_p(
&self,
tok: &CudaSlice<u32>,
p: &CudaSlice<f32>,
out: &mut CudaSlice<u32>,
slot: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("pack_tok_p");
let sl = slot as i32;
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(tok).arg(p).arg(out).arg(&sl);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn tok_map_u32(
&self,
tok: &mut CudaSlice<u32>,
map: &CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("tok_map_u32");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(tok).arg(map);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn spec_assemble_verify(
&self,
tokp: &CudaSlice<u32>,
pend: &CudaSlice<u32>,
d2t: Option<&CudaSlice<u32>>,
vtok: &mut CudaSlice<u32>,
brk: &mut CudaSlice<u32>,
p_min: f32,
k: usize,
pmin0: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_assemble_verify");
let (ki, pm) = (k as i32, if pmin0 { 1i32 } else { 0i32 });
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
match d2t {
Some(m) => {
b.arg(tokp)
.arg(pend)
.arg(m)
.arg(vtok)
.arg(brk)
.arg(&p_min)
.arg(&ki)
.arg(&pm);
unsafe {
b.launch(cfg)?;
}
}
None => {
let null: u64 = 0;
b.arg(tokp)
.arg(pend)
.arg(&null)
.arg(vtok)
.arg(brk)
.arg(&p_min)
.arg(&ki)
.arg(&pm);
unsafe {
b.launch(cfg)?;
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn ssm_conv_ring_rebuild_dc(
&self,
qkv_tm: &CudaSlice<f32>,
ring_old: &CudaSlice<f32>,
conv_state: &mut CudaSlice<f32>,
conv_dim: usize,
acc: &CudaSlice<u32>,
base: usize,
t_v: usize,
d_conv: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("ssm_conv_ring_rebuild_f32_dc");
let n = conv_dim * (d_conv - 1);
let cfg = LaunchConfig::for_num_elems(n as u32);
let (cd, b0, tv, dc) = (conv_dim as i32, base as i32, t_v as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm)
.arg(ring_old)
.arg(conv_state)
.arg(&cd)
.arg(acc)
.arg(&b0)
.arg(&tv)
.arg(&dc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_scan_s128_dc(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
n_head: usize,
acc: &CudaSlice<u32>,
base: usize,
t_v: usize,
scale: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_scan_s128_dc");
const S_V: u32 = 128;
const WARP: u32 = 32;
const COLS_PER_BLOCK: u32 = 4;
let cfg = LaunchConfig {
grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
block_dim: (WARP, COLS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (h, b0, tv) = (n_head as i32, base as i32, t_v as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(g)
.arg(beta)
.arg(state_in)
.arg(state_out)
.arg(o)
.arg(&h)
.arg(acc)
.arg(&b0)
.arg(&tv)
.arg(&scale);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn spec_rollback_kv(
&self,
len_ptrs: &CudaSlice<u64>,
saved: &CudaSlice<i32>,
acc: &CudaSlice<u32>,
base: usize,
n_layer: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_rollback_kv");
let (b, nl) = (base as i32, n_layer as i32);
let cfg = LaunchConfig {
grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
block_dim: (64, 1, 1),
shared_mem_bytes: 0,
};
let __s_bl = self.gpu.stream();
let mut bl = __s_bl.launch_builder(&f);
bl.arg(len_ptrs).arg(saved).arg(acc).arg(&b).arg(&nl);
unsafe {
bl.launch(cfg)?;
}
Ok(())
}
pub fn spec_fork_valid(
&self,
acc: &CudaSlice<u32>,
optimistic_pending: u32,
valid: &mut CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_fork_valid");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
};
let __s_bl = self.gpu.stream();
let mut bl = __s_bl.launch_builder(&f);
bl.arg(acc).arg(&optimistic_pending).arg(valid);
unsafe {
bl.launch(cfg)?;
}
Ok(())
}
pub fn spec_fork_reconcile_kv(
&self,
len_ptrs: &CudaSlice<u64>,
saved: &CudaSlice<i32>,
acc: &CudaSlice<u32>,
valid: &CudaSlice<u32>,
base: usize,
n_layer: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_fork_reconcile_kv");
let (b, nl) = (base as i32, n_layer as i32);
let cfg = LaunchConfig {
grid_dim: (n_layer.div_ceil(64) as u32, 1, 1),
block_dim: (64, 1, 1),
shared_mem_bytes: 0,
};
let __s_bl = self.gpu.stream();
let mut bl = __s_bl.launch_builder(&f);
bl.arg(len_ptrs)
.arg(saved)
.arg(acc)
.arg(valid)
.arg(&b)
.arg(&nl);
unsafe {
bl.launch(cfg)?;
}
Ok(())
}
pub fn spec_fork_restore_f32(
&self,
snapshot: &CudaSlice<f32>,
state: &mut CudaSlice<f32>,
valid: &CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
assert_eq!(
snapshot.len(),
state.len(),
"fork recurrent snapshot shape mismatch"
);
let f = self.func("spec_fork_restore_f32");
let n = state.len() as i32;
let blocks = state.len().div_ceil(256).min(65535).max(1) as u32;
let cfg = LaunchConfig {
grid_dim: (blocks, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_bl = self.gpu.stream();
let mut bl = __s_bl.launch_builder(&f);
bl.arg(snapshot).arg(state).arg(valid).arg(&n);
unsafe {
bl.launch(cfg)?;
}
Ok(())
}
pub fn spec_seed_gather(
&self,
vx: &CudaSlice<f32>,
fill_prev: &CudaSlice<f32>,
acc: &CudaSlice<u32>,
h_seed: &mut CudaSlice<f32>,
base: usize,
n_embd: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_seed_gather");
let (b, ne) = (base as i32, n_embd as i32);
let cfg = LaunchConfig {
grid_dim: (n_embd.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_bl = self.gpu.stream();
let mut bl = __s_bl.launch_builder(&f);
bl.arg(vx)
.arg(fill_prev)
.arg(acc)
.arg(h_seed)
.arg(&b)
.arg(&ne);
unsafe {
bl.launch(cfg)?;
}
Ok(())
}
pub fn spec_accept_greedy(
&self,
preds: &CudaSlice<u32>,
draft: &CudaSlice<u32>,
last_pred: u32,
base: usize,
k_round: usize,
out: &mut CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("spec_accept_greedy");
let (b, k) = (base as i32, k_round as i32);
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_bl = self.gpu.stream();
let mut bl = __s_bl.launch_builder(&f);
bl.arg(preds)
.arg(draft)
.arg(&last_pred)
.arg(&b)
.arg(&k)
.arg(out);
unsafe {
bl.launch(cfg)?;
}
Ok(())
}
pub fn gumbel_perturb(
&self,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
n: usize,
seed: u64,
stream_pos: u32,
temp: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gumbel_perturb_f32");
let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(&mut *y)
.arg(&ni)
.arg(&slo)
.arg(&shi)
.arg(&stream_pos)
.arg(&temp);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn mask_logits_col(
&self,
logits: &mut CudaSlice<f32>,
mask: &CudaSlice<u32>,
col: usize,
n: usize,
mask_words: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("mask_logits_f32");
let (ci, ni, mw) = (col as i32, n as i32, mask_words as i32);
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256).min(1024) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&mut *logits).arg(mask).arg(&ci).arg(&ni).arg(&mw);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gumbel_perturb_col(
&self,
x: &CudaSlice<f32>,
col: usize,
y: &mut CudaSlice<f32>,
n: usize,
seed: u64,
stream_pos: u32,
temp: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gumbel_perturb_f32");
let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
let col_view = x.slice(col * n..(col + 1) * n);
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&col_view)
.arg(&mut *y)
.arg(&ni)
.arg(&slo)
.arg(&shi)
.arg(&stream_pos)
.arg(&temp);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gumbel_perturb_filtered_col(
&self,
x: &CudaSlice<f32>,
col: usize,
y: &mut CudaSlice<f32>,
n: usize,
seed: u64,
stream_pos: u32,
temp: f32,
stat_max: &CudaSlice<f32>,
stat_th: &CudaSlice<f32>,
stat_idx: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gumbel_perturb_filtered_col_f32");
let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
let (ci, si) = (col as i32, stat_idx as i32);
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(&ci)
.arg(&mut *y)
.arg(&ni)
.arg(&slo)
.arg(&shi)
.arg(&stream_pos)
.arg(&temp)
.arg(stat_max)
.arg(stat_th)
.arg(&si);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn sctr_inc(&self, ctr: &mut CudaSlice<u32>) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("memra_sctr_inc");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&mut *ctr);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gumbel_perturb_ctr(
&self,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
n: usize,
seed: u64,
ctr: &CudaSlice<u32>,
temp: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gumbel_perturb_ctr_f32");
let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256) as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(&mut *y)
.arg(&ni)
.arg(&slo)
.arg(&shi)
.arg(ctr)
.arg(&temp);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn softmax_gather(
&self,
x: &CudaSlice<f32>,
row_stride: usize,
ids: &CudaSlice<u32>,
rows: &CudaSlice<i32>,
out: &mut CudaSlice<f32>,
n: usize,
npair: usize,
temp: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("softmax_gather_f32");
let (ni, rs) = (n as i32, row_stride as i64);
let np = npair as i32;
let cfg = LaunchConfig {
grid_dim: (npair as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(&rs)
.arg(ids)
.arg(rows)
.arg(&mut *out)
.arg(&ni)
.arg(&np)
.arg(&temp);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn residual_sample(
&self,
p: &CudaSlice<f32>,
q: Option<&CudaSlice<f32>>,
n: usize,
temp: f32,
seed: u64,
stream_pos: u32,
out_tok: &mut CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("residual_sample_f32");
let (ni, slo, shi) = (n as i32, (seed & 0xFFFF_FFFF) as u32, (seed >> 32) as u32);
let nth = 1024u32;
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (nth, 1, 1),
shared_mem_bytes: 0,
};
let has_q: i32 = q.is_some() as i32;
let qbuf = q.unwrap_or(p); let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(p)
.arg(qbuf)
.arg(&has_q)
.arg(&ni)
.arg(&temp)
.arg(&slo)
.arg(&shi)
.arg(&stream_pos)
.arg(&mut *out_tok);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn with_moe_cache<R>(
&self,
max_block_bytes: usize,
f: impl FnOnce(
&mut crate::moe_cache::MoeSlotCache,
&Engine,
) -> Result<R, Box<dyn std::error::Error>>,
) -> Result<R, Box<dyn std::error::Error>> {
let mut guard = self.moe_cache.lock().unwrap();
if guard.is_none() {
*guard = Some(crate::moe_cache::MoeSlotCache::new(self, max_block_bytes)?);
}
let cache = guard.as_mut().unwrap();
f(cache, self)
}
pub fn freeze_moe_cache(&self) {
if let Some(cache) = self.moe_cache.lock().unwrap().as_mut() {
cache.freeze();
}
}
pub fn export_moe_residency(&self) -> Option<Vec<(u16, u8, u16)>> {
self.moe_cache
.lock()
.unwrap()
.as_ref()
.map(crate::moe_cache::MoeSlotCache::export_residency)
}
pub(crate) fn moe_cache_frozen(&self) -> bool {
self.moe_cache
.lock()
.unwrap()
.as_ref()
.is_some_and(crate::moe_cache::MoeSlotCache::is_frozen)
}
pub fn frozen_cpu_experts_prefer_tokenwise_prime(&self) -> bool {
crate::cpu_experts::configured()
&& self.moe_cache_frozen()
&& std::env::var("MEMRA_CPU_EXPERT_BATCHED_PRIME").as_deref() != Ok("1")
}
pub(crate) fn configure_moe_cache_layout(&self, block_bytes: Vec<usize>) {
assert!(
self.moe_cache.lock().unwrap().is_none(),
"MoE cache layout configured after cache construction"
);
*self.moe_cache_layout.lock().unwrap() = Some(block_bytes);
}
pub(crate) fn moe_cache_layout(&self) -> Option<Vec<usize>> {
self.moe_cache_layout.lock().unwrap().clone()
}
pub fn moe_cache_enabled() -> bool {
std::env::var("MEMRA_MOE_CACHE").as_deref() != Ok("0")
}
pub fn moe_cache_stats(&self) -> Option<(u64, u64, u64, usize)> {
let guard = self.moe_cache.lock().unwrap();
guard
.as_ref()
.map(|c| (c.hits, c.misses, c.staged_bytes, c.n_slots()))
}
pub fn cpu_expert_stats(
&self,
) -> Option<(u64, u64, u64, u64, u64, u64, u64, u64, u64, u64, u64)> {
crate::cpu_experts::configured().then(crate::cpu_experts::stats)
}
pub fn cpu_expert_predictor_stats(&self) -> (u64, u64) {
crate::cpu_experts::predictor_stats()
}
pub fn cpu_expert_exposed_wait_ns(&self) -> Option<u64> {
crate::cpu_experts::configured().then(crate::cpu_experts::exposed_wait_ns)
}
pub fn cpu_expert_gpu_residency_stats(&self) -> Option<(u64, u64, u64)> {
crate::cpu_experts::configured().then(crate::cpu_experts::incomplete_gpu_residency_stats)
}
pub fn moe_pread_stats(&self) -> Option<(u64, u64, u64, u64, u64, u64, u64)> {
let guard = self.moe_cache.lock().unwrap();
guard
.as_ref()
.and_then(|cache| cache.pread_stats())
.map(|stats| {
(
stats.reads,
stats.bytes,
stats.read_errors,
stats.short_reads,
stats.fallbacks,
stats.buffer_waits,
stats.ring_full,
)
})
}
pub fn spill_config_fallbacks(&self) -> u64 {
crate::spill_pread::config_fallbacks()
}
pub fn moe_cache_reset_counters(&self) {
if let Some(c) = self.moe_cache.lock().unwrap().as_mut() {
c.reset_counters();
}
}
pub fn htod_bytes(&self, v: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
Ok(self.gpu.stream().clone_htod(v)?)
}
pub fn htod_bytes_padded(
&self,
v: &[u8],
pad: usize,
) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
let mut d = self.alloc_u8_uninit(v.len() + pad)?;
{
let mut view = d.slice_mut(0..v.len());
self.gpu.stream().memcpy_htod(v, &mut view)?;
}
Ok(d)
}
pub fn copy_into(
&self,
dst: &mut CudaSlice<f32>,
off: usize,
src: &CudaSlice<f32>,
len: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let mut view = dst.slice_mut(off..off + len);
self.gpu
.stream()
.memcpy_dtod(&src.slice(0..len), &mut view)?;
Ok(())
}
pub fn copy_u8_into(
&self,
dst: &mut CudaSlice<u8>,
off: usize,
src: &CudaSlice<u8>,
len: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let mut view = dst.slice_mut(off..off + len);
self.gpu
.stream()
.memcpy_dtod(&src.slice(0..len), &mut view)?;
Ok(())
}
pub fn copy_u8_range_into(
&self,
dst: &mut CudaSlice<u8>,
dst_off: usize,
src: &CudaSlice<u8>,
src_off: usize,
len: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let mut dst_view = dst.slice_mut(dst_off..dst_off + len);
self.gpu
.stream()
.memcpy_dtod(&src.slice(src_off..src_off + len), &mut dst_view)?;
Ok(())
}
pub fn prepare_kv_append(
&self,
kv: &mut crate::cache::KvLayer,
retain_from: usize,
append_rows: usize,
) -> Result<usize, Box<dyn std::error::Error>> {
let Some(plan) = kv
.ring
.as_ref()
.map(|ring| ring.append_plan(kv.len, retain_from, append_rows))
.transpose()?
else {
return Ok(kv.len);
};
match plan {
crate::cache::KvRingAppend::Contiguous { write_row } => Ok(write_row),
crate::cache::KvRingAppend::Rebase {
src_row,
keep_rows,
new_base,
write_row,
} => {
if keep_rows > 0 {
let k_len = keep_rows * kv.k_tok_bytes;
let v_len = keep_rows * kv.v_tok_bytes;
let mut k_tmp = self.alloc_u8_uninit(k_len)?;
let mut v_tmp = self.alloc_u8_uninit(v_len)?;
self.copy_u8_range_into(&mut k_tmp, 0, &kv.k, src_row * kv.k_tok_bytes, k_len)?;
self.copy_u8_range_into(&mut v_tmp, 0, &kv.v, src_row * kv.v_tok_bytes, v_len)?;
self.copy_u8_into(&mut kv.k, 0, &k_tmp, k_len)?;
self.copy_u8_into(&mut kv.v, 0, &v_tmp, v_len)?;
}
kv.ring.as_mut().unwrap().apply_rebase(new_base);
Ok(write_row)
}
}
}
pub fn htod_u8_into(
&self,
dst: &mut CudaSlice<u8>,
off: usize,
src: &[u8],
) -> Result<(), Box<dyn std::error::Error>> {
let mut view = dst.slice_mut(off..off + src.len());
self.gpu.stream().memcpy_htod(src, &mut view)?;
Ok(())
}
pub fn view<'a>(&self, b: &'a CudaSlice<f32>, len: usize) -> cudarc::driver::CudaView<'a, f32> {
b.slice(0..len)
}
pub fn view_u8_range<'a>(
&self,
b: &'a CudaSlice<u8>,
start: usize,
end: usize,
) -> cudarc::driver::CudaView<'a, u8> {
b.slice(start..end)
}
pub fn view_u8<'a>(
&self,
b: &'a CudaSlice<u8>,
len: usize,
) -> cudarc::driver::CudaView<'a, u8> {
b.slice(0..len)
}
pub fn append_kv_quantized(
&self,
k_row: &CudaSlice<f32>,
v_row: &CudaSlice<f32>,
kc: &mut CudaSlice<u8>,
vc: &mut CudaSlice<u8>,
t: usize,
kv_dim_k: usize,
kv_dim_v: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = if g {
self.func_g("append_quantize_kv_q8_0_q5_1")
} else {
self.func("append_quantize_kv_q8_0_q5_1")
};
let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k_row)
.arg(v_row)
.arg(kc)
.arg(vc)
.arg(&ti)
.arg(&kdk)
.arg(&kdv)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn append_kv_quantized_dc(
&self,
k_row: &CudaSlice<f32>,
v_row: &CudaSlice<f32>,
kc: &mut CudaSlice<u8>,
vc: &mut CudaSlice<u8>,
t_dev: &CudaSlice<i32>,
kv_dim_k: usize,
kv_dim_v: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pk, _g0) = k_row.device_ptr(s);
let (pv, _g1) = v_row.device_ptr(s);
let (pkc, _g2) = kc.device_ptr_mut(s);
let (pvc, _g3) = vc.device_ptr_mut(s);
let (pt, _g4) = t_dev.device_ptr(s);
let mut ps = [
&pk as *const _ as *mut std::ffi::c_void,
&pv as *const _ as *mut _,
&pkc as *const _ as *mut _,
&pvc as *const _ as *mut _,
&pt as *const _ as *mut _,
&kdk as *const _ as *mut _,
&kdv as *const _ as *mut _,
&ktb as *const _ as *mut _,
&vtb as *const _ as *mut _,
];
unsafe {
self.launch_pdl_flash(
g,
"append_quantize_kv_q8_0_q5_1_dc",
(nblk, 1, 1),
(32, 1, 1),
0,
&mut ps,
)?;
}
return Ok(());
}
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")
};
let cfg = LaunchConfig {
grid_dim: (nblk, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k_row)
.arg(v_row)
.arg(kc)
.arg(vc)
.arg(t_dev)
.arg(&kdk)
.arg(&kdv)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn append_kv_quantized_rows(
&self,
k_rows: &CudaSlice<f32>,
v_rows: &CudaSlice<f32>,
kc: &mut CudaSlice<u8>,
vc: &mut CudaSlice<u8>,
t0: usize,
t: usize,
kv_dim_k: usize,
kv_dim_v: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
if std::env::var("MEMRA_PRIME_APPEND_LOOP").is_ok() {
for i in 0..t {
let k_row = k_rows.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
let v_row = v_rows.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
self.append_kv_quantized_view(
&k_row,
&v_row,
kc,
vc,
t0 + i,
kv_dim_k,
kv_dim_v,
k_tok_bytes,
v_tok_bytes,
g,
)?;
}
return Ok(());
}
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")
};
let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk, t as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (t0i, kdk, kdv) = (t0 as i32, kv_dim_k as i32, kv_dim_v as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k_rows)
.arg(v_rows)
.arg(kc)
.arg(vc)
.arg(&t0i)
.arg(&kdk)
.arg(&kdv)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn inc_seqlen(&self, p: &mut CudaSlice<i32>) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("inc_i32");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(p);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn append_kv_quantized_view(
&self,
k_row: &cudarc::driver::CudaView<f32>,
v_row: &cudarc::driver::CudaView<f32>,
kc: &mut CudaSlice<u8>,
vc: &mut CudaSlice<u8>,
t: usize,
kv_dim_k: usize,
kv_dim_v: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = if g {
self.func_g("append_quantize_kv_q8_0_q5_1")
} else {
self.func("append_quantize_kv_q8_0_q5_1")
};
let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (ti, kdk, kdv) = (t as i32, kv_dim_k as i32, kv_dim_v as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k_row)
.arg(v_row)
.arg(kc)
.arg(vc)
.arg(&ti)
.arg(&kdk)
.arg(&kdv)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn copy_view_into(
&self,
dst: &mut CudaSlice<f32>,
off: usize,
src: &cudarc::driver::CudaView<f32>,
len: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let mut view = dst.slice_mut(off..off + len);
self.gpu
.stream()
.memcpy_dtod(&src.slice(0..len), &mut view)?;
Ok(())
}
pub fn clone_dtod(
&self,
src: &CudaSlice<f32>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let mut dst = self.gpu.stream().alloc_zeros::<f32>(src.len())?;
self.gpu.stream().memcpy_dtod(src, &mut dst)?;
Ok(dst)
}
pub fn dtod_copy_view(
&self,
src: &cudarc::driver::CudaView<f32>,
dst: &mut CudaSlice<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
self.gpu.stream().memcpy_dtod(src, dst)?;
Ok(())
}
pub fn dtod_copy_view_i8(
&self,
src: &cudarc::driver::CudaView<i8>,
dst: &mut CudaSlice<i8>,
) -> Result<(), Box<dyn std::error::Error>> {
self.gpu.stream().memcpy_dtod(src, dst)?;
Ok(())
}
pub fn dtod_copy_into(
&self,
src: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
offset: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let n = src.len();
let mut dv = dst.slice_mut(offset..offset + n);
self.gpu.stream().memcpy_dtod(src, &mut dv)?;
Ok(())
}
pub fn uninit_i8(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
self.alloc_uninit::<i8>(n)
}
pub fn qmatvec(
&self,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("qmatvec_f32");
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,
};
let (inf, outf, mi, qt, rb) =
(in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(w)
.arg(x)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&qt)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
let s = self.gpu.stream().alloc_zeros::<u8>(n)?;
self.keep_if_capturing(&s);
Ok(s)
}
pub fn alloc_u8_uninit(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
let s = unsafe { self.gpu.stream().alloc::<u8>(n)? };
self.keep_if_capturing(&s);
Ok(s)
}
pub fn memset_zeros_view(
&self,
dst: &mut cudarc::driver::CudaViewMut<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
self.gpu.stream().memset_zeros(dst)?;
Ok(())
}
pub fn stage_expert(
&self,
host_bytes: &[u8],
scratch: &mut CudaSlice<u8>,
off: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let mut dst = scratch.slice_mut(off..off + host_bytes.len()); self.gpu.stream().memcpy_htod(host_bytes, &mut dst)?; Ok(())
}
pub fn moe_router_topk(
&self,
logits: &CudaSlice<f32>,
t: usize,
n_expert: usize,
n_used: usize,
) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let f = self.func("moe_router_topk_f32");
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),
shared_mem_bytes: 0,
};
let (ne, nu) = (n_expert as i32, n_used as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(logits)
.arg(&mut sel_idx)
.arg(&mut sel_w)
.arg(&ne)
.arg(&nu);
unsafe {
b.launch(cfg)?;
}
Ok((sel_idx, sel_w))
}
pub fn moe_router_topk_scaled(
&self,
logits: &CudaSlice<f32>,
t: usize,
n_expert: usize,
n_used: usize,
ex_scale: &CudaSlice<f32>,
) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let f = self.func("moe_router_topk_scaled_f32");
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),
shared_mem_bytes: 0,
};
let (ne, nu) = (n_expert as i32, n_used as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(logits)
.arg(&mut sel_idx)
.arg(&mut sel_w)
.arg(&ne)
.arg(&nu)
.arg(ex_scale);
unsafe {
b.launch(cfg)?;
}
Ok((sel_idx, sel_w))
}
pub fn moe_router_topk_host(
&self,
logits: &CudaSlice<f32>,
t: usize,
n_expert: usize,
n_used: usize,
) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
let f = self.func("moe_router_topk_f32");
let n = t * n_used;
let mut sel_idx = self.alloc_uninit::<i32>(n)?;
let mut sel_w = self.alloc_uninit::<f32>(n)?;
let cfg = LaunchConfig {
grid_dim: (t as u32, 1, 1),
block_dim: (n_expert as u32, 1, 1),
shared_mem_bytes: 0,
};
let (ne, nu) = (n_expert as i32, n_used as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(logits)
.arg(&mut sel_idx)
.arg(&mut sel_w)
.arg(&ne)
.arg(&nu);
unsafe {
b.launch(cfg)?;
}
let bytes = n * 8;
let mut guard = self.router_stage.lock().unwrap();
if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
*guard = Some(PinnedStage::new(bytes.max(4096))?);
}
let stage = guard.as_mut().unwrap();
let (si, sw) = unsafe {
(
std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n),
)
};
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()))
}
#[allow(clippy::too_many_arguments)]
pub fn moe_router_sigmoid_topk(
&self,
logits: &CudaSlice<f32>,
t: usize,
n_expert: usize,
n_used: usize,
active_count: usize,
correction_bias: &CudaSlice<f32>,
active: &CudaSlice<u8>,
scaling_factor: f32,
route_norm: bool,
) -> Result<(CudaSlice<i32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
if n_expert == 0 || n_expert > 1024 || n_used == 0 || n_used > n_expert {
return Err(format!(
"sigmoid router shape unsupported: n_expert={n_expert}, n_used={n_used}",
)
.into());
}
if logits.len() < t * n_expert
|| correction_bias.len() != n_expert
|| active.len() != n_expert
{
return Err(format!(
"sigmoid router buffer mismatch: logits={} bias={} active={} expected logits>={} row={}",
logits.len(), correction_bias.len(), active.len(), t * n_expert, n_expert,
).into());
}
let f = self.func("moe_router_sigmoid_topk_f32");
let mut sel_idx = self.alloc_uninit::<i32>(t * n_used)?;
let mut sel_w = self.alloc_uninit::<f32>(t * n_used)?;
let threads = n_expert.div_ceil(32) * 32;
let cfg = LaunchConfig {
grid_dim: (t as u32, 1, 1),
block_dim: (threads as u32, 1, 1),
shared_mem_bytes: 0,
};
let (ne, nu, rn) = (n_expert as i32, n_used as i32, i32::from(route_norm));
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(logits)
.arg(correction_bias)
.arg(active)
.arg(&mut sel_idx)
.arg(&mut sel_w)
.arg(&ne)
.arg(&nu)
.arg(&scaling_factor)
.arg(&rn);
unsafe {
b.launch(cfg)?;
}
Ok((sel_idx, sel_w))
}
#[allow(clippy::too_many_arguments)]
pub fn moe_router_sigmoid_topk_host(
&self,
logits: &CudaSlice<f32>,
t: usize,
n_expert: usize,
n_used: usize,
active_count: usize,
correction_bias: &CudaSlice<f32>,
active: &CudaSlice<u8>,
scaling_factor: f32,
route_norm: bool,
) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
let (sel_idx, sel_w) = self.moe_router_sigmoid_topk(
logits,
t,
n_expert,
n_used,
active_count,
correction_bias,
active,
scaling_factor,
route_norm,
)?;
let n = t * n_used;
let bytes = n * 8;
let mut guard = self.router_stage.lock().unwrap();
if guard.as_ref().map(|p| p.cap < bytes).unwrap_or(true) {
*guard = Some(PinnedStage::new(bytes.max(4096))?);
}
let stage = guard.as_mut().unwrap();
let (si, sw) = unsafe {
(
std::slice::from_raw_parts_mut(stage.ptr as *mut i32, n),
std::slice::from_raw_parts_mut(stage.ptr.add(n * 4) as *mut f32, n),
)
};
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()))
}
pub fn stage_expert_async(
&self,
host_bytes: &[u8],
scratch: &mut CudaSlice<u8>,
off: usize,
) -> Result<cudarc::driver::CudaEvent, Box<dyn std::error::Error>> {
let mut dst = scratch.slice_mut(off..off + host_bytes.len());
self.copy_stream.memcpy_htod(host_bytes, &mut dst)?;
Ok(self.copy_stream.record_event(None)?)
}
pub fn compute_wait(
&self,
ev: &cudarc::driver::CudaEvent,
) -> Result<(), Box<dyn std::error::Error>> {
self.gpu.stream().wait(ev)?;
Ok(())
}
pub fn qmatvec_view(
&self,
w: &CudaSlice<u8>,
range: std::ops::Range<usize>,
x: &cudarc::driver::CudaView<f32>,
m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("qmatvec_f32");
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,
};
let (inf, outf, mi, qt, rb) =
(in_f as i32, out_f as i32, m as i32, qtype, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&wv)
.arg(x)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&qt)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_gate_up_silu8_q8(
&self,
gp: WPtr8,
up: WPtr8,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
n_ff: usize,
n_used: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_gate_up_silu8_q8");
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),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&gp)
.arg(&up)
.arg(aq)
.arg(ad)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_down8_fma_q8(
&self,
dp: WPtr8,
w: F32x8,
aq2: &CudaSlice<i8>,
ad2: &CudaSlice<f32>,
dst: &mut cudarc::driver::CudaViewMut<f32>,
in_f: usize,
out_f: usize,
n_used: usize,
qt: i32,
rb: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("moe_down8_fma_q8");
let cfg = LaunchConfig {
grid_dim: (out_f as u32, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, nu, rbi) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&dp)
.arg(&w)
.arg(aq2)
.arg(ad2)
.arg(dst)
.arg(&inf)
.arg(&outf)
.arg(&nu)
.arg(&qt)
.arg(&rbi);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn qmatvec_expert_q8(
&self,
w: &CudaSlice<u8>,
range: std::ops::Range<usize>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("qmatvec_expert_q8");
let wv = w.slice(range);
let mut y = self.alloc_uninit::<f32>(m * out_f)?;
const ROWS: u32 = 4; let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, m as u32, 1),
block_dim: (32, ROWS, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi, rbi) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&wv)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&qtype)
.arg(&rbi);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn moe_gate_up_silu8(
&self,
gp: WPtr8,
up: WPtr8,
x: &cudarc::driver::CudaView<f32>,
in_f: usize,
n_ff: usize,
n_used: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_gate_up_silu8_f32");
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),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (inf, nff, rbg, rbu) = (in_f as i32, n_ff as i32, rb_g as i64, rb_u as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&gp)
.arg(&up)
.arg(x)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_down8_fma_into(
&self,
dp: WPtr8,
w: F32x8,
act: &CudaSlice<f32>,
dst: &mut cudarc::driver::CudaViewMut<f32>,
in_f: usize,
out_f: usize,
n_used: usize,
qt: i32,
rb: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("moe_down8_fma_f32");
let cfg = LaunchConfig {
grid_dim: (out_f as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, nu, rbv) = (in_f as i32, out_f as i32, n_used as i32, rb as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&dp)
.arg(&w)
.arg(act)
.arg(dst)
.arg(&inf)
.arg(&outf)
.arg(&nu)
.arg(&qt)
.arg(&rbv);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
pub fn moe_pairs_matvec_q8(
&self,
table: &CudaSlice<u64>,
proj: i32,
pair_tok: &CudaSlice<i32>,
pair_ex: &CudaSlice<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
out_f: usize,
n_expert: usize,
n_pairs: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_pairs_matvec_q8");
let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
const ROWS: u32 = 4;
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_pairs as u32, 1),
block_dim: (32, ROWS, 1),
shared_mem_bytes: 0,
};
let (inf, outf, ne, np, rbi) = (
in_f as i32,
out_f as i32,
n_expert as i32,
n_pairs as i32,
row_bytes as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(&proj)
.arg(pair_tok)
.arg(pair_ex)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&ne)
.arg(&np)
.arg(&qtype)
.arg(&rbi);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_pairs_matvec_q8_em(
&self,
table: &CudaSlice<u64>,
proj: i32,
ex_ids: &CudaSlice<i32>,
ex_off: &CudaSlice<i32>,
ex_pairs: &CudaSlice<i32>,
pair_tok: &CudaSlice<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
out_f: usize,
n_expert: usize,
n_active: usize,
n_pairs: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_pairs_matvec_q8_em");
let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
const ROWS: u32 = 4;
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
block_dim: (32, ROWS, 1),
shared_mem_bytes: 0,
};
let (inf, outf, ne, na, rbi) = (
in_f as i32,
out_f as i32,
n_expert as i32,
n_active as i32,
row_bytes as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(&proj)
.arg(ex_ids)
.arg(ex_off)
.arg(ex_pairs)
.arg(pair_tok)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&ne)
.arg(&na)
.arg(&qtype)
.arg(&rbi);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_pairs_matvec_q8_dec(
&self,
table: &CudaSlice<u64>,
proj: i32,
ex_ids: &CudaSlice<i32>,
ex_off: &CudaSlice<i32>,
ex_pairs: &CudaSlice<i32>,
pair_tok: &CudaSlice<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
out_f: usize,
n_expert: usize,
n_active: usize,
n_pairs: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_pairs_matvec_q8_dec");
let mut y = self.alloc_uninit::<f32>(n_pairs * out_f)?;
const ROWS: u32 = 4;
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + ROWS - 1) / ROWS, n_active as u32, 1),
block_dim: (32, ROWS, 1),
shared_mem_bytes: 0,
};
let (inf, outf, ne, na, rbi) = (
in_f as i32,
out_f as i32,
n_expert as i32,
n_active as i32,
row_bytes as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(&proj)
.arg(ex_ids)
.arg(ex_off)
.arg(ex_pairs)
.arg(pair_tok)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&ne)
.arg(&na)
.arg(&qtype)
.arg(&rbi);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn moe_pairs_gelu_mul(
&self,
gate: &CudaSlice<f32>,
up: &CudaSlice<f32>,
n: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_pairs_gelu_mul");
let mut act = self.alloc_uninit::<f32>(n)?;
let cfg = LaunchConfig::for_num_elems(n as u32);
let nl = n as i64;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate).arg(up).arg(&mut act).arg(&nl);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
pub fn moe_pairs_silu_mul(
&self,
gate: &CudaSlice<f32>,
up: &CudaSlice<f32>,
n: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_pairs_silu_mul");
let mut act = self.alloc_uninit::<f32>(n)?;
let cfg = LaunchConfig::for_num_elems(n as u32);
let nl = n as i64;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate).arg(up).arg(&mut act).arg(&nl);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_pairs_scatter(
&self,
y_down: &CudaSlice<f32>,
pair_w: &CudaSlice<f32>,
tok_pair_off: &CudaSlice<i32>,
tok_pair_ids: &CudaSlice<i32>,
moe_out: &mut CudaSlice<f32>,
t: usize,
n_embd: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("moe_pairs_scatter");
let cfg = LaunchConfig {
grid_dim: (((n_embd + 255) / 256) as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let ne = n_embd as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(y_down)
.arg(pair_w)
.arg(tok_pair_off)
.arg(tok_pair_ids)
.arg(moe_out)
.arg(&ne);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn moe_gate_up_gelu8_dev_q8(
&self,
table: &CudaSlice<u64>,
sel: &cudarc::driver::CudaView<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
n_ff: usize,
n_used: usize,
n_expert: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
let (inf, nff, ne, rbg, rbu) = (
in_f as i32,
n_ff as i32,
n_expert as i32,
rb_g as i64,
rb_u as i64,
);
let f = self.func("moe_gate_up_gelu8_dev_q8");
let cfg = LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(aq)
.arg(ad)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&ne)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_gate_up_gelu8_dev_q8_rows(
&self,
table: &CudaSlice<u64>,
sel: &CudaSlice<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
t: usize,
in_f: usize,
n_ff: usize,
n_used: usize,
n_expert: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
let (inf, nff, ne, rbg, rbu, nu) = (
in_f as i32,
n_ff as i32,
n_expert as i32,
rb_g as i64,
rb_u as i64,
n_used as i32,
);
let f = self.func("moe_gate_up_gelu8_dev_q8_rows");
let cfg = LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, t as u32),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(aq)
.arg(ad)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&ne)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu)
.arg(&nu);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_gate_up_gelu8_dev_q8_csr(
&self,
table: &CudaSlice<u64>,
sel: &CudaSlice<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
n_pairs: usize,
in_f: usize,
n_ff: usize,
n_used: usize,
n_expert: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
let (inf, nff, ne, rbg, rbu, nu, npi) = (
in_f as i32,
n_ff as i32,
n_expert as i32,
rb_g as i64,
rb_u as i64,
n_used as i32,
n_pairs as i32,
);
let f = self.func("moe_gate_up_gelu8_dev_q8_csr");
let cfg = LaunchConfig {
grid_dim: (n_ff as u32, n_pairs as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(aq)
.arg(ad)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&ne)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu)
.arg(&nu)
.arg(&npi);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_down8_fma_dev_q8_rows_g(
&self,
table: &CudaSlice<u64>,
sel: &CudaSlice<i32>,
w: &CudaSlice<f32>,
aq2: &CudaSlice<i8>,
ad2: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
t: usize,
in_f: usize,
out_f: usize,
n_used: usize,
n_expert: usize,
qt: i32,
rb: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let (inf, outf, nu, ne, rbi) = (
in_f as i32,
out_f as i32,
n_used as i32,
n_expert as i32,
rb as i64,
);
let step_b1_w8 = t == 1 && in_f == 1280 && out_f == 4096 && n_used == 8 && qt == QT_IQ4_XS;
let f = self.func(if step_b1_w8 {
"moe_down8_fma_dev_q8_rows_w8"
} else {
"moe_down8_fma_dev_q8_rows_g"
});
let cfg = LaunchConfig {
grid_dim: (out_f as u32, 1, t as u32),
block_dim: (32, if step_b1_w8 { 8 } else { 1 }, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(w)
.arg(aq2)
.arg(ad2)
.arg(dst)
.arg(&inf)
.arg(&outf)
.arg(&nu)
.arg(&ne)
.arg(&qt)
.arg(&rbi);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn rp_probe_q4(&self, m: usize) -> Result<(f64, f64), Box<dyn std::error::Error>> {
let (out_f, in_f) = (2048usize, 2816usize);
let nblk = in_f / 32;
let mut seed = 0x9E3779B97F4A7C15u64;
let mut rng = move || {
seed = seed
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(seed >> 33) as u8
};
let mut w = vec![0u8; out_f * nblk * 18];
for b in w.iter_mut() {
*b = rng();
}
for r in 0..out_f {
for g in 0..nblk {
let off = (r * nblk + g) * 18;
w[off] = 0x00;
w[off + 1] = 0x2C; }
}
let qplane = out_f * nblk * 16;
let mut wrp = vec![0u8; w.len()];
for r in 0..out_f {
for g in 0..nblk {
let src = &w[(r * nblk + g) * 18..(r * nblk + g) * 18 + 18];
wrp[qplane + (r * nblk + g) * 2..qplane + (r * nblk + g) * 2 + 2]
.copy_from_slice(&src[0..2]);
wrp[(r * nblk + g) * 16..(r * nblk + g) * 16 + 16].copy_from_slice(&src[2..18]);
}
}
let w_d = self.htod_bytes(&w)?;
let wrp_d = self.htod_bytes(&wrp)?;
let mut aq = vec![0i8; m * in_f];
for v in aq.iter_mut() {
*v = rng() as i8;
}
let aq_d = self.htod_i8(&aq)?;
let ad_d = self.htod(&vec![0.03125f32; m * nblk])?;
let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
const RPB: u32 = 4;
let cfg = LaunchConfig {
grid_dim: ((out_f as u32).div_ceil(RPB), 1, 1),
block_dim: (32, RPB, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
let (rb, qp) = ((nblk * 18) as i64, qplane as i64);
let fb = self.func("qmatvec_q4_0_mmvq_b4");
let fr = self.func("qmatvec_q4_0_mmvq_b4_rp");
{
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&fb);
b.arg(&w_d)
.arg(&aq_d)
.arg(&ad_d)
.arg(&mut y0)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&fr);
b.arg(&wrp_d)
.arg(&aq_d)
.arg(&ad_d)
.arg(&mut y1)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&qp);
unsafe {
b.launch(cfg)?;
}
}
self.gpu.stream().synchronize()?;
let (h0, h1) = (self.dtoh(&y0)?, self.dtoh(&y1)?);
let nd = h0
.iter()
.zip(&h1)
.filter(|(a, b)| a.to_bits() != b.to_bits())
.count();
if nd != 0 {
return Err(format!("rp twin not bitwise: {nd}/{} diffs", h0.len()).into());
}
let mut time = |rp: bool| -> Result<f64, Box<dyn std::error::Error>> {
self.gpu.stream().synchronize()?;
let t0 = std::time::Instant::now();
for _ in 0..500 {
if rp {
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&fr);
b.arg(&wrp_d)
.arg(&aq_d)
.arg(&ad_d)
.arg(&mut y1)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&qp);
unsafe {
b.launch(cfg)?;
}
} else {
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&fb);
b.arg(&w_d)
.arg(&aq_d)
.arg(&ad_d)
.arg(&mut y0)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
}
}
self.gpu.stream().synchronize()?;
Ok(t0.elapsed().as_secs_f64() * 1e6 / 500.0)
};
let _ = time(false)?;
let _ = time(true)?; Ok((time(false)?, time(true)?))
}
pub fn build_q4_rp4(
&self,
t: &mut crate::model::GpuTensor,
) -> Result<(), Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let GpuTensor::Quant {
bytes,
qtype,
row_bytes,
ne,
rp4,
..
} = t
else {
return Ok(());
};
if *qtype != QT_Q4_0 || rp4.is_some() || ne.len() != 2 {
return Ok(());
}
let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 18 {
return Ok(());
}
let nblk = in_f / 32;
let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 18)?;
let f = self.func("q4_0_split_rp_build");
let n = (out_f * nblk) as i32;
let cfg = LaunchConfig {
grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (of, nb) = (out_f as i32, nblk as i32);
let _ = n;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
unsafe {
b.launch(cfg)?;
}
*rp4 = Some(dst);
Ok(())
}
pub fn build_q8_rp4(
&self,
t: &mut crate::model::GpuTensor,
) -> Result<(), Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let GpuTensor::Quant {
bytes,
qtype,
row_bytes,
ne,
rp4,
..
} = t
else {
return Ok(());
};
if *qtype != QT_Q8_0 || rp4.is_some() || ne.len() != 2 {
return Ok(());
}
let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
if in_f % 32 != 0 || *row_bytes != (in_f / 32) * 34 {
return Ok(());
}
*rp4 = Some(self.build_q8_rp4_raw(bytes, in_f, out_f)?);
Ok(())
}
pub fn build_q8_rp4_raw(
&self,
bytes: &CudaSlice<u8>,
in_f: usize,
out_f: usize,
) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
assert!(in_f % 32 == 0);
let nblk = in_f / 32;
let mut dst = self.alloc_uninit::<u8>(out_f * nblk * 34)?;
let f = self.func("q8_0_split_rp_build");
let cfg = LaunchConfig {
grid_dim: (((out_f * nblk) as u32).div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (of, nb) = (out_f as i32, nblk as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
unsafe {
b.launch(cfg)?;
}
Ok(dst)
}
pub fn build_q4k_rp4(
&self,
t: &mut crate::model::GpuTensor,
) -> Result<(), Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let GpuTensor::Quant {
bytes,
qtype,
row_bytes,
ne,
rp4,
..
} = t
else {
return Ok(());
};
if *qtype != QT_Q4_K || rp4.is_some() || ne.len() != 2 {
return Ok(());
}
let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 144 {
return Ok(());
}
*rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q4_K)?);
Ok(())
}
pub fn build_q6k_rp4(
&self,
t: &mut crate::model::GpuTensor,
) -> Result<(), Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let GpuTensor::Quant {
bytes,
qtype,
row_bytes,
ne,
rp4,
..
} = t
else {
return Ok(());
};
if *qtype != QT_Q6_K || rp4.is_some() || ne.len() != 2 {
return Ok(());
}
let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
if in_f % 256 != 0 || *row_bytes != (in_f / 256) * 210 {
return Ok(());
}
*rp4 = Some(self.build_kq_rp4_raw(bytes, in_f, out_f, QT_Q6_K)?);
Ok(())
}
pub fn build_kq_rp4_raw(
&self,
bytes: &CudaSlice<u8>,
in_f: usize,
out_f: usize,
qtype: i32,
) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
assert!(in_f % 256 == 0);
let nsbk = in_f / 256;
let (sb_bytes, kname) = match qtype {
QT_Q4_K => (144usize, "q4_K_split_rp_build"),
QT_Q6_K => (210usize, "q6_K_split_rp_build"),
_ => return Err(format!("build_kq_rp4_raw: qtype {qtype} has no rp mirror").into()),
};
let mut dst = self.alloc_uninit::<u8>(out_f * nsbk * sb_bytes)?;
let f = self.func(kname);
let cfg = LaunchConfig {
grid_dim: (((out_f * nsbk) as u32).div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (of, nb) = (out_f as i32, nsbk as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&*bytes).arg(&mut dst).arg(&of).arg(&nb);
unsafe {
b.launch(cfg)?;
}
Ok(dst)
}
pub fn kqrp_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| match std::env::var("MEMRA_KQRP").as_deref() {
Ok("0") => false,
Ok(_) => true,
Err(_) => cfg!(memra_hopper_mma),
})
}
pub fn build_q4_rp_swap(
&self,
t: &mut crate::model::GpuTensor,
) -> Result<bool, Box<dyn std::error::Error>> {
self.build_q4_rp4(t)?;
self.gpu.stream().synchronize()?; use crate::model::GpuTensor;
let GpuTensor::Quant { bytes, rp4, rp, .. } = t else {
return Ok(false);
};
match rp4.take() {
Some(split) => {
*bytes = split; *rp = true;
Ok(true)
}
None => Ok(false),
}
}
pub fn q4rp_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_Q4RP")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub fn copy_rows_strided(
&self,
src: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
row_elems: usize,
n_rows: usize,
src_stride: usize,
src_off: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("copy_rows_strided_f32");
let cfg = LaunchConfig {
grid_dim: (((row_elems as u32 + 255) / 256).max(1), n_rows as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (re, nr) = (row_elems as i32, n_rows as i32);
let (st, off) = (src_stride as i64, src_off as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src)
.arg(&mut *dst)
.arg(&re)
.arg(&nr)
.arg(&st)
.arg(&off);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn u32_set_k(
&self,
dst: &mut CudaSlice<u32>,
v: u32,
idx: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("u32_set_k");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
};
let ii = idx as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(dst).arg(&v).arg(&ii);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn i32_add_k(
&self,
d: &mut CudaSlice<i32>,
v: i32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("i32_add_k");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(d).arg(&v);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn i32_iota_from(
&self,
ctr: &CudaSlice<i32>,
dst: &mut CudaSlice<i32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("i32_iota_from");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(ctr).arg(dst).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn u32_map_k(
&self,
buf: &mut CudaSlice<u32>,
map: &CudaSlice<u32>,
idx: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("u32_map_k");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
};
let ii = idx as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(buf).arg(map).arg(&ii);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn u32_pack2(
&self,
a: &CudaSlice<u32>,
off_a: usize,
n1: usize,
b_in: &CudaSlice<u32>,
n2: usize,
out: &mut CudaSlice<u32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("u32_pack2");
let cfg = LaunchConfig::for_num_elems((n1 + n2) as u32);
let (oa, i1, i2) = (off_a as i32, n1 as i32, n2 as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a).arg(&oa).arg(&i1).arg(b_in).arg(&i2).arg(out);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn moe_w_exscale(
&self,
w: &mut CudaSlice<f32>,
sel: &CudaSlice<i32>,
s: &CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("moe_w_exscale");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(w).arg(sel).arg(s).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn moe_w_scale_by_expert(
&self,
w: &mut CudaSlice<f32>,
sel: &CudaSlice<i32>,
macros: &CudaSlice<f32>,
n_expert: usize,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("moe_w_scale_by_expert");
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(64) as u32, 1, 1),
block_dim: (64, 1, 1),
shared_mem_bytes: 0,
};
let (ne, nn) = (n_expert as i32, n as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(w).arg(sel).arg(macros).arg(&ne).arg(&nn);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn moe_gate_up_silu8_dev_q8(
&self,
table: &CudaSlice<u64>,
sel: &cudarc::driver::CudaView<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
n_ff: usize,
n_used: usize,
n_expert: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
macros: &CudaSlice<f32>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
static GU: std::sync::OnceLock<(String, u32)> = std::sync::OnceLock::new();
let (mode, wpb) = GU.get_or_init(|| {
let mode = std::env::var("MEMRA_MOE_DEVQ8_GU").unwrap_or_default();
let wpb = std::env::var("MEMRA_MOE_DEVQ8_WPB")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(4u32)
.clamp(1, 16);
(mode, wpb)
});
let (mode, wpb) = (mode.as_str(), *wpb);
let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
let (inf, nff, ne, rbg, rbu) = (
in_f as i32,
n_ff as i32,
n_expert as i32,
rb_g as i64,
rb_u as i64,
);
let (f, cfg) = match mode {
"1" | "2" | "4" => {
let rpw: u32 = mode.parse().unwrap();
let f = self.func(match rpw {
1 => "moe_gate_up_silu8_dev_q8_r1",
2 => "moe_gate_up_silu8_dev_q8_r2",
_ => "moe_gate_up_silu8_dev_q8_r4",
});
let rows_per_block = (rpw * wpb) as usize;
let gx = n_ff.div_ceil(rows_per_block) as u32;
(
f,
LaunchConfig {
grid_dim: (gx, n_used as u32, 1),
block_dim: (32, wpb, 1),
shared_mem_bytes: 0,
},
)
}
"j8" if n_used <= 32 => (
self.func("moe_gate_up_silu8_dev_q8_j8"),
LaunchConfig {
grid_dim: (n_ff as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
),
"vsm2" => {
let f = self.func("moe_gate_up_silu8_dev_q8_vsm2");
let sh = (rb_g + rb_u) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
sh as i32,
)?;
(
f,
LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: sh,
},
)
}
"vsm" => {
let f = self.func("moe_gate_up_silu8_dev_q8_vsm");
let sh = (rb_g + rb_u) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
sh as i32,
)?;
(
f,
LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: sh,
},
)
}
"sg" => (
self.func("moe_gate_up_silu8_dev_q8_sg"),
LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
},
),
"j8sg" if n_used <= 32 => (
self.func("moe_gate_up_silu8_dev_q8_j8sg"),
LaunchConfig {
grid_dim: (n_ff as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
),
"u64" if in_f == 2048 => (
self.func("moe_gate_up_silu8_dev_q8_u64"),
LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
},
),
"gs4" if in_f == 2048 => (
self.func("moe_gate_up_silu8_dev_q8_gs4"),
LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 4, 1),
shared_mem_bytes: 0,
},
),
"v" | "" => (
self.func("moe_gate_up_silu8_dev_q8_v"),
LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
},
),
"s2" => (
self.func("moe_gate_up_silu8_dev_q8_s2"),
LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 2, 1),
shared_mem_bytes: 0,
},
),
"s2z" => {
let rz = wpb.min(16); (
self.func("moe_gate_up_silu8_dev_q8_s2z"),
LaunchConfig {
grid_dim: (n_ff.div_ceil(rz as usize) as u32, n_used as u32, 1),
block_dim: (32, 2, rz),
shared_mem_bytes: 0,
},
)
}
_ => (
self.func("moe_gate_up_silu8_dev_q8"),
LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
},
),
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(aq)
.arg(ad)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&ne)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu)
.arg(macros);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_down8_fma_dev_q8(
&self,
table: &CudaSlice<u64>,
sel: &cudarc::driver::CudaView<i32>,
w: &cudarc::driver::CudaView<f32>,
aq2: &CudaSlice<i8>,
ad2: &CudaSlice<f32>,
dst: &mut cudarc::driver::CudaViewMut<f32>,
in_f: usize,
out_f: usize,
n_used: usize,
n_expert: usize,
qt: i32,
rb: usize,
) -> Result<(), Box<dyn std::error::Error>> {
static DOWN: std::sync::OnceLock<String> = std::sync::OnceLock::new();
let mode = DOWN.get_or_init(|| std::env::var("MEMRA_MOE_DEVQ8_DOWN").unwrap_or_default());
let (inf, outf, nu, ne, rbi) = (
in_f as i32,
out_f as i32,
n_used as i32,
n_expert as i32,
rb as i64,
);
let (f, cfg) = match mode.as_str() {
m @ ("1" | "2" | "4") if n_used <= 8 => {
let rpw: usize = m.parse().unwrap();
let f = self.func(match rpw {
1 => "moe_down8_fma_dev_q8_w8r1",
2 => "moe_down8_fma_dev_q8_w8r2",
_ => "moe_down8_fma_dev_q8_w8r4",
});
(
f,
LaunchConfig {
grid_dim: (out_f.div_ceil(rpw) as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
)
}
"h2" if in_f == 512 => (
self.func("moe_down8_fma_dev_q8_h2"),
LaunchConfig {
grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
},
),
"" if in_f == 704 && n_used <= 8 => (
self.func("moe_down8_fma_dev_q8_w8r2"),
LaunchConfig {
grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
),
"w8h2v" | "" if in_f == 512 && n_used <= 8 => (
self.func("moe_down8_fma_dev_q8_w8h2v"),
LaunchConfig {
grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
),
"w8h2r2v" if in_f == 512 && n_used <= 8 => (
self.func("moe_down8_fma_dev_q8_w8h2r2v"),
LaunchConfig {
grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
),
"w8h2r2" if in_f == 512 && n_used <= 8 => (
self.func("moe_down8_fma_dev_q8_w8h2r2"),
LaunchConfig {
grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
),
"w8h2" if in_f == 512 && n_used <= 8 => (
self.func("moe_down8_fma_dev_q8_w8h2"),
LaunchConfig {
grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
),
_ => (
self.func("moe_down8_fma_dev_q8"),
LaunchConfig {
grid_dim: (out_f as u32, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
},
),
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(w)
.arg(aq2)
.arg(ad2)
.arg(dst)
.arg(&inf)
.arg(&outf)
.arg(&nu)
.arg(&ne)
.arg(&qt)
.arg(&rbi);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn moe_gate_up_silu8_dev_q8_rows(
&self,
table: &CudaSlice<u64>,
sel: &CudaSlice<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
t: usize,
in_f: usize,
n_ff: usize,
n_used: usize,
n_expert: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
macros: &CudaSlice<f32>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_gate_up_silu8_dev_q8_v_rows");
let mut act = self.alloc_uninit::<f32>(t * n_used * n_ff)?;
let cfg = LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, t as u32),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (inf, nff, ne, nu, rbg, rbu) = (
in_f as i32,
n_ff as i32,
n_expert as i32,
n_used as i32,
rb_g as i64,
rb_u as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(aq)
.arg(ad)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&ne)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu)
.arg(&nu)
.arg(macros);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_down8_fma_dev_q8_rows(
&self,
table: &CudaSlice<u64>,
sel: &CudaSlice<i32>,
w: &CudaSlice<f32>,
aq2: &CudaSlice<i8>,
ad2: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
t: usize,
in_f: usize,
out_f: usize,
n_used: usize,
n_expert: usize,
qt: i32,
rb: usize,
) -> Result<(), Box<dyn std::error::Error>> {
assert!(
in_f == 512 && n_used <= 8,
"down rows twin is w8h2v shape-gated"
);
let f = self.func("moe_down8_fma_dev_q8_w8h2v_rows");
let cfg = LaunchConfig {
grid_dim: (out_f.div_ceil(2) as u32, 1, t as u32),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
};
let (inf, outf, nu, ne, rbi) = (
in_f as i32,
out_f as i32,
n_used as i32,
n_expert as i32,
rb as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(w)
.arg(aq2)
.arg(ad2)
.arg(dst)
.arg(&inf)
.arg(&outf)
.arg(&nu)
.arg(&ne)
.arg(&qt)
.arg(&rbi);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn moe_gate_up_silu8_dev_q8_csr(
&self,
table: &CudaSlice<u64>,
sel: &CudaSlice<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
n_pairs: usize,
in_f: usize,
n_ff: usize,
n_used: usize,
n_expert: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_gate_up_silu8_dev_q8_csr_iq4");
let mut act = self.alloc_uninit::<f32>(n_pairs * n_ff)?;
let cfg = LaunchConfig {
grid_dim: (n_ff as u32, n_pairs as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (inf, nff, ne, nu, npi, rbg, rbu) = (
in_f as i32,
n_ff as i32,
n_expert as i32,
n_used as i32,
n_pairs as i32,
rb_g as i64,
rb_u as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(aq)
.arg(ad)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&ne)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu)
.arg(&nu)
.arg(&npi);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_down8_fma_dev_q8_variant(
&self,
variant: &str,
table: &CudaSlice<u64>,
sel: &cudarc::driver::CudaView<i32>,
w: &cudarc::driver::CudaView<f32>,
aq2: &CudaSlice<i8>,
ad2: &CudaSlice<f32>,
dst: &mut cudarc::driver::CudaViewMut<f32>,
in_f: usize,
out_f: usize,
n_used: usize,
n_expert: usize,
qt: i32,
rb: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let (inf, outf, nu, ne, rbi) = (
in_f as i32,
out_f as i32,
n_used as i32,
n_expert as i32,
rb as i64,
);
let (f, cfg) = match variant {
"w8h2" | "w8h2v" => (
self.func(if variant == "w8h2" {
"moe_down8_fma_dev_q8_w8h2"
} else {
"moe_down8_fma_dev_q8_w8h2v"
}),
LaunchConfig {
grid_dim: (out_f.div_ceil(2) as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
),
"w8h2r2" | "w8h2r2v" => (
self.func(if variant == "w8h2r2" {
"moe_down8_fma_dev_q8_w8h2r2"
} else {
"moe_down8_fma_dev_q8_w8h2r2v"
}),
LaunchConfig {
grid_dim: (out_f.div_ceil(4) as u32, 1, 1),
block_dim: (32, n_used as u32, 1),
shared_mem_bytes: 0,
},
),
_ => (
self.func("moe_down8_fma_dev_q8"),
LaunchConfig {
grid_dim: (out_f as u32, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
},
),
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(w)
.arg(aq2)
.arg(ad2)
.arg(dst)
.arg(&inf)
.arg(&outf)
.arg(&nu)
.arg(&ne)
.arg(&qt)
.arg(&rbi);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn moe_gate_up_silu8_dev_q8_variant(
&self,
variant: &str,
table: &CudaSlice<u64>,
sel: &cudarc::driver::CudaView<i32>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
n_ff: usize,
n_used: usize,
n_expert: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let mut act = self.alloc_uninit::<f32>(n_used * n_ff)?;
let (inf, nff, ne, rbg, rbu) = (
in_f as i32,
n_ff as i32,
n_expert as i32,
rb_g as i64,
rb_u as i64,
);
let f = self.func(if variant == "v" {
"moe_gate_up_silu8_dev_q8_v"
} else {
"moe_gate_up_silu8_dev_q8"
});
let cfg = LaunchConfig {
grid_dim: (n_ff as u32, n_used as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(aq)
.arg(ad)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&ne)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
pub fn moe_gate_up_silu8_dev(
&self,
table: &CudaSlice<u64>,
sel: &cudarc::driver::CudaView<i32>,
x: &cudarc::driver::CudaView<f32>,
in_f: usize,
n_ff: usize,
n_used: usize,
n_expert: usize,
qt_g: i32,
qt_u: i32,
rb_g: usize,
rb_u: usize,
macros: &CudaSlice<f32>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("moe_gate_up_silu8_dev");
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),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (inf, nff, ne, rbg, rbu) = (
in_f as i32,
n_ff as i32,
n_expert as i32,
rb_g as i64,
rb_u as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(x)
.arg(&mut act)
.arg(&inf)
.arg(&nff)
.arg(&ne)
.arg(&qt_g)
.arg(&qt_u)
.arg(&rbg)
.arg(&rbu)
.arg(macros);
unsafe {
b.launch(cfg)?;
}
Ok(act)
}
#[allow(clippy::too_many_arguments)]
pub fn moe_down8_fma_dev(
&self,
table: &CudaSlice<u64>,
sel: &cudarc::driver::CudaView<i32>,
w: &cudarc::driver::CudaView<f32>,
act: &CudaSlice<f32>,
dst: &mut cudarc::driver::CudaViewMut<f32>,
in_f: usize,
out_f: usize,
n_used: usize,
n_expert: usize,
qt: i32,
rb: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("moe_down8_fma_dev");
let cfg = LaunchConfig {
grid_dim: (out_f as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, nu, ne, rbv) = (
in_f as i32,
out_f as i32,
n_used as i32,
n_expert as i32,
rb as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(table)
.arg(sel)
.arg(w)
.arg(act)
.arg(dst)
.arg(&inf)
.arg(&outf)
.arg(&nu)
.arg(&ne)
.arg(&qt)
.arg(&rbv);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn axpy_into(
&self,
src: &CudaSlice<f32>,
alpha: f32,
dst: &mut cudarc::driver::CudaViewMut<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("axpy_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let (a, ni) = (alpha, n as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src).arg(dst).arg(&a).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn add_scaled_rows(
&self,
src: &CudaSlice<f32>,
scale: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("add_scaled_rows_f32");
let cfg = LaunchConfig::for_num_elems((ncols * nrows) as u32);
let (nc, nr) = (ncols as i32, nrows as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src).arg(scale).arg(dst).arg(&nc).arg(&nr);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gather_rows(
&self,
src: &CudaSlice<f32>,
idx: &CudaSlice<i32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
m_e: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gather_rows_f32");
let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
let (nc, me) = (ncols as i32, m_e as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src).arg(idx).arg(dst).arg(&nc).arg(&me);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn scatter_slot(
&self,
src: &CudaSlice<f32>,
tok_idx: &CudaSlice<i32>,
slot_idx: &CudaSlice<i32>,
weight: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
wbuf: &mut CudaSlice<f32>,
ncols: usize,
n_used: usize,
m_e: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("scatter_add_slot_f32");
let cfg = LaunchConfig::for_num_elems((m_e * ncols) as u32);
let (nc, nu, me) = (ncols as i32, n_used as i32, m_e as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src)
.arg(tok_idx)
.arg(slot_idx)
.arg(weight)
.arg(dst)
.arg(wbuf)
.arg(&nc)
.arg(&nu)
.arg(&me);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn reduce_slots(
&self,
slots: &CudaSlice<f32>,
wbuf: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
n_used: usize,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("reduce_slots_f32");
let cfg = LaunchConfig::for_num_elems((t * ncols) as u32);
let (nc, nu, ti) = (ncols as i32, n_used as i32, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(slots).arg(wbuf).arg(dst).arg(&nc).arg(&nu).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn quantize_q8_1_view(
&self,
x: &cudarc::driver::CudaView<f32>,
m: usize,
in_f: usize,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let f = self.func("quantize_q8_1");
let nblk = in_f / 32;
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);
let (inf, mi) = (in_f as i32, m as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
unsafe {
b.launch(cfg)?;
}
Ok((q, d))
}
pub fn quantize_q8_1(
&self,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let nblk = in_f / 32;
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);
let (inf, mi) = (in_f as i32, m as i32);
if Self::pdl_on() && Self::pdl_wb_on() {
{
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (px, _g0) = x.device_ptr(s);
let (pq, _g1) = q.device_ptr_mut(s);
let (pd, _g2) = d.device_ptr_mut(s);
let mut ps = [
&px as *const _ as *mut std::ffi::c_void,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&inf as *const _ as *mut _,
&mi as *const _ as *mut _,
];
unsafe {
self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?;
}
}
return Ok((q, d));
}
let f = self.func("quantize_q8_1");
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(&mut q).arg(&mut d).arg(&inf).arg(&mi);
unsafe {
b.launch(cfg)?;
}
Ok((q, d))
}
pub fn quantize_fp4_act(
&self,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
) -> Result<(CudaSlice<u32>, CudaSlice<u8>), Box<dyn std::error::Error>> {
let f = self.func("quantize_fp4_act");
let nb16 = in_f / 16;
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);
let (inf, mi) = (in_f as i32, m as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(&mut aq4).arg(&mut ad4).arg(&inf).arg(&mi);
unsafe {
b.launch(cfg)?;
}
Ok((aq4, ad4))
}
pub fn qmatvec_gemm_nvfp4_fp4(
&self,
bytes: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
scale: f32,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
assert!(
in_f % 64 == 0,
"FP4 GEMM requires in_f % 64 == 0, got {in_f}"
);
let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
let mut y = self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)?;
if scale != 1.0 {
self.scale_inplace(&mut y, scale, m * out_f)?;
}
Ok(y)
}
fn fp4_gemm_launch(
&self,
bytes: &CudaSlice<u8>,
aq4: &CudaSlice<u32>,
ad4: &CudaSlice<u8>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("qmatvec_gemm_nvfp4_fp4");
let mut y = self.alloc_uninit::<f32>(m * out_f)?; const BM: u32 = 64;
const BN: u32 = 256;
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + BM - 1) / BM, (m as u32 + BN - 1) / BN, 1),
block_dim: (32, 4, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(aq4)
.arg(ad4)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn qmatvec_gemm_nvfp4_fp4_raw(
&self,
bytes: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
assert!(
in_f % 64 == 0,
"FP4 GEMM requires in_f % 64 == 0, got {in_f}"
);
let (aq4, ad4) = self.quantize_fp4_act(x, m, in_f)?;
self.fp4_gemm_launch(bytes, &aq4, &ad4, m, in_f, out_f, row_bytes)
}
pub fn qmatvec_q8_0_fast(
&self,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
let f = self.func("qmatvec_q8_0_dp4a");
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,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(w)
.arg(&aq)
.arg(&ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
#[allow(non_snake_case)] pub fn qmatvec_q4_K_fast(
&self,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
let f = self.func("qmatvec_q4_K_dp4a");
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,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(w)
.arg(&aq)
.arg(&ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
#[allow(non_snake_case)] pub fn qmatvec_q6_K_fast(
&self,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
let f = self.func("qmatvec_q6_K_dp4a");
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,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(w)
.arg(&aq)
.arg(&ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
#[allow(non_snake_case)] pub fn qmatvec_q5_K_fast(
&self,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
self.qmatvec_dp4a_named("qmatvec_q5_K_dp4a", w, x, m, in_f, out_f, row_bytes)
}
#[allow(non_snake_case)] pub fn qmatvec_q3_K_fast(
&self,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
self.qmatvec_dp4a_named("qmatvec_q3_K_dp4a", w, x, m, in_f, out_f, row_bytes)
}
pub fn qmatvec_nvfp4_fast_rp(
&self,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
assert!(
in_f % 64 == 0,
"NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
);
self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a_rp", w, x, m, in_f, out_f, row_bytes)
}
pub fn qmatvec_nvfp4_fast(
&self,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
assert!(
in_f % 64 == 0,
"NVFP4 dp4a requires in_f % 64 == 0, got {in_f}"
);
self.qmatvec_dp4a_named("qmatvec_nvfp4_dp4a", w, x, m, in_f, out_f, row_bytes)
}
#[allow(non_snake_case)] pub fn qmatvec_iq4_XS_fast(
&self,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
self.qmatvec_dp4a_named("qmatvec_iq4_XS_dp4a", w, x, m, in_f, out_f, row_bytes)
}
fn qmatvec_dp4a_named(
&self,
name: &str,
w: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
let f = self.func(name);
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,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(w)
.arg(&aq)
.arg(&ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn htod(&self, v: &[f32]) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
Ok(self.gpu.stream().clone_htod(v)?)
}
pub fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
Ok(self.gpu.stream().clone_htod(v)?)
}
pub fn htod_i8(&self, v: &[i8]) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
Ok(self.gpu.stream().clone_htod(v)?)
}
pub fn htod_u64(&self, v: &[u64]) -> Result<CudaSlice<u64>, Box<dyn std::error::Error>> {
Ok(self.gpu.stream().clone_htod(v)?)
}
pub fn dtoh_view(
&self,
d: &cudarc::driver::CudaView<f32>,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let v = self.gpu.stream().clone_dtoh(d)?;
self.gpu.stream().synchronize()?;
Ok(v)
}
pub fn dtoh(&self, d: &CudaSlice<f32>) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let v = self.gpu.stream().clone_dtoh(d)?;
self.gpu.stream().synchronize()?;
Ok(v)
}
pub fn dtoh_pair(
&self,
a: &CudaSlice<f32>,
b: &CudaSlice<f32>,
) -> Result<(Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
let av = self.gpu.stream().clone_dtoh(a)?;
let bv = self.gpu.stream().clone_dtoh(b)?;
self.gpu.stream().synchronize()?;
Ok((av, bv))
}
pub fn dtoh_i32(&self, d: &CudaSlice<i32>) -> Result<Vec<i32>, Box<dyn std::error::Error>> {
let v = self.gpu.stream().clone_dtoh(d)?;
self.gpu.stream().synchronize()?;
Ok(v)
}
pub fn dtoh_u8(&self, d: &CudaSlice<u8>) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
let v = self.gpu.stream().clone_dtoh(d)?;
self.gpu.stream().synchronize()?;
Ok(v)
}
pub fn dtoh_u8_view(
&self,
d: &cudarc::driver::CudaView<u8>,
) -> Result<Vec<u8>, Box<dyn std::error::Error>> {
let v = self.gpu.stream().clone_dtoh(d)?;
self.gpu.stream().synchronize()?;
Ok(v)
}
pub fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let s = self.gpu.stream().alloc_zeros::<f32>(n)?;
self.keep_if_capturing(&s);
Ok(s)
}
pub fn prob_of_token_device(
&self,
logits: &CudaSlice<f32>,
tok: &CudaSlice<u32>,
n_vocab: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let nb = ARGMAX_NB;
let mut part = self.alloc_uninit::<f32>(nb)?;
let mut p = self.alloc_uninit::<f32>(1)?;
let f1 = self.func("prob_of_token_partial_f32");
let cfg1 = LaunchConfig {
grid_dim: (nb as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let nv = n_vocab as i32;
let __s_b1 = self.gpu.stream();
let mut b1 = __s_b1.launch_builder(&f1);
b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
unsafe {
b1.launch(cfg1)?;
}
let f2 = self.func("prob_of_token_final_f32");
let cfg2 = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let nbi = nb as i32;
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f2);
b2.arg(&part).arg(&mut p).arg(&nbi);
unsafe {
b2.launch(cfg2)?;
}
Ok(p)
}
pub fn prob_of_token_device_col(
&self,
logits: &CudaSlice<f32>,
tok_all: &CudaSlice<u32>,
tok_idx: usize,
p_out: &mut CudaSlice<f32>,
p_idx: usize,
n_vocab: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let tok_v = tok_all.slice(tok_idx..tok_idx + 1);
let mut p_v = p_out.slice_mut(p_idx..p_idx + 1);
let nb = ARGMAX_NB;
let mut part = self.alloc_uninit::<f32>(nb)?;
let f1 = self.func("prob_of_token_partial_f32");
let cfg1 = LaunchConfig {
grid_dim: (nb as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let nv = n_vocab as i32;
let __s_b1 = self.gpu.stream();
let mut b1 = __s_b1.launch_builder(&f1);
b1.arg(logits).arg(&tok_v).arg(&mut part).arg(&nv);
unsafe {
b1.launch(cfg1)?;
}
let f2 = self.func("prob_of_token_final_f32");
let cfg2 = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let nbi = nb as i32;
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f2);
b2.arg(&part).arg(&mut p_v).arg(&nbi);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
pub fn prob_of_token_device_into(
&self,
logits: &CudaSlice<f32>,
tok: &CudaSlice<u32>,
p_out: &mut CudaSlice<f32>,
n_vocab: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let nb = ARGMAX_NB;
let mut part = self.alloc_uninit::<f32>(nb)?;
let f1 = self.func("prob_of_token_partial_f32");
let cfg1 = LaunchConfig {
grid_dim: (nb as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let nv = n_vocab as i32;
let __s_b1 = self.gpu.stream();
let mut b1 = __s_b1.launch_builder(&f1);
b1.arg(logits).arg(tok).arg(&mut part).arg(&nv);
unsafe {
b1.launch(cfg1)?;
}
let f2 = self.func("prob_of_token_final_f32");
let cfg2 = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let nbi = nb as i32;
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f2);
b2.arg(&part).arg(p_out).arg(&nbi);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
pub fn argmax_token_device(
&self,
logits: &CudaSlice<f32>,
n_vocab: usize,
) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
let mut tok = unsafe { self.gpu.stream().alloc::<u32>(1)? };
self.argmax_token_device_into(logits, &mut tok, n_vocab)?;
Ok(tok)
}
pub fn argmax_token_device_into(
&self,
logits: &CudaSlice<f32>,
tok: &mut CudaSlice<u32>,
n_vocab: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let nb = ARGMAX_NB;
let f1 = self.func("argmax_partial_f32");
let f2 = self.func("argmax_final_f32");
let mut guard = self.argmax_partials.lock().unwrap();
if guard.is_none() {
let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
*guard = Some((pv, pi));
}
let (part_v, part_i) = guard.as_mut().unwrap();
let nv = n_vocab as i32;
let nbi = nb as i32;
let cfg1 = LaunchConfig {
grid_dim: (nb as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b1 = self.gpu.stream();
let mut b1 = __s_b1.launch_builder(&f1);
b1.arg(logits).arg(&mut *part_v).arg(&mut *part_i).arg(&nv);
unsafe {
b1.launch(cfg1)?;
}
let cfg2 = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f2);
b2.arg(&*part_v).arg(&*part_i).arg(tok).arg(&nbi);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
pub fn argmax_token_device_col(
&self,
logits: &CudaSlice<f32>,
col: usize,
n_vocab: usize,
toks: &mut CudaSlice<u32>,
out_idx: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let nb = ARGMAX_NB;
let f1 = self.func("argmax_partial_f32");
let f2 = self.func("argmax_final_f32");
let mut guard = self.argmax_partials.lock().unwrap();
if guard.is_none() {
let pv = self.gpu.stream().alloc_zeros::<f32>(nb)?;
let pi = self.gpu.stream().alloc_zeros::<i32>(nb)?;
*guard = Some((pv, pi));
}
let (part_v, part_i) = guard.as_mut().unwrap();
let col_view = logits.slice(col * n_vocab..(col + 1) * n_vocab);
let nv = n_vocab as i32;
let nbi = nb as i32;
let cfg1 = LaunchConfig {
grid_dim: (nb as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b1 = self.gpu.stream();
let mut b1 = __s_b1.launch_builder(&f1);
b1.arg(&col_view)
.arg(&mut *part_v)
.arg(&mut *part_i)
.arg(&nv);
unsafe {
b1.launch(cfg1)?;
}
let mut tok_view = toks.slice_mut(out_idx..out_idx + 1);
let cfg2 = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f2);
b2.arg(&*part_v).arg(&*part_i).arg(&mut tok_view).arg(&nbi);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
pub fn htod_u32_v(&self, v: &[u32]) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
Ok(self.gpu.stream().clone_htod(v)?)
}
pub fn dtoh_u32(&self, d: &CudaSlice<u32>) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
let v = self.gpu.stream().clone_dtoh(d)?;
self.gpu.stream().synchronize()?;
Ok(v)
}
pub fn htod_u32_into(
&self,
dst: &mut CudaSlice<u32>,
src: &[u32],
) -> Result<(), Box<dyn std::error::Error>> {
let mut view = dst.slice_mut(0..src.len());
self.gpu.stream().memcpy_htod(src, &mut view)?;
Ok(())
}
pub fn htod_i32_into(
&self,
dst: &mut CudaSlice<i32>,
src: &[i32],
) -> Result<(), Box<dyn std::error::Error>> {
let mut view = dst.slice_mut(0..src.len());
self.gpu.stream().memcpy_htod(src, &mut view)?;
Ok(())
}
pub fn alloc_u32_zeroed(&self, n: usize) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
let s = self.gpu.stream().alloc_zeros::<u32>(n)?;
self.keep_if_capturing(&s);
Ok(s)
}
pub fn embed_gather_device_into(
&self,
embd: &CudaSlice<u8>,
token_d: &CudaSlice<u32>,
x_out: &mut CudaSlice<f32>,
n_embd: usize,
qtype: i32,
row_bytes: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("embed_gather_u32");
let cfg = LaunchConfig {
grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(embd)
.arg(token_d)
.arg(x_out)
.arg(&ne)
.arg(&qt)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn dtoh_i32_one(&self, d: &CudaSlice<i32>) -> Result<i32, Box<dyn std::error::Error>> {
let v = self.gpu.stream().clone_dtoh(d)?;
self.gpu.stream().synchronize()?;
Ok(v[0])
}
pub fn i32_set_k(
&self,
dst: &mut CudaSlice<i32>,
v: i32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("i32_set_k");
let cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
};
let idx = 0i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(dst).arg(&v).arg(&idx);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn set_i32_one(
&self,
d: &mut CudaSlice<i32>,
v: i32,
) -> Result<(), Box<dyn std::error::Error>> {
self.gpu.stream().memcpy_htod(&[v], d)?;
Ok(())
}
pub fn set_u32_one(
&self,
d: &mut CudaSlice<u32>,
v: u32,
) -> Result<(), Box<dyn std::error::Error>> {
self.gpu.stream().memcpy_htod(&[v], d)?;
Ok(())
}
pub fn dtoh_u32_one(&self, d: &CudaSlice<u32>) -> Result<u32, Box<dyn std::error::Error>> {
let v = self.gpu.stream().clone_dtoh(d)?;
self.gpu.stream().synchronize()?;
Ok(v[0])
}
pub fn upload_u8(&self, bytes: &[u8]) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
Ok(self.gpu.stream().clone_htod(bytes)?)
}
pub fn embed_gather_device(
&self,
embd: &CudaSlice<u8>,
token_d: &CudaSlice<u32>,
n_embd: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("embed_gather_u32");
let mut x = self.alloc_uninit::<f32>(n_embd)?;
let cfg = LaunchConfig {
grid_dim: (((n_embd as u32 + 255) / 256).max(1), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (ne, qt, rb) = (n_embd as i32, qtype, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(embd)
.arg(token_d)
.arg(&mut x)
.arg(&ne)
.arg(&qt)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(x)
}
pub fn embed_gather_device_t(
&self,
embd: &CudaSlice<u8>,
tokens: &[u32],
n_embd: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let t = tokens.len();
let tok_d = self.gpu.stream().clone_htod(tokens)?;
let f = self.func("embed_gather_u32_t");
let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
let cfg = LaunchConfig {
grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(embd)
.arg(&tok_d)
.arg(&mut x)
.arg(&ne)
.arg(&qt)
.arg(&rb)
.arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(x)
}
pub fn embed_gather_device_tv(
&self,
embd: &CudaSlice<u8>,
tok_v: &cudarc::driver::CudaView<u32>,
t: usize,
n_embd: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("embed_gather_u32_t");
let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
let cfg = LaunchConfig {
grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(embd)
.arg(tok_v)
.arg(&mut x)
.arg(&ne)
.arg(&qt)
.arg(&rb)
.arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(x)
}
pub fn embed_gather_device_td(
&self,
embd: &CudaSlice<u8>,
tok_d: &CudaSlice<u32>,
t: usize,
n_embd: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("embed_gather_u32_t");
let mut x = self.alloc_uninit::<f32>(t * n_embd)?;
let cfg = LaunchConfig {
grid_dim: (((n_embd as u32 + 255) / 256).max(1), t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (ne, qt, rb, ti) = (n_embd as i32, qtype, row_bytes as i64, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(embd)
.arg(tok_d)
.arg(&mut x)
.arg(&ne)
.arg(&qt)
.arg(&rb)
.arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(x)
}
#[inline]
fn keep_if_capturing<T: cudarc::driver::DeviceRepr + Send + 'static>(&self, s: &CudaSlice<T>) {
if self
.capture_keep_on
.load(std::sync::atomic::Ordering::Relaxed)
{
self.capture_keep.lock().unwrap().push(Box::new(s.clone()));
}
}
fn alloc_uninit<T: cudarc::driver::DeviceRepr + Send + 'static>(
&self,
n: usize,
) -> Result<CudaSlice<T>, Box<dyn std::error::Error>> {
let mut s = unsafe { self.gpu.stream().alloc::<T>(n)? };
{
static Z: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
if *Z.get_or_init(|| std::env::var("MEMRA_DEBUG_ZERO_ALLOCS").as_deref() == Ok("1")) {
use cudarc::driver::DevicePtrMut;
let n_bytes = s.len() * std::mem::size_of::<T>();
let stream = self.gpu.stream();
let (p_, _g) = s.device_ptr_mut(&stream);
unsafe {
cudarc::driver::sys::cuMemsetD8Async(p_, 0, n_bytes, stream.cu_stream())
.result()?;
}
}
}
self.keep_if_capturing(&s);
Ok(s)
}
pub fn uninit_q8_pair(
&self,
n: usize,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
Ok((
self.alloc_uninit::<i8>(n)?,
self.alloc_uninit::<f32>(n / 32)?,
))
}
pub fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
self.alloc_uninit::<f32>(n)
}
pub fn alloc_i8_uninit(&self, n: usize) -> Result<CudaSlice<i8>, Box<dyn std::error::Error>> {
self.alloc_uninit::<i8>(n)
}
#[allow(clippy::too_many_arguments)]
pub fn rms_norm3(
&self,
x: &CudaSlice<f32>,
w0: &CudaSlice<f32>,
w1: &CudaSlice<f32>,
w2: &CudaSlice<f32>,
d0: &mut CudaSlice<f32>,
d1: &mut CudaSlice<f32>,
d2: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rms_norm3_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(w0)
.arg(w1)
.arg(w2)
.arg(d0)
.arg(d1)
.arg(d2)
.arg(&nc)
.arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn qkvnorm_w_on_prefill(rows: usize, ncols: usize) -> bool {
static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*WARP_ON.get_or_init(|| {
std::env::var("MEMRA_QKVNORM_W")
.map(|v| v != "0")
.unwrap_or(true)
}) && ncols % 4 == 0
&& rows >= 64
}
#[allow(clippy::too_many_arguments)]
pub fn rms_norm_qkv_w4b(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
wq: &CudaSlice<f32>,
wk: &CudaSlice<f32>,
wv: &CudaSlice<f32>,
dq: &mut CudaSlice<f32>,
dk: &mut CudaSlice<f32>,
dv: &mut CudaSlice<f32>,
dvb: &mut CudaSlice<u8>,
ncols: usize,
rq: usize,
rk: usize,
eps: f32,
vf16: bool,
) -> Result<(), Box<dyn std::error::Error>> {
assert!(ncols % 4 == 0 && rq + 2 * rk >= 64);
let f = self.func("rms_norm_qkv_w4b_f32");
let rows = (rq + 2 * rk) as u32;
let cfg = LaunchConfig {
grid_dim: (rows.div_ceil(8), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
let vf = vf16 as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(wq)
.arg(wk)
.arg(wv)
.arg(dq)
.arg(dk)
.arg(dv)
.arg(&mut *dvb)
.arg(&nc)
.arg(&rqi)
.arg(&rki)
.arg(&rvi)
.arg(&e)
.arg(&vf);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn rms_norm_qkv(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
wq: &CudaSlice<f32>,
wk: &CudaSlice<f32>,
wv: &CudaSlice<f32>,
dq: &mut CudaSlice<f32>,
dk: &mut CudaSlice<f32>,
dv: &mut CudaSlice<f32>,
ncols: usize,
rq: usize,
rk: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
static WARP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let warp_on = *WARP_ON.get_or_init(|| {
std::env::var("MEMRA_QKVNORM_W")
.map(|v| v != "0")
.unwrap_or(true)
});
if warp_on && ncols % 4 == 0 && rq + 2 * rk >= 64 {
let f = self.func("rms_norm_qkv_w4_f32");
let rows = (rq + 2 * rk) as u32;
let cfg = LaunchConfig {
grid_dim: (rows.div_ceil(8), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (nc, rqi, rki, rvi, e) = (ncols as i32, rq as i32, rk as i32, rk as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(wq)
.arg(wk)
.arg(wv)
.arg(dq)
.arg(dk)
.arg(dv)
.arg(&nc)
.arg(&rqi)
.arg(&rki)
.arg(&rvi)
.arg(&e);
unsafe {
b.launch(cfg)?;
}
return Ok(());
}
let f = self.func("rms_norm_qkv_f32");
let grid = (rq + 2 * rk) as u32;
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, rqi, rki, e) = (ncols as i32, rq as i32, rk as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(wq)
.arg(wk)
.arg(wv)
.arg(dq)
.arg(dk)
.arg(dv)
.arg(&nc)
.arg(&rqi)
.arg(&rki)
.arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn rms_norm2x(
&self,
a: &CudaSlice<f32>,
bb: &CudaSlice<f32>,
wa: &CudaSlice<f32>,
wb: &CudaSlice<f32>,
da: &mut CudaSlice<f32>,
db: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rms_norm2x_f32");
let cfg = LaunchConfig {
grid_dim: (2 * nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a)
.arg(bb)
.arg(wa)
.arg(wb)
.arg(da)
.arg(db)
.arg(&nc)
.arg(&nr)
.arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn softcap(
&self,
y: &mut CudaSlice<f32>,
cap: f32,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("softcap_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(y).arg(&cap).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn mask_ids_rows(
&self,
y: &mut CudaSlice<f32>,
ids: &CudaSlice<i32>,
n_ids: usize,
n_vocab: usize,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("mask_ids_rows_f32");
let cfg = LaunchConfig::for_num_elems((n_ids * t) as u32);
let (ni, nv, ti) = (n_ids as i32, n_vocab as i32, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(y).arg(ids).arg(&ni).arg(&nv).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn add_scale_rms_norm(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
c: f32,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("add_scale_rms_norm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, e2) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a)
.arg(b_in)
.arg(&c)
.arg(w)
.arg(res)
.arg(dst)
.arg(&nc)
.arg(&e2);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn add_scale_rms_norm_q8_1(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
c: f32,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let (nc, e2) = (ncols as i32, eps);
if Self::pdl_on() && Self::pdl_wb_on() {
{
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pa, _g0) = a.device_ptr(s);
let (pb, _g1) = b_in.device_ptr(s);
let (pw, _g2) = w.device_ptr(s);
let (pr, _g3) = res.device_ptr_mut(s);
let (pq, _g4) = out_q.device_ptr_mut(s);
let (pd, _g5) = out_d.device_ptr_mut(s);
let mut ps = [
&pa as *const _ as *mut std::ffi::c_void,
&pb as *const _ as *mut _,
&c as *const _ as *mut _,
&pw as *const _ as *mut _,
&pr as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
&e2 as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"add_scale_rms_norm_q8_1",
(nrows as u32, 1, 1),
(rms_block(), 1, 1),
&mut ps,
)?;
}
}
return Ok((out_q, out_d));
}
let f = self.func("add_scale_rms_norm_q8_1");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a)
.arg(b_in)
.arg(&c)
.arg(w)
.arg(res)
.arg(&mut out_q)
.arg(&mut out_d)
.arg(&nc)
.arg(&e2);
unsafe {
b.launch(cfg)?;
}
Ok((out_q, out_d))
}
#[allow(clippy::too_many_arguments)]
pub fn add_scale_rms_norm_q8_1_into(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
c: f32,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
out_q: &mut CudaSlice<i8>,
out_d: &mut CudaSlice<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
let (nc, e2) = (ncols as i32, eps);
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pa, _g0) = a.device_ptr(s);
let (pb, _g1) = b_in.device_ptr(s);
let (pw, _g2) = w.device_ptr(s);
let (pr, _g3) = res.device_ptr_mut(s);
let (pq, _g4) = out_q.device_ptr_mut(s);
let (pd, _g5) = out_d.device_ptr_mut(s);
let mut ps = [
&pa as *const _ as *mut std::ffi::c_void,
&pb as *const _ as *mut _,
&c as *const _ as *mut _,
&pw as *const _ as *mut _,
&pr as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
&e2 as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"add_scale_rms_norm_q8_1",
(nrows as u32, 1, 1),
(rms_block(), 1, 1),
&mut ps,
)?;
}
return Ok(());
}
let f = self.func("add_scale_rms_norm_q8_1");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a)
.arg(b_in)
.arg(&c)
.arg(w)
.arg(res)
.arg(&mut *out_q)
.arg(&mut *out_d)
.arg(&nc)
.arg(&e2);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn rms_pre_add_scale_rms_norm_q8_1(
&self,
a: &CudaSlice<f32>,
wa: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
c: f32,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let (nc, e2) = (ncols as i32, eps);
if Self::pdl_on() {
{
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pa, _g0) = a.device_ptr(s);
let (pwa, _g1) = wa.device_ptr(s);
let (pb, _g2) = b_in.device_ptr(s);
let (pw, _g3) = w.device_ptr(s);
let (pr, _g4) = res.device_ptr_mut(s);
let (pq, _g5) = out_q.device_ptr_mut(s);
let (pd, _g6) = out_d.device_ptr_mut(s);
let mut ps = [
&pa as *const _ as *mut std::ffi::c_void,
&pwa as *const _ as *mut _,
&pb as *const _ as *mut _,
&c as *const _ as *mut _,
&pw as *const _ as *mut _,
&pr as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
&e2 as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"rms_pre_add_scale_rms_norm_q8_1",
(nrows as u32, 1, 1),
(rms_block(), 1, 1),
&mut ps,
)?;
}
}
return Ok((out_q, out_d));
}
let f = self.func("rms_pre_add_scale_rms_norm_q8_1");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
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);
unsafe {
b.launch(cfg)?;
}
Ok((out_q, out_d))
}
pub fn gelu_tanh_mul_q8_1(
&self,
gate: &CudaSlice<f32>,
up: &cudarc::driver::CudaView<f32>,
act: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
debug_assert!(ncols % 128 == 0);
let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let nc = ncols as i32;
if Self::pdl_on() {
{
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pg, _g0) = gate.device_ptr(s);
let (pu, _g1) = up.device_ptr(s);
let (pact, _g2) = act.device_ptr_mut(s);
let (pq, _g3) = out_q.device_ptr_mut(s);
let (pd, _g4) = out_d.device_ptr_mut(s);
let mut ps = [
&pg as *const _ as *mut std::ffi::c_void,
&pu as *const _ as *mut _,
&pact as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"gelu_tanh_mul_q8_1",
(nrows as u32, 1, 1),
(rms_block(), 1, 1),
&mut ps,
)?;
}
}
return Ok((out_q, out_d));
}
let f = self.func("gelu_tanh_mul_q8_1");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate)
.arg(up)
.arg(act)
.arg(&mut out_q)
.arg(&mut out_d)
.arg(&nc);
unsafe {
b.launch(cfg)?;
}
Ok((out_q, out_d))
}
#[allow(clippy::too_many_arguments)]
pub fn gelu_tanh_mul_q8_1_into(
&self,
gate: &CudaSlice<f32>,
up: &cudarc::driver::CudaView<f32>,
act: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
out_q: &mut CudaSlice<i8>,
out_d: &mut CudaSlice<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert!(ncols % 128 == 0);
debug_assert!(out_q.len() >= nrows * ncols && out_d.len() >= nrows * (ncols / 32));
let nc = ncols as i32;
if Self::pdl_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pg, _g0) = gate.device_ptr(s);
let (pu, _g1) = up.device_ptr(s);
let (pact, _g2) = act.device_ptr_mut(s);
let (pq, _g3) = out_q.device_ptr_mut(s);
let (pd, _g4) = out_d.device_ptr_mut(s);
let mut ps = [
&pg as *const _ as *mut std::ffi::c_void,
&pu as *const _ as *mut _,
&pact as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"gelu_tanh_mul_q8_1",
(nrows as u32, 1, 1),
(rms_block(), 1, 1),
&mut ps,
)?;
}
return Ok(());
}
let f = self.func("gelu_tanh_mul_q8_1");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate)
.arg(up)
.arg(&mut *act)
.arg(&mut *out_q)
.arg(&mut *out_d)
.arg(&nc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn add_rms_norm3_q8z(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
w0: &CudaSlice<f32>,
w1: &CudaSlice<f32>,
w2: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
out1: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<
(
(CudaSlice<i8>, CudaSlice<f32>),
(CudaSlice<i8>, CudaSlice<f32>),
),
Box<dyn std::error::Error>,
> {
let mut q0 = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut d0 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let mut q2 = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut d2 = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let f = self.func("add_rms_norm3_q8z_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, e2) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a)
.arg(b_in)
.arg(w0)
.arg(w1)
.arg(w2)
.arg(res)
.arg(&mut q0)
.arg(&mut d0)
.arg(out1)
.arg(&mut q2)
.arg(&mut d2)
.arg(&nc)
.arg(&e2);
unsafe {
b.launch(cfg)?;
}
Ok(((q0, d0), (q2, d2)))
}
#[allow(clippy::too_many_arguments)]
pub fn add_rms_norm3(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
w0: &CudaSlice<f32>,
w1: &CudaSlice<f32>,
w2: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
d0: &mut CudaSlice<f32>,
d1: &mut CudaSlice<f32>,
d2: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("add_rms_norm3_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, e2) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a)
.arg(b_in)
.arg(w0)
.arg(w1)
.arg(w2)
.arg(res)
.arg(d0)
.arg(d1)
.arg(d2)
.arg(&nc)
.arg(&e2);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn add_scale(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
c: f32,
dst: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("add_scale_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a).arg(b_in).arg(&c).arg(dst).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn rms_norm(
&self,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let (nc, e) = (ncols as i32, eps);
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (px, _g0) = x.device_ptr(s);
let (pw, _g1) = w.device_ptr(s);
let (pd, _g2) = dst.device_ptr_mut(s);
let mut ps = [
&px as *const _ as *mut std::ffi::c_void,
&pw as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
&e as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"rms_norm_f32",
(nrows as u32, 1, 1),
(rms_block(), 1, 1),
&mut ps,
)?;
}
return Ok(());
}
let f = self.func("rms_norm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn rms_norm_decode(
&self,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rms_norm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (1024, 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(w).arg(dst).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn rms_norm_q8_1(
&self,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let nblk = ncols / 32;
let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
let (nc, e) = (ncols as i32, eps);
if Self::pdl_on() {
{
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (px, _g0) = x.device_ptr(s);
let (pw, _g1) = w.device_ptr(s);
let (pq, _g2) = q.device_ptr_mut(s);
let (pd, _g3) = d.device_ptr_mut(s);
let mut ps = [
&px as *const _ as *mut std::ffi::c_void,
&pw as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
&e as *const _ as *mut _,
];
unsafe {
self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1), &mut ps)?;
}
}
return Ok((q, d));
}
let f = self.func("rms_norm_q8_1");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (1024, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(w).arg(&mut q).arg(&mut d).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok((q, d))
}
pub fn rms_norm_q8_1_into(
&self,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
q: &mut CudaSlice<i8>,
d: &mut CudaSlice<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
let nblk = ncols / 32;
debug_assert!(q.len() >= nrows * ncols && d.len() >= nrows * nblk);
let (nc, e) = (ncols as i32, eps);
if Self::pdl_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (px, _g0) = x.device_ptr(s);
let (pw, _g1) = w.device_ptr(s);
let (pq, _g2) = q.device_ptr_mut(s);
let (pd, _g3) = d.device_ptr_mut(s);
let mut ps = [
&px as *const _ as *mut std::ffi::c_void,
&pw as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
&e as *const _ as *mut _,
];
unsafe {
self.launch_pdl("rms_norm_q8_1", (nrows as u32, 1, 1), (1024, 1, 1), &mut ps)?;
}
return Ok(());
}
let f = self.func("rms_norm_q8_1");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (1024, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(w).arg(&mut *q).arg(&mut *d).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn quantize_q8_1_into(
&self,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
q: &mut CudaSlice<i8>,
d: &mut CudaSlice<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
let nblk = in_f / 32;
debug_assert!(q.len() >= m * in_f && d.len() >= m * nblk);
let cfg = LaunchConfig::for_num_elems((m * in_f) as u32);
let (inf, mi) = (in_f as i32, m as i32);
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (px, _g0) = x.device_ptr(s);
let (pq, _g1) = q.device_ptr_mut(s);
let (pd, _g2) = d.device_ptr_mut(s);
let mut ps = [
&px as *const _ as *mut std::ffi::c_void,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&inf as *const _ as *mut _,
&mi as *const _ as *mut _,
];
unsafe {
self.launch_pdl("quantize_q8_1", cfg.grid_dim, cfg.block_dim, &mut ps)?;
}
return Ok(());
}
let f = self.func("quantize_q8_1");
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(&mut *q).arg(&mut *d).arg(&inf).arg(&mi);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn add_rms_norm_q8_1(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let nblk = ncols / 32;
let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut d = self.alloc_uninit::<f32>(nrows * nblk)?;
let f = self.func("add_rms_norm_q8_1");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (1024, 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_bld = self.gpu.stream();
let mut bld = __s_bld.launch_builder(&f);
bld.arg(a)
.arg(b_in)
.arg(w)
.arg(res)
.arg(&mut q)
.arg(&mut d)
.arg(&nc)
.arg(&e);
unsafe {
bld.launch(cfg)?;
}
Ok((q, d))
}
pub fn add_rms_norm(
&self,
a: &CudaSlice<f32>,
b: &CudaSlice<f32>,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let (nc, e) = (ncols as i32, eps);
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pa, _g0) = a.device_ptr(s);
let (pb, _g1) = b.device_ptr(s);
let (pw, _g2) = w.device_ptr(s);
let (pr, _g3) = res.device_ptr_mut(s);
let (pd, _g4) = dst.device_ptr_mut(s);
let mut ps = [
&pa as *const _ as *mut std::ffi::c_void,
&pb as *const _ as *mut _,
&pw as *const _ as *mut _,
&pr as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
&e as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"add_rms_norm_f32",
(nrows as u32, 1, 1),
(rms_block(), 1, 1),
&mut ps,
)?;
}
return Ok(());
}
let f = self.func("add_rms_norm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f);
b2.arg(a)
.arg(b)
.arg(w)
.arg(&mut *res)
.arg(&mut *dst)
.arg(&nc)
.arg(&e);
unsafe {
b2.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn rms_pre_add_rms_norm(
&self,
a: &CudaSlice<f32>,
wa: &CudaSlice<f32>,
b: &CudaSlice<f32>,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rms_pre_add_rms_norm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f);
b2.arg(a)
.arg(wa)
.arg(b)
.arg(w)
.arg(&mut *res)
.arg(&mut *dst)
.arg(&nc)
.arg(&e);
unsafe {
b2.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn rms_pre_add_rms_norm_q8z(
&self,
a: &CudaSlice<f32>,
wa: &CudaSlice<f32>,
b: &CudaSlice<f32>,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
debug_assert!(ncols % 128 == 0);
let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let (nc, e) = (ncols as i32, eps);
if Self::pdl_on() {
{
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pa, _g0) = a.device_ptr(s);
let (pwa, _g1) = wa.device_ptr(s);
let (pb, _g2) = b.device_ptr(s);
let (pw, _g3) = w.device_ptr(s);
let (pr, _g4) = res.device_ptr_mut(s);
let (pdst, _g5) = dst.device_ptr_mut(s);
let (pq, _g6) = out_q.device_ptr_mut(s);
let (pd, _g7) = out_d.device_ptr_mut(s);
let mut ps = [
&pa as *const _ as *mut std::ffi::c_void,
&pwa as *const _ as *mut _,
&pb as *const _ as *mut _,
&pw as *const _ as *mut _,
&pr as *const _ as *mut _,
&pdst as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&nc as *const _ as *mut _,
&e as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"rms_pre_add_rms_norm_q8z_f32",
(nrows as u32, 1, 1),
(rms_block(), 1, 1),
&mut ps,
)?;
}
}
return Ok((out_q, out_d));
}
let f = self.func("rms_pre_add_rms_norm_q8z_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f);
b2.arg(a)
.arg(wa)
.arg(b)
.arg(w)
.arg(&mut *res)
.arg(&mut *dst)
.arg(&mut out_q)
.arg(&mut out_d)
.arg(&nc)
.arg(&e);
unsafe {
b2.launch(cfg)?;
}
Ok((out_q, out_d))
}
pub fn build_q4_out_concat3(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
w2: &crate::model::GpuTensor,
) -> Result<Option<crate::model::GpuTensor>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let part = |w: &GpuTensor| -> Option<(usize, usize)> {
match w {
GpuTensor::Quant {
qtype,
row_bytes,
rp,
..
} if *qtype == QT_Q4_0 && !*rp => Some((*row_bytes, w.out_features())),
_ => None,
}
};
let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (part(w0), part(w1), part(w2))
else {
return Ok(None);
};
if rb0 != rb1
|| rb0 != rb2
|| w0.in_features() != w1.in_features()
|| w0.in_features() != w2.in_features()
{
return Ok(None);
}
fn bytes_of(w: &crate::model::GpuTensor) -> &CudaSlice<u8> {
match w {
crate::model::GpuTensor::Quant { bytes, .. } => bytes,
_ => unreachable!(),
}
}
let (b0, b1, b2) = (bytes_of(w0), bytes_of(w1), bytes_of(w2));
let total = rb0 * (o0 + o1 + o2);
let mut cat = self.alloc_u8(total)?;
self.copy_u8_into(&mut cat, 0, b0, rb0 * o0)?;
self.copy_u8_into(&mut cat, rb0 * o0, b1, rb1 * o1)?;
self.copy_u8_into(&mut cat, rb0 * (o0 + o1), b2, rb2 * o2)?;
Ok(Some(GpuTensor::Quant {
bytes: cat,
qtype: QT_Q4_0,
row_bytes: rb0,
ne: vec![w0.in_features() as u64, (o0 + o1 + o2) as u64],
scale: 1.0,
rp: false,
#[cfg(memra_cutlass)]
cutlass: None,
fp8: None,
blk: None,
rp4: None,
f16: None,
}))
}
#[allow(clippy::too_many_arguments)]
pub fn rms_norm_qkv_rope_cat(
&self,
qkv: &CudaSlice<f32>,
wq: &CudaSlice<f32>,
wk: &CudaSlice<f32>,
wv: &CudaSlice<f32>,
q: &mut CudaSlice<f32>,
k: &mut CudaSlice<f32>,
v: &mut CudaSlice<f32>,
head_dim: usize,
rq: usize,
rk: usize,
pos: &CudaSlice<i32>,
nh_q: usize,
nh_k: usize,
base: f32,
freq_scale: f32,
ff: Option<&CudaSlice<f32>>,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let rows = rq + rk + rk;
let theta_scale = base.powf(-2.0 / head_dim as f32);
let (nc, rqi, rki, nhq, nhk) = (
head_dim as i32,
rq as i32,
rk as i32,
nh_q as i32,
nh_k as i32,
);
if Self::pdl_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pqkv, _g0) = qkv.device_ptr(s);
let (pwq, _g1) = wq.device_ptr(s);
let (pwk, _g2) = wk.device_ptr(s);
let (pwv, _g3) = wv.device_ptr(s);
let (pq, _g4) = q.device_ptr_mut(s);
let (pk, _g5) = k.device_ptr_mut(s);
let (pv, _g6) = v.device_ptr_mut(s);
let (ppos, _g7) = pos.device_ptr(s);
let (pff, _g8) = match ff {
Some(t) => {
let (p, g) = t.device_ptr(s);
(p, Some(g))
}
None => (0, None),
};
let mut ps = [
&pqkv as *const _ as *mut std::ffi::c_void,
&pwq as *const _ as *mut _,
&pwk as *const _ as *mut _,
&pwv as *const _ as *mut _,
&pq as *const _ as *mut _,
&pk as *const _ as *mut _,
&pv as *const _ as *mut _,
&nc as *const _ as *mut _,
&rqi as *const _ as *mut _,
&rki as *const _ as *mut _,
&ppos as *const _ as *mut _,
&nhq as *const _ as *mut _,
&nhk as *const _ as *mut _,
&theta_scale as *const _ as *mut _,
&freq_scale as *const _ as *mut _,
&pff as *const _ as *mut _,
&eps as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"rms_norm_qkv_rope_cat_f32",
(rows as u32, 1, 1),
(rms_block(), 1, 1),
&mut ps,
)?;
}
return Ok(());
}
let f = self.func("rms_norm_qkv_rope_cat_f32");
let cfg = LaunchConfig {
grid_dim: (rows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
match ff {
Some(t) => {
b.arg(qkv)
.arg(wq)
.arg(wk)
.arg(wv)
.arg(&mut *q)
.arg(&mut *k)
.arg(&mut *v)
.arg(&nc)
.arg(&rqi)
.arg(&rki)
.arg(pos)
.arg(&nhq)
.arg(&nhk)
.arg(&theta_scale)
.arg(&freq_scale)
.arg(t)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
}
None => {
let null: u64 = 0;
b.arg(qkv)
.arg(wq)
.arg(wk)
.arg(wv)
.arg(&mut *q)
.arg(&mut *k)
.arg(&mut *v)
.arg(&nc)
.arg(&rqi)
.arg(&rki)
.arg(pos)
.arg(&nhq)
.arg(&nhk)
.arg(&theta_scale)
.arg(&freq_scale)
.arg(&null)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn rms_norm_qkv_rope(
&self,
q0: &CudaSlice<f32>,
k0: &CudaSlice<f32>,
v0: &CudaSlice<f32>,
wq: &CudaSlice<f32>,
wk: &CudaSlice<f32>,
wv: &CudaSlice<f32>,
q: &mut CudaSlice<f32>,
k: &mut CudaSlice<f32>,
v: &mut CudaSlice<f32>,
head_dim: usize,
rq: usize,
rk: usize,
pos: &CudaSlice<i32>,
nh_q: usize,
nh_k: usize,
base: f32,
freq_scale: f32,
ff: Option<&CudaSlice<f32>>,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rms_norm_qkv_rope_f32");
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,
};
let theta_scale = base.powf(-2.0 / head_dim as f32);
let (nc, rqi, rki, nhq, nhk) = (
head_dim as i32,
rq as i32,
rk as i32,
nh_q as i32,
nh_k as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
match ff {
Some(t) => {
b.arg(q0)
.arg(k0)
.arg(v0)
.arg(wq)
.arg(wk)
.arg(wv)
.arg(&mut *q)
.arg(&mut *k)
.arg(&mut *v)
.arg(&nc)
.arg(&rqi)
.arg(&rki)
.arg(pos)
.arg(&nhq)
.arg(&nhk)
.arg(&theta_scale)
.arg(&freq_scale)
.arg(t)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
}
None => {
let null: u64 = 0;
b.arg(q0)
.arg(k0)
.arg(v0)
.arg(wq)
.arg(wk)
.arg(wv)
.arg(&mut *q)
.arg(&mut *k)
.arg(&mut *v)
.arg(&nc)
.arg(&rqi)
.arg(&rki)
.arg(pos)
.arg(&nhq)
.arg(&nhk)
.arg(&theta_scale)
.arg(&freq_scale)
.arg(&null)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn rms_norm_qkv_rope_append_dc(
&self,
q0: &CudaSlice<f32>,
k0: &CudaSlice<f32>,
v0: &CudaSlice<f32>,
wq: &CudaSlice<f32>,
wk: &CudaSlice<f32>,
wv: &CudaSlice<f32>,
q: &mut CudaSlice<f32>,
k: &mut CudaSlice<f32>,
v: &mut CudaSlice<f32>,
head_dim: usize,
rq: usize,
rk: usize,
pos: &CudaSlice<i32>,
nh_q: usize,
nh_k: usize,
base: f32,
freq_scale: f32,
ff: Option<&CudaSlice<f32>>,
eps: f32,
kc: &mut CudaSlice<u8>,
vc: &mut CudaSlice<u8>,
t_dev: &CudaSlice<i32>,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let rows = rq + rk + rk;
let theta_scale = base.powf(-2.0 / head_dim as f32);
let (nc, rqi, rki, nhq, nhk) = (
head_dim as i32,
rq as i32,
rk as i32,
nh_q as i32,
nh_k as i32,
);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (p0, _a0) = q0.device_ptr(s);
let (p1, _a1) = k0.device_ptr(s);
let (p2, _a2) = v0.device_ptr(s);
let (pwq, _a3) = wq.device_ptr(s);
let (pwk, _a4) = wk.device_ptr(s);
let (pwv, _a5) = wv.device_ptr(s);
let (pq, _a6) = q.device_ptr_mut(s);
let (pk, _a7) = k.device_ptr_mut(s);
let (pv, _a8) = v.device_ptr_mut(s);
let (pp, _a9) = pos.device_ptr(s);
let pff: u64 = match ff {
Some(t) => {
let (p, _gg) = t.device_ptr(s);
p as u64
}
None => 0,
};
let (pkc, _a10) = kc.device_ptr_mut(s);
let (pvc, _a11) = vc.device_ptr_mut(s);
let (pt, _a12) = t_dev.device_ptr(s);
let mut ps = [
&p0 as *const _ as *mut std::ffi::c_void,
&p1 as *const _ as *mut _,
&p2 as *const _ as *mut _,
&pwq as *const _ as *mut _,
&pwk as *const _ as *mut _,
&pwv as *const _ as *mut _,
&pq as *const _ as *mut _,
&pk as *const _ as *mut _,
&pv as *const _ as *mut _,
&nc as *const _ as *mut _,
&rqi as *const _ as *mut _,
&rki as *const _ as *mut _,
&pp as *const _ as *mut _,
&nhq as *const _ as *mut _,
&nhk as *const _ as *mut _,
&theta_scale as *const _ as *mut _,
&freq_scale as *const _ as *mut _,
&pff as *const _ as *mut _,
&eps as *const _ as *mut _,
&pkc as *const _ as *mut _,
&pvc as *const _ as *mut _,
&pt as *const _ as *mut _,
&ktb as *const _ as *mut _,
&vtb as *const _ as *mut _,
];
unsafe {
self.launch_pdl_flash(
g,
"rms_norm_qkv_rope_append_dc_f32",
(rows as u32, 1, 1),
(rms_block(), 1, 1),
0,
&mut ps,
)?;
}
return Ok(());
}
let f = if g {
self.func_g("rms_norm_qkv_rope_append_dc_f32")
} else {
self.func("rms_norm_qkv_rope_append_dc_f32")
};
let cfg = LaunchConfig {
grid_dim: (rows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
match ff {
Some(t) => {
b.arg(q0)
.arg(k0)
.arg(v0)
.arg(wq)
.arg(wk)
.arg(wv)
.arg(&mut *q)
.arg(&mut *k)
.arg(&mut *v)
.arg(&nc)
.arg(&rqi)
.arg(&rki)
.arg(pos)
.arg(&nhq)
.arg(&nhk)
.arg(&theta_scale)
.arg(&freq_scale)
.arg(t)
.arg(&eps)
.arg(&mut *kc)
.arg(&mut *vc)
.arg(t_dev)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
}
None => {
let null: u64 = 0;
b.arg(q0)
.arg(k0)
.arg(v0)
.arg(wq)
.arg(wk)
.arg(wv)
.arg(&mut *q)
.arg(&mut *k)
.arg(&mut *v)
.arg(&nc)
.arg(&rqi)
.arg(&rki)
.arg(pos)
.arg(&nhq)
.arg(&nhk)
.arg(&theta_scale)
.arg(&freq_scale)
.arg(&null)
.arg(&eps)
.arg(&mut *kc)
.arg(&mut *vc)
.arg(t_dev)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
}
}
Ok(())
}
pub fn add_q8_1(
&self,
a: &CudaSlice<f32>,
b: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
debug_assert!(ncols % 128 == 0);
let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let f = self.func("add_q8_1_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let nc = ncols as i32;
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f);
b2.arg(a)
.arg(b)
.arg(&mut *res)
.arg(&mut out_q)
.arg(&mut out_d)
.arg(&nc);
unsafe {
b2.launch(cfg)?;
}
Ok((out_q, out_d))
}
pub fn rms_pre_add_q8_1(
&self,
a: &CudaSlice<f32>,
wa: &CudaSlice<f32>,
b: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
debug_assert!(ncols % 128 == 0);
let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let f = self.func("rms_pre_add_q8_1_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, ep) = (ncols as i32, eps);
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&f);
b2.arg(a)
.arg(wa)
.arg(b)
.arg(&mut *res)
.arg(&mut out_q)
.arg(&mut out_d)
.arg(&nc)
.arg(&ep);
unsafe {
b2.launch(cfg)?;
}
Ok((out_q, out_d))
}
pub fn l2_v2_on(ncols: usize) -> bool {
ncols == 128 && std::env::var("MEMRA_L2_V2").as_deref() != Ok("0")
}
pub fn l2_norm_pp(
&self,
x: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
dst16: Option<&mut CudaSlice<u8>>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
if Self::l2_v2_on(ncols) {
let f = self.func("l2_norm_pp_v2_f32");
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,
};
let (nc, nr, e) = (ncols as i32, nrows as i32, eps);
let d16: u64 = match dst16 {
Some(d) => self.addr_u8(d),
None => 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(dst).arg(&d16).arg(&nc).arg(&nr).arg(&e);
unsafe {
b.launch(cfg)?;
}
return Ok(());
}
self.l2_norm(x, dst, ncols, nrows, eps)
}
pub fn l2_norm(
&self,
x: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("l2_norm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(dst).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn l2_norm_decode(
&self,
x: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("l2_norm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(dst).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn rope_neox(
&self,
x: &mut CudaSlice<f32>,
pos: &CudaSlice<i32>,
head_dim: usize,
n_dims: usize,
n_heads: usize,
n_tokens: usize,
freq_base: f32,
freq_scale: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rope_neox_f32");
let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
let grid = (n_heads * n_tokens) as u32;
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: ((head_dim / 2) as u32, 1, 1),
shared_mem_bytes: 0,
};
let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(pos)
.arg(&hd)
.arg(&nd)
.arg(&nh)
.arg(&theta_scale)
.arg(&freq_scale);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn rope_neox_ff(
&self,
x: &mut CudaSlice<f32>,
pos: &CudaSlice<i32>,
head_dim: usize,
n_dims: usize,
n_heads: usize,
n_tokens: usize,
freq_base: f32,
freq_scale: f32,
ff: &CudaSlice<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rope_neox_ff_f32");
let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
let grid = (n_heads * n_tokens) as u32;
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: ((head_dim / 2) as u32, 1, 1),
shared_mem_bytes: 0,
};
let (hd, nd, nh) = (head_dim as i32, n_dims as i32, n_heads as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x)
.arg(pos)
.arg(&hd)
.arg(&nd)
.arg(&nh)
.arg(&theta_scale)
.arg(&freq_scale)
.arg(ff);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn rope_neox2(
&self,
q: &mut CudaSlice<f32>,
k: &mut CudaSlice<f32>,
pos: &CudaSlice<i32>,
head_dim: usize,
n_dims: usize,
nh_q: usize,
nh_k: usize,
n_tokens: usize,
freq_base: f32,
freq_scale: f32,
ff: Option<&CudaSlice<f32>>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rope_neox2_f32");
let theta_scale = (freq_base).powf(-2.0 / n_dims as f32);
let grid = ((nh_q + nh_k) * n_tokens) as u32;
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: ((head_dim / 2) as u32, 1, 1),
shared_mem_bytes: 0,
};
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,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(pos)
.arg(&hd)
.arg(&nd)
.arg(&nq)
.arg(&nk)
.arg(&nt)
.arg(&theta_scale)
.arg(&freq_scale);
match ff {
Some(ffv) => {
b.arg(ffv);
unsafe {
b.launch(cfg)?;
}
}
None => {
let null: u64 = 0;
b.arg(&null);
unsafe {
b.launch(cfg)?;
}
}
}
Ok(())
}
pub fn gelu_tanh_mul(
&self,
gate: &CudaSlice<f32>,
up: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gelu_tanh_mul_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate).arg(up).arg(dst).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn silu_mul(
&self,
gate: &CudaSlice<f32>,
up: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("silu_mul_f32");
let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate).arg(up).arg(dst).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn silu_mul_f16out(
&self,
gate: &CudaSlice<f32>,
up: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
dst16: &mut CudaSlice<u8>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("silu_mul_f16out_f32");
let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate).arg(up).arg(dst).arg(dst16).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn silu_mul_scaled(
&self,
gate: &CudaSlice<f32>,
up: &CudaSlice<f32>,
gs: f32,
us: f32,
dst: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("silu_mul_scaled_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let (gsf, usf) = (gs, us);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate).arg(up).arg(&gsf).arg(&usf).arg(dst).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn swigluoai_mul_scaled(
&self,
gate: &CudaSlice<f32>,
up: &CudaSlice<f32>,
gs: f32,
us: f32,
alpha: f32,
limit: f32,
dst: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("swigluoai_mul_scaled_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate)
.arg(up)
.arg(&gs)
.arg(&us)
.arg(&alpha)
.arg(&limit)
.arg(dst)
.arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn silu_mul_scaled_q8_1(
&self,
gate: &CudaSlice<f32>,
up: &CudaSlice<f32>,
gs: f32,
us: f32,
n: usize,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let f = self.func("silu_mul_scaled_q8_1");
let nblk = n / 32;
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);
let (gsf, usf, ni) = (gs, us, n as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate)
.arg(up)
.arg(&gsf)
.arg(&usf)
.arg(&mut aq)
.arg(&mut ad)
.arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok((aq, ad))
}
pub fn add(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("add_f32");
let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
let ni = n as i32;
let __s_bld = self.gpu.stream();
let mut bld = __s_bld.launch_builder(&f);
bld.arg(a).arg(b_in).arg(dst).arg(&ni);
unsafe {
bld.launch(cfg)?;
}
Ok(())
}
pub fn mul(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("mul_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_bld = self.gpu.stream();
let mut bld = __s_bld.launch_builder(&f);
bld.arg(a).arg(b_in).arg(dst).arg(&ni);
unsafe {
bld.launch(cfg)?;
}
Ok(())
}
pub fn matmul(
&self,
w: &crate::model::GpuTensor,
x: &CudaSlice<f32>,
m: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let in_f = w.in_features();
let out_f = w.out_features();
#[allow(non_snake_case)]
let GEMM_M_THRESHOLD = if self.verify_exact_on() {
usize::MAX
} else {
16usize
};
const GEMM_MIN_OUT_F: usize = 128; if m >= GEMM_M_THRESHOLD {
if let Some(y) = self.try_fp8_gemm(w, x, m)? {
return Ok(y);
}
if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? {
return Ok(y);
}
if let Some(y) = self.try_f16_gemm(w, x, m)? {
return Ok(y);
}
}
if let GpuTensor::Quant { qtype, .. } = w {
if *qtype == QT_F8_E4M3_BLK {
if m >= GEMM_M_THRESHOLD {
if let Some(y) = self.try_e4m3_blk_prefill(w, x, m)? {
return Ok(y);
}
}
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? {
return Ok(y);
}
}
}
if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.mmq_supports(w) {
return self.qmatvec_mmq(w, x, m);
}
if m >= GEMM_M_THRESHOLD && out_f >= GEMM_MIN_OUT_F && self.gemm_supports(w) {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
return self.qmatvec_gemm(w, &aq, &ad, m);
}
if m >= GEMM_M_THRESHOLD {
if let Some(y) = self.try_fp4_gemm(w, x, m, in_f, out_f)? {
return Ok(y);
}
}
let fast = std::env::var("MEMRA_FAST").as_deref() != Ok("0");
if m == 1 && fast {
if let GpuTensor::Quant {
bytes,
qtype,
row_bytes,
rp,
rp4,
scale,
..
} = w
{
if self.mmvq_supports(*qtype) {
let (bytes, rp) = match rp4 {
Some(m4) => (m4, true),
None => (bytes, *rp),
};
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
return self.qmatvec_mmvq(
bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, rp,
);
}
}
}
if (2..=16).contains(&m)
&& fast
&& std::env::var("MEMRA_NO_BATCHED").is_err()
&& (m <= 4 || Self::b8_enabled())
{
let m_ok = m <= 8
|| matches!(w, GpuTensor::Quant { qtype, .. }
if *qtype == QT_Q4_0 || *qtype == QT_Q6_K || *qtype == QT_F8_E4M3
|| *qtype == QT_NVFP4 || *qtype == QT_Q4_K || *qtype == QT_Q5_K || *qtype == QT_Q8_0);
if m_ok {
if let GpuTensor::Quant {
bytes,
qtype,
row_bytes,
rp,
rp4,
..
} = w
{
if self.batched_supports(*qtype) && self.mmvq_supports(*qtype) {
let (bytes, rp) = match rp4 {
Some(m4) => (m4, true),
None => (bytes, *rp),
};
let mcols = Self::batched_mcols(m);
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
let mut y = self.qmatvec_mmvq_batched(
bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, mcols, 1.0, rp,
)?;
if let GpuTensor::Quant { scale, .. } = w {
if *scale != 1.0 {
self.scale_inplace(&mut y, *scale, m * out_f)?;
}
}
return Ok(y);
}
}
}
}
if fast {
if let GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
..
} = w
{
if *qtype == QT_F8_E4M3 {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
return self.qmatvec_mmvq(
bytes, &aq, &ad, m, in_f, out_f, *qtype, *row_bytes, *scale, false,
);
}
}
}
let mut y = match w {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
..
} if fast && *qtype == QT_Q8_0 => {
self.qmatvec_q8_0_fast(bytes, x, m, in_f, out_f, *row_bytes)?
}
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
..
} if fast && *qtype == QT_Q4_K => {
self.qmatvec_q4_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
}
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
..
} if fast && *qtype == QT_Q6_K => {
self.qmatvec_q6_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
}
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
..
} if fast && *qtype == QT_Q5_K => {
self.qmatvec_q5_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
}
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
..
} if fast && *qtype == QT_Q3_K => {
self.qmatvec_q3_K_fast(bytes, x, m, in_f, out_f, *row_bytes)?
}
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
rp,
..
} if fast && *qtype == QT_NVFP4 => self.qmatvec_dp4a_named(
if *rp {
"qmatvec_nvfp4_dp4a_rp"
} else {
"qmatvec_nvfp4_dp4a"
},
bytes,
x,
m,
in_f,
out_f,
*row_bytes,
)?,
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
..
} if fast && *qtype == QT_IQ4_XS && Self::iq_fast_enabled() => {
self.qmatvec_iq4_XS_fast(bytes, x, m, in_f, out_f, *row_bytes)?
}
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
rp,
..
} =>
{
self.qmatvec(
bytes,
x,
m,
in_f,
out_f,
if *rp && *qtype == QT_NVFP4 {
QT_NVFP4_RP
} else {
*qtype
},
*row_bytes,
)?
}
GpuTensor::Float { data, .. } => self.linear(x, data, m, in_f, out_f)?,
GpuTensor::FloatBf16 { data, .. } => {
self.linear_bf16_chunked(x, data, m, in_f, out_f, false)?
}
};
if let GpuTensor::Quant { scale, .. } = w {
if *scale != 1.0 {
self.scale_inplace(&mut y, *scale, m * out_f)?;
}
}
Ok(y)
}
pub fn uses_q8_1_fast(&self, w: &crate::model::GpuTensor) -> bool {
use crate::model::GpuTensor;
if std::env::var("MEMRA_FAST").as_deref() == Ok("0") {
return false;
}
match w {
GpuTensor::Quant { qtype, .. } => {
matches!(
*qtype,
QT_Q8_0
| QT_Q4_K
| QT_Q6_K
| QT_Q5_K
| QT_Q3_K
| QT_NVFP4
| QT_F8_E4M3
| QT_F8_E4M3_BLK
| QT_Q4_0
) || (*qtype == QT_IQ4_XS && Self::iq_fast_enabled())
}
GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
}
}
pub fn matmul_pre(
&self,
w: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
x_fallback: &CudaSlice<f32>,
m: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let x_raw_ok = x_fallback.len() >= m * w.in_features();
if m >= 16 && x_raw_ok && !self.verify_exact_on() {
if let Some(y) = self.try_fp8_gemm(w, x_fallback, m)? {
return Ok(y);
}
if let Some(y) = self.try_fp8_blk_mmq(w, x_fallback, m)? {
return Ok(y);
}
if let Some(y) = self.try_f16_gemm(w, x_fallback, m)? {
return Ok(y);
}
}
if m >= 16 && x_raw_ok && !self.verify_exact_on() {
if let Some(y) = self.try_e4m3_blk_prefill(w, x_fallback, m)? {
return Ok(y);
}
}
if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
return Ok(y);
}
if m >= 16
&& w.out_features() >= 128
&& self.mmq_supports(w)
&& !self.verify_exact_on()
&& x_raw_ok
{
return self.qmatvec_mmq(w, x_fallback, m);
}
if m >= 16 && x_raw_ok && !self.verify_exact_on() {
if let Some(y) =
self.try_fp4_gemm(w, x_fallback, m, w.in_features(), w.out_features())?
{
return Ok(y);
}
}
if m >= 16 && self.gemm_supports(w) && !self.verify_exact_on() {
return self.qmatvec_gemm(w, aq, ad, m);
}
if !self.uses_q8_1_fast(w) {
return self.matmul(w, x_fallback, m);
}
let in_f = w.in_features();
let out_f = w.out_features();
let (bytes, qtype, row_bytes, scale, rp) = match w {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
rp,
..
} => (bytes, *qtype, *row_bytes, *scale, *rp),
_ => unreachable!("uses_q8_1_fast guaranteed Quant"),
};
let (mbytes, mrp) = match w {
GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
_ => (bytes, rp),
};
if m == 1 && self.mmvq_supports(qtype) {
return self.qmatvec_mmvq(mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, mrp);
}
if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
&& std::env::var("MEMRA_NO_BATCHED").is_err()
&& (m <= 4 || Self::b8_enabled())
&& (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_NVFP4
|| qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_F8_E4M3 || qtype == QT_Q8_0)
{
let mcols = Self::batched_mcols(m);
return self.qmatvec_mmvq_batched(
mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, mrp,
);
}
if qtype == QT_F8_E4M3 || qtype == QT_Q4_0 {
let (b2, r2) = if qtype == QT_Q4_0 {
(mbytes, mrp)
} else {
(bytes, rp)
};
return self.qmatvec_mmvq(b2, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, r2);
}
let name = match qtype {
QT_Q8_0 => "qmatvec_q8_0_dp4a",
QT_Q4_K => "qmatvec_q4_K_dp4a",
QT_Q6_K => "qmatvec_q6_K_dp4a",
QT_Q5_K => "qmatvec_q5_K_dp4a",
QT_Q3_K => "qmatvec_q3_K_dp4a",
QT_NVFP4 => {
if rp {
"qmatvec_nvfp4_dp4a_rp"
} else {
"qmatvec_nvfp4_dp4a"
}
}
QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
_ => unreachable!(),
};
let f = self.func(name);
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,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
if scale != 1.0 {
self.scale_inplace(&mut y, scale, m * out_f)?;
}
Ok(y)
}
pub fn matmul_decode_exact(
&self,
w: &crate::model::GpuTensor,
x: &CudaSlice<f32>,
m: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
if let GpuTensor::Float { data, .. } = w {
return self.linear_decode_exact(x, data, m, w.in_features(), w.out_features());
}
if let GpuTensor::FloatBf16 { data, .. } = w {
let (in_f, out_f) = (w.in_features(), w.out_features());
return self.linear_bf16_chunked(x, data, m, in_f, out_f, true);
}
if !self.uses_q8_1_fast(w) {
return self.matmul(w, x, m);
}
let in_f = w.in_features();
let out_f = w.out_features();
let (bytes, qtype, row_bytes, scale, rp) = match w {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
rp,
..
} => (bytes, *qtype, *row_bytes, *scale, *rp),
_ => return self.matmul(w, x, m),
};
let (bytes, rp) = match w {
GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
_ => (bytes, rp),
};
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
if let Some(y) = self.try_e4m3_blk_pre(w, &aq, &ad, m)? {
return Ok(y);
}
if (2..=16).contains(&m) && self.batched_supports(qtype) && self.mmvq_supports(qtype)
&& std::env::var("MEMRA_NO_BATCHED").is_err()
&& (m <= 4 || Self::b8_enabled())
&& (m <= 8 || qtype == QT_Q4_0 || qtype == QT_Q6_K || qtype == QT_F8_E4M3
|| qtype == QT_NVFP4 || qtype == QT_Q4_K || qtype == QT_Q5_K || qtype == QT_Q8_0)
{
let mcols = Self::batched_mcols(m);
return self.qmatvec_mmvq_batched(
bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp,
);
}
if self.mmvq_supports(qtype) {
return self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
}
self.matmul_pre(w, &aq, &ad, x, m)
}
pub fn matmul_decode_exact_pre(
&self,
w: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
debug_assert!(
self.uses_q8_1_fast(w),
"matmul_decode_exact_pre: caller must guarantee q8_1-fast"
);
if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
return Ok(y);
}
let in_f = w.in_features();
let out_f = w.out_features();
let (bytes, qtype, row_bytes, scale, rp) = match w {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
rp,
..
} => (bytes, *qtype, *row_bytes, *scale, *rp),
_ => {
return Err(
"matmul_decode_exact_pre: Quant tensor required (q8_1-fast contract)".into(),
);
}
};
let (bytes, rp) = match w {
GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
_ => (bytes, rp),
};
if (2..=16).contains(&m)
&& self.batched_supports(qtype)
&& self.mmvq_supports(qtype)
&& std::env::var("MEMRA_NO_BATCHED").is_err()
&& (m <= 4 || Self::b8_enabled())
&& (m <= 8
|| qtype == QT_Q4_0
|| qtype == QT_Q6_K
|| qtype == QT_F8_E4M3
|| qtype == QT_NVFP4
|| qtype == QT_Q4_K
|| qtype == QT_Q5_K
|| qtype == QT_Q8_0)
{
let mcols = Self::batched_mcols(m);
return self.qmatvec_mmvq_batched(
bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, mcols, scale, rp,
);
}
if self.mmvq_supports(qtype) {
return self.qmatvec_mmvq(bytes, aq, ad, m, in_f, out_f, qtype, row_bytes, scale, rp);
}
let x0 = self.zeros(0)?;
self.matmul_pre(w, aq, ad, &x0, m)
}
pub fn matmul_decode_exact_dual_pre(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>>
{
use crate::model::GpuTensor;
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let on = *ON.get_or_init(|| {
std::env::var("MEMRA_SPEC_DUAL_T")
.map(|v| v != "0")
.unwrap_or(true)
});
if !on
|| !(2..=7).contains(&m)
|| std::env::var("MEMRA_NO_BATCHED").is_ok()
|| !self.uses_q8_1_fast(w0)
|| !self.uses_q8_1_fast(w1)
{
return Ok(None);
}
if !self.mmvq_supports(QT_NVFP4) {
return Ok(None);
}
let (in_f, out_f) = (w0.in_features(), w0.out_features());
if w1.in_features() != in_f || w1.out_features() != out_f {
return Ok(None);
}
let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
(
GpuTensor::Quant {
bytes: b0,
qtype: q0,
row_bytes: rb0,
scale: s0,
rp: rp0,
rp4: None,
..
},
GpuTensor::Quant {
bytes: b1,
qtype: q1,
row_bytes: rb1,
scale: s1,
rp: rp1,
rp4: None,
..
},
) if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 => {
(b0, b1, *rb0, *s0, *s1, *rp0)
}
_ => return Ok(None),
};
if m > 4 && !(rp && Self::b8_enabled() && std::env::var("MEMRA_B567").as_deref() != Ok("0"))
{
return Ok(None);
}
let (y0, y1) =
self.qmatvec_batched_dual_raw(b0, b1, aq, ad, m, in_f, out_f, row_bytes, rp)?;
Ok(Some(((y0, s0), (y1, s1))))
}
pub fn matmul_decode_exact_dual(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
x: &CudaSlice<f32>,
m: usize,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let on = *ON.get_or_init(|| {
std::env::var("MEMRA_SPEC_DUAL_T")
.map(|v| v != "0")
.unwrap_or(true)
});
if !on
|| !(2..=4).contains(&m)
|| std::env::var("MEMRA_NO_BATCHED").is_ok()
|| !self.uses_q8_1_fast(w0)
|| !self.uses_q8_1_fast(w1)
{
return Ok(None);
}
if !self.mmvq_supports(QT_NVFP4) {
return Ok(None);
}
let (in_f, out_f) = (w0.in_features(), w0.out_features());
if w1.in_features() != in_f || w1.out_features() != out_f {
return Ok(None);
}
let (b0, b1, row_bytes, s0, s1, rp) = match (w0, w1) {
(
GpuTensor::Quant {
bytes: b0,
qtype: q0,
row_bytes: rb0,
scale: s0,
rp: rp0,
rp4: None,
..
},
GpuTensor::Quant {
bytes: b1,
qtype: q1,
row_bytes: rb1,
scale: s1,
rp: rp1,
rp4: None,
..
},
) if *q0 == QT_NVFP4 && *q1 == QT_NVFP4 && rb0 == rb1 && rp0 == rp1 => {
(b0, b1, *rb0, *s0, *s1, *rp0)
}
_ => return Ok(None),
};
if std::env::var("MEMRA_DEBUG").is_ok() {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| eprintln!("[memra] dual gate+up batched ENGAGED (m={m} rp={rp})"));
}
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
let (y0, y1) =
self.qmatvec_batched_dual_raw(b0, b1, &aq, &ad, m, in_f, out_f, row_bytes, rp)?;
let mut y0 = y0;
let mut y1 = y1;
if s0 != 1.0 {
self.scale_inplace(&mut y0, s0, m * out_f)?;
}
if s1 != 1.0 {
self.scale_inplace(&mut y1, s1, m * out_f)?;
}
Ok(Some((y0, y1)))
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_batched_dual_raw(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
rp: bool,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4;
let mcols = Self::batched_mcols(m);
let tiny_rp1 = rp
&& mcols == 4
&& out_f <= 128
&& std::env::var("MEMRA_NVFP4_AUX_DUAL").as_deref() != Ok("0");
let (name, rows_per_block) = if tiny_rp1 {
("qmatvec_nvfp4_mmvq_dual_b4_rp", ROWS_PER_BLOCK)
} else {
match (mcols, rp, m) {
(2, false, _) => ("qmatvec_nvfp4_mmvq_dual_b2", ROWS_PER_BLOCK),
(4, false, _) => ("qmatvec_nvfp4_mmvq_dual_b4_r2", ROWS_PER_BLOCK * 2),
(2, true, _) => ("qmatvec_nvfp4_mmvq_dual_b2_rp", ROWS_PER_BLOCK),
(4, true, _) => ("qmatvec_nvfp4_mmvq_dual_b4_rpr2", ROWS_PER_BLOCK * 2),
(8, true, 5) => ("qmatvec_nvfp4_mmvq_dual_b5_rpr2", ROWS_PER_BLOCK * 2),
(8, true, 6) => ("qmatvec_nvfp4_mmvq_dual_b6_rpr2", ROWS_PER_BLOCK * 2),
(8, true, 7) => ("qmatvec_nvfp4_mmvq_dual_b7_rpr2", ROWS_PER_BLOCK * 2),
_ => {
return Err(
format!("qmatvec_batched_dual_raw: no dual kernel for m {m}").into(),
);
}
}
};
let f = self.func(name);
let mut y0 = self.alloc_uninit::<f32>(m * out_f)?;
let mut y1 = self.alloc_uninit::<f32>(m * out_f)?;
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok((y0, y1))
}
pub fn matmul_pre_dual_noscale(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<Option<((CudaSlice<f32>, f32), (CudaSlice<f32>, f32))>, Box<dyn std::error::Error>>
{
use crate::model::GpuTensor;
if m != 1 || !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
return Ok(None);
}
if !self.mmvq_supports(QT_NVFP4) {
return Ok(None);
}
let (in_f, out_f) = (w0.in_features(), w0.out_features());
if w1.in_features() != in_f || w1.out_features() != out_f {
return Ok(None);
}
let no_mirror =
|w: &crate::model::GpuTensor| !matches!(w, GpuTensor::Quant { rp4: Some(_), .. });
if self.q8_ffn_fuse2_on()
&& no_mirror(w0)
&& no_mirror(w1)
&& let Some([p0, p1]) = self.q8_fused_params(&[w0, w1])
{
let (y0, y1) = self.q8_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2)?;
return Ok(Some(((y0, 1.0), (y1, 1.0))));
}
if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
let (y0, y1) =
self.e4m3_fused2_core(p0.0, p1.0, aq, ad, in_f, p0.1, p1.1, p0.2, 1.0, 1.0)?;
return Ok(Some(((y0, p0.3), (y1, p1.3))));
}
let (b0, q0, rb0, s0, rp0) = match w0 {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
rp,
..
} => (bytes, *qtype, *row_bytes, *scale, *rp),
_ => return Ok(None),
};
let (b1, q1, rb1, s1, rp1) = match w1 {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
rp,
..
} => (bytes, *qtype, *row_bytes, *scale, *rp),
_ => return Ok(None),
};
if q0 != QT_NVFP4 || q1 != QT_NVFP4 || rb0 != rb1 || rp0 != rp1 {
return Ok(None);
}
const ROWS_PER_BLOCK: u32 = 4; const RPW: u32 = 2;
let rows_per_block = ROWS_PER_BLOCK * RPW;
let f = self.func(if rp0 {
"qmatvec_nvfp4_mmvq_dual_mr2_rp"
} else {
"qmatvec_nvfp4_mmvq_dual_mr2"
});
let mut y0 = self.alloc_uninit::<f32>(out_f)?;
let mut y1 = self.alloc_uninit::<f32>(out_f)?;
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 2, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, rb0 as i64);
let one = 1.0f32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb)
.arg(&one)
.arg(&one);
unsafe {
b.launch(cfg)?;
}
Ok(Some(((y0, s0), (y1, s1))))
}
pub fn matmul_q8_fused2(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
return Ok(Some(self.e4m3_fused2_core(
p0.0,
p1.0,
aq,
ad,
w0.in_features(),
p0.1,
p1.1,
p0.2,
p0.3,
p1.3,
)?));
}
let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
return Ok(None);
};
Ok(Some(self.q8_fused2_core(
p0.0,
p1.0,
aq,
ad,
w0.in_features(),
p0.1,
p1.1,
p0.2,
)?))
}
#[allow(clippy::too_many_arguments)]
fn q8_fused2_core(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
out0: usize,
out1: usize,
row_bytes: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
let f = self.func("qmatvec_q8_0_mmvq_fused2");
let mut y0 = self.alloc_uninit::<f32>(out0)?;
let mut y1 = self.alloc_uninit::<f32>(out1)?;
let cfg = LaunchConfig {
grid_dim: (nb0 + nb1, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&inf)
.arg(&o0)
.arg(&o1)
.arg(&rbl);
unsafe {
b.launch(cfg)?;
}
Ok((y0, y1))
}
pub fn matmul_q8_fused2_x(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
x: &CudaSlice<f32>,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
if !self.uses_q8_1_fast(w0) || !self.uses_q8_1_fast(w1) {
return Ok(None);
}
if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
return Ok(Some(self.e4m3_fused2_core(
p0.0,
p1.0,
&aq,
&ad,
w0.in_features(),
p0.1,
p1.1,
p0.2,
p0.3,
p1.3,
)?));
}
let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
return Ok(None);
};
let (aq, ad) = self.quantize_q8_1(x, 1, w0.in_features())?;
Ok(Some(self.q8_fused2_core(
p0.0,
p1.0,
&aq,
&ad,
w0.in_features(),
p0.1,
p1.1,
p0.2,
)?))
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_q8_fused2_raw(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
x: &CudaSlice<f32>,
in_f: usize,
out0: usize,
out1: usize,
row_bytes: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
self.q8_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes)
}
pub fn matmul_q4_fused3(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
w2: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
{
use crate::model::GpuTensor;
let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
match w {
GpuTensor::Quant {
qtype, row_bytes, ..
} if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
_ => None,
}
};
let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2)) else {
return Ok(None);
};
if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
return Ok(None);
}
fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
match w {
GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
Some(m) => (m, true),
None => (bytes, *rp),
},
_ => unreachable!(),
}
}
let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
if rp0 != rp1 || rp1 != rp2 {
return Ok(None);
}
let rp = rp0;
let rpb: u32 = 4;
let mr1 = rp && Self::q40_mr1_on();
let nb = |o: usize| {
if mr1 {
(o as u32).div_ceil(rpb)
} else {
(o as u32).div_ceil(2).div_ceil(rpb)
}
};
let grid = nb(o0) + nb(o1) + nb(o2);
let mut y0 = self.alloc_uninit::<f32>(o0)?;
let mut y1 = self.alloc_uninit::<f32>(o1)?;
let mut y2 = self.alloc_uninit::<f32>(o2)?;
let f = self.func(if mr1 {
"qmatvec_q4_0_mmvq_fused3_mr1_rp"
} else if rp {
"qmatvec_q4_0_mmvq_fused3_rp"
} else {
"qmatvec_q4_0_mmvq_fused3"
});
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (32, rpb, 1),
shared_mem_bytes: 0,
};
let inf = w0.in_features() as i32;
let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
{
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (p0, _g0) = b0.device_ptr(s);
let (p1, _g1) = b1.device_ptr(s);
let (p2, _g2) = b2.device_ptr(s);
let (paq, _g3) = aq.device_ptr(s);
let (pad, _g4) = ad.device_ptr(s);
let (py0, _g5) = y0.device_ptr_mut(s);
let (py1, _g6) = y1.device_ptr_mut(s);
let (py2, _g7) = y2.device_ptr_mut(s);
let mut ps = [
&p0 as *const _ as *mut std::ffi::c_void,
&p1 as *const _ as *mut _,
&p2 as *const _ as *mut _,
&paq as *const _ as *mut _,
&pad as *const _ as *mut _,
&py0 as *const _ as *mut _,
&py1 as *const _ as *mut _,
&py2 as *const _ as *mut _,
&inf as *const _ as *mut _,
&oo0 as *const _ as *mut _,
&oo1 as *const _ as *mut _,
&oo2 as *const _ as *mut _,
&r0 as *const _ as *mut _,
&r1 as *const _ as *mut _,
&r2 as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"qmatvec_q4_0_mmvq_fused3_mr1_rp",
(grid, 1, 1),
(32, rpb, 1),
&mut ps,
)?;
}
}
return Ok(Some((y0, y1, y2)));
}
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(b2)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&mut y2)
.arg(&inf)
.arg(&oo0)
.arg(&oo1)
.arg(&oo2)
.arg(&r0)
.arg(&r1)
.arg(&r2);
unsafe {
b.launch(cfg)?;
}
Ok(Some((y0, y1, y2)))
}
#[allow(clippy::too_many_arguments)]
pub fn matmul_q4_fused3_into(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
w2: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
y0: &mut CudaSlice<f32>,
y1: &mut CudaSlice<f32>,
y2: &mut CudaSlice<f32>,
) -> Result<bool, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
match w {
GpuTensor::Quant {
qtype, row_bytes, ..
} if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
_ => None,
}
};
let (Some((rb0, o0)), Some((rb1, o1)), Some((rb2, o2))) = (q4(w0), q4(w1), q4(w2)) else {
return Ok(false);
};
if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
return Ok(false);
}
fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
match w {
GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
Some(m) => (m, true),
None => (bytes, *rp),
},
_ => unreachable!(),
}
}
let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
if rp0 != rp1 || rp1 != rp2 {
return Ok(false);
}
let rp = rp0;
let rpb: u32 = 4;
let mr1 = rp && Self::q40_mr1_on();
let nb = |o: usize| {
if mr1 {
(o as u32).div_ceil(rpb)
} else {
(o as u32).div_ceil(2).div_ceil(rpb)
}
};
let grid = nb(o0) + nb(o1) + nb(o2);
debug_assert!(y0.len() >= o0 && y1.len() >= o1 && y2.len() >= o2);
let f = self.func(if mr1 {
"qmatvec_q4_0_mmvq_fused3_mr1_rp"
} else if rp {
"qmatvec_q4_0_mmvq_fused3_rp"
} else {
"qmatvec_q4_0_mmvq_fused3"
});
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (32, rpb, 1),
shared_mem_bytes: 0,
};
let inf = w0.in_features() as i32;
let (oo0, oo1, oo2) = (o0 as i32, o1 as i32, o2 as i32);
let (r0, r1, r2) = (rb0 as i64, rb1 as i64, rb2 as i64);
if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (p0, _g0) = b0.device_ptr(s);
let (p1, _g1) = b1.device_ptr(s);
let (p2, _g2) = b2.device_ptr(s);
let (paq, _g3) = aq.device_ptr(s);
let (pad, _g4) = ad.device_ptr(s);
let (py0, _g5) = y0.device_ptr_mut(s);
let (py1, _g6) = y1.device_ptr_mut(s);
let (py2, _g7) = y2.device_ptr_mut(s);
let mut ps = [
&p0 as *const _ as *mut std::ffi::c_void,
&p1 as *const _ as *mut _,
&p2 as *const _ as *mut _,
&paq as *const _ as *mut _,
&pad as *const _ as *mut _,
&py0 as *const _ as *mut _,
&py1 as *const _ as *mut _,
&py2 as *const _ as *mut _,
&inf as *const _ as *mut _,
&oo0 as *const _ as *mut _,
&oo1 as *const _ as *mut _,
&oo2 as *const _ as *mut _,
&r0 as *const _ as *mut _,
&r1 as *const _ as *mut _,
&r2 as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"qmatvec_q4_0_mmvq_fused3_mr1_rp",
(grid, 1, 1),
(32, rpb, 1),
&mut ps,
)?;
}
return Ok(true);
}
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(b2)
.arg(aq)
.arg(ad)
.arg(&mut *y0)
.arg(&mut *y1)
.arg(&mut *y2)
.arg(&inf)
.arg(&oo0)
.arg(&oo1)
.arg(&oo2)
.arg(&r0)
.arg(&r1)
.arg(&r2);
unsafe {
b.launch(cfg)?;
}
Ok(true)
}
pub fn matmul_q4_fused2(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
match w {
GpuTensor::Quant {
qtype, row_bytes, ..
} if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
_ => None,
}
};
let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else {
return Ok(None);
};
if w0.in_features() != w1.in_features() {
return Ok(None);
}
fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
match w {
GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
Some(m) => (m, true),
None => (bytes, *rp),
},
_ => unreachable!(),
}
}
let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
if rp0 != rp1 {
return Ok(None);
}
let rp = rp0;
let rpb: u32 = 4;
let mr1 = rp && Self::q40_mr1_on();
let nb = |o: usize| {
if mr1 {
(o as u32).div_ceil(rpb)
} else {
(o as u32).div_ceil(2).div_ceil(rpb)
}
};
let grid = nb(o0) + nb(o1);
let mut y0 = self.alloc_uninit::<f32>(o0)?;
let mut y1 = self.alloc_uninit::<f32>(o1)?;
let f = self.func(if mr1 {
"qmatvec_q4_0_mmvq_fused2_mr1_rp"
} else if rp {
"qmatvec_q4_0_mmvq_fused2_rp"
} else {
"qmatvec_q4_0_mmvq_fused2"
});
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (32, rpb, 1),
shared_mem_bytes: 0,
};
let inf = w0.in_features() as i32;
let (oo0, oo1) = (o0 as i32, o1 as i32);
let (r0, r1) = (rb0 as i64, rb1 as i64);
if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
{
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (p0, _g0) = b0.device_ptr(s);
let (p1, _g1) = b1.device_ptr(s);
let (paq, _g2) = aq.device_ptr(s);
let (pad, _g3) = ad.device_ptr(s);
let (py0, _g4) = y0.device_ptr_mut(s);
let (py1, _g5) = y1.device_ptr_mut(s);
let mut ps = [
&p0 as *const _ as *mut std::ffi::c_void,
&p1 as *const _ as *mut _,
&paq as *const _ as *mut _,
&pad as *const _ as *mut _,
&py0 as *const _ as *mut _,
&py1 as *const _ as *mut _,
&inf as *const _ as *mut _,
&oo0 as *const _ as *mut _,
&oo1 as *const _ as *mut _,
&r0 as *const _ as *mut _,
&r1 as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"qmatvec_q4_0_mmvq_fused2_mr1_rp",
(grid, 1, 1),
(32, rpb, 1),
&mut ps,
)?;
}
}
return Ok(Some((y0, y1)));
}
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&inf)
.arg(&oo0)
.arg(&oo1)
.arg(&r0)
.arg(&r1);
unsafe {
b.launch(cfg)?;
}
Ok(Some((y0, y1)))
}
pub fn matmul_q4_fused2_into(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
y0: &mut CudaSlice<f32>,
y1: &mut CudaSlice<f32>,
) -> Result<bool, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
match w {
GpuTensor::Quant {
qtype, row_bytes, ..
} if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
_ => None,
}
};
let (Some((rb0, o0)), Some((rb1, o1))) = (q4(w0), q4(w1)) else {
return Ok(false);
};
if w0.in_features() != w1.in_features() {
return Ok(false);
}
fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
match w {
GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
Some(m) => (m, true),
None => (bytes, *rp),
},
_ => unreachable!(),
}
}
let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
if rp0 != rp1 {
return Ok(false);
}
let rp = rp0;
let rpb: u32 = 4;
let mr1 = rp && Self::q40_mr1_on();
let nb = |o: usize| {
if mr1 {
(o as u32).div_ceil(rpb)
} else {
(o as u32).div_ceil(2).div_ceil(rpb)
}
};
let grid = nb(o0) + nb(o1);
debug_assert!(y0.len() >= o0 && y1.len() >= o1);
let f = self.func(if mr1 {
"qmatvec_q4_0_mmvq_fused2_mr1_rp"
} else if rp {
"qmatvec_q4_0_mmvq_fused2_rp"
} else {
"qmatvec_q4_0_mmvq_fused2"
});
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (32, rpb, 1),
shared_mem_bytes: 0,
};
let inf = w0.in_features() as i32;
let (oo0, oo1) = (o0 as i32, o1 as i32);
let (r0, r1) = (rb0 as i64, rb1 as i64);
if mr1 && Self::pdl_on() && Self::pdl_mmvq_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (p0, _g0) = b0.device_ptr(s);
let (p1, _g1) = b1.device_ptr(s);
let (paq, _g2) = aq.device_ptr(s);
let (pad, _g3) = ad.device_ptr(s);
let (py0, _g4) = y0.device_ptr_mut(s);
let (py1, _g5) = y1.device_ptr_mut(s);
let mut ps = [
&p0 as *const _ as *mut std::ffi::c_void,
&p1 as *const _ as *mut _,
&paq as *const _ as *mut _,
&pad as *const _ as *mut _,
&py0 as *const _ as *mut _,
&py1 as *const _ as *mut _,
&inf as *const _ as *mut _,
&oo0 as *const _ as *mut _,
&oo1 as *const _ as *mut _,
&r0 as *const _ as *mut _,
&r1 as *const _ as *mut _,
];
unsafe {
self.launch_pdl(
"qmatvec_q4_0_mmvq_fused2_mr1_rp",
(grid, 1, 1),
(32, rpb, 1),
&mut ps,
)?;
}
return Ok(true);
}
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(aq)
.arg(ad)
.arg(&mut *y0)
.arg(&mut *y1)
.arg(&inf)
.arg(&oo0)
.arg(&oo1)
.arg(&r0)
.arg(&r1);
unsafe {
b.launch(cfg)?;
}
Ok(true)
}
pub fn matmul_q4_fused2_batched(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
if m < 2 || m > 8 {
return Ok(None);
}
let q4 = |w: &GpuTensor| -> Option<(usize, usize)> {
match w {
GpuTensor::Quant {
qtype, row_bytes, ..
} if *qtype == QT_Q4_0 => Some((*row_bytes, w.out_features())),
_ => None,
}
};
let (Some((rb0, o0)), Some((_rb1, o1))) = (q4(w0), q4(w1)) else {
return Ok(None);
};
if w0.in_features() != w1.in_features() {
return Ok(None);
}
fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
match w {
GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
Some(mr) => (mr, true),
None => (bytes, *rp),
},
_ => unreachable!(),
}
}
let ((b0, rp0), (b1, rp1)) = (eff(w0), eff(w1));
if !rp0 || !rp1 {
return Ok(None);
}
let mcols = Self::batched_mcols(m);
let rpb: u32 = 4;
let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
let grid = nb(o0) + nb(o1);
let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
let f = self.func(match mcols {
2 => "qmatvec_q4_0_mmvq_b2_f2_rp",
4 => "qmatvec_q4_0_mmvq_b4_f2_rp",
_ => "qmatvec_q4_0_mmvq_b8_f2_rp",
});
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (32, rpb, 1),
shared_mem_bytes: 0,
};
let inf = w0.in_features() as i32;
let (oo0, oo1, mi) = (o0 as i32, o1 as i32, m as i32);
let rb = rb0 as i64;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&inf)
.arg(&oo0)
.arg(&oo1)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(Some((y0, y1)))
}
#[allow(clippy::too_many_arguments)]
pub fn matmul_q4_fused3_batched(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
w2: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
{
use crate::model::GpuTensor;
if m < 2 || m > 8 {
return Ok(None);
}
let q4 = |w: &GpuTensor| -> Option<usize> {
match w {
GpuTensor::Quant { qtype, .. } if *qtype == QT_Q4_0 => Some(w.out_features()),
_ => None,
}
};
let (Some(o0), Some(o1), Some(o2)) = (q4(w0), q4(w1), q4(w2)) else {
return Ok(None);
};
if w0.in_features() != w1.in_features() || w0.in_features() != w2.in_features() {
return Ok(None);
}
fn eff(w: &GpuTensor) -> (&CudaSlice<u8>, bool) {
match w {
GpuTensor::Quant { bytes, rp4, rp, .. } => match rp4 {
Some(mr) => (mr, true),
None => (bytes, *rp),
},
_ => unreachable!(),
}
}
let ((b0, rp0), (b1, rp1), (b2, rp2)) = (eff(w0), eff(w1), eff(w2));
if !rp0 || !rp1 || !rp2 {
return Ok(None);
}
let mcols = Self::batched_mcols(m);
let rpb: u32 = 4;
let nb = |o: usize| (o as u32).div_ceil(2 * rpb);
let grid = nb(o0) + nb(o1) + nb(o2);
let mut y0 = self.alloc_uninit::<f32>(m * o0)?;
let mut y1 = self.alloc_uninit::<f32>(m * o1)?;
let mut y2 = self.alloc_uninit::<f32>(m * o2)?;
let f = self.func(match mcols {
2 => "qmatvec_q4_0_mmvq_b2_f3_rp",
4 => "qmatvec_q4_0_mmvq_b4_f3_rp",
_ => "qmatvec_q4_0_mmvq_b8_f3_rp",
});
let cfg = LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (32, rpb, 1),
shared_mem_bytes: 0,
};
let inf = w0.in_features() as i32;
let (oo0, oo1, oo2, mi) = (o0 as i32, o1 as i32, o2 as i32, m as i32);
let rb = 0i64;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(b2)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&mut y2)
.arg(&inf)
.arg(&oo0)
.arg(&oo1)
.arg(&oo2)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(Some((y0, y1, y2)))
}
pub fn matmul_q8_fused3(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
w2: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
{
if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
return Ok(Some(self.e4m3_fused3_core(
p0.0,
p1.0,
p2.0,
aq,
ad,
w0.in_features(),
p0.1,
p1.1,
p2.1,
p0.2,
p0.3,
p1.3,
p2.3,
)?));
}
let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else {
return Ok(None);
};
Ok(Some(self.q8_fused3_core(
p0.0,
p1.0,
p2.0,
aq,
ad,
w0.in_features(),
p0.1,
p1.1,
p2.1,
p0.2,
)?))
}
#[allow(clippy::too_many_arguments)]
fn q8_fused3_core(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
b2: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
out0: usize,
out1: usize,
out2: usize,
row_bytes: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4;
let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
let f = self.func("qmatvec_q8_0_mmvq_fused3");
let mut y0 = self.alloc_uninit::<f32>(out0)?;
let mut y1 = self.alloc_uninit::<f32>(out1)?;
let mut y2 = self.alloc_uninit::<f32>(out2)?;
let cfg = LaunchConfig {
grid_dim: (nb0 + nb1 + nb2, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, o0, o1, o2, rbl) = (
in_f as i32,
out0 as i32,
out1 as i32,
out2 as i32,
row_bytes as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(b2)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&mut y2)
.arg(&inf)
.arg(&o0)
.arg(&o1)
.arg(&o2)
.arg(&rbl);
unsafe {
b.launch(cfg)?;
}
Ok((y0, y1, y2))
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_q8_fused3_raw(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
b2: &CudaSlice<u8>,
x: &CudaSlice<f32>,
in_f: usize,
out0: usize,
out1: usize,
out2: usize,
row_bytes: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
self.q8_fused3_core(b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes)
}
pub fn matmul_q8_fused2_t(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
if !(2..=8).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() {
return Ok(None);
}
if let Some([p0, p1]) = self.e4m3_fused_params(&[w0, w1]) {
if m > 4 && !Self::b8_enabled() {
return Ok(None);
}
return Ok(Some(self.e4m3_fused2_t_core(
p0.0,
p1.0,
aq,
ad,
m,
w0.in_features(),
p0.1,
p1.1,
p0.2,
p0.3,
p1.3,
)?));
}
let Some([p0, p1]) = self.q8_fused_params(&[w0, w1]) else {
return Ok(None);
};
Ok(Some(self.q8_fused2_t_core(
p0.0,
p1.0,
aq,
ad,
m,
w0.in_features(),
p0.1,
p1.1,
p0.2,
)?))
}
#[allow(clippy::too_many_arguments)]
fn q8_fused2_t_core(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out0: usize,
out1: usize,
row_bytes: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
let f = self.func(match Self::batched_mcols(m) {
2 => "qmatvec_q8_0_mmvq_fused2_b2",
4 => "qmatvec_q8_0_mmvq_fused2_b4",
_ => "qmatvec_q8_0_mmvq_fused2_b8",
});
let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
let cfg = LaunchConfig {
grid_dim: (nb0 + nb1, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, o0, o1, mi, rbl) = (
in_f as i32,
out0 as i32,
out1 as i32,
m as i32,
row_bytes as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&inf)
.arg(&o0)
.arg(&o1)
.arg(&mi)
.arg(&rbl);
unsafe {
b.launch(cfg)?;
}
Ok((y0, y1))
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_q8_fused2_t_raw(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out0: usize,
out1: usize,
row_bytes: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
self.q8_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes)
}
#[allow(clippy::too_many_arguments)]
pub fn matmul_q8_fused3_t(
&self,
w0: &crate::model::GpuTensor,
w1: &crate::model::GpuTensor,
w2: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<Option<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>>
{
if !(2..=4).contains(&m) || std::env::var("MEMRA_NO_BATCHED").is_ok() {
return Ok(None);
}
if let Some([p0, p1, p2]) = self.e4m3_fused_params(&[w0, w1, w2]) {
return Ok(Some(self.e4m3_fused3_t_core(
p0.0,
p1.0,
p2.0,
aq,
ad,
m,
w0.in_features(),
p0.1,
p1.1,
p2.1,
p0.2,
p0.3,
p1.3,
p2.3,
)?));
}
let Some([p0, p1, p2]) = self.q8_fused_params(&[w0, w1, w2]) else {
return Ok(None);
};
Ok(Some(self.q8_fused3_t_core(
p0.0,
p1.0,
p2.0,
aq,
ad,
m,
w0.in_features(),
p0.1,
p1.1,
p2.1,
p0.2,
)?))
}
#[allow(clippy::too_many_arguments)]
fn q8_fused3_t_core(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
b2: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out0: usize,
out1: usize,
out2: usize,
row_bytes: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4;
let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
let f = self.func(if Self::batched_mcols(m) == 2 {
"qmatvec_q8_0_mmvq_fused3_b2"
} else {
"qmatvec_q8_0_mmvq_fused3_b4"
});
let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
let cfg = LaunchConfig {
grid_dim: (nb0 + nb1 + nb2, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, o0, o1, o2, mi, rbl) = (
in_f as i32,
out0 as i32,
out1 as i32,
out2 as i32,
m as i32,
row_bytes as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(b2)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&mut y2)
.arg(&inf)
.arg(&o0)
.arg(&o1)
.arg(&o2)
.arg(&mi)
.arg(&rbl);
unsafe {
b.launch(cfg)?;
}
Ok((y0, y1, y2))
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_q8_fused3_t_raw(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
b2: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out0: usize,
out1: usize,
out2: usize,
row_bytes: usize,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
self.q8_fused3_t_core(b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes)
}
pub fn q8_ffn_fuse2_on(&self) -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_Q8_FFN_FUSE2").as_deref() != Ok("0"))
}
#[allow(clippy::type_complexity)]
fn q8_fused_params<'w, const N: usize>(
&self,
ws: &[&'w crate::model::GpuTensor; N],
) -> Option<[(&'w CudaSlice<u8>, usize, usize); N]> {
use crate::model::GpuTensor;
if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") {
return None;
}
if std::env::var("MEMRA_Q8_DUAL").is_ok_and(|v| v == "0") {
return None;
}
let in_f = ws[0].in_features();
let mut out: [Option<(&CudaSlice<u8>, usize, usize)>; N] = [None; N];
for (i, w) in ws.iter().enumerate() {
match w {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
..
} if *qtype == QT_Q8_0 && *scale == 1.0 && w.in_features() == in_f => {
out[i] = Some((bytes, w.out_features(), *row_bytes))
}
_ => return None,
}
}
Some(out.map(|o| o.unwrap()))
}
pub fn e4m3_dual_on(&self) -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_E4M3_DUAL").as_deref() != Ok("0"))
}
#[allow(clippy::type_complexity)]
fn e4m3_fused_params<'w, const N: usize>(
&self,
ws: &[&'w crate::model::GpuTensor; N],
) -> Option<[(&'w CudaSlice<u8>, usize, usize, f32); N]> {
use crate::model::GpuTensor;
if !self.e4m3_dual_on() {
return None;
}
let in_f = ws[0].in_features();
let mut out: [Option<(&CudaSlice<u8>, usize, usize, f32)>; N] = [None; N];
for (i, w) in ws.iter().enumerate() {
match w {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
rp,
rp4,
..
} if *qtype == QT_F8_E4M3
&& w.in_features() == in_f
&& *row_bytes == in_f
&& !*rp
&& rp4.is_none() =>
{
out[i] = Some((bytes, w.out_features(), *row_bytes, *scale))
}
_ => return None,
}
}
Some(out.map(|o| o.unwrap()))
}
#[allow(clippy::too_many_arguments)]
fn e4m3_fused2_core(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
out0: usize,
out1: usize,
row_bytes: usize,
ws0: f32,
ws1: f32,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4; let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
let f = self.func("qmatvec_e4m3_mmvq_fused2");
let mut y0 = self.alloc_uninit::<f32>(out0)?;
let mut y1 = self.alloc_uninit::<f32>(out1)?;
let cfg = LaunchConfig {
grid_dim: (nb0 + nb1, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, o0, o1, rbl) = (in_f as i32, out0 as i32, out1 as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&inf)
.arg(&o0)
.arg(&o1)
.arg(&rbl)
.arg(&ws0)
.arg(&ws1);
unsafe {
b.launch(cfg)?;
}
Ok((y0, y1))
}
#[allow(clippy::too_many_arguments)]
fn e4m3_fused3_core(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
b2: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
in_f: usize,
out0: usize,
out1: usize,
out2: usize,
row_bytes: usize,
ws0: f32,
ws1: f32,
ws2: f32,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4;
let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
let f = self.func("qmatvec_e4m3_mmvq_fused3");
let mut y0 = self.alloc_uninit::<f32>(out0)?;
let mut y1 = self.alloc_uninit::<f32>(out1)?;
let mut y2 = self.alloc_uninit::<f32>(out2)?;
let cfg = LaunchConfig {
grid_dim: (nb0 + nb1 + nb2, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, o0, o1, o2, rbl) = (
in_f as i32,
out0 as i32,
out1 as i32,
out2 as i32,
row_bytes as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(b2)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&mut y2)
.arg(&inf)
.arg(&o0)
.arg(&o1)
.arg(&o2)
.arg(&rbl)
.arg(&ws0)
.arg(&ws1)
.arg(&ws2);
unsafe {
b.launch(cfg)?;
}
Ok((y0, y1, y2))
}
#[allow(clippy::too_many_arguments)]
fn e4m3_fused2_t_core(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out0: usize,
out1: usize,
row_bytes: usize,
ws0: f32,
ws1: f32,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4;
let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
let f = self.func(match Self::batched_mcols(m) {
2 => "qmatvec_e4m3_mmvq_fused2_b2",
4 => "qmatvec_e4m3_mmvq_fused2_b4",
_ => "qmatvec_e4m3_mmvq_fused2_b8",
});
let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
let cfg = LaunchConfig {
grid_dim: (nb0 + nb1, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, o0, o1, mi, rbl) = (
in_f as i32,
out0 as i32,
out1 as i32,
m as i32,
row_bytes as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&inf)
.arg(&o0)
.arg(&o1)
.arg(&mi)
.arg(&rbl);
unsafe {
b.launch(cfg)?;
}
if ws0 != 1.0 {
self.scale_inplace(&mut y0, ws0, m * out0)?;
}
if ws1 != 1.0 {
self.scale_inplace(&mut y1, ws1, m * out1)?;
}
Ok((y0, y1))
}
#[allow(clippy::too_many_arguments)]
fn e4m3_fused3_t_core(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
b2: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out0: usize,
out1: usize,
out2: usize,
row_bytes: usize,
ws0: f32,
ws1: f32,
ws2: f32,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4;
let nb0 = (out0 as u32).div_ceil(ROWS_PER_BLOCK);
let nb1 = (out1 as u32).div_ceil(ROWS_PER_BLOCK);
let nb2 = (out2 as u32).div_ceil(ROWS_PER_BLOCK);
let f = self.func(if Self::batched_mcols(m) == 2 {
"qmatvec_e4m3_mmvq_fused3_b2"
} else {
"qmatvec_e4m3_mmvq_fused3_b4"
});
let mut y0 = self.alloc_uninit::<f32>(m * out0)?;
let mut y1 = self.alloc_uninit::<f32>(m * out1)?;
let mut y2 = self.alloc_uninit::<f32>(m * out2)?;
let cfg = LaunchConfig {
grid_dim: (nb0 + nb1 + nb2, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, o0, o1, o2, mi, rbl) = (
in_f as i32,
out0 as i32,
out1 as i32,
out2 as i32,
m as i32,
row_bytes as i64,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(b0)
.arg(b1)
.arg(b2)
.arg(aq)
.arg(ad)
.arg(&mut y0)
.arg(&mut y1)
.arg(&mut y2)
.arg(&inf)
.arg(&o0)
.arg(&o1)
.arg(&o2)
.arg(&mi)
.arg(&rbl);
unsafe {
b.launch(cfg)?;
}
if ws0 != 1.0 {
self.scale_inplace(&mut y0, ws0, m * out0)?;
}
if ws1 != 1.0 {
self.scale_inplace(&mut y1, ws1, m * out1)?;
}
if ws2 != 1.0 {
self.scale_inplace(&mut y2, ws2, m * out2)?;
}
Ok((y0, y1, y2))
}
pub fn qmatvec_e4m3_blk_mmvq(
&self,
bytes: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
scales: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
scale_cols: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
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, scale_cols, &mut y,
)?;
Ok(y)
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_e4m3_blk_mmvq_into(
&self,
bytes: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
scales: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
scale_cols: usize,
y: &mut CudaSlice<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4; let f = self.func("qmatvec_e4m3_blk_mmvq");
let cfg = LaunchConfig {
grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), m as u32, 1),
block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
let (inf, outf, mi, rb, sc) = (
in_f as i32,
out_f as i32,
m as i32,
row_bytes as i64,
scale_cols as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(scales)
.arg(&mut *y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb)
.arg(&sc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_e4m3_blk_mmvq_batched(
&self,
bytes: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
scales: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
scale_cols: usize,
mcols: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4; debug_assert!(mcols >= m, "blk batched: mcols {mcols} < m {m}");
let name = match mcols {
2 => "qmatvec_e4m3_blk_mmvq_b2",
4 => "qmatvec_e4m3_blk_mmvq_b4",
8 => "qmatvec_e4m3_blk_mmvq_b8",
16 => "qmatvec_e4m3_blk_mmvq_b16",
_ => {
return Err(
format!("qmatvec_e4m3_blk_mmvq_batched: no kernel for mcols {mcols}").into(),
);
}
};
let mut y = self.alloc_uninit::<f32>(m * out_f)?;
let f = self.func(name);
let cfg = LaunchConfig {
grid_dim: ((out_f as u32).div_ceil(ROWS_PER_BLOCK), 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi, rb, sc) = (
in_f as i32,
out_f as i32,
m as i32,
row_bytes as i64,
scale_cols as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(scales)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb)
.arg(&sc);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_e4m3_blk_batched_raw(
&self,
bytes: &CudaSlice<u8>,
x: &CudaSlice<f32>,
scales: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
scale_cols: usize,
mcols: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
self.qmatvec_e4m3_blk_mmvq_batched(
bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols, mcols,
)
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_e4m3_blk_mmvq_raw(
&self,
bytes: &CudaSlice<u8>,
x: &CudaSlice<f32>,
scales: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
scale_cols: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
self.qmatvec_e4m3_blk_mmvq(
bytes, &aq, &ad, scales, m, in_f, out_f, row_bytes, scale_cols,
)
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_e4m3_fused2_raw(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
x: &CudaSlice<f32>,
in_f: usize,
out0: usize,
out1: usize,
row_bytes: usize,
ws0: f32,
ws1: f32,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
self.e4m3_fused2_core(b0, b1, &aq, &ad, in_f, out0, out1, row_bytes, ws0, ws1)
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_e4m3_fused3_raw(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
b2: &CudaSlice<u8>,
x: &CudaSlice<f32>,
in_f: usize,
out0: usize,
out1: usize,
out2: usize,
row_bytes: usize,
ws0: f32,
ws1: f32,
ws2: f32,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, 1, in_f)?;
self.e4m3_fused3_core(
b0, b1, b2, &aq, &ad, in_f, out0, out1, out2, row_bytes, ws0, ws1, ws2,
)
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_e4m3_fused2_t_raw(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out0: usize,
out1: usize,
row_bytes: usize,
ws0: f32,
ws1: f32,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
self.e4m3_fused2_t_core(b0, b1, &aq, &ad, m, in_f, out0, out1, row_bytes, ws0, ws1)
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_e4m3_fused3_t_raw(
&self,
b0: &CudaSlice<u8>,
b1: &CudaSlice<u8>,
b2: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out0: usize,
out1: usize,
out2: usize,
row_bytes: usize,
ws0: f32,
ws1: f32,
ws2: f32,
) -> Result<(CudaSlice<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
self.e4m3_fused3_t_core(
b0, b1, b2, &aq, &ad, m, in_f, out0, out1, out2, row_bytes, ws0, ws1, ws2,
)
}
fn try_e4m3_blk_pre(
&self,
w: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
if let GpuTensor::Quant {
bytes,
qtype,
row_bytes,
blk: Some(g),
..
} = w
{
if *qtype == QT_F8_E4M3_BLK {
if (2..=16).contains(&m)
&& std::env::var("MEMRA_NO_BATCHED").is_err()
&& (m <= 4 || Self::b8_enabled())
{
let mcols = Self::batched_mcols(m);
return Ok(Some(self.qmatvec_e4m3_blk_mmvq_batched(
bytes,
aq,
ad,
&g.scales,
m,
w.in_features(),
w.out_features(),
*row_bytes,
g.cols,
mcols,
)?));
}
return Ok(Some(self.qmatvec_e4m3_blk_mmvq(
bytes,
aq,
ad,
&g.scales,
m,
w.in_features(),
w.out_features(),
*row_bytes,
g.cols,
)?));
}
}
Ok(None)
}
fn try_e4m3_blk_prefill(
&self,
w: &crate::model::GpuTensor,
x: &CudaSlice<f32>,
m: usize,
) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let GpuTensor::Quant {
bytes,
qtype,
blk: Some(g),
..
} = w
else {
return Ok(None);
};
if *qtype != QT_F8_E4M3_BLK {
return Ok(None);
}
if let Some(y) = self.try_fp8_blk_mmq(w, x, m)? {
return Ok(Some(y));
}
let (in_f, out_f) = (w.in_features(), w.out_features());
let slab = self.fp8_blk_dequant_q8_0_dev(bytes, &g.scales, out_f, in_f)?;
let tmp = GpuTensor::Quant {
bytes: slab,
qtype: QT_Q8_0,
row_bytes: in_f / 32 * 34,
ne: vec![in_f as u64, out_f as u64],
scale: 1.0,
rp: false,
#[cfg(memra_cutlass)]
cutlass: None,
fp8: None,
blk: None,
f16: None,
rp4: None,
};
Ok(Some(self.matmul(&tmp, x, m)?))
}
pub fn matmul_pre_noscale(
&self,
w: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<Option<(CudaSlice<f32>, f32)>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
if m == 1 {
if let Some(y) = self.try_e4m3_blk_pre(w, aq, ad, m)? {
return Ok(Some((y, 1.0)));
}
}
if m != 1 || !self.uses_q8_1_fast(w) {
return Ok(None);
}
let in_f = w.in_features();
let out_f = w.out_features();
let (bytes, qtype, row_bytes, scale, rp) = match w {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
rp,
..
} => (bytes, *qtype, *row_bytes, *scale, *rp),
_ => return Ok(None),
};
if self.mmvq_supports(qtype) {
let (mbytes, mrp) = match w {
GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
_ => (bytes, rp),
};
let y = self.qmatvec_mmvq(
mbytes, aq, ad, m, in_f, out_f, qtype, row_bytes, 1.0, mrp,
)?;
return Ok(Some((y, scale)));
}
let name = match qtype {
QT_Q8_0 => "qmatvec_q8_0_dp4a",
QT_Q4_K => "qmatvec_q4_K_dp4a",
QT_Q6_K => "qmatvec_q6_K_dp4a",
QT_Q5_K => "qmatvec_q5_K_dp4a",
QT_Q3_K => "qmatvec_q3_K_dp4a",
QT_NVFP4 => {
if rp {
"qmatvec_nvfp4_dp4a_rp"
} else {
"qmatvec_nvfp4_dp4a"
}
}
QT_IQ4_XS => "qmatvec_iq4_XS_dp4a",
_ => return Ok(None),
};
let f = self.func(name);
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,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(Some((y, scale)))
}
pub fn mmvq_supports(&self, qtype: i32) -> bool {
if qtype == QT_F8_E4M3 {
return true;
}
if std::env::var("MEMRA_MMVQ").as_deref() == Ok("0") {
return false;
}
matches!(
qtype,
QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_Q4_0
)
}
pub fn qmatvec_mmvq(
&self,
bytes: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
scale: f32,
rp: bool,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
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,
)?;
Ok(y)
}
#[allow(clippy::too_many_arguments)]
pub fn qmatvec_mmvq_into(
&self,
bytes: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
scale: f32,
rp: bool,
y: &mut CudaSlice<f32>,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert!(y.len() >= m * out_f);
const ROWS_PER_BLOCK: u32 = 4; if qtype == QT_Q8_0
&& rp
&& m == 1
&& out_f >= 64
&& (out_f as u32).div_ceil(ROWS_PER_BLOCK) < 4 * self.sm_count() as u32
&& {
static G2: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*G2.get_or_init(|| std::env::var("MEMRA_Q80_G2").as_deref() != Ok("0"))
}
{
let f = self.func("qmatvec_q8_0_mmvq_rp_g2");
let cfg = LaunchConfig {
grid_dim: ((out_f as u32).div_ceil(2), 1, 1),
block_dim: (32, 2, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, 1i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(&mut *y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
if scale != 1.0 {
self.scale_inplace(y, scale, out_f)?;
}
return Ok(());
}
let mut mr: u32 = if m == 1 && (qtype == QT_NVFP4 || qtype == QT_Q5_K) {
2
} else {
1
};
if m == 1 && qtype == QT_Q4_0 {
static Q40MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
mr = *Q40MR.get_or_init(|| {
std::env::var("MEMRA_Q40_MR")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(1)
});
}
let q5_mode = std::env::var("MEMRA_Q5K_ISSUE").ok();
let q5_force = q5_mode.as_deref() == Some("2");
let q5_il = qtype == QT_Q5_K
&& m == 1
&& (q5_force || q5_mode.as_deref().map(|v| v != "0").unwrap_or(true));
if q5_il && !q5_force && out_f > 65536 {
mr = 1;
}
if qtype == QT_Q4_0 && rp && mr != 1 {
mr = 2;
}
if qtype == QT_Q8_0 && rp {
static Q80MR: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
mr = *Q80MR.get_or_init(|| {
std::env::var("MEMRA_Q80_MR")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(1)
});
}
let name = match (qtype, mr, rp) {
(QT_NVFP4, 2, false) => "qmatvec_nvfp4_mmvq_mr2",
(QT_NVFP4, 2, true) => "qmatvec_nvfp4_mmvq_mr2_rp",
(QT_NVFP4, _, true) => "qmatvec_nvfp4_mmvq_rp",
(QT_Q4_0, 1, true) => "qmatvec_q4_0_mmvq_rp",
(QT_Q4_0, _, true) => "qmatvec_q4_0_mmvq_mr2_rp",
(QT_Q5_K, 2, _) => {
if q5_il {
"qmatvec_q5_K_mmvq_mr2_il"
} else {
"qmatvec_q5_K_mmvq_mr2"
}
}
(QT_Q8_0, 2, true) => "qmatvec_q8_0_mmvq_mr2_rp",
(QT_Q8_0, _, true)
if in_f % 1024 == 0 && {
static CA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*CA.get_or_init(|| std::env::var("MEMRA_Q80_CA").as_deref() == Ok("1"))
} =>
{
"qmatvec_q8_0_mmvq_rpca"
}
(QT_Q8_0, _, true) => "qmatvec_q8_0_mmvq_rp",
(QT_Q8_0, _, _) => "qmatvec_q8_0_mmvq",
(QT_Q4_K, _, true) => "qmatvec_q4_K_mmvq_rp",
(QT_Q6_K, _, true) => "qmatvec_q6_K_mmvq_rp",
(QT_Q4_K, _, _) => "qmatvec_q4_K_mmvq",
(QT_Q4_0, 2, false) => "qmatvec_q4_0_mmvq_mr2",
(QT_Q4_0, _, false) => "qmatvec_q4_0_mmvq",
(QT_Q5_K, _, _) => {
if q5_il {
"qmatvec_q5_K_mmvq_il"
} else {
"qmatvec_q5_K_mmvq"
}
}
(QT_Q6_K, _, _) => "qmatvec_q6_K_mmvq",
(QT_NVFP4, _, false) => "qmatvec_nvfp4_mmvq",
(QT_F8_E4M3, _, _) => "qmatvec_e4m3_mmvq",
_ => panic!("qmatvec_mmvq: qtype {qtype} has no MMVQ kernel"),
};
let f = self.func(name);
let rows_per_block = ROWS_PER_BLOCK * mr;
let cfg = LaunchConfig {
grid_dim: (
(out_f as u32 + rows_per_block - 1) / rows_per_block,
m as u32,
1,
),
block_dim: (32, ROWS_PER_BLOCK, 1), shared_mem_bytes: 0, };
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
if qtype == QT_NVFP4 || qtype == QT_F8_E4M3 {
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(&mut *y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb)
.arg(&scale);
unsafe {
b.launch(cfg)?;
}
} else if Self::pdl_on()
&& Self::pdl_mmvq_on()
&& matches!(
name,
"qmatvec_q4_0_mmvq_rp" | "qmatvec_q6_K_mmvq" | "qmatvec_q6_K_mmvq_rp"
)
{
{
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pw, _g0) = bytes.device_ptr(s);
let (paq, _g1) = aq.device_ptr(s);
let (pad, _g2) = ad.device_ptr(s);
let (py, _g3) = y.device_ptr_mut(s);
let mut ps = [
&pw as *const _ as *mut std::ffi::c_void,
&paq as *const _ as *mut _,
&pad as *const _ as *mut _,
&py as *const _ as *mut _,
&inf as *const _ as *mut _,
&outf as *const _ as *mut _,
&mi as *const _ as *mut _,
&rb as *const _ as *mut _,
];
unsafe {
self.launch_pdl(name, cfg.grid_dim, cfg.block_dim, &mut ps)?;
}
}
if scale != 1.0 {
self.scale_inplace(y, scale, m * out_f)?;
}
} else {
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(&mut *y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
if scale != 1.0 {
self.scale_inplace(y, scale, m * out_f)?;
}
}
Ok(())
}
pub fn qmatvec_mmvq_raw(
&self,
bytes: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
rp: bool,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
self.qmatvec_mmvq(bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, 1.0, rp)
}
pub fn batched_supports(&self, qtype: i32) -> bool {
matches!(
qtype,
QT_Q8_0 | QT_Q4_K | QT_Q5_K | QT_Q6_K | QT_NVFP4 | QT_F8_E4M3 | QT_Q4_0
)
}
pub fn iq_fast_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_IQ_FAST")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub fn b8_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_B8").map(|v| v != "0").unwrap_or(true))
}
pub fn batched_mcols(m: usize) -> usize {
if m == 2 {
2
} else if m <= 4 {
4
} else if m <= 8 {
8
} else {
16
}
}
fn batched_kernel_name(qtype: i32, mcols: usize) -> Option<&'static str> {
Some(match (qtype, mcols) {
(QT_Q8_0, 2) => "qmatvec_q8_0_mmvq_b2",
(QT_Q8_0, 4) => "qmatvec_q8_0_mmvq_b4",
(QT_Q8_0, 8) => "qmatvec_q8_0_mmvq_b8",
(QT_Q8_0, 16) => "qmatvec_q8_0_mmvq_b16",
(QT_Q4_K, 2) => "qmatvec_q4_K_mmvq_b2",
(QT_Q4_K, 4) => "qmatvec_q4_K_mmvq_b4",
(QT_Q4_K, 8) => "qmatvec_q4_K_mmvq_b8",
(QT_Q4_K, 16) => "qmatvec_q4_K_mmvq_b16",
(QT_Q5_K, 2) => "qmatvec_q5_K_mmvq_b2",
(QT_Q5_K, 4) => "qmatvec_q5_K_mmvq_b4",
(QT_Q5_K, 8) => "qmatvec_q5_K_mmvq_b8",
(QT_Q5_K, 16) => "qmatvec_q5_K_mmvq_b16",
(QT_Q6_K, 2) => "qmatvec_q6_K_mmvq_b2",
(QT_Q6_K, 4) => "qmatvec_q6_K_mmvq_b4",
(QT_Q6_K, 8) => "qmatvec_q6_K_mmvq_b8",
(QT_Q6_K, 16) => "qmatvec_q6_K_mmvq_b16",
(QT_NVFP4, 2) => "qmatvec_nvfp4_mmvq_b2",
(QT_NVFP4, 4) => "qmatvec_nvfp4_mmvq_b4",
(QT_NVFP4, 8) => "qmatvec_nvfp4_mmvq_b8",
(QT_NVFP4, 16) => "qmatvec_nvfp4_mmvq_b16",
(QT_F8_E4M3, 2) => "qmatvec_e4m3_mmvq_b2",
(QT_F8_E4M3, 4) => "qmatvec_e4m3_mmvq_b4",
(QT_F8_E4M3, 8) => "qmatvec_e4m3_mmvq_b8",
(QT_F8_E4M3, 16) => "qmatvec_e4m3_mmvq_b16",
(QT_Q4_0, 2) => "qmatvec_q4_0_mmvq_b2",
(QT_Q4_0, 4) => "qmatvec_q4_0_mmvq_b4",
(QT_Q4_0, 8) => "qmatvec_q4_0_mmvq_b8",
(QT_Q4_0, 16) => "qmatvec_q4_0_mmvq_b16",
_ => return None,
})
}
pub fn sm_count(&self) -> i32 {
static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
*SMS.get_or_init(|| {
use cudarc::driver::sys::CUdevice_attribute_enum as A;
self.gpu
.ctx
.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT)
.unwrap_or(82)
})
}
pub fn batched_variant(
&self,
_m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
mcols: usize,
rp: bool,
) -> &'static str {
if qtype == QT_Q8_0 {
return if rp { "rp" } else { "base" };
}
static BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
let bv = *BV.get_or_init(|| match std::env::var("MEMRA_MMVQ_BV").as_deref() {
Ok("base") => "base",
Ok("pf") => "pf",
Ok("r2") => "r2",
Ok("r2w8") => "r2w8",
Ok("pfr2") => "pfr2",
Ok("ca") => "ca",
Ok("car2") => "car2",
Ok("rp") => "rp",
Ok("rpr2") => "rpr2",
Ok("rpr2w8") => "rpr2w8",
Ok("rpca") => "rpca",
Ok("rpcar2") => "rpcar2",
Ok("rpsc") => "rpsc",
Ok("rpms") => "rpms",
Ok("rpmsc") => "rpmsc",
Ok("rpks") => "rpks",
Ok("rpksc") => "rpksc",
_ => "auto",
});
let ca_ok = qtype == QT_NVFP4 && (row_bytes % 16 == 0) && (in_f % 1024 == 0);
static KS_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let ks_on = *KS_ON.get_or_init(|| std::env::var("MEMRA_KS").as_deref() != Ok("0"));
let sc_ok = ks_on && qtype == QT_NVFP4 && (in_f % 256 == 0) && (in_f / 64 <= 272);
let ks_ok = ks_on && qtype == QT_NVFP4 && (in_f % 512 == 0) && (in_f / 64 <= 272);
static SMS: std::sync::OnceLock<i32> = std::sync::OnceLock::new();
let sms = *SMS.get_or_init(|| {
use cudarc::driver::sys::CUdevice_attribute_enum as A;
self.gpu
.ctx
.attribute(A::CU_DEVICE_ATTRIBUTE_MULTIPROCESSOR_COUNT)
.unwrap_or(82)
});
let kq_r2 = matches!(qtype, QT_Q4_K | QT_Q5_K | QT_Q6_K);
static KQBV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
let kq_bv = *KQBV.get_or_init(|| match std::env::var("MEMRA_KQ_BV").as_deref() {
Ok("base") => "base",
Ok("r2") => "r2",
Ok("r2w8") => "r2w8",
_ => "auto",
});
let variant: &'static str = if qtype == QT_Q4_0 {
static Q40BV: std::sync::OnceLock<&'static str> = std::sync::OnceLock::new();
let q40 = *Q40BV.get_or_init(|| match std::env::var("MEMRA_Q40_BV").as_deref() {
Ok("base") => "base",
Ok("r2") => "r2",
Ok("ms") => "ms",
Ok("sm") => "sm",
Ok("la") => "la",
_ => "auto",
});
let v = if q40 != "auto" {
q40
} else if (out_f as u32).div_ceil(8) >= 4 * sms as u32 {
"r2"
} else {
"base"
};
if rp {
match v {
"ms" => "r2ms_rp",
"sm" => "r2sm_rp",
"la" => "r2la_rp",
"r2" => "r2_rp",
_ => "rp",
}
} else if matches!(v, "ms" | "sm" | "la") {
"r2"
} else {
v
}
} else if qtype != QT_NVFP4 && !kq_r2 {
"base"
} else if kq_r2 && rp {
"rp"
} else if kq_r2 {
if kq_bv != "auto" {
if kq_bv == "r2w8" && mcols != 4 {
"r2"
} else {
kq_bv
}
} else if bv != "auto" {
match bv {
"r2" | "pfr2" | "rpr2" | "car2" => "r2",
"r2w8" | "rpr2w8" => {
if mcols != 4 {
"r2"
} else {
"r2w8"
}
}
_ => "base", }
} else {
let blocks = (out_f + 7) / 8;
let waves = blocks as f64 / (7 * sms as usize) as f64;
let filled = blocks >= 4 * sms as usize;
let use_r2 = if qtype == QT_Q4_K {
filled
} else {
waves >= 2.0
};
if use_r2 { "r2" } else { "base" }
}
} else if bv != "auto" {
let v = if bv == "r2w8" && mcols == 2 {
"r2"
} else if bv == "ca" && (!ca_ok || mcols == 8) {
"pf"
} else if bv == "car2" && (!ca_ok || mcols == 8) {
"r2"
} else if bv == "pfr2" && mcols == 8 {
"r2"
} else if (bv == "rpr2w8" || bv == "rpr2") && mcols == 2 {
"rpr2"
}
else if (bv == "rpca" || bv == "rpcar2") && (!ca_ok || mcols == 8) {
if mcols == 8 { "rpr2w8" } else { "rpr2" }
} else if bv == "rpcar2" && mcols == 2 {
"rpca"
}
else if (bv == "rpsc" || bv == "rpmsc") && !sc_ok {
"rpr2"
} else if (bv == "rpks" || bv == "rpksc") && !ks_ok {
"rpr2"
} else {
bv
};
if rp {
match v {
"base" | "pf" | "ca" | "rp" => "rp",
"r2" | "pfr2" | "car2" | "rpr2" => "rpr2",
"r2w8" | "rpr2w8" => {
if mcols == 2 {
"rpr2"
} else {
"rpr2w8"
}
}
other => other, }
} else {
v
}
} else if mcols == 8 {
if rp {
if sc_ok { "rpsc" } else { "rpr2w8" }
} else {
"r2w8"
}
} else if mcols >= 4 {
let blocks = (out_f + 7) / 8;
let r7 = 7 * sms as usize;
let r8 = 8 * sms as usize;
let waves = blocks as f64 / r7 as f64;
let filled = blocks >= 4 * sms as usize;
if filled && blocks.div_ceil(r8) < blocks.div_ceil(r7) {
if rp { "rpr2w8" } else { "r2w8" }
} else if waves >= 2.0 || (waves <= 1.0 && filled) {
if rp { "rpr2" } else { "r2" }
} else {
if rp { "rp" } else { "pf" }
}
} else if in_f >= 6144 {
if rp { "rpr2" } else { "r2" }
} else if rp {
let waves = ((out_f + 7) / 8) as f64 / (7 * sms as usize) as f64;
if sc_ok && waves >= 0.9 && waves <= 1.1 {
"rpsc"
} else {
"rp"
}
} else {
"base"
};
variant
}
pub fn qmatvec_mmvq_batched(
&self,
bytes: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
mcols: usize,
scale: f32,
rp: bool,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
const ROWS_PER_BLOCK: u32 = 4;
let forced: Option<&'static str> = {
static V: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
V.get_or_init(|| std::env::var("MEMRA_BVAR").ok())
.as_deref()
.map(|s| Box::leak(s.to_string().into_boxed_str()) as &'static str)
};
let variant = match forced {
Some(v) if !rp || v.contains("rp") => v,
_ => self.batched_variant(m, in_f, out_f, qtype, row_bytes, mcols, rp),
};
let base_name = Self::batched_kernel_name(qtype, mcols).ok_or_else(|| {
format!("qmatvec_mmvq_batched: no kernel for qtype {qtype} mcols {mcols}")
})?;
let variant = if mcols == 16 {
if rp { "rp" } else { "base" }
} else {
variant
};
static B567: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let b567 = *B567.get_or_init(|| std::env::var("MEMRA_B567").as_deref() != Ok("0"));
if b567
&& qtype == QT_NVFP4
&& rp
&& mcols == 8
&& (5..=7).contains(&m)
&& matches!(variant, "rpsc" | "rpr2w8")
{
let f = self.func(&format!("qmatvec_nvfp4_mmvq_b{m}_{variant}"));
let rows_per_block = ROWS_PER_BLOCK * 2; let mut y = self.alloc_uninit::<f32>(m * out_f)?;
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
if scale != 1.0 {
self.scale_inplace(&mut y, scale, m * out_f)?;
}
return Ok(y);
}
let (name, rows_per_block): (std::borrow::Cow<'static, str>, u32) = match variant {
"base" => (base_name.into(), ROWS_PER_BLOCK),
"pf" => (format!("{base_name}_pf").into(), ROWS_PER_BLOCK),
"ca" => (format!("{base_name}_ca").into(), ROWS_PER_BLOCK),
"rp" => (format!("{base_name}_rp").into(), ROWS_PER_BLOCK),
"rpca" => (format!("{base_name}_rpca").into(), ROWS_PER_BLOCK), "rpks" => (format!("{base_name}_rpks").into(), ROWS_PER_BLOCK),
"rpksc" => (format!("{base_name}_rpksc").into(), ROWS_PER_BLOCK),
"rpms" => (format!("{base_name}_rpms").into(), ROWS_PER_BLOCK),
"rpmsc" => (format!("{base_name}_rpmsc").into(), ROWS_PER_BLOCK),
"r2ms_rp" => (format!("{base_name}_r2ms_rp").into(), ROWS_PER_BLOCK),
"r2sm_rp" => (format!("{base_name}_r2sm_rp").into(), ROWS_PER_BLOCK * 2),
"r2la_rp" => (format!("{base_name}_r2la_rp").into(), ROWS_PER_BLOCK * 2),
v => (format!("{base_name}_{v}").into(), ROWS_PER_BLOCK * 2), };
debug_assert!(
!rp || name.contains("_rp"),
"rp weight dispatched to a GGUF-layout kernel"
);
let f = self.func(&name);
let mut y = self.alloc_uninit::<f32>(m * out_f)?;
let smem = if name.contains("_r2sm_rp") {
(mcols * 32 * 9 * 4 + mcols * 32 * 4) as u32
} else {
0
};
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + rows_per_block - 1) / rows_per_block, 1, 1),
block_dim: (32, ROWS_PER_BLOCK, 1),
shared_mem_bytes: smem,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
if scale != 1.0 {
self.scale_inplace(&mut y, scale, m * out_f)?;
}
Ok(y)
}
pub fn qmatvec_batched_raw(
&self,
bytes: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
mcols: usize,
rp: bool,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
self.qmatvec_mmvq_batched(
bytes, &aq, &ad, m, in_f, out_f, qtype, row_bytes, mcols, 1.0, rp,
)
}
pub fn qmatvec_nvfp4_batched_raw(
&self,
bytes: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
row_bytes: usize,
mcols: usize,
rp: bool,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
self.qmatvec_batched_raw(bytes, x, m, in_f, out_f, QT_NVFP4, row_bytes, mcols, rp)
}
fn try_fp4_gemm(
&self,
w: &crate::model::GpuTensor,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
if cfg!(memra_portable_cuda) {
return Ok(None);
}
if std::env::var("MEMRA_FP4").is_err() {
return Ok(None);
}
#[cfg(memra_cutlass)]
if m >= 128 && std::env::var("MEMRA_FP4_CUTLASS").is_ok() {
if let GpuTensor::Quant {
bytes,
qtype,
scale,
row_bytes,
cutlass,
..
} = w
{
if *qtype == QT_NVFP4 && in_f % 64 == 0 {
if let Some(cw) = cutlass {
let y = self.cutlass_fp4_gemm(
&cw.b_packed,
&cw.sfb_swizzled,
x,
*scale,
m,
out_f,
in_f,
)?;
return Ok(Some(y));
} else if std::env::var("MEMRA_FP4_CUTLASS_OTF").is_ok() {
let (b_packed, sfb_sw) =
self.build_cutlass_weight(bytes, out_f, in_f, *row_bytes)?;
let y =
self.cutlass_fp4_gemm(&b_packed, &sfb_sw, x, *scale, m, out_f, in_f)?;
return Ok(Some(y));
}
}
}
}
if let GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
rp,
..
} = w
{
if *qtype == QT_NVFP4 && in_f % 64 == 0 && !*rp {
let y =
self.qmatvec_gemm_nvfp4_fp4(bytes, x, m, in_f, out_f, *row_bytes, *scale)?;
return Ok(Some(y));
}
}
Ok(None)
}
pub fn rms_norm_f16out(
&self,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
dst16: &mut CudaSlice<u8>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rms_norm_f16out_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(w).arg(dst).arg(dst16).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn add_rms_norm_f16out(
&self,
a: &CudaSlice<f32>,
b: &CudaSlice<f32>,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
dst16: &mut CudaSlice<u8>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("add_rms_norm_f16out_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (rms_block(), 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(a)
.arg(b)
.arg(w)
.arg(res)
.arg(dst)
.arg(dst16)
.arg(&nc)
.arg(&e);
unsafe {
lb.launch(cfg)?;
}
Ok(())
}
pub fn matmul_group_xh(
&self,
ws: &[&crate::model::GpuTensor],
x: &CudaSlice<f32>,
xh: &CudaSlice<u8>,
m: usize,
) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
let mut out = Vec::with_capacity(ws.len());
let in_f = ws[0].in_features();
for w in ws {
if w.in_features() == in_f && m >= 16 && !self.verify_exact_on() {
if let Some(y) = self.try_f16_gemm_pre(w, xh, m)? {
out.push(y);
continue;
}
}
out.push(self.matmul(w, x, m)?);
}
Ok(out)
}
pub fn gdn_pad_mask(
&self,
beta: &mut CudaSlice<f32>,
g_log: &mut CudaSlice<f32>,
len_d: &CudaSlice<i32>,
h: usize,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_pad_mask_f32");
let cfg = LaunchConfig::for_num_elems((t * h) as u32);
let (hi, ti) = (h as i32, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(beta).arg(g_log).arg(len_d).arg(&hi).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn row_gather_dev(
&self,
src: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
len_d: &CudaSlice<i32>,
ncols: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("row_gather_dev_f32");
let cfg = LaunchConfig::for_num_elems(ncols as u32);
let nc = ncols as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src).arg(dst).arg(len_d).arg(&nc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn matmul_group(
&self,
ws: &[&crate::model::GpuTensor],
x: &CudaSlice<f32>,
m: usize,
) -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let mut out = Vec::with_capacity(ws.len());
let any_mirror = ws
.iter()
.any(|w| matches!(w, GpuTensor::Quant { f16: Some(_), .. }));
if m >= 16 && any_mirror && !self.verify_exact_on() {
let in_f = ws[0].in_features();
let xh = self.f16_act(x, m * in_f, in_f)?;
for w in ws {
if w.in_features() == in_f {
if let Some(y) = self.try_f16_gemm_pre(w, &xh, m)? {
out.push(y);
continue;
}
}
out.push(self.matmul(w, x, m)?);
}
return Ok(out);
}
for w in ws {
out.push(self.matmul(w, x, m)?);
}
Ok(out)
}
pub fn matmul_group_multi(
&self,
ws: &[&crate::model::GpuTensor],
xs: &[&CudaSlice<f32>],
ms: &[usize],
) -> Result<Vec<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
assert_eq!(xs.len(), ms.len());
let in_f = ws[0].in_features();
let total: usize = ms.iter().sum();
let mut xcat = self.uninit(total * in_f)?;
let mut off = 0usize;
for (x, &m) in xs.iter().zip(ms) {
self.copy_into(&mut xcat, off * in_f, x, m * in_f)?;
off += m;
}
let ys = self.matmul_group(ws, &xcat, total)?;
let mut out: Vec<Vec<CudaSlice<f32>>> = (0..xs.len()).map(|_| Vec::new()).collect();
for (w, y) in ws.iter().zip(ys) {
let out_f = w.out_features();
let mut off = 0usize;
for (s, &m) in ms.iter().enumerate() {
let mut ys_s = self.uninit(m * out_f)?;
let src = y.slice(off * out_f..(off + m) * out_f);
self.gpu.stream().memcpy_dtod(&src, &mut ys_s)?;
out[s].push(ys_s);
off += m;
}
}
Ok(out)
}
pub fn gemm_supports(&self, w: &crate::model::GpuTensor) -> bool {
use crate::model::GpuTensor;
if !legacy_quant_gemm_allowed(
cfg!(memra_portable_cuda),
cfg!(memra_hopper_mma),
std::env::var_os("MEMRA_NO_GEMM").is_some(),
) {
return false;
}
match w {
GpuTensor::Quant { qtype, .. } => {
matches!(*qtype, QT_Q8_0 | QT_Q4_K | QT_Q6_K | QT_Q5_K | QT_Q4_0)
|| (*qtype == QT_NVFP4 && w.in_features() % 64 == 0)
}
GpuTensor::Float { .. } | GpuTensor::FloatBf16 { .. } => false,
}
}
pub fn qmatvec_gemm(
&self,
w: &crate::model::GpuTensor,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
let in_f = w.in_features();
let out_f = w.out_features();
let (bytes, qtype, row_bytes, scale, rp) = match w {
GpuTensor::Quant {
bytes,
qtype,
row_bytes,
scale,
rp,
..
} => (bytes, *qtype, *row_bytes, *scale, *rp),
_ => unreachable!("gemm_supports guaranteed Quant"),
};
if cfg!(memra_hopper_mma) && qtype == QT_Q8_0 && out_f % 64 == 0 && wgmma_gemm_enabled() {
if let GpuTensor::Quant { rp4: Some(m4), .. } = w {
let mut y = self.qmatvec_gemm_q8_0_wgmma_raw(m4, aq, ad, m, in_f, out_f)?;
if scale != 1.0 {
self.scale_inplace(&mut y, scale, m * out_f)?;
}
return Ok(y);
}
}
let name = match qtype {
QT_Q8_0 => "qmatvec_gemm_q8_0",
QT_Q4_K => "qmatvec_gemm_q4_K",
QT_Q4_0 => {
if rp {
"qmatvec_gemm_q4_0_rp"
} else {
"qmatvec_gemm_q4_0"
}
}
QT_Q5_K => "qmatvec_gemm_q5_K",
QT_Q6_K => "qmatvec_gemm_q6_K",
QT_NVFP4 => {
if rp {
"qmatvec_gemm_nvfp4_rp"
} else {
"qmatvec_gemm_nvfp4"
}
}
_ => unreachable!(),
};
let f = self.func(name);
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);
let k1_tile = if is_k1 {
k1_launch_override().unwrap_or((128, 128, 8))
} else {
(128, 128, 8)
};
let (bm, bn): (u32, u32) = if is_k1 {
(k1_tile.0, k1_tile.1)
} else {
(64, 256)
};
let warps: u32 = if is_k1 {
k1_tile.2
} else {
match qtype {
QT_NVFP4 => 8,
_ => 4,
}
};
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
block_dim: (32, warps, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
if scale != 1.0 {
self.scale_inplace(&mut y, scale, m * out_f)?;
}
Ok(y)
}
pub fn qmatvec_gemm_raw(
&self,
bytes: &CudaSlice<u8>,
x: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
qtype: i32,
row_bytes: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let (aq, ad) = self.quantize_q8_1(x, m, in_f)?;
let name = match qtype {
QT_Q8_0 => "qmatvec_gemm_q8_0",
QT_Q4_K => "qmatvec_gemm_q4_K",
QT_Q4_0 => "qmatvec_gemm_q4_0",
QT_Q5_K => "qmatvec_gemm_q5_K",
QT_Q6_K => "qmatvec_gemm_q6_K",
QT_NVFP4 => "qmatvec_gemm_nvfp4",
QT_NVFP4_RP => "qmatvec_gemm_nvfp4_rp",
_ => panic!("qmatvec_gemm_raw: qtype {qtype} has no GEMM kernel"),
};
let f = self.func(name);
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);
let k1_tile = if is_k1 {
k1_launch_override().unwrap_or((128, 128, 8))
} else {
(128, 128, 8)
};
let (bm, bn): (u32, u32) = if is_k1 {
(k1_tile.0, k1_tile.1)
} else {
(64, 256)
};
let warps: u32 = if is_k1 {
k1_tile.2
} else {
match qtype {
QT_NVFP4 | QT_NVFP4_RP => 8,
_ => 4,
}
};
let cfg = LaunchConfig {
grid_dim: ((out_f as u32 + bm - 1) / bm, (m as u32 + bn - 1) / bn, 1),
block_dim: (32, warps, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi, rb) = (in_f as i32, out_f as i32, m as i32, row_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(bytes)
.arg(&aq)
.arg(&ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi)
.arg(&rb);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn qmatvec_gemm_q8_0_wgmma_raw(
&self,
rp4: &CudaSlice<u8>,
aq: &CudaSlice<i8>,
ad: &CudaSlice<f32>,
m: usize,
in_f: usize,
out_f: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
assert!(
out_f % 64 == 0 && in_f % 32 == 0,
"wgmma GEMM needs out_f%64==0, in_f%32==0"
);
let f = self.func("qmatvec_gemm_q8_0_wgmma");
let mut y = self.alloc_uninit::<f32>(m * out_f)?; let cfg = LaunchConfig {
grid_dim: ((out_f / 64) as u32, (m as u32).div_ceil(64), 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (inf, outf, mi) = (in_f as i32, out_f as i32, m as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(rp4)
.arg(aq)
.arg(ad)
.arg(&mut y)
.arg(&inf)
.arg(&outf)
.arg(&mi);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn scale_inplace(
&self,
y: &mut CudaSlice<f32>,
s: f32,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("scale_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let (sf, ni) = (s, n as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(y).arg(&sf).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn bf16_to_f32(
&self,
data: &cudarc::driver::CudaView<'_, u8>,
n: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let mut out = self.alloc_uninit::<f32>(n)?;
let f = self.func("bf16_to_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(data).arg(&mut out).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(out)
}
fn linear_bf16_chunked(
&self,
x: &CudaSlice<f32>,
data: &CudaSlice<u8>,
m: usize,
in_f: usize,
out_f: usize,
exact: bool,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
const CHUNK_BYTES: usize = 256 << 20;
let chunk_rows = (CHUNK_BYTES / (in_f * 4)).max(1).min(out_f);
if chunk_rows >= out_f {
let wf32 = self.bf16_to_f32(&data.slice(0..in_f * out_f * 2), in_f * out_f)?;
return if exact {
self.linear_decode_exact(x, &wf32, m, in_f, out_f)
} else {
self.linear(x, &wf32, m, in_f, out_f)
};
}
let mut y = self.alloc_uninit::<f32>(m * out_f)?;
let mut r0 = 0usize;
while r0 < out_f {
let rows = chunk_rows.min(out_f - r0);
let wslice = data.slice(r0 * in_f * 2..(r0 + rows) * in_f * 2);
let wf32 = self.bf16_to_f32(&wslice, in_f * rows)?;
let yc = if exact {
self.linear_decode_exact(x, &wf32, m, in_f, rows)?
} else {
self.linear(x, &wf32, m, in_f, rows)?
};
for mi in 0..m {
let src = yc.slice(mi * rows..(mi + 1) * rows);
let mut dst = y.slice_mut(mi * out_f + r0..mi * out_f + r0 + rows);
self.gpu.stream().memcpy_dtod(&src, &mut dst)?;
}
r0 += rows;
}
Ok(y)
}
pub fn linear_decode_exact(
&self,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
m_tokens: usize,
in_f: usize,
out_f: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
if m_tokens == 1 {
return self.linear(x, w, 1, in_f, out_f);
}
let xv = self.view(x, m_tokens * in_f);
let mut y = self.alloc_uninit::<f32>(m_tokens * out_f)?;
for t in 0..m_tokens {
let row = xv.slice(t * in_f..(t + 1) * in_f);
let mut xr = self.alloc_uninit::<f32>(in_f)?;
self.copy_view_into(&mut xr, 0, &row, in_f)?;
let yr = self.linear(&xr, w, 1, in_f, out_f)?;
self.copy_into(&mut y, t * out_f, &yr, out_f)?;
}
Ok(y)
}
pub fn linear(
&self,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
m_tokens: usize,
in_f: usize,
out_f: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
use cudarc::cublaslt::{Matmul, MatmulConfig};
let mut c = self.alloc_uninit::<f32>(m_tokens * out_f)?; let cfg = MatmulConfig {
transa: true,
transb: false,
transc: false,
m: out_f as u64,
n: m_tokens as u64,
k: in_f as u64,
alpha: 1.0,
lda: in_f as i64,
ldb: in_f as i64,
beta: 0.0,
ldc: out_f as i64,
stride_a: None,
stride_b: None,
stride_c: None,
stride_bias: None,
batch_size: None,
};
unsafe {
self.gpu.blas.matmul(cfg, w, x, &mut c, None, None)?;
}
Ok(c)
}
pub fn sdpa_naive(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("sdpa_naive_f32");
let cfg = LaunchConfig {
grid_dim: (n_head as u32, t as u32, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: (t_kv * 4) as u32,
};
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,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn sdpa_naive_w(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
window: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("sdpa_naive_w_f32");
let cfg = LaunchConfig {
grid_dim: (n_head as u32, t as u32, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: (t_kv * 4) as u32,
};
let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
t_kv as i32,
causal as i32,
window as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&wi);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn sdpa_naive_view(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<f32>,
v: &cudarc::driver::CudaView<f32>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("sdpa_naive_f32");
let cfg = LaunchConfig {
grid_dim: (n_head as u32, t as u32, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: (t_kv * 4) as u32,
};
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,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_dequant_kv_view_f32(
&self,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
kf: &mut CudaSlice<f32>,
vf: &mut CudaSlice<f32>,
kv_dim_k: usize,
kv_dim_v: usize,
t_kv: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = if g {
self.func_g("fa_dequant_kv_ws_f32")
} else {
self.func("fa_dequant_kv_ws_f32")
};
let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk.max(1), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k)
.arg(v)
.arg(&mut *kf)
.arg(&mut *vf)
.arg(&kdk)
.arg(&kdv)
.arg(&tkvi)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn sdpa_naive_quantized_view(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
k_tok_bytes: usize,
v_tok_bytes: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let kv_dim = n_head_kv * head_dim;
let mut kf = self.uninit(t_kv * kv_dim)?;
let mut vf = self.uninit(t_kv * kv_dim)?;
let f = self.func("fa_dequant_kv_ws_f32");
let total = (2 * t_kv * kv_dim) as u64;
let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk.max(1), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k)
.arg(v)
.arg(&mut kf)
.arg(&mut vf)
.arg(&kv_dim_i)
.arg(&kv_dim_i)
.arg(&t_kv_i)
.arg(&k_tok_bytes_i)
.arg(&v_tok_bytes_i);
unsafe { b.launch(cfg)? };
self.sdpa_naive(
q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
)
}
#[allow(clippy::too_many_arguments)]
pub fn sdpa_naive_w_quantized_view(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
window: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let kv_dim = n_head_kv * head_dim;
let mut kf = self.uninit(t_kv * kv_dim)?;
let mut vf = self.uninit(t_kv * kv_dim)?;
let f = self.func("fa_dequant_kv_ws_f32");
let total = (2 * t_kv * kv_dim) as u64;
let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk.max(1), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (kv_dim_i, t_kv_i) = (kv_dim as i32, t_kv as i32);
let (k_tok_bytes_i, v_tok_bytes_i) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k)
.arg(v)
.arg(&mut kf)
.arg(&mut vf)
.arg(&kv_dim_i)
.arg(&kv_dim_i)
.arg(&t_kv_i)
.arg(&k_tok_bytes_i)
.arg(&v_tok_bytes_i);
unsafe { b.launch(cfg)? };
self.sdpa_naive_w(
q, &kf, &vf, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
)
}
pub fn fa_prefill(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
) -> Result<(), Box<dyn std::error::Error>> {
if portable_mma_gated() {
return self.sdpa_naive(
q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
);
}
let fa3_on = head_dim == 256
&& causal
&& t == t_kv
&& match std::env::var("MEMRA_FA3").as_deref() {
Ok("0") => false,
Ok("1") => true,
_ => cfg!(memra_hopper_mma),
};
if fa3_on {
let n = t * n_head * head_dim;
let nkv = t * n_head_kv * head_dim;
let mut q16 = self.alloc_u8_uninit(n * 2)?;
let mut k16 = self.alloc_u8_uninit(nkv * 2)?;
let mut v16 = self.alloc_u8_uninit(nkv * 2)?;
self.f32_to_bf16_into(q, &mut q16, n)?;
self.f32_to_bf16_into(k, &mut k16, nkv)?;
self.f32_to_bf16_into(v, &mut v16, nkv)?;
let rc = {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let stream = self.gpu.stream();
let (qp, _g1) = q16.device_ptr(&stream);
let (kp, _g2) = k16.device_ptr(&stream);
let (vp, _g3) = v16.device_ptr(&stream);
let (op, _g4) = o.device_ptr_mut(&stream);
unsafe {
memra_fa3_prefill(
qp as *const core::ffi::c_void,
kp as *const core::ffi::c_void,
vp as *const core::ffi::c_void,
op as *mut f32,
t as i32,
n_head as i32,
n_head_kv as i32,
head_dim as i32,
scale,
stream.cu_stream() as *mut core::ffi::c_void,
)
}
};
if rc != 0 {
return Err(format!("memra_fa3_prefill rc={rc}").into());
}
return Ok(());
}
static FA_P1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let fa_p1 = *FA_P1.get_or_init(|| std::env::var("MEMRA_FA_P1").as_deref() == Ok("1"));
if fa_p1 && head_dim == 256 && !std::env::var("MEMRA_FA_FLOOR").is_ok() {
const BLOCK_Q: usize = 64;
const BKX: usize = 32;
let f = self.func("fa_prefill_bf16_p1");
let shmem = (2 * (2 * BKX * head_dim + BLOCK_Q * BKX)
+ 4 * (BLOCK_Q * BKX + 2 * BLOCK_Q)) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (
(t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
n_head as u32,
1,
),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
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,
);
let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&qb)
.arg(&kb)
.arg(&vb)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz);
unsafe {
b.launch(cfg)?;
}
return Ok(());
}
const BK: usize = 32;
let w2 = std::env::var("MEMRA_FA_PP_W2").as_deref() == Ok("1");
let (block_q, warps, w2_sfx): (usize, u32, &str) =
if w2 { (32, 2, "_w2") } else { (64, 4, "") };
let hd_sfx = fa_hd_suffix(head_dim)?;
let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
let bf16kv = !floor && !w2 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0");
let (kb16, vb16) = if bf16kv {
let n = t_kv * n_head_kv * head_dim;
let mut kb = self.alloc_u8_uninit(n * 2)?;
let mut vb = self.alloc_u8_uninit(n * 2)?;
let fcv = self.func("f32_to_bf16_bulk");
let ni = n as i64;
let cfgc = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&fcv);
b.arg(k).arg(&mut kb).arg(&ni);
unsafe {
b.launch(cfgc)?;
}
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&fcv);
b.arg(v).arg(&mut vb).arg(&ni);
unsafe {
b.launch(cfgc)?;
}
(Some(kb), Some(vb))
} else {
(None, None)
};
let f = self.func(&if bf16kv {
format!("fa_prefill_bf16kv_pp{hd_sfx}")
} else {
format!(
"fa_prefill_f32{}{}{hd_sfx}",
if floor { "" } else { "_pp" },
if floor { "" } else { w2_sfx }
)
});
let kv_stages = if bf16kv { 2 } else { 1 };
let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
+ 4 * (block_q * BK + 2 * block_q)) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (
(t as u32 + block_q as u32 - 1) / block_q as u32,
n_head as u32,
1,
),
block_dim: (32, warps, 1),
shared_mem_bytes: shmem,
};
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,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q);
match (&kb16, &vb16) {
(Some(kb), Some(vb)) => {
b.arg(kb).arg(vb);
}
_ => {
b.arg(k).arg(v);
}
}
b.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_prefill_w(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
window: usize,
) -> Result<(), Box<dyn std::error::Error>> {
if portable_mma_gated() {
return self.sdpa_naive_w(
q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal, window,
);
}
static FAW_F32: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let faw_f32 =
*FAW_F32.get_or_init(|| std::env::var("MEMRA_FAW_STAGE").as_deref() == Ok("f32"));
let floor = std::env::var("MEMRA_FA_FLOOR").is_ok();
self.fa_prefill_w_arm(
q,
k,
v,
o,
head_dim,
n_head,
n_head_kv,
t,
t_kv,
scale,
causal,
window,
floor || faw_f32,
floor,
)
}
#[allow(clippy::too_many_arguments)]
pub fn fa_prefill_w_pre(
&self,
qb: &CudaSlice<u8>,
kb: &CudaSlice<u8>,
vb: &CudaSlice<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
window: usize,
v_f16: bool,
) -> Result<(), Box<dyn std::error::Error>> {
const BLOCK_Q: usize = 64;
const BK: usize = 32;
debug_assert_eq!(head_dim, 256);
let hp = fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
debug_assert!(!v_f16 || hp, "f16 V emitted but the SWA hp arm is off");
if hp {
const BLOCK_QH: usize = 32;
let mut vguard = self.fa_vf16_scratch.lock().unwrap();
let vh: &CudaSlice<u8> = if v_f16 {
vb
} else {
let n = t_kv * n_head_kv * head_dim;
if vguard.as_ref().map(|b| b.len() < n * 2).unwrap_or(true) {
*vguard = Some(self.alloc_uninit::<u8>(n * 2)?);
}
self.bf16_to_f16_into(vb, n, vguard.as_mut().unwrap())?;
vguard.as_ref().unwrap()
};
let f = self.func("fa_prefill_w_bf16_p1h2");
let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK) + 4 * (2 * BLOCK_QH)) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
t_kv as i32,
causal as i32,
window as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qb)
.arg(kb)
.arg(vh)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&wi);
unsafe {
b.launch(cfg)?;
}
return Ok(());
}
let f = self.func("fa_prefill_w_bf16_p1");
let shmem =
(2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (
(t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
n_head as u32,
1,
),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
t_kv as i32,
causal as i32,
window as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qb)
.arg(kb)
.arg(vb)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&wi);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_prefill_w_arm(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
window: usize,
f32_stage: bool,
floor: bool,
) -> Result<(), Box<dyn std::error::Error>> {
const BLOCK_Q: usize = 64;
const BK: usize = 32;
debug_assert_eq!(head_dim, 256, "fa_prefill_w is stamped hd256 only");
static P1_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let p1 = !floor
&& !f32_stage
&& *P1_ON.get_or_init(|| {
std::env::var("MEMRA_FAW_P1")
.map(|v| v != "0")
.unwrap_or(true)
});
let hp =
p1 && fa_f16pv_on() && faw_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
if hp {
const BLOCK_QH: usize = 32;
let f = self.func("fa_prefill_w_bf16_p1h2");
let shmem = (2 * (2 * BK * head_dim + 2 * BLOCK_QH * BK) + 4 * (2 * BLOCK_QH)) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: ((t as u32).div_ceil(BLOCK_QH as u32), (n_head / 2) as u32, 1),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
t_kv as i32,
causal as i32,
window as i32,
);
let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
let vh = self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&qb)
.arg(&kb)
.arg(&vh)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&wi);
unsafe {
b.launch(cfg)?;
}
return Ok(());
}
if p1 {
let f = self.func("fa_prefill_w_bf16_p1");
let shmem =
(2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (
(t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
n_head as u32,
1,
),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
t_kv as i32,
causal as i32,
window as i32,
);
let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&qb)
.arg(&kb)
.arg(&vb)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&wi);
unsafe {
b.launch(cfg)?;
}
return Ok(());
}
static G4_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let g4 = !floor
&& !f32_stage
&& n_head_kv == 1
&& n_head % 4 == 0
&& *G4_ON.get_or_init(|| {
std::env::var("MEMRA_FAW_G4")
.map(|v| v != "0")
.unwrap_or(true)
});
if g4 {
const SP_M: usize = 16;
static O2_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let o2 = *O2_ON.get_or_init(|| {
std::env::var("MEMRA_FAW_O2")
.map(|v| v != "0")
.unwrap_or(true)
});
let f = self.func(if o2 {
"fa_prefill_w_bf16_g4o2"
} else {
"fa_prefill_w_bf16_g4"
});
let shmem = if o2 {
(2 * (4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M)) as u32
} else {
(2 * (2 * BK * head_dim + 4 * SP_M * head_dim + 4 * SP_M * BK) + 4 * (4 * SP_M))
as u32
};
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: ((t as u32).div_ceil(SP_M as u32), (n_head / 4) as u32, 1),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
t_kv as i32,
causal as i32,
window as i32,
);
let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&qb)
.arg(&kb)
.arg(&vb)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&wi);
unsafe {
b.launch(cfg)?;
}
return Ok(());
}
let f = self.func(if floor {
"fa_prefill_w_f32"
} else if f32_stage {
"fa_prefill_w_f32_pp"
} else {
"fa_prefill_w_bf16_pp"
});
let shmem =
(2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (
(t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
n_head as u32,
1,
),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
let (hd, nh, nhkv, ti, tkvi, cz, wi) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
t as i32,
t_kv as i32,
causal as i32,
window as i32,
);
if f32_stage {
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&wi);
unsafe {
b.launch(cfg)?;
}
} else {
let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&qb)
.arg(&kb)
.arg(&vb)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&wi);
unsafe {
b.launch(cfg)?;
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_prefill_hd512(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
) -> Result<(), Box<dyn std::error::Error>> {
if portable_mma_gated() {
return self.sdpa_naive(
q, k, v, o, head_dim, n_head, n_head_kv, t, t_kv, scale, causal,
);
}
static F32_STAGE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let f32_stage =
*F32_STAGE.get_or_init(|| std::env::var("MEMRA_FA512_STAGE").as_deref() == Ok("f32"));
static SP_ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let sp = !f32_stage
&& *SP_ON.get_or_init(|| {
std::env::var("MEMRA_FA512_SP")
.map(|v| v != "0")
.unwrap_or(true)
});
self.fa_prefill_hd512_arm(
q,
k,
v,
o,
head_dim,
n_head,
n_head_kv,
t,
t_kv,
scale,
causal,
f32_stage,
sp,
sp && fa_f16pv_on(),
)
}
#[allow(clippy::too_many_arguments)]
pub fn fa_prefill_hd512_pre(
&self,
qb: &CudaSlice<u8>,
kb: &CudaSlice<u8>,
vb: &CudaSlice<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
v_f16: bool,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert_eq!(head_dim, 512);
const SP_M: usize = 16;
const BKS: usize = 32;
let f16pv = fa_f16pv_on();
let nw = if f16pv { fa512_wide_warps() } else { 2 };
let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
debug_assert!(!v_f16 || f16pv, "f16 V emitted without the door on");
let mut vguard = self.fa_vf16_scratch.lock().unwrap();
let vref: &CudaSlice<u8> = if f16pv && !v_f16 {
let n = t_kv * n_head_kv * head_dim;
let need = n * 2;
if vguard.as_ref().map(|b| b.len() < need).unwrap_or(true) {
*vguard = Some(self.alloc_uninit::<u8>(need)?);
}
let dst = vguard.as_mut().unwrap();
self.bf16_to_f16_into(vb, n, dst)?;
vguard.as_ref().unwrap()
} else {
vb
};
let f = self.func(if hp {
"fa_prefill_bf16_hd512_sp16h2"
} else {
match (f16pv, nw) {
(true, 4) => "fa_prefill_bf16_hd512_sp16w4",
(true, _) => "fa_prefill_bf16_hd512_sp16",
_ => "fa_prefill_bf16_hd512_sp",
}
});
let (nwarp, npart) = if hp {
(4usize, 4usize)
} else if nw > 2 {
(nw, nw)
} else {
(2, 1)
};
let shmem = if hp {
(2 * (2 * BKS * head_dim + 2 * SP_M * BKS) + 4 * (2 * npart * SP_M * BKS + 2 * SP_M))
as u32
} else {
(2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
+ 4 * (npart * SP_M * BKS + SP_M)) as u32
};
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let grid_y = if hp {
(n_head / 2) as u32
} else {
n_head as u32
};
let cfg = LaunchConfig {
grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
block_dim: (32, nwarp as u32, 1),
shared_mem_bytes: shmem,
};
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,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qb)
.arg(kb)
.arg(vref)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_prefill_hd512_arm(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
f32_stage: bool,
sp: bool,
f16pv: bool,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert_eq!(head_dim, 512, "fa_prefill_hd512 is hd512 only");
if sp && !f32_stage {
const SP_M: usize = 16;
const BKS: usize = 32;
let nw = if f16pv { fa512_wide_warps() } else { 2 };
let hp = f16pv && fa512_hp_on() && n_head % 2 == 0 && (n_head / n_head_kv) % 2 == 0;
let f = self.func(if hp {
"fa_prefill_bf16_hd512_sp16h2"
} else {
match (f16pv, nw) {
(true, 4) => "fa_prefill_bf16_hd512_sp16w4",
(true, _) => "fa_prefill_bf16_hd512_sp16",
_ => "fa_prefill_bf16_hd512_sp",
}
});
let (nwarp, npart) = if hp {
(4usize, 4usize)
} else if nw > 2 {
(nw, nw)
} else {
(2, 1)
};
let shmem = if hp {
(2 * (2 * BKS * head_dim + 2 * SP_M * BKS)
+ 4 * (2 * npart * SP_M * BKS + 2 * SP_M)) as u32
} else {
(2 * (SP_M * head_dim + 2 * BKS * head_dim + SP_M * BKS)
+ 4 * (npart * SP_M * BKS + SP_M)) as u32
};
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let grid_y = if hp {
(n_head / 2) as u32
} else {
n_head as u32
};
let cfg = LaunchConfig {
grid_dim: ((t as u32).div_ceil(SP_M as u32), grid_y, 1),
block_dim: (32, nwarp as u32, 1),
shared_mem_bytes: shmem,
};
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,
);
let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
let vb = if f16pv {
self.f32_to_f16(v, t_kv * n_head_kv * head_dim)?
} else {
self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&qb)
.arg(&kb)
.arg(&vb)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz);
unsafe {
b.launch(cfg)?;
}
return Ok(());
}
const BLOCK_Q: usize = 32;
const BK: usize = 32;
const HALF: usize = 256;
let f = self.func(if f32_stage {
"fa_prefill_f32_hd512"
} else {
"fa_prefill_bf16_hd512"
});
let shmem = (2 * (BLOCK_Q * head_dim + BK * head_dim + BK * HALF + BLOCK_Q * BK)
+ 4 * BLOCK_Q) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (
(t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
n_head as u32,
2,
),
block_dim: (32, 2, 1),
shared_mem_bytes: shmem,
};
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,
);
if f32_stage {
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz);
unsafe {
b.launch(cfg)?;
}
} else {
let qb = self.f32_to_bf16(q, t * n_head * head_dim)?;
let kb = self.f32_to_bf16(k, t_kv * n_head_kv * head_dim)?;
let vb = self.f32_to_bf16(v, t_kv * n_head_kv * head_dim)?;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(&qb)
.arg(&kb)
.arg(&vb)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz);
unsafe {
b.launch(cfg)?;
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn rope_neox2_bf16e(
&self,
q: &mut CudaSlice<f32>,
k: &mut CudaSlice<f32>,
qb: &mut CudaSlice<u8>,
kb: &mut CudaSlice<u8>,
pos: &CudaSlice<i32>,
head_dim: usize,
n_dims: usize,
nh_q: usize,
nh_k: usize,
n_tokens: usize,
base: f32,
freq_scale: f32,
ff: Option<&CudaSlice<f32>>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("rope_neox2_bf16e_f32");
let rows = ((nh_q + nh_k) * n_tokens) as u32;
let cfg = LaunchConfig {
grid_dim: (rows, 1, 1),
block_dim: ((head_dim / 2) as u32, 1, 1),
shared_mem_bytes: 0,
};
let theta_scale = base.powf(-2.0 / n_dims as f32);
let (hd, nd, nhq, nhk, nt) = (
head_dim as i32,
n_dims as i32,
nh_q as i32,
nh_k as i32,
n_tokens as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
match ff {
Some(t) => {
b.arg(&mut *q)
.arg(&mut *k)
.arg(&mut *qb)
.arg(&mut *kb)
.arg(pos)
.arg(&hd)
.arg(&nd)
.arg(&nhq)
.arg(&nhk)
.arg(&nt)
.arg(&theta_scale)
.arg(&freq_scale)
.arg(t);
unsafe {
b.launch(cfg)?;
}
}
None => {
let null: u64 = 0;
b.arg(&mut *q)
.arg(&mut *k)
.arg(&mut *qb)
.arg(&mut *kb)
.arg(pos)
.arg(&hd)
.arg(&nd)
.arg(&nhq)
.arg(&nhk)
.arg(&nt)
.arg(&theta_scale)
.arg(&freq_scale)
.arg(&null);
unsafe {
b.launch(cfg)?;
}
}
}
Ok(())
}
pub fn f32_to_bf16(
&self,
x: &CudaSlice<f32>,
n: usize,
) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
assert!(n % 4 == 0, "f32_to_bf16 requires n % 4 == 0, got {n}");
let mut y = self.alloc_uninit::<u8>(n * 2)?;
let f = self.func("f32_to_bf16_flat");
let n_i = n as i64;
let cfg = LaunchConfig {
grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(&mut y).arg(&n_i);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn f32_to_f16(
&self,
x: &CudaSlice<f32>,
n: usize,
) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
assert!(n % 4 == 0, "f32_to_f16 requires n % 4 == 0, got {n}");
let mut y = self.alloc_uninit::<u8>(n * 2)?;
let f = self.func("f32_to_f16_flat");
let n_i = n as i64;
let cfg = LaunchConfig {
grid_dim: (((n / 4) as u32).div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(&mut y).arg(&n_i);
unsafe {
b.launch(cfg)?;
}
Ok(y)
}
pub fn bf16_to_f16(
&self,
xb: &CudaSlice<u8>,
n: usize,
) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
let mut y = self.alloc_uninit::<u8>(n * 2)?;
self.bf16_to_f16_into(xb, n, &mut y)?;
Ok(y)
}
pub fn bf16_to_f16_into(
&self,
xb: &CudaSlice<u8>,
n: usize,
y: &mut CudaSlice<u8>,
) -> Result<(), Box<dyn std::error::Error>> {
assert!(n % 2 == 0, "bf16_to_f16 requires n % 2 == 0, got {n}");
assert!(y.len() >= n * 2);
let f = self.func("bf16_to_f16_flat");
let n2 = (n / 2) as i64;
let cfg = LaunchConfig {
grid_dim: (((n / 2) as u32).div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(xb).arg(y).arg(&n2);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_prefill_vl8(
&self,
seqs: &[FaSeqVl],
head_dim: usize,
n_head: usize,
n_head_kv: usize,
scale: f32,
) -> Result<(), Box<dyn std::error::Error>> {
const BK: usize = 32;
let b = seqs.len();
assert!(b >= 1 && b <= 8);
let mut packed = [FaSeqVl::default(); 8];
packed[..b].copy_from_slice(seqs);
let v = FaVl8(packed);
let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
let ept = (n_head_kv * head_dim) as i32;
{
let f = self.func("fa_mirror_vl");
let max_n = (max_t as i64) * ept as i64;
let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
for which in 0..2i32 {
let cfg = LaunchConfig {
grid_dim: (blocks, 1, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&ept).arg(&which);
unsafe {
lb.launch(cfg)?;
}
}
}
let hd_sfx = fa_hd_suffix(head_dim)?;
let f = self.func(&format!("fa_prefill_bf16kv_vl{hd_sfx}"));
let block_q = 64usize;
let kv_stages = 2usize;
let shmem = (2 * (kv_stages * 2 * BK * head_dim + block_q * BK)
+ 4 * (block_q * BK + 2 * block_q)) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (max_t.div_ceil(block_q as u32), n_head as u32, b as u32),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&hd).arg(&nh).arg(&nhkv).arg(&scale);
unsafe {
lb.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn attn_pre_vl8(
&self,
seqs: &[AttnPreVl],
wq: &CudaSlice<f32>,
wk: &CudaSlice<f32>,
head_dim: usize,
rope_dims: usize,
n_head: usize,
n_head_kv: usize,
eps: f32,
freq_base: f32,
freq_scale: f32,
kv_dim_k: usize,
kv_dim_v: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let b = seqs.len();
assert!(b >= 1 && b <= 8);
let mut packed = [AttnPreVl::default(); 8];
packed[..b].copy_from_slice(seqs);
let v = AttnPreVl8(packed);
let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
{
let f = self.func("q_gate_split_vl");
let n = max_t * (n_head * head_dim) as u32;
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256), 1, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&hd).arg(&nh);
unsafe {
lb.launch(cfg)?;
}
}
{
let f = self.func("attn_rms_vl");
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,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v)
.arg(wq)
.arg(wk)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&eps);
unsafe {
lb.launch(cfg)?;
}
}
{
let f = self.func("attn_rope_vl");
let theta_scale = freq_base.powf(-2.0 / rope_dims as f32);
let nd = rope_dims as i32;
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,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v)
.arg(&hd)
.arg(&nd)
.arg(&nh)
.arg(&nhkv)
.arg(&theta_scale)
.arg(&freq_scale);
unsafe {
lb.launch(cfg)?;
}
}
{
let f = self.func("append_kv_vl");
let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk, max_t, b as u32),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&kdk).arg(&kdv).arg(&ktb).arg(&vtb);
unsafe {
lb.launch(cfg)?;
}
}
Ok(())
}
pub fn fa_prefill_view(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
if portable_mma_gated() {
return self.sdpa_naive_quantized_view(
q,
k,
v,
o,
head_dim,
n_head,
n_head_kv,
t,
t_kv,
scale,
causal,
k_tok_bytes,
v_tok_bytes,
);
}
const BLOCK_Q: usize = 64;
const BK: usize = 32;
let name = format!("fa_prefill_q{}", fa_hd_suffix(head_dim)?);
let f = if g {
self.func_g(&name)
} else {
self.func(&name)
};
let shmem =
(2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (
(t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
n_head as u32,
1,
),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
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,
);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_prefill_view_ws(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
if portable_mma_gated() {
return self.sdpa_naive_quantized_view(
q,
k,
v,
o,
head_dim,
n_head,
n_head_kv,
t,
t_kv,
scale,
causal,
k_tok_bytes,
v_tok_bytes,
);
}
const BLOCK_Q: usize = 64;
const BK: usize = 32;
let kv_dim_k = n_head_kv * head_dim;
let kv_dim_v = n_head_kv * head_dim;
let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
let mut guard = self.prime_deqw_ws.lock().unwrap();
let need_grow = match guard.as_ref() {
Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
None => true,
};
if need_grow {
let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
let (ck, cv) = guard
.as_ref()
.map(|(a, b)| (a.len(), b.len()))
.unwrap_or((0, 0));
*guard = Some((
self.alloc_u8(grow(ck, k_ws_bytes))?,
self.alloc_u8(grow(cv, v_ws_bytes))?,
));
}
let (kw, vw) = guard.as_mut().unwrap();
{
let f = if g {
self.func_g("fa_dequant_kv_ws_bf16")
} else {
self.func("fa_dequant_kv_ws_bf16")
};
let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk.max(1), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k)
.arg(v)
.arg(&mut *kw)
.arg(&mut *vw)
.arg(&kdk)
.arg(&kdv)
.arg(&tkvi)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
}
let db = std::env::var("MEMRA_PRIME_DEQW_DB")
.map(|v| v != "0")
.unwrap_or(true);
{
let hd_sfx = fa_hd_suffix(head_dim)?;
let f = self.func(&format!(
"fa_prefill_qw{}{hd_sfx}",
if db { "_db" } else { "" }
));
let shmem = if db {
(2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
} else {
(2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
};
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (
(t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
n_head as u32,
1,
),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
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,
);
let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(&*kw)
.arg(&*vw)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&kdk)
.arg(&kdv);
unsafe {
b.launch(cfg)?;
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_prefill_view_ws_w_hd128(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t: usize,
t_kv: usize,
scale: f32,
causal: bool,
window: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
) -> Result<(), Box<dyn std::error::Error>> {
assert_eq!(
head_dim, 128,
"fa_prefill_view_ws_w_hd128: only the hd128 twin is stamped"
);
if portable_mma_gated() {
return self.sdpa_naive_w_quantized_view(
q,
k,
v,
o,
head_dim,
n_head,
n_head_kv,
t,
t_kv,
scale,
causal,
window,
k_tok_bytes,
v_tok_bytes,
);
}
const BLOCK_Q: usize = 64;
const BK: usize = 32;
let kv_dim_k = n_head_kv * head_dim;
let kv_dim_v = n_head_kv * head_dim;
let k_ws_bytes = t_kv * kv_dim_k * 2; let v_ws_bytes = t_kv * kv_dim_v * 2;
let mut guard = self.prime_deqw_ws.lock().unwrap();
let need_grow = match guard.as_ref() {
Some((kw, vw)) => kw.len() < k_ws_bytes || vw.len() < v_ws_bytes,
None => true,
};
if need_grow {
let grow = |cur: usize, need: usize| if cur >= need { cur } else { need };
let (ck, cv) = guard
.as_ref()
.map(|(a, b)| (a.len(), b.len()))
.unwrap_or((0, 0));
*guard = Some((
self.alloc_u8(grow(ck, k_ws_bytes))?,
self.alloc_u8(grow(cv, v_ws_bytes))?,
));
}
let (kw, vw) = guard.as_mut().unwrap();
{
let f = self.func("fa_dequant_kv_ws_bf16");
let total = (t_kv * (kv_dim_k + kv_dim_v)) as u64;
let nblk = ((total + 255) / 256).min(65535 * 16) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk.max(1), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (kdk, kdv, tkvi) = (kv_dim_k as i32, kv_dim_v as i32, t_kv as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k)
.arg(v)
.arg(&mut *kw)
.arg(&mut *vw)
.arg(&kdk)
.arg(&kdv)
.arg(&tkvi)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
}
let db = std::env::var("MEMRA_PRIME_DEQW_DB")
.map(|v| v != "0")
.unwrap_or(true);
{
let f = self.func(if db {
"fa_prefill_qw_db_w_hd128"
} else {
"fa_prefill_qw_w_hd128"
});
let shmem = if db {
(2 * (4 * BK * head_dim + BLOCK_Q * BK) + 4 * BLOCK_Q) as u32
} else {
(2 * (2 * BK * head_dim + BLOCK_Q * BK) + 4 * (BLOCK_Q * BK + 2 * BLOCK_Q)) as u32
};
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (
(t as u32 + BLOCK_Q as u32 - 1) / BLOCK_Q as u32,
n_head as u32,
1,
),
block_dim: (32, 4, 1),
shared_mem_bytes: shmem,
};
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,
);
let (kdk, kdv, wnd) = (kv_dim_k as i32, kv_dim_v as i32, window as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(&*kw)
.arg(&*vw)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&ti)
.arg(&tkvi)
.arg(&scale)
.arg(&cz)
.arg(&kdk)
.arg(&kdv)
.arg(&wnd);
unsafe {
b.launch(cfg)?;
}
}
Ok(())
}
pub fn fa_decode(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t_kv: usize,
scale: f32,
k_tok_bytes: usize,
v_tok_bytes: usize,
) -> Result<(), Box<dyn std::error::Error>> {
self.fa_decode_kvmod(
q,
k,
v,
o,
head_dim,
n_head,
n_head_kv,
t_kv,
scale,
k_tok_bytes,
v_tok_bytes,
false,
)
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
fn fa_decode_scalar_unified(
&self,
q: &cudarc::driver::CudaView<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut cudarc::driver::CudaViewMut<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t_kv_host: usize,
t_kv_dev: Option<&CudaSlice<i32>>,
scale: f32,
n_splits: usize,
split_keys: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
part_o: &mut CudaSlice<f32>,
part_m: &mut CudaSlice<f32>,
part_l: &mut CudaSlice<f32>,
q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
) -> Result<(), Box<dyn std::error::Error>> {
let f = if g {
self.func_g("fa_decode_f32")
} else {
self.fa_func("fa_decode_f32", head_dim)
};
let cfg = LaunchConfig {
grid_dim: (n_head as u32, n_splits as u32, 1),
block_dim: (head_dim as u32, 1, 1),
shared_mem_bytes: (4 * (head_dim + 32)) as u32,
};
let (hd, nh, nhkv, nsp) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
n_splits as i32,
);
let (ktb, vtb, tkvi, ski) = (
k_tok_bytes as i64,
v_tok_bytes as i64,
t_kv_host as i32,
split_keys as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
match t_kv_dev {
Some(d) => {
b.arg(q)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&tkvi)
.arg(d)
.arg(&scale)
.arg(&nsp)
.arg(&ski)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
}
None => {
let null: u64 = 0;
b.arg(q)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&tkvi)
.arg(&null)
.arg(&scale)
.arg(&nsp)
.arg(&ski)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
}
}
let cfg2 = LaunchConfig {
grid_dim: (n_head as u32, 1, 1),
block_dim: (head_dim as u32, 1, 1),
shared_mem_bytes: 0,
};
if let Some((oq, od)) = q8_out {
let fc = if g {
self.func_g("fa_decode_combine_q8_1")
} else {
self.fa_func("fa_decode_combine_q8_1", head_dim)
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(oq)
.arg(od)
.arg(&hd)
.arg(&nh)
.arg(&nsp);
unsafe {
b2.launch(cfg2)?;
}
return Ok(());
}
let fc = if g {
self.func_g("fa_decode_combine_f32")
} else {
self.fa_func("fa_decode_combine_f32", head_dim)
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nsp);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
pub fn fa_decode_kvmod(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t_kv: usize,
scale: f32,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let q_view = q.as_view();
let mut o_view = o.as_view_mut();
self.fa_decode_kvmod_view(
&q_view,
k,
v,
&mut o_view,
head_dim,
n_head,
n_head_kv,
t_kv,
scale,
k_tok_bytes,
v_tok_bytes,
g,
)
}
#[allow(clippy::too_many_arguments)]
pub fn fa_decode_kvmod_view(
&self,
q: &cudarc::driver::CudaView<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut cudarc::driver::CudaViewMut<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t_kv: usize,
scale: f32,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let mut fa_vec = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
if g && head_dim == 256 && !fa_v4_at(t_kv) {
fa_vec = false;
}
let sp = fa_split_keys(t_kv, n_head_kv);
let n_splits = if fa_vec {
((t_kv + sp - 1) / sp).max(1)
} else {
((t_kv + 255) / 256).max(1)
};
let o_len = n_head * n_splits * head_dim;
let ml_len = n_head * n_splits;
let mut part_guard = self.fa_part_pool.lock().unwrap();
if part_guard
.as_ref()
.map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
.unwrap_or(true)
{
let old = part_guard.take();
let (co, cm) = old
.as_ref()
.map(|pp| (pp.0.len(), pp.1.len()))
.unwrap_or((0, 0));
if let Some(old) = old {
self.fa_part_retired.lock().unwrap().push(old);
}
if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
eprintln!(
"[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
co, o_len, cm, ml_len
);
}
*part_guard = Some((
self.alloc_uninit::<f32>(o_len.max(2 * co))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
));
}
let pg = part_guard.as_mut().unwrap();
self.gpu
.stream()
.memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
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,
);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
let fa512_min = fa512_min_tkv();
let deep = fa_vec
&& head_dim == 256
&& fa_v4_at(t_kv)
&& !g
&& fa_deep_at(t_kv)
&& !matches!(fa_v4_mode(), "noB3" | "stage");
let (f, cfg) = if fa_vec && head_dim == 512 && t_kv >= fa512_min {
let gqa = (n_head / n_head_kv).max(1) as u32;
let fv = self.fa_func("fa_decode_vec_q_dpl16", head_dim);
(
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: 0,
},
)
} else if fa_vec && head_dim <= 256 {
let gqa = (n_head / n_head_kv).max(1) as u32;
static SMEM_TKV: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
let smem_tkv = *SMEM_TKV.get_or_init(|| {
std::env::var("MEMRA_FA_SMEM_TKV")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or_else(|| {
FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
})
});
if fa_v4_at(t_kv) && head_dim == 256 {
let v4name = match fa_v4_mode() {
"noB3" => "fa_decode_vec_q_v4_noB3", "stage" => "fa_decode_vec_q_v4_stage", _ if deep => "fa_decode_vec_q_v4_deep",
_ => "fa_decode_vec_q_v4",
};
let fv = if g {
self.func_g(v4name)
} else {
self.func(v4name)
};
let shmem = (if deep { 12160 } else { 11520 }
+ 32 * head_dim * if g { 1 } else { 2 }) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
fv.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
(
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
},
)
} else if fa_v3_active(head_dim) {
let fv = if g {
self.func_g("fa_decode_vec_q_v3")
} else {
self.func("fa_decode_vec_q_v3")
};
let shmem = (32 * head_dim * 2) as u32; (
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
},
)
} else if fa_v2_on() {
let fv = if g {
self.func_g("fa_decode_vec_q_v2")
} else {
self.func("fa_decode_vec_q_v2")
};
let shmem = (2 * 32 * head_dim * 2) as u32; (
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
},
)
} else if smem_tkv > 0 && t_kv >= smem_tkv && !g && !(head_dim == 512 && Self::gkv_on())
{
let fv = if g {
self.func_g("fa_decode_vec_q_smem")
} else {
self.func("fa_decode_vec_q_smem")
};
let shmem = (2 * 32 * head_dim * 2) as u32; use cudarc::driver::sys::CUfunction_attribute_enum as A;
fv.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
(
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
},
)
} else {
let fv = if g {
self.func_g("fa_decode_vec_q")
} else {
self.func("fa_decode_vec_q")
};
(
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: 0,
},
)
}
} else {
return self.fa_decode_scalar_unified(
q,
k,
v,
o,
head_dim,
n_head,
n_head_kv,
t_kv,
None,
scale,
n_splits,
if fa_vec { sp } else { 256 },
k_tok_bytes,
v_tok_bytes,
g,
part_o,
part_m,
part_l,
None,
);
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&tkvi)
.arg(&scale)
.arg(&nsp)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
let (fc, cfg2) = (
if g {
self.func_g("fa_decode_combine_f32")
} else {
self.fa_func("fa_decode_combine_f32", head_dim)
},
LaunchConfig {
grid_dim: (n_head as u32, 1, 1),
block_dim: (head_dim as u32, 1, 1),
shared_mem_bytes: 0,
},
);
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nsp);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_decode_batch_seqs_v4(
&self,
q: &CudaSlice<f32>,
kv_ptrs: &cudarc::driver::CudaView<u64>,
pos_seq: &CudaSlice<i32>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
b_n: usize,
t_kv_max: usize,
scale: f32,
split_keys: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert!(head_dim == 256, "seqs twin is v4-stamped (hd256 only)");
let n_splits_max = (t_kv_max + split_keys - 1) / split_keys;
let o_len = b_n * n_head * n_splits_max * head_dim;
let ml_len = b_n * n_head * n_splits_max;
let mut part_guard = self.fa_part_pool.lock().unwrap();
if part_guard
.as_ref()
.map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
.unwrap_or(true)
{
let old = part_guard.take();
let (co, cm) = old
.as_ref()
.map(|pp| (pp.0.len(), pp.1.len()))
.unwrap_or((0, 0));
if let Some(old) = old {
self.fa_part_retired.lock().unwrap().push(old);
}
if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
eprintln!(
"[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
co, o_len, cm, ml_len
);
}
*part_guard = Some((
self.alloc_uninit::<f32>(o_len.max(2 * co))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
));
}
let pg = part_guard.as_mut().unwrap();
self.gpu
.stream()
.memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
let (nspm, spk) = (n_splits_max as i32, split_keys as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let gqa = (n_head / n_head_kv).max(1) as u32;
let f = self.func("fa_decode_vec_q_seqs_v4");
let shmem = (11520 + 32 * head_dim * 2) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits_max as u32, b_n as u32),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
};
{
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(kv_ptrs)
.arg(pos_seq)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&scale)
.arg(&nspm)
.arg(&spk)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
}
let fc = self.func("fa_decode_combine_seqs");
let cfg2 = LaunchConfig {
grid_dim: (n_head as u32, b_n as u32, 1),
block_dim: (head_dim as u32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(pos_seq)
.arg(&nspm)
.arg(&spk);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn append_kv_quantized_seqs(
&self,
k_rows: &CudaSlice<f32>,
v_rows: &CudaSlice<f32>,
kv_ptrs: &cudarc::driver::CudaView<u64>,
pos_seq: &CudaSlice<i32>,
b_n: usize,
kv_dim_k: usize,
kv_dim_v: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("append_quantize_kv_q8_0_q5_1_seqs");
let nblk = (kv_dim_k.max(kv_dim_v) / 32) as u32;
let cfg = LaunchConfig {
grid_dim: (nblk, b_n as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let (kdk, kdv) = (kv_dim_k as i32, kv_dim_v as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k_rows)
.arg(v_rows)
.arg(kv_ptrs)
.arg(pos_seq)
.arg(&kdk)
.arg(&kdv)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn fa_rows_eligible(&self, base_len: usize, head_dim: usize) -> bool {
std::env::var("MEMRA_NO_FA_VEC").is_err()
&& std::env::var("MEMRA_FA_ROWS_OFF").is_err()
&& base_len + 1 >= fa_vec_min_tkv()
&& head_dim <= 256
&& head_dim % 32 == 0
}
#[allow(clippy::too_many_arguments)]
pub fn fa_decode_rows(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
base_len: usize,
t: usize,
scale: f32,
k_tok_bytes: usize,
v_tok_bytes: usize,
base_dev: Option<(&CudaSlice<i32>, i32)>,
kv_shared: bool,
g: bool,
mut q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert!(base_len + 1 >= fa_vec_min_tkv() && head_dim <= 512 && head_dim % 32 == 0);
let t_kv_max = base_len + t; let mut sp = fa_split_keys(t_kv_max, n_head_kv); if head_dim == 512 {
static SP512: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
let v = *SP512.get_or_init(|| {
std::env::var("MEMRA_FA_SP512")
.ok()
.and_then(|x| x.parse().ok())
.unwrap_or(0)
});
sp = if v >= 8 {
v
} else {
FA_SP512_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
};
}
let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let gqa = (n_head / n_head_kv).max(1) as u32;
let mut groups: Vec<(usize, usize, usize)> = Vec::new(); if head_dim == 512 || fa_split_keys(base_len + 1, n_head_kv) == sp {
groups.push((0, t, sp));
} else {
let mut r0 = 0usize;
while r0 < t {
let sp_g = fa_split_keys(base_len + r0 + 1, n_head_kv);
let mut r1 = r0 + 1;
while r1 < t && fa_split_keys(base_len + r1 + 1, n_head_kv) == sp_g {
r1 += 1;
}
groups.push((r0, r1 - r0, sp_g));
r0 = r1;
}
}
static SMEM_TKV_R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
let smem_tkv = *SMEM_TKV_R.get_or_init(|| {
std::env::var("MEMRA_FA_SMEM_TKV")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
});
let v4 = fa_v4_at(base_len + t) && head_dim == 256;
let v3 = fa_v3_active(head_dim);
let smem_rows =
head_dim <= 256 && !v3 && !fa_v2_on() && smem_tkv > 0 && t_kv_max >= smem_tkv;
let _ = kv_shared;
let i2 = head_dim == 512 && std::env::var("MEMRA_FA_I2").as_deref() != Ok("0");
static TB512: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let tb512 = head_dim == 512
&& sp <= 32
&& n_head / n_head_kv.max(1) <= 16
&& *TB512.get_or_init(|| std::env::var("MEMRA_FA_TB512").as_deref() != Ok("0"));
let fname = if tb512 {
"fa_decode_vec_q_rows_v4_512_tb"
} else if i2 {
"fa_decode_vec_q_rows_dpl16_i2"
} else if head_dim == 512 {
"fa_decode_vec_q_rows_dpl16"
}
else if v4 {
"fa_decode_vec_q_rows_v4"
} else if v3 {
"fa_decode_vec_q_rows_v3"
} else if fa_v2_on() {
"fa_decode_vec_q_rows_v2"
} else if smem_rows {
"fa_decode_vec_q_rows_smem"
} else {
"fa_decode_vec_q_rows"
};
let f = if head_dim == 512 {
self.fa_func(fname, head_dim)
} else if g {
self.func_g(if smem_rows {
"fa_decode_vec_q_rows"
} else {
fname
})
} else {
self.func(fname)
};
let shmem = if tb512 {
let gk = Self::gkv_on();
let sh =
(8192 + 1024 + 32 * 512 + 32 * 64 + 32 * head_dim * if gk { 1 } else { 2 }) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
sh as i32,
)?;
sh
} else if v4 || v3 || smem_rows || fa_v2_on() {
let sh = (if v4 {
11520 + 32 * head_dim * if g { 1 } else { 2 }
} else if v3 {
32 * head_dim * 2
} else {
2 * 32 * head_dim * 2
}) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
sh as i32,
)?;
sh
} else {
0
};
for &(r0, t_g, sp_g) in &groups {
let n_splits_g = (base_len + r0 + t_g).div_ceil(sp_g);
let (nspm, spk) = (n_splits_g as i32, sp_g as i32);
let base_i = (base_len + r0) as i32;
let o_len = t_g * n_head * n_splits_g * head_dim;
let ml_len = t_g * n_head * n_splits_g;
let mut part_guard = self.fa_part_pool.lock().unwrap();
if part_guard
.as_ref()
.map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
.unwrap_or(true)
{
let old = part_guard.take();
let (co, cm) = old
.as_ref()
.map(|pp| (pp.0.len(), pp.1.len()))
.unwrap_or((0, 0));
if let Some(old) = old {
self.fa_part_retired.lock().unwrap().push(old);
}
if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
eprintln!(
"[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
co, o_len, cm, ml_len
);
}
*part_guard = Some((
self.alloc_uninit::<f32>(o_len.max(2 * co))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
));
}
let pg = part_guard.as_mut().unwrap();
self.gpu
.stream()
.memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
let (part_o, part_m, part_l) = (&mut *part_o, &mut *part_m, &mut *part_l);
let qv = self.view(q, t * n_head * head_dim);
let q_g = qv.slice(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
let cfg = LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits_g as u32, t_g as u32),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
};
{
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
if tb512 {
let (bd, plus) =
base_dev.expect("hd512 rows twin requires a device base counter");
let plus_g = plus + r0 as i32;
let nr = t_g as i32;
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pq, _b0) = q_g.device_ptr(s);
let (pk, _b1) = k.device_ptr(s);
let (pv, _b2) = v.device_ptr(s);
let (po, _b3) = part_o.device_ptr_mut(s);
let (pm, _b4) = part_m.device_ptr_mut(s);
let (pl, _b5) = part_l.device_ptr_mut(s);
let (pb, _b6) = bd.device_ptr(s);
let mut ps = [
&pq as *const _ as *mut std::ffi::c_void,
&pk as *const _ as *mut _,
&pv as *const _ as *mut _,
&po as *const _ as *mut _,
&pm as *const _ as *mut _,
&pl as *const _ as *mut _,
&hd as *const _ as *mut _,
&nh as *const _ as *mut _,
&nhkv as *const _ as *mut _,
&pb as *const _ as *mut _,
&plus_g as *const _ as *mut _,
&scale as *const _ as *mut _,
&nspm as *const _ as *mut _,
&spk as *const _ as *mut _,
&ktb as *const _ as *mut _,
&vtb as *const _ as *mut _,
&nr as *const _ as *mut _,
];
unsafe {
self.launch_pdl_flash(
Self::gkv_on(),
"fa_decode_vec_q_rows_v4_512_tb",
(n_head_kv as u32, n_splits_g as u32, 1),
(32, gqa, 1),
shmem,
&mut ps,
)?;
}
} else {
let cfg_tb = LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits_g as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
};
b.arg(&q_g)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(bd)
.arg(&plus_g)
.arg(&scale)
.arg(&nspm)
.arg(&spk)
.arg(&ktb)
.arg(&vtb)
.arg(&nr);
unsafe {
b.launch(cfg_tb)?;
}
}
} else if head_dim == 512 {
let (bd, plus) =
base_dev.expect("hd512 rows twin requires a device base counter");
let plus_g = plus + r0 as i32;
b.arg(&q_g)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(bd)
.arg(&plus_g)
.arg(&scale)
.arg(&nspm)
.arg(&spk)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
} else {
b.arg(&q_g)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(&base_i)
.arg(&scale)
.arg(&nspm)
.arg(&spk)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
}
}
let cfg2 = LaunchConfig {
grid_dim: (n_head as u32, t_g as u32, 1),
block_dim: (head_dim as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut o_g = o.slice_mut(r0 * n_head * head_dim..(r0 + t_g) * n_head * head_dim);
if head_dim == 512 {
let (bd, plus) = base_dev.unwrap();
let plus_g = plus + r0 as i32;
if let Some((oq, od)) = q8_out.as_mut() {
debug_assert!(t == 1, "rows q8 emit is a t=1 decode arm");
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (po, _g0) = part_o.device_ptr(s);
let (pm, _g1) = part_m.device_ptr(s);
let (pl, _g2) = part_l.device_ptr(s);
let (pq, _g3) = oq.device_ptr_mut(s);
let (pd, _g4) = od.device_ptr_mut(s);
let (pb, _g5) = bd.device_ptr(s);
let mut ps = [
&po as *const _ as *mut std::ffi::c_void,
&pm as *const _ as *mut _,
&pl as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&hd as *const _ as *mut _,
&nh as *const _ as *mut _,
&pb as *const _ as *mut _,
&plus_g as *const _ as *mut _,
&nspm as *const _ as *mut _,
&spk as *const _ as *mut _,
];
unsafe {
self.launch_pdl_flash(
Self::gkv_on(),
"fa_decode_combine_rows_dc_q8_1",
cfg2.grid_dim,
cfg2.block_dim,
0,
&mut ps,
)?;
}
continue;
}
let fc = self.fa_func("fa_decode_combine_rows_dc_q8_1", head_dim);
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(&mut **oq)
.arg(&mut **od)
.arg(&hd)
.arg(&nh)
.arg(bd)
.arg(&plus_g)
.arg(&nspm)
.arg(&spk);
unsafe {
b2.launch(cfg2)?;
}
continue;
}
let fc = self.fa_func("fa_decode_combine_rows_dc", head_dim);
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(&mut o_g)
.arg(&hd)
.arg(&nh)
.arg(bd)
.arg(&plus_g)
.arg(&nspm)
.arg(&spk);
unsafe {
b2.launch(cfg2)?;
}
} else {
assert!(
q8_out.is_none(),
"rows q8 emit requires the hd512 dc combine"
);
let fc = self.func("fa_decode_combine_rows");
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(&mut o_g)
.arg(&hd)
.arg(&nh)
.arg(&base_i)
.arg(&nspm)
.arg(&spk);
unsafe {
b2.launch(cfg2)?;
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_decode_rows_w(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
base_dev: &CudaSlice<i32>,
base_plus: i32,
t: usize,
scale: f32,
window: usize,
k_tok_bytes: usize,
v_tok_bytes: usize,
q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert!(head_dim == 256);
let sp = {
static SPW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
let v = *SPW.get_or_init(|| {
std::env::var("MEMRA_FA_SPW")
.ok()
.and_then(|x| x.parse().ok())
.unwrap_or(0)
});
if v >= 8 {
v
} else {
FA_SPW_DEFAULT.load(std::sync::atomic::Ordering::Relaxed)
}
};
let n_splits_max = (window + sp - 1) / sp;
let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
let (nspm, spk, wini) = (n_splits_max as i32, sp as i32, window as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let gqa = (n_head / n_head_kv).max(1) as u32;
let o_len = t * n_head * n_splits_max * head_dim;
let ml_len = t * n_head * n_splits_max;
let mut part_guard = self.fa_part_pool.lock().unwrap();
if part_guard
.as_ref()
.map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
.unwrap_or(true)
{
let old = part_guard.take();
let (co, cm) = old
.as_ref()
.map(|pp| (pp.0.len(), pp.1.len()))
.unwrap_or((0, 0));
if let Some(old) = old {
self.fa_part_retired.lock().unwrap().push(old);
}
if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
eprintln!(
"[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
co, o_len, cm, ml_len
);
}
*part_guard = Some((
self.alloc_uninit::<f32>(o_len.max(2 * co))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
));
}
let pg = part_guard.as_mut().unwrap();
self.gpu
.stream()
.memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
static SMEM_TKV_W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
let smem_tkv = *SMEM_TKV_W.get_or_init(|| {
std::env::var("MEMRA_FA_SMEM_TKV")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or_else(|| FA_SMEM_TKV_DEFAULT.load(std::sync::atomic::Ordering::Relaxed))
});
use cudarc::driver::sys::CUfunction_attribute_enum as A;
let wg = Self::wkv_on();
let sp2 =
gqa <= 4 && fa_v4_at(window) && std::env::var("MEMRA_FA_SPW2").as_deref() != Ok("0");
if sp2 {
let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pq, _b0) = q.device_ptr(s);
let (pk, _b1) = k.device_ptr(s);
let (pv, _b2) = v.device_ptr(s);
let (po, _b3) = part_o.device_ptr_mut(s);
let (pm, _b4) = part_m.device_ptr_mut(s);
let (pl, _b5) = part_l.device_ptr_mut(s);
let (pb, _b6) = base_dev.device_ptr(s);
let mut ps = [
&pq as *const _ as *mut std::ffi::c_void,
&pk as *const _ as *mut _,
&pv as *const _ as *mut _,
&po as *const _ as *mut _,
&pm as *const _ as *mut _,
&pl as *const _ as *mut _,
&hd as *const _ as *mut _,
&nh as *const _ as *mut _,
&nhkv as *const _ as *mut _,
&pb as *const _ as *mut _,
&base_plus as *const _ as *mut _,
&scale as *const _ as *mut _,
&nspm as *const _ as *mut _,
&spk as *const _ as *mut _,
&ktb as *const _ as *mut _,
&vtb as *const _ as *mut _,
&wini as *const _ as *mut _,
];
unsafe {
self.launch_pdl_flash(
wg,
"fa_decode_vec_q_rows_v4_w_sp",
(n_head_kv as u32, n_splits_max as u32, t as u32),
(32, gqa + 1, 1),
sh,
&mut ps,
)?;
}
} else {
let f = if wg {
self.func_g("fa_decode_vec_q_rows_v4_w_sp")
} else {
self.func("fa_decode_vec_q_rows_v4_w_sp")
};
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
sh as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
block_dim: (32, gqa + 1, 1),
shared_mem_bytes: sh,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(base_dev)
.arg(&base_plus)
.arg(&scale)
.arg(&nspm)
.arg(&spk)
.arg(&ktb)
.arg(&vtb)
.arg(&wini);
unsafe {
b.launch(cfg)?;
}
}
} else {
if fa_v4_at(window) && Self::pdl_on() && Self::pdl_wb_on() {
let sh = (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32;
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (pq, _b0) = q.device_ptr(s);
let (pk, _b1) = k.device_ptr(s);
let (pv, _b2) = v.device_ptr(s);
let (po, _b3) = part_o.device_ptr_mut(s);
let (pm, _b4) = part_m.device_ptr_mut(s);
let (pl, _b5) = part_l.device_ptr_mut(s);
let (pb, _b6) = base_dev.device_ptr(s);
let mut ps = [
&pq as *const _ as *mut std::ffi::c_void,
&pk as *const _ as *mut _,
&pv as *const _ as *mut _,
&po as *const _ as *mut _,
&pm as *const _ as *mut _,
&pl as *const _ as *mut _,
&hd as *const _ as *mut _,
&nh as *const _ as *mut _,
&nhkv as *const _ as *mut _,
&pb as *const _ as *mut _,
&base_plus as *const _ as *mut _,
&scale as *const _ as *mut _,
&nspm as *const _ as *mut _,
&spk as *const _ as *mut _,
&ktb as *const _ as *mut _,
&vtb as *const _ as *mut _,
&wini as *const _ as *mut _,
];
unsafe {
self.launch_pdl_flash(
wg,
"fa_decode_vec_q_rows_v4_w",
(n_head_kv as u32, n_splits_max as u32, t as u32),
(32, gqa, 1),
sh,
&mut ps,
)?;
}
} else {
let pick = |name: &str| {
if wg {
self.func_g(name)
} else {
self.func(name)
}
};
let (f, sh) = if fa_v4_at(window) {
let f = pick("fa_decode_vec_q_rows_v4_w");
(f, (11520 + 32 * head_dim * if wg { 1 } else { 2 }) as u32)
} else if smem_tkv > 0 && window >= smem_tkv {
(
pick("fa_decode_vec_q_rows_smem_w"),
(2 * 32 * head_dim * 2) as u32,
)
} else {
(pick("fa_decode_vec_q_rows_reg_w"), 0u32)
};
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
sh as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
block_dim: (32, gqa, 1),
shared_mem_bytes: sh,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(base_dev)
.arg(&base_plus)
.arg(&scale)
.arg(&nspm)
.arg(&spk)
.arg(&ktb)
.arg(&vtb)
.arg(&wini);
unsafe {
b.launch(cfg)?;
}
}
}
let cfg2 = LaunchConfig {
grid_dim: (n_head as u32, t as u32, 1),
block_dim: (head_dim as u32, 1, 1),
shared_mem_bytes: 0,
};
if let Some((oq, od)) = q8_out {
if Self::pdl_on() && Self::pdl_wb_on() {
use cudarc::driver::{DevicePtr, DevicePtrMut};
let s = &self.gpu.stream();
let (po, _g0) = part_o.device_ptr(s);
let (pm, _g1) = part_m.device_ptr(s);
let (pl, _g2) = part_l.device_ptr(s);
let (pq, _g3) = oq.device_ptr_mut(s);
let (pd, _g4) = od.device_ptr_mut(s);
let mut ps = [
&po as *const _ as *mut std::ffi::c_void,
&pm as *const _ as *mut _,
&pl as *const _ as *mut _,
&pq as *const _ as *mut _,
&pd as *const _ as *mut _,
&hd as *const _ as *mut _,
&nh as *const _ as *mut _,
&nspm as *const _ as *mut _,
&spk as *const _ as *mut _,
&wini as *const _ as *mut _,
];
unsafe {
self.launch_pdl_flash(
wg,
"fa_decode_combine_rows_w_q8_1",
cfg2.grid_dim,
cfg2.block_dim,
0,
&mut ps,
)?;
}
return Ok(());
}
let fc = if wg {
self.func_g("fa_decode_combine_rows_w_q8_1")
} else {
self.func("fa_decode_combine_rows_w_q8_1")
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(oq)
.arg(od)
.arg(&hd)
.arg(&nh)
.arg(&nspm)
.arg(&spk)
.arg(&wini);
unsafe {
b2.launch(cfg2)?;
}
return Ok(());
}
let fc = if wg {
self.func_g("fa_decode_combine_rows_w")
} else {
self.func("fa_decode_combine_rows_w")
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nspm)
.arg(&spk)
.arg(&wini);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn fa_decode_rows_dc(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
base_dev: &CudaSlice<i32>,
t_kv_upper: usize,
t: usize,
scale: f32,
k_tok_bytes: usize,
v_tok_bytes: usize,
base_plus: i32,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let v4 = head_dim == 256 && fa_v4_at(t_kv_upper);
assert!(
v4 || fa_v3_active(head_dim),
"stream fa rows requires the v3 or v4 lane"
);
assert!(v4 || base_plus == 0, "v3_dc kernel takes no plus arg");
if v4 {
let sp = fa_split_keys(t_kv_upper, n_head_kv);
let n_splits_max = (t_kv_upper + sp - 1) / sp;
let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
let (nspm, spk) = (n_splits_max as i32, sp as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let gqa = (n_head / n_head_kv).max(1) as u32;
let o_len = t * n_head * n_splits_max * head_dim;
let ml_len = t * n_head * n_splits_max;
let mut part_guard = self.fa_part_pool.lock().unwrap();
if part_guard
.as_ref()
.map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
.unwrap_or(true)
{
let old = part_guard.take();
let (co, cm) = old
.as_ref()
.map(|pp| (pp.0.len(), pp.1.len()))
.unwrap_or((0, 0));
if let Some(old) = old {
self.fa_part_retired.lock().unwrap().push(old);
}
if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
eprintln!(
"[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
co, o_len, cm, ml_len
);
}
*part_guard = Some((
self.alloc_uninit::<f32>(o_len.max(2 * co))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
));
}
let pg = part_guard.as_mut().unwrap();
self.gpu
.stream()
.memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
let f = if g {
self.func_g("fa_decode_vec_q_rows_v4_dc")
} else {
self.func("fa_decode_vec_q_rows_v4_dc")
};
let sh = (11520 + 32 * head_dim * if g { 1 } else { 2 }) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
sh as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
block_dim: (32, gqa, 1),
shared_mem_bytes: sh,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(base_dev)
.arg(&base_plus)
.arg(&scale)
.arg(&nspm)
.arg(&spk)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
let fc = self.func("fa_decode_combine_rows_dc");
let cfg2 = LaunchConfig {
grid_dim: (n_head as u32, t as u32, 1),
block_dim: (head_dim as u32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(base_dev)
.arg(&base_plus)
.arg(&nspm)
.arg(&spk);
unsafe {
b2.launch(cfg2)?;
}
return Ok(());
}
let sp = fa_split_keys(t_kv_upper, n_head_kv);
let n_splits_max = (t_kv_upper + sp - 1) / sp;
let (hd, nh, nhkv) = (head_dim as i32, n_head as i32, n_head_kv as i32);
let (nspm, spk) = (n_splits_max as i32, sp as i32);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let gqa = (n_head / n_head_kv).max(1) as u32;
let o_len = t * n_head * n_splits_max * head_dim;
let ml_len = t * n_head * n_splits_max;
let mut part_guard = self.fa_part_pool.lock().unwrap();
if part_guard
.as_ref()
.map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
.unwrap_or(true)
{
let old = part_guard.take();
let (co, cm) = old
.as_ref()
.map(|pp| (pp.0.len(), pp.1.len()))
.unwrap_or((0, 0));
if let Some(old) = old {
self.fa_part_retired.lock().unwrap().push(old);
}
if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
eprintln!(
"[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
co, o_len, cm, ml_len
);
}
*part_guard = Some((
self.alloc_uninit::<f32>(o_len.max(2 * co))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
));
}
let pg = part_guard.as_mut().unwrap();
self.gpu
.stream()
.memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
let f = self.func("fa_decode_vec_q_rows_v3_dc");
let sh = (32 * head_dim * 2) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
f.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
sh as i32,
)?;
let cfg = LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits_max as u32, t as u32),
block_dim: (32, gqa, 1),
shared_mem_bytes: sh,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(base_dev)
.arg(&scale)
.arg(&nspm)
.arg(&spk)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
let fc = self.func("fa_decode_combine_rows_dc");
let cfg2 = LaunchConfig {
grid_dim: (n_head as u32, t as u32, 1),
block_dim: (head_dim as u32, 1, 1),
shared_mem_bytes: 0,
};
let plus0 = 0i32;
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(base_dev)
.arg(&plus0)
.arg(&nspm)
.arg(&spk);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
pub fn fa_decode_dc(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t_kv_dev: &CudaSlice<i32>,
bucket_max: usize,
scale: f32,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
) -> Result<(), Box<dyn std::error::Error>> {
self.fa_decode_dc_q8(
q,
k,
v,
o,
head_dim,
n_head,
n_head_kv,
t_kv_dev,
bucket_max,
scale,
k_tok_bytes,
v_tok_bytes,
g,
None,
)
}
#[allow(clippy::too_many_arguments)]
pub fn fa_decode_dc_q8(
&self,
q: &CudaSlice<f32>,
k: &cudarc::driver::CudaView<u8>,
v: &cudarc::driver::CudaView<u8>,
o: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
n_head_kv: usize,
t_kv_dev: &CudaSlice<i32>,
bucket_max: usize,
scale: f32,
k_tok_bytes: usize,
v_tok_bytes: usize,
g: bool,
q8_out: Option<(&mut CudaSlice<i8>, &mut CudaSlice<f32>)>,
) -> Result<(), Box<dyn std::error::Error>> {
let mut fa_vec =
std::env::var("MEMRA_NO_FA_VEC").is_err() && bucket_max >= fa_vec_min_tkv();
if g && head_dim == 256 && !fa_v4_at(bucket_max) {
fa_vec = false;
} let sp = fa_split_keys(bucket_max, n_head_kv);
let n_splits = if fa_vec {
((bucket_max + sp - 1) / sp).max(1)
} else {
((bucket_max + 255) / 256).max(1)
};
let o_len = n_head * n_splits * head_dim;
let ml_len = n_head * n_splits;
let mut part_guard = self.fa_part_pool.lock().unwrap();
if part_guard
.as_ref()
.map(|pp| pp.0.len() < o_len || pp.1.len() < ml_len)
.unwrap_or(true)
{
let old = part_guard.take();
let (co, cm) = old
.as_ref()
.map(|pp| (pp.0.len(), pp.1.len()))
.unwrap_or((0, 0));
if let Some(old) = old {
self.fa_part_retired.lock().unwrap().push(old);
}
if std::env::var("MEMRA_DEBUG_FAPOOL").is_ok() {
eprintln!(
"[fa-pool] REALLOC o {} -> {} ml {} -> {} (old retired)",
co, o_len, cm, ml_len
);
}
*part_guard = Some((
self.alloc_uninit::<f32>(o_len.max(2 * co))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
self.alloc_uninit::<f32>(ml_len.max(2 * cm))?,
));
}
let pg = part_guard.as_mut().unwrap();
self.gpu
.stream()
.memset_zeros(&mut pg.0.slice_mut(0..o_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.1.slice_mut(0..ml_len))?;
self.gpu
.stream()
.memset_zeros(&mut pg.2.slice_mut(0..ml_len))?;
let (part_o, part_m, part_l) = (&mut pg.0, &mut pg.1, &mut pg.2);
let (hd, nh, nhkv, nsp) = (
head_dim as i32,
n_head as i32,
n_head_kv as i32,
n_splits as i32,
);
let (ktb, vtb) = (k_tok_bytes as i64, v_tok_bytes as i64);
let fa_vec = fa_vec && head_dim <= 512 && head_dim % 32 == 0;
let deep = fa_vec
&& head_dim == 256
&& fa_v4_at(bucket_max)
&& !g
&& fa_deep_at(bucket_max)
&& !matches!(fa_v4_mode(), "noB3" | "stage");
let (f, cfg) = if fa_vec
&& head_dim == 512
&& bucket_max >= {
static FA512_MIN_DC: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*FA512_MIN_DC.get_or_init(|| {
std::env::var("MEMRA_FA512_MIN")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(512)
})
} {
let gqa = (n_head / n_head_kv).max(1) as u32;
(
self.fa_func("fa_decode_vec_q_dpl16_dc", head_dim),
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: 0,
},
)
} else if fa_vec && head_dim == 512 {
let q_view = q.as_view();
let mut o_view = o.as_view_mut();
return self.fa_decode_scalar_unified(
&q_view,
k,
v,
&mut o_view,
head_dim,
n_head,
n_head_kv,
0,
Some(t_kv_dev),
scale,
n_splits,
sp,
k_tok_bytes,
v_tok_bytes,
g,
&mut *part_o,
&mut *part_m,
&mut *part_l,
q8_out,
);
} else if fa_vec && head_dim == 256 && fa_v4_at(bucket_max) {
let gqa = (n_head / n_head_kv).max(1) as u32;
let fv = if g {
self.func_g("fa_decode_vec_q_v4_dc")
} else if deep {
self.func("fa_decode_vec_q_v4_deep_dc")
} else {
self.func("fa_decode_vec_q_v4_dc")
};
let shmem =
(if deep { 12160 } else { 11520 } + 32 * head_dim * if g { 1 } else { 2 }) as u32;
use cudarc::driver::sys::CUfunction_attribute_enum as A;
fv.set_attribute(
A::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
shmem as i32,
)?;
(
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
},
)
} else if fa_vec && fa_v3_active(head_dim) {
let gqa = (n_head / n_head_kv).max(1) as u32;
let fv = if g {
self.func_g("fa_decode_vec_q_v3_dc")
} else {
self.func("fa_decode_vec_q_v3_dc")
};
let shmem = (32 * head_dim * 2) as u32; (
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
},
)
} else if fa_vec && fa_v2_on() {
let gqa = (n_head / n_head_kv).max(1) as u32;
let fv = if g {
self.func_g("fa_decode_vec_q_v2_dc")
} else {
self.func("fa_decode_vec_q_v2_dc")
};
let shmem = (2 * 32 * head_dim * 2) as u32; (
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: shmem,
},
)
} else if fa_vec {
let gqa = (n_head / n_head_kv).max(1) as u32;
let fv = if g {
self.func_g("fa_decode_vec_q_dc")
} else {
self.func("fa_decode_vec_q_dc")
};
(
fv,
LaunchConfig {
grid_dim: (n_head_kv as u32, n_splits as u32, 1),
block_dim: (32, gqa, 1),
shared_mem_bytes: 0,
},
)
} else {
let q_view = q.as_view();
let mut o_view = o.as_view_mut();
return self.fa_decode_scalar_unified(
&q_view,
k,
v,
&mut o_view,
head_dim,
n_head,
n_head_kv,
0,
Some(t_kv_dev),
scale,
n_splits,
if fa_vec { sp } else { 256 },
k_tok_bytes,
v_tok_bytes,
g,
&mut *part_o,
&mut *part_m,
&mut *part_l,
q8_out,
);
};
let ski = sp as i32; let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(&mut *part_o)
.arg(&mut *part_m)
.arg(&mut *part_l)
.arg(&hd)
.arg(&nh)
.arg(&nhkv)
.arg(t_kv_dev)
.arg(&scale)
.arg(&nsp)
.arg(&ski)
.arg(&ktb)
.arg(&vtb);
unsafe {
b.launch(cfg)?;
}
let cfg2 = LaunchConfig {
grid_dim: (n_head as u32, 1, 1),
block_dim: (head_dim as u32, 1, 1),
shared_mem_bytes: 0,
};
if let Some((oq, od)) = q8_out {
let fc = if g {
self.func_g("fa_decode_combine_q8_1")
} else {
self.fa_func("fa_decode_combine_q8_1", head_dim)
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(oq)
.arg(od)
.arg(&hd)
.arg(&nh)
.arg(&nsp);
unsafe {
b2.launch(cfg2)?;
}
return Ok(());
}
let fc = if g {
self.func_g("fa_decode_combine_f32")
} else {
self.fa_func("fa_decode_combine_f32", head_dim)
};
let __s_b2 = self.gpu.stream();
let mut b2 = __s_b2.launch_builder(&fc);
b2.arg(&*part_o)
.arg(&*part_m)
.arg(&*part_l)
.arg(o)
.arg(&hd)
.arg(&nh)
.arg(&nsp);
unsafe {
b2.launch(cfg2)?;
}
Ok(())
}
pub fn fa_geom_eager(
&self,
t_kv: usize,
head_dim: usize,
n_head_kv: usize,
g: bool,
) -> (bool, usize) {
let fa_ok = std::env::var("MEMRA_NO_FA_VEC").is_err() && t_kv >= fa_vec_min_tkv();
let vec512 = fa_ok && head_dim == 512 && t_kv >= fa512_min_tkv();
let mut fa_vec = vec512 || (fa_ok && head_dim <= 256 && head_dim % 32 == 0);
if g && head_dim == 256 && !fa_v4_at(t_kv) {
fa_vec = false;
}
let sp = fa_split_keys(t_kv, n_head_kv);
let n_splits = if fa_vec {
((t_kv + sp - 1) / sp).max(1)
} else {
((t_kv + 255) / 256).max(1)
};
(fa_vec, n_splits)
}
pub fn fa_bucket_key(
&self,
t_kv: usize,
head_dim: usize,
n_head_kv: usize,
g: bool,
) -> (bool, usize) {
self.fa_geom_eager(t_kv, head_dim, n_head_kv, g)
}
pub fn capture_graph_retained<F>(
&self,
step: F,
) -> Result<
(
cudarc::driver::CudaGraph,
Vec<Box<dyn std::any::Any + Send>>,
),
Box<dyn std::error::Error>,
>
where
F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
{
use cudarc::driver::sys::CUgraphInstantiate_flags;
self.capture_graph_retained_flags(
CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
step,
)
}
pub fn capture_graph_retained_flags<F>(
&self,
flags: cudarc::driver::sys::CUgraphInstantiate_flags,
mut step: F,
) -> Result<
(
cudarc::driver::CudaGraph,
Vec<Box<dyn std::any::Any + Send>>,
),
Box<dyn std::error::Error>,
>
where
F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
{
use cudarc::driver::sys::CUstreamCaptureMode;
self.capture_keep.lock().unwrap().clear();
let was_tracking = self.gpu.ctx.is_event_tracking();
if was_tracking {
unsafe {
self.gpu.ctx.disable_event_tracking();
}
}
let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
self.capture_keep_on
.store(true, std::sync::atomic::Ordering::Relaxed);
let w = (|| {
step(self)?;
step(self)
})();
self.capture_keep_on
.store(false, std::sync::atomic::Ordering::Relaxed);
w?;
self.gpu.stream().synchronize()?;
self.gpu
.stream()
.begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
let r = step(self);
let g = self.gpu.stream().end_capture(flags);
r?;
let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
graph.upload()?;
Ok(graph)
};
let result = run();
self.capture_keep_on
.store(false, std::sync::atomic::Ordering::Relaxed);
if was_tracking {
unsafe {
self.gpu.ctx.enable_event_tracking();
}
}
let keeper = std::mem::take(&mut *self.capture_keep.lock().unwrap());
Ok((result?, keeper))
}
pub fn capture_graph<F>(
&self,
mut step: F,
) -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>>
where
F: FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>,
{
use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
let was_tracking = self.gpu.ctx.is_event_tracking();
if was_tracking {
unsafe {
self.gpu.ctx.disable_event_tracking();
}
}
let iflag = {
static F: std::sync::OnceLock<CUgraphInstantiate_flags> = std::sync::OnceLock::new();
*F.get_or_init(|| match std::env::var("MEMRA_GRAPH_IFLAG").as_deref() {
Ok("upload") => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_UPLOAD,
Ok("priority") => {
CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY
}
_ => CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
})
};
let ct = {
static T: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*T.get_or_init(|| std::env::var("MEMRA_GRAPH_CAPTIME").as_deref() == Ok("1"))
};
let warmups = {
static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*W.get_or_init(|| {
std::env::var("MEMRA_GRAPH_WARMUPS")
.ok()
.and_then(|v| v.parse().ok())
.filter(|n| *n >= 1)
.unwrap_or(1)
})
};
let mut run = || -> Result<cudarc::driver::CudaGraph, Box<dyn std::error::Error>> {
let t_w = std::time::Instant::now();
for _ in 0..warmups {
step(self)?;
}
self.gpu.stream().synchronize()?;
let ms_warm = t_w.elapsed().as_secs_f64() * 1e3;
let t_c = std::time::Instant::now();
self.gpu
.stream()
.begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
let r = step(self);
let ms_body = t_c.elapsed().as_secs_f64() * 1e3;
let t_i = std::time::Instant::now();
let g = self.gpu.stream().end_capture(iflag);
let ms_inst = t_i.elapsed().as_secs_f64() * 1e3;
r?;
let graph = g?.ok_or("capture produced no graph (stream was not capturing)")?;
let t_u = std::time::Instant::now();
graph.upload()?;
if ct {
println!(
"[graph-captime] warmup2x {ms_warm:.2} ms capture-body {ms_body:.2} ms \
instantiate {ms_inst:.2} ms upload {:.2} ms",
t_u.elapsed().as_secs_f64() * 1e3
);
}
Ok(graph)
};
let result = run();
if was_tracking {
unsafe {
self.gpu.ctx.enable_event_tracking();
}
}
result
}
pub fn gdn_scan_s128_view(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
state_in: &cudarc::driver::CudaView<f32>,
state_out: &mut cudarc::driver::CudaViewMut<f32>,
o: &mut CudaSlice<f32>,
n_head: usize,
t: usize,
scale: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_scan_s128");
const S_V: u32 = 128;
const WARP: u32 = 32;
const COLS: u32 = 4;
let cfg = LaunchConfig {
grid_dim: (n_head as u32, 1, S_V / COLS),
block_dim: (WARP, COLS, 1),
shared_mem_bytes: 0,
};
let (h, ti) = (n_head as i32, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
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);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn ssm_conv1d_view(
&self,
x: &cudarc::driver::CudaView<f32>,
w: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
conv_dim: usize,
t: usize,
d_conv: usize,
silu: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("ssm_conv1d_silu_f32");
let cfg = LaunchConfig {
grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn ssm_conv1d_tm(
&self,
qkv_tm: &CudaSlice<f32>,
w: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
conv_dim: usize,
t: usize,
d_conv: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("ssm_conv1d_tm_f32");
let cfg = LaunchConfig {
grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn ssm_conv1d_tm_state(
&self,
qkv_tm: &CudaSlice<f32>,
conv_state: &mut CudaSlice<f32>,
w: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
conv_dim: usize,
t: usize,
d_conv: usize,
) -> Result<(), Box<dyn std::error::Error>> {
self.ssm_conv1d_tm_state_pad(qkv_tm, conv_state, w, y, conv_dim, t, d_conv, None)
}
#[allow(clippy::too_many_arguments)]
pub fn ssm_conv1d_tm_state_pad(
&self,
qkv_tm: &CudaSlice<f32>,
conv_state: &mut CudaSlice<f32>,
w: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
conv_dim: usize,
t: usize,
d_conv: usize,
pad_len: Option<&CudaSlice<i32>>,
) -> Result<(), Box<dyn std::error::Error>> {
assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
let ring_old = if t < d_conv - 1 {
Some(self.clone_dtod(conv_state)?)
} else {
None
};
{
let f = self.func("ssm_conv1d_tm_state_f32");
let cfg = LaunchConfig {
grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm)
.arg(&*conv_state)
.arg(w)
.arg(y)
.arg(&cd)
.arg(&ti)
.arg(&dc);
unsafe {
b.launch(cfg)?;
}
}
match (ring_old, pad_len) {
(None, Some(len_d)) => {
let f = self.func("ssm_conv_ring_update_dev_f32");
let n = conv_dim * (d_conv - 1);
let cfg = LaunchConfig::for_num_elems(n as u32);
let (cd, dc) = (conv_dim as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
unsafe {
b.launch(cfg)?;
}
}
(None, None) => {
let f = self.func("ssm_conv_ring_update_f32");
let n = conv_dim * (d_conv - 1);
let cfg = LaunchConfig::for_num_elems(n as u32);
let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
unsafe {
b.launch(cfg)?;
}
}
(Some(old), _) => {
self.ssm_conv_ring_rebuild(qkv_tm, &old, conv_state, conv_dim, t, d_conv)?
}
}
Ok(())
}
pub fn ssm_conv1d_tm_state_pad_v(
&self,
qkv_tm: &cudarc::driver::CudaView<f32>,
conv_state: &mut CudaSlice<f32>,
w: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
conv_dim: usize,
t: usize,
d_conv: usize,
pad_len: Option<&CudaSlice<i32>>,
) -> Result<(), Box<dyn std::error::Error>> {
assert!(t >= 1, "ssm_conv1d_tm_state requires T >= 1");
let ring_old = if t < d_conv - 1 {
Some(self.clone_dtod(conv_state)?)
} else {
None
};
{
let f = self.func("ssm_conv1d_tm_state_f32");
let cfg = LaunchConfig {
grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm)
.arg(&*conv_state)
.arg(w)
.arg(y)
.arg(&cd)
.arg(&ti)
.arg(&dc);
unsafe {
b.launch(cfg)?;
}
}
match (ring_old, pad_len) {
(None, Some(len_d)) => {
let f = self.func("ssm_conv_ring_update_dev_f32");
let n = conv_dim * (d_conv - 1);
let cfg = LaunchConfig::for_num_elems(n as u32);
let (cd, dc) = (conv_dim as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
unsafe {
b.launch(cfg)?;
}
}
(None, None) => {
let f = self.func("ssm_conv_ring_update_f32");
let n = conv_dim * (d_conv - 1);
let cfg = LaunchConfig::for_num_elems(n as u32);
let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
unsafe {
b.launch(cfg)?;
}
}
(Some(_), _) => unreachable!(
"ssm_conv1d_tm_state_pad_v: T < d_conv-1 has no view path (PRIME_MIN_T gates it)"
),
}
Ok(())
}
pub fn ssm_conv_ring_rebuild(
&self,
qkv_tm: &CudaSlice<f32>,
ring_old: &CudaSlice<f32>,
conv_state: &mut CudaSlice<f32>,
conv_dim: usize,
tc: usize,
d_conv: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("ssm_conv_ring_rebuild_f32");
let n = conv_dim * (d_conv - 1);
let cfg = LaunchConfig::for_num_elems(n as u32);
let (cd, ti, dc) = (conv_dim as i32, tc as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm)
.arg(ring_old)
.arg(conv_state)
.arg(&cd)
.arg(&ti)
.arg(&dc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_prep_decode(
&self,
conv_out: &CudaSlice<f32>,
beta_raw: &CudaSlice<f32>,
alpha: &CudaSlice<f32>,
dt_bias: &CudaSlice<f32>,
a: &CudaSlice<f32>,
q_l2: &mut CudaSlice<f32>,
k_l2: &mut CudaSlice<f32>,
v_g: &mut CudaSlice<f32>,
beta: &mut CudaSlice<f32>,
g_log: &mut CudaSlice<f32>,
d_state: usize,
num_v: usize,
num_k: usize,
key_dim: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_prep_decode_f32");
let cfg = LaunchConfig {
grid_dim: (num_v as u32, 1, 1),
block_dim: (32, 4, 1),
shared_mem_bytes: 0,
};
let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(conv_out)
.arg(beta_raw)
.arg(alpha)
.arg(dt_bias)
.arg(a)
.arg(q_l2)
.arg(k_l2)
.arg(v_g)
.arg(beta)
.arg(g_log)
.arg(&ds)
.arg(&nv)
.arg(&nk)
.arg(&kd)
.arg(&eps);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn ssm_conv1d_gdn(
&self,
qkv_tm: &CudaSlice<f32>,
w: &CudaSlice<f32>,
q_g: &mut CudaSlice<f32>,
k_g: &mut CudaSlice<f32>,
v_g: &mut CudaSlice<f32>,
conv_dim: usize,
t: usize,
d_conv: usize,
d_state: usize,
num_v: usize,
num_k: usize,
key_dim: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("ssm_conv1d_gdn_f32");
let cfg = LaunchConfig {
grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
let (ds, nv, nk, kd) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm)
.arg(w)
.arg(q_g)
.arg(k_g)
.arg(v_g)
.arg(&cd)
.arg(&ti)
.arg(&dc)
.arg(&ds)
.arg(&nv)
.arg(&nk)
.arg(&kd);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn ssm_conv1d(
&self,
x: &CudaSlice<f32>,
w: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
conv_dim: usize,
t: usize,
d_conv: usize,
silu: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("ssm_conv1d_silu_f32");
let cfg = LaunchConfig {
grid_dim: (conv_dim as u32, ((t as u32 + 255) / 256).max(1), 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (cd, ti, dc, s) = (conv_dim as i32, t as i32, d_conv as i32, silu as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(w).arg(y).arg(&cd).arg(&ti).arg(&dc).arg(&s);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gdn_scan_s128(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
n_head: usize,
t: usize,
scale: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_scan_s128");
const S_V: u32 = 128;
const WARP: u32 = 32;
const COLS_PER_BLOCK: u32 = 4;
let cfg = LaunchConfig {
grid_dim: (n_head as u32, 1, S_V / COLS_PER_BLOCK),
block_dim: (WARP, COLS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let (h, ti) = (n_head as i32, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
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);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn ssm_conv1d_fused_decode_b(
&self,
qkv_cols: &CudaSlice<f32>,
conv_state_ptrs: &cudarc::driver::CudaView<u64>,
w: &CudaSlice<f32>,
conv_outs: &mut CudaSlice<f32>,
conv_dim: usize,
d_conv: usize,
b_n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("ssm_conv1d_fused_decode_b_f32");
let cfg = LaunchConfig {
grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (cd, dc) = (conv_dim as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_cols)
.arg(conv_state_ptrs)
.arg(w)
.arg(conv_outs)
.arg(&cd)
.arg(&dc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_prep_decode_b(
&self,
conv_outs: &CudaSlice<f32>,
beta_raws: &CudaSlice<f32>,
alphas: &CudaSlice<f32>,
dt_bias: &CudaSlice<f32>,
a: &CudaSlice<f32>,
q_l2: &mut CudaSlice<f32>,
k_l2: &mut CudaSlice<f32>,
v_g: &mut CudaSlice<f32>,
beta: &mut CudaSlice<f32>,
g_log: &mut CudaSlice<f32>,
d_state: usize,
num_v: usize,
num_k: usize,
key_dim: usize,
eps: f32,
conv_dim: usize,
b_n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_prep_decode_b_f32");
let cfg = LaunchConfig {
grid_dim: (num_v as u32, 1, b_n as u32),
block_dim: (32, 4, 1),
shared_mem_bytes: 0,
};
let (ds, nv, nk, kd, cd) = (
d_state as i32,
num_v as i32,
num_k as i32,
key_dim as i32,
conv_dim as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(conv_outs)
.arg(beta_raws)
.arg(alphas)
.arg(dt_bias)
.arg(a)
.arg(q_l2)
.arg(k_l2)
.arg(v_g)
.arg(beta)
.arg(g_log)
.arg(&ds)
.arg(&nv)
.arg(&nk)
.arg(&kd)
.arg(&eps)
.arg(&cd);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_scan_s128_batched(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
state_in_ptrs: &cudarc::driver::CudaView<u64>,
state_out_ptrs: &cudarc::driver::CudaView<u64>,
o: &mut CudaSlice<f32>,
n_head: usize,
b_n: usize,
scale: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_scan_s128_b");
const S_V: u32 = 128;
const WARP: u32 = 32;
const COLS_PER_BLOCK: u32 = 4;
let cfg = LaunchConfig {
grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
block_dim: (WARP, COLS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let h = n_head as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(g)
.arg(beta)
.arg(state_in_ptrs)
.arg(state_out_ptrs)
.arg(o)
.arg(&h)
.arg(&scale);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn ssm_conv1d_fused_decode_b_view(
&self,
qkv_cols: &cudarc::driver::CudaView<f32>,
conv_state_ptrs: &cudarc::driver::CudaView<u64>,
w: &CudaSlice<f32>,
conv_outs: &mut CudaSlice<f32>,
conv_dim: usize,
d_conv: usize,
b_n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("ssm_conv1d_fused_decode_b_f32");
let cfg = LaunchConfig {
grid_dim: (((conv_dim + 255) / 256) as u32, 1, b_n as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (cd, dc) = (conv_dim as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_cols)
.arg(conv_state_ptrs)
.arg(w)
.arg(conv_outs)
.arg(&cd)
.arg(&dc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_prep_decode_b_view(
&self,
conv_outs: &CudaSlice<f32>,
beta_raws: &cudarc::driver::CudaView<f32>,
alphas: &cudarc::driver::CudaView<f32>,
dt_bias: &CudaSlice<f32>,
a: &CudaSlice<f32>,
q_l2: &mut CudaSlice<f32>,
k_l2: &mut CudaSlice<f32>,
v_g: &mut CudaSlice<f32>,
beta: &mut CudaSlice<f32>,
g_log: &mut CudaSlice<f32>,
d_state: usize,
num_v: usize,
num_k: usize,
key_dim: usize,
eps: f32,
conv_dim: usize,
b_n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_prep_decode_b_f32");
let cfg = LaunchConfig {
grid_dim: (num_v as u32, 1, b_n as u32),
block_dim: (32, 4, 1),
shared_mem_bytes: 0,
};
let (ds, nv, nk, kd, cd) = (
d_state as i32,
num_v as i32,
num_k as i32,
key_dim as i32,
conv_dim as i32,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(conv_outs)
.arg(beta_raws)
.arg(alphas)
.arg(dt_bias)
.arg(a)
.arg(q_l2)
.arg(k_l2)
.arg(v_g)
.arg(beta)
.arg(g_log)
.arg(&ds)
.arg(&nv)
.arg(&nk)
.arg(&kd)
.arg(&eps)
.arg(&cd);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_scan_s128_batched_view(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
state_in_ptrs: &cudarc::driver::CudaView<u64>,
state_out_ptrs: &cudarc::driver::CudaView<u64>,
o: &mut cudarc::driver::CudaViewMut<f32>,
n_head: usize,
b_n: usize,
scale: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_scan_s128_b");
const S_V: u32 = 128;
const WARP: u32 = 32;
const COLS_PER_BLOCK: u32 = 4;
let cfg = LaunchConfig {
grid_dim: (n_head as u32, b_n as u32, S_V / COLS_PER_BLOCK),
block_dim: (WARP, COLS_PER_BLOCK, 1),
shared_mem_bytes: 0,
};
let h = n_head as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(v)
.arg(g)
.arg(beta)
.arg(state_in_ptrs)
.arg(state_out_ptrs)
.arg(o)
.arg(&h)
.arg(&scale);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gdn_chunked_enabled() -> bool {
static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*E.get_or_init(|| {
std::env::var("MEMRA_GDN_CHUNKED")
.map(|v| v != "0")
.unwrap_or(true)
})
}
pub fn gdn_chunk_size() -> usize {
static C: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*C.get_or_init(|| {
let c: usize = std::env::var("MEMRA_GDN_CHUNK")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(32);
c.clamp(32, 128) / 32 * 32
})
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments, clippy::type_complexity)]
#[allow(clippy::too_many_arguments)]
pub fn gdn_chunk_k123(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
wb16: Option<&mut CudaSlice<u8>>,
n_head: usize,
t: usize,
c: usize,
hk: usize,
k2w: Option<(&CudaSlice<u8>, &CudaSlice<u8>, &mut CudaSlice<u8>)>,
) -> Result<
(
CudaSlice<f32>,
CudaSlice<f32>,
CudaSlice<f32>,
CudaSlice<f32>,
),
Box<dyn std::error::Error>,
> {
const D: usize = 128;
let h = n_head;
let nc = (t + c - 1) / c;
let (hi, ti, ci) = (h as i32, t as i32, c as i32);
let mut gcum = self.uninit(t * h)?;
let mut a = self.uninit(nc * h * c * c)?;
let mut p = self.uninit(nc * h * c * c)?;
let mut u = self.uninit(nc * h * c * D)?;
let mut w = self.uninit(nc * h * c * D)?;
{
let f = self.func("gdn_chunk_cumgate_f32");
let cfg = LaunchConfig {
grid_dim: (nc as u32, h as u32, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(g).arg(&mut gcum).arg(&hi).arg(&ti).arg(&ci);
unsafe {
b.launch(cfg)?;
}
}
if let Some((qb, kb, pb)) = k2w {
assert!(c == 32, "gdn_k2_wgmma is a C==32 tile");
let f = self.func("gdn_k2_wgmma");
let cfg = LaunchConfig {
grid_dim: (nc as u32, h as u32, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let hki = hk as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qb)
.arg(kb)
.arg(&gcum)
.arg(beta)
.arg(&mut a)
.arg(&mut *pb)
.arg(&hi)
.arg(&ti)
.arg(&ci)
.arg(&hki);
unsafe {
b.launch(cfg)?;
}
} else if c <= 64 && !portable_mma_gated() {
let f = self.func("gdn_chunk_attn_f32");
let jt = ((c + 31) / 32) as u32;
let cfg = LaunchConfig {
grid_dim: (nc as u32, h as u32, jt),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let hki = hk as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(&gcum)
.arg(beta)
.arg(&mut a)
.arg(&mut p)
.arg(&hi)
.arg(&ti)
.arg(&ci)
.arg(&hki);
unsafe {
b.launch(cfg)?;
}
} else {
assert!(
hk == h,
"generic K2 is broadcast-only (de-broadcast rides C==32)"
);
let f = self.func("gdn_chunk_attn_g_f32");
let cfg = LaunchConfig {
grid_dim: (nc as u32, h as u32, 1),
block_dim: (32, 8, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(k)
.arg(&gcum)
.arg(beta)
.arg(&mut a)
.arg(&mut p)
.arg(&hi)
.arg(&ti)
.arg(&ci);
unsafe {
b.launch(cfg)?;
}
}
{
let cfg = LaunchConfig {
grid_dim: (nc as u32, h as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
match c {
32 | 64 => {
let f = self.func(if c == 32 {
"gdn_chunk_solve32_f32"
} else {
"gdn_chunk_solve64_f32"
});
let wb: u64 = match wb16 {
Some(d) => self.addr_u8(d),
None => 0,
};
let hki = hk as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(v)
.arg(k)
.arg(&a)
.arg(&gcum)
.arg(&mut u)
.arg(&mut w)
.arg(&wb)
.arg(&hi)
.arg(&ti)
.arg(&hki);
unsafe {
b.launch(cfg)?;
}
}
_ => {
assert!(hk == h, "generic K3 is broadcast-only");
let f = self.func("gdn_chunk_solve_f32");
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(v)
.arg(k)
.arg(&a)
.arg(&gcum)
.arg(&mut u)
.arg(&mut w)
.arg(&hi)
.arg(&ti)
.arg(&ci);
unsafe {
b.launch(cfg)?;
}
}
}
}
Ok((gcum, p, u, w))
}
pub fn gdn_db_on() -> bool {
std::env::var("MEMRA_GDN_DB").as_deref() != Ok("0")
}
pub fn gdn_mma_enabled(&self, c: usize) -> bool {
!portable_mma_gated()
&& c == 32
&& match std::env::var("MEMRA_GDN_MMA").as_deref() {
Ok("1") => true,
Ok("0") => false,
_ => cfg!(memra_hopper_mma),
}
}
pub fn gdn_wgmma_on(&self, c: usize) -> bool {
self.gdn_mma_enabled(c)
&& match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
Ok("0") => false,
Ok("1") => true,
_ => cfg!(memra_hopper_mma),
}
}
#[allow(clippy::too_many_arguments)]
pub fn ssm_conv1d_gdn_state_pad(
&self,
qkv_tm: &cudarc::driver::CudaView<f32>,
conv_state: &mut CudaSlice<f32>,
w: &CudaSlice<f32>,
q_g: &mut CudaSlice<f32>,
k_g: &mut CudaSlice<f32>,
v_g: &mut CudaSlice<f32>,
conv_dim: usize,
t: usize,
d_conv: usize,
d_state: usize,
num_v: usize,
num_k: usize,
key_dim: usize,
hk: usize,
pad_len: Option<&CudaSlice<i32>>,
) -> Result<(), Box<dyn std::error::Error>> {
assert!(
t >= d_conv - 1,
"fused state conv requires T >= pad (PRIME_MIN_T gates)"
);
{
let f = self.func("ssm_conv1d_gdn_state_f32");
let cfg = LaunchConfig {
grid_dim: (((conv_dim + 255) / 256) as u32, t as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
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,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm)
.arg(&*conv_state)
.arg(w)
.arg(q_g)
.arg(k_g)
.arg(v_g)
.arg(&cd)
.arg(&ti)
.arg(&dc)
.arg(&ds)
.arg(&nv)
.arg(&nk)
.arg(&kd)
.arg(&hki);
unsafe {
b.launch(cfg)?;
}
}
match pad_len {
Some(len_d) => {
let f = self.func("ssm_conv_ring_update_dev_f32");
let n = conv_dim * (d_conv - 1);
let cfg = LaunchConfig::for_num_elems(n as u32);
let (cd, dc) = (conv_dim as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm).arg(conv_state).arg(len_d).arg(&cd).arg(&dc);
unsafe {
b.launch(cfg)?;
}
}
None => {
let f = self.func("ssm_conv_ring_update_f32");
let n = conv_dim * (d_conv - 1);
let cfg = LaunchConfig::for_num_elems(n as u32);
let (cd, ti, dc) = (conv_dim as i32, t as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_tm).arg(conv_state).arg(&cd).arg(&ti).arg(&dc);
unsafe {
b.launch(cfg)?;
}
}
}
Ok(())
}
pub fn gdn_chunk_alloc(
&self,
n_head: usize,
t: usize,
c: usize,
hk: usize,
) -> Result<GdnChunkBufs, Box<dyn std::error::Error>> {
const D: usize = 128;
assert!(
c == 32,
"gdn_chunk_alloc: varlen chain is the C==32 mma pair"
);
let h = n_head;
let nc = (t + c - 1) / c;
Ok(GdnChunkBufs {
gcum: self.uninit(t * h)?,
a: self.uninit(nc * h * c * c)?,
p: self.uninit(nc * h * c * c)?,
u: self.uninit(nc * h * c * D)?,
w: self.uninit(nc * h * c * D)?,
kb16: self.alloc_u8_uninit(t * hk * D * 2)?,
wb16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
y16: self.alloc_u8_uninit(nc * h * c * D * 2)?,
ssnap16: self.alloc_u8_uninit(nc * h * D * D * 2)?,
qb16: self.alloc_u8_uninit(t * hk * D * 2)?,
pb16: self.alloc_u8_uninit(nc * h * c * c * 2)?,
o: self.uninit(D * h * t)?,
t,
nc,
})
}
pub fn f32_to_bf16_v(
&self,
x: &cudarc::driver::CudaView<f32>,
dst: &mut CudaSlice<u8>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("f32_to_bf16_bulk");
let ni = n as i64;
let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(dst).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn f32_to_bf16_into(
&self,
x: &CudaSlice<f32>,
dst: &mut CudaSlice<u8>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("f32_to_bf16_bulk");
let ni = n as i64;
let cfg = LaunchConfig::for_num_elems((n as u32).div_ceil(4));
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(dst).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gdn_chunk_k123_vl8(
&self,
seqs: &[GdnSeqVl],
n_head: usize,
hk: usize,
wq: Option<&GdnWVl8>,
) -> Result<(), Box<dyn std::error::Error>> {
let b = seqs.len();
assert!(b >= 1 && b <= 8, "gdn_chunk_k123_vl8: 1..=8 sequences");
let mut packed = [GdnSeqVl::default(); 8];
packed[..b].copy_from_slice(seqs);
let v = GdnVl8(packed);
let (hi, ci) = (n_head as i32, 32i32);
let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
{
let f = self.func("gdn_chunk_cumgate_vl");
let cfg = LaunchConfig {
grid_dim: (max_nc, n_head as u32, b as u32),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&hi).arg(&ci);
unsafe {
lb.launch(cfg)?;
}
}
let hki = hk as i32;
if let Some(w) = wq {
let f = self.func("gdn_k2_wgmma_vl");
let cfg = LaunchConfig {
grid_dim: (max_nc, n_head as u32, b as u32),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(w).arg(&hi).arg(&ci).arg(&hki);
unsafe {
lb.launch(cfg)?;
}
} else {
let f = self.func("gdn_chunk_attn_vl");
let cfg = LaunchConfig {
grid_dim: (max_nc, n_head as u32, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
unsafe {
lb.launch(cfg)?;
}
}
{
let f = self.func("gdn_chunk_solve32_vl");
let cfg = LaunchConfig {
grid_dim: (max_nc, n_head as u32, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
unsafe {
lb.launch(cfg)?;
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gdn_prep_vl8(
&self,
seqs: &[GdnPrepVl],
conv_w: &CudaSlice<f32>,
dt_bias: &CudaSlice<f32>,
a: &CudaSlice<f32>,
conv_dim: usize,
d_conv: usize,
d_state: usize,
num_v: usize,
num_k: usize,
key_dim: usize,
hk: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let b = seqs.len();
assert!(b >= 1 && b <= 8);
let mut packed = [GdnPrepVl::default(); 8];
packed[..b].copy_from_slice(seqs);
let v = GdnPrepVl8(packed);
let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
let (cdi, dci) = (conv_dim as i32, d_conv as i32);
let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
assert!(
conv_fuse || hk == num_v,
"de-broadcast requires the fused conv"
);
if conv_fuse {
let f = self.func("ssm_conv1d_gdn_state_vl");
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,
};
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,
);
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v)
.arg(conv_w)
.arg(&cdi)
.arg(&dci)
.arg(&dsi)
.arg(&nvi)
.arg(&nki)
.arg(&kdi)
.arg(&hki);
unsafe {
lb.launch(cfg)?;
}
} else {
let f = self.func("ssm_conv1d_tm_state_vl");
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,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(conv_w).arg(&cdi).arg(&dci);
unsafe {
lb.launch(cfg)?;
}
}
{
let f = self.func("ssm_conv_ring_update_vl");
let n = (conv_dim * (d_conv - 1)) as u32;
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256), 1, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&cdi).arg(&dci);
unsafe {
lb.launch(cfg)?;
}
}
if !conv_fuse {
let f = self.func("qkv_to_gdn_repack_vl");
let n = max_t * (num_v * d_state) as u32;
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256), 1, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (dsi, nvi, nki, kdi) = (d_state as i32, num_v as i32, num_k as i32, key_dim as i32);
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&dsi).arg(&nvi).arg(&nki).arg(&kdi);
unsafe {
lb.launch(cfg)?;
}
}
if Self::l2_v2_on(d_state) {
let f = self.func("gdn_l2_v2_vl");
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,
};
let (dsi, nvi) = (d_state as i32, hk as i32);
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
unsafe {
lb.launch(cfg)?;
}
} else {
let f = self.func("gdn_l2_vl");
let cfg = LaunchConfig {
grid_dim: (max_t * hk as u32, 2, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let (dsi, nvi) = (d_state as i32, hk as i32);
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&dsi).arg(&nvi).arg(&eps);
unsafe {
lb.launch(cfg)?;
}
}
{
let f = self.func("gdn_gate_prep_vl");
let n = max_t * num_v as u32;
let cfg = LaunchConfig {
grid_dim: (n.div_ceil(256), 1, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let nvi = num_v as i32;
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(dt_bias).arg(a).arg(&nvi);
unsafe {
lb.launch(cfg)?;
}
}
Ok(())
}
pub fn gdn_mirror_vl8(
&self,
seqs: &[GdnSeqVl],
n_head: usize,
which: i32,
hk: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let b = seqs.len();
assert!(b >= 1 && b <= 8);
let mut packed = [GdnSeqVl::default(); 8];
packed[..b].copy_from_slice(seqs);
let v = GdnVl8(packed);
let ept = (if which == 0 { hk } else { n_head } * 128) as i32;
let max_n = seqs
.iter()
.map(|s| {
if which == 0 {
s.t as i64 * ept as i64
} else {
s.nc as i64 * ept as i64 * 32
}
})
.max()
.unwrap();
let f = self.func("gdn_mirror_vl");
let blocks = ((max_n as u32).div_ceil(4)).div_ceil(256);
let cfg = LaunchConfig {
grid_dim: (blocks, 1, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&ept).arg(&which);
unsafe {
lb.launch(cfg)?;
}
Ok(())
}
pub fn gdn_tail_vl8(
&self,
seqs: &[GdnPrepVl],
norm_w: &CudaSlice<f32>,
d_state: usize,
num_v: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let b = seqs.len();
assert!(b >= 1 && b <= 8);
let mut packed = [GdnPrepVl::default(); 8];
packed[..b].copy_from_slice(seqs);
let v = GdnPrepVl8(packed);
let max_t = seqs.iter().map(|s| s.t).max().unwrap() as u32;
let f = self.func("gated_rmsnorm_f16out_vl");
let cfg = LaunchConfig {
grid_dim: (max_t * num_v as u32, 1, b as u32),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (dsi, nvi) = (d_state as i32, num_v as i32);
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(norm_w).arg(&dsi).arg(&nvi).arg(&eps);
unsafe {
lb.launch(cfg)?;
}
Ok(())
}
pub fn addr_f32(&self, x: &CudaSlice<f32>) -> u64 {
use cudarc::driver::DevicePtr;
let s = self.gpu.stream();
let (p, _g) = x.device_ptr(&s);
p as u64
}
pub fn addr_f32_mut(&self, x: &mut CudaSlice<f32>) -> u64 {
use cudarc::driver::DevicePtrMut;
let s = self.gpu.stream();
let (p, _g) = x.device_ptr_mut(&s);
p as u64
}
pub fn addr_f32v(&self, x: &cudarc::driver::CudaView<f32>) -> u64 {
use cudarc::driver::DevicePtr;
let s = self.gpu.stream();
let (p, _g) = x.device_ptr(&s);
p as u64
}
pub fn addr_u8(&self, x: &CudaSlice<u8>) -> u64 {
use cudarc::driver::DevicePtr;
let s = self.gpu.stream();
let (p, _g) = x.device_ptr(&s);
p as u64
}
pub fn gdn_chunk_vl8(
&self,
seqs: &[GdnSeqVl],
n_head: usize,
scale: f32,
hk: usize,
wq: Option<&GdnWVl8>,
) -> Result<(), Box<dyn std::error::Error>> {
const NSPLIT: u32 = 4;
let b = seqs.len();
assert!(b >= 1 && b <= 8, "gdn_chunk_vl8: 1..=8 sequences");
let mut packed = [GdnSeqVl::default(); 8];
packed[..b].copy_from_slice(seqs);
let v = GdnVl8(packed);
let (hi, ci) = (n_head as i32, 32i32);
let max_nc = seqs.iter().map(|a| a.nc).max().unwrap() as u32;
let hki = hk as i32;
if let Some(w) = wq {
let f = self.func("gdn_k45_wgmma_vl");
let cfg = LaunchConfig {
grid_dim: (n_head as u32, NSPLIT, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(w).arg(&scale).arg(&hi).arg(&ci).arg(&hki);
unsafe {
lb.launch(cfg)?;
}
let _ = max_nc;
return Ok(());
}
{
let f = self.func("gdn_chunk_state_mma_vl");
let cfg = LaunchConfig {
grid_dim: (n_head as u32, NSPLIT, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&hi).arg(&ci).arg(&hki);
unsafe {
lb.launch(cfg)?;
}
}
{
let f = self.func("gdn_chunk_output_mma_vl");
let cfg = LaunchConfig {
grid_dim: (max_nc, n_head as u32, b as u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_lb = self.gpu.stream();
let mut lb = __s_lb.launch_builder(&f);
lb.arg(&v).arg(&hi).arg(&ci).arg(&scale).arg(&hki);
unsafe {
lb.launch(cfg)?;
}
}
Ok(())
}
pub fn gdn_scan_chunked(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
kb16_pre: Option<&CudaSlice<u8>>,
qb16_pre: Option<&CudaSlice<u8>>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
n_head: usize,
t: usize,
scale: f32,
c: usize,
hk: usize,
) -> Result<(), Box<dyn std::error::Error>> {
const D: usize = 128;
const NSPLIT: u32 = 4;
assert!(c >= 1 && c <= 128, "gdn_scan_chunked: C must be in 1..=128");
let h = n_head;
let nc = (t + c - 1) / c;
let (hi, ti, ci) = (h as i32, t as i32, c as i32);
let gdn_mma_pre = !portable_mma_gated()
&& c == 32
&& match std::env::var("MEMRA_GDN_MMA").as_deref() {
Ok("1") => true,
Ok("0") => false,
_ => cfg!(memra_hopper_mma),
};
let mut wb16_pre: Option<CudaSlice<u8>> = if gdn_mma_pre {
Some(self.alloc_u8_uninit(nc * h * c * D * 2)?)
} else {
None
};
let gdn_wgmma_pre = gdn_mma_pre
&& match std::env::var("MEMRA_GDN_WGMMA").as_deref() {
Ok("0") => false,
Ok("1") => true,
_ => cfg!(memra_hopper_mma),
};
let nk = t * hk * D;
let mut kb16_local: Option<CudaSlice<u8>> = None;
if gdn_mma_pre && kb16_pre.is_none() {
let mut kb = self.alloc_u8_uninit(nk * 2)?;
let f = self.func("f32_to_bf16_bulk");
let n2 = nk as i64;
let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k).arg(&mut kb).arg(&n2);
unsafe {
b.launch(cfg2)?;
}
kb16_local = Some(kb);
}
let kb16_ref0: Option<&CudaSlice<u8>> = kb16_local.as_ref().or(kb16_pre);
if let Some(kb) = kb16_pre {
assert!(kb.len() >= nk * 2, "kb16_pre too small");
}
let mut qb16: Option<CudaSlice<u8>> = None;
let mut pb16: Option<CudaSlice<u8>> = None;
if gdn_wgmma_pre {
if qb16_pre.is_none() {
let mut qb = self.alloc_u8_uninit(nk * 2)?;
let f = self.func("f32_to_bf16_bulk");
let n2 = nk as i64;
let cfg2 = LaunchConfig::for_num_elems((nk as u32).div_ceil(4));
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q).arg(&mut qb).arg(&n2);
unsafe {
b.launch(cfg2)?;
}
qb16 = Some(qb);
} else if let Some(qb) = qb16_pre {
assert!(qb.len() >= nk * 2, "qb16_pre too small");
}
pb16 = Some(self.alloc_u8_uninit(nc * h * c * c * 2)?);
}
let qb16_ref0: Option<&CudaSlice<u8>> = qb16.as_ref().or(qb16_pre);
let k2w = if gdn_wgmma_pre {
Some((
*qb16_ref0.as_ref().unwrap(),
*kb16_ref0.as_ref().unwrap(),
pb16.as_mut().unwrap(),
))
} else {
None
};
let (gcum, p, u, w) =
self.gdn_chunk_k123(q, k, v, g, beta, wb16_pre.as_mut(), n_head, t, c, hk, k2w)?;
let _ = &w;
let mut y = self.uninit(nc * h * c * D)?;
let mut ssnap = self.uninit(nc * h * D * D)?; let gdn_mma = !portable_mma_gated()
&& c == 32
&& match std::env::var("MEMRA_GDN_MMA").as_deref() {
Ok("1") => true,
Ok("0") => false,
_ => cfg!(memra_hopper_mma),
};
if gdn_mma {
let wb16 = wb16_pre
.take()
.expect("mma path pre-allocates wb16 (K3 store fold)");
let kb16_ref: &CudaSlice<u8> = kb16_ref0.expect("mma path pre-builds kb16 above K123");
if gdn_wgmma_pre {
let qb16 = qb16_ref0.unwrap();
let pb16 = pb16.as_ref().unwrap();
{
let f = self.func("gdn_k45_wgmma");
let cfg = LaunchConfig {
grid_dim: (h as u32, 4, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let hki = hk as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(kb16_ref)
.arg(&gcum)
.arg(beta)
.arg(&u)
.arg(&wb16)
.arg(qb16)
.arg(pb16)
.arg(o)
.arg(&scale)
.arg(state_in)
.arg(&mut *state_out)
.arg(&hi)
.arg(&ti)
.arg(&ci)
.arg(&hki);
unsafe {
b.launch(cfg)?;
}
}
return Ok(());
}
let mut y16 = self.alloc_u8_uninit(nc * h * c * D * 2)?;
let mut ssnap16 = self.alloc_u8_uninit(nc * h * D * D * 2)?;
{
let f = self.func("gdn_chunk_state_mma");
let cfg = LaunchConfig {
grid_dim: (h as u32, NSPLIT, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let hki = hk as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(kb16_ref)
.arg(&gcum)
.arg(beta)
.arg(&u)
.arg(&wb16)
.arg(&mut y16)
.arg(&mut ssnap16)
.arg(state_in)
.arg(&mut *state_out)
.arg(&hi)
.arg(&ti)
.arg(&ci)
.arg(&hki);
unsafe {
b.launch(cfg)?;
}
}
{
let f = self.func("gdn_chunk_output_mma");
let jt = ((c + 31) / 32) as u32;
let cfg = LaunchConfig {
grid_dim: (nc as u32, h as u32, jt),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let hki = hk as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(&gcum)
.arg(&p)
.arg(&y16)
.arg(&ssnap16)
.arg(o)
.arg(&hi)
.arg(&ti)
.arg(&ci)
.arg(&scale)
.arg(&hki);
unsafe {
b.launch(cfg)?;
}
}
return Ok(());
}
{
let f = self.func("gdn_chunk_state_f32");
let cfg = LaunchConfig {
grid_dim: (h as u32, NSPLIT, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(k)
.arg(&gcum)
.arg(beta)
.arg(&u)
.arg(&w)
.arg(&mut y)
.arg(&mut ssnap)
.arg(state_in)
.arg(&mut *state_out)
.arg(&hi)
.arg(&ti)
.arg(&ci);
unsafe {
b.launch(cfg)?;
}
}
{
let f = self.func("gdn_chunk_output_f32");
let jt = ((c + 31) / 32) as u32;
let cfg = LaunchConfig {
grid_dim: (nc as u32, h as u32, jt),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(q)
.arg(&gcum)
.arg(&p)
.arg(&y)
.arg(&ssnap)
.arg(o)
.arg(&hi)
.arg(&ti)
.arg(&ci)
.arg(&scale);
unsafe {
b.launch(cfg)?;
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
pub fn gdn_scan_prefill(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
kb16_pre: Option<&CudaSlice<u8>>,
qb16_pre: Option<&CudaSlice<u8>>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
n_head: usize,
t: usize,
scale: f32,
hk: usize,
) -> Result<(), Box<dyn std::error::Error>> {
if std::env::var("MEMRA_GDN_DIFF").is_ok() && t >= 16 {
assert!(hk == n_head, "GDN_DIFF oracle is broadcast-only");
return self.gdn_scan_diff(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale);
}
if Self::gdn_chunked_enabled() && t >= 16 {
self.gdn_scan_chunked(
q,
k,
v,
g,
beta,
kb16_pre,
qb16_pre,
state_in,
state_out,
o,
n_head,
t,
scale,
Self::gdn_chunk_size(),
hk,
)
} else {
assert!(
hk == n_head,
"s128 scan is broadcast-only (prep guarantees by predicate)"
);
self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)
}
}
#[allow(clippy::too_many_arguments)]
fn gdn_scan_diff(
&self,
q: &CudaSlice<f32>,
k: &CudaSlice<f32>,
v: &CudaSlice<f32>,
g: &CudaSlice<f32>,
beta: &CudaSlice<f32>,
state_in: &CudaSlice<f32>,
state_out: &mut CudaSlice<f32>,
o: &mut CudaSlice<f32>,
n_head: usize,
t: usize,
scale: f32,
) -> Result<(), Box<dyn std::error::Error>> {
static CALL: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
let call = CALL.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let mut o_c = self.uninit(o.len())?;
let mut st_c = self.uninit(state_out.len())?;
self.gdn_scan_chunked(
q,
k,
v,
g,
beta,
None,
None,
state_in,
&mut st_c,
&mut o_c,
n_head,
t,
scale,
Self::gdn_chunk_size(),
n_head,
)?;
self.gdn_scan_s128(q, k, v, g, beta, state_in, state_out, o, n_head, t, scale)?;
let (oh_s, oh_c) = (self.dtoh(o)?, self.dtoh(&o_c)?);
let (sh_s, sh_c) = (self.dtoh(state_out)?, self.dtoh(&st_c)?);
let stats = |a: &[f32], b: &[f32]| -> (f32, f32, f64) {
let mut max_abs = 0f32;
let mut max_rel = 0f32;
let mut sum_rel = 0f64;
for (x, y) in a.iter().zip(b) {
let ad = (x - y).abs();
let rel = ad / x.abs().max(y.abs()).max(1e-3);
if ad > max_abs {
max_abs = ad;
}
if rel > max_rel {
max_rel = rel;
}
sum_rel += rel as f64;
}
(max_abs, max_rel, sum_rel / a.len() as f64)
};
let (o_ma, o_mr, o_mean) = stats(&oh_s, &oh_c);
let (s_ma, s_mr, s_mean) = stats(&sh_s, &sh_c);
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} | \
state: max_abs={s_ma:.3e} max_rel={s_mr:.3e} mean_rel={s_mean:.3e}",
Self::gdn_chunk_size()
);
Ok(())
}
pub fn gdn_glog(
&self,
alpha: &CudaSlice<f32>,
dt_bias: &CudaSlice<f32>,
a: &CudaSlice<f32>,
g_log: &mut CudaSlice<f32>,
n_head: usize,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_glog_f32");
let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
let (h, ti) = (n_head as i32, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn sigmoid_v(
&self,
x: &cudarc::driver::CudaView<f32>,
y: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("sigmoid_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(y).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gdn_glog_v(
&self,
alpha: &cudarc::driver::CudaView<f32>,
dt_bias: &CudaSlice<f32>,
a: &CudaSlice<f32>,
g_log: &mut CudaSlice<f32>,
n_head: usize,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gdn_glog_f32");
let cfg = LaunchConfig::for_num_elems((n_head * t) as u32);
let (h, ti) = (n_head as i32, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(alpha).arg(dt_bias).arg(a).arg(g_log).arg(&h).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn sigmoid(
&self,
x: &CudaSlice<f32>,
y: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("sigmoid_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(x).arg(y).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn sig_mul_f16out(
&self,
a: &CudaSlice<f32>,
g: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
dst16: &mut CudaSlice<u8>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("sig_mul_f16out_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a).arg(g).arg(dst).arg(dst16).arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn attn_head_gate(
&self,
a: &CudaSlice<f32>,
g: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
dst16: Option<&mut CudaSlice<u8>>,
head_dim: usize,
n_head: usize,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("attn_head_gate_f32");
let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
let d16: u64 = match dst16 {
Some(d) => self.addr_u8(d),
None => 0,
};
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a)
.arg(g)
.arg(dst)
.arg(&d16)
.arg(&hd)
.arg(&nh)
.arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn swiglu_clamped_mul_scaled(
&self,
gate: &CudaSlice<f32>,
up: &CudaSlice<f32>,
gs: f32,
us: f32,
limit: f32,
dst: &mut CudaSlice<f32>,
n: usize,
) -> Result<(), Box<dyn std::error::Error>> {
debug_assert!(
limit > 1e-6,
"swiglu_clamped needs a live limit; use silu_mul_scaled"
);
let f = self.func("swiglu_clamped_mul_scaled_f32");
let cfg = LaunchConfig::for_num_elems(n as u32);
let ni = n as i32;
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(gate)
.arg(up)
.arg(&gs)
.arg(&us)
.arg(&limit)
.arg(dst)
.arg(&ni);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gated_rmsnorm(
&self,
o: &CudaSlice<f32>,
w: &CudaSlice<f32>,
z: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gated_rmsnorm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gated_rmsnorm_f16out(
&self,
o: &CudaSlice<f32>,
w: &CudaSlice<f32>,
z: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
dst16: &mut CudaSlice<u8>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gated_rmsnorm_f16out_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn add_rms_norm_zq8(
&self,
a: &CudaSlice<f32>,
b_in: &CudaSlice<f32>,
w: &CudaSlice<f32>,
res: &mut CudaSlice<f32>,
z: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
assert!(ncols % 32 == 0);
let mut q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let f = self.func("add_rms_norm_zq8");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (1024, 1, 1),
shared_mem_bytes: 0,
};
let (nc, ep) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(a)
.arg(b_in)
.arg(w)
.arg(res)
.arg(z)
.arg(&mut q)
.arg(&mut d)
.arg(&nc)
.arg(&ep);
unsafe {
b.launch(cfg)?;
}
Ok((q, d))
}
pub fn gated_rmsnorm_zv(
&self,
o: &CudaSlice<f32>,
w: &CudaSlice<f32>,
z: &cudarc::driver::CudaView<f32>,
dst: &mut CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gated_rmsnorm_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(o).arg(w).arg(z).arg(dst).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gated_rmsnorm_f16out_zv(
&self,
o: &CudaSlice<f32>,
w: &CudaSlice<f32>,
z: &cudarc::driver::CudaView<f32>,
dst: &mut CudaSlice<f32>,
dst16: &mut CudaSlice<u8>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("gated_rmsnorm_f16out_f32");
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (nc, e) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(o).arg(w).arg(z).arg(dst).arg(dst16).arg(&nc).arg(&e);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn gated_rmsnorm_q8_1(
&self,
o: &CudaSlice<f32>,
w: &CudaSlice<f32>,
z: &CudaSlice<f32>,
ncols: usize,
nrows: usize,
eps: f32,
) -> Result<(CudaSlice<i8>, CudaSlice<f32>), Box<dyn std::error::Error>> {
assert!(ncols % 32 == 0);
let f = self.func("gated_rmsnorm_q8_1");
let mut out_q = self.alloc_uninit::<i8>(nrows * ncols)?;
let mut out_d = self.alloc_uninit::<f32>(nrows * (ncols / 32))?;
let cfg = LaunchConfig {
grid_dim: (nrows as u32, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let (nc, ep) = (ncols as i32, eps);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(o)
.arg(w)
.arg(z)
.arg(&mut out_q)
.arg(&mut out_d)
.arg(&nc)
.arg(&ep);
unsafe {
b.launch(cfg)?;
}
Ok((out_q, out_d))
}
pub fn transpose(
&self,
inp: &CudaSlice<f32>,
rows: usize,
cols: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let f = self.func("transpose_f32");
let mut out = self.zeros(rows * cols)?;
let cfg = LaunchConfig::for_num_elems((rows * cols) as u32);
let (r, c) = (rows as i32, cols as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(inp).arg(&mut out).arg(&r).arg(&c);
unsafe {
b.launch(cfg)?;
}
Ok(out)
}
pub fn repeat_heads(
&self,
inp: &CudaSlice<f32>,
out: &mut CudaSlice<f32>,
head_dim: usize,
n_in: usize,
n_out: usize,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("repeat_heads_f32");
let cfg = LaunchConfig::for_num_elems((head_dim * n_out * t) as u32);
let (hd, ni, no, ti) = (head_dim as i32, n_in as i32, n_out as i32, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(inp).arg(out).arg(&hd).arg(&ni).arg(&no).arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn q_gate_split(
&self,
qf: &CudaSlice<f32>,
q_out: &mut CudaSlice<f32>,
gate_out: &mut CudaSlice<f32>,
head_dim: usize,
n_head: usize,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("q_gate_split_f32");
let cfg = LaunchConfig::for_num_elems((head_dim * n_head * t) as u32);
let (hd, nh, ti) = (head_dim as i32, n_head as i32, t as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qf)
.arg(q_out)
.arg(gate_out)
.arg(&hd)
.arg(&nh)
.arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn qkv_to_gdn_repack(
&self,
conv_out: &CudaSlice<f32>,
q_g: &mut CudaSlice<f32>,
k_g: &mut CudaSlice<f32>,
v_g: &mut CudaSlice<f32>,
d_state: usize,
num_v: usize,
num_k: usize,
key_dim: usize,
t: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("qkv_to_gdn_repack_f32");
let cfg = LaunchConfig::for_num_elems((d_state * num_v * t) as u32);
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,
);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(conv_out)
.arg(q_g)
.arg(k_g)
.arg(v_g)
.arg(&ds)
.arg(&nv)
.arg(&nk)
.arg(&kd)
.arg(&ti);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn conv_left_pad(
&self,
src: &CudaSlice<f32>,
dst: &mut CudaSlice<f32>,
conv_dim: usize,
t: usize,
pad: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("conv_left_pad_f32");
let cfg = LaunchConfig::for_num_elems((conv_dim * t) as u32);
let (cd, ti, p) = (conv_dim as i32, t as i32, pad as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(src).arg(dst).arg(&cd).arg(&ti).arg(&p);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn conv_assemble_and_roll(
&self,
qkv_col: &CudaSlice<f32>,
conv_state: &mut CudaSlice<f32>,
conv_in: &mut CudaSlice<f32>,
conv_dim: usize,
pad: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("conv_assemble_and_roll_f32");
let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
let (cd, p) = (conv_dim as i32, pad as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_col).arg(conv_state).arg(conv_in).arg(&cd).arg(&p);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn ssm_conv1d_fused_decode(
&self,
qkv_col: &CudaSlice<f32>,
conv_state: &mut CudaSlice<f32>,
w: &CudaSlice<f32>,
conv_out: &mut CudaSlice<f32>,
conv_dim: usize,
d_conv: usize,
) -> Result<(), Box<dyn std::error::Error>> {
let f = self.func("ssm_conv1d_fused_decode_f32");
let cfg = LaunchConfig::for_num_elems(conv_dim as u32);
let (cd, dc) = (conv_dim as i32, d_conv as i32);
let __s_b = self.gpu.stream();
let mut b = __s_b.launch_builder(&f);
b.arg(qkv_col)
.arg(conv_state)
.arg(w)
.arg(conv_out)
.arg(&cd)
.arg(&dc);
unsafe {
b.launch(cfg)?;
}
Ok(())
}
pub fn slice_range(
&self,
src: &CudaSlice<f32>,
start: usize,
len: usize,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
let host = self.gpu.stream().clone_dtoh(src)?;
self.gpu.stream().synchronize()?;
Ok(self.htod(&host[start..start + len])?)
}
}
#[cfg(test)]
mod target_dispatch_tests {
use super::legacy_quant_gemm_allowed;
#[test]
fn legacy_quant_gemm_arch_policy_honors_the_escape_hatch() {
assert!(legacy_quant_gemm_allowed(false, false, false));
assert!(!legacy_quant_gemm_allowed(false, false, true));
assert!(!legacy_quant_gemm_allowed(true, false, false));
assert!(!legacy_quant_gemm_allowed(true, false, true));
assert!(legacy_quant_gemm_allowed(true, true, false));
assert!(!legacy_quant_gemm_allowed(true, true, true));
}
#[cfg(all(memra_portable_cuda, not(memra_hopper_mma)))]
#[test]
fn portable_build_disables_legacy_quant_gemm_without_an_env_override() {
assert!(!legacy_quant_gemm_allowed(
cfg!(memra_portable_cuda),
cfg!(memra_hopper_mma),
false
));
}
#[cfg(memra_hopper_mma)]
#[test]
fn hopper_mma_build_re_admits_legacy_quant_gemm() {
assert!(legacy_quant_gemm_allowed(
cfg!(memra_portable_cuda),
cfg!(memra_hopper_mma),
false
));
assert!(super::portable_mma_gated() == false);
}
}
impl memra_kv::KvDev for Engine {
fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
Engine::zeros(self, n)
}
fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
Engine::uninit(self, n)
}
fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
Engine::alloc_u8(self, n)
}
fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>> {
Engine::htod_i32(self, v)
}
fn clone_dtod(
&self,
src: &CudaSlice<f32>,
) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
Engine::clone_dtod(self, src)
}
fn copy_into(
&self,
dst: &mut CudaSlice<f32>,
off: usize,
src: &CudaSlice<f32>,
len: usize,
) -> Result<(), Box<dyn std::error::Error>> {
Engine::copy_into(self, dst, off, src, len)
}
fn set_i32_one(
&self,
d: &mut CudaSlice<i32>,
v: i32,
) -> Result<(), Box<dyn std::error::Error>> {
Engine::set_i32_one(self, d, v)
}
}