[llvm-branch-commits] [flang] [flang][Lower] Reassociate nested assignment sums (PR #217287)
Tom Eccles via llvm-branch-commits
llvm-branch-commits at lists.llvm.org
Thu Aug 20 03:54:27 PDT 2026
https://github.com/tblah updated https://github.com/llvm/llvm-project/pull/217287
>From b69782f0435adfc42ebeb44109e655b1472803e3 Mon Sep 17 00:00:00 2001
From: Tom Eccles <tom.eccles at arm.com>
Date: Tue, 11 Aug 2026 10:25:50 +0100
Subject: [PATCH] [flang][Lower] Reassociate nested assignment sums
Part 6/6 of generalisations requested in #207377
Search eligible assignment RHS expressions top-down for scalar REAL
and COMPLEX additive spines. Rewrite a qualifying node once; otherwise
recurse through ordinary operations, conversions, conditional wrappers,
and separate arguments of pure calls. Preserve effect-sensitive
boundaries.
There is no known effect on benchmark results as a result of this patch.
Assisted-by: Codex
---
flang/include/flang/Evaluate/tools.h | 5 +-
flang/lib/Evaluate/tools.cpp | 135 +++++++++++++-
flang/lib/Lower/Bridge.cpp | 2 +-
.../split-sum-expression-tree-lowering.f90 | 168 +++++++++++++++++-
4 files changed, 299 insertions(+), 11 deletions(-)
diff --git a/flang/include/flang/Evaluate/tools.h b/flang/include/flang/Evaluate/tools.h
index cce39602c05ef..8cfc6adfcc941 100644
--- a/flang/include/flang/Evaluate/tools.h
+++ b/flang/include/flang/Evaluate/tools.h
@@ -1131,8 +1131,9 @@ bool HasVolatileOrAsynchronousSymbol(const Expr<SomeType> &expr);
bool CanBuildSplitSumExpressionTree(
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(
+// Try to rewrite eligible scalar real or complex sums within an expression as
+// split sum expression trees.
+std::optional<Expr<SomeType>> TryBuildSplitSumExpressionTrees(
const Expr<SomeType> &expr);
// Utilities for attaching the location of the declaration of a symbol
diff --git a/flang/lib/Evaluate/tools.cpp b/flang/lib/Evaluate/tools.cpp
index 71e72c739d760..23a2f42ca66cd 100644
--- a/flang/lib/Evaluate/tools.cpp
+++ b/flang/lib/Evaluate/tools.cpp
@@ -10,6 +10,7 @@
#include "flang/Common/idioms.h"
#include "flang/Common/type-kinds.h"
#include "flang/Evaluate/characteristics.h"
+#include "flang/Evaluate/rewrite.h"
#include "flang/Evaluate/traverse.h"
#include "flang/Parser/message.h"
#include "flang/Semantics/tools.h"
@@ -1536,6 +1537,132 @@ static std::optional<Expr<SomeType>> tryBuildSplitSumExpressionTree(
return std::nullopt;
}
+template <typename> struct IsExpr : std::false_type {};
+template <typename T> struct IsExpr<Expr<T>> : std::true_type {};
+
+template <typename> struct IsFunctionRef : std::false_type {};
+template <typename T> struct IsFunctionRef<FunctionRef<T>> : std::true_type {};
+
+template <typename> struct IsConditionalExpr : std::false_type {};
+template <typename T>
+struct IsConditionalExpr<ConditionalExpr<T>> : std::true_type {};
+
+class SplitSumExpressionTreeRewriter {
+public:
+ std::optional<Expr<SomeType>> Rewrite(const Expr<SomeType> &expr) {
+ Expr<SomeType> rewritten{rewriteExpr(Expr<SomeType>{expr})};
+ if (changed_)
+ return rewritten;
+ return std::nullopt;
+ }
+
+private:
+ template <typename T>
+ static std::optional<Expr<T>> tryRewriteCurrent(const Expr<T> &) {
+ return std::nullopt;
+ }
+
+ template <common::TypeCategory CAT, int KIND>
+ static std::optional<NumericExpr<CAT, KIND>> tryRewriteCurrent(
+ const NumericExpr<CAT, KIND> &expr) {
+ if constexpr (CAT == common::TypeCategory::Real ||
+ CAT == common::TypeCategory::Complex)
+ return tryBuildSplitSumExpressionTree(expr);
+ return std::nullopt;
+ }
+
+ template <typename D, std::size_t... Is>
+ D rewriteOperation(D &&op, std::index_sequence<Is...>) {
+ return D{rewriteExpr(std::move(op.template operand<Is>()))...};
+ }
+
+ template <typename T, std::size_t... Is>
+ Extremum<T> rewriteOperation(Extremum<T> &&op, std::index_sequence<Is...>) {
+ return Extremum<T>{
+ op.ordering, rewriteExpr(std::move(op.template operand<Is>()))...};
+ }
+
+ template <int KIND, std::size_t... Is>
+ ComplexComponent<KIND> rewriteOperation(
+ ComplexComponent<KIND> &&op, std::index_sequence<Is...>) {
+ return ComplexComponent<KIND>{op.isImaginaryPart,
+ rewriteExpr(std::move(op.template operand<Is>()))...};
+ }
+
+ template <int KIND, std::size_t... Is>
+ LogicalOperation<KIND> rewriteOperation(
+ LogicalOperation<KIND> &&op, std::index_sequence<Is...>) {
+ return LogicalOperation<KIND>{op.logicalOperator,
+ rewriteExpr(std::move(op.template operand<Is>()))...};
+ }
+
+ template <typename T, std::size_t... Is>
+ Relational<T> rewriteOperation(
+ Relational<T> &&op, std::index_sequence<Is...>) {
+ return Relational<T>{
+ op.opr, rewriteExpr(std::move(op.template operand<Is>()))...};
+ }
+
+ Relational<SomeType> rewriteRelational(Relational<SomeType> &&relational) {
+ return common::visit(
+ [&](auto &&typed) -> Relational<SomeType> {
+ using RelationalType = std::decay_t<decltype(typed)>;
+ return Relational<SomeType>{rewriteOperation(std::move(typed),
+ std::make_index_sequence<RelationalType::operands>{})};
+ },
+ std::move(relational.u));
+ }
+
+ template <typename T>
+ ConditionalExpr<T> rewriteConditional(ConditionalExpr<T> &&conditional) {
+ return ConditionalExpr<T>{rewriteExpr(std::move(conditional.condition())),
+ rewriteExpr(std::move(conditional.thenValue())),
+ rewriteExpr(std::move(conditional.elseValue()))};
+ }
+
+ template <typename T> FunctionRef<T> rewriteFunction(FunctionRef<T> &&ref) {
+ if (!ref.proc().IsPure())
+ return std::move(ref);
+ for (std::optional<ActualArgument> &maybeArg : ref.arguments()) {
+ if (maybeArg)
+ if (Expr<SomeType> * argExpr{maybeArg->UnwrapExpr()})
+ *argExpr = rewriteExpr(std::move(*argExpr));
+ }
+ return std::move(ref);
+ }
+
+ template <typename T> Expr<T> rewriteExpr(Expr<T> &&expr) {
+ if (std::optional<Expr<T>> rewritten{tryRewriteCurrent(expr)}) {
+ changed_ = true;
+ return std::move(*rewritten);
+ }
+ return common::visit(
+ [&](auto &&node) -> Expr<T> {
+ using Node = std::decay_t<decltype(node)>;
+ if constexpr (IsExpr<Node>::value) {
+ return Expr<T>{rewriteExpr(std::move(node))};
+ } else if constexpr (IsFunctionRef<Node>::value) {
+ return Expr<T>{rewriteFunction(std::move(node))};
+ } else if constexpr (IsConditionalExpr<Node>::value) {
+ return Expr<T>{rewriteConditional(std::move(node))};
+ } else if constexpr (std::is_same_v<Node, Relational<SomeType>>) {
+ return Expr<T>{rewriteRelational(std::move(node))};
+ } else if constexpr (rewrite::is_operation_v<Node>) {
+ if constexpr (std::is_same_v<Node, Parentheses<T>>)
+ return Expr<T>{std::move(node)};
+ else
+ return Expr<T>{rewriteOperation(
+ std::move(node), std::make_index_sequence<Node::operands>{})};
+ } else {
+ return Expr<T>{std::move(node)};
+ }
+ },
+ std::move(expr.u));
+ }
+
+ bool changed_{false};
+};
+
} // namespace
bool CanBuildSplitSumExpressionTree(FoldingContext &context,
@@ -1546,13 +1673,9 @@ bool CanBuildSplitSumExpressionTree(FoldingContext &context,
!HasVolatileOrAsynchronousSymbol(lhs);
}
-std::optional<Expr<SomeType>> TryBuildSplitSumExpressionTree(
+std::optional<Expr<SomeType>> TryBuildSplitSumExpressionTrees(
const Expr<SomeType> &expr) {
- return common::visit(
- [&](const auto &typedExpr) -> std::optional<Expr<SomeType>> {
- return tryBuildSplitSumExpressionTree(typedExpr);
- },
- expr.u);
+ return SplitSumExpressionTreeRewriter{}.Rewrite(expr);
}
bool IsArraySection(const Expr<SomeType> &expr) {
diff --git a/flang/lib/Lower/Bridge.cpp b/flang/lib/Lower/Bridge.cpp
index f9becabaa1fe2..d605259713007 100644
--- a/flang/lib/Lower/Bridge.cpp
+++ b/flang/lib/Lower/Bridge.cpp
@@ -5444,7 +5444,7 @@ class FirConverter : public Fortran::lower::AbstractConverter {
Fortran::evaluate::CanBuildSplitSumExpressionTree(
getFoldingContext(), assign.lhs, assign.rhs)) {
rewritten =
- Fortran::evaluate::TryBuildSplitSumExpressionTree(assign.rhs);
+ Fortran::evaluate::TryBuildSplitSumExpressionTrees(assign.rhs);
if (rewritten)
rhsExpr = &*rewritten;
}
diff --git a/flang/test/Lower/split-sum-expression-tree-lowering.f90 b/flang/test/Lower/split-sum-expression-tree-lowering.f90
index f681f4ebe45fe..4b9b49219d8b8 100644
--- a/flang/test/Lower/split-sum-expression-tree-lowering.f90
+++ b/flang/test/Lower/split-sum-expression-tree-lowering.f90
@@ -752,8 +752,8 @@ subroutine guard_call(x,a,b,c,d,e)
! 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.
+! A qualifying outer sum is rewritten once. The pure call remains an opaque
+! term, so its additive argument retains source order.
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
@@ -837,6 +837,170 @@ subroutine guard_nested_impure_call(x,a,b,c,d,e)
! NO-REWRITE: %[[RES:.*]] = arith.addf %[[HEAD_BC]], %[[DE]]
! NO-REWRITE: hlfir.assign %[[RES]]
+! Default: d * real((a+b)+c,8)
+! Rewritten: d * real(c+(a+b),8)
+subroutine nested_conversion_operand(x,a,b,c,d)
+ real(8) :: x,d
+ real(4) :: a,b,c
+ x = d * real(a+b+c,8)
+end
+
+! SPLIT-LABEL: func.func @_QPnested_conversion_operand
+! SPLIT: %[[DV:.*]] = fir.load
+! SPLIT: %[[CV:.*]] = fir.load
+! SPLIT: %[[AV:.*]] = fir.load
+! SPLIT: %[[BV:.*]] = fir.load
+! SPLIT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! SPLIT: %[[SUM:.*]] = arith.addf %[[CV]], %[[AB]]
+! SPLIT: %[[CONVERT:.*]] = fir.convert %[[SUM]]
+! SPLIT: %[[RES:.*]] = arith.mulf %[[DV]], %[[CONVERT]]
+! SPLIT: hlfir.assign %[[RES]]
+
+! DEFAULT-LABEL: func.func @_QPnested_conversion_operand
+! DEFAULT: %[[DV:.*]] = fir.load
+! DEFAULT: %[[AV:.*]] = fir.load
+! DEFAULT: %[[BV:.*]] = fir.load
+! DEFAULT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! DEFAULT: %[[CV:.*]] = fir.load
+! DEFAULT: %[[SUM:.*]] = arith.addf %[[AB]], %[[CV]]
+! DEFAULT: %[[CONVERT:.*]] = fir.convert %[[SUM]]
+! DEFAULT: %[[RES:.*]] = arith.mulf %[[DV]], %[[CONVERT]]
+! DEFAULT: hlfir.assign %[[RES]]
+
+! Default: sqrt((a+b)+c)
+! Rewritten: sqrt(c+(a+b))
+subroutine nested_pure_call_argument(x,a,b,c)
+ real(8) :: x,a,b,c
+ x = sqrt(a+b+c)
+end
+
+! SPLIT-LABEL: func.func @_QPnested_pure_call_argument
+! SPLIT: %[[CV:.*]] = fir.load
+! SPLIT: %[[AV:.*]] = fir.load
+! SPLIT: %[[BV:.*]] = fir.load
+! SPLIT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! SPLIT: %[[SUM:.*]] = arith.addf %[[CV]], %[[AB]]
+! SPLIT: %[[SQRT:.*]] = math.sqrt %[[SUM]]
+! SPLIT: hlfir.assign %[[SQRT]]
+
+! DEFAULT-LABEL: func.func @_QPnested_pure_call_argument
+! DEFAULT: %[[AV:.*]] = fir.load
+! DEFAULT: %[[BV:.*]] = fir.load
+! DEFAULT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! DEFAULT: %[[CV:.*]] = fir.load
+! DEFAULT: %[[SUM:.*]] = arith.addf %[[AB]], %[[CV]]
+! DEFAULT: %[[SQRT:.*]] = math.sqrt %[[SUM]]
+! DEFAULT: hlfir.assign %[[SQRT]]
+
+! Default: atan2((a+b)+c,(d+e)+f)
+! Rewritten: atan2(c+(a+b),f+(d+e))
+subroutine nested_separate_call_arguments(x,a,b,c,d,e,f)
+ real(8) :: x,a,b,c,d,e,f
+ x = atan2(a+b+c,d+e+f)
+end
+
+! SPLIT-LABEL: func.func @_QPnested_separate_call_arguments
+! SPLIT: %[[CV:.*]] = fir.load
+! SPLIT: %[[AV:.*]] = fir.load
+! SPLIT: %[[BV:.*]] = fir.load
+! SPLIT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! SPLIT: %[[FIRST:.*]] = arith.addf %[[CV]], %[[AB]]
+! SPLIT: %[[FV:.*]] = fir.load
+! SPLIT: %[[DV:.*]] = fir.load
+! SPLIT: %[[EV:.*]] = fir.load
+! SPLIT: %[[DE:.*]] = arith.addf %[[DV]], %[[EV]]
+! SPLIT: %[[SECOND:.*]] = arith.addf %[[FV]], %[[DE]]
+! SPLIT: math.atan2 %[[FIRST]], %[[SECOND]]
+
+! Default: (flag ? (a+b)+c : d)
+! Rewritten: (flag ? c+(a+b) : d)
+subroutine nested_conditional_branch(x,flag,a,b,c,d)
+ real(8) :: x,a,b,c,d
+ logical :: flag
+ x = (flag ? a+b+c : d)
+end
+
+! SPLIT-LABEL: func.func @_QPnested_conditional_branch
+! SPLIT: fir.if
+! SPLIT: %[[CV:.*]] = fir.load
+! SPLIT: %[[AV:.*]] = fir.load
+! SPLIT: %[[BV:.*]] = fir.load
+! SPLIT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! SPLIT: %[[SUM:.*]] = arith.addf %[[CV]], %[[AB]]
+! SPLIT: fir.result %[[SUM]]
+
+! DEFAULT-LABEL: func.func @_QPnested_conditional_branch
+! DEFAULT: fir.if
+! DEFAULT: %[[AV:.*]] = fir.load
+! DEFAULT: %[[BV:.*]] = fir.load
+! DEFAULT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! DEFAULT: %[[CV:.*]] = fir.load
+! DEFAULT: %[[SUM:.*]] = arith.addf %[[AB]], %[[CV]]
+! DEFAULT: fir.result %[[SUM]]
+
+! Default: ((a+b)+c > d ? e : f)
+! Rewritten: (c+(a+b) > d ? e : f)
+subroutine nested_relational_operand(x,a,b,c,d,e,f)
+ real(8) :: x,a,b,c,d,e,f
+ x = (a+b+c > d ? e : f)
+end
+
+! SPLIT-LABEL: func.func @_QPnested_relational_operand
+! SPLIT: %[[CV:.*]] = fir.load
+! SPLIT: %[[AV:.*]] = fir.load
+! SPLIT: %[[BV:.*]] = fir.load
+! SPLIT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! SPLIT: %[[SUM:.*]] = arith.addf %[[CV]], %[[AB]]
+! SPLIT: arith.cmpf ogt, %[[SUM]]
+
+! DEFAULT-LABEL: func.func @_QPnested_relational_operand
+! DEFAULT: %[[AV:.*]] = fir.load
+! DEFAULT: %[[BV:.*]] = fir.load
+! DEFAULT: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! DEFAULT: %[[CV:.*]] = fir.load
+! DEFAULT: %[[SUM:.*]] = arith.addf %[[AB]], %[[CV]]
+! DEFAULT: arith.cmpf ogt, %[[SUM]]
+
+subroutine guard_parenthesized_call_argument(x,a,b,c)
+ real(8) :: x,a,b,c
+ x = sqrt((a+b+c))
+end
+
+! NO-REWRITE-LABEL: func.func @_QPguard_parenthesized_call_argument
+! NO-REWRITE: %[[AV:.*]] = fir.load
+! NO-REWRITE: %[[BV:.*]] = fir.load
+! NO-REWRITE: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! NO-REWRITE: %[[CV:.*]] = fir.load
+! NO-REWRITE: %[[SUM:.*]] = arith.addf %[[AB]], %[[CV]]
+! NO-REWRITE: %[[PAREN:.*]] = hlfir.no_reassoc %[[SUM]]
+! NO-REWRITE: %[[SQRT:.*]] = math.sqrt %[[PAREN]]
+! NO-REWRITE: hlfir.assign %[[SQRT]]
+
+subroutine guard_short_call_argument(x,a,b)
+ real(8) :: x,a,b
+ x = sqrt(a+b)
+end
+
+! NO-REWRITE-LABEL: func.func @_QPguard_short_call_argument
+! NO-REWRITE: %[[AV:.*]] = fir.load
+! NO-REWRITE: %[[BV:.*]] = fir.load
+! NO-REWRITE: %[[SUM:.*]] = arith.addf %[[AV]], %[[BV]]
+! NO-REWRITE: %[[SQRT:.*]] = math.sqrt %[[SUM]]
+! NO-REWRITE: hlfir.assign %[[SQRT]]
+
+subroutine guard_non_assignment_context(a,b,c)
+ real(8) :: a,b,c
+ call consume(a+b+c)
+end
+
+! NO-REWRITE-LABEL: func.func @_QPguard_non_assignment_context
+! NO-REWRITE: %[[AV:.*]] = fir.load
+! NO-REWRITE: %[[BV:.*]] = fir.load
+! NO-REWRITE: %[[AB:.*]] = arith.addf %[[AV]], %[[BV]]
+! NO-REWRITE: %[[CV:.*]] = fir.load
+! NO-REWRITE: %[[SUM:.*]] = arith.addf %[[AB]], %[[CV]]
+! NO-REWRITE: fir.call @_QPconsume
+
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)
More information about the llvm-branch-commits
mailing list