[Mlir-commits] [mlir] [MLIR][NVVM] Add S2G and Reduce override NVVM Dialect ops (PR #216481)
Rajat Bajpai
llvmlistbot at llvm.org
Sat Aug 22 06:46:43 PDT 2026
================
@@ -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");
----------------
rajatbajpai wrote:
makes sense, merged.
https://github.com/llvm/llvm-project/pull/216481
More information about the Mlir-commits
mailing list