Skip to content

Commit 36606b0

Browse files
authored
Merge pull request #10 from szymon-zadworny/macros
Macros
2 parents 43f9081 + 28b15b5 commit 36606b0

9 files changed

Lines changed: 286 additions & 50 deletions

File tree

Cargo.lock

Lines changed: 82 additions & 12 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

Cargo.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,3 @@
11
[workspace]
22
resolver = "3"
3-
members = ["oneapi-rs","oneapi-rs-sys"]
3+
members = ["oneapi-rs", "oneapi-rs-derive","oneapi-rs-sys"]

oneapi-rs-derive/Cargo.toml

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
[package]
2+
name = "oneapi-rs-derive"
3+
version = "0.1.0"
4+
edition = "2024"
5+
6+
[lib]
7+
proc-macro = true
8+
9+
[dependencies]
10+
proc-macro-crate = "3.5.0"
11+
proc-macro2 = "1.0.107"
12+
quote = "1.0.47"
13+
syn = "3.0.3"

oneapi-rs-derive/src/lib.rs

Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,85 @@
1+
use proc_macro::TokenStream;
2+
use proc_macro_crate::{FoundCrate, crate_name};
3+
use quote::{format_ident, quote};
4+
use syn::{
5+
Data, DataStruct, DeriveInput, Error, Field, Ident, LitInt, WhereClause, parse_macro_input,
6+
parse_quote,
7+
};
8+
9+
fn find_oneapi() -> Ident {
10+
let crate_name =
11+
crate_name("oneapi-rs").expect("Expected oneapi-rs to be present in Cargo.toml");
12+
match crate_name {
13+
FoundCrate::Itself => format_ident!("oneapi_rs"),
14+
FoundCrate::Name(name) => format_ident!("{name}"),
15+
}
16+
}
17+
18+
/// Derive macro generating an impl of the `KernelArgumentList` trait for a given struct.
19+
#[proc_macro_derive(KernelArgumentList)]
20+
pub fn derive_kernel_argument_list(input: TokenStream) -> TokenStream {
21+
let mut input = parse_macro_input!(input as DeriveInput);
22+
let oneapi = find_oneapi();
23+
24+
let Data::Struct(data) = &input.data else {
25+
return Error::new_spanned(input, "This derive macro only works on structs.")
26+
.into_compile_error()
27+
.into();
28+
};
29+
expand_where_clause(input.generics.make_where_clause(), data, &oneapi);
30+
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
31+
32+
let ident = input.ident;
33+
let argc = data.fields.len();
34+
let members = data.fields.members();
35+
36+
let expanded = quote! {
37+
unsafe impl #impl_generics #oneapi::kernel::KernelArgumentList<#argc>
38+
for #ident #ty_generics #where_clause {
39+
unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] {
40+
[ #(unsafe { self.#members.as_raw_arg() }),* ]
41+
}
42+
}
43+
};
44+
45+
TokenStream::from(expanded)
46+
}
47+
48+
fn expand_where_clause(where_clause: &mut WhereClause, data: &DataStruct, oneapi: &Ident) {
49+
for Field { ty, .. } in &data.fields {
50+
where_clause
51+
.predicates
52+
.push(parse_quote!(#ty: #oneapi::kernel::KernelArgument));
53+
}
54+
}
55+
56+
fn get_single_tuple_impl(argc: usize) -> proc_macro2::TokenStream {
57+
let iter = { 0..argc }.map(syn::Index::from);
58+
let types = { 0..argc }
59+
.map(|i| format_ident!("T{i}"))
60+
.collect::<Vec<_>>();
61+
62+
quote! {
63+
unsafe impl<#(#types),*> crate::kernel::KernelArgumentList<#argc> for (#(#types),*)
64+
where #(#types: crate::kernel::KernelArgument),* {
65+
unsafe fn as_raw_arg_list(&self) -> [&[u8]; #argc] {
66+
[ #(unsafe { self.#iter.as_raw_arg() }),* ]
67+
}
68+
}
69+
}
70+
}
71+
72+
/// A macro that generates tuples from 2..N which implement the `KernelArgumentList` trait.
73+
#[proc_macro]
74+
pub fn impl_arg_list_for_tuples(input: TokenStream) -> TokenStream {
75+
let input = parse_macro_input!(input as LitInt);
76+
let argc = input.base10_parse::<usize>().unwrap();
77+
78+
let impls = { 2..=argc }.map(get_single_tuple_impl);
79+
80+
let expanded = quote! {
81+
#(#impls)*
82+
};
83+
84+
TokenStream::from(expanded)
85+
}

oneapi-rs/Cargo.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ allocator-api2 = "0.4.0"
99
bytemuck = "1.25.1"
1010
cxx = "1.0.194"
1111
oneapi-rs-sys = { path = "../oneapi-rs-sys" }
12+
oneapi-rs-derive = { path = "../oneapi-rs-derive" }
1213
pin-project = "1.1.13"
1314
thiserror = "2.0.18"
1415

oneapi-rs/examples/kernel_launch.rs

Lines changed: 5 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -6,13 +6,7 @@
66
// SPDX-License-Identifier: MIT OR Apache-2.0
77
//
88

9-
use oneapi_rs::{
10-
buffer::Buffer,
11-
kernel::{KernelArgument, KernelArgumentList},
12-
queue::Queue,
13-
range::NdRange,
14-
usm::{SharedAllocator, UsmAllocator},
15-
};
9+
use oneapi_rs::{queue::Queue, range::NdRange};
1610

1711
static IOTA_SRC: &str = r#"
1812
#include <sycl/sycl.hpp>
@@ -21,46 +15,23 @@ namespace syclexp = sycl::ext::oneapi::experimental;
2115
2216
extern "C"
2317
SYCL_EXT_ONEAPI_FUNCTION_PROPERTY((syclexp::nd_range_kernel<1>))
24-
void iota(float start, float *ptr) {
18+
void iota(double start, double *ptr) {
2519
size_t id = syclext::this_work_item::get_nd_item<1>().get_global_linear_id();
26-
ptr[id] = start + static_cast<float>(id);
20+
ptr[id] = start + static_cast<double>(id);
2721
}
2822
"#;
2923

30-
struct IotaArgs<'a> {
31-
start: f32,
32-
buffer: &'a mut Buffer<f32, UsmAllocator<SharedAllocator>>,
33-
}
34-
35-
unsafe impl<'a> KernelArgumentList<2> for IotaArgs<'a> {
36-
unsafe fn as_raw_arg_list(&self) -> [&[u8]; 2] {
37-
return [unsafe { self.start.as_raw_arg() }, unsafe {
38-
self.buffer.as_raw_arg()
39-
}];
40-
}
41-
}
42-
4324
fn main() {
4425
let mut queue = Queue::new();
45-
let mut buffer = queue.alloc_shared::<f32>(1024).wait();
26+
let mut buffer = queue.alloc_shared::<f64>(1024).wait();
4627

4728
let kernel = queue
4829
.get_context()
4930
.create_kernel_bundle_from_source(IOTA_SRC)
5031
.build()
5132
.get_kernel("iota");
5233

53-
unsafe {
54-
queue.launch(
55-
NdRange::new([1024], [16]),
56-
&kernel,
57-
IotaArgs {
58-
start: 3.14,
59-
buffer: &mut buffer,
60-
},
61-
)
62-
}
63-
.wait();
34+
unsafe { queue.launch(NdRange::new([1024], [16]), &kernel, (3.14, &mut buffer)) }.wait();
6435

6536
for e in buffer.iter() {
6637
print!("{e} ");

0 commit comments

Comments
 (0)