[Mlir-commits] [mlir] [MLIR][NVVM] Add S2G and Reduce override NVVM Dialect ops (PR #216481)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Sat Aug 15 03:42:41 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir-llvm

@llvm/pr-subscribers-mlir

Author: Rajat Bajpai (rajatbajpai)

<details>
<summary>Changes</summary>

This change adds S2G and Reduction NVVM Dialect operations with tensor map override capability.

---

Patch is 98.20 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/216481.diff


6 Files Affected:

- (modified) mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td (+212) 
- (modified) mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp (+200) 
- (added) mlir/test/Target/LLVMIR/nvvm/nvvmir-invalid/tma_reduce_override_invalid.mlir (+34) 
- (added) mlir/test/Target/LLVMIR/nvvm/nvvmir-invalid/tma_store_override_invalid.mlir (+50) 
- (added) mlir/test/Target/LLVMIR/nvvm/tma_store_override.mlir (+154) 
- (added) mlir/test/Target/LLVMIR/nvvm/tma_store_reduce_override.mlir (+372) 


``````````diff
diff --git a/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td b/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td
index 46c38bcb5475d..b74be0a427ee0 100644
--- a/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td
+++ b/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td
@@ -4383,6 +4383,10 @@ def NVVM_CpAsyncBulkTensorGlobalToSharedClusterOp :
   }];
 }
 
+//===----------------------------------------------------------------------===//
+// NVVM S2G Ops
+//===----------------------------------------------------------------------===//
+
 def NVVM_CpAsyncBulkTensorSharedCTAToGlobalOp : 
   NVVM_PTXBuilder_Op<"cp.async.bulk.tensor.global.shared.cta",
   [AttrSizedOperandSegments]>,
@@ -4443,6 +4447,107 @@ def NVVM_CpAsyncBulkTensorSharedCTAToGlobalOp :
   }];
 }
 
