Skip to content

Commit d624fe1

Browse files
Add dependent events support for memset
1 parent b6a48cf commit d624fe1

4 files changed

Lines changed: 50 additions & 6 deletions

File tree

oneapi-rs-sys/include/queue.hpp

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,12 +13,22 @@
1313

1414
#include <memory>
1515

16+
namespace sycl_shims {
17+
struct EventPtr;
18+
} // namespace sycl_shims
19+
1620
namespace sycl_shims::queue {
1721
std::unique_ptr<Queue> new_queue();
1822
std::unique_ptr<Queue> new_queue_immediate();
1923
std::unique_ptr<Queue> new_queue_from_device(Device const &);
2024
std::unique_ptr<Queue> clone(Queue const &);
21-
std::unique_ptr<Event> memset(std::unique_ptr<Queue> &, std::uint8_t * ptr, int value, std::size_t num_bytes);
25+
std::unique_ptr<Event> memset(
26+
std::unique_ptr<Queue> &,
27+
std::uint8_t * ptr,
28+
int value,
29+
std::size_t num_bytes,
30+
rust::Vec<EventPtr> dep_events
31+
);
2232
std::unique_ptr<Event> barrier(std::unique_ptr<Queue> &);
2333
void wait(std::unique_ptr<Queue> &);
2434
} // namespace sycl_shims::queue

oneapi-rs-sys/src/queue-sys.rs

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,12 @@
88

99
#[cxx::bridge(namespace = "sycl_shims::queue")]
1010
pub mod ffi {
11+
#[namespace = "sycl_shims"]
12+
extern "C++" {
13+
include!("oneapi-rs-sys/src/types-sys.rs.h");
14+
type EventPtr = crate::types::ffi::EventPtr;
15+
}
16+
1117
unsafe extern "C++" {
1218
include!("oneapi-rs-sys/include/queue.hpp");
1319

@@ -22,7 +28,13 @@ pub mod ffi {
2228
fn new_queue_immediate() -> UniquePtr<Queue>;
2329
fn new_queue_from_device(device: &Device) -> UniquePtr<Queue>;
2430
fn clone(queue: &Queue) -> UniquePtr<Queue>;
25-
unsafe fn memset(queue: &mut UniquePtr<Queue>, ptr: *mut u8, value: i32, num_bytes: usize) -> UniquePtr<Event>;
31+
unsafe fn memset(
32+
queue: &mut UniquePtr<Queue>,
33+
ptr: *mut u8,
34+
value: i32,
35+
num_bytes: usize,
36+
dep_events: Vec<EventPtr>
37+
) -> UniquePtr<Event>;
2638
fn barrier(queue: &mut UniquePtr<Queue>) -> UniquePtr<Event>;
2739
fn wait(queue: &mut UniquePtr<Queue>);
2840
}

oneapi-rs-sys/src/queue.cpp

Lines changed: 11 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -30,8 +30,17 @@ std::unique_ptr<Queue> new_queue_from_device(Device const & device) {
3030
std::unique_ptr<Queue> clone(Queue const & queue) {
3131
return std::make_unique<Queue>(sycl::queue(queue));
3232
}
33-
std::unique_ptr<Event> memset(std::unique_ptr<Queue> & queue, std::uint8_t * ptr, int value, std::size_t num_bytes) {
34-
return std::make_unique<Event>(queue->memset(ptr, value, num_bytes));
33+
std::unique_ptr<Event> memset(
34+
std::unique_ptr<Queue> & queue,
35+
std::uint8_t * ptr,
36+
int value,
37+
std::size_t num_bytes,
38+
rust::Vec<EventPtr> dep_events
39+
) {
40+
std::vector<sycl::event> deps;
41+
for (auto&& e: dep_events)
42+
deps.push_back(std::move(*e.ptr.release()));
43+
return std::make_unique<Event>(queue->memset(ptr, value, num_bytes, deps));
3544
}
3645
std::unique_ptr<Event> barrier(std::unique_ptr<Queue> & queue) {
3746
return std::make_unique<Event>(queue->ext_oneapi_submit_barrier());

oneapi-rs/src/queue.rs

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

99
use bytemuck::Pod;
10-
use oneapi_rs_sys::queue::ffi;
10+
use oneapi_rs_sys::{queue::ffi, types::ffi::EventPtr};
1111

1212
use crate::{buffer::{Buffer, BufferInProgress}, device::Device, event::Event, usm::{HostAllocator, SharedAllocator, UsmAlloc, UsmAllocator}};
1313

@@ -60,9 +60,22 @@ impl Queue {
6060
}
6161

6262
pub unsafe fn memset<T, A: UsmAlloc>(&mut self, buffer: &mut Buffer<T, A>, value: i32) -> Event {
63+
unsafe { self.memset_with_deps(buffer, value, &[]) }
64+
}
65+
66+
pub unsafe fn memset_with_deps<T, A: UsmAlloc>(
67+
&mut self,
68+
buffer: &mut Buffer<T, A>,
69+
value: i32,
70+
dep_events: &[&Event]
71+
) -> Event {
6372
let ptr = buffer.get_byte_ptr();
6473
let num_bytes = buffer.get_byte_size();
65-
unsafe { ffi::memset(&mut self.0, ptr, value, num_bytes) }.into()
74+
let dep_events = dep_events
75+
.iter()
76+
.map(|e| EventPtr { ptr: (*e).clone().0 })
77+
.collect::<Vec<_>>();
78+
unsafe { ffi::memset(&mut self.0, ptr, value, num_bytes, dep_events) }.into()
6679
}
6780

6881
pub fn barrier(&mut self) -> Event {

0 commit comments

Comments
 (0)