Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions oneapi-rs-sys/build.rs
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,7 @@ fn get_compiler_path() -> Result<PathBuf, Error> {
if let Ok(path) = std::env::var("CMPLR_ROOT") {
let path = PathBuf::from(path).join("bin/icpx");
if path.exists() {
return Ok(path)
return Ok(path);
}
}
if let Ok(path) = which("icpx") {
Expand All @@ -81,5 +81,8 @@ fn get_compiler_path() -> Result<PathBuf, Error> {
return Ok(path);
}

Err(Error::new(std::io::ErrorKind::NotFound, "No DPC++ compiler found"))
Err(Error::new(
std::io::ErrorKind::NotFound,
"No DPC++ compiler found",
))
}
5 changes: 2 additions & 3 deletions oneapi-rs-sys/include/device.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,12 +8,11 @@

#pragma once

#include <memory>

#include "oneapi-rs-sys/include/types.hpp"
#include "rust/cxx.h"

#include <memory>
#include <vector>

namespace sycl_shims {
struct DevicePtr;
struct PlatformPtr;
Expand Down
7 changes: 4 additions & 3 deletions oneapi-rs-sys/include/event.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,19 +8,20 @@

#pragma once

#include <memory>

#include "oneapi-rs-sys/include/types.hpp"
#include "oneapi-rs-sys/src/event-sys.rs.h"
#include "rust/cxx.h"

#include <memory>

namespace sycl_shims {
enum class EventCommandStatus : std::uint8_t;
} // namespace sycl_shims

namespace sycl_shims::event {
void wait(std::unique_ptr<Event> &);
void register_callback(std::unique_ptr<Queue> &, Event const &, SharedWaker const *);
void register_callback(std::unique_ptr<Queue> &, Event const &,
SharedWaker const *);
EventCommandStatus get_command_execution_status(Event const &);
std::unique_ptr<Event> clone(Event const &);
} // namespace sycl_shims::event
15 changes: 7 additions & 8 deletions oneapi-rs-sys/include/platform.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,13 +8,12 @@

#pragma once

#include "rust/cxx.h"
#include "oneapi-rs-sys/include/types.hpp"
#include <memory>

#include <sycl/sycl.hpp>

#include <memory>
#include <vector>
#include "oneapi-rs-sys/include/types.hpp"
#include "rust/cxx.h"

namespace sycl_shims {
struct DevicePtr;
Expand All @@ -23,8 +22,8 @@ struct PlatformPtr;

namespace sycl_shims::platform {
rust::Vec<PlatformPtr> get_platforms();
rust::Vec<DevicePtr> get_devices(Platform const&);
rust::String get_version(Platform const&);
rust::String get_name(Platform const&);
rust::String get_vendor(Platform const&);
rust::Vec<DevicePtr> get_devices(Platform const &);
rust::String get_version(Platform const &);
rust::String get_name(Platform const &);
rust::String get_vendor(Platform const &);
} // namespace sycl_shims::platform
14 changes: 5 additions & 9 deletions oneapi-rs-sys/include/queue.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,11 +8,11 @@

#pragma once

#include <memory>

#include "oneapi-rs-sys/include/types.hpp"
#include "rust/cxx.h"

#include <memory>

namespace sycl_shims {
struct EventPtr;
} // namespace sycl_shims
Expand All @@ -22,13 +22,9 @@ std::unique_ptr<Queue> new_queue();
std::unique_ptr<Queue> new_queue_immediate();
std::unique_ptr<Queue> new_queue_from_device(Device const &);
std::unique_ptr<Queue> clone(Queue const &);
std::unique_ptr<Event> memset(
std::unique_ptr<Queue> &,
std::uint8_t * ptr,
int value,
std::size_t num_bytes,
rust::Vec<EventPtr>
);
std::unique_ptr<Event> memset(std::unique_ptr<Queue> &, std::uint8_t *ptr,
int value, std::size_t num_bytes,
rust::Vec<EventPtr>);
std::unique_ptr<Event> barrier(std::unique_ptr<Queue> &, rust::Vec<EventPtr>);
void wait(std::unique_ptr<Queue> &);
} // namespace sycl_shims::queue
2 changes: 1 addition & 1 deletion oneapi-rs-sys/include/types.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -15,4 +15,4 @@ using Device = sycl::device;
using Platform = sycl::platform;
using Queue = sycl::queue;
using Event = sycl::event;
}
} // namespace sycl_shims
15 changes: 9 additions & 6 deletions oneapi-rs-sys/include/usm.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,14 +8,17 @@

#pragma once

#include <memory>

