use cudarc::driver::{CudaSlice, DevicePtr, DevicePtrMut};
unsafe extern "C" {
fn memra_fp8_pp_gemm(
w_e4m3: *const core::ffi::c_void,
x_f32: *const f32,
xq_e4m3: *mut core::ffi::c_void,
scales: *mut f32,
y_f32: *mut f32,
m: i32,
n: i32,
k: i32,
w_scale: f32,
ws: *mut core::ffi::c_void,
ws_bytes: usize,
stream: *mut core::ffi::c_void,
) -> i32;
}
pub fn pp_fp8_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_PP_FP8")
.map(|v| v == "1")
.unwrap_or(false)
})
}
pub fn st_e4m3_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_ST_E4M3").as_deref() != Ok("0"))
}
pub fn st_e4m3_blk_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
st_e4m3_enabled() && std::env::var("MEMRA_ST_E4M3_BLK").as_deref() != Ok("0")
})
}
static BLK_NATIVE_NAN_REFUSED: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
pub fn note_blk_native_nan_refused() {
BLK_NATIVE_NAN_REFUSED.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
pub fn blk_native_nan_refused() -> usize {
BLK_NATIVE_NAN_REFUSED.load(std::sync::atomic::Ordering::Relaxed)
}
pub struct Fp8Scratch {
pub xq: CudaSlice<u8>,
pub scales: CudaSlice<f32>,
pub ws: CudaSlice<u8>,
cap_xq: usize,
}
const FP8_WS_BYTES: usize = 64 << 20;
impl crate::Engine {
pub fn try_fp8_gemm(
&self,
w: &crate::model::GpuTensor,
x: &CudaSlice<f32>,
m: usize,
) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
if crate::portable_mma_gated() {
return Ok(None);
}
let (w_bytes, w_scale, ne) = match w {
GpuTensor::Quant {
qtype,
bytes,
scale,
ne,
..
} if *qtype == crate::QT_F8_E4M3 => (bytes, *scale, ne),
GpuTensor::Quant {
fp8: Some(f8), ne, ..
} if pp_fp8_enabled() && f8.blk.is_none() => (&f8.bytes, f8.scale, ne),
_ => return Ok(None),
};
let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
let need_xq = m * in_f;
let mut guard = self.fp8_scratch.lock().unwrap();
if guard.is_none() {
*guard = Some(Fp8Scratch {
xq: self.alloc_u8_uninit(need_xq)?,
scales: self.alloc_uninit::<f32>(4)?,
ws: self.alloc_u8_uninit(FP8_WS_BYTES)?,
cap_xq: need_xq,
});
}
let s = guard.as_mut().unwrap();
if need_xq > s.cap_xq {
s.xq = self.alloc_u8_uninit(need_xq)?;
s.cap_xq = need_xq;
}
let mut y = self.uninit(m * out_f)?; let rc = {
let stream = self.gpu.stream();
let (w_p, _gw) = w_bytes.device_ptr(&stream);
let (x_p, _gx) = x.device_ptr(&stream);
let (q_p, _gq) = s.xq.device_ptr_mut(&stream);
let (sc_p, _gs) = s.scales.device_ptr_mut(&stream);
let (y_p, _gy) = y.device_ptr_mut(&stream);
let (ws_p, _gws) = s.ws.device_ptr_mut(&stream);
unsafe {
memra_fp8_pp_gemm(
w_p as *const core::ffi::c_void,
x_p as *const f32,
q_p as *mut core::ffi::c_void,
sc_p as *mut f32,
y_p as *mut f32,
m as i32,
out_f as i32,
in_f as i32,
w_scale,
ws_p as *mut core::ffi::c_void,
FP8_WS_BYTES,
stream.cu_stream() as *mut core::ffi::c_void,
)
}
};
if rc != 0 {
return Err(format!(
"memra_fp8_pp_gemm rc={rc} (m={m} n={out_f} k={in_f}; 1xxxx=cudaError quant chain, \
2xxxx=no cublasLt algo, 3xxxx=matmul status)"
)
.into());
}
Ok(Some(y))
}
}
pub fn fp8_mmq_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_FP8_MMQ")
.map(|v| v == "1")
.unwrap_or(false)
})
}
pub fn fp8_blk_mmq_native_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("MEMRA_FP8_MMQ").as_deref() != Ok("0"))
}
static FP8_MMQ_NAN_OK: std::sync::Mutex<Option<std::collections::HashMap<u64, bool>>> =
std::sync::Mutex::new(None);
impl crate::Engine {
pub fn try_fp8_blk_mmq(
&self,
w: &crate::model::GpuTensor,
x: &CudaSlice<f32>,
m: usize,
) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
use crate::model::GpuTensor;
FP8_MMQ_ENTRIES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if crate::portable_mma_gated() {
FP8_MMQ_GATE_OFF.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(None);
}
let (f8_bytes, f8_scale, blk, ne) = match w {
GpuTensor::Quant {
fp8: Some(f8), ne, ..
} if f8.blk.is_some() => {
if !fp8_mmq_enabled() {
FP8_MMQ_GATE_OFF.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(None);
}
(&f8.bytes, f8.scale, f8.blk.as_ref().unwrap(), ne)
}
GpuTensor::Quant {
bytes, qtype, scale, blk: Some(g), ne, ..
} if *qtype == crate::QT_F8_E4M3_BLK => {
if !fp8_blk_mmq_native_enabled() {
FP8_MMQ_GATE_OFF.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(None);
}
(bytes, *scale, g, ne)
}
_ => {
FP8_MMQ_NO_OPERAND.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(None);
}
};
if ne.len() != 2 {
FP8_MMQ_BAD_SHAPE.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(None);
}
let (in_f, out_f) = (ne[0] as usize, ne[1] as usize);
if in_f % 16 != 0
|| blk.rows != out_f.div_ceil(128)
|| blk.cols != in_f.div_ceil(128)
|| f8_bytes.len() < out_f * in_f
|| x.len() < m * in_f
{
FP8_MMQ_BAD_SHAPE.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(None);
}
if f8_scale != 1.0 {
FP8_MMQ_BAD_SCALE.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(None);
}
{
let key = {
let stream = self.gpu.stream();
let (p, _g) = f8_bytes.device_ptr(&stream);
p as u64
};
let mut guard = FP8_MMQ_NAN_OK.lock().unwrap();
let map = guard.get_or_insert_with(std::collections::HashMap::new);
let ok = match map.get(&key) {
Some(v) => *v,
None => {
let v = self.fp8_blk_nan_count(f8_bytes)? == 0;
map.insert(key, v);
v
}
};
if !ok {
FP8_MMQ_NAN_REFUSED.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(None);
}
}
let y = self.qmatvec_mmq_fp8_blk(f8_bytes, &blk.scales, x, m, in_f, out_f)?;
FP8_MMQ_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Ok(Some(y))
}
}
static FP8_MMQ_HITS: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
static FP8_MMQ_NO_OPERAND: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
static FP8_MMQ_BAD_SHAPE: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
static FP8_MMQ_BAD_SCALE: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
static FP8_MMQ_NAN_REFUSED: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
static FP8_MMQ_ENTRIES: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
static FP8_MMQ_GATE_OFF: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
pub fn fp8_mmq_hits() -> usize {
FP8_MMQ_HITS.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn fp8_mmq_ledger() -> (usize, usize, usize, usize, usize, usize, usize) {
use std::sync::atomic::Ordering::Relaxed;
(
FP8_MMQ_ENTRIES.load(Relaxed),
FP8_MMQ_GATE_OFF.load(Relaxed),
FP8_MMQ_HITS.load(Relaxed),
FP8_MMQ_NO_OPERAND.load(Relaxed),
FP8_MMQ_BAD_SHAPE.load(Relaxed),
FP8_MMQ_BAD_SCALE.load(Relaxed),
FP8_MMQ_NAN_REFUSED.load(Relaxed),
)
}
unsafe extern "C" {
fn memra_fp8_blk_q8_0_bytes(out_dim: i32, in_dim: i32) -> usize;
fn memra_fp8_blk_dequant_q8_0(
f8_weights: *const core::ffi::c_void,
blk_scales: *const f32,
out_q8: *mut core::ffi::c_void,
out_dim: i32,
in_dim: i32,
stream: *mut core::ffi::c_void,
) -> i32;
}
pub fn fp8_blk_gpu_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
std::env::var("MEMRA_FP8_BLK_GPU")
.map(|v| v == "1")
.unwrap_or(false)
})
}
impl crate::Engine {
pub fn fp8_blk_q8_0_bytes(out_f: usize, in_f: usize) -> usize {
unsafe { memra_fp8_blk_q8_0_bytes(out_f as i32, in_f as i32) }
}
pub fn fp8_blk_dequant_q8_0(
&self,
f8: &[u8],
grid: &[f32],
out_f: usize,
in_f: usize,
) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
let (rows, cols) = (out_f.div_ceil(128), in_f.div_ceil(128));
if f8.len() != out_f * in_f {
return Err(format!(
"fp8_blk_dequant_q8_0: f8 len {} != out_f*in_f {}",
f8.len(),
out_f * in_f
)
.into());
}
if grid.len() != rows * cols {
return Err(format!(
"fp8_blk_dequant_q8_0: grid len {} != rows*cols {rows}*{cols}",
grid.len()
)
.into());
}
let need = Self::fp8_blk_q8_0_bytes(out_f, in_f);
if need == 0 {
return Err(format!(
"fp8_blk_dequant_q8_0: bad dims out_f={out_f} in_f={in_f} (in_f must be %32)"
)
.into());
}
let src = self.htod_bytes(f8)?;
let scales = self.htod(grid)?;
let dst = self.fp8_blk_dequant_q8_0_dev(&src, &scales, out_f, in_f)?;
self.gpu.stream().synchronize()?;
Ok(dst)
}
pub fn fp8_blk_dequant_q8_0_dev(
&self,
f8: &CudaSlice<u8>,
grid: &CudaSlice<f32>,
out_f: usize,
in_f: usize,
) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
let (rows, cols) = (out_f.div_ceil(128), in_f.div_ceil(128));
if f8.len() < out_f * in_f {
return Err(format!(
"fp8_blk_dequant_q8_0_dev: f8 len {} < out_f*in_f {}",
f8.len(),
out_f * in_f
)
.into());
}
if grid.len() < rows * cols {
return Err(format!(
"fp8_blk_dequant_q8_0_dev: grid len {} < rows*cols {rows}*{cols}",
grid.len()
)
.into());
}
let need = Self::fp8_blk_q8_0_bytes(out_f, in_f);
if need == 0 {
return Err(format!(
"fp8_blk_dequant_q8_0_dev: bad dims out_f={out_f} in_f={in_f} (in_f must be %32)"
)
.into());
}
let mut dst = self.alloc_u8_uninit(need)?;
let rc = {
let stream = self.gpu.stream();
let (s_p, _gs) = f8.device_ptr(&stream);
let (g_p, _gg) = grid.device_ptr(&stream);
let (d_p, _gd) = dst.device_ptr_mut(&stream);
unsafe {
memra_fp8_blk_dequant_q8_0(
s_p as *const core::ffi::c_void,
g_p as *const f32,
d_p as *mut core::ffi::c_void,
out_f as i32,
in_f as i32,
stream.cu_stream() as *mut core::ffi::c_void,
)
}
};
if rc != 0 {
return Err(format!(
"memra_fp8_blk_dequant_q8_0 rc={rc} (out_f={out_f} in_f={in_f}; 1=bad dims, \
else cudaError_t)"
)
.into());
}
Ok(dst)
}
}