https://github.com/adurang created https://github.com/llvm/llvm-project/pull/221270
None >From e09a0da04b33c18f988df09c13d60c9a797800bc Mon Sep 17 00:00:00 2001 From: "Duran, Alex" <[email protected]> Date: Thu, 3 Sep 2026 10:32:33 -0700 Subject: [PATCH 1/2] [OFFLOAD] Initialize Platforms and Devices lazily --- offload/liboffload/src/OffloadImpl.cpp | 308 ++++++++++++++++++------- 1 file changed, 226 insertions(+), 82 deletions(-) diff --git a/offload/liboffload/src/OffloadImpl.cpp b/offload/liboffload/src/OffloadImpl.cpp index 6e608e3f965c4..670179c21f054 100644 --- a/offload/liboffload/src/OffloadImpl.cpp +++ b/offload/liboffload/src/OffloadImpl.cpp @@ -56,9 +56,11 @@ struct ol_platform_impl_t { /// Initialize the associated plugin and devices. llvm::Error init(); - /// Direct access to the plugin, may be uninitialized if accessed here. + bool Initialized = false; std::unique_ptr<GenericPluginTy> Plugin; + bool isInitialized() const { return Initialized; } + llvm::SmallVector<std::unique_ptr<ol_device_impl_t>> Devices; }; @@ -66,20 +68,58 @@ struct ol_platform_impl_t { // we add some additional data here for now to avoid churn in the plugin // interface. struct ol_device_impl_t { - ol_device_impl_t(int DeviceNum, GenericDeviceTy *Device, - ol_platform_impl_t &Platform, InfoTreeNode &&DevInfo) - : DeviceNum(DeviceNum), Device(Device), Platform(Platform), - Info(std::forward<InfoTreeNode>(DevInfo)) {} - + ol_device_impl_t(int DeviceNum, ol_platform_impl_t &Platform) + : DeviceNum(DeviceNum), Platform(Platform) {} int DeviceNum; - GenericDeviceTy *Device; ol_platform_impl_t &Platform; + + llvm::Error init() { + if (!Platform.isInitialized()) { + if (auto Err = Platform.init()) + return Err; + } + + if (llvm::Error Err = Platform.Plugin->initDevice(DeviceNum)) + return Err; + + Device = &Platform.Plugin->getDevice(DeviceNum); + llvm::Expected<InfoTreeNode> InfoOrErr = Device->obtainInfo(); + if (!InfoOrErr) + return InfoOrErr.takeError(); + Info = std::move(*InfoOrErr); + + return llvm::Error::success(); + } + + llvm::Expected<GenericDeviceTy *> getDevice() { + if (!Device) { + if (llvm::Error Err = init()) + return Err; + } + + return Device; + } + + llvm::Expected<InfoTreeNode &> getInfo() { + if (!Device) { + if (llvm::Error Err = init()) + return Err; + } + + return Info; + } +private: + GenericDeviceTy *Device = nullptr; InfoTreeNode Info; }; llvm::Error ol_platform_impl_t::destroy() { return Plugin->deinit(); } llvm::Error ol_platform_impl_t::init() { + if (Initialized) + return llvm::Error::success(); + Initialized = true; + if (!Plugin) return llvm::Error::success(); @@ -87,15 +127,7 @@ llvm::Error ol_platform_impl_t::init() { return Err; for (auto Id = 0, End = Plugin->getNumDevices(); Id != End; Id++) { - if (llvm::Error Err = Plugin->initDevice(Id)) - return 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))); + Devices.emplace_back(std::make_unique<ol_device_impl_t>(Id, *this)); } return llvm::Error::success(); @@ -189,10 +221,12 @@ struct ol_context_impl_t { return nullptr; auto &Bucket = It->second; + GenericDeviceTy *DeviceImpl = llvm::cantFail(Device->getDevice()); + // As queues are pulled and popped from this list, longer running queues // naturally bubble to the start of the array. Hence looping backwards. for (auto Q = Bucket.rbegin(); Q != Bucket.rend(); Q++) { - if (!Device->Device->hasPendingWork(*Q)) { + if (!DeviceImpl->hasPendingWork(*Q)) { auto OutstandingQueue = *Q; *Q = Bucket.back(); Bucket.pop_back(); @@ -214,8 +248,13 @@ struct ol_context_impl_t { llvm::Error Result = Plugin::success(); for (auto &Bucket : OutstandingQueues) { auto *Device = Bucket.first; + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) { + Result = llvm::joinErrors(std::move(Result), DeviceOrErr.takeError()); + continue; + } for (auto *AI : Bucket.second) - if (auto Err = Device->Device->synchronize(AI, /*Release=*/true)) + if (auto Err = (*DeviceOrErr)->synchronize(AI, /*Release=*/true)) Result = llvm::joinErrors(std::move(Result), std::move(Err)); } OutstandingQueues.clear(); @@ -326,14 +365,6 @@ 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; - } - Context.TracingEnabled = std::getenv("OFFLOAD_TRACE"); Context.ValidationEnabled = !std::getenv("OFFLOAD_DISABLE_VALIDATION"); @@ -378,7 +409,8 @@ Error olShutDown_impl() { for (auto &Platform : OldContext->Platforms) { // Host plugin is nullptr and has no deinit - if (!Platform->Plugin || !Platform->Plugin->is_initialized()) + if (!Platform->isInitialized() || !Platform->Plugin || + !Platform->Plugin->is_initialized()) continue; if (auto Res = Platform->destroy()) @@ -476,8 +508,11 @@ Error olGetDeviceInfoImplDetail(ol_device_handle_t Device, // AMD doesn't provide the global memory size (trivially) with the device info // struct, so use the plugin interface case OL_DEVICE_INFO_GLOBAL_MEM_SIZE: { + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); uint64_t Mem; - if (auto Err = Device->Device->getDeviceMemorySize(Mem)) + if (auto Err = (*DeviceOrErr)->getDeviceMemorySize(Mem)) return Err; return Info.write<uint64_t>(Mem); } break; @@ -490,7 +525,11 @@ Error olGetDeviceInfoImplDetail(ol_device_handle_t Device, return createOffloadError(ErrorCode::INVALID_ENUMERATION, "getDeviceInfo enum '%i' is invalid", PropName); - auto EntryOpt = Device->Info.get(static_cast<DeviceInfo>(PropName)); + auto InfoOrErr = Device->getInfo(); + if (!InfoOrErr) + return InfoOrErr.takeError(); + + auto EntryOpt = InfoOrErr->get(static_cast<DeviceInfo>(PropName)); if (!EntryOpt) return makeError(ErrorCode::UNIMPLEMENTED, "plugin did not provide a response for this information"); @@ -603,6 +642,8 @@ 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) { + if (auto Err = Platform->init()) + return Err; for (auto &Device : Platform->Devices) { if (!Callback(Device.get(), UserData)) { return Error::success(); @@ -626,14 +667,20 @@ Error olCreateContext_impl(size_t DevicesCount, ol_device_handle_t *Devices, ErrorCode::INVALID_DEVICE, "all devices in a context must belong to the same platform"); DeviceList.push_back(Devices[I]); - PluginDevices.push_back(Devices[I]->Device); + auto DeviceOrErr = Devices[I]->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + PluginDevices.push_back(*DeviceOrErr); } // The host plugin has no GenericPluginTy instance; skip the plugin-side // context in that case and just record the device set. std::unique_ptr<plugin::PluginContextTy> PluginCtx; if (Platform->Plugin) { - auto PluginCtxOrErr = Platform->Plugin->createPluginContext(PluginDevices); + if (auto Err = Platform->init()) + return Err; + auto PluginCtxOrErr = + Platform->Plugin->createPluginContext(PluginDevices); if (!PluginCtxOrErr) return PluginCtxOrErr.takeError(); PluginCtx = std::move(*PluginCtxOrErr); @@ -698,13 +745,18 @@ constexpr size_t MAX_ALLOC_TRIES = 50; Error olMemAllocImplHelper(ol_device_handle_t Device, ol_alloc_type_t Type, size_t Size, size_t Alignment, void **AllocationOut) { + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + GenericDeviceTy *DeviceImpl = *DeviceOrErr; + SmallVector<void *> Rejects; // Repeat the allocation up to a certain amount of times. If it happens to // already be allocated (e.g. by a device from another vendor) throw it away // and try again. for (size_t Count = 0; Count < MAX_ALLOC_TRIES; Count++) { - auto NewAlloc = Device->Device->dataAlloc( + auto NewAlloc = DeviceImpl->dataAlloc( Size, nullptr, convertOlToPluginAllocTy(Type), Alignment); if (!NewAlloc) return NewAlloc.takeError(); @@ -736,7 +788,7 @@ Error olMemAllocImplHelper(ol_device_handle_t Device, ol_alloc_type_t Type, for (void *R : Rejects) if (auto Err = - Device->Device->dataDelete(R, convertOlToPluginAllocTy(Type))) + DeviceImpl->dataDelete(R, convertOlToPluginAllocTy(Type))) return Err; return Error::success(); } @@ -800,8 +852,12 @@ Error olMemFree_impl(void *Address) { Bases.erase(std::lower_bound(Bases.begin(), Bases.end(), Address)); } + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + if (auto Res = - Device->Device->dataDelete(Address, convertOlToPluginAllocTy(Type))) + (*DeviceOrErr)->dataDelete(Address, convertOlToPluginAllocTy(Type))) return Res; return Error::success(); @@ -868,16 +924,20 @@ Error olCreateQueue_impl(ol_context_handle_t Context, ol_device_handle_t Device, auto CreatedQueue = std::make_unique<ol_queue_impl_t>(nullptr, Context, Device); + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + auto OutstandingQueue = Context->getOutstandingQueue(Device); if (OutstandingQueue) { // The queue is empty, but we still need to sync it to release any temporary // memory allocations or do other cleanup. if (auto Err = - Device->Device->synchronize(OutstandingQueue, /*Release=*/false)) + (*DeviceOrErr)->synchronize(OutstandingQueue, /*Release=*/false)) return Err; CreatedQueue->AsyncInfo = OutstandingQueue; } else if (auto Err = Context->PluginCtx->initAsyncInfo( - *Device->Device, &(CreatedQueue->AsyncInfo))) { + **DeviceOrErr, &(CreatedQueue->AsyncInfo))) { return Err; } @@ -888,17 +948,22 @@ Error olCreateQueue_impl(ol_context_handle_t Context, ol_device_handle_t Device, Error olDestroyQueue_impl(ol_queue_handle_t Queue) { auto *Device = Queue->Device; auto *Context = Queue->Context; + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + auto *DeviceImpl = *DeviceOrErr; + // This is safe; as soon as olDestroyQueue is called it is not possible to add // any more work to the queue, so if it's finished now it will remain finished // forever. - auto Res = Device->Device->hasPendingWork(Queue->AsyncInfo); + auto Res = DeviceImpl->hasPendingWork(Queue->AsyncInfo); if (!Res) return Res.takeError(); if (!*Res) { // The queue is complete, so sync it and throw it back into the pool. - if (auto Err = Device->Device->synchronize(Queue->AsyncInfo, - /*Release=*/true)) + if (auto Err = DeviceImpl->synchronize(Queue->AsyncInfo, + /*Release=*/true)) return Err; } else { // The queue still has outstanding work. Store it so we can check it later. @@ -915,7 +980,10 @@ Error olSyncQueue_impl(ol_queue_handle_t Queue) { // We don't need to release the queue and we would like the ability for // other offload threads to submit work concurrently, so pass "false" here // so we don't release the underlying queue object. - if (auto Err = Queue->Device->Device->synchronize(Queue->AsyncInfo, false)) + auto DeviceOrErr = Queue->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + if (auto Err = (*DeviceOrErr)->synchronize(Queue->AsyncInfo, false)) return Err; } @@ -924,7 +992,10 @@ Error olSyncQueue_impl(ol_queue_handle_t Queue) { Error olWaitEvents_impl(ol_queue_handle_t Queue, ol_event_handle_t *Events, size_t NumEvents) { - auto *Device = Queue->Device->Device; + auto DeviceOrErr = Queue->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + auto *Device = *DeviceOrErr; for (size_t I = 0; I < NumEvents; I++) { auto *Event = Events[I]; @@ -956,7 +1027,10 @@ Error olGetQueueInfoImplDetail(ol_queue_handle_t Queue, case OL_QUEUE_INFO_CONTEXT: return Info.write<ol_context_handle_t>(Queue->Context); case OL_QUEUE_INFO_EMPTY: { - auto Pending = Queue->Device->Device->hasPendingWork(Queue->AsyncInfo); + auto DeviceOrErr = Queue->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + auto Pending = (*DeviceOrErr)->hasPendingWork(Queue->AsyncInfo); if (auto Err = Pending.takeError()) return Err; return Info.write<bool>(!*Pending); @@ -984,7 +1058,11 @@ Error olSyncEvent_impl(ol_event_handle_t Event) { if (!Event->EventInfo) return Plugin::success(); - if (auto Res = Event->Device->Device->syncEvent(Event->EventInfo)) + auto DeviceOrErr = Event->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + + if (auto Res = (*DeviceOrErr)->syncEvent(Event->EventInfo)) return Res; return Error::success(); @@ -1004,7 +1082,11 @@ Error olGetEventElapsedTime_impl(ol_event_handle_t StartEvent, ErrorCode::INVALID_DEVICE, "StartEvent and EndEvent must belong to the same device"); - auto ElapsedTimeOrErr = StartEvent->Device->Device->getEventElapsedTime( + auto DeviceOrErr = StartEvent->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + + auto ElapsedTimeOrErr = (*DeviceOrErr)->getEventElapsedTime( StartEvent->EventInfo, EndEvent->EventInfo); if (!ElapsedTimeOrErr) return ElapsedTimeOrErr.takeError(); @@ -1014,10 +1096,14 @@ Error olGetEventElapsedTime_impl(ol_event_handle_t StartEvent, } Error olDestroyEvent_impl(ol_event_handle_t Event) { - if (Event->EventInfo) - if (auto Res = Event->Device->Device->destroyEvent(Event->EventInfo, - Event->ProfilingEnabled)) + if (Event->EventInfo) { + auto DeviceOrErr = Event->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + if (auto Res = (*DeviceOrErr)->destroyEvent(Event->EventInfo, + Event->ProfilingEnabled)) return Res; + } return olDestroy(Event); } @@ -1037,8 +1123,11 @@ Error olGetEventInfoImplDetail(ol_event_handle_t Event, if (!Event->EventInfo) return Info.write<bool>(true); - auto Res = Queue->Device->Device->isEventComplete(Event->EventInfo, - Queue->AsyncInfo); + auto DeviceOrErr = Queue->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + auto Res = (*DeviceOrErr)->isEventComplete(Event->EventInfo, + Queue->AsyncInfo); if (auto Err = Res.takeError()) return Err; return Info.write<bool>(*Res); @@ -1067,14 +1156,18 @@ Error olCreateEvent_impl(ol_queue_handle_t Queue, ol_event_flags_t Flags, auto Event = std::make_unique<ol_event_impl_t>(nullptr, Queue->Device, Queue, EnableProfiling); - if (auto Err = Queue->Device->Device->createEvent(&Event->EventInfo, - EnableProfiling)) + auto DeviceOrErr = Queue->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + auto *DeviceImpl = *DeviceOrErr; + + if (auto Err = DeviceImpl->createEvent(&Event->EventInfo, EnableProfiling)) return Err; - if (auto Err = Queue->Device->Device->recordEvent( + if (auto Err = DeviceImpl->recordEvent( Event->EventInfo, Queue->AsyncInfo, EnableProfiling)) { if (Event->EventInfo) { - if (auto DestroyErr = Queue->Device->Device->destroyEvent( + if (auto DestroyErr = DeviceImpl->destroyEvent( Event->EventInfo, EnableProfiling)) return joinErrors(std::move(Err), std::move(DestroyErr)); } @@ -1092,14 +1185,26 @@ Error olMemcpy_impl(ol_queue_handle_t Queue, void *DstPtr, bool IsDstHost = DstDevice->Platform.BackendType == OL_PLATFORM_BACKEND_HOST; bool IsSrcHost = SrcDevice->Platform.BackendType == OL_PLATFORM_BACKEND_HOST; + auto DstDeviceImplOrErr = DstDevice->getDevice(); + if (!DstDeviceImplOrErr) + return DstDeviceImplOrErr.takeError(); + auto SrcDeviceImplOrErr = SrcDevice->getDevice(); + if (!SrcDeviceImplOrErr) + return SrcDeviceImplOrErr.takeError(); + auto *DstDeviceImpl = *DstDeviceImplOrErr; + auto *SrcDeviceImpl = *SrcDeviceImplOrErr; + if (IsDstHost && IsSrcHost) { if (!Queue) { std::memcpy(DstPtr, SrcPtr, Size); return Error::success(); } - return Queue->Device->Device->dataMemcpy(DstPtr, SrcPtr, Size, - Queue->AsyncInfo); + auto QueueDeviceOrErr = Queue->Device->getDevice(); + if (!QueueDeviceOrErr) + return QueueDeviceOrErr.takeError(); + return (*QueueDeviceOrErr)->dataMemcpy(DstPtr, SrcPtr, Size, + Queue->AsyncInfo); } // If no queue is given the memcpy will be synchronous @@ -1107,18 +1212,19 @@ Error olMemcpy_impl(ol_queue_handle_t Queue, void *DstPtr, if (IsDstHost) { if (auto Res = - SrcDevice->Device->dataRetrieve(DstPtr, SrcPtr, Size, QueueImpl)) + SrcDeviceImpl->dataRetrieve(DstPtr, SrcPtr, Size, QueueImpl)) return Res; } else if (IsSrcHost) { if (auto Res = - DstDevice->Device->dataSubmit(DstPtr, SrcPtr, Size, QueueImpl)) + DstDeviceImpl->dataSubmit(DstPtr, SrcPtr, Size, QueueImpl)) return Res; - } else if (SrcDevice->Platform.Plugin == DstDevice->Platform.Plugin && + } else if (SrcDevice->Platform.Plugin == + DstDevice->Platform.Plugin && SrcDevice->Platform.Plugin->isDataExchangable( - SrcDevice->Device->getDeviceId(), - DstDevice->Device->getDeviceId())) { - if (auto Res = SrcDevice->Device->dataExchange(SrcPtr, *DstDevice->Device, - DstPtr, Size, QueueImpl)) + SrcDeviceImpl->getDeviceId(), + DstDeviceImpl->getDeviceId())) { + if (auto Res = SrcDeviceImpl->dataExchange(SrcPtr, *DstDeviceImpl, + DstPtr, Size, QueueImpl)) return Res; } else { if (Queue) @@ -1129,9 +1235,9 @@ Error olMemcpy_impl(ol_queue_handle_t Queue, void *DstPtr, if (!Buffer) return createOffloadError(ErrorCode::OUT_OF_RESOURCES, "Couldn't allocate a buffer for transfer"); - Error Res = SrcDevice->Device->dataRetrieve(Buffer, SrcPtr, Size, nullptr); + Error Res = SrcDeviceImpl->dataRetrieve(Buffer, SrcPtr, Size, nullptr); if (!Res) - Res = DstDevice->Device->dataSubmit(DstPtr, Buffer, Size, nullptr); + Res = DstDeviceImpl->dataSubmit(DstPtr, Buffer, Size, nullptr); free(Buffer); return Res; @@ -1142,8 +1248,11 @@ Error olMemcpy_impl(ol_queue_handle_t Queue, void *DstPtr, Error olMemFill_impl(ol_queue_handle_t Queue, void *Ptr, size_t PatternSize, const void *PatternPtr, size_t FillSize) { - return Queue->Device->Device->dataFill(Ptr, PatternPtr, PatternSize, FillSize, - Queue->AsyncInfo); + auto DeviceOrErr = Queue->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + return (*DeviceOrErr)->dataFill(Ptr, PatternPtr, PatternSize, FillSize, + Queue->AsyncInfo); } Error olMemPrefetch_impl(ol_queue_handle_t Queue, size_t Count, @@ -1153,8 +1262,11 @@ Error olMemPrefetch_impl(ol_queue_handle_t Queue, size_t Count, return Error::success(); bool ToHost = (Flags & OL_MEM_MIGRATION_FLAG_DEVICE_TO_HOST) != 0; - return Queue->Device->Device->dataPrefetch(Count, Mems, Sizes, ToHost, - Queue->AsyncInfo); + auto DeviceOrErr = Queue->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + return (*DeviceOrErr)->dataPrefetch(Count, Mems, Sizes, ToHost, + Queue->AsyncInfo); } Error olCreateProgram_impl(ol_context_handle_t Context, @@ -1165,8 +1277,13 @@ Error olCreateProgram_impl(ol_context_handle_t Context, "device does not belong to the given context"); StringRef Buffer(reinterpret_cast<const char *>(ProgData), ProgDataSize); - Expected<plugin::DeviceImageTy *> Res = Device->Device->loadBinary( - Device->Device->Plugin, Buffer, Context->PluginCtx.get()); + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + auto *DeviceImpl = *DeviceOrErr; + + Expected<plugin::DeviceImageTy *> Res = DeviceImpl->loadBinary( + DeviceImpl->Plugin, Buffer, Context->PluginCtx.get()); if (!Res) return Res.takeError(); assert(*Res && "loadBinary returned nullptr"); @@ -1178,9 +1295,12 @@ Error olCreateProgram_impl(ol_context_handle_t Context, Error olIsValidBinary_impl(ol_device_handle_t Device, const void *ProgData, size_t ProgDataSize, bool *IsValid) { StringRef Buffer(reinterpret_cast<const char *>(ProgData), ProgDataSize); - *IsValid = Device->Device ? Device->Device->Plugin.isDeviceCompatible( - Device->Device->getDeviceId(), Buffer) - : false; + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + auto *DeviceImpl = *DeviceOrErr; + *IsValid = + DeviceImpl->Plugin.isDeviceCompatible(DeviceImpl->getDeviceId(), Buffer); return Error::success(); } @@ -1205,7 +1325,11 @@ Error olCalculateOptimalOccupancy_impl(ol_device_handle_t Device, "provided symbol is not a kernel"); auto *KernelImpl = std::get<GenericKernelTy *>(Kernel->PluginImpl); - auto Res = KernelImpl->maxGroupSize(*Device->Device, DynamicMemSize); + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + + auto Res = KernelImpl->maxGroupSize(**DeviceOrErr, DynamicMemSize); if (auto Err = Res.takeError()) return Err; @@ -1222,7 +1346,10 @@ Error olGetKernelMaxCooperativeGroupCount_impl( return createOffloadError(ErrorCode::SYMBOL_KIND, "provided symbol is not a kernel"); - GenericDeviceTy *DeviceImpl = Device->Device; + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + GenericDeviceTy *DeviceImpl = *DeviceOrErr; auto *KernelImpl = std::get<GenericKernelTy *>(Kernel->PluginImpl); // Extract work group size from LaunchSizeArgs @@ -1248,7 +1375,6 @@ Error olLaunchKernel_impl(ol_queue_handle_t Queue, ol_device_handle_t Device, const ol_kernel_launch_prop_t *Properties, size_t NumArgs, void **ArgPtrs, const size_t *ArgSizes) { - auto *DeviceImpl = Device->Device; if (Queue && Device != Queue->Device) { return createOffloadError( ErrorCode::INVALID_DEVICE, @@ -1259,6 +1385,11 @@ Error olLaunchKernel_impl(ol_queue_handle_t Queue, ol_device_handle_t Device, return createOffloadError(ErrorCode::SYMBOL_KIND, "provided symbol is not a kernel"); + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + auto *DeviceImpl = *DeviceOrErr; + auto *QueueImpl = Queue ? Queue->AsyncInfo : nullptr; KernelLaunchArgsTy LaunchArgs{}; LaunchArgs.NumArgs = static_cast<uint32_t>(NumArgs); @@ -1438,13 +1569,20 @@ Error olGetSymbolInfoSize_impl(ol_symbol_handle_t Symbol, Error olLaunchHostFunction_impl(ol_queue_handle_t Queue, ol_host_function_cb_t Callback, void *UserData) { - return Queue->Device->Device->enqueueHostCall(Callback, UserData, - Queue->AsyncInfo); + auto DeviceOrErr = Queue->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + return (*DeviceOrErr)->enqueueHostCall(Callback, UserData, + Queue->AsyncInfo); } Error olMemRegister_impl(ol_device_handle_t Device, void *Ptr, size_t Size, ol_memory_register_flags_t Flags, void **LockedPtr) { - Expected<void *> LockedPtrOrErr = Device->Device->registerMemory( + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + + Expected<void *> LockedPtrOrErr = (*DeviceOrErr)->registerMemory( Ptr, Size, Flags & OL_MEMORY_REGISTER_FLAG_LOCK_MEMORY); if (!LockedPtrOrErr) return LockedPtrOrErr.takeError(); @@ -1456,14 +1594,20 @@ Error olMemRegister_impl(ol_device_handle_t Device, void *Ptr, size_t Size, Error olMemUnregister_impl(ol_device_handle_t Device, void *Ptr, ol_memory_register_flags_t Flags) { - return Device->Device->unregisterMemory( + auto DeviceOrErr = Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + return (*DeviceOrErr)->unregisterMemory( Ptr, Flags & OL_MEMORY_REGISTER_FLAG_UNLOCK_MEMORY); } Error olQueryQueue_impl(ol_queue_handle_t Queue, bool *IsQueueWorkCompleted) { if (Queue->AsyncInfo->Queue) { - if (auto Err = Queue->Device->Device->queryAsync(Queue->AsyncInfo, false, - IsQueueWorkCompleted)) + auto DeviceOrErr = Queue->Device->getDevice(); + if (!DeviceOrErr) + return DeviceOrErr.takeError(); + if (auto Err = (*DeviceOrErr)->queryAsync(Queue->AsyncInfo, false, + IsQueueWorkCompleted)) return Err; } else if (IsQueueWorkCompleted) { // No underlying queue means there's no work to complete. >From 5310f78f51e7d3f7ca593608ff487cbaa125f2b0 Mon Sep 17 00:00:00 2001 From: "Duran, Alex" <[email protected]> Date: Fri, 4 Sep 2026 08:59:52 -0700 Subject: [PATCH 2/2] don't initialize devices when validating the image --- offload/liboffload/src/OffloadImpl.cpp | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/offload/liboffload/src/OffloadImpl.cpp b/offload/liboffload/src/OffloadImpl.cpp index 670179c21f054..62f9878e7cd2b 100644 --- a/offload/liboffload/src/OffloadImpl.cpp +++ b/offload/liboffload/src/OffloadImpl.cpp @@ -1295,12 +1295,8 @@ Error olCreateProgram_impl(ol_context_handle_t Context, Error olIsValidBinary_impl(ol_device_handle_t Device, const void *ProgData, size_t ProgDataSize, bool *IsValid) { StringRef Buffer(reinterpret_cast<const char *>(ProgData), ProgDataSize); - auto DeviceOrErr = Device->getDevice(); - if (!DeviceOrErr) - return DeviceOrErr.takeError(); - auto *DeviceImpl = *DeviceOrErr; *IsValid = - DeviceImpl->Plugin.isDeviceCompatible(DeviceImpl->getDeviceId(), Buffer); + Device->Platform.Plugin->isDeviceCompatible(Device->DeviceNum, Buffer); return Error::success(); } _______________________________________________ llvm-branch-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/llvm-branch-commits
