[clang] [llvm] [clang][NVPTX] Add overloaded fmul intrinsics (PR #224546)
via llvm-commits
llvm-commits at lists.llvm.org
Fri Sep 18 00:49:37 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-llvm-ir
Author: Srinivasa Ravi (Wolfram70)
<details>
<summary>Changes</summary>
Follows https://github.com/llvm/llvm-project/pull/221681.
This change adds the following overloaded `fmul` intrinsics with NVPTX codegen:
- `llvm.nvvm.fmul`
- `llvm.nvvm.fmul.ftz`
- `llvm.nvvm.fmul.sat`
- `llvm.nvvm.fmul.ftz.sat`
The rounding mode is passed in as an `i32` immediate operand. Auto-upgrades the older non-overloaded intrinsics to the new ones, and updates clang builtins and CIR codegen to lower to the new intrinsics.
In the interest of completion, this also:
- Adds intrinsics support for lowering some multiplications that were omitted earlier (`f16/f16x2` without saturation, `bf16/bf16x2` multiplications, and `f32` with saturation), and support for the `f32x2` type.
- Adds tests for constant folding of these intrinsics with the newly supported scalar types.
PTX Spec Reference:
https://docs.nvidia.com/cuda/developer-preview/13.4/parallel-thread-execution/index.html#floating-point-instructions-mul
Assisted-by: Claude Opus 5
---
Patch is 137.25 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/224546.diff
25 Files Affected:
- (modified) clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp (+73-36)
- (modified) clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp (+71-35)
- (modified) clang/test/CIR/CodeGenCUDA/builtins-nvvm-math.cu (+24)
- (modified) clang/test/CodeGen/builtins-nvptx.c (+4-4)
- (modified) llvm/docs/NVPTXUsage.md (+46-10)
- (modified) llvm/include/llvm/IR/IntrinsicsNVVM.td (+13-17)
- (modified) llvm/include/llvm/IR/NVVMIntrinsicUtils.h (+23-32)
- (modified) llvm/lib/Analysis/ConstantFolding.cpp (+16-50)
- (modified) llvm/lib/IR/AutoUpgrade.cpp (+35-18)
- (modified) llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp (+25-17)
- (modified) llvm/lib/Target/NVPTX/NVPTXInstrInfo.td (+10)
- (modified) llvm/lib/Target/NVPTX/NVPTXIntrinsics.td (+36-30)
- (modified) llvm/test/Assembler/auto_upgrade_nvvm_intrinsics.ll (+24)
- (added) llvm/test/CodeGen/NVPTX/bf16-mul.ll (+33)
- (removed) llvm/test/CodeGen/NVPTX/f16-mul-sat.ll (-63)
- (added) llvm/test/CodeGen/NVPTX/f16-mul.ll (+123)
- (modified) llvm/test/CodeGen/NVPTX/fp-arith-sat.ll (+33)
- (added) llvm/test/CodeGen/NVPTX/fp-mul-f32x2.ll (+125)
- (added) llvm/test/CodeGen/NVPTX/fp-mul-invalid.ll (+47)
- (added) llvm/test/CodeGen/NVPTX/fp-mul.ll (+58)
- (modified) llvm/test/Transforms/InstCombine/NVPTX/nvvm-intrins.ll (+9-9)
- (modified) llvm/test/Transforms/InstSimplify/const-fold-nvvm-mul.ll (+712-120)
- (added) llvm/test/Verifier/NVPTX/fmul.ll (+16)
- (modified) llvm/test/Verifier/intrinsic-bad-arg-type1.ll (+2-2)
- (modified) llvm/unittests/IR/IntrinsicsTest.cpp (+1-1)
``````````diff
diff --git a/clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp b/clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp
index b81faf65414c1..31bd0e820d7f3 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp
@@ -69,11 +69,11 @@ static mlir::Value emitUnaryNVVMIntrinsic(CIRGenFunction &cgf,
.getResult();
}
-/// Emit a CIR LLVMIntrinsicCallOp for an NVVM fadd intrinsic, which takes the
-/// rounding mode as a trailing operand.
-static mlir::Value emitNVVMFAdd(CIRGenFunction &cgf, const CallExpr *expr,
- llvm::StringRef intrinsicName,
- llvm::APFloat::roundingMode rm) {
+/// Emit a CIR LLVMIntrinsicCallOp for an NVVM fadd/fmul intrinsic, which takes
+/// the rounding mode as a trailing operand.
+static mlir::Value emitNVVMFPArith(CIRGenFunction &cgf, const CallExpr *expr,
+ llvm::StringRef intrinsicName,
+ llvm::APFloat::roundingMode rm) {
auto &builder = cgf.getBuilder();
mlir::Location loc = cgf.getLoc(expr->getExprLoc());
mlir::Value lhs = cgf.emitScalarExpr(expr->getArg(0));
@@ -735,59 +735,96 @@ CIRGenFunction::emitNVPTXBuiltinExpr(unsigned builtinId, const CallExpr *expr) {
return emitUnaryNVVMIntrinsic(*this, expr, "nvvm.ex2.approx.ftz");
case NVPTX::BI__nvvm_add_rn_f:
case NVPTX::BI__nvvm_add_rn_d:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd",
- llvm::APFloat::rmNearestTiesToEven);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd",
+ llvm::APFloat::rmNearestTiesToEven);
case NVPTX::BI__nvvm_add_rz_f:
case NVPTX::BI__nvvm_add_rz_d:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd", llvm::APFloat::rmTowardZero);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd",
+ llvm::APFloat::rmTowardZero);
case NVPTX::BI__nvvm_add_rm_f:
case NVPTX::BI__nvvm_add_rm_d:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd",
- llvm::APFloat::rmTowardNegative);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd",
+ llvm::APFloat::rmTowardNegative);
case NVPTX::BI__nvvm_add_rp_f:
case NVPTX::BI__nvvm_add_rp_d:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd",
- llvm::APFloat::rmTowardPositive);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd",
+ llvm::APFloat::rmTowardPositive);
case NVPTX::BI__nvvm_add_rn_ftz_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz",
- llvm::APFloat::rmNearestTiesToEven);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.ftz",
+ llvm::APFloat::rmNearestTiesToEven);
case NVPTX::BI__nvvm_add_rz_ftz_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz",
- llvm::APFloat::rmTowardZero);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.ftz",
+ llvm::APFloat::rmTowardZero);
case NVPTX::BI__nvvm_add_rm_ftz_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz",
- llvm::APFloat::rmTowardNegative);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.ftz",
+ llvm::APFloat::rmTowardNegative);
case NVPTX::BI__nvvm_add_rp_ftz_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz",
- llvm::APFloat::rmTowardPositive);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.ftz",
+ llvm::APFloat::rmTowardPositive);
case NVPTX::BI__nvvm_add_rn_sat_f:
case NVPTX::BI__nvvm_add_rn_sat_f16:
case NVPTX::BI__nvvm_add_rn_sat_v2f16:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat",
- llvm::APFloat::rmNearestTiesToEven);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.sat",
+ llvm::APFloat::rmNearestTiesToEven);
case NVPTX::BI__nvvm_add_rz_sat_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat",
- llvm::APFloat::rmTowardZero);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.sat",
+ llvm::APFloat::rmTowardZero);
case NVPTX::BI__nvvm_add_rm_sat_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat",
- llvm::APFloat::rmTowardNegative);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.sat",
+ llvm::APFloat::rmTowardNegative);
case NVPTX::BI__nvvm_add_rp_sat_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat",
- llvm::APFloat::rmTowardPositive);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.sat",
+ llvm::APFloat::rmTowardPositive);
case NVPTX::BI__nvvm_add_rn_ftz_sat_f:
case NVPTX::BI__nvvm_add_rn_ftz_sat_f16:
case NVPTX::BI__nvvm_add_rn_ftz_sat_v2f16:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz.sat",
- llvm::APFloat::rmNearestTiesToEven);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.ftz.sat",
+ llvm::APFloat::rmNearestTiesToEven);
case NVPTX::BI__nvvm_add_rz_ftz_sat_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz.sat",
- llvm::APFloat::rmTowardZero);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.ftz.sat",
+ llvm::APFloat::rmTowardZero);
case NVPTX::BI__nvvm_add_rm_ftz_sat_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz.sat",
- llvm::APFloat::rmTowardNegative);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.ftz.sat",
+ llvm::APFloat::rmTowardNegative);
case NVPTX::BI__nvvm_add_rp_ftz_sat_f:
- return emitNVVMFAdd(*this, expr, "nvvm.fadd.ftz.sat",
- llvm::APFloat::rmTowardPositive);
+ return emitNVVMFPArith(*this, expr, "nvvm.fadd.ftz.sat",
+ llvm::APFloat::rmTowardPositive);
+ case NVPTX::BI__nvvm_mul_rn_f:
+ case NVPTX::BI__nvvm_mul_rn_d:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul",
+ llvm::APFloat::rmNearestTiesToEven);
+ case NVPTX::BI__nvvm_mul_rz_f:
+ case NVPTX::BI__nvvm_mul_rz_d:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul",
+ llvm::APFloat::rmTowardZero);
+ case NVPTX::BI__nvvm_mul_rm_f:
+ case NVPTX::BI__nvvm_mul_rm_d:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul",
+ llvm::APFloat::rmTowardNegative);
+ case NVPTX::BI__nvvm_mul_rp_f:
+ case NVPTX::BI__nvvm_mul_rp_d:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul",
+ llvm::APFloat::rmTowardPositive);
+ case NVPTX::BI__nvvm_mul_rn_ftz_f:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul.ftz",
+ llvm::APFloat::rmNearestTiesToEven);
+ case NVPTX::BI__nvvm_mul_rz_ftz_f:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul.ftz",
+ llvm::APFloat::rmTowardZero);
+ case NVPTX::BI__nvvm_mul_rm_ftz_f:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul.ftz",
+ llvm::APFloat::rmTowardNegative);
+ case NVPTX::BI__nvvm_mul_rp_ftz_f:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul.ftz",
+ llvm::APFloat::rmTowardPositive);
+ case NVPTX::BI__nvvm_mul_rn_sat_f16:
+ case NVPTX::BI__nvvm_mul_rn_sat_v2f16:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul.sat",
+ llvm::APFloat::rmNearestTiesToEven);
+ case NVPTX::BI__nvvm_mul_rn_ftz_sat_f16:
+ case NVPTX::BI__nvvm_mul_rn_ftz_sat_v2f16:
+ return emitNVVMFPArith(*this, expr, "nvvm.fmul.ftz.sat",
+ llvm::APFloat::rmNearestTiesToEven);
case NVPTX::BI__nvvm_ldg_h:
case NVPTX::BI__nvvm_ldg_h2:
cgm.errorNYI(expr->getSourceRange(),
diff --git a/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp b/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp
index 76f6757326eca..830d2e2bb8f43 100644
--- a/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp
+++ b/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp
@@ -433,9 +433,9 @@ static Value *MakeFMAOOB(unsigned IntrinsicID, llvm::Type *Ty,
CGF.EmitScalarExpr(E->getArg(2))});
}
-static Value *MakeFAdd(unsigned IntrinsicID, APFloat::roundingMode RM,
- unsigned BuiltinID, const CallExpr *E,
- CodeGenFunction &CGF) {
+static Value *MakeFPArith(unsigned IntrinsicID, APFloat::roundingMode RM,
+ unsigned BuiltinID, const CallExpr *E,
+ CodeGenFunction &CGF) {
llvm::Type *Ty = CGF.ConvertType(E->getType());
return MakeHalfType(CGF.CGM.getIntrinsic(IntrinsicID, Ty), BuiltinID, E, CGF,
{CGF.Builder.getInt32(static_cast<int>(RM))});
@@ -1241,60 +1241,96 @@ Value *CodeGenFunction::EmitNVPTXBuiltinExpr(unsigned BuiltinID,
EmitScalarExpr(E->getArg(0)));
case NVPTX::BI__nvvm_add_rn_f:
case NVPTX::BI__nvvm_add_rn_d:
- return MakeFAdd(Intrinsic::nvvm_fadd, APFloat::rmNearestTiesToEven,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd, APFloat::rmNearestTiesToEven,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rz_f:
case NVPTX::BI__nvvm_add_rz_d:
- return MakeFAdd(Intrinsic::nvvm_fadd, APFloat::rmTowardZero, BuiltinID, E,
- *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd, APFloat::rmTowardZero, BuiltinID,
+ E, *this);
case NVPTX::BI__nvvm_add_rm_f:
case NVPTX::BI__nvvm_add_rm_d:
- return MakeFAdd(Intrinsic::nvvm_fadd, APFloat::rmTowardNegative, BuiltinID,
- E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd, APFloat::rmTowardNegative,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rp_f:
case NVPTX::BI__nvvm_add_rp_d:
- return MakeFAdd(Intrinsic::nvvm_fadd, APFloat::rmTowardPositive, BuiltinID,
- E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd, APFloat::rmTowardPositive,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rn_ftz_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_ftz, APFloat::rmNearestTiesToEven,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_ftz, APFloat::rmNearestTiesToEven,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rz_ftz_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_ftz, APFloat::rmTowardZero, BuiltinID,
- E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_ftz, APFloat::rmTowardZero,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rm_ftz_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_ftz, APFloat::rmTowardNegative,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_ftz, APFloat::rmTowardNegative,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rp_ftz_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_ftz, APFloat::rmTowardPositive,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_ftz, APFloat::rmTowardPositive,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rn_sat_f:
case NVPTX::BI__nvvm_add_rn_sat_f16:
case NVPTX::BI__nvvm_add_rn_sat_v2f16:
- return MakeFAdd(Intrinsic::nvvm_fadd_sat, APFloat::rmNearestTiesToEven,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_sat, APFloat::rmNearestTiesToEven,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rz_sat_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_sat, APFloat::rmTowardZero, BuiltinID,
- E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_sat, APFloat::rmTowardZero,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rm_sat_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_sat, APFloat::rmTowardNegative,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_sat, APFloat::rmTowardNegative,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rp_sat_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_sat, APFloat::rmTowardPositive,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_sat, APFloat::rmTowardPositive,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rn_ftz_sat_f:
case NVPTX::BI__nvvm_add_rn_ftz_sat_f16:
case NVPTX::BI__nvvm_add_rn_ftz_sat_v2f16:
- return MakeFAdd(Intrinsic::nvvm_fadd_ftz_sat, APFloat::rmNearestTiesToEven,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_ftz_sat,
+ APFloat::rmNearestTiesToEven, BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rz_ftz_sat_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_ftz_sat, APFloat::rmTowardZero,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_ftz_sat, APFloat::rmTowardZero,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rm_ftz_sat_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_ftz_sat, APFloat::rmTowardNegative,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_ftz_sat, APFloat::rmTowardNegative,
+ BuiltinID, E, *this);
case NVPTX::BI__nvvm_add_rp_ftz_sat_f:
- return MakeFAdd(Intrinsic::nvvm_fadd_ftz_sat, APFloat::rmTowardPositive,
- BuiltinID, E, *this);
+ return MakeFPArith(Intrinsic::nvvm_fadd_ftz_sat, APFloat::rmTowardPositive,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_mul_rn_f:
+ case NVPTX::BI__nvvm_mul_rn_d:
+ return MakeFPArith(Intrinsic::nvvm_fmul, APFloat::rmNearestTiesToEven,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_mul_rz_f:
+ case NVPTX::BI__nvvm_mul_rz_d:
+ return MakeFPArith(Intrinsic::nvvm_fmul, APFloat::rmTowardZero, BuiltinID,
+ E, *this);
+ case NVPTX::BI__nvvm_mul_rm_f:
+ case NVPTX::BI__nvvm_mul_rm_d:
+ return MakeFPArith(Intrinsic::nvvm_fmul, APFloat::rmTowardNegative,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_mul_rp_f:
+ case NVPTX::BI__nvvm_mul_rp_d:
+ return MakeFPArith(Intrinsic::nvvm_fmul, APFloat::rmTowardPositive,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_mul_rn_ftz_f:
+ return MakeFPArith(Intrinsic::nvvm_fmul_ftz, APFloat::rmNearestTiesToEven,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_mul_rz_ftz_f:
+ return MakeFPArith(Intrinsic::nvvm_fmul_ftz, APFloat::rmTowardZero,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_mul_rm_ftz_f:
+ return MakeFPArith(Intrinsic::nvvm_fmul_ftz, APFloat::rmTowardNegative,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_mul_rp_ftz_f:
+ return MakeFPArith(Intrinsic::nvvm_fmul_ftz, APFloat::rmTowardPositive,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_mul_rn_sat_f16:
+ case NVPTX::BI__nvvm_mul_rn_sat_v2f16:
+ return MakeFPArith(Intrinsic::nvvm_fmul_sat, APFloat::rmNearestTiesToEven,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_mul_rn_ftz_sat_f16:
+ case NVPTX::BI__nvvm_mul_rn_ftz_sat_v2f16:
+ return MakeFPArith(Intrinsic::nvvm_fmul_ftz_sat,
+ APFloat::rmNearestTiesToEven, BuiltinID, E, *this);
case NVPTX::BI__nvvm_ldg_h:
case NVPTX::BI__nvvm_ldg_h2:
return MakeLdg(*this, E);
diff --git a/clang/test/CIR/CodeGenCUDA/builtins-nvvm-math.cu b/clang/test/CIR/CodeGenCUDA/builtins-nvvm-math.cu
index 69df1d376f7f7..9660598f680d7 100644
--- a/clang/test/CIR/CodeGenCUDA/builtins-nvvm-math.cu
+++ b/clang/test/CIR/CodeGenCUDA/builtins-nvvm-math.cu
@@ -87,3 +87,27 @@ __device__ double test_add_rz_d(double x, double y) {
__device__ float test_add_rm_ftz_sat_f(float x, float y) {
return __nvvm_add_rm_ftz_sat_f(x, y);
}
+
+// CIR-LABEL: @_Z13test_mul_rn_fff
+// CIR: cir.call_llvm_intrinsic "nvvm.fmul" {{.*}} : (!cir.float, !cir.float, !s32i) -> !cir.float
+// LLVM-LABEL: @_Z13test_mul_rn_fff
+// LLVM: call {{.*}}float @llvm.nvvm.fmul.f32(float {{.*}}, float {{.*}}, /* rnd=rn */ i32 1)
+__device__ float test_mul_rn_f(float x, float y) {
+ return __nvvm_mul_rn_f(x, y);
+}
+
+// CIR-LABEL: @_Z13test_mul_rz_ddd
+// CIR: cir.call_llvm_intrinsic "nvvm.fmul" {{.*}} : (!cir.double, !cir.double, !s32i) -> !cir.double
+// LLVM-LABEL: @_Z13test_mul_rz_ddd
+// LLVM: call {{.*}}double @llvm.nvvm.fmul.f64(double {{.*}}, double {{.*}}, /* rnd=rz */ i32 0)
+__device__ double test_mul_rz_d(double x, double y) {
+ return __nvvm_mul_rz_d(x, y);
+}
+
+// CIR-LABEL: @_Z17test_mul_rp_ftz_fff
+// CIR: cir.call_llvm_intrinsic "nvvm.fmul.ftz" {{.*}} : (!cir.float, !cir.float, !s32i) -> !cir.float
+// LLVM-LABEL: @_Z17test_mul_rp_ftz_fff
+// LLVM: call {{.*}}float @llvm.nvvm.fmul.ftz.f32(float {{.*}}, float {{.*}}, /* rnd=rp */ i32 2)
+__device__ float test_mul_rp_ftz_f(float x, float y) {
+ return __nvvm_mul_rp_ftz_f(x, y);
+}
diff --git a/clang/test/CodeGen/builtins-nvptx.c b/clang/test/CodeGen/builtins-nvptx.c
index 607095bce481c..a6c1a9c3e7272 100644
--- a/clang/test/CodeGen/builtins-nvptx.c
+++ b/clang/test/CodeGen/builtins-nvptx.c
@@ -1700,13 +1700,13 @@ __device__ void nvvm_add_mul_f16_sat() {
// CHECK: call <2 x half> @llvm.nvvm.fadd.ftz.sat.v2f16({{.*}}i32 1)
__nvvm_add_rn_ftz_sat_v2f16(F16X2, F16X2_2);
- // CHECK: call half @llvm.nvvm.mul.rn.sat.f16
+ // CHECK: call half @llvm.nvvm.fmul.sat.f16({{.*}}i32 1)
__nvvm_mul_rn_sat_f16(F16, F16_2);
- // CHECK: call half @llvm.nvvm.mul.rn.ftz.sat.f16
+ // CHECK: call half @llvm.nvvm.fmul.ftz.sat.f16({{.*}}i32 1)
__nvvm_mul_rn_ftz_sat_f16(F16, F16_2);
- // CHECK: call <2 x half> @llvm.nvvm.mul.rn.sat.v2f16
+ // CHECK: call <2 x half> @llvm.nvvm.fmul.sat.v2f16({{.*}}i32 1)
__nvvm_mul_rn_sat_v2f16(F16X2, F16X2_2);
- // CHECK: call <2 x half> @llvm.nvvm.mul.rn.ftz.sat.v2f16
+ // CHECK: call <2 x half> @llvm.nvvm.fmul.ftz.sat.v2f16({{.*}}i32 1)
__nvvm_mul_rn_ftz_sat_v2f16(F16X2, F16X2_2);
// CHECK: ret void
diff --git a/llvm/docs/NVPTXUsage.md b/llvm/docs/NVPTXUsage.md
index 67bf8ea482933..1386046f7b85d 100644
--- a/llvm/docs/NVPTXUsage.md
+++ b/llvm/docs/NVPTXUsage.md
@@ -1403,29 +1403,65 @@ PTX instruction. The supported combinations are:
- None
```
-#### '`llvm.nvvm.mul.*`' Half-precision Intrinsics
+#### '`llvm.nvvm.fmul.*`' Intrinsics
##### Syntax:
-```llvm
-declare half @llvm.nvvm.mul.rn.sat.f16(half %a, half %b)
-declare <2 x half> @llvm.nvvm.mul.rn.sat.v2f16(<2 x half> %a, <2 x half> %b)
+This is an overloaded intrinsic. The '`.ftz`' and '`.sat`' modifiers are
+optional.
-declare half @llvm.nvvm.mul.rn.ftz.sat.f16(half %a, half %b)
-declare <2 x half> @llvm.nvvm.mul.rn.ftz.sat.v2f16(<2 x half> %a, <2 x half> %b)
+```llvm
+declare half @llvm.nvvm.fmul{.ftz}{.sat}.f16(half %a, half %b, i32 immarg %rnd)
+declare <2 x half> @llvm.nvvm.fmul{.ftz}{.sat}.v2f16(<2 x half> %a, <2 x half> %b, i32 immarg %rnd)
+declare bfloat @llvm.nvvm.fmul.bf16(bfloat %a, bfloat %b, i32 immarg %rnd)
+declare <2 x bfloat> @llvm.nvvm.fmul.v2bf16(<2 x bfloat> %a, <2 x bfloat> %b, i32 immarg %rnd)
+declare float @llvm.nvvm.fmul{.ftz}{.sat}.f32(float %a, float %b, i32 immarg %rnd)
+declare <2 x float> @llvm.nvvm.fmul{.ftz}.v2f32(<2 x float> %a, <2 x float> %b, i32 immarg %rnd)
+declare double @llvm.nvvm.fmul.f64(double %a, double %b, i32 immarg %rnd)
```
##### Overview:
-The '`llvm.nvvm.mul.*`' intrinsics perform a multiplication operation with
-the specified rounding mode and modifiers.
+The '`llvm.nvvm.fmul.*`' intrinsics multiply `%a` and `%b` using the rounding
+mode selected by `%rnd` and the modifiers present in the intrinsic name. They
+correspond directly to the `mul` PTX instruction.
##### Semantics:
-The '`.sat`' modifier performs a saturating multiplication where the result is
-clamped to `[0.0, 1.0]` and `NaN` results are flushed to `+0.0f`.
+`%rnd` selects the rounding mode applied to the result, see
+{ref}`fp-rounding-modes`.
+
The '`.ftz`' modifier flushes subnormal inputs and results to sign-preserving
zero.
+The '`.sat`' modifier performs a saturating multiplication where the result is
+cl...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/224546
More information about the llvm-commits
mailing list