#include "oneapi-rs-sys/include/types.hpp"
#include "rust/cxx.h"

#include <memory>

namespace sycl_shims::usm {
std::uint8_t* aligned_alloc_device(std::size_t alignment, std::size_t num_bytes, Queue const &);
std::uint8_t* aligned_alloc_host(std::size_t alignment, std::size_t num_bytes, Queue const &);
std::uint8_t* aligned_alloc_shared(std::size_t alignment, std::size_t num_bytes, Queue const &);
void free(std::uint8_t*, Queue const &);
std::uint8_t *aligned_alloc_device(std::size_t alignment, std::size_t num_bytes,
Queue const &);
std::uint8_t *aligned_alloc_host(std::size_t alignment, std::size_t num_bytes,
Queue const &);
std::uint8_t *aligned_alloc_shared(std::size_t alignment, std::size_t num_bytes,
Queue const &);
void free(std::uint8_t *, Queue const &);
} // namespace sycl_shims::usm
6 changes: 5 additions & 1 deletion oneapi-rs-sys/src/event-sys.rs
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,11 @@ pub mod ffi {
type Queue = crate::types::ffi::Queue;

fn wait(event: &mut UniquePtr<Event>);
unsafe fn register_callback(queue: &mut UniquePtr<Queue>, event: &Event, waker: *const SharedWaker);
unsafe fn register_callback(
queue: &mut UniquePtr<Queue>,
event: &Event,
waker: *const SharedWaker,
);
fn get_command_execution_status(event: &Event) -> EventCommandStatus;
fn clone(event: &Event) -> UniquePtr<Event>;
}
Expand Down
31 changes: 16 additions & 15 deletions oneapi-rs-sys/src/event.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -13,27 +13,28 @@ using sycl::info::event_command_status;
namespace syclintel = sycl::ext::intel;

