[llvm] [InstCombine] Drop one-use check in matchesSquareSum for (a*a) (PR #228020)

via llvm-commits llvm-commits at lists.llvm.org
Thu Oct 1 03:15:55 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-llvm-transforms

Author: konoongg (konoongg)

<details>
<summary>Changes</summary>

The pattern `(a * a) + (((a * 2) + b) * b)` is equivalent to `(a + b) * (a + b)`. The current implementation requires the `a * a` subexpression to have a single use, which prevents the fold when `a * a` is used elsewhere. This check is unnecessary: we can still perform the fold and keep the original `a * a` instruction for other uses.

For example, with an extra use of `a * a`:

  %1 = mul i32 %a, %a
  %2 = mul i32 %a, 2
  %3 = add i32 %2, %b
  %4 = mul i32 %3, %b
  %5 = add i32 %4, %1
  %6 = mul i32 %1, %1

After this change, InstCombine produces:

  %1 = add i32 %a, %b
  %2 = mul i32 %1, %1
  %3 = mul i32 %a, %a
  %4 = mul i32 %3, %3

This reduces the instruction count from 6 to 4, even though `a * a` is now computed both explicitly (for the extra use) and implicitly as part of `(a + b) * (a + b)`.

The same applies to the floating-point (fadd) variant. Both are
updated together.

Tests: llvm/test/Transforms/InstCombine/add.ll llvm/test/Transforms/InstCombine/fadd.ll

---
Full diff: https://github.com/llvm/llvm-project/pull/228020.diff


3 Files Affected:

- (modified) llvm/lib/Transforms/InstCombine/InstCombineAddSub.cpp (+1-1) 
- (modified) llvm/test/Transforms/InstCombine/add.ll (+10-14) 
- (modified) llvm/test/Transforms/InstCombine/fadd.ll (+5-7) 


``````````diff
diff --git a/llvm/lib/Transforms/InstCombine/InstCombineAddSub.cpp b/llvm/lib/Transforms/InstCombine/InstCombineAddSub.cpp
index 317d587ddedd6..bff47f0103914 100644
--- a/llvm/lib/Transforms/InstCombine/InstCombineAddSub.cpp
+++ b/llvm/lib/Transforms/InstCombine/InstCombineAddSub.cpp
@@ -1052,7 +1052,7 @@ static bool matchesSquareSum(BinaryOperator &I, Mul2Rhs M2Rhs, Value *&A,
 
   // (a * a) + (((a * 2) + b) * b)
   if (match(&I, m_c_BinOp(
-                    AddOp, m_OneUse(m_BinOp(MulOp, m_Value(A), m_Deferred(A))),
+                    AddOp, m_BinOp(MulOp, m_Value(A), m_Deferred(A)),
                     m_OneUse(m_c_BinOp(
                         MulOp,
                         m_c_BinOp(AddOp, m_BinOp(Mul2Op, m_Deferred(A), M2Rhs),
diff --git a/llvm/test/Transforms/InstCombine/add.ll b/llvm/test/Transforms/InstCombine/add.ll
index 22472b0114fd7..992cf69f3cba3 100644
--- a/llvm/test/Transforms/InstCombine/add.ll
+++ b/llvm/test/Transforms/InstCombine/add.ll
@@ -4445,13 +4445,11 @@ define i32 @add_reduce_sqr_sum_not_one_use(i32 %a, i32 %b) {
 
 define i32 @add_reduce_sqr_sum_not_one_use2(i32 %a, i32 %b) {
 ; CHECK-LABEL: @add_reduce_sqr_sum_not_one_use2(
-; CHECK-NEXT:    [[A_SQ:%.*]] = mul nsw i32 [[A:%.*]], [[A]]
-; CHECK-NEXT:    [[TWO_A:%.*]] = shl i32 [[A]], 1
-; CHECK-NEXT:    [[TWO_A_PLUS_B:%.*]] = add i32 [[TWO_A]], [[B:%.*]]
-; CHECK-NEXT:    [[MUL:%.*]] = mul i32 [[TWO_A_PLUS_B]], [[B]]
-; CHECK-NEXT:    tail call void @fake_func(i32 [[A_SQ]])
-; CHECK-NEXT:    [[ADD:%.*]] = add i32 [[MUL]], [[A_SQ]]
-; CHECK-NEXT:    ret i32 [[ADD]]
+; CHECK-NEXT:    [[A_SQ:%.*]] = mul nsw i32 [[A:%.*]], [[A:%.*]]
+; CHECK-NEXT:    tail call void @fake_func(i32 [[A_SQ:%.*]])
+; CHECK-NEXT:    [[AB:%.*]] = add i32 [[A:%.*]], [[B:%.*]]
+; CHECK-NEXT:    [[AB_SQ:%.*]] = mul i32 [[AB:%.*]], [[AB:%.*]]
+; CHECK-NEXT:    ret i32 [[AB_SQ:%.*]]
 ;
   %a_sq = mul nsw i32 %a, %a
   %two_a = shl i32 %a, 1
@@ -4484,13 +4482,11 @@ define i32 @add_reduce_sqr_sum_order2_not_one_use(i32 %a, i32 %b) {
 
 define i32 @add_reduce_sqr_sum_order2_not_one_use2(i32 %a, i32 %b) {
 ; CHECK-LABEL: @add_reduce_sqr_sum_order2_not_one_use2(
-; CHECK-NEXT:    [[A_SQ:%.*]] = mul nsw i32 [[A:%.*]], [[A]]
-; CHECK-NEXT:    [[TWOA:%.*]] = shl i32 [[A]], 1
-; CHECK-NEXT:    [[TWOAB1:%.*]] = add i32 [[TWOA]], [[B:%.*]]
-; CHECK-NEXT:    [[TWOAB_B2:%.*]] = mul i32 [[TWOAB1]], [[B]]
-; CHECK-NEXT:    tail call void @fake_func(i32 [[A_SQ]])
-; CHECK-NEXT:    [[AB2:%.*]] = add i32 [[A_SQ]], [[TWOAB_B2]]
-; CHECK-NEXT:    ret i32 [[AB2]]
+; CHECK-NEXT:    [[A_SQ:%.*]] = mul nsw i32 [[A:%.*]], [[A:%.*]]
+; CHECK-NEXT:    tail call void @fake_func(i32 [[A_SQ:%.*]])
+; CHECK-NEXT:    [[AB:%.*]] = add i32 [[A:%.*]], [[B:%.*]]
+; CHECK-NEXT:    [[AB_SQ:%.*]] = mul i32 [[AB:%.*]], [[AB:%.*]]
+; CHECK-NEXT:    ret i32 [[AB_SQ:%.*]]
 ;
   %a_sq = mul nsw i32 %a, %a
   %twoa = mul i32 %a, 2
diff --git a/llvm/test/Transforms/InstCombine/fadd.ll b/llvm/test/Transforms/InstCombine/fadd.ll
index 2311685db410d..c3f5cc742dd28 100644
--- a/llvm/test/Transforms/InstCombine/fadd.ll
+++ b/llvm/test/Transforms/InstCombine/fadd.ll
@@ -829,13 +829,11 @@ define float @fadd_reduce_sqr_sum_varA_not_one_use1(float %a, float %b) {
 
 define float @fadd_reduce_sqr_sum_varA_not_one_use2(float %a, float %b) {
 ; CHECK-LABEL: @fadd_reduce_sqr_sum_varA_not_one_use2(
-; CHECK-NEXT:    [[A_SQ:%.*]] = fmul float [[A:%.*]], [[A]]
-; CHECK-NEXT:    [[TWO_A:%.*]] = fmul float [[A]], 2.000000e+00
-; CHECK-NEXT:    [[TWO_A_PLUS_B:%.*]] = fadd float [[TWO_A]], [[B:%.*]]
-; CHECK-NEXT:    [[MUL:%.*]] = fmul float [[TWO_A_PLUS_B]], [[B]]
-; CHECK-NEXT:    [[ADD:%.*]] = fadd reassoc nsz float [[MUL]], [[A_SQ]]
-; CHECK-NEXT:    tail call void @fake_func(float [[A_SQ]])
-; CHECK-NEXT:    ret float [[ADD]]
+; CHECK-NEXT:    [[A_SQ:%.*]] = fmul float [[A:%.*]], [[A:%.*]]
+; CHECK-NEXT:    [[AB:%.*]] = fadd reassoc nsz float [[A:%.*]], [[B:%.*]]
+; CHECK-NEXT:    [[AB_SQ:%.*]] = fmul reassoc nsz float [[AB:%.*]], [[AB:%.*]]
+; CHECK-NEXT:    tail call void @fake_func(float [[A_SQ:%.*]])
+; CHECK-NEXT:    ret float [[AB_SQ:%.*]]
 ;
   %a_sq = fmul float %a, %a
   %two_a = fmul float %a, 2.0

``````````

</details>


https://github.com/llvm/llvm-project/pull/228020


More information about the llvm-commits mailing list