[llvm] [offload] add handling of memory alignment to MemoryManagerTy (PR #218418)

via llvm-commits llvm-commits at lists.llvm.org
Mon Aug 24 07:04:13 PDT 2026


https://github.com/EuphoricThinking updated https://github.com/llvm/llvm-project/pull/218418

>From fe46d367fea54c4557fae682ed09fb7d65183c55 Mon Sep 17 00:00:00 2001
From: Agata Momot <agata.momot at intel.com>
Date: Thu, 6 Aug 2026 13:07:54 +0200
Subject: [PATCH] [offload] add handling of memory alignment to MemoryManagerTy

Previously, memory alignment was handled internally by the device.
The MemoryManagerTy either redirected large allocation requests directly
to the device or stored pointers to the smaller memory chunks in pools.
The information about the alignment expected by the user was just passed by
the MemoryManagerTy to the device allocators. However, most vendors do not
provide an API for specifying the alignment of the memory allocation
(except for the Level Zero plugin), allowing only for verification of whether
the returned pointer is correctly aligned. In this patch, MemoryManagerTy
allocates excess memory and performs pointer arithmetic on the pointer returned
by the device allocator, so that the final pointer is aligned accordingly
regardless of the plugin.

NodeTy has two new fields:
- `BasePtr`: the pointer to the originally allocated memory
- `Ptr`: the pointer returned to the caller, possibly after pointer arithmetic
operations on BasePtr for ensuring the requested alignment.

Pointer arithmetic is performed when a specific alignment is requested
or when the chosen pointer from the memory pool is being reused after
an aligned allocation, since Ptrs are not reset to BasePtrs during deallocation.

The Alignment parameter is preserved in the plugin-side implementations of
allocate() for validation purposes, since the device allocations might have
granularity which is not compatible with the requested alignment.

The presented changes are part of the series of patches introducing memory
alignment. The next patch will remove the doubling of the memory alignment
management in the MemoryManagerTy and the Level Zero plugin.
---
 offload/plugins-nextgen/amdgpu/src/rtl.cpp    | 14 +---
 .../common/include/MemoryManager.h            | 74 +++++++++++++++----
 offload/plugins-nextgen/cuda/src/rtl.cpp      | 14 +---
 3 files changed, 62 insertions(+), 40 deletions(-)

diff --git a/offload/plugins-nextgen/amdgpu/src/rtl.cpp b/offload/plugins-nextgen/amdgpu/src/rtl.cpp
index 2ba4d80978f4a..e4f4568f73574 100644
--- a/offload/plugins-nextgen/amdgpu/src/rtl.cpp
+++ b/offload/plugins-nextgen/amdgpu/src/rtl.cpp
@@ -339,7 +339,7 @@ struct AMDGPUMemoryPoolTy {
     // compared with the alignment of the memory allocated using the given pool.
     // If the default alignment is greater than or equal to the alignment
     // requested by the user, it would still meet the user's requirements.
-    if (Alignment > 0 && Alignment > PoolAllocationAlignment) {
+    if (Alignment > PoolAllocationAlignment) {
       return Plugin::error(ErrorCode::UNSUPPORTED,
                            "requested alignment (%lu) larger than maximum "
                            "supported pool alignment (%lu)",
@@ -349,18 +349,6 @@ struct AMDGPUMemoryPoolTy {
     hsa_status_t Status =
         hsa_amd_memory_pool_allocate(MemoryPool, Size, 0, PtrStorage);
 
-    if (Alignment > 0 && !isAddrAligned(Align(Alignment), *PtrStorage)) {
-      if (auto FreeErr = deallocate(*PtrStorage)) {
-        return Plugin::error(ErrorCode::UNKNOWN,
-                             "Failure in deallcation of the incorrectly "
-                             "aligned pointer; requested alignemnt: %lu",
-                             Alignment);
-      }
-
-      return Plugin::error(ErrorCode::UNSUPPORTED,
-                           "unsupported alignment size");
-    }
-
     return Plugin::check(Status, "error in hsa_amd_memory_pool_allocate: %s");
   }
 
diff --git a/offload/plugins-nextgen/common/include/MemoryManager.h b/offload/plugins-nextgen/common/include/MemoryManager.h
index 9712a461ad974..4d4f8168095d8 100644
--- a/offload/plugins-nextgen/common/include/MemoryManager.h
+++ b/offload/plugins-nextgen/common/include/MemoryManager.h
@@ -106,13 +106,17 @@ class MemoryManagerTy {
 
   /// A structure stores the meta data of a target pointer
   struct NodeTy {
-    /// Memory size
+    /// Final memory size, including the alignment
     const size_t Size;
-    /// Target pointer
+    /// Target pointer, returned to the caller after adjustments related to the
+    /// memory alignment (moving the pointer)
     void *Ptr;
+    /// Pointer to the originally allocated memory
+    void *BasePtr;
 
     /// Constructor
-    NodeTy(size_t Size, void *Ptr) : Size(Size), Ptr(Ptr) {}
+    NodeTy(size_t FinalSize, void *Ptr, void *BasePtr)
+        : Size(FinalSize), Ptr(Ptr), BasePtr(BasePtr) {}
   };
 
   /// To make \p NodePtrTy ordered when they're put into \p std::multiset.
@@ -171,7 +175,7 @@ class MemoryManagerTy {
       if (List.empty())
         continue;
       for (const NodeTy &N : List) {
-        if (auto Err = deleteOnDevice(N.Ptr))
+        if (auto Err = deleteOnDevice(N.BasePtr))
           return Err;
         RemoveList.push_back(N.Ptr);
       }
@@ -218,6 +222,19 @@ class MemoryManagerTy {
     return TgtPtr;
   }
 
+  void changeKeyPtr(void *FinalPtr, NodeTy *NodePtr) {
+    std::lock_guard<std::mutex> LG(MapTableLock);
+    auto NodeHandle = PtrToNodeTable.extract(NodePtr->Ptr);
+    NodeHandle.mapped().Ptr = FinalPtr;
+    NodeHandle.key() = FinalPtr;
+    PtrToNodeTable.insert(std::move(NodeHandle));
+  }
+
+  void *alignPointer(void *PtrToBeAligned, size_t Alignment) {
+    uintptr_t AlignedPointer = (uintptr_t)PtrToBeAligned;
+    return (void *)((AlignedPointer + Alignment - 1) & ~(Alignment - 1));
+  }
+
 public:
   static constexpr size_t DefaultSizeThreshold = 1U << 13;
 
@@ -234,8 +251,8 @@ class MemoryManagerTy {
   /// Destructor
   ~MemoryManagerTy() {
     for (auto &PtrToNode : PtrToNodeTable) {
-      assert(PtrToNode.second.Ptr && "nullptr in map table");
-      if (auto Err = deleteOnDevice(PtrToNode.second.Ptr))
+      assert(PtrToNode.second.BasePtr && "nullptr in map table");
+      if (auto Err = deleteOnDevice(PtrToNode.second.BasePtr))
         REPORT() << "Failure to delete memory: " << toString(std::move(Err));
     }
   }
@@ -264,9 +281,22 @@ class MemoryManagerTy {
       ODBG(OLDT_Alloc) << "Got target pointer " << *TgtPtrOrErr
                        << ". Return directly.";
 
+      if (Alignment > 0 && !isAddrAligned(Align(Alignment), *TgtPtrOrErr)) {
+        auto AlignErr = make_error<StringError>(
+            "Allocated address is misaligned", inconvertibleErrorCode());
+        if (auto FreeErr = deleteOnDevice(*TgtPtrOrErr)) {
+          return joinErrors(std::move(FreeErr), std::move(AlignErr));
+        }
+
+        return AlignErr;
+      }
+
       return *TgtPtrOrErr;
     }
 
+    if (Alignment > 0) {
+      Size += Alignment - 1;
+    }
     NodeTy *NodePtr = nullptr;
 
     // Try to get a node from FreeList
@@ -274,22 +304,32 @@ class MemoryManagerTy {
       const int B = findBucket(Size);
       FreeListTy &List = FreeLists[B];
 
-      NodeTy TempNode(Size, nullptr);
+      NodeTy TempNode(Size, nullptr, nullptr);
       std::lock_guard<std::mutex> LG(FreeListLocks[B]);
-      auto [First, Last] = List.equal_range(TempNode);
+      auto Itr = List.find(TempNode);
 
-      auto Itr = std::find_if(First, Last, [Alignment](const NodeTy &N) {
-        return Alignment == 0 || isAddrAligned(Align(Alignment), N.Ptr);
-      });
-      if (Itr != Last) {
+      if (Itr != List.end()) {
         NodePtr = &Itr->get();
         List.erase(Itr);
       }
     }
 
-    if (NodePtr != nullptr)
+    if (NodePtr != nullptr) {
       ODBG(OLDT_Alloc) << "Find one node " << NodePtr << " in the bucket.";
 
+      if (Alignment > 0) {
+        void *AlignedPointer = alignPointer(NodePtr->BasePtr, Alignment);
+
+        if (AlignedPointer != NodePtr->Ptr) {
+          changeKeyPtr(AlignedPointer, NodePtr);
+        }
+      } else {
+        if (NodePtr->Ptr != NodePtr->BasePtr) {
+          changeKeyPtr(NodePtr->BasePtr, NodePtr);
+        }
+      }
+    }
+
     // We cannot find a valid node in FreeLists. Let's allocate on device and
     // create a node for it.
     if (NodePtr == nullptr) {
@@ -305,10 +345,16 @@ class MemoryManagerTy {
       if (TgtPtr == nullptr)
         return nullptr;
 
+      void *BasePtr = TgtPtr;
+      if (Alignment > 0) {
+        TgtPtr = alignPointer(TgtPtr, Alignment);
+      }
+
       // Create a new node and add it into the map table
       {
         std::lock_guard<std::mutex> Guard(MapTableLock);
-        auto Itr = PtrToNodeTable.emplace(TgtPtr, NodeTy(Size, TgtPtr));
+        auto Itr =
+            PtrToNodeTable.emplace(TgtPtr, NodeTy(Size, TgtPtr, BasePtr));
         NodePtr = &Itr.first->second;
       }
 
diff --git a/offload/plugins-nextgen/cuda/src/rtl.cpp b/offload/plugins-nextgen/cuda/src/rtl.cpp
index 0f666ffb65e8f..cd3ac71c44397 100644
--- a/offload/plugins-nextgen/cuda/src/rtl.cpp
+++ b/offload/plugins-nextgen/cuda/src/rtl.cpp
@@ -590,7 +590,7 @@ struct CUDADeviceTy : public GenericDeviceTy {
     CUdeviceptr DevicePtr;
     CUresult Res;
 
-    if (Alignment > 0 && Alignment > Granularity) {
+    if (Alignment > Granularity) {
       return Plugin::error(ErrorCode::UNSUPPORTED,
                            "requested alignment (%lu) larger than maximum "
                            "supported alignment (%lu)",
@@ -615,18 +615,6 @@ struct CUDADeviceTy : public GenericDeviceTy {
     if (auto Err = Plugin::check(Res, "error in cuMemAlloc[Host|Managed]: %s"))
       return std::move(Err);
 
-    if (Alignment > 0 && !isAddrAligned(Align(Alignment), MemAlloc)) {
-      if (auto FreeErr = free(MemAlloc, Kind)) {
-        return Plugin::error(ErrorCode::UNKNOWN,
-                             "Failure in deallcation of the incorrectly "
-                             "aligned pointer; requested alignemnt: %lu",
-                             Alignment);
-      }
-
-      return Plugin::error(ErrorCode::UNSUPPORTED,
-                           "unsupported alignment size");
-    }
-
     return MemAlloc;
   }
 



More information about the llvm-commits mailing list