+def NVVM_CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp :
+  NVVM_VoidIntrinsicOp<"cp.async.bulk.tensor.global.shared.cta.override",
+                       [AttrSizedOperandSegments]> {
+    let arguments = (ins
+      LLVM_AnyPointer:$tmaDesc,
+      LLVM_PointerShared:$srcMem,
+      LLVM_PointerGlobal:$overrideAdrr,
+      Variadic<I32>:$coordinates,
+      Variadic<I16>:$tensorSize,
+      Variadic<I32>:$lowerStride,
+      Optional<I16>:$upperStride,
+      Optional<I64>:$l2CacheHint,
+      DefaultValuedAttr<TMAStoreModeAttr, "TMAStoreMode::TILE">:$mode);
+
+    let summary = "Async bulk tensor copy from shared::cta to global memory with "
+                  "tensor-map field overrides";
+    let description = [{
+      Initiates an asynchronous copy of tensor data from shared::cta memory to global
+      memory while overriding specific fields of the opaque tensor-map object with
+      explicit operands. It corresponds to the `cp.async.bulk.tensor.[1-5]d.*` PTX
+      instructions qualified with `.override::*`.
+
+      The `mode` attribute selects the store mode. The override variant is selected by
+      which optional operands are provided:
+
+      - `override.addr` (`.override::global_address`): only `overrideAdrr` is given; the
+        global base address from the tensor-map is replaced. Supported in `TILE` (1D–5D),
+        `IM2COL`/`IM2COL_W` (3D–5D), and `TILE_SCATTER4` (2D, requires 5 coordinates).
+      - `override.addr.dim` (1D, `TILE` only): `overrideAdrr` plus `tensorSize` (one
+        element); also overrides the tensor global dimension.
+      - `override.addr.dim.stride` (2D–5D, `TILE` only): `overrideAdrr` plus `tensorSize`,
+        `lowerStride`, and `upperStride`; also overrides the global dimensions and strides.
+        The effective global stride is
+        `global_stride[i] = ((%stride{i} + (%upper_stride{i} << 32)) << 4)`.
+
+      `overrideAdrr` must be 16B aligned (a runtime error is raised otherwise) and the
+      memory range `[%override_addr, %override_addr + 128 KiB)` must be allocated and
+      accessible during execution. When overriding dimensions/strides, the base-address
+      override is mandatory and the tensor start coordinates must be zero; otherwise the
+      behavior is undefined.
+
+      The optional `l2CacheHint` specifies a cache-eviction policy for the access.
+
+      Examples:
+
+      // override.addr (TILE, 1D)
+      ```mlir
+      nvvm.cp.async.bulk.tensor.global.shared.cta.override %tma_desc, %src, %override_addr, box[%d0] {
+        mode = #nvvm.tma_store_mode<tile>
+      } : !llvm.ptr, !llvm.ptr<3>, !llvm.ptr<1>
+      ```
+
+      // override.addr (IM2COL, 3D) with L2 cache hint
+      ```mlir
+      nvvm.cp.async.bulk.tensor.global.shared.cta.override %tma_desc, %src, %override_addr, box[%d0, %d1, %d2]
+        l2_cache_hint = %ch {
+        mode = #nvvm.tma_store_mode<im2col>
+      } : !llvm.ptr, !llvm.ptr<3>, !llvm.ptr<1>
+      ```
+
+      // override.addr (TILE_SCATTER4, 2D) — 5 coordinates: x0, y0, y1, y2, y3
+      ```mlir
+      nvvm.cp.async.bulk.tensor.global.shared.cta.override %tma_desc, %src, %override_addr, box[%x0, %y0, %y1, %y2, %y3] {
+        mode = #nvvm.tma_store_mode<tile_scatter4>
+      } : !llvm.ptr, !llvm.ptr<3>, !llvm.ptr<1>
+      ```
+
+      // override.addr.dim (TILE, 1D)
+      ```mlir
+      nvvm.cp.async.bulk.tensor.global.shared.cta.override %tma_desc, %src, %override_addr,
+        box[%d0] tensor_size[%ts0] {
+        mode = #nvvm.tma_store_mode<tile>
+      } : !llvm.ptr, !llvm.ptr<3>, !llvm.ptr<1>
+      ```
+
+      // override.addr.dim.stride (TILE, 2D)
+      ```mlir
+      nvvm.cp.async.bulk.tensor.global.shared.cta.override %tma_desc, %src, %override_addr,
+        box[%d0, %d1] tensor_size[%ts0, %ts1] lower_stride[%lstrd0] upper_stride[%ustrd] {
+        mode = #nvvm.tma_store_mode<tile>
+      } : !llvm.ptr, !llvm.ptr<3>, !llvm.ptr<1>
+      ```
+
+      [For more information, see PTX ISA](https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#data-movement-and-conversion-instructions-cp-async-bulk-tensor)
+    }];
+
+    let assemblyFormat = [{
+      $tmaDesc `,`
+      $srcMem `,`
+      $overrideAdrr `,`
+      `box` `[`$coordinates `]`
+      (`tensor_size` `[`$tensorSize^ `]`)?
+      (`lower_stride` `[`$lowerStride^ `]`)?
+      (`upper_stride` `[`$upperStride^ `]`)?
+      (`l2_cache_hint` `=` $l2CacheHint^)?
+      attr-dict  `:` type($tmaDesc) `,` type($srcMem) `,` type($overrideAdrr)
+    }];
+
+    let hasVerifier = 1;
+}
+
 //===----------------------------------------------------------------------===//
 // NVVM Prefetch Op
 //===----------------------------------------------------------------------===//
@@ -4583,6 +4688,10 @@ def NVVM_CpAsyncBulkTensorPrefetchOp :
   let hasVerifier = 1;
 }
 
+//===----------------------------------------------------------------------===//
+// NVVM Reduction Ops
+//===----------------------------------------------------------------------===//
+
 // List of Reduction Ops supported with TMA Store
 def TMAReduxKindAdd : I32EnumAttrCase<"ADD", 0, "add">;
 def TMAReduxKindMin : I32EnumAttrCase<"MIN", 1, "min">;
@@ -4652,6 +4761,109 @@ def NVVM_CpAsyncBulkTensorReduceOp :
   }];
 }
 
