[flang-commits] [flang] [flang][Lower] Reassociate sums containing pure calls (PR #217286)
Tom Eccles via flang-commits
flang-commits at lists.llvm.org
Thu Aug 20 03:36:27 PDT 2026
https://github.com/tblah updated https://github.com/llvm/llvm-project/pull/217286
>From c49a1b21b2ce3730179b5b55f05d5626bf0f533a Mon Sep 17 00:00:00 2001
From: Tom Eccles <tom.eccles at arm.com>
Date: Mon, 10 Aug 2026 17:57:09 +0100
Subject: [PATCH 1/2] [flang][Lower] Reassociate sums containing pure calls
Part 5/6 of generalisations requested in #207377
Allow pure procedure references to remain opaque terms while splitting
an additive expression. Continue rejecting impure calls and expressions
that reference volatile or asynchronous objects.
There are no known effects on benchmark scores as a result of this
patch.
Assisted-by: Codex
---
flang/lib/Evaluate/tools.cpp | 12 ++-
.../split-sum-expression-tree-lowering.f90 | 87 +++++++++++++++++++
2 files changed, 98 insertions(+), 1 deletion(-)
diff --git a/flang/lib/Evaluate/tools.cpp b/flang/lib/Evaluate/tools.cpp
index 4924da89a2318..e10a8b37babf6 100644
--- a/flang/lib/Evaluate/tools.cpp
+++ b/flang/lib/Evaluate/tools.cpp
@@ -1354,6 +1354,16 @@ struct HasProcedureRefHelper : public AnyTraverse<HasProcedureRefHelper> {
bool operator()(const ProcedureRef &) const { return true; }
};
+struct HasImpureProcedureRefHelper
+ : public AnyTraverse<HasImpureProcedureRefHelper> {
+ using Base = AnyTraverse<HasImpureProcedureRefHelper>;
+ HasImpureProcedureRefHelper() : Base{*this} {}
+ using Base::operator();
+ bool operator()(const ProcedureRef &ref) const {
+ return ref.proc().IsPure() ? Base::operator()(ref) : true;
+ }
+};
+
struct HasVolatileOrAsynchronousSymbolHelper
: public AnyTraverse<HasVolatileOrAsynchronousSymbolHelper> {
using Base = AnyTraverse<HasVolatileOrAsynchronousSymbolHelper>;
@@ -1541,7 +1551,7 @@ static std::optional<Expr<SomeType>> tryBuildSplitSumExpressionTree(
bool CanBuildSplitSumExpressionTree(
const Expr<SomeType> &lhs, const Expr<SomeType> &rhs) {
return rhs.Rank() == 0 && lhs.Rank() == 0 && !HasVectorSubscript(rhs) &&
- !HasVectorSubscript(lhs) && !HasProcedureRef(rhs) &&
+ !HasVectorSubscript(lhs) && !HasImpureProcedureRefHelper{}(rhs) &&
!HasProcedureRef(lhs) && !HasVolatileOrAsynchronousSymbol(rhs) &&
!HasVolatileOrAsynchronousSymbol(lhs);
}
diff --git a/flang/test/Lower/split-sum-expression-tree-lowering.f90 b/flang/test/Lower/split-sum-expression-tree-lowering.f90
index 186fa183cce07..f681f4ebe45fe 100644
--- a/flang/test/Lower/split-sum-expression-tree-lowering.f90
+++ b/flang/test/Lower/split-sum-expression-tree-lowering.f90
@@ -750,6 +750,93 @@ subroutine guard_call(x,a,b,c,d,e)
! NO-REWRITE: %[[RES:.*]] = arith.addf %[[XFOOBC]], %[[DE]]
! NO-REWRITE: hlfir.assign %[[RES]] to %[[X]]#0
+! Default: (((x + sqrt((a+b)+c)) + d*e) + f*g)
+! Rewritten: ((d*e + f*g) + (x + sqrt((a+b)+c)))
+! The pure call is an opaque outer term. Its additive argument retains source
+! order in this stage.
+subroutine eligible_pure_call(x,a,b,c,d,e,f,g)
+ real(8) :: x,a,b,c,d,e,f,g
+ x = x + sqrt(a+b+c) + d*e + f*g
+end
+
+! SPLIT-LABEL: func.func @_QPeligible_pure_call
+! SPLIT: %[[DV:.*]] = fir.load
+! SPLIT: %[[EV:.*]] = fir.load
+! SPLIT: %[[DE:.*]] = arith.mulf %[[DV]], %[[EV]]
+! SPLIT: %[[FV:.*]] = fir.load
+! SPLIT: %[[GV:.*]] = fir.load
+! SPLIT: %[[FG:.*]] = arith.mulf %[[FV]], %[[GV]]
+! SPLIT: %[[TAIL:.*]] = arith.addf %[[DE]], %[[FG]]
+! SPLIT: %[[XV:.*]] = fir.load
+! SPLIT: %[[AV:.*]] = fir.load
+! SPLIT: %[[BV:.*]] = fir.load
+! SPLIT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! SPLIT: %[[CV:.*]] = fir.load
+! SPLIT: %[[ABC:.*]] = arith.addf %[[AB]], %[[CV]]
+! SPLIT: %[[CALL:.*]] = math.sqrt %[[ABC]]
+! SPLIT: %[[HEAD:.*]] = arith.addf %[[XV]], %[[CALL]]
+! SPLIT: %[[RES:.*]] = arith.addf %[[TAIL]], %[[HEAD]]
+! SPLIT: hlfir.assign %[[RES]]
+
+! DEFAULT-LABEL: func.func @_QPeligible_pure_call
+! DEFAULT: %[[XV:.*]] = fir.load
+! DEFAULT: %[[AV:.*]] = fir.load
+! DEFAULT: %[[BV:.*]] = fir.load
+! DEFAULT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! DEFAULT: %[[CV:.*]] = fir.load
+! DEFAULT: %[[ABC:.*]] = arith.addf %[[AB]], %[[CV]]
+! DEFAULT: %[[CALL:.*]] = math.sqrt %[[ABC]]
+! DEFAULT: %[[HEAD:.*]] = arith.addf %[[XV]], %[[CALL]]
+! DEFAULT: %[[DV:.*]] = fir.load
+! DEFAULT: %[[EV:.*]] = fir.load
+! DEFAULT: %[[DE:.*]] = arith.mulf %[[DV]], %[[EV]]
+! DEFAULT: %[[HEAD_DE:.*]] = arith.addf %[[HEAD]], %[[DE]]
+! DEFAULT: %[[FV:.*]] = fir.load
+! DEFAULT: %[[GV:.*]] = fir.load
+! DEFAULT: %[[FG:.*]] = arith.mulf %[[FV]], %[[GV]]
+! DEFAULT: %[[RES:.*]] = arith.addf %[[HEAD_DE]], %[[FG]]
+! DEFAULT: hlfir.assign %[[RES]]
+
+subroutine guard_pure_call_volatile_arg(x,v,a,b,c,d)
+ real(8) :: x,a,b,c,d
+ real(8), volatile :: v
+ x = x + sqrt(v) + a*b + c*d
+end
+
+! NO-REWRITE-LABEL: func.func @_QPguard_pure_call_volatile_arg
+! NO-REWRITE: %[[XV:.*]] = fir.load
+! NO-REWRITE: %[[CALL:.*]] = math.sqrt
+! NO-REWRITE: %[[HEAD:.*]] = arith.addf %[[XV]], %[[CALL]]
+! NO-REWRITE: %[[AV:.*]] = fir.load
+! NO-REWRITE: %[[BV:.*]] = fir.load
+! NO-REWRITE: %[[AB:.*]] = arith.mulf %[[AV]], %[[BV]]
+! NO-REWRITE: %[[HEAD_AB:.*]] = arith.addf %[[HEAD]], %[[AB]]
+! NO-REWRITE: %[[CV:.*]] = fir.load
+! NO-REWRITE: %[[DV:.*]] = fir.load
+! NO-REWRITE: %[[CD:.*]] = arith.mulf %[[CV]], %[[DV]]
+! NO-REWRITE: %[[RES:.*]] = arith.addf %[[HEAD_AB]], %[[CD]]
+! NO-REWRITE: hlfir.assign %[[RES]]
+
+subroutine guard_nested_impure_call(x,a,b,c,d,e)
+ real(8) :: x,a,b,c,d,e,foo
+ x = x + sqrt(foo(a)) + b*c + d*e
+end
+
+! NO-REWRITE-LABEL: func.func @_QPguard_nested_impure_call
+! NO-REWRITE: %[[XV:.*]] = fir.load
+! NO-REWRITE: %[[IMPURE:.*]] = fir.call @_QPfoo
+! NO-REWRITE: %[[PURE:.*]] = math.sqrt %[[IMPURE]]
+! NO-REWRITE: %[[HEAD:.*]] = arith.addf %[[XV]], %[[PURE]]
+! NO-REWRITE: %[[BV:.*]] = fir.load
+! NO-REWRITE: %[[CV:.*]] = fir.load
+! NO-REWRITE: %[[BC:.*]] = arith.mulf %[[BV]], %[[CV]]
+! NO-REWRITE: %[[HEAD_BC:.*]] = arith.addf %[[HEAD]], %[[BC]]
+! NO-REWRITE: %[[DV:.*]] = fir.load
+! NO-REWRITE: %[[EV:.*]] = fir.load
+! NO-REWRITE: %[[DE:.*]] = arith.mulf %[[DV]], %[[EV]]
+! NO-REWRITE: %[[RES:.*]] = arith.addf %[[HEAD_BC]], %[[DE]]
+! NO-REWRITE: hlfir.assign %[[RES]]
+
subroutine guard_array(n,x,a,b,c,d,e,f)
integer :: n
real(8) :: x(n),a(n),b(n),c(n),d(n),e(n),f(n)
>From 85d10f4387e2a6661ebd2418990785d6c3359ecc Mon Sep 17 00:00:00 2001
From: Tom Eccles <tom.eccles at arm.com>
Date: Thu, 20 Aug 2026 11:34:32 +0100
Subject: [PATCH 2/2] Re-use FindImpureCall()
---
flang/include/flang/Evaluate/tools.h | 2 +-
flang/lib/Evaluate/tools.cpp | 14 ++------------
flang/lib/Lower/Bridge.cpp | 4 ++--
3 files changed, 5 insertions(+), 15 deletions(-)
diff --git a/flang/include/flang/Evaluate/tools.h b/flang/include/flang/Evaluate/tools.h
index c877ec5f5705b..cce39602c05ef 100644
--- a/flang/include/flang/Evaluate/tools.h
+++ b/flang/include/flang/Evaluate/tools.h
@@ -1129,7 +1129,7 @@ bool HasVolatileOrAsynchronousSymbol(const Expr<SomeType> &expr);
// Can a scalar real or complex RHS expression in an assignment be rewritten
// as a split sum expression tree?
bool CanBuildSplitSumExpressionTree(
- const Expr<SomeType> &lhs, const Expr<SomeType> &rhs);
+ FoldingContext &, const Expr<SomeType> &lhs, const Expr<SomeType> &rhs);
// Try to rewrite a scalar real or complex sum as a split sum expression tree.
std::optional<Expr<SomeType>> TryBuildSplitSumExpressionTree(
diff --git a/flang/lib/Evaluate/tools.cpp b/flang/lib/Evaluate/tools.cpp
index e10a8b37babf6..71e72c739d760 100644
--- a/flang/lib/Evaluate/tools.cpp
+++ b/flang/lib/Evaluate/tools.cpp
@@ -1354,16 +1354,6 @@ struct HasProcedureRefHelper : public AnyTraverse<HasProcedureRefHelper> {
bool operator()(const ProcedureRef &) const { return true; }
};
-struct HasImpureProcedureRefHelper
- : public AnyTraverse<HasImpureProcedureRefHelper> {
- using Base = AnyTraverse<HasImpureProcedureRefHelper>;
- HasImpureProcedureRefHelper() : Base{*this} {}
- using Base::operator();
- bool operator()(const ProcedureRef &ref) const {
- return ref.proc().IsPure() ? Base::operator()(ref) : true;
- }
-};
-
struct HasVolatileOrAsynchronousSymbolHelper
: public AnyTraverse<HasVolatileOrAsynchronousSymbolHelper> {
using Base = AnyTraverse<HasVolatileOrAsynchronousSymbolHelper>;
@@ -1548,10 +1538,10 @@ static std::optional<Expr<SomeType>> tryBuildSplitSumExpressionTree(
} // namespace
-bool CanBuildSplitSumExpressionTree(
+bool CanBuildSplitSumExpressionTree(FoldingContext &context,
const Expr<SomeType> &lhs, const Expr<SomeType> &rhs) {
return rhs.Rank() == 0 && lhs.Rank() == 0 && !HasVectorSubscript(rhs) &&
- !HasVectorSubscript(lhs) && !HasImpureProcedureRefHelper{}(rhs) &&
+ !HasVectorSubscript(lhs) && !FindImpureCall(context, rhs) &&
!HasProcedureRef(lhs) && !HasVolatileOrAsynchronousSymbol(rhs) &&
!HasVolatileOrAsynchronousSymbol(lhs);
}
diff --git a/flang/lib/Lower/Bridge.cpp b/flang/lib/Lower/Bridge.cpp
index e2db2d20a4b4f..f9becabaa1fe2 100644
--- a/flang/lib/Lower/Bridge.cpp
+++ b/flang/lib/Lower/Bridge.cpp
@@ -5441,8 +5441,8 @@ class FirConverter : public Fortran::lower::AbstractConverter {
const Fortran::lower::SomeExpr *rhsExpr = &assign.rhs;
std::optional<Fortran::lower::SomeExpr> rewritten;
if (bridge.getLoweringOptions().getSplitSumExpressionTree() &&
- Fortran::evaluate::CanBuildSplitSumExpressionTree(assign.lhs,
- assign.rhs)) {
+ Fortran::evaluate::CanBuildSplitSumExpressionTree(
+ getFoldingContext(), assign.lhs, assign.rhs)) {
rewritten =
Fortran::evaluate::TryBuildSplitSumExpressionTree(assign.rhs);
if (rewritten)
More information about the flang-commits
mailing list