diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/ProjectAggregateExpressionsForCse.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/ProjectAggregateExpressionsForCse.java index 1f3ed8c0be7448..9f869118800bc6 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/ProjectAggregateExpressionsForCse.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/processor/post/ProjectAggregateExpressionsForCse.java @@ -29,6 +29,8 @@ import org.apache.doris.nereids.trees.expressions.Slot; import org.apache.doris.nereids.trees.expressions.SlotReference; import org.apache.doris.nereids.trees.expressions.functions.agg.AggregateFunction; +import org.apache.doris.nereids.trees.plans.AggMode; +import org.apache.doris.nereids.trees.plans.AggPhase; import org.apache.doris.nereids.trees.plans.Plan; import org.apache.doris.nereids.trees.plans.algebra.Aggregate; import org.apache.doris.nereids.trees.plans.physical.AbstractPhysicalPlan; @@ -52,6 +54,16 @@ /** * create project under aggregate to enable CSE + * + *

For one-phase aggregates whose child is a PhysicalDistribute + * (aggregate -> distribute -> scan), the CSE project is inserted below the + * distribute so that the distribution-key slots stay intact and the exchange + * only carries the (already pruned) aggregate input. The translator's bucketed + * fusion (fusing one-phase aggregate + distribute into BucketedAggregationNode) + * builds directly on the distribute's child, so the fused plan naturally + * becomes BucketedAgg(sum(x), max(x)) -> Project(a+b AS x) -> scan and the + * common aggregate argument is evaluated once per row instead of once per + * aggregate function.

*/ public class ProjectAggregateExpressionsForCse extends PlanPostProcessor { @Override @@ -59,7 +71,8 @@ public Plan visitPhysicalHashAggregate(PhysicalHashAggregate agg aggregate = (PhysicalHashAggregate) super.visit(aggregate, ctx); // for multi-phases aggregate, only process the 1st phase aggregate - if (aggregate.child() instanceof PhysicalDistribute || aggregate.child() instanceof Aggregate) { + // Bucketed agg is always single-phase, but keep the same guard for safety. + if (aggregate.child() instanceof Aggregate) { return aggregate; } @@ -154,6 +167,60 @@ public Plan visitPhysicalHashAggregate(PhysicalHashAggregate agg aggregate = (PhysicalHashAggregate) aggregate .withAggOutput(aggOutputReplaced) .withChildren(project); + } else if (aggregate.child() instanceof PhysicalDistribute) { + // One-phase (INPUT_TO_RESULT) aggregate over a distribute + // (aggregate -> distribute -> scan): insert the CSE project between + // the distribute and its child, instead of between the aggregate and + // the distribute. This keeps the aggregate's child as a distribute + // (so bucketed fusion and the property machinery still see the same + // shape), and the project lands inside the scan + // fragment, so the common aggregate argument is computed once per row + // before the exchange. After bucketed fusion bypasses the distribute, + // the executed plan is BucketedAgg(sum(x), max(x)) -> Project(a+b AS x) + // -> scan. + // + // Only the one-phase shape reaches here with complex aggregate + // arguments: two-phase GLOBAL aggregates (BUFFER_TO_RESULT) reference + // the local phase's intermediate slots, so no CSE candidate is + // extracted for them anyway. Guard explicitly anyway to keep the + // intent clear and to stay safe if a future aggregate function + // surfaces a non-slot argument on the GLOBAL phase. + if (!(aggregate instanceof PhysicalHashAggregate)) { + return aggregate; + } + PhysicalHashAggregate hashAggregate = + (PhysicalHashAggregate) aggregate; + if (hashAggregate.getAggPhase() != AggPhase.GLOBAL + || hashAggregate.getAggMode() != AggMode.INPUT_TO_RESULT) { + return aggregate; + } + PhysicalDistribute distribute = (PhysicalDistribute) aggregate.child(); + List projections = new ArrayList<>(); + projections.addAll(inputSlots); + projections.addAll(cseCandidates.values()); + List projectOutput = new ImmutableList.Builder() + .addAll(inputSlots) + .addAll(slotMap.values()) + .build(); + LogicalProperties projectLogicalProperties = new LogicalProperties( + () -> projectOutput, + () -> DataTrait.EMPTY_TRAIT + ); + AbstractPhysicalPlan distributeChild = ((AbstractPhysicalPlan) distribute.child()); + PhysicalProperties projectPhysicalProperties = ChildOutputPropertyDeriver.computeProjectOutputProperties( + projections, distributeChild.getPhysicalProperties()); + PhysicalProject project = new PhysicalProject<>(projections, Optional.empty(), + projectLogicalProperties, + projectPhysicalProperties, + distributeChild.getStats(), + distribute.child()); + // withChildren keeps the distribution spec and physical properties of the + // distribute unchanged; its output now comes from the CSE project, which + // still carries every distribution-key slot (the group-by slots are part + // of inputSlots above). + PhysicalDistribute newDistribute = distribute.withChildren(ImmutableList.of(project)); + return (Plan) aggregate.withAggOutput(aggOutputReplaced) + .withChildren(newDistribute); } else { List projections = new ArrayList<>(); projections.addAll(inputSlots); diff --git a/regression-test/data/nereids_rules_p0/agg_strategy/cse_agg_distribute.out b/regression-test/data/nereids_rules_p0/agg_strategy/cse_agg_distribute.out new file mode 100644 index 00000000000000..91465208722b1f --- /dev/null +++ b/regression-test/data/nereids_rules_p0/agg_strategy/cse_agg_distribute.out @@ -0,0 +1,5 @@ +-- This file is automatically generated. You should know what you did if you want to edit this +-- !one_phase_join_result -- +g1 33 19 33 19 +g2 22 15 22 15 + diff --git a/regression-test/suites/nereids_rules_p0/agg_strategy/cse_agg_distribute.groovy b/regression-test/suites/nereids_rules_p0/agg_strategy/cse_agg_distribute.groovy new file mode 100644 index 00000000000000..dc8f8c8ee64c6c --- /dev/null +++ b/regression-test/suites/nereids_rules_p0/agg_strategy/cse_agg_distribute.groovy @@ -0,0 +1,69 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +suite("cse_agg_distribute") { + sql "SET enable_nereids_planner=true" + sql "SET enable_fallback_to_original_planner=false" + sql "SET runtime_filter_mode=OFF" + + sql "DROP TABLE IF EXISTS cse_agg_distribute_tbl" + sql """ + CREATE TABLE cse_agg_distribute_tbl ( + id int, + grp varchar(20), + a int, + b int + ) DUPLICATE KEY(id) + DISTRIBUTED BY HASH(id) BUCKETS 3 + PROPERTIES('replication_num' = '1') + """ + sql """ INSERT INTO cse_agg_distribute_tbl VALUES + (1, 'g1', 1, 2), + (2, 'g2', 3, 4), + (3, 'g1', 5, 6), + (4, 'g2', 7, 8), + (5, 'g1', 9, 10) + """ + + // SUM(a+b) and MAX(a+b) share the same argument, so the aggregate-argument + // CSE must extract "a+b" into a project node and make both functions + // reference the extracted slot, instead of re-evaluating a+b per function. + String query = "SELECT grp, SUM(a+b), MAX(a+b) FROM cse_agg_distribute_tbl GROUP BY grp" + + // --------------------------------------------------------------------- + // one-phase aggregate over a distribute (the aggregate is a join child, + // so the distribute is required by the join): the CSE project must be + // inserted below the distribute, keeping the distribution-key slots + // intact. Both aggregates must reference the extracted slot (4 + // occurrences: SUM/MAX of each side). + // --------------------------------------------------------------------- + sql "set agg_phase=1" + // branch-4.2 has no bucketed hash aggregate, so enable_bucketed_hash_agg does not exist here + String joinQuery = """ + SELECT t1.grp, t1.s, t1.m, t2.s2, t2.m2 FROM + (SELECT grp, SUM(a+b) s, MAX(a+b) m FROM cse_agg_distribute_tbl GROUP BY grp) t1 + JOIN (SELECT grp, SUM(a+b) s2, MAX(a+b) m2 FROM cse_agg_distribute_tbl GROUP BY grp) t2 + ON t1.grp = t2.grp + """ + explain { + sql("${joinQuery}") + contains("VEXCHANGE") + contains("VSELECT") + multiContains("cast(a as BIGINT) + cast(b as BIGINT))[#", 4) + } + order_qt_one_phase_join_result """${joinQuery} ORDER BY t1.grp""" +}