use super::device::GpuDevice;
use super::kernels::MambaKernels;
use std::sync::Arc;
pub struct GpuCtx {
pub stream: Arc<cudarc::driver::CudaStream>,
pub kernels: MambaKernels,
pub blas: cudarc::cublas::CudaBlas,
pub _blas_workspace: cudarc::driver::CudaSlice<u8>,
}
impl GpuCtx {
pub fn new(device: &GpuDevice) -> Result<Self, String> {
let stream = device.fork_stream()?;
let arch = GpuDevice::nvrtc_arch(device.compute_capability);
let kernels = MambaKernels::compile(device.context(), arch)?;
let (blas, ws) = device.create_cublas(&stream)?;
Ok(Self {
stream,
kernels,
blas,
_blas_workspace: ws,
})
}
pub fn disable_tf32(&self) {
unsafe {
cudarc::cublas::sys::cublasSetMathMode(
*self.blas.handle(),
cudarc::cublas::sys::cublasMath_t::CUBLAS_DEFAULT_MATH,
);
}
}
}