mod attention;
mod fp8_recipe;
mod gemm;
#[cfg(test)]
mod fp8_gemm_range_tests_3728;
#[cfg(test)]
mod fp8_recipe_tests_3807;
use super::super::*;
pub(crate) fn cublas_gemm_threshold() -> u32 {
std::env::var("CUBLAS_GEMM_THRESHOLD")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(4)
}
const F32_TO_F16_PTX: &str = r#"
.version 7.5
.target sm_75
.address_size 64
.visible .entry f32_to_f16(
.param .u64 param_dst,
.param .u64 param_src,
.param .u32 param_count
) {
.reg .u64 %rd<5>;
.reg .u32 %r<4>;
.reg .f32 %f0;
.reg .b16 %h0;
.reg .pred %p0;
ld.param.u64 %rd0, [param_dst];
ld.param.u64 %rd1, [param_src];
ld.param.u32 %r0, [param_count];
mov.u32 %r1, %tid.x;
mov.u32 %r2, %ctaid.x;
mov.u32 %r3, %ntid.x;
mad.lo.u32 %r1, %r2, %r3, %r1;
setp.ge.u32 %p0, %r1, %r0;
@%p0 bra L_DONE;
cvt.u64.u32 %rd2, %r1;
shl.b64 %rd3, %rd2, 2;
add.u64 %rd3, %rd1, %rd3;
ld.global.f32 %f0, [%rd3];
cvt.rn.f16.f32 %h0, %f0;
shl.b64 %rd4, %rd2, 1;
add.u64 %rd4, %rd0, %rd4;
st.global.b16 [%rd4], %h0;
L_DONE:
ret;
}
"#;
const F16_TO_F32_PTX: &str = r#"
.version 7.5
.target sm_75
.address_size 64
.visible .entry f16_to_f32(
.param .u64 param_dst,
.param .u64 param_src,
.param .u32 param_count
) {
.reg .u64 %rd<5>;
.reg .u32 %r<4>;
.reg .f32 %f0;
.reg .b16 %h0;
.reg .pred %p0;
ld.param.u64 %rd0, [param_dst];
ld.param.u64 %rd1, [param_src];
ld.param.u32 %r0, [param_count];
mov.u32 %r1, %tid.x;
mov.u32 %r2, %ctaid.x;
mov.u32 %r3, %ntid.x;
mad.lo.u32 %r1, %r2, %r3, %r1;
setp.ge.u32 %p0, %r1, %r0;
@%p0 bra L_DONE;
// Load FP16 (16 bits)
cvt.u64.u32 %rd2, %r1;
shl.b64 %rd3, %rd2, 1;
add.u64 %rd3, %rd1, %rd3;
ld.global.b16 %h0, [%rd3];
// Convert FP16 -> FP32 (hardware instruction)
cvt.f32.f16 %f0, %h0;
// Store as FP32 (4 bytes)
shl.b64 %rd4, %rd2, 2;
add.u64 %rd4, %rd0, %rd4;
st.global.f32 [%rd4], %f0;
L_DONE:
ret;
}
"#;
impl CudaExecutor {
pub(crate) fn ensure_cublas(&mut self) -> Result<(), GpuError> {
if self.cublas_handle.is_some() {
return Ok(());
}
let handle = trueno_gpu::driver::CublasHandle::new(&self.context)?;
handle.set_stream(&self.stream)?;
self.cublas_handle = Some(handle);
Ok(())
}
pub(crate) fn ensure_cublas_workspace(&mut self) -> Result<(), GpuError> {
if self.cublas_workspace.is_some() {
return Ok(());
}
self.ensure_cublas()?;
const WORKSPACE_SIZE: usize = 32 * 1024 * 1024;
let workspace = GpuBuffer::<u8>::new(&self.context, WORKSPACE_SIZE)?;
let handle = self.cublas_handle.as_ref().expect("cublas initialized");
handle.set_workspace(workspace.as_ptr(), WORKSPACE_SIZE)?;
eprintln!(
"[PMAT-063] cuBLAS workspace: {} MB pre-allocated for graph capture",
WORKSPACE_SIZE / 1024 / 1024
);
self.cublas_workspace = Some(workspace);
Ok(())
}
fn ensure_dequant_scratch(&mut self, n: u32, k: u32) -> Result<(), GpuError> {
let needed = n as usize * k as usize;
if self.dequant_scratch_size >= needed {
return Ok(());
}
self.dequant_scratch = Some(GpuBuffer::new(&self.context, needed)?);
self.dequant_scratch_size = needed;
Ok(())
}
pub(crate) fn ensure_fp16_activation_scratch(&mut self, count: usize) -> Result<(), GpuError> {
if self.fp16_activation_scratch_size >= count {
return Ok(());
}
self.fp16_activation_scratch = Some(GpuBuffer::new(&self.context, count)?);
self.fp16_activation_scratch_size = count;
Ok(())
}
pub(crate) fn convert_f32_to_f16(
&mut self,
src_ptr: u64,
dst_ptr: u64,
count: u32,
) -> Result<(), GpuError> {
if !self.modules.contains_key("f32_to_f16") {
let module = self.compile_ptx(F32_TO_F16_PTX)?;
self.modules.insert("f32_to_f16".to_string(), module);
}
let module = self.modules.get_mut("f32_to_f16").expect("just inserted");
let config = LaunchConfig::linear(count, 256);
let mut dst = dst_ptr;
let mut src = src_ptr;
let mut cnt = count;
unsafe {
self.stream.launch_kernel(
module,
"f32_to_f16",
&config,
&mut [
std::ptr::from_mut(&mut dst) as *mut std::ffi::c_void,
std::ptr::from_mut(&mut src) as *mut std::ffi::c_void,
std::ptr::from_mut(&mut cnt) as *mut std::ffi::c_void,
],
)?;
}
Ok(())
}
fn ensure_fp8_activation_scratch(&mut self, count: usize) -> Result<(), GpuError> {
if self.fp8_activation_scratch_size >= count {
return Ok(());
}
self.fp8_activation_scratch = Some(GpuBuffer::new(&self.context, count)?);
self.fp8_activation_scratch_size = count;
self.fp8_act_cache.invalidate();
Ok(())
}
fn convert_f16_to_f32(
&mut self,
src_ptr: u64,
dst_ptr: u64,
count: u32,
) -> Result<(), GpuError> {
if !self.modules.contains_key("f16_to_f32") {
let module = self.compile_ptx(F16_TO_F32_PTX)?;
self.modules.insert("f16_to_f32".to_string(), module);
}
let module = self.modules.get_mut("f16_to_f32").expect("just inserted");
let config = LaunchConfig::linear(count, 256);
let mut dst = dst_ptr;
let mut src = src_ptr;
let mut cnt = count;
unsafe {
self.stream.launch_kernel(
module,
"f16_to_f32",
&config,
&mut [
std::ptr::from_mut(&mut dst) as *mut std::ffi::c_void,
std::ptr::from_mut(&mut src) as *mut std::ffi::c_void,
std::ptr::from_mut(&mut cnt) as *mut std::ffi::c_void,
],
)?;
}
Ok(())
}
fn zero_fp8_activation_tail(&mut self, dst_ptr: u64, written: u32) -> Result<(), GpuError> {
if !Self::fp8_pad_zero_enabled() {
return Ok(());
}
let size = self.fp8_activation_scratch_size;
let written = written as usize;
let Some(buf) = self.fp8_activation_scratch.as_mut() else {
return Ok(());
};
if buf.as_ptr() != dst_ptr || size <= written {
return Ok(());
}
let zeros = vec![0u8; size - written];
buf.copy_from_host_at(&zeros, written)?;
Ok(())
}
fn fp8_pad_zero_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("APR_FP8_PAD_ZERO").as_deref() == Ok("1"))
}
}