Skip to content

Commit 2fa3f90

Browse files
cargo fmt
1 parent 1021932 commit 2fa3f90

8 files changed

Lines changed: 97 additions & 46 deletions

File tree

oneapi-rs-sys/src/kernel-bundle-sys.rs

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -23,9 +23,14 @@ pub mod ffi {
2323
#[namespace = "sycl_shims"]
2424
type Kernel = crate::types::ffi::Kernel;
2525

26-
fn create_kernel_bundle_from_source(ctxt: &Context, source: &str)
27-
-> UniquePtr<SourceKernelBundle>;
26+
fn create_kernel_bundle_from_source(
27+
ctxt: &Context,
28+
source: &str,
29+
) -> UniquePtr<SourceKernelBundle>;
2830
fn build(source: &mut UniquePtr<SourceKernelBundle>) -> UniquePtr<ExecutableKernelBundle>;
29-
fn get_kernel(bundle: &mut UniquePtr<ExecutableKernelBundle>, name: &str) -> UniquePtr<Kernel>;
31+
fn get_kernel(
32+
bundle: &mut UniquePtr<ExecutableKernelBundle>,
33+
name: &str,
34+
) -> UniquePtr<Kernel>;
3035
}
3136
}

oneapi-rs-sys/src/lib.rs

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,4 +29,3 @@ pub mod context;
2929

3030
#[path = "kernel-bundle-sys.rs"]
3131
pub mod kernel_bundle;
32-

