From 0e26964b04522eaf9ee4a799113c5566d96d86b0 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 09:07:43 +0000 Subject: [PATCH 01/16] Add a new proc-macro crate --- Cargo.lock | 39 +++++++++++++++++++++++++++---------- Cargo.toml | 2 +- oneapi-rs-derive/Cargo.toml | 11 +++++++++++ oneapi-rs-derive/src/lib.rs | 14 +++++++++++++ 4 files changed, 55 insertions(+), 11 deletions(-) create mode 100644 oneapi-rs-derive/Cargo.toml create mode 100644 oneapi-rs-derive/src/lib.rs diff --git a/Cargo.lock b/Cargo.lock index 480b569..dde5de3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -94,7 +94,7 @@ dependencies = [ "proc-macro2", "quote", "scratch", - "syn", + "syn 2.0.118", ] [[package]] @@ -108,7 +108,7 @@ dependencies = [ "indexmap", "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -126,7 +126,7 @@ dependencies = [ "indexmap", "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -203,7 +203,7 @@ checksum = "2d6d3cde68c518367be28956066ddfef33813991b77a55005a69dae04bf3b10b" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -285,6 +285,14 @@ dependencies = [ "tokio", ] +[[package]] +name = "oneapi-rs-derive" +version = "0.1.0" +dependencies = [ + "quote", + "syn 3.0.3", +] + [[package]] name = "oneapi-rs-sys" version = "0.1.0" @@ -312,7 +320,7 @@ checksum = "c96395f0a926bc13b1c17622aaddda1ecb55d49c8f1bf9777e4d877800a43f8b" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -332,9 +340,9 @@ dependencies = [ [[package]] name = "quote" -version = "1.0.46" +version = "1.0.47" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dfbc457d0c7a0759a614551b11a6409e5951f6c7537be1f1b7682b9ae9230368" +checksum = "1fbf4db142a473a8d80c26bbf18454ed458bf8d26c8219c331daecfdbd079001" dependencies = [ "proc-macro2", ] @@ -372,7 +380,7 @@ checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -404,6 +412,17 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "syn" +version = "3.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "53e9bae58849f64dfa4f5d5ae372c8341f7305f82a3868709269343628b659a3" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + [[package]] name = "termcolor" version = "1.4.1" @@ -430,7 +449,7 @@ checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -451,7 +470,7 @@ checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 0e6831f..b4bc193 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,3 +1,3 @@ [workspace] resolver = "3" -members = ["oneapi-rs","oneapi-rs-sys"] +members = ["oneapi-rs", "oneapi-rs-derive","oneapi-rs-sys"] diff --git a/oneapi-rs-derive/Cargo.toml b/oneapi-rs-derive/Cargo.toml new file mode 100644 index 0000000..ed30b2e --- /dev/null +++ b/oneapi-rs-derive/Cargo.toml @@ -0,0 +1,11 @@ +[package] +name = "oneapi-rs-derive" +version = "0.1.0" +edition = "2024" + +[lib] +proc-macro = true + +[dependencies] +quote = "1.0.47" +syn = "3.0.3" diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs new file mode 100644 index 0000000..b93cf3f --- /dev/null +++ b/oneapi-rs-derive/src/lib.rs @@ -0,0 +1,14 @@ +pub fn add(left: u64, right: u64) -> u64 { + left + right +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn it_works() { + let result = add(2, 2); + assert_eq!(result, 4); + } +} From 976e91b7d2946f7bff5e3ca3158af6b4b0749136 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 09:12:29 +0000 Subject: [PATCH 02/16] Add basic macro structure --- oneapi-rs-derive/src/lib.rs | 22 +++++++++++----------- 1 file changed, 11 insertions(+), 11 deletions(-) diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index b93cf3f..d41d53f 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -1,14 +1,14 @@ -pub fn add(left: u64, right: u64) -> u64 { - left + right -} +use proc_macro::TokenStream; +use quote::quote; +use syn::{parse_macro_input, DeriveInput}; + +#[proc_macro_derive(KernelArgumentList)] +pub fn my_macro(input: TokenStream) -> TokenStream { + let input = parse_macro_input!(input as DeriveInput); -#[cfg(test)] -mod tests { - use super::*; + let expanded = quote! { + // ... + }; - #[test] - fn it_works() { - let result = add(2, 2); - assert_eq!(result, 4); - } + TokenStream::from(expanded) } From e006eceea3c01d171178d83cc5fc08bf5abfd0e6 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 09:55:17 +0000 Subject: [PATCH 03/16] Add basic derive macro implementation for structs --- oneapi-rs-derive/src/lib.rs | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index d41d53f..361023d 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -1,13 +1,24 @@ use proc_macro::TokenStream; use quote::quote; -use syn::{parse_macro_input, DeriveInput}; +use syn::{Data, DeriveInput, parse_macro_input}; #[proc_macro_derive(KernelArgumentList)] pub fn my_macro(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as DeriveInput); + let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); + + let Data::Struct(data) = input.data else { panic!() }; + let ident = input.ident; + let argc = data.fields.len(); + let members = data.fields.members(); + let expanded = quote! { - // ... + unsafe impl #impl_generics KernelArgumentList<#argc> for #ident #ty_generics #where_clause { + unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] { + [ #(self.#members.as_raw_arg()),* ] + } + } }; TokenStream::from(expanded) From aaad0f8fad30c458462b48bde58ff8d760fbe5b7 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 10:08:35 +0000 Subject: [PATCH 04/16] Add derive macro to main crate --- Cargo.lock | 1 + oneapi-rs/Cargo.toml | 1 + oneapi-rs/examples/kernel_launch.rs | 9 +-------- oneapi-rs/src/kernel.rs | 2 ++ 4 files changed, 5 insertions(+), 8 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index dde5de3..3a92192 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -279,6 +279,7 @@ dependencies = [ "allocator-api2", "bytemuck", "cxx", + "oneapi-rs-derive", "oneapi-rs-sys", "pin-project", "thiserror", diff --git a/oneapi-rs/Cargo.toml b/oneapi-rs/Cargo.toml index 4a9876e..ed5df30 100644 --- a/oneapi-rs/Cargo.toml +++ b/oneapi-rs/Cargo.toml @@ -9,6 +9,7 @@ allocator-api2 = "0.4.0" bytemuck = "1.25.1" cxx = "1.0.194" oneapi-rs-sys = { path = "../oneapi-rs-sys" } +oneapi-rs-derive = { path = "../oneapi-rs-derive" } pin-project = "1.1.13" thiserror = "2.0.18" diff --git a/oneapi-rs/examples/kernel_launch.rs b/oneapi-rs/examples/kernel_launch.rs index 6224b27..4325385 100644 --- a/oneapi-rs/examples/kernel_launch.rs +++ b/oneapi-rs/examples/kernel_launch.rs @@ -27,19 +27,12 @@ void iota(float start, float *ptr) { } "#; +#[derive(KernelArgumentList)] struct IotaArgs<'a> { start: f32, buffer: &'a mut Buffer>, } -unsafe impl<'a> KernelArgumentList<2> for IotaArgs<'a> { - unsafe fn as_raw_arg_list(&self) -> [&[u8]; 2] { - return [unsafe { self.start.as_raw_arg() }, unsafe { - self.buffer.as_raw_arg() - }]; - } -} - fn main() { let mut queue = Queue::new(); let mut buffer = queue.alloc_shared::(1024).wait(); diff --git a/oneapi-rs/src/kernel.rs b/oneapi-rs/src/kernel.rs index 0e5d6a4..f9f0947 100644 --- a/oneapi-rs/src/kernel.rs +++ b/oneapi-rs/src/kernel.rs @@ -63,3 +63,5 @@ unsafe impl KernelArgument for T { pub unsafe trait KernelArgumentList { unsafe fn as_raw_arg_list(&self) -> [&[u8]; ARGC]; } + +pub use oneapi_rs_derive::KernelArgumentList; From bf885a0f29f09c20bfea8a31ad6c071ab42de24d Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 13:45:00 +0000 Subject: [PATCH 05/16] Add tuple argument list support --- Cargo.lock | 5 +++-- oneapi-rs-derive/Cargo.toml | 1 + oneapi-rs-derive/src/lib.rs | 34 ++++++++++++++++++++++++++--- oneapi-rs/examples/kernel_launch.rs | 27 ++++------------------- oneapi-rs/src/buffer.rs | 22 ++++++++++++++++--- oneapi-rs/src/kernel.rs | 4 ++++ 6 files changed, 62 insertions(+), 31 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 3a92192..de6c610 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -290,6 +290,7 @@ dependencies = [ name = "oneapi-rs-derive" version = "0.1.0" dependencies = [ + "proc-macro2", "quote", "syn 3.0.3", ] @@ -332,9 +333,9 @@ checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" [[package]] name = "proc-macro2" -version = "1.0.106" +version = "1.0.107" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +checksum = "985e7ec9bb745e6ce6535b544d84d6cd6f7ad8bd711c398938ae983b91a766d9" dependencies = [ "unicode-ident", ] diff --git a/oneapi-rs-derive/Cargo.toml b/oneapi-rs-derive/Cargo.toml index ed30b2e..1ca1601 100644 --- a/oneapi-rs-derive/Cargo.toml +++ b/oneapi-rs-derive/Cargo.toml @@ -7,5 +7,6 @@ edition = "2024" proc-macro = true [dependencies] +proc-macro2 = "1.0.107" quote = "1.0.47" syn = "3.0.3" diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index 361023d..8b7f5b3 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -1,6 +1,6 @@ use proc_macro::TokenStream; -use quote::quote; -use syn::{Data, DeriveInput, parse_macro_input}; +use quote::{format_ident, quote}; +use syn::{Data, DeriveInput, LitInt, parse_macro_input}; #[proc_macro_derive(KernelArgumentList)] pub fn my_macro(input: TokenStream) -> TokenStream { @@ -16,10 +16,38 @@ pub fn my_macro(input: TokenStream) -> TokenStream { let expanded = quote! { unsafe impl #impl_generics KernelArgumentList<#argc> for #ident #ty_generics #where_clause { unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] { - [ #(self.#members.as_raw_arg()),* ] + [ #(unsafe { self.#members.as_raw_arg() }),* ] } } }; TokenStream::from(expanded) } + +fn get_single_tuple_impl(argc: usize) -> proc_macro2::TokenStream { + let iter = {0..argc}.map(syn::Index::from); + let types = {0..argc}.map(|i| format_ident!("T{i}")).collect::>(); + + quote! { + unsafe impl<#(#types),*> KernelArgumentList<#argc> for (#(#types),*) + where #(#types: KernelArgument),* { + unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] { + [ #(unsafe { self.#iter.as_raw_arg() }),* ] + } + } + } +} + +#[proc_macro] +pub fn impl_arg_list_for_tuples(input: TokenStream) -> TokenStream { + let input = parse_macro_input!(input as LitInt); + let argc = input.base10_parse::().unwrap(); + + let impls = {2..argc}.map(get_single_tuple_impl); + + let expanded = quote! { + #(#impls)* + }; + + TokenStream::from(expanded) +} diff --git a/oneapi-rs/examples/kernel_launch.rs b/oneapi-rs/examples/kernel_launch.rs index 4325385..d9fcb4a 100644 --- a/oneapi-rs/examples/kernel_launch.rs +++ b/oneapi-rs/examples/kernel_launch.rs @@ -7,11 +7,8 @@ // use oneapi_rs::{ - buffer::Buffer, - kernel::{KernelArgument, KernelArgumentList}, queue::Queue, range::NdRange, - usm::{SharedAllocator, UsmAllocator}, }; static IOTA_SRC: &str = r#" @@ -21,21 +18,15 @@ namespace syclexp = sycl::ext::oneapi::experimental; extern "C" SYCL_EXT_ONEAPI_FUNCTION_PROPERTY((syclexp::nd_range_kernel<1>)) -void iota(float start, float *ptr) { +void iota(double start, double *ptr) { size_t id = syclext::this_work_item::get_nd_item<1>().get_global_linear_id(); - ptr[id] = start + static_cast(id); + ptr[id] = start + static_cast(id); } "#; -#[derive(KernelArgumentList)] -struct IotaArgs<'a> { - start: f32, - buffer: &'a mut Buffer>, -} - fn main() { 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() @@ -43,17 +34,7 @@ fn main() { .build() .get_kernel("iota"); - unsafe { - queue.launch( - NdRange::new([1024], [16]), - &kernel, - IotaArgs { - start: 3.14, - buffer: &mut buffer, - }, - ) - } - .wait(); + unsafe { queue.launch(NdRange::new([1024], [16]), &kernel, (3.14, &mut buffer)) }.wait(); for e in buffer.iter() { print!("{e} "); diff --git a/oneapi-rs/src/buffer.rs b/oneapi-rs/src/buffer.rs index d9d7a16..6d67adf 100644 --- a/oneapi-rs/src/buffer.rs +++ b/oneapi-rs/src/buffer.rs @@ -68,6 +68,12 @@ impl Buffer { pub(crate) fn get_byte_size(&self) -> usize { self.layout.size() } + + unsafe fn as_raw_arg_impl(&self) -> &[u8] { + let data_ptr: *const NonNull<_> = &self.data; + let cast_ptr = data_ptr as *const u8; + unsafe { slice::from_raw_parts(cast_ptr, std::mem::size_of_val(&cast_ptr)) } + } } impl Deref for Buffer { @@ -143,8 +149,18 @@ impl IntoFuture for EnqueuedBuffer { unsafe impl KernelArgument for Buffer { unsafe fn as_raw_arg(&self) -> &[u8] { - let data_ptr: *const NonNull<_> = &self.data; - let cast_ptr = data_ptr as *const u8; - unsafe { slice::from_raw_parts(cast_ptr, std::mem::size_of_val(&cast_ptr)) } + unsafe { self.as_raw_arg_impl() } + } +} + +unsafe impl KernelArgument for &Buffer { + unsafe fn as_raw_arg(&self) -> &[u8] { + unsafe { self.as_raw_arg_impl() } + } +} + +unsafe impl KernelArgument for &mut Buffer { + unsafe fn as_raw_arg(&self) -> &[u8] { + unsafe { self.as_raw_arg_impl() } } } diff --git a/oneapi-rs/src/kernel.rs b/oneapi-rs/src/kernel.rs index f9f0947..c87ead3 100644 --- a/oneapi-rs/src/kernel.rs +++ b/oneapi-rs/src/kernel.rs @@ -65,3 +65,7 @@ pub unsafe trait KernelArgumentList { } pub use oneapi_rs_derive::KernelArgumentList; + +use oneapi_rs_derive::impl_arg_list_for_tuples; + +impl_arg_list_for_tuples!{16} From 910cbf37374b5c3abf7030f406a38027a8eb47b4 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 14:23:11 +0000 Subject: [PATCH 06/16] Add trait bounds to derive macro --- oneapi-rs-derive/src/lib.rs | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index 8b7f5b3..2eec8a5 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -1,14 +1,15 @@ use proc_macro::TokenStream; use quote::{format_ident, quote}; -use syn::{Data, DeriveInput, LitInt, parse_macro_input}; +use syn::{Data, DataStruct, DeriveInput, Field, LitInt, WhereClause, parse_macro_input, parse_quote}; #[proc_macro_derive(KernelArgumentList)] pub fn my_macro(input: TokenStream) -> TokenStream { - let input = parse_macro_input!(input as DeriveInput); + let mut input = parse_macro_input!(input as DeriveInput); + let Data::Struct(data) = &input.data else { panic!() }; + expand_where_clause(input.generics.make_where_clause(), data); let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); - let Data::Struct(data) = input.data else { panic!() }; let ident = input.ident; let argc = data.fields.len(); let members = data.fields.members(); @@ -24,6 +25,12 @@ pub fn my_macro(input: TokenStream) -> TokenStream { TokenStream::from(expanded) } +fn expand_where_clause(where_clause: &mut WhereClause, data: &DataStruct) { + for Field { ty, .. } in &data.fields { + where_clause.predicates.push(parse_quote!(#ty: oneapi_rs::kernel::KernelArgument)); + } +} + fn get_single_tuple_impl(argc: usize) -> proc_macro2::TokenStream { let iter = {0..argc}.map(syn::Index::from); let types = {0..argc}.map(|i| format_ident!("T{i}")).collect::>(); From 171372da75788b52e12063f3cfeb2ee18f091bac Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 14:32:01 +0000 Subject: [PATCH 07/16] Use full module paths inside macros --- oneapi-rs-derive/src/lib.rs | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index 2eec8a5..3ea4663 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -15,7 +15,7 @@ pub fn my_macro(input: TokenStream) -> TokenStream { let members = data.fields.members(); let expanded = quote! { - unsafe impl #impl_generics KernelArgumentList<#argc> for #ident #ty_generics #where_clause { + unsafe impl #impl_generics oneapi_rs::kernel::KernelArgumentList<#argc> for #ident #ty_generics #where_clause { unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] { [ #(unsafe { self.#members.as_raw_arg() }),* ] } @@ -36,8 +36,8 @@ fn get_single_tuple_impl(argc: usize) -> proc_macro2::TokenStream { let types = {0..argc}.map(|i| format_ident!("T{i}")).collect::>(); quote! { - unsafe impl<#(#types),*> KernelArgumentList<#argc> for (#(#types),*) - where #(#types: KernelArgument),* { + unsafe impl<#(#types),*> crate::kernel::KernelArgumentList<#argc> for (#(#types),*) + where #(#types: crate::kernel::KernelArgument),* { unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] { [ #(unsafe { self.#iter.as_raw_arg() }),* ] } From a08a47860f2865fbfd8cb392d5bc52d73db982d4 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 14:37:23 +0000 Subject: [PATCH 08/16] Rename derive macro function --- oneapi-rs-derive/src/lib.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index 3ea4663..0eb4ef7 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -3,7 +3,7 @@ use quote::{format_ident, quote}; use syn::{Data, DataStruct, DeriveInput, Field, LitInt, WhereClause, parse_macro_input, parse_quote}; #[proc_macro_derive(KernelArgumentList)] -pub fn my_macro(input: TokenStream) -> TokenStream { +pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream { let mut input = parse_macro_input!(input as DeriveInput); let Data::Struct(data) = &input.data else { panic!() }; From 2242309e6833c4f3ccc6abe97194a232848645c2 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 14:51:52 +0000 Subject: [PATCH 09/16] Make tuple impl generation range inclusive --- oneapi-rs-derive/src/lib.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index 0eb4ef7..ea7a6ac 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -50,7 +50,7 @@ pub fn impl_arg_list_for_tuples(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as LitInt); let argc = input.base10_parse::().unwrap(); - let impls = {2..argc}.map(get_single_tuple_impl); + let impls = {2..=argc}.map(get_single_tuple_impl); let expanded = quote! { #(#impls)* From 99918fbb37ad9f1c1549ecac99ddeb8c952165fa Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 14:56:25 +0000 Subject: [PATCH 10/16] Add blanket KernelArgumentList impl for single arguments --- oneapi-rs/src/kernel.rs | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/oneapi-rs/src/kernel.rs b/oneapi-rs/src/kernel.rs index c87ead3..62f214e 100644 --- a/oneapi-rs/src/kernel.rs +++ b/oneapi-rs/src/kernel.rs @@ -64,6 +64,12 @@ pub unsafe trait KernelArgumentList { unsafe fn as_raw_arg_list(&self) -> [&[u8]; ARGC]; } +unsafe impl KernelArgumentList<1> for T { + unsafe fn as_raw_arg_list(&self) -> [&[u8]; 1] { + [ unsafe { self.as_raw_arg() } ] + } +} + pub use oneapi_rs_derive::KernelArgumentList; use oneapi_rs_derive::impl_arg_list_for_tuples; From 8373c70a3b1e14795c1cd847a9533729ffbe6755 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 14:58:05 +0000 Subject: [PATCH 11/16] cargo fmt --- oneapi-rs-derive/src/lib.rs | 20 ++++++++++++++------ oneapi-rs/examples/kernel_launch.rs | 5 +---- oneapi-rs/src/kernel.rs | 4 ++-- 3 files changed, 17 insertions(+), 12 deletions(-) diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index ea7a6ac..6f284b7 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -1,12 +1,16 @@ use proc_macro::TokenStream; use quote::{format_ident, quote}; -use syn::{Data, DataStruct, DeriveInput, Field, LitInt, WhereClause, parse_macro_input, parse_quote}; +use syn::{ + Data, DataStruct, DeriveInput, Field, LitInt, WhereClause, parse_macro_input, parse_quote, +}; #[proc_macro_derive(KernelArgumentList)] pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream { let mut input = parse_macro_input!(input as DeriveInput); - let Data::Struct(data) = &input.data else { panic!() }; + let Data::Struct(data) = &input.data else { + panic!() + }; expand_where_clause(input.generics.make_where_clause(), data); let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); @@ -27,13 +31,17 @@ pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream { fn expand_where_clause(where_clause: &mut WhereClause, data: &DataStruct) { for Field { ty, .. } in &data.fields { - where_clause.predicates.push(parse_quote!(#ty: oneapi_rs::kernel::KernelArgument)); + where_clause + .predicates + .push(parse_quote!(#ty: oneapi_rs::kernel::KernelArgument)); } } fn get_single_tuple_impl(argc: usize) -> proc_macro2::TokenStream { - let iter = {0..argc}.map(syn::Index::from); - let types = {0..argc}.map(|i| format_ident!("T{i}")).collect::>(); + let iter = { 0..argc }.map(syn::Index::from); + let types = { 0..argc } + .map(|i| format_ident!("T{i}")) + .collect::>(); quote! { unsafe impl<#(#types),*> crate::kernel::KernelArgumentList<#argc> for (#(#types),*) @@ -50,7 +58,7 @@ pub fn impl_arg_list_for_tuples(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as LitInt); let argc = input.base10_parse::().unwrap(); - let impls = {2..=argc}.map(get_single_tuple_impl); + let impls = { 2..=argc }.map(get_single_tuple_impl); let expanded = quote! { #(#impls)* diff --git a/oneapi-rs/examples/kernel_launch.rs b/oneapi-rs/examples/kernel_launch.rs index d9fcb4a..c7fe823 100644 --- a/oneapi-rs/examples/kernel_launch.rs +++ b/oneapi-rs/examples/kernel_launch.rs @@ -6,10 +6,7 @@ // SPDX-License-Identifier: MIT OR Apache-2.0 // -use oneapi_rs::{ - queue::Queue, - range::NdRange, -}; +use oneapi_rs::{queue::Queue, range::NdRange}; static IOTA_SRC: &str = r#" #include diff --git a/oneapi-rs/src/kernel.rs b/oneapi-rs/src/kernel.rs index 62f214e..c4f62f2 100644 --- a/oneapi-rs/src/kernel.rs +++ b/oneapi-rs/src/kernel.rs @@ -66,7 +66,7 @@ pub unsafe trait KernelArgumentList { unsafe impl KernelArgumentList<1> for T { unsafe fn as_raw_arg_list(&self) -> [&[u8]; 1] { - [ unsafe { self.as_raw_arg() } ] + [unsafe { self.as_raw_arg() }] } } @@ -74,4 +74,4 @@ pub use oneapi_rs_derive::KernelArgumentList; use oneapi_rs_derive::impl_arg_list_for_tuples; -impl_arg_list_for_tuples!{16} +impl_arg_list_for_tuples! {16} From da96fd48f7fa18bca75a91fd62e7c28e80c20055 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Fri, 24 Jul 2026 15:12:49 +0000 Subject: [PATCH 12/16] Add error message to derive macro --- oneapi-rs-derive/src/lib.rs | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index 6f284b7..87f0071 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -1,7 +1,8 @@ use proc_macro::TokenStream; use quote::{format_ident, quote}; use syn::{ - Data, DataStruct, DeriveInput, Field, LitInt, WhereClause, parse_macro_input, parse_quote, + Data, DataStruct, DeriveInput, Error, Field, LitInt, WhereClause, parse_macro_input, + parse_quote, }; #[proc_macro_derive(KernelArgumentList)] @@ -9,7 +10,9 @@ pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream { let mut input = parse_macro_input!(input as DeriveInput); let Data::Struct(data) = &input.data else { - panic!() + return Error::new_spanned(input, "This derive macro only works on structs.") + .into_compile_error() + .into(); }; expand_where_clause(input.generics.make_where_clause(), data); let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); @@ -19,7 +22,8 @@ pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream { let members = data.fields.members(); let expanded = quote! { - unsafe impl #impl_generics oneapi_rs::kernel::KernelArgumentList<#argc> for #ident #ty_generics #where_clause { + unsafe impl #impl_generics oneapi_rs::kernel::KernelArgumentList<#argc> + for #ident #ty_generics #where_clause { unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] { [ #(unsafe { self.#members.as_raw_arg() }),* ] } From d853e11e61f84973f4d427df6a4c4ca5c02e8bbb Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Mon, 27 Jul 2026 12:33:43 +0000 Subject: [PATCH 13/16] Add macro documentation --- oneapi-rs-derive/src/lib.rs | 2 ++ 1 file changed, 2 insertions(+) diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index 87f0071..8d26e90 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -5,6 +5,7 @@ use syn::{ parse_quote, }; +/// Derive macro generating an impl of the `KernelArgumentList` trait for a given struct. #[proc_macro_derive(KernelArgumentList)] pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream { let mut input = parse_macro_input!(input as DeriveInput); @@ -57,6 +58,7 @@ fn get_single_tuple_impl(argc: usize) -> proc_macro2::TokenStream { } } +/// A macro that generates tuples from 2..N which implement the `KernelArgumentList` trait. #[proc_macro] pub fn impl_arg_list_for_tuples(input: TokenStream) -> TokenStream { let input = parse_macro_input!(input as LitInt); From 757b0b49789a93f769ad68b06db63bedd99d9465 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Mon, 27 Jul 2026 12:33:53 +0000 Subject: [PATCH 14/16] Add derive macro example --- oneapi-rs/examples/kernel_launch_derive.rs | 62 ++++++++++++++++++++++ 1 file changed, 62 insertions(+) create mode 100644 oneapi-rs/examples/kernel_launch_derive.rs diff --git a/oneapi-rs/examples/kernel_launch_derive.rs b/oneapi-rs/examples/kernel_launch_derive.rs new file mode 100644 index 0000000..8be1e8a --- /dev/null +++ b/oneapi-rs/examples/kernel_launch_derive.rs @@ -0,0 +1,62 @@ +// +// Copyright (C) 2026 Intel Corporation +// +// Under the MIT License or the Apache License v2.0. +// See LICENSE-MIT and LICENSE-APACHE for license information. +// SPDX-License-Identifier: MIT OR Apache-2.0 +// + +use oneapi_rs::{ + buffer::Buffer, + kernel::{KernelArgument, KernelArgumentList}, + queue::Queue, + range::NdRange, + usm::{SharedAllocator, UsmAllocator}, +}; + +static IOTA_SRC: &str = r#" +#include +namespace syclext = sycl::ext::oneapi; +namespace syclexp = sycl::ext::oneapi::experimental; + +extern "C" +SYCL_EXT_ONEAPI_FUNCTION_PROPERTY((syclexp::nd_range_kernel<1>)) +void iota(double start, double *ptr) { + size_t id = syclext::this_work_item::get_nd_item<1>().get_global_linear_id(); + ptr[id] = start + static_cast(id); +} +"#; + +#[derive(KernelArgumentList)] +struct IotaArgs<'a> { + start: f64, + ptr: &'a mut Buffer>, +} + +fn main() { + let mut queue = Queue::new(); + let mut buffer = queue.alloc_shared::(1024).wait(); + + let kernel = queue + .get_context() + .create_kernel_bundle_from_source(IOTA_SRC) + .build() + .get_kernel("iota"); + + unsafe { + queue.launch( + NdRange::new([1024], [16]), + &kernel, + IotaArgs { + start: 3.14, + ptr: &mut buffer, + }, + ) + } + .wait(); + + for e in buffer.iter() { + print!("{e} "); + } + println!(); +} From 06db8bb22323f7a6a4ac6f1a73d6f99abf192205 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Mon, 27 Jul 2026 12:46:25 +0000 Subject: [PATCH 15/16] Resolve oneapi_rs crate name at macro-expansion time --- Cargo.lock | 49 +++++++++++++++++++++++++++++++++++++ oneapi-rs-derive/Cargo.toml | 1 + oneapi-rs-derive/src/lib.rs | 20 +++++++++++---- 3 files changed, 65 insertions(+), 5 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index de6c610..e304268 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -290,6 +290,7 @@ dependencies = [ name = "oneapi-rs-derive" version = "0.1.0" dependencies = [ + "proc-macro-crate", "proc-macro2", "quote", "syn 3.0.3", @@ -331,6 +332,15 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" +[[package]] +name = "proc-macro-crate" +version = "3.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e67ba7e9b2b56446f1d419b1d807906278ffa1a658a8a5d8a39dcb1f5a78614f" +dependencies = [ + "toml_edit", +] + [[package]] name = "proc-macro2" version = "1.0.107" @@ -475,6 +485,36 @@ dependencies = [ "syn 2.0.118", ] +[[package]] +name = "toml_datetime" +version = "1.1.1+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3165f65f62e28e0115a00b2ebdd37eb6f3b641855f9d636d3cd4103767159ad7" +dependencies = [ + "serde_core", +] + +[[package]] +name = "toml_edit" +version = "0.25.13+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6975367e4d2ef766d86af01ffad14b622fecc8d4357a998fbc4deb6e9bacaf9b" +dependencies = [ + "indexmap", + "toml_datetime", + "toml_parser", + "winnow", +] + +[[package]] +name = "toml_parser" +version = "1.1.2+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2abe9b86193656635d2411dc43050282ca48aa31c2451210f4202550afb7526" +dependencies = [ + "winnow", +] + [[package]] name = "unicode-ident" version = "1.0.24" @@ -519,3 +559,12 @@ checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" dependencies = [ "windows-link", ] + +[[package]] +name = "winnow" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "23b97319f7b8343df12cc98938e5c3eb436064524c8d2b4e30a1d3a36eecdf81" +dependencies = [ + "memchr", +] diff --git a/oneapi-rs-derive/Cargo.toml b/oneapi-rs-derive/Cargo.toml index 1ca1601..ff171cf 100644 --- a/oneapi-rs-derive/Cargo.toml +++ b/oneapi-rs-derive/Cargo.toml @@ -7,6 +7,7 @@ edition = "2024" proc-macro = true [dependencies] +proc-macro-crate = "3.5.0" proc-macro2 = "1.0.107" quote = "1.0.47" syn = "3.0.3" diff --git a/oneapi-rs-derive/src/lib.rs b/oneapi-rs-derive/src/lib.rs index 8d26e90..7f99a62 100644 --- a/oneapi-rs-derive/src/lib.rs +++ b/oneapi-rs-derive/src/lib.rs @@ -1,21 +1,31 @@ use proc_macro::TokenStream; +use proc_macro_crate::{FoundCrate, crate_name}; use quote::{format_ident, quote}; use syn::{ - Data, DataStruct, DeriveInput, Error, Field, LitInt, WhereClause, parse_macro_input, + Data, DataStruct, DeriveInput, Error, Field, Ident, LitInt, WhereClause, parse_macro_input, parse_quote, }; +fn find_oneapi() -> Ident { + let crate_name = crate_name("oneapi_rs").expect("oneapi_rs is present in Cargo.toml"); + match crate_name { + FoundCrate::Itself => format_ident!("crate"), + FoundCrate::Name(name) => format_ident!("{name}"), + } +} + /// Derive macro generating an impl of the `KernelArgumentList` trait for a given struct. #[proc_macro_derive(KernelArgumentList)] pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream { let mut input = parse_macro_input!(input as DeriveInput); + let oneapi = find_oneapi(); let Data::Struct(data) = &input.data else { return Error::new_spanned(input, "This derive macro only works on structs.") .into_compile_error() .into(); }; - expand_where_clause(input.generics.make_where_clause(), data); + expand_where_clause(input.generics.make_where_clause(), data, &oneapi); let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); let ident = input.ident; @@ -23,7 +33,7 @@ pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream { let members = data.fields.members(); let expanded = quote! { - unsafe impl #impl_generics oneapi_rs::kernel::KernelArgumentList<#argc> + unsafe impl #impl_generics #oneapi::kernel::KernelArgumentList<#argc> for #ident #ty_generics #where_clause { unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] { [ #(unsafe { self.#members.as_raw_arg() }),* ] @@ -34,11 +44,11 @@ pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream { TokenStream::from(expanded) } -fn expand_where_clause(where_clause: &mut WhereClause, data: &DataStruct) { +fn expand_where_clause(where_clause: &mut WhereClause, data: &DataStruct, oneapi: &Ident) { for Field { ty, .. } in &data.fields { where_clause .predicates - .push(parse_quote!(#ty: oneapi_rs::kernel::KernelArgument)); + .push(parse_quote!(#ty: #oneapi::kernel::KernelArgument)); } } From 2d3f0bdafca77775006806784328d6fcc8568ae1 Mon Sep 17 00:00:00 2001 From: Szymon Zadworny Date: Mon, 27 Jul 2026 14:44:41 +0000 Subject: [PATCH 16/16] Add support for zero-sized argument lists --- oneapi-rs/src/kernel.rs | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/oneapi-rs/src/kernel.rs b/oneapi-rs/src/kernel.rs index c4f62f2..bf7a278 100644 --- a/oneapi-rs/src/kernel.rs +++ b/oneapi-rs/src/kernel.rs @@ -64,6 +64,12 @@ pub unsafe trait KernelArgumentList { unsafe fn as_raw_arg_list(&self) -> [&[u8]; ARGC]; } +unsafe impl KernelArgumentList<0> for T { + unsafe fn as_raw_arg_list(&self) -> [&[u8]; 0] { + [] + } +} + unsafe impl KernelArgumentList<1> for T { unsafe fn as_raw_arg_list(&self) -> [&[u8]; 1] { [unsafe { self.as_raw_arg() }]