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 <[email protected]> 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) _______________________________________________ llvm-branch-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/llvm-branch-commits
