[llvm] [offload] add context to olMemAlloc* (PR #222677)

Alex Duran via llvm-commits llvm-commits at lists.llvm.org
Fri Sep 11 05:35:16 PDT 2026


================
@@ -701,176 +711,133 @@ TargetAllocTy convertOlToPluginAllocTy(ol_alloc_type_t Type) {
   }
 }
 
-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) {
-  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(
-        Size, nullptr, convertOlToPluginAllocTy(Type), Alignment);
-    if (!NewAlloc)
-      return NewAlloc.takeError();
-
-    void *NewEnd = &static_cast<char *>(*NewAlloc)[Size];
-    auto &AllocBases = OffloadContext::get().AllocBases;
-    auto &AllocInfoMap = OffloadContext::get().AllocInfoMap;
-    {
-      std::lock_guard<std::mutex> Lock(OffloadContext::get().AllocInfoMapMutex);
-
-      // Check that this memory region doesn't overlap another one
-      // That is, the start of this allocation needs to be after another
-      // allocation's end point, and the end of this allocation needs to be
-      // before the next one's start.
-      // `Gap` is the first alloc who ends after the new alloc's start point.
-      auto Gap =
-          std::lower_bound(AllocBases.begin(), AllocBases.end(), *NewAlloc,
-                           [&](const void *Iter, const void *Val) {
-                             return AllocInfoMap.at(Iter).End <= Val;
-                           });
-      if (Gap == AllocBases.end() || NewEnd <= AllocInfoMap.at(*Gap).Start) {
-        // Success, no conflict
-        AllocInfoMap.insert_or_assign(
-            *NewAlloc, AllocInfo{Device, Type, *NewAlloc, NewEnd});
-        AllocBases.insert(
-            std::lower_bound(AllocBases.begin(), AllocBases.end(), *NewAlloc),
-            *NewAlloc);
-        *AllocationOut = *NewAlloc;
-
-        for (void *R : Rejects)
-          if (auto Err =
-                  Device->Device->dataDelete(R, convertOlToPluginAllocTy(Type)))
-            return Err;
-        return Error::success();
-      }
-
-      // To avoid the next attempt allocating the same memory we just freed, we
-      // hold onto it until we complete the allocation
-      Rejects.push_back(*NewAlloc);
-    }
+ol_alloc_type_t convertPluginToOlAllocTy(TargetAllocTy Kind) {
+  switch (Kind) {
+  case TARGET_ALLOC_HOST:
+    return OL_ALLOC_TYPE_HOST;
+  case TARGET_ALLOC_SHARED:
+    return OL_ALLOC_TYPE_MANAGED;
+  case TARGET_ALLOC_DEVICE:
+  case TARGET_ALLOC_DEFAULT:
+    return OL_ALLOC_TYPE_DEVICE;
   }
-
-  // We've tried multiple times, and can't allocate a non-overlapping region.
-  return createOffloadError(ErrorCode::BACKEND_FAILURE,
-                            "failed to allocate non-overlapping memory");
+  llvm_unreachable("unhandled TargetAllocTy");
 }
 
-Error olMemAlloc_impl(ol_device_handle_t Device, ol_alloc_type_t Type,
-                      size_t Size, void **AllocationOut) {
+Error olMemAlloc_impl(ol_context_handle_t Context, ol_device_handle_t Device,
+                      ol_alloc_type_t Type, size_t Size, void **AllocationOut) {
   if (Type == OL_ALLOC_TYPE_HOST)
     return createOffloadError(ErrorCode::INVALID_ENUMERATION,
                               "use olMemAllocHost for host allocations");
-  return olMemAllocImplHelper(Device, Type, Size, /*Alignment=*/0,
-                              AllocationOut);
+  auto AllocOrErr =
+      Context->allocate(Device, static_cast<int64_t>(Size),
+                        convertOlToPluginAllocTy(Type), /*Alignment=*/0);
+  if (!AllocOrErr)
+    return AllocOrErr.takeError();
+  *AllocationOut = *AllocOrErr;
+  return Error::success();
 }
 
-Error olMemAllocHost_impl(ol_device_handle_t Device, size_t Size,
+Error olMemAllocHost_impl(ol_context_handle_t Context,
+                          ol_device_handle_t Device, size_t Size,
                           void **AllocationOut) {
-  return olMemAllocImplHelper(Device, OL_ALLOC_TYPE_HOST, Size,
-                              /*Alignment=*/0, AllocationOut);
+  auto AllocOrErr = Context->allocate(Device, static_cast<int64_t>(Size),
+                                      TARGET_ALLOC_HOST, /*Alignment=*/0);
+  if (!AllocOrErr)
+    return AllocOrErr.takeError();
+  *AllocationOut = *AllocOrErr;
+  return Error::success();
 }
 
-Error olMemAllocAligned_impl(ol_device_handle_t Device, ol_alloc_type_t Type,
+Error olMemAllocAligned_impl(ol_context_handle_t Context,
+                             ol_device_handle_t Device, ol_alloc_type_t Type,
                              size_t Size, size_t Alignment,
                              void **AllocationOut) {
   if (Type == OL_ALLOC_TYPE_HOST)
     return createOffloadError(ErrorCode::INVALID_ENUMERATION,
                               "use olMemAllocAlignedHost for host allocations");
-  return olMemAllocImplHelper(Device, Type, Size, Alignment, AllocationOut);
+  auto AllocOrErr =
+      Context->allocate(Device, static_cast<int64_t>(Size),
+                        convertOlToPluginAllocTy(Type), Alignment);
+  if (!AllocOrErr)
+    return AllocOrErr.takeError();
+  *AllocationOut = *AllocOrErr;
+  return Error::success();
 }
 
-Error olMemAllocAlignedHost_impl(ol_device_handle_t Device, size_t Size,
+Error olMemAllocAlignedHost_impl(ol_context_handle_t Context,
+                                 ol_device_handle_t Device, size_t Size,
                                  size_t Alignment, void **AllocationOut) {
-  return olMemAllocImplHelper(Device, OL_ALLOC_TYPE_HOST, Size, Alignment,
-                              AllocationOut);
+  auto AllocOrErr = Context->allocate(Device, static_cast<int64_t>(Size),
+                                      TARGET_ALLOC_HOST, Alignment);
+  if (!AllocOrErr)
+    return AllocOrErr.takeError();
+  *AllocationOut = *AllocOrErr;
+  return Error::success();
 }
 
-Error olMemFree_impl(void *Address) {
-  ol_device_handle_t Device;
-  ol_alloc_type_t Type;
-  {
-    std::lock_guard<std::mutex> Lock(OffloadContext::get().AllocInfoMapMutex);
-    if (!OffloadContext::get().AllocInfoMap.contains(Address))
-      return createOffloadError(ErrorCode::INVALID_ARGUMENT,
-                                "address is not a known allocation");
-
-    auto AllocInfo = OffloadContext::get().AllocInfoMap.at(Address);
-    Device = AllocInfo.Device;
-    Type = AllocInfo.Type;
-    OffloadContext::get().AllocInfoMap.erase(Address);
-
-    auto &Bases = OffloadContext::get().AllocBases;
-    Bases.erase(std::lower_bound(Bases.begin(), Bases.end(), Address));
-  }
-
-  if (auto Res =
-          Device->Device->dataDelete(Address, convertOlToPluginAllocTy(Type)))
-    return Res;
-
-  return Error::success();
+Error olMemFree_impl(ol_context_handle_t Context, void *Address) {
+  return Context->deallocate(Address);
----------------
adurang wrote:

How is this supposed to work without knowing the allocation type as before?

https://github.com/llvm/llvm-project/pull/222677


More information about the llvm-commits mailing list