use std::ffi::c_void;
use std::mem::size_of;
use cudarc::cublaslt::result as lt;
use cudarc::cublaslt::sys;
use kime_tensor::{Error, Result};
fn err(e: lt::CublasError) -> Error {
Error::Device(format!("cublasLt: {e:?}"))
}
pub(crate) struct Handle(pub(crate) sys::cublasLtHandle_t);
unsafe impl Send for Handle {}
unsafe impl Sync for Handle {}
impl Handle {
pub(crate) fn new() -> Result<Self> {
lt::create_handle().map(Self).map_err(err)
}
}
impl Drop for Handle {
fn drop(&mut self) {
let _ = unsafe { lt::destroy_handle(self.0) };
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Ty {
F16,
F32,
}
impl Ty {
fn cuda(self) -> sys::cudaDataType {
match self {
Ty::F16 => sys::cudaDataType::CUDA_R_16F,
Ty::F32 => sys::cudaDataType::CUDA_R_32F,
}
}
}
pub(crate) struct Gemm {
desc: sys::cublasLtMatmulDesc_t,
a: sys::cublasLtMatrixLayout_t,
b: sys::cublasLtMatrixLayout_t,
c: sys::cublasLtMatrixLayout_t,
algo: sys::cublasLtMatmulAlgo_t,
beta: f32,
pub(crate) w: u64,
pub(crate) x: u64,
pub(crate) y: u64,
}
unsafe impl Send for Gemm {}
impl Gemm {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
h: &Handle,
(m, k, n): (usize, usize, usize),
ab: Ty,
c: Ty,
accumulate: bool,
workspace: usize,
(w, x, y): (u64, u64, u64),
) -> Result<Self> {
let desc =
lt::create_matmul_desc(sys::cublasComputeType_t::CUBLAS_COMPUTE_32F, Ty::F32.cuda())
.map_err(err)?;
let mut g = Self {
desc,
a: std::ptr::null_mut(),
b: std::ptr::null_mut(),
c: std::ptr::null_mut(),
algo: unsafe { std::mem::zeroed() },
beta: if accumulate { 1.0 } else { 0.0 },
w,
x,
y,
};
let op: i32 = 1;
unsafe {
lt::set_matmul_desc_attribute(
g.desc,
sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_TRANSA,
(&raw const op).cast::<c_void>(),
size_of::<i32>(),
)
.map_err(err)?;
}
let (m, k, n) = (m as u64, k as u64, n as u64);
g.a = lt::create_matrix_layout(ab.cuda(), k, n, k as i64).map_err(err)?;
g.b = lt::create_matrix_layout(ab.cuda(), k, m, k as i64).map_err(err)?;
g.c = lt::create_matrix_layout(c.cuda(), n, m, n as i64).map_err(err)?;
let pref = lt::create_matmul_pref().map_err(err)?;
let ws = workspace as u64;
let found = unsafe {
let set = lt::set_matmul_pref_attribute(
pref,
sys::cublasLtMatmulPreferenceAttributes_t::CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
(&raw const ws).cast::<c_void>(),
size_of::<u64>(),
);
let found = set.and_then(|()| {
lt::get_matmul_algo_heuristic(h.0, g.desc, g.a, g.b, g.c, g.c, pref)
});
let _ = lt::destroy_matmul_pref(pref);
found
};
g.algo = found.map_err(err)?.algo;
Ok(g)
}
pub(crate) unsafe fn run(
&self,
h: &Handle,
workspace: u64,
size: usize,
stream: sys::cudaStream_t,
) -> Result<()> {
let alpha = 1f32;
unsafe {
lt::matmul(
h.0,
self.desc,
(&raw const alpha).cast(),
(&raw const self.beta).cast(),
self.w as *const c_void,
self.a,
self.x as *const c_void,
self.b,
self.y as *const c_void,
self.c,
self.y as *mut c_void,
self.c,
&raw const self.algo,
workspace as *mut c_void,
size,
stream,
)
.map_err(err)
}
}
}
impl Drop for Gemm {
fn drop(&mut self) {
unsafe {
for l in [self.a, self.b, self.c] {
if !l.is_null() {
let _ = lt::destroy_matrix_layout(l);
}
}
let _ = lt::destroy_matmul_desc(self.desc);
}
}
}