This is an automated email from the ASF dual-hosted git repository.

hongze pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/incubator-gluten.git


The following commit(s) were added to refs/heads/main by this push:
     new 7f416da938 [VL] Remove unnecessary vanilla Spark compatibility code 
for VeloxCollectSet function (#7590)
7f416da938 is described below

commit 7f416da938e2185524da23596c0388a34c14d694
Author: Hongze Zhang <[email protected]>
AuthorDate: Fri Oct 18 15:14:59 2024 +0800

    [VL] Remove unnecessary vanilla Spark compatibility code for 
VeloxCollectSet function (#7590)
---
 .../apache/gluten/expression/aggregate/VeloxCollect.scala    |  8 ++------
 .../org/apache/gluten/extension/CollectRewriteRule.scala     | 12 +++---------
 2 files changed, 5 insertions(+), 15 deletions(-)

diff --git 
a/backends-velox/src/main/scala/org/apache/gluten/expression/aggregate/VeloxCollect.scala
 
b/backends-velox/src/main/scala/org/apache/gluten/expression/aggregate/VeloxCollect.scala
index c35020fab3..dc41bbc4fc 100644
--- 
a/backends-velox/src/main/scala/org/apache/gluten/expression/aggregate/VeloxCollect.scala
+++ 
b/backends-velox/src/main/scala/org/apache/gluten/expression/aggregate/VeloxCollect.scala
@@ -29,6 +29,8 @@ abstract class VeloxCollect(child: Expression)
 
   override def dataType: DataType = ArrayType(child.dataType, false)
 
+  override def nullable: Boolean = false
+
   override def aggBufferAttributes: Seq[AttributeReference] = Seq(buffer)
 
   override lazy val initialValues: Seq[Expression] = 
Seq(Literal.create(Array(), dataType))
@@ -49,10 +51,6 @@ abstract class VeloxCollect(child: Expression)
 
 case class VeloxCollectSet(child: Expression) extends VeloxCollect(child) {
 
-  // Velox's collect_set implementation allows null output. Thus we usually 
wrap
-  // the function to enforce non-null output. See 
CollectRewriteRule#ensureNonNull.
-  override def nullable: Boolean = true
-
   override lazy val evaluateExpression: Expression =
     ArrayDistinct(buffer)
 
@@ -64,8 +62,6 @@ case class VeloxCollectSet(child: Expression) extends 
VeloxCollect(child) {
 
 case class VeloxCollectList(child: Expression) extends VeloxCollect(child) {
 
-  override def nullable: Boolean = false
-
   override val evaluateExpression: Expression = buffer
 
   override def prettyName: String = "velox_collect_list"
diff --git 
a/backends-velox/src/main/scala/org/apache/gluten/extension/CollectRewriteRule.scala
 
b/backends-velox/src/main/scala/org/apache/gluten/extension/CollectRewriteRule.scala
index d182d0aae7..86defb2fcb 100644
--- 
a/backends-velox/src/main/scala/org/apache/gluten/extension/CollectRewriteRule.scala
+++ 
b/backends-velox/src/main/scala/org/apache/gluten/extension/CollectRewriteRule.scala
@@ -21,7 +21,7 @@ import 
org.apache.gluten.expression.aggregate.{VeloxCollectList, VeloxCollectSet
 import org.apache.gluten.utils.LogicalPlanSelector
 
 import org.apache.spark.sql.SparkSession
-import org.apache.spark.sql.catalyst.expressions.{And, Expression, IsNotNull, 
WindowExpression}
+import org.apache.spark.sql.catalyst.expressions.{Expression, WindowExpression}
 import org.apache.spark.sql.catalyst.expressions.aggregate._
 import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, LogicalPlan, 
Window}
 import org.apache.spark.sql.catalyst.rules.Rule
@@ -72,15 +72,9 @@ case class CollectRewriteRule(spark: SparkSession) extends 
Rule[LogicalPlan] {
 object CollectRewriteRule {
   private object ToVeloxCollect {
     def unapply(expr: Expression): Option[Expression] = expr match {
-      case aggExpr @ AggregateExpression(s: CollectSet, _, _, filter, _) if 
has[VeloxCollectSet] =>
-        // 1. Replace null result from VeloxCollectSet with empty array to 
align with
-        //    vanilla Spark.
-        // 2. Filter out null inputs from VeloxCollectSet to align with 
vanilla Spark.
-        //
-        // Since https://github.com/apache/incubator-gluten/pull/4805
-        val newFilter = (filter ++ Some(IsNotNull(s.child))).reduceOption(And)
+      case aggExpr @ AggregateExpression(s: CollectSet, _, _, _, _) if 
has[VeloxCollectSet] =>
         val newAggExpr =
-          aggExpr.copy(aggregateFunction = VeloxCollectSet(s.child), filter = 
newFilter)
+          aggExpr.copy(aggregateFunction = VeloxCollectSet(s.child))
         Some(newAggExpr)
       case aggExpr @ AggregateExpression(l: CollectList, _, _, _, _) if 
has[VeloxCollectList] =>
         val newAggExpr = aggExpr.copy(VeloxCollectList(l.child))


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to