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

Łukasz Plewa via llvm-commits llvm-commits at lists.llvm.org
Tue Sep 22 05:38:25 PDT 2026


https://github.com/lplewa updated https://github.com/llvm/llvm-project/pull/224930

>From 60b5a048bb92964ebe59871c572440eb150508cb Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?=C5=81ukasz=20Plewa?= <lukasz.plewa at intel.com>
Date: Fri, 18 Sep 2026 13:33:45 +0200
Subject: [PATCH 1/3] [offload] add context to olMemAlloc* (#222677)

This is the last part of the context refactor. This patch:
- Adds context param to all memory allocation functions.
- Routes ptr info through the plugins instead of a global map.
- Removes allocation tracking from liboffload.
- olGetMemInfo(OL_MEM_INFO_DEVICE) now returns INVALID_ARGUMENT for host
allocations, which have no per-device affinity.

Assisted-by: Claude Opus 4.7
---
 libsycl/src/detail/queue_impl.cpp             |  18 +-
 libsycl/src/usm_functions.cpp                 |  17 +-
 libsycl/unittests/handler/test_helpers.hpp    |   8 +-
 libsycl/unittests/mock/helpers.cpp            |  38 ++-
 libsycl/unittests/mock/helpers.hpp            |  23 +-
 libsycl/unittests/mock/mock.cpp               |  38 +--
 libsycl/unittests/queue/memcpy.cpp            |  10 +-
 libsycl/unittests/usm/alloc.cpp               | 105 +++----
 .../tools/llvm-gpu-loader/llvm-gpu-loader.cpp |  22 +-
 llvm/tools/llvm-gpu-loader/llvm-gpu-loader.h  |   8 +-
 .../languages/kernel/src/LanguageRuntime.cpp  |  13 +-
 offload/liboffload/API/Memory.td              |  39 ++-
 offload/liboffload/src/OffloadImpl.cpp        | 261 ++++++++----------
 .../amdgpu/dynamic_hsa/hsa_ext_amd.h          |   2 +
 offload/plugins-nextgen/amdgpu/src/rtl.cpp    |  70 +++++
 .../common/include/PluginInterface.h          | 105 +++++--
 .../common/src/PluginInterface.cpp            | 195 ++++++++-----
 .../cuda/dynamic_cuda/cuda.cpp                |   1 +
 .../plugins-nextgen/cuda/dynamic_cuda/cuda.h  |  18 +-
 offload/plugins-nextgen/cuda/src/rtl.cpp      |  55 ++++
 offload/plugins-nextgen/host/src/rtl.cpp      |  49 ++++
 .../level_zero/include/L0Memory.h             |  29 +-
 .../level_zero/include/L0Plugin.h             |  16 ++
 .../level_zero/src/L0Context.cpp              |   3 +-
 .../level_zero/src/L0Device.cpp               |   3 +-
 .../level_zero/src/L0Memory.cpp               |  14 +-
 .../level_zero/src/L0Plugin.cpp               | 102 ++++++-
 .../include/mathtest/DeviceContext.hpp        |   8 +-
 .../include/mathtest/DeviceResources.hpp      |  15 +-
 .../Conformance/lib/DeviceContext.cpp         |   7 +-
 .../Conformance/lib/DeviceResources.cpp       |   5 +-
 .../event/olGetEventElapsedTime.cpp           |   4 +-
 .../OffloadAPI/kernel/olLaunchKernel.cpp      |  74 ++---
 .../OffloadAPI/memory/olGetMemInfo.cpp        |  71 +++--
 .../OffloadAPI/memory/olGetMemInfoSize.cpp    |  16 +-
 .../OffloadAPI/memory/olMemAlloc.cpp          |  29 +-
 .../OffloadAPI/memory/olMemAllocAligned.cpp   |  41 +--
 .../unittests/OffloadAPI/memory/olMemFill.cpp |  50 ++--
 .../unittests/OffloadAPI/memory/olMemFree.cpp |  13 +-
 .../OffloadAPI/memory/olMemPrefetch.cpp       |  35 ++-
 .../unittests/OffloadAPI/memory/olMemcpy.cpp  |  71 ++---
 .../OffloadAPI/queue/olLaunchHostFunction.cpp |   4 +-
 .../OffloadAPI/queue/olWaitEvents.cpp         |   6 +-
 43 files changed, 1132 insertions(+), 579 deletions(-)

diff --git a/libsycl/src/detail/queue_impl.cpp b/libsycl/src/detail/queue_impl.cpp
index 9d9ccb34f625b..d410bc2c7c795 100644
--- a/libsycl/src/detail/queue_impl.cpp
+++ b/libsycl/src/detail/queue_impl.cpp
@@ -145,15 +145,19 @@ void QueueImpl::submitKernelImpl(DeviceKernelInfo &KernelInfo, void *ArgData,
       createEvent(std::move(MCurrentSubmitInfo.DepEvents));
 }
 
-static ol_device_handle_t getAllocDevice(const void *ptr) {
+static ol_device_handle_t getAllocDevice(ol_context_handle_t Context,
+                                         const void *ptr) {
   // TODO: consider caching this information to avoid querying it every time.
   ol_device_handle_t Device{};
   [[maybe_unused]] ol_result_t Result =
-      callNoCheck(olGetMemInfo, ptr, OL_MEM_INFO_DEVICE,
+      callNoCheck(olGetMemInfo, Context, ptr, OL_MEM_INFO_DEVICE,
                   sizeof(ol_device_handle_t), &Device);
   if (detail::isFailed(Result)) {
-    // If liboffload could not find the allocation, assume it is a host one.
-    if (Result->Code == OL_ERRC_NOT_FOUND) {
+    // NOT_FOUND: the pointer isn't a liboffload allocation at all (plain host
+    // malloc). INVALID_ARGUMENT: it's a liboffload host allocation, which has
+    // no per-device affinity. Either way, route through the host device.
+    if (Result->Code == OL_ERRC_NOT_FOUND ||
+        Result->Code == OL_ERRC_INVALID_ARGUMENT) {
       return getHostOLDevice();
     }
     checkAndThrow(Result);
@@ -174,8 +178,10 @@ QueueImpl::memcpy(void *Dest, const void *Src, std::size_t NumBytes,
     throw sycl::exception(sycl::make_error_code(sycl::errc::invalid),
                           "Nullptr argument in memcpy operation");
 
-  ol_device_handle_t DestOLDevice = getAllocDevice(Dest);
-  ol_device_handle_t SrcOLDevice = getAllocDevice(Src);
+  ol_device_handle_t DestOLDevice =
+      getAllocDevice(MContext->getOLHandleRef(), Dest);
+  ol_device_handle_t SrcOLDevice =
+      getAllocDevice(MContext->getOLHandleRef(), Src);
 
   handleEventDependencies(DepEvents);
   callAndThrow(olMemcpy, MOffloadQueue, Dest, DestOLDevice, Src, SrcOLDevice,
diff --git a/libsycl/src/usm_functions.cpp b/libsycl/src/usm_functions.cpp
index 6354b61fbf976..699fdeec79e22 100644
--- a/libsycl/src/usm_functions.cpp
+++ b/libsycl/src/usm_functions.cpp
@@ -8,6 +8,7 @@
 
 #include <sycl/__impl/usm_functions.hpp>
 
+#include <detail/context_impl.hpp>
 #include <detail/device_impl.hpp>
 #include <detail/offload/offload_utils.hpp>
 
@@ -153,19 +154,21 @@ void *aligned_alloc(std::size_t alignment, std::size_t numBytes,
 
   void *Ptr{};
   auto OLDevice = detail::getSyclObjImpl(syclDevice)->getOLHandle();
+  auto OLContext = detail::getSyclObjImpl(syclContext)->getOLHandleRef();
 
   ol_result_t Result{};
   if (alignment == 0) {
     Result =
         kind == usm::alloc::host
-            ? detail::callNoCheck(olMemAllocHost, OLDevice, numBytes, &Ptr)
-            : detail::callNoCheck(olMemAlloc, OLDevice,
+            ? detail::callNoCheck(olMemAllocHost, OLContext, OLDevice, numBytes,
+                                  &Ptr)
+            : detail::callNoCheck(olMemAlloc, OLContext, OLDevice,
                                   detail::getOlAllocType(kind), numBytes, &Ptr);
   } else {
     Result = kind == usm::alloc::host
-                 ? detail::callNoCheck(olMemAllocAlignedHost, OLDevice,
-                                       numBytes, alignment, &Ptr)
-                 : detail::callNoCheck(olMemAllocAligned, OLDevice,
+                 ? detail::callNoCheck(olMemAllocAlignedHost, OLContext,
+                                       OLDevice, numBytes, alignment, &Ptr)
+                 : detail::callNoCheck(olMemAllocAligned, OLContext, OLDevice,
                                        detail::getOlAllocType(kind), numBytes,
                                        alignment, &Ptr);
   }
@@ -194,8 +197,8 @@ void *malloc(std::size_t numBytes, const queue &syclQueue, usm::alloc kind,
 // SYCL 2020 4.8.3.6. Memory deallocation functions.
 
 void free(void *ptr, const context &ctxt) {
-  std::ignore = ctxt;
-  detail::callAndThrow(olMemFree, ptr);
+  auto OLContext = detail::getSyclObjImpl(ctxt)->getOLHandleRef();
+  detail::callAndThrow(olMemFree, OLContext, ptr);
 }
 
 void free(void *ptr, const queue &q) { return free(ptr, q.get_context()); }
diff --git a/libsycl/unittests/handler/test_helpers.hpp b/libsycl/unittests/handler/test_helpers.hpp
index ad43374df79cd..1c383e747390e 100644
--- a/libsycl/unittests/handler/test_helpers.hpp
+++ b/libsycl/unittests/handler/test_helpers.hpp
@@ -18,12 +18,12 @@ inline void expectDeviceMemoryInfo(mock::MockWrapper &Mock,
                                    const std::vector<const void *> ExpectedPtrs,
                                    ol_device_handle_t Device, int Count) {
   EXPECT_CALL(Mock.get(),
-              olGetMemInfo(::testing::_, OL_MEM_INFO_DEVICE,
+              olGetMemInfo(::testing::_, ::testing::_, OL_MEM_INFO_DEVICE,
                            sizeof(ol_device_handle_t), ::testing::_))
       .Times(Count)
-      .WillRepeatedly([ExpectedPtrs, Device](const void *Ptr, ol_mem_info_t,
-                                             size_t,
-                                             void *PropValue) -> ol_result_t {
+      .WillRepeatedly([ExpectedPtrs, Device](
+                          ol_context_handle_t, const void *Ptr, ol_mem_info_t,
+                          size_t, void *PropValue) -> ol_result_t {
         EXPECT_NE(std::find(ExpectedPtrs.begin(), ExpectedPtrs.end(), Ptr),
                   ExpectedPtrs.end());
         *(static_cast<ol_device_handle_t *>(PropValue)) = Device;
diff --git a/libsycl/unittests/mock/helpers.cpp b/libsycl/unittests/mock/helpers.cpp
index f1fcc3fcf7eb4..bf0baf51b1917 100644
--- a/libsycl/unittests/mock/helpers.cpp
+++ b/libsycl/unittests/mock/helpers.cpp
@@ -343,8 +343,10 @@ void mock::MockLiboffload::initDefault() {
       });
 
   ON_CALL(*this, olGetMemInfo)
-      .WillByDefault([this](const void *Ptr, ol_mem_info_t PropName,
-                            size_t PropSize, void *PropValue) -> ol_result_t {
+      .WillByDefault([this](ol_context_handle_t Context, const void *Ptr,
+                            ol_mem_info_t PropName, size_t PropSize,
+                            void *PropValue) -> ol_result_t {
+        EXPECT_NE(Context, nullptr);
         EXPECT_NE(Ptr, nullptr);
         // Other properties are not used by the runtime yet
         EXPECT_EQ(PropName, OL_MEM_INFO_DEVICE);
@@ -361,8 +363,10 @@ void mock::MockLiboffload::initDefault() {
       });
 
   ON_CALL(*this, olMemAlloc)
-      .WillByDefault([](ol_device_handle_t Device, ol_alloc_type_t Type,
-                        size_t Size, void **AllocationOut) -> ol_result_t {
+      .WillByDefault([](ol_context_handle_t Context, ol_device_handle_t Device,
+                        ol_alloc_type_t Type, size_t Size,
+                        void **AllocationOut) -> ol_result_t {
+        EXPECT_NE(Context, nullptr);
         EXPECT_NE(Device, nullptr);
         EXPECT_NE(Type, OL_ALLOC_TYPE_HOST);
         EXPECT_GT(Size, 0);
@@ -372,8 +376,9 @@ void mock::MockLiboffload::initDefault() {
       });
 
   ON_CALL(*this, olMemAllocHost)
-      .WillByDefault([](ol_device_handle_t Device, size_t Size,
-                        void **AllocationOut) -> ol_result_t {
+      .WillByDefault([](ol_context_handle_t Context, ol_device_handle_t Device,
+                        size_t Size, void **AllocationOut) -> ol_result_t {
+        EXPECT_NE(Context, nullptr);
         EXPECT_NE(Device, nullptr);
         EXPECT_GT(Size, 0);
         EXPECT_NE(AllocationOut, nullptr);
@@ -381,17 +386,22 @@ void mock::MockLiboffload::initDefault() {
         return OL_SUCCESS;
       });
 
-  ON_CALL(*this, olMemFree).WillByDefault([](void *Address) -> ol_result_t {
-    EXPECT_NE(Address, nullptr);
-    mock::releaseDummyHandle(Address);
-    return OL_SUCCESS;
-  });
+  ON_CALL(*this, olMemFree)
+      .WillByDefault(
+          [](ol_context_handle_t Context, void *Address) -> ol_result_t {
+            EXPECT_NE(Context, nullptr);
+            EXPECT_NE(Address, nullptr);
+            mock::releaseDummyHandle(Address);
+            return OL_SUCCESS;
+          });
 
   ON_CALL(*this, olMemAllocAligned)
-      .WillByDefault([this](ol_device_handle_t Device,
+      .WillByDefault([this](ol_context_handle_t Context,
+                            ol_device_handle_t Device,
                             ol_alloc_type_t AllocType, size_t Size,
                             size_t Alignment,
                             void **AllocationOut) -> ol_result_t {
+        EXPECT_NE(Context, nullptr);
         EXPECT_NE(Device, nullptr);
         EXPECT_TRUE(AllocType == OL_ALLOC_TYPE_DEVICE ||
                     AllocType == OL_ALLOC_TYPE_MANAGED);
@@ -407,9 +417,11 @@ void mock::MockLiboffload::initDefault() {
       });
 
   ON_CALL(*this, olMemAllocAlignedHost)
-      .WillByDefault([this](ol_device_handle_t Device, size_t Size,
+      .WillByDefault([this](ol_context_handle_t Context,
+                            ol_device_handle_t Device, size_t Size,
                             size_t Alignment,
                             void **AllocationOut) -> ol_result_t {
+        EXPECT_NE(Context, nullptr);
         EXPECT_NE(Device, nullptr);
         EXPECT_GT(Size, 0);
         EXPECT_GT(Alignment, 0);
diff --git a/libsycl/unittests/mock/helpers.hpp b/libsycl/unittests/mock/helpers.hpp
index d70a0649d790f..b814df4a55c4b 100644
--- a/libsycl/unittests/mock/helpers.hpp
+++ b/libsycl/unittests/mock/helpers.hpp
@@ -134,20 +134,23 @@ class MockLiboffload {
               (ol_queue_handle_t Queue, size_t Count, const void **Mems,
                const size_t *Sizes, ol_mem_migration_flags_t Flags));
   MOCK_METHOD(ol_result_t, olGetMemInfo,
-              (const void *Ptr, ol_mem_info_t PropName, size_t PropSize,
-               void *PropValue));
+              (ol_context_handle_t Context, const void *Ptr,
+               ol_mem_info_t PropName, size_t PropSize, void *PropValue));
   MOCK_METHOD(ol_result_t, olMemAlloc,
-              (ol_device_handle_t Device, ol_alloc_type_t Type, size_t Size,
-               void **AllocationOut));
+              (ol_context_handle_t Context, ol_device_handle_t Device,
+               ol_alloc_type_t Type, size_t Size, void **AllocationOut));
   MOCK_METHOD(ol_result_t, olMemAllocHost,
-              (ol_device_handle_t Device, size_t Size, void **AllocationOut));
-  MOCK_METHOD(ol_result_t, olMemFree, (void *Address));
+              (ol_context_handle_t Context, ol_device_handle_t Device,
+               size_t Size, void **AllocationOut));
+  MOCK_METHOD(ol_result_t, olMemFree,
+              (ol_context_handle_t Context, void *Address));
   MOCK_METHOD(ol_result_t, olMemAllocAligned,
-              (ol_device_handle_t Device, ol_alloc_type_t AllocType,
-               size_t Size, size_t Alignment, void **AllocationOut));
-  MOCK_METHOD(ol_result_t, olMemAllocAlignedHost,
-              (ol_device_handle_t Device, size_t Size, size_t Alignment,
+              (ol_context_handle_t Context, ol_device_handle_t Device,
+               ol_alloc_type_t AllocType, size_t Size, size_t Alignment,
                void **AllocationOut));
+  MOCK_METHOD(ol_result_t, olMemAllocAlignedHost,
+              (ol_context_handle_t Context, ol_device_handle_t Device,
+               size_t Size, size_t Alignment, void **AllocationOut));
 
   ol_result_t makeEmptyStrError(ol_errc_t Code) {
     auto [Iterator, Flag] =
diff --git a/libsycl/unittests/mock/mock.cpp b/libsycl/unittests/mock/mock.cpp
index eb3bbcb7bacba..e8c26f922cdfc 100644
--- a/libsycl/unittests/mock/mock.cpp
+++ b/libsycl/unittests/mock/mock.cpp
@@ -131,25 +131,29 @@ ol_result_t olMemPrefetch(ol_queue_handle_t Queue, size_t Count,
                                                  Flags);
 }
 
-ol_result_t olGetMemInfo(const void *Ptr, ol_mem_info_t PropName,
-                         size_t PropSize, void *PropValue) {
-  return mock::getMockLiboffload().olGetMemInfo(Ptr, PropName, PropSize,
-                                                PropValue);
+ol_result_t olGetMemInfo(ol_context_handle_t Context, const void *Ptr,
+                         ol_mem_info_t PropName, size_t PropSize,
+                         void *PropValue) {
+  return mock::getMockLiboffload().olGetMemInfo(Context, Ptr, PropName,
+                                                PropSize, PropValue);
 }
 
-ol_result_t olMemAlloc(ol_device_handle_t Device, ol_alloc_type_t Type,
-                       size_t Size, void **AllocationOut) {
-  return mock::getMockLiboffload().olMemAlloc(Device, Type, Size,
+ol_result_t olMemAlloc(ol_context_handle_t Context, ol_device_handle_t Device,
+                       ol_alloc_type_t Type, size_t Size,
+                       void **AllocationOut) {
+  return mock::getMockLiboffload().olMemAlloc(Context, Device, Type, Size,
                                               AllocationOut);
 }
 
-ol_result_t olMemAllocHost(ol_device_handle_t Device, size_t Size,
+ol_result_t olMemAllocHost(ol_context_handle_t Context,
+                           ol_device_handle_t Device, size_t Size,
                            void **AllocationOut) {
-  return mock::getMockLiboffload().olMemAllocHost(Device, Size, AllocationOut);
+  return mock::getMockLiboffload().olMemAllocHost(Context, Device, Size,
+                                                  AllocationOut);
 }
 
-ol_result_t olMemFree(void *Address) {
-  return mock::getMockLiboffload().olMemFree(Address);
+ol_result_t olMemFree(ol_context_handle_t Context, void *Address) {
+  return mock::getMockLiboffload().olMemFree(Context, Address);
 }
 
 ol_result_t olCreateEvent(ol_queue_handle_t Queue, ol_event_flags_t Flags,
@@ -161,15 +165,17 @@ ol_result_t olDestroyEvent(ol_event_handle_t Event) {
   return mock::getMockLiboffload().olDestroyEvent(Event);
 }
 
-ol_result_t olMemAllocAligned(ol_device_handle_t Device,
+ol_result_t olMemAllocAligned(ol_context_handle_t Context,
+                              ol_device_handle_t Device,
                               ol_alloc_type_t AllocType, size_t Size,
                               size_t Alignment, void **OutPtr) {
-  return mock::getMockLiboffload().olMemAllocAligned(Device, AllocType, Size,
-                                                     Alignment, OutPtr);
+  return mock::getMockLiboffload().olMemAllocAligned(Context, Device, AllocType,
+                                                     Size, Alignment, OutPtr);
 }
 
-ol_result_t olMemAllocAlignedHost(ol_device_handle_t Device, size_t Size,
+ol_result_t olMemAllocAlignedHost(ol_context_handle_t Context,
+                                  ol_device_handle_t Device, size_t Size,
                                   size_t Alignment, void **OutPtr) {
-  return mock::getMockLiboffload().olMemAllocAlignedHost(Device, Size,
+  return mock::getMockLiboffload().olMemAllocAlignedHost(Context, Device, Size,
                                                          Alignment, OutPtr);
 }
diff --git a/libsycl/unittests/queue/memcpy.cpp b/libsycl/unittests/queue/memcpy.cpp
index db977f79d5bef..8184e5808e06e 100644
--- a/libsycl/unittests/queue/memcpy.cpp
+++ b/libsycl/unittests/queue/memcpy.cpp
@@ -26,11 +26,13 @@ TEST(Queue, Memcpy) {
   ol_device_handle_t OLDev =
       detail::getSyclObjImpl(Q.get_device())->getOLHandle();
 
-  EXPECT_CALL(Mock.get(), olGetMemInfo(_, OL_MEM_INFO_DEVICE,
+  EXPECT_CALL(Mock.get(), olGetMemInfo(_, _, OL_MEM_INFO_DEVICE,
                                        sizeof(ol_device_handle_t), _))
       .Times(NMemcpies * 2)
-      .WillRepeatedly([&](const void *Ptr, ol_mem_info_t PropName,
-                          size_t PropSize, void *PropValue) -> ol_result_t {
+      .WillRepeatedly([&](ol_context_handle_t Context, const void *Ptr,
+                          ol_mem_info_t PropName, size_t PropSize,
+                          void *PropValue) -> ol_result_t {
+        EXPECT_NE(Context, nullptr);
         EXPECT_TRUE(Ptr == SrcPtr || Ptr == DstPtr);
         bool IsHostPtr = Ptr == SrcPtr ? IsSrcHostPtr : IsDstHostPtr;
         if (IsHostPtr)
@@ -73,7 +75,7 @@ TEST(Queue, MemcpyZeroBytes) {
   mock::MockWrapper Mock;
   queue Q;
   EXPECT_CALL(Mock.get(), olWaitEvents(_, _, 1)).Times(1);
-  EXPECT_CALL(Mock.get(), olGetMemInfo(_, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olGetMemInfo(_, _, _, _, _)).Times(0);
   EXPECT_CALL(Mock.get(), olMemcpy(_, _, _, _, _, _)).Times(0);
   event Event = Q.memcpy(nullptr, nullptr, 0);
   Q.memcpy(nullptr, nullptr, 0, Event);
diff --git a/libsycl/unittests/usm/alloc.cpp b/libsycl/unittests/usm/alloc.cpp
index 6773e86236775..1e80c336d18c6 100644
--- a/libsycl/unittests/usm/alloc.cpp
+++ b/libsycl/unittests/usm/alloc.cpp
@@ -32,23 +32,24 @@ TEST(USMFunctions, DeviceAllocation) {
   context Ctx = Q.get_context();
   ol_device_handle_t OLDev = detail::getSyclObjImpl(Dev)->getOLHandle();
 
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, _, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAlloc(OLDev, OL_ALLOC_TYPE_DEVICE, NumBytes, _))
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, _, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(),
+              olMemAlloc(_, OLDev, OL_ALLOC_TYPE_DEVICE, NumBytes, _))
       .Times(1);
   void *Ptr1 = malloc_device(NumBytes, Dev, Ctx);
   EXPECT_NE(Ptr1, nullptr);
 
-  EXPECT_CALL(Mock.get(), olMemFree(Ptr1)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, Ptr1)).Times(1);
   free(Ptr1, Ctx);
 
-  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(OLDev, OL_ALLOC_TYPE_DEVICE,
+  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, OLDev, OL_ALLOC_TYPE_DEVICE,
                                             NumBytes, Alignment, _))
       .Times(1);
   void *Ptr2 = aligned_alloc_device(Alignment, NumBytes, Q);
   EXPECT_NE(Ptr2, nullptr);
 
-  EXPECT_CALL(Mock.get(), olMemFree(Ptr2)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, Ptr2)).Times(1);
   free(Ptr2, Q);
 }
 
@@ -59,21 +60,22 @@ TEST(USMFunctions, HostAllocation) {
   ol_device_handle_t OLDev =
       detail::getSyclObjImpl(Q.get_device())->getOLHandle();
 
-  EXPECT_CALL(Mock.get(), olMemAllocAlignedHost(_, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocHost(OLDev, NumBytes, _)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemAllocAlignedHost(_, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocHost(_, OLDev, NumBytes, _)).Times(1);
   void *Ptr1 = malloc_host(NumBytes, Ctx);
   EXPECT_NE(Ptr1, nullptr);
 
-  EXPECT_CALL(Mock.get(), olMemFree(Ptr1)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, Ptr1)).Times(1);
   free(Ptr1, Ctx);
 
-  EXPECT_CALL(Mock.get(), olMemAllocHost(_, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocAlignedHost(OLDev, NumBytes, Alignment, _))
+  EXPECT_CALL(Mock.get(), olMemAllocHost(_, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(),
+              olMemAllocAlignedHost(_, OLDev, NumBytes, Alignment, _))
       .Times(1);
   void *Ptr2 = aligned_alloc_host(Alignment, NumBytes, Ctx);
   EXPECT_NE(Ptr2, nullptr);
 
-  EXPECT_CALL(Mock.get(), olMemFree(Ptr2)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, Ptr2)).Times(1);
   free(Ptr2, Ctx);
 }
 
@@ -84,23 +86,24 @@ TEST(USMFunctions, SharedAllocation) {
   context Ctx = Q.get_context();
   ol_device_handle_t OLDev = detail::getSyclObjImpl(Dev)->getOLHandle();
 
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, _, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAlloc(OLDev, OL_ALLOC_TYPE_MANAGED, NumBytes, _))
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, _, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(),
+              olMemAlloc(_, OLDev, OL_ALLOC_TYPE_MANAGED, NumBytes, _))
       .Times(1);
   void *Ptr1 = malloc_shared(NumBytes, Dev, Ctx);
   EXPECT_NE(Ptr1, nullptr);
 
-  EXPECT_CALL(Mock.get(), olMemFree(Ptr1)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, Ptr1)).Times(1);
   free(Ptr1, Ctx);
 
-  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(OLDev, OL_ALLOC_TYPE_MANAGED,
+  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, OLDev, OL_ALLOC_TYPE_MANAGED,
                                             NumBytes, Alignment, _))
       .Times(1);
   void *Ptr2 = aligned_alloc_shared(Alignment, NumBytes, Q);
   EXPECT_NE(Ptr2, nullptr);
 
-  EXPECT_CALL(Mock.get(), olMemFree(Ptr2)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, Ptr2)).Times(1);
   free(Ptr2, Q);
 }
 
@@ -110,10 +113,10 @@ TEST(USMFunctions, ZeroByteAllocation) {
   device Dev = Q.get_device();
   context Ctx = Q.get_context();
 
-  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocHost(_, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, _, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocAlignedHost(_, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocHost(_, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, _, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocAlignedHost(_, _, _, _, _)).Times(0);
 
   EXPECT_EQ(malloc_device(0, Dev, Ctx), nullptr);
   EXPECT_EQ(malloc_shared(0, Dev, Ctx), nullptr);
@@ -129,21 +132,21 @@ TEST(USMFunctions, InvalidAlignment) {
 
   constexpr size_t NonPowerOf2Alignment = 3;
 
-  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocHost(_, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocHost(_, _, _, _)).Times(0);
 
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(OLDev, OL_ALLOC_TYPE_DEVICE,
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, OLDev, OL_ALLOC_TYPE_DEVICE,
                                             NumBytes, NonPowerOf2Alignment, _))
       .Times(1);
   EXPECT_EQ(aligned_alloc_device(NonPowerOf2Alignment, NumBytes, Dev, Ctx),
             nullptr);
 
-  EXPECT_CALL(Mock.get(),
-              olMemAllocAlignedHost(OLDev, NumBytes, NonPowerOf2Alignment, _))
+  EXPECT_CALL(Mock.get(), olMemAllocAlignedHost(_, OLDev, NumBytes,
+                                                NonPowerOf2Alignment, _))
       .Times(1);
   EXPECT_EQ(aligned_alloc_host(NonPowerOf2Alignment, NumBytes, Ctx), nullptr);
 
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(OLDev, OL_ALLOC_TYPE_MANAGED,
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, OLDev, OL_ALLOC_TYPE_MANAGED,
                                             NumBytes, NonPowerOf2Alignment, _))
       .Times(1);
   EXPECT_EQ(aligned_alloc_shared(NonPowerOf2Alignment, NumBytes, Dev, Ctx),
@@ -159,27 +162,29 @@ TEST(USMFunctions, ZeroAlignmentSucceeds) {
 
   constexpr size_t ZeroAlignment = 0;
 
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, _, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocAlignedHost(_, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, _, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocAlignedHost(_, _, _, _, _)).Times(0);
 
-  EXPECT_CALL(Mock.get(), olMemAlloc(OLDev, OL_ALLOC_TYPE_DEVICE, NumBytes, _))
+  EXPECT_CALL(Mock.get(),
+              olMemAlloc(_, OLDev, OL_ALLOC_TYPE_DEVICE, NumBytes, _))
       .Times(1);
   void *Ptr1 = aligned_alloc_device(ZeroAlignment, NumBytes, Dev, Ctx);
   EXPECT_NE(Ptr1, nullptr);
-  EXPECT_CALL(Mock.get(), olMemFree(Ptr1)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, Ptr1)).Times(1);
   free(Ptr1, Ctx);
 
-  EXPECT_CALL(Mock.get(), olMemAllocHost(OLDev, NumBytes, _)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemAllocHost(_, OLDev, NumBytes, _)).Times(1);
   void *Ptr2 = aligned_alloc_host(ZeroAlignment, NumBytes, Ctx);
   EXPECT_NE(Ptr2, nullptr);
-  EXPECT_CALL(Mock.get(), olMemFree(Ptr2)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, Ptr2)).Times(1);
   free(Ptr2, Ctx);
 
-  EXPECT_CALL(Mock.get(), olMemAlloc(OLDev, OL_ALLOC_TYPE_MANAGED, NumBytes, _))
+  EXPECT_CALL(Mock.get(),
+              olMemAlloc(_, OLDev, OL_ALLOC_TYPE_MANAGED, NumBytes, _))
       .Times(1);
   void *Ptr3 = aligned_alloc_shared(ZeroAlignment, NumBytes, Dev, Ctx);
   EXPECT_NE(Ptr3, nullptr);
-  EXPECT_CALL(Mock.get(), olMemFree(Ptr3)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, Ptr3)).Times(1);
   free(Ptr3, Ctx);
 }
 
@@ -194,54 +199,54 @@ TEST(USMFunctions, TemplatedAlignment) {
   context Ctx = Q.get_context();
   ol_device_handle_t OLDev = detail::getSyclObjImpl(Dev)->getOLHandle();
 
-  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _)).Times(0);
-  EXPECT_CALL(Mock.get(), olMemAllocHost(_, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAlloc(_, _, _, _, _)).Times(0);
+  EXPECT_CALL(Mock.get(), olMemAllocHost(_, _, _, _)).Times(0);
 
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(OLDev, OL_ALLOC_TYPE_DEVICE,
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, OLDev, OL_ALLOC_TYPE_DEVICE,
                                             sizeof(Over), alignof(Over), _))
       .Times(1);
   Over *P1 = aligned_alloc_device<Over>(1, 1, Dev, Ctx);
   EXPECT_NE(P1, nullptr);
-  EXPECT_CALL(Mock.get(), olMemFree(P1)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, P1)).Times(1);
   free(P1, Ctx);
 
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(OLDev, OL_ALLOC_TYPE_DEVICE,
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, OLDev, OL_ALLOC_TYPE_DEVICE,
                                             sizeof(Over), alignof(Over), _))
       .Times(1);
   Over *P2 = malloc_device<Over>(1, Dev, Ctx);
   EXPECT_NE(P2, nullptr);
-  EXPECT_CALL(Mock.get(), olMemFree(P2)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, P2)).Times(1);
   free(P2, Ctx);
 
   EXPECT_CALL(Mock.get(),
-              olMemAllocAlignedHost(OLDev, sizeof(Over), alignof(Over), _))
+              olMemAllocAlignedHost(_, OLDev, sizeof(Over), alignof(Over), _))
       .Times(1);
   Over *P3 = aligned_alloc_host<Over>(1, 1, Ctx);
   EXPECT_NE(P3, nullptr);
-  EXPECT_CALL(Mock.get(), olMemFree(P3)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, P3)).Times(1);
   free(P3, Ctx);
 
   EXPECT_CALL(Mock.get(),
-              olMemAllocAlignedHost(OLDev, sizeof(Over), alignof(Over), _))
+              olMemAllocAlignedHost(_, OLDev, sizeof(Over), alignof(Over), _))
       .Times(1);
   Over *P4 = malloc_host<Over>(1, Ctx);
   EXPECT_NE(P4, nullptr);
-  EXPECT_CALL(Mock.get(), olMemFree(P4)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, P4)).Times(1);
   free(P4, Ctx);
 
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(OLDev, OL_ALLOC_TYPE_MANAGED,
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, OLDev, OL_ALLOC_TYPE_MANAGED,
                                             sizeof(Over), alignof(Over), _))
       .Times(1);
   Over *P5 = aligned_alloc_shared<Over>(1, 1, Dev, Ctx);
   EXPECT_NE(P5, nullptr);
-  EXPECT_CALL(Mock.get(), olMemFree(P5)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, P5)).Times(1);
   free(P5, Ctx);
 
-  EXPECT_CALL(Mock.get(), olMemAllocAligned(OLDev, OL_ALLOC_TYPE_MANAGED,
+  EXPECT_CALL(Mock.get(), olMemAllocAligned(_, OLDev, OL_ALLOC_TYPE_MANAGED,
                                             sizeof(Over), alignof(Over), _))
       .Times(1);
   Over *P6 = malloc_shared<Over>(1, Dev, Ctx);
   EXPECT_NE(P6, nullptr);
-  EXPECT_CALL(Mock.get(), olMemFree(P6)).Times(1);
+  EXPECT_CALL(Mock.get(), olMemFree(_, P6)).Times(1);
   free(P6, Ctx);
 }
diff --git a/llvm/tools/llvm-gpu-loader/llvm-gpu-loader.cpp b/llvm/tools/llvm-gpu-loader/llvm-gpu-loader.cpp
index 55b59faeefdb4..d8addd06e239c 100644
--- a/llvm/tools/llvm-gpu-loader/llvm-gpu-loader.cpp
+++ b/llvm/tools/llvm-gpu-loader/llvm-gpu-loader.cpp
@@ -93,6 +93,7 @@ static cl::list<std::string> Args(cl::ConsumeAfter,
     handleError(Err, __LINE__);
 
 static void *copyArgumentVector(int Argc, const char **Argv,
+                                ol_context_handle_t Context,
                                 ol_device_handle_t Device) {
   size_t ArgSize = sizeof(char *) * (Argc + 1);
   size_t StringLen = 0;
@@ -101,7 +102,7 @@ static void *copyArgumentVector(int Argc, const char **Argv,
 
   // We allocate enough space for a null terminated array and all the strings.
   void *DevArgv;
-  OFFLOAD_ERR(olMemAllocHost(Device, ArgSize + StringLen, &DevArgv));
+  OFFLOAD_ERR(olMemAllocHost(Context, Device, ArgSize + StringLen, &DevArgv));
   if (!DevArgv)
     handleError(
         createStringError("Failed to allocate memory for environment."));
@@ -120,12 +121,13 @@ static void *copyArgumentVector(int Argc, const char **Argv,
   return DevArgv;
 }
 
-void *copyEnvironment(const char **Envp, ol_device_handle_t Device) {
+void *copyEnvironment(const char **Envp, ol_context_handle_t Context,
+                      ol_device_handle_t Device) {
   int Envc = 0;
   for (const char **Env = Envp; *Env != 0; ++Env)
     ++Envc;
 
-  return copyArgumentVector(Envc, Envp, Device);
+  return copyArgumentVector(Envc, Envp, Context, Device);
 }
 
 ol_device_handle_t findDevice(MemoryBufferRef Binary) {
@@ -263,12 +265,14 @@ int main(int argc, const char **argv, const char **envp) {
   OFFLOAD_ERR(olCreateQueue(Context, Device, &Queue));
 
   int DevArgc = static_cast<int>(NewArgv.size());
-  void *DevArgv = copyArgumentVector(NewArgv.size(), NewArgv.begin(), Device);
-  void *DevEnvp = copyEnvironment(envp, Device);
+  void *DevArgv =
+      copyArgumentVector(NewArgv.size(), NewArgv.begin(), Context, Device);
+  void *DevEnvp = copyEnvironment(envp, Context, Device);
 
   void *DevRet;
   int Zero = 0;
-  OFFLOAD_ERR(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, sizeof(int), &DevRet));
+  OFFLOAD_ERR(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, sizeof(int), &DevRet));
   OFFLOAD_ERR(olMemcpy(Queue, DevRet, Device, &Zero, Host, sizeof(int)));
 
   uint32_t Dims = (BlocksZ > 1) ? 3 : (BlocksY > 1) ? 2 : 1;
@@ -299,9 +303,9 @@ int main(int argc, const char **argv, const char **envp) {
   OFFLOAD_ERR(olMemcpy(Queue, &Ret, Host, DevRet, Device, sizeof(int)));
   OFFLOAD_ERR(olSyncQueue(Queue));
 
-  OFFLOAD_ERR(olMemFree(DevRet));
-  OFFLOAD_ERR(olMemFree(DevArgv));
-  OFFLOAD_ERR(olMemFree(DevEnvp));
+  OFFLOAD_ERR(olMemFree(Context, DevRet));
+  OFFLOAD_ERR(olMemFree(Context, DevArgv));
+  OFFLOAD_ERR(olMemFree(Context, DevEnvp));
   OFFLOAD_ERR(olDestroyQueue(Queue));
   OFFLOAD_ERR(olDestroyContext(Context));
   OFFLOAD_ERR(olDestroyProgram(Program));
diff --git a/llvm/tools/llvm-gpu-loader/llvm-gpu-loader.h b/llvm/tools/llvm-gpu-loader/llvm-gpu-loader.h
index c1dc25cc76cbc..55bb0c3ee1b8e 100644
--- a/llvm/tools/llvm-gpu-loader/llvm-gpu-loader.h
+++ b/llvm/tools/llvm-gpu-loader/llvm-gpu-loader.h
@@ -153,13 +153,15 @@ ol_result_t (*olDestroyQueue)(ol_queue_handle_t Queue);
 
 ol_result_t (*olSyncQueue)(ol_queue_handle_t Queue);
 
-ol_result_t (*olMemAlloc)(ol_device_handle_t Device, ol_alloc_type_t Type,
+ol_result_t (*olMemAlloc)(ol_context_handle_t Context,
+                          ol_device_handle_t Device, ol_alloc_type_t Type,
                           size_t Size, void **AllocationOut);
 
-ol_result_t (*olMemAllocHost)(ol_device_handle_t Device, size_t Size,
+ol_result_t (*olMemAllocHost)(ol_context_handle_t Context,
+                              ol_device_handle_t Device, size_t Size,
                               void **AllocationOut);
 
-ol_result_t (*olMemFree)(void *Address);
+ol_result_t (*olMemFree)(ol_context_handle_t Context, void *Address);
 
 ol_result_t (*olMemcpy)(ol_queue_handle_t Queue, void *DstPtr,
                         ol_device_handle_t DstDevice, const void *SrcPtr,
diff --git a/offload/languages/kernel/src/LanguageRuntime.cpp b/offload/languages/kernel/src/LanguageRuntime.cpp
index 6dab19fc2fd8c..b4eb18c59de30 100644
--- a/offload/languages/kernel/src/LanguageRuntime.cpp
+++ b/offload/languages/kernel/src/LanguageRuntime.cpp
@@ -34,12 +34,15 @@ using namespace llvm::offload;
 Error_t Malloc(void **DevPtr, size_t Size) {
   ThreadStateTy &ThreadState = ThreadStateTy::get();
   ol_device_handle_t Device = ThreadState.getDefaultDevice();
-  ol_result_t Result = olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, DevPtr);
+  ol_context_handle_t Context = StateTy::get().getContext();
+  ol_result_t Result =
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, DevPtr);
   return convertAndSetLastError(Result);
 }
 
 Error_t Free(void *DevPtr) {
-  ol_result_t Result = olMemFree(DevPtr);
+  ol_context_handle_t Context = StateTy::get().getContext();
+  ol_result_t Result = olMemFree(Context, DevPtr);
   return convertAndSetLastError(Result);
 }
 
@@ -121,7 +124,8 @@ Error_t SetDevice(int DeviceNo) {
 Error_t HostAlloc(void **Ptr, size_t Size, unsigned int Flags) {
   ThreadStateTy &ThreadState = ThreadStateTy::get();
   ol_device_handle_t Device = ThreadState.getDefaultDevice();
-  ol_result_t Result = olMemAllocHost(Device, Size, Ptr);
+  ol_context_handle_t Context = StateTy::get().getContext();
+  ol_result_t Result = olMemAllocHost(Context, Device, Size, Ptr);
   return convertAndSetLastError(Result);
 }
 
@@ -130,7 +134,8 @@ Error_t MallocHost(void **Ptr, size_t Size) {
 }
 
 Error_t FreeHost(void *Ptr) {
-  ol_result_t Result = olMemFree(Ptr);
+  ol_context_handle_t Context = StateTy::get().getContext();
+  ol_result_t Result = olMemFree(Context, Ptr);
   return convertAndSetLastError(Result);
 }
 
diff --git a/offload/liboffload/API/Memory.td b/offload/liboffload/API/Memory.td
index 17eabeccce9e0..565f3fc62e056 100644
--- a/offload/liboffload/API/Memory.td
+++ b/offload/liboffload/API/Memory.td
@@ -36,10 +36,11 @@ def ol_memory_register_flag_t : Enum {
 def olMemAlloc : Function {
   let desc = "Creates a memory allocation on the specified device.";
   let details = [
-      "All liboffload allocations share a single virtual address range. There is no risk of multiple devices returning equal pointers to different memory.",
+      "The allocation is scoped to `Context` and `Device` must belong to it.",
       "This function can only be used to create device or managed allocations. To create a host allocation use `olMemAllocHost`."
   ];
   let params = [
+    Param<"ol_context_handle_t", "Context", "handle of the context", PARAM_IN>,
     Param<"ol_device_handle_t", "Device", "handle of the device to allocate on", PARAM_IN>,
     Param<"ol_alloc_type_t", "Type", "type of the allocation. Must be either `OL_ALLOC_TYPE_DEVICE` or `OL_ALLOC_TYPE_MANAGED`", PARAM_IN>,
     Param<"size_t", "Size", "size of the allocation in bytes", PARAM_IN>,
@@ -51,6 +52,9 @@ def olMemAlloc : Function {
     ]>,
     Return<"OL_ERRC_INVALID_ENUMERATION", [
       "`Type == OL_ALLOC_TYPE_HOST`"
+    ]>,
+    Return<"OL_ERRC_INVALID_DEVICE", [
+      "Device does not belong to `Context`"
     ]>
   ];
 }
@@ -58,9 +62,10 @@ def olMemAlloc : Function {
 def olMemAllocHost : Function {
   let desc = "Creates a host memory allocation accessible from the specified device.";
   let details = [
-      "All liboffload allocations share a single virtual address range. There is no risk of multiple devices returning equal pointers to different memory."
+      "The allocation is scoped to `Context` and `Device` must belong to it."
   ];
   let params = [
+    Param<"ol_context_handle_t", "Context", "handle of the context", PARAM_IN>,
     Param<"ol_device_handle_t", "Device", "handle of the device to allocate on", PARAM_IN>,
     Param<"size_t", "Size", "size of the allocation in bytes", PARAM_IN>,
     Param<"void**", "AllocationOut", "output for the allocated pointer", PARAM_OUT>
@@ -68,6 +73,9 @@ def olMemAllocHost : Function {
   let returns = [
     Return<"OL_ERRC_INVALID_SIZE", [
       "`Size == 0`"
+    ]>,
+    Return<"OL_ERRC_INVALID_DEVICE", [
+      "Device does not belong to `Context`"
     ]>
   ];
 }
@@ -75,10 +83,11 @@ def olMemAllocHost : Function {
 def olMemAllocAligned : Function {
   let desc = "Creates a memory allocation on the specified device with the specified alignment.";
   let details = [
-      "All liboffload allocations share a single virtual address range. There is no risk of multiple devices returning equal pointers to different memory.",
+      "The allocation is scoped to `Context` and `Device` must belong to it.",
       "This function can only be used to create device or managed allocations. To create a host allocation use `olMemAllocAlignedHost`."
   ];
   let params = [
+    Param<"ol_context_handle_t", "Context", "handle of the context", PARAM_IN>,
     Param<"ol_device_handle_t", "Device", "handle of the device to allocate on", PARAM_IN>,
     Param<"ol_alloc_type_t", "Type", "type of the allocation. Must be either `OL_ALLOC_TYPE_DEVICE` or `OL_ALLOC_TYPE_MANAGED`", PARAM_IN>,
     Param<"size_t", "Size", "size of the allocation in bytes", PARAM_IN>,
@@ -100,15 +109,19 @@ def olMemAllocAligned : Function {
     Return<"OL_ERRC_INVALID_ENUMERATION", [
       "`Type == OL_ALLOC_TYPE_HOST`"
     ]>,
+    Return<"OL_ERRC_INVALID_DEVICE", [
+      "Device does not belong to `Context`"
+    ]>,
   ];
 }
 
 def olMemAllocAlignedHost : Function {
   let desc = "Creates a host memory allocation with the specified alignment.";
   let details = [
-      "All liboffload allocations share a single virtual address range. There is no risk of multiple devices returning equal pointers to different memory."
+      "The allocation is scoped to `Context` and `Device` must belong to it."
   ];
   let params = [
+    Param<"ol_context_handle_t", "Context", "handle of the context", PARAM_IN>,
     Param<"ol_device_handle_t", "Device", "handle of the device to allocate on", PARAM_IN>,
     Param<"size_t", "Size", "size of the allocation in bytes", PARAM_IN>,
     Param<"size_t", "Alignment",
@@ -126,15 +139,23 @@ def olMemAllocAlignedHost : Function {
     Return<"OL_ERRC_INVALID_ARGUMENT", [
       "`(Alignment & (Alignment - 1)) != 0`"
     ]>,
+    Return<"OL_ERRC_INVALID_DEVICE", [
+      "Device does not belong to `Context`"
+    ]>,
   ];
 }
 
 def olMemFree : Function {
   let desc = "Frees a memory allocation previously made by an olMemAlloc* function.";
   let params = [
+    Param<"ol_context_handle_t", "Context", "handle of the context the allocation was made in", PARAM_IN>,
     Param<"void*", "Address", "address of the allocation to free", PARAM_IN>,
   ];
-  let returns = [];
+  let returns = [
+    Return<"OL_ERRC_INVALID_CONTEXT", [
+      "The allocation was not made in `Context`"
+    ]>
+  ];
 }
 
 def ol_mem_info_t : Enum {
@@ -153,8 +174,10 @@ def olGetMemInfo : Function {
   let details = [
     "`olGetMemInfoSize` can be used to query the storage size required for the given query.",
     "The provided pointer can point to any location inside the allocation.",
+    "The allocation must belong to `Context`.",
   ];
   let params = [
+    Param<"ol_context_handle_t", "Context", "handle of the context the allocation was made in", PARAM_IN>,
     Param<"const void *", "Ptr", "pointer to the allocated memory", PARAM_IN>,
     Param<"ol_mem_info_t", "PropName", "type of the info to retrieve", PARAM_IN>,
     Param<"size_t", "PropSize", "the number of bytes pointed to by PropValue.", PARAM_IN>,
@@ -168,7 +191,7 @@ def olGetMemInfo : Function {
       "`PropSize == 0`",
       "If `PropSize` is less than the real number of bytes needed to return the info."
     ]>,
-    Return<"OL_ERRC_NOT_FOUND", ["memory was not allocated by liboffload"]>
+    Return<"OL_ERRC_NOT_FOUND", ["memory was not allocated in `Context`"]>
   ];
 }
 
@@ -176,14 +199,16 @@ def olGetMemInfoSize : Function {
   let desc = "Returns the storage size of the given queue query.";
   let details = [
     "The provided pointer can point to any location inside the allocation.",
+    "The allocation must belong to `Context`.",
   ];
   let params = [
+    Param<"ol_context_handle_t", "Context", "handle of the context the allocation was made in", PARAM_IN>,
     Param<"const void *", "Ptr", "pointer to the allocated memory", PARAM_IN>,
     Param<"ol_mem_info_t", "PropName", "type of the info to query", PARAM_IN>,
     Param<"size_t*", "PropSizeRet", "pointer to the number of bytes required to store the query", PARAM_OUT>
   ];
   let returns = [
-    Return<"OL_ERRC_NOT_FOUND", ["memory was not allocated by liboffload"]>
+    Return<"OL_ERRC_NOT_FOUND", ["memory was not allocated in `Context`"]>
   ];
 }
 
diff --git a/offload/liboffload/src/OffloadImpl.cpp b/offload/liboffload/src/OffloadImpl.cpp
index 11c07483b38aa..d090a06967751 100644
--- a/offload/liboffload/src/OffloadImpl.cpp
+++ b/offload/liboffload/src/OffloadImpl.cpp
@@ -167,8 +167,32 @@ struct ol_context_impl_t {
   llvm::SmallVector<ol_device_handle_t> Devices;
   std::unique_ptr<plugin::PluginContextTy> PluginCtx;
 
-  bool contains(ol_device_handle_t Device) const {
-    return llvm::is_contained(Devices, Device);
+  llvm::Error requireDevice(ol_device_handle_t Device) const {
+    if (!llvm::is_contained(Devices, Device))
+      return createOffloadError(ErrorCode::INVALID_DEVICE,
+                                "device does not belong to the given context");
+    return llvm::Error::success();
+  }
+
+  ol_device_handle_t findDevice(plugin::GenericDeviceTy *Device) const {
+    for (auto *D : Devices)
+      if (D->Device == Device)
+        return D;
+    return nullptr;
+  }
+
+  llvm::Expected<void *> allocate(ol_device_handle_t Device, int64_t Size,
+                                  TargetAllocTy Kind, size_t Alignment) {
+    if (auto Err = requireDevice(Device))
+      return std::move(Err);
+    return PluginCtx->allocate(*Device->Device, Size, /*HostPtr=*/nullptr, Kind,
+                               Alignment);
+  }
+
+  llvm::Error deallocate(void *Ptr) { return PluginCtx->deallocate(Ptr); }
+
+  llvm::Expected<PluginAllocInfoTy> getAllocInfo(const void *Ptr) {
+    return PluginCtx->getAllocInfo(Ptr);
   }
 
   /// Queues destroyed while still busy, keyed by owning device. Per-context
@@ -241,14 +265,6 @@ struct ol_context_impl_t {
 namespace llvm {
 namespace offload {
 
-struct AllocInfo {
-  ol_device_handle_t Device;
-  ol_alloc_type_t Type;
-  void *Start;
-  // One byte past the end
-  void *End;
-};
-
 // Global shared state for liboffload
 struct OffloadContext;
 // This pointer is non-null if and only if the context is valid and fully
@@ -263,11 +279,6 @@ struct OffloadContext {
 
   bool TracingEnabled = false;
   bool ValidationEnabled = true;
-  DenseMap<void *, AllocInfo> AllocInfoMap{};
-  std::mutex AllocInfoMapMutex{};
-  // Partitioned list of memory base addresses. Each element in this list is a
-  // key in AllocInfoMap
-  SmallVector<void *> AllocBases{};
   SmallVector<std::unique_ptr<ol_platform_impl_t>, 4> Platforms{};
   size_t RefCount;
 
@@ -701,176 +712,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);
 }
 
-Error olGetMemInfoImplDetail(const void *Ptr, ol_mem_info_t PropName,
-                             size_t PropSize, void *PropValue,
-                             size_t *PropSizeRet) {
+Error olGetMemInfoImplDetail(ol_context_handle_t Context, const void *Ptr,
+                             ol_mem_info_t PropName, size_t PropSize,
+                             void *PropValue, size_t *PropSizeRet) {
   InfoWriter Info(PropSize, PropValue, PropSizeRet);
-  std::lock_guard<std::mutex> Lock(OffloadContext::get().AllocInfoMapMutex);
-
-  auto &AllocBases = OffloadContext::get().AllocBases;
-  auto &AllocInfoMap = OffloadContext::get().AllocInfoMap;
-  const AllocInfo *Alloc = nullptr;
-  if (AllocInfoMap.contains(Ptr)) {
-    // Fast case, we have been given the base pointer directly
-    Alloc = &AllocInfoMap.at(Ptr);
-  } else {
-    // Slower case, we need to look up the base pointer first
-    // Find the first memory allocation whose end is after the target pointer,
-    // and then check to see if it is in range
-    auto Loc = std::lower_bound(AllocBases.begin(), AllocBases.end(), Ptr,
-                                [&](const void *Iter, const void *Val) {
-                                  return AllocInfoMap.at(Iter).End <= Val;
-                                });
-    if (Loc == AllocBases.end() || Ptr < AllocInfoMap.at(*Loc).Start)
-      return Plugin::error(ErrorCode::NOT_FOUND,
-                           "allocated memory information not found");
-    Alloc = &AllocInfoMap.at(*Loc);
-  }
+
+  auto AllocOrErr = Context->getAllocInfo(Ptr);
+  if (!AllocOrErr)
+    return AllocOrErr.takeError();
+  const auto &Alloc = *AllocOrErr;
 
   switch (PropName) {
-  case OL_MEM_INFO_DEVICE:
-    return Info.write<ol_device_handle_t>(Alloc->Device);
+  case OL_MEM_INFO_DEVICE: {
+    // OL_MEM_INFO_DEVICE is not meaningful for host allocations: a host pool
+    // allocation has no per-device affinity. This does not affect the
+    // size-only query (PropValue == nullptr): the answer is always
+    // sizeof(ol_device_handle_t).
+    if (PropValue && Alloc.Kind == TARGET_ALLOC_HOST)
+      return createOffloadError(
+          ErrorCode::INVALID_ARGUMENT,
+          "OL_MEM_INFO_DEVICE is not valid for host allocations");
+    if (PropValue) {
+      ol_device_handle_t OlDev = Context->findDevice(Alloc.Device);
+      if (!OlDev)
+        return createOffloadError(ErrorCode::NOT_FOUND,
+                                  "allocation device not part of this context");
+      return Info.write<ol_device_handle_t>(OlDev);
+    }
+    return Info.write<ol_device_handle_t>(nullptr);
+  }
   case OL_MEM_INFO_BASE:
-    return Info.write<void *>(Alloc->Start);
+    return Info.write<void *>(Alloc.Base);
   case OL_MEM_INFO_SIZE:
-    return Info.write<size_t>(static_cast<char *>(Alloc->End) -
-                              static_cast<char *>(Alloc->Start));
+    return Info.write<size_t>(Alloc.Size);
   case OL_MEM_INFO_TYPE:
-    return Info.write<ol_alloc_type_t>(Alloc->Type);
+    return Info.write<ol_alloc_type_t>(convertPluginToOlAllocTy(Alloc.Kind));
   default:
     return createOffloadError(ErrorCode::INVALID_ENUMERATION,
                               "olGetMemInfo enum '%i' is invalid", PropName);
   }
 }
 
-Error olGetMemInfo_impl(const void *Ptr, ol_mem_info_t PropName,
-                        size_t PropSize, void *PropValue) {
-  return olGetMemInfoImplDetail(Ptr, PropName, PropSize, PropValue, nullptr);
+Error olGetMemInfo_impl(ol_context_handle_t Context, const void *Ptr,
+                        ol_mem_info_t PropName, size_t PropSize,
+                        void *PropValue) {
+  return olGetMemInfoImplDetail(Context, Ptr, PropName, PropSize, PropValue,
+                                nullptr);
 }
 
-Error olGetMemInfoSize_impl(const void *Ptr, ol_mem_info_t PropName,
-                            size_t *PropSizeRet) {
-  return olGetMemInfoImplDetail(Ptr, PropName, 0, nullptr, PropSizeRet);
+Error olGetMemInfoSize_impl(ol_context_handle_t Context, const void *Ptr,
+                            ol_mem_info_t PropName, size_t *PropSizeRet) {
+  return olGetMemInfoImplDetail(Context, Ptr, PropName, 0, nullptr,
+                                PropSizeRet);
 }
 
 Error olCreateQueue_impl(ol_context_handle_t Context, ol_device_handle_t Device,
                          ol_queue_handle_t *Queue) {
-  if (!Context->contains(Device))
-    return createOffloadError(ErrorCode::INVALID_DEVICE,
-                              "device does not belong to the given context");
+  if (auto Err = Context->requireDevice(Device))
+    return Err;
 
   auto CreatedQueue =
       std::make_unique<ol_queue_impl_t>(nullptr, Context, Device);
@@ -1167,9 +1135,8 @@ Error olMemPrefetch_impl(ol_queue_handle_t Queue, size_t Count,
 Error olCreateProgram_impl(ol_context_handle_t Context,
                            ol_device_handle_t Device, const void *ProgData,
                            size_t ProgDataSize, ol_program_handle_t *Program) {
-  if (!Context->contains(Device))
-    return createOffloadError(ErrorCode::INVALID_DEVICE,
-                              "device does not belong to the given context");
+  if (auto Err = Context->requireDevice(Device))
+    return Err;
 
   StringRef Buffer(reinterpret_cast<const char *>(ProgData), ProgDataSize);
   Expected<plugin::DeviceImageTy *> Res = Device->Device->loadBinary(
diff --git a/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa_ext_amd.h b/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa_ext_amd.h
index d26f9248e27ef..c736ec0759841 100644
--- a/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa_ext_amd.h
+++ b/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa_ext_amd.h
@@ -172,6 +172,8 @@ typedef struct hsa_amd_pointer_info_s {
   void* agentBaseAddress;
   void* hostBaseAddress;
   size_t sizeInBytes;
+  void *userData;
+  hsa_agent_t agentOwner;
 } hsa_amd_pointer_info_t;
 
 typedef enum {
diff --git a/offload/plugins-nextgen/amdgpu/src/rtl.cpp b/offload/plugins-nextgen/amdgpu/src/rtl.cpp
index 1af5d4950a8d8..92afb883840ef 100644
--- a/offload/plugins-nextgen/amdgpu/src/rtl.cpp
+++ b/offload/plugins-nextgen/amdgpu/src/rtl.cpp
@@ -3986,6 +3986,24 @@ struct AMDGPUPluginContextTy final : public PluginContextTy {
     // TODO: Implement this function.
     return Plugin::success();
   }
+
+  Expected<void *> allocate(GenericDeviceTy &Device, int64_t Size,
+                            void *HostPtr, TargetAllocTy Kind,
+                            size_t Alignment) override;
+  Error deallocate(GenericDeviceTy &Device, void *Ptr,
+                   TargetAllocTy Kind) override;
+  Expected<PluginAllocInfoTy> getAllocInfo(const void *Ptr) override;
+
+private:
+  // Track each allocation's Kind so we can tell HOST from SHARED — HSA
+  // backs both with the same host fine-grained pool.
+  // TODO: drop once TARGET_ALLOC_SHARED has its own backing pool.
+  struct AllocInfo {
+    TargetAllocTy Kind;
+    GenericDeviceTy *Device;
+  };
+  llvm::DenseMap<const void *, AllocInfo> Allocations;
+  std::mutex AllocationsMutex;
 };
 
 /// Class implementing the AMDGPU-specific functionalities of the plugin.
@@ -4295,6 +4313,58 @@ struct AMDGPUPluginTy final : public GenericPluginTy {
   AMDHostDeviceTy *HostDevice;
 };
 
+Expected<void *> AMDGPUPluginContextTy::allocate(GenericDeviceTy &Device,
+                                                 int64_t Size, void *HostPtr,
+                                                 TargetAllocTy Kind,
+                                                 size_t Alignment) {
+  auto PtrOrErr =
+      PluginContextTy::allocate(Device, Size, HostPtr, Kind, Alignment);
+  if (!PtrOrErr || !*PtrOrErr)
+    return PtrOrErr;
+  std::lock_guard<std::mutex> Lock(AllocationsMutex);
+  Allocations[*PtrOrErr] = {Kind, &Device};
+  return PtrOrErr;
+}
+
+Error AMDGPUPluginContextTy::deallocate(GenericDeviceTy &Device, void *Ptr,
+                                        TargetAllocTy Kind) {
+  // Erase before base deallocate: once Ptr returns to the MM freelist a
+  // concurrent alloc could reuse it and re-populate Allocations. On failure
+  // Ptr is in an undetermined state (maybe freed, maybe not) so we don't
+  // re-add it either.
+  {
+    std::lock_guard<std::mutex> Lock(AllocationsMutex);
+    Allocations.erase(Ptr);
+  }
+  return PluginContextTy::deallocate(Device, Ptr, Kind);
+}
+
+Expected<PluginAllocInfoTy>
+AMDGPUPluginContextTy::getAllocInfo(const void *Ptr) {
+  // HSA gives the base of the region containing Ptr, so interior pointers
+  // resolve to the same tracker entry.
+  hsa_amd_pointer_info_t HsaInfo{};
+  HsaInfo.size = sizeof(hsa_amd_pointer_info_t);
+  hsa_status_t Status = hsa_amd_pointer_info(
+      const_cast<void *>(Ptr), &HsaInfo, /*Allocator=*/nullptr,
+      /*num_agents_accessible=*/nullptr, /*accessible=*/nullptr);
+  if (auto Err = Plugin::check(Status, "error in hsa_amd_pointer_info: %s"))
+    return std::move(Err);
+
+  AllocInfo Info;
+  {
+    std::lock_guard<std::mutex> Lock(AllocationsMutex);
+    auto It = Allocations.find(HsaInfo.agentBaseAddress);
+    if (It == Allocations.end())
+      return Plugin::error(ErrorCode::NOT_FOUND,
+                           "pointer is not a known allocation in this context");
+    Info = It->second;
+  }
+
+  return PluginAllocInfoTy{Info.Device, Info.Kind, HsaInfo.agentBaseAddress,
+                           HsaInfo.sizeInBytes};
+}
+
 Error AMDGPUKernelTy::launchImpl(GenericDeviceTy &GenericDevice,
                                  uint32_t NumThreads[3], uint32_t NumBlocks[3],
                                  uint32_t DynBlockMemSize,
diff --git a/offload/plugins-nextgen/common/include/PluginInterface.h b/offload/plugins-nextgen/common/include/PluginInterface.h
index 0881fa3fd7627..bc3b6db0a1ec1 100644
--- a/offload/plugins-nextgen/common/include/PluginInterface.h
+++ b/offload/plugins-nextgen/common/include/PluginInterface.h
@@ -900,6 +900,14 @@ class PinnedAllocationMapTy {
   }
 };
 
+/// Description of an allocation: owning device, kind, base address and size.
+struct PluginAllocInfoTy {
+  GenericDeviceTy *Device;
+  TargetAllocTy Kind;
+  void *Base;
+  size_t Size;
+};
+
 /// A plugin-side context grouping a set of devices.
 struct PluginContextTy {
   PluginContextTy(GenericPluginTy &Plugin,
@@ -911,7 +919,7 @@ struct PluginContextTy {
   PluginContextTy(PluginContextTy &&) = delete;
   PluginContextTy &operator=(PluginContextTy &&) = delete;
 
-  virtual ~PluginContextTy() = default;
+  virtual ~PluginContextTy();
 
   /// Release resources owned by this context. Called from olDestroyContext
   /// before the object is destroyed so that errors are propagated instead of
@@ -926,9 +934,63 @@ struct PluginContextTy {
   virtual Error initAsyncInfoImpl(GenericDeviceTy &Device,
                                   AsyncInfoWrapperTy &AsyncInfoWrapper) = 0;
 
+  /// Allocate Size bytes of Kind memory accessible from Device. HostPtr is an
+  /// optional hint (e.g. for pinned-buffer registration); pass nullptr when
+  /// unused.
+  virtual llvm::Expected<void *> allocate(GenericDeviceTy &Device, int64_t Size,
+                                          void *HostPtr, TargetAllocTy Kind,
+                                          size_t Alignment);
+
+  /// Free a pointer returned by allocate; resolves owner/kind via
+  /// getAllocInfo. Requires a non-empty device set, so this is only valid on
+  /// user-created contexts (not on the per-plugin default context, which
+  /// carries no devices).
+  virtual llvm::Error deallocate(void *Ptr);
+
+  /// Free a pointer when the caller already knows the owning device and kind.
+  virtual llvm::Error deallocate(GenericDeviceTy &Device, void *Ptr,
+                                 TargetAllocTy Kind);
+
+  /// Look up the allocation containing Ptr. Returns NOT_FOUND when Ptr is not
+  /// known to this context. Only valid on user-created contexts.
+  virtual llvm::Expected<PluginAllocInfoTy> getAllocInfo(const void *Ptr) = 0;
+
 protected:
   GenericPluginTy &Plugin;
   llvm::SmallVector<GenericDeviceTy *> Devices;
+
+private:
+  MemoryManagerTy *getDeviceMemoryManagerFor(GenericDeviceTy &Device,
+                                             TargetAllocTy Kind);
+  MemoryManagerTy *getHostMemoryManager();
+
+  llvm::DenseMap<std::pair<GenericDeviceTy *, int>,
+                 std::unique_ptr<MemoryManagerTy>>
+      DeviceMemoryManagers;
+  std::unique_ptr<MemoryManagerTy> HostMemoryManager;
+  std::mutex MemoryManagersMutex;
+};
+
+/// Default plugin context: a device-less placeholder used as the per-plugin
+/// MemoryManager dispatcher for libomptarget. Allocations made through this
+/// context are pooled but not tracked per-pointer, so getAllocInfo is not
+/// supported.
+struct DefaultPluginContextTy : public PluginContextTy {
+  DefaultPluginContextTy(GenericPluginTy &Plugin)
+      : PluginContextTy(Plugin, llvm::ArrayRef<GenericDeviceTy *>{}) {}
+
+  llvm::Expected<PluginAllocInfoTy>
+  getAllocInfo(const void * /*Ptr*/) override {
+    return Plugin::error(error::ErrorCode::UNSUPPORTED,
+                         "getAllocInfo is not supported on the default "
+                         "plugin context");
+  }
+
+  Error initAsyncInfoImpl(GenericDeviceTy & /*Device*/,
+                          AsyncInfoWrapperTy & /*AsyncInfoWrapper*/) override {
+    return Plugin::error(error::ErrorCode::UNSUPPORTED,
+                         "default plugin context cannot initialize async info");
+  }
 };
 
 /// Class implementing common functionalities of offload devices. Each plugin
@@ -1407,6 +1469,9 @@ struct GenericDeviceTy : public DeviceAllocatorTy {
   /// deallocated by the allocator.
   llvm::SmallVector<DeviceImageTy *> LoadedImages;
 
+  /// Per device setting of MemoryManager's Threshold
+  virtual size_t getMemoryManagerSizeThreshold() { return 0; }
+
 private:
   /// Get and set the stack size and heap size for the device. If not used, the
   /// plugin can implement the setters as no-op and setting the output
@@ -1417,16 +1482,6 @@ struct GenericDeviceTy : public DeviceAllocatorTy {
   /// only necessary for unhosted targets like the GPU.
   virtual bool shouldSetupRPCServer() const { return false; }
 
-  /// Pointer to the device memory manager or nullptr if not available.
-  MemoryManagerTy *MemoryManager;
-  /// Memory managers for the host and shared allocation kinds or nullptr if not
-  /// available.
-  MemoryManagerTy *HostMemoryManager;
-  MemoryManagerTy *SharedMemoryManager;
-
-  /// Per device setting of MemoryManager's Threshold
-  virtual size_t getMemoryManagerSizeThreshold() { return 0; }
-
   virtual Expected<bool> isAccessiblePtrImpl(const void *Ptr, size_t Size) {
     return false;
   }
@@ -1459,20 +1514,6 @@ struct GenericDeviceTy : public DeviceAllocatorTy {
   /// Record and replay manager.
   RecordReplayTy *RecordReplay = nullptr;
 
-  /// Return the memory manager for the given allocation kind.
-  MemoryManagerTy *getMemoryManagerFor(TargetAllocTy Kind) {
-    switch (Kind) {
-    case TARGET_ALLOC_DEFAULT:
-    case TARGET_ALLOC_DEVICE:
-      return MemoryManager;
-    case TARGET_ALLOC_HOST:
-      return HostMemoryManager;
-    case TARGET_ALLOC_SHARED:
-      return SharedMemoryManager;
-    }
-    return nullptr;
-  }
-
 protected:
   /// Environment variables defined by the LLVM OpenMP implementation
   /// regarding the initial number of streams and events.
@@ -1691,6 +1732,16 @@ struct GenericPluginTy {
   virtual Expected<std::unique_ptr<PluginContextTy>>
   createPluginContext(llvm::ArrayRef<GenericDeviceTy *> Devices) = 0;
 
+  /// Create the per-plugin default context returned by getDefaultContext.
+  /// The default builds a plain PluginContextTy with no devices, which is
+  /// sufficient for plugins where the default is just a MemoryManager
+  /// dispatcher.
+  virtual Expected<std::unique_ptr<PluginContextTy>>
+  createDefaultPluginContext();
+
+  /// Return the default context that services Device.
+  virtual PluginContextTy &getDefaultContext(GenericDeviceTy &Device);
+
 protected:
   /// Indicate whether a device id is valid.
   bool isValidDeviceId(int32_t DeviceId) const {
@@ -1900,6 +1951,10 @@ struct GenericPluginTy {
 
   /// The interface between the plugin and the GPU for host services.
   RPCServerTy *RPCServer;
+
+  /// The default plugin context returned by getDefaultContext, used by
+  /// libomptarget.
+  std::unique_ptr<PluginContextTy> DefaultContext;
 };
 
 /// Auxiliary interface class for GenericDeviceResourceManagerTy. This class
diff --git a/offload/plugins-nextgen/common/src/PluginInterface.cpp b/offload/plugins-nextgen/common/src/PluginInterface.cpp
index e555321e1b96b..00d04e4a1f2de 100644
--- a/offload/plugins-nextgen/common/src/PluginInterface.cpp
+++ b/offload/plugins-nextgen/common/src/PluginInterface.cpp
@@ -448,8 +448,7 @@ uint32_t GenericKernelTy::getEffectiveNumBlocks(
 GenericDeviceTy::GenericDeviceTy(GenericPluginTy &Plugin, int32_t DeviceId,
                                  int32_t NumDevices,
                                  const llvm::omp::GV &OMPGridValues)
-    : Plugin(Plugin), MemoryManager(nullptr), HostMemoryManager(nullptr),
-      SharedMemoryManager(nullptr), OMP_TeamLimit("OMP_TEAM_LIMIT"),
+    : Plugin(Plugin), OMP_TeamLimit("OMP_TEAM_LIMIT"),
       OMP_NumTeams("OMP_NUM_TEAMS"),
       OMP_TeamsThreadLimit("OMP_TEAMS_THREAD_LIMIT"),
       OMPX_DebugKind("LIBOMPTARGET_DEVICE_RTL_DEBUG"),
@@ -557,24 +556,6 @@ Error GenericDeviceTy::init(GenericPluginTy &Plugin) {
     GridValues.GV_Max_WG_Size =
         std::min(GridValues.GV_Max_WG_Size, uint32_t(OMP_TeamsThreadLimit));
 
-  // Enable the memory manager if required. Leave the pool disabled while
-  // allocation traces are requested, so that we don't mask use-after-free
-  // (since they don't fault if the memory is still in the pool).
-  auto [ThresholdMM, EnableMM] = MemoryManagerTy::getSizeThresholdFromEnv();
-  if (EnableMM && !OMPX_TrackAllocationTraces) {
-    if (ThresholdMM == 0)
-      ThresholdMM = getMemoryManagerSizeThreshold();
-    MemoryManager = new MemoryManagerTy(*this, ThresholdMM);
-  }
-  if (!OMPX_TrackAllocationTraces) {
-    // Keep the threshold for pooling sizes conservative since we're dealing
-    // with pinned memory for the host.
-    HostMemoryManager = new MemoryManagerTy(
-        *this, MemoryManagerTy::DefaultSizeThreshold, TARGET_ALLOC_HOST);
-    SharedMemoryManager = new MemoryManagerTy(
-        *this, MemoryManagerTy::DefaultSizeThreshold, TARGET_ALLOC_SHARED);
-  }
-
   return Plugin::success();
 }
 
@@ -616,18 +597,6 @@ Error GenericDeviceTy::deinit(GenericPluginTy &Plugin) {
       return Err;
   LoadedImages.clear();
 
-  // Delete the memory manager before deinitializing the device. Otherwise,
-  // we may delete device allocations after the device is deinitialized.
-  if (MemoryManager)
-    delete MemoryManager;
-  MemoryManager = nullptr;
-  if (HostMemoryManager)
-    delete HostMemoryManager;
-  HostMemoryManager = nullptr;
-  if (SharedMemoryManager)
-    delete SharedMemoryManager;
-  SharedMemoryManager = nullptr;
-
   if (RecordReplay) {
     if (auto Err = RecordReplay->deinit())
       return Err;
@@ -1016,30 +985,20 @@ Expected<void *> GenericDeviceTy::dataAlloc(int64_t Size, void *HostPtr,
   if (RecordReplay && RecordReplay->isRecordingOrReplaying())
     return RecordReplay->allocate(Size);
 
-  if (MemoryManagerTy *MM = getMemoryManagerFor(Kind)) {
-    auto AllocOrErr = MM->allocate(Size, HostPtr, Alignment);
-    if (!AllocOrErr)
-      return AllocOrErr.takeError();
-    Alloc = *AllocOrErr;
-    if (!Alloc)
-      return Plugin::error(ErrorCode::OUT_OF_RESOURCES,
-                           "failed to allocate from memory manager");
-  } else {
-    auto AllocOrErr = allocate(Size, HostPtr, Kind, Alignment);
-    if (!AllocOrErr)
-      return AllocOrErr.takeError();
-    Alloc = *AllocOrErr;
-    if (!Alloc)
-      return Plugin::error(ErrorCode::OUT_OF_RESOURCES,
-                           "failed to allocate from device allocator");
+  auto AllocOrErr = allocate(Size, HostPtr, Kind, Alignment);
+  if (!AllocOrErr)
+    return AllocOrErr.takeError();
+  Alloc = *AllocOrErr;
+  if (!Alloc)
+    return Plugin::error(ErrorCode::OUT_OF_RESOURCES,
+                         "failed to allocate from device allocator");
 
-    if (Alignment > 0 && !isAddrAligned(Align(Alignment), Alloc)) {
-      if (auto Err = free(Alloc, Kind))
-        return Err;
+  if (Alignment > 0 && !isAddrAligned(Align(Alignment), Alloc)) {
+    if (auto Err = free(Alloc, Kind))
+      return Err;
 
-      return Plugin::error(ErrorCode::UNSUPPORTED,
-                           "device allocator returned a misaligned pointer");
-    }
+    return Plugin::error(ErrorCode::UNSUPPORTED,
+                         "device allocator returned a misaligned pointer");
   }
 
   // Report error if the memory manager or the device allocator did not return
@@ -1105,12 +1064,8 @@ Error GenericDeviceTy::dataDelete(void *TgtPtr, TargetAllocTy Kind) {
     ATI->DeallocationTrace = StackTrace;
   }
 
-  if (MemoryManagerTy *MM = getMemoryManagerFor(Kind)) {
-    if (auto Err = MM->free(TgtPtr))
-      return Err;
-  } else if (auto Err = free(TgtPtr, Kind)) {
+  if (auto Err = free(TgtPtr, Kind))
     return Err;
-  }
 
   return Plugin::success();
 }
@@ -1214,6 +1169,103 @@ Error PluginContextTy::initAsyncInfo(GenericDeviceTy &Device,
   return Err;
 }
 
+PluginContextTy::~PluginContextTy() = default;
+
+MemoryManagerTy *
+PluginContextTy::getDeviceMemoryManagerFor(GenericDeviceTy &Device,
+                                           TargetAllocTy Kind) {
+  if (Device.OMPX_TrackAllocationTraces)
+    return nullptr;
+
+  if (Kind == TARGET_ALLOC_DEFAULT)
+    Kind = TARGET_ALLOC_DEVICE;
+  assert((Kind == TARGET_ALLOC_DEVICE || Kind == TARGET_ALLOC_SHARED) &&
+         "host allocations are not device-bound");
+
+  size_t Threshold;
+  if (Kind == TARGET_ALLOC_DEVICE) {
+    auto [EnvThreshold, EnableMM] = MemoryManagerTy::getSizeThresholdFromEnv();
+    if (!EnableMM)
+      return nullptr;
+    Threshold =
+        EnvThreshold ? EnvThreshold : Device.getMemoryManagerSizeThreshold();
+  } else {
+    Threshold = MemoryManagerTy::DefaultSizeThreshold;
+  }
+
+  std::pair<GenericDeviceTy *, int> Key{&Device, static_cast<int>(Kind)};
+
+  std::lock_guard<std::mutex> Lock(MemoryManagersMutex);
+  auto It = DeviceMemoryManagers.find(Key);
+  if (It != DeviceMemoryManagers.end())
+    return It->second.get();
+
+  auto Manager = std::make_unique<MemoryManagerTy>(Device, Threshold, Kind);
+  auto *Raw = Manager.get();
+  DeviceMemoryManagers[Key] = std::move(Manager);
+  return Raw;
+}
+
+MemoryManagerTy *PluginContextTy::getHostMemoryManager() {
+  if (Devices.empty())
+    return nullptr;
+  if (Devices.front()->OMPX_TrackAllocationTraces)
+    return nullptr;
+
+  std::lock_guard<std::mutex> Lock(MemoryManagersMutex);
+  if (HostMemoryManager)
+    return HostMemoryManager.get();
+
+  HostMemoryManager = std::make_unique<MemoryManagerTy>(
+      *Devices.front(), MemoryManagerTy::DefaultSizeThreshold,
+      TARGET_ALLOC_HOST);
+  return HostMemoryManager.get();
+}
+
+Expected<void *> PluginContextTy::allocate(GenericDeviceTy &Device,
+                                           int64_t Size, void *HostPtr,
+                                           TargetAllocTy Kind,
+                                           size_t Alignment) {
+  MemoryManagerTy *MM = (Kind == TARGET_ALLOC_HOST)
+                            ? getHostMemoryManager()
+                            : getDeviceMemoryManagerFor(Device, Kind);
+  if (MM)
+    return MM->allocate(Size, HostPtr, Alignment);
+  return Device.dataAlloc(Size, HostPtr, Kind, Alignment);
+}
+
+Error PluginContextTy::deallocate(void *Ptr) {
+  assert(!Devices.empty() && "context constructed without devices");
+  auto InfoOrErr = getAllocInfo(Ptr);
+  if (!InfoOrErr)
+    return InfoOrErr.takeError();
+  GenericDeviceTy *OwnerDevice = InfoOrErr->Device;
+  if (!OwnerDevice)
+    OwnerDevice = Devices.front();
+  return deallocate(*OwnerDevice, Ptr, InfoOrErr->Kind);
+}
+
+Error PluginContextTy::deallocate(GenericDeviceTy &Device, void *Ptr,
+                                  TargetAllocTy Kind) {
+  MemoryManagerTy *MM = (Kind == TARGET_ALLOC_HOST)
+                            ? getHostMemoryManager()
+                            : getDeviceMemoryManagerFor(Device, Kind);
+  if (MM)
+    return MM->free(Ptr);
+  return Device.dataDelete(Ptr, Kind);
+}
+
+PluginContextTy &
+GenericPluginTy::getDefaultContext(GenericDeviceTy & /*Device*/) {
+  assert(DefaultContext && "default context not initialized");
+  return *DefaultContext;
+}
+
+Expected<std::unique_ptr<PluginContextTy>>
+GenericPluginTy::createDefaultPluginContext() {
+  return std::make_unique<DefaultPluginContextTy>(*this);
+}
+
 Error GenericDeviceTy::enqueueHostCall(void (*Callback)(void *), void *UserData,
                                        __tgt_async_info *AsyncInfo) {
   AsyncInfoWrapperTy AsyncInfoWrapper(*this, AsyncInfo);
@@ -1341,12 +1393,20 @@ Error GenericPluginTy::init() {
   RPCServer = new RPCServerTy(*this);
   assert(RPCServer && "Invalid RPC server");
 
+  auto DefaultCtxOrErr = createDefaultPluginContext();
+  if (!DefaultCtxOrErr)
+    return DefaultCtxOrErr.takeError();
+  DefaultContext = std::move(*DefaultCtxOrErr);
+
   return Plugin::success();
 }
 
 Error GenericPluginTy::deinit() {
   assert(Initialized && "Plugin was not initialized!");
 
+  // Release context-held resources before the devices that back them.
+  DefaultContext.reset();
+
   // Deinitialize all active devices.
   for (int32_t DeviceId = 0; DeviceId < NumDevices; ++DeviceId) {
     if (Devices[DeviceId]) {
@@ -1570,29 +1630,28 @@ int32_t GenericPluginTy::load_binary(int32_t DeviceId,
 
 void *GenericPluginTy::data_alloc(int32_t DeviceId, int64_t Size, void *HostPtr,
                                   int32_t Kind) {
-  auto AllocOrErr = getDevice(DeviceId).dataAlloc(
-      Size, HostPtr, (TargetAllocTy)Kind, /*Alignment=*/0);
+  auto &Device = getDevice(DeviceId);
+  auto AllocOrErr = getDefaultContext(Device).allocate(
+      Device, Size, HostPtr, static_cast<TargetAllocTy>(Kind),
+      /*Alignment=*/0);
   if (!AllocOrErr) {
-    auto Err = AllocOrErr.takeError();
     REPORT() << "Failure to allocate device memory: "
-             << toString(std::move(Err));
+             << toString(AllocOrErr.takeError());
     return nullptr;
   }
   assert(*AllocOrErr && "Null pointer upon successful allocation");
-
   return *AllocOrErr;
 }
 
 int32_t GenericPluginTy::data_delete(int32_t DeviceId, void *TgtPtr,
                                      int32_t Kind) {
-  auto Err =
-      getDevice(DeviceId).dataDelete(TgtPtr, static_cast<TargetAllocTy>(Kind));
-  if (Err) {
+  auto &Device = getDevice(DeviceId);
+  if (auto Err = getDefaultContext(Device).deallocate(
+          Device, TgtPtr, static_cast<TargetAllocTy>(Kind))) {
     REPORT() << "Failure to deallocate device pointer " << TgtPtr << ": "
              << toString(std::move(Err));
     return OFFLOAD_FAIL;
   }
-
   return OFFLOAD_SUCCESS;
 }
 
diff --git a/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.cpp b/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.cpp
index e3a854d5690c1..4c7c22da2f046 100644
--- a/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.cpp
+++ b/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.cpp
@@ -70,6 +70,7 @@ DLWRAP(cuMemFreeAsync, 2)
 
 DLWRAP(cuMemPrefetchAsync, 4)
 DLWRAP(cuPointerGetAttribute, 3)
+DLWRAP(cuPointerGetAttributes, 4)
 
 DLWRAP(cuModuleGetFunction, 3)
 DLWRAP(cuModuleGetGlobal, 4)
diff --git a/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.h b/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.h
index a7524d417dded..50902f8116149 100644
--- a/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.h
+++ b/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.h
@@ -481,11 +481,27 @@ CUresult cuMemFreeAsync(CUdeviceptr, CUstream);
 
 CUresult cuMemPrefetchAsync(CUdeviceptr, size_t, CUdevice, CUstream);
 
+typedef enum CUmemorytype_enum {
+  CU_MEMORYTYPE_HOST = 0x01,
+  CU_MEMORYTYPE_DEVICE = 0x02,
+  CU_MEMORYTYPE_ARRAY = 0x03,
+  CU_MEMORYTYPE_UNIFIED = 0x04,
+} CUmemorytype;
+
 typedef enum CUpointer_attribute_enum {
-  CU_POINTER_ATTRIBUTE_IS_MANAGED = 8
+  CU_POINTER_ATTRIBUTE_CONTEXT = 1,
+  CU_POINTER_ATTRIBUTE_MEMORY_TYPE = 2,
+  CU_POINTER_ATTRIBUTE_DEVICE_POINTER = 3,
+  CU_POINTER_ATTRIBUTE_HOST_POINTER = 4,
+  CU_POINTER_ATTRIBUTE_IS_MANAGED = 8,
+  CU_POINTER_ATTRIBUTE_DEVICE_ORDINAL = 9,
+  CU_POINTER_ATTRIBUTE_RANGE_START_ADDR = 11,
+  CU_POINTER_ATTRIBUTE_RANGE_SIZE = 12,
 } CUpointer_attribute;
 
 CUresult cuPointerGetAttribute(void *, CUpointer_attribute, CUdeviceptr);
+CUresult cuPointerGetAttributes(unsigned int, CUpointer_attribute *, void **,
+                                CUdeviceptr);
 
 CUresult cuModuleGetFunction(CUfunction *, CUmodule, const char *);
 CUresult cuModuleGetGlobal(CUdeviceptr *, size_t *, CUmodule, const char *);
diff --git a/offload/plugins-nextgen/cuda/src/rtl.cpp b/offload/plugins-nextgen/cuda/src/rtl.cpp
index 8f1a4b29eef8b..2dc63a2bea8d3 100644
--- a/offload/plugins-nextgen/cuda/src/rtl.cpp
+++ b/offload/plugins-nextgen/cuda/src/rtl.cpp
@@ -1674,6 +1674,61 @@ struct CUDAPluginContextTy final : public PluginContextTy {
     CUstream Stream;
     return CUDADevice.getStream(AsyncInfoWrapper, Stream);
   }
+
+  Expected<PluginAllocInfoTy> getAllocInfo(const void *Ptr) override {
+    if (Devices.empty())
+      return Plugin::error(error::ErrorCode::NOT_FOUND,
+                           "pointer is not a known allocation in this context");
+
+    // Any device in the context can service the query; use the first.
+    auto &Ctx0 = static_cast<CUDADeviceTy &>(*Devices.front());
+    if (auto Err = Ctx0.setContext())
+      return std::move(Err);
+
+    CUdeviceptr CUPtr = reinterpret_cast<CUdeviceptr>(Ptr);
+
+    CUpointer_attribute Attrs[] = {
+        CU_POINTER_ATTRIBUTE_MEMORY_TYPE,
+        CU_POINTER_ATTRIBUTE_IS_MANAGED,
+        CU_POINTER_ATTRIBUTE_DEVICE_ORDINAL,
+        CU_POINTER_ATTRIBUTE_RANGE_START_ADDR,
+        CU_POINTER_ATTRIBUTE_RANGE_SIZE,
+    };
+    unsigned MemType = 0;
+    int IsManaged = 0;
+    int Ordinal = -1;
+    CUdeviceptr RangeStart = 0;
+    size_t RangeSize = 0;
+    void *Data[] = {&MemType, &IsManaged, &Ordinal, &RangeStart, &RangeSize};
+    if (CUresult Res = cuPointerGetAttributes(sizeof(Attrs) / sizeof(Attrs[0]),
+                                              Attrs, Data, CUPtr))
+      return Plugin::error(error::ErrorCode::NOT_FOUND,
+                           "cuPointerGetAttributes failed: %d", Res);
+
+    TargetAllocTy Kind = TARGET_ALLOC_DEVICE;
+    if (IsManaged)
+      Kind = TARGET_ALLOC_SHARED;
+    else if (MemType == CU_MEMORYTYPE_HOST)
+      Kind = TARGET_ALLOC_HOST;
+
+    // Ordinal is the CUDA driver ordinal (matches CUdevice); compare against
+    // that rather than the offload-side device index, which can differ under
+    // CUDA_VISIBLE_DEVICES or plugin-side device filtering.
+    GenericDeviceTy *OwnerDevice = nullptr;
+    for (auto *D : Devices) {
+      auto &CD = static_cast<CUDADeviceTy &>(*D);
+      if (static_cast<int>(CD.getCUDADevice()) == Ordinal) {
+        OwnerDevice = D;
+        break;
+      }
+    }
+    if (!OwnerDevice)
+      return Plugin::error(error::ErrorCode::NOT_FOUND,
+                           "allocation owner is not a device of this context");
+
+    return PluginAllocInfoTy{OwnerDevice, Kind,
+                             reinterpret_cast<void *>(RangeStart), RangeSize};
+  }
 };
 
 /// Class implementing the CUDA-specific functionalities of the plugin.
diff --git a/offload/plugins-nextgen/host/src/rtl.cpp b/offload/plugins-nextgen/host/src/rtl.cpp
index 55ada2f82c360..a594cf0fbc6ed 100644
--- a/offload/plugins-nextgen/host/src/rtl.cpp
+++ b/offload/plugins-nextgen/host/src/rtl.cpp
@@ -12,6 +12,8 @@
 
 #include <cassert>
 #include <cstddef>
+#include <map>
+#include <mutex>
 #include <string>
 #include <unordered_map>
 
@@ -488,6 +490,53 @@ struct GenELF64PluginContextTy final : public PluginContextTy {
   Error initAsyncInfoImpl(GenericDeviceTy &, AsyncInfoWrapperTy &) override {
     return Plugin::success();
   }
+
+  Expected<void *> allocate(GenericDeviceTy &Device, int64_t Size,
+                            void *HostPtr, TargetAllocTy Kind,
+                            size_t Alignment) override {
+    auto PtrOrErr =
+        PluginContextTy::allocate(Device, Size, HostPtr, Kind, Alignment);
+    if (!PtrOrErr)
+      return PtrOrErr.takeError();
+    void *Ptr = *PtrOrErr;
+    std::lock_guard<std::mutex> Lock(AllocsMutex);
+    Allocs[Ptr] = Entry{&Device, Kind, static_cast<size_t>(Size)};
+    return Ptr;
+  }
+
+  Error deallocate(GenericDeviceTy &Device, void *Ptr,
+                   TargetAllocTy Kind) override {
+    {
+      std::lock_guard<std::mutex> Lock(AllocsMutex);
+      Allocs.erase(Ptr);
+    }
+    return PluginContextTy::deallocate(Device, Ptr, Kind);
+  }
+
+  Expected<PluginAllocInfoTy> getAllocInfo(const void *Ptr) override {
+    std::lock_guard<std::mutex> Lock(AllocsMutex);
+    auto It = Allocs.upper_bound(const_cast<void *>(Ptr));
+    if (It == Allocs.begin())
+      return Plugin::error(error::ErrorCode::NOT_FOUND,
+                           "pointer is not a known allocation in this context");
+    --It;
+    void *Base = It->first;
+    const Entry &E = It->second;
+    if (reinterpret_cast<const char *>(Ptr) >=
+        reinterpret_cast<char *>(Base) + E.Size)
+      return Plugin::error(error::ErrorCode::NOT_FOUND,
+                           "pointer is not a known allocation in this context");
+    return PluginAllocInfoTy{E.Device, E.Kind, Base, E.Size};
+  }
+
+private:
+  struct Entry {
+    GenericDeviceTy *Device;
+    TargetAllocTy Kind;
+    size_t Size;
+  };
+  std::map<void *, Entry> Allocs;
+  std::mutex AllocsMutex;
 };
 
 /// Class implementing the plugin functionalities for GenELF64.
diff --git a/offload/plugins-nextgen/level_zero/include/L0Memory.h b/offload/plugins-nextgen/level_zero/include/L0Memory.h
index 8f2f3422b770e..db7623afbb63e 100644
--- a/offload/plugins-nextgen/level_zero/include/L0Memory.h
+++ b/offload/plugins-nextgen/level_zero/include/L0Memory.h
@@ -250,13 +250,20 @@ class MemAllocatorTy {
     /// Remove allocation information for the given memory location.
     bool remove(void *Ptr, MemAllocInfoTy *Removed = nullptr);
 
-    /// Finds allocation information for the given memory location.
+    /// Finds allocation information for the given memory location. Ptr may
+    /// point anywhere inside the allocation.
     const MemAllocInfoTy *find(void *Ptr) const {
-      auto AllocInfo = Map.find(Ptr);
-      if (AllocInfo == Map.end())
+      if (Map.empty())
         return nullptr;
-      else
-        return &AllocInfo->second;
+      auto I = Map.upper_bound(Ptr);
+      if (I == Map.begin())
+        return nullptr;
+      --I;
+      uintptr_t PtrAsInt = reinterpret_cast<uintptr_t>(Ptr);
+      uintptr_t Base = reinterpret_cast<uintptr_t>(I->first);
+      if (PtrAsInt >= Base + I->second.ReqSize)
+        return nullptr;
+      return &I->second;
     }
 
     /// Check if the map contains the given pointer and offset.
@@ -287,6 +294,11 @@ class MemAllocatorTy {
 
   /// L0 context to use.
   const L0ContextTy *L0Context = nullptr;
+  /// ze_context used for allocations. Normally matches
+  /// L0Context->getZeContext(), but for pools owned by a user-created
+  /// plugin context this holds that context's ze_context so memory ends
+  /// up in the ze_context the caller's queues use.
+  ze_context_handle_t ZeContext = nullptr;
   /// L0 device to use.
   L0DeviceTy *Device = nullptr;
   /// Whether the device supports large memory allocation.
@@ -377,8 +389,11 @@ class MemAllocatorTy {
   MemAllocatorTy &operator=(const MemAllocatorTy &&) = delete;
   ~MemAllocatorTy() = default;
 
-  Error initDevicePools(L0DeviceTy &L0Device, const L0OptionsTy &Option);
-  Error initHostPool(L0ContextTy &Driver, const L0OptionsTy &Option);
+  Error initDevicePools(L0DeviceTy &L0Device, const L0OptionsTy &Option,
+                        ze_context_handle_t ZeCtx);
+  Error initHostPool(L0ContextTy &Driver, const L0OptionsTy &Option,
+                     ze_context_handle_t ZeCtx);
+  ze_context_handle_t getZeContext() const { return ZeContext; }
   void updateMaxAllocSize(L0DeviceTy &L0Device);
 
   /// Release resources and report statistics if requested.
diff --git a/offload/plugins-nextgen/level_zero/include/L0Plugin.h b/offload/plugins-nextgen/level_zero/include/L0Plugin.h
index 82e47e052f339..d6656cb4cf792 100644
--- a/offload/plugins-nextgen/level_zero/include/L0Plugin.h
+++ b/offload/plugins-nextgen/level_zero/include/L0Plugin.h
@@ -42,6 +42,17 @@ class LevelZeroPluginContextTy final : public PluginContextTy {
   Error initAsyncInfoImpl(GenericDeviceTy &Device,
                           AsyncInfoWrapperTy &AsyncInfoWrapper) override;
 
+  llvm::Expected<void *> allocate(GenericDeviceTy &Device, int64_t Size,
+                                  void *HostPtr, TargetAllocTy Kind,
+                                  size_t Alignment) override;
+  llvm::Error deallocate(GenericDeviceTy &Device, void *Ptr,
+                         TargetAllocTy Kind) override;
+  Expected<PluginAllocInfoTy> getAllocInfo(const void *Ptr) override;
+
+  /// Initialize per-plugin-context memory allocators. Runs the pool
+  /// probe L0 calls up-front so the first user allocation is not delayed.
+  Error initAllocators();
+
   /// Pop an idle queue for \p Device from the cache, or create a new one.
   Expected<L0QueueTy *> takeCachedQueue(L0DeviceTy *Device) {
     return QueueCache.getQueue(*Device);
@@ -56,6 +67,11 @@ class LevelZeroPluginContextTy final : public PluginContextTy {
   bool OwnsZeContext;
 
   L0QueueCacheTy QueueCache;
+
+  /// Per-plugin-context allocators; scoped to this context's ze_context.
+  llvm::DenseMap<L0DeviceTy *, std::unique_ptr<MemAllocatorTy>>
+      DeviceAllocators;
+  std::unique_ptr<MemAllocatorTy> HostAllocator;
 };
 
 /// Class implementing the LevelZero specific functionalities of the plugin.
diff --git a/offload/plugins-nextgen/level_zero/src/L0Context.cpp b/offload/plugins-nextgen/level_zero/src/L0Context.cpp
index 13afa04887689..5899611d6c7d5 100644
--- a/offload/plugins-nextgen/level_zero/src/L0Context.cpp
+++ b/offload/plugins-nextgen/level_zero/src/L0Context.cpp
@@ -76,7 +76,8 @@ Error L0ContextTy::init() {
     CleanupOnError();
     return Err;
   }
-  if (auto Err = HostMemAllocator.initHostPool(*this, Plugin.getOptions())) {
+  if (auto Err = HostMemAllocator.initHostPool(*this, Plugin.getOptions(),
+                                               zeContext)) {
     if (auto DeinitErr = EventPool.deinit())
       Err = joinErrors(std::move(Err), std::move(DeinitErr));
     CleanupOnError();
diff --git a/offload/plugins-nextgen/level_zero/src/L0Device.cpp b/offload/plugins-nextgen/level_zero/src/L0Device.cpp
index f6900731ce80f..e41268d7eb843 100644
--- a/offload/plugins-nextgen/level_zero/src/L0Device.cpp
+++ b/offload/plugins-nextgen/level_zero/src/L0Device.cpp
@@ -197,7 +197,8 @@ Error L0DeviceTy::initImpl(GenericPluginTy &Plugin) {
     return QueueGroupInfoOrErr.takeError();
   QueueConfig = *QueueGroupInfoOrErr;
 
-  if (auto Err = MemAllocator.initDevicePools(*this, Options))
+  if (auto Err = MemAllocator.initDevicePools(*this, Options,
+                                              L0Context.getZeContext()))
     return Err;
   L0Context.getHostMemAllocator().updateMaxAllocSize(*this);
   reportDeviceInfo();
diff --git a/offload/plugins-nextgen/level_zero/src/L0Memory.cpp b/offload/plugins-nextgen/level_zero/src/L0Memory.cpp
index c8098d37b2ef2..4dbebffb1fcf5 100644
--- a/offload/plugins-nextgen/level_zero/src/L0Memory.cpp
+++ b/offload/plugins-nextgen/level_zero/src/L0Memory.cpp
@@ -76,7 +76,7 @@ Error MemAllocatorTy::MemPoolTy::init(int32_t Kind, MemAllocatorTy *AllocatorIn,
   PoolSizeMax = UserPoolSize << 20; // Covert MB to B.
   PoolSize = 0;
 
-  auto Context = Allocator->L0Context->getZeContext();
+  auto Context = Allocator->ZeContext;
   const auto Device = Allocator->Device;
 
   // Check page size used for this allocation kind to decide minimum.
@@ -380,11 +380,13 @@ bool MemAllocatorTy::MemAllocInfoMapTy::remove(void *Ptr,
 }
 
 Error MemAllocatorTy::initDevicePools(L0DeviceTy &L0Device,
-                                      const L0OptionsTy &Options) {
+                                      const L0OptionsTy &Options,
+                                      ze_context_handle_t ZeCtx) {
   SupportsLargeMem = L0Device.supportsLargeMem();
   IsHostMem = false;
   Device = &L0Device;
   L0Context = &L0Device.getL0Context();
+  ZeContext = ZeCtx;
   for (auto Kind : {TARGET_ALLOC_DEVICE, TARGET_ALLOC_SHARED}) {
     if (Options.MemPoolConfig[Kind].Use) {
       std::lock_guard<std::mutex> Lock(Mtx);
@@ -404,10 +406,12 @@ Error MemAllocatorTy::initDevicePools(L0DeviceTy &L0Device,
 }
 
 Error MemAllocatorTy::initHostPool(L0ContextTy &Driver,
-                                   const L0OptionsTy &Option) {
+                                   const L0OptionsTy &Option,
+                                   ze_context_handle_t ZeCtx) {
   SupportsLargeMem = Driver.supportsLargeMem();
   IsHostMem = true;
   L0Context = &Driver;
+  ZeContext = ZeCtx;
   if (Option.MemPoolConfig[TARGET_ALLOC_HOST].Use) {
     std::lock_guard<std::mutex> Lock(Mtx);
     Pools[TARGET_ALLOC_HOST] = std::make_unique<MemPoolTy>();
@@ -654,7 +658,7 @@ Expected<void *> MemAllocatorTy::allocFromL0(size_t Size, size_t Align,
   }
 
   auto zeDevice = Device ? Device->getZeDevice() : nullptr;
-  auto zeContext = L0Context->getZeContext();
+  auto zeContext = ZeContext;
   bool MakeResident = false;
   switch (Kind) {
   case TARGET_ALLOC_DEVICE:
@@ -690,7 +694,7 @@ Expected<void *> MemAllocatorTy::allocFromL0(size_t Size, size_t Align,
 }
 
 Error MemAllocatorTy::deallocFromL0(void *Ptr) {
-  CALL_ZE_RET_ERROR(zeMemFree, L0Context->getZeContext(), Ptr);
+  CALL_ZE_RET_ERROR(zeMemFree, ZeContext, Ptr);
   ODBG(OLDT_Alloc) << "Freed device pointer " << Ptr;
   return Plugin::success();
 }
diff --git a/offload/plugins-nextgen/level_zero/src/L0Plugin.cpp b/offload/plugins-nextgen/level_zero/src/L0Plugin.cpp
index 25105711b4904..f9cf3b7eb59cf 100644
--- a/offload/plugins-nextgen/level_zero/src/L0Plugin.cpp
+++ b/offload/plugins-nextgen/level_zero/src/L0Plugin.cpp
@@ -265,6 +265,16 @@ Error LevelZeroPluginContextTy::initAsyncInfoImpl(
 Error LevelZeroPluginContextTy::deinit() {
   if (auto Err = QueueCache.deinit())
     return Err;
+  // Tear down allocators before their ze_context.
+  for (auto &KV : DeviceAllocators)
+    if (auto Err = KV.second->deinit())
+      return Err;
+  DeviceAllocators.clear();
+  if (HostAllocator) {
+    if (auto Err = HostAllocator->deinit())
+      return Err;
+    HostAllocator.reset();
+  }
   if (OwnsZeContext && ZeContext) {
     CALL_ZE_RET_ERROR(zeContextDestroy, ZeContext);
     ZeContext = nullptr;
@@ -273,6 +283,91 @@ Error LevelZeroPluginContextTy::deinit() {
   return Plugin::success();
 }
 
+Error LevelZeroPluginContextTy::initAllocators() {
+  const auto &Options = static_cast<LevelZeroPluginTy &>(Plugin).getOptions();
+  for (auto *D : Devices) {
+    auto &L0Device = static_cast<L0DeviceTy &>(*D);
+    auto Alloc = std::make_unique<MemAllocatorTy>();
+    if (auto Err = Alloc->initDevicePools(L0Device, Options, ZeContext))
+      return Err;
+    DeviceAllocators.try_emplace(&L0Device, std::move(Alloc));
+  }
+  if (Devices.empty())
+    return Plugin::success();
+  auto &First = static_cast<L0DeviceTy &>(*Devices.front());
+  HostAllocator = std::make_unique<MemAllocatorTy>();
+  if (auto Err =
+          HostAllocator->initHostPool(First.getL0Context(), Options, ZeContext))
+    return Err;
+  // Host MaxAllocSize = min over devices, matching L0ContextTy's driver pool.
+  for (auto *D : Devices)
+    HostAllocator->updateMaxAllocSize(static_cast<L0DeviceTy &>(*D));
+  return Plugin::success();
+}
+
+Expected<void *> LevelZeroPluginContextTy::allocate(GenericDeviceTy &Device,
+                                                    int64_t Size,
+                                                    void * /*HostPtr*/,
+                                                    TargetAllocTy Kind,
+                                                    size_t Alignment) {
+  MemAllocatorTy *Allocator = nullptr;
+  int32_t ResolvedKind = Kind;
+  if (Kind == TARGET_ALLOC_HOST) {
+    if (!HostAllocator)
+      return Plugin::error(ErrorCode::INVALID_ARGUMENT,
+                           "host allocator not initialized");
+    Allocator = HostAllocator.get();
+  } else {
+    if (ResolvedKind == TARGET_ALLOC_DEFAULT)
+      ResolvedKind = TARGET_ALLOC_DEVICE;
+    auto &L0Device = static_cast<L0DeviceTy &>(Device);
+    auto It = DeviceAllocators.find(&L0Device);
+    if (It == DeviceAllocators.end())
+      return Plugin::error(ErrorCode::INVALID_DEVICE,
+                           "device is not part of this context");
+    Allocator = It->second.get();
+  }
+  return Allocator->alloc(Size, Alignment, ResolvedKind, /*Offset=*/0,
+                          /*UserAlloc=*/true, /*DevMalloc=*/false,
+                          /*MemAdvice=*/
+                          std::numeric_limits<uint32_t>::max(),
+                          AllocOptionTy::ALLOC_OPT_NONE);
+}
+
+Error LevelZeroPluginContextTy::deallocate(GenericDeviceTy &Device, void *Ptr,
+                                           TargetAllocTy Kind) {
+  if (Kind == TARGET_ALLOC_HOST) {
+    if (!HostAllocator)
+      return Plugin::error(ErrorCode::NOT_FOUND,
+                           "no host allocation tracked in this context");
+    return HostAllocator->dealloc(Ptr);
+  }
+  auto &L0Device = static_cast<L0DeviceTy &>(Device);
+  auto It = DeviceAllocators.find(&L0Device);
+  if (It == DeviceAllocators.end())
+    return Plugin::error(ErrorCode::INVALID_DEVICE,
+                         "device is not part of this context");
+  return It->second->dealloc(Ptr);
+}
+
+Expected<PluginAllocInfoTy>
+LevelZeroPluginContextTy::getAllocInfo(const void *Ptr) {
+  void *Raw = const_cast<void *>(Ptr);
+  for (auto &KV : DeviceAllocators) {
+    if (auto *Info = KV.second->getAllocInfo(Raw))
+      return PluginAllocInfoTy{KV.first, static_cast<TargetAllocTy>(Info->Kind),
+                               Info->Base, Info->ReqSize};
+  }
+  if (HostAllocator) {
+    if (auto *Info = HostAllocator->getAllocInfo(Raw))
+      return PluginAllocInfoTy{nullptr, static_cast<TargetAllocTy>(Info->Kind),
+                               Info->Base, Info->ReqSize};
+  }
+
+  return Plugin::error(error::ErrorCode::NOT_FOUND,
+                       "pointer is not a known allocation in this context");
+}
+
 Expected<std::unique_ptr<PluginContextTy>>
 LevelZeroPluginTy::createPluginContext(
     llvm::ArrayRef<GenericDeviceTy *> Devices) {
@@ -314,8 +409,11 @@ LevelZeroPluginTy::createPluginContext(
     OwnsZeContext = true;
   }
 
-  return std::make_unique<LevelZeroPluginContextTy>(*this, Devices, Driver,
-                                                    ZeContext, OwnsZeContext);
+  auto Ctx = std::make_unique<LevelZeroPluginContextTy>(
+      *this, Devices, Driver, ZeContext, OwnsZeContext);
+  if (auto Err = Ctx->initAllocators())
+    return std::move(Err);
+  return Ctx;
 }
 
 } // namespace llvm::omp::target::plugin
diff --git a/offload/unittests/Conformance/include/mathtest/DeviceContext.hpp b/offload/unittests/Conformance/include/mathtest/DeviceContext.hpp
index 2e23cd3f4b19a..883f4f1d778a2 100644
--- a/offload/unittests/Conformance/include/mathtest/DeviceContext.hpp
+++ b/offload/unittests/Conformance/include/mathtest/DeviceContext.hpp
@@ -40,7 +40,8 @@ const llvm::SetVector<llvm::StringRef> &getPlatforms();
 
 namespace detail {
 
-void allocManagedMemory(ol_device_handle_t DeviceHandle, std::size_t Size,
+void allocManagedMemory(ol_context_handle_t Context,
+                        ol_device_handle_t DeviceHandle, std::size_t Size,
                         void **AllocationOut) noexcept;
 } // namespace detail
 
@@ -62,10 +63,11 @@ class DeviceContext {
   ManagedBuffer<T> createManagedBuffer(std::size_t Size) const noexcept {
     void *UntypedAddress = nullptr;
 
-    detail::allocManagedMemory(DeviceHandle, Size * sizeof(T), &UntypedAddress);
+    detail::allocManagedMemory(Context, DeviceHandle, Size * sizeof(T),
+                               &UntypedAddress);
     T *TypedAddress = static_cast<T *>(UntypedAddress);
 
-    return ManagedBuffer<T>(TypedAddress, Size);
+    return ManagedBuffer<T>(Context, TypedAddress, Size);
   }
 
   [[nodiscard]] llvm::Expected<std::shared_ptr<DeviceImage>>
diff --git a/offload/unittests/Conformance/include/mathtest/DeviceResources.hpp b/offload/unittests/Conformance/include/mathtest/DeviceResources.hpp
index 860448afa3a01..00bfebacd9048 100644
--- a/offload/unittests/Conformance/include/mathtest/DeviceResources.hpp
+++ b/offload/unittests/Conformance/include/mathtest/DeviceResources.hpp
@@ -29,7 +29,7 @@ class DeviceContext;
 
 namespace detail {
 
-void freeDeviceMemory(void *Address) noexcept;
+void freeDeviceMemory(ol_context_handle_t Context, void *Address) noexcept;
 } // namespace detail
 
 //===----------------------------------------------------------------------===//
@@ -40,14 +40,14 @@ template <typename T> class [[nodiscard]] ManagedBuffer {
 public:
   ~ManagedBuffer() noexcept {
     if (Address)
-      detail::freeDeviceMemory(Address);
+      detail::freeDeviceMemory(Context, Address);
   }
 
   ManagedBuffer(const ManagedBuffer &) = delete;
   ManagedBuffer &operator=(const ManagedBuffer &) = delete;
 
   ManagedBuffer(ManagedBuffer &&Other) noexcept
-      : Address(Other.Address), Size(Other.Size) {
+      : Context(Other.Context), Address(Other.Address), Size(Other.Size) {
     Other.Address = nullptr;
     Other.Size = 0;
   }
@@ -57,8 +57,9 @@ template <typename T> class [[nodiscard]] ManagedBuffer {
       return *this;
 
     if (Address)
-      detail::freeDeviceMemory(Address);
+      detail::freeDeviceMemory(Context, Address);
 
+    Context = Other.Context;
     Address = Other.Address;
     Size = Other.Size;
 
@@ -85,9 +86,11 @@ template <typename T> class [[nodiscard]] ManagedBuffer {
 private:
   friend class DeviceContext;
 
-  explicit ManagedBuffer(T *Address, std::size_t Size) noexcept
-      : Address(Address), Size(Size) {}
+  explicit ManagedBuffer(ol_context_handle_t Context, T *Address,
+                         std::size_t Size) noexcept
+      : Context(Context), Address(Address), Size(Size) {}
 
+  ol_context_handle_t Context = nullptr;
   T *Address = nullptr;
   std::size_t Size = 0;
 };
diff --git a/offload/unittests/Conformance/lib/DeviceContext.cpp b/offload/unittests/Conformance/lib/DeviceContext.cpp
index d81f5f0c5867a..62b265043a2eb 100644
--- a/offload/unittests/Conformance/lib/DeviceContext.cpp
+++ b/offload/unittests/Conformance/lib/DeviceContext.cpp
@@ -154,11 +154,12 @@ const llvm::SetVector<llvm::StringRef> &mathtest::getPlatforms() {
   return Platforms;
 }
 
-void detail::allocManagedMemory(ol_device_handle_t DeviceHandle,
+void detail::allocManagedMemory(ol_context_handle_t Context,
+                                ol_device_handle_t DeviceHandle,
                                 std::size_t Size,
                                 void **AllocationOut) noexcept {
-  OL_CHECK(
-      olMemAlloc(DeviceHandle, OL_ALLOC_TYPE_MANAGED, Size, AllocationOut));
+  OL_CHECK(olMemAlloc(Context, DeviceHandle, OL_ALLOC_TYPE_MANAGED, Size,
+                      AllocationOut));
 }
 
 //===----------------------------------------------------------------------===//
diff --git a/offload/unittests/Conformance/lib/DeviceResources.cpp b/offload/unittests/Conformance/lib/DeviceResources.cpp
index d1c7b90e751e6..8ef41d66d62cf 100644
--- a/offload/unittests/Conformance/lib/DeviceResources.cpp
+++ b/offload/unittests/Conformance/lib/DeviceResources.cpp
@@ -24,9 +24,10 @@ using namespace mathtest;
 // Helpers
 //===----------------------------------------------------------------------===//
 
-void detail::freeDeviceMemory(void *Address) noexcept {
+void detail::freeDeviceMemory(ol_context_handle_t Context,
+                              void *Address) noexcept {
   if (Address)
-    OL_CHECK(olMemFree(Address));
+    OL_CHECK(olMemFree(Context, Address));
 }
 
 //===----------------------------------------------------------------------===//
diff --git a/offload/unittests/OffloadAPI/event/olGetEventElapsedTime.cpp b/offload/unittests/OffloadAPI/event/olGetEventElapsedTime.cpp
index 356f04bfd9bdb..d1b25f5f669f7 100644
--- a/offload/unittests/OffloadAPI/event/olGetEventElapsedTime.cpp
+++ b/offload/unittests/OffloadAPI/event/olGetEventElapsedTime.cpp
@@ -28,13 +28,13 @@ struct olGetEventElapsedTimeTest : OffloadQueueTest {
     LaunchArgs.NumGroups = {1, 1, 1};
     LaunchArgs.DynSharedMemory = 0;
 
-    ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+    ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                               LaunchArgs.GroupSize.x * sizeof(uint32_t), &Mem));
   }
 
   void TearDown() override {
     if (Mem)
-      ASSERT_SUCCESS(olMemFree(Mem));
+      ASSERT_SUCCESS(olMemFree(Context, Mem));
     if (Program)
       ASSERT_SUCCESS(olDestroyProgram(Program));
     RETURN_ON_FATAL_FAILURE(OffloadQueueTest::TearDown());
diff --git a/offload/unittests/OffloadAPI/kernel/olLaunchKernel.cpp b/offload/unittests/OffloadAPI/kernel/olLaunchKernel.cpp
index 488021ded63be..f9cc29d1e9d0c 100644
--- a/offload/unittests/OffloadAPI/kernel/olLaunchKernel.cpp
+++ b/offload/unittests/OffloadAPI/kernel/olLaunchKernel.cpp
@@ -60,7 +60,7 @@ KERNEL_MULTI_TEST(Global, global, "write", "read")
 
 TEST_P(olLaunchKernelFooTest, Success) {
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(uint32_t), &Mem));
 
   void *ArgPtrs[] = {&Mem};
@@ -76,15 +76,16 @@ TEST_P(olLaunchKernelFooTest, Success) {
     ASSERT_EQ(Data[i], i);
   }
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 TEST_P(olLaunchKernelFooTest, SuccessThreaded) {
   threadify([&](size_t) {
     void *DevAlloc, *HstAlloc;
     size_t Size = LaunchArgs.GroupSize.x * sizeof(uint32_t);
-    ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &DevAlloc));
-    ASSERT_SUCCESS(olMemAllocHost(Device, Size, &HstAlloc));
+    ASSERT_SUCCESS(
+        olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &DevAlloc));
+    ASSERT_SUCCESS(olMemAllocHost(Context, Device, Size, &HstAlloc));
 
     void *ArgPtrs[] = {&DevAlloc};
     size_t ArgSizes[] = {sizeof(DevAlloc)};
@@ -101,8 +102,8 @@ TEST_P(olLaunchKernelFooTest, SuccessThreaded) {
       ASSERT_EQ(Data[i], i);
     }
 
-    ASSERT_SUCCESS(olMemFree(DevAlloc));
-    ASSERT_SUCCESS(olMemFree(HstAlloc));
+    ASSERT_SUCCESS(olMemFree(Context, DevAlloc));
+    ASSERT_SUCCESS(olMemFree(Context, HstAlloc));
   });
 }
 
@@ -115,7 +116,7 @@ TEST_P(olLaunchKernelNoArgsTest, Success) {
 
 TEST_P(olLaunchKernelMultiArgsTest, Success) {
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(int), &Mem));
 
   char A = 3;
@@ -134,7 +135,7 @@ TEST_P(olLaunchKernelMultiArgsTest, Success) {
   for (uint32_t i = 0; i < LaunchArgs.GroupSize.x; i++)
     ASSERT_EQ(Data[i], A + C + static_cast<int>(i));
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 struct Foo {
@@ -144,7 +145,7 @@ struct Foo {
 
 TEST_P(olLaunchKernelCompositeTest, Success) {
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(uint32_t), &Mem));
 
   uint8_t N = 1;
@@ -163,12 +164,12 @@ TEST_P(olLaunchKernelCompositeTest, Success) {
   for (uint32_t i = 0; i < LaunchArgs.GroupSize.x; i++)
     ASSERT_EQ(Data[i], N + F.a + F.b + i);
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 TEST_P(olLaunchKernelFooTest, SuccessSynchronous) {
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(uint32_t), &Mem));
 
   void *ArgPtrs[] = {&Mem};
@@ -182,7 +183,7 @@ TEST_P(olLaunchKernelFooTest, SuccessSynchronous) {
     ASSERT_EQ(Data[i], i);
   }
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 TEST_P(olLaunchKernelByteTest, Success) {
@@ -204,7 +205,7 @@ TEST_P(olLaunchKernelLocalMemTest, Success) {
   LaunchArgs.DynSharedMemory = 64 * sizeof(uint32_t);
 
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * LaunchArgs.NumGroups.x *
                                 sizeof(uint32_t),
                             &Mem));
@@ -221,7 +222,7 @@ TEST_P(olLaunchKernelLocalMemTest, Success) {
   for (uint32_t i = 0; i < LaunchArgs.GroupSize.x * LaunchArgs.NumGroups.x; i++)
     ASSERT_EQ(Data[i], (i % 64) * 2);
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 TEST_P(olLaunchKernelLocalMemReductionTest, Success) {
@@ -231,7 +232,7 @@ TEST_P(olLaunchKernelLocalMemReductionTest, Success) {
   LaunchArgs.DynSharedMemory = 64 * sizeof(uint32_t);
 
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.NumGroups.x * sizeof(uint32_t), &Mem));
 
   void *ArgPtrs[] = {&Mem};
@@ -246,7 +247,7 @@ TEST_P(olLaunchKernelLocalMemReductionTest, Success) {
   for (uint32_t i = 0; i < LaunchArgs.NumGroups.x; i++)
     ASSERT_EQ(Data[i], 2 * LaunchArgs.GroupSize.x);
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 TEST_P(olLaunchKernelLocalMemStaticTest, Success) {
@@ -254,7 +255,7 @@ TEST_P(olLaunchKernelLocalMemStaticTest, Success) {
   LaunchArgs.DynSharedMemory = 0;
 
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.NumGroups.x * sizeof(uint32_t), &Mem));
 
   void *ArgPtrs[] = {&Mem};
@@ -269,7 +270,7 @@ TEST_P(olLaunchKernelLocalMemStaticTest, Success) {
   for (uint32_t i = 0; i < LaunchArgs.NumGroups.x; i++)
     ASSERT_EQ(Data[i], 2 * LaunchArgs.GroupSize.x);
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 // The test intends to verify the correctness of the current implementation of
@@ -280,9 +281,10 @@ TEST_P(olLaunchKernelSingleCounterSyncEventTest, SuccessSyncEvent) {
 
   size_t Size = sizeof(uint32_t);
 
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size,
+                            &InitValuePassed));
   ASSERT_SUCCESS(
-      olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &InitValuePassed));
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &ResNum));
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &ResNum));
 
   uint32_t HostInitVal = 0;
   ASSERT_SUCCESS(
@@ -319,8 +321,8 @@ TEST_P(olLaunchKernelSingleCounterSyncEventTest, SuccessSyncEvent) {
 
   ASSERT_EQ(FinalResVal, NumberToAdd * LoopRange);
 
-  ASSERT_SUCCESS(olMemFree(InitValuePassed));
-  ASSERT_SUCCESS(olMemFree(ResNum));
+  ASSERT_SUCCESS(olMemFree(Context, InitValuePassed));
+  ASSERT_SUCCESS(olMemFree(Context, ResNum));
 }
 
 // The test checks the correctness of the synchronization between queues using
@@ -338,10 +340,12 @@ TEST_P(olLaunchKernelSingleCounterSyncEventTest, SuccessTwoQueues) {
 
   size_t Size = sizeof(uint32_t);
 
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size,
+                            &InitValuePassed));
   ASSERT_SUCCESS(
-      olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &InitValuePassed));
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &ResNum1));
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &ResNum2));
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &ResNum1));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &ResNum2));
 
   uint32_t HostInitVal = 0;
   ASSERT_SUCCESS(
@@ -388,14 +392,14 @@ TEST_P(olLaunchKernelSingleCounterSyncEventTest, SuccessTwoQueues) {
 
   ASSERT_EQ(FinalResVal, 2 * NumberToAdd * LoopRange);
 
-  ASSERT_SUCCESS(olMemFree(InitValuePassed));
-  ASSERT_SUCCESS(olMemFree(ResNum1));
-  ASSERT_SUCCESS(olMemFree(ResNum2));
+  ASSERT_SUCCESS(olMemFree(Context, InitValuePassed));
+  ASSERT_SUCCESS(olMemFree(Context, ResNum1));
+  ASSERT_SUCCESS(olMemFree(Context, ResNum2));
 }
 
 TEST_P(olLaunchKernelGlobalTest, Success) {
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(uint32_t), &Mem));
 
   void *ArgPtrs[] = {&Mem};
@@ -413,7 +417,7 @@ TEST_P(olLaunchKernelGlobalTest, Success) {
     ASSERT_EQ(Data[i], i * 2);
   }
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 TEST_P(olLaunchKernelGlobalTest, InvalidNotAKernel) {
@@ -429,7 +433,7 @@ TEST_P(olLaunchKernelGlobalCtorTest, Success) {
   SKIP_KNOWN_FAILURE(LevelZero{"unsupported feature"});
 
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(uint32_t), &Mem));
 
   void *ArgPtrs[] = {&Mem};
@@ -444,7 +448,7 @@ TEST_P(olLaunchKernelGlobalCtorTest, Success) {
     ASSERT_EQ(Data[i], i + 100);
   }
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 TEST_P(olLaunchKernelGlobalDtorTest, Success) {
@@ -458,8 +462,8 @@ TEST_P(olLaunchKernelGlobalDtorTest, Success) {
 
 TEST_P(olLaunchKernelGridSizeTest, Success) {
   void *Mem;
-  ASSERT_SUCCESS(
-      olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, 6 * sizeof(uint32_t), &Mem));
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
+                            6 * sizeof(uint32_t), &Mem));
 
   uint32_t *NumBlocks = static_cast<uint32_t *>(Mem);
   uint32_t *NumThreads = static_cast<uint32_t *>(Mem) + 3;
@@ -493,5 +497,5 @@ TEST_P(olLaunchKernelGridSizeTest, Success) {
     ASSERT_EQ(NumThreads[2], LaunchArgs.GroupSize.z);
   }
 
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
diff --git a/offload/unittests/OffloadAPI/memory/olGetMemInfo.cpp b/offload/unittests/OffloadAPI/memory/olGetMemInfo.cpp
index f7b192b08b88b..3fa35ea8eb45c 100644
--- a/offload/unittests/OffloadAPI/memory/olGetMemInfo.cpp
+++ b/offload/unittests/OffloadAPI/memory/olGetMemInfo.cpp
@@ -18,13 +18,13 @@ struct olGetMemInfoAllocTypeTest : OffloadDeviceTestWithParam<ol_alloc_type_t> {
         OffloadDeviceTestWithParam<ol_alloc_type_t>::SetUp());
     AllocType = getTestParam();
     if (AllocType == OL_ALLOC_TYPE_HOST)
-      ASSERT_SUCCESS(olMemAllocHost(Device, SIZE, &Ptr));
+      ASSERT_SUCCESS(olMemAllocHost(Context, Device, SIZE, &Ptr));
     else
-      ASSERT_SUCCESS(olMemAlloc(Device, AllocType, SIZE, &Ptr));
+      ASSERT_SUCCESS(olMemAlloc(Context, Device, AllocType, SIZE, &Ptr));
   }
 
   void TearDown() override {
-    ASSERT_SUCCESS(olMemFree(Ptr));
+    ASSERT_SUCCESS(olMemFree(Context, Ptr));
     RETURN_ON_FATAL_FAILURE(
         OffloadDeviceTestWithParam<ol_alloc_type_t>::TearDown());
   }
@@ -37,11 +37,12 @@ struct olGetMemInfoAllocTypeTest : OffloadDeviceTestWithParam<ol_alloc_type_t> {
 struct olGetMemInfoTest : OffloadDeviceTest {
   void SetUp() override {
     RETURN_ON_FATAL_FAILURE(OffloadDeviceTest::SetUp());
-    ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, SIZE, &Ptr));
+    ASSERT_SUCCESS(
+        olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, SIZE, &Ptr));
   }
 
   void TearDown() override {
-    ASSERT_SUCCESS(olMemFree(Ptr));
+    ASSERT_SUCCESS(olMemFree(Context, Ptr));
     RETURN_ON_FATAL_FAILURE(OffloadDeviceTest::TearDown());
   }
 
@@ -54,65 +55,91 @@ OFFLOAD_TESTS_INSTANTIATE_DEVICE_FIXTURE_WITH_PARAM(
 OFFLOAD_TESTS_INSTANTIATE_DEVICE_FIXTURE(olGetMemInfoTest);
 
 TEST_P(olGetMemInfoAllocTypeTest, SuccessDevice) {
+  // Host-pool allocations have no per-device affinity, so querying
+  // OL_MEM_INFO_DEVICE is invalid for them.
   ol_device_handle_t RetrievedDevice;
-  ASSERT_SUCCESS(olGetMemInfo(Ptr, OL_MEM_INFO_DEVICE, sizeof(RetrievedDevice),
-                              &RetrievedDevice));
+  if (AllocType == OL_ALLOC_TYPE_HOST) {
+    ASSERT_ERROR(OL_ERRC_INVALID_ARGUMENT,
+                 olGetMemInfo(Context, Ptr, OL_MEM_INFO_DEVICE,
+                              sizeof(RetrievedDevice), &RetrievedDevice));
+    return;
+  }
+  ASSERT_SUCCESS(olGetMemInfo(Context, Ptr, OL_MEM_INFO_DEVICE,
+                              sizeof(RetrievedDevice), &RetrievedDevice));
   ASSERT_EQ(RetrievedDevice, Device);
 }
 
 TEST_P(olGetMemInfoAllocTypeTest, SuccessBase) {
   void *RetrievedBase;
-  ASSERT_SUCCESS(olGetMemInfo(Ptr, OL_MEM_INFO_BASE, sizeof(RetrievedBase),
-                              &RetrievedBase));
+  ASSERT_SUCCESS(olGetMemInfo(Context, Ptr, OL_MEM_INFO_BASE,
+                              sizeof(RetrievedBase), &RetrievedBase));
   ASSERT_EQ(RetrievedBase, Ptr);
 }
 
 TEST_P(olGetMemInfoAllocTypeTest, SuccessSize) {
   size_t RetrievedSize;
-  ASSERT_SUCCESS(olGetMemInfo(Ptr, OL_MEM_INFO_SIZE, sizeof(RetrievedSize),
-                              &RetrievedSize));
+  ASSERT_SUCCESS(olGetMemInfo(Context, Ptr, OL_MEM_INFO_SIZE,
+                              sizeof(RetrievedSize), &RetrievedSize));
   ASSERT_EQ(RetrievedSize, SIZE);
 }
 
 TEST_P(olGetMemInfoAllocTypeTest, SuccessType) {
   ol_alloc_type_t RetrievedType;
-  ASSERT_SUCCESS(olGetMemInfo(Ptr, OL_MEM_INFO_TYPE, sizeof(RetrievedType),
-                              &RetrievedType));
+  ASSERT_SUCCESS(olGetMemInfo(Context, Ptr, OL_MEM_INFO_TYPE,
+                              sizeof(RetrievedType), &RetrievedType));
   ASSERT_EQ(RetrievedType, getTestParam());
 }
 
+TEST_P(olGetMemInfoAllocTypeTest, SuccessInteriorPointer) {
+  // The spec allows Ptr to point anywhere inside the allocation; the query
+  // must still resolve to the allocation's base and full size.
+  void *Interior = reinterpret_cast<char *>(Ptr) + SIZE / 2;
+
+  void *RetrievedBase;
+  ASSERT_SUCCESS(olGetMemInfo(Context, Interior, OL_MEM_INFO_BASE,
+                              sizeof(RetrievedBase), &RetrievedBase));
+  ASSERT_EQ(RetrievedBase, Ptr);
+
+  size_t RetrievedSize;
+  ASSERT_SUCCESS(olGetMemInfo(Context, Interior, OL_MEM_INFO_SIZE,
+                              sizeof(RetrievedSize), &RetrievedSize));
+  ASSERT_EQ(RetrievedSize, SIZE);
+}
+
 TEST_P(olGetMemInfoTest, InvalidNotFound) {
   // Assuming that we aren't unlucky and happen to get 0x1234 as a random
   // pointer
   void *RetrievedBase;
   ASSERT_ERROR(OL_ERRC_NOT_FOUND,
-               olGetMemInfo(reinterpret_cast<void *>(0x1234), OL_MEM_INFO_BASE,
-                            sizeof(RetrievedBase), &RetrievedBase));
+               olGetMemInfo(Context, reinterpret_cast<void *>(0x1234),
+                            OL_MEM_INFO_BASE, sizeof(RetrievedBase),
+                            &RetrievedBase));
 }
 
 TEST_P(olGetMemInfoTest, InvalidNullPtr) {
   ol_device_handle_t RetrievedDevice;
   ASSERT_ERROR(OL_ERRC_INVALID_NULL_POINTER,
-               olGetMemInfo(nullptr, OL_MEM_INFO_DEVICE,
+               olGetMemInfo(Context, nullptr, OL_MEM_INFO_DEVICE,
                             sizeof(RetrievedDevice), &RetrievedDevice));
 }
 
 TEST_P(olGetMemInfoTest, InvalidSizeZero) {
   ol_device_handle_t RetrievedDevice;
-  ASSERT_ERROR(OL_ERRC_INVALID_SIZE,
-               olGetMemInfo(Ptr, OL_MEM_INFO_DEVICE, 0, &RetrievedDevice));
+  ASSERT_ERROR(
+      OL_ERRC_INVALID_SIZE,
+      olGetMemInfo(Context, Ptr, OL_MEM_INFO_DEVICE, 0, &RetrievedDevice));
 }
 
 TEST_P(olGetMemInfoTest, InvalidSizeSmall) {
   ol_device_handle_t RetrievedDevice;
   ASSERT_ERROR(OL_ERRC_INVALID_SIZE,
-               olGetMemInfo(Ptr, OL_MEM_INFO_DEVICE,
+               olGetMemInfo(Context, Ptr, OL_MEM_INFO_DEVICE,
                             sizeof(RetrievedDevice) - 1, &RetrievedDevice));
 }
 
 TEST_P(olGetMemInfoTest, InvalidNullPointerPropValue) {
   ol_device_handle_t RetrievedDevice;
-  ASSERT_ERROR(
-      OL_ERRC_INVALID_NULL_POINTER,
-      olGetMemInfo(Ptr, OL_MEM_INFO_DEVICE, sizeof(RetrievedDevice), nullptr));
+  ASSERT_ERROR(OL_ERRC_INVALID_NULL_POINTER,
+               olGetMemInfo(Context, Ptr, OL_MEM_INFO_DEVICE,
+                            sizeof(RetrievedDevice), nullptr));
 }
diff --git a/offload/unittests/OffloadAPI/memory/olGetMemInfoSize.cpp b/offload/unittests/OffloadAPI/memory/olGetMemInfoSize.cpp
index 7e0ffaacea19d..ff8cf410ffb73 100644
--- a/offload/unittests/OffloadAPI/memory/olGetMemInfoSize.cpp
+++ b/offload/unittests/OffloadAPI/memory/olGetMemInfoSize.cpp
@@ -15,11 +15,12 @@ struct olGetMemInfoSizeTypesTest : olPropertyTest<ol_mem_info_t> {
 
   void SetUp() override {
     RETURN_ON_FATAL_FAILURE(olPropertyTest<ol_mem_info_t>::SetUp());
-    ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, 0x1024, &Ptr));
+    ASSERT_SUCCESS(olMemAlloc(this->Context, this->Device, OL_ALLOC_TYPE_DEVICE,
+                              0x1024, &Ptr));
   }
 
   void TearDown() override {
-    ASSERT_SUCCESS(olMemFree(Ptr));
+    ASSERT_SUCCESS(olMemFree(this->Context, Ptr));
     RETURN_ON_FATAL_FAILURE(olPropertyTest<ol_mem_info_t>::TearDown());
   }
 
@@ -35,11 +36,12 @@ struct olGetMemInfoSizeTest : OffloadDeviceTest {
 
   void SetUp() override {
     RETURN_ON_FATAL_FAILURE(OffloadDeviceTest::SetUp());
-    ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, 0x1024, &Ptr));
+    ASSERT_SUCCESS(olMemAlloc(this->Context, this->Device, OL_ALLOC_TYPE_DEVICE,
+                              0x1024, &Ptr));
   }
 
   void TearDown() override {
-    ASSERT_SUCCESS(olMemFree(Ptr));
+    ASSERT_SUCCESS(olMemFree(this->Context, Ptr));
     RETURN_ON_FATAL_FAILURE(OffloadDeviceTest::TearDown());
   }
 
@@ -50,17 +52,17 @@ OFFLOAD_TESTS_INSTANTIATE_DEVICE_FIXTURE(olGetMemInfoSizeTest);
 
 TEST_P(olGetMemInfoSizeTypesTest, Success) {
   size_t Size = 0;
-  ASSERT_SUCCESS(olGetMemInfoSize(Ptr, Property, &Size));
+  ASSERT_SUCCESS(olGetMemInfoSize(Context, Ptr, Property, &Size));
   ASSERT_EQ(Size, PropertySize);
 }
 
 TEST_P(olGetMemInfoSizeTest, InvalidSymbolInfoEnumeration) {
   size_t Size = 0;
   ASSERT_ERROR(OL_ERRC_INVALID_ENUMERATION,
-               olGetMemInfoSize(Ptr, OL_MEM_INFO_FORCE_UINT32, &Size));
+               olGetMemInfoSize(Context, Ptr, OL_MEM_INFO_FORCE_UINT32, &Size));
 }
 
 TEST_P(olGetMemInfoSizeTest, InvalidNullPointer) {
   ASSERT_ERROR(OL_ERRC_INVALID_NULL_POINTER,
-               olGetMemInfoSize(Ptr, OL_MEM_INFO_DEVICE, nullptr));
+               olGetMemInfoSize(Context, Ptr, OL_MEM_INFO_DEVICE, nullptr));
 }
diff --git a/offload/unittests/OffloadAPI/memory/olMemAlloc.cpp b/offload/unittests/OffloadAPI/memory/olMemAlloc.cpp
index 3dfd8de06eea6..81d52e0e1e16d 100644
--- a/offload/unittests/OffloadAPI/memory/olMemAlloc.cpp
+++ b/offload/unittests/OffloadAPI/memory/olMemAlloc.cpp
@@ -17,10 +17,10 @@ struct olMemAllocAllocTypesTest : OffloadDeviceTestWithParam<ol_alloc_type_t> {
   ol_result_t allocateDeviceOrHost(size_t Size, void **Alloc) {
     ol_alloc_type_t AllocType = getTestParam();
     if (AllocType == OL_ALLOC_TYPE_HOST) {
-      return olMemAllocHost(this->Device, Size, Alloc);
+      return olMemAllocHost(this->Context, this->Device, Size, Alloc);
     }
 
-    return olMemAlloc(this->Device, AllocType, Size, Alloc);
+    return olMemAlloc(this->Context, this->Device, AllocType, Size, Alloc);
   }
 };
 
@@ -32,7 +32,7 @@ TEST_P(olMemAllocAllocTypesTest, Success) {
   void *Alloc = nullptr;
   ASSERT_SUCCESS(allocateDeviceOrHost(DefaultAllocSize, &Alloc));
   ASSERT_NE(Alloc, nullptr);
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 TEST_P(olMemAllocTest, SuccessAllocMany) {
@@ -43,10 +43,11 @@ TEST_P(olMemAllocTest, SuccessAllocMany) {
     void *Alloc = nullptr;
     ol_alloc_type_t AllocType = AllocTypes[I % 3];
     if (AllocType == OL_ALLOC_TYPE_HOST) {
-      ASSERT_SUCCESS(olMemAllocHost(Device, DefaultAllocSize * I, &Alloc));
+      ASSERT_SUCCESS(
+          olMemAllocHost(Context, Device, DefaultAllocSize * I, &Alloc));
     } else {
       ASSERT_SUCCESS(
-          olMemAlloc(Device, AllocType, DefaultAllocSize * I, &Alloc));
+          olMemAlloc(Context, Device, AllocType, DefaultAllocSize * I, &Alloc));
     }
     ASSERT_NE(Alloc, nullptr);
 
@@ -54,34 +55,36 @@ TEST_P(olMemAllocTest, SuccessAllocMany) {
   }
 
   for (auto *A : Allocs) {
-    olMemFree(A);
+    olMemFree(Context, A);
   }
 }
 
 TEST_P(olMemAllocTest, InvalidNullDevice) {
   void *Alloc = nullptr;
-  ASSERT_ERROR(OL_ERRC_INVALID_NULL_HANDLE,
-               olMemAlloc(nullptr, OL_ALLOC_TYPE_DEVICE, 1024, &Alloc));
+  ASSERT_ERROR(
+      OL_ERRC_INVALID_NULL_HANDLE,
+      olMemAlloc(Context, nullptr, OL_ALLOC_TYPE_DEVICE, 1024, &Alloc));
 }
 
 TEST_P(olMemAllocTest, InvalidNullDeviceHost) {
   void *Alloc = nullptr;
   ASSERT_ERROR(OL_ERRC_INVALID_NULL_HANDLE,
-               olMemAllocHost(nullptr, 1024, &Alloc));
+               olMemAllocHost(Context, nullptr, 1024, &Alloc));
 }
 
 TEST_P(olMemAllocTest, InvalidNullOutPtr) {
-  ASSERT_ERROR(OL_ERRC_INVALID_NULL_POINTER,
-               olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, 1024, nullptr));
+  ASSERT_ERROR(
+      OL_ERRC_INVALID_NULL_POINTER,
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, 1024, nullptr));
 }
 
 TEST_P(olMemAllocTest, InvalidNullOutPtrHost) {
   ASSERT_ERROR(OL_ERRC_INVALID_NULL_POINTER,
-               olMemAllocHost(Device, 1024, nullptr));
+               olMemAllocHost(Context, Device, 1024, nullptr));
 }
 
 TEST_P(olMemAllocTest, InvalidHostType) {
   void *Alloc = nullptr;
   ASSERT_ERROR(OL_ERRC_INVALID_ENUMERATION,
-               olMemAlloc(Device, OL_ALLOC_TYPE_HOST, 1024, &Alloc));
+               olMemAlloc(Context, Device, OL_ALLOC_TYPE_HOST, 1024, &Alloc));
 }
diff --git a/offload/unittests/OffloadAPI/memory/olMemAllocAligned.cpp b/offload/unittests/OffloadAPI/memory/olMemAllocAligned.cpp
index a894dde5ec77c..0bb7d759a129b 100644
--- a/offload/unittests/OffloadAPI/memory/olMemAllocAligned.cpp
+++ b/offload/unittests/OffloadAPI/memory/olMemAllocAligned.cpp
@@ -19,10 +19,12 @@ struct olMemAllocAlignedTypesTest
                                    void **Alloc) {
     ol_alloc_type_t AllocType = getTestParam();
     if (AllocType == OL_ALLOC_TYPE_HOST) {
-      return olMemAllocAlignedHost(this->Device, Size, Alignment, Alloc);
+      return olMemAllocAlignedHost(this->Context, this->Device, Size, Alignment,
+                                   Alloc);
     }
 
-    return olMemAllocAligned(this->Device, AllocType, Size, Alignment, Alloc);
+    return olMemAllocAligned(this->Context, this->Device, AllocType, Size,
+                             Alignment, Alloc);
   }
 };
 
@@ -40,11 +42,12 @@ TEST_P(olMemAllocAlignedTest, SuccessAllocMany) {
     void *Alloc = nullptr;
     ol_alloc_type_t AllocType = AllocTypes[I % 3];
     if (AllocType == OL_ALLOC_TYPE_HOST) {
-      ASSERT_SUCCESS(olMemAllocAlignedHost(Device, DefaultAllocSize * I,
-                                           DefaultAlignment, &Alloc));
+      ASSERT_SUCCESS(olMemAllocAlignedHost(
+          Context, Device, DefaultAllocSize * I, DefaultAlignment, &Alloc));
     } else {
-      ASSERT_SUCCESS(olMemAllocAligned(Device, AllocType, DefaultAllocSize * I,
-                                       DefaultAlignment, &Alloc));
+      ASSERT_SUCCESS(olMemAllocAligned(Context, Device, AllocType,
+                                       DefaultAllocSize * I, DefaultAlignment,
+                                       &Alloc));
     }
     ASSERT_NE(Alloc, nullptr);
 
@@ -52,43 +55,43 @@ TEST_P(olMemAllocAlignedTest, SuccessAllocMany) {
   }
 
   for (auto *A : Allocs) {
-    olMemFree(A);
+    olMemFree(Context, A);
   }
 }
 
 TEST_P(olMemAllocAlignedTest, InvalidNullDevice) {
   void *Alloc = nullptr;
   ASSERT_ERROR(OL_ERRC_INVALID_NULL_HANDLE,
-               olMemAllocAligned(nullptr, OL_ALLOC_TYPE_DEVICE, 1024,
+               olMemAllocAligned(Context, nullptr, OL_ALLOC_TYPE_DEVICE, 1024,
                                  DefaultAlignment, &Alloc));
 }
 
 TEST_P(olMemAllocAlignedTest, InvalidNullOutPtr) {
   ASSERT_ERROR(OL_ERRC_INVALID_NULL_POINTER,
-               olMemAllocAligned(Device, OL_ALLOC_TYPE_DEVICE, 1024,
+               olMemAllocAligned(Context, Device, OL_ALLOC_TYPE_DEVICE, 1024,
                                  DefaultAlignment, nullptr));
 }
 
 TEST_P(olMemAllocAlignedTest, InvalidAlignmentZero) {
   void *Alloc = nullptr;
 
-  ASSERT_ERROR(
-      OL_ERRC_INVALID_ARGUMENT,
-      olMemAllocAligned(Device, OL_ALLOC_TYPE_DEVICE, 1024, 0, &Alloc));
+  ASSERT_ERROR(OL_ERRC_INVALID_ARGUMENT,
+               olMemAllocAligned(Context, Device, OL_ALLOC_TYPE_DEVICE, 1024, 0,
+                                 &Alloc));
 }
 
 TEST_P(olMemAllocAlignedTest, InvalidAlignmentNotAPowerOfTwo) {
   void *Alloc = nullptr;
 
-  ASSERT_ERROR(
-      OL_ERRC_INVALID_ARGUMENT,
-      olMemAllocAligned(Device, OL_ALLOC_TYPE_DEVICE, 1024, 3, &Alloc));
+  ASSERT_ERROR(OL_ERRC_INVALID_ARGUMENT,
+               olMemAllocAligned(Context, Device, OL_ALLOC_TYPE_DEVICE, 1024, 3,
+                                 &Alloc));
 }
 
 TEST_P(olMemAllocAlignedTest, InvalidHostType) {
   void *Alloc = nullptr;
   ASSERT_ERROR(OL_ERRC_INVALID_ENUMERATION,
-               olMemAllocAligned(Device, OL_ALLOC_TYPE_HOST, 1024,
+               olMemAllocAligned(Context, Device, OL_ALLOC_TYPE_HOST, 1024,
                                  DefaultAlignment, &Alloc));
 }
 
@@ -100,7 +103,7 @@ TEST_P(olMemAllocAlignedTest, CudaExceedDefaultAlignment) {
   void *Alloc = nullptr;
   // The default page size for cuda is 64 KB.
   ASSERT_ERROR(OL_ERRC_UNSUPPORTED,
-               olMemAllocAligned(Device, OL_ALLOC_TYPE_DEVICE, 1024,
+               olMemAllocAligned(Context, Device, OL_ALLOC_TYPE_DEVICE, 1024,
                                  1024 * 64 * 64 * 64, &Alloc));
   ASSERT_EQ(Alloc, nullptr);
 }
@@ -116,7 +119,7 @@ TEST_P(olMemAllocAlignedTypesTest, SuccessAllocDifferentAlignments) {
     SCOPED_TRACE("alignment: " + std::to_string(Alignment));
     ASSERT_SUCCESS(allocateDeviceOrHost(DefaultAllocSize, Alignment, &Alloc));
     ASSERT_NE(Alloc, nullptr);
-    olMemFree(Alloc);
+    olMemFree(Context, Alloc);
   }
 }
 
@@ -142,6 +145,6 @@ TEST_P(olMemAllocAlignedTypesTest, SuccessMemcpyDiferentAlignments) {
       ASSERT_EQ(Val, 42);
     }
 
-    ASSERT_SUCCESS(olMemFree(Alloc));
+    ASSERT_SUCCESS(olMemFree(Context, Alloc));
   }
 }
diff --git a/offload/unittests/OffloadAPI/memory/olMemFill.cpp b/offload/unittests/OffloadAPI/memory/olMemFill.cpp
index 467a551c48c94..57812dc4db052 100644
--- a/offload/unittests/OffloadAPI/memory/olMemFill.cpp
+++ b/offload/unittests/OffloadAPI/memory/olMemFill.cpp
@@ -27,7 +27,8 @@ struct olMemFillTest : OffloadQueueTest {
     }
 
     void *Alloc;
-    ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+    ASSERT_SUCCESS(
+        olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
     PatternTy Pattern = PatternVal;
     ASSERT_SUCCESS(olMemFill(Queue, Alloc, sizeof(Pattern), &Pattern, Size));
@@ -43,7 +44,7 @@ struct olMemFillTest : OffloadQueueTest {
       ASSERT_EQ(AllocPtr[i], Pattern);
     }
 
-    olMemFree(Alloc);
+    olMemFree(Context, Alloc);
   }
 };
 OFFLOAD_TESTS_INSTANTIATE_DEVICE_FIXTURE(olMemFillTest);
@@ -78,7 +79,8 @@ TEST_P(olMemFillTest, Success32Enqueue) {
 TEST_P(olMemFillTest, SuccessLarge) {
   constexpr size_t Size = 1024;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   struct PatternT {
     uint64_t A;
@@ -96,7 +98,7 @@ TEST_P(olMemFillTest, SuccessLarge) {
     ASSERT_EQ(AllocPtr[i].B, UINT64_MAX);
   }
 
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 TEST_P(olMemFillTest, SuccessLargeEnqueue) {
@@ -105,7 +107,8 @@ TEST_P(olMemFillTest, SuccessLargeEnqueue) {
   ManuallyTriggeredTask Manual;
   ASSERT_SUCCESS(Manual.enqueue(Queue));
 
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   struct PatternT {
     uint64_t A;
@@ -124,13 +127,14 @@ TEST_P(olMemFillTest, SuccessLargeEnqueue) {
     ASSERT_EQ(AllocPtr[i].B, UINT64_MAX);
   }
 
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 TEST_P(olMemFillTest, SuccessLargeByteAligned) {
   constexpr size_t Size = 17 * 64;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   struct __attribute__((packed)) PatternT {
     uint64_t A;
@@ -150,7 +154,7 @@ TEST_P(olMemFillTest, SuccessLargeByteAligned) {
     ASSERT_EQ(AllocPtr[i].C, 255);
   }
 
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 TEST_P(olMemFillTest, SuccessLargeByteAlignedEnqueue) {
@@ -159,7 +163,8 @@ TEST_P(olMemFillTest, SuccessLargeByteAlignedEnqueue) {
   ManuallyTriggeredTask Manual;
   ASSERT_SUCCESS(Manual.enqueue(Queue));
 
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   struct __attribute__((packed)) PatternT {
     uint64_t A;
@@ -180,33 +185,35 @@ TEST_P(olMemFillTest, SuccessLargeByteAlignedEnqueue) {
     ASSERT_EQ(AllocPtr[i].C, 255);
   }
 
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 TEST_P(olMemFillTest, InvalidSizeNotMultipleOfPatternSize) {
   constexpr size_t Size = 1025;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   uint16_t Pattern = 0x4242;
   ASSERT_ERROR(OL_ERRC_INVALID_SIZE,
                olMemFill(Queue, Alloc, sizeof(Pattern), &Pattern, Size));
 
   olSyncQueue(Queue);
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 TEST_P(olMemFillTest, InvalidPatternSizeLargerThanFillSize) {
   constexpr size_t Size = 4;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   uint64_t Pattern = 0x4242424242424242;
   ASSERT_ERROR(OL_ERRC_INVALID_SIZE,
                olMemFill(Queue, Alloc, sizeof(Pattern), &Pattern, Size));
 
   olSyncQueue(Queue);
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 // Even though L0, CUDA and HSA do not support non-power-of-two patterns,
@@ -219,7 +226,8 @@ static constexpr std::array<unsigned char, 3> FallbackPattern = {0x11, 0x22,
 TEST_P(olMemFillTest, SuccessNonPow2PatternManaged) {
   constexpr size_t Size = FallbackPattern.size() * 1000;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   ASSERT_SUCCESS(olMemFill(Queue, Alloc, FallbackPattern.size(),
                            FallbackPattern.data(), Size));
@@ -229,13 +237,14 @@ TEST_P(olMemFillTest, SuccessNonPow2PatternManaged) {
   for (size_t I = 0; I < Size; I++)
     ASSERT_EQ(AllocPtr[I], FallbackPattern[I % FallbackPattern.size()]);
 
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 TEST_P(olMemFillTest, SuccessNonPow2PatternDevice) {
   constexpr size_t Size = FallbackPattern.size() * 1000;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
 
   ASSERT_SUCCESS(olMemFill(Queue, Alloc, FallbackPattern.size(),
                            FallbackPattern.data(), Size));
@@ -247,13 +256,14 @@ TEST_P(olMemFillTest, SuccessNonPow2PatternDevice) {
   for (size_t I = 0; I < Size; I++)
     ASSERT_EQ(HostBuf[I], FallbackPattern[I % FallbackPattern.size()]);
 
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 TEST_P(olMemFillTest, SuccessNonPow2PatternDeviceSmall) {
   constexpr size_t Size = FallbackPattern.size() * 2;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
 
   ASSERT_SUCCESS(olMemFill(Queue, Alloc, FallbackPattern.size(),
                            FallbackPattern.data(), Size));
@@ -265,5 +275,5 @@ TEST_P(olMemFillTest, SuccessNonPow2PatternDeviceSmall) {
   for (size_t I = 0; I < Size; I++)
     ASSERT_EQ(HostBuf[I], FallbackPattern[I % FallbackPattern.size()]);
 
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
diff --git a/offload/unittests/OffloadAPI/memory/olMemFree.cpp b/offload/unittests/OffloadAPI/memory/olMemFree.cpp
index 0bb7475bee1ae..6df2ba637f5a5 100644
--- a/offload/unittests/OffloadAPI/memory/olMemFree.cpp
+++ b/offload/unittests/OffloadAPI/memory/olMemFree.cpp
@@ -22,16 +22,17 @@ TEST_P(olMemFreeAllocTypesTest, Success) {
   void *Alloc = nullptr;
   ol_alloc_type_t AllocType = getTestParam();
   if (AllocType == OL_ALLOC_TYPE_HOST) {
-    ASSERT_SUCCESS(olMemAllocHost(Device, 1024, &Alloc));
+    ASSERT_SUCCESS(olMemAllocHost(Context, Device, 1024, &Alloc));
   } else {
-    ASSERT_SUCCESS(olMemAlloc(Device, AllocType, 1024, &Alloc));
+    ASSERT_SUCCESS(olMemAlloc(Context, Device, AllocType, 1024, &Alloc));
   }
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemFreeTest, InvalidNullPtr) {
   void *Alloc = nullptr;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, 1024, &Alloc));
-  ASSERT_ERROR(OL_ERRC_INVALID_NULL_POINTER, olMemFree(nullptr));
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, 1024, &Alloc));
+  ASSERT_ERROR(OL_ERRC_INVALID_NULL_POINTER, olMemFree(Context, nullptr));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
diff --git a/offload/unittests/OffloadAPI/memory/olMemPrefetch.cpp b/offload/unittests/OffloadAPI/memory/olMemPrefetch.cpp
index ba446ee4cd7ec..3571521676b2d 100644
--- a/offload/unittests/OffloadAPI/memory/olMemPrefetch.cpp
+++ b/offload/unittests/OffloadAPI/memory/olMemPrefetch.cpp
@@ -16,7 +16,8 @@ OFFLOAD_TESTS_INSTANTIATE_DEVICE_FIXTURE(olMemPrefetchTest);
 TEST_P(olMemPrefetchTest, SuccessHostToDevice) {
   constexpr size_t Size = 1024;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   std::memset(Alloc, 0x42, Size);
 
@@ -29,13 +30,14 @@ TEST_P(olMemPrefetchTest, SuccessHostToDevice) {
   for (size_t I = 0; I < Size; I++)
     ASSERT_EQ(static_cast<uint8_t *>(Alloc)[I], 0x42);
 
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemPrefetchTest, SuccessDeviceToHost) {
   constexpr size_t Size = 1024;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   std::memset(Alloc, 0x21, Size);
 
@@ -51,7 +53,7 @@ TEST_P(olMemPrefetchTest, SuccessDeviceToHost) {
   for (size_t I = 0; I < Size; I++)
     ASSERT_EQ(static_cast<uint8_t *>(Alloc)[I], 0x21);
 
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemPrefetchTest, SuccessMultiple) {
@@ -61,7 +63,8 @@ TEST_P(olMemPrefetchTest, SuccessMultiple) {
   const void *Mems[NumAllocs];
   size_t Sizes[NumAllocs];
   for (size_t I = 0; I < NumAllocs; I++) {
-    ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Allocs[I]));
+    ASSERT_SUCCESS(
+        olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Allocs[I]));
     std::memset(Allocs[I], static_cast<int>(0x10 + I), Size);
     Mems[I] = Allocs[I];
     Sizes[I] = Size;
@@ -75,7 +78,7 @@ TEST_P(olMemPrefetchTest, SuccessMultiple) {
     for (size_t J = 0; J < Size; J++)
       ASSERT_EQ(static_cast<uint8_t *>(Allocs[I])[J],
                 static_cast<uint8_t>(0x10 + I));
-    ASSERT_SUCCESS(olMemFree(Allocs[I]));
+    ASSERT_SUCCESS(olMemFree(Context, Allocs[I]));
   }
 }
 
@@ -88,7 +91,8 @@ TEST_P(olMemPrefetchTest, SuccessZeroCount) {
 TEST_P(olMemPrefetchTest, SuccessZeroSize) {
   constexpr size_t Size = 1024;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   const void *Mems[] = {Alloc};
   const size_t Sizes[] = {0};
@@ -96,7 +100,7 @@ TEST_P(olMemPrefetchTest, SuccessZeroSize) {
                                OL_MEM_MIGRATION_FLAG_HOST_TO_DEVICE));
   ASSERT_SUCCESS(olSyncQueue(Queue));
 
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemPrefetchTest, SuccessUnsupportedAllocType) {
@@ -104,7 +108,8 @@ TEST_P(olMemPrefetchTest, SuccessUnsupportedAllocType) {
   // contract the hint must be silently ignored and the call must still succeed.
   constexpr size_t Size = 1024;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
 
   const void *Mems[] = {Alloc};
   const size_t Sizes[] = {Size};
@@ -112,20 +117,21 @@ TEST_P(olMemPrefetchTest, SuccessUnsupportedAllocType) {
                                OL_MEM_MIGRATION_FLAG_HOST_TO_DEVICE));
   ASSERT_SUCCESS(olSyncQueue(Queue));
 
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemPrefetchTest, InvalidFlags) {
   constexpr size_t Size = 1024;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   const void *Mems[] = {Alloc};
   const size_t Sizes[] = {Size};
   ASSERT_ERROR(OL_ERRC_INVALID_ENUMERATION,
                olMemPrefetch(Queue, 1, Mems, Sizes, 0xdeadbeef));
 
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemPrefetchTest, InvalidNullMems) {
@@ -138,12 +144,13 @@ TEST_P(olMemPrefetchTest, InvalidNullMems) {
 TEST_P(olMemPrefetchTest, InvalidNullSizes) {
   constexpr size_t Size = 1024;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED, Size, &Alloc));
 
   const void *Mems[] = {Alloc};
   ASSERT_ERROR(OL_ERRC_INVALID_NULL_POINTER,
                olMemPrefetch(Queue, 1, Mems, nullptr,
                              OL_MEM_MIGRATION_FLAG_HOST_TO_DEVICE));
 
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
diff --git a/offload/unittests/OffloadAPI/memory/olMemcpy.cpp b/offload/unittests/OffloadAPI/memory/olMemcpy.cpp
index 176ac0f2ec48d..f10c0ffc09141 100644
--- a/offload/unittests/OffloadAPI/memory/olMemcpy.cpp
+++ b/offload/unittests/OffloadAPI/memory/olMemcpy.cpp
@@ -44,11 +44,12 @@ OFFLOAD_TESTS_INSTANTIATE_DEVICE_FIXTURE(olMemcpyGlobalTest);
 TEST_P(olMemcpyTest, SuccessHtoD) {
   constexpr size_t Size = 1024;
   void *Alloc;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
   std::vector<uint8_t> Input(Size, 42);
   ASSERT_SUCCESS(olMemcpy(Queue, Alloc, Device, Input.data(), Host, Size));
   olSyncQueue(Queue);
-  olMemFree(Alloc);
+  olMemFree(Context, Alloc);
 }
 
 TEST_P(olMemcpyTest, SuccessDtoH) {
@@ -57,14 +58,15 @@ TEST_P(olMemcpyTest, SuccessDtoH) {
   std::vector<uint8_t> Input(Size, 42);
   std::vector<uint8_t> Output(Size, 0);
 
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
   ASSERT_SUCCESS(olMemcpy(Queue, Alloc, Device, Input.data(), Host, Size));
   ASSERT_SUCCESS(olMemcpy(Queue, Output.data(), Host, Alloc, Device, Size));
   ASSERT_SUCCESS(olSyncQueue(Queue));
   for (uint8_t Val : Output) {
     ASSERT_EQ(Val, 42);
   }
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemcpyTest, SuccessDtoD) {
@@ -74,8 +76,10 @@ TEST_P(olMemcpyTest, SuccessDtoD) {
   std::vector<uint8_t> Input(Size, 42);
   std::vector<uint8_t> Output(Size, 0);
 
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &AllocA));
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &AllocB));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &AllocA));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &AllocB));
   ASSERT_SUCCESS(olMemcpy(Queue, AllocA, Device, Input.data(), Host, Size));
   ASSERT_SUCCESS(olMemcpy(Queue, AllocB, Device, AllocA, Device, Size));
   ASSERT_SUCCESS(olMemcpy(Queue, Output.data(), Host, AllocB, Device, Size));
@@ -83,8 +87,8 @@ TEST_P(olMemcpyTest, SuccessDtoD) {
   for (uint8_t Val : Output) {
     ASSERT_EQ(Val, 42);
   }
-  ASSERT_SUCCESS(olMemFree(AllocA));
-  ASSERT_SUCCESS(olMemFree(AllocB));
+  ASSERT_SUCCESS(olMemFree(Context, AllocA));
+  ASSERT_SUCCESS(olMemFree(Context, AllocB));
 }
 
 TEST_P(olMemcpyTest, SuccessHtoHSync) {
@@ -108,7 +112,8 @@ TEST_P(olMemcpyTest, SuccessHtoHQueuedOrdering) {
   std::vector<uint8_t> Copied(Size, 0);
   std::vector<uint8_t> Output(Size, 0);
 
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
   ASSERT_SUCCESS(olMemcpy(Queue, Alloc, Device, Input.data(), Host, Size));
   ASSERT_SUCCESS(
       olMemcpy(Queue, Intermediate.data(), Host, Alloc, Device, Size));
@@ -121,7 +126,7 @@ TEST_P(olMemcpyTest, SuccessHtoHQueuedOrdering) {
   for (uint8_t Val : Output)
     ASSERT_EQ(Val, 42);
 
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemcpyTest, SuccessHtoHQueuedOrderingHostAlloc) {
@@ -132,11 +137,12 @@ TEST_P(olMemcpyTest, SuccessHtoHQueuedOrderingHostAlloc) {
   void *Copied;
   void *Output;
 
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
-  ASSERT_SUCCESS(olMemAllocHost(Device, Size, &Input));
-  ASSERT_SUCCESS(olMemAllocHost(Device, Size, &Intermediate));
-  ASSERT_SUCCESS(olMemAllocHost(Device, Size, &Copied));
-  ASSERT_SUCCESS(olMemAllocHost(Device, Size, &Output));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
+  ASSERT_SUCCESS(olMemAllocHost(Context, Device, Size, &Input));
+  ASSERT_SUCCESS(olMemAllocHost(Context, Device, Size, &Intermediate));
+  ASSERT_SUCCESS(olMemAllocHost(Context, Device, Size, &Copied));
+  ASSERT_SUCCESS(olMemAllocHost(Context, Device, Size, &Output));
 
   std::memset(Input, 42, Size);
   std::memset(Intermediate, 0, Size);
@@ -153,11 +159,11 @@ TEST_P(olMemcpyTest, SuccessHtoHQueuedOrderingHostAlloc) {
   for (size_t I = 0; I < Size; ++I)
     ASSERT_EQ(static_cast<uint8_t *>(Output)[I], 42);
 
-  ASSERT_SUCCESS(olMemFree(Output));
-  ASSERT_SUCCESS(olMemFree(Copied));
-  ASSERT_SUCCESS(olMemFree(Intermediate));
-  ASSERT_SUCCESS(olMemFree(Input));
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Output));
+  ASSERT_SUCCESS(olMemFree(Context, Copied));
+  ASSERT_SUCCESS(olMemFree(Context, Intermediate));
+  ASSERT_SUCCESS(olMemFree(Context, Input));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemcpyTest, SuccessDtoHSync) {
@@ -166,13 +172,14 @@ TEST_P(olMemcpyTest, SuccessDtoHSync) {
   std::vector<uint8_t> Input(Size, 42);
   std::vector<uint8_t> Output(Size, 0);
 
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
+  ASSERT_SUCCESS(
+      olMemAlloc(Context, Device, OL_ALLOC_TYPE_DEVICE, Size, &Alloc));
   ASSERT_SUCCESS(olMemcpy(nullptr, Alloc, Device, Input.data(), Host, Size));
   ASSERT_SUCCESS(olMemcpy(nullptr, Output.data(), Host, Alloc, Device, Size));
   for (uint8_t Val : Output) {
     ASSERT_EQ(Val, 42);
   }
-  ASSERT_SUCCESS(olMemFree(Alloc));
+  ASSERT_SUCCESS(olMemFree(Context, Alloc));
 }
 
 TEST_P(olMemcpyTest, SuccessSizeZero) {
@@ -196,14 +203,14 @@ TEST_P(olMemcpyTest, SuccessHtoHQueuedSizeZero) {
 
 TEST_P(olMemcpyGlobalTest, SuccessRoundTrip) {
   void *SourceMem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             64 * sizeof(uint32_t), &SourceMem));
   uint32_t *SourceData = (uint32_t *)SourceMem;
   for (auto I = 0; I < 64; I++)
     SourceData[I] = I;
 
   void *DestMem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             64 * sizeof(uint32_t), &DestMem));
 
   ASSERT_SUCCESS(
@@ -217,13 +224,13 @@ TEST_P(olMemcpyGlobalTest, SuccessRoundTrip) {
   for (uint32_t I = 0; I < 64; I++)
     ASSERT_EQ(DestData[I], I);
 
-  ASSERT_SUCCESS(olMemFree(DestMem));
-  ASSERT_SUCCESS(olMemFree(SourceMem));
+  ASSERT_SUCCESS(olMemFree(Context, DestMem));
+  ASSERT_SUCCESS(olMemFree(Context, SourceMem));
 }
 
 TEST_P(olMemcpyGlobalTest, SuccessWrite) {
   void *SourceMem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(uint32_t),
                             &SourceMem));
   uint32_t *SourceData = (uint32_t *)SourceMem;
@@ -231,7 +238,7 @@ TEST_P(olMemcpyGlobalTest, SuccessWrite) {
     SourceData[I] = I;
 
   void *DestMem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(uint32_t),
                             &DestMem));
   void *ArgPtrs[] = {&DestMem};
@@ -248,13 +255,13 @@ TEST_P(olMemcpyGlobalTest, SuccessWrite) {
   for (uint32_t I = 0; I < 64; I++)
     ASSERT_EQ(DestData[I], I);
 
-  ASSERT_SUCCESS(olMemFree(DestMem));
-  ASSERT_SUCCESS(olMemFree(SourceMem));
+  ASSERT_SUCCESS(olMemFree(Context, DestMem));
+  ASSERT_SUCCESS(olMemFree(Context, SourceMem));
 }
 
 TEST_P(olMemcpyGlobalTest, SuccessRead) {
   void *DestMem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(uint32_t),
                             &DestMem));
 
@@ -269,5 +276,5 @@ TEST_P(olMemcpyGlobalTest, SuccessRead) {
   for (uint32_t I = 0; I < 64; I++)
     ASSERT_EQ(DestData[I], I * 2);
 
-  ASSERT_SUCCESS(olMemFree(DestMem));
+  ASSERT_SUCCESS(olMemFree(Context, DestMem));
 }
diff --git a/offload/unittests/OffloadAPI/queue/olLaunchHostFunction.cpp b/offload/unittests/OffloadAPI/queue/olLaunchHostFunction.cpp
index 1dedf7ecaaefd..92a696331e618 100644
--- a/offload/unittests/OffloadAPI/queue/olLaunchHostFunction.cpp
+++ b/offload/unittests/OffloadAPI/queue/olLaunchHostFunction.cpp
@@ -60,7 +60,7 @@ TEST_P(olLaunchHostFunctionKernelTest, SuccessBlocking) {
   ASSERT_SUCCESS(olCreateQueue(Context, Device, &Queue));
 
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             LaunchArgs.GroupSize.x * sizeof(uint32_t), &Mem));
 
   uint32_t *Data = (uint32_t *)Mem;
@@ -98,7 +98,7 @@ TEST_P(olLaunchHostFunctionKernelTest, SuccessBlocking) {
   }
 
   ASSERT_SUCCESS(olDestroyQueue(Queue));
-  ASSERT_SUCCESS(olMemFree(Mem));
+  ASSERT_SUCCESS(olMemFree(Context, Mem));
 }
 
 TEST_P(olLaunchHostFunctionTest, InvalidNullCallback) {
diff --git a/offload/unittests/OffloadAPI/queue/olWaitEvents.cpp b/offload/unittests/OffloadAPI/queue/olWaitEvents.cpp
index 112785c1a066f..b9fc9af02fabd 100644
--- a/offload/unittests/OffloadAPI/queue/olWaitEvents.cpp
+++ b/offload/unittests/OffloadAPI/queue/olWaitEvents.cpp
@@ -36,7 +36,7 @@ TEST_P(olWaitEventsTest, Success) {
   ol_event_handle_t Events[NUM_KERNELS];
 
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             NUM_KERNELS * sizeof(uint32_t), &Mem));
   uint32_t Idx = 0;
   void *ArgPtrs[] = {&Idx, &Mem};
@@ -71,7 +71,7 @@ TEST_P(olWaitEventsTest, SuccessSingleQueue) {
   ASSERT_SUCCESS(olCreateQueue(Context, Device, &Queue));
 
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             NUM_KERNELS * sizeof(uint32_t), &Mem));
   uint32_t Idx = 0;
   void *ArgPtrs[] = {&Idx, &Mem};
@@ -102,7 +102,7 @@ TEST_P(olWaitEventsTest, SuccessMultipleEvents) {
   ol_event_handle_t Events[NUM_KERNELS];
 
   void *Mem;
-  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED,
+  ASSERT_SUCCESS(olMemAlloc(Context, Device, OL_ALLOC_TYPE_MANAGED,
                             NUM_KERNELS * sizeof(uint32_t), &Mem));
   uint32_t Idx = 0;
   void *ArgPtrs[] = {&Idx, &Mem};

>From 5dc5f1046785aa296394550a1a91fce76a9a5721 Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?=C5=81ukasz=20Plewa?= <lukasz.plewa at intel.com>
Date: Sun, 20 Sep 2026 15:25:04 +0200
Subject: [PATCH 2/3] [offload] bypass memory manager when record-replay is
 active

Skip the MM pool (and the free path) when the device has an active
RecordReplay instance so dataAlloc's RR shortcut services these
allocations.
---
 .../plugins-nextgen/common/src/PluginInterface.cpp    | 11 +++++++++++
 1 file changed, 11 insertions(+)

diff --git a/offload/plugins-nextgen/common/src/PluginInterface.cpp b/offload/plugins-nextgen/common/src/PluginInterface.cpp
index 00d04e4a1f2de..61b8b5fe1155b 100644
--- a/offload/plugins-nextgen/common/src/PluginInterface.cpp
+++ b/offload/plugins-nextgen/common/src/PluginInterface.cpp
@@ -1226,6 +1226,12 @@ Expected<void *> PluginContextTy::allocate(GenericDeviceTy &Device,
                                            int64_t Size, void *HostPtr,
                                            TargetAllocTy Kind,
                                            size_t Alignment) {
+  // Record-replay hands out interior pointers into a preallocated slab so
+  // recorded kernels can re-execute at their original addresses; the MM pool
+  // must be bypassed for those allocations to reach the RR bump allocator.
+  if (auto *RR = Device.getRecordReplay(); RR && RR->isRecordingOrReplaying())
+    return Device.dataAlloc(Size, HostPtr, Kind, Alignment);
+
   MemoryManagerTy *MM = (Kind == TARGET_ALLOC_HOST)
                             ? getHostMemoryManager()
                             : getDeviceMemoryManagerFor(Device, Kind);
@@ -1247,6 +1253,11 @@ Error PluginContextTy::deallocate(void *Ptr) {
 
 Error PluginContextTy::deallocate(GenericDeviceTy &Device, void *Ptr,
                                   TargetAllocTy Kind) {
+  // Symmetric with allocate: record-replay allocations never entered the MM
+  // pool, so route their free through dataDelete's RR shortcut.
+  if (auto *RR = Device.getRecordReplay(); RR && RR->isRecordingOrReplaying())
+    return Device.dataDelete(Ptr, Kind);
+
   MemoryManagerTy *MM = (Kind == TARGET_ALLOC_HOST)
                             ? getHostMemoryManager()
                             : getDeviceMemoryManagerFor(Device, Kind);

>From 91b5cec4a21f65d1d11f680d1ef1e49862035f2f Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?=C5=81ukasz=20Plewa?= <lukasz.plewa at intel.com>
Date: Mon, 21 Sep 2026 12:09:00 +0200
Subject: [PATCH 3/3] review fix

---
 offload/plugins-nextgen/level_zero/src/L0Plugin.cpp | 6 +++---
 1 file changed, 3 insertions(+), 3 deletions(-)

diff --git a/offload/plugins-nextgen/level_zero/src/L0Plugin.cpp b/offload/plugins-nextgen/level_zero/src/L0Plugin.cpp
index f9cf3b7eb59cf..a09ca28ad6df8 100644
--- a/offload/plugins-nextgen/level_zero/src/L0Plugin.cpp
+++ b/offload/plugins-nextgen/level_zero/src/L0Plugin.cpp
@@ -353,9 +353,9 @@ Error LevelZeroPluginContextTy::deallocate(GenericDeviceTy &Device, void *Ptr,
 Expected<PluginAllocInfoTy>
 LevelZeroPluginContextTy::getAllocInfo(const void *Ptr) {
   void *Raw = const_cast<void *>(Ptr);
-  for (auto &KV : DeviceAllocators) {
-    if (auto *Info = KV.second->getAllocInfo(Raw))
-      return PluginAllocInfoTy{KV.first, static_cast<TargetAllocTy>(Info->Kind),
+  for (const auto &[Device, Allocator] : DeviceAllocators) {
+    if (auto *Info = Allocator->getAllocInfo(Raw))
+      return PluginAllocInfoTy{Device, static_cast<TargetAllocTy>(Info->Kind),
                                Info->Base, Info->ReqSize};
   }
   if (HostAllocator) {



More information about the llvm-commits mailing list