comphead commented on code in PR #6455:
URL: https://github.com/apache/datafusion-comet/pull/6455#discussion_r4167088404
##########
spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala:
##########
@@ -1141,6 +1159,48 @@ object QueryPlanSerde extends Logging with CometExprShim
with CometTypeShim {
case _ => false
}
+ /**
+ * Whether `expr` consumes a decimal result of a dispatched DSv2 scalar
function anywhere in its
+ * argument trees, including through intermediate expressions and container
access.
+ *
+ * Spark does not rescale such a result to the type the function declares,
or write null when it
+ * does not fit, until it writes a row. An expression around the call reads
the `Decimal` the
+ * function returned. The dispatcher has to write an Arrow vector of the
declared type, so it
+ * rescales and nulls at its own output, and a native expression over that
output would read
+ * something else. `IS NULL` of a value that does not fit is false in Spark,
a cast to string
+ * keeps the function's scale, and `hash` reads the unscaled value at that
scale (#6425). So
+ * such an expression runs in the same kernel as the call, where Spark's own
code reads the
+ * value the function returned. Checking only immediate children would let
an intermediate
+ * expression, such as `abs(call)` or `call[0]`, normalize the decimal
before its parent reads
+ * it. `Alias` is skipped because it computes nothing: the call under it is
the root, and Spark
+ * writes a root as a row.
+ */
+ private def readsDispatchedDsv2Decimal(expr: Expression): Boolean =
+ !isStructuralExpr(expr) &&
expr.children.exists(_.exists(isDispatchedDsv2DecimalCall))
Review Comment:
Would it be enough to follow only decimal-typed children here? A node whose
type has no decimal, such as `IS NULL`, a cast to string or `hash`, writes a
value that Arrow holds exactly, so the kernel can end there. As written,
`upper(name) = 'X' AND decfn.ns.as_money(i) IS NULL` should dispatch the whole
`AND` including `upper`, and `sum(CASE WHEN decfn.ns.as_money(i) IS NULL THEN 1
ELSE 0 END)` should fall back even though the aggregate only sees an `INT`. A
sibling the dispatcher declines would also take the operator back to Spark. The
scan also walks the full subtree at every level of the top-down conversion,
which a decimal-only walk would bound. Something like
`isDispatchedDsv2DecimalCall(c) || (SupportLevel.containsType(c.dataType,
classOf[DecimalType]) && readsDispatchedDsv2Decimal(c))` per child, shared with
the aggregate check, is what I have in mind. I have not run it, so this is an
expectation. A test with one of those shapes would settle it.
##########
spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala:
##########
@@ -1141,6 +1159,48 @@ object QueryPlanSerde extends Logging with CometExprShim
with CometTypeShim {
case _ => false
}
+ /**
+ * Whether `expr` consumes a decimal result of a dispatched DSv2 scalar
function anywhere in its
+ * argument trees, including through intermediate expressions and container
access.
+ *
+ * Spark does not rescale such a result to the type the function declares,
or write null when it
+ * does not fit, until it writes a row. An expression around the call reads
the `Decimal` the
+ * function returned. The dispatcher has to write an Arrow vector of the
declared type, so it
+ * rescales and nulls at its own output, and a native expression over that
output would read
+ * something else. `IS NULL` of a value that does not fit is false in Spark,
a cast to string
+ * keeps the function's scale, and `hash` reads the unscaled value at that
scale (#6425). So
+ * such an expression runs in the same kernel as the call, where Spark's own
code reads the
+ * value the function returned. Checking only immediate children would let
an intermediate
+ * expression, such as `abs(call)` or `call[0]`, normalize the decimal
before its parent reads
+ * it. `Alias` is skipped because it computes nothing: the call under it is
the root, and Spark
+ * writes a root as a row.
+ */
+ private def readsDispatchedDsv2Decimal(expr: Expression): Boolean =
+ !isStructuralExpr(expr) &&
expr.children.exists(_.exists(isDispatchedDsv2DecimalCall))
+
+ private def isDispatchedDsv2DecimalCall(expr: Expression): Boolean = {
+ val dispatchedDsv2Call = expr match {
+ case i: Invoke =>
+ i.targetObject match {
+ case Literal(_: ScalarFunction[_], _) => true
+ case _ => false
+ }
+ case s: StaticInvoke =>
+ classOf[ScalarFunction[_]].isAssignableFrom(s.staticObject) &&
+ CometStaticInvoke.runsInDispatcher(s)
+ case _ => false
+ }
+ dispatchedDsv2Call && containsDecimal(expr.dataType)
+ }
+
+ private def containsDecimal(dataType: DataType): Boolean = dataType match {
+ case _: DecimalType => true
+ case ArrayType(elementType, _) => containsDecimal(elementType)
+ case MapType(keyType, valueType, _) => containsDecimal(keyType) ||
containsDecimal(valueType)
+ case StructType(fields) => fields.exists(f => containsDecimal(f.dataType))
+ case _ => false
+ }
Review Comment:
`SupportLevel.containsType(dt, classOf[DecimalType])` already walks array
elements, struct fields and map keys and values at every level, so
`containsDecimal` looks like a copy of it. Would it work to call that from
`isDispatchedDsv2DecimalCall` and drop this helper?
##########
spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala:
##########
@@ -2354,6 +2361,227 @@ class CometCodegenSuite
Invoke(target, "twice", StringType,
Seq(Literal(UTF8String.fromString("ab"), StringType)))
assert(runKernel(folded, 1)(_.getUTF8String(0).toString) === "abab")
}
+
+ /**
+ * Runs `f` with [[CometCodegenSuite.DecimalFunctionCatalog]] registered as
`decfn` and `values`
+ * in `t (i INT)`, for the #6425 tests.
+ */
+ private def withDecimalFunctions(values: Any*)(f: => Unit): Unit = {
+ withSQLConf(
+ "spark.sql.catalog.decfn" ->
classOf[CometCodegenSuite.DecimalFunctionCatalog].getName) {
+ withTable("t") {
+ sql("CREATE TABLE t (i INT) USING parquet")
+ // One file, so the kernel sees every row in one batch.
+ sql(
+ "INSERT INTO t SELECT /*+ REPARTITION(1) */ * FROM VALUES " +
+ values.map(v => s"($v)").mkString(", ") + " AS v(i)")
+ f
+ }
+ }
+ }
+
+ private def dec(s: String) = if (s == null) null else new
java.math.BigDecimal(s)
+
+ test("decimal results of a DSv2 function are rescaled to the declared type
(#6425)") {
+ // Spark lowers a call to a DSv2 function with an instance `invoke` method
to `Invoke`, and one
+ // with a static `invoke` to `StaticInvoke`. The dispatcher runs both.
`as_money` and
+ // `as_wide_money` return `Decimal(i)` at scale 0, one declaring
`DECIMAL(10, 2)` and one
+ // `DECIMAL(20, 12)`, which covers both of the dispatcher's decimal
writers. `static_as_money`
+ // is `as_money` with a static `invoke`. Spark's row writer rescales the
value with
+ // `changePrecision` and writes null when it does not fit: 100000000 and
-100000000 have nine
+ // integer digits and both types allow eight. Spark adds no overflow check
around the call, so
+ // that null does not depend on ANSI mode. `map` is itself dispatched, so
its value goes
+ // through the kernel's map writer. `mills_as_money` returns `i`
thousandths, at scale 3, into
+ // `DECIMAL(7, 2)`, so the rescale drops a digit and `changePrecision`
rounds half up: -1.005
+ // becomes -1.01 and 1.004 becomes 1.00. 99999.999 has the five integer
digits the type allows,
+ // but rounds up to 100000.00, which has six, so it is null.
+ //
+ // Each case is `(i, as_money, as_wide_money, mills_as_money)`.
`static_as_money` and the `map`
+ // value match `as_money`.
+ val cases = Seq[(Any, String, String, String)](
+ (3, "3.00", "3.000000000000", "0.00"),
+ (-7, "-7.00", "-7.000000000000", "-0.01"),
+ (null, null, null, null),
+ (99999999, "99999999.00", "99999999.000000000000", null),
+ (100000000, null, null, null),
+ (-99999999, "-99999999.00", "-99999999.000000000000", null),
+ (-100000000, null, null, null),
+ (-1005, "-1005.00", "-1005.000000000000", "-1.01"),
+ (1004, "1004.00", "1004.000000000000", "1.00"),
+ (5, "5.00", "5.000000000000", "0.01"))
+ val expected = cases.map { case (i, money, wide, mills) =>
+ Row(i, dec(money), dec(money), dec(wide), Map("k" -> dec(money)),
dec(mills))
+ }
+ withDecimalFunctions(cases.map(_._1): _*) {
+ for (ansi <- Seq("true", "false")) {
+ withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi) {
+ val df = sql(
+ "SELECT i, decfn.ns.as_money(i), decfn.ns.static_as_money(i), " +
+ "decfn.ns.as_wide_money(i), map('k', decfn.ns.as_money(i)), " +
+ "decfn.ns.mills_as_money(i) FROM t")
+ assertCodegenRan {
+ checkSparkAnswerAndImpl(df, dispatched = Seq("invoke",
"staticinvoke"))
+ }
+ checkAnswer(df, expected)
+ }
+ }
+ }
+ }
+
+ test("decimals in a DSv2 function's array and struct results are rescaled
(#6425)") {
+ // Each function returns `Decimal(i)`, at scale 0, in every decimal of its
result, so the
+ // kernel's array and struct writers have to rescale them, as Spark's
`UnsafeArrayWriter` and
+ // `UnsafeRowWriter` do. `money_array`'s element and `money_struct`'s `m`
field are
+ // `DECIMAL(10, 2)`, so both are null at 100000000. The writers skip the
null check for a
+ // non-nullable child, so `non_null_money_array`'s element and
`money_struct`'s `non_null_m`
+ // field cover that path. They are `DECIMAL(12, 2)`, which holds any `INT`.
+ //
+ // `array(...)` or `named_struct(...)` around a scalar call would not
reach these writers:
+ // Comet evaluates both natively, and dispatches only the call.
Review Comment:
This comment looks stale. Since `readsDispatchedDsv2Decimal` runs a consumer
in the call's kernel, `array(decfn.ns.as_money(i))` and `named_struct('m',
decfn.ns.as_money(i))` should now be dispatched whole and reach these writers,
instead of being evaluated natively with only the call dispatched. If that is
right, the comment could say so, and one query with each shape would cover the
new route, which I could not find in the tests.
##########
spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala:
##########
@@ -2354,6 +2361,227 @@ class CometCodegenSuite
Invoke(target, "twice", StringType,
Seq(Literal(UTF8String.fromString("ab"), StringType)))
assert(runKernel(folded, 1)(_.getUTF8String(0).toString) === "abab")
}
+
+ /**
+ * Runs `f` with [[CometCodegenSuite.DecimalFunctionCatalog]] registered as
`decfn` and `values`
+ * in `t (i INT)`, for the #6425 tests.
+ */
+ private def withDecimalFunctions(values: Any*)(f: => Unit): Unit = {
+ withSQLConf(
+ "spark.sql.catalog.decfn" ->
classOf[CometCodegenSuite.DecimalFunctionCatalog].getName) {
+ withTable("t") {
+ sql("CREATE TABLE t (i INT) USING parquet")
+ // One file, so the kernel sees every row in one batch.
+ sql(
+ "INSERT INTO t SELECT /*+ REPARTITION(1) */ * FROM VALUES " +
+ values.map(v => s"($v)").mkString(", ") + " AS v(i)")
+ f
+ }
+ }
+ }
+
+ private def dec(s: String) = if (s == null) null else new
java.math.BigDecimal(s)
+
+ test("decimal results of a DSv2 function are rescaled to the declared type
(#6425)") {
Review Comment:
Most of these tests only need a catalog class, so they might fit a SQL file
under `spark/src/test/resources/sql-tests/expressions/` instead of Scala. `--
Config: spark.sql.catalog.decfn=<fixture class>` registers a catalog
(`iceberg/metadata_column_partition.sql` does), `-- ConfigMatrix:
spark.sql.ansi.enabled=true,false` replaces the ANSI loop, `query
expect_dispatch(invoke, staticinvoke)` and `query expect_fallback(aggregates
the decimal result of a DSv2 function)` cover the dispatcher claims, and a
`CREATE TABLE ... USING parquet AS SELECT` statement covers the write boundary
(`concat.sql` has one). Those modes already compare with Spark, so the explicit
`Row(...)` lists would not be needed. The expression, transitive and aggregate
tests use the same four rows, so they could share one file and one setup. Only
the non-nullable test needs to stay in Scala, because the two engines' error
messages differ. The fixture classes would move out of the `CometCodegenSuite`
companion so th
e file can name them.
##########
spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala:
##########
@@ -2373,6 +2601,86 @@ object CometCodegenSuite {
class NotSerializableTarget {
def twice(s: UTF8String): UTF8String = UTF8String.fromString(s.toString +
s.toString)
}
+
+ /**
+ * DSv2 function catalog for the #6425 tests. Each function returns its
argument as a `Decimal`
+ * whose scale need not match the type it declares. `mills_as_money` returns
its argument as
+ * thousandths, and the rest at scale 0.
+ */
+ class DecimalFunctionCatalog extends FunctionCatalog {
+ private val money = DecimalType(10, 2)
+ // Holds any `INT`.
+ private val intMoney = DecimalType(12, 2)
+ private val functions: Map[String, UnboundFunction] = Map(
+ "as_money" -> new IntAsDecimalFunction(money),
+ "static_as_money" -> new StaticAsMoneyFunction,
+ "as_wide_money" -> new IntAsDecimalFunction(DecimalType(20, 12)),
+ "mills_as_money" -> new IntAsDecimalFunction(DecimalType(7, 2),
valueScale = 3),
+ "non_null_money" -> new IntAsDecimalFunction(money, nullable = false),
+ "money_array" -> new IntAsDecimalFunction(ArrayType(money, containsNull
= true)),
+ "non_null_money_array" ->
+ new IntAsDecimalFunction(ArrayType(intMoney, containsNull = false)),
+ "money_struct" -> new IntAsDecimalFunction(
+ new StructType().add("m", money).add("non_null_m", intMoney, nullable
= false)))
+ private var catalogName: String = _
+
+ override def initialize(name: String, options: CaseInsensitiveStringMap):
Unit =
+ catalogName = name
+
+ override def name(): String = catalogName
+
+ override def listFunctions(namespace: Array[String]): Array[Identifier] =
+ functions.keys.map(Identifier.of(namespace, _)).toArray
+
+ override def loadFunction(ident: Identifier): UnboundFunction =
+ functions.getOrElse(ident.name(), throw new
NoSuchFunctionException(ident))
+ }
+
+ /**
+ * Returns its `INT` argument as the unscaled value of a `Decimal` at
`valueScale`, whatever
+ * scale `declared` has. For an array or struct type, every decimal in the
result holds that
+ * value: the array has one element, and each field of the struct has it.
`invoke` is an
+ * instance method, so Spark lowers a call to `Invoke`. The function binds
to itself.
+ */
+ class IntAsDecimalFunction(declared: DataType, valueScale: Int = 0,
nullable: Boolean = true)
+ extends UnboundFunction
+ with ScalarFunction[Any] {
+ override def name(): String = "int_as_decimal"
+ override def description(): String = s"int -> ${declared.sql}, at scale
$valueScale"
+ override def bind(inputType: StructType): BoundFunction = this
+ override def inputTypes(): Array[DataType] = Array(IntegerType)
+ override def resultType(): DataType = declared
+ override def isResultNullable(): Boolean = nullable
+ def invoke(v: Int): Any = valueOf(declared, v)
+ override def produceResult(input: InternalRow): Any =
invoke(input.getInt(0))
Review Comment:
`produceResult` is never called here. Spark first looks for the magic
`invoke` method and lowers to `Invoke` or `StaticInvoke`, and only uses
`produceResult` through `ApplyFunctionExpression` when no `invoke` exists
(`V2ExpressionUtils.resolveScalarFunction`). Both overrides, here and in
`StaticAsMoneyFunction`, and the `InternalRow` import could go.
##########
spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala:
##########
@@ -231,15 +231,43 @@ private[codegen] object CometBatchKernelCodegenOutput
extends CometTypeShim {
val set = if (nested) "setSafe" else "set"
OutputEmit("", s"$targetVec.$set($idx, $source);")
case dt: DecimalType =>
+ // Rescale to the declared type, and write null when the value does not
fit, as Spark's
+ // `UnsafeRowWriter` and `UnsafeArrayWriter` do in `write(ordinal,
Decimal, precision,
+ // scale)`. A Spark expression already produces its declared precision
and scale, but a
+ // DSv2 function called through `Invoke` / `StaticInvoke` can return a
`Decimal` of any
+ // scale (#6425). Like Spark's writers, this rescales the value in
place, and leaves it
+ // untouched when it does not fit. Unlike them, it does not test
`source` for null: the
+ // callers write null values themselves, and skip that test only for a
type that is not
+ // nullable.
+ //
+ // Only the kernel's own output sees the null. Spark rescales such a
value when it writes a
+ // row, and an expression around the call reads the value the function
returned, so
+ // `QueryPlanSerde` dispatches that expression with the call, and falls
an aggregate back.
+ //
+ // The precision and scale test repeats `changePrecision`'s own fast
path. It keeps the call
+ // off the common path, so the JIT can still scalar-replace the
`Decimal` that an input
+ // getter allocates. With the bare call, passing a `DECIMAL(18, 2)`
column through took
+ // about half as long again per row.
+ //
// DecimalOutputShortFastPath: precision <= 18 fits in a signed long, so
pass the unscaled
// value to `setSafe(int, long)` and skip the BigDecimal allocation.
Review Comment:
The rule that Spark rescales a DSv2 result only when it writes a row, so a
consumer runs in the kernel with the call and an aggregate falls back, is now
explained here, at the aggregate check, in the `readsDispatchedDsv2Decimal`
scaladoc and in the `getSupportLevel` scaladoc in `statics.scala`. Keeping the
full explanation on `readsDispatchedDsv2Decimal` with one-line pointers
elsewhere would leave one place to update. Here, three points seem enough: the
rescale mirrors `UnsafeRowWriter`, the caller owns the null check, and the
guard keeps `changePrecision` off the hot path. The new `changePrecision`
assertions in `CometCodegenSourceSuite` repeat what the behavior tests prove,
so matching the guard (`.precision() == 18`) there might pin the part that only
a source test can.
--
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]