use std::ffi::c_void;
use std::sync::Arc;
use cudarc::driver::{LaunchConfig, PushKernelArg};
use onnx_runtime_ep_api::{EpError, Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::{DataType, Node};
use crate::blas::{GemmDtype, GemmEx, WORKSPACE_BYTES, gemm_ex};
use crate::error::{driver_err, not_implemented};
use crate::runtime::{CudaRuntime, cuptr};
const BIAS_SRC: &str = r#"
extern "C" __global__ void gemm_bias_f32(
float* y, // [m, n] row-major, in/out
const float* c, // broadcastable bias
const int m,
const int n,
const int c_row_stride, // 0 if C broadcasts over rows, else n_c
const int c_col_stride, // 0 if C broadcasts over cols, else 1
const float beta)
{
const long total = (long)m * n;
for (long idx = (long)blockIdx.x * blockDim.x + threadIdx.x;
idx < total; idx += (long)gridDim.x * blockDim.x) {
const int row = (int)(idx / n);
const int col = (int)(idx % n);
const float cv = c[(long)row * c_row_stride + (long)col * c_col_stride];
y[idx] += beta * cv;
}
}
"#;
const BIAS_MODULE: &str = "gemm_bias_f32";
const BIAS_ENTRY: &str = "gemm_bias_f32";
const BIAS_BLOCK: u32 = 256;
pub struct GemmFactory {
pub runtime: Arc<CudaRuntime>,
}
impl KernelFactory for GemmFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let alpha = node.attr("alpha").and_then(|a| a.as_float()).unwrap_or(1.0);
let beta = node.attr("beta").and_then(|a| a.as_float()).unwrap_or(1.0);
let trans_a = node.attr("transA").and_then(|a| a.as_int()).unwrap_or(0) != 0;
let trans_b = node.attr("transB").and_then(|a| a.as_int()).unwrap_or(0) != 0;
Ok(Box::new(GemmKernel {
runtime: self.runtime.clone(),
alpha,
beta,
trans_a,
trans_b,
}))
}
}
#[derive(Debug)]
pub struct GemmKernel {
runtime: Arc<CudaRuntime>,
alpha: f32,
beta: f32,
trans_a: bool,
trans_b: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) struct GemmPlan {
pub(super) m: usize,
pub(super) k: usize,
pub(super) n: usize,
pub(super) transa: bool,
pub(super) transb: bool,
pub(super) ldb_operand: usize,
pub(super) lda_operand: usize,
}
pub(super) fn plan_gemm(
a: &[usize],
b: &[usize],
trans_a: bool,
trans_b: bool,
) -> Result<GemmPlan> {
if a.len() != 2 || b.len() != 2 {
return Err(not_implemented(format!(
"Gemm with operand ranks {}D x {}D (Gemm requires 2-D A and B)",
a.len(),
b.len()
)));
}
let (ra, ca) = (a[0], a[1]);
let (rb, cb) = (b[0], b[1]);
let (m, ka) = if trans_a { (ca, ra) } else { (ra, ca) };
let (kb, n) = if trans_b { (cb, rb) } else { (rb, cb) };
if ka != kb {
return Err(EpError::KernelFailed(format!(
"cuda_ep Gemm: inner dimensions disagree — A' is [{m},{ka}] but B' is [{kb},{n}] \
(A {a:?} transA={trans_a}, B {b:?} transB={trans_b})"
)));
}
Ok(GemmPlan {
m,
k: ka,
n,
transa: trans_b,
transb: trans_a,
ldb_operand: cb,
lda_operand: ca,
})
}
fn bias_strides(c: &[usize], m: usize, n: usize) -> Result<(i32, i32)> {
let (cr, cc) = match c.len() {
0 => (1usize, 1usize),
1 => (1usize, c[0]),
2 => (c[0], c[1]),
_ => {
return Err(not_implemented(format!(
"Gemm bias C rank {} (bias must broadcast to [M,N]; rank <= 2)",
c.len()
)));
}
};
if (cr != 1 && cr != m) || (cc != 1 && cc != n) {
return Err(EpError::KernelFailed(format!(
"cuda_ep Gemm: bias C {c:?} does not broadcast to [M={m}, N={n}] \
(each C dim must be 1, M, or N)"
)));
}
let row_stride = if cr == 1 { 0 } else { cc as i32 };
let col_stride = if cc == 1 { 0 } else { 1 };
Ok((row_stride, col_stride))
}
impl GemmKernel {
fn run(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
if !(2..=3).contains(&inputs.len()) || outputs.len() != 1 {
return Err(EpError::KernelFailed(format!(
"cuda_ep Gemm: expected 2 inputs (A,B) or 3 (A,B,C) and 1 output, \
got {} inputs and {} outputs",
inputs.len(),
outputs.len()
)));
}
let a = &inputs[0];
let b = &inputs[1];
let bias = inputs.get(2).filter(|c| !c.is_absent());
for (name, dt) in [("A", a.dtype), ("B", b.dtype), ("Y", outputs[0].dtype)] {
if dt != DataType::Float32 {
return Err(not_implemented(format!(
"Gemm with {name} dtype {dt:?} (this slice is f32-only; f16/bf16 pending)"
)));
}
}
for (name, contiguous) in [
("A", a.is_contiguous()),
("B", b.is_contiguous()),
("Y", outputs[0].is_contiguous()),
] {
if !contiguous {
return Err(not_implemented(format!(
"Gemm with a non-contiguous (strided) {name}; materialise it (insert a copy) \
before the Gemm"
)));
}
}
let plan = plan_gemm(a.shape, b.shape, self.trans_a, self.trans_b)?;
let expected_out = plan.m * plan.n;
if outputs[0].numel() != expected_out {
return Err(EpError::KernelFailed(format!(
"cuda_ep Gemm: output has {} elements, expected {} ([M={},N={}])",
outputs[0].numel(),
expected_out,
plan.m,
plan.n
)));
}
let bias_plan = match bias {
None => None,
Some(c) => {
if c.dtype != DataType::Float32 {
return Err(not_implemented(format!(
"Gemm bias C dtype {:?} (this slice is f32-only)",
c.dtype
)));
}
if !c.is_contiguous() {
return Err(not_implemented(
"Gemm with a non-contiguous (strided) bias C; materialise it first",
));
}
let (rs, cs) = bias_strides(c.shape, plan.m, plan.n)?;
Some((c, rs, cs))
}
};
crate::trace::record_kernel_metrics(inputs, outputs, || {
let mut flops = (plan.m as u64)
.saturating_mul(plan.n as u64)
.saturating_mul(plan.k as u64)
.saturating_mul(2);
if bias_plan.is_some() && self.beta != 0.0 {
flops = flops.saturating_add(
(plan.m as u64)
.saturating_mul(plan.n as u64)
.saturating_mul(2),
);
}
flops
});
let a_ptr = cuptr(a.data_ptr::<u8>() as *const c_void);
let b_ptr = cuptr(b.data_ptr::<u8>() as *const c_void);
let y_ptr = cuptr(outputs[0].data_ptr_mut::<u8>() as *const c_void);
let workspace = self.runtime.alloc_raw(WORKSPACE_BYTES)?;
let params = GemmEx {
dtype: GemmDtype::F32,
transa: plan.transa,
transb: plan.transb,
m: plan.n,
n: plan.m,
k: plan.k,
alpha: self.alpha,
beta: 0.0,
a: b_ptr,
lda: plan.ldb_operand,
b: a_ptr,
ldb: plan.lda_operand,
c: y_ptr,
ldc: plan.n,
epilogue: None,
};
let gemm_res = unsafe {
gemm_ex(
self.runtime.blas(),
self.runtime.stream_ptr(),
¶ms,
workspace,
WORKSPACE_BYTES,
)
};
let bias_res = gemm_res.and_then(|()| {
if let Some((c, rs, cs)) = bias_plan {
if self.beta != 0.0 {
self.apply_bias(y_ptr, c, plan.m, plan.n, rs, cs)
} else {
Ok(())
}
} else {
Ok(())
}
});
let synced = bias_res.and_then(|()| self.runtime.synchronize());
let free = unsafe { self.runtime.free_raw(workspace) };
synced.and(free)
}
fn apply_bias(
&self,
y_ptr: cudarc::driver::sys::CUdeviceptr,
c: &TensorView,
m: usize,
n: usize,
row_stride: i32,
col_stride: i32,
) -> Result<()> {
let c_ptr = cuptr(c.data_ptr::<u8>() as *const c_void);
let total = m * n;
let (m_i, n_i) = (
i32::try_from(m)
.map_err(|_| EpError::KernelFailed(format!("cuda_ep Gemm: M={m} exceeds i32")))?,
i32::try_from(n)
.map_err(|_| EpError::KernelFailed(format!("cuda_ep Gemm: N={n} exceeds i32")))?,
);
let beta = self.beta;
let func = self
.runtime
.nvrtc_function(BIAS_MODULE, BIAS_SRC, BIAS_ENTRY)?;
let blocks = total.div_ceil(BIAS_BLOCK as usize).clamp(1, 65_535) as u32;
let cfg = LaunchConfig {
grid_dim: (blocks, 1, 1),
block_dim: (BIAS_BLOCK, 1, 1),
shared_mem_bytes: 0,
};
let stream = self.runtime.stream();
let mut builder = stream.launch_builder(&func);
builder
.arg(&y_ptr)
.arg(&c_ptr)
.arg(&m_i)
.arg(&n_i)
.arg(&row_stride)
.arg(&col_stride)
.arg(&beta);
unsafe { builder.launch(cfg) }
.map(|_| ())
.map_err(|e| driver_err("launch gemm_bias_f32", e))
}
}
impl Kernel for GemmKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.run(inputs, outputs)
}
fn supports_strided_input(&self, _idx: usize) -> bool {
false
}
fn capture_support(&self) -> onnx_runtime_ep_api::CaptureSupport {
onnx_runtime_ep_api::CaptureSupport::unsupported(
"per-call workspace allocation/free is not capturable",
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn plan_no_transpose() {
let p = plan_gemm(&[2, 3], &[3, 4], false, false).unwrap();
assert_eq!((p.m, p.k, p.n), (2, 3, 4));
assert!(!p.transa && !p.transb);
assert_eq!(p.ldb_operand, 4);
assert_eq!(p.lda_operand, 3);
}
#[test]
fn plan_trans_a() {
let p = plan_gemm(&[3, 2], &[3, 4], true, false).unwrap();
assert_eq!((p.m, p.k, p.n), (2, 3, 4));
assert!(p.transb, "transA maps to cuBLASLt transb");
assert!(!p.transa);
}
#[test]
fn plan_trans_b() {
let p = plan_gemm(&[2, 3], &[4, 3], false, true).unwrap();
assert_eq!((p.m, p.k, p.n), (2, 3, 4));
assert!(p.transa, "transB maps to cuBLASLt transa");
assert!(!p.transb);
}
#[test]
fn plan_inner_mismatch_is_plain_error() {
let e = plan_gemm(&[2, 3], &[5, 4], false, false).unwrap_err();
let msg = format!("{e}");
assert!(msg.contains("inner dimensions disagree"), "{msg}");
assert!(!msg.contains("not implemented"), "{msg}");
}
#[test]
fn plan_rejects_non_2d() {
let e = plan_gemm(&[2, 3, 4], &[4, 5], false, false).unwrap_err();
assert!(format!("{e}").contains("requires 2-D"), "{e}");
}
#[test]
fn bias_strides_scalar_and_vectors() {
assert_eq!(bias_strides(&[], 2, 4).unwrap(), (0, 0));
assert_eq!(bias_strides(&[4], 2, 4).unwrap(), (0, 1));
assert_eq!(bias_strides(&[1], 2, 4).unwrap(), (0, 0));
assert_eq!(bias_strides(&[2, 4], 2, 4).unwrap(), (4, 1));
assert_eq!(bias_strides(&[2, 1], 2, 4).unwrap(), (1, 0));
assert_eq!(bias_strides(&[1, 4], 2, 4).unwrap(), (0, 1));
}
#[test]
fn bias_strides_rejects_non_broadcastable() {
let e = bias_strides(&[3, 4], 2, 4).unwrap_err();
assert!(format!("{e}").contains("does not broadcast"), "{e}");
}
}