comphead commented on code in PR #6455:
URL: https://github.com/apache/datafusion-comet/pull/6455#discussion_r4167024122


##########
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/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/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.



##########
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.



-- 
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]

Reply via email to