This is an automated email from the ASF dual-hosted git repository.
morrySnow pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/doris.git
The following commit(s) were added to refs/heads/master by this push:
new 749290cb041 [opt](repeat) deduplicate grouping scalar functions from
projetion inrepeat (#65880)
749290cb041 is described below
commit 749290cb04107ae08e6b33980495f81b35874467
Author: morrySnow <[email protected]>
AuthorDate: Tue Jul 28 10:14:59 2026 +0800
[opt](repeat) deduplicate grouping scalar functions from projetion inrepeat
(#65880)
### What problem does this PR solve?
Related PR: #34196
Problem Summary:
when normalize repeat, we maybe generate duplicate grouping functions
because the same grouping functions appear in projection more than once.
---
.../nereids/rules/analysis/NormalizeRepeat.java | 34 +++++++++-------------
.../doris/nereids/trees/plans/algebra/Repeat.java | 2 +-
.../apache/doris/nereids/util/ExpressionUtils.java | 4 +--
.../rules/analysis/NormalizeRepeatTest.java | 30 +++++++++++++++++++
4 files changed, 47 insertions(+), 23 deletions(-)
diff --git
a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeat.java
b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeat.java
index 1f00cc3e619..bbabf96f7f1 100644
---
a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeat.java
+++
b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeat.java
@@ -129,8 +129,7 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory
{
}
private static void checkGroupingSetsSize(LogicalRepeat<Plan> repeat) {
- Set<Expression> flattenGroupingSetExpr = ImmutableSet.copyOf(
- ExpressionUtils.flatExpressions(repeat.getGroupingSets()));
+ Set<Expression> flattenGroupingSetExpr =
ExpressionUtils.flatExpressions(repeat.getGroupingSets());
if (flattenGroupingSetExpr.size() >
LogicalRepeat.MAX_GROUPING_SETS_NUM) {
throw new AnalysisException(
"Too many sets in GROUP BY clause, the max grouping sets
item is "
@@ -139,7 +138,9 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory
{
}
private static LogicalAggregate<Plan> normalizeRepeat(LogicalRepeat<Plan>
repeat) {
- Set<Expression> needToSlotsGroupingExpr =
collectNeedToSlotGroupingExpr(repeat);
+ // grouping sets should be pushed down, e.g. grouping sets((k + 1)),
+ // we should push down the `k + 1` to the bottom plan
+ Set<Expression> needToSlotsGroupingExpr =
ExpressionUtils.flatExpressions(repeat.getGroupingSets());
NormalizeToSlotContext groupingExprContext = buildContext(repeat,
needToSlotsGroupingExpr);
Map<Expression, NormalizeToSlotTriplet> groupingExprMap =
groupingExprContext.getNormalizeToSlotMap();
Map<Expression, Alias> existsAlias = getExistsAlias(repeat,
groupingExprMap);
@@ -159,7 +160,8 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory
{
// rewrite grouping scalar function to virtual slots
// rewrite the arguments of agg function to slots
List<NamedExpression> normalizedAggOutput = Lists.newArrayList();
- List<NamedExpression> groupingFunctions = Lists.newArrayList();
+ // use a map to deduplicate grouping scalar functions in projection
+ Map<GroupingScalarFunction, NamedExpression> groupingFunctions =
Maps.newHashMap();
for (Expression expr : repeat.getOutputExpressions()) {
Expression rewrittenExpr = expr.rewriteDownShortCircuit(
e ->
normalizeAggFuncChildrenAndGroupingScalarFunc(argsContext, e,
groupingFunctions));
@@ -172,17 +174,16 @@ public class NormalizeRepeat extends
OneAnalysisRuleFactory {
Set<SlotReference> aggUsedSlots = ExpressionUtils.collect(
normalizedAggOutput, expr ->
expr.getClass().equals(SlotReference.class));
- Set<Slot> groupingSetsUsedSlot = ImmutableSet.copyOf(
- ExpressionUtils.flatExpressions(normalizedGroupingSets));
+ Set<Slot> groupingSetsUsedSlot =
ExpressionUtils.flatExpressions(normalizedGroupingSets);
SetView<SlotReference> aggUsedSlotNotInGroupBy
- = Sets.difference(Sets.difference(aggUsedSlots,
groupingFunctions.stream()
+ = Sets.difference(Sets.difference(aggUsedSlots,
groupingFunctions.values().stream()
.map(NamedExpression::toSlot).collect(Collectors.toSet())),
groupingSetsUsedSlot);
List<NamedExpression> normalizedRepeatOutput =
ImmutableList.<NamedExpression>builder()
.addAll(groupingSetsUsedSlot)
.addAll(aggUsedSlotNotInGroupBy)
- .addAll(groupingFunctions)
+ .addAll(groupingFunctions.values())
.build();
// 3 parts need push down:
@@ -218,7 +219,7 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory
{
List<Expression> normalizedAggGroupBy =
ImmutableList.<Expression>builder()
.addAll(groupingSetsUsedSlot)
-
.addAll(groupingFunctions.stream().map(NamedExpression::toSlot).collect(Collectors.toList()))
+
.addAll(groupingFunctions.values().stream().map(NamedExpression::toSlot).collect(Collectors.toList()))
.add(groupingId)
.build();
@@ -227,13 +228,6 @@ public class NormalizeRepeat extends
OneAnalysisRuleFactory {
Optional.of(normalizedRepeat), normalizedRepeat);
}
- private static Set<Expression>
collectNeedToSlotGroupingExpr(LogicalRepeat<Plan> repeat) {
- // grouping sets should be pushed down, e.g. grouping sets((k + 1)),
- // we should push down the `k + 1` to the bottom plan
- return ImmutableSet.copyOf(
- ExpressionUtils.flatExpressions(repeat.getGroupingSets()));
- }
-
private static Set<Expression>
collectNeedToSlotArgsOfGroupingScalarFuncAndAggFunc(LogicalRepeat<Plan> repeat)
{
Set<GroupingScalarFunction> groupingScalarFunctions =
ExpressionUtils.collect(
repeat.getOutputExpressions(),
GroupingScalarFunction.class::isInstance);
@@ -300,7 +294,7 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory
{
private static NormalizeToSlotContext buildContextWithAlias(Repeat<?
extends Plan> repeat,
Map<Expression, Alias> existsAliasMap, Collection<? extends
Expression> sourceExpressions) {
- List<Expression> groupingSetExpressions =
ExpressionUtils.flatExpressions(repeat.getGroupingSets());
+ Set<Expression> groupingSetExpressions =
ExpressionUtils.flatExpressions(repeat.getGroupingSets());
Map<Expression, NormalizeToSlotTriplet> normalizeToSlotMap =
Maps.newLinkedHashMap();
for (Expression expression : sourceExpressions) {
Optional<NormalizeToSlotTriplet> pushDownTriplet;
@@ -325,7 +319,7 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory
{
}
private static Expression
normalizeAggFuncChildrenAndGroupingScalarFunc(NormalizeToSlotContext context,
- Expression expr, List<NamedExpression> groupingSetExpressions) {
+ Expression expr, Map<GroupingScalarFunction, NamedExpression>
groupingSetExpressions) {
if (expr instanceof AggregateFunction) {
AggregateFunction function = (AggregateFunction) expr;
List<Expression> normalizedRealExpressions =
context.normalizeToUseSlotRef(function.getArguments());
@@ -336,8 +330,8 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory
{
List<Expression> normalizedRealExpressions =
context.normalizeToUseSlotRef(function.getArguments());
function = function.withChildren(normalizedRealExpressions);
// eliminate GroupingScalarFunction and replace to
VirtualSlotReference
- Alias alias = new Alias(function,
Repeat.generateVirtualSlotName(function));
- groupingSetExpressions.add(alias);
+ NamedExpression alias =
groupingSetExpressions.computeIfAbsent(function,
+ f -> new Alias(f, Repeat.generateVirtualSlotName(f)));
return alias.toSlot();
} else {
return expr;
diff --git
a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/algebra/Repeat.java
b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/algebra/Repeat.java
index 7a7f1f1f3c2..2e0dde6c305 100644
---
a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/algebra/Repeat.java
+++
b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/algebra/Repeat.java
@@ -51,7 +51,7 @@ public interface Repeat<CHILD_PLAN extends Plan> extends
Aggregate<CHILD_PLAN> {
@Override
default List<Expression> getGroupByExpressions() {
- return ExpressionUtils.flatExpressions(getGroupingSets());
+ return
ImmutableList.copyOf(ExpressionUtils.flatExpressions(getGroupingSets()));
}
@Override
diff --git
a/fe/fe-core/src/main/java/org/apache/doris/nereids/util/ExpressionUtils.java
b/fe/fe-core/src/main/java/org/apache/doris/nereids/util/ExpressionUtils.java
index 6d318cafd9c..fd7e761bc6e 100644
---
a/fe/fe-core/src/main/java/org/apache/doris/nereids/util/ExpressionUtils.java
+++
b/fe/fe-core/src/main/java/org/apache/doris/nereids/util/ExpressionUtils.java
@@ -851,13 +851,13 @@ public class ExpressionUtils {
}
/** flatExpressions */
- public static <E extends Expression> List<E> flatExpressions(List<List<E>>
expressionLists) {
+ public static <E extends Expression> Set<E> flatExpressions(List<List<E>>
expressionLists) {
int num = 0;
for (List<E> expressionList : expressionLists) {
num += expressionList.size();
}
- ImmutableList.Builder<E> flatten =
ImmutableList.builderWithExpectedSize(num);
+ ImmutableSet.Builder<E> flatten =
ImmutableSet.builderWithExpectedSize(num);
for (List<E> expressionList : expressionLists) {
flatten.addAll(expressionList);
}
diff --git
a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeatTest.java
b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeatTest.java
index e2c874aad17..c1ba24c0e2d 100644
---
a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeatTest.java
+++
b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeatTest.java
@@ -21,8 +21,10 @@ import org.apache.doris.nereids.trees.expressions.Alias;
import org.apache.doris.nereids.trees.expressions.Slot;
import org.apache.doris.nereids.trees.expressions.functions.agg.Sum;
import org.apache.doris.nereids.trees.expressions.functions.scalar.GroupingId;
+import
org.apache.doris.nereids.trees.expressions.functions.scalar.GroupingScalarFunction;
import org.apache.doris.nereids.trees.plans.Plan;
import org.apache.doris.nereids.trees.plans.algebra.Repeat.RepeatType;
+import org.apache.doris.nereids.trees.plans.logical.LogicalAggregate;
import org.apache.doris.nereids.trees.plans.logical.LogicalOlapScan;
import org.apache.doris.nereids.trees.plans.logical.LogicalRepeat;
import org.apache.doris.nereids.util.MemoPatternMatchSupported;
@@ -31,6 +33,7 @@ import org.apache.doris.nereids.util.PlanChecker;
import org.apache.doris.nereids.util.PlanConstructor;
import com.google.common.collect.ImmutableList;
+import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.Test;
public class NormalizeRepeatTest implements MemoPatternMatchSupported {
@@ -92,4 +95,31 @@ public class NormalizeRepeatTest implements
MemoPatternMatchSupported {
logicalAggregate(logicalRepeat(logicalOlapScan()))
);
}
+
+ @Test
+ public void testDeduplicateGroupingScalarFunctions() {
+ Slot id = scan1.getOutput().get(0);
+ Slot name = scan1.getOutput().get(1);
+ Alias firstGroupingId = new Alias(new GroupingId(name),
"first_grouping_id");
+ Alias secondGroupingId = new Alias(new GroupingId(name),
"second_grouping_id");
+ LogicalRepeat<Plan> repeat = new LogicalRepeat<>(
+ ImmutableList.of(ImmutableList.of(id)),
+ ImmutableList.of(id, firstGroupingId, secondGroupingId),
+ RepeatType.GROUPING_SETS,
+ scan1
+ );
+
+ LogicalAggregate<Plan> normalizedAggregate =
NormalizeRepeat.doNormalize(repeat);
+ LogicalRepeat<Plan> normalizedRepeat = (LogicalRepeat<Plan>)
normalizedAggregate.child();
+ long groupingFunctionCount =
normalizedRepeat.getOutputExpressions().stream()
+ .filter(expression -> expression instanceof Alias
+ && ((Alias) expression).child() instanceof
GroupingScalarFunction)
+ .count();
+
+ Assertions.assertEquals(1, groupingFunctionCount);
+ Assertions.assertEquals(3,
normalizedAggregate.getOutputExpressions().size());
+ Assertions.assertEquals(
+ ((Alias)
normalizedAggregate.getOutputExpressions().get(1)).child(),
+ ((Alias)
normalizedAggregate.getOutputExpressions().get(2)).child());
+ }
}
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]