diff --git a/Cargo.lock b/Cargo.lock index 480b569..e304268 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]] @@ -279,12 +279,23 @@ dependencies = [ "allocator-api2", "bytemuck", "cxx", + "oneapi-rs-derive", "oneapi-rs-sys", "pin-project", "thiserror", "tokio", ] +[[package]] +name = "oneapi-rs-derive" +version = "0.1.0" +dependencies = [ + "proc-macro-crate", + "proc-macro2", + "quote", + "syn 3.0.3", +] + [[package]] name = "oneapi-rs-sys" version = "0.1.0" @@ -312,7 +323,7 @@ checksum = "c96395f0a926bc13b1c17622aaddda1ecb55d49c8f1bf9777e4d877800a43f8b" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -321,20 +332,29 @@ 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.106" +version = "1.0.107" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +checksum = "985e7ec9bb745e6ce6535b544d84d6cd6f7ad8bd711c398938ae983b91a766d9" dependencies = [ "unicode-ident", ] [[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 +392,7 @@ checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -404,6 +424,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 +461,7 @@ checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.118", ] [[package]] @@ -451,7 +482,37 @@ checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" dependencies = [ "proc-macro2", "quote", - "syn", + "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]] @@ -498,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/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..ff171cf --- /dev/null +++ b/oneapi-rs-derive/Cargo.toml @@ -0,0 +1,13 @@ +[package] +name = "oneapi-rs-derive" +version = "0.1.0" +edition = "2024" + +[lib] +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 new file mode 100644 index 0000000..7f99a62 --- /dev/null +++ b/oneapi-rs-derive/src/lib.rs @@ -0,0 +1,84 @@ +use proc_macro::TokenStream; +use proc_macro_crate::{FoundCrate, crate_name}; +use quote::{format_ident, quote}; +use syn::{ + 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, &oneapi); + let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl(); + + let ident = input.ident; + let argc = data.fields.len(); + let members = data.fields.members(); + + let expanded = quote! { + 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() }),* ] + } + } + }; + + TokenStream::from(expanded) +} + +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::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::>(); + + quote! { + 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() }),* ] + } + } + } +} + +/// 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); + 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/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..c7fe823 100644 --- a/oneapi-rs/examples/kernel_launch.rs +++ b/oneapi-rs/examples/kernel_launch.rs @@ -6,13 +6,7 @@ // SPDX-License-Identifier: MIT OR Apache-2.0 // -use oneapi_rs::{ - buffer::Buffer, - kernel::{KernelArgument, KernelArgumentList}, - queue::Queue, - range::NdRange, - usm::{SharedAllocator, UsmAllocator}, -}; +use oneapi_rs::{queue::Queue, range::NdRange}; static IOTA_SRC: &str = r#" #include @@ -21,28 +15,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); } "#; -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(); + let mut buffer = queue.alloc_shared::(1024).wait(); let kernel = queue .get_context() @@ -50,17 +31,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/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!(); +} 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 0e5d6a4..bf7a278 100644 --- a/oneapi-rs/src/kernel.rs +++ b/oneapi-rs/src/kernel.rs @@ -63,3 +63,21 @@ unsafe impl KernelArgument for T { 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() }] + } +} + +pub use oneapi_rs_derive::KernelArgumentList; + +use oneapi_rs_derive::impl_arg_list_for_tuples; + +impl_arg_list_for_tuples! {16}