use crate::error::{DriverError, IntoResult};
use crate::simt::context::CudaContext;
use std::borrow::Cow;
use std::ffi::{c_void, CString};
use std::mem::MaybeUninit;
use std::sync::Arc;
#[derive(Debug)]
pub struct CudaModule {
pub(crate) cu_module: cuda_bindings::CUmodule,
pub(crate) ctx: Arc<CudaContext>,
}
unsafe impl Send for CudaModule {}
unsafe impl Sync for CudaModule {}
impl Drop for CudaModule {
fn drop(&mut self) {
self.ctx.record_err(self.ctx.bind_to_thread());
self.ctx
.record_err(unsafe { cuda_bindings::cuModuleUnload(self.cu_module).result() });
}
}
impl CudaContext {
pub fn load_module_from_ptx_src(
self: &Arc<Self>,
ptx_src: &str,
) -> Result<Arc<CudaModule>, DriverError> {
self.bind_to_thread()?;
let c_src = CString::new(ptx_src).unwrap();
let cu_module = unsafe {
let mut cu_module = MaybeUninit::uninit();
cuda_bindings::cuModuleLoadData(cu_module.as_mut_ptr(), c_src.as_ptr() as *const _)
.result()?;
cu_module.assume_init()
};
Ok(Arc::new(CudaModule {
cu_module,
ctx: self.clone(),
}))
}
pub fn load_module_from_image(
self: &Arc<Self>,
image: &[u8],
) -> Result<Arc<CudaModule>, DriverError> {
self.bind_to_thread()?;
let image = null_terminated_image(image);
let cu_module = unsafe {
let mut cu_module = MaybeUninit::uninit();
cuda_bindings::cuModuleLoadData(cu_module.as_mut_ptr(), image.as_ptr() as *const _)
.result()?;
cu_module.assume_init()
};
Ok(Arc::new(CudaModule {
cu_module,
ctx: self.clone(),
}))
}
pub fn load_module_from_file(
self: &Arc<Self>,
filename: &str,
) -> Result<Arc<CudaModule>, DriverError> {
self.bind_to_thread()?;
let c_str = CString::new(filename).unwrap();
let mut cu_module = MaybeUninit::uninit();
let cu_module = unsafe {
cuda_bindings::cuModuleLoad(cu_module.as_mut_ptr(), c_str.as_ptr()).result()?;
cu_module.assume_init()
};
Ok(Arc::new(CudaModule {
cu_module,
ctx: self.clone(),
}))
}
}
fn null_terminated_image(image: &[u8]) -> Cow<'_, [u8]> {
if image.last() == Some(&0) {
Cow::Borrowed(image)
} else {
let mut owned = Vec::with_capacity(image.len() + 1);
owned.extend_from_slice(image);
owned.push(0);
Cow::Owned(owned)
}
}
#[derive(Debug, Clone)]
pub struct CudaFunction {
pub(crate) cu_function: cuda_bindings::CUfunction,
#[allow(unused)]
pub(crate) module: Arc<CudaModule>,
}
unsafe impl Send for CudaFunction {}
unsafe impl Sync for CudaFunction {}
impl CudaModule {
pub fn context(&self) -> &Arc<CudaContext> {
&self.ctx
}
pub unsafe fn cu_module(&self) -> cuda_bindings::CUmodule {
self.cu_module
}
pub fn load_function(self: &Arc<Self>, fn_name: &str) -> Result<CudaFunction, DriverError> {
self.ctx.bind_to_thread()?;
let c_name = CString::new(fn_name).unwrap();
let cu_function = unsafe {
let mut cu_function = MaybeUninit::uninit();
cuda_bindings::cuModuleGetFunction(
cu_function.as_mut_ptr(),
self.cu_module,
c_name.as_ptr(),
)
.result()?;
cu_function.assume_init()
};
Ok(CudaFunction {
cu_function,
module: self.clone(),
})
}
}
#[derive(Clone, Copy, Debug)]
pub struct ConstantHandle {
pub(crate) dptr: cuda_bindings::CUdeviceptr,
}
impl ConstantHandle {
pub unsafe fn from_raw(dptr: cuda_bindings::CUdeviceptr) -> Self {
Self { dptr }
}
}
impl ConstantHandle {
pub unsafe fn write_async(
&self,
stream: &crate::CudaStream,
src: *const u8,
num_bytes: usize,
) -> Result<(), DriverError> {
stream.context().bind_to_thread()?;
unsafe {
crate::simt::memory::memcpy_htod_async(self.dptr, src, num_bytes, stream.cu_stream())
}
}
pub fn write_async_staged(
&self,
stream: &crate::CudaStream,
bytes: Box<[MaybeUninit<u8>]>,
) -> Result<(), DriverError> {
stream.context().bind_to_thread()?;
let num_bytes = bytes.len();
if num_bytes == 0 {
return Ok(());
}
unsafe {
crate::simt::memory::memcpy_htod_async(
self.dptr,
bytes.as_ptr() as *const u8,
num_bytes,
stream.cu_stream(),
)?;
}
unsafe extern "C" fn drop_staged_bytes(callback: *mut c_void) {
drop(unsafe {
Box::<Box<[MaybeUninit<u8>]>>::from_raw(callback as *mut Box<[MaybeUninit<u8>]>)
});
}
let callback_data = Box::into_raw(Box::new(bytes)) as *mut c_void;
let callback_result = unsafe {
cuda_bindings::cuLaunchHostFunc(
stream.cu_stream(),
Some(drop_staged_bytes),
callback_data,
)
}
.result();
if let Err(err) = callback_result {
let staged = unsafe {
Box::<Box<[MaybeUninit<u8>]>>::from_raw(
callback_data as *mut Box<[MaybeUninit<u8>]>,
)
};
if let Err(sync_err) = stream.synchronize() {
Box::leak(staged);
return Err(sync_err);
}
drop(staged);
return Err(err);
}
Ok(())
}
pub unsafe fn write_blocking(
&self,
module: &Arc<CudaModule>,
src: *const u8,
num_bytes: usize,
) -> Result<(), DriverError> {
unsafe { module.copy_bytes_to_device_global_sync(self.dptr, src, num_bytes) }
}
}
impl CudaModule {
pub fn get_global(
self: &Arc<Self>,
name: &str,
) -> Result<(cuda_bindings::CUdeviceptr, usize), DriverError> {
self.ctx.bind_to_thread()?;
let c_name = CString::new(name).unwrap();
let mut dptr = MaybeUninit::<cuda_bindings::CUdeviceptr>::uninit();
let mut size = MaybeUninit::<usize>::uninit();
unsafe {
cuda_bindings::cuModuleGetGlobal_v2(
dptr.as_mut_ptr(),
size.as_mut_ptr(),
self.cu_module,
c_name.as_ptr(),
)
.result()?;
Ok((dptr.assume_init(), size.assume_init()))
}
}
pub unsafe fn copy_bytes_to_device_global_sync(
self: &Arc<Self>,
dptr: cuda_bindings::CUdeviceptr,
src: *const u8,
num_bytes: usize,
) -> Result<(), DriverError> {
self.ctx.bind_to_thread()?;
unsafe { crate::simt::memory::memcpy_htod_sync(dptr, src, num_bytes) }
}
}
impl CudaFunction {
fn attribute(
&self,
attribute: cuda_bindings::CUfunction_attribute,
) -> Result<u32, DriverError> {
self.context().bind_to_thread()?;
let mut value = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuFuncGetAttribute(value.as_mut_ptr(), attribute, self.cu_function)
.result()?;
u32::try_from(value.assume_init())
.map_err(|_| DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE))
}
}
pub fn context(&self) -> &Arc<CudaContext> {
self.module.context()
}
pub fn max_threads_per_block(&self) -> Result<u32, DriverError> {
self.attribute(
cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_MAX_THREADS_PER_BLOCK,
)
}
pub fn static_shared_memory_bytes(&self) -> Result<u32, DriverError> {
self.attribute(cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_SHARED_SIZE_BYTES)
}
pub fn max_dynamic_shared_memory_bytes(&self) -> Result<u32, DriverError> {
self.attribute(
cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
)
}
pub fn num_registers(&self) -> Result<u32, DriverError> {
self.attribute(cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_NUM_REGS)
}
pub fn local_size_bytes(&self) -> Result<u32, DriverError> {
self.attribute(cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_LOCAL_SIZE_BYTES)
}
pub fn const_size_bytes(&self) -> Result<u32, DriverError> {
self.attribute(cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_CONST_SIZE_BYTES)
}
pub fn required_cluster_dimensions(&self) -> Result<Option<(u32, u32, u32)>, DriverError> {
let required = (
self.attribute(
cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_REQUIRED_CLUSTER_WIDTH,
)?,
self.attribute(
cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_REQUIRED_CLUSTER_HEIGHT,
)?,
self.attribute(
cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_REQUIRED_CLUSTER_DEPTH,
)?,
);
match required {
(0, 0, 0) => Ok(None),
(x, y, z) if x != 0 && y != 0 && z != 0 => Ok(Some(required)),
_ => Err(DriverError(
cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE,
)),
}
}
pub(crate) fn set_max_dynamic_shared_memory_bytes(
&self,
bytes: u32,
) -> Result<(), DriverError> {
self.context().bind_to_thread()?;
let bytes = i32::try_from(bytes)
.map_err(|_| DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE))?;
unsafe {
cuda_bindings::cuFuncSetAttribute(
self.cu_function,
cuda_bindings::CUfunction_attribute_enum_CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
bytes,
)
}
.result()
}
pub fn max_active_blocks_per_multiprocessor(
&self,
block_threads: u32,
dynamic_shared_memory_bytes: u32,
) -> Result<u32, DriverError> {
self.context().bind_to_thread()?;
let block_threads = i32::try_from(block_threads)
.map_err(|_| DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE))?;
let mut blocks = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuOccupancyMaxActiveBlocksPerMultiprocessor(
blocks.as_mut_ptr(),
self.cu_function,
block_threads,
dynamic_shared_memory_bytes as usize,
)
.result()?;
u32::try_from(blocks.assume_init())
.map_err(|_| DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE))
}
}
pub fn max_potential_cluster_size(
&self,
grid_dim: (u32, u32, u32),
block_dim: (u32, u32, u32),
dynamic_shared_memory_bytes: u32,
) -> Result<u32, DriverError> {
self.context().bind_to_thread()?;
let config = cuda_bindings::CUlaunchConfig_st {
gridDimX: grid_dim.0,
gridDimY: grid_dim.1,
gridDimZ: grid_dim.2,
blockDimX: block_dim.0,
blockDimY: block_dim.1,
blockDimZ: block_dim.2,
sharedMemBytes: dynamic_shared_memory_bytes,
hStream: std::ptr::null_mut(),
attrs: std::ptr::null_mut(),
numAttrs: 0,
};
let mut cluster_size = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuOccupancyMaxPotentialClusterSize(
cluster_size.as_mut_ptr(),
self.cu_function,
&config,
)
.result()?;
u32::try_from(cluster_size.assume_init())
.map_err(|_| DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE))
}
}
pub fn max_active_clusters(
&self,
grid_dim: (u32, u32, u32),
block_dim: (u32, u32, u32),
dynamic_shared_memory_bytes: u32,
cluster_dim: (u32, u32, u32),
) -> Result<u32, DriverError> {
self.context().bind_to_thread()?;
let mut cluster_attribute: cuda_bindings::CUlaunchAttribute_st =
unsafe { std::mem::zeroed() };
unsafe {
let base = &mut cluster_attribute as *mut _ as *mut u8;
(base as *mut u32).write(
cuda_bindings::CUlaunchAttributeID_enum_CU_LAUNCH_ATTRIBUTE_CLUSTER_DIMENSION,
);
let dimensions = base.add(8) as *mut u32;
dimensions.write(cluster_dim.0);
dimensions.add(1).write(cluster_dim.1);
dimensions.add(2).write(cluster_dim.2);
}
let config = cuda_bindings::CUlaunchConfig_st {
gridDimX: grid_dim.0,
gridDimY: grid_dim.1,
gridDimZ: grid_dim.2,
blockDimX: block_dim.0,
blockDimY: block_dim.1,
blockDimZ: block_dim.2,
sharedMemBytes: dynamic_shared_memory_bytes,
hStream: std::ptr::null_mut(),
attrs: &mut cluster_attribute,
numAttrs: 1,
};
let mut active_clusters = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuOccupancyMaxActiveClusters(
active_clusters.as_mut_ptr(),
self.cu_function,
&config,
)
.result()?;
u32::try_from(active_clusters.assume_init())
.map_err(|_| DriverError(cuda_bindings::cudaError_enum_CUDA_ERROR_INVALID_VALUE))
}
}
pub unsafe fn cu_function(&self) -> cuda_bindings::CUfunction {
self.cu_function
}
}