[llvm-branch-commits] [llvm] [OFFLOAD] Add olIterateCompatibleDevices API (PR #226441)
Alex Duran via llvm-branch-commits
llvm-branch-commits at lists.llvm.org
Tue Sep 29 08:03:06 PDT 2026
https://github.com/adurang updated https://github.com/llvm/llvm-project/pull/226441
>From c3eb05b51860f7ba30e6a6550a100811c1b21f34 Mon Sep 17 00:00:00 2001
From: "Duran, Alex" <alejandro.duran at intel.com>
Date: Fri, 25 Sep 2026 07:24:14 -0700
Subject: [PATCH 1/4] [offload] Initilize Platforms lazily
---
offload/liboffload/API/Common.td | 13 ++-
offload/liboffload/src/OffloadImpl.cpp | 109 +++++++++++++-----
.../Conformance/lib/DeviceContext.cpp | 6 +-
.../OffloadAPI/common/Environment.cpp | 4 +-
offload/unittests/OffloadAPI/init/olInit.cpp | 13 +++
5 files changed, 111 insertions(+), 34 deletions(-)
diff --git a/offload/liboffload/API/Common.td b/offload/liboffload/API/Common.td
index f09ee57961d67..5ed15cd4aca63 100644
--- a/offload/liboffload/API/Common.td
+++ b/offload/liboffload/API/Common.td
@@ -149,12 +149,17 @@ def ol_dimensions_t : Struct {
];
}
+def OL_ALL_PLATFORMS : Macro {
+ let desc = "Value for ol_init_args_t::NumPlatforms that requests all available platforms to be initialized eagerly.";
+ let value = "UINT32_MAX";
+}
+
def ol_init_args_t : Struct {
let desc = "Configuration arguments for olInit.";
let members = [
StructMember<"size_t", "Size", "Size of this struct, used for ABI compatibility. Must be set to sizeof(ol_init_args_t) by the caller.">,
- StructMember<"uint32_t", "NumPlatforms", "Number of entries in the Platforms array.">,
- StructMember<"const ol_platform_backend_t*", "Platforms", "Pointer to an array of platform backends to initialize.">
+ StructMember<"uint32_t", "NumPlatforms", "Number of entries in the Platforms array, or OL_ALL_PLATFORMS to eagerly initialize all available platforms.">,
+ StructMember<"const ol_platform_backend_t*", "Platforms", "Pointer to an array of platform backends to make available and initialize eagerly.">
];
}
@@ -168,7 +173,9 @@ def olInit : Function {
let details = [
"This must be the first API call made by a user of the Offload library",
"Each call will increment an internal reference count that is decremented by `olShutDown`",
- "If InitArgs is NULL, default configuration is used which initializes all available platforms"
+ "If InitArgs is NULL or NumPlatforms is 0, all available platforms are made available and are lazily initialized on their first use",
+ "If NumPlatforms is OL_ALL_PLATFORMS, all available platforms are made available and are initialized eagerly by this call",
+ "Otherwise only the platforms listed in InitArgs (plus the host platform) are made available and are initialized eagerly by this call"
];
let params = [
Param<"const ol_init_args_t*", "InitArgs", "Optional pointer to initialization configuration. NULL uses defaults.", PARAM_IN_OPTIONAL>
diff --git a/offload/liboffload/src/OffloadImpl.cpp b/offload/liboffload/src/OffloadImpl.cpp
index 66093d43b56e8..445e486b6b750 100644
--- a/offload/liboffload/src/OffloadImpl.cpp
+++ b/offload/liboffload/src/OffloadImpl.cpp
@@ -31,6 +31,21 @@ struct ol_platform_impl_t {
: BackendType(BackendType), Plugin(std::move(Plugin)) {}
ol_platform_backend_t BackendType;
+ /// Get the plugin, lazily initializing it if necessary.
+ llvm::Expected<GenericPluginTy *> getPlugin() {
+ if (llvm::Error Err = init())
+ return Err;
+ return Plugin.get();
+ }
+
+ /// Get the device list, lazily initializing it if necessary.
+ llvm::Expected<llvm::SmallVector<std::unique_ptr<ol_device_impl_t>> &>
+ getDevices() {
+ if (llvm::Error Err = init())
+ return Err;
+ return Devices;
+ }
+
/// Complete all pending work for this platform and perform any needed
/// cleanup.
///
@@ -44,6 +59,11 @@ struct ol_platform_impl_t {
/// Direct access to the plugin, may be uninitialized if accessed here.
std::unique_ptr<GenericPluginTy> Plugin;
+private:
+ std::once_flag Initialized;
+ /// Most plugins don't allow to retry initialization after failure, so
+ /// remember that it failed.
+ bool InitFailed = false;
llvm::SmallVector<std::unique_ptr<ol_device_impl_t>> Devices;
};
@@ -62,27 +82,51 @@ struct ol_device_impl_t {
InfoTreeNode Info;
};
-llvm::Error ol_platform_impl_t::destroy() { return Plugin->deinit(); }
+llvm::Error ol_platform_impl_t::destroy() {
+ if (!Plugin || !Plugin->is_initialized())
+ return llvm::Error::success();
+
+ Devices.clear();
+ return Plugin->deinit();
+}
llvm::Error ol_platform_impl_t::init() {
- if (!Plugin)
- return llvm::Error::success();
+ std::unique_ptr<llvm::Error> Storage;
+
+ // This can be called concurrently, make sure we only do the actual
+ // initialization once.
+ std::call_once(Initialized, [&]() {
+ // FIXME: Need better handling for the host platform.
+ if (!Plugin)
+ return;
+
+ auto SetError = [&](llvm::Error Err) {
+ InitFailed = true;
+ Storage = std::make_unique<llvm::Error>(std::move(Err));
+ };
- if (llvm::Error Err = Plugin->init())
- return Err;
+ if (llvm::Error Err = Plugin->init())
+ return SetError(std::move(Err));
- for (auto Id = 0, End = Plugin->getNumDevices(); Id != End; Id++) {
- if (llvm::Error Err = Plugin->initDevice(Id))
- return Err;
+ for (auto Id = 0, End = Plugin->getNumDevices(); Id != End; Id++) {
+ if (llvm::Error Err = Plugin->initDevice(Id))
+ return SetError(std::move(Err));
- GenericDeviceTy *Device = &Plugin->getDevice(Id);
- llvm::Expected<InfoTreeNode> Info = Device->obtainInfo();
- if (llvm::Error Err = Info.takeError())
- return Err;
- Devices.emplace_back(std::make_unique<ol_device_impl_t>(Id, Device, *this,
- std::move(*Info)));
- }
+ GenericDeviceTy *Device = &Plugin->getDevice(Id);
+ llvm::Expected<InfoTreeNode> Info = Device->obtainInfo();
+ if (llvm::Error Err = Info.takeError())
+ return SetError(std::move(Err));
+ Devices.emplace_back(std::make_unique<ol_device_impl_t>(
+ Id, Device, *this, std::move(*Info)));
+ }
+ });
+ if (Storage)
+ return std::move(*Storage);
+ if (InitFailed)
+ return createOffloadError(
+ ErrorCode::BACKEND_FAILURE,
+ "platform not available because initialization failure");
return llvm::Error::success();
}
@@ -304,8 +348,9 @@ ol_platform_backend_t pluginNameToBackend(StringRef Name) {
#include "Shared/Targets.def"
Error initPlugins(OffloadContext &Context, const ol_init_args_t *InitArgs) {
+ bool InitAll = InitArgs && InitArgs->NumPlatforms == OL_ALL_PLATFORMS;
SmallSet<ol_platform_backend_t, 0> Requested;
- if (InitArgs && InitArgs->NumPlatforms > 0)
+ if (InitArgs && !InitAll && InitArgs->NumPlatforms > 0)
for (uint32_t I = 0; I < InitArgs->NumPlatforms; I++)
Requested.insert(InitArgs->Platforms[I]);
@@ -322,12 +367,14 @@ Error initPlugins(OffloadContext &Context, const ol_init_args_t *InitArgs) {
} while (false);
#include "Shared/Targets.def"
- // Eagerly initialize all of the plugins and devices. We need to make sure
- // that the platform is initialized at a consistent point to maintain the
- // expected teardown order in the vendor libraries.
- for (auto &Platform : Context.Platforms) {
- if (Error Err = Platform->init())
- return Err;
+ // If platforms were explicitly requested, eagerly initialize all the created
+ // ones (the requested ones and the host). Otherwise they are initialized
+ // lazily on their first use.
+ if (InitAll || !Requested.empty()) {
+ for (auto &Platform : Context.Platforms) {
+ if (Error Err = Platform->init())
+ return Err;
+ }
}
Context.TracingEnabled = std::getenv("OFFLOAD_TRACE");
@@ -348,7 +395,8 @@ Error olInit_impl(const ol_init_args_t *InitArgs) {
if (InitArgs->Size < sizeof(ol_init_args_t))
return createOffloadError(ErrorCode::INVALID_SIZE,
"ol_init_args_t Size field is too small");
- if (InitArgs->NumPlatforms > 0 && !InitArgs->Platforms)
+ if (InitArgs->NumPlatforms > 0 &&
+ InitArgs->NumPlatforms != OL_ALL_PLATFORMS && !InitArgs->Platforms)
return createOffloadError(ErrorCode::INVALID_NULL_POINTER,
"NumPlatforms > 0 but Platforms is null");
}
@@ -373,10 +421,6 @@ Error olShutDown_impl() {
auto *OldContext = OffloadContextVal.exchange(nullptr);
for (auto &Platform : OldContext->Platforms) {
- // Host plugin is nullptr and has no deinit
- if (!Platform->Plugin || !Platform->Plugin->is_initialized())
- continue;
-
if (auto Res = Platform->destroy())
Result = joinErrors(std::move(Result), std::move(Res));
}
@@ -428,7 +472,11 @@ Error olGetPlatformInfoSize_impl(ol_platform_handle_t Platform,
Error olPlatformRegisterRPCCallback_impl(ol_platform_handle_t Platform,
ol_platform_rpc_cb_t Callback) {
- Platform->Plugin->getRPCServer().registerCallback(Callback);
+ auto PluginOrErr = Platform->getPlugin();
+ if (!PluginOrErr)
+ return PluginOrErr.takeError();
+
+ (*PluginOrErr)->getRPCServer().registerCallback(Callback);
return Error::success();
}
@@ -602,7 +650,10 @@ Error olGetDeviceInfoSize_impl(ol_device_handle_t Device,
Error olIterateDevices_impl(ol_device_iterate_cb_t Callback, void *UserData) {
for (auto &Platform : OffloadContext::get().Platforms) {
- for (auto &Device : Platform->Devices) {
+ auto DevicesOrErr = Platform->getDevices();
+ if (!DevicesOrErr)
+ return DevicesOrErr.takeError();
+ for (auto &Device : *DevicesOrErr) {
if (!Callback(Device.get(), UserData)) {
return Error::success();
}
diff --git a/offload/unittests/Conformance/lib/DeviceContext.cpp b/offload/unittests/Conformance/lib/DeviceContext.cpp
index 62b265043a2eb..ac963fcc273bd 100644
--- a/offload/unittests/Conformance/lib/DeviceContext.cpp
+++ b/offload/unittests/Conformance/lib/DeviceContext.cpp
@@ -48,7 +48,11 @@ namespace {
// The static 'Wrapper' instance ensures olInit() is called once at program
// startup and olShutDown() is called once at program termination
struct OffloadInitWrapper {
- OffloadInitWrapper() { OL_CHECK(olInit(nullptr)); }
+ OffloadInitWrapper() {
+ ol_init_args_t Args = OL_INIT_ARGS_INIT;
+ Args.NumPlatforms = OL_ALL_PLATFORMS;
+ OL_CHECK(olInit(&Args));
+ }
~OffloadInitWrapper() { OL_CHECK(olShutDown()); }
};
static OffloadInitWrapper Wrapper{};
diff --git a/offload/unittests/OffloadAPI/common/Environment.cpp b/offload/unittests/OffloadAPI/common/Environment.cpp
index 89660a5d6a7b4..69bcf3ea4fb75 100644
--- a/offload/unittests/OffloadAPI/common/Environment.cpp
+++ b/offload/unittests/OffloadAPI/common/Environment.cpp
@@ -22,7 +22,9 @@ using namespace llvm;
#ifndef DISABLE_WRAPPER
struct OffloadInitWrapper {
OffloadInitWrapper() {
- if (ol_result_t Res = olInit(nullptr)) {
+ ol_init_args_t Args = OL_INIT_ARGS_INIT;
+ Args.NumPlatforms = OL_ALL_PLATFORMS;
+ if (ol_result_t Res = olInit(&Args)) {
errs() << "olInit failed: "
<< (Res->Details ? Res->Details : "(no details)") << " (code "
<< Res->Code << ")\n";
diff --git a/offload/unittests/OffloadAPI/init/olInit.cpp b/offload/unittests/OffloadAPI/init/olInit.cpp
index 4c74122e89d42..d629931f930d2 100644
--- a/offload/unittests/OffloadAPI/init/olInit.cpp
+++ b/offload/unittests/OffloadAPI/init/olInit.cpp
@@ -42,6 +42,19 @@ TEST_F(olInitTest, WithInitArgs) {
ASSERT_SUCCESS(olShutDown());
}
+TEST_F(olInitTest, WithInitArgsNoPlatforms) {
+ ol_init_args_t Args = OL_INIT_ARGS_INIT;
+ ASSERT_SUCCESS(olInit(&Args));
+ ASSERT_SUCCESS(olShutDown());
+}
+
+TEST_F(olInitTest, WithInitArgsAllPlatforms) {
+ ol_init_args_t Args = OL_INIT_ARGS_INIT;
+ Args.NumPlatforms = OL_ALL_PLATFORMS;
+ ASSERT_SUCCESS(olInit(&Args));
+ ASSERT_SUCCESS(olShutDown());
+}
+
TEST_F(olInitTest, InvalidSize) {
ol_init_args_t Args = OL_INIT_ARGS_INIT;
Args.Size = 0;
>From 6b3116af6bdeb689f43e37d502e848bb48a4e4a4 Mon Sep 17 00:00:00 2001
From: "Duran, Alex" <alejandro.duran at intel.com>
Date: Mon, 7 Sep 2026 02:42:28 -0700
Subject: [PATCH 2/4] [OFFLOAD] add olIteratePlatforms
---
offload/liboffload/API/Platform.td | 23 ++++++++++
offload/liboffload/src/OffloadImpl.cpp | 11 +++++
.../platform/olIteratePlatforms.cpp | 45 +++++++++++++++++++
3 files changed, 79 insertions(+)
create mode 100644 offload/unittests/OffloadAPI/platform/olIteratePlatforms.cpp
diff --git a/offload/liboffload/API/Platform.td b/offload/liboffload/API/Platform.td
index 62810e8fdb7ca..65efec2b8af4a 100644
--- a/offload/liboffload/API/Platform.td
+++ b/offload/liboffload/API/Platform.td
@@ -97,3 +97,26 @@ def olPlatformRegisterRPCCallback : Function {
"RPC callback function pointer", PARAM_IN>];
let returns = [Return<"OL_ERRC_INVALID_PLATFORM">, Return<"OL_ERRC_SUCCESS">];
}
+
+def ol_platform_iterate_cb_t : FptrTypedef {
+ let desc = "User-provided function to be used with `olIteratePlatforms`";
+ let params = [
+ Param<"ol_platform_handle_t", "Platform", "the platform handle of the current iteration", PARAM_IN>,
+ Param<"void*", "UserData", "optional user data", PARAM_IN_OPTIONAL>
+ ];
+ let return = "bool";
+}
+
+def olIteratePlatforms : Function {
+ let desc = "Iterates over all available platforms, calling the callback for each platform.";
+ let details = [
+ "If the user-provided callback returns `false`, the iteration is stopped."
+ ];
+ let params = [
+ Param<"ol_platform_iterate_cb_t", "Callback", "User-provided function called for each available platform", PARAM_IN>,
+ Param<"void*", "UserData", "Optional user data to pass to the callback", PARAM_IN_OPTIONAL>
+ ];
+ let returns = [
+ Return<"OL_ERRC_INVALID_PLATFORM">
+ ];
+}
diff --git a/offload/liboffload/src/OffloadImpl.cpp b/offload/liboffload/src/OffloadImpl.cpp
index 445e486b6b750..8e6f352f2a5c2 100644
--- a/offload/liboffload/src/OffloadImpl.cpp
+++ b/offload/liboffload/src/OffloadImpl.cpp
@@ -480,6 +480,17 @@ Error olPlatformRegisterRPCCallback_impl(ol_platform_handle_t Platform,
return Error::success();
}
+Error olIteratePlatforms_impl(ol_platform_iterate_cb_t Callback,
+ void *UserData) {
+ for (auto &Platform : OffloadContext::get().Platforms) {
+ if (!Callback(Platform.get(), UserData)) {
+ return Error::success();
+ }
+ }
+
+ return Error::success();
+}
+
Error olGetDeviceInfoImplDetail(ol_device_handle_t Device,
ol_device_info_t PropName, size_t PropSize,
void *PropValue, size_t *PropSizeRet) {
diff --git a/offload/unittests/OffloadAPI/platform/olIteratePlatforms.cpp b/offload/unittests/OffloadAPI/platform/olIteratePlatforms.cpp
new file mode 100644
index 0000000000000..4f8f4e35df23b
--- /dev/null
+++ b/offload/unittests/OffloadAPI/platform/olIteratePlatforms.cpp
@@ -0,0 +1,45 @@
+//===------- Offload API tests - olIteratePlatforms -----------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#include "../common/Fixtures.hpp"
+#include <OffloadAPI.h>
+#include <gtest/gtest.h>
+
+using olIteratePlatformsTest = OffloadTest;
+
+TEST_F(olIteratePlatformsTest, SuccessEmptyCallback) {
+ ASSERT_SUCCESS(olIteratePlatforms(
+ [](ol_platform_handle_t, void *) { return false; }, nullptr));
+}
+
+TEST_F(olIteratePlatformsTest, SuccessGetPlatform) {
+ uint32_t PlatformCount = 0;
+ ol_platform_handle_t Platform = nullptr;
+
+ ASSERT_SUCCESS(olIteratePlatforms(
+ [](ol_platform_handle_t, void *Data) {
+ auto Count = static_cast<uint32_t *>(Data);
+ *Count += 1;
+ return true;
+ },
+ &PlatformCount));
+
+ if (PlatformCount == 0) {
+ GTEST_SKIP() << "No available platforms.";
+ }
+
+ ASSERT_SUCCESS(olIteratePlatforms(
+ [](ol_platform_handle_t P, void *Data) {
+ auto PlatformPtr = static_cast<ol_platform_handle_t *>(Data);
+ *PlatformPtr = P;
+ return true;
+ },
+ &Platform));
+
+ ASSERT_NE(Platform, nullptr);
+}
>From 90a30430472a71e85a133e160549bb60bad39d4a Mon Sep 17 00:00:00 2001
From: "Duran, Alex" <alejandro.duran at intel.com>
Date: Mon, 7 Sep 2026 02:42:34 -0700
Subject: [PATCH 3/4] [offload][omp] Load plugins through liboffload
---
offload/include/PluginManager.h | 3 ++-
offload/liboffload/src/OffloadImpl.cpp | 6 +++++
offload/libompaccsupport/PluginManager.cpp | 28 ++++++++++++----------
3 files changed, 24 insertions(+), 13 deletions(-)
diff --git a/offload/include/PluginManager.h b/offload/include/PluginManager.h
index 6c6fdebe76dff..eea8b62a8c39d 100644
--- a/offload/include/PluginManager.h
+++ b/offload/include/PluginManager.h
@@ -13,6 +13,7 @@
#ifndef OMPTARGET_PLUGIN_MANAGER_H
#define OMPTARGET_PLUGIN_MANAGER_H
+#include "OffloadAPI.h"
#include "PluginInterface.h"
#include "DeviceImage.h"
@@ -155,7 +156,7 @@ struct PluginManager {
llvm::SmallVector<__tgt_bin_desc *> DelayedBinDesc;
// List of all plugins, in use or not.
- llvm::SmallVector<std::unique_ptr<GenericPluginTy>> Plugins;
+ llvm::SmallVector<GenericPluginTy *> Plugins;
// Mapping of plugins to the OpenMP device identifier.
llvm::DenseMap<std::pair<const GenericPluginTy *, int32_t>, int32_t>
diff --git a/offload/liboffload/src/OffloadImpl.cpp b/offload/liboffload/src/OffloadImpl.cpp
index 8e6f352f2a5c2..def07f5fa0725 100644
--- a/offload/liboffload/src/OffloadImpl.cpp
+++ b/offload/liboffload/src/OffloadImpl.cpp
@@ -1493,5 +1493,11 @@ Error olQueryQueue_impl(ol_queue_handle_t Queue, bool *IsQueueWorkCompleted) {
return Error::success();
}
+// Temporary helpers to help transition of libomptarget to liboffload
+extern "C" GenericPluginTy *
+__ol_tgt_GetPluginFromPlatform(ol_platform_handle_t Platform) {
+ return Platform->Plugin.get();
+}
+
} // namespace offload
} // namespace llvm
diff --git a/offload/libompaccsupport/PluginManager.cpp b/offload/libompaccsupport/PluginManager.cpp
index 5f67a204f4203..514d7347b5cec 100644
--- a/offload/libompaccsupport/PluginManager.cpp
+++ b/offload/libompaccsupport/PluginManager.cpp
@@ -31,9 +31,8 @@ using namespace llvm::omp::target::debug;
PluginManager *PM = nullptr;
-// Every plugin exports this method to create an instance of the plugin type.
-#define PLUGIN_TARGET(Name) extern "C" GenericPluginTy *createPlugin_##Name();
-#include "Shared/Targets.def"
+extern "C" GenericPluginTy *
+__ol_tgt_GetPluginFromPlatform(ol_platform_handle_t Platform);
void PluginManager::init() {
TIMESCOPE();
@@ -43,14 +42,20 @@ void PluginManager::init() {
}
ODBG(ODT_Init) << "Loading RTLs";
-
- // Attempt to create an instance of each supported plugin.
-#define PLUGIN_TARGET(Name) \
- do { \
- Plugins.emplace_back( \
- std::unique_ptr<GenericPluginTy>(createPlugin_##Name())); \
- } while (false);
-#include "Shared/Targets.def"
+ if (ol_result_t Res = olInit(nullptr))
+ REPORT() << "Failed to initialize liboffload: " << Res->Details;
+
+ if (ol_result_t Res = olIteratePlatforms(
+ [](ol_platform_handle_t Platform, void *Data) {
+ auto *PM = static_cast<PluginManager *>(Data);
+ auto *Plugin = __ol_tgt_GetPluginFromPlatform(Platform);
+ ODBG(ODT_Init) << "Adding plugin " << Plugin->getName()
+ << " from liboffload";
+ PM->Plugins.push_back(Plugin);
+ return true;
+ },
+ this))
+ REPORT() << "Failed to iterate platforms: " << Res->Details;
ODBG(ODT_Init) << "RTLs loaded!";
}
@@ -73,7 +78,6 @@ void PluginManager::deinit() {
std::string InfoMsg = toString(std::move(Err));
ODBG(ODT_Deinit) << "Failed to deinit plugin: " << InfoMsg;
}
- Plugin.release();
}
ODBG(ODT_Deinit) << "RTLs unloaded!";
>From 99ac64c1b433e95ca0df8bddfb8b9f2670d60189 Mon Sep 17 00:00:00 2001
From: "Duran, Alex" <alejandro.duran at intel.com>
Date: Mon, 7 Sep 2026 03:17:04 -0700
Subject: [PATCH 4/4] [OFFLOAD] Add olIterateCompatibleDevices API
---
offload/liboffload/API/Program.td | 15 ++++
offload/liboffload/src/OffloadImpl.cpp | 27 ++++++-
offload/unittests/OffloadAPI/CMakeLists.txt | 1 +
.../platform/olIteratePlatforms.cpp | 3 +-
.../program/olIterateCompatibleDevices.cpp | 73 +++++++++++++++++++
5 files changed, 115 insertions(+), 4 deletions(-)
create mode 100644 offload/unittests/OffloadAPI/program/olIterateCompatibleDevices.cpp
diff --git a/offload/liboffload/API/Program.td b/offload/liboffload/API/Program.td
index ecc8fb73dad3f..9f945872aec3c 100644
--- a/offload/liboffload/API/Program.td
+++ b/offload/liboffload/API/Program.td
@@ -46,6 +46,21 @@ def olIsValidBinary : Function {
let returns = [];
}
+def olIterateCompatibleDevices : Function {
+ let desc = "Iterates over all available devices that are compatible with the binary image pointed to by `ProgData`, calling the callback for each device.";
+ let details = [
+ "The provided `ProgData` will not be loaded onto any device",
+ "If the user-provided callback returns `false`, the iteration is stopped."
+ ];
+ let params = [
+ Param<"const void*", "ProgData", "pointer to the program binary data", PARAM_IN>,
+ Param<"size_t", "ProgDataSize", "size of the program binary in bytes", PARAM_IN>,
+ Param<"ol_device_iterate_cb_t", "Callback", "User-provided function called for each compatible device", PARAM_IN>,
+ Param<"void*", "UserData", "Optional user data to pass to the callback", PARAM_IN_OPTIONAL>
+ ];
+ let returns = [];
+}
+
def olDestroyProgram : Function {
let desc = "Destroy the program and free all underlying resources.";
let details = [];
diff --git a/offload/liboffload/src/OffloadImpl.cpp b/offload/liboffload/src/OffloadImpl.cpp
index def07f5fa0725..c14688caa085d 100644
--- a/offload/liboffload/src/OffloadImpl.cpp
+++ b/offload/liboffload/src/OffloadImpl.cpp
@@ -665,9 +665,32 @@ Error olIterateDevices_impl(ol_device_iterate_cb_t Callback, void *UserData) {
if (!DevicesOrErr)
return DevicesOrErr.takeError();
for (auto &Device : *DevicesOrErr) {
- if (!Callback(Device.get(), UserData)) {
+ if (!Callback(Device.get(), UserData))
+ return Error::success();
+ }
+ }
+
+ return Error::success();
+}
+
+Error olIterateCompatibleDevices_impl(const void *ProgData, size_t ProgDataSize,
+ ol_device_iterate_cb_t Callback,
+ void *UserData) {
+ StringRef Buffer(reinterpret_cast<const char *>(ProgData), ProgDataSize);
+
+ for (auto &Platform : OffloadContext::get().Platforms) {
+ if (!Platform->Plugin || !Platform->Plugin->isPluginCompatible(Buffer))
+ continue;
+ auto DevicesOrErr = Platform->getDevices();
+ if (!DevicesOrErr)
+ return DevicesOrErr.takeError();
+ for (auto &Device : *DevicesOrErr) {
+ if (!Device->Platform.Plugin->isDeviceCompatible(Device->DeviceNum,
+ Buffer))
+ continue;
+
+ if (!Callback(Device.get(), UserData))
return Error::success();
- }
}
}
diff --git a/offload/unittests/OffloadAPI/CMakeLists.txt b/offload/unittests/OffloadAPI/CMakeLists.txt
index 292ea1eb4852f..8bed11be4d85e 100644
--- a/offload/unittests/OffloadAPI/CMakeLists.txt
+++ b/offload/unittests/OffloadAPI/CMakeLists.txt
@@ -49,6 +49,7 @@ add_offload_unittest("platform"
add_offload_unittest("program"
program/olCreateProgram.cpp
program/olIsValidBinary.cpp
+ program/olIterateCompatibleDevices.cpp
program/olDestroyProgram.cpp)
add_offload_unittest("queue"
diff --git a/offload/unittests/OffloadAPI/platform/olIteratePlatforms.cpp b/offload/unittests/OffloadAPI/platform/olIteratePlatforms.cpp
index 4f8f4e35df23b..a4446c583dab1 100644
--- a/offload/unittests/OffloadAPI/platform/olIteratePlatforms.cpp
+++ b/offload/unittests/OffloadAPI/platform/olIteratePlatforms.cpp
@@ -29,9 +29,8 @@ TEST_F(olIteratePlatformsTest, SuccessGetPlatform) {
},
&PlatformCount));
- if (PlatformCount == 0) {
+ if (PlatformCount == 0)
GTEST_SKIP() << "No available platforms.";
- }
ASSERT_SUCCESS(olIteratePlatforms(
[](ol_platform_handle_t P, void *Data) {
diff --git a/offload/unittests/OffloadAPI/program/olIterateCompatibleDevices.cpp b/offload/unittests/OffloadAPI/program/olIterateCompatibleDevices.cpp
new file mode 100644
index 0000000000000..afa26c57d6728
--- /dev/null
+++ b/offload/unittests/OffloadAPI/program/olIterateCompatibleDevices.cpp
@@ -0,0 +1,73 @@
+//===------- Offload API tests - olIterateCompatibleDevices -------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#include "../common/Fixtures.hpp"
+#include <OffloadAPI.h>
+#include <gtest/gtest.h>
+
+using olIterateCompatibleDevicesTest = OffloadDeviceTest;
+OFFLOAD_TESTS_INSTANTIATE_DEVICE_FIXTURE(olIterateCompatibleDevicesTest);
+
+TEST_P(olIterateCompatibleDevicesTest, Success) {
+ std::unique_ptr<llvm::MemoryBuffer> DeviceBin;
+ ASSERT_TRUE(TestEnvironment::loadDeviceBinary("foo", Device, DeviceBin));
+ ASSERT_GE(DeviceBin->getBufferSize(), 0lu);
+
+ struct CallbackDataTy {
+ ol_device_handle_t ExpectedDevice;
+ bool Found = false;
+ } CallbackData{Device};
+
+ ASSERT_SUCCESS(olIterateCompatibleDevices(
+ DeviceBin->getBufferStart(), DeviceBin->getBufferSize(),
+ [](ol_device_handle_t D, void *UserData) {
+ auto *Data = static_cast<CallbackDataTy *>(UserData);
+ if (D == Data->ExpectedDevice)
+ Data->Found = true;
+ return true;
+ },
+ &CallbackData));
+
+ ASSERT_TRUE(CallbackData.Found);
+}
+
+TEST_P(olIterateCompatibleDevicesTest, SuccessStopIteration) {
+ std::unique_ptr<llvm::MemoryBuffer> DeviceBin;
+ ASSERT_TRUE(TestEnvironment::loadDeviceBinary("foo", Device, DeviceBin));
+ ASSERT_GE(DeviceBin->getBufferSize(), 0lu);
+
+ uint32_t CallCount = 0;
+ ASSERT_SUCCESS(olIterateCompatibleDevices(
+ DeviceBin->getBufferStart(), DeviceBin->getBufferSize(),
+ [](ol_device_handle_t, void *UserData) {
+ auto *Count = static_cast<uint32_t *>(UserData);
+ *Count += 1;
+ return false;
+ },
+ &CallCount));
+
+ ASSERT_EQ(CallCount, 1u);
+}
+
+TEST_P(olIterateCompatibleDevicesTest, EmptyBinary) {
+ std::unique_ptr<llvm::MemoryBuffer> DeviceBin;
+ ASSERT_TRUE(TestEnvironment::loadDeviceBinary("foo", Device, DeviceBin));
+ ASSERT_GE(DeviceBin->getBufferSize(), 0lu);
+
+ uint32_t CallCount = 0;
+ ASSERT_SUCCESS(olIterateCompatibleDevices(
+ DeviceBin->getBufferStart(), 0,
+ [](ol_device_handle_t, void *UserData) {
+ auto *Count = static_cast<uint32_t *>(UserData);
+ *Count += 1;
+ return true;
+ },
+ &CallCount));
+
+ ASSERT_EQ(CallCount, 0u);
+}
More information about the llvm-branch-commits
mailing list