diff --git a/gluten-substrait/src/main/scala/org/apache/gluten/execution/BatchScanExecTransformer.scala b/gluten-substrait/src/main/scala/org/apache/gluten/execution/BatchScanExecTransformer.scala index a0c3bb87575..54fd01877a0 100644 --- a/gluten-substrait/src/main/scala/org/apache/gluten/execution/BatchScanExecTransformer.scala +++ b/gluten-substrait/src/main/scala/org/apache/gluten/execution/BatchScanExecTransformer.scala @@ -68,6 +68,8 @@ case class BatchScanExecTransformer( runtimeFilters = QueryPlan.normalizePredicates( runtimeFilters.filterNot(_ == DynamicPruningExpression(Literal.TrueLiteral)), output), + keyGroupedPartitioning = keyGroupedPartitioning.map( + _.map(QueryPlan.normalizeExpressions(_, output))), pushDownFilters = pushDownFilters.map(QueryPlan.normalizePredicates(_, output)) ) } @@ -194,12 +196,15 @@ abstract class BatchScanExecTransformerBase( override def equals(other: Any): Boolean = other match { case other: BatchScanExecTransformerBase => - this.pushDownFilters == other.pushDownFilters && super.equals(other) + this.keyGroupedPartitioning == other.keyGroupedPartitioning && + this.pushDownFilters == other.pushDownFilters && + super.equals(other) case _ => false } - override def hashCode(): Int = Objects.hashCode(batch, runtimeFilters, pushDownFilters) + override def hashCode(): Int = + Objects.hashCode(batch, runtimeFilters, keyGroupedPartitioning, pushDownFilters) /** Return a copy of this scan with a new output schema. */ def withOutput(newOutput: Seq[AttributeReference]): BatchScanExecTransformerBase diff --git a/gluten-substrait/src/main/scala/org/apache/gluten/execution/ScanTransformerFactory.scala b/gluten-substrait/src/main/scala/org/apache/gluten/execution/ScanTransformerFactory.scala index 711f8c69086..b2839aa43b1 100644 --- a/gluten-substrait/src/main/scala/org/apache/gluten/execution/ScanTransformerFactory.scala +++ b/gluten-substrait/src/main/scala/org/apache/gluten/execution/ScanTransformerFactory.scala @@ -44,6 +44,8 @@ object ScanTransformerFactory { batchScanExec.output, batchScanExec.scan, batchScanExec.runtimeFilters, + keyGroupedPartitioning = + SparkShimLoader.getSparkShims.getKeyGroupedPartitioning(batchScanExec), table = SparkShimLoader.getSparkShims.getBatchScanExecTable(batchScanExec) ) }