[llvm] [offload] Fix kernel record/replay and add extensible mechanism (PR #190588)

Kevin Sala Penades via llvm-commits llvm-commits at lists.llvm.org
Mon Apr 6 01:19:57 PDT 2026


https://github.com/kevinsala created https://github.com/llvm/llvm-project/pull/190588

WIP

>From 54f0600ceb1358c488172f7dbe0a43e8f817792f Mon Sep 17 00:00:00 2001
From: Kevin Sala <salapenades1 at llnl.gov>
Date: Thu, 5 Mar 2026 21:43:15 -0800
Subject: [PATCH] [offload] Fix kernel record/replay and add extensible
 mechanism

---
 offload/libomptarget/device.cpp               |  15 +-
 offload/libomptarget/omptarget.cpp            |  10 +-
 .../amdgpu/dynamic_hsa/hsa.cpp                |   7 +
 .../plugins-nextgen/amdgpu/dynamic_hsa/hsa.h  |   7 +
 .../amdgpu/dynamic_hsa/hsa_ext_amd.h          |  37 ++
 offload/plugins-nextgen/amdgpu/src/rtl.cpp    |  96 +++-
 offload/plugins-nextgen/common/CMakeLists.txt |   1 +
 .../common/include/PluginInterface.h          | 107 ++++-
 .../common/include/RecordReplay.h             | 224 +++++++++
 .../common/src/PluginInterface.cpp            | 425 ++----------------
 .../common/src/RecordReplay.cpp               | 245 ++++++++++
 offload/plugins-nextgen/cuda/src/rtl.cpp      | 174 +++----
 .../kernelreplay/llvm-omp-kernel-replay.cpp   |  34 +-
 13 files changed, 840 insertions(+), 542 deletions(-)
 create mode 100644 offload/plugins-nextgen/common/include/RecordReplay.h
 create mode 100644 offload/plugins-nextgen/common/src/RecordReplay.cpp

diff --git a/offload/libomptarget/device.cpp b/offload/libomptarget/device.cpp
index 12c15aea1ad68..8fcf1c5d390dd 100644
--- a/offload/libomptarget/device.cpp
+++ b/offload/libomptarget/device.cpp
@@ -91,10 +91,21 @@ llvm::Error DeviceTy::init() {
   if (OMPX_RecordKernel) {
     // Enables saving the device memory kernel output post execution if set.
     BoolEnvar OMPX_ReplaySaveOutput("LIBOMPTARGET_RR_SAVE_OUTPUT", false);
+    Int64Envar OMPX_RecordMemSize("LIBOMPTARGET_RR_MEM_SIZE",
+                                  8 * 1024 * 1024 * 1024ULL);
+    Int32Envar OMPX_RecordDevice("LIBOMPTARGET_RR_DEVICE", 0);
+    if (OMPX_RecordDevice != RTLDeviceID)
+      return llvm::Error::success();
 
     uint64_t ReqPtrArgOffset;
-    RTL->initialize_record_replay(RTLDeviceID, 0, nullptr, true,
-                                  OMPX_ReplaySaveOutput, ReqPtrArgOffset);
+    Ret =
+        RTL->initialize_record_replay(RTLDeviceID, OMPX_RecordMemSize, nullptr,
+                                      /*IsRecord=*/true, /*IsNative=*/true,
+                                      OMPX_ReplaySaveOutput, ReqPtrArgOffset);
+    if (Ret != OFFLOAD_SUCCESS)
+      return error::createOffloadError(error::ErrorCode::BACKEND_FAILURE,
+                                       "failed to initialize RR in device %d\n",
+                                       DeviceID);
   }
 
   return llvm::Error::success();
diff --git a/offload/libomptarget/omptarget.cpp b/offload/libomptarget/omptarget.cpp
index 5fe8dd705b5c5..7b95908fe6010 100644
--- a/offload/libomptarget/omptarget.cpp
+++ b/offload/libomptarget/omptarget.cpp
@@ -2381,9 +2381,9 @@ int target(ident_t *Loc, DeviceTy &Device, void *HostPtr,
 int target_activate_rr(DeviceTy &Device, uint64_t MemorySize, void *VAddr,
                        bool IsRecord, bool SaveOutput,
                        uint64_t &ReqPtrArgOffset) {
-  return Device.RTL->initialize_record_replay(Device.DeviceID, MemorySize,
-                                              VAddr, IsRecord, SaveOutput,
-                                              ReqPtrArgOffset);
+  return Device.RTL->initialize_record_replay(
+      Device.DeviceID, MemorySize, VAddr, IsRecord, /*IsNative=*/true,
+      SaveOutput, ReqPtrArgOffset);
 }
 
 /// Executes a kernel using pre-recorded information for loading to
@@ -2431,7 +2431,11 @@ int target_replay(ident_t *Loc, DeviceTy &Device, void *HostPtr,
   KernelArgs.NumArgs = NumArgs;
   KernelArgs.Tripcount = LoopTripCount;
   KernelArgs.NumTeams[0] = NumTeams;
+  KernelArgs.NumTeams[1] = 1;
+  KernelArgs.NumTeams[2] = 1;
   KernelArgs.ThreadLimit[0] = ThreadLimit;
+  KernelArgs.ThreadLimit[1] = 1;
+  KernelArgs.ThreadLimit[2] = 1;
 
   int Ret = Device.launchKernel(TgtEntryPtr, TgtArgs, TgtOffsets, KernelArgs,
                                 AsyncInfo);
diff --git a/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa.cpp b/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa.cpp
index 279a296dd1618..5c7ec186b0ceb 100644
--- a/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa.cpp
+++ b/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa.cpp
@@ -75,6 +75,13 @@ DLWRAP(hsa_amd_profiling_set_profiler_enabled, 2)
 DLWRAP(hsa_code_object_reader_create_from_memory, 3)
 DLWRAP(hsa_code_object_reader_destroy, 1)
 DLWRAP(hsa_executable_load_agent_code_object, 5)
+DLWRAP(hsa_amd_vmem_address_reserve, 4)
+DLWRAP(hsa_amd_vmem_address_free, 2)
+DLWRAP(hsa_amd_vmem_handle_create, 5)
+DLWRAP(hsa_amd_vmem_handle_release, 1)
+DLWRAP(hsa_amd_vmem_map, 5)
+DLWRAP(hsa_amd_vmem_unmap, 2)
+DLWRAP(hsa_amd_vmem_set_access, 4)
 
 DLWRAP_FINALIZE()
 
diff --git a/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa.h b/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa.h
index f6e3337ddb3f4..258c7234251d5 100644
--- a/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa.h
+++ b/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa.h
@@ -116,6 +116,13 @@ typedef struct hsa_isa_s {
   uint64_t handle;
 } hsa_isa_t;
 
+typedef enum {
+  HSA_ACCESS_PERMISSION_NONE = 0,
+  HSA_ACCESS_PERMISSION_RO = 1,
+  HSA_ACCESS_PERMISSION_WO = 2,
+  HSA_ACCESS_PERMISSION_RW = 3
+} hsa_access_permission_t;
+
 hsa_status_t hsa_system_get_info(hsa_system_info_t attribute, void *value);
 
 hsa_status_t hsa_agent_get_info(hsa_agent_t agent, hsa_agent_info_t attribute,
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 7ff77f8e2a2fa..330a14a8ed452 100644
--- a/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa_ext_amd.h
+++ b/offload/plugins-nextgen/amdgpu/dynamic_hsa/hsa_ext_amd.h
@@ -163,6 +163,20 @@ typedef struct hsa_amd_pointer_info_s {
   size_t sizeInBytes;
 } hsa_amd_pointer_info_t;
 
+typedef enum {
+  MEMORY_TYPE_NONE,
+  MEMORY_TYPE_PINNED,
+} hsa_amd_memory_type_t;
+
+typedef struct hsa_amd_vmem_alloc_handle_s {
+  uint64_t handle;
+} hsa_amd_vmem_alloc_handle_t;
+
+typedef struct hsa_amd_memory_access_desc_s {
+  hsa_access_permission_t permissions;
+  hsa_agent_t agent_handle;
+} hsa_amd_memory_access_desc_t;
+
 hsa_status_t hsa_amd_pointer_info(const void* ptr,
                                           hsa_amd_pointer_info_t* info,
                                           void* (*alloc)(size_t),
@@ -181,6 +195,29 @@ hsa_amd_profiling_get_dispatch_time(hsa_agent_t agent, hsa_signal_t signal,
 hsa_status_t hsa_amd_profiling_set_profiler_enabled(hsa_queue_t *queue,
                                                     int enable);
 
+hsa_status_t hsa_amd_vmem_address_reserve(void **va, size_t size,
+                                          uint64_t address, uint64_t flags);
+
+hsa_status_t hsa_amd_vmem_address_free(void *va, size_t size);
+
+hsa_status_t
+hsa_amd_vmem_handle_create(hsa_amd_memory_pool_t pool, size_t size,
+                           hsa_amd_memory_type_t type, uint64_t flags,
+                           hsa_amd_vmem_alloc_handle_t *memory_handle);
+
+hsa_status_t
+hsa_amd_vmem_handle_release(hsa_amd_vmem_alloc_handle_t memory_handle);
+
+hsa_status_t hsa_amd_vmem_map(void *va, size_t size, size_t in_offset,
+                              hsa_amd_vmem_alloc_handle_t memory_handle,
+                              uint64_t flags);
+
+hsa_status_t hsa_amd_vmem_unmap(void *va, size_t size);
+
+hsa_status_t hsa_amd_vmem_set_access(void *va, size_t size,
+                                     const hsa_amd_memory_access_desc_t *desc,
+                                     size_t desc_cnt);
+
 #ifdef __cplusplus
 }
 #endif
diff --git a/offload/plugins-nextgen/amdgpu/src/rtl.cpp b/offload/plugins-nextgen/amdgpu/src/rtl.cpp
index 688f643b5b829..9ceca83cf6de5 100644
--- a/offload/plugins-nextgen/amdgpu/src/rtl.cpp
+++ b/offload/plugins-nextgen/amdgpu/src/rtl.cpp
@@ -290,7 +290,7 @@ struct AMDGPUMemoryPoolTy {
     if (auto Err = getAttr(HSA_AMD_MEMORY_POOL_INFO_GLOBAL_FLAGS, GlobalFlags))
       return Err;
 
-    return Plugin::success();
+    return getAttr(HSA_AMD_MEMORY_POOL_INFO_RUNTIME_ALLOC_GRANULE, Granule);
   }
 
   /// Getter of the HSA memory pool.
@@ -320,6 +320,9 @@ struct AMDGPUMemoryPoolTy {
     return (GlobalFlags & HSA_AMD_MEMORY_POOL_GLOBAL_FLAG_KERNARG_INIT);
   }
 
+  /// Get the page size.
+  size_t getGranule() const { return Granule; }
+
   /// Allocate memory on the memory pool.
   Error allocate(size_t Size, void **PtrStorage) {
     hsa_status_t Status =
@@ -400,6 +403,9 @@ struct AMDGPUMemoryPoolTy {
   /// The global flags of memory pool. Only valid if the memory pool belongs to
   /// the global segment.
   uint32_t GlobalFlags;
+
+  /// The page size in this memory pool.
+  size_t Granule;
 };
 
 /// Class that implements a memory manager that gets memory from a specific
@@ -2320,6 +2326,86 @@ struct AMDGPUDeviceTy : public GenericDeviceTy, AMDGenericDeviceTy {
     return Plugin::success();
   }
 
+  /// Suggest a virtual address for device memory mapping.
+  void *getSuggestedVirtualAddress() override {
+    return reinterpret_cast<void *>(0x1534f7e00000ULL);
+  }
+
+  /// Allocate \p Size bytes on the device and hints the backend to map it to
+  /// virtual address \p VAddr. The function returns the allocated virtual
+  /// address. The memory must be deallocated through
+  /// GenericDeviceTy::deallocateWithVirtualAddress().
+  Expected<void *> allocateWithVirtualAddress(uint64_t Size,
+                                              void *VAddr) override {
+    uint64_t ExpectedVAddr = 0;
+    if (VAddr != nullptr)
+      ExpectedVAddr = static_cast<uint64_t>(reinterpret_cast<uintptr_t>(VAddr));
+
+    // Transparently round up to a multiple of the page size.
+    auto *Pool = CoarseGrainedMemoryPools[0];
+    Size = utils::roundUp(Size, (uint64_t)Pool->getGranule());
+
+    // Reserve the virtual address range.
+    hsa_status_t Status =
+        hsa_amd_vmem_address_reserve(&VAddr, Size, ExpectedVAddr, 0);
+    if (auto Err = Plugin::check(Status,
+                                 "error in hsa_amd_vmem_address_reserve: %s\n"))
+      return Err;
+
+    // Create a handle of the allocation.
+    hsa_amd_vmem_alloc_handle_t Handle;
+    Status = hsa_amd_vmem_handle_create(Pool->get(), Size, MEMORY_TYPE_PINNED,
+                                        0, &Handle);
+    if (auto Err =
+            Plugin::check(Status, "error in hsa_amd_vmem_handle_create: %s\n"))
+      return Err;
+
+    // Map the virtual address range to the memory allocation.
+    Status = hsa_amd_vmem_map(VAddr, Size, 0, Handle, 0);
+    if (auto Err = Plugin::check(Status, "error in hsa_amd_vmem_map: %s\n"))
+      return Err;
+
+    // Set the memory access properties for the allocation.
+    hsa_amd_memory_access_desc_t Desc;
+    Desc.agent_handle = Agent;
+    Desc.permissions = HSA_ACCESS_PERMISSION_RW;
+    Status = hsa_amd_vmem_set_access(VAddr, Size, &Desc, 1);
+    if (auto Err =
+            Plugin::check(Status, "error in hsa_amd_vmem_set_access: %s\n"))
+      return Err;
+
+    // Register the virtual address range in the tracker.
+    if (auto Err = VMemTracker.registerReservation(VAddr, Size, Handle))
+      return Err;
+
+    return VAddr;
+  }
+
+  /// Deallocate device memory \p VAddr, which was allocated through
+  /// GenericDeviceTy::allocateWithVirtualAddress(), and unmap the virtual
+  /// address range.
+  Error deallocateWithVirtualAddress(void *VAddr, uint64_t) override {
+    // Unregister the virtual address range and obtain the information about
+    // the reservation.
+    auto InfoOrErr = VMemTracker.unregisterReservation(VAddr);
+    if (!InfoOrErr)
+      return InfoOrErr.takeError();
+
+    auto [Size, Handle] = *InfoOrErr;
+
+    hsa_status_t Status = hsa_amd_vmem_unmap(VAddr, Size);
+    if (auto Err = Plugin::check(Status, "error in hsa_amd_vmem_unmap: %s\n"))
+      return Err;
+
+    Status = hsa_amd_vmem_handle_release(Handle);
+    if (auto Err =
+            Plugin::check(Status, "error in hsa_amd_vmem_handle_release: %s\n"))
+      return Err;
+
+    Status = hsa_amd_vmem_address_free(VAddr, Size);
+    return Plugin::check(Status, "error in hsa_amd_vmem_address_free: %s\n");
+  }
+
   Error unloadBinaryImpl(DeviceImageTy *Image) override {
     AMDGPUDeviceImageTy &AMDImage = static_cast<AMDGPUDeviceImageTy &>(*Image);
 
@@ -3204,10 +3290,7 @@ struct AMDGPUDeviceTy : public GenericDeviceTy, AMDGenericDeviceTy {
       if (Status == HSA_STATUS_SUCCESS)
         PoolNode.add("Allocatable", TmpBool);
 
-      Status = Pool->getAttrRaw(HSA_AMD_MEMORY_POOL_INFO_RUNTIME_ALLOC_GRANULE,
-                                TmpSt);
-      if (Status == HSA_STATUS_SUCCESS)
-        PoolNode.add("Runtime Alloc Granule", TmpSt, "bytes");
+      PoolNode.add("Runtime Alloc Granule", Pool->getGranule(), "bytes");
 
       Status = Pool->getAttrRaw(
           HSA_AMD_MEMORY_POOL_INFO_RUNTIME_ALLOC_ALIGNMENT, TmpSt);
@@ -3502,6 +3585,9 @@ struct AMDGPUDeviceTy : public GenericDeviceTy, AMDGenericDeviceTy {
   /// True is the system is configured with XNACK-Enabled.
   /// False otherwise.
   bool IsXnackEnabled = false;
+
+  /// Tracker for virtual address reservations.
+  VMemTrackerTy<hsa_amd_vmem_alloc_handle_t> VMemTracker;
 };
 
 Error AMDGPUDeviceImageTy::loadExecutable(const AMDGPUDeviceTy &Device) {
diff --git a/offload/plugins-nextgen/common/CMakeLists.txt b/offload/plugins-nextgen/common/CMakeLists.txt
index f57f4d3a84972..cd150d1bf9298 100644
--- a/offload/plugins-nextgen/common/CMakeLists.txt
+++ b/offload/plugins-nextgen/common/CMakeLists.txt
@@ -13,6 +13,7 @@ add_library(PluginCommon OBJECT
   src/PluginInterface.cpp
   src/GlobalHandler.cpp
   src/JIT.cpp
+  src/RecordReplay.cpp
   src/RPC.cpp
   src/OffloadError.cpp
   src/Utils/ELF.cpp
diff --git a/offload/plugins-nextgen/common/include/PluginInterface.h b/offload/plugins-nextgen/common/include/PluginInterface.h
index 7990b09d59c69..757f117314b87 100644
--- a/offload/plugins-nextgen/common/include/PluginInterface.h
+++ b/offload/plugins-nextgen/common/include/PluginInterface.h
@@ -34,13 +34,18 @@
 #include "MemoryManager.h"
 #include "OffloadError.h"
 #include "RPC.h"
+#include "RecordReplay.h"
 #include "omptarget.h"
 
 #ifdef OMPT_SUPPORT
 #include "omp-tools.h"
 #endif
 
+#include "llvm/ADT/DenseMapInfo.h"
+#include "llvm/ADT/DenseSet.h"
+#include "llvm/ADT/Hashing.h"
 #include "llvm/ADT/SmallVector.h"
+#include "llvm/ADT/StableHashing.h"
 #include "llvm/Frontend/OpenMP/OMPConstants.h"
 #include "llvm/Frontend/OpenMP/OMPGridValues.h"
 #include "llvm/Support/Allocator.h"
@@ -61,7 +66,6 @@ namespace plugin {
 struct GenericPluginTy;
 struct GenericKernelTy;
 struct GenericDeviceTy;
-struct RecordReplayTy;
 template <typename ResourceRef> class GenericDeviceResourceManagerTy;
 
 namespace Plugin {
@@ -313,6 +317,45 @@ struct DynBlockMemConfTy {
   void *FallbackPtr = nullptr;
 };
 
+/// Tracker of virtual memory address reservations.
+template <typename HandleTy> class VMemTrackerTy {
+  struct EntryTy {
+    uint64_t Size;
+    HandleTy Handle;
+  };
+
+  /// Map of virtual memory address reservations.
+  DenseMap<void *, EntryTy> VMemMap;
+
+  /// Mutex for safe access to the map.
+  std::mutex Mutex;
+
+public:
+  /// Register a new virtual address reservation.
+  Error registerReservation(void *VAddr, uint64_t Size, HandleTy Handle) {
+    std::lock_guard<std::mutex> Lock(Mutex);
+    auto It = VMemMap.find(VAddr);
+    if (It != VMemMap.end())
+      return Plugin::error(error::ErrorCode::INVALID_ARGUMENT,
+                           "virtual address already reserved");
+    VMemMap[VAddr] = {Size, Handle};
+    return Plugin::success();
+  }
+
+  /// Unregister a virtual address reservation and return its information.
+  Expected<std::pair<uint64_t, HandleTy>> unregisterReservation(void *VAddr) {
+    std::lock_guard<std::mutex> Lock(Mutex);
+    auto It = VMemMap.find(VAddr);
+    if (It == VMemMap.end())
+      return Plugin::error(error::ErrorCode::INVALID_ARGUMENT,
+                           "virtual address not reserved");
+    uint64_t Size = It->second.Size;
+    HandleTy Handle = It->second.Handle;
+    VMemMap.erase(It);
+    return std::make_pair(Size, Handle);
+  }
+};
+
 /// Class wrapping a __tgt_device_image and its offload entry table on a
 /// specific device. This class is responsible for storing and managing
 /// the offload entries for an image on a device.
@@ -372,7 +415,8 @@ struct GenericKernelTy {
   /// one used to initialize the kernel.
   Error launch(GenericDeviceTy &GenericDevice, void **ArgPtrs,
                ptrdiff_t *ArgOffsets, KernelArgsTy &KernelArgs,
-               AsyncInfoWrapperTy &AsyncInfoWrapper) const;
+               AsyncInfoWrapperTy &AsyncInfoWrapper,
+               RecordReplayTy::HandleTy *RRHandle = nullptr) const;
   virtual Error launchImpl(GenericDeviceTy &GenericDevice,
                            uint32_t NumThreads[3], uint32_t NumBlocks[3],
                            uint32_t DynBlockMemSize, KernelArgsTy &KernelArgs,
@@ -778,6 +822,27 @@ struct GenericDeviceTy : public DeviceAllocatorTy {
   GenericDeviceTy(GenericPluginTy &Plugin, int32_t DeviceId, int32_t NumDevices,
                   const llvm::omp::GV &GridValues);
 
+  /// Suggest a virtual address for device memory mapping.
+  virtual void *getSuggestedVirtualAddress() { return nullptr; }
+
+  /// Allocate \p Size bytes on the device and hints the backend to map it to
+  /// virtual address \p VAddr. The function returns the allocated virtual
+  /// address. The memory must be deallocated through
+  /// GenericDeviceTy::deallocateWithVirtualAddress().
+  virtual Expected<void *> allocateWithVirtualAddress(uint64_t Size,
+                                                      void *VAddr = nullptr) {
+    return Plugin::error(error::ErrorCode::UNSUPPORTED,
+                         "allocate with virtual address not supported");
+  }
+
+  /// Deallocate device memory \p VAddr, which was allocated through
+  /// GenericDeviceTy::allocateWithVirtualAddress(), and unmap the virtual
+  /// address range.
+  virtual Error deallocateWithVirtualAddress(void *VAddr, uint64_t Size) {
+    return Plugin::error(error::ErrorCode::UNSUPPORTED,
+                         "allocate with virtual address not supported");
+  }
+
   /// Get the device identifier within the corresponding plugin. Notice that
   /// this id is not unique between different plugins; they may overlap.
   int32_t getDeviceId() const { return DeviceId; }
@@ -1156,6 +1221,26 @@ struct GenericDeviceTy : public DeviceAllocatorTy {
     return ATI;
   }
 
+  Error initRecordReplay(int64_t Size, void *VAddr, bool IsRecord,
+                         bool IsNative, bool SaveOutput) {
+    if (RecordReplay)
+      return Plugin::error(error::ErrorCode::INVALID_ARGUMENT,
+                           "RR already initialized");
+    if (!IsNative)
+      return Plugin::error(error::ErrorCode::UNSUPPORTED,
+                           "non-native RR not available");
+
+    RecordReplayTy::StatusTy Status = IsRecord
+                                          ? RecordReplayTy::StatusTy::Recording
+                                          : RecordReplayTy::StatusTy::Replaying;
+
+    RecordReplay = new NativeRecordReplayTy(Status, SaveOutput, *this);
+
+    return RecordReplay->init(Size, VAddr);
+  }
+
+  RecordReplayTy *getRecordReplay() { return RecordReplay; }
+
   /// Map to record kernel have been launchedl, for error reporting purposes.
   ProtectedObj<KernelTraceInfoRecordTy> KernelLaunchTraces;
 
@@ -1218,6 +1303,9 @@ struct GenericDeviceTy : public DeviceAllocatorTy {
   /// Indicate whether failures when locking mapped buffers should be ignored.
   bool IgnoreLockMappedFailures;
 
+  /// Record and replay manager.
+  RecordReplayTy *RecordReplay = nullptr;
+
 protected:
   /// Environment variables defined by the LLVM OpenMP implementation
   /// regarding the initial number of streams and events.
@@ -1281,8 +1369,7 @@ struct GenericPluginTy {
 
   /// Construct a plugin instance.
   GenericPluginTy(Triple::ArchType TA)
-      : GlobalHandler(nullptr), JIT(TA), RPCServer(nullptr),
-        RecordReplay(nullptr) {}
+      : GlobalHandler(nullptr), JIT(TA), RPCServer(nullptr) {}
 
   virtual ~GenericPluginTy() {}
 
@@ -1374,11 +1461,6 @@ struct GenericPluginTy {
   virtual Error deinitRPCDoorbell() { return Plugin::success(); }
 
   /// Get a reference to the record and replay interface for the plugin.
-  RecordReplayTy &getRecordReplay() {
-    assert(RecordReplay && "RR interface not initialized");
-    return *RecordReplay;
-  }
-
   /// Initialize a device within the plugin.
   Error initDevice(int32_t DeviceId);
 
@@ -1471,8 +1553,8 @@ struct GenericPluginTy {
 
   /// Initializes the record and replay mechanism inside the plugin.
   int32_t initialize_record_replay(int32_t DeviceId, int64_t MemorySize,
-                                   void *VAddr, bool isRecord, bool SaveOutput,
-                                   uint64_t &ReqPtrArgOffset);
+                                   void *VAddr, bool IsRecord, bool IsNative,
+                                   bool SaveOutput, uint64_t &ReqPtrArgOffset);
 
   /// Loads the associated binary into the plugin and returns a handle to it.
   int32_t load_binary(int32_t DeviceId, __tgt_device_image *TgtImage,
@@ -1648,9 +1730,6 @@ struct GenericPluginTy {
 
   /// The interface between the plugin and the GPU for host services.
   RPCServerTy *RPCServer;
-
-  /// The interface between the plugin and the GPU for host services.
-  RecordReplayTy *RecordReplay;
 };
 
 /// Auxiliary interface class for GenericDeviceResourceManagerTy. This class
diff --git a/offload/plugins-nextgen/common/include/RecordReplay.h b/offload/plugins-nextgen/common/include/RecordReplay.h
new file mode 100644
index 0000000000000..41145dfe97144
--- /dev/null
+++ b/offload/plugins-nextgen/common/include/RecordReplay.h
@@ -0,0 +1,224 @@
+//===- RecordReplay.h - Record Replay interface ---------------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+//
+//===----------------------------------------------------------------------===//
+
+#ifndef OPENMP_LIBOMPTARGET_PLUGINS_NEXTGEN_COMMON_RECORDREPLAY_H
+#define OPENMP_LIBOMPTARGET_PLUGINS_NEXTGEN_COMMON_RECORDREPLAY_H
+
+#include <cstddef>
+#include <cstdint>
+#include <mutex>
+#include <unordered_set>
+
+#include "Shared/APITypes.h"
+#include "Shared/EnvironmentVar.h"
+#include "Shared/Utils.h"
+
+#include "OffloadError.h"
+
+#include "llvm/ADT/Hashing.h"
+#include "llvm/ADT/SmallVector.h"
+#include "llvm/ADT/StableHashing.h"
+#include "llvm/Support/Error.h"
+#include "llvm/Support/ErrorHandling.h"
+#include "llvm/Support/MemoryBufferRef.h"
+#include "llvm/Support/raw_ostream.h"
+
+namespace llvm {
+namespace omp {
+namespace target {
+namespace plugin {
+
+struct GenericKernelTy;
+struct GenericDeviceTy;
+
+struct RecordReplayTy {
+protected:
+  struct InstanceTy;
+
+public:
+  /// Describes the state of the record replay mechanism.
+  enum StatusTy { Deactivated = 0, Recording, Replaying };
+
+  /// Describes the format of the recording and replaying.
+  enum FormatTy { Native = 0 };
+
+  struct HandleTy {
+    const InstanceTy *Instance = nullptr;
+    bool Active = false;
+  };
+
+protected:
+  /// Address and size of record replay memory space.
+  void *StartAddr = nullptr;
+  uint64_t TotalSize = 0;
+  uint64_t CurrentSize = 0;
+  std::mutex AllocationLock;
+
+  /// Status of the record or replay.
+  StatusTy Status;
+
+  /// Whether the record replay should save a memory snapshot after a kernel
+  /// execution.
+  bool SaveOutput;
+
+  /// Reference to the corresponding device.
+  GenericDeviceTy &Device;
+
+  /// The information for a global.
+  struct GlobalEntryTy {
+    std::string Name;
+    uint64_t Size;
+    void *Addr;
+  };
+
+  /// List of all globals mapped to the device.
+  llvm::SmallVector<GlobalEntryTy> GlobalEntries;
+
+  // An instance of a kernel record replay.
+  struct InstanceTy {
+    /// The launch configuration parameters.
+    uint32_t NumTeams = 0;
+    uint32_t NumThreads = 0;
+    uint32_t SharedMemorySize = 0;
+
+    /// The hashes representing the kernel and the launch configuration.
+    size_t KernelHash = 0;
+    size_t LaunchConfigHash = 0;
+
+    /// The number of occurrences during the execution.
+    mutable size_t Occurrences = 0;
+
+    InstanceTy(StringRef KernelName, uint32_t NumTeams, uint32_t NumThreads,
+               uint32_t SharedMemorySize)
+        : NumTeams(NumTeams), NumThreads(NumThreads),
+          SharedMemorySize(SharedMemorySize) {
+      KernelHash = stable_hash_name(KernelName);
+      LaunchConfigHash =
+          stable_hash_combine((stable_hash)NumTeams, (stable_hash)NumThreads,
+                              (stable_hash)SharedMemorySize);
+    }
+
+    bool operator==(const InstanceTy &Other) const {
+      return (KernelHash == Other.KernelHash &&
+              LaunchConfigHash == Other.LaunchConfigHash &&
+              NumTeams == Other.NumTeams && NumThreads == Other.NumThreads &&
+              SharedMemorySize == Other.SharedMemorySize);
+    }
+  };
+
+  struct InstanceHasher {
+    std::size_t operator()(const InstanceTy &I) const {
+      llvm::stable_hash H =
+          llvm::stable_hash_combine(I.KernelHash, I.LaunchConfigHash);
+      return static_cast<std::size_t>(H);
+    }
+  };
+
+  /// Tracker of record replay instances.
+  std::unordered_set<InstanceTy, InstanceHasher> Instances;
+  std::mutex InstancesLock;
+
+public:
+  RecordReplayTy(StatusTy Status, bool SaveOutput, GenericDeviceTy &Device)
+      : Status(Status), SaveOutput(SaveOutput), Device(Device) {}
+
+  virtual ~RecordReplayTy() = default;
+
+  /// Initialize kernel record replay for the corresponding device.
+  Error init(uint64_t MemSize, void *VAddr);
+  Error deinit();
+
+  bool isRecording() const { return Status == StatusTy::Recording; }
+  bool isReplaying() const { return Status == StatusTy::Replaying; }
+  bool isRecordingOrReplaying() const { return isRecording() || isReplaying(); }
+  bool shouldRecordPrologue() const { return isRecording(); }
+  bool shouldRecordEpilogue() const {
+    return isRecordingOrReplaying() && SaveOutput;
+  }
+
+  /// Add information about a global.
+  void addGlobal(const char *Name, uint64_t Size, void *Addr) {
+    GlobalEntries.emplace_back(GlobalEntryTy{Name, Size, Addr});
+  }
+
+  /// Record the prologue and return the handle. This phase can include the
+  /// recording of memory snapshot, the record descriptor and the globals.
+  Expected<HandleTy>
+  recordPrologue(const GenericKernelTy &Kernel, const KernelArgsTy &KernelArgs,
+                 const KernelLaunchParamsTy &LaunchParams, uint32_t NumTeams[3],
+                 uint32_t NumThreads[3], uint32_t SharedMemorySize);
+
+  /// Record the epilogue, which can include the memory snapshot when recording
+  /// or replaying.
+  Error recordEpilogue(const GenericKernelTy &Kernel, HandleTy Handle);
+
+  /// Allocates device memory from the record replay space.
+  void *allocate(uint64_t Size);
+
+private:
+  /// Register an instance and return a reference and whether it was registered
+  /// as a new instance.
+  std::pair<const InstanceTy &, bool>
+  registerInstance(StringRef KernelName, uint32_t NumTeams, uint32_t NumThreads,
+                   uint32_t SharedMemorySize);
+
+  /// The interface that should be provided by kernel record replay
+  /// implementations.
+  virtual Error
+  recordPrologueImpl(const GenericKernelTy &Kernel, const InstanceTy &Instance,
+                     const KernelArgsTy &KernelArgs,
+                     const KernelLaunchParamsTy &LaunchParams) = 0;
+  virtual Error recordEpilogueImpl(const GenericKernelTy &Kernel,
+                                   const InstanceTy &Instance) = 0;
+  virtual Error recordDescriptorImpl(const GenericKernelTy &Kernel,
+                                     const InstanceTy &Instance,
+                                     const KernelArgsTy &KernelArgs,
+                                     const KernelLaunchParamsTy &LaunchParams,
+                                     uint32_t NumTeams[3],
+                                     uint32_t NumThreads[3],
+                                     uint32_t SharedMemorySize) = 0;
+};
+
+/// The native kernel record replay support.
+struct NativeRecordReplayTy : public RecordReplayTy {
+  NativeRecordReplayTy(StatusTy Status, bool SaveOutput,
+                       GenericDeviceTy &Device)
+      : RecordReplayTy(Status, SaveOutput, Device) {}
+
+private:
+  Error recordPrologueImpl(const GenericKernelTy &Kernel,
+                           const InstanceTy &Instance,
+                           const KernelArgsTy &KernelArgs,
+                           const KernelLaunchParamsTy &LaunchParams) override;
+  Error recordEpilogueImpl(const GenericKernelTy &Kernel,
+                           const InstanceTy &Instance) override;
+  Error recordDescriptorImpl(const GenericKernelTy &Kernel,
+                             const InstanceTy &Instance,
+                             const KernelArgsTy &KernelArgs,
+                             const KernelLaunchParamsTy &LaunchParams,
+                             uint32_t NumTeams[3], uint32_t NumThreads[3],
+                             uint32_t SharedMemorySize) override;
+
+  /// Record a memory snapshot on a file.
+  Error recordSnapshot(StringRef Filename);
+
+  /// Record the globals on a file.
+  Error recordGlobals(StringRef Filename);
+
+  /// Record the device image on a file.
+  Error recordImage(const GenericKernelTy &Kernel, StringRef Filename);
+};
+
+} // namespace plugin
+} // namespace target
+} // namespace omp
+} // namespace llvm
+
+#endif // OPENMP_LIBOMPTARGET_PLUGINS_COMMON_RECORDREPLAY_H
diff --git a/offload/plugins-nextgen/common/src/PluginInterface.cpp b/offload/plugins-nextgen/common/src/PluginInterface.cpp
index 3420678cac98b..a0ed19c54d7fd 100644
--- a/offload/plugins-nextgen/common/src/PluginInterface.cpp
+++ b/offload/plugins-nextgen/common/src/PluginInterface.cpp
@@ -29,7 +29,6 @@
 #include "llvm/Bitcode/BitcodeReader.h"
 #include "llvm/Frontend/OpenMP/OMPConstants.h"
 #include "llvm/Support/Error.h"
-#include "llvm/Support/JSON.h"
 #include "llvm/Support/MathExtras.h"
 #include "llvm/Support/MemoryBuffer.h"
 #include "llvm/Support/Signals.h"
@@ -45,339 +44,6 @@ using namespace plugin;
 using namespace error;
 using namespace llvm::offload::debug;
 
-// TODO: Fix any thread safety issues for multi-threaded kernel recording.
-namespace llvm::omp::target::plugin {
-struct RecordReplayTy {
-
-  // Describes the state of the record replay mechanism.
-  enum RRStatusTy { RRDeactivated = 0, RRRecording, RRReplaying };
-
-private:
-  // Memory pointers for recording, replaying memory.
-  void *MemoryStart = nullptr;
-  void *MemoryPtr = nullptr;
-  size_t MemorySize = 0;
-  size_t TotalSize = 0;
-  GenericDeviceTy *Device = nullptr;
-  std::mutex AllocationLock;
-
-  RRStatusTy Status = RRDeactivated;
-  bool ReplaySaveOutput = false;
-  bool UsedVAMap = false;
-  uintptr_t MemoryOffset = 0;
-
-  // A list of all globals mapped to the device.
-  struct GlobalEntry {
-    const char *Name;
-    uint64_t Size;
-    void *Addr;
-  };
-  llvm::SmallVector<GlobalEntry> GlobalEntries{};
-
-  Expected<void *> suggestAddress(uint64_t MaxMemoryAllocation) {
-    // Get a valid pointer address for this system
-    auto AddrOrErr =
-        Device->allocate(1024, /*HstPtr=*/nullptr, TARGET_ALLOC_DEFAULT);
-    if (!AddrOrErr)
-      return AddrOrErr.takeError();
-
-    void *Addr = *AddrOrErr;
-    if (auto Err = Device->free(Addr))
-      return std::move(Err);
-
-    // Align Address to MaxMemoryAllocation
-    Addr = (void *)utils::alignPtr((Addr), MaxMemoryAllocation);
-    return Addr;
-  }
-
-  Error preAllocateVAMemory(uint64_t MaxMemoryAllocation, void *VAddr) {
-    size_t ASize = MaxMemoryAllocation;
-
-    if (!VAddr && isRecording()) {
-      auto VAddrOrErr = suggestAddress(MaxMemoryAllocation);
-      if (!VAddrOrErr)
-        return VAddrOrErr.takeError();
-      VAddr = *VAddrOrErr;
-    }
-
-    ODBG(OLDT_Alloc) << "Request " << MaxMemoryAllocation
-                     << " bytes allocated at " << VAddr;
-
-    if (auto Err = Device->memoryVAMap(&MemoryStart, VAddr, &ASize))
-      return Err;
-
-    if (isReplaying() && VAddr != MemoryStart) {
-      return Plugin::error(ErrorCode::INVALID_ARGUMENT,
-                           "record-Replay cannot assign the"
-                           "requested recorded address (%p, %p)",
-                           VAddr, MemoryStart);
-    }
-
-    INFO(OMP_INFOTYPE_PLUGIN_KERNEL, Device->getDeviceId(),
-         "Allocated %" PRIu64 " bytes at %p for replay.\n", ASize, MemoryStart);
-
-    MemoryPtr = MemoryStart;
-    MemorySize = 0;
-    TotalSize = ASize;
-    UsedVAMap = true;
-    return Plugin::success();
-  }
-
-  Error preAllocateHeuristic(uint64_t MaxMemoryAllocation,
-                             uint64_t RequiredMemoryAllocation, void *VAddr) {
-    const size_t MAX_MEMORY_ALLOCATION = MaxMemoryAllocation;
-    constexpr size_t STEP = 1024 * 1024 * 1024ULL;
-    MemoryStart = nullptr;
-    for (TotalSize = MAX_MEMORY_ALLOCATION; TotalSize > 0; TotalSize -= STEP) {
-      auto MemoryStartOrErr =
-          Device->allocate(TotalSize, /*HstPtr=*/nullptr, TARGET_ALLOC_DEFAULT);
-      if (!MemoryStartOrErr)
-        return MemoryStartOrErr.takeError();
-      MemoryStart = *MemoryStartOrErr;
-      if (MemoryStart)
-        break;
-    }
-    if (!MemoryStart)
-      return Plugin::error(ErrorCode::INVALID_ARGUMENT,
-                           "allocating record/replay memory");
-
-    if (VAddr && VAddr != MemoryStart)
-      MemoryOffset = uintptr_t(VAddr) - uintptr_t(MemoryStart);
-
-    MemoryPtr = MemoryStart;
-    MemorySize = 0;
-
-    // Check if we need adjustment.
-    if (MemoryOffset > 0 &&
-        TotalSize >= RequiredMemoryAllocation + MemoryOffset) {
-      // If we are off but "before" the required address and with enough space,
-      // we just "allocate" the offset to match the required address.
-      MemoryPtr = (char *)MemoryPtr + MemoryOffset;
-      MemorySize += MemoryOffset;
-      MemoryOffset = 0;
-      assert(MemoryPtr == VAddr && "Expected offset adjustment to work");
-    } else if (MemoryOffset) {
-      // If we are off and in a situation we cannot just "waste" memory to force
-      // a match, we hope adjusting the arguments is sufficient.
-      REPORT() << "WARNING Failed to allocate replay memory at required "
-               << "location " << VAddr << ", got " << MemoryStart
-               << ", trying to offset argument pointers by " << MemoryOffset;
-    }
-
-    INFO(OMP_INFOTYPE_PLUGIN_KERNEL, Device->getDeviceId(),
-         "Allocated %" PRIu64 " bytes at %p for replay.\n", TotalSize,
-         MemoryStart);
-
-    return Plugin::success();
-  }
-
-  Error preallocateDeviceMemory(uint64_t DeviceMemorySize, void *ReqVAddr) {
-    if (Device->supportVAManagement()) {
-      auto Err = preAllocateVAMemory(DeviceMemorySize, ReqVAddr);
-      if (Err) {
-        REPORT() << "WARNING VA mapping failed, fallback to heuristic: "
-                 << "(Error: " << toString(std::move(Err)) << ")";
-      }
-    }
-
-    uint64_t DevMemSize;
-    if (Device->getDeviceMemorySize(DevMemSize))
-      return Plugin::error(ErrorCode::UNKNOWN,
-                           "cannot determine Device Memory Size");
-
-    return preAllocateHeuristic(DevMemSize, DeviceMemorySize, ReqVAddr);
-  }
-
-  void dumpDeviceMemory(StringRef Filename) {
-    ErrorOr<std::unique_ptr<WritableMemoryBuffer>> DeviceMemoryMB =
-        WritableMemoryBuffer::getNewUninitMemBuffer(MemorySize);
-    if (!DeviceMemoryMB)
-      report_fatal_error("Error creating MemoryBuffer for device memory");
-
-    auto Err = Device->dataRetrieve(DeviceMemoryMB.get()->getBufferStart(),
-                                    MemoryStart, MemorySize, nullptr);
-    if (Err)
-      report_fatal_error("Error retrieving data for target pointer");
-
-    StringRef DeviceMemory(DeviceMemoryMB.get()->getBufferStart(), MemorySize);
-    std::error_code EC;
-    raw_fd_ostream OS(Filename, EC);
-    if (EC)
-      report_fatal_error("Error dumping memory to file " + Filename + " :" +
-                         EC.message());
-    OS << DeviceMemory;
-    OS.close();
-  }
-
-public:
-  bool isRecording() const { return Status == RRStatusTy::RRRecording; }
-  bool isReplaying() const { return Status == RRStatusTy::RRReplaying; }
-  bool isRecordingOrReplaying() const {
-    return (Status != RRStatusTy::RRDeactivated);
-  }
-  void setStatus(RRStatusTy Status) { this->Status = Status; }
-  bool isSaveOutputEnabled() const { return ReplaySaveOutput; }
-  void addEntry(const char *Name, uint64_t Size, void *Addr) {
-    GlobalEntries.emplace_back(GlobalEntry{Name, Size, Addr});
-  }
-
-  void saveImage(const char *Name, const DeviceImageTy &Image) {
-    SmallString<128> ImageName = {Name, ".image"};
-    std::error_code EC;
-    raw_fd_ostream OS(ImageName, EC);
-    if (EC)
-      report_fatal_error("Error saving image : " + StringRef(EC.message()));
-    OS << Image.getMemoryBuffer().getBuffer();
-    OS.close();
-  }
-
-  void dumpGlobals(StringRef Filename, DeviceImageTy &Image) {
-    int32_t Size = 0;
-
-    for (auto &OffloadEntry : GlobalEntries) {
-      if (!OffloadEntry.Size)
-        continue;
-      // Get the total size of the string and entry including the null byte.
-      Size += std::strlen(OffloadEntry.Name) + 1 + sizeof(uint32_t) +
-              OffloadEntry.Size;
-    }
-
-    ErrorOr<std::unique_ptr<WritableMemoryBuffer>> GlobalsMB =
-        WritableMemoryBuffer::getNewUninitMemBuffer(Size);
-    if (!GlobalsMB)
-      report_fatal_error("Error creating MemoryBuffer for globals memory");
-
-    void *BufferPtr = GlobalsMB.get()->getBufferStart();
-    for (auto &OffloadEntry : GlobalEntries) {
-      if (!OffloadEntry.Size)
-        continue;
-
-      int32_t NameLength = std::strlen(OffloadEntry.Name) + 1;
-      memcpy(BufferPtr, OffloadEntry.Name, NameLength);
-      BufferPtr = utils::advancePtr(BufferPtr, NameLength);
-
-      *((uint32_t *)(BufferPtr)) = OffloadEntry.Size;
-      BufferPtr = utils::advancePtr(BufferPtr, sizeof(uint32_t));
-
-      auto Err = Plugin::success();
-      {
-        if (auto Err = Device->dataRetrieve(BufferPtr, OffloadEntry.Addr,
-                                            OffloadEntry.Size, nullptr))
-          report_fatal_error("Error retrieving data for global");
-      }
-      if (Err)
-        report_fatal_error("Error retrieving data for global");
-      BufferPtr = utils::advancePtr(BufferPtr, OffloadEntry.Size);
-    }
-    assert(BufferPtr == GlobalsMB->get()->getBufferEnd() &&
-           "Buffer over/under-filled.");
-    assert(Size == utils::getPtrDiff(BufferPtr,
-                                     GlobalsMB->get()->getBufferStart()) &&
-           "Buffer size mismatch");
-
-    StringRef GlobalsMemory(GlobalsMB.get()->getBufferStart(), Size);
-    std::error_code EC;
-    raw_fd_ostream OS(Filename, EC);
-    OS << GlobalsMemory;
-    OS.close();
-  }
-
-  void saveKernelDescr(const char *Name, KernelLaunchParamsTy LaunchParams,
-                       int32_t NumArgs, uint64_t NumTeamsClause,
-                       uint32_t ThreadLimitClause, uint64_t LoopTripCount) {
-    json::Object JsonKernelInfo;
-    JsonKernelInfo["Name"] = Name;
-    JsonKernelInfo["NumArgs"] = NumArgs;
-    JsonKernelInfo["NumTeamsClause"] = NumTeamsClause;
-    JsonKernelInfo["ThreadLimitClause"] = ThreadLimitClause;
-    JsonKernelInfo["LoopTripCount"] = LoopTripCount;
-    JsonKernelInfo["DeviceMemorySize"] = MemorySize;
-    JsonKernelInfo["DeviceId"] = Device->getDeviceId();
-    JsonKernelInfo["BumpAllocVAStart"] = (intptr_t)MemoryStart;
-
-    json::Array JsonArgPtrs;
-    for (int I = 0; I < NumArgs; ++I)
-      JsonArgPtrs.push_back((intptr_t)LaunchParams.Ptrs[I]);
-    JsonKernelInfo["ArgPtrs"] = json::Value(std::move(JsonArgPtrs));
-
-    json::Array JsonArgOffsets;
-    for (int I = 0; I < NumArgs; ++I)
-      JsonArgOffsets.push_back(0);
-    JsonKernelInfo["ArgOffsets"] = json::Value(std::move(JsonArgOffsets));
-
-    SmallString<128> JsonFilename = {Name, ".json"};
-    std::error_code EC;
-    raw_fd_ostream JsonOS(JsonFilename.str(), EC);
-    if (EC)
-      report_fatal_error("Error saving kernel json file : " +
-                         StringRef(EC.message()));
-    JsonOS << json::Value(std::move(JsonKernelInfo));
-    JsonOS.close();
-  }
-
-  void saveKernelInput(const char *Name, DeviceImageTy &Image) {
-    SmallString<128> GlobalsFilename = {Name, ".globals"};
-    dumpGlobals(GlobalsFilename, Image);
-
-    SmallString<128> MemoryFilename = {Name, ".memory"};
-    dumpDeviceMemory(MemoryFilename);
-  }
-
-  void saveKernelOutputInfo(const char *Name) {
-    SmallString<128> OutputFilename = {
-        Name, (isRecording() ? ".original.output" : ".replay.output")};
-    dumpDeviceMemory(OutputFilename);
-  }
-
-  void *alloc(uint64_t Size) {
-    assert(MemoryStart && "Expected memory has been pre-allocated");
-    void *Alloc = nullptr;
-    constexpr int Alignment = 16;
-    // Assumes alignment is a power of 2.
-    int64_t AlignedSize = (Size + (Alignment - 1)) & (~(Alignment - 1));
-    std::lock_guard<std::mutex> LG(AllocationLock);
-    Alloc = MemoryPtr;
-    MemoryPtr = (char *)MemoryPtr + AlignedSize;
-    MemorySize += AlignedSize;
-    ODBG(OLDT_Alloc) << "Memory Allocator return " << Alloc;
-    return Alloc;
-  }
-
-  Error init(GenericDeviceTy *Device, uint64_t MemSize, void *VAddr,
-             RRStatusTy Status, bool SaveOutput, uint64_t &ReqPtrArgOffset) {
-    this->Device = Device;
-    this->Status = Status;
-    this->ReplaySaveOutput = SaveOutput;
-
-    if (auto Err = preallocateDeviceMemory(MemSize, VAddr))
-      return Err;
-
-    INFO(OMP_INFOTYPE_PLUGIN_KERNEL, Device->getDeviceId(),
-         "Record Replay Initialized (%p)"
-         " as starting address, %lu Memory Size"
-         " and set on status %s\n",
-         MemoryStart, TotalSize,
-         Status == RRStatusTy::RRRecording ? "Recording" : "Replaying");
-
-    // Tell the user to offset pointer arguments as the memory allocation does
-    // not match.
-    ReqPtrArgOffset = MemoryOffset;
-    return Plugin::success();
-  }
-
-  Error deinit() {
-    if (UsedVAMap) {
-      if (auto Err = Device->memoryVAUnMap(MemoryStart, TotalSize))
-        return Err;
-    } else {
-      if (auto Err = Device->free(MemoryStart))
-        return Err;
-    }
-    return Plugin::success();
-  }
-};
-} // namespace llvm::omp::target::plugin
-
 AsyncInfoWrapperTy::AsyncInfoWrapperTy(GenericDeviceTy &Device,
                                        __tgt_async_info *AsyncInfoPtr)
     : Device(Device),
@@ -441,7 +107,8 @@ GenericKernelTy::getKernelLaunchEnvironment(
   // Ctor/Dtor have no arguments, replaying uses the original kernel launch
   // environment. Older versions of the compiler do not generate a kernel
   // launch environment.
-  if (GenericDevice.Plugin.getRecordReplay().isReplaying() ||
+  if ((GenericDevice.getRecordReplay() &&
+       GenericDevice.getRecordReplay()->isReplaying()) ||
       KernelArgs.Version < OMP_KERNEL_ARG_MIN_VERSION_WITH_DYN_PTR)
     return nullptr;
 
@@ -564,7 +231,8 @@ GenericKernelTy::prepareBlockMemory(GenericDeviceTy &GenericDevice,
 
 Error GenericKernelTy::launch(GenericDeviceTy &GenericDevice, void **ArgPtrs,
                               ptrdiff_t *ArgOffsets, KernelArgsTy &KernelArgs,
-                              AsyncInfoWrapperTy &AsyncInfoWrapper) const {
+                              AsyncInfoWrapperTy &AsyncInfoWrapper,
+                              RecordReplayTy::HandleTy *RRHandle) const {
   llvm::SmallVector<void *, 16> Args;
   llvm::SmallVector<void *, 16> Ptrs;
 
@@ -608,13 +276,15 @@ Error GenericKernelTy::launch(GenericDeviceTy &GenericDevice, void **ArgPtrs,
 
   // Record the kernel description after we modified the argument count and num
   // blocks/threads.
-  RecordReplayTy &RecordReplay = GenericDevice.Plugin.getRecordReplay();
-  if (RecordReplay.isRecording()) {
-    RecordReplay.saveImage(getName(), getImage());
-    RecordReplay.saveKernelInput(getName(), getImage());
-    RecordReplay.saveKernelDescr(getName(), LaunchParams, KernelArgs.NumArgs,
-                                 NumBlocks[0], NumThreads[0],
-                                 KernelArgs.Tripcount);
+  RecordReplayTy *RecordReplay = GenericDevice.getRecordReplay();
+  if (RecordReplay) {
+    auto RRHandleOrErr =
+        RecordReplay->recordPrologue(*this, KernelArgs, LaunchParams, NumBlocks,
+                                     NumThreads, DynBlockMemConf.NativeSize);
+    if (!RRHandleOrErr)
+      return RRHandleOrErr.takeError();
+    if (RRHandle)
+      *RRHandle = *RRHandleOrErr;
   }
 
   if (auto Err =
@@ -926,10 +596,12 @@ Error GenericDeviceTy::deinit(GenericPluginTy &Plugin) {
     delete MemoryManager;
   MemoryManager = nullptr;
 
-  RecordReplayTy &RecordReplay = Plugin.getRecordReplay();
-  if (RecordReplay.isRecordingOrReplaying())
-    if (auto Err = RecordReplay.deinit())
+  if (RecordReplay) {
+    if (auto Err = RecordReplay->deinit())
       return Err;
+    delete RecordReplay;
+    RecordReplay = nullptr;
+  }
 
   if (RPCServer)
     if (auto Err = RPCServer->deinitDevice(*this))
@@ -1270,8 +942,8 @@ Expected<void *> GenericDeviceTy::dataAlloc(int64_t Size, void *HostPtr,
                                             TargetAllocTy Kind) {
   void *Alloc = nullptr;
 
-  if (Plugin.getRecordReplay().isRecordingOrReplaying())
-    return Plugin.getRecordReplay().alloc(Size);
+  if (RecordReplay && RecordReplay->isRecordingOrReplaying())
+    return RecordReplay->allocate(Size);
 
   switch (Kind) {
   case TARGET_ALLOC_DEFAULT:
@@ -1335,7 +1007,7 @@ Expected<void *> GenericDeviceTy::dataAlloc(int64_t Size, void *HostPtr,
 
 Error GenericDeviceTy::dataDelete(void *TgtPtr, TargetAllocTy Kind) {
   // Free is a noop when recording or replaying.
-  if (Plugin.getRecordReplay().isRecordingOrReplaying())
+  if (RecordReplay && RecordReplay->isRecordingOrReplaying())
     return Plugin::success();
 
   // Keep track of the deallocation stack if we track allocation traces.
@@ -1434,9 +1106,8 @@ Error GenericDeviceTy::launchKernel(void *EntryPtr, void **ArgPtrs,
                                     ptrdiff_t *ArgOffsets,
                                     KernelArgsTy &KernelArgs,
                                     __tgt_async_info *AsyncInfo) {
-  AsyncInfoWrapperTy AsyncInfoWrapper(
-      *this,
-      Plugin.getRecordReplay().isRecordingOrReplaying() ? nullptr : AsyncInfo);
+  AsyncInfoWrapperTy AsyncInfoWrapper(*this,
+                                      RecordReplay ? nullptr : AsyncInfo);
 
   GenericKernelTy &GenericKernel =
       *reinterpret_cast<GenericKernelTy *>(EntryPtr);
@@ -1453,16 +1124,16 @@ Error GenericDeviceTy::launchKernel(void *EntryPtr, void **ArgPtrs,
         .emplace(&GenericKernel, std::move(StackTrace), AsyncInfo);
   }
 
+  RecordReplayTy::HandleTy RRHandle;
   auto Err = GenericKernel.launch(*this, ArgPtrs, ArgOffsets, KernelArgs,
-                                  AsyncInfoWrapper);
+                                  AsyncInfoWrapper, &RRHandle);
 
   // 'finalize' here to guarantee next record-replay actions are in-sync
   AsyncInfoWrapper.finalize(Err);
 
-  RecordReplayTy &RecordReplay = Plugin.getRecordReplay();
-  if (RecordReplay.isRecordingOrReplaying() &&
-      RecordReplay.isSaveOutputEnabled())
-    RecordReplay.saveKernelOutputInfo(GenericKernel.getName());
+  if (RecordReplay)
+    if (auto Err = RecordReplay->recordEpilogue(GenericKernel, RRHandle))
+      return Err;
 
   return Err;
 }
@@ -1605,9 +1276,6 @@ Error GenericPluginTy::init() {
   RPCServer = new RPCServerTy(*this);
   assert(RPCServer && "Invalid RPC server");
 
-  RecordReplay = new RecordReplayTy();
-  assert(RecordReplay && "Invalid RR interface");
-
   return Plugin::success();
 }
 
@@ -1633,9 +1301,6 @@ Error GenericPluginTy::deinit() {
     delete RPCServer;
   }
 
-  if (RecordReplay)
-    delete RecordReplay;
-
   // Perform last deinitializations on the plugin.
   if (Error Err = deinitImpl())
     return Err;
@@ -1798,25 +1463,17 @@ int32_t GenericPluginTy::is_data_exchangable(int32_t SrcDeviceId,
   return isDataExchangable(SrcDeviceId, DstDeviceId);
 }
 
-int32_t GenericPluginTy::initialize_record_replay(int32_t DeviceId,
-                                                  int64_t MemorySize,
-                                                  void *VAddr, bool isRecord,
-                                                  bool SaveOutput,
-                                                  uint64_t &ReqPtrArgOffset) {
+int32_t GenericPluginTy::initialize_record_replay(
+    int32_t DeviceId, int64_t MemorySize, void *VAddr, bool IsRecord,
+    bool IsNative, bool SaveOutput, uint64_t &ReqPtrArgOffset) {
   GenericDeviceTy &Device = getDevice(DeviceId);
-  RecordReplayTy::RRStatusTy Status =
-      isRecord ? RecordReplayTy::RRStatusTy::RRRecording
-               : RecordReplayTy::RRStatusTy::RRReplaying;
-
-  if (auto Err = RecordReplay->init(&Device, MemorySize, VAddr, Status,
-                                    SaveOutput, ReqPtrArgOffset)) {
-    REPORT() << "WARNING RR did not initialize RR-properly with " << MemorySize
-             << " bytes (Error: " << toString(std::move(Err)) << ")";
-    RecordReplay->setStatus(RecordReplayTy::RRStatusTy::RRDeactivated);
-
-    if (!isRecord) {
-      return OFFLOAD_FAIL;
-    }
+
+  if (auto Err = Device.initRecordReplay(MemorySize, VAddr, IsRecord, IsNative,
+                                         SaveOutput)) {
+    REPORT() << "Failure to initialize RR with " << MemorySize
+             << " bytes on device " << DeviceId << ": "
+             << toString(std::move(Err));
+    return OFFLOAD_FAIL;
   }
   return OFFLOAD_SUCCESS;
 }
@@ -2185,9 +1842,9 @@ int32_t GenericPluginTy::get_global(__tgt_device_binary Binary, uint64_t Size,
   assert(DevicePtr && "Invalid device global's address");
 
   // Save the loaded globals if we are recording.
-  RecordReplayTy &RecordReplay = Device.Plugin.getRecordReplay();
-  if (RecordReplay.isRecording())
-    RecordReplay.addEntry(Name, Size, *DevicePtr);
+  RecordReplayTy *RecordReplay = Device.getRecordReplay();
+  if (RecordReplay && RecordReplay->isRecording())
+    RecordReplay->addGlobal(Name, Size, *DevicePtr);
 
   return OFFLOAD_SUCCESS;
 }
diff --git a/offload/plugins-nextgen/common/src/RecordReplay.cpp b/offload/plugins-nextgen/common/src/RecordReplay.cpp
new file mode 100644
index 0000000000000..5d35c73da236d
--- /dev/null
+++ b/offload/plugins-nextgen/common/src/RecordReplay.cpp
@@ -0,0 +1,245 @@
+#include "PluginInterface.h"
+
+#include "Shared/APITypes.h"
+
+#include "ErrorReporting.h"
+#include "Shared/Utils.h"
+
+#include "llvm/Support/Error.h"
+#include "llvm/Support/JSON.h"
+#include "llvm/Support/MemoryBuffer.h"
+#include "llvm/Support/raw_ostream.h"
+
+#include <cstdint>
+#include <filesystem>
+#include <functional>
+
+using namespace llvm;
+using namespace omp;
+using namespace target;
+using namespace plugin;
+using namespace error;
+
+Error RecordReplayTy::init(uint64_t MemSize, void *VAddr) {
+  if (!VAddr)
+    VAddr = Device.getSuggestedVirtualAddress();
+
+  auto StartAddrOrErr = Device.allocateWithVirtualAddress(MemSize, VAddr);
+  if (!StartAddrOrErr)
+    return StartAddrOrErr.takeError();
+  if (!*StartAddrOrErr)
+    return Plugin::error(ErrorCode::OUT_OF_RESOURCES, "allocating memory");
+
+  StartAddr = *StartAddrOrErr;
+  TotalSize = MemSize;
+
+  INFO(OMP_INFOTYPE_PLUGIN_KERNEL, Device.getDeviceId(),
+       "Record initialized with starting address %p, "
+       "memory size %lu bytes and status %s\n",
+       StartAddr, TotalSize,
+       Status == StatusTy::Recording ? "recording" : "replaying");
+
+  return Plugin::success();
+}
+
+Error RecordReplayTy::deinit() {
+  if (StartAddr)
+    return Device.deallocateWithVirtualAddress(StartAddr, TotalSize);
+  return Plugin::success();
+}
+
+std::pair<const RecordReplayTy::InstanceTy &, bool>
+RecordReplayTy::registerInstance(StringRef KernelName, uint32_t NumTeams,
+                                 uint32_t NumThreads,
+                                 uint32_t SharedMemorySize) {
+  std::lock_guard<std::mutex> LG(InstancesLock);
+  auto [It, Inserted] =
+      Instances.emplace(KernelName, NumTeams, NumThreads, SharedMemorySize);
+  // Increase the number of occurrences.
+  It->Occurrences += 1;
+  return {*It, Inserted};
+}
+
+void *RecordReplayTy::allocate(uint64_t Size) {
+  assert(StartAddr && "Expected memory has been pre-allocated");
+  constexpr int Alignment = 16;
+  // Assumes alignment is a power of 2.
+  int64_t AlignedSize = (Size + (Alignment - 1)) & (~(Alignment - 1));
+  std::lock_guard<std::mutex> LG(AllocationLock);
+  void *Alloc = (char *)StartAddr + CurrentSize;
+  CurrentSize += AlignedSize;
+  ODBG(OLDT_Alloc) << "Memory Allocator return " << Alloc;
+  return Alloc;
+}
+
+Expected<RecordReplayTy::HandleTy> RecordReplayTy::recordPrologue(
+    const GenericKernelTy &Kernel, const KernelArgsTy &KernelArgs,
+    const KernelLaunchParamsTy &LaunchParams, uint32_t NumTeams[3],
+    uint32_t NumThreads[3], uint32_t SharedMemorySize) {
+  if (!isRecordingOrReplaying())
+    return HandleTy{nullptr, false};
+
+  // Register the instance and avoid recording if it is inactive or replaying.
+  auto [Instance, First] = registerInstance(Kernel.getName(), NumTeams[0],
+                                            NumThreads[0], SharedMemorySize);
+
+  HandleTy Handle{&Instance, First};
+  if (isReplaying() || !First)
+    return Handle;
+
+  if (auto Err =
+          recordDescriptorImpl(Kernel, Instance, KernelArgs, LaunchParams,
+                               NumTeams, NumThreads, SharedMemorySize))
+    return Err;
+
+  if (auto Err = recordPrologueImpl(Kernel, Instance, KernelArgs, LaunchParams))
+    return Err;
+
+  return Handle;
+}
+
+Error RecordReplayTy::recordEpilogue(const GenericKernelTy &Kernel,
+                                     HandleTy Handle) {
+  if (!shouldRecordEpilogue() || !Handle.Active)
+    return Plugin::success();
+
+  return recordEpilogueImpl(Kernel, *Handle.Instance);
+}
+
+Error NativeRecordReplayTy::recordPrologueImpl(
+    const GenericKernelTy &Kernel, const InstanceTy &Instance,
+    const KernelArgsTy &KernelArgs, const KernelLaunchParamsTy &LaunchParams) {
+  SmallString<128> SnapshotFilename = {Kernel.getName(), ".memory"};
+  if (auto Err = recordSnapshot(SnapshotFilename))
+    return Err;
+
+  SmallString<128> GlobalsFilename = {Kernel.getName(), ".globals"};
+  if (auto Err = recordGlobals(GlobalsFilename))
+    return Err;
+
+  SmallString<128> ImageFilename = {Kernel.getName(), ".image"};
+  return recordImage(Kernel, ImageFilename);
+}
+
+Error NativeRecordReplayTy::recordEpilogueImpl(const GenericKernelTy &Kernel,
+                                               const InstanceTy &Instance) {
+  SmallString<128> SnapshotFilename = {
+      Kernel.getName(),
+      (isRecording() ? ".original.output" : ".replay.output")};
+  return recordSnapshot(SnapshotFilename);
+}
+
+Error NativeRecordReplayTy::recordDescriptorImpl(
+    const GenericKernelTy &Kernel, const InstanceTy &Instance,
+    const KernelArgsTy &KernelArgs, const KernelLaunchParamsTy &LaunchParams,
+    uint32_t NumTeams[3], uint32_t NumThreads[3], uint32_t SharedMemorySize) {
+  json::Object JsonKernelInfo;
+  JsonKernelInfo["Name"] = Kernel.getName();
+  JsonKernelInfo["NumArgs"] = KernelArgs.NumArgs;
+  JsonKernelInfo["NumTeamsClause"] = NumTeams[0];
+  JsonKernelInfo["ThreadLimitClause"] = NumThreads[0];
+  JsonKernelInfo["LoopTripCount"] = KernelArgs.Tripcount;
+  JsonKernelInfo["DeviceMemorySize"] = CurrentSize;
+  JsonKernelInfo["DeviceId"] = Device.getDeviceId();
+  JsonKernelInfo["VAllocAddr"] = (intptr_t)StartAddr;
+  JsonKernelInfo["VAllocSize"] = TotalSize;
+
+  json::Array JsonArgPtrs;
+  for (uint32_t I = 0; I < KernelArgs.NumArgs; ++I)
+    JsonArgPtrs.push_back((intptr_t)(*(void **)LaunchParams.Ptrs[I]));
+  JsonKernelInfo["ArgPtrs"] = json::Value(std::move(JsonArgPtrs));
+
+  json::Array JsonArgOffsets;
+  for (uint32_t I = 0; I < KernelArgs.NumArgs; ++I)
+    JsonArgOffsets.push_back(0);
+  JsonKernelInfo["ArgOffsets"] = json::Value(std::move(JsonArgOffsets));
+
+  SmallString<128> JsonFilename = {Kernel.getName(), ".json"};
+  std::error_code EC;
+  raw_fd_ostream JsonOS(JsonFilename.str(), EC);
+  if (EC)
+    return Plugin::error(ErrorCode::UNKNOWN, "saving kernel json file");
+  JsonOS << json::Value(std::move(JsonKernelInfo));
+  JsonOS.close();
+  return Plugin::success();
+}
+
+Error NativeRecordReplayTy::recordSnapshot(StringRef Filename) {
+  ErrorOr<std::unique_ptr<WritableMemoryBuffer>> DeviceMemoryMB =
+      WritableMemoryBuffer::getNewUninitMemBuffer(CurrentSize);
+  if (!DeviceMemoryMB)
+    return Plugin::error(ErrorCode::UNKNOWN,
+                         "creating MemoryBuffer for device memory");
+
+  if (auto Err = Device.dataRetrieve(DeviceMemoryMB.get()->getBufferStart(),
+                                     StartAddr, CurrentSize, nullptr))
+    return Err;
+
+  StringRef DeviceMemory(DeviceMemoryMB.get()->getBufferStart(), CurrentSize);
+  std::error_code EC;
+  raw_fd_ostream OS(Filename, EC);
+  if (EC)
+    return Plugin::error(ErrorCode::UNKNOWN, "dumping memory to file");
+  OS << DeviceMemory;
+  OS.close();
+  return Plugin::success();
+}
+
+Error NativeRecordReplayTy::recordImage(const GenericKernelTy &Kernel,
+                                        StringRef Filename) {
+  std::error_code EC;
+  raw_fd_ostream OS(Filename, EC);
+  if (EC)
+    return Plugin::error(ErrorCode::UNKNOWN, "saving image");
+  OS << Kernel.getImage().getMemoryBuffer().getBuffer();
+  OS.close();
+  return Plugin::success();
+}
+
+Error NativeRecordReplayTy::recordGlobals(StringRef Filename) {
+  int32_t Size = 0;
+
+  for (auto &OffloadEntry : GlobalEntries) {
+    if (!OffloadEntry.Size)
+      continue;
+    // Get the total size of the string and entry including the null byte.
+    Size +=
+        OffloadEntry.Name.length() + 1 + sizeof(uint32_t) + OffloadEntry.Size;
+  }
+
+  ErrorOr<std::unique_ptr<WritableMemoryBuffer>> GlobalsMB =
+      WritableMemoryBuffer::getNewUninitMemBuffer(Size);
+  if (!GlobalsMB)
+    return Plugin::error(ErrorCode::UNKNOWN,
+                         "creating MemoryBuffer for globals memory");
+
+  void *BufferPtr = GlobalsMB.get()->getBufferStart();
+  for (auto &OffloadEntry : GlobalEntries) {
+    if (!OffloadEntry.Size)
+      continue;
+
+    int32_t NameLength = OffloadEntry.Name.length() + 1;
+    memcpy(BufferPtr, OffloadEntry.Name.data(), NameLength);
+    BufferPtr = utils::advancePtr(BufferPtr, NameLength);
+
+    *((uint32_t *)(BufferPtr)) = OffloadEntry.Size;
+    BufferPtr = utils::advancePtr(BufferPtr, sizeof(uint32_t));
+
+    if (auto Err = Device.dataRetrieve(BufferPtr, OffloadEntry.Addr,
+                                       OffloadEntry.Size, nullptr))
+      return Err;
+    BufferPtr = utils::advancePtr(BufferPtr, OffloadEntry.Size);
+  }
+  assert(BufferPtr == GlobalsMB->get()->getBufferEnd() &&
+         "Buffer over/under-filled.");
+  assert(Size ==
+             utils::getPtrDiff(BufferPtr, GlobalsMB->get()->getBufferStart()) &&
+         "Buffer size mismatch");
+
+  StringRef GlobalsMemory(GlobalsMB.get()->getBufferStart(), Size);
+  std::error_code EC;
+  raw_fd_ostream OS(Filename, EC);
+  OS << GlobalsMemory;
+  OS.close();
+  return Plugin::success();
+}
diff --git a/offload/plugins-nextgen/cuda/src/rtl.cpp b/offload/plugins-nextgen/cuda/src/rtl.cpp
index 7a47f2ce7e5aa..87bcd68abf505 100644
--- a/offload/plugins-nextgen/cuda/src/rtl.cpp
+++ b/offload/plugins-nextgen/cuda/src/rtl.cpp
@@ -47,38 +47,6 @@ struct CUDAKernelTy;
 struct CUDADeviceTy;
 struct CUDAPluginTy;
 
-#if (defined(CUDA_VERSION) && (CUDA_VERSION < 11000))
-/// Forward declarations for all Virtual Memory Management
-/// related data structures and functions. This is necessary
-/// for older cuda versions.
-typedef void *CUmemGenericAllocationHandle;
-typedef void *CUmemAllocationProp;
-typedef void *CUmemAccessDesc;
-typedef void *CUmemAllocationGranularity_flags;
-CUresult cuMemAddressReserve(CUdeviceptr *ptr, size_t size, size_t alignment,
-                             CUdeviceptr addr, unsigned long long flags) {}
-CUresult cuMemMap(CUdeviceptr ptr, size_t size, size_t offset,
-                  CUmemGenericAllocationHandle handle,
-                  unsigned long long flags) {}
-CUresult cuMemCreate(CUmemGenericAllocationHandle *handle, size_t size,
-                     const CUmemAllocationProp *prop,
-                     unsigned long long flags) {}
-CUresult cuMemSetAccess(CUdeviceptr ptr, size_t size,
-                        const CUmemAccessDesc *desc, size_t count) {}
-CUresult
-cuMemGetAllocationGranularity(size_t *granularity,
-                              const CUmemAllocationProp *prop,
-                              CUmemAllocationGranularity_flags option) {}
-#endif
-
-#if (defined(CUDA_VERSION) && (CUDA_VERSION < 11020))
-// Forward declarations of asynchronous memory management functions. This is
-// necessary for older versions of CUDA.
-CUresult cuMemAllocAsync(CUdeviceptr *ptr, size_t, CUstream) { *ptr = 0; }
-
-CUresult cuMemFreeAsync(CUdeviceptr dptr, CUstream hStream) {}
-#endif
-
 /// Class implementing the CUDA device images properties.
 struct CUDADeviceImageTy : public DeviceImageTy {
   /// Create the CUDA image with the id and the target image pointer.
@@ -679,125 +647,99 @@ struct CUDADeviceTy : public GenericDeviceTy {
     return Plugin::check(Res, "error in cuStreamSynchronize: %s");
   }
 
-  /// CUDA support VA management
-  bool supportVAManagement() const override {
-#if (defined(CUDA_VERSION) && (CUDA_VERSION >= 11000))
-    return true;
-#else
-    return false;
-#endif
+  /// Suggest a virtual address for device memory mapping.
+  void *getSuggestedVirtualAddress() override {
+    return reinterpret_cast<void *>(0x153940000000ULL);
   }
 
-  /// Allocates \p RSize bytes (rounded up to page size) and hints the cuda
-  /// driver to map it to \p VAddr. The obtained address is stored in \p Addr.
-  /// At return \p RSize contains the actual size
-  Error memoryVAMap(void **Addr, void *VAddr, size_t *RSize) override {
-    CUdeviceptr DVAddr = reinterpret_cast<CUdeviceptr>(VAddr);
-    auto IHandle = DeviceMMaps.find(DVAddr);
-    size_t Size = *RSize;
-
-    if (Size == 0)
-      return Plugin::error(ErrorCode::INVALID_ARGUMENT,
-                           "memory Map Size must be larger than 0");
+  /// Allocate \p Size bytes on the device and hints the backend to map it to
+  /// virtual address \p VAddr. The function returns the allocated virtual
+  /// address. The memory must be deallocated through
+  /// GenericDeviceTy::deallocateWithVirtualAddress().
+  Expected<void *> allocateWithVirtualAddress(uint64_t Size,
+                                              void *VAddr) override {
+    CUdeviceptr ExpectedVAddr = 0;
+    if (VAddr != nullptr)
+      ExpectedVAddr = reinterpret_cast<CUdeviceptr>(VAddr);
 
-    // Check if we have already mapped this address
-    if (IHandle != DeviceMMaps.end())
-      return Plugin::error(ErrorCode::INVALID_ARGUMENT,
-                           "address already memory mapped");
-
-    CUmemAllocationProp Prop = {};
+    // Get the page size in the device.
     size_t Granularity = 0;
-
-    size_t Free, Total;
-    CUresult Res = cuMemGetInfo(&Free, &Total);
-    if (auto Err = Plugin::check(Res, "Error in cuMemGetInfo: %s"))
-      return Err;
-
-    if (Size >= Free) {
-      *Addr = nullptr;
-      return Plugin::error(
-          ErrorCode::OUT_OF_RESOURCES,
-          "cannot map memory size larger than the available device memory");
-    }
-
-    // currently NVidia only supports pinned device types
+    CUmemAllocationProp Prop = {};
     Prop.type = CU_MEM_ALLOCATION_TYPE_PINNED;
     Prop.location.type = CU_MEM_LOCATION_TYPE_DEVICE;
-
     Prop.location.id = DeviceId;
-    cuMemGetAllocationGranularity(&Granularity, &Prop,
-                                  CU_MEM_ALLOC_GRANULARITY_MINIMUM);
+    CUresult Res = cuMemGetAllocationGranularity(
+        &Granularity, &Prop, CU_MEM_ALLOC_GRANULARITY_MINIMUM);
     if (auto Err =
             Plugin::check(Res, "error in cuMemGetAllocationGranularity: %s"))
       return Err;
-
     if (Granularity == 0)
       return Plugin::error(ErrorCode::INVALID_ARGUMENT,
-                           "wrong device Page size");
+                           "wrong device page size");
 
-    // Ceil to page size.
+    // Transparently round up to a multiple of the page size.
     Size = utils::roundUp(Size, Granularity);
 
-    // Create a handler of our allocation
-    CUmemGenericAllocationHandle AHandle;
-    Res = cuMemCreate(&AHandle, Size, &Prop, 0);
-    if (auto Err = Plugin::check(Res, "error in cuMemCreate: %s"))
-      return Err;
-
+    // Reserve the virtual address range.
     CUdeviceptr DevPtr = 0;
-    Res = cuMemAddressReserve(&DevPtr, Size, 0, DVAddr, 0);
+    Res = cuMemAddressReserve(&DevPtr, Size, 0, ExpectedVAddr, 0);
     if (auto Err = Plugin::check(Res, "error in cuMemAddressReserve: %s"))
       return Err;
 
-    Res = cuMemMap(DevPtr, Size, 0, AHandle, 0);
-    if (auto Err = Plugin::check(Res, "error in cuMemMap: %s"))
+    // Create a handle of the allocation.
+    CUmemGenericAllocationHandle Handle;
+    Res = cuMemCreate(&Handle, Size, &Prop, 0);
+    if (auto Err = Plugin::check(Res, "error in cuMemCreate: %s"))
       return Err;
 
-    CUmemAccessDesc ADesc = {};
-    ADesc.location.type = CU_MEM_LOCATION_TYPE_DEVICE;
-    ADesc.location.id = DeviceId;
-    ADesc.flags = CU_MEM_ACCESS_FLAGS_PROT_READWRITE;
+    // Map the virtual address range to the memory allocation.
+    Res = cuMemMap(DevPtr, Size, 0, Handle, 0);
+    if (auto Err = Plugin::check(Res, "error in cuMemMap: %s"))
+      return Err;
 
-    // Sets address
-    Res = cuMemSetAccess(DevPtr, Size, &ADesc, 1);
+    // Set the memory access properties for the allocation.
+    CUmemAccessDesc Desc = {};
+    Desc.location.type = CU_MEM_LOCATION_TYPE_DEVICE;
+    Desc.location.id = DeviceId;
+    Desc.flags = CU_MEM_ACCESS_FLAGS_PROT_READWRITE;
+    Res = cuMemSetAccess(DevPtr, Size, &Desc, 1);
     if (auto Err = Plugin::check(Res, "error in cuMemSetAccess: %s"))
       return Err;
 
-    *Addr = reinterpret_cast<void *>(DevPtr);
-    *RSize = Size;
-    DeviceMMaps.insert({DevPtr, AHandle});
-    return Plugin::success();
-  }
+    VAddr = reinterpret_cast<void *>(DevPtr);
 
-  /// De-allocates device memory and Unmaps the Virtual Addr
-  Error memoryVAUnMap(void *VAddr, size_t Size) override {
-    CUdeviceptr DVAddr = reinterpret_cast<CUdeviceptr>(VAddr);
-    auto IHandle = DeviceMMaps.find(DVAddr);
-    // Mapping does not exist
-    if (IHandle == DeviceMMaps.end()) {
-      return Plugin::error(ErrorCode::INVALID_ARGUMENT,
-                           "addr is not MemoryMapped");
-    }
+    // Register the virtual address range in the tracker.
+    if (auto Err = VMemTracker.registerReservation(VAddr, Size, Handle))
+      return Err;
 
-    if (IHandle == DeviceMMaps.end())
-      return Plugin::error(ErrorCode::INVALID_ARGUMENT,
-                           "addr is not MemoryMapped");
+    return VAddr;
+  }
+
+  /// Deallocate device memory \p VAddr, which was allocated through
+  /// GenericDeviceTy::allocateWithVirtualAddress(), and unmap the virtual
+  /// address range.
+  Error deallocateWithVirtualAddress(void *VAddr, uint64_t) override {
+    // Unregister the virtual address range and obtain the information about
+    // the reservation.
+    auto InfoOrErr = VMemTracker.unregisterReservation(VAddr);
+    if (!InfoOrErr)
+      return InfoOrErr.takeError();
 
-    CUmemGenericAllocationHandle &AllocHandle = IHandle->second;
+    auto [Size, Handle] = *InfoOrErr;
+    CUdeviceptr DevAddr = reinterpret_cast<CUdeviceptr>(VAddr);
 
-    CUresult Res = cuMemUnmap(DVAddr, Size);
+    CUresult Res = cuMemUnmap(DevAddr, Size);
     if (auto Err = Plugin::check(Res, "error in cuMemUnmap: %s"))
       return Err;
 
-    Res = cuMemRelease(AllocHandle);
+    Res = cuMemRelease(Handle);
     if (auto Err = Plugin::check(Res, "error in cuMemRelease: %s"))
       return Err;
 
-    Res = cuMemAddressFree(DVAddr, Size);
+    Res = cuMemAddressFree(DevAddr, Size);
     if (auto Err = Plugin::check(Res, "error in cuMemAddressFree: %s"))
       return Err;
 
-    DeviceMMaps.erase(IHandle);
     return Plugin::success();
   }
 
@@ -1459,9 +1401,6 @@ struct CUDADeviceTy : public GenericDeviceTy {
   /// The CUDA device handler.
   CUdevice Device = CU_DEVICE_INVALID;
 
-  /// The memory mapped addresses and their handles
-  std::unordered_map<CUdeviceptr, CUmemGenericAllocationHandle> DeviceMMaps;
-
   /// The compute capability of the corresponding CUDA device.
   struct ComputeCapabilityTy {
     uint32_t Major;
@@ -1474,6 +1413,9 @@ struct CUDADeviceTy : public GenericDeviceTy {
   /// The maximum number of warps that can be resident on all the SMs
   /// simultaneously.
   uint32_t HardwareParallelism = 0;
+
+  /// Tracker for virtual address reservations.
+  VMemTrackerTy<CUmemGenericAllocationHandle> VMemTracker;
 };
 
 Error CUDAKernelTy::launchImpl(GenericDeviceTy &GenericDevice,
diff --git a/offload/tools/kernelreplay/llvm-omp-kernel-replay.cpp b/offload/tools/kernelreplay/llvm-omp-kernel-replay.cpp
index 718e9e78304e8..5dc0bda86a99c 100644
--- a/offload/tools/kernelreplay/llvm-omp-kernel-replay.cpp
+++ b/offload/tools/kernelreplay/llvm-omp-kernel-replay.cpp
@@ -91,11 +91,14 @@ int main(int argc, char **argv) {
   for (auto It : *TgtArgOffsetsArray)
     TgtArgOffsets.push_back(static_cast<ptrdiff_t>(It.getAsInteger().value()));
 
-  void *BAllocStart = reinterpret_cast<void *>(
-      JsonKernelInfo->getAsObject()->getInteger("BumpAllocVAStart").value());
-
-  llvm::offloading::EntryTy KernelEntry = {~0U,     0, 0, 0,      nullptr,
-                                           nullptr, 0, 0, nullptr};
+  void *VAllocAddr = reinterpret_cast<void *>(
+      JsonKernelInfo->getAsObject()->getInteger("VAllocAddr").value());
+  uint64_t VAllocSize =
+      JsonKernelInfo->getAsObject()->getInteger("VAllocSize").value();
+
+  llvm::offloading::EntryTy KernelEntry = {
+      0x0, 0x1,    object::OffloadKind::OFK_OpenMP, 0, nullptr, nullptr, 0,
+      0,   nullptr};
   std::string KernelEntryName = KernelFunc.value().str();
   KernelEntry.SymbolName = const_cast<char *>(KernelEntryName.c_str());
   // Anything non-zero works to uniquely identify the kernel.
@@ -119,11 +122,6 @@ int main(int argc, char **argv) {
   Desc.HostEntriesEnd = &KernelEntry + 1;
   Desc.DeviceImages = &DeviceImage;
 
-  auto DeviceMemorySizeJson =
-      JsonKernelInfo->getAsObject()->getInteger("DeviceMemorySize");
-  // Set device memory size to the ceiling of GB granularity.
-  uint64_t DeviceMemorySize = std::ceil(DeviceMemorySizeJson.value());
-
   auto DeviceIdJson = JsonKernelInfo->getAsObject()->getInteger("DeviceId");
   // TODO: Print warning if the user overrides the device id in the json file.
   int32_t DeviceId = (DeviceIdOpt > -1 ? DeviceIdOpt : DeviceIdJson.value());
@@ -134,8 +132,8 @@ int main(int argc, char **argv) {
   __tgt_register_lib(&Desc);
 
   uint64_t ReqPtrArgOffset = 0;
-  int Rc = __tgt_activate_record_replay(DeviceId, DeviceMemorySize, BAllocStart,
-                                        false, VerifyOpt, ReqPtrArgOffset);
+  int Rc = __tgt_activate_record_replay(DeviceId, VAllocSize, VAllocAddr, false,
+                                        VerifyOpt, ReqPtrArgOffset);
 
   if (Rc != OMP_TGT_SUCCESS) {
     report_fatal_error("Cannot activate record replay\n");
@@ -150,8 +148,8 @@ int main(int argc, char **argv) {
 
   // On AMD for currently unknown reasons we cannot copy memory mapped data to
   // device. This is a work-around.
-  uint8_t *recored_data = new uint8_t[DeviceMemoryMB.get()->getBufferSize()];
-  std::memcpy(recored_data,
+  uint8_t *RecordedData = new uint8_t[DeviceMemoryMB.get()->getBufferSize()];
+  std::memcpy(RecordedData,
               const_cast<char *>(DeviceMemoryMB.get()->getBuffer().data()),
               DeviceMemoryMB.get()->getBufferSize());
 
@@ -160,15 +158,15 @@ int main(int argc, char **argv) {
     for (auto *&Arg : TgtArgs) {
       auto ArgInt = uintptr_t(Arg);
       // Try to find pointer arguments.
-      if (ArgInt < uintptr_t(BAllocStart) ||
-          ArgInt >= uintptr_t(BAllocStart) + DeviceMemorySize)
+      if (ArgInt < uintptr_t(VAllocAddr) ||
+          ArgInt >= uintptr_t(VAllocAddr) + VAllocSize)
         continue;
       Arg = reinterpret_cast<void *>(ArgInt - ReqPtrArgOffset);
     }
   }
 
   __tgt_target_kernel_replay(
-      /*Loc=*/nullptr, DeviceId, KernelEntry.Address, (char *)recored_data,
+      /*Loc=*/nullptr, DeviceId, KernelEntry.Address, (char *)RecordedData,
       DeviceMemoryMB.get()->getBufferSize(), TgtArgs.data(),
       TgtArgOffsets.data(), NumArgs.value(), NumTeams, NumThreads,
       LoopTripCount.value());
@@ -198,7 +196,7 @@ int main(int argc, char **argv) {
                 "verify!\n";
   }
 
-  delete[] recored_data;
+  delete[] RecordedData;
 
   return 0;
 }



More information about the llvm-commits mailing list