[Mlir-commits] [clang] [llvm] [mlir] [clang][NVPTX][MLIR][NVVM] Add overloaded fadd intrinsics (PR #217336)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Wed Aug 19 06:13:25 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-nvptx
Author: Srinivasa Ravi (Wolfram70)
<details>
<summary>Changes</summary>
This change adds the following overloaded `fadd` intrinsics with
NVPTX codegen:
- `llvm.nvvm.fadd`
- `llvm.nvvm.fadd.ftz`
- `llvm.nvvm.fadd.sat`
- `llvm.nvvm.fadd.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, CIR codegen, and MLIR NVVM ops
to lower to the new intrinsics.
In the interest of completion, this also adds:
- Intrinsics support for lowering some half-precision additions
that were omitted earlier (`f16/f16x2` without saturation, and
`bf16/bf16x2` additions), and also adds support for the `f32x2`
type.
- 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-add
---
Patch is 218.58 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/217336.diff
37 Files Affected:
- (modified) clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp (+72)
- (modified) clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp (+67-1)
- (modified) clang/test/CIR/CodeGenCUDA/builtins-nvvm-math.cu (+24)
- (modified) clang/test/CodeGen/builtins-nvptx.c (+13-13)
- (modified) llvm/docs/NVPTXUsage.md (+73-10)
- (modified) llvm/include/llvm/IR/IntrinsicsNVVM.td (+12-22)
- (modified) llvm/include/llvm/IR/NVVMIntrinsicUtils.h (+34-30)
- (modified) llvm/lib/Analysis/ConstantFolding.cpp (+33-43)
- (modified) llvm/lib/IR/AutoUpgrade.cpp (+45)
- (modified) llvm/lib/IR/NVVMIntrinsicUtils.cpp (+5)
- (modified) llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp (+85-21)
- (modified) llvm/lib/Target/NVPTX/NVPTXIntrinsics.td (+91-43)
- (modified) llvm/test/Assembler/auto_upgrade_nvvm_intrinsics.ll (+24)
- (added) llvm/test/CodeGen/NVPTX/bf16-add.ll (+33)
- (added) llvm/test/CodeGen/NVPTX/bf16-sub.ll (+35)
- (removed) llvm/test/CodeGen/NVPTX/f16-add-sat.ll (-63)
- (added) llvm/test/CodeGen/NVPTX/f16-add.ll (+123)
- (removed) llvm/test/CodeGen/NVPTX/f16-sub-sat.ll (-69)
- (added) llvm/test/CodeGen/NVPTX/f16-sub.ll (+133)
- (added) llvm/test/CodeGen/NVPTX/fp-add-f32x2.ll (+60)
- (added) llvm/test/CodeGen/NVPTX/fp-add-invalid.ll (+47)
- (modified) llvm/test/CodeGen/NVPTX/fp-arith-sat.ll (+16-16)
- (added) llvm/test/CodeGen/NVPTX/fp-fold-sub-f32x2.ll (+64)
- (modified) llvm/test/CodeGen/NVPTX/fp-fold-sub.ll (+10-10)
- (modified) llvm/test/CodeGen/NVPTX/mixed-precision-fp.ll (+32-32)
- (modified) llvm/test/Transforms/InstCombine/NVPTX/nvvm-intrins.ll (+9-9)
- (modified) llvm/test/Transforms/InstSimplify/const-fold-nvvm-add.ll (+658-114)
- (added) llvm/test/Verifier/NVPTX/fadd.ll (+16)
- (modified) llvm/unittests/IR/IntrinsicsTest.cpp (+1-1)
- (modified) mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp (-8)
- (modified) mlir/lib/Target/LLVMIR/Dialect/NVVM/NVVMToLLVMIRTranslation.cpp (+26-77)
- (modified) mlir/test/Target/LLVMIR/nvvm/addf/addf.mlir (+34-32)
- (modified) mlir/test/Target/LLVMIR/nvvm/addf/addf_invalid.mlir (-10)
- (modified) mlir/test/Target/LLVMIR/nvvm/addf/addf_vector.mlir (+97-158)
- (modified) mlir/test/Target/LLVMIR/nvvm/subf/subf.mlir (+35-32)
- (modified) mlir/test/Target/LLVMIR/nvvm/subf/subf_invalid.mlir (-10)
- (modified) mlir/test/Target/LLVMIR/nvvm/subf/subf_vector.mlir (+111-171)
``````````diff
diff --git a/clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp b/clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp
index 2220639876695..68eb0cf1206f1 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenBuiltinNVPTX.cpp
@@ -49,6 +49,23 @@ 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) {
+ auto &builder = cgf.getBuilder();
+ mlir::Location loc = cgf.getLoc(expr->getExprLoc());
+ mlir::Value lhs = cgf.emitScalarExpr(expr->getArg(0));
+ mlir::Value rhs = cgf.emitScalarExpr(expr->getArg(1));
+ mlir::Value rnd =
+ builder.getConstInt(loc, builder.getSInt32Ty(), static_cast<int>(rm));
+ return cir::LLVMIntrinsicCallOp::create(builder, loc,
+ builder.getStringAttr(intrinsicName),
+ lhs.getType(), {lhs, rhs, rnd})
+ .getResult();
+}
+
static mlir::Value makeScopedAtomicRMW(CIRGenFunction &cgf,
const CallExpr *expr,
cir::AtomicFetchKind kind,
@@ -794,6 +811,61 @@ CIRGenFunction::emitNVPTXBuiltinExpr(unsigned builtinId, const CallExpr *expr) {
return emitUnaryNVVMIntrinsic(*this, expr, "nvvm.ex2.approx");
case NVPTX::BI__nvvm_ex2_approx_ftz_f:
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);
+ case NVPTX::BI__nvvm_add_rz_f:
+ case NVPTX::BI__nvvm_add_rz_d:
+ return emitNVVMFAdd(*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);
+ case NVPTX::BI__nvvm_add_rp_f:
+ case NVPTX::BI__nvvm_add_rp_d:
+ return emitNVVMFAdd(*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);
+ case NVPTX::BI__nvvm_add_rz_ftz_f:
+ return emitNVVMFAdd(*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);
+ case NVPTX::BI__nvvm_add_rp_ftz_f:
+ return emitNVVMFAdd(*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);
+ case NVPTX::BI__nvvm_add_rz_sat_f:
+ return emitNVVMFAdd(*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);
+ case NVPTX::BI__nvvm_add_rp_sat_f:
+ return emitNVVMFAdd(*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.sat.ftz",
+ llvm::APFloat::rmNearestTiesToEven);
+ case NVPTX::BI__nvvm_add_rz_ftz_sat_f:
+ return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat.ftz",
+ llvm::APFloat::rmTowardZero);
+ case NVPTX::BI__nvvm_add_rm_ftz_sat_f:
+ return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat.ftz",
+ llvm::APFloat::rmTowardNegative);
+ case NVPTX::BI__nvvm_add_rp_ftz_sat_f:
+ return emitNVVMFAdd(*this, expr, "nvvm.fadd.sat.ftz",
+ llvm::APFloat::rmTowardPositive);
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 64fdae9d8934d..a0577a349885b 100644
--- a/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp
+++ b/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp
@@ -395,7 +395,8 @@ static Value *MakeCpAsync(unsigned IntrinsicID, unsigned IntrinsicIDS,
}
static Value *MakeHalfType(Function *Intrinsic, unsigned BuiltinID,
- const CallExpr *E, CodeGenFunction &CGF) {
+ const CallExpr *E, CodeGenFunction &CGF,
+ ArrayRef<Value *> TrailingArgs = {}) {
SmallVector<Value *, 16> Args;
auto *FTy = Intrinsic->getFunctionType();
unsigned ICEArguments = 0;
@@ -410,6 +411,7 @@ static Value *MakeHalfType(Function *Intrinsic, unsigned BuiltinID,
ArgValue = CGF.Builder.CreateBitCast(ArgValue, PTy);
Args.push_back(ArgValue);
}
+ Args.append(TrailingArgs.begin(), TrailingArgs.end());
return CGF.Builder.CreateCall(Intrinsic, Args);
}
@@ -427,6 +429,14 @@ 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) {
+ llvm::Type *Ty = CGF.ConvertType(E->getType());
+ return MakeHalfType(CGF.CGM.getIntrinsic(IntrinsicID, Ty), BuiltinID, E, CGF,
+ {CGF.Builder.getInt32(static_cast<int>(RM))});
+}
+
} // namespace
Value *CodeGenFunction::EmitNVPTXBuiltinExpr(unsigned BuiltinID,
@@ -1134,6 +1144,62 @@ Value *CodeGenFunction::EmitNVPTXBuiltinExpr(unsigned BuiltinID,
case NVPTX::BI__nvvm_ex2_approx_ftz_f:
return Builder.CreateUnaryIntrinsic(Intrinsic::nvvm_ex2_approx_ftz,
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);
+ case NVPTX::BI__nvvm_add_rz_f:
+ case NVPTX::BI__nvvm_add_rz_d:
+ return MakeFAdd(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);
+ case NVPTX::BI__nvvm_add_rp_f:
+ case NVPTX::BI__nvvm_add_rp_d:
+ return MakeFAdd(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);
+ case NVPTX::BI__nvvm_add_rz_ftz_f:
+ return MakeFAdd(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);
+ case NVPTX::BI__nvvm_add_rp_ftz_f:
+ return MakeFAdd(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);
+ case NVPTX::BI__nvvm_add_rz_sat_f:
+ return MakeFAdd(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);
+ case NVPTX::BI__nvvm_add_rp_sat_f:
+ return MakeFAdd(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_sat_ftz, APFloat::rmNearestTiesToEven,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_add_rz_ftz_sat_f:
+ return MakeFAdd(Intrinsic::nvvm_fadd_sat_ftz, APFloat::rmTowardZero,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_add_rm_ftz_sat_f:
+ return MakeFAdd(Intrinsic::nvvm_fadd_sat_ftz, APFloat::rmTowardNegative,
+ BuiltinID, E, *this);
+ case NVPTX::BI__nvvm_add_rp_ftz_sat_f:
+ return MakeFAdd(Intrinsic::nvvm_fadd_sat_ftz, APFloat::rmTowardPositive,
+ 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 a2f0edb6d073b..c2f4d19322cfe 100644
--- a/clang/test/CIR/CodeGenCUDA/builtins-nvvm-math.cu
+++ b/clang/test/CIR/CodeGenCUDA/builtins-nvvm-math.cu
@@ -63,3 +63,27 @@ __device__ double test_ex2_approx_d(double x) {
__device__ float test_ex2_approx_ftz_f(float x) {
return __nvvm_ex2_approx_ftz_f(x);
}
+
+// CIR-LABEL: @_Z13test_add_rn_fff
+// CIR: cir.call_llvm_intrinsic "nvvm.fadd" {{.*}} : (!cir.float, !cir.float, !s32i) -> !cir.float
+// LLVM-LABEL: @_Z13test_add_rn_fff
+// LLVM: call {{.*}}float @llvm.nvvm.fadd.f32(float {{.*}}, float {{.*}}, /* rnd=rn */ i32 1)
+__device__ float test_add_rn_f(float x, float y) {
+ return __nvvm_add_rn_f(x, y);
+}
+
+// CIR-LABEL: @_Z13test_add_rz_ddd
+// CIR: cir.call_llvm_intrinsic "nvvm.fadd" {{.*}} : (!cir.double, !cir.double, !s32i) -> !cir.double
+// LLVM-LABEL: @_Z13test_add_rz_ddd
+// LLVM: call {{.*}}double @llvm.nvvm.fadd.f64(double {{.*}}, double {{.*}}, /* rnd=rz */ i32 0)
+__device__ double test_add_rz_d(double x, double y) {
+ return __nvvm_add_rz_d(x, y);
+}
+
+// CIR-LABEL: @_Z21test_add_rm_ftz_sat_fff
+// CIR: cir.call_llvm_intrinsic "nvvm.fadd.sat.ftz" {{.*}} : (!cir.float, !cir.float, !s32i) -> !cir.float
+// LLVM-LABEL: @_Z21test_add_rm_ftz_sat_fff
+// LLVM: call {{.*}}float @llvm.nvvm.fadd.sat.ftz.f32(float {{.*}}, float {{.*}}, /* rnd=rm */ i32 3)
+__device__ float test_add_rm_ftz_sat_f(float x, float y) {
+ return __nvvm_add_rm_ftz_sat_f(x, y);
+}
diff --git a/clang/test/CodeGen/builtins-nvptx.c b/clang/test/CodeGen/builtins-nvptx.c
index 87be7b46aad8e..53fc5aa8d6d8c 100644
--- a/clang/test/CodeGen/builtins-nvptx.c
+++ b/clang/test/CodeGen/builtins-nvptx.c
@@ -245,7 +245,7 @@ __device__ void nvvm_math(float f1, float f2, double d1, double d2) {
float t3 = __nvvm_sqrt_rn_f(f1);
// CHECK: call float @llvm.nvvm.rcp.rn.f
float t4 = __nvvm_rcp_rn_f(f2);
-// CHECK: call float @llvm.nvvm.add.rn.f
+// CHECK: call float @llvm.nvvm.fadd.f32({{.*}}i32 1)
float t5 = __nvvm_add_rn_f(f1, f2);
// CHECK: call double @llvm.nvvm.fmax.d
@@ -1548,21 +1548,21 @@ __device__ void nvvm_min_max_sm86() {
// CHECK-LABEL: nvvm_add_fma_f32_sat
__device__ void nvvm_add_fma_f32_sat() {
- // CHECK: call float @llvm.nvvm.add.rn.sat.f
+ // CHECK: call float @llvm.nvvm.fadd.sat.f32({{.*}}i32 1)
__nvvm_add_rn_sat_f(1.0f, 2.0f);
- // CHECK: call float @llvm.nvvm.add.rn.ftz.sat.f
+ // CHECK: call float @llvm.nvvm.fadd.sat.ftz.f32({{.*}}i32 1)
__nvvm_add_rn_ftz_sat_f(1.0f, 2.0f);
- // CHECK: call float @llvm.nvvm.add.rz.sat.f
+ // CHECK: call float @llvm.nvvm.fadd.sat.f32({{.*}}i32 0)
__nvvm_add_rz_sat_f(1.0f, 2.0f);
- // CHECK: call float @llvm.nvvm.add.rz.ftz.sat.f
+ // CHECK: call float @llvm.nvvm.fadd.sat.ftz.f32({{.*}}i32 0)
__nvvm_add_rz_ftz_sat_f(1.0f, 2.0f);
- // CHECK: call float @llvm.nvvm.add.rm.sat.f
+ // CHECK: call float @llvm.nvvm.fadd.sat.f32({{.*}}i32 3)
__nvvm_add_rm_sat_f(1.0f, 2.0f);
- // CHECK: call float @llvm.nvvm.add.rm.ftz.sat.f
+ // CHECK: call float @llvm.nvvm.fadd.sat.ftz.f32({{.*}}i32 3)
__nvvm_add_rm_ftz_sat_f(1.0f, 2.0f);
- // CHECK: call float @llvm.nvvm.add.rp.sat.f
+ // CHECK: call float @llvm.nvvm.fadd.sat.f32({{.*}}i32 2)
__nvvm_add_rp_sat_f(1.0f, 2.0f);
- // CHECK: call float @llvm.nvvm.add.rp.ftz.sat.f
+ // CHECK: call float @llvm.nvvm.fadd.sat.ftz.f32({{.*}}i32 2)
__nvvm_add_rp_ftz_sat_f(1.0f, 2.0f);
// CHECK: call float @llvm.nvvm.fma.rn.sat.f
@@ -1592,13 +1592,13 @@ __device__ void nvvm_add_fma_f32_sat() {
// CHECK-LABEL: nvvm_add_mul_f16_sat
__device__ void nvvm_add_mul_f16_sat() {
- // CHECK: call half @llvm.nvvm.add.rn.sat.f16
+ // CHECK: call half @llvm.nvvm.fadd.sat.f16({{.*}}i32 1)
__nvvm_add_rn_sat_f16(F16, F16_2);
- // CHECK: call half @llvm.nvvm.add.rn.ftz.sat.f16
+ // CHECK: call half @llvm.nvvm.fadd.sat.ftz.f16({{.*}}i32 1)
__nvvm_add_rn_ftz_sat_f16(F16, F16_2);
- // CHECK: call <2 x half> @llvm.nvvm.add.rn.sat.v2f16
+ // CHECK: call <2 x half> @llvm.nvvm.fadd.sat.v2f16({{.*}}i32 1)
__nvvm_add_rn_sat_v2f16(F16X2, F16X2_2);
- // CHECK: call <2 x half> @llvm.nvvm.add.rn.ftz.sat.v2f16
+ // CHECK: call <2 x half> @llvm.nvvm.fadd.sat.ftz.v2f16({{.*}}i32 1)
__nvvm_add_rn_ftz_sat_v2f16(F16X2, F16X2_2);
// CHECK: call half @llvm.nvvm.mul.rn.sat.f16
diff --git a/llvm/docs/NVPTXUsage.md b/llvm/docs/NVPTXUsage.md
index 8924a44e43a8e..9212a6894e28d 100644
--- a/llvm/docs/NVPTXUsage.md
+++ b/llvm/docs/NVPTXUsage.md
@@ -1174,6 +1174,33 @@ For more information, see [PTX ISA](https://docs.nvidia.com/cuda/parallel-thread
### Arithmetic Intrinsics
+Some of these intrinsics take the rounding mode as an `i32` immediate operand
+instead of encoding it in the intrinsic name. The accepted values match the
+`llvm::RoundingMode` enumeration and are described in the following table:
+
+(fp-rounding-modes)=
+
+```{list-table} Floating-Point Rounding Modes
+:widths: 15 15 70
+:header-rows: 1
+
+ * - Value
+ - Rounding Mode
+ - Description
+ * - 0
+ - `rz`
+ - Round towards zero
+ * - 1
+ - `rn`
+ - Round to nearest, with ties to even
+ * - 2
+ - `rp`
+ - Round towards positive infinity
+ * - 3
+ - `rm`
+ - Round towards negative infinity
+```
+
#### '`llvm.nvvm.fabs.*`' Intrinsic
##### Syntax:
@@ -1273,29 +1300,65 @@ used in the '`llvm.nvvm.idp4a.[us].u`' variants, while sign-extension is used
with '`llvm.nvvm.idp4a.[us].s`' variants. The dot product of these 4-element
vectors is added to `%c` to produce the return.
-#### '`llvm.nvvm.add.*`' Half-precision Intrinsics
+#### '`llvm.nvvm.fadd.*`' Intrinsics
##### Syntax:
-```llvm
-declare half @llvm.nvvm.add.rn.sat.f16(half %a, half %b)
-declare <2 x half> @llvm.nvvm.add.rn.sat.v2f16(<2 x half> %a, <2 x half> %b)
+This is an overloaded intrinsic. The '`.sat`' and '`.ftz`' modifiers are
+optional.
-declare half @llvm.nvvm.add.rn.ftz.sat.f16(half %a, half %b)
-declare <2 x half> @llvm.nvvm.add.rn.ftz.sat.v2f16(<2 x half> %a, <2 x half> %b)
+```llvm
+declare half @llvm.nvvm.fadd{.sat}{.ftz}.f16(half %a, half %b, i32 immarg %rnd)
+declare <2 x half> @llvm.nvvm.fadd{.sat}{.ftz}.v2f16(<2 x half> %a, <2 x half> %b, i32 immarg %rnd)
+declare bfloat @llvm.nvvm.fadd.bf16(bfloat %a, bfloat %b, i32 immarg %rnd)
+declare <2 x bfloat> @llvm.nvvm.fadd.v2bf16(<2 x bfloat> %a, <2 x bfloat> %b, i32 immarg %rnd)
+declare float @llvm.nvvm.fadd{.sat}{.ftz}.f32(float %a, float %b, i32 immarg %rnd)
+declare <2 x float> @llvm.nvvm.fadd{.ftz}.v2f32(<2 x float> %a, <2 x float> %b, i32 immarg %rnd)
+declare double @llvm.nvvm.fadd.f64(double %a, double %b, i32 immarg %rnd)
```
##### Overview:
-The '`llvm.nvvm.add.*`' intrinsics perform an addition operation with the
-specified rounding mode and modifiers.
+The '`llvm.nvvm.fadd.*`' intrinsics add `%a` and `%b` using the rounding mode
+selected by `%rnd` and the modifiers present in the intrinsic name. They
+correspond directly to the `add` PTX instruction.
##### Semantics:
-The '`.sat`' modifier performs a saturating addition 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 addition where the result is
+clamped to `[0.0, 1.0]` and `NaN` results are flushed to `+0.0f`.
+
+Not every combination of operand type, rounding mode and modifier maps to a
+PTX instruction. The supported combinations are:
+
+```{list-table}
+:widths: 25 25 25 25
+:header-rows: 1
+
+ * - Operand Type
+ - Rounding Modes
+ - Modifiers
+ * - `half`, `<2 x half>`
+ - `rn`
+ - `.sat`, `.ftz`
+ * - `bfloat`, `<2 x bfloat>`
+ - `rn`
+ - None
+ * - `float`
+ - `rn`, `rz`, `rp`, `rm`
+ - `.sat`, `.ftz`
+ * - `<2 x float>`
+ - `rn`, `rz`, `rp`, `rm`
+ - `.ftz`
+ * - `double`
+ - `rn`, `rz`, `rp`, `rm`
+ - None
+```
#### '`llvm.nvvm.mul.*`' Half-precision Intrinsics
diff --git a/llvm/include/llvm/IR/IntrinsicsNVVM.td b/llvm/include/llvm/IR/IntrinsicsNVVM.td
index 4357ad367d269..2f26ce3765dbc 100644
--- a/llvm/include/llvm/IR/IntrinsicsNVVM.td
+++ b/llvm/include/llvm/IR/IntrinsicsNVVM.td
@@ -1676,31 +1676,21 @@ let TargetPrefix = "nvvm" in {
}
//
- // Add
+ // FAdd
//
let IntrProperties = [IntrNoMem, IntrSpeculatable, Commutative,
- IntrNoCreateUndefOrPoison] in {
- foreach rnd = ["_rn", "_rz", "_rm", "_rp"] in {
- foreach ftz = ["", "_ftz"] in {
- foreach sat = ["", "_sat"] in {
- def int_nvvm_add # rnd # ftz # sat # _f : NVVMBuiltin,
- DefaultAttrsIntrinsic<[llvm_float_ty], [llvm_float_ty, llvm_float_ty]>;
- } // sat
- } // ftz
- def int_nvvm_add # rnd # _d : NVVMBuiltin,
- DefaultAttrsIntrinsic<[llvm_double_ty], [llvm_double_ty, llvm_double_ty]>;
- }
-
- foreach ftz = ["", "_ftz"] in {
- def int_nvvm_add_rn # ftz # _sat_f16 : NVVMBuiltin,
- DefaultAttrsIntrinsic<[llvm_half_ty], [llvm_half_ty, llvm_half_ty]>;
-
- def int_nvvm_add_rn # ftz # _sat_v2f16 : NVVMBuiltin,
- DefaultAttrsIntrinsic<[llvm_v2f16_ty], [llvm_v2f16_ty, llvm_v2f16_ty]>;
-
- } // ftz
- }
+ IntrNoCreateUndefOrPoison, ImmArg<ArgIndex<2>>,
+ Range<ArgIndex<2>, 0, 4>,
+ ArgInfo<ArgIndex<2>,
+ [ArgName<"rnd">,
+ ImmArgPrinter<"printFAddRoundingMode">]>] in
+ foreach sat = ["", "_sat"] in
+ foreach ftz = ["", "_ftz"] in
+ def int_nvvm_fadd # sat # ftz :
+ DefaultAttrsIntrinsic<[llvm_anyfloat_ty],
+ [LLVMMatchType<0>, LLVMMatchType<0>,
+ llvm_i32_ty]>;
//
// Dot Product
diff --git a/llvm/include/llvm/IR/NVVMIntrinsicUtils.h b/llvm/include/llvm/IR/NVVMIntrinsicUtils.h
index b38ebc3e2b309..b52a277e8de60 100644
--- a/llvm/include/llvm/IR/NVVMIntrinsicUtils.h
+++ b/llvm/include/llvm/IR/NVVMIntrinsicUtils.h
@@ -160,6 +160,7 @@ LLVM_ABI void printTensormapSwizzleAtomicity(raw_ostream &OS,
const Constant *ImmArgVal);
LLVM_ABI void printTensormapFillMode(raw_ostream &OS,
const Constant *ImmArgVal);
+LLVM_ABI void printFAddRoundingMode(raw_ostream &OS, const Constant *ImmArgVal);
inline bool FPToIntegerIntrinsicShouldFTZ(Intrinsic::ID IntrinsicID) {
switch (IntrinsicID) {
@@ -594,45 +595,48 @@ inline DenormalMode GetNVVMDenormMode(bool ShouldFTZ) {
inline bool FAddShouldFTZ(Intrinsic::ID IntrinsicID) {
switch (IntrinsicID) {
- case Intrinsic::nvvm_add_rm_ftz_f:
- case Intrinsic::nvvm_add_rn_ftz_f:
- case Intrinsic::nvvm_add_rp_ftz_f:
- case Intrinsic::nvvm_add_rz_ftz_f:
+ case Intrinsic::nvvm_fadd_ftz:
+ ...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/217336
More information about the Mlir-commits
mailing list