[llvm] [NVPTX] Fold vector FMAs in the IR peephole pass (PR #224018)
via llvm-commits
llvm-commits at lists.llvm.org
Wed Sep 16 06:50:30 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-nvptx
Author: Daniel Donenfeld (daniel-donenfeld)
<details>
<summary>Changes</summary>
Enable folding vector FMUL+FADD/FSUB sequences into vector FMA operations. This enables the generation of the fma.f32x2 operation on sm_100 and higher.
---
Full diff: https://github.com/llvm/llvm-project/pull/224018.diff
2 Files Affected:
- (modified) llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp (+4-3)
- (modified) llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll (+81)
``````````diff
diff --git a/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp b/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp
index bd16c7213b1e7..9db516928ff2c 100644
--- a/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp
@@ -10,7 +10,7 @@
// run late in the NVPTX IR pass pipeline just before the instruction selection.
//
// Currently, it implements the following transformation(s):
-// 1. FMA folding (float/double types):
+// 1. FMA folding (float/double types and vectors thereof):
// Transforms FMUL+FADD/FSUB sequences into FMA intrinsics when the
// 'contract' fast-math flag is present. Supported patterns:
// - fadd(fmul(a, b), c) => fma(a, b, c)
@@ -125,8 +125,9 @@ static bool foldFMA(Function &F) {
if (!BI->hasAllowContract())
continue;
- // Only float and double are supported.
- if (!BI->getType()->isFloatTy() && !BI->getType()->isDoubleTy())
+ // Float, double, and vectors thereof are supported.
+ Type *ScalarTy = BI->getType()->getScalarType();
+ if (!ScalarTy->isFloatTy() && !ScalarTy->isDoubleTy())
continue;
if (tryFoldBinaryFMul(BI))
diff --git a/llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll b/llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll
index 6d9ad8d3ad436..6c2cf8e5afd84 100644
--- a/llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll
+++ b/llvm/test/CodeGen/NVPTX/nvptx-fold-fma.ll
@@ -245,3 +245,84 @@ define double @test_fadd_fmul_c_double(double %a, double %b, double %c) {
%add = fadd contract double %mul, %c
ret double %add
}
+
+
+; fadd(fmul(a, b), c) => fma(a, b, c)
+define <2 x float> @test_fadd_fmul_c_v2f32(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
+; CHECK-LABEL: define <2 x float> @test_fadd_fmul_c_v2f32(
+; CHECK-SAME: <2 x float> [[A:%.*]], <2 x float> [[B:%.*]], <2 x float> [[C:%.*]]) {
+; CHECK-NEXT: [[ADD:%.*]] = call contract <2 x float> @llvm.fma.v2f32(<2 x float> [[A]], <2 x float> [[B]], <2 x float> [[C]])
+; CHECK-NEXT: ret <2 x float> [[ADD]]
+;
+ %mul = fmul contract <2 x float> %a, %b
+ %add = fadd contract <2 x float> %mul, %c
+ ret <2 x float> %add
+}
+
+
+; fadd(c, fmul(a, b)) => fma(a, b, c)
+define <2 x float> @test_fadd_c_fmul_v2f32(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
+; CHECK-LABEL: define <2 x float> @test_fadd_c_fmul_v2f32(
+; CHECK-SAME: <2 x float> [[A:%.*]], <2 x float> [[B:%.*]], <2 x float> [[C:%.*]]) {
+; CHECK-NEXT: [[ADD:%.*]] = call contract <2 x float> @llvm.fma.v2f32(<2 x float> [[A]], <2 x float> [[B]], <2 x float> [[C]])
+; CHECK-NEXT: ret <2 x float> [[ADD]]
+;
+ %mul = fmul contract <2 x float> %a, %b
+ %add = fadd contract <2 x float> %c, %mul
+ ret <2 x float> %add
+}
+
+
+; fsub(fmul(a, b), c) => fma(a, b, fneg(c))
+define <2 x float> @test_fsub_fmul_c_v2f32(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
+; CHECK-LABEL: define <2 x float> @test_fsub_fmul_c_v2f32(
+; CHECK-SAME: <2 x float> [[A:%.*]], <2 x float> [[B:%.*]], <2 x float> [[C:%.*]]) {
+; CHECK-NEXT: [[TMP1:%.*]] = fneg contract <2 x float> [[C]]
+; CHECK-NEXT: [[SUB:%.*]] = call contract <2 x float> @llvm.fma.v2f32(<2 x float> [[A]], <2 x float> [[B]], <2 x float> [[TMP1]])
+; CHECK-NEXT: ret <2 x float> [[SUB]]
+;
+ %mul = fmul contract <2 x float> %a, %b
+ %sub = fsub contract <2 x float> %mul, %c
+ ret <2 x float> %sub
+}
+
+
+; fsub(c, fmul(a, b)) => fma(fneg(a), b, c)
+define <2 x float> @test_fsub_c_fmul_v2f32(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
+; CHECK-LABEL: define <2 x float> @test_fsub_c_fmul_v2f32(
+; CHECK-SAME: <2 x float> [[A:%.*]], <2 x float> [[B:%.*]], <2 x float> [[C:%.*]]) {
+; CHECK-NEXT: [[TMP1:%.*]] = fneg contract <2 x float> [[A]]
+; CHECK-NEXT: [[SUB:%.*]] = call contract <2 x float> @llvm.fma.v2f32(<2 x float> [[TMP1]], <2 x float> [[B]], <2 x float> [[C]])
+; CHECK-NEXT: ret <2 x float> [[SUB]]
+;
+ %mul = fmul contract <2 x float> %a, %b
+ %sub = fsub contract <2 x float> %c, %mul
+ ret <2 x float> %sub
+}
+
+
+; fadd(fmul(a, b), c) => fma(a, b, c)
+define <2 x double> @test_fadd_fmul_c_v2f64(<2 x double> %a, <2 x double> %b, <2 x double> %c) {
+; CHECK-LABEL: define <2 x double> @test_fadd_fmul_c_v2f64(
+; CHECK-SAME: <2 x double> [[A:%.*]], <2 x double> [[B:%.*]], <2 x double> [[C:%.*]]) {
+; CHECK-NEXT: [[ADD:%.*]] = call contract <2 x double> @llvm.fma.v2f64(<2 x double> [[A]], <2 x double> [[B]], <2 x double> [[C]])
+; CHECK-NEXT: ret <2 x double> [[ADD]]
+;
+ %mul = fmul contract <2 x double> %a, %b
+ %add = fadd contract <2 x double> %mul, %c
+ ret <2 x double> %add
+}
+
+
+; fsub(fmul(a, b), c) => fma(a, b, fneg(c))
+define <2 x double> @test_fsub_fmul_c_v2f64(<2 x double> %a, <2 x double> %b, <2 x double> %c) {
+; CHECK-LABEL: define <2 x double> @test_fsub_fmul_c_v2f64(
+; CHECK-SAME: <2 x double> [[A:%.*]], <2 x double> [[B:%.*]], <2 x double> [[C:%.*]]) {
+; CHECK-NEXT: [[TMP1:%.*]] = fneg contract <2 x double> [[C]]
+; CHECK-NEXT: [[SUB:%.*]] = call contract <2 x double> @llvm.fma.v2f64(<2 x double> [[A]], <2 x double> [[B]], <2 x double> [[TMP1]])
+; CHECK-NEXT: ret <2 x double> [[SUB]]
+;
+ %mul = fmul contract <2 x double> %a, %b
+ %sub = fsub contract <2 x double> %mul, %c
+ ret <2 x double> %sub
+}
``````````
</details>
https://github.com/llvm/llvm-project/pull/224018
More information about the llvm-commits
mailing list