mbutrovich commented on code in PR #5420:
URL: https://github.com/apache/datafusion-comet/pull/5420#discussion_r4189976294
##########
spark/src/main/scala/org/apache/spark/sql/comet/operators.scala:
##########
@@ -1979,6 +1979,30 @@ trait CometBaseAggregate {
case _ => false
})
+ protected def hasMaxPrecisionDecimalAvg(op: BaseAggregateExec): Boolean =
+ op.aggregateExpressions.exists(_.aggregateFunction match {
+ case avg: Average =>
+ avg.sumDataType match {
+ case decimal: DecimalType => decimal.precision ==
DecimalType.MAX_PRECISION
+ case _ => false
+ }
+ case _ => false
+ })
+
+ protected def aggregateSupportLevel(op: BaseAggregateExec): SupportLevel = {
+ val unsupportedAverage = op.groupingExpressions.isEmpty &&
hasMaxPrecisionDecimalAvg(op)
+
+ if (unsupportedAverage) {
+ // Spark's global buffer can retain a wider sum until division; native
overflow is sticky.
+ // Both stages must fall back because decimal AVG buffers cannot cross
engines.
+ Unsupported(
+ Some(
+ "Ungrouped AVG on DECIMAL with maximum-precision intermediate state
is not supported"))
+ } else {
+ Compatible()
+ }
+ }
Review Comment:
#6041 landed after this fallback was designed, and it handles the same Spark
behavior for decimal `SUM` differently. `SumDecimalAccumulator` keeps an
unbounded `i256` running sum and checks the precision only in `state` and
`evaluate`
([`sum_decimal.rs`](https://github.com/apache/datafusion-comet/blob/0ac4dadae70838d8675bfd846c2d1c7617c25c78/native/spark-expr/src/agg_funcs/sum_decimal.rs#L170-L176)).
`CometHashAggregateExec` then falls back only when codegen is off
([`operators.scala`](https://github.com/apache/datafusion-comet/blob/0ac4dadae70838d8675bfd846c2d1c7617c25c78/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala#L2434-L2453)).
That is the one case where Spark's ungrouped buffer is an `UnsafeRow`, which
nulls the sum as soon as it leaves the precision and keeps it null.
Spark's `Average` adds decimals with the same `DecimalAddNoOverflowCheck` as
`Sum`
([`Average.scala`](https://github.com/apache/spark/blob/e221b56be7b6d9e48e107fc4d1cf0c15f02700f8/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/Average.scala#L85-L88),
[`Sum.scala`](https://github.com/apache/spark/blob/e221b56be7b6d9e48e107fc4d1cf0c15f02700f8/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/Sum.scala#L90-L93)).
As I read it, the same design would work for `AvgDecimalAccumulator`, and
`avg()` already has an `i256` path for the division. That would keep ungrouped
wide-decimal `AVG` native, along with the sibling `MIN`, `MAX`, `COUNT` and
`SUM` that this fallback takes with it. `SumDecimalAccumulator` also serves
ever-expanding window frames, so it would let the new expanding-frame `AVG`
guard in
[`CometWindowExec`](https://github.com/apache/datafusion-comet/blob/86f25149fb21a472a67d7fa3c0d929fc2931f6c5/spark/src/main/scal
a/org/apache/spark/sql/comet/CometWindowExec.scala#L359-L369) go too.
Could `AvgDecimalAccumulator` follow #6041 here, with
`hasMaxPrecisionDecimalAvg` added to the `codegenOff` check in place of this
operator-wide fallback? If you'd rather keep this PR to the current fix, could
you open an issue for it and link it from the new Aggregation section in
`operators.md`?
##########
native/spark-expr/src/agg_funcs/avg_decimal.rs:
##########
@@ -671,31 +609,281 @@ impl GroupsAccumulator for AvgDecimalGroupsAccumulator {
}
/// Returns the `sum`/`count` as a i128 Decimal128 with
-/// target_scale and target_precision and return None if overflows.
+/// target_scale and target_precision, returning the rounded wide value on
overflow.
///
/// * sum: The total sum value stored as Decimal128 with sum_scale
/// * count: total count, stored as a i128 (*NOT* a Decimal128 value)
/// * target_min: The minimum output value possible to represent with the
target precision
/// * target_max: The maximum output value possible to represent with the
target precision
/// * scaler: scale factor for avg
#[inline(always)]
-fn avg(sum: i128, count: i128, target_min: i128, target_max: i128, scaler:
i128) -> Option<i128> {
- if let Some(value) = sum.checked_mul(scaler) {
- // `sum / count` with ROUND_HALF_UP
+fn avg(
+ sum: i128,
+ count: i128,
+ target_min: i128,
+ target_max: i128,
+ scaler: i128,
+) -> std::result::Result<i128, i256> {
+ // The scaled numerator can exceed i128 while both the sum and final
average fit.
+ // Use the inexpensive path when it fits, and widen only that intermediate
otherwise.
+ let rounded = if let Some(value) = sum.checked_mul(scaler) {
let (div, rem) = value.div_rem(&count);
let half = div_ceil(count, 2);
- let half_neg = half.neg_wrapping();
- let new_value = match value >= 0 {
+ let rounded = match value >= 0 {
true if rem >= half => div.add_wrapping(1),
- false if rem <= half_neg => div.sub_wrapping(1),
+ false if rem <= half.neg_wrapping() => div.sub_wrapping(1),
_ => div,
};
- if new_value >= target_min && new_value <= target_max {
- Some(new_value)
+ i256::from_i128(rounded)
+ } else {
+ // A Decimal128 numerator times the Decimal128 scale factor fits
within i256.
+ let value = i256::from_i128(sum) * i256::from_i128(scaler);
+ let divisor = i256::from_i128(count);
+ let (div, rem) = (value / divisor, value % divisor);
+ let half = i256::from_i128(div_ceil(count, 2));
+ if rem >= half {
+ div + i256::ONE
+ } else if rem <= -half {
+ div - i256::ONE
} else {
- None
+ div
}
+ };
+ if rounded >= i256::from_i128(target_min) && rounded <=
i256::from_i128(target_max) {
+ Ok(rounded.to_i128().expect("bounded Decimal128 average"))
} else {
- None
+ Err(rounded)
+ }
+}
+
+fn avg_result_overflow_error(value: i256, precision: u8, scale: i8) ->
crate::SparkError {
+ let unscaled = value.to_string();
+ let digits = unscaled.trim_start_matches('-').len();
+ crate::SparkError::NumericValueOutOfRange {
+ value: format_decimal_str(&unscaled, digits, scale),
+ precision,
+ scale,
+ }
+}
Review Comment:
Could this use `Decimal256Type::format_decimal`, the way
`SumDecimalAccumulator::evaluate` builds the same error
([`sum_decimal.rs`](https://github.com/apache/datafusion-comet/blob/0ac4dadae70838d8675bfd846c2d1c7617c25c78/native/spark-expr/src/agg_funcs/sum_decimal.rs#L288-L292))?
I compared the two locally on zero, positive, negative and widened `i256`
values at scales 0, 4, 22 and 38, and they produce the same string. The imports
would swap `format_decimal_str` for `Decimal256Type` and
`DECIMAL256_MAX_PRECISION`.
```suggestion
fn avg_result_overflow_error(value: i256, precision: u8, scale: i8) ->
crate::SparkError {
crate::SparkError::NumericValueOutOfRange {
value: Decimal256Type::format_decimal(value,
DECIMAL256_MAX_PRECISION, scale),
precision,
scale,
}
}
```
--
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]