use cudarc::cublaslt::{CudaBlasLT, Matmul, MatmulConfig};
use cudarc::driver::{CudaContext, CudaStream, sys as cu};
use std::sync::Arc;
pub use memra_gguf;
pub fn cpu_linear(x: &[f32], w: &[f32], m: usize, in_f: usize, out_f: usize) -> Vec<f32> {
assert_eq!(x.len(), m * in_f);
assert_eq!(w.len(), out_f * in_f);
let mut y = vec![0f32; m * out_f];
for t in 0..m {
for o in 0..out_f {
let mut acc = 0f32;
let xr = &x[t * in_f..t * in_f + in_f];
let wr = &w[o * in_f..o * in_f + in_f];
for i in 0..in_f {
acc += xr[i] * wr[i];
}
y[t * out_f + o] = acc;
}
}
y
}
pub struct Gpu {
pub ctx: Arc<CudaContext>,
stream: Arc<CudaStream>,
blas: Arc<CudaBlasLT>,
phase: std::sync::Mutex<Option<[(Arc<CudaStream>, Arc<CudaBlasLT>); 2]>>,
}
struct StreamBinding {
stream: Arc<CudaStream>,
blas: Arc<CudaBlasLT>,
}
thread_local! {
static STREAM_OVERRIDE: std::cell::RefCell<Vec<StreamBinding>> =
const { std::cell::RefCell::new(Vec::new()) };
}
thread_local! {
static DECODE_PHASE: std::cell::Cell<Option<usize>> = const { std::cell::Cell::new(None) };
}
thread_local! {
static RANK0_REDIRECT: std::cell::RefCell<Option<(usize, Arc<CudaStream>, Arc<CudaBlasLT>)>> =
const { std::cell::RefCell::new(None) };
}
pub fn set_rank0_redirect(binding: Option<(usize, Arc<CudaStream>, Arc<CudaBlasLT>)>) {
RANK0_REDIRECT.with(|c| *c.borrow_mut() = binding);
}
pub struct Rank0RedirectGuard(());
pub fn rank0_redirect_scope(
ordinal: usize,
stream: Arc<CudaStream>,
blas: Arc<CudaBlasLT>,
) -> Rank0RedirectGuard {
set_rank0_redirect(Some((ordinal, stream, blas)));
Rank0RedirectGuard(())
}
impl Drop for Rank0RedirectGuard {
fn drop(&mut self) {
set_rank0_redirect(None);
}
}
fn rank0_redirect_for(ordinal: usize) -> Option<(Arc<CudaStream>, Arc<CudaBlasLT>)> {
RANK0_REDIRECT.with(|c| {
c.borrow()
.as_ref()
.filter(|(o, ..)| *o == ordinal)
.map(|(_, s, b)| (s.clone(), b.clone()))
})
}
pub fn set_decode_phase(p: Option<usize>) {
DECODE_PHASE.with(|c| c.set(p));
}
pub fn decode_phase() -> Option<usize> {
DECODE_PHASE.with(|c| c.get())
}
pub struct StreamOverride(());
pub struct GpuMainOverride {
stream: Option<StreamOverride>,
expected_ctx: cu::CUcontext,
}
pub fn push_stream_override(stream: Arc<CudaStream>, blas: Arc<CudaBlasLT>) -> StreamOverride {
STREAM_OVERRIDE.with(|o| o.borrow_mut().push(StreamBinding { stream, blas }));
StreamOverride(())
}
impl Drop for StreamOverride {
fn drop(&mut self) {
STREAM_OVERRIDE.with(|o| {
o.borrow_mut().pop();
});
}
}
impl Drop for GpuMainOverride {
fn drop(&mut self) {
drop(self.stream.take());
let set_rc = unsafe { cu::cuCtxSetCurrent(self.expected_ctx) };
let mut popped = std::ptr::null_mut();
let pop_rc = unsafe { cu::cuCtxPopCurrent_v2(&mut popped) };
if set_rc != cu::CUresult::CUDA_SUCCESS
|| pop_rc != cu::CUresult::CUDA_SUCCESS
|| popped != self.expected_ctx
{
let message = format!(
"rank-local CUDA context restore failed: set_rc={set_rc:?} pop_rc={pop_rc:?} \
expected={:?} popped={popped:?}",
self.expected_ctx,
);
if std::thread::panicking() {
eprintln!("{message}");
} else {
panic!("{message}");
}
}
}
}
impl Gpu {
#[inline]
pub fn stream(&self) -> Arc<CudaStream> {
STREAM_OVERRIDE
.with(|o| o.borrow().last().map(|binding| binding.stream.clone()))
.unwrap_or_else(|| self.stream.clone())
}
#[inline]
pub fn blas(&self) -> Arc<CudaBlasLT> {
STREAM_OVERRIDE
.with(|o| o.borrow().last().map(|binding| binding.blas.clone()))
.unwrap_or_else(|| self.blas.clone())
}
#[inline]
pub fn main_stream(&self) -> &Arc<CudaStream> {
&self.stream
}
pub fn enter_main(&self) -> Result<GpuMainOverride, Box<dyn std::error::Error>> {
let rc = unsafe { cu::cuCtxPushCurrent_v2(self.ctx.cu_ctx()) };
if rc != cu::CUresult::CUDA_SUCCESS {
return Err(format!("rank-local CUDA context push failed: {rc:?}").into());
}
let (stream, blas) = if let Some(pair) = rank0_redirect_for(self.ctx.ordinal()) {
pair
} else {
match decode_phase() {
Some(p) => self.phase_pair(p)?,
None => (self.stream.clone(), self.blas.clone()),
}
};
Ok(GpuMainOverride {
stream: Some(push_stream_override(stream, blas)),
expected_ctx: self.ctx.cu_ctx(),
})
}
pub fn phase_pair(
&self,
p: usize,
) -> Result<(Arc<CudaStream>, Arc<CudaBlasLT>), Box<dyn std::error::Error>> {
let mut guard = self
.phase
.lock()
.map_err(|_| "gpu phase lock is poisoned")?;
if guard.is_none() {
let s0 = self.ctx.new_stream()?;
let s1 = self.ctx.new_stream()?;
let b0 = Arc::new(CudaBlasLT::new(s0.clone())?);
let b1 = Arc::new(CudaBlasLT::new(s1.clone())?);
*guard = Some([(s0, b0), (s1, b1)]);
}
let arr = guard.as_ref().expect("armed above");
Ok((arr[p & 1].0.clone(), arr[p & 1].1.clone()))
}
}
impl Gpu {
pub fn new(ordinal: usize) -> Result<Self, Box<dyn std::error::Error>> {
let ctx = CudaContext::new(ordinal)?;
let stream = ctx.new_stream()?;
unsafe {
use cudarc::driver::sys;
let dev = ctx.cu_device();
let mut pool: sys::CUmemoryPool = std::ptr::null_mut();
if sys::cuDeviceGetDefaultMemPool(&mut pool, dev) == sys::CUresult::CUDA_SUCCESS
&& !pool.is_null()
{
let off: std::os::raw::c_int = 0;
let _ = sys::cuMemPoolSetAttribute(
pool,
sys::CUmemPool_attribute::CU_MEMPOOL_ATTR_REUSE_ALLOW_OPPORTUNISTIC,
&off as *const _ as *mut std::os::raw::c_void,
);
let on: std::os::raw::c_int = 1;
let _ = sys::cuMemPoolSetAttribute(
pool,
sys::CUmemPool_attribute::CU_MEMPOOL_ATTR_REUSE_ALLOW_INTERNAL_DEPENDENCIES,
&on as *const _ as *mut std::os::raw::c_void,
);
let thresh: u64 = u64::MAX;
let _ = sys::cuMemPoolSetAttribute(
pool,
sys::CUmemPool_attribute::CU_MEMPOOL_ATTR_RELEASE_THRESHOLD,
&thresh as *const _ as *mut std::os::raw::c_void,
);
}
}
let blas = Arc::new(CudaBlasLT::new(stream.clone())?);
Ok(Self {
ctx,
stream,
blas,
phase: std::sync::Mutex::new(None),
})
}
pub fn linear_f32(
&self,
x: &cudarc::driver::CudaSlice<f32>,
w: &cudarc::driver::CudaSlice<f32>,
m_tokens: usize,
in_f: usize,
out_f: usize,
) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
let stream = self.stream();
let mut c = stream.alloc_zeros::<f32>(m_tokens * out_f)?;
let cfg = MatmulConfig {
transa: true, transb: false,
transc: false,
m: out_f as u64,
n: m_tokens as u64,
k: in_f as u64,
alpha: 1.0,
lda: in_f as i64, ldb: in_f as i64, beta: 0.0,
ldc: out_f as i64, stride_a: None,
stride_b: None,
stride_c: None,
stride_bias: None,
batch_size: None,
};
let blas = self.blas();
unsafe {
blas.matmul(cfg, w, x, &mut c, None, None)?;
}
let y = stream.clone_dtoh(&c)?;
stream.synchronize()?;
Ok(y)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cpu_linear_tiny() {
let x = vec![1.0, 2.0];
let w = vec![1.0, 0.0, 0.0, 1.0]; let y = cpu_linear(&x, &w, 1, 2, 2);
assert_eq!(y, vec![1.0, 2.0]);
let w2 = vec![1.0, 1.0, 2.0, 0.0];
let y2 = cpu_linear(&x, &w2, 1, 2, 2);
assert_eq!(y2, vec![3.0, 2.0]);
}
}