[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