use crate::backend::kernels::{GpuBuffer, GpuKernelExecutor, KernelInfo};
use crate::error::{NdimageError, NdimageResult};
use scirs2_core::ndarray::{Array, ArrayView2};
use scirs2_core::numeric::{Float, FromPrimitive};
use std::collections::HashMap;
use std::ffi::c_void;
use std::fmt::Debug;
use std::marker::PhantomData;
use std::sync::{Arc, Mutex};
use oxicuda_driver::ffi::{CUdeviceptr, CUstream};
use oxicuda_driver::loader::try_driver;
use oxicuda_driver::{Context, Device, Function, Module};
use oxicuda_memory::DeviceBuffer;
pub trait GpuContext: Send + Sync {
fn name(&self) -> &str;
fn device_count(&self) -> usize;
fn current_device(&self) -> usize;
fn memory_info(&self) -> (usize, usize); }
fn cuda_err(context: &str, error: oxicuda_driver::CudaError) -> NdimageError {
NdimageError::ComputationError(format!("{context}: {error}"))
}
fn nvrtc_err(error: oxicuda_nvrtc::NvrtcError) -> NdimageError {
match error {
oxicuda_nvrtc::NvrtcError::Unavailable { .. } => NdimageError::NotImplementedError(
"CUDA kernel JIT compilation requires the NVRTC runtime library (libnvrtc), \
which is not available on this system. oxicuda-ptx provides pure-Rust PTX \
generation but not CUDA-C runtime compilation."
.into(),
),
oxicuda_nvrtc::NvrtcError::Compilation { code, msg, log } => {
NdimageError::ComputationError(format!(
"CUDA kernel compilation failed (nvrtc error {code}: {msg}):\n{log}"
))
}
other => {
NdimageError::ComputationError(format!("CUDA kernel JIT compilation failed: {other}"))
}
}
}
pub struct CudaBuffer<T>
where
T: Send + Sync,
{
buffer: DeviceBuffer<u8>,
size: usize,
phantom: PhantomData<T>,
}
impl<T: Send + Sync + 'static> CudaBuffer<T> {
pub fn new(size: usize) -> NdimageResult<Self> {
let byte_size = size
.checked_mul(std::mem::size_of::<T>())
.ok_or_else(|| NdimageError::ComputationError("CUDA buffer size overflow".into()))?;
let buffer = if byte_size == 0 {
unsafe { DeviceBuffer::<u8>::from_raw(0, 0) }
} else {
DeviceBuffer::<u8>::alloc(byte_size).map_err(|e| cuda_err("CUDA malloc failed", e))?
};
Ok(Self {
buffer,
size,
phantom: PhantomData,
})
}
pub fn from_host_data(data: &[T]) -> NdimageResult<Self> {
let mut buffer = Self::new(data.len())?;
buffer.copy_from_host(data)?;
Ok(buffer)
}
fn device_ptr(&self) -> CUdeviceptr {
self.buffer.as_device_ptr()
}
}
impl<T: Send + Sync + 'static> GpuBuffer<T> for CudaBuffer<T> {
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
self
}
fn size(&self) -> usize {
self.size
}
fn copy_from_host(&mut self, data: &[T]) -> NdimageResult<()> {
if data.len() != self.size {
return Err(NdimageError::InvalidInput("Data size mismatch".to_string()));
}
if self.size == 0 {
return Ok(());
}
let bytes = unsafe {
std::slice::from_raw_parts(data.as_ptr() as *const u8, std::mem::size_of_val(data))
};
self.buffer
.copy_from_host(bytes)
.map_err(|e| cuda_err("CUDA memcpy failed", e))
}
fn copy_to_host(&self, data: &mut [T]) -> NdimageResult<()> {
if data.len() != self.size {
return Err(NdimageError::InvalidInput("Data size mismatch".to_string()));
}
if self.size == 0 {
return Ok(());
}
let byte_len = std::mem::size_of_val(data);
let bytes =
unsafe { std::slice::from_raw_parts_mut(data.as_mut_ptr() as *mut u8, byte_len) };
self.buffer
.copy_to_host(bytes)
.map_err(|e| cuda_err("CUDA memcpy failed", e))
}
}
pub struct CudaContext {
device_id: i32,
compute_capability: (i32, i32),
max_threads_per_block: i32,
max_shared_memory: usize,
context: Arc<Context>,
}
impl CudaContext {
pub fn new(deviceid: Option<usize>) -> NdimageResult<Self> {
let device_id = deviceid.unwrap_or(0) as i32;
oxicuda_driver::init().map_err(|e| cuda_err("Failed to initialize CUDA driver", e))?;
let device_count =
Device::count().map_err(|e| cuda_err("Failed to get CUDA device count", e))?;
if device_id >= device_count {
return Err(NdimageError::InvalidInput(format!(
"CUDA device {device_id} not found. Only {device_count} devices available"
)));
}
let device =
Device::get(device_id).map_err(|e| cuda_err("Failed to get CUDA device", e))?;
let context = Arc::new(
Context::new(&device).map_err(|e| cuda_err("Failed to create CUDA context", e))?,
);
let compute_capability = device.compute_capability().unwrap_or((7, 5));
let max_threads_per_block = device.max_threads_per_block().unwrap_or(1024);
let max_shared_memory = device
.max_shared_memory_per_block()
.map(|v| v as usize)
.unwrap_or(49_152);
Ok(Self {
device_id,
compute_capability,
max_threads_per_block,
max_shared_memory,
context,
})
}
pub fn compute_capability(&self) -> (i32, i32) {
self.compute_capability
}
pub fn max_threads_per_block(&self) -> i32 {
self.max_threads_per_block
}
pub fn max_shared_memory(&self) -> usize {
self.max_shared_memory
}
fn get_compilation_options(&self) -> Vec<String> {
let mut options = vec![
format!(
"--gpu-architecture=compute_{}{}",
self.compute_capability.0, self.compute_capability.1
),
"--fmad=true".to_string(),
"--use_fast_math".to_string(),
"--restrict".to_string(),
];
if self.compute_capability >= (7, 0) {
options.push("--extra-device-vectorization".to_string());
}
if self.compute_capability >= (8, 0) {
options.push("--allow-unsupported-compiler".to_string());
}
options
}
pub fn compile_kernel(&self, source: &str, kernelname: &str) -> NdimageResult<CudaKernel> {
{
let cache = KERNEL_CACHE.lock().map_err(|_| {
NdimageError::ComputationError("Failed to acquire kernel cache lock".into())
})?;
if let Some(kernel) = cache.get(kernelname) {
return Ok(kernel.clone());
}
}
let cuda_source = convert_opencl_to_cuda(source);
let options = self.get_compilation_options();
let option_refs: Vec<&str> = options.iter().map(String::as_str).collect();
let ptx = oxicuda_nvrtc::compile_to_ptx(&cuda_source, kernelname, &option_refs)
.map_err(nvrtc_err)?;
let module = Module::from_ptx(ptx.as_str())
.map_err(|e| cuda_err("Failed to load CUDA module", e))?;
let function = module
.get_function(kernelname)
.map_err(|e| cuda_err(&format!("Failed to get CUDA function '{kernelname}'"), e))?;
let kernel = CudaKernel {
name: kernelname.to_string(),
module: Arc::new(module),
function,
ptx_code: ptx.as_str().as_bytes().to_vec(),
};
{
let mut cache = KERNEL_CACHE.lock().map_err(|_| {
NdimageError::ComputationError(
"Failed to acquire kernel cache lock for insertion".into(),
)
})?;
cache.insert(kernelname.to_string(), kernel.clone());
}
Ok(kernel)
}
}
impl GpuContext for CudaContext {
fn name(&self) -> &str {
"CUDA"
}
fn device_count(&self) -> usize {
Device::count().map(|c| c as usize).unwrap_or(0)
}
fn current_device(&self) -> usize {
self.device_id as usize
}
fn memory_info(&self) -> (usize, usize) {
oxicuda_memory::memory_info()
.map(|m| (m.used(), m.total))
.unwrap_or((0, 0))
}
}
#[derive(Clone)]
pub struct CudaKernel {
name: String,
module: Arc<Module>,
function: Function,
ptx_code: Vec<u8>,
}
lazy_static::lazy_static! {
static ref KERNEL_CACHE: Arc<Mutex<HashMap<String, CudaKernel>>> = Arc::new(Mutex::new(HashMap::new()));
}
pub struct CudaExecutor {
context: Arc<CudaContext>,
stream: oxicuda_driver::Stream,
}
impl CudaExecutor {
pub fn new(context: Arc<CudaContext>) -> NdimageResult<Self> {
let stream = oxicuda_driver::Stream::new(&context.context)
.map_err(|e| cuda_err("Failed to create CUDA stream", e))?;
Ok(Self { context, stream })
}
}
impl<T> GpuKernelExecutor<T> for CudaExecutor
where
T: Float + FromPrimitive + Debug + Clone + Send + Sync + 'static,
{
fn execute_kernel(
&self,
kernel: &KernelInfo,
inputs: &[&dyn GpuBuffer<T>],
outputs: &[&mut dyn GpuBuffer<T>],
work_size: &[usize],
params: &[T],
) -> NdimageResult<()> {
let cuda_kernel = self
.context
.compile_kernel(&kernel.source, &kernel.entry_point)?;
let (grid_dim, block_dim) = calculate_launch_config(work_size, kernel.work_dimensions);
let mut dev_ptrs: Vec<CUdeviceptr> = Vec::with_capacity(inputs.len() + outputs.len());
for input in inputs {
let cuda_buf = input
.as_any()
.downcast_ref::<CudaBuffer<T>>()
.ok_or_else(|| NdimageError::InvalidInput("Expected CUDA buffer".into()))?;
dev_ptrs.push(cuda_buf.device_ptr());
}
for output in outputs {
let cuda_buf = output
.as_any()
.downcast_ref::<CudaBuffer<T>>()
.ok_or_else(|| NdimageError::InvalidInput("Expected CUDA buffer".into()))?;
dev_ptrs.push(cuda_buf.device_ptr());
}
let mut param_storage: Vec<T> = params.to_vec();
let mut kernel_args: Vec<*mut c_void> =
Vec::with_capacity(dev_ptrs.len() + param_storage.len());
for dp in &mut dev_ptrs {
kernel_args.push(dp as *mut CUdeviceptr as *mut c_void);
}
for param in &mut param_storage {
kernel_args.push(param as *mut T as *mut c_void);
}
let api = try_driver().map_err(|e| cuda_err("CUDA driver unavailable", e))?;
let launch_rc = unsafe {
(api.cu_launch_kernel)(
cuda_kernel.function.raw(),
grid_dim.0,
grid_dim.1,
grid_dim.2,
block_dim.0,
block_dim.1,
block_dim.2,
0, self.stream.raw(),
kernel_args.as_mut_ptr(),
std::ptr::null_mut(),
)
};
oxicuda_driver::check(launch_rc).map_err(|e| cuda_err("CUDA kernel launch failed", e))?;
let sync_rc = unsafe { (api.cu_stream_synchronize)(self.stream.raw()) };
oxicuda_driver::check(sync_rc).map_err(|e| cuda_err("CUDA stream sync failed", e))?;
Ok(())
}
}
pub struct CudaOperations {
context: Arc<CudaContext>,
executor: CudaExecutor,
}
impl CudaOperations {
pub fn new(deviceid: Option<usize>) -> NdimageResult<Self> {
let context = Arc::new(CudaContext::new(deviceid)?);
let executor = CudaExecutor::new(context.clone())?;
Ok(Self { context, executor })
}
pub fn context(&self) -> &Arc<CudaContext> {
&self.context
}
pub fn gaussian_filter_2d<T>(
&self,
input: &ArrayView2<T>,
sigma: [T; 2],
) -> NdimageResult<Array<T, scirs2_core::ndarray::Ix2>>
where
T: Float + FromPrimitive + Debug + Clone + Default + Send + Sync + 'static,
{
crate::backend::kernels::gpu_gaussian_filter_2d(input, sigma, &self.executor)
}
pub fn convolve_2d<T>(
&self,
input: &ArrayView2<T>,
kernel: &ArrayView2<T>,
) -> NdimageResult<Array<T, scirs2_core::ndarray::Ix2>>
where
T: Float + FromPrimitive + Debug + Clone + Default + Send + Sync + 'static,
{
crate::backend::kernels::gpu_convolve_2d(input, kernel, &self.executor)
}
pub fn median_filter_2d<T>(
&self,
input: &ArrayView2<T>,
size: [usize; 2],
) -> NdimageResult<Array<T, scirs2_core::ndarray::Ix2>>
where
T: Float + FromPrimitive + Debug + Clone + Default + Send + Sync + 'static,
{
crate::backend::kernels::gpu_median_filter_2d(input, size, &self.executor)
}
pub fn erosion_2d<T>(
&self,
input: &ArrayView2<T>,
structure: &ArrayView2<bool>,
) -> NdimageResult<Array<T, scirs2_core::ndarray::Ix2>>
where
T: Float + FromPrimitive + Debug + Clone + Default + Send + Sync + 'static,
{
crate::backend::kernels::gpu_erosion_2d(input, structure, &self.executor)
}
}
#[allow(dead_code)]
pub fn allocate_gpu_buffer<T>(data: &[T]) -> NdimageResult<Box<dyn GpuBuffer<T>>>
where
T: Send + Sync + 'static,
{
Ok(Box::new(CudaBuffer::from_host_data(data)?))
}
#[allow(dead_code)]
pub fn allocate_gpu_buffer_empty<T>(size: usize) -> NdimageResult<Box<dyn GpuBuffer<T>>>
where
T: Send + Sync + 'static,
{
Ok(Box::new(CudaBuffer::<T>::new(size)?))
}
pub struct CudaMemoryManager {
buffer_pools: HashMap<usize, Vec<CUdeviceptr>>,
total_allocated: usize,
max_pool_size: usize,
}
impl CudaMemoryManager {
pub fn new(_max_poolsize: usize) -> Self {
Self {
buffer_pools: HashMap::new(),
total_allocated: 0,
max_pool_size: _max_poolsize,
}
}
pub fn allocate_buffer(&mut self, size: usize) -> NdimageResult<*mut c_void> {
if let Some(pool) = self.buffer_pools.get_mut(&size) {
if let Some(dptr) = pool.pop() {
return Ok(dptr as usize as *mut c_void);
}
}
let api = try_driver().map_err(|e| cuda_err("CUDA driver unavailable", e))?;
let mut dptr: CUdeviceptr = 0;
let rc = unsafe { (api.cu_mem_alloc_v2)(&mut dptr, size) };
oxicuda_driver::check(rc).map_err(|e| cuda_err("CUDA malloc failed", e))?;
self.total_allocated += size;
Ok(dptr as usize as *mut c_void)
}
#[allow(clippy::not_unsafe_ptr_arg_deref)]
pub fn deallocate_buffer(&mut self, ptr: *mut c_void, size: usize) -> NdimageResult<()> {
let dptr = ptr as usize as CUdeviceptr;
let pool = self.buffer_pools.entry(size).or_default();
if pool.len() < self.max_pool_size {
pool.push(dptr);
} else {
let api = try_driver().map_err(|e| cuda_err("CUDA driver unavailable", e))?;
let rc = unsafe { (api.cu_mem_free_v2)(dptr) };
oxicuda_driver::check(rc).map_err(|e| cuda_err("CUDA free failed", e))?;
self.total_allocated = self.total_allocated.saturating_sub(size);
}
Ok(())
}
pub fn get_memory_stats(&self) -> (usize, usize) {
let pooled_memory: usize = self
.buffer_pools
.iter()
.map(|(size, pool)| size * pool.len())
.sum();
(self.total_allocated, pooled_memory)
}
pub fn clear_pools(&mut self) -> NdimageResult<()> {
let api = try_driver().ok();
for (size, pool) in self.buffer_pools.drain() {
for dptr in pool {
if let Some(api) = api {
let rc = unsafe { (api.cu_mem_free_v2)(dptr) };
oxicuda_driver::check(rc)
.map_err(|e| cuda_err("CUDA free failed during pool clear", e))?;
}
self.total_allocated = self.total_allocated.saturating_sub(size);
}
}
Ok(())
}
}
impl Drop for CudaMemoryManager {
fn drop(&mut self) {
let _ = self.clear_pools();
}
}
pub struct AdvancedCudaExecutor {
context: Arc<CudaContext>,
stream: oxicuda_driver::Stream,
memory_manager: Mutex<CudaMemoryManager>,
execution_stats: Mutex<ExecutionStats>,
}
#[derive(Default)]
struct ExecutionStats {
kernel_launches: u64,
total_execution_time: f64,
memory_transfers: u64,
total_transfer_time: f64,
}
impl AdvancedCudaExecutor {
pub fn new(context: Arc<CudaContext>) -> NdimageResult<Self> {
let stream = oxicuda_driver::Stream::new(&context.context)
.map_err(|e| cuda_err("Failed to create CUDA stream", e))?;
Ok(Self {
context,
stream,
memory_manager: Mutex::new(CudaMemoryManager::new(10)), execution_stats: Mutex::new(ExecutionStats::default()),
})
}
pub fn context(&self) -> &Arc<CudaContext> {
&self.context
}
pub fn stream(&self) -> *mut c_void {
self.stream.raw().0
}
pub fn get_execution_stats(&self) -> NdimageResult<(u64, f64, u64, f64)> {
let stats = self
.execution_stats
.lock()
.map_err(|_| NdimageError::ComputationError("Failed to acquire stats lock".into()))?;
Ok((
stats.kernel_launches,
stats.total_execution_time,
stats.memory_transfers,
stats.total_transfer_time,
))
}
pub fn get_memory_stats(&self) -> NdimageResult<(usize, usize)> {
let memory_manager = self.memory_manager.lock().map_err(|_| {
NdimageError::ComputationError("Failed to acquire memory manager lock".into())
})?;
Ok(memory_manager.get_memory_stats())
}
pub fn allocate_managed_buffer<T>(&self, size: usize) -> NdimageResult<CudaManagedBuffer<T>> {
let mut memory_manager = self.memory_manager.lock().map_err(|_| {
NdimageError::ComputationError("Failed to acquire memory manager lock".into())
})?;
let byte_size = size * std::mem::size_of::<T>();
let device_ptr = memory_manager.allocate_buffer(byte_size)?;
Ok(CudaManagedBuffer {
device_ptr,
size,
byte_size,
phantom: PhantomData,
})
}
}
pub struct CudaManagedBuffer<T> {
device_ptr: *mut c_void,
size: usize,
byte_size: usize,
phantom: PhantomData<T>,
}
impl<T> CudaManagedBuffer<T> {
#[allow(clippy::not_unsafe_ptr_arg_deref)]
pub fn copy_from_host_async(&self, data: &[T], stream: *mut c_void) -> NdimageResult<()> {
if data.len() != self.size {
return Err(NdimageError::InvalidInput("Data size mismatch".to_string()));
}
let api = try_driver().map_err(|e| cuda_err("CUDA driver unavailable", e))?;
let dptr = self.device_ptr as usize as CUdeviceptr;
let rc = unsafe {
(api.cu_memcpy_htod_async_v2)(
dptr,
data.as_ptr() as *const c_void,
self.byte_size,
CUstream(stream),
)
};
oxicuda_driver::check(rc).map_err(|e| cuda_err("CUDA async memcpy failed", e))
}
#[allow(clippy::not_unsafe_ptr_arg_deref)]
pub fn copy_to_host_async(&self, data: &mut [T], stream: *mut c_void) -> NdimageResult<()> {
if data.len() != self.size {
return Err(NdimageError::InvalidInput("Data size mismatch".to_string()));
}
let api = try_driver().map_err(|e| cuda_err("CUDA driver unavailable", e))?;
let dptr = self.device_ptr as usize as CUdeviceptr;
let rc = unsafe {
(api.cu_memcpy_dtoh_async_v2)(
data.as_mut_ptr() as *mut c_void,
dptr,
self.byte_size,
CUstream(stream),
)
};
oxicuda_driver::check(rc).map_err(|e| cuda_err("CUDA async memcpy failed", e))
}
}
#[allow(dead_code)]
fn convert_opencl_to_cuda(source: &str) -> String {
let mut cuda_source = source.to_string();
cuda_source = cuda_source.replace("__kernel", "extern \"C\" __global__");
cuda_source = cuda_source.replace("__global ", "");
cuda_source = cuda_source.replace("__local", "__shared__");
cuda_source = cuda_source.replace("__constant", "__constant__");
cuda_source = cuda_source.replace("get_global_id(0)", "blockIdx.x * blockDim.x + threadIdx.x");
cuda_source = cuda_source.replace("get_global_id(1)", "blockIdx.y * blockDim.y + threadIdx.y");
cuda_source = cuda_source.replace("get_global_id(2)", "blockIdx.z * blockDim.z + threadIdx.z");
cuda_source = cuda_source.replace("get_local_id(0)", "threadIdx.x");
cuda_source = cuda_source.replace("get_local_id(1)", "threadIdx.y");
cuda_source = cuda_source.replace("get_local_id(2)", "threadIdx.z");
cuda_source = cuda_source.replace("get_group_id(0)", "blockIdx.x");
cuda_source = cuda_source.replace("get_group_id(1)", "blockIdx.y");
cuda_source = cuda_source.replace("get_group_id(2)", "blockIdx.z");
cuda_source = cuda_source.replace("get_local_size(0)", "blockDim.x");
cuda_source = cuda_source.replace("get_local_size(1)", "blockDim.y");
cuda_source = cuda_source.replace("get_local_size(2)", "blockDim.z");
cuda_source = cuda_source.replace("get_global_size(0)", "gridDim.x * blockDim.x");
cuda_source = cuda_source.replace("get_global_size(1)", "gridDim.y * blockDim.y");
cuda_source = cuda_source.replace("get_global_size(2)", "gridDim.z * blockDim.z");
cuda_source = cuda_source.replace("barrier(CLK_LOCAL_MEM_FENCE)", "__syncthreads()");
cuda_source = cuda_source.replace("barrier(CLK_GLOBAL_MEM_FENCE)", "__threadfence()");
cuda_source = cuda_source.replace("clamp(", "fminf(fmaxf(");
cuda_source = cuda_source.replace("mix(", "lerp(");
cuda_source = cuda_source.replace("mad(", "fmaf(");
cuda_source = cuda_source.replace("atomic_add(", "atomicAdd(");
cuda_source = cuda_source.replace("atomic_sub(", "atomicSub(");
cuda_source = cuda_source.replace("atomic_inc(", "atomicInc(");
cuda_source = cuda_source.replace("atomic_dec(", "atomicDec(");
cuda_source = cuda_source.replace("atomic_min(", "atomicMin(");
cuda_source = cuda_source.replace("atomic_max(", "atomicMax(");
cuda_source = cuda_source.replace("atomic_and(", "atomicAnd(");
cuda_source = cuda_source.replace("atomic_or(", "atomicOr(");
cuda_source = cuda_source.replace("atomic_xor(", "atomicXor(");
if !cuda_source.contains("#include") {
cuda_source = format!(
"#include <cuda_runtime.h>\n#include <device_launch_parameters.h>\n\n{cuda_source}"
);
}
cuda_source
}
#[allow(dead_code)]
fn calculate_launch_config(
work_size: &[usize],
dimensions: usize,
) -> ((u32, u32, u32), (u32, u32, u32)) {
calculate_launch_config_advanced(work_size, dimensions, 1024, (65535, 65535, 65535))
}
#[allow(dead_code)]
fn calculate_launch_config_advanced(
work_size: &[usize],
dimensions: usize,
max_threads_per_block: usize,
max_grid_size: (u32, u32, u32),
) -> ((u32, u32, u32), (u32, u32, u32)) {
let block_size = match dimensions {
1 => {
let optimal_size = if work_size[0] < 128 {
64
} else if work_size[0] < 512 {
128
} else if work_size[0] < 2048 {
256
} else {
512
};
(optimal_size.min(max_threads_per_block), 1, 1)
}
2 => {
let total_threads = max_threads_per_block.min(1024);
let aspect_ratio = work_size[0] as f64 / work_size[1] as f64;
let (bx, by) = if aspect_ratio > 2.0 {
(32, total_threads / 32) } else if aspect_ratio < 0.5 {
(total_threads / 32, 32) } else {
let sqrt_threads = (total_threads as f64).sqrt() as usize;
let power_of_2 = 1 << (sqrt_threads as f64).log2().floor() as usize;
(power_of_2, total_threads / power_of_2)
};
(bx, by, 1)
}
3 => {
let total_threads = max_threads_per_block.min(512); let cube_root = (total_threads as f64).powf(1.0 / 3.0) as usize;
let optimal_dim = 1 << (cube_root as f64).log2().floor() as usize;
let remaining = total_threads / (optimal_dim * optimal_dim);
(optimal_dim, optimal_dim, remaining.max(1))
}
_ => (256, 1, 1), };
let grid_size = match dimensions {
1 => {
let blocks =
((work_size[0] + block_size.0 - 1) / block_size.0).min(max_grid_size.0 as usize);
(blocks as u32, 1, 1)
}
2 => {
let blocks_x =
((work_size[0] + block_size.0 - 1) / block_size.0).min(max_grid_size.0 as usize);
let blocks_y =
((work_size[1] + block_size.1 - 1) / block_size.1).min(max_grid_size.1 as usize);
(blocks_x as u32, blocks_y as u32, 1)
}
3 => {
let blocks_x =
((work_size[0] + block_size.0 - 1) / block_size.0).min(max_grid_size.0 as usize);
let blocks_y =
((work_size[1] + block_size.1 - 1) / block_size.1).min(max_grid_size.1 as usize);
let blocks_z =
((work_size[2] + block_size.2 - 1) / block_size.2).min(max_grid_size.2 as usize);
(blocks_x as u32, blocks_y as u32, blocks_z as u32)
}
_ => {
let blocks =
((work_size[0] + block_size.0 - 1) / block_size.0).min(max_grid_size.0 as usize);
(blocks as u32, 1, 1)
}
};
(
grid_size,
(
block_size.0 as u32,
block_size.1 as u32,
block_size.2 as u32,
),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[ignore] fn test_cudacontext_creation() {
let context = CudaContext::new(None);
assert!(context.is_ok());
if let Ok(ctx) = context {
assert_eq!(ctx.device_id, 0);
assert!(ctx.device_count() > 0);
}
}
#[test]
#[ignore] fn test_cuda_buffer_allocation() {
let buffer = CudaBuffer::<f32>::new(1024);
assert!(buffer.is_ok());
if let Ok(buf) = buffer {
assert_eq!(buf.size(), 1024);
}
}
#[test]
fn convert_opencl_to_cuda_translates_qualifiers() {
let src = "__kernel void k(__global float* a) { int i = get_global_id(0); a[i] = 0.0f; }";
let out = convert_opencl_to_cuda(src);
assert!(out.contains("extern \"C\" __global__"));
assert!(out.contains("blockIdx.x * blockDim.x + threadIdx.x"));
assert!(out.contains("#include <cuda_runtime.h>"));
}
#[test]
fn launch_config_covers_work_items() {
let (grid, block) = calculate_launch_config(&[1024, 768], 2);
assert!(block.0 >= 1 && block.1 >= 1 && block.2 == 1);
assert!(grid.0 >= 1 && grid.1 >= 1);
let (grid1, block1) = calculate_launch_config(&[4096], 1);
assert_eq!(block1.1, 1);
assert_eq!(block1.2, 1);
assert!(grid1.0 >= 1);
}
#[test]
fn nvrtc_probe_never_panics() {
let _ = oxicuda_nvrtc::is_available();
}
#[test]
fn cuda_context_creation_never_panics_without_gpu() {
let _ = CudaContext::new(None);
}
#[test]
fn allocate_gpu_buffer_degrades_without_gpu() {
let _ = allocate_gpu_buffer::<f32>(&[1.0_f32, 2.0, 3.0]);
let _ = allocate_gpu_buffer_empty::<f32>(16);
}
}