use super::buffers::{GpuBuffer, GradSlice};
use super::context::GpuCtx;
use super::dtype::WeightDtype;
use super::launch::grid_1d;
use cudarc::driver::PushKernelArg;
use std::ffi::{c_int, c_void};
pub fn gpu_sgemm_forward_raw(
ctx: &GpuCtx,
y: &mut GpuBuffer,
x: &GpuBuffer,
w_ptr: cudarc::driver::sys::CUdeviceptr,
bias_ptr: Option<cudarc::driver::sys::CUdeviceptr>,
dims: (usize, usize, usize),
) -> Result<(), String> {
let (batch, n_in, n_out) = dims;
if ctx.batch_invariant() {
return super::sgemm_bi::sgemm_bi_forward(
&ctx.stream,
&ctx.kernels,
y,
x,
w_ptr,
bias_ptr.unwrap_or(0),
(batch, n_in, n_out),
);
}
let beta = if let Some(b_ptr) = bias_ptr {
let b_i = batch as i32;
let n_i = n_out as i32;
let y_ptr = y.cached_ptr();
let mut builder = ctx.stream.launch_builder(&ctx.kernels.bias_broadcast);
builder.arg(&y_ptr); builder.arg(&b_ptr);
builder.arg(&b_i);
builder.arg(&n_i);
unsafe { builder.launch(grid_1d(batch * n_out)) }
.map_err(|e| format!("bias_broadcast_raw: {:?}", e))?;
1.0f32
} else {
0.0f32
};
let alpha: f32 = 1.0;
let w_raw = w_ptr as *const f32;
let x_raw = x.raw_ptr(&ctx.stream) as *const f32;
let y_raw = y.raw_ptr(&ctx.stream) as *mut f32;
unsafe {
cudarc::cublas::result::sgemm(
*ctx.blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
n_out as c_int,
batch as c_int,
n_in as c_int,
&alpha as *const f32,
w_raw,
n_out as c_int,
x_raw,
n_in as c_int,
&beta as *const f32,
y_raw,
n_out as c_int,
)
.map_err(|e| format!("cuBLAS sgemm_forward_raw failed: {e:?}"))?;
}
Ok(())
}
pub fn gpu_sgemm_forward_ptr(
ctx: &GpuCtx,
y: &mut GpuBuffer,
x_ptr: cudarc::driver::sys::CUdeviceptr,
w_ptr: cudarc::driver::sys::CUdeviceptr,
bias_ptr: Option<cudarc::driver::sys::CUdeviceptr>,
dims: (usize, usize, usize),
) -> Result<(), String> {
let (batch, n_in, n_out) = dims;
let beta = if let Some(b_ptr) = bias_ptr {
let b_i = batch as i32;
let n_i = n_out as i32;
let y_ptr = y.cached_ptr();
let mut builder = ctx.stream.launch_builder(&ctx.kernels.bias_broadcast);
builder.arg(&y_ptr);
builder.arg(&b_ptr);
builder.arg(&b_i);
builder.arg(&n_i);
unsafe { builder.launch(grid_1d(batch * n_out)) }
.map_err(|e| format!("bias_broadcast_ptr: {:?}", e))?;
1.0f32
} else {
0.0f32
};
let alpha: f32 = 1.0;
let w_raw = w_ptr as *const f32;
let x_raw = x_ptr as *const f32;
let y_raw = y.raw_ptr(&ctx.stream) as *mut f32;
unsafe {
cudarc::cublas::result::sgemm(
*ctx.blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
n_out as c_int,
batch as c_int,
n_in as c_int,
&alpha as *const f32,
w_raw,
n_out as c_int,
x_raw,
n_in as c_int,
&beta as *const f32,
y_raw,
n_out as c_int,
)
.map_err(|e| format!("cuBLAS sgemm_forward_ptr failed: {e:?}"))?;
}
Ok(())
}
pub fn gpu_sgemm_backward_dx_raw(
ctx: &GpuCtx,
dx: &mut GpuBuffer,
dy: &GpuBuffer,
w_ptr: cudarc::driver::sys::CUdeviceptr,
batch: usize,
n_in: usize,
n_out: usize,
) -> Result<(), String> {
if ctx.batch_invariant() {
return super::sgemm_bi::sgemm_bi_backward_dx(
&ctx.stream,
&ctx.kernels,
dx,
dy,
w_ptr,
(batch, n_in, n_out),
);
}
let alpha: f32 = 1.0;
let beta: f32 = 0.0;
let w_raw = w_ptr as *const f32;
let dy_raw = dy.raw_ptr(&ctx.stream) as *const f32;
let dx_raw = dx.raw_ptr(&ctx.stream) as *mut f32;
unsafe {
cudarc::cublas::result::sgemm(
*ctx.blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_T,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
n_in as c_int,
batch as c_int,
n_out as c_int,
&alpha as *const f32,
w_raw,
n_out as c_int,
dy_raw,
n_out as c_int,
&beta as *const f32,
dx_raw,
n_in as c_int,
)
.map_err(|e| format!("cuBLAS sgemm_backward_dx_raw failed: {e:?}"))?;
}
Ok(())
}
pub fn gpu_sgemm_backward_dw_grad(
ctx: &GpuCtx,
dw: &GradSlice,
dy: &GpuBuffer,
x_saved: &GpuBuffer,
batch: usize,
n_in: usize,
n_out: usize,
) -> Result<(), String> {
if ctx.batch_invariant() {
return super::sgemm_bi::sgemm_bi_backward_dw(
&ctx.stream,
&ctx.kernels,
dw.ptr(),
dy,
x_saved,
(batch, n_in, n_out),
);
}
let alpha: f32 = 1.0;
let beta: f32 = 1.0;
let dy_ptr = dy.raw_ptr(&ctx.stream) as *const f32;
let x_ptr = x_saved.raw_ptr(&ctx.stream) as *const f32;
let dw_ptr = dw.ptr() as *mut f32;
unsafe {
cudarc::cublas::result::sgemm(
*ctx.blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_T,
n_out as c_int,
n_in as c_int,
batch as c_int,
&alpha as *const f32,
dy_ptr,
n_out as c_int,
x_ptr,
n_in as c_int,
&beta as *const f32,
dw_ptr,
n_out as c_int,
)
.map_err(|e| format!("cuBLAS sgemm_backward_dw_grad failed: {e:?}"))?;
}
Ok(())
}
pub fn gpu_sgemm_backward_dw_grad_typed(
ctx: &GpuCtx,
dw: &GradSlice,
dy: TypedPtr,
x_saved: TypedPtr,
batch: usize,
n_in: usize,
n_out: usize,
) -> Result<(), String> {
debug_assert_eq!(
dy.dtype, x_saved.dtype,
"cuBLAS GemmEx requires A.dtype == B.dtype"
);
debug_assert!(
dy.dtype != WeightDtype::F32 || !ctx.batch_invariant(),
"f32 TypedPtr under the batch-invariant flag would silently take \
non-deterministic cuBLAS — use gpu_sgemm_backward_dw_grad instead"
);
if ctx.batch_invariant() && dy.dtype != WeightDtype::F32 {
return bi_sgemm_backward_dw_typed(ctx, dw.ptr(), dy, x_saved, (batch, n_in, n_out));
}
let alpha: f32 = 1.0;
let beta: f32 = 1.0;
unsafe {
cudarc::cublas::result::gemm_ex(
*ctx.blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_T,
n_out as c_int,
n_in as c_int,
batch as c_int,
&alpha as *const f32 as *const c_void,
dy.ptr as *const c_void,
dy.dtype.cuda_data_type(),
n_out as c_int,
x_saved.ptr as *const c_void,
x_saved.dtype.cuda_data_type(),
n_in as c_int,
&beta as *const f32 as *const c_void,
dw.ptr() as *mut c_void,
cudarc::cublas::sys::cudaDataType::CUDA_R_32F,
n_out as c_int,
dy.dtype.compute_type(),
cudarc::cublas::sys::cublasGemmAlgo_t::CUBLAS_GEMM_DEFAULT,
)
.map_err(|e| format!("cuBLAS gemm_ex backward dW typed failed: {e:?}"))?;
}
Ok(())
}
pub fn gpu_gemm_ex_backward_dx_typed(
ctx: &GpuCtx,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
batch: usize,
n_in: usize,
n_out: usize,
) -> Result<(), String> {
debug_assert_eq!(
dy.dtype, w.dtype,
"cuBLAS GemmEx requires A.dtype == B.dtype"
);
debug_assert_eq!(
dx.dtype, dy.dtype,
"typed dX GEMM: dx.dtype must match dy/w for PEDANTIC path"
);
debug_assert!(
dx.dtype != WeightDtype::F32 || !ctx.batch_invariant(),
"f32 TypedPtr under the batch-invariant flag would silently take \
non-deterministic cuBLAS — use gpu_sgemm_backward_dx_raw instead"
);
if ctx.batch_invariant() && dx.dtype != WeightDtype::F32 {
return bi_sgemm_backward_dx_typed(ctx, dx, dy, w, (batch, n_in, n_out));
}
let alpha: f32 = 1.0;
let beta: f32 = 0.0;
unsafe {
cudarc::cublas::result::gemm_ex(
*ctx.blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_T,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
n_in as c_int,
batch as c_int,
n_out as c_int,
&alpha as *const f32 as *const c_void,
w.ptr as *const c_void,
w.dtype.cuda_data_type(),
n_out as c_int,
dy.ptr as *const c_void,
dy.dtype.cuda_data_type(),
n_out as c_int,
&beta as *const f32 as *const c_void,
dx.ptr as *mut c_void,
dx.dtype.cuda_data_type(),
n_in as c_int,
dy.dtype.compute_type(),
cudarc::cublas::sys::cublasGemmAlgo_t::CUBLAS_GEMM_DEFAULT,
)
.map_err(|e| format!("cuBLAS gemm_ex backward dX typed failed: {e:?}"))?;
}
Ok(())
}
fn bi_upcast_to_f32(
ctx: &GpuCtx,
src: TypedPtr,
dst_ptr: cudarc::driver::sys::CUdeviceptr,
n: usize,
) -> Result<(), String> {
let kernel = match src.dtype {
WeightDtype::Bf16 => &ctx.kernels.cast_bf16_to_f32,
WeightDtype::F16 => &ctx.kernels.cast_f16_to_f32,
WeightDtype::F32 => return Err("bi_upcast_to_f32: src is already f32".into()),
};
let n_i = n as i32;
let src_ptr = src.ptr;
let mut b = ctx.stream.launch_builder(kernel);
b.arg(&dst_ptr);
b.arg(&src_ptr);
b.arg(&n_i);
unsafe { b.launch(grid_1d(n)) }
.map(|_| ())
.map_err(|e| format!("bi_upcast_to_f32: {e:?}"))
}
fn bi_downcast_from_f32(
ctx: &GpuCtx,
dst: TypedPtr,
src_ptr: cudarc::driver::sys::CUdeviceptr,
n: usize,
) -> Result<(), String> {
let kernel = match dst.dtype {
WeightDtype::Bf16 => &ctx.kernels.cast_f32_to_bf16,
WeightDtype::F16 => &ctx.kernels.cast_f32_to_f16,
WeightDtype::F32 => return Err("bi_downcast_from_f32: dst is already f32".into()),
};
let n_i = n as i32;
let dst_ptr = dst.ptr;
let mut b = ctx.stream.launch_builder(kernel);
b.arg(&dst_ptr);
b.arg(&src_ptr);
b.arg(&n_i);
unsafe { b.launch(grid_1d(n)) }
.map(|_| ())
.map_err(|e| format!("bi_downcast_from_f32: {e:?}"))
}
pub fn bi_sgemm_forward_typed(
ctx: &GpuCtx,
y: TypedPtr,
x: TypedPtr,
w: TypedPtr,
bias_ptr: cudarc::driver::sys::CUdeviceptr,
dims: (usize, usize, usize),
) -> Result<(), String> {
if ctx.bi_tensor_cores() {
match super::sgemm_bi::sgemm_bi_forward_tc(
&ctx.stream,
&ctx.kernels,
y,
x,
w,
bias_ptr,
dims,
) {
Ok(_tile) => return Ok(()),
Err(e) if e.starts_with("UNCOVERED") => {}
Err(e) => return Err(e),
}
}
match super::sgemm_bi::sgemm_bi_forward_typed(
&ctx.stream,
&ctx.kernels,
y,
x,
w,
bias_ptr,
dims,
) {
Ok(()) => return Ok(()),
Err(e) if e.starts_with("UNCOVERED") => {}
Err(e) => return Err(e),
}
let (m, k, n) = dims;
ctx.with_bi_upcast_scratch((m * k, k * n, m * n), |xs, ws, ys| {
bi_upcast_to_f32(ctx, x, xs.cached_ptr(), m * k)?;
bi_upcast_to_f32(ctx, w, ws.cached_ptr(), k * n)?;
super::sgemm_bi::sgemm_bi_forward(
&ctx.stream,
&ctx.kernels,
ys,
xs,
ws.cached_ptr(),
bias_ptr,
dims,
)?;
bi_downcast_from_f32(ctx, y, ys.cached_ptr(), m * n)
})
}
pub fn bi_sgemm_backward_dw_typed(
ctx: &GpuCtx,
dw_ptr: cudarc::driver::sys::CUdeviceptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
) -> Result<(), String> {
if ctx.bi_tensor_cores() {
match super::sgemm_bi::sgemm_bi_backward_dw_tc(
&ctx.stream,
&ctx.kernels,
dw_ptr,
dy,
x_saved,
dims,
) {
Ok(_tile) => return Ok(()),
Err(e) if e.starts_with("UNCOVERED") => {}
Err(e) => return Err(e),
}
}
match super::sgemm_bi::sgemm_bi_backward_dw_typed(
&ctx.stream,
&ctx.kernels,
dw_ptr,
dy,
x_saved,
dims,
) {
Ok(()) => return Ok(()),
Err(e) if e.starts_with("UNCOVERED") => {}
Err(e) => return Err(e),
}
let (m, k, n) = dims;
ctx.with_bi_upcast_scratch((m * n, m * k, 0), |dys, xs, _| {
bi_upcast_to_f32(ctx, dy, dys.cached_ptr(), m * n)?;
bi_upcast_to_f32(ctx, x_saved, xs.cached_ptr(), m * k)?;
super::sgemm_bi::sgemm_bi_backward_dw(&ctx.stream, &ctx.kernels, dw_ptr, dys, xs, dims)
})
}
pub fn bi_sgemm_backward_dx_typed(
ctx: &GpuCtx,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<(), String> {
if ctx.bi_tensor_cores() {
match super::sgemm_bi::sgemm_bi_backward_dx_tc(&ctx.stream, &ctx.kernels, dx, dy, w, dims) {
Ok(_tile) => return Ok(()),
Err(e) if e.starts_with("UNCOVERED") => {}
Err(e) => return Err(e),
}
}
match super::sgemm_bi::sgemm_bi_backward_dx_typed(&ctx.stream, &ctx.kernels, dx, dy, w, dims) {
Ok(()) => return Ok(()),
Err(e) if e.starts_with("UNCOVERED") => {}
Err(e) => return Err(e),
}
let (m, k, n) = dims;
ctx.with_bi_upcast_scratch((m * n, k * n, m * k), |dys, ws, dxs| {
bi_upcast_to_f32(ctx, dy, dys.cached_ptr(), m * n)?;
bi_upcast_to_f32(ctx, w, ws.cached_ptr(), k * n)?;
super::sgemm_bi::sgemm_bi_backward_dx(
&ctx.stream,
&ctx.kernels,
dxs,
dys,
ws.cached_ptr(),
dims,
)?;
bi_downcast_from_f32(ctx, dx, dxs.cached_ptr(), m * k)
})
}
pub fn gpu_sgemm_backward_grad_raw(
ctx: &GpuCtx,
dx: &mut GpuBuffer,
grads: (&GradSlice, Option<&GradSlice>),
dy: &GpuBuffer,
x_saved: &GpuBuffer,
w_ptr: cudarc::driver::sys::CUdeviceptr,
dims: (usize, usize, usize),
) -> Result<(), String> {
let (dw, db) = grads;
let (batch, n_in, n_out) = dims;
gpu_sgemm_backward_dw_grad(ctx, dw, dy, x_saved, batch, n_in, n_out)?;
gpu_sgemm_backward_dx_raw(ctx, dx, dy, w_ptr, batch, n_in, n_out)?;
if let Some(db) = db {
let b_i = batch as i32;
let n_i = n_out as i32;
let db_ptr = db.ptr();
let dy_ptr = dy.cached_ptr();
let mut builder = ctx.stream.launch_builder(&ctx.kernels.colsum_accumulate);
builder.arg(&db_ptr);
builder.arg(&dy_ptr);
builder.arg(&b_i);
builder.arg(&n_i);
unsafe { builder.launch(grid_1d(n_out)) }
.map_err(|e| format!("colsum_accumulate_grad_raw: {:?}", e))?;
}
Ok(())
}
pub fn gpu_gemm_forward_dispatch(
ctx: &GpuCtx,
y: &mut GpuBuffer,
x: &GpuBuffer,
w_ptr: cudarc::driver::sys::CUdeviceptr,
w_dtype: WeightDtype,
bias_ptr: Option<cudarc::driver::sys::CUdeviceptr>,
dims: (usize, usize, usize),
) -> Result<(), String> {
match w_dtype {
WeightDtype::F32 => gpu_sgemm_forward_raw(ctx, y, x, w_ptr, bias_ptr, dims),
WeightDtype::F16 | WeightDtype::Bf16 => {
let (batch, n_in, _) = dims;
let half_bytes = batch * n_in * w_dtype.size_bytes();
ctx.ensure_half_staging(half_bytes)?;
let half_ptr = ctx.half_staging_ptr();
let n = (batch * n_in) as i32;
let src_ptr = x.cached_ptr();
let kernel = match w_dtype {
WeightDtype::Bf16 => &ctx.kernels.cast_f32_to_bf16,
WeightDtype::F16 => &ctx.kernels.cast_f32_to_f16,
_ => unreachable!(),
};
let mut builder = ctx.stream.launch_builder(kernel);
builder.arg(&half_ptr);
builder.arg(&src_ptr);
builder.arg(&n);
unsafe { builder.launch(grid_1d(batch * n_in)) }
.map_err(|e| format!("cast_f32_to_half: {e:?}"))?;
gpu_gemm_ex_forward_raw(
ctx,
y,
TypedPtr {
ptr: half_ptr,
dtype: w_dtype,
},
TypedPtr {
ptr: w_ptr,
dtype: w_dtype,
},
bias_ptr,
dims,
)
}
}
}
pub fn gpu_sgemm_tied_lm_head_raw(
ctx: &GpuCtx,
logits_ptr: cudarc::driver::sys::CUdeviceptr,
temporal_ptr: cudarc::driver::sys::CUdeviceptr,
embed_ptr: cudarc::driver::sys::CUdeviceptr,
batch: usize,
d_model: usize,
vocab_padded: usize,
) -> Result<(), String> {
gpu_sgemm_tied_lm_head_blas(
&ctx.blas,
logits_ptr,
temporal_ptr,
embed_ptr,
batch,
d_model,
vocab_padded,
)
}
pub fn gpu_sgemm_tied_lm_head_blas(
blas: &cudarc::cublas::CudaBlas,
logits_ptr: cudarc::driver::sys::CUdeviceptr,
temporal_ptr: cudarc::driver::sys::CUdeviceptr,
embed_ptr: cudarc::driver::sys::CUdeviceptr,
batch: usize,
d_model: usize,
vocab_padded: usize,
) -> Result<(), String> {
let alpha: f32 = 1.0;
let beta: f32 = 0.0;
unsafe {
cudarc::cublas::result::sgemm(
*blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_T,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
vocab_padded as c_int,
batch as c_int,
d_model as c_int,
&alpha as *const f32,
embed_ptr as *const f32,
d_model as c_int,
temporal_ptr as *const f32,
d_model as c_int,
&beta as *const f32,
logits_ptr as *mut f32,
vocab_padded as c_int,
)
.map_err(|e| format!("cuBLAS tied sgemm failed: {e:?}"))?;
}
Ok(())
}
#[derive(Copy, Clone)]
pub struct TypedPtr {
pub ptr: cudarc::driver::sys::CUdeviceptr,
pub dtype: WeightDtype,
}
#[derive(Copy, Clone)]
pub struct TiedLmDims {
pub batch: usize,
pub d_model: usize,
pub vocab_padded: usize,
}
pub fn gpu_gemm_ex_tied_lm_head_raw(
ctx: &GpuCtx,
logits_ptr: cudarc::driver::sys::CUdeviceptr,
temporal_ptr: cudarc::driver::sys::CUdeviceptr,
embed_ptr: cudarc::driver::sys::CUdeviceptr,
dtype: WeightDtype,
dims: TiedLmDims,
) -> Result<(), String> {
gpu_gemm_ex_tied_lm_head_blas(&ctx.blas, logits_ptr, temporal_ptr, embed_ptr, dtype, dims)
}
pub fn gpu_gemm_ex_tied_lm_head_blas(
blas: &cudarc::cublas::CudaBlas,
logits_ptr: cudarc::driver::sys::CUdeviceptr,
temporal_ptr: cudarc::driver::sys::CUdeviceptr,
embed_ptr: cudarc::driver::sys::CUdeviceptr,
dtype: WeightDtype,
dims: TiedLmDims,
) -> Result<(), String> {
let TiedLmDims {
batch,
d_model,
vocab_padded,
} = dims;
let alpha: f32 = 1.0;
let beta: f32 = 0.0;
unsafe {
cudarc::cublas::result::gemm_ex(
*blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_T,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
vocab_padded as c_int,
batch as c_int,
d_model as c_int,
&alpha as *const f32 as *const c_void,
embed_ptr as *const c_void,
dtype.cuda_data_type(),
d_model as c_int,
temporal_ptr as *const c_void,
dtype.cuda_data_type(),
d_model as c_int,
&beta as *const f32 as *const c_void,
logits_ptr as *mut c_void,
cudarc::cublas::sys::cudaDataType::CUDA_R_32F,
vocab_padded as c_int,
dtype.compute_type(),
cudarc::cublas::sys::cublasGemmAlgo_t::CUBLAS_GEMM_DEFAULT,
)
.map_err(|e| format!("cuBLAS tied gemm_ex failed: {e:?}"))?;
}
Ok(())
}
pub fn gpu_gemm_ex_forward_raw(
ctx: &GpuCtx,
y: &mut GpuBuffer,
x: TypedPtr,
w: TypedPtr,
bias_ptr: Option<cudarc::driver::sys::CUdeviceptr>,
dims: (usize, usize, usize),
) -> Result<(), String> {
gpu_gemm_typed_forward_raw(
ctx,
TypedPtr {
ptr: y.cached_ptr(),
dtype: WeightDtype::F32,
},
x,
w,
bias_ptr,
dims,
)
}
pub fn gpu_gemm_typed_raw_no_bias(
blas: &cudarc::cublas::CudaBlas,
c: TypedPtr,
x: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<(), String> {
let (batch, n_in, n_out) = dims;
let alpha: f32 = 1.0;
let beta: f32 = 0.0;
unsafe {
cudarc::cublas::result::gemm_ex(
*blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
n_out as c_int,
batch as c_int,
n_in as c_int,
&alpha as *const f32 as *const c_void,
w.ptr as *const c_void,
w.dtype.cuda_data_type(),
n_out as c_int,
x.ptr as *const c_void,
x.dtype.cuda_data_type(),
n_in as c_int,
&beta as *const f32 as *const c_void,
c.ptr as *mut c_void,
c.dtype.cuda_data_type(),
n_out as c_int,
w.dtype.compute_type(),
cudarc::cublas::sys::cublasGemmAlgo_t::CUBLAS_GEMM_DEFAULT,
)
.map_err(|e| format!("cuBLAS gemm_ex typed (no-bias) failed: {e:?}"))?;
}
Ok(())
}
fn pick_bi_gemm(
ctx: &GpuCtx,
a_dtype: WeightDtype,
b_dtype: WeightDtype,
c_dtype: WeightDtype,
) -> Option<&cudarc::driver::CudaFunction> {
if a_dtype != b_dtype {
return None;
}
match (a_dtype, c_dtype) {
(WeightDtype::Bf16, WeightDtype::Bf16) => Some(&ctx.kernels.gemm_bi_bf16_bf16),
(WeightDtype::F16, WeightDtype::F16) => Some(&ctx.kernels.gemm_bi_f16_f16),
(WeightDtype::Bf16, WeightDtype::F32) => Some(&ctx.kernels.gemm_bi_bf16_f32),
(WeightDtype::F16, WeightDtype::F32) => Some(&ctx.kernels.gemm_bi_f16_f32),
(WeightDtype::F32, WeightDtype::F32) => Some(&ctx.kernels.gemm_bi_f32_f32),
_ => None,
}
}
struct BiGemmArgs {
c: cudarc::driver::sys::CUdeviceptr,
a: cudarc::driver::sys::CUdeviceptr,
b: cudarc::driver::sys::CUdeviceptr,
bias: cudarc::driver::sys::CUdeviceptr,
alpha: f32,
beta: f32,
m: i32,
n: i32,
k: i32,
}
#[allow(dead_code)]
fn launch_bi_gemm(
ctx: &GpuCtx,
kernel: &cudarc::driver::CudaFunction,
args: BiGemmArgs,
) -> Result<(), String> {
const BLOCK_M: i32 = 64;
const BLOCK_N: i32 = 64;
const THREADS: u32 = 256;
let num_pid_m = (args.m + BLOCK_M - 1) / BLOCK_M;
let num_pid_n = (args.n + BLOCK_N - 1) / BLOCK_N;
let grid = (num_pid_m as u32) * (num_pid_n as u32);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (THREADS, 1, 1),
shared_mem_bytes: 0,
};
let lda = args.k;
let ldb = args.n;
let ldc = args.n;
let mut builder = ctx.stream.launch_builder(kernel);
builder.arg(&args.c);
builder.arg(&args.a);
builder.arg(&args.b);
builder.arg(&args.bias);
builder.arg(&args.alpha);
builder.arg(&args.beta);
builder.arg(&args.m);
builder.arg(&args.n);
builder.arg(&args.k);
builder.arg(&lda);
builder.arg(&ldb);
builder.arg(&ldc);
unsafe { builder.launch(cfg) }.map_err(|e| format!("gemm_bi launch failed: {e:?}"))?;
Ok(())
}
fn pick_bi_matvec(
ctx: &GpuCtx,
a_dtype: WeightDtype,
b_dtype: WeightDtype,
c_dtype: WeightDtype,
) -> Option<&cudarc::driver::CudaFunction> {
if a_dtype != b_dtype {
return None;
}
match (a_dtype, c_dtype) {
(WeightDtype::Bf16, WeightDtype::Bf16) => Some(&ctx.kernels.matvec_bi_bf16_bf16),
(WeightDtype::F16, WeightDtype::F16) => Some(&ctx.kernels.matvec_bi_f16_f16),
(WeightDtype::Bf16, WeightDtype::F32) => Some(&ctx.kernels.matvec_bi_bf16_f32),
(WeightDtype::F16, WeightDtype::F32) => Some(&ctx.kernels.matvec_bi_f16_f32),
(WeightDtype::F32, WeightDtype::F32) => Some(&ctx.kernels.matvec_bi_f32_f32),
_ => None,
}
}
fn launch_bi_matvec(
ctx: &GpuCtx,
kernel: &cudarc::driver::CudaFunction,
args: BiGemmArgs,
io_dtype: WeightDtype,
) -> Result<(), String> {
const BLOCK_N_MV: i32 = 32;
const THREADS_PER_BLOCK: i32 = 256;
let a_bytes = (args.k as u32) * (io_dtype.size_bytes() as u32);
let smem_bytes = (a_bytes + 15) & !15;
let num_pid_n = (args.n + BLOCK_N_MV - 1) / BLOCK_N_MV;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (num_pid_n as u32, args.m as u32, 1),
block_dim: (THREADS_PER_BLOCK as u32, 1, 1),
shared_mem_bytes: smem_bytes,
};
let lda = args.k;
let ldb = args.n;
let ldc = args.n;
let mut builder = ctx.stream.launch_builder(kernel);
builder.arg(&args.c);
builder.arg(&args.a);
builder.arg(&args.b);
builder.arg(&args.bias);
builder.arg(&args.alpha);
builder.arg(&args.beta);
builder.arg(&args.m);
builder.arg(&args.n);
builder.arg(&args.k);
builder.arg(&lda);
builder.arg(&ldb);
builder.arg(&ldc);
unsafe { builder.launch(cfg) }.map_err(|e| format!("matvec_bi launch failed: {e:?}"))?;
Ok(())
}
pub fn gpu_gemm_typed_forward_raw(
ctx: &GpuCtx,
c: TypedPtr,
x: TypedPtr,
w: TypedPtr,
bias_ptr: Option<cudarc::driver::sys::CUdeviceptr>,
dims: (usize, usize, usize),
) -> Result<(), String> {
let (batch, n_in, n_out) = dims;
let _ = pick_bi_gemm(ctx, x.dtype, w.dtype, c.dtype);
if ctx.batch_invariant()
&& c.dtype != WeightDtype::F32
&& c.dtype == x.dtype
&& x.dtype == w.dtype
&& batch >= 128
&& n_out >= 2
{
return bi_sgemm_forward_typed(ctx, c, x, w, bias_ptr.unwrap_or(0), dims);
}
if ctx.batch_invariant()
&& let Some(kernel) = pick_bi_matvec(ctx, x.dtype, w.dtype, c.dtype)
{
let bias_arg = bias_ptr.unwrap_or(0);
return launch_bi_matvec(
ctx,
kernel,
BiGemmArgs {
c: c.ptr,
a: x.ptr,
b: w.ptr,
bias: bias_arg,
alpha: 1.0,
beta: 0.0,
m: batch as i32,
n: n_out as i32,
k: n_in as i32,
},
x.dtype,
);
}
let beta = if let Some(b_ptr) = bias_ptr {
let b_i = batch as i32;
let n_i = n_out as i32;
let c_ptr = c.ptr;
let bias_kernel = match c.dtype {
WeightDtype::F32 => &ctx.kernels.bias_broadcast,
d => ctx.kernels.bias_broadcast_typed.get(d),
};
let mut builder = ctx.stream.launch_builder(bias_kernel);
builder.arg(&c_ptr);
builder.arg(&b_ptr);
builder.arg(&b_i);
builder.arg(&n_i);
unsafe { builder.launch(grid_1d(batch * n_out)) }
.map_err(|e| format!("bias_broadcast_typed: {:?}", e))?;
1.0f32
} else {
0.0f32
};
let alpha: f32 = 1.0;
unsafe {
cudarc::cublas::result::gemm_ex(
*ctx.blas.handle(),
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_N,
n_out as c_int,
batch as c_int,
n_in as c_int,
&alpha as *const f32 as *const c_void,
w.ptr as *const c_void,
w.dtype.cuda_data_type(),
n_out as c_int,
x.ptr as *const c_void,
x.dtype.cuda_data_type(),
n_in as c_int,
&beta as *const f32 as *const c_void,
c.ptr as *mut c_void,
c.dtype.cuda_data_type(),
n_out as c_int,
w.dtype.compute_type(),
cudarc::cublas::sys::cublasGemmAlgo_t::CUBLAS_GEMM_DEFAULT,
)
.map_err(|e| format!("cuBLAS gemm_ex typed failed: {e:?}"))?;
}
Ok(())
}