1use crate::cuda;
20use crate::errors::{Error, Result};
21
22pub type DevicePtr = cuda::CUdeviceptr;
24
25pub fn alloc(byte_size: usize) -> Result<DevicePtr> {
27 let mut ptr: DevicePtr = 0;
28 let status = unsafe { cuda::cuMemAlloc_v2(&mut ptr, byte_size) };
29 if status != cuda::cudaError_enum_CUDA_SUCCESS {
30 return Err(Error::from_cuda_error(status).into());
31 }
32 Ok(ptr)
33}
34
35pub fn free(ptr: DevicePtr) -> Result<()> {
37 let status = unsafe { cuda::cuMemFree_v2(ptr) };
38 if status != cuda::cudaError_enum_CUDA_SUCCESS {
39 return Err(Error::from_cuda_error(status).into());
40 }
41 Ok(())
42}
43
44pub unsafe fn copy_h_to_d<T>(dst: DevicePtr, src: *const T, count: usize) -> Result<()> {
50 let status =
51 unsafe { cuda::cuMemcpyHtoD_v2(dst, src.cast(), count * std::mem::size_of::<T>()) };
52 if status != cuda::cudaError_enum_CUDA_SUCCESS {
53 return Err(Error::from_cuda_error(status).into());
54 }
55 Ok(())
56}
57
58pub unsafe fn copy_d_to_h<T>(dst: *mut T, src: DevicePtr, count: usize) -> Result<()> {
64 let status =
65 unsafe { cuda::cuMemcpyDtoH_v2(dst.cast(), src, count * std::mem::size_of::<T>()) };
66 if status != cuda::cudaError_enum_CUDA_SUCCESS {
67 return Err(Error::from_cuda_error(status).into());
68 }
69 Ok(())
70}
71
72pub fn alloc_host<T>(n_elems: usize) -> Result<*mut T> {
78 let mut ptr: *mut std::ffi::c_void = std::ptr::null_mut();
79 let status = unsafe { cuda::cuMemAllocHost_v2(&mut ptr, n_elems * std::mem::size_of::<T>()) };
80 if status != cuda::cudaError_enum_CUDA_SUCCESS {
81 return Err(Error::from_cuda_error(status).into());
82 }
83 Ok(ptr.cast())
84}
85
86pub unsafe fn free_host<T>(ptr: *mut T) -> Result<()> {
91 let status = unsafe { cuda::cuMemFreeHost(ptr.cast()) };
92 if status != cuda::cudaError_enum_CUDA_SUCCESS {
93 return Err(Error::from_cuda_error(status).into());
94 }
95 Ok(())
96}
97
98pub fn copy_rows_d_to_d(
105 dst: DevicePtr,
106 dst_stride_bytes: usize,
107 src: DevicePtr,
108 src_stride_bytes: usize,
109 row_bytes: usize,
110 num_rows: usize,
111) -> Result<()> {
112 let params = cuda::CUDA_MEMCPY2D {
113 srcMemoryType: cuda::CUmemorytype_enum_CU_MEMORYTYPE_DEVICE,
114 srcDevice: src,
115 srcPitch: src_stride_bytes,
116 dstMemoryType: cuda::CUmemorytype_enum_CU_MEMORYTYPE_DEVICE,
117 dstDevice: dst,
118 dstPitch: dst_stride_bytes,
119 WidthInBytes: row_bytes,
120 Height: num_rows,
121 ..Default::default()
122 };
123 let status = unsafe { cuda::cuMemcpy2D_v2(¶ms) };
124 if status != cuda::cudaError_enum_CUDA_SUCCESS {
125 return Err(Error::from_cuda_error(status).into());
126 }
127 Ok(())
128}