sunchao commented on code in PR #4565:
URL: https://github.com/apache/datafusion-comet/pull/4565#discussion_r4169583187


##########
spark/src/main/scala/org/apache/spark/sql/comet/operators.scala:
##########
@@ -2430,24 +2434,69 @@ object CometObjectHashAggregateExec
   }
 }
 
-case class CometHashAggregateExec(
-    override val nativeOp: Operator,
-    override val originalPlan: SparkPlan,
-    override val output: Seq[Attribute],
-    groupingExpressions: Seq[NamedExpression],
-    aggregateExpressions: Seq[AggregateExpression],
-    aggregateAttributes: Seq[Attribute],
-    resultExpressions: Seq[NamedExpression],
-    input: Seq[Attribute],
-    child: SparkPlan,
-    override val serializedPlanOpt: SerializedPlan)
+object CometSortAggregateExec extends CometBaseAggregate[SortAggregateExec] {
+
+  // Spark plans SortAggregate for TypedImperativeAggregate functions whose 
intermediate buffer
+  // formats differ between Spark and Comet, the same risk as 
ObjectHashAggregate.
+  override protected def requiresCometShuffle: Boolean = true
+
+  override protected def operatorSupportLevel(op: SortAggregateExec): 
SupportLevel = {
+    // Spark's sort aggregation buffers in an UnsafeRow, which latches a 
decimal sum that leaves
+    // the precision, only when every buffer field is mutable and codegen does 
not apply;
+    // otherwise the sum stays unbounded. Comet latches only when grouped, so 
decline rather
+    // than track every combination.
+    if (hasMaxPrecisionDecimalSum(op)) {

Review Comment:
   [P2] Could this overflow guard also cover decimal `AVG`? Cache one partition 
ordered by `(g, ord)` containing one group with `DECIMAL(38,38)` values `0.6, 
0.6, -0.4` and a string `label` equal to `'x'`. With Comet shuffle and native 
caching enabled, `SELECT g, avg(v), first(label) FROM t GROUP BY g` becomes 
eligible for native sort aggregation. Spark uses a generic buffer because of 
the string `FIRST`, allowing the intermediate sum `1.2` to recover to `0.8`, 
and returns `0.26666666666666666666666666666666666667` in both ANSI modes. The 
native `AvgDecimal` path instead returns `NULL` in legacy mode and an overflow 
error in ANSI mode. This query previously stayed on Spark, so registering the 
operator exposes a new wrong-result/query-failure case. Reject affected decimal 
averages here until accumulation matches Spark, and add a regression for both 
ANSI settings.
   
   Evidence: Freshly compiled `/tmp/pr4565-c66-review/C66Validation.scala` 
against Spark 4.1.3. With AQE disabled, it cached rows `(g,ord,v,label) = 
(1,1,0.6,'x'), (1,2,0.6,'x'), (1,3,-0.4,'x')` after casting `v` to 
`DECIMAL(38,38)` and sorting one partition by `(g,ord)`. Both ANSI settings 
produced two `SortAggregateExec` nodes and the expected average. Freshly 
compiled `/tmp/pr4565-c66-review/native.rs`, including this head’s 
`avg_decimal.rs`, against DataFusion 55.1.0 and Arrow 59.3.0. Equivalent 
Partial/Final aggregation with `FIRST` and the new conditional output-sort 
logic returned `NULL` in legacy mode and `DecimalSumOverflow { function_name: 
"avg" }` in ANSI mode. Logs are `spark.log` and `native.log` in that directory. 
Code-path inspection confirms `CometAverage` accepts this type, while the new 
operator guard checks only `Sum`. This is a bounded component reproduction, not 
an end-to-end Comet 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]

Reply via email to