[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