[llvm] [Offload] Add `cuMemPrefetchBatchAsync` into CUDA plugin. (PR #222650)

via llvm-commits llvm-commits at lists.llvm.org
Thu Sep 10 06:47:50 PDT 2026


https://github.com/blazej-smorawski updated https://github.com/llvm/llvm-project/pull/222650

>From 31e9af95528e53ccb6700267a4c97b94d269f162 Mon Sep 17 00:00:00 2001
From: blazej-smorawski <blazej.smorawski at intel.com>
Date: Wed, 9 Sep 2026 12:15:47 +0200
Subject: [PATCH 1/2] [Offload] Add batched prefetch to CUDA plugin

---
 .../cuda/dynamic_cuda/cuda.cpp                | 16 +++++-
 .../plugins-nextgen/cuda/dynamic_cuda/cuda.h  | 22 ++++++-
 .../plugins-nextgen/cuda/src/cuda_compat.h    | 25 ++++++++
 offload/plugins-nextgen/cuda/src/rtl.cpp      | 57 ++++++++++++++-----
 4 files changed, 105 insertions(+), 15 deletions(-)
 create mode 100644 offload/plugins-nextgen/cuda/src/cuda_compat.h

diff --git a/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.cpp b/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.cpp
index e3a854d5690c1..04190f55fdc2a 100644
--- a/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.cpp
+++ b/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.cpp
@@ -21,6 +21,7 @@
 #include <memory>
 #include <string>
 #include <unordered_map>
+#include <unordered_set>
 
 using namespace llvm::offload::debug;
 
@@ -69,6 +70,7 @@ DLWRAP(cuMemFreeHost, 1)
 DLWRAP(cuMemFreeAsync, 2)
 
 DLWRAP(cuMemPrefetchAsync, 4)
+DLWRAP(cuMemPrefetchBatchAsync, 8)
 DLWRAP(cuPointerGetAttribute, 3)
 
 DLWRAP(cuModuleGetFunction, 3)
@@ -146,6 +148,9 @@ static bool checkForCUDA() {
       {"cuDevicePrimaryCtxSetFlags", "cuDevicePrimaryCtxSetFlags_v2"},
   };
 
+  // Set of APIs that might not be supported in older versions of CUDA
+  std::unordered_set<std::string> OptionalAPIs = {"cuMemPrefetchBatchAsync"};
+
   const char *CudaLib = DYNAMIC_CUDA_PATH;
   std::string ErrMsg;
   auto DynlibHandle = std::make_unique<llvm::sys::DynamicLibrary>(
@@ -175,8 +180,17 @@ static bool checkForCUDA() {
     if (P == nullptr) {
       ODBG(OLDT_Init) << "Unable to find '" << Sym << "' in '" << CudaLib
                       << "'!";
-      return false;
+
+      // Check if the missing symbols is in optional list
+      if (OptionalAPIs.find(Sym) == OptionalAPIs.end())
+        return false;
+
+      // Leave the API as nullptr, should be guarded with
+      // api_helper::canCall<>()
+      *dlwrap::pointer(I) = nullptr;
+      continue;
     }
+
     ODBG(OLDT_Init) << "Implementing " << Sym << " with dlsym(" << Sym
                     << ") -> " << P;
 
diff --git a/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.h b/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.h
index a7524d417dded..8e9b6b5469d6a 100644
--- a/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.h
+++ b/offload/plugins-nextgen/cuda/dynamic_cuda/cuda.h
@@ -16,6 +16,10 @@
 #include <cstddef>
 #include <cstdint>
 
+#if defined(__cplusplus)
+extern "C" {
+#endif
+
 #define cuDeviceTotalMem cuDeviceTotalMem_v2
 #define cuModuleGetGlobal cuModuleGetGlobal_v2
 #define cuMemGetInfo cuMemGetInfo_v2
@@ -59,7 +63,14 @@ typedef enum CUmemAccess_flags_enum {
 
 typedef enum CUmemLocationType_enum {
   CU_MEM_LOCATION_TYPE_INVALID = 0x0,
-  CU_MEM_LOCATION_TYPE_DEVICE = 0x1,
+  CU_MEM_LOCATION_TYPE_DEVICE =
+      0x1, /**< Location is a device location, thus id is a device ordinal */
+  CU_MEM_LOCATION_TYPE_HOST = 0x2, /**< Location is host, id is ignored */
+  CU_MEM_LOCATION_TYPE_HOST_NUMA =
+      0x3, /**< Location is a host NUMA node, thus id is a host NUMA node id */
+  CU_MEM_LOCATION_TYPE_HOST_NUMA_CURRENT =
+      0x4, /**< Location is a host NUMA node of the current thread, id is
+              ignored */
   CU_MEM_LOCATION_TYPE_MAX = 0x7FFFFFFF
 } CUmemLocationType;
 
@@ -480,6 +491,11 @@ CUresult cuMemFreeHost(void *);
 CUresult cuMemFreeAsync(CUdeviceptr, CUstream);
 
 CUresult cuMemPrefetchAsync(CUdeviceptr, size_t, CUdevice, CUstream);
+CUresult cuMemPrefetchBatchAsync(CUdeviceptr *dptrs, size_t *sizes,
+                                 size_t count, CUmemLocation *prefetchLocs,
+                                 size_t *prefetchLocIdxs,
+                                 size_t numPrefetchLocs,
+                                 unsigned long long flags, CUstream hStream);
 
 typedef enum CUpointer_attribute_enum {
   CU_POINTER_ATTRIBUTE_IS_MANAGED = 8
@@ -541,4 +557,8 @@ CUresult cuOccupancyMaxPotentialBlockSize(int *, int *, CUfunction,
 CUresult cuOccupancyMaxActiveBlocksPerMultiprocessor(int *, CUfunction, int,
                                                      size_t);
 
+#if defined(__cplusplus)
+} // extern "C"
 #endif
+
+#endif // DYNAMIC_CUDA_CUDA_H_INCLUDED
diff --git a/offload/plugins-nextgen/cuda/src/cuda_compat.h b/offload/plugins-nextgen/cuda/src/cuda_compat.h
new file mode 100644
index 0000000000000..fa2238d603f2f
--- /dev/null
+++ b/offload/plugins-nextgen/cuda/src/cuda_compat.h
@@ -0,0 +1,25 @@
+//===--- cuda/src/cuda_compat.h -------------------------------------------===//
+//
+// 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
+//
+//===----------------------------------------------------------------------===//
+//
+// CUDA compatibility layer enabling us to compile using new APIs.
+//
+//===----------------------------------------------------------------------===//
+
+#ifndef OPENMP_LIBOMPTARGET_PLUGINS_NEXTGEN_CUDA_CUDACOMPAT_H
+#define OPENMP_LIBOMPTARGET_PLUGINS_NEXTGEN_CUDA_CUDACOMPAT_H
+
+#include "APIHelpers.h"
+
+#include <cuda.h>
+
+API_HELPER_OPTIONAL(CUresult, cuMemPrefetchBatchAsync, CUdeviceptr *dptrs,
+                    size_t *sizes, size_t count, CUmemLocation *prefetchLocs,
+                    size_t *prefetchLocIdxs, size_t numPrefetchLocs,
+                    unsigned long long flags, CUstream hStream)
+
+#endif
diff --git a/offload/plugins-nextgen/cuda/src/rtl.cpp b/offload/plugins-nextgen/cuda/src/rtl.cpp
index 72e5dcf115fe9..506651a7e03ee 100644
--- a/offload/plugins-nextgen/cuda/src/rtl.cpp
+++ b/offload/plugins-nextgen/cuda/src/rtl.cpp
@@ -10,12 +10,16 @@
 //
 //===----------------------------------------------------------------------===//
 
+#include <algorithm>
 #include <cassert>
 #include <cstddef>
 #include <cuda.h>
 #include <string>
 #include <unordered_map>
 
+#include "APIHelpers.h"
+#include "cuda_compat.h"
+
 #include "Shared/APITypes.h"
 #include "Shared/Debug.h"
 #include "Shared/Environment.h"
@@ -26,6 +30,7 @@
 #include "PluginInterface.h"
 #include "Utils/ELF.h"
 
+#include "llvm/ADT/SmallVector.h"
 #include "llvm/ADT/StringExtras.h"
 #include "llvm/BinaryFormat/ELF.h"
 #include "llvm/Frontend/OpenMP/OMPConstants.h"
@@ -920,9 +925,6 @@ struct CUDADeviceTy : public GenericDeviceTy {
   }
 
   /// Prefetch managed memory to the device or back to the host.
-  // TODO: switch to cuMemPrefetchBatchAsync once the minimum supported CUDA
-  // driver is 13 or newer. That entry point takes the (Mems, Sizes, Count)
-  // arrays directly and lets the driver batch the migration.
   Error dataPrefetchImpl(size_t Count, const void **Mems, const size_t *Sizes,
                          bool ToHost,
                          AsyncInfoWrapperTy &AsyncInfoWrapper) override {
@@ -941,25 +943,54 @@ struct CUDADeviceTy : public GenericDeviceTy {
         !ConcurrentManagedAccess)
       return Plugin::success();
 
-    CUstream Stream;
-    if (auto Err = getStream(AsyncInfoWrapper, Stream))
-      return Err;
-
-    CUdevice Dst = ToHost ? CU_DEVICE_CPU : Device;
-    for (size_t I = 0; I < Count; I++) {
-      if (Sizes[I] == 0)
+    llvm::SmallVector<size_t, 8> FilteredSizes;
+    llvm::SmallVector<CUdeviceptr, 8> FilteredPtrs;
+    for (size_t MemoryPtrIndex = 0; MemoryPtrIndex < Count; MemoryPtrIndex++) {
+      if (Sizes[MemoryPtrIndex] == 0)
         continue;
 
       // Prefetch only works with USM (managed) memory; ignore the hint
       // otherwise.
       unsigned int IsManaged = 0;
       if (cuPointerGetAttribute(&IsManaged, CU_POINTER_ATTRIBUTE_IS_MANAGED,
-                                (CUdeviceptr)Mems[I]) != CUDA_SUCCESS ||
+                                (CUdeviceptr)Mems[MemoryPtrIndex]) !=
+              CUDA_SUCCESS ||
           !IsManaged)
         continue;
 
-      CUresult Res =
-          cuMemPrefetchAsync((CUdeviceptr)Mems[I], Sizes[I], Dst, Stream);
+      FilteredSizes.push_back(Sizes[MemoryPtrIndex]);
+      FilteredPtrs.push_back(reinterpret_cast<CUdeviceptr>(Mems[MemoryPtrIndex]));
+    }
+
+    if (FilteredPtrs.size() == 0)
+      return Plugin::success();
+
+    CUstream Stream;
+    if (auto Err = getStream(AsyncInfoWrapper, Stream))
+      return Err;
+
+    if (api_helper::canCall<cuMemPrefetchBatchAsync>()) {
+      CUmemLocation Loc{};
+      if (ToHost)
+        Loc = {.type = CU_MEM_LOCATION_TYPE_HOST, .id = 0};
+      else
+        Loc = {.type = CU_MEM_LOCATION_TYPE_DEVICE, .id = Device};
+
+      size_t LocIdxs = 0;
+      CUresult Res = cuMemPrefetchBatchAsync(
+          FilteredPtrs.data(), FilteredSizes.data(),
+          FilteredPtrs.size(), &Loc, &LocIdxs, 1, 0, Stream);
+      if (auto Err = Plugin::check(Res, "error in cuMemPrefetchBatchAsync: %s"))
+        return Err;
+
+      return Plugin::success();
+    }
+
+    // Fallback path for CUDA < 13
+    CUdevice Dst = ToHost ? CU_DEVICE_CPU : Device;
+    for (size_t I = 0; I < FilteredPtrs.size(); I++) {
+      CUresult Res = cuMemPrefetchAsync((CUdeviceptr)FilteredPtrs[I],
+                                        FilteredSizes[I], Dst, Stream);
       if (auto Err = Plugin::check(Res, "error in cuMemPrefetchAsync: %s"))
         return Err;
     }

>From 6df205d327579c8605d88693aea60feef76e779a Mon Sep 17 00:00:00 2001
From: blazej-smorawski <blazej.smorawski at intel.com>
Date: Thu, 10 Sep 2026 15:46:46 +0200
Subject: [PATCH 2/2] Fix format

---
 offload/plugins-nextgen/cuda/src/rtl.cpp | 7 ++++---
 1 file changed, 4 insertions(+), 3 deletions(-)

diff --git a/offload/plugins-nextgen/cuda/src/rtl.cpp b/offload/plugins-nextgen/cuda/src/rtl.cpp
index 506651a7e03ee..d1395685caf0b 100644
--- a/offload/plugins-nextgen/cuda/src/rtl.cpp
+++ b/offload/plugins-nextgen/cuda/src/rtl.cpp
@@ -959,7 +959,8 @@ struct CUDADeviceTy : public GenericDeviceTy {
         continue;
 
       FilteredSizes.push_back(Sizes[MemoryPtrIndex]);
-      FilteredPtrs.push_back(reinterpret_cast<CUdeviceptr>(Mems[MemoryPtrIndex]));
+      FilteredPtrs.push_back(
+          reinterpret_cast<CUdeviceptr>(Mems[MemoryPtrIndex]));
     }
 
     if (FilteredPtrs.size() == 0)
@@ -978,8 +979,8 @@ struct CUDADeviceTy : public GenericDeviceTy {
 
       size_t LocIdxs = 0;
       CUresult Res = cuMemPrefetchBatchAsync(
-          FilteredPtrs.data(), FilteredSizes.data(),
-          FilteredPtrs.size(), &Loc, &LocIdxs, 1, 0, Stream);
+          FilteredPtrs.data(), FilteredSizes.data(), FilteredPtrs.size(), &Loc,
+          &LocIdxs, 1, 0, Stream);
       if (auto Err = Plugin::check(Res, "error in cuMemPrefetchBatchAsync: %s"))
         return Err;
 



More information about the llvm-commits mailing list