use std::error::Error;
use std::ffi::{c_int, c_uint, c_void};
use std::mem::size_of;
use cuda_driver_sys::{CUcontext, cuCtxCreate_v2, cuCtxDestroy_v2, cuCtxPushCurrent_v2, cuCtxSynchronize, CUdevice, CUdevice_attribute_enum, cuDeviceGet, cuDeviceGetAttribute, CUdeviceptr, CUfunction, cuInit, cuLaunchKernel, cuMemAlloc_v2, cuMemcpyDtoH_v2, cuMemcpyHtoD_v2, cuMemFree_v2, cuMemsetD8_v2, CUresult};
pub struct CudaEnv {
ctx: CUcontext,
device: CUdevice
}
impl CudaEnv {
pub unsafe fn new(flags: c_uint, ordinal: c_int) -> Result<CudaEnv, Box<dyn Error>> {
let mut ctx: CUcontext = std::ptr::null_mut();
let mut device: CUdevice = CUdevice::default();
let mut result = cuInit(flags);
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error initializing CUDA : {:?}", result).into());
}
result = cuDeviceGet(&mut device, ordinal);
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error getting CUDA device : {:?}", result).into());
}
result = cuCtxCreate_v2(&mut ctx, 0, device);
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error creating CUDA context : {:?}", result).into());
}
result = cuCtxPushCurrent_v2(ctx);
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error pushing CUDA context : {:?}", result).into());
}
Ok(CudaEnv {
ctx,
device
})
}
pub unsafe fn get_max_threads_per_block(&self) -> u32 {
let mut max_threads_per_block: i32 = 0;
cuDeviceGetAttribute(&mut max_threads_per_block as *mut c_int, CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_MAX_THREADS_PER_BLOCK, self.device);
max_threads_per_block as u32
}
pub unsafe fn get_max_block_per_grid(&self) -> u32 {
let mut max_block_per_grid: i32 = 0;
cuDeviceGetAttribute(&mut max_block_per_grid as *mut c_int, CUdevice_attribute_enum::CU_DEVICE_ATTRIBUTE_MAX_GRID_DIM_X, self.device);
max_block_per_grid as u32
}
pub unsafe fn get_block_and_grid_dim(data_len: usize, max_threads: u32, max_block: u32) -> ((u32, u32, u32), (u32, u32, u32)) {
if is_perfect_square(max_threads) {
let block_dim = (max_threads as f32).sqrt() as u32;
let grid_dim = (data_len as f32 / max_threads as f32).ceil() as u32;
if grid_dim <= max_block {
return ((block_dim, block_dim, 1), (grid_dim, 1, 1));
}
} else {
let block_dim = max_threads;
let grid_dim = (data_len as f32 / max_threads as f32).ceil() as u32;
if grid_dim <= max_block {
return ((block_dim, 1, 1), (grid_dim, 1, 1));
}
}
((1, 1, 1), (1, 1, 1))
}
pub unsafe fn launch(&self, function: CUfunction, args: &[*mut c_void], grid_size: (u32, u32, u32), bloc_size: (u32, u32, u32)) -> Result<(), Box<dyn Error>> {
let result = cuLaunchKernel(
function,
grid_size.0, grid_size.1, grid_size.2,
bloc_size.0, bloc_size.1, bloc_size.2,
0,
std::ptr::null_mut(),
args.as_ptr() as *mut _,
std::ptr::null_mut()
);
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error launching kernel : {:?}", result).into());
}
let result = cuCtxSynchronize();
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error synchronizing kernel : {:?}", result).into());
}
Ok(())
}
pub unsafe fn allocate(&self, size: usize) -> Result<CUdeviceptr, Box<dyn Error>> {
let mut device_ptr: CUdeviceptr = CUdeviceptr::default();
let result = cuMemAlloc_v2(&mut device_ptr, size);
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error allocating memory : {:?}", result).into());
}
Ok(device_ptr)
}
pub unsafe fn copy_host_to_device<T>(&self, device_ptr: CUdeviceptr, data: &[T]) -> Result<(), Box<dyn Error>> {
let result = cuMemcpyHtoD_v2(device_ptr, data.as_ptr() as *const c_void, data.len() * size_of::<T>());
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error copy data into device : {:?}", result).into());
}
Ok(())
}
pub unsafe fn copy_device_to_host<T>(&self, data: &mut [T], device_ptr: CUdeviceptr) -> Result<(), Box<dyn Error>> {
let result = cuMemcpyDtoH_v2(data.as_mut_ptr() as *mut c_void, device_ptr, data.len() * size_of::<T>());
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error copy data into device : {:?}", result).into());
}
Ok(())
}
pub unsafe fn set_empty_device_data(&self, device_ptr: CUdeviceptr, size: usize) -> Result<(), Box<dyn Error>> {
let result = cuMemsetD8_v2(device_ptr, 0, size);
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error setting empty data into device : {:?}", result).into());
}
Ok(())
}
pub unsafe fn free_all_data(&self, device_ptrs: &[CUdeviceptr]) -> Result<(), Box<dyn Error>> {
for device_ptr in device_ptrs {
self.free_data(*device_ptr)?;
}
Ok(())
}
pub unsafe fn free_data(&self, device_ptr: CUdeviceptr) -> Result<(), Box<dyn Error>> {
let result = cuMemFree_v2(device_ptr);
if result != CUresult::CUDA_SUCCESS {
return Err(format!("Error freeing data : {:?}", result).into());
}
Ok(())
}
pub unsafe fn free(&self) -> Result<(), Box<dyn Error>> {
cuCtxDestroy_v2(self.ctx);
Ok(())
}
}
fn is_perfect_square(num: u32) -> bool {
let sqrt_num = (num as f64).sqrt() as u32;
return sqrt_num * sqrt_num == num;
}