+def NVVM_CpAsyncBulkTensorReduceOverrideAddrOp :
+  NVVM_VoidIntrinsicOp<"cp.async.bulk.tensor.reduce.override",
+                       [AttrSizedOperandSegments]> {
+    let arguments = (ins
+      LLVM_AnyPointer:$tmaDesc,
+      LLVM_PointerShared:$srcMem,
+      LLVM_PointerGlobal:$overrideAdrr,
+      Variadic<I32>:$coordinates,
+      Variadic<I16>:$tensorSize,
+      Variadic<I32>:$lowerStride,
+      Optional<I16>:$upperStride,
+      Optional<I64>:$l2CacheHint,
+      TMAReduxKindAttr:$redKind,
+      DefaultValuedAttr<TMAStoreModeAttr, "TMAStoreMode::TILE">:$mode
+      );
+
+    let summary = "Async bulk tensor reduction from shared::cta to global memory with "
+                  "tensor-map field overrides";
+    let description = [{
+      Initiates an asynchronous reduction of tensor data in global memory with the tensor
+      data in shared::cta memory, while overriding specific fields of the opaque tensor-map
+      object with explicit operands. It corresponds to the
+      `cp.reduce.async.bulk.tensor.[1-5]d.global.shared::cta.*` PTX instructions qualified
+      with `.override::*`.
+
+      The `mode` attribute selects the store mode and `redKind` selects the reduction
+      operation (`ADD`, `MIN`, `MAX`, `INC`, `DEC`, `AND`, `OR`, `XOR`) that combines the
+      source data in shared memory with the destination data in global memory. The override
+      variant is selected by which optional operands are provided:
+
+      - `override.addr` (`.override::global_address`): only `overrideAdrr` is given; the
+        global base address from the tensor-map is replaced. Supported in `TILE` (1D–5D),
+        `IM2COL` and `IM2COL_W` (3D–5D). `TILE_SCATTER4` is not supported.
+      - `override.addr.dim` (1D, `TILE` only): `overrideAdrr` plus `tensorSize` (one
+        element); also overrides the tensor global dimension.
+      - `override.addr.dim.stride` (2D–5D, `TILE` only): `overrideAdrr` plus `tensorSize`,
+        `lowerStride`, and `upperStride`; also overrides the global dimensions and strides.
+        The effective global stride is
+        `global_stride[i] = ((%stride{i} + (%upper_stride{i} << 32)) << 4)`.
+
+      `overrideAdrr` must be 16B aligned (a runtime error is raised otherwise) and the
+      memory range `[%override_addr, %override_addr + 128 KiB)` must be allocated and
+      accessible during execution. When overriding dimensions/strides, the base-address
+      override is mandatory and the tensor start coordinates must be zero; otherwise the
+      behavior is undefined.
+
+      The optional `l2CacheHint` specifies a cache-eviction policy for the access.
+
+      Examples:
+
+      // override.addr (TILE, 1D) with ADD reduction
+      ```mlir
+      nvvm.cp.async.bulk.tensor.reduce.override %tma_desc, %src, %override_addr, box[%d0] {
+        redKind = #nvvm.tma_redux_kind<add>,
+        mode = #nvvm.tma_store_mode<tile>
+      } : !llvm.ptr, !llvm.ptr<3>, !llvm.ptr<1>
+      ```
+
+      // override.addr (IM2COL, 3D) with L2 cache hint and MIN reduction
+      ```mlir
+      nvvm.cp.async.bulk.tensor.reduce.override %tma_desc, %src, %override_addr, box[%d0, %d1, %d2]
+        l2_cache_hint = %ch {
+        redKind = #nvvm.tma_redux_kind<min>,
+        mode = #nvvm.tma_store_mode<im2col>
+      } : !llvm.ptr, !llvm.ptr<3>, !llvm.ptr<1>
+      ```
+
+      // override.addr.dim (TILE, 1D)
+      ```mlir
+      nvvm.cp.async.bulk.tensor.reduce.override %tma_desc, %src, %override_addr,
+        box[%d0] tensor_size[%ts0] {
+        redKind = #nvvm.tma_redux_kind<add>,
+        mode = #nvvm.tma_store_mode<tile>
+      } : !llvm.ptr, !llvm.ptr<3>, !llvm.ptr<1>
+      ```
+
+      // override.addr.dim.stride (TILE, 2D)
+      ```mlir
+      nvvm.cp.async.bulk.tensor.reduce.override %tma_desc, %src, %override_addr,
+        box[%d0, %d1] tensor_size[%ts0, %ts1] lower_stride[%lstrd0] upper_stride[%ustrd] {
+        redKind = #nvvm.tma_redux_kind<and>,
+        mode = #nvvm.tma_store_mode<tile>
+      } : !llvm.ptr, !llvm.ptr<3>, !llvm.ptr<1>
+      ```
+
+      [For more information, see PTX ISA](https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#data-movement-and-conversion-instructions-cp-reduce-async-bulk-tensor)
+    }];
+
+    let assemblyFormat = [{
+      $tmaDesc `,`
+      $srcMem `,`
+      $overrideAdrr `,`
+      `box` `[`$coordinates `]`
+      (`tensor_size` `[`$tensorSize^ `]`)?
+      (`lower_stride` `[`$lowerStride^ `]`)?
+      (`upper_stride` `[`$upperStride^ `]`)?
+      (`l2_cache_hint` `=` $l2CacheHint^)?
+      attr-dict  `:` type($tmaDesc) `,` type($srcMem) `,` type($overrideAdrr)
+    }];
+
+    let hasVerifier = 1;
+}
+
 def NVVM_CpAsyncBulkGlobalToSharedClusterOp :
   NVVM_Op<"cp.async.bulk.shared.cluster.global", [AttrSizedOperandSegments]> {
   let summary = "Async bulk copy from global to Shared {cta or cluster} memory";
diff --git a/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp b/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp
index 31fed0b25990d..f43475eeb8d99 100644
--- a/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp
+++ b/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp
@@ -120,6 +120,32 @@ static LogicalResult cpAsyncBulkTensorCommonVerifier(size_t tensorDims,
   return success();
 }
 
+LogicalResult CpAsyncBulkTensorOverrideAddrCommonVerifier(
+    OperandRange coordinates, OperandRange tensorSize, OperandRange lowerStride,
+    Value upperStride, bool isTile, Location loc) {
+  LogicalResult res = success();
+  if (!tensorSize.empty() && coordinates.size() != tensorSize.size())
+    res =
+        emitError(loc, "Expected coordinates size to be equal to tensor size");
+
+  if (!lowerStride.empty() && lowerStride.size() != tensorSize.size() - 1)
+    res = emitError(
+        loc,
+        "Expected lower_stride size to be equal to one less than tensor size");
+
+  if (!lowerStride.empty() != static_cast<bool>(upperStride))
+    res = emitError(loc,
+                    "Expected lower_stride and upper_stride to be either both "
+                    "present or both absent");
+
+  bool isDimStride = tensorSize.size() > 0;
+  if (!isTile && isDimStride)
+    res = emitError(
+        loc, "Only tile mode supports override address with dim and stride");
+
+  return res;
+}
+
 LogicalResult CpAsyncBulkTensorSharedCTAToGlobalOp::verify() {
   TMAStoreMode mode = getMode();
   // We lower through inline-ptx when getPredicate() is true.
@@ -146,6 +172,24 @@ LogicalResult CpAsyncBulkTensorSharedCTAToGlobalOp::verify() {
   return success();
 }
 
+LogicalResult CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp::verify() {
+  TMAStoreMode mode = getMode();
+  bool isIm2Col = mode == TMAStoreMode::IM2COL;
+  bool isTile = mode == TMAStoreMode::TILE;
+
+  LogicalResult commonRes = cpAsyncBulkTensorCommonVerifier(
+      getCoordinates().size(), isIm2Col, 0, getLoc());
+
+  LogicalResult overrideAddrRes = CpAsyncBulkTensorOverrideAddrCommonVerifier(
+      getCoordinates(), getTensorSize(), getLowerStride(), getUpperStride(),
+      isTile, getLoc());
+
+  if (mode == TMAStoreMode::TILE_SCATTER4 && getCoordinates().size() != 5)
+    overrideAddrRes = emitError("Mode tile scatter4 expects 5 coordinates");
+
+  return failed(commonRes) || failed(overrideAddrRes) ? failure() : success();
+}
+
 LogicalResult CpAsyncOp::verify() {
   if (getModifier() != LoadCacheModifierKind::CG &&
       getModifier() != LoadCacheModifierKind::CA)
@@ -245,6 +289,20 @@ LogicalResult CpAsyncBulkTensorReduceOp::verify() {
   return success();
 }
 
+LogicalResult CpAsyncBulkTensorReduceOverrideAddrOp::verify() {
+  bool isIm2Col = getMode() == TMAStoreMode::IM2COL;
+  bool isTile = getMode() == TMAStoreMode::TILE;
+
+  LogicalResult commonRes = cpAsyncBulkTensorCommonVerifier(
+      getCoordinates().size(), isIm2Col, 0, getLoc());
+
+  LogicalResult overrideAddrRes = CpAsyncBulkTensorOverrideAddrCommonVerifier(
+      getCoordinates(), getTensorSize(), getLowerStride(), getUpperStride(),
+      isTile, getLoc());
+
+  return failed(commonRes) || failed(overrideAddrRes) ? failure() : success();
+}
+
 LogicalResult CpAsyncBulkGlobalToSharedClusterOp::verify() {
   bool isSharedCTA = isPtrInSharedCTASpace(getDstMem());
   if (isSharedCTA && getMulticastMask())
@@ -4414,6 +4472,78 @@ CpAsyncBulkTensorSharedCTAToGlobalOp::getIntrinsicIDAndArgs(
   return {id, std::move(args)};
 }
 
+NVVM::IDArgPair
+CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp::getIntrinsicIDAndArgs(
+    Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
+  auto thisOp =
+      cast<NVVM::CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp>(op);
+
+  llvm::SmallVector<llvm::Value *> args;
+  args.push_back(mt.lookupValue(thisOp.getSrcMem()));
+  args.push_back(mt.lookupValue(thisOp.getTmaDesc()));
+  args.push_back(mt.lookupValue(thisOp.getOverrideAdrr()));
+  for (Value v : thisOp.getTensorSize())
+    args.push_back(mt.lookupValue(v));
+  for (Value v : thisOp.getLowerStride())
+    args.push_back(mt.lookupValue(v));
+  if (thisOp.getUpperStride())
+    args.push_back(mt.lookupValue(thisOp.getUpperStride()));
+  for (Value v : thisOp.getCoordinates())
+    args.push_back(mt.lookupValue(v));
+
+  mlir::Value cacheHint = thisOp.getL2CacheHint();
+  const bool hasCacheHint = static_cast<bool>(cacheHint);
+  args.push_back(hasCacheHint ? mt.lookupValue(cacheHint)
+                              : builder.getInt64(0));
+  args.push_back(builder.getInt1(hasCacheHint));
+
+  using namespace llvm::Intrinsic;
+  const unsigned NI = not_intrinsic;
+  // clang-format off
+  // override_addr variants, indexed [mode][dim].
+  static constexpr ID IDTable[][6] = {
+      {NI, nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_1d,
+       nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_2d,
+       nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_3d,
+       nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_4d,
+       nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_5d},
+      {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_override_addr_3d,
+       nvvm_cp_async_bulk_tensor_s2g_im2col_override_addr_4d,
+       nvvm_cp_async_bulk_tensor_s2g_im2col_override_addr_5d},
+      {NI, NI, NI, NI, NI,
+       nvvm_cp_async_bulk_tensor_s2g_tile_scatter4_override_addr_2d},
+      {NI, NI, NI, nvvm_cp_async_bulk_tensor_s2g_im2col_w_override_addr_3d,
+       nvvm_cp_async_bulk_tensor_s2g_im2col_w_override_addr_4d,
+       nvvm_cp_async_bulk_tensor_s2g_im2col_w_override_addr_5d}};
+
+  // Tile-only override_addr_dim (1D) / override_addr_dim_stride (2D-5D)
+  // variants, indexed [dim].
+  static constexpr ID dimStrideIDTable[] = {
+      NI, nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_1d,
+      nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_2d,
+      nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_3d,
+      nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_4d,
+      nvvm_cp_async_bulk_tensor_s2g_tile_override_addr_dim_stride_5d};
+  // clang-format on
+
+  size_t mode = static_cast<size_t>(thisOp.getMode());
+  size_t dim = thisOp.getCoordinates().size();
+  bool isDimStride = !thisOp.getTensorSize().empty();
+
+  assert(mode < std::size(IDTable) &&
+         "Invalid mode for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
+  assert(dim < std::size(IDTable[mode]) &&
+         "Invalid dim for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
+  assert(dim < std::size(dimStrideIDTable) &&
+         "Invalid dim for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
+
+  ID intrinsicID = isDimStride ? dimStrideIDTable[dim] : IDTable[mode][dim];
+  assert(
+      intrinsicID != NI &&
+      "Invalid intrinsic for CpAsyncBulkTensorSharedCTAToGlobalOverrideAddrOp");
+  return {intrinsicID, std::move(args)};
+}
+
 NVVM::IDArgPair CpAsyncBulkTensorReduceOp::getIntrinsicIDAndArgs(
     Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
   auto thisOp = cast<NVVM::CpAsyncBulkTensorReduceOp>(op);
@@ -4460,6 +4590,76 @@ NVVM::IDArgPair CpAsyncBulkTensorReduceOp::getIntrinsicIDAndArgs(
   return {intrinsicID, std::move(args)};
 }
 
+NVVM::IDArgPair CpAsyncBulkTensorReduceOverrideAddrOp::getIntrinsicIDAndArgs(
+    Operation &op, LLVM::ModuleTranslation &mt, llvm::IRBuilderBase &builder) {
+  auto thisOp = cast<NVVM::CpAsyncBulkTensorReduceOverrideAddrOp>(op);
+
+  llvm::SmallVector<llvm::Value *> args;
+  args.push_back(mt.lookupValue(thisOp.getSrcMem()));
+  args.push_back(mt.lookupValue(thisOp.getTmaDesc()));
+  args.push_back(mt.lookupValue(thisOp.getOverrideAdrr()));
+
+  for (Value v : thisOp.getTensorSize())
+    args.push_back(mt.lookupValue(v));
+  for (Value v : thisOp.getLowerStride())
+    args.push_back(mt.lookupValue(v));
+  if (thisOp.getUpperStride())
+    args.push_back(mt.lookupValue(thisOp.getUpperStride()));
+  for (Value v : thisOp.getCoordinates())
+    args.push_back(mt.lookupValue(v));
+
+  mlir::Value cacheHint = thisOp.getL2CacheHint();
+  const bool hasCacheHint = static_cast<bool>(cacheHint);
+  args.push_back(hasCacheHint ? mt.lookupValue(cacheHint)
+                              : builder.getInt64(0));
+  args.push_back(builder.getInt32(static_cast<uint32_t>(thisOp.getRedKind())));
+  args.push_back(builder.getInt1(hasCacheHint));
+
+  using namespace llvm::Intrinsic;
+  const unsigned NI = not_intrinsic;
+  // clang-format off
+// override_addr variants, indexed [mode][dim].
+static constexpr ID IDTable[][6] = {
+    {NI, nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_1d,
+     nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_2d,
+     nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_3d,
+     nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_4d,
+     nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_5d},
+    {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_override_addr_3d,
+     nvvm_cp_async_bulk_tensor_reduce_im2col_override_addr_4d,
+     nvvm_cp_async_bulk_tensor_reduce_im2col_override_addr_5d},
+    {NI, NI, NI, NI, NI, NI}, // scatter4 not supported for reduce
+    {NI, NI, NI, nvvm_cp_async_bulk_tensor_reduce_im2col_w_override_addr_3d,
+     nvvm_cp_async_bulk_tensor_reduce_im2col_w_override_addr_4d,
+     nvvm_cp_async_bulk_tensor_reduce_im2col_w_override_addr_5d}};
+
+// Tile-only override_addr_dim (1D) / override_addr_dim_stride (2D-5D)
+// variants, indexed [dim].
+static constexpr ID dimStrideIDTable[] = {
+    NI, nvvm_cp_async_bulk_tensor_reduce_tile_override_addr_dim_1...
[truncated]

``````````

</details>


https://github.com/llvm/llvm-project/pull/216481


More information about the Mlir-commits mailing list