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
1 change: 1 addition & 0 deletions oneapi-rs-sys/build.rs
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ fn main() {
];

let cpp_headers = [
"include/utils.hpp",
"include/types.hpp",
"include/platform.hpp",
"include/device.hpp",
Expand Down
40 changes: 40 additions & 0 deletions oneapi-rs-sys/include/utils.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
//
// Copyright (C) 2026 Intel Corporation
//
// Under the MIT License or the Apache License v2.0.
// See LICENSE-MIT and LICENSE-APACHE for license information.
// SPDX-License-Identifier: MIT OR Apache-2.0
//

#pragma once

#include <memory>
#include <type_traits>
#include <utility>
#include <vector>

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

namespace sycl_shims::utils {
template <typename T>
using UnwrappedPtr = std::remove_reference_t<decltype(*std::declval<T>().ptr)>;

template <typename T>
std::vector<UnwrappedPtr<T>> vec_to_vector(rust::Vec<T> &&vec) {
std::vector<UnwrappedPtr<T>> vector;
for (auto &&e : vec)
vector.push_back(std::move(*e.ptr));

return vector;
}

template <typename T>
rust::Vec<T> vector_to_vec(std::vector<UnwrappedPtr<T>> &&vector) {
rust::Vec<T> vec;
for (auto &&e : vector)
vec.push_back(T{std::make_unique<UnwrappedPtr<T>>(e)});

return vec;
}
} // namespace sycl_shims::utils
8 changes: 4 additions & 4 deletions oneapi-rs-sys/src/context.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -7,13 +7,13 @@
//

#include "oneapi-rs-sys/include/context.hpp"
#include "oneapi-rs-sys/include/utils.hpp"
#include "oneapi-rs-sys/src/context-sys.rs.h"

using sycl_shims::utils::vec_to_vector;

namespace sycl_shims::context {
std::unique_ptr<Context> new_context(rust::Vec<DevicePtr> devices) {
std::vector<sycl::device> raw_devices;
for (auto &&d : devices)
raw_devices.push_back(std::move(*d.ptr));
return std::make_unique<Context>(raw_devices);
return std::make_unique<Context>(vec_to_vector(std::move(devices)));
}
} // namespace sycl_shims::context
9 changes: 3 additions & 6 deletions oneapi-rs-sys/src/device.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -7,18 +7,15 @@
//

#include "oneapi-rs-sys/include/device.hpp"
#include "oneapi-rs-sys/include/utils.hpp"
#include "oneapi-rs-sys/src/device-sys.rs.h"

using sycl_shims::utils::vector_to_vec;
using dt = sycl::info::device_type;

namespace sycl_shims::device {
rust::Vec<DevicePtr> get_devices() {
rust::Vec<DevicePtr> devices;

for (auto &&device : sycl::device::get_devices())
devices.push_back(DevicePtr{std::make_unique<Device>(device)});

return devices;
return vector_to_vec<DevicePtr>(sycl::device::get_devices());
}

DeviceType get_device_type(Device const &device) {
Expand Down
17 changes: 5 additions & 12 deletions oneapi-rs-sys/src/platform.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -7,25 +7,18 @@
//

#include "oneapi-rs-sys/include/platform.hpp"
#include "oneapi-rs-sys/include/utils.hpp"
#include "oneapi-rs-sys/src/platform-sys.rs.h"

using sycl_shims::utils::vector_to_vec;

namespace sycl_shims::platform {
rust::Vec<PlatformPtr> get_platforms() {
rust::Vec<PlatformPtr> platforms;

for (auto &&platform : sycl::platform::get_platforms())
platforms.push_back(PlatformPtr{std::make_unique<Platform>(platform)});

return platforms;
return vector_to_vec<PlatformPtr>(sycl::platform::get_platforms());
}

rust::Vec<DevicePtr> get_devices(Platform const &platform) {
rust::Vec<DevicePtr> devices;

for (auto &&device : platform.get_devices())
devices.push_back(DevicePtr{std::make_unique<Device>(device)});

return devices;
return vector_to_vec<DevicePtr>(platform.get_devices());
}

rust::String get_version(Platform const &platform) {
Expand Down
21 changes: 8 additions & 13 deletions oneapi-rs-sys/src/queue.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,12 @@
//

#include "oneapi-rs-sys/include/queue.hpp"
#include "oneapi-rs-sys/include/utils.hpp"
#include "oneapi-rs-sys/src/queue-sys.rs.h"

using sycl::ext::intel::property::queue::immediate_command_list;
using sycl::property::queue::in_order;
using sycl_shims::utils::vec_to_vector;

namespace syclexp = sycl::ext::oneapi::experimental;

Expand Down Expand Up @@ -39,18 +41,14 @@ std::unique_ptr<Queue> clone(Queue const &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::vector<sycl::event> deps;
for (auto &&e : dep_events)
deps.push_back(std::move(*e.ptr));
return std::make_unique<Event>(queue->memset(ptr, value, num_bytes, deps));
return std::make_unique<Event>(queue->memset(
ptr, value, num_bytes, vec_to_vector(std::move(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)
deps.push_back(std::move(*e.ptr));
return std::make_unique<Event>(queue->ext_oneapi_submit_barrier(deps));
return std::make_unique<Event>(
queue->ext_oneapi_submit_barrier(vec_to_vector(std::move(dep_events))));
}

void wait(std::unique_ptr<Queue> &queue) { queue->wait(); }
Expand Down Expand Up @@ -102,10 +100,7 @@ launch_3d(std::unique_ptr<Queue> &queue, Range3 global_size, Range3 local_size,
std::unique_ptr<Event> memcpy(std::unique_ptr<Queue> &queue, std::uint8_t *dest,
std::uint8_t const *src, std::size_t num_bytes,
rust::Vec<EventPtr> dep_events) {
std::vector<sycl::event> deps;
for (auto &&e : dep_events)
deps.push_back(std::move(*e.ptr));

return std::make_unique<Event>(queue->memcpy(dest, src, num_bytes, deps));
return std::make_unique<Event>(queue->memcpy(
dest, src, num_bytes, vec_to_vector(std::move(dep_events))));
}
} // namespace sycl_shims::queue