[llvm-branch-commits] [llvm] [offload][omp] Move OpenMP KLE to libomptarget (PR #223761)

Alex Duran via llvm-branch-commits llvm-branch-commits at lists.llvm.org
Mon Sep 21 07:13:33 PDT 2026


https://github.com/adurang updated https://github.com/llvm/llvm-project/pull/223761

>From 55af1dc71c499d2c41f9c6e86002d4bd1f25dd9e Mon Sep 17 00:00:00 2001
From: "Duran, Alex" <alejandro.duran at intel.com>
Date: Tue, 15 Sep 2026 09:40:11 -0700
Subject: [PATCH] [offload][omp] Move OpenMP KLE to libomptarget

Move preparations related to OpenMP KLE and
dynamicCGroupMem fallback out of the plugins and
into libomptarget.

Resructure Device::launch as it grew too large.
---
 offload/include/device.h                      |   1 +
 offload/libompaccsupport/PluginManager.cpp    |   1 +
 offload/libompaccsupport/device.cpp           | 352 +++++++++++++++---
 .../common/include/PluginInterface.h          |  55 +--
 .../common/src/PluginInterface.cpp            | 146 +-------
 .../common/src/RecordReplay.cpp               |   2 +-
 6 files changed, 301 insertions(+), 256 deletions(-)

diff --git a/offload/include/device.h b/offload/include/device.h
index 9c5fa061ec608..c4bcec5a71805 100644
--- a/offload/include/device.h
+++ b/offload/include/device.h
@@ -49,6 +49,7 @@ struct KernelLaunchInfoTy {
   uint32_t MaxNumThreads = 0;
   uint32_t PreferredNumThreads = 0;
   uint32_t ReductionDataSize = 0;
+  uint32_t StaticBlockMemSize = 0;
   llvm::omp::OMPTgtExecModeFlags Mode = llvm::omp::OMP_TGT_EXEC_MODE_BARE;
 
   bool isBareMode() const { return Mode == llvm::omp::OMP_TGT_EXEC_MODE_BARE; }
diff --git a/offload/libompaccsupport/PluginManager.cpp b/offload/libompaccsupport/PluginManager.cpp
index db68b1c6fb9c4..4bf468d8ec5bc 100644
--- a/offload/libompaccsupport/PluginManager.cpp
+++ b/offload/libompaccsupport/PluginManager.cpp
@@ -520,6 +520,7 @@ static int loadImagesOntoDevice(DeviceTy &Device) {
                   ? std::max(Cfg.MinThreads,
                              int32_t(GenericDevice.getDefaultNumThreads()))
                   : GenericDevice.getDefaultNumThreads();
+          LaunchInfo.StaticBlockMemSize = Kernel->getStaticBlockMemSize();
 
           Device.setKernelLaunchInfo(DeviceEntry.Address, LaunchInfo);
         }
diff --git a/offload/libompaccsupport/device.cpp b/offload/libompaccsupport/device.cpp
index 14ab45709d0d5..a198a4a900835 100644
--- a/offload/libompaccsupport/device.cpp
+++ b/offload/libompaccsupport/device.cpp
@@ -389,6 +389,152 @@ static void resolveKernelLaunchParams(void **const TgtArgs,
   LaunchArgs.Args = &Ptrs[0];
 }
 
+namespace {
+/// Configuration of dynamic block memory needed for launching a kernel.
+struct DynBlockMemConfTy {
+  /// The size of the dynamic block memory buffer.
+  uint32_t Size = 0;
+  /// The size of dynamic shared memory natively provided by the device.
+  uint32_t NativeSize = 0;
+  /// The fallback that was triggered (if any).
+  DynCGroupMemFallbackType Fallback = DynCGroupMemFallbackType::None;
+  /// The fallback pointer if global memory was used as alternative.
+  void *FallbackPtr = nullptr;
+};
+} // namespace
+
+/// Prepare the block memory buffer requested for the kernel and execute the
+/// specified fallback if necessary.
+static llvm::Expected<DynBlockMemConfTy>
+prepareBlockMemory(GenericDeviceTy &GenericDevice,
+                   const KernelLaunchInfoTy &KernelEnv, uint32_t DynCGroupMem,
+                   DynCGroupMemFallbackType DynCGroupMemFallback,
+                   uint32_t NumBlocks) {
+  uint32_t MaxBlockMemSize = GenericDevice.getMaxBlockSharedMemSize();
+  uint32_t DynBlockMemSize = DynCGroupMem;
+  uint32_t TotalBlockMemSize = KernelEnv.StaticBlockMemSize + DynBlockMemSize;
+  uint32_t DynNativeBlockMemSize = DynBlockMemSize;
+  void *DynFallbackPtr = nullptr;
+
+  // No enough block memory to cover the static one. Cannot run the kernel.
+  if (KernelEnv.StaticBlockMemSize > MaxBlockMemSize)
+    return error::createOffloadError(
+        error::ErrorCode::INVALID_ARGUMENT,
+        "Static block memory size exceeds maximum");
+  // No enough block memory to cover dynamic one, and the fallback is aborting.
+  if (DynCGroupMemFallback == DynCGroupMemFallbackType::Abort &&
+      TotalBlockMemSize > MaxBlockMemSize)
+    return error::createOffloadError(
+        error::ErrorCode::INVALID_ARGUMENT,
+        "Requested block memory size (static + dynamic) exceeds maximum");
+
+  DynCGroupMemFallbackType DynFallback = DynCGroupMemFallbackType::None;
+  if (DynBlockMemSize && TotalBlockMemSize > MaxBlockMemSize) {
+    // Launch without native dynamic block memory.
+    DynNativeBlockMemSize = 0;
+    DynFallback = DynCGroupMemFallback;
+    if (DynFallback != DynCGroupMemFallbackType::DefaultMem) {
+      // Do not provide any memory as fallback.
+      DynBlockMemSize = 0;
+    } else {
+      // Get global memory as fallback.
+      auto AllocOrErr = GenericDevice.dataAlloc(
+          NumBlocks * DynBlockMemSize,
+          /*HostPtr=*/nullptr, TARGET_ALLOC_DEVICE, /*Alignment=*/0);
+      if (!AllocOrErr)
+        return AllocOrErr.takeError();
+      DynFallbackPtr = *AllocOrErr;
+    }
+  }
+  return DynBlockMemConfTy{DynBlockMemSize, DynNativeBlockMemSize, DynFallback,
+                           DynFallbackPtr};
+}
+
+static void freeAfterSynchronization(GenericDeviceTy &GenericDevice,
+                                     AsyncInfoTy &AsyncInfo, void *Ptr,
+                                     TargetAllocTy Kind) {
+  AsyncInfo.addPostProcessingFunction([&GenericDevice, Ptr, Kind]() -> int {
+    if (auto Err = GenericDevice.dataDelete(Ptr, Kind)) {
+      REPORT() << "Failure to free device memory " << Ptr << ": "
+               << toString(std::move(Err));
+      return OFFLOAD_FAIL;
+    }
+    return OFFLOAD_SUCCESS;
+  });
+}
+
+/// Return a device pointer to a new kernel launch environment, or null if
+/// this launch has no reserved dyn_ptr slot to store one in. \p NumBlocks0 is
+/// the number of blocks for this launch and is used to size the reduction
+/// buffer.
+static llvm::Expected<KernelLaunchEnvironmentTy *> getKernelLaunchEnvironment(
+    GenericDeviceTy &GenericDevice, const KernelLaunchArgsTy &LaunchArgs,
+    const KernelLaunchInfoTy &KernelEnv,
+    const DynBlockMemConfTy &DynBlockMemConf, uint32_t DynCGroupMem,
+    void **DynPtrSlot, AsyncInfoTy &AsyncInfo, uint32_t NumBlocks0) {
+  // Ctor/Dtor have no arguments, replaying uses the original kernel launch
+  // environment, and launches with no reserved dyn_ptr slot (e.g. older
+  // compiler versions, or non-OpenMP launches) have nowhere to store one.
+  if ((GenericDevice.getRecordReplay() &&
+       GenericDevice.getRecordReplay()->isReplaying()) ||
+      !DynPtrSlot)
+    return nullptr;
+
+  const bool NeedsReductionBuffer = KernelEnv.ReductionDataSize != 0;
+  if (NeedsReductionBuffer && LaunchArgs.OmpABIVersion < OMP_KERNEL_ARG_VERSION)
+    return error::createOffloadError(
+        error::ErrorCode::INVALID_BINARY,
+        "kernel was built against an older OpenMP kernel-launch-environment "
+        "ABI (v%u); current runtime requires v%u for cross-team reductions",
+        LaunchArgs.OmpABIVersion, OMP_KERNEL_ARG_VERSION);
+  if (!NeedsReductionBuffer && !DynCGroupMem)
+    return reinterpret_cast<KernelLaunchEnvironmentTy *>(~0);
+
+  auto AllocOrErr = GenericDevice.dataAlloc(
+      sizeof(KernelLaunchEnvironmentTy),
+      /*HostPtr=*/nullptr, TARGET_ALLOC_DEVICE, /*Alignment=*/0);
+  if (!AllocOrErr)
+    return AllocOrErr.takeError();
+
+  // Remember to free the memory later.
+  freeAfterSynchronization(GenericDevice, AsyncInfo, *AllocOrErr,
+                           TARGET_ALLOC_DEVICE);
+
+  // Use the KLE in the __tgt_async_info to ensure a stable address for the
+  // async data transfer.
+  auto &LocalKLE =
+      static_cast<__tgt_async_info *>(AsyncInfo)->KernelLaunchEnvironment;
+  LocalKLE = KernelLaunchEnvironmentTy{};
+  LocalKLE.DynCGroupMemSize = DynBlockMemConf.Size;
+  LocalKLE.DynCGroupMemFbPtr = DynBlockMemConf.FallbackPtr;
+  LocalKLE.DynCGroupMemFb = DynBlockMemConf.Fallback;
+  LocalKLE.ReductionBuffer = nullptr;
+
+  if (NeedsReductionBuffer) {
+    // Use number of teams many buffer elements.
+    auto ReductionAllocOrErr = GenericDevice.dataAlloc(
+        uint64_t(KernelEnv.ReductionDataSize) * NumBlocks0,
+        /*HostPtr=*/nullptr, TARGET_ALLOC_DEVICE, /*Alignment=*/0);
+    if (!ReductionAllocOrErr)
+      return ReductionAllocOrErr.takeError();
+    LocalKLE.ReductionBuffer = *ReductionAllocOrErr;
+    // Remember to free the memory later.
+    freeAfterSynchronization(GenericDevice, AsyncInfo, *ReductionAllocOrErr,
+                             TARGET_ALLOC_DEVICE);
+  }
+
+  INFO(OMP_INFOTYPE_DATA_TRANSFER, GenericDevice.getDeviceId(),
+       "Copying data from host to device, HstPtr=" DPxMOD ", TgtPtr=" DPxMOD
+       ", Size=%" PRId64 ", Name=KernelLaunchEnv\n",
+       DPxPTR(&LocalKLE), DPxPTR(*AllocOrErr),
+       sizeof(KernelLaunchEnvironmentTy));
+
+  if (auto Err = GenericDevice.dataSubmit(
+          *AllocOrErr, &LocalKLE, sizeof(KernelLaunchEnvironmentTy), AsyncInfo))
+    return Err;
+  return static_cast<KernelLaunchEnvironmentTy *>(*AllocOrErr);
+}
+
 /// Get the effective number of threads for the kernel based on the
 /// user-defined number of threads.
 static uint32_t getEffectiveNumThreads(GenericDeviceTy &GenericDevice,
@@ -505,32 +651,27 @@ getEffectiveNumBlocks(GenericDeviceTy &GenericDevice, uint32_t UserNumBlocks,
                   GenericDevice.getBlockLimit(EffectiveNumThreads));
 }
 
-// Run region on device
-int32_t DeviceTy::launchKernel(void *TgtEntryPtr, void **TgtVarsPtr,
-                               ptrdiff_t *TgtOffsets, KernelArgsTy &KernelArgs,
-                               KernelReplayOutcomeTy *ReplayOutcome,
-                               AsyncInfoTy &AsyncInfo) {
-  llvm::SmallVector<void *> Args, Ptrs;
-  llvm::SmallVector<int64_t> ArgSizes;
-
+/// Build the base KernelLaunchArgsTy for a launch from the public
+/// KernelArgsTy and the kernel's cached launch-geometry properties.
+static KernelLaunchArgsTy buildLaunchArgs(const KernelArgsTy &KernelArgs,
+                                          KernelReplayOutcomeTy *ReplayOutcome,
+                                          const KernelLaunchInfoTy &KernelEnv) {
   KernelLaunchArgsTy LaunchArgs;
   LaunchArgs.OmpABIVersion = KernelArgs.Version;
   LaunchArgs.ReplayOutcome = ReplayOutcome;
   LaunchArgs.ArgSizes = KernelArgs.ArgSizes;
   LaunchArgs.Tripcount = KernelArgs.Tripcount;
-  LaunchArgs.DynCGroupMem = KernelArgs.DynCGroupMem;
   llvm::copy(KernelArgs.UserNumBlocks, LaunchArgs.UserNumBlocks);
   llvm::copy(KernelArgs.UserThreadLimit, LaunchArgs.UserThreadLimit);
   LaunchArgs.Flags.Cooperative = KernelArgs.Flags.Cooperative;
-  LaunchArgs.Flags.DynCGroupMemFallback = KernelArgs.Flags.DynCGroupMemFallback;
-
-  KernelLaunchInfoTy KernelEnv = getKernelLaunchInfo(TgtEntryPtr);
-  LaunchArgs.KernelEnvironment.ReductionDataSize = KernelEnv.ReductionDataSize;
-  LaunchArgs.KernelEnvironment.MaxNumThreads = KernelEnv.MaxNumThreads;
-
-  const bool StrictBlocks = KernelArgs.Flags.StrictBlocks;
-  const bool StrictThreads = KernelArgs.Flags.StrictThreads;
+  LaunchArgs.MaxNumThreads = KernelEnv.MaxNumThreads;
+  return LaunchArgs;
+}
 
+/// Assert the launch geometry invariants expected by the plugin layer.
+static void checkLaunchInvariants(const KernelLaunchArgsTy &LaunchArgs,
+                                  const KernelArgsTy &KernelArgs,
+                                  const KernelLaunchInfoTy &KernelEnv) {
   // Multidimensional is only supported with bare mode for now.
   assert(KernelEnv.isBareMode() ||
          LaunchArgs.UserThreadLimit[1] == 1 &&
@@ -540,41 +681,58 @@ int32_t DeviceTy::launchKernel(void *TgtEntryPtr, void **TgtVarsPtr,
              "Non-bare mode should only use the first thread and block "
              "dimensions");
 
-  assert(!StrictBlocks ||
+  assert(!KernelArgs.Flags.StrictBlocks ||
          LaunchArgs.UserNumBlocks[0] > 0 && LaunchArgs.UserNumBlocks[1] > 0 &&
              LaunchArgs.UserNumBlocks[2] > 0 &&
              "Strict requires number of blocks greater than zero");
-  assert(!StrictThreads ||
+  assert(!KernelArgs.Flags.StrictThreads ||
          LaunchArgs.UserThreadLimit[0] > 0 &&
              LaunchArgs.UserThreadLimit[1] > 0 &&
              LaunchArgs.UserThreadLimit[2] > 0 &&
              "Strict requires number of threads greater than zero");
+}
+
+/// Calculate or adjust, in place, the effective number of threads and blocks
+/// for the first dimension, unless the caller requested strict counts.
+static void adjustEffectiveGeometry(GenericDeviceTy &GenericDevice,
+                                    KernelLaunchArgsTy &LaunchArgs,
+                                    const KernelArgsTy &KernelArgs,
+                                    const KernelLaunchInfoTy &KernelEnv) {
+  const bool StrictBlocks = KernelArgs.Flags.StrictBlocks;
+  const bool StrictThreads = KernelArgs.Flags.StrictThreads;
+  if (StrictThreads && StrictBlocks)
+    return;
+
+  assert(!KernelEnv.isBareMode() &&
+         "bare kernel launches must request strict thread/block counts");
 
   // Record whether the user actually requested a thread limit (thread_limit
   // clause) before possibly overwriting UserThreadLimit[0] below with the
   // computed effective value.
   const bool ThreadLimitFromUser = LaunchArgs.UserThreadLimit[0] > 0;
 
-  // Calculate or adjust the effective number of threads and blocks for the
-  // first dimension, if the caller didn't request strict counts.
-  if (!StrictThreads || !StrictBlocks) {
-    assert(!KernelEnv.isBareMode() &&
-           "bare kernel launches must request strict thread/block counts");
-
-    GenericDeviceTy &GenericDevice = RTL->getDevice(RTLDeviceID);
-    uint32_t EffectiveNumThreads = LaunchArgs.UserThreadLimit[0];
-    if (!StrictThreads)
-      EffectiveNumThreads =
-          getEffectiveNumThreads(GenericDevice, EffectiveNumThreads, KernelEnv);
+  uint32_t EffectiveNumThreads = LaunchArgs.UserThreadLimit[0];
+  if (!StrictThreads)
+    EffectiveNumThreads =
+        getEffectiveNumThreads(GenericDevice, EffectiveNumThreads, KernelEnv);
 
-    if (!StrictBlocks)
-      LaunchArgs.UserNumBlocks[0] = getEffectiveNumBlocks(
-          GenericDevice, LaunchArgs.UserNumBlocks[0], LaunchArgs.Tripcount,
-          EffectiveNumThreads, StrictThreads, ThreadLimitFromUser, KernelEnv);
+  if (!StrictBlocks)
+    LaunchArgs.UserNumBlocks[0] = getEffectiveNumBlocks(
+        GenericDevice, LaunchArgs.UserNumBlocks[0], LaunchArgs.Tripcount,
+        EffectiveNumThreads, StrictThreads, ThreadLimitFromUser, KernelEnv);
 
-    LaunchArgs.UserThreadLimit[0] = EffectiveNumThreads;
-  }
+  LaunchArgs.UserThreadLimit[0] = EffectiveNumThreads;
+}
 
+/// Flatten the kernel arguments into \p LaunchArgs.Args. Returns the address
+/// of the element reserved for the kernel launch environment (dyn_ptr), or
+/// null if this launch has no such slot.
+static void **resolveArgsAndDynPtrSlot(KernelArgsTy &KernelArgs,
+                                       void **TgtVarsPtr, ptrdiff_t *TgtOffsets,
+                                       llvm::SmallVector<void *> &Args,
+                                       llvm::SmallVector<void *> &Ptrs,
+                                       llvm::SmallVector<int64_t> &ArgSizes,
+                                       KernelLaunchArgsTy &LaunchArgs) {
   if (KernelArgs.Flags.IsCUDA) {
     // Kernel languages (CUDA/HIP) pass an already-flattened argument-pointer
     // array through KernelArgs.ArgPtrs instead of using the OpenMP
@@ -583,34 +741,106 @@ int32_t DeviceTy::launchKernel(void *TgtEntryPtr, void **TgtVarsPtr,
         reinterpret_cast<KernelLaunchParamsTy *>(KernelArgs.ArgPtrs);
     LaunchArgs.NumArgs = LaunchParams->NumArgs;
     LaunchArgs.Args = LaunchParams->Args;
-  } else {
-    resolveKernelLaunchParams(TgtVarsPtr, TgtOffsets, KernelArgs.NumArgs, Args,
-                              Ptrs, LaunchArgs);
-    // The dyn_ptr slot is reserved by the host (version >= 4) or by
-    // upgradeKernelArgs (version 3) as the last element of the argument
-    // array. Version 3 device kernels expect it first instead, so rotate it
-    // to the front to match that ABI.
-    if (KernelArgs.NumArgs > 0 &&
-        KernelArgs.Version >= OMP_KERNEL_ARG_MIN_VERSION_WITH_DYN_PTR) {
-      if (KernelArgs.Version == OMP_KERNEL_ARG_MIN_VERSION_WITH_DYN_PTR) {
-        std::rotate(Args.begin(), Args.end() - 1, Args.end());
-        LaunchArgs.DynPtrSlot = &Args[0];
-
-        // Keep ArgSizes in sync with the rotated Args, if present.
-        if (LaunchArgs.ArgSizes) {
-          ArgSizes.assign(LaunchArgs.ArgSizes,
-                          LaunchArgs.ArgSizes + KernelArgs.NumArgs);
-          std::rotate(ArgSizes.begin(), ArgSizes.end() - 1, ArgSizes.end());
-          LaunchArgs.ArgSizes = ArgSizes.data();
-        }
-      } else {
-        LaunchArgs.DynPtrSlot = &Args[KernelArgs.NumArgs - 1];
-      }
-    }
+    return nullptr;
+  }
+
+  resolveKernelLaunchParams(TgtVarsPtr, TgtOffsets, KernelArgs.NumArgs, Args,
+                            Ptrs, LaunchArgs);
+
+  if (KernelArgs.NumArgs == 0 ||
+      KernelArgs.Version < OMP_KERNEL_ARG_MIN_VERSION_WITH_DYN_PTR)
+    return nullptr;
+
+  // The dyn_ptr slot is reserved by the host (version >= 4) or by
+  // upgradeKernelArgs (version 3) as the last element of the argument array.
+  // Version 3 device kernels expect it first instead, so rotate it to the
+  // front to match that ABI.
+  if (KernelArgs.Version != OMP_KERNEL_ARG_MIN_VERSION_WITH_DYN_PTR)
+    return &Args[KernelArgs.NumArgs - 1];
+
+  std::rotate(Args.begin(), Args.end() - 1, Args.end());
+
+  // Keep ArgSizes in sync with the rotated Args, if present.
+  if (LaunchArgs.ArgSizes) {
+    ArgSizes.assign(LaunchArgs.ArgSizes,
+                    LaunchArgs.ArgSizes + KernelArgs.NumArgs);
+    std::rotate(ArgSizes.begin(), ArgSizes.end() - 1, ArgSizes.end());
+    LaunchArgs.ArgSizes = ArgSizes.data();
+  }
+  return &Args[0];
+}
+
+/// Compute the dynamic block-memory configuration for this launch, filling in
+/// \p LaunchArgs.DynCGroupMem with the native size to request, and, if this
+/// launch has a reserved dyn_ptr slot (\p DynPtrSlot), the device-side kernel
+/// launch environment.
+static llvm::Error
+prepareDynamicLaunchState(GenericDeviceTy &GenericDevice,
+                          const KernelLaunchInfoTy &KernelEnv,
+                          KernelLaunchArgsTy &LaunchArgs, uint32_t DynCGroupMem,
+                          DynCGroupMemFallbackType DynCGroupMemFallback,
+                          void **DynPtrSlot, AsyncInfoTy &AsyncInfo) {
+  uint32_t NumBlocksTotal = LaunchArgs.UserNumBlocks[0] *
+                            LaunchArgs.UserNumBlocks[1] *
+                            LaunchArgs.UserNumBlocks[2];
+  auto DynBlockMemConfOrErr =
+      prepareBlockMemory(GenericDevice, KernelEnv, DynCGroupMem,
+                         DynCGroupMemFallback, NumBlocksTotal);
+  if (!DynBlockMemConfOrErr)
+    return DynBlockMemConfOrErr.takeError();
+
+  DynBlockMemConfTy &DynBlockMemConf = *DynBlockMemConfOrErr;
+  LaunchArgs.DynCGroupMem = DynBlockMemConf.NativeSize;
+  if (DynBlockMemConf.FallbackPtr)
+    freeAfterSynchronization(GenericDevice, AsyncInfo,
+                             DynBlockMemConf.FallbackPtr, TARGET_ALLOC_DEVICE);
+
+  auto KernelLaunchEnvOrErr = getKernelLaunchEnvironment(
+      GenericDevice, LaunchArgs, KernelEnv, DynBlockMemConf, DynCGroupMem,
+      DynPtrSlot, AsyncInfo, LaunchArgs.UserNumBlocks[0]);
+  if (!KernelLaunchEnvOrErr)
+    return KernelLaunchEnvOrErr.takeError();
+
+  // Fill in the kernel launch environment (dyn_ptr) if this launch has a
+  // reserved slot for it. When replaying, getKernelLaunchEnvironment()
+  // returns null so the recorded value already in the slot is preserved.
+  if (DynPtrSlot && *KernelLaunchEnvOrErr)
+    *DynPtrSlot = *KernelLaunchEnvOrErr;
+
+  return llvm::Error::success();
+}
+
+// Run region on device
+int32_t DeviceTy::launchKernel(void *TgtEntryPtr, void **TgtVarsPtr,
+                               ptrdiff_t *TgtOffsets, KernelArgsTy &KernelArgs,
+                               KernelReplayOutcomeTy *ReplayOutcome,
+                               AsyncInfoTy &AsyncInfo) {
+  llvm::SmallVector<void *> Args, Ptrs;
+  llvm::SmallVector<int64_t> ArgSizes;
+
+  GenericDeviceTy &GenericDevice = RTL->getDevice(RTLDeviceID);
+  KernelLaunchInfoTy KernelEnv = getKernelLaunchInfo(TgtEntryPtr);
+  KernelLaunchArgsTy LaunchArgs =
+      buildLaunchArgs(KernelArgs, ReplayOutcome, KernelEnv);
+
+  checkLaunchInvariants(LaunchArgs, KernelArgs, KernelEnv);
+  adjustEffectiveGeometry(GenericDevice, LaunchArgs, KernelArgs, KernelEnv);
+
+  void **DynPtrSlot = resolveArgsAndDynPtrSlot(
+      KernelArgs, TgtVarsPtr, TgtOffsets, Args, Ptrs, ArgSizes, LaunchArgs);
+
+  auto DynCGroupMemFallback = static_cast<DynCGroupMemFallbackType>(
+      KernelArgs.Flags.DynCGroupMemFallback);
+  if (auto Err = prepareDynamicLaunchState(
+          GenericDevice, KernelEnv, LaunchArgs, KernelArgs.DynCGroupMem,
+          DynCGroupMemFallback, DynPtrSlot, AsyncInfo)) {
+    REPORT() << "Failure to prepare launch state for kernel " << TgtEntryPtr
+             << ": " << toString(std::move(Err));
+    return OFFLOAD_FAIL;
   }
 
   auto *Kernel = reinterpret_cast<GenericKernelTy *>(TgtEntryPtr);
-  INFO(OMP_INFOTYPE_PLUGIN_KERNEL, RTL->getDevice(RTLDeviceID).getDeviceId(),
+  INFO(OMP_INFOTYPE_PLUGIN_KERNEL, GenericDevice.getDeviceId(),
        "Launching kernel %s with [%u,%u,%u] blocks and [%u,%u,%u] threads in "
        "%s mode\n",
        Kernel->getName(), LaunchArgs.UserNumBlocks[0],
diff --git a/offload/plugins-nextgen/common/include/PluginInterface.h b/offload/plugins-nextgen/common/include/PluginInterface.h
index 0bd2558b0c204..b72e9a582d701 100644
--- a/offload/plugins-nextgen/common/include/PluginInterface.h
+++ b/offload/plugins-nextgen/common/include/PluginInterface.h
@@ -320,18 +320,6 @@ struct InfoTreeNode {
   }
 };
 
-/// Configuration of dynamic block memory needed for launching a kernel.
-struct DynBlockMemConfTy {
-  /// The size of the dynamic block memory buffer.
-  uint32_t Size = 0;
-  /// The size of dynamic shared memory natively provided by the device.
-  uint32_t NativeSize = 0;
-  /// The fallback that was triggered (if any).
-  DynCGroupMemFallbackType Fallback = DynCGroupMemFallbackType::None;
-  /// The fallback pointer if global memory was used as alternative.
-  void *FallbackPtr = nullptr;
-};
-
 /// Tracker of virtual memory address reservations.
 template <typename HandleTy> class VMemTrackerTy {
   struct EntryTy {
@@ -446,11 +434,6 @@ struct KernelLaunchArgsTy {
   /// Size of the argument data in bytes, one entry per \p Args element,
   /// possibly null.
   int64_t *ArgSizes = nullptr;
-  /// Address of the element of \p Args reserved for the kernel launch
-  /// environment (dyn_ptr), or null if this launch has no such slot. The
-  /// caller owns the storage it points into; the plugin fills it in once it
-  /// has computed the actual (device-side) value.
-  void **DynPtrSlot = nullptr;
   /// Tripcount for the teams / distribute loop, 0 otherwise.
   uint64_t Tripcount = 0;
   /// Amount of dynamic cgroup memory requested.
@@ -459,18 +442,12 @@ struct KernelLaunchArgsTy {
   uint32_t UserNumBlocks[3] = {0, 0, 0};
   /// User-requested number of threads (for x,y,z dimension).
   uint32_t UserThreadLimit[3] = {0, 0, 0};
-  struct {
-    /// Size in bytes of a single cross-team reduction buffer element for
-    /// this kernel, or 0 if the kernel does not need a reduction buffer.
-    uint32_t ReductionDataSize = 0;
-    /// Maximum number of threads per block that this kernel may use.
-    uint32_t MaxNumThreads = 0;
-  } KernelEnvironment;
+  /// Maximum number of threads per block that this kernel may use.
+  uint32_t MaxNumThreads = 0;
   struct {
     uint64_t Cooperative : 1; // Was this kernel spawned as cooperative.
-    uint64_t DynCGroupMemFallback : 2; // The fallback for dynamic cgroup mem.
-    uint64_t Unused : 61;
-  } Flags = {0, 0, 0};
+    uint64_t Unused : 63;
+  } Flags = {0, 0};
   /// Set by the caller when replaying a previously recorded kernel launch, so
   /// the plugin can report the outcome back; null for a normal launch.
   KernelReplayOutcomeTy *ReplayOutcome = nullptr;
@@ -492,10 +469,7 @@ struct GenericKernelTy {
 
   /// Launch the kernel on the specific device. The device must be the same
   /// one used to initialize the kernel. \p LaunchArgs.Args is the flattened
-  /// argument-pointer array to pass to the kernel, with any offsets already
-  /// resolved. \p LaunchArgs.DynPtrSlot, if non-null, points at the element
-  /// of it reserved for the kernel launch environment (dyn_ptr); the caller
-  /// owns the storage it points into.
+  /// argument-pointer array to pass to the kernel, with any offsets.
   Error launch(GenericDeviceTy &GenericDevice, KernelLaunchArgsTy &LaunchArgs,
                AsyncInfoWrapperTy &AsyncInfoWrapper) const;
   virtual Error launchImpl(GenericDeviceTy &GenericDevice,
@@ -534,15 +508,6 @@ struct GenericKernelTy {
     return *ImagePtr;
   }
 
-  /// Return a device pointer to a new kernel launch environment.
-  ///
-  /// \p NumBlocks0 is the number of blocks for this launch and is used to size
-  /// the reduction buffer.
-  Expected<KernelLaunchEnvironmentTy *> getKernelLaunchEnvironment(
-      GenericDeviceTy &GenericDevice, const KernelLaunchArgsTy &LaunchArgs,
-      const DynBlockMemConfTy &DynBlockMemConf,
-      AsyncInfoWrapperTy &AsyncInfoWrapper, uint32_t NumBlocks0) const;
-
   /// Indicate whether an execution mode is valid.
   static bool isValidExecutionMode(OMPTgtExecModeFlags ExecutionMode) {
     switch (ExecutionMode) {
@@ -570,13 +535,6 @@ struct GenericKernelTy {
                                        uint32_t NumBlocks[3]) const;
 
 private:
-  /// Prepare the block memory buffer requested for the kernel and execute the
-  /// specified fallback if necessary.
-  Expected<DynBlockMemConfTy>
-  prepareBlockMemory(GenericDeviceTy &GenericDevice,
-                     const KernelLaunchArgsTy &LaunchArgs,
-                     uint32_t NumBlocks) const;
-
   /// The kernel name.
   std::string Name;
 
@@ -586,9 +544,6 @@ struct GenericKernelTy {
 protected:
   /// The static memory sized per block.
   uint32_t StaticBlockMemSize = 0;
-
-  /// The prototype kernel launch environment.
-  KernelLaunchEnvironmentTy KernelLaunchEnvironment;
 };
 
 /// Information about an allocation, when it has been allocated, and when/if it
diff --git a/offload/plugins-nextgen/common/src/PluginInterface.cpp b/offload/plugins-nextgen/common/src/PluginInterface.cpp
index 467244b5c9a77..6bced40873243 100644
--- a/offload/plugins-nextgen/common/src/PluginInterface.cpp
+++ b/offload/plugins-nextgen/common/src/PluginInterface.cpp
@@ -75,78 +75,6 @@ Error GenericKernelTy::init(GenericDeviceTy &GenericDevice,
   return initImpl(GenericDevice, Image);
 }
 
-Expected<KernelLaunchEnvironmentTy *>
-GenericKernelTy::getKernelLaunchEnvironment(
-    GenericDeviceTy &GenericDevice, const KernelLaunchArgsTy &LaunchArgs,
-    const DynBlockMemConfTy &DynBlockMemConf,
-    AsyncInfoWrapperTy &AsyncInfoWrapper, uint32_t NumBlocks0) const {
-  // Ctor/Dtor have no arguments, replaying uses the original kernel launch
-  // environment, and launches with no reserved dyn_ptr slot (e.g. older
-  // compiler versions, or non-OpenMP launches) have nowhere to store one.
-  if ((GenericDevice.getRecordReplay() &&
-       GenericDevice.getRecordReplay()->isReplaying()) ||
-      !LaunchArgs.DynPtrSlot)
-    return nullptr;
-
-  const bool NeedsReductionBuffer =
-      LaunchArgs.KernelEnvironment.ReductionDataSize != 0;
-  if (NeedsReductionBuffer && LaunchArgs.OmpABIVersion < OMP_KERNEL_ARG_VERSION)
-    return Plugin::error(ErrorCode::INVALID_BINARY,
-                         "kernel was built against an older OpenMP "
-                         "kernel-launch-environment ABI (v%u); current "
-                         "runtime requires v%u for cross-team reductions",
-                         LaunchArgs.OmpABIVersion, OMP_KERNEL_ARG_VERSION);
-  if (!NeedsReductionBuffer && !LaunchArgs.DynCGroupMem)
-    return reinterpret_cast<KernelLaunchEnvironmentTy *>(~0);
-
-  auto AllocOrErr = GenericDevice.dataAlloc(
-      sizeof(KernelLaunchEnvironmentTy),
-      /*HostPtr=*/nullptr, TargetAllocTy::TARGET_ALLOC_DEVICE, /*Alignment=*/0);
-  if (!AllocOrErr)
-    return AllocOrErr.takeError();
-
-  // Remember to free the memory later.
-  AsyncInfoWrapper.freeAllocationAfterSynchronization(
-      *AllocOrErr, TargetAllocTy::TARGET_ALLOC_DEVICE);
-
-  /// Use the KLE in the __tgt_async_info to ensure a stable address for the
-  /// async data transfer.
-  auto &LocalKLE = (*AsyncInfoWrapper).KernelLaunchEnvironment;
-  LocalKLE = KernelLaunchEnvironment;
-
-  LocalKLE.DynCGroupMemSize = DynBlockMemConf.Size;
-  LocalKLE.DynCGroupMemFbPtr = DynBlockMemConf.FallbackPtr;
-  LocalKLE.DynCGroupMemFb = DynBlockMemConf.Fallback;
-  LocalKLE.ReductionBuffer = nullptr;
-
-  if (NeedsReductionBuffer) {
-    // Use number of teams many buffer elements.
-    auto AllocOrErr = GenericDevice.dataAlloc(
-        uint64_t(LaunchArgs.KernelEnvironment.ReductionDataSize) * NumBlocks0,
-        /*HostPtr=*/nullptr, TargetAllocTy::TARGET_ALLOC_DEVICE,
-        /*Alignment=*/0);
-    if (!AllocOrErr)
-      return AllocOrErr.takeError();
-    LocalKLE.ReductionBuffer = *AllocOrErr;
-    // Remember to free the memory later.
-    AsyncInfoWrapper.freeAllocationAfterSynchronization(
-        *AllocOrErr, TargetAllocTy::TARGET_ALLOC_DEVICE);
-  }
-
-  INFO(OMP_INFOTYPE_DATA_TRANSFER, GenericDevice.getDeviceId(),
-       "Copying data from host to device, HstPtr=" DPxMOD ", TgtPtr=" DPxMOD
-       ", Size=%" PRId64 ", Name=KernelLaunchEnv\n",
-       DPxPTR(&LocalKLE), DPxPTR(*AllocOrErr),
-       sizeof(KernelLaunchEnvironmentTy));
-
-  auto Err = GenericDevice.dataSubmit(*AllocOrErr, &LocalKLE,
-                                      sizeof(KernelLaunchEnvironmentTy),
-                                      AsyncInfoWrapper);
-  if (Err)
-    return Err;
-  return static_cast<KernelLaunchEnvironmentTy *>(*AllocOrErr);
-}
-
 Error GenericKernelTy::printLaunchInfo(GenericDeviceTy &GenericDevice,
                                        const KernelLaunchArgsTy &LaunchArgs,
                                        uint32_t NumThreads[3],
@@ -161,53 +89,6 @@ Error GenericKernelTy::printLaunchInfoDetails(
   return Plugin::success();
 }
 
-Expected<DynBlockMemConfTy>
-GenericKernelTy::prepareBlockMemory(GenericDeviceTy &GenericDevice,
-                                    const KernelLaunchArgsTy &LaunchArgs,
-                                    uint32_t NumBlocks) const {
-  uint32_t MaxBlockMemSize = GenericDevice.getMaxBlockSharedMemSize();
-  uint32_t DynBlockMemSize = LaunchArgs.DynCGroupMem;
-  uint32_t TotalBlockMemSize = StaticBlockMemSize + DynBlockMemSize;
-  uint32_t DynNativeBlockMemSize = DynBlockMemSize;
-  void *DynFallbackPtr = nullptr;
-
-  // No enough block memory to cover the static one. Cannot run the kernel.
-  if (StaticBlockMemSize > MaxBlockMemSize)
-    return Plugin::error(ErrorCode::INVALID_ARGUMENT,
-                         "Static block memory size exceeds maximum");
-  // No enough block memory to cover dynamic one, and the fallback is aborting.
-  if (static_cast<DynCGroupMemFallbackType>(
-          LaunchArgs.Flags.DynCGroupMemFallback) ==
-          DynCGroupMemFallbackType::Abort &&
-      TotalBlockMemSize > MaxBlockMemSize)
-    return Plugin::error(
-        ErrorCode::INVALID_ARGUMENT,
-        "Requested block memory size (static + dynamic) exceeds maximum");
-
-  DynCGroupMemFallbackType DynFallback = DynCGroupMemFallbackType::None;
-  if (DynBlockMemSize && TotalBlockMemSize > MaxBlockMemSize) {
-    // Launch without native dynamic block memory.
-    DynNativeBlockMemSize = 0;
-    DynFallback = static_cast<DynCGroupMemFallbackType>(
-        LaunchArgs.Flags.DynCGroupMemFallback);
-    if (DynFallback != DynCGroupMemFallbackType::DefaultMem) {
-      // Do not provide any memory as fallback.
-      DynBlockMemSize = 0;
-    } else {
-      // Get global memory as fallback.
-      auto AllocOrErr = GenericDevice.dataAlloc(
-          NumBlocks * DynBlockMemSize,
-          /*HostPtr=*/nullptr, TargetAllocTy::TARGET_ALLOC_DEVICE,
-          /*Alignment=*/0);
-      if (!AllocOrErr)
-        return AllocOrErr.takeError();
-      DynFallbackPtr = *AllocOrErr;
-    }
-  }
-  return DynBlockMemConfTy{DynBlockMemSize, DynNativeBlockMemSize, DynFallback,
-                           DynFallbackPtr};
-}
-
 Error GenericKernelTy::launch(GenericDeviceTy &GenericDevice,
                               KernelLaunchArgsTy &LaunchArgs,
                               AsyncInfoWrapperTy &AsyncInfoWrapper) const {
@@ -218,29 +99,6 @@ Error GenericKernelTy::launch(GenericDeviceTy &GenericDevice,
                                     LaunchArgs.UserNumBlocks[1],
                                     LaunchArgs.UserNumBlocks[2]};
 
-  auto DynBlockMemConfOrErr = prepareBlockMemory(
-      GenericDevice, LaunchArgs,
-      EffectiveNumBlocks[0] * EffectiveNumBlocks[1] * EffectiveNumBlocks[2]);
-  if (!DynBlockMemConfOrErr)
-    return DynBlockMemConfOrErr.takeError();
-
-  DynBlockMemConfTy &DynBlockMemConf = *DynBlockMemConfOrErr;
-  if (DynBlockMemConf.FallbackPtr)
-    AsyncInfoWrapper.freeAllocationAfterSynchronization(
-        DynBlockMemConf.FallbackPtr, TargetAllocTy::TARGET_ALLOC_DEVICE);
-
-  auto KernelLaunchEnvOrErr =
-      getKernelLaunchEnvironment(GenericDevice, LaunchArgs, DynBlockMemConf,
-                                 AsyncInfoWrapper, EffectiveNumBlocks[0]);
-  if (!KernelLaunchEnvOrErr)
-    return KernelLaunchEnvOrErr.takeError();
-
-  // Fill in the kernel launch environment (dyn_ptr) if this launch has a
-  // reserved slot for it. When replaying, getKernelLaunchEnvironment()
-  // returns null so the recorded value already in the slot is preserved.
-  if (LaunchArgs.DynPtrSlot && *KernelLaunchEnvOrErr)
-    *LaunchArgs.DynPtrSlot = *KernelLaunchEnvOrErr;
-
   if (auto Err = printLaunchInfo(GenericDevice, LaunchArgs, EffectiveNumThreads,
                                  EffectiveNumBlocks))
     return Err;
@@ -255,7 +113,7 @@ Error GenericKernelTy::launch(GenericDeviceTy &GenericDevice,
     // Record the kernel prologue data before kernel launch.
     auto RRHandleOrErr = RecordReplay->recordPrologue(
         *this, LaunchArgs, EffectiveNumBlocks, EffectiveNumThreads,
-        DynBlockMemConf.NativeSize);
+        LaunchArgs.DynCGroupMem);
     if (!RRHandleOrErr)
       return RRHandleOrErr.takeError();
     RRHandle = *RRHandleOrErr;
@@ -263,7 +121,7 @@ Error GenericKernelTy::launch(GenericDeviceTy &GenericDevice,
 
   if (auto Err =
           launchImpl(GenericDevice, EffectiveNumThreads, EffectiveNumBlocks,
-                     DynBlockMemConf.NativeSize, LaunchArgs, AsyncInfoWrapper))
+                     LaunchArgs.DynCGroupMem, LaunchArgs, AsyncInfoWrapper))
     return Err;
 
   if (RecordReplay) {
diff --git a/offload/plugins-nextgen/common/src/RecordReplay.cpp b/offload/plugins-nextgen/common/src/RecordReplay.cpp
index 768e85338a140..f97c26ceb2c52 100644
--- a/offload/plugins-nextgen/common/src/RecordReplay.cpp
+++ b/offload/plugins-nextgen/common/src/RecordReplay.cpp
@@ -279,7 +279,7 @@ Error NativeRecordReplayTy::recordDescImpl(
 
   // Export minimum and maximum for allowed number of threads. If zero, it means
   // there was no restriction provided by the program.
-  uint32_t MaxThreads = LaunchArgs.KernelEnvironment.MaxNumThreads;
+  uint32_t MaxThreads = LaunchArgs.MaxNumThreads;
   json::Array JsonThreadsLimits;
   JsonThreadsLimits.push_back(1);
   JsonThreadsLimits.push_back(MaxThreads);



More information about the llvm-branch-commits mailing list