1#![allow(clippy::drop_non_drop)] use core::ffi::c_void;
16use libloading::{Library, Symbol};
17use once_cell::sync::OnceCell;
18
19pub type ZeResult = u32;
22pub const ZE_RESULT_SUCCESS: ZeResult = 0;
23
24pub type ZeDriverHandle = *mut c_void;
25pub type ZeDeviceHandle = *mut c_void;
26
27pub const ZE_STRUCTURE_TYPE_DEVICE_PROPERTIES: u32 = 0x1;
28
29pub const ZE_DEVICE_TYPE_GPU: u32 = 1;
30pub const ZE_DEVICE_TYPE_CPU: u32 = 2;
31pub const ZE_DEVICE_TYPE_FPGA: u32 = 3;
32pub const ZE_DEVICE_TYPE_MCA: u32 = 4;
33pub const ZE_DEVICE_TYPE_VPU: u32 = 5;
34
35pub const ZE_MAX_DEVICE_NAME: usize = 256;
36
37#[repr(C)]
38#[derive(Clone, Copy)]
39pub struct ZeDeviceUuid {
40 pub id: [u8; 16],
41}
42
43#[repr(C)]
44#[derive(Clone, Copy)]
45pub struct ZeDeviceProperties {
46 pub stype: u32,
47 pub p_next: *mut c_void,
48 pub type_: u32,
49 pub vendor_id: u32,
50 pub device_id: u32,
51 pub subdevice_id: u32,
52 pub core_clock_rate: u32,
53 pub max_mem_alloc_size: u64,
54 pub max_hardware_contexts: u32,
55 pub max_command_queue_priority: u32,
56 pub num_threads_per_eu: u32,
57 pub physical_eu_simd_width: u32,
58 pub num_eus_per_subslice: u32,
59 pub num_subslices_per_slice: u32,
60 pub num_slices: u32,
61 pub timer_resolution: u64,
62 pub timestamp_valid_bits: u32,
63 pub kernel_timestamp_valid_bits: u32,
64 pub uuid: ZeDeviceUuid,
65 pub name: [core::ffi::c_char; ZE_MAX_DEVICE_NAME],
66}
67
68pub type ZeContextHandle = *mut c_void;
69pub type ZeCommandQueueHandle = *mut c_void;
70pub type ZeCommandListHandle = *mut c_void;
71
72pub const ZE_STRUCTURE_TYPE_CONTEXT_DESC: u32 = 0x2;
73pub const ZE_STRUCTURE_TYPE_COMMAND_QUEUE_DESC: u32 = 0x3;
74pub const ZE_STRUCTURE_TYPE_COMMAND_LIST_DESC: u32 = 0x4;
75pub const ZE_STRUCTURE_TYPE_DEVICE_MEM_ALLOC_DESC: u32 = 0xc;
76pub const ZE_STRUCTURE_TYPE_HOST_MEM_ALLOC_DESC: u32 = 0xd;
77pub const ZE_STRUCTURE_TYPE_MODULE_DESC: u32 = 0xf;
78pub const ZE_STRUCTURE_TYPE_KERNEL_DESC: u32 = 0x10;
79
80pub const ZE_MODULE_FORMAT_IL_SPIRV: u32 = 0;
81pub const ZE_MODULE_FORMAT_NATIVE: u32 = 1;
82
83pub type ZeModuleHandle = *mut c_void;
84pub type ZeKernelHandle = *mut c_void;
85
86#[repr(C)]
87pub struct ZeDeviceMemAllocDesc {
88 pub stype: u32,
89 pub p_next: *const c_void,
90 pub flags: u32,
91 pub ordinal: u32,
92}
93
94#[repr(C)]
95pub struct ZeHostMemAllocDesc {
96 pub stype: u32,
97 pub p_next: *const c_void,
98 pub flags: u32,
99}
100
101#[repr(C)]
102pub struct ZeModuleDesc {
103 pub stype: u32,
104 pub p_next: *const c_void,
105 pub format: u32,
106 pub input_size: usize,
107 pub p_input_module: *const u8,
108 pub p_build_flags: *const core::ffi::c_char,
109 pub p_constants: *const c_void,
110}
111
112#[repr(C)]
113pub struct ZeKernelDesc {
114 pub stype: u32,
115 pub p_next: *const c_void,
116 pub flags: u32,
117 pub p_kernel_name: *const core::ffi::c_char,
118}
119
120#[repr(C)]
121#[derive(Clone, Copy)]
122pub struct ZeGroupCount {
123 pub group_count_x: u32,
124 pub group_count_y: u32,
125 pub group_count_z: u32,
126}
127
128pub const ZE_COMMAND_QUEUE_MODE_DEFAULT: u32 = 0;
129pub const ZE_COMMAND_QUEUE_PRIORITY_NORMAL: u32 = 0;
130
131#[repr(C)]
132pub struct ZeContextDesc {
133 pub stype: u32,
134 pub p_next: *const c_void,
135 pub flags: u32,
136}
137
138#[repr(C)]
139pub struct ZeCommandQueueDesc {
140 pub stype: u32,
141 pub p_next: *const c_void,
142 pub ordinal: u32,
143 pub index: u32,
144 pub flags: u32,
145 pub mode: u32,
146 pub priority: u32,
147}
148
149#[repr(C)]
150pub struct ZeCommandListDesc {
151 pub stype: u32,
152 pub p_next: *const c_void,
153 pub command_queue_group_ordinal: u32,
154 pub flags: u32,
155}
156
157type ZeInitFn = unsafe extern "C" fn(flags: u32) -> ZeResult;
158pub type ZeContextCreateFn = unsafe extern "C" fn(
159 h_driver: ZeDriverHandle,
160 desc: *const ZeContextDesc,
161 ph_context: *mut ZeContextHandle,
162) -> ZeResult;
163pub type ZeContextDestroyFn = unsafe extern "C" fn(h_context: ZeContextHandle) -> ZeResult;
164pub type ZeCommandQueueCreateFn = unsafe extern "C" fn(
165 h_context: ZeContextHandle,
166 h_device: ZeDeviceHandle,
167 desc: *const ZeCommandQueueDesc,
168 ph_queue: *mut ZeCommandQueueHandle,
169) -> ZeResult;
170pub type ZeCommandQueueDestroyFn = unsafe extern "C" fn(h_queue: ZeCommandQueueHandle) -> ZeResult;
171pub type ZeCommandListCreateFn = unsafe extern "C" fn(
172 h_context: ZeContextHandle,
173 h_device: ZeDeviceHandle,
174 desc: *const ZeCommandListDesc,
175 ph_list: *mut ZeCommandListHandle,
176) -> ZeResult;
177pub type ZeCommandListDestroyFn = unsafe extern "C" fn(h_list: ZeCommandListHandle) -> ZeResult;
178type ZeDriverGetFn =
179 unsafe extern "C" fn(p_count: *mut u32, ph_drivers: *mut ZeDriverHandle) -> ZeResult;
180type ZeDeviceGetFn = unsafe extern "C" fn(
181 h_driver: ZeDriverHandle,
182 p_count: *mut u32,
183 ph_devices: *mut ZeDeviceHandle,
184) -> ZeResult;
185type ZeDeviceGetPropertiesFn = unsafe extern "C" fn(
186 h_device: ZeDeviceHandle,
187 p_properties: *mut ZeDeviceProperties,
188) -> ZeResult;
189
190pub type ZeMemAllocDeviceFn = unsafe extern "C" fn(
191 h_context: ZeContextHandle,
192 device_desc: *const ZeDeviceMemAllocDesc,
193 size: usize,
194 alignment: usize,
195 h_device: ZeDeviceHandle,
196 pptr: *mut *mut c_void,
197) -> ZeResult;
198pub type ZeMemAllocSharedFn = unsafe extern "C" fn(
199 h_context: ZeContextHandle,
200 device_desc: *const ZeDeviceMemAllocDesc,
201 host_desc: *const ZeHostMemAllocDesc,
202 size: usize,
203 alignment: usize,
204 h_device: ZeDeviceHandle,
205 pptr: *mut *mut c_void,
206) -> ZeResult;
207pub type ZeMemFreeFn =
208 unsafe extern "C" fn(h_context: ZeContextHandle, ptr: *mut c_void) -> ZeResult;
209pub type ZeModuleCreateFn = unsafe extern "C" fn(
210 h_context: ZeContextHandle,
211 h_device: ZeDeviceHandle,
212 desc: *const ZeModuleDesc,
213 ph_module: *mut ZeModuleHandle,
214 ph_build_log: *mut *mut c_void,
215) -> ZeResult;
216pub type ZeModuleDestroyFn = unsafe extern "C" fn(h_module: ZeModuleHandle) -> ZeResult;
217pub type ZeKernelCreateFn = unsafe extern "C" fn(
218 h_module: ZeModuleHandle,
219 desc: *const ZeKernelDesc,
220 ph_kernel: *mut ZeKernelHandle,
221) -> ZeResult;
222pub type ZeKernelDestroyFn = unsafe extern "C" fn(h_kernel: ZeKernelHandle) -> ZeResult;
223pub type ZeKernelSetGroupSizeFn =
224 unsafe extern "C" fn(h_kernel: ZeKernelHandle, gx: u32, gy: u32, gz: u32) -> ZeResult;
225pub type ZeKernelSetArgumentValueFn = unsafe extern "C" fn(
226 h_kernel: ZeKernelHandle,
227 arg_index: u32,
228 arg_size: usize,
229 p_arg_value: *const c_void,
230) -> ZeResult;
231pub type ZeCommandListAppendLaunchKernelFn = unsafe extern "C" fn(
232 h_list: ZeCommandListHandle,
233 h_kernel: ZeKernelHandle,
234 p_launch_args: *const ZeGroupCount,
235 h_signal_event: *mut c_void,
236 num_wait_events: u32,
237 ph_wait_events: *mut *mut c_void,
238) -> ZeResult;
239pub type ZeCommandListAppendMemoryCopyFn = unsafe extern "C" fn(
240 h_list: ZeCommandListHandle,
241 dstptr: *mut c_void,
242 srcptr: *const c_void,
243 size: usize,
244 h_signal_event: *mut c_void,
245 num_wait_events: u32,
246 ph_wait_events: *mut *mut c_void,
247) -> ZeResult;
248pub type ZeCommandListCloseFn = unsafe extern "C" fn(h_list: ZeCommandListHandle) -> ZeResult;
249pub type ZeCommandListResetFn = unsafe extern "C" fn(h_list: ZeCommandListHandle) -> ZeResult;
250pub type ZeCommandQueueExecuteCommandListsFn = unsafe extern "C" fn(
251 h_queue: ZeCommandQueueHandle,
252 num_lists: u32,
253 ph_lists: *const ZeCommandListHandle,
254 h_fence: *mut c_void,
255) -> ZeResult;
256pub type ZeCommandQueueSynchronizeFn =
257 unsafe extern "C" fn(h_queue: ZeCommandQueueHandle, timeout: u64) -> ZeResult;
258
259static LIB: OnceCell<Option<Loaded>> = OnceCell::new();
262
263pub struct Loaded {
264 _lib: Library,
265 pub ze_driver_get: ZeDriverGetFn,
266 pub ze_device_get: ZeDeviceGetFn,
267 pub ze_device_get_properties: ZeDeviceGetPropertiesFn,
268 pub ze_context_create: ZeContextCreateFn,
269 pub ze_context_destroy: ZeContextDestroyFn,
270 pub ze_command_queue_create: ZeCommandQueueCreateFn,
271 pub ze_command_queue_destroy: ZeCommandQueueDestroyFn,
272 pub ze_command_list_create: ZeCommandListCreateFn,
273 pub ze_command_list_destroy: ZeCommandListDestroyFn,
274 pub ze_mem_alloc_device: ZeMemAllocDeviceFn,
275 pub ze_mem_alloc_shared: ZeMemAllocSharedFn,
276 pub ze_mem_free: ZeMemFreeFn,
277 pub ze_module_create: ZeModuleCreateFn,
278 pub ze_module_destroy: ZeModuleDestroyFn,
279 pub ze_kernel_create: ZeKernelCreateFn,
280 pub ze_kernel_destroy: ZeKernelDestroyFn,
281 pub ze_kernel_set_group_size: ZeKernelSetGroupSizeFn,
282 pub ze_kernel_set_argument_value: ZeKernelSetArgumentValueFn,
283 pub ze_command_list_append_launch_kernel: ZeCommandListAppendLaunchKernelFn,
284 pub ze_command_list_append_memory_copy: ZeCommandListAppendMemoryCopyFn,
285 pub ze_command_list_close: ZeCommandListCloseFn,
286 pub ze_command_list_reset: ZeCommandListResetFn,
287 pub ze_command_queue_execute_command_lists: ZeCommandQueueExecuteCommandListsFn,
288 pub ze_command_queue_synchronize: ZeCommandQueueSynchronizeFn,
289}
290
291unsafe impl Send for Loaded {}
292unsafe impl Sync for Loaded {}
293
294#[cfg(target_os = "windows")]
295const LIB_CANDIDATES: &[&str] = &["ze_loader.dll"];
296#[cfg(any(target_os = "linux", target_os = "android"))]
297const LIB_CANDIDATES: &[&str] = &["libze_loader.so.1", "libze_loader.so"];
298#[cfg(not(any(target_os = "windows", target_os = "linux", target_os = "android")))]
299const LIB_CANDIDATES: &[&str] = &[];
300
301pub(crate) fn loaded() -> Option<&'static Loaded> {
302 LIB.get_or_init(load).as_ref()
303}
304
305fn load() -> Option<Loaded> {
306 for name in LIB_CANDIDATES {
307 let lib = match unsafe { Library::new(*name) } {
308 Ok(l) => l,
309 Err(_) => continue,
310 };
311 unsafe {
312 let init: Symbol<ZeInitFn> = match lib.get(b"zeInit\0") {
313 Ok(s) => s,
314 Err(_) => continue,
315 };
316 if init(0) != ZE_RESULT_SUCCESS {
317 continue;
318 }
319 let dg: Symbol<ZeDriverGetFn> = match lib.get(b"zeDriverGet\0") {
320 Ok(s) => s,
321 Err(_) => continue,
322 };
323 let devg: Symbol<ZeDeviceGetFn> = match lib.get(b"zeDeviceGet\0") {
324 Ok(s) => s,
325 Err(_) => continue,
326 };
327 let devp: Symbol<ZeDeviceGetPropertiesFn> = match lib.get(b"zeDeviceGetProperties\0") {
328 Ok(s) => s,
329 Err(_) => continue,
330 };
331 let ze_driver_get = *dg;
332 let ze_device_get = *devg;
333 let ze_device_get_properties = *devp;
334 drop(dg);
335 drop(devg);
336 drop(devp);
337 drop(init);
338
339 let ctx_create: Symbol<ZeContextCreateFn> = match lib.get(b"zeContextCreate\0") {
340 Ok(s) => s,
341 Err(_) => continue,
342 };
343 let ctx_destroy: Symbol<ZeContextDestroyFn> = match lib.get(b"zeContextDestroy\0") {
344 Ok(s) => s,
345 Err(_) => continue,
346 };
347 let q_create: Symbol<ZeCommandQueueCreateFn> = match lib.get(b"zeCommandQueueCreate\0")
348 {
349 Ok(s) => s,
350 Err(_) => continue,
351 };
352 let q_destroy: Symbol<ZeCommandQueueDestroyFn> =
353 match lib.get(b"zeCommandQueueDestroy\0") {
354 Ok(s) => s,
355 Err(_) => continue,
356 };
357 let l_create: Symbol<ZeCommandListCreateFn> = match lib.get(b"zeCommandListCreate\0") {
358 Ok(s) => s,
359 Err(_) => continue,
360 };
361 let l_destroy: Symbol<ZeCommandListDestroyFn> = match lib.get(b"zeCommandListDestroy\0")
362 {
363 Ok(s) => s,
364 Err(_) => continue,
365 };
366
367 let ze_context_create = *ctx_create;
368 let ze_context_destroy = *ctx_destroy;
369 let ze_command_queue_create = *q_create;
370 let ze_command_queue_destroy = *q_destroy;
371 let ze_command_list_create = *l_create;
372 let ze_command_list_destroy = *l_destroy;
373 drop(ctx_create);
374 drop(ctx_destroy);
375 drop(q_create);
376 drop(q_destroy);
377 drop(l_create);
378 drop(l_destroy);
379
380 macro_rules! sym {
381 ($ty:ty, $name:literal) => {{
382 let s: Symbol<$ty> = match lib.get($name) {
383 Ok(s) => s,
384 Err(_) => continue,
385 };
386 let f = *s;
387 drop(s);
388 f
389 }};
390 }
391 let ze_mem_alloc_device = sym!(ZeMemAllocDeviceFn, b"zeMemAllocDevice\0");
392 let ze_mem_alloc_shared = sym!(ZeMemAllocSharedFn, b"zeMemAllocShared\0");
393 let ze_mem_free = sym!(ZeMemFreeFn, b"zeMemFree\0");
394 let ze_module_create = sym!(ZeModuleCreateFn, b"zeModuleCreate\0");
395 let ze_module_destroy = sym!(ZeModuleDestroyFn, b"zeModuleDestroy\0");
396 let ze_kernel_create = sym!(ZeKernelCreateFn, b"zeKernelCreate\0");
397 let ze_kernel_destroy = sym!(ZeKernelDestroyFn, b"zeKernelDestroy\0");
398 let ze_kernel_set_group_size = sym!(ZeKernelSetGroupSizeFn, b"zeKernelSetGroupSize\0");
399 let ze_kernel_set_argument_value =
400 sym!(ZeKernelSetArgumentValueFn, b"zeKernelSetArgumentValue\0");
401 let ze_command_list_append_launch_kernel = sym!(
402 ZeCommandListAppendLaunchKernelFn,
403 b"zeCommandListAppendLaunchKernel\0"
404 );
405 let ze_command_list_append_memory_copy = sym!(
406 ZeCommandListAppendMemoryCopyFn,
407 b"zeCommandListAppendMemoryCopy\0"
408 );
409 let ze_command_list_close = sym!(ZeCommandListCloseFn, b"zeCommandListClose\0");
410 let ze_command_list_reset = sym!(ZeCommandListResetFn, b"zeCommandListReset\0");
411 let ze_command_queue_execute_command_lists = sym!(
412 ZeCommandQueueExecuteCommandListsFn,
413 b"zeCommandQueueExecuteCommandLists\0"
414 );
415 let ze_command_queue_synchronize =
416 sym!(ZeCommandQueueSynchronizeFn, b"zeCommandQueueSynchronize\0");
417
418 return Some(Loaded {
419 _lib: lib,
420 ze_driver_get,
421 ze_device_get,
422 ze_device_get_properties,
423 ze_context_create,
424 ze_context_destroy,
425 ze_command_queue_create,
426 ze_command_queue_destroy,
427 ze_command_list_create,
428 ze_command_list_destroy,
429 ze_mem_alloc_device,
430 ze_mem_alloc_shared,
431 ze_mem_free,
432 ze_module_create,
433 ze_module_destroy,
434 ze_kernel_create,
435 ze_kernel_destroy,
436 ze_kernel_set_group_size,
437 ze_kernel_set_argument_value,
438 ze_command_list_append_launch_kernel,
439 ze_command_list_append_memory_copy,
440 ze_command_list_close,
441 ze_command_list_reset,
442 ze_command_queue_execute_command_lists,
443 ze_command_queue_synchronize,
444 });
445 }
446 }
447 None
448}
449
450#[derive(Debug, Clone)]
453pub struct EnumeratedDevice {
454 pub ordinal: u32,
455 pub type_: u32,
456 pub vendor_id: u32,
457 pub device_id: u32,
458 pub name: String,
459 pub core_clock_khz: u32,
460 pub max_mem_alloc_size: u64,
461 pub num_slices: u32,
462 pub num_subslices_per_slice: u32,
463 pub num_eus_per_subslice: u32,
464}
465
466pub fn is_available() -> bool {
467 loaded().is_some()
468}
469
470pub fn enumerate() -> Vec<EnumeratedDevice> {
471 let Some(l) = loaded() else { return Vec::new() };
472 let mut out = Vec::new();
473 let mut ordinal = 0u32;
474 unsafe {
475 let mut driver_count: u32 = 0;
476 if (l.ze_driver_get)(&mut driver_count, core::ptr::null_mut()) != ZE_RESULT_SUCCESS
477 || driver_count == 0
478 {
479 return out;
480 }
481 let mut drivers = vec![core::ptr::null_mut::<c_void>(); driver_count as usize];
482 if (l.ze_driver_get)(&mut driver_count, drivers.as_mut_ptr()) != ZE_RESULT_SUCCESS {
483 return out;
484 }
485 for driver in drivers.into_iter().take(driver_count as usize) {
486 let mut dev_count: u32 = 0;
487 if (l.ze_device_get)(driver, &mut dev_count, core::ptr::null_mut()) != ZE_RESULT_SUCCESS
488 || dev_count == 0
489 {
490 continue;
491 }
492 let mut devs = vec![core::ptr::null_mut::<c_void>(); dev_count as usize];
493 if (l.ze_device_get)(driver, &mut dev_count, devs.as_mut_ptr()) != ZE_RESULT_SUCCESS {
494 continue;
495 }
496 for dev in devs.into_iter().take(dev_count as usize) {
497 let mut props: ZeDeviceProperties = core::mem::zeroed();
498 props.stype = ZE_STRUCTURE_TYPE_DEVICE_PROPERTIES;
499 if (l.ze_device_get_properties)(dev, &mut props) != ZE_RESULT_SUCCESS {
500 continue;
501 }
502 out.push(EnumeratedDevice {
503 ordinal,
504 type_: props.type_,
505 vendor_id: props.vendor_id,
506 device_id: props.device_id,
507 name: c_name_to_string(&props.name),
508 core_clock_khz: props.core_clock_rate.saturating_mul(1000),
509 max_mem_alloc_size: props.max_mem_alloc_size,
510 num_slices: props.num_slices,
511 num_subslices_per_slice: props.num_subslices_per_slice,
512 num_eus_per_subslice: props.num_eus_per_subslice,
513 });
514 ordinal += 1;
515 }
516 }
517 }
518 out
519}
520
521fn c_name_to_string(raw: &[core::ffi::c_char]) -> String {
522 let bytes: Vec<u8> = raw
523 .iter()
524 .take_while(|&&c| c != 0)
525 .map(|&c| c as u8)
526 .collect();
527 String::from_utf8_lossy(&bytes).into_owned()
528}
529
530#[cfg(test)]
531mod tests {
532 use super::*;
533
534 #[test]
535 fn probe_does_not_panic() {
536 let _ = is_available();
537 let _ = enumerate();
538 }
539}