[llvm] [offload] Initilize Platforms lazily (PR #226438)
via llvm-commits
llvm-commits at lists.llvm.org
Fri Sep 25 08:11:00 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-offload
Author: Alex Duran (adurang)
<details>
<summary>Changes</summary>
Add support to Initialize liboffload Platforms when required instead of eagerly on olInit. This allows for users like libomptarget that can filter the Platforms before olInit to still not initialize all of them.
The approach taken was based on the one implemeted by @<!-- -->jhuber6 in #<!-- -->163272 while adding the possibility to request Platforms to be initialized eagerly for the cases where init/deinit happens in constructors/destructors. Unit tests were updated where necessary to request this.
Assisted by Claude.
Co-authored-by: Joseph Huber <huberjn@<!-- -->outlook.com>
---
Full diff: https://github.com/llvm/llvm-project/pull/226438.diff
5 Files Affected:
- (modified) offload/liboffload/API/Common.td (+10-3)
- (modified) offload/liboffload/src/OffloadImpl.cpp (+80-29)
- (modified) offload/unittests/Conformance/lib/DeviceContext.cpp (+5-1)
- (modified) offload/unittests/OffloadAPI/common/Environment.cpp (+3-1)
- (modified) offload/unittests/OffloadAPI/init/olInit.cpp (+13)
``````````diff
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;
``````````
</details>
https://github.com/llvm/llvm-project/pull/226438
More information about the llvm-commits
mailing list