[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