oneapi-rs-sys/src/types-sys.rs

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -80,15 +80,15 @@ pub mod ffi {
8080

8181
// cxx doesn't support const generic parameters
8282
struct Range1 {
83-
data: [u64; 1]
83+
data: [u64; 1],
8484
}
8585

8686
struct Range2 {
87-
data: [u64; 2]
87+
data: [u64; 2],
8888
}
8989

9090
struct Range3 {
91-
data: [u64; 3]
91+
data: [u64; 3],
9292
}
9393

9494
impl UniquePtr<Device> {}

oneapi-rs/examples/kernel_launch.rs

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

9-
use oneapi_rs::{buffer::Buffer, kernel_bundle::{KernelArgument, KernelArgumentList, NdRange, Range}, queue::Queue, usm::{SharedAllocator, UsmAllocator}};
9+
use oneapi_rs::{
10+
buffer::Buffer,
11+
kernel_bundle::{KernelArgument, KernelArgumentList, NdRange},
12+
queue::Queue,
13+
usm::{SharedAllocator, UsmAllocator},
14+
};
1015

1116
static IOTA_SRC: &str = r#"
1217
#include <sycl/sycl.hpp>
@@ -23,28 +28,38 @@ void iota(float start, float *ptr) {
2328

2429
struct IotaArgs<'a> {
2530
start: f32,
26-
buffer: &'a mut Buffer<f32, UsmAllocator<SharedAllocator>>
31+
buffer: &'a mut Buffer<f32, UsmAllocator<SharedAllocator>>,
2732
}
2833

2934
unsafe impl<'a> KernelArgumentList<2> for IotaArgs<'a> {
3035
unsafe fn as_raw_arg_list(&self) -> [&[u8]; 2] {
31-
return [
32-
unsafe { self.start.as_raw_arg() },
33-
unsafe { self.buffer.as_raw_arg() }
34-
]
36+
return [unsafe { self.start.as_raw_arg() }, unsafe {
37+
self.buffer.as_raw_arg()
38+
}];
3539
}
3640
}
3741

3842
fn main() {
3943
let mut queue = Queue::new();
4044
let mut buffer = queue.alloc_shared::<f32>(1024).wait();
4145

42-
let kernel = queue.get_context()
46+
let kernel = queue
47+
.get_context()
4348
.create_kernel_bundle_from_source(IOTA_SRC)
4449
.build()
4550
.get_kernel("iota");
4651

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

4964
for e in buffer.iter() {
5065
print!("{e} ");

oneapi-rs/src/buffer.rs

Lines changed: 4 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,9 @@ use bytemuck::Pod;
1919
use pin_project::pin_project;
2020

2121
use crate::{
22-
event::{Event, EventFuture}, kernel_bundle::KernelArgument, usm::UsmAlloc,
22+
event::{Event, EventFuture},
23+
kernel_bundle::KernelArgument,
24+
usm::UsmAlloc,
2325
};
2426

2527
/// The Buffer struct defines a shared array of one, two or three dimensions that can be used
@@ -143,11 +145,6 @@ unsafe impl<T: Pod, A: UsmAlloc> KernelArgument for Buffer<T, A> {
143145
unsafe fn as_raw_arg(&self) -> &[u8] {
144146
let data_ptr: *const NonNull<_> = &self.data;
145147
let cast_ptr = data_ptr as *const u8;
146-
unsafe {
147-
slice::from_raw_parts(
148-
cast_ptr,
149-
std::mem::size_of::<*mut u8>()
150-
)
151-
}
148+
unsafe { slice::from_raw_parts(cast_ptr, std::mem::size_of::<*mut u8>()) }
152149
}
153150
}

oneapi-rs/src/context.rs

Lines changed: 4 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_sys::{kernel_bundle, context::ffi, types::ffi::DevicePtr};
9+
use oneapi_rs_sys::{context::ffi, kernel_bundle, types::ffi::DevicePtr};
1010

1111
use crate::{device::Device, kernel_bundle::SourceKernelBundle};
1212

@@ -24,7 +24,9 @@ impl Context {
2424
pub fn new(devices: &[&Device]) -> Self {
2525
let devices = devices
2626
.iter()
27-
.map(|d| DevicePtr { ptr: (*d).clone().0 })
27+
.map(|d| DevicePtr {
28+
ptr: (*d).clone().0,
29+
})
2830
.collect::<Vec<_>>();
2931

3032
ffi::new_context(devices).into()

oneapi-rs/src/kernel_bundle.rs

Lines changed: 53 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -61,65 +61,99 @@ pub unsafe trait KernelArgumentList<const ARGC: usize> {
6161
unsafe fn as_raw_arg_list(&self) -> [&[u8]; ARGC];
6262
}
6363

64-
6564
pub type Range<const DIMENSIONS: usize = 1> = [u64; DIMENSIONS];
6665

6766
pub struct NdRange<const DIMENSIONS: usize = 1> {
6867
pub group_size: Range<DIMENSIONS>,
69-
pub local_size: Range<DIMENSIONS>
68+
pub local_size: Range<DIMENSIONS>,
7069
}
7170

7271
impl<const DIMENSIONS: usize> NdRange<DIMENSIONS> {
7372
pub fn new(group_size: Range<DIMENSIONS>, local_size: Range<DIMENSIONS>) -> Self {
7473
Self {
7574
group_size,
76-
local_size
75+
local_size,
7776
}
7877
}
7978
}
8079

8180
pub trait ValidDimension {
82-
unsafe fn launch<const ARGC: usize>(&self, queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList<ARGC>) -> Event;
81+
unsafe fn launch<const ARGC: usize>(
82+
&self,
83+
queue: &mut Queue,
84+
kernel: &Kernel,
85+
args: impl KernelArgumentList<ARGC>,
86+
) -> Event;
8387
}
8488

8589
impl ValidDimension for NdRange<1> {
86-
unsafe fn launch<const ARGC: usize>(&self, queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList<ARGC>) -> Event {
90+
unsafe fn launch<const ARGC: usize>(
91+
&self,
92+
queue: &mut Queue,
93+
kernel: &Kernel,
94+
args: impl KernelArgumentList<ARGC>,
95+
) -> Event {
8796
unsafe {
8897
oneapi_rs_sys::queue::ffi::launch_1d(
8998
&mut queue.0,
90-
types::ffi::Range1 { data: self.group_size },
91-
types::ffi::Range1 { data: self.local_size },
99+
types::ffi::Range1 {
100+
data: self.group_size,
101+
},
102+
types::ffi::Range1 {
103+
data: self.local_size,
104+
},
92105
&kernel.0,
93-
&args.as_raw_arg_list()
106+
&args.as_raw_arg_list(),
94107
)
95-
}.into()
108+
}
109+
.into()
96110
}
97111
}
98112

99113
impl ValidDimension for NdRange<2> {
100-
unsafe fn launch<const ARGC: usize>(&self, queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList<ARGC>) -> Event {
114+
unsafe fn launch<const ARGC: usize>(
115+
&self,
116+
queue: &mut Queue,
117+
kernel: &Kernel,
118+
args: impl KernelArgumentList<ARGC>,
119+
) -> Event {
101120
unsafe {
102121
oneapi_rs_sys::queue::ffi::launch_2d(
103122
&mut queue.0,
104-
types::ffi::Range2 { data: self.group_size },
105-
types::ffi::Range2 { data: self.local_size },
123+
types::ffi::Range2 {
124+
data: self.group_size,
125+
},
126+
types::ffi::Range2 {
127+
data: self.local_size,
128+
},
106129
&kernel.0,
107-
&args.as_raw_arg_list()
130+
&args.as_raw_arg_list(),
108131
)
109-
}.into()
132+
}
133+
.into()
110134
}
111135
}
112136

113137
impl ValidDimension for NdRange<3> {
114-
unsafe fn launch<const ARGC: usize>(&self, queue: &mut Queue, kernel: &Kernel, args: impl KernelArgumentList<ARGC>) -> Event {
138+
unsafe fn launch<const ARGC: usize>(
139+
&self,
140+
queue: &mut Queue,
141+
kernel: &Kernel,
142+
args: impl KernelArgumentList<ARGC>,
143+
) -> Event {
115144
unsafe {
116145
oneapi_rs_sys::queue::ffi::launch_3d(
117146
&mut queue.0,
118-
types::ffi::Range3 { data: self.group_size },
119-
types::ffi::Range3 { data: self.local_size },
147+
types::ffi::Range3 {
148+
data: self.group_size,
149+
},
150+
types::ffi::Range3 {
151+
data: self.local_size,
152+
},
120153
&kernel.0,
121-
&args.as_raw_arg_list()
154+
&args.as_raw_arg_list(),
122155
)
123-
}.into()
156+
}
157+
.into()
124158
}
125159
}

oneapi-rs/src/lib.rs

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -7,12 +7,11 @@
77
//
88

99
pub mod buffer;
10+
pub mod context;
1011
pub mod device;
1112
pub mod event;
1213
pub mod info;
14+
pub mod kernel_bundle;
1315
pub mod platform;
1416
pub mod queue;
1517
pub mod usm;
16-
pub mod context;
17-
pub mod kernel_bundle;
18-

0 commit comments

Comments
 (0)