use std::ffi::{c_char, c_void, CString};
use std::ptr::{null, null_mut};
use std::sync::OnceLock;
use libloading::Library;
use crate::gpu::GpuError;
#[allow(non_camel_case_types)]
pub type cl_int = i32;
#[allow(non_camel_case_types)]
pub type cl_uint = u32;
#[allow(non_camel_case_types)]
pub type cl_ulong = u64;
#[allow(non_camel_case_types)]
pub type cl_bool = cl_uint;
#[allow(non_camel_case_types)]
pub type cl_bitfield = cl_ulong;
#[allow(non_camel_case_types)]
pub type size_t = usize;
#[allow(non_camel_case_types)]
pub type cl_platform_id = *mut c_void;
#[allow(non_camel_case_types)]
pub type cl_device_id = *mut c_void;
#[allow(non_camel_case_types)]
pub type cl_context = *mut c_void;
#[allow(non_camel_case_types)]
pub type cl_command_queue = *mut c_void;
#[allow(non_camel_case_types)]
pub type cl_mem = *mut c_void;
#[allow(non_camel_case_types)]
pub type cl_program = *mut c_void;
#[allow(non_camel_case_types)]
pub type cl_kernel = *mut c_void;
#[allow(non_camel_case_types)]
pub type cl_event = *mut c_void;
#[allow(non_camel_case_types)]
pub type cl_device_type = cl_bitfield;
#[allow(non_camel_case_types)]
pub type cl_mem_flags = cl_bitfield;
#[allow(non_camel_case_types)]
pub type cl_command_queue_properties = cl_bitfield;
#[allow(non_camel_case_types)]
pub type cl_queue_properties = cl_ulong;
#[allow(non_camel_case_types)]
pub type cl_context_properties = isize;
#[allow(non_camel_case_types)]
pub type cl_mem_info = cl_uint;
#[allow(non_camel_case_types)]
pub type cl_program_build_info = cl_uint;
pub const CL_SUCCESS: cl_int = 0;
pub const CL_TRUE: cl_bool = 1;
pub const CL_BLOCKING: cl_bool = 1;
pub const CL_DEVICE_TYPE_GPU: cl_device_type = 1 << 2;
pub const CL_QUEUE_PROFILING_ENABLE: cl_command_queue_properties = 1 << 1;
pub const CL_QUEUE_PROPERTIES: cl_uint = 0x1093;
pub const CL_MEM_READ_WRITE: cl_mem_flags = 1 << 0;
pub const CL_MEM_SIZE: cl_mem_info = 0x1102;
pub const CL_PROGRAM_BUILD_LOG: cl_program_build_info = 0x1183;
type FnGetPlatformIDs = unsafe extern "C" fn(cl_uint, *mut cl_platform_id, *mut cl_uint) -> cl_int;
type FnGetDeviceIDs = unsafe extern "C" fn(
cl_platform_id,
cl_device_type,
cl_uint,
*mut cl_device_id,
*mut cl_uint,
) -> cl_int;
type FnCreateContext = unsafe extern "C" fn(
*const cl_context_properties,
cl_uint,
*const cl_device_id,
Option<unsafe extern "C" fn(*const c_char, *const c_void, size_t, *mut c_void)>,
*mut c_void,
*mut cl_int,
) -> cl_context;
type FnCreateCommandQueueWithProperties = unsafe extern "C" fn(
cl_context,
cl_device_id,
*const cl_queue_properties,
*mut cl_int,
) -> cl_command_queue;
type FnCreateCommandQueue = unsafe extern "C" fn(
cl_context,
cl_device_id,
cl_command_queue_properties,
*mut cl_int,
) -> cl_command_queue;
type FnCreateProgramWithSource = unsafe extern "C" fn(
cl_context,
cl_uint,
*const *const c_char,
*const size_t,
*mut cl_int,
) -> cl_program;
type FnBuildProgram = unsafe extern "C" fn(
cl_program,
cl_uint,
*const cl_device_id,
*const c_char,
Option<unsafe extern "C" fn(cl_program, *mut c_void)>,
*mut c_void,
) -> cl_int;
type FnGetProgramBuildInfo = unsafe extern "C" fn(
cl_program,
cl_device_id,
cl_program_build_info,
size_t,
*mut c_void,
*mut size_t,
) -> cl_int;
type FnCreateKernel = unsafe extern "C" fn(cl_program, *const c_char, *mut cl_int) -> cl_kernel;
type FnCreateBuffer =
unsafe extern "C" fn(cl_context, cl_mem_flags, size_t, *mut c_void, *mut cl_int) -> cl_mem;
type FnEnqueueWriteBuffer = unsafe extern "C" fn(
cl_command_queue,
cl_mem,
cl_bool,
size_t,
size_t,
*const c_void,
cl_uint,
*const cl_event,
*mut cl_event,
) -> cl_int;
type FnEnqueueReadBuffer = unsafe extern "C" fn(
cl_command_queue,
cl_mem,
cl_bool,
size_t,
size_t,
*mut c_void,
cl_uint,
*const cl_event,
*mut cl_event,
) -> cl_int;
type FnSetKernelArg = unsafe extern "C" fn(cl_kernel, cl_uint, size_t, *const c_void) -> cl_int;
type FnEnqueueNDRangeKernel = unsafe extern "C" fn(
cl_command_queue,
cl_kernel,
cl_uint,
*const size_t,
*const size_t,
*const size_t,
cl_uint,
*const cl_event,
*mut cl_event,
) -> cl_int;
type FnGetMemObjectInfo =
unsafe extern "C" fn(cl_mem, cl_mem_info, size_t, *mut c_void, *mut size_t) -> cl_int;
type FnFinish = unsafe extern "C" fn(cl_command_queue) -> cl_int;
type FnRelease = unsafe extern "C" fn(*mut c_void) -> cl_int;
enum QueueCreator {
WithProperties(FnCreateCommandQueueWithProperties),
Legacy(FnCreateCommandQueue),
}
pub struct OpenClApi {
_lib: Library,
cl_get_platform_ids: FnGetPlatformIDs,
cl_get_device_ids: FnGetDeviceIDs,
cl_create_context: FnCreateContext,
queue_creator: QueueCreator,
cl_create_program_with_source: FnCreateProgramWithSource,
cl_build_program: FnBuildProgram,
cl_get_program_build_info: FnGetProgramBuildInfo,
cl_create_kernel: FnCreateKernel,
cl_create_buffer: FnCreateBuffer,
cl_enqueue_write_buffer: FnEnqueueWriteBuffer,
cl_enqueue_read_buffer: FnEnqueueReadBuffer,
cl_set_kernel_arg: FnSetKernelArg,
cl_enqueue_nd_range_kernel: FnEnqueueNDRangeKernel,
cl_get_mem_object_info: FnGetMemObjectInfo,
cl_finish: FnFinish,
cl_release_context: FnRelease,
cl_release_command_queue: FnRelease,
cl_release_program: FnRelease,
cl_release_kernel: FnRelease,
cl_release_mem_object: FnRelease,
}
unsafe impl Send for OpenClApi {}
unsafe impl Sync for OpenClApi {}
#[derive(Debug)]
pub enum OpenClLoadError {
LibraryNotFound(String),
SymbolMissing(&'static str),
}
impl std::fmt::Display for OpenClLoadError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
OpenClLoadError::LibraryNotFound(msg) => {
write!(f, "OpenCL ICD library could not be loaded: {msg}")
}
OpenClLoadError::SymbolMissing(name) => {
write!(f, "OpenCL ICD is missing required symbol `{name}`")
}
}
}
}
impl std::error::Error for OpenClLoadError {}
impl OpenClApi {
pub fn load() -> Result<Self, OpenClLoadError> {
let lib = unsafe {
Library::new("libOpenCL.so.1")
.or_else(|_| Library::new("libOpenCL.so"))
.or_else(|_| Library::new("OpenCL"))
}
.map_err(|e| OpenClLoadError::LibraryNotFound(e.to_string()))?;
unsafe {
macro_rules! sym {
($bytes:expr, $name:expr) => {
*lib.get($bytes)
.map_err(|_| OpenClLoadError::SymbolMissing($name))?
};
}
let queue_creator = if let Ok(f) = lib
.get::<FnCreateCommandQueueWithProperties>(b"clCreateCommandQueueWithProperties\0")
{
QueueCreator::WithProperties(*f)
} else if let Ok(f) = lib.get::<FnCreateCommandQueue>(b"clCreateCommandQueue\0") {
QueueCreator::Legacy(*f)
} else {
return Err(OpenClLoadError::SymbolMissing(
"clCreateCommandQueueWithProperties/clCreateCommandQueue",
));
};
let cl_get_platform_ids: FnGetPlatformIDs =
sym!(b"clGetPlatformIDs\0", "clGetPlatformIDs");
let cl_get_device_ids: FnGetDeviceIDs = sym!(b"clGetDeviceIDs\0", "clGetDeviceIDs");
let cl_create_context: FnCreateContext = sym!(b"clCreateContext\0", "clCreateContext");
let cl_create_program_with_source: FnCreateProgramWithSource =
sym!(b"clCreateProgramWithSource\0", "clCreateProgramWithSource");
let cl_build_program: FnBuildProgram = sym!(b"clBuildProgram\0", "clBuildProgram");
let cl_get_program_build_info: FnGetProgramBuildInfo =
sym!(b"clGetProgramBuildInfo\0", "clGetProgramBuildInfo");
let cl_create_kernel: FnCreateKernel = sym!(b"clCreateKernel\0", "clCreateKernel");
let cl_create_buffer: FnCreateBuffer = sym!(b"clCreateBuffer\0", "clCreateBuffer");
let cl_enqueue_write_buffer: FnEnqueueWriteBuffer =
sym!(b"clEnqueueWriteBuffer\0", "clEnqueueWriteBuffer");
let cl_enqueue_read_buffer: FnEnqueueReadBuffer =
sym!(b"clEnqueueReadBuffer\0", "clEnqueueReadBuffer");
let cl_set_kernel_arg: FnSetKernelArg = sym!(b"clSetKernelArg\0", "clSetKernelArg");
let cl_enqueue_nd_range_kernel: FnEnqueueNDRangeKernel =
sym!(b"clEnqueueNDRangeKernel\0", "clEnqueueNDRangeKernel");
let cl_get_mem_object_info: FnGetMemObjectInfo =
sym!(b"clGetMemObjectInfo\0", "clGetMemObjectInfo");
let cl_finish: FnFinish = sym!(b"clFinish\0", "clFinish");
let cl_release_context: FnRelease = sym!(b"clReleaseContext\0", "clReleaseContext");
let cl_release_command_queue: FnRelease =
sym!(b"clReleaseCommandQueue\0", "clReleaseCommandQueue");
let cl_release_program: FnRelease = sym!(b"clReleaseProgram\0", "clReleaseProgram");
let cl_release_kernel: FnRelease = sym!(b"clReleaseKernel\0", "clReleaseKernel");
let cl_release_mem_object: FnRelease =
sym!(b"clReleaseMemObject\0", "clReleaseMemObject");
Ok(Self {
cl_get_platform_ids,
cl_get_device_ids,
cl_create_context,
queue_creator,
cl_create_program_with_source,
cl_build_program,
cl_get_program_build_info,
cl_create_kernel,
cl_create_buffer,
cl_enqueue_write_buffer,
cl_enqueue_read_buffer,
cl_set_kernel_arg,
cl_enqueue_nd_range_kernel,
cl_get_mem_object_info,
cl_finish,
cl_release_context,
cl_release_command_queue,
cl_release_program,
cl_release_kernel,
cl_release_mem_object,
_lib: lib,
})
}
}
pub fn platform_ids(&self) -> Result<Vec<cl_platform_id>, GpuError> {
let mut count: cl_uint = 0;
let status = unsafe { (self.cl_get_platform_ids)(0, null_mut(), &mut count) };
if status != CL_SUCCESS {
return Err(GpuError::Other(format!(
"clGetPlatformIDs (count) failed with status {status}"
)));
}
if count == 0 {
return Ok(Vec::new());
}
let mut ids: Vec<cl_platform_id> = vec![null_mut(); count as usize];
let status = unsafe { (self.cl_get_platform_ids)(count, ids.as_mut_ptr(), null_mut()) };
if status != CL_SUCCESS {
return Err(GpuError::Other(format!(
"clGetPlatformIDs failed with status {status}"
)));
}
Ok(ids)
}
pub fn device_ids(&self, device_type: cl_device_type) -> Result<Vec<cl_device_id>, GpuError> {
let platforms = self.platform_ids()?;
let mut all: Vec<cl_device_id> = Vec::new();
for platform in platforms {
let mut count: cl_uint = 0;
let status = unsafe {
(self.cl_get_device_ids)(platform, device_type, 0, null_mut(), &mut count)
};
if status != CL_SUCCESS || count == 0 {
continue;
}
let mut ids: Vec<cl_device_id> = vec![null_mut(); count as usize];
let status = unsafe {
(self.cl_get_device_ids)(platform, device_type, count, ids.as_mut_ptr(), null_mut())
};
if status == CL_SUCCESS {
all.extend_from_slice(&ids);
}
}
Ok(all)
}
pub fn create_context(&self, device: cl_device_id) -> Result<cl_context, GpuError> {
let devices = [device];
let mut err: cl_int = 0;
let ctx = unsafe {
(self.cl_create_context)(null(), 1, devices.as_ptr(), None, null_mut(), &mut err)
};
if err != CL_SUCCESS || ctx.is_null() {
return Err(GpuError::Other(format!(
"clCreateContext failed with status {err}"
)));
}
Ok(ctx)
}
pub fn create_command_queue(
&self,
ctx: cl_context,
device: cl_device_id,
) -> Result<cl_command_queue, GpuError> {
let mut err: cl_int = 0;
let queue = match self.queue_creator {
QueueCreator::WithProperties(create) => {
let props: [cl_queue_properties; 3] = [
CL_QUEUE_PROPERTIES as cl_queue_properties,
CL_QUEUE_PROFILING_ENABLE,
0,
];
unsafe { create(ctx, device, props.as_ptr(), &mut err) }
}
QueueCreator::Legacy(create) => {
unsafe { create(ctx, device, CL_QUEUE_PROFILING_ENABLE, &mut err) }
}
};
if err != CL_SUCCESS || queue.is_null() {
return Err(GpuError::Other(format!(
"clCreateCommandQueue failed with status {err}"
)));
}
Ok(queue)
}
pub fn build_program(
&self,
ctx: cl_context,
device: cl_device_id,
source: &str,
) -> Result<cl_program, GpuError> {
let src_bytes = source.as_bytes();
let strings: [*const c_char; 1] = [src_bytes.as_ptr() as *const c_char];
let lengths: [size_t; 1] = [src_bytes.len()];
let mut err: cl_int = 0;
let program = unsafe {
(self.cl_create_program_with_source)(
ctx,
1,
strings.as_ptr(),
lengths.as_ptr(),
&mut err,
)
};
if err != CL_SUCCESS || program.is_null() {
return Err(GpuError::KernelCompilationError(format!(
"clCreateProgramWithSource failed with status {err}"
)));
}
let devices = [device];
let build_status = unsafe {
(self.cl_build_program)(program, 1, devices.as_ptr(), null(), None, null_mut())
};
if build_status != CL_SUCCESS {
let log = self.program_build_log(program, device);
unsafe {
(self.cl_release_program)(program);
}
return Err(GpuError::KernelCompilationError(format!(
"clBuildProgram failed with status {build_status}: {log}"
)));
}
Ok(program)
}
fn program_build_log(&self, program: cl_program, device: cl_device_id) -> String {
let mut size: size_t = 0;
let status = unsafe {
(self.cl_get_program_build_info)(
program,
device,
CL_PROGRAM_BUILD_LOG,
0,
null_mut(),
&mut size,
)
};
if status != CL_SUCCESS || size == 0 {
return String::new();
}
let mut buffer = vec![0u8; size];
let status = unsafe {
(self.cl_get_program_build_info)(
program,
device,
CL_PROGRAM_BUILD_LOG,
size,
buffer.as_mut_ptr() as *mut c_void,
null_mut(),
)
};
if status != CL_SUCCESS {
return String::new();
}
while buffer.last() == Some(&0) {
buffer.pop();
}
String::from_utf8_lossy(&buffer).into_owned()
}
pub fn create_kernel(&self, program: cl_program, name: &str) -> Result<cl_kernel, GpuError> {
let c_name = CString::new(name)
.map_err(|_| GpuError::Other(format!("invalid OpenCL kernel name `{name}`")))?;
let mut err: cl_int = 0;
let kernel = unsafe { (self.cl_create_kernel)(program, c_name.as_ptr(), &mut err) };
if err != CL_SUCCESS || kernel.is_null() {
return Err(GpuError::Other(format!(
"clCreateKernel(`{name}`) failed with status {err}"
)));
}
Ok(kernel)
}
pub fn create_buffer(
&self,
ctx: cl_context,
flags: cl_mem_flags,
size: usize,
) -> Result<cl_mem, GpuError> {
let mut err: cl_int = 0;
let mem = unsafe { (self.cl_create_buffer)(ctx, flags, size, null_mut(), &mut err) };
if err != CL_SUCCESS || mem.is_null() {
return Err(GpuError::OutOfMemory(format!(
"clCreateBuffer({size} bytes) failed with status {err}"
)));
}
Ok(mem)
}
pub fn enqueue_write(
&self,
queue: cl_command_queue,
mem: cl_mem,
offset: usize,
data: &[u8],
) -> Result<(), GpuError> {
let status = unsafe {
(self.cl_enqueue_write_buffer)(
queue,
mem,
CL_BLOCKING,
offset,
data.len(),
data.as_ptr() as *const c_void,
0,
null(),
null_mut(),
)
};
check(status, "clEnqueueWriteBuffer")
}
pub fn enqueue_read(
&self,
queue: cl_command_queue,
mem: cl_mem,
offset: usize,
data: &mut [u8],
) -> Result<(), GpuError> {
let status = unsafe {
(self.cl_enqueue_read_buffer)(
queue,
mem,
CL_BLOCKING,
offset,
data.len(),
data.as_mut_ptr() as *mut c_void,
0,
null(),
null_mut(),
)
};
check(status, "clEnqueueReadBuffer")
}
pub fn set_arg_bytes(
&self,
kernel: cl_kernel,
index: cl_uint,
bytes: &[u8],
) -> Result<(), GpuError> {
let status = unsafe {
(self.cl_set_kernel_arg)(kernel, index, bytes.len(), bytes.as_ptr() as *const c_void)
};
check(status, "clSetKernelArg")
}
pub fn set_arg_mem(
&self,
kernel: cl_kernel,
index: cl_uint,
mem: &cl_mem,
) -> Result<(), GpuError> {
let status = unsafe {
(self.cl_set_kernel_arg)(
kernel,
index,
std::mem::size_of::<cl_mem>(),
mem as *const cl_mem as *const c_void,
)
};
check(status, "clSetKernelArg")
}
pub fn enqueue_nd_range(
&self,
queue: cl_command_queue,
kernel: cl_kernel,
global: &[usize],
local: Option<&[usize]>,
) -> Result<(), GpuError> {
let work_dim = global.len() as cl_uint;
let local_ptr = match local {
Some(l) => l.as_ptr(),
None => null(),
};
let status = unsafe {
(self.cl_enqueue_nd_range_kernel)(
queue,
kernel,
work_dim,
null(),
global.as_ptr(),
local_ptr,
0,
null(),
null_mut(),
)
};
check(status, "clEnqueueNDRangeKernel")
}
pub fn finish(&self, queue: cl_command_queue) -> Result<(), GpuError> {
let status = unsafe { (self.cl_finish)(queue) };
check(status, "clFinish")
}
#[allow(dead_code)]
pub fn mem_size(&self, mem: cl_mem) -> Result<usize, GpuError> {
let mut size: size_t = 0;
let status = unsafe {
(self.cl_get_mem_object_info)(
mem,
CL_MEM_SIZE,
std::mem::size_of::<size_t>(),
&mut size as *mut size_t as *mut c_void,
null_mut(),
)
};
if status != CL_SUCCESS {
return Err(GpuError::Other(format!(
"clGetMemObjectInfo(CL_MEM_SIZE) failed with status {status}"
)));
}
Ok(size)
}
fn release_context(&self, ctx: cl_context) {
unsafe {
(self.cl_release_context)(ctx);
}
}
fn release_command_queue(&self, queue: cl_command_queue) {
unsafe {
(self.cl_release_command_queue)(queue);
}
}
fn release_program(&self, program: cl_program) {
unsafe {
(self.cl_release_program)(program);
}
}
fn release_kernel(&self, kernel: cl_kernel) {
unsafe {
(self.cl_release_kernel)(kernel);
}
}
fn release_mem_object(&self, mem: cl_mem) {
unsafe {
(self.cl_release_mem_object)(mem);
}
}
}
fn check(status: cl_int, op: &str) -> Result<(), GpuError> {
if status == CL_SUCCESS {
Ok(())
} else {
Err(GpuError::Other(format!("{op} failed with status {status}")))
}
}
static OPENCL_API: OnceLock<Option<OpenClApi>> = OnceLock::new();
pub fn api() -> Option<&'static OpenClApi> {
OPENCL_API.get_or_init(|| OpenClApi::load().ok()).as_ref()
}
pub struct ClContext(pub cl_context);
impl Drop for ClContext {
fn drop(&mut self) {
if self.0.is_null() {
return;
}
if let Some(a) = api() {
a.release_context(self.0);
}
}
}
pub struct ClQueue(pub cl_command_queue);
impl Drop for ClQueue {
fn drop(&mut self) {
if self.0.is_null() {
return;
}
if let Some(a) = api() {
a.release_command_queue(self.0);
}
}
}
pub struct ClProgram(pub cl_program);
impl Drop for ClProgram {
fn drop(&mut self) {
if self.0.is_null() {
return;
}
if let Some(a) = api() {
a.release_program(self.0);
}
}
}
pub struct ClKernel(pub cl_kernel);
impl Drop for ClKernel {
fn drop(&mut self) {
if self.0.is_null() {
return;
}
if let Some(a) = api() {
a.release_kernel(self.0);
}
}
}
pub struct ClBuffer {
pub mem: cl_mem,
pub size: usize,
}
impl Drop for ClBuffer {
fn drop(&mut self) {
if self.mem.is_null() {
return;
}
if let Some(a) = api() {
a.release_mem_object(self.mem);
}
}
}
unsafe impl Send for ClContext {}
unsafe impl Sync for ClContext {}
unsafe impl Send for ClQueue {}
unsafe impl Sync for ClQueue {}
unsafe impl Send for ClProgram {}
unsafe impl Sync for ClProgram {}
unsafe impl Send for ClKernel {}
unsafe impl Sync for ClKernel {}
unsafe impl Send for ClBuffer {}
unsafe impl Sync for ClBuffer {}