From 65c57520c699067a38e6be5ee03fbc342bf1e299 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Fri, 31 Jul 2026 22:02:23 +0800 Subject: [PATCH 1/4] fix: normalize floating-point values in native collect_set --- .../org/apache/comet/serde/aggregates.scala | 20 ++++-- .../spark/sql/comet/CometExecUtils.scala | 7 +- .../expressions/aggregate/collect_set.sql | 37 ----------- .../collect_set_floating_fallback.sql | 49 ++++++++++++++ .../aggregate/collect_set_normalization.sql | 64 +++++++++++++++++++ 5 files changed, 133 insertions(+), 44 deletions(-) create mode 100644 spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_floating_fallback.sql create mode 100644 spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_normalization.sql diff --git a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala index 316744be1f6..0c236999706 100644 --- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala +++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala @@ -22,13 +22,14 @@ package org.apache.comet.serde import scala.jdk.CollectionConverters._ import org.apache.spark.sql.catalyst.expressions.{Attribute, Cast, Literal} -import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, ApproximatePercentile, Average, BitAndAgg, BitOrAgg, BitXorAgg, BloomFilterAggregate, CentralMomentAgg, CollectList, CollectSet, Corr, Count, Covariance, CovPopulation, CovSample, First, HyperLogLogPlusPlus, Last, Max, Min, Percentile, StddevPop, StddevSamp, Sum, VariancePop, VarianceSamp} +import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, ApproximatePercentile, Average, BitAndAgg, BitOrAgg, BitXorAgg, BloomFilterAggregate, CentralMomentAgg, CollectList, CollectSet, Corr, Count, Covariance, CovPopulation, CovSample, First, HyperLogLogPlusPlus, Last, Max, Min, Partial, Percentile, StddevPop, StddevSamp, Sum, VariancePop, VarianceSamp} import org.apache.spark.sql.catalyst.util.ArrayData +import org.apache.spark.sql.comet.CometExecUtils import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{BinaryType, BooleanType, ByteType, DataType, DateType, DecimalType, DoubleType, FloatType, IntegerType, LongType, NumericType, ShortType, StringType, TimestampNTZType, TimestampType} import org.apache.comet.CometConf.COMET_EXEC_STRICT_FLOATING_POINT -import org.apache.comet.CometSparkSessionExtensions.{isSpark41Plus, withFallbackReason} +import org.apache.comet.CometSparkSessionExtensions.{isSpark41Plus, isSpark42Plus, withFallbackReason} import org.apache.comet.expressions.CometEvalMode import org.apache.comet.serde.QueryPlanSerde.{evalModeToProto, exprToProto, serializeDataType} import org.apache.comet.shims.{CometCollectShim, CometEvalModeUtil} @@ -852,9 +853,9 @@ object CometBloomFilterAggregate extends CometAggregateExpressionSerde[BloomFilt object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { override def getIncompatibleReasons(): Seq[String] = Seq( - "Comet deduplicates NaN values (treats `NaN == NaN`) while Spark treats each NaN as a" + - s" distinct value. When `${COMET_EXEC_STRICT_FLOATING_POINT.key}=true`, `collect_set`" + - " on floating-point types falls back to Spark unless" + + "Before Spark 4.2, Comet deduplicates NaN values (treats `NaN == NaN`) while Spark treats" + + s" each NaN as a distinct value. When `${COMET_EXEC_STRICT_FLOATING_POINT.key}=true`," + + " `collect_set` on floating-point types falls back to Spark on those versions unless" + " `spark.comet.expression.CollectSet.allowIncompatible=true` is set.") override def getSupportLevel(expr: CollectSet): SupportLevel = { @@ -865,6 +866,8 @@ object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { // analysis time, and CometCollectShim.ignoreNulls hardcodes true, making this a no-op. if (!CometCollectShim.ignoreNulls(expr)) { Unsupported(Some("collect_set with RESPECT NULLS (ignoreNulls = false) is not supported")) + } else if (isSpark42Plus) { + Compatible() } else { SupportLevel .strictFloatingPointReason( @@ -882,7 +885,12 @@ object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { inputs: Seq[Attribute], binding: Boolean, conf: SQLConf): Option[ExprOuterClass.AggExpr] = { - val child = expr.children.head + val child = + if (isSpark42Plus && aggExpr.mode == Partial) { + CometExecUtils.normalizeFloatingNumbers(expr.children.head) + } else { + expr.children.head + } val childExpr = exprToProto(child, inputs, binding) val dataType = serializeDataType(expr.dataType) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometExecUtils.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometExecUtils.scala index 8145e563f49..a549cbe1ed5 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometExecUtils.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometExecUtils.scala @@ -24,7 +24,8 @@ import scala.reflect.ClassTag import org.apache.spark.{Partition, SparkContext, TaskContext} import org.apache.spark.rdd.RDD -import org.apache.spark.sql.catalyst.expressions.{Attribute, NamedExpression, SortOrder} +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, NamedExpression, SortOrder} +import org.apache.spark.sql.catalyst.optimizer.NormalizeFloatingNumbers import org.apache.spark.sql.comet.execution.arrow.CometArrowStream import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.SparkPlan @@ -36,6 +37,10 @@ import org.apache.comet.serde.QueryPlanSerde.{exprToProto, serializeDataType} object CometExecUtils { + /** Expose Spark's package-private recursive floating-point normalizer to Comet serde. */ + def normalizeFloatingNumbers(expr: Expression): Expression = + NormalizeFloatingNumbers.normalize(expr) + /** * Create an empty RDD with the given number of partitions. */ diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set.sql index cd528272af4..1fca2925de3 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set.sql @@ -15,7 +15,6 @@ -- specific language governing permissions and limitations -- under the License. --- Config: spark.comet.exec.strictFloatingPoint=true -- ConfigMatrix: parquet.enable.dictionary=false,true -- ============================================================ @@ -147,42 +146,6 @@ SELECT grp, sort_array(collect_set(i)) FROM cs_src_intbig GROUP BY grp ORDER BY query SELECT grp, sort_array(collect_set(bi)) FROM cs_src_intbig GROUP BY grp ORDER BY grp --- ============================================================ --- Float (with NULLs, NaN, Inf, -Inf, +0, -0) --- Comet deduplicates NaN while Spark does not; with --- strictFloatingPoint=true collect_set falls back to Spark. --- ============================================================ - -statement -CREATE TABLE cs_src_float(v float, grp string) USING parquet - -statement -INSERT INTO cs_src_float VALUES - (1.5, 'a'), (2.5, 'a'), (1.5, 'a'), (NULL, 'a'), - (CAST('NaN' AS FLOAT), 'b'), (CAST('NaN' AS FLOAT), 'b'), (1.0, 'b'), - (CAST('Infinity' AS FLOAT), 'c'), (CAST('-Infinity' AS FLOAT), 'c'), (CAST('Infinity' AS FLOAT), 'c'), - (CAST(0.0 AS FLOAT), 'd'), (CAST(-0.0 AS FLOAT), 'd'), (1.0, 'd'), (NULL, 'd') - -query expect_fallback(not fully compatible with Spark) -SELECT grp, sort_array(collect_set(v)) FROM cs_src_float GROUP BY grp ORDER BY grp - --- ============================================================ --- Double (with NULLs, NaN, Inf, -Inf, +0, -0) --- ============================================================ - -statement -CREATE TABLE cs_src_double(v double, grp string) USING parquet - -statement -INSERT INTO cs_src_double VALUES - (1.1, 'a'), (2.2, 'a'), (1.1, 'a'), (NULL, 'a'), - (CAST('NaN' AS DOUBLE), 'b'), (CAST('NaN' AS DOUBLE), 'b'), (1.0, 'b'), - (CAST('Infinity' AS DOUBLE), 'c'), (CAST('-Infinity' AS DOUBLE), 'c'), (CAST('Infinity' AS DOUBLE), 'c'), - (0.0, 'd'), (-0.0, 'd'), (1.0, 'd'), (NULL, 'd') - -query expect_fallback(not fully compatible with Spark) -SELECT grp, sort_array(collect_set(v)) FROM cs_src_double GROUP BY grp ORDER BY grp - -- ============================================================ -- String (with NULLs and empty string) -- ============================================================ diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_floating_fallback.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_floating_fallback.sql new file mode 100644 index 00000000000..a321cc323dc --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_floating_fallback.sql @@ -0,0 +1,49 @@ +-- 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. + +-- MaxSparkVersion: 4.1 +-- Config: spark.comet.exec.strictFloatingPoint=true + +statement +CREATE TABLE cs_fallback_float(v float, grp string) USING parquet + +statement +INSERT INTO cs_fallback_float VALUES + (1.5, 'a'), (2.5, 'a'), (1.5, 'a'), (NULL, 'a'), + (CAST('NaN' AS FLOAT), 'b'), (CAST('NaN' AS FLOAT), 'b'), (1.0, 'b'), + (CAST('Infinity' AS FLOAT), 'c'), (CAST('-Infinity' AS FLOAT), 'c'), + (CAST('Infinity' AS FLOAT), 'c'), + (CAST(0.0 AS FLOAT), 'd'), (CAST(-0.0 AS FLOAT), 'd'), (1.0, 'd'), (NULL, 'd') + +query expect_fallback(not fully compatible with Spark) +SELECT grp, sort_array(collect_set(v)) +FROM cs_fallback_float GROUP BY grp ORDER BY grp + +statement +CREATE TABLE cs_fallback_double(v double, grp string) USING parquet + +statement +INSERT INTO cs_fallback_double VALUES + (1.1, 'a'), (2.2, 'a'), (1.1, 'a'), (NULL, 'a'), + (CAST('NaN' AS DOUBLE), 'b'), (CAST('NaN' AS DOUBLE), 'b'), (1.0, 'b'), + (CAST('Infinity' AS DOUBLE), 'c'), (CAST('-Infinity' AS DOUBLE), 'c'), + (CAST('Infinity' AS DOUBLE), 'c'), + (0.0, 'd'), (-0.0, 'd'), (1.0, 'd'), (NULL, 'd') + +query expect_fallback(not fully compatible with Spark) +SELECT grp, sort_array(collect_set(v)) +FROM cs_fallback_double GROUP BY grp ORDER BY grp diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_normalization.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_normalization.sql new file mode 100644 index 00000000000..6350e9c0b2d --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_normalization.sql @@ -0,0 +1,64 @@ +-- 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. + +-- MinSparkVersion: 4.2 +-- Config: spark.comet.exec.strictFloatingPoint=true + +statement +CREATE TABLE cs_norm_scalar(grp string, f float, d double) USING parquet + +statement +INSERT INTO cs_norm_scalar VALUES + ('nan', CAST('NaN' AS FLOAT), CAST('NaN' AS DOUBLE)), + ('nan', CAST('NaN' AS FLOAT), CAST('NaN' AS DOUBLE)), + ('zero', CAST(-0.0 AS FLOAT), CAST(-0.0 AS DOUBLE)), + ('zero', CAST(0.0 AS FLOAT), CAST(0.0 AS DOUBLE)), + ('mixed', CAST('NaN' AS FLOAT), CAST('NaN' AS DOUBLE)), + ('mixed', CAST('NaN' AS FLOAT), CAST('NaN' AS DOUBLE)), + ('mixed', CAST(-0.0 AS FLOAT), CAST(-0.0 AS DOUBLE)), + ('mixed', CAST(0.0 AS FLOAT), CAST(0.0 AS DOUBLE)), + ('mixed', CAST(1.0 AS FLOAT), CAST(1.0 AS DOUBLE)) + +-- Repartition so equal values must also deduplicate across partial buffers. +query +SELECT grp, size(collect_set(f)), size(collect_set(d)) +FROM (SELECT /*+ REPARTITION(3) */ * FROM cs_norm_scalar) +GROUP BY grp +ORDER BY grp + +-- Spark 4.2 canonicalizes the surviving signed zero to positive zero. +query +SELECT collect_set(f), collect_set(d) FROM cs_norm_scalar WHERE grp = 'zero' + +statement +CREATE TABLE cs_norm_nested( + grp string, + s struct, + a array) USING parquet + +statement +INSERT INTO cs_norm_nested VALUES + ('nan', named_struct('v', CAST('NaN' AS DOUBLE)), array(CAST('NaN' AS FLOAT))), + ('nan', named_struct('v', CAST('NaN' AS DOUBLE)), array(CAST('NaN' AS FLOAT))), + ('zero', named_struct('v', CAST(-0.0 AS DOUBLE)), array(CAST(-0.0 AS FLOAT))), + ('zero', named_struct('v', CAST(0.0 AS DOUBLE)), array(CAST(0.0 AS FLOAT))) + +query +SELECT grp, size(collect_set(s)), size(collect_set(a)) +FROM (SELECT /*+ REPARTITION(3) */ * FROM cs_norm_nested) +GROUP BY grp +ORDER BY grp From 3c58054666ab4733416bac5ad82d4814555e990e Mon Sep 17 00:00:00 2001 From: peterxcli Date: Sun, 2 Aug 2026 11:24:40 +0800 Subject: [PATCH 2/4] fix: complete collect_set normalization handling --- .../org/apache/comet/serde/aggregates.scala | 30 ++++--- .../collect_set_floating_fallback.sql | 4 +- .../aggregate/collect_set_normalization.sql | 78 +++++++++++++++---- 3 files changed, 84 insertions(+), 28 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala index 0c236999706..e7fcacbfe1d 100644 --- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala +++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala @@ -22,7 +22,7 @@ package org.apache.comet.serde import scala.jdk.CollectionConverters._ import org.apache.spark.sql.catalyst.expressions.{Attribute, Cast, Literal} -import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, ApproximatePercentile, Average, BitAndAgg, BitOrAgg, BitXorAgg, BloomFilterAggregate, CentralMomentAgg, CollectList, CollectSet, Corr, Count, Covariance, CovPopulation, CovSample, First, HyperLogLogPlusPlus, Last, Max, Min, Partial, Percentile, StddevPop, StddevSamp, Sum, VariancePop, VarianceSamp} +import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, ApproximatePercentile, Average, BitAndAgg, BitOrAgg, BitXorAgg, BloomFilterAggregate, CentralMomentAgg, CollectList, CollectSet, Complete, Corr, Count, Covariance, CovPopulation, CovSample, First, HyperLogLogPlusPlus, Last, Max, Min, Partial, Percentile, StddevPop, StddevSamp, Sum, VariancePop, VarianceSamp} import org.apache.spark.sql.catalyst.util.ArrayData import org.apache.spark.sql.comet.CometExecUtils import org.apache.spark.sql.internal.SQLConf @@ -852,11 +852,19 @@ object CometBloomFilterAggregate extends CometAggregateExpressionSerde[BloomFilt object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { - override def getIncompatibleReasons(): Seq[String] = Seq( - "Before Spark 4.2, Comet deduplicates NaN values (treats `NaN == NaN`) while Spark treats" + - s" each NaN as a distinct value. When `${COMET_EXEC_STRICT_FLOATING_POINT.key}=true`," + - " `collect_set` on floating-point types falls back to Spark on those versions unless" + - " `spark.comet.expression.CollectSet.allowIncompatible=true` is set.") + override def getIncompatibleReasons(): Seq[String] = { + if (isSpark42Plus) { + Nil + } else { + Seq( + "Before Spark 4.2, Comet deduplicates NaN values (treats `NaN == NaN`) while Spark" + + " treats each NaN as a distinct value. Comet treats -0.0 and 0.0 as distinct while" + + " Spark treats them as equal." + + s" When `${COMET_EXEC_STRICT_FLOATING_POINT.key}=true`, `collect_set` on" + + " floating-point types falls back to Spark on those versions unless" + + " `spark.comet.expression.CollectSet.allowIncompatible=true` is set.") + } + } override def getSupportLevel(expr: CollectSet): SupportLevel = { // The native path always drops null inputs. Spark 4.2 added an `ignoreNulls` field to @@ -873,7 +881,7 @@ object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { .strictFloatingPointReason( expr.children.head.dataType, "collect_set on floating-point types " + - "(Comet deduplicates NaN values while Spark treats each NaN as distinct)") + "(Comet deduplicates NaN values and distinguishes -0.0 from 0.0, unlike Spark)") .map(reason => Incompatible(Some(reason))) .getOrElse(Compatible()) } @@ -885,12 +893,12 @@ object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { inputs: Seq[Attribute], binding: Boolean, conf: SQLConf): Option[ExprOuterClass.AggExpr] = { - val child = - if (isSpark42Plus && aggExpr.mode == Partial) { + val child = aggExpr.mode match { + case Partial | Complete if isSpark42Plus => CometExecUtils.normalizeFloatingNumbers(expr.children.head) - } else { + case _ => expr.children.head - } + } val childExpr = exprToProto(child, inputs, binding) val dataType = serializeDataType(expr.dataType) diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_floating_fallback.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_floating_fallback.sql index a321cc323dc..55bdf13c271 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_floating_fallback.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_floating_fallback.sql @@ -27,7 +27,7 @@ INSERT INTO cs_fallback_float VALUES (CAST('NaN' AS FLOAT), 'b'), (CAST('NaN' AS FLOAT), 'b'), (1.0, 'b'), (CAST('Infinity' AS FLOAT), 'c'), (CAST('-Infinity' AS FLOAT), 'c'), (CAST('Infinity' AS FLOAT), 'c'), - (CAST(0.0 AS FLOAT), 'd'), (CAST(-0.0 AS FLOAT), 'd'), (1.0, 'd'), (NULL, 'd') + (CAST(0.0 AS FLOAT), 'd'), (CAST('-0.0' AS FLOAT), 'd'), (1.0, 'd'), (NULL, 'd') query expect_fallback(not fully compatible with Spark) SELECT grp, sort_array(collect_set(v)) @@ -42,7 +42,7 @@ INSERT INTO cs_fallback_double VALUES (CAST('NaN' AS DOUBLE), 'b'), (CAST('NaN' AS DOUBLE), 'b'), (1.0, 'b'), (CAST('Infinity' AS DOUBLE), 'c'), (CAST('-Infinity' AS DOUBLE), 'c'), (CAST('Infinity' AS DOUBLE), 'c'), - (0.0, 'd'), (-0.0, 'd'), (1.0, 'd'), (NULL, 'd') + (0.0, 'd'), (CAST('-0.0' AS DOUBLE), 'd'), (1.0, 'd'), (NULL, 'd') query expect_fallback(not fully compatible with Spark) SELECT grp, sort_array(collect_set(v)) diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_normalization.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_normalization.sql index 6350e9c0b2d..1f0ce012378 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_normalization.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/collect_set_normalization.sql @@ -17,48 +17,96 @@ -- MinSparkVersion: 4.2 -- Config: spark.comet.exec.strictFloatingPoint=true +-- ConfigMatrix: parquet.enable.dictionary=false,true statement CREATE TABLE cs_norm_scalar(grp string, f float, d double) USING parquet statement INSERT INTO cs_norm_scalar VALUES + ('ordinary', CAST(1.5 AS FLOAT), CAST(1.1 AS DOUBLE)), + ('ordinary', CAST(2.5 AS FLOAT), CAST(2.2 AS DOUBLE)), + ('ordinary', CAST(1.5 AS FLOAT), CAST(1.1 AS DOUBLE)), + ('ordinary', CAST(NULL AS FLOAT), CAST(NULL AS DOUBLE)), ('nan', CAST('NaN' AS FLOAT), CAST('NaN' AS DOUBLE)), ('nan', CAST('NaN' AS FLOAT), CAST('NaN' AS DOUBLE)), - ('zero', CAST(-0.0 AS FLOAT), CAST(-0.0 AS DOUBLE)), + ('nan', CAST(1.0 AS FLOAT), CAST(1.0 AS DOUBLE)), + ('infinity', CAST('Infinity' AS FLOAT), CAST('Infinity' AS DOUBLE)), + ('infinity', CAST('-Infinity' AS FLOAT), CAST('-Infinity' AS DOUBLE)), + ('infinity', CAST('Infinity' AS FLOAT), CAST('Infinity' AS DOUBLE)), + ('zero', CAST('-0.0' AS FLOAT), CAST('-0.0' AS DOUBLE)), ('zero', CAST(0.0 AS FLOAT), CAST(0.0 AS DOUBLE)), - ('mixed', CAST('NaN' AS FLOAT), CAST('NaN' AS DOUBLE)), - ('mixed', CAST('NaN' AS FLOAT), CAST('NaN' AS DOUBLE)), - ('mixed', CAST(-0.0 AS FLOAT), CAST(-0.0 AS DOUBLE)), - ('mixed', CAST(0.0 AS FLOAT), CAST(0.0 AS DOUBLE)), - ('mixed', CAST(1.0 AS FLOAT), CAST(1.0 AS DOUBLE)) + ('zero', CAST(1.0 AS FLOAT), CAST(1.0 AS DOUBLE)), + ('zero', CAST(NULL AS FLOAT), CAST(NULL AS DOUBLE)) -- Repartition so equal values must also deduplicate across partial buffers. query -SELECT grp, size(collect_set(f)), size(collect_set(d)) +SELECT grp, sort_array(collect_set(f)), sort_array(collect_set(d)) FROM (SELECT /*+ REPARTITION(3) */ * FROM cs_norm_scalar) GROUP BY grp ORDER BY grp --- Spark 4.2 canonicalizes the surviving signed zero to positive zero. +-- Exercise the no-GROUP BY aggregate shape while merging partial buffers. query -SELECT collect_set(f), collect_set(d) FROM cs_norm_scalar WHERE grp = 'zero' +SELECT sort_array(collect_set(f)), sort_array(collect_set(d)) +FROM ( + SELECT /*+ REPARTITION(3) */ f, d + FROM cs_norm_scalar + WHERE grp IN ('nan', 'zero') +) statement CREATE TABLE cs_norm_nested( grp string, s struct, - a array) USING parquet + a array, + deep_a array>, + deep_s struct>) USING parquet statement INSERT INTO cs_norm_nested VALUES - ('nan', named_struct('v', CAST('NaN' AS DOUBLE)), array(CAST('NaN' AS FLOAT))), - ('nan', named_struct('v', CAST('NaN' AS DOUBLE)), array(CAST('NaN' AS FLOAT))), - ('zero', named_struct('v', CAST(-0.0 AS DOUBLE)), array(CAST(-0.0 AS FLOAT))), - ('zero', named_struct('v', CAST(0.0 AS DOUBLE)), array(CAST(0.0 AS FLOAT))) + ('nan', + named_struct('v', CAST('NaN' AS DOUBLE)), + array(CAST('NaN' AS FLOAT)), + array(named_struct('v', CAST('NaN' AS DOUBLE))), + named_struct('a', array(CAST('NaN' AS DOUBLE)))), + ('nan', + named_struct('v', CAST('NaN' AS DOUBLE)), + array(CAST('NaN' AS FLOAT)), + array(named_struct('v', CAST('NaN' AS DOUBLE))), + named_struct('a', array(CAST('NaN' AS DOUBLE)))), + ('zero', + named_struct('v', CAST('-0.0' AS DOUBLE)), + array(CAST('-0.0' AS FLOAT)), + array(named_struct('v', CAST('-0.0' AS DOUBLE))), + named_struct('a', array(CAST('-0.0' AS DOUBLE)))), + ('zero', + named_struct('v', CAST(0.0 AS DOUBLE)), + array(CAST(0.0 AS FLOAT)), + array(named_struct('v', CAST(0.0 AS DOUBLE))), + named_struct('a', array(CAST(0.0 AS DOUBLE)))), + ('null', + CAST(NULL AS STRUCT), + CAST(NULL AS ARRAY), + CAST(NULL AS ARRAY>), + CAST(NULL AS STRUCT>)), + ('null', + named_struct('v', CAST(NULL AS DOUBLE)), + array(CAST(NULL AS FLOAT)), + array(CAST(NULL AS STRUCT)), + named_struct('a', array(CAST(NULL AS DOUBLE)))), + ('null', + named_struct('v', CAST(NULL AS DOUBLE)), + array(CAST(NULL AS FLOAT)), + array(CAST(NULL AS STRUCT)), + named_struct('a', array(CAST(NULL AS DOUBLE)))) query -SELECT grp, size(collect_set(s)), size(collect_set(a)) +SELECT grp, + sort_array(collect_set(s)), + sort_array(collect_set(a)), + sort_array(collect_set(deep_a)), + sort_array(collect_set(deep_s)) FROM (SELECT /*+ REPARTITION(3) */ * FROM cs_norm_nested) GROUP BY grp ORDER BY grp From 014c23ee6a86fd4f97c4cfa5ed99ea2112b9b083 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Sun, 9 Aug 2026 01:25:31 +0800 Subject: [PATCH 3/4] add unsupport reason and explain why adding the spark 4.2 gate --- .../scala/org/apache/comet/serde/aggregates.scala | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala index 778c02391fd..49424046990 100644 --- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala +++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala @@ -845,6 +845,15 @@ object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { } } + override def getUnsupportedReasons(): Seq[String] = + if (isSpark42Plus) { + Seq( + "`collect_set` with `RESPECT NULLS` falls back to Spark, since the native " + + "implementation always drops null inputs.") + } else { + Nil + } + override def getSupportLevel(expr: CollectSet): SupportLevel = { // The native path always drops null inputs. Spark 4.2 added an `ignoreNulls` field to // CollectSet that `RESPECT NULLS` sets to false, preserving nulls in the result; Comet @@ -873,6 +882,9 @@ object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { binding: Boolean, conf: SQLConf): Option[ExprOuterClass.AggExpr] = { val child = aggExpr.mode match { + // Spark 4.2 introduced this normalization. Keep older versions unchanged to avoid adding + // JVM codegen dispatch for nested arrays; their floating-point behavior is documented as + // incompatible. case Partial | Complete if isSpark42Plus => CometExecUtils.normalizeFloatingNumbers(expr.children.head) case _ => From c9b5f608ddb505d2baba545f5238f177ca6cf413 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Fri, 28 Aug 2026 03:29:16 +0800 Subject: [PATCH 4/4] docs: address review feedback on normalization visibility and Spark-internal risk - Document the array-input JVM codegen dispatch trade-off in getCompatibleNotes() so it appears in the generated Compatibility Guide. - Reference SPARK-57298 in the serde comment and explain why double normalization cannot occur (normalization lives in CollectSet's buffer conversion, not the plan; needNormalize short-circuits on KnownFloatingPointNormalized). - Expand the CometExecUtils.normalizeFloatingNumbers doc comment with the compatibility risk of the private[sql] API and the compile-time failure signal across Spark profiles. Co-Authored-By: Claude Fable 5 --- .../org/apache/comet/serde/aggregates.scala | 22 ++++++++++++++++--- .../spark/sql/comet/CometExecUtils.scala | 10 ++++++++- 2 files changed, 28 insertions(+), 4 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala index 49424046990..a87f12bb4db 100644 --- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala +++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala @@ -831,6 +831,19 @@ object CometBloomFilterAggregate extends CometAggregateExpressionSerde[BloomFilt object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { + override def getCompatibleNotes(): Seq[String] = + if (isSpark42Plus) { + Seq( + "On Spark 4.2+, `collect_set` inputs are normalized with Spark's recursive" + + " floating-point normalizer to match SPARK-57298. For array inputs containing" + + " floating-point values the normalizer produces `ArrayTransform`, which Comet" + + " executes through the JVM codegen dispatcher instead of fully natively (scalar and" + + " struct inputs stay native). When `spark.comet.exec.scalaUDF.codegen.enabled=false`," + + " those array cases fall back to Spark.") + } else { + Nil + } + override def getIncompatibleReasons(): Seq[String] = { if (isSpark42Plus) { Nil @@ -882,9 +895,12 @@ object CometCollectSet extends CometAggregateExpressionSerde[CollectSet] { binding: Boolean, conf: SQLConf): Option[ExprOuterClass.AggExpr] = { val child = aggExpr.mode match { - // Spark 4.2 introduced this normalization. Keep older versions unchanged to avoid adding - // JVM codegen dispatch for nested arrays; their floating-point behavior is documented as - // incompatible. + // Spark 4.2 (SPARK-57298) normalizes NaN and -0.0 inside CollectSet's buffer conversion + // (`convertToBufferElement`/`eval` in collect.scala), so the normalization is invisible in + // the plan Comet receives and must be reapplied here on the raw input. `normalize` is + // idempotent: it short-circuits on `KnownFloatingPointNormalized`, so an already-normalized + // child is never wrapped twice. Keep older versions unchanged to avoid adding JVM codegen + // dispatch for nested arrays; their floating-point behavior is documented as incompatible. case Partial | Complete if isSpark42Plus => CometExecUtils.normalizeFloatingNumbers(expr.children.head) case _ => diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometExecUtils.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometExecUtils.scala index a549cbe1ed5..07c44fd195b 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometExecUtils.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometExecUtils.scala @@ -37,7 +37,15 @@ import org.apache.comet.serde.QueryPlanSerde.{exprToProto, serializeDataType} object CometExecUtils { - /** Expose Spark's package-private recursive floating-point normalizer to Comet serde. */ + /** + * Expose Spark's package-private recursive floating-point normalizer to Comet serde. + * + * Compatibility note: `NormalizeFloatingNumbers.normalize` is `private[sql]` with no stability + * guarantee. It has existed with this signature in every Spark version Comet supports (3.4 + * through 4.2). Because this is a direct compile-time reference (not reflection), any rename or + * signature change in a future Spark version fails the build for that profile loudly; if that + * happens, shim this method per Spark version like other `Shim*` classes. + */ def normalizeFloatingNumbers(expr: Expression): Expression = NormalizeFloatingNumbers.normalize(expr)