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


##########
spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala:
##########
@@ -231,15 +231,39 @@ 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.
+      //
+      // 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.
+      val dec = ctx.freshName("dec")
+      val (precision, scale) = (dt.precision, dt.scale)
       val write =
-        if (dt.precision <= Decimal.MAX_LONG_DIGITS) {
-          s"$targetVec.setSafe($idx, $source.toUnscaledLong());"
+        if (precision <= Decimal.MAX_LONG_DIGITS) {
+          s"$targetVec.setSafe($idx, $dec.toUnscaledLong());"
         } else {
-          s"$targetVec.setSafe($idx, $source.toJavaBigDecimal());"
+          s"$targetVec.setSafe($idx, $dec.toJavaBigDecimal());"
         }
-      OutputEmit("", write)
+      OutputEmit(
+        "",
+        s"""org.apache.spark.sql.types.Decimal $dec = $source;
+           |if (($dec.precision() == $precision && $dec.scale() == $scale) ||
+           |    $dec.changePrecision($precision, $scale)) {
+           |  $write
+           |} else {
+           |  $targetVec.setNull($idx);

Review Comment:
   Fixed in 2949bda8a7. You're right that the kernel's null reached `IS NULL`. 
An expression that takes one of these calls as an argument now runs in the same 
kernel as the call, so Spark's own code reads the value the function returned, 
and the operator falls back to Spark if the dispatcher declines it. I added 
your `IS NULL` case to `CometCodegenSuite`, along with a cast to a wider 
decimal, a comparison, a cast to string and `hash`. The last two differ from 
Spark even for values that fit, because Spark reads the function's scale 0 
value, so the problem was wider than the overflow case. All of them fail 
without the change.
   
   Aggregates can't run in the kernel, so `count`, `max` and `sum` over one of 
these calls now fall back to Spark. `count` of a value that doesn't fit was 
another shape the old writer got right by accident: Spark counts it, and the 
kernel's null didn't. Group, sort, window and join keys already matched Spark 
when I checked them, because Spark writes those into rows too.
   



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