diff --git a/oneapi-rs-sys/src/event-sys.rs b/oneapi-rs-sys/src/event-sys.rs index 6158ac8..52bf7a1 100644 --- a/oneapi-rs-sys/src/event-sys.rs +++ b/oneapi-rs-sys/src/event-sys.rs @@ -25,7 +25,7 @@ pub mod ffi { #[namespace = "sycl_shims"] type Queue = crate::types::ffi::Queue; - fn wait(event: &mut UniquePtr); + fn wait(event: &mut UniquePtr) -> Result<()>; unsafe fn register_callback( queue: &mut UniquePtr, event: &Event, diff --git a/oneapi-rs-sys/src/kernel-bundle-sys.rs b/oneapi-rs-sys/src/kernel-bundle-sys.rs index 2e53ab4..832f8ff 100644 --- a/oneapi-rs-sys/src/kernel-bundle-sys.rs +++ b/oneapi-rs-sys/src/kernel-bundle-sys.rs @@ -26,11 +26,15 @@ pub mod ffi { fn create_kernel_bundle_from_source( ctxt: &Context, source: &str, - ) -> UniquePtr; - fn build(source: &mut UniquePtr) -> UniquePtr; + ) -> Result>; + + fn build( + source: &mut UniquePtr, + ) -> Result>; + fn get_kernel( bundle: &mut UniquePtr, name: &str, - ) -> UniquePtr; + ) -> Result>; } } diff --git a/oneapi-rs-sys/src/queue-sys.rs b/oneapi-rs-sys/src/queue-sys.rs index 5e77ef0..3de169d 100644 --- a/oneapi-rs-sys/src/queue-sys.rs +++ b/oneapi-rs-sys/src/queue-sys.rs @@ -45,8 +45,12 @@ pub mod ffi { dep_events: Vec, ) -> UniquePtr; - fn barrier(queue: &mut UniquePtr, dep_events: Vec) -> UniquePtr; - fn wait(queue: &mut UniquePtr); + fn barrier( + queue: &mut UniquePtr, + dep_events: Vec, + ) -> Result>; + + fn wait(queue: &mut UniquePtr) -> Result<()>; unsafe fn launch_1d( queue: &mut UniquePtr, @@ -54,7 +58,7 @@ pub mod ffi { local_size: Range1, kernel: &Kernel, args: &[&[u8]], - ) -> UniquePtr; + ) -> Result>; unsafe fn launch_2d( queue: &mut UniquePtr, @@ -62,7 +66,7 @@ pub mod ffi { local_size: Range2, kernel: &Kernel, args: &[&[u8]], - ) -> UniquePtr; + ) -> Result>; unsafe fn launch_3d( queue: &mut UniquePtr, @@ -70,7 +74,7 @@ pub mod ffi { local_size: Range3, kernel: &Kernel, args: &[&[u8]], - ) -> UniquePtr; + ) -> Result>; unsafe fn memcpy( queue: &mut UniquePtr, diff --git a/oneapi-rs/examples/kernel_launch.rs b/oneapi-rs/examples/kernel_launch.rs index ce8aa2c..9de73df 100644 --- a/oneapi-rs/examples/kernel_launch.rs +++ b/oneapi-rs/examples/kernel_launch.rs @@ -22,15 +22,15 @@ void iota(float start, float *ptr) { "#; #[tokio::main] -async fn main() { +async fn main() -> oneapi_rs::Result<()> { let mut queue = Queue::new(); - let mut device_buffer = queue.alloc_device::(1024).await; + let mut device_buffer = queue.alloc_device::(1024).await?; let kernel = queue .get_context() - .create_kernel_bundle_from_source(IOTA_SRC) - .build() - .get_kernel("iota"); + .create_kernel_bundle_from_source(IOTA_SRC)? + .build()? + .get_kernel("iota")?; unsafe { queue.launch( @@ -38,15 +38,17 @@ async fn main() { &kernel, (3.14_f32, &mut device_buffer), ) - } - .await; + }? + .await?; - let mut host_buffer = queue.alloc_host::(1024).await; + let mut host_buffer = queue.alloc_host::(1024).await?; - queue.copy(&device_buffer, &mut host_buffer).await; + queue.copy(&device_buffer, &mut host_buffer).await?; for e in host_buffer.iter() { print!("{e} "); } println!(); + + Ok(()) } diff --git a/oneapi-rs/examples/kernel_launch_derive.rs b/oneapi-rs/examples/kernel_launch_derive.rs index 66e628d..c763f7d 100644 --- a/oneapi-rs/examples/kernel_launch_derive.rs +++ b/oneapi-rs/examples/kernel_launch_derive.rs @@ -27,15 +27,15 @@ struct IotaArgs<'a> { ptr: &'a mut SharedBuffer, } -fn main() { +fn main() -> oneapi_rs::Result<()> { let mut queue = Queue::new(); - let mut buffer = queue.alloc_shared::(1024).wait(); + let mut buffer = queue.alloc_shared::(1024).wait()?; let kernel = queue .get_context() - .create_kernel_bundle_from_source(IOTA_SRC) - .build() - .get_kernel("iota"); + .create_kernel_bundle_from_source(IOTA_SRC)? + .build()? + .get_kernel("iota")?; unsafe { queue.launch( @@ -46,11 +46,13 @@ fn main() { ptr: &mut buffer, }, ) - } - .wait(); + }? + .wait()?; for e in buffer.iter() { print!("{e} "); } println!(); + + Ok(()) } diff --git a/oneapi-rs/src/buffer.rs b/oneapi-rs/src/buffer.rs index fb4cfae..be86382 100644 --- a/oneapi-rs/src/buffer.rs +++ b/oneapi-rs/src/buffer.rs @@ -19,6 +19,7 @@ use bytemuck::Pod; use pin_project::pin_project; use crate::{ + Result, event::{Event, EventFuture}, kernel::KernelArgument, usm::{ @@ -121,9 +122,8 @@ impl EnqueuedBuffer { impl EnqueuedBuffer { /// Waits for [`Buffer`] initialization to finish. - pub fn wait(mut self) -> Buffer { - self.event.wait(); - self.buffer + pub fn wait(mut self) -> Result> { + self.event.wait().map(|_| self.buffer) } } @@ -140,17 +140,17 @@ pub struct BufferFuture { } impl Future for BufferFuture { - type Output = Buffer; + type Output = Result>; fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { let this = self.project(); this.event_future .poll(cx) - .map(|_| this.buffer.take().unwrap()) + .map(|result| result.map(|_| this.buffer.take().unwrap())) } } impl IntoFuture for EnqueuedBuffer { - type Output = Buffer; + type Output = Result>; type IntoFuture = BufferFuture; fn into_future(self) -> Self::IntoFuture { diff --git a/oneapi-rs/src/context.rs b/oneapi-rs/src/context.rs index 0dfff64..17f7a42 100644 --- a/oneapi-rs/src/context.rs +++ b/oneapi-rs/src/context.rs @@ -8,7 +8,7 @@ use oneapi_rs_sys::{context::ffi, kernel_bundle, types::ffi::DevicePtr}; -use crate::{device::Device, kernel::SourceKernelBundle}; +use crate::{Result, device::Device, kernel::SourceKernelBundle}; /// A context represents the runtime data structures and state required by a SYCL backend API /// to interact with a group of devices associated with a platform. @@ -32,7 +32,7 @@ impl Context { ffi::new_context(devices).into() } - pub fn create_kernel_bundle_from_source(&self, source: &str) -> SourceKernelBundle { - kernel_bundle::ffi::create_kernel_bundle_from_source(&self.0, source).into() + pub fn create_kernel_bundle_from_source(&self, source: &str) -> Result { + kernel_bundle::ffi::create_kernel_bundle_from_source(&self.0, source).map(Into::into) } } diff --git a/oneapi-rs/src/event.rs b/oneapi-rs/src/event.rs index 9e8e1c3..f2a9af1 100644 --- a/oneapi-rs/src/event.rs +++ b/oneapi-rs/src/event.rs @@ -16,13 +16,13 @@ use oneapi_rs_sys::{event::ffi, types::SharedWaker}; use pin_project::pin_project; -use crate::{info::InfoTarget, private::Sealed, queue::Queue}; +use crate::{Result, info::InfoTarget, private::Sealed, queue::Queue}; pub struct Event(pub(crate) cxx::UniquePtr); impl Event { - pub fn wait(&mut self) { - ffi::wait(&mut self.0); + pub fn wait(&mut self) -> Result<()> { + ffi::wait(&mut self.0) } } @@ -50,7 +50,7 @@ pub struct EventFuture { } impl Future for EventFuture { - type Output = (); + type Output = Result<()>; fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { let this = self.project(); @@ -67,7 +67,8 @@ impl Future for EventFuture { } else { // Quick check before registering to avoid wasting time if this.shared.done.load(Relaxed) { - return Poll::Ready(()); + // The event finished - waiting for it returns immediately + return Poll::Ready(this.event.wait()); } this.shared.waker.register(cx.waker()); @@ -76,7 +77,8 @@ impl Future for EventFuture { // Check the event again to avoid a race condition // https://docs.rs/futures/latest/futures/task/struct.AtomicWaker.html#examples if this.shared.done.load(Relaxed) { - Poll::Ready(()) + // The event finished - waiting for it returns immediately + Poll::Ready(this.event.wait()) } else { Poll::Pending } @@ -84,7 +86,7 @@ impl Future for EventFuture { } impl IntoFuture for Event { - type Output = (); + type Output = Result<()>; type IntoFuture = EventFuture; fn into_future(self) -> Self::IntoFuture { diff --git a/oneapi-rs/src/kernel.rs b/oneapi-rs/src/kernel.rs index 60b425e..7d82819 100644 --- a/oneapi-rs/src/kernel.rs +++ b/oneapi-rs/src/kernel.rs @@ -6,6 +6,7 @@ // SPDX-License-Identifier: MIT OR Apache-2.0 // +use crate::Result; use bytemuck::Pod; use oneapi_rs_sys::{kernel_bundle::ffi, types}; @@ -19,8 +20,8 @@ impl From> for SourceKernelBundle } impl SourceKernelBundle { - pub fn build(&mut self) -> ExecutableKernelBundle { - ffi::build(&mut self.0).into() + pub fn build(&mut self) -> Result { + ffi::build(&mut self.0).map(Into::into) } } @@ -34,8 +35,8 @@ impl From> for ExecutableKern } impl ExecutableKernelBundle { - pub fn get_kernel(&mut self, name: &str) -> Kernel { - ffi::get_kernel(&mut self.0, name).into() + pub fn get_kernel(&mut self, name: &str) -> Result { + ffi::get_kernel(&mut self.0, name).map(Into::into) } } diff --git a/oneapi-rs/src/lib.rs b/oneapi-rs/src/lib.rs index 46e5dad..68b11ef 100644 --- a/oneapi-rs/src/lib.rs +++ b/oneapi-rs/src/lib.rs @@ -103,6 +103,9 @@ pub mod queue; pub mod range; pub mod usm; +pub type SyclError = cxx::Exception; +pub type Result = std::result::Result; + mod private { pub trait Sealed {} } diff --git a/oneapi-rs/src/queue.rs b/oneapi-rs/src/queue.rs index 6235cdf..2ff117a 100644 --- a/oneapi-rs/src/queue.rs +++ b/oneapi-rs/src/queue.rs @@ -6,6 +6,7 @@ // SPDX-License-Identifier: MIT OR Apache-2.0 // +use crate::Result; use bytemuck::Pod; use oneapi_rs_sys::{queue::ffi, types::ffi::EventPtr}; @@ -121,24 +122,24 @@ impl Queue { } /// Submits a barrier to the queue. - pub fn barrier(&mut self) -> Event { - self.barrier_with_deps(&[]) + pub fn barrier(&mut self) -> Result { + self.barrier_with_deps(&[]).map(Into::into) } /// Submits a barrier to the queue after all specified events finish. - pub fn barrier_with_deps(&mut self, dep_events: &[&Event]) -> Event { + pub fn barrier_with_deps(&mut self, dep_events: &[&Event]) -> Result { let dep_events = dep_events .iter() .map(|e| EventPtr { ptr: (*e).clone().0, }) .collect::>(); - ffi::barrier(&mut self.0, dep_events).into() + ffi::barrier(&mut self.0, dep_events).map(Into::into) } /// Performs a blocking wait for the completion of all enqueued tasks in the queue. - pub fn wait(&mut self) { - ffi::wait(&mut self.0); + pub fn wait(&mut self) -> Result<()> { + ffi::wait(&mut self.0) } /// Enqueues a kernel object to the queue as an ND-range kernel, using the number of work-items @@ -148,7 +149,7 @@ impl Queue { nd_range: NdRange, kernel: &Kernel, args: impl KernelArgumentList, - ) -> Event + ) -> Result where NdRange: ValidDimension, { diff --git a/oneapi-rs/src/range.rs b/oneapi-rs/src/range.rs index cfb656e..1dba104 100644 --- a/oneapi-rs/src/range.rs +++ b/oneapi-rs/src/range.rs @@ -6,6 +6,7 @@ // SPDX-License-Identifier: MIT OR Apache-2.0 // +use crate::Result; use oneapi_rs_sys::types; use crate::{ @@ -45,7 +46,7 @@ pub trait ValidDimension: Sealed { queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList, - ) -> Event; + ) -> Result; } impl Sealed for NdRange<1> {} @@ -55,7 +56,7 @@ impl ValidDimension for NdRange<1> { queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList, - ) -> Event { + ) -> Result { unsafe { oneapi_rs_sys::queue::ffi::launch_1d( &mut queue.0, @@ -69,7 +70,7 @@ impl ValidDimension for NdRange<1> { &args.as_raw_arg_list(), ) } - .into() + .map(Into::into) } } @@ -80,7 +81,7 @@ impl ValidDimension for NdRange<2> { queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList, - ) -> Event { + ) -> Result { unsafe { oneapi_rs_sys::queue::ffi::launch_2d( &mut queue.0, @@ -94,7 +95,7 @@ impl ValidDimension for NdRange<2> { &args.as_raw_arg_list(), ) } - .into() + .map(Into::into) } } @@ -105,7 +106,7 @@ impl ValidDimension for NdRange<3> { queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList, - ) -> Event { + ) -> Result { unsafe { oneapi_rs_sys::queue::ffi::launch_3d( &mut queue.0, @@ -119,6 +120,6 @@ impl ValidDimension for NdRange<3> { &args.as_raw_arg_list(), ) } - .into() + .map(Into::into) } }