sunchao commented on code in PR #6563:
URL: https://github.com/apache/datafusion-comet/pull/6563#discussion_r4200038527
##########
spark/src/main/scala/org/apache/comet/serde/arrays.scala:
##########
@@ -540,53 +540,143 @@ object CometSlice extends CometExpressionSerde[Slice] {
private[comet] object ArraySetSupport {
val floatingPointReason: String =
- "Floating-point elements match Spark's signed-zero and NaN semantics
natively only on " +
- "Spark 4.2.0, whose optimizer normalizes the arguments (SPARK-54918)"
-
- // The native kernels match Spark only when the plan has already normalized
the arguments, and
- // only Spark 4.2.0 does that (SPARK-54918). Earlier releases keep flat
signed zeros apart.
- // From 4.0.5, 4.1.4 and 4.2.1, SPARK-59602 normalizes during evaluation
instead, which the
- // native kernels do not match for NaN payloads or nested zeros. A top-level
- // KnownFloatingPointNormalized marker cannot replace the version check:
Spark also normalizes
- // CreateArray, If, CaseWhen, and Coalesce recursively without wrapping the
resulting array.
- def normalizesArgumentsInPlan(version: String): Boolean =
- Utils.majorMinorPatchVersion(version).contains((4, 2, 0))
+ "Floating-point elements match Spark's signed-zero semantics natively only
on Spark " +
+ "4.0.5+, 4.1.4+ and 4.2+, which treat -0.0 and 0.0 as one value in these
functions " +
+ "(SPARK-54918, SPARK-59602)"
+
+ val collationReason: String =
+ "Elements that hold both a floating-point value and a non-UTF8_BINARY
collated string fall " +
+ "back to Spark, which compares the strings under their collation, while
Comet's native " +
+ "kernels compare their raw bytes"
+
+ // Spark 4.2.0 normalizes the arguments of these functions in the plan
(SPARK-54918), and 4.0.5,
+ // 4.1.4 and 4.2.1 normalize while evaluating them (SPARK-59602). Either
way, Spark treats -0.0
+ // and 0.0, and every NaN, as one value at any depth, which the spark_
variants match. Earlier
+ // releases keep -0.0 and 0.0 apart in a flat array. The check reads the
version rather than a
+ // KnownFloatingPointNormalized marker, because SPARK-59602 adds no marker,
and SPARK-54918
+ // normalizes CreateArray, If, CaseWhen and Coalesce without wrapping the
resulting array.
+ def normalizesFloats(version: String): Boolean =
+ Utils.majorMinorPatchVersion(version).exists {
+ case (4, 0, patch) => patch >= 5
+ case (4, 1, patch) => patch >= 4
+ case (major, minor, _) => major > 4 || (major == 4 && minor >= 2)
+ }
def supportLevel(dataType: DataType): SupportLevel = {
- if (SupportLevel.containsType(dataType, classOf[FloatType],
classOf[DoubleType]) &&
- !normalizesArgumentsInPlan(SPARK_VERSION)) {
+ if (hasFloats(dataType) && hasNonDefaultStringCollation(dataType)) {
+ // The spark_ variants normalize the floats, then DataFusion compares
the elements, strings
+ // included, by their bytes. Collated strings without floats take the
plain DataFusion
+ // functions: https://github.com/apache/datafusion-comet/issues/6470.
+ Incompatible(Some(collationReason))
+ } else if (hasFloats(dataType) && !normalizesFloats(SPARK_VERSION)) {
Incompatible(Some(floatingPointReason))
} else {
Compatible()
}
}
+
+ // DataFusion folds -0.0 into 0.0 only in a flat float array and compares
NaNs by their bits.
+ // The spark_ variants normalize floats at any depth first, as Spark does.
+ def function(name: String, dataType: DataType): String =
+ if (hasFloats(dataType)) s"spark_$name" else name
+
+ private def hasFloats(dataType: DataType): Boolean =
+ SupportLevel.containsType(dataType, classOf[FloatType],
classOf[DoubleType])
}
// Use projection fallback to avoid codegen dispatch overhead for array-valued
results.
// The native implementation remains available through opt-in.
-object CometArrayDistinct extends
CometScalarFunction[ArrayDistinct]("array_distinct") {
- override def getIncompatibleReasons(): Seq[String] =
Seq(ArraySetSupport.floatingPointReason)
+object CometArrayDistinct extends CometExpressionSerde[ArrayDistinct] {
+ override def getIncompatibleReasons(): Seq[String] =
+ Seq(ArraySetSupport.floatingPointReason, ArraySetSupport.collationReason)
override def getSupportLevel(expr: ArrayDistinct): SupportLevel =
ArraySetSupport.supportLevel(expr.dataType)
+
+ override def convert(
+ expr: ArrayDistinct,
+ inputs: Seq[Attribute],
+ binding: Boolean): Option[ExprOuterClass.Expr] = {
+ val childProto = exprToProtoInternal(expr.child, inputs, binding)
+ scalarFunctionExprToProto(
+ ArraySetSupport.function("array_distinct", expr.dataType),
+ childProto)
+ }
}
object CometArrayUnion extends CometExpressionSerde[ArrayUnion] {
- override def getIncompatibleReasons(): Seq[String] =
Seq(ArraySetSupport.floatingPointReason)
+
+ /**
+ * Spark's `ArrayUnion` is a `BinaryExpression`: for a NULL left array it
returns NULL without
+ * evaluating the right operand. The native function evaluates both operands
over the whole
+ * batch first, so a right operand that throws, such as `slice(b, 0, 1)` or
an ANSI cast, fails
+ * on rows Spark never evaluates it for. `convert` reproduces the
short-circuit with a `CASE
+ * WHEN <left> IS NOT NULL` guard, as `CometElementAt` does. A column or
literal cannot throw,
+ * so it needs no guard. The guard serializes the left operand twice, which
a stateful operand
+ * cannot survive, so that shape stays on Spark. See
+ * https://github.com/apache/datafusion-comet/issues/6613.
+ */
+ private val eagerRightOperandReason: String =
+ "a nullable nondeterministic left operand: native array_union evaluates
the right operand " +
+ "over the whole batch, where Spark skips it on the rows whose left
operand is NULL"
+
+ /** True when `convert` has to guard the call to reproduce Spark's NULL
short-circuit. */
+ private def needsNullGuard(expr: ArrayUnion): Boolean =
+ expr.left.nullable && !expr.right.isInstanceOf[Attribute] &&
!expr.right.isInstanceOf[Literal]
+
+ override def getIncompatibleReasons(): Seq[String] =
+ Seq(ArraySetSupport.floatingPointReason, ArraySetSupport.collationReason)
+
+ override def getUnsupportedReasons(): Seq[String] =
Seq(eagerRightOperandReason)
override def getSupportLevel(expr: ArrayUnion): SupportLevel =
- ArraySetSupport.supportLevel(expr.dataType)
+ if (needsNullGuard(expr) && !expr.left.deterministic) {
+ Unsupported(Some(eagerRightOperandReason))
+ } else {
+ ArraySetSupport.supportLevel(expr.dataType)
+ }
override def convert(
expr: ArrayUnion,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val leftArrayExprProto = exprToProtoInternal(expr.children.head, inputs,
binding)
- val rightArrayExprProto = exprToProtoInternal(expr.children(1), inputs,
binding)
+ val leftArrayExprProto = exprToProtoInternal(expr.left, inputs, binding)
+ val rightArrayExprProto = exprToProtoInternal(expr.right, inputs, binding)
val arraysUnionScalarExpr =
- scalarFunctionExprToProto("array_union", leftArrayExprProto,
rightArrayExprProto)
- arraysUnionScalarExpr
+ scalarFunctionExprToProto(
+ ArraySetSupport.function("array_union", expr.dataType),
+ leftArrayExprProto,
+ rightArrayExprProto)
+ if (!needsNullGuard(expr)) {
+ arraysUnionScalarExpr
+ } else {
+ // DataFusion's CaseExpr evaluates the THEN branch only on the rows the
guard selects.
+ val isNotNullExpr = createUnaryExpr(
Review Comment:
[P2] [P2] Avoid duplicating the entire left subtree in each null guard. With
nullable `ARRAY<INT>` columns, build `e0 = a` and `e(n+1) = array_union(e(n),
slice(b, 1, 1))`. This second serialization of `expr.left` expands 12 nested
unions into 4,095 union nodes, versus 12 at the base revision. On a non-null
row, the guard evaluates the subtree and the THEN branch evaluates it again,
producing the same exponential execution count. Results remain `[1,2]` for
`a=[1]`, `b=[2]`, but previously native queries incur substantially larger
plans and computation. Please evaluate or materialize the left value once and
reuse it for both the null check and union.
Evidence: A bounded harness compiling byte-for-byte verified base/head
`CometArrayUnion` bodies with protobuf builders produced 4/8/12 union nodes at
base versus 15/255/4095 at head. Spark 4.1.3 optimization retained the original
4/8/12 nullable-left unions, and interpreted plus generated evaluation returned
`[1,2]`. A freshly compiled DataFusion 55.1.0 physical-expression probe using
the actual slice implementation confirmed 12 versus 4,095 union invocations for
one non-null row. Comet’s array-valued CASE path delegates to this lazy
`CaseExpr` behavior. This is a split serde/Spark/native reproduction, not a
full Comet integration run.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]