use core::ffi::c_int;
use std::ffi::c_void;
use cudarc::cublaslt::{result, sys};
use cudarc::driver::sys::CUdeviceptr;
use onnx_runtime_ep_api::{EpError, Result};
use crate::error::cublas_err;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum GemmDtype {
F32,
F16,
Bf16,
}
impl GemmDtype {
fn data_type(self) -> sys::cudaDataType {
match self {
GemmDtype::F32 => sys::cudaDataType_t::CUDA_R_32F,
GemmDtype::F16 => sys::cudaDataType_t::CUDA_R_16F,
GemmDtype::Bf16 => sys::cudaDataType_t::CUDA_R_16BF,
}
}
fn compute_type(self) -> sys::cublasComputeType_t {
sys::cublasComputeType_t::CUBLAS_COMPUTE_32F
}
}
#[derive(Debug)]
pub struct CublasLt {
handle: sys::cublasLtHandle_t,
}
unsafe impl Send for CublasLt {}
unsafe impl Sync for CublasLt {}
impl CublasLt {
pub fn new() -> Result<Self> {
let handle = result::create_handle().map_err(|e| cublas_err("cublasLtCreate", e))?;
Ok(Self { handle })
}
}
impl Drop for CublasLt {
fn drop(&mut self) {
if !self.handle.is_null() {
unsafe {
let _ = result::destroy_handle(self.handle);
}
self.handle = std::ptr::null_mut();
}
}
}
struct MatrixLayout(sys::cublasLtMatrixLayout_t);
impl MatrixLayout {
fn new(dtype: sys::cudaDataType, rows: u64, cols: u64, ld: i64) -> Result<Self> {
let h = result::create_matrix_layout(dtype, rows, cols, ld)
.map_err(|e| cublas_err("cublasLtMatrixLayoutCreate", e))?;
Ok(Self(h))
}
fn set_batch(&self, count: c_int, stride: i64) -> Result<()> {
unsafe {
result::set_matrix_layout_attribute(
self.0,
sys::cublasLtMatrixLayoutAttribute_t::CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT,
(&count) as *const c_int as *const c_void,
std::mem::size_of::<c_int>(),
)
.map_err(|e| cublas_err("set BATCH_COUNT", e))?;
result::set_matrix_layout_attribute(
self.0,
sys::cublasLtMatrixLayoutAttribute_t::CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
(&stride) as *const i64 as *const c_void,
std::mem::size_of::<i64>(),
)
.map_err(|e| cublas_err("set STRIDED_BATCH_OFFSET", e))?;
}
Ok(())
}
}
impl Drop for MatrixLayout {
fn drop(&mut self) {
unsafe {
let _ = result::destroy_matrix_layout(self.0);
}
}
}
struct MatmulDesc(sys::cublasLtMatmulDesc_t);
impl MatmulDesc {
fn new(compute: sys::cublasComputeType_t, scale: sys::cudaDataType) -> Result<Self> {
let h = result::create_matmul_desc(compute, scale)
.map_err(|e| cublas_err("cublasLtMatmulDescCreate", e))?;
Ok(Self(h))
}
fn set_epilogue(&self, epilogue: GemmEpilogue) -> Result<()> {
let kind = epilogue.kind.as_cublas();
let bias = epilogue.bias;
unsafe {
result::set_matmul_desc_attribute(
self.0,
sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_BIAS_POINTER,
(&bias) as *const CUdeviceptr as *const c_void,
std::mem::size_of::<CUdeviceptr>(),
)
.map_err(|e| cublas_err("set MATMUL_DESC_BIAS_POINTER", e))?;
result::set_matmul_desc_attribute(
self.0,
sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_EPILOGUE,
(&kind) as *const sys::cublasLtEpilogue_t as *const c_void,
std::mem::size_of::<sys::cublasLtEpilogue_t>(),
)
.map_err(|e| cublas_err("set MATMUL_DESC_EPILOGUE", e))
}
}
}
impl Drop for MatmulDesc {
fn drop(&mut self) {
unsafe {
let _ = result::destroy_matmul_desc(self.0);
}
}
}
struct MatmulPref(sys::cublasLtMatmulPreference_t);
impl MatmulPref {
fn new(workspace_bytes: usize) -> Result<Self> {
let h = result::create_matmul_pref()
.map_err(|e| cublas_err("cublasLtMatmulPreferenceCreate", e))?;
unsafe {
result::set_matmul_pref_attribute(
h,
sys::cublasLtMatmulPreferenceAttributes_t::CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
(&workspace_bytes) as *const usize as *const c_void,
std::mem::size_of::<usize>(),
)
.map_err(|e| cublas_err("set MAX_WORKSPACE_BYTES", e))?;
}
Ok(Self(h))
}
}
impl Drop for MatmulPref {
fn drop(&mut self) {
unsafe {
let _ = result::destroy_matmul_pref(self.0);
}
}
}
pub struct GemmParams {
pub dtype: GemmDtype,
pub a: CUdeviceptr,
pub b: CUdeviceptr,
pub c: CUdeviceptr,
pub m: usize,
pub k: usize,
pub n: usize,
pub batch: usize,
pub a_batch_stride: usize,
pub b_batch_stride: usize,
pub epilogue: Option<GemmEpilogue>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum GemmEpilogueKind {
Bias,
ReluBias,
GeluBias,
}
impl GemmEpilogueKind {
fn as_cublas(self) -> sys::cublasLtEpilogue_t {
match self {
Self::Bias => sys::cublasLtEpilogue_t::CUBLASLT_EPILOGUE_BIAS,
Self::ReluBias => sys::cublasLtEpilogue_t::CUBLASLT_EPILOGUE_RELU_BIAS,
Self::GeluBias => sys::cublasLtEpilogue_t::CUBLASLT_EPILOGUE_GELU_BIAS,
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct GemmEpilogue {
pub kind: GemmEpilogueKind,
pub bias: CUdeviceptr,
}
pub const WORKSPACE_BYTES: usize = 32 * 1024 * 1024;
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm(
handle: &CublasLt,
stream: cudarc::driver::sys::CUstream,
p: &GemmParams,
workspace: CUdeviceptr,
workspace_bytes: usize,
) -> Result<()> {
if p.m == 0 || p.n == 0 || p.k == 0 || p.batch == 0 {
return Err(EpError::KernelFailed(format!(
"cuda_ep MatMul: degenerate GEMM dims M={} K={} N={} batch={}",
p.m, p.n, p.k, p.batch
)));
}
let dt = p.dtype.data_type();
let (m, n, k) = (p.n as u64, p.m as u64, p.k as u64); let (lda, ldb, ldc) = (p.n as i64, p.k as i64, p.n as i64);
let a_layout = MatrixLayout::new(dt, m, k, lda)?;
let b_layout = MatrixLayout::new(dt, k, n, ldb)?;
let c_layout = MatrixLayout::new(dt, m, n, ldc)?;
if p.batch > 1 {
let count = i32::try_from(p.batch).map_err(|_| {
EpError::KernelFailed(format!("cuda_ep MatMul: batch {} exceeds i32", p.batch))
})?;
a_layout.set_batch(count, p.b_batch_stride as i64)?; b_layout.set_batch(count, p.a_batch_stride as i64)?; c_layout.set_batch(count, (p.m * p.n) as i64)?; }
let desc = MatmulDesc::new(p.dtype.compute_type(), sys::cudaDataType_t::CUDA_R_32F)?;
if let Some(epilogue) = p.epilogue {
desc.set_epilogue(epilogue)?;
}
let pref = MatmulPref::new(workspace_bytes)?;
let heuristic = unsafe {
result::get_matmul_algo_heuristic(
handle.handle,
desc.0,
a_layout.0,
b_layout.0,
c_layout.0,
c_layout.0,
pref.0,
)
}
.map_err(|e| {
cublas_err(
&format!(
"no cuBLASLt algorithm for MatMul M={} K={} N={} batch={} dtype={:?}",
p.m, p.k, p.n, p.batch, p.dtype
),
e,
)
})?;
let alpha: f32 = 1.0;
let beta: f32 = 0.0;
unsafe {
result::matmul(
handle.handle,
desc.0,
(&alpha) as *const f32 as *const c_void,
(&beta) as *const f32 as *const c_void,
p.b as *const c_void, a_layout.0,
p.a as *const c_void, b_layout.0,
p.c as *const c_void, c_layout.0,
p.c as *mut c_void, c_layout.0,
(&heuristic.algo) as *const sys::cublasLtMatmulAlgo_t,
workspace as *mut c_void,
workspace_bytes,
stream as sys::cudaStream_t,
)
}
.map_err(|e| cublas_err("cublasLtMatmul", e))?;
Ok(())
}
pub struct GemmEx {
pub dtype: GemmDtype,
pub transa: bool,
pub transb: bool,
pub m: usize,
pub n: usize,
pub k: usize,
pub alpha: f32,
pub beta: f32,
pub a: CUdeviceptr,
pub lda: usize,
pub b: CUdeviceptr,
pub ldb: usize,
pub c: CUdeviceptr,
pub ldc: usize,
pub epilogue: Option<GemmEpilogue>,
}
const CUBLAS_OP_N: i32 = 0;
const CUBLAS_OP_T: i32 = 1;
impl MatmulDesc {
fn set_transpose(
&self,
attr: sys::cublasLtMatmulDescAttributes_t,
transpose: bool,
) -> Result<()> {
let op: i32 = if transpose { CUBLAS_OP_T } else { CUBLAS_OP_N };
unsafe {
result::set_matmul_desc_attribute(
self.0,
attr,
(&op) as *const i32 as *const c_void,
std::mem::size_of::<i32>(),
)
.map_err(|e| cublas_err("set MATMUL_DESC_TRANS", e))
}
}
}
pub unsafe fn gemm_ex(
handle: &CublasLt,
stream: cudarc::driver::sys::CUstream,
p: &GemmEx,
workspace: CUdeviceptr,
workspace_bytes: usize,
) -> Result<()> {
if p.m == 0 || p.n == 0 || p.k == 0 {
return Err(EpError::KernelFailed(format!(
"cuda_ep attention GEMM: degenerate dims M={} N={} K={}",
p.m, p.n, p.k
)));
}
let dt = p.dtype.data_type();
let (a_rows, a_cols) = if p.transa {
(p.k as u64, p.m as u64)
} else {
(p.m as u64, p.k as u64)
};
let (b_rows, b_cols) = if p.transb {
(p.n as u64, p.k as u64)
} else {
(p.k as u64, p.n as u64)
};
let a_layout = MatrixLayout::new(dt, a_rows, a_cols, p.lda as i64)?;
let b_layout = MatrixLayout::new(dt, b_rows, b_cols, p.ldb as i64)?;
let c_layout = MatrixLayout::new(dt, p.m as u64, p.n as u64, p.ldc as i64)?;
let desc = MatmulDesc::new(p.dtype.compute_type(), sys::cudaDataType_t::CUDA_R_32F)?;
desc.set_transpose(
sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_TRANSA,
p.transa,
)?;
desc.set_transpose(
sys::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_TRANSB,
p.transb,
)?;
if let Some(epilogue) = p.epilogue {
desc.set_epilogue(epilogue)?;
}
let pref = MatmulPref::new(workspace_bytes)?;
let heuristic = unsafe {
result::get_matmul_algo_heuristic(
handle.handle,
desc.0,
a_layout.0,
b_layout.0,
c_layout.0,
c_layout.0,
pref.0,
)
}
.map_err(|e| {
cublas_err(
&format!(
"no cuBLASLt algorithm for attention GEMM M={} N={} K={} transa={} transb={} dtype={:?}",
p.m, p.n, p.k, p.transa, p.transb, p.dtype
),
e,
)
})?;
let alpha = p.alpha;
let beta = p.beta;
unsafe {
result::matmul(
handle.handle,
desc.0,
(&alpha) as *const f32 as *const c_void,
(&beta) as *const f32 as *const c_void,
p.a as *const c_void,
a_layout.0,
p.b as *const c_void,
b_layout.0,
p.c as *const c_void,
c_layout.0,
p.c as *mut c_void,
c_layout.0,
(&heuristic.algo) as *const sys::cublasLtMatmulAlgo_t,
workspace as *mut c_void,
workspace_bytes,
stream as sys::cudaStream_t,
)
}
.map_err(|e| cublas_err("cublasLtMatmul (attention)", e))?;
Ok(())
}