[llvm] [mlir] [LLVM][NVPTX] Add Rubin extensions for G2S Tensor intrinsics (PR #220029)
Rajat Bajpai via llvm-commits
llvm-commits at lists.llvm.org
Sat Sep 5 03:40:03 PDT 2026
================
@@ -5941,37 +5999,51 @@ void llvm::UpgradeIntrinsicCall(CallBase *CI, Function *NewFn) {
CI->eraseFromParent();
return;
}
- case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_3d:
- case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_4d:
- case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_im2col_5d:
- case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_1d:
- case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_2d:
- case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_3d:
- case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_4d:
- case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_tile_5d: {
+ // clang-format off
+#define G2S_CLUSTER_CASE(ID_SUFFIX, NAME) \
+ case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_##ID_SUFFIX:
+ NVVM_TMA_G2S_MODES(G2S_CLUSTER_CASE)
+#undef G2S_CLUSTER_CASE
+ {
SmallVector<Value *, 16> Args(CI->args());
-
- // Create AddrSpaceCast to shared_cluster if needed.
- // This handles case (1) in shouldUpgradeNVPTXTMAG2SIntrinsics().
unsigned AS = CI->getArgOperand(0)->getType()->getPointerAddressSpace();
if (AS == NVPTXAS::ADDRESS_SPACE_SHARED)
Args[0] = Builder.CreateAddrSpaceCast(
Args[0], Builder.getPtrTy(NVPTXAS::ADDRESS_SPACE_SHARED_CLUSTER));
- // Attach the flag argument for cta_group, with a
- // default value of 0. This handles case (2) in
- // shouldUpgradeNVPTXTMAG2SIntrinsics().
- size_t NumArgs = CI->arg_size();
- Value *FlagArg = CI->getArgOperand(NumArgs - 3);
- if (!FlagArg->getType()->isIntegerTy(1))
- Args.push_back(ConstantInt::get(Builder.getInt32Ty(), 0));
+ // Append the missing trailing arguments with default values (cta_group,
+ // validate_pattern).
+ while (Args.size() < NewFn->getFunctionType()->getNumParams())
+ Args.push_back(Builder.getInt32(0));
+
+ NewCall = Builder.CreateCall(NewFn, Args);
+ NewCall->takeName(CI);
+ CI->replaceAllUsesWith(NewCall);
+ CI->eraseFromParent();
+ return;
+ }
+
+#define G2S_CTA_CASE(ID_SUFFIX, NAME) \
+ case Intrinsic::nvvm_cp_async_bulk_tensor_g2s_cta_##ID_SUFFIX:
+ NVVM_TMA_G2S_MODES(G2S_CTA_CASE)
+#undef G2S_CTA_CASE
+ {
+ SmallVector<Value *, 16> Args(CI->args());
+ // Append the missing trailing validate_pattern argument with default
+ // value 0.
+ assert(Args.size() + 1 == NewFn->getFunctionType()->getNumParams() &&
+ "expected only the trailing validate_pattern to be missing");
+ Args.push_back(Builder.getInt32(0));
NewCall = Builder.CreateCall(NewFn, Args);
NewCall->takeName(CI);
CI->replaceAllUsesWith(NewCall);
CI->eraseFromParent();
return;
}
+// clang-format on
----------------
rajatbajpai wrote:
Sure, done.
https://github.com/llvm/llvm-project/pull/220029
More information about the llvm-commits
mailing list