use crate::error::{Error, Result};
use std::sync::Arc;
use cudarc::driver::{CudaContext, CudaFunction, CudaModule};
pub(crate) type CUptr = cudarc::driver::sys::CUdeviceptr;
pub(crate) struct HalfKernel {
pub bf16: CudaFunction,
pub f16: CudaFunction,
}
impl HalfKernel {
pub(crate) fn get(&self, dt: crate::Dtype) -> &CudaFunction {
match dt {
crate::Dtype::Bf16 => &self.bf16,
crate::Dtype::F16 => &self.f16,
crate::Dtype::F32 => unreachable!("HalfKernel has no f32 variant"),
}
}
}
pub(crate) struct Kernels {
pub sgemm_nn: CudaFunction,
pub sgemm_tn: CudaFunction,
pub sgemm_nt: CudaFunction,
pub sgemm_nn_slim: CudaFunction,
pub sgemm_tn_slim: CudaFunction,
pub sgemm_nt_slim: CudaFunction,
pub sgemm_nn_gemv: CudaFunction,
pub sgemm_tn_gemv: CudaFunction,
pub sgemm_nt_gemv: CudaFunction,
pub sgemm_nn_ultra_thin: CudaFunction,
pub sgemm_nn_narrow: CudaFunction,
pub sgemm_nn_narrow_small: CudaFunction,
pub sgemm_tn_narrow: CudaFunction,
pub sgemm_nt_narrow: CudaFunction,
pub sgemm_nn_splitk32_partial: CudaFunction,
pub sgemm_nn_splitk_slim_partial: CudaFunction,
pub sgemm_tn_splitm_partial: CudaFunction,
pub sgemm_splitk_reduce: CudaFunction,
pub sgemm_splitm_reduce: CudaFunction,
pub sgemm_dx_col_gemv: CudaFunction,
pub sgemm_transpose_f32_2d: CudaFunction,
pub sgemm_nn_gemv_typed: HalfKernel,
pub sgemm_tn_gemv_typed: HalfKernel,
pub sgemm_nt_gemv_typed: HalfKernel,
pub sgemm_nn_ultra_thin_typed: HalfKernel,
pub sgemm_nn_narrow_typed: HalfKernel,
pub sgemm_nn_narrow_small_typed: HalfKernel,
pub sgemm_tn_narrow_typed: HalfKernel,
pub sgemm_nt_narrow_typed: HalfKernel,
pub sgemm_nn_big_typed: HalfKernel,
pub sgemm_tn_big_typed: HalfKernel,
pub sgemm_nt_big_typed: HalfKernel,
pub sgemm_nn_tc_typed: HalfKernel,
pub sgemm_tn_tc_typed: HalfKernel,
pub sgemm_nt_tc_typed: HalfKernel,
pub sgemm_nn_tc64_typed: HalfKernel,
pub sgemm_tn_tc64_typed: HalfKernel,
pub sgemm_nt_tc64_typed: HalfKernel,
pub cast_f32_to_bf16: CudaFunction,
pub cast_bf16_to_f32: CudaFunction,
pub cast_f32_to_f16: CudaFunction,
pub cast_f16_to_f32: CudaFunction,
splitk_scratch: std::sync::OnceLock<(cudarc::driver::CudaSlice<f32>, CUptr)>,
transpose_scratch: std::sync::OnceLock<(cudarc::driver::CudaSlice<f32>, CUptr)>,
_module: Arc<CudaModule>,
}
pub(crate) fn nvrtc_arch(cc: (u32, u32)) -> &'static str {
match cc {
(12, _) => "sm_120",
(10, _) => "sm_100",
(9, _) => "sm_90",
(8, 9) => "sm_89",
(8, 6) => "sm_86",
(8, 0) => "sm_80",
_ => {
if cc.0 > 12 {
"sm_120"
} else {
"sm_80"
}
}
}
}
fn cuda_include_paths() -> Vec<String> {
let mut candidates: Vec<String> = Vec::new();
for var in ["CUDA_HOME", "CUDA_PATH", "CUDA_ROOT"] {
if let Ok(p) = std::env::var(var) {
candidates.push(format!("{p}/include"));
}
}
for std_path in [
"/usr/local/cuda/include",
"/usr/local/cuda-13.2/include",
"/usr/local/cuda-12.8/include",
"/usr/local/cuda-12.6/include",
"/usr/local/cuda-12.4/include",
"/opt/cuda/include",
] {
candidates.push(std_path.to_string());
}
candidates
.into_iter()
.filter(|p| std::path::Path::new(p).join("cuda_fp16.h").exists())
.collect()
}
fn nvrtc_version() -> (i32, i32) {
let mut major: core::ffi::c_int = 0;
let mut minor: core::ffi::c_int = 0;
let rc = unsafe { cudarc::nvrtc::sys::nvrtcVersion(&mut major, &mut minor) };
if rc == cudarc::nvrtc::sys::nvrtcResult::NVRTC_SUCCESS {
(major, minor)
} else {
(0, 0)
}
}
fn kernel_cache_dir() -> Option<std::path::PathBuf> {
match std::env::var("SGEMM_BI_KERNEL_CACHE") {
Ok(v) if matches!(v.trim(), "0" | "off" | "OFF") => None,
Ok(v) if !v.trim().is_empty() => Some(std::path::PathBuf::from(v.trim())),
_ => {
let base = std::env::var("XDG_CACHE_HOME")
.map(std::path::PathBuf::from)
.or_else(|_| {
std::env::var("HOME").map(|h| std::path::PathBuf::from(h).join(".cache"))
})
.unwrap_or_else(|_| std::env::temp_dir());
Some(base.join("sgemm-bi").join("kernels"))
}
}
}
fn cache_key(material: &str) -> String {
let fnv = |seed: u64| -> u64 {
let mut h = seed;
for b in material.as_bytes() {
h ^= u64::from(*b);
h = h.wrapping_mul(0x0000_0100_0000_01b3);
}
h
};
format!(
"{:016x}{:016x}",
fnv(0xcbf2_9ce4_8422_2325),
fnv(0x6c62_272e_07bb_0142)
)
}
impl Kernels {
pub(crate) fn compile(
ctx: &Arc<CudaContext>,
_stream: &Arc<cudarc::driver::CudaStream>,
arch: &'static str,
) -> Result<Self> {
let sources = [
include_str!("../kernels/prelude.cuh"),
include_str!("../kernels/casts.cu"),
include_str!("../kernels/sgemm_bi.cu"),
];
let combined: String = sources.join("\n");
let group_m: usize = match arch {
"sm_80" | "sm_86" | "sm_87" => 8,
_ => 16,
};
let option_strings = vec![
"--fmad=true".to_string(),
"--extra-device-vectorization".to_string(),
format!("-DSGB_GROUP_M={group_m}"),
];
let opts = cudarc::nvrtc::CompileOptions {
arch: Some(arch),
options: option_strings.clone(),
include_paths: cuda_include_paths(),
..Default::default()
};
let (nv_major, nv_minor) = nvrtc_version();
let key = cache_key(&format!(
"{combined}\u{1f}{arch}\u{1f}{option_strings:?}\u{1f}nvrtc{nv_major}.{nv_minor}"
));
let cache_path = kernel_cache_dir().map(|d| d.join(format!("sgemm-bi-{key}.ptx")));
let mut module = None;
if let Some(path) = &cache_path
&& let Ok(src) = std::fs::read_to_string(path)
{
match ctx.load_module(cudarc::nvrtc::Ptx::from_src(src)) {
Ok(m) => module = Some(m),
Err(_) => {
let _ = std::fs::remove_file(path);
}
}
}
let module = match module {
Some(m) => m,
None => {
let ptx = cudarc::nvrtc::compile_ptx_with_opts(combined, opts)
.map_err(|e| Error::Cuda(format!("NVRTC compile failed: {e:?}")))?;
if let Some(path) = &cache_path
&& let Some(dir) = path.parent()
&& std::fs::create_dir_all(dir).is_ok()
{
let tmp = path.with_extension(format!("tmp-{}", std::process::id()));
if std::fs::write(&tmp, ptx.to_src()).is_ok() {
let _ = std::fs::rename(&tmp, path);
}
}
ctx.load_module(ptx)
.map_err(|e| Error::Cuda(format!("module load failed: {e:?}")))?
}
};
let get = |name: &str| -> Result<CudaFunction> {
module
.load_function(name)
.map_err(|e| Error::Cuda(format!("load_function({name}): {e:?}")))
};
let load_half = |base: &str| -> Result<HalfKernel> {
Ok(HalfKernel {
bf16: get(&format!("{base}_bf16"))?,
f16: get(&format!("{base}_f16"))?,
})
};
let set_dynsmem = |f: &CudaFunction, bytes: i32, name: &str| -> Result<()> {
f.set_attribute(
cudarc::driver::sys::CUfunction_attribute_enum::CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
bytes,
)
.map_err(|e| Error::Cuda(format!("set MAX_DYNAMIC_SHARED for {name}: {e:?}")))
};
let load_half_dynsmem = |base: &str, bytes: i32| -> Result<HalfKernel> {
let k = load_half(base)?;
set_dynsmem(&k.bf16, bytes, base)?;
set_dynsmem(&k.f16, bytes, base)?;
Ok(k)
};
let get_dynsmem = |name: &str, bytes: i32| -> Result<CudaFunction> {
let f = get(name)?;
set_dynsmem(&f, bytes, name)?;
Ok(f)
};
Ok(Self {
sgemm_nn: get_dynsmem("sgemm_bi_nn", 34 * 1024)?,
sgemm_tn: get_dynsmem("sgemm_bi_tn", 34 * 1024)?,
sgemm_nt: get_dynsmem("sgemm_bi_nt", 34 * 1024)?,
sgemm_nn_slim: get("sgemm_bi_nn_slim")?,
sgemm_tn_slim: get("sgemm_bi_tn_slim")?,
sgemm_nt_slim: get("sgemm_bi_nt_slim")?,
sgemm_nn_gemv: get("sgemm_bi_nn_gemv")?,
sgemm_tn_gemv: get("sgemm_bi_tn_gemv")?,
sgemm_nt_gemv: get("sgemm_bi_nt_gemv")?,
sgemm_nn_ultra_thin: get("sgemm_bi_nn_ultra_thin")?,
sgemm_nn_narrow: get("sgemm_bi_nn_narrow")?,
sgemm_nn_narrow_small: get("sgemm_bi_nn_narrow_small")?,
sgemm_tn_narrow: get("sgemm_bi_tn_narrow")?,
sgemm_nt_narrow: get("sgemm_bi_nt_narrow")?,
sgemm_nn_splitk32_partial: get("sgemm_bi_nn_splitk32_partial")?,
sgemm_nn_splitk_slim_partial: get("sgemm_bi_nn_splitk_slim_partial")?,
sgemm_tn_splitm_partial: get_dynsmem("sgemm_bi_tn_splitm_partial", 34 * 1024)?,
sgemm_splitk_reduce: get("sgemm_bi_splitk_reduce")?,
sgemm_splitm_reduce: get("sgemm_bi_splitm_reduce")?,
sgemm_dx_col_gemv: get("sgemm_bi_dx_col_gemv")?,
sgemm_transpose_f32_2d: get("sgemm_transpose_f32_2d")?,
sgemm_nn_gemv_typed: load_half("sgemm_bi_nn_gemv")?,
sgemm_tn_gemv_typed: load_half("sgemm_bi_tn_gemv")?,
sgemm_nt_gemv_typed: load_half("sgemm_bi_nt_gemv")?,
sgemm_nn_ultra_thin_typed: load_half("sgemm_bi_nn_ultra_thin")?,
sgemm_nn_narrow_typed: load_half("sgemm_bi_nn_narrow")?,
sgemm_nn_narrow_small_typed: load_half("sgemm_bi_nn_narrow_small")?,
sgemm_tn_narrow_typed: load_half("sgemm_bi_tn_narrow")?,
sgemm_nt_narrow_typed: load_half("sgemm_bi_nt_narrow")?,
sgemm_nn_big_typed: load_half_dynsmem("sgemm_bi_nn_big", 34 * 1024)?,
sgemm_tn_big_typed: load_half_dynsmem("sgemm_bi_tn_big", 34 * 1024)?,
sgemm_nt_big_typed: load_half_dynsmem("sgemm_bi_nt_big", 34 * 1024)?,
sgemm_nn_tc_typed: load_half_dynsmem("sgemm_bi_nn_tc", 75_776)?,
sgemm_tn_tc_typed: load_half_dynsmem("sgemm_bi_tn_tc", 75_776)?,
sgemm_nt_tc_typed: load_half_dynsmem("sgemm_bi_nt_tc", 75_776)?,
sgemm_nn_tc64_typed: load_half("sgemm_bi_nn_tc64")?,
sgemm_tn_tc64_typed: load_half("sgemm_bi_tn_tc64")?,
sgemm_nt_tc64_typed: load_half("sgemm_bi_nt_tc64")?,
cast_f32_to_bf16: get("sgb_cast_f32_to_bf16")?,
cast_bf16_to_f32: get("sgb_cast_bf16_to_f32")?,
cast_f32_to_f16: get("sgb_cast_f32_to_f16")?,
cast_f16_to_f32: get("sgb_cast_f16_to_f32")?,
splitk_scratch: std::sync::OnceLock::new(),
transpose_scratch: std::sync::OnceLock::new(),
_module: module,
})
}
pub(crate) fn splitk_ptr(&self, stream: &Arc<cudarc::driver::CudaStream>) -> Result<CUptr> {
if self.splitk_scratch.get().is_none() {
let buf = stream
.alloc_zeros::<f32>(1 << 23)
.map_err(|e| Error::Cuda(format!("splitk_scratch alloc: {e:?}")))?;
let ptr = {
use cudarc::driver::DevicePtr;
let (p, _g) = buf.device_ptr(stream);
p
};
let _ = self.splitk_scratch.set((buf, ptr));
}
self.splitk_scratch
.get()
.map(|(_, p)| *p)
.ok_or_else(|| Error::Cuda("splitk_scratch cell empty after init".into()))
}
pub(crate) fn transpose_ptr(&self, stream: &Arc<cudarc::driver::CudaStream>) -> Result<CUptr> {
if self.transpose_scratch.get().is_none() {
let buf = stream
.alloc_zeros::<f32>(1 << 22)
.map_err(|e| Error::Cuda(format!("transpose_scratch alloc: {e:?}")))?;
let ptr = {
use cudarc::driver::DevicePtr;
let (p, _g) = buf.device_ptr(stream);
p
};
let _ = self.transpose_scratch.set((buf, ptr));
}
self.transpose_scratch
.get()
.map(|(_, p)| *p)
.ok_or_else(|| Error::Cuda("transpose_scratch cell empty after init".into()))
}
}