namespace sycl_shims::event {
void wait(std::unique_ptr<Event> & event) {
event->wait();
}
std::unique_ptr<Event> clone(Event const & event) {
void wait(std::unique_ptr<Event> &event) { event->wait(); }
std::unique_ptr<Event> clone(Event const &event) {
return std::make_unique<Event>(sycl::event(event));
}
EventCommandStatus get_command_execution_status(Event const & event) {

EventCommandStatus get_command_execution_status(Event const &event) {
auto status = event.get_info<sycl::info::event::command_execution_status>();
switch (status) {
case event_command_status::submitted:
return EventCommandStatus::Submitted;
case event_command_status::running:
return EventCommandStatus::Running;
case event_command_status::complete:
return EventCommandStatus::Complete;
default:
return EventCommandStatus::Unknown;
case event_command_status::submitted:
return EventCommandStatus::Submitted;
case event_command_status::running:
return EventCommandStatus::Running;
case event_command_status::complete:
return EventCommandStatus::Complete;
default:
return EventCommandStatus::Unknown;
}
}
void register_callback(std::unique_ptr<Queue> & queue, Event const & event, SharedWaker const * waker) {
queue->submit([=](sycl::handler& cgh) {

void register_callback(std::unique_ptr<Queue> &queue, Event const &event,
SharedWaker const *waker) {
queue->submit([=](sycl::handler &cgh) {
cgh.depends_on(event);
cgh.host_task([=]() { waker->wake(); });
});
Expand Down
2 changes: 1 addition & 1 deletion oneapi-rs-sys/src/queue-sys.rs
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ pub mod ffi {
ptr: *mut u8,
value: i32,
num_bytes: usize,
dep_events: Vec<EventPtr>
dep_events: Vec<EventPtr>,
) -> UniquePtr<Event>;
fn barrier(queue: &mut UniquePtr<Queue>, dep_events: Vec<EventPtr>) -> UniquePtr<Event>;
fn wait(queue: &mut UniquePtr<Queue>);
Expand Down
43 changes: 20 additions & 23 deletions oneapi-rs-sys/src/queue.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -9,46 +9,43 @@
#include "oneapi-rs-sys/include/queue.hpp"
#include "oneapi-rs-sys/src/queue-sys.rs.h"

using sycl::property::queue::in_order;
using sycl::ext::intel::property::queue::immediate_command_list;
using sycl::property::queue::in_order;

namespace sycl_shims::queue {
std::unique_ptr<Queue> new_queue() {
return std::make_unique<Queue>(sycl::queue({
in_order()
}));
return std::make_unique<Queue>(sycl::queue({in_order()}));
}

std::unique_ptr<Queue> new_queue_immediate() {
return std::make_unique<Queue>(sycl::queue({
in_order(),
immediate_command_list()
}));
return std::make_unique<Queue>(
sycl::queue({in_order(), immediate_command_list()}));
}
std::unique_ptr<Queue> new_queue_from_device(Device const & device) {

std::unique_ptr<Queue> new_queue_from_device(Device const &device) {
return std::make_unique<Queue>(sycl::queue(device, {in_order()}));
}
std::unique_ptr<Queue> clone(Queue const & queue) {

std::unique_ptr<Queue> clone(Queue const &queue) {
return std::make_unique<Queue>(sycl::queue(queue));
}
std::unique_ptr<Event> memset(
std::unique_ptr<Queue> & queue,
std::uint8_t * ptr,
int value,
std::size_t num_bytes,
rust::Vec<EventPtr> dep_events
) {

std::unique_ptr<Event> memset(std::unique_ptr<Queue> &queue, std::uint8_t *ptr,
int value, std::size_t num_bytes,
rust::Vec<EventPtr> dep_events) {
std::vector<sycl::event> deps;
for (auto&& e: dep_events)
for (auto &&e : dep_events)
deps.push_back(std::move(*e.ptr));
return std::make_unique<Event>(queue->memset(ptr, value, num_bytes, deps));
}
std::unique_ptr<Event> barrier(std::unique_ptr<Queue> & queue, rust::Vec<EventPtr> dep_events) {

std::unique_ptr<Event> barrier(std::unique_ptr<Queue> &queue,
rust::Vec<EventPtr> dep_events) {
std::vector<sycl::event> deps;
for (auto&& e: dep_events)
for (auto &&e : dep_events)
deps.push_back(std::move(*e.ptr));
return std::make_unique<Event>(queue->ext_oneapi_submit_barrier(deps));
}
void wait(std::unique_ptr<Queue> & queue) {
queue->wait();
}

void wait(std::unique_ptr<Queue> &queue) { queue->wait(); }
} // namespace sycl_shims::queue
18 changes: 9 additions & 9 deletions oneapi-rs-sys/src/types-sys.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,20 +6,20 @@
// SPDX-License-Identifier: MIT OR Apache-2.0
//

use std::sync::atomic::{Ordering::Relaxed, AtomicBool};
use std::sync::atomic::{AtomicBool, Ordering::Relaxed};

use futures::task::AtomicWaker;

pub struct SharedWaker {
pub waker: AtomicWaker,
pub done: AtomicBool
pub done: AtomicBool,
}

impl SharedWaker {
pub fn new() -> Self {
Self {
waker: AtomicWaker::new(),
done: AtomicBool::new(false)
done: AtomicBool::new(false),
}
}

Expand All @@ -44,15 +44,15 @@ pub mod ffi {
// https://github.com/dtolnay/cxx/issues/774#issuecomment-808674945
// We must use pointer wrapper structs instead.
struct DevicePtr {
ptr: UniquePtr<Device>
ptr: UniquePtr<Device>,
}

struct PlatformPtr {
ptr: UniquePtr<Platform>
ptr: UniquePtr<Platform>,
}

struct EventPtr {
ptr: UniquePtr<Event>
ptr: UniquePtr<Event>,
}

#[derive(Debug)]
Expand All @@ -63,15 +63,15 @@ pub mod ffi {
Custom,
Automatic,
All,
Unimplemented
Unimplemented,
}

#[derive(Debug)]
enum EventCommandStatus {
Submitted,
Running,
Complete,
Unknown
Unknown,
}

impl UniquePtr<Device> {}
Expand Down
18 changes: 15 additions & 3 deletions oneapi-rs-sys/src/usm-sys.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,21 @@ pub mod ffi {

extern "C++" {
include!("oneapi-rs-sys/include/usm.hpp");
unsafe fn aligned_alloc_device(alignment: usize, num_bytes: usize, queue: &Queue) -> Result<*mut u8>;
unsafe fn aligned_alloc_host(alignment: usize, num_bytes: usize, queue: &Queue) -> Result<*mut u8>;
unsafe fn aligned_alloc_shared(alignment: usize, num_bytes: usize, queue: &Queue) -> Result<*mut u8>;
unsafe fn aligned_alloc_device(
alignment: usize,
num_bytes: usize,
queue: &Queue,
) -> Result<*mut u8>;
unsafe fn aligned_alloc_host(
alignment: usize,
num_bytes: usize,
queue: &Queue,
) -> Result<*mut u8>;
unsafe fn aligned_alloc_shared(
alignment: usize,
num_bytes: usize,
queue: &Queue,
) -> Result<*mut u8>;
unsafe fn free(ptr: *mut u8, queue: &Queue);
}
}
Loading