Skip to content

Commit 9e3b755

Browse files
Add Rust NdRange support
1 parent dc08aab commit 9e3b755

3 files changed

Lines changed: 99 additions & 5 deletions

File tree

oneapi-rs/examples/kernel_launch.rs

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

9-
use oneapi_rs::{buffer::Buffer, kernel_bundle::{KernelArgument, KernelArgumentList}, queue::Queue, usm::{SharedAllocator, UsmAllocator}};
9+
use oneapi_rs::{buffer::Buffer, kernel_bundle::{KernelArgument, KernelArgumentList, NdRange, Range}, queue::Queue, usm::{SharedAllocator, UsmAllocator}};
1010

1111
static IOTA_SRC: &str = r#"
1212
#include <sycl/sycl.hpp>
@@ -44,7 +44,7 @@ fn main() {
4444
.build()
4545
.get_kernel("iota");
4646

47-
unsafe { queue.launch(&kernel, IotaArgs { start: 3.14, buffer: &mut buffer }) }.wait();
47+
unsafe { queue.launch(NdRange::new([1024], [16]), &kernel, IotaArgs { start: 3.14, buffer: &mut buffer }) }.wait();
4848

4949
for e in buffer.iter() {
5050
print!("{e} ");

oneapi-rs/src/kernel_bundle.rs

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

9+
use std::marker::PhantomData;
10+
911
use bytemuck::Pod;
1012
use oneapi_rs_sys::{kernel_bundle::ffi, types};
1113

14+
use crate::{event::Event, queue::Queue};
15+
1216
pub struct SourceKernelBundle(pub(crate) cxx::UniquePtr<types::ffi::SourceKernelBundle>);
1317

1418
impl From<cxx::UniquePtr<types::ffi::SourceKernelBundle>> for SourceKernelBundle {
@@ -58,3 +62,80 @@ unsafe impl<T: Pod> KernelArgument for T {
5862
pub unsafe trait KernelArgumentList<const ARGC: usize> {
5963
unsafe fn as_raw_arg_list(&self) -> [&[u8]; ARGC];
6064
}
65+
66+
67+
pub type Range<const DIMENSIONS: usize = 1> = [u64; DIMENSIONS];
68+
69+
pub struct NdRange<const DIMENSIONS: usize = 1> {
70+
pub group_size: Range<DIMENSIONS>,
71+
pub local_size: Range<DIMENSIONS>
72+
}
73+
74+
impl<const DIMENSIONS: usize> NdRange<DIMENSIONS> {
75+
pub fn new(group_size: Range<DIMENSIONS>, local_size: Range<DIMENSIONS>) -> Self {
76+
Self {
77+
group_size,
78+
local_size
79+
}
80+
}
81+
}
82+
83+
pub trait ValidDimension {
84+
unsafe fn launch<const ARGC: usize>(&self, queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList<ARGC>) -> Event;
85+
}
86+
87+
impl ValidDimension for NdRange<1> {
88+
unsafe fn launch<const ARGC: usize>(&self, queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList<ARGC>) -> Event {
89+
unsafe {
90+
oneapi_rs_sys::queue::ffi::launch_1d(
91+
&mut queue.0,
92+
self.group_size[0],
93+
self.local_size[0],
94+
&kernel.0,
95+
&args.as_raw_arg_list()
96+
)
97+
}.into()
98+
}
99+
}
100+
101+
impl ValidDimension for NdRange<2> {
102+
unsafe fn launch<const ARGC: usize>(&self, queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList<ARGC>) -> Event {
103+
unsafe {
104+
oneapi_rs_sys::queue::ffi::launch_2d(
105+
&mut queue.0,
106+
types::ffi::Range2 {
107+
x: self.group_size[0],
108+
y: self.group_size[1],
109+
},
110+
types::ffi::Range2 {
111+
x: self.local_size[0],
112+
y: self.local_size[1],
113+
},
114+
&kernel.0,
115+
&args.as_raw_arg_list()
116+
)
117+
}.into()
118+
}
119+
}
120+
121+
impl ValidDimension for NdRange<3> {
122+
unsafe fn launch<const ARGC: usize>(&self, queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList<ARGC>) -> Event {
123+
unsafe {
124+
oneapi_rs_sys::queue::ffi::launch_3d(
125+
&mut queue.0,
126+
types::ffi::Range3 {
127+
x: self.group_size[0],
128+
y: self.group_size[1],
129+
z: self.group_size[2],
130+
},
131+
types::ffi::Range3 {
132+
x: self.local_size[0],
133+
y: self.local_size[1],
134+
z: self.local_size[2],
135+
},
136+
&kernel.0,
137+
&args.as_raw_arg_list()
138+
)
139+
}.into()
140+
}
141+
}

oneapi-rs/src/queue.rs

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,12 @@ use bytemuck::Pod;
1010
use oneapi_rs_sys::{queue::ffi, types::ffi::EventPtr};
1111

1212
use crate::{
13-
buffer::{Buffer, EnqueuedBuffer}, context::Context, device::Device, event::Event, kernel_bundle::{Kernel, KernelArgumentList}, usm::{HostAllocator, SharedAllocator, UsmAlloc, UsmAllocator},
13+
buffer::{Buffer, EnqueuedBuffer},
14+
context::Context,
15+
device::Device,
16+
event::Event,
17+
kernel_bundle::{Kernel, KernelArgumentList, NdRange, ValidDimension},
18+
usm::{HostAllocator, SharedAllocator, UsmAlloc, UsmAllocator},
1419
};
1520

1621
/// The `Queue` connects a host program to a single device. Programs submit tasks to a device via the
@@ -128,8 +133,16 @@ impl Queue {
128133
ffi::wait(&mut self.0);
129134
}
130135

131-
pub unsafe fn launch<const ARGC: usize>(&mut self, kernel: &Kernel, args: impl KernelArgumentList<ARGC>) -> Event {
132-
unsafe { ffi::launch(&mut self.0, &kernel.0, &args.as_raw_arg_list()) }.into()
136+
pub unsafe fn launch<const ARGC: usize, const DIMENSIONS: usize>(
137+
&mut self,
138+
nd_range: NdRange<DIMENSIONS>,
139+
kernel: &Kernel,
140+
args: impl KernelArgumentList<ARGC>,
141+
) -> Event
142+
where
143+
NdRange<DIMENSIONS>: ValidDimension,
144+
{
145+
unsafe { nd_range.launch(self, kernel, args) }
133146
}
134147
}
135148

0 commit comments

Comments
 (0)