gabotechs commented on code in PR #25337:
URL: https://github.com/apache/datafusion/pull/25337#discussion_r4152679411


##########
datafusion/functions-aggregate/src/approx_percentile_cont.rs:
##########
@@ -405,10 +448,19 @@ impl ApproxPercentileAccumulator {
 
 impl Accumulator for ApproxPercentileAccumulator {
     fn state(&mut self) -> Result<Vec<ScalarValue>> {
-        Ok(self.digest.to_scalar_state().into_iter().collect())
+        let mut state: Vec<ScalarValue> =
+            self.digest.to_scalar_state().into_iter().collect();
+        if self.should_track_percentile {
+            state.push(ScalarValue::Float64(self.percentile.get().ok()));
+        }
+        Ok(state)

Review Comment:
   It'd be nice if we could avoid these if-driven-development patterns.
   
   I see that `should_track_percentile` is here just so that the 
`approx_median` function can disable it and afford to carry just q less state 
parameter (6 instead of 7). IMHO, I don't think it's worth it, the performance 
impact is going to be neglibible and removing it would allow us to simplify the 
code.
   
   



##########
datafusion/functions-aggregate/src/percentile_cont.rs:
##########
@@ -273,27 +282,36 @@ impl AggregateUDFImpl for PercentileCont {
     }
 }
 
-fn get_percentile(args: &AccumulatorArgs) -> Result<f64> {
-    let percentile = validate_percentile_expr(&args.exprs[1], 
"PERCENTILE_CONT")?;
+fn get_percentile_param(args: &AccumulatorArgs) -> Result<(PercentileParam, 
bool)> {
+    let percentile = PercentileParam::try_new(&args.exprs[1], 
"PERCENTILE_CONT")?;
 
     let is_descending = args
         .order_bys
         .first()
         .map(|sort_expr| sort_expr.options.descending)
         .unwrap_or(false);
 
-    let percentile = if is_descending {
+    Ok((percentile, is_descending))
+}
+
+/// The percentile to use, applying the `1.0 - p` flip for descending
+/// `WITHIN GROUP (ORDER BY ... DESC)`.
+fn effective_percentile(
+    percentile: &PercentileParam,
+    is_descending: bool,
+) -> Result<f64> {
+    let percentile = percentile.get()?;
+    Ok(if is_descending {
         1.0 - percentile
     } else {
         percentile
-    };
-
-    Ok(percentile)
+    })
 }
 
 pub fn create_percentile_accumulator(
     name: &str,
-    percentile: f64,
+    percentile: PercentileParam,
+    is_descending: bool,
     input_dt: &DataType,

Review Comment:
   The fact that `PercentileParam` and `is_descending` always need to go 
threaded together makes me think that `is_descencing` should be a field in 
`PercentileParam`



##########
datafusion/functions-aggregate/src/approx_percentile_cont.rs:
##########
@@ -441,9 +493,16 @@ impl Accumulator for ApproxPercentileAccumulator {
             return Ok(());
         }
 
-        let states = (0..states[0].len())
+        if self.should_track_percentile
+            && let Some(percentile_array) = states.get(6)

Review Comment:
   When referring to hardcoded numbers, prefer leaving a comment about why this 
is `6` and not something else.
   
   If you think it's appropriate, I'd even extract it to a constant with a nice 
comment.



##########
datafusion/functions-aggregate/src/approx_percentile_cont.rs:
##########
@@ -337,31 +343,68 @@ impl AggregateUDFImpl for ApproxPercentileCont {
 #[derive(Debug)]
 pub struct ApproxPercentileAccumulator {
     digest: TDigest,
-    percentile: f64,
+    percentile: PercentileParam,
+    is_descending: bool,
     return_type: DataType,
+    should_track_percentile: bool,
 }
 
 impl ApproxPercentileAccumulator {
-    pub fn new(percentile: f64, return_type: DataType) -> Self {
+    pub(crate) fn new(
+        percentile: PercentileParam,
+        is_descending: bool,
+        return_type: DataType,
+    ) -> Self {
         Self {
             digest: TDigest::new(DEFAULT_MAX_SIZE),
             percentile,
+            is_descending,
             return_type,
+            should_track_percentile: false,
         }

Review Comment:
   I see there's an inconsistent usage of the new params the 
`ApproxPercentileAccumulator` accepts:
   - `is_descencing` is required as a new argument in the constructor methods
   - `should_track_percentile` is accessed through a builder pattern
   - `aggregate_fn_name` and `state` are passed wrapped in a new 
`PercentileParam` struct
   
   How about just having them inside the `PercentileParam` struct?



##########
datafusion/sqllogictest/test_files/aggregate.slt:
##########
@@ -190,7 +190,7 @@ SELECT c1, approx_percentile_cont(0.95, -1000) WITHIN GROUP 
(ORDER BY c3) AS c3_
 statement error Function 'approx_percentile_cont' failed to match any signature
 SELECT approx_percentile_cont(0.95, c1) WITHIN GROUP (ORDER BY c3) FROM 
aggregate_test_100
 
-statement error DataFusion error: Error during planning: Percentile value for 
'APPROX_PERCENTILE_CONT' must be a literal
+statement error DataFusion error: Error during planning: Percentile value for 
'APPROX_PERCENTILE_CONT' must be constant across the aggregation, found 
differing values

Review Comment:
   :+1: nice



##########
datafusion/sqllogictest/test_files/aggregate_memory_spill.slt:
##########
@@ -227,11 +227,18 @@ FROM (
   GROUP BY (v * 7) % 100000
 )
 ----
-<slt:ignore>
-06)----------AggregateExec: mode=FinalPartitioned, gby=[t.v * Int64(7) % 
Int64(100000)@0 as t.v * Int64(7) % Int64(100000)], aggr=[sum(t.v)], 
metrics=[<slt:ignore>spilled_rows=<slt:ignore>K,<slt:ignore>]
-<slt:ignore>
-08)--------------AggregateExec: mode=Partial, gby=[v@0 * 7 % 100000 as t.v * 
Int64(7) % Int64(100000)], aggr=[sum(t.v)], 
metrics=[<slt:ignore>early_emit_count=<slt:ignore>]
-<slt:ignore>
+Plan with Metrics
+01)ProjectionExec: expr=[count(Int64(1))@0 as count(*), sum(total)@1 as 
sum(total)], metrics=[output_rows=1, elapsed_compute=<slt:ignore>, 
output_bytes=16.0 B, output_batches=1, expr_0_eval_time=<slt:ignore>, 
expr_1_eval_time=<slt:ignore>]
+02)--AggregateExec: mode=Final, gby=[], aggr=[count(Int64(1)), sum(total)], 
metrics=[output_rows=1, elapsed_compute=<slt:ignore>, output_bytes=16.0 B, 
output_batches=1, agg_expr_0_arguments_time=<slt:ignore>, 
agg_expr_0_evaluate_time=<slt:ignore>, agg_expr_0_merge_time=<slt:ignore>, 
agg_expr_1_arguments_time=<slt:ignore>, agg_expr_1_evaluate_time=<slt:ignore>, 
agg_expr_1_merge_time=<slt:ignore>]
+03)----CoalescePartitionsExec, metrics=[output_rows=4, 
elapsed_compute=<slt:ignore>, output_bytes=64.0 B, output_batches=4]
+04)------AggregateExec: mode=Partial, gby=[], aggr=[count(Int64(1)), 
sum(total)], metrics=[output_rows=4, elapsed_compute=<slt:ignore>, 
output_bytes=64.0 B, output_batches=4, agg_expr_0_arguments_time=<slt:ignore>, 
agg_expr_0_state_time=<slt:ignore>, agg_expr_0_update_time=<slt:ignore>, 
agg_expr_1_arguments_time=<slt:ignore>, agg_expr_1_state_time=<slt:ignore>, 
agg_expr_1_update_time=<slt:ignore>]
+05)--------ProjectionExec: expr=[sum(t.v)@1 as total], 
metrics=[output_rows=100.0 K, elapsed_compute=<slt:ignore>, output_bytes=787.4 
KB, output_batches=788, expr_0_eval_time=<slt:ignore>]
+06)----------AggregateExec: mode=FinalPartitioned, gby=[t.v * Int64(7) % 
Int64(100000)@0 as t.v * Int64(7) % Int64(100000)], aggr=[sum(t.v)], 
metrics=[output_rows=100.0 K, elapsed_compute=<slt:ignore>, 
output_bytes=<slt:ignore> MB, output_batches=788, spill_count=<slt:ignore>, 
spilled_bytes=<slt:ignore> MB, spilled_rows=<slt:ignore> K, 
agg_expr_0_arguments_time=<slt:ignore>, agg_expr_0_evaluate_time=<slt:ignore>, 
agg_expr_0_merge_time=<slt:ignore>, agg_expr_0_state_time=<slt:ignore>, 
aggregate_arguments_time=<slt:ignore>, aggregation_time=<slt:ignore>, 
emitting_time=<slt:ignore>, time_calculating_group_ids=<slt:ignore>]
+07)------------RepartitionExec: partitioning=Hash([t.v * Int64(7) % 
Int64(100000)@0], 4), input_partitions=4, metrics=[output_rows=100.0 K, 
elapsed_compute=<slt:ignore>, output_bytes=<slt:ignore> KB, output_batches=784, 
spill_count=<slt:ignore>, spilled_bytes=<slt:ignore> KB, 
spilled_rows=<slt:ignore>, fetch_time=<slt:ignore>, 
repartition_time=<slt:ignore>, send_time=<slt:ignore>]
+08)--------------AggregateExec: mode=Partial, gby=[v@0 * 7 % 100000 as t.v * 
Int64(7) % Int64(100000)], aggr=[sum(t.v)], metrics=[output_rows=100.0 K, 
elapsed_compute=<slt:ignore>, output_bytes=<slt:ignore> MB, output_batches=782, 
spill_count=0, spilled_bytes=0.0 B, spilled_rows=0, 
early_emit_count=<slt:ignore>, skipped_aggregation_rows=0, 
agg_expr_0_arguments_time=<slt:ignore>, 
agg_expr_0_convert_to_state_time=<slt:ignore>, 
agg_expr_0_state_time=<slt:ignore>, agg_expr_0_update_time=<slt:ignore>, 
aggregate_arguments_time=<slt:ignore>, aggregation_time=<slt:ignore>, 
emitting_time=<slt:ignore>, time_calculating_group_ids=<slt:ignore>, 
reduction_factor=100% (100.0 K/100.0 K)]
+09)----------------RepartitionExec: partitioning=RoundRobinBatch(4), 
input_partitions=1, maintains_sort_order=true, metrics=[output_rows=100.0 K, 
elapsed_compute=<slt:ignore>, output_bytes=<slt:ignore> KB, output_batches=782, 
spill_count=<slt:ignore>, spilled_bytes=<slt:ignore> KB, 
spilled_rows=<slt:ignore> K, fetch_time=<slt:ignore>, 
repartition_time=<slt:ignore>, send_time=<slt:ignore>]
+10)------------------ProjectionExec: expr=[value@0 as v], 
metrics=[output_rows=100.0 K, elapsed_compute=<slt:ignore>, output_bytes=782.0 
KB, output_batches=782, expr_0_eval_time=<slt:ignore>]
+11)--------------------LazyMemoryExec: partitions=1, 
batch_generators=[generate_series: start=1, end=100000, batch_size=128], 
metrics=[output_rows=100.0 K, elapsed_compute=<slt:ignore>, output_bytes=782.0 
KB, output_batches=782]
 

Review Comment:
   For this one, I'd probably respect the <slt:ignore> tags present in the 
original test assertion. I imagine however put it there had their reasons 
(flaky CI maybe?)



##########
datafusion/sqllogictest/test_files/aggregate.slt:
##########
@@ -190,7 +190,7 @@ SELECT c1, approx_percentile_cont(0.95, -1000) WITHIN GROUP 
(ORDER BY c3) AS c3_
 statement error Function 'approx_percentile_cont' failed to match any signature
 SELECT approx_percentile_cont(0.95, c1) WITHIN GROUP (ORDER BY c3) FROM 
aggregate_test_100
 
-statement error DataFusion error: Error during planning: Percentile value for 
'APPROX_PERCENTILE_CONT' must be a literal
+statement error DataFusion error: Error during planning: Percentile value for 
'APPROX_PERCENTILE_CONT' must be constant across the aggregation, found 
differing values

Review Comment:
   We are missing some sqllogictests here, I don't see any existing tests 
exercising the new path. Some tests cases Astra came up with:
   
   > Could we add a few SLTs covering the new runtime resolution and state 
merging paths? The current additions mostly use p = 0.5, and the grouped tests 
use the same percentile for every group.
   >  - Different percentiles in different groups. For (g, x, p) = (1, 10, 
0.0), (1, 20, 0.0), (2, 30, 1.0), (2, 40, 1.0), grouping by g should return 10 
and 40. This checks that the percentile is constant per group, without 
requiring it to be identical across groups.
   > -  A percentile that changes between batches. Use a multi-batch source 
where each batch has a constant percentile, but later batches use a different 
one. This should error even though each individual batch passes the min/max 
check. A generate_series source with a small batch_size would help exercise 
this.
   > - Percentile propagation and validation across partial/final aggregation. 
Test both matching percentiles across partitions and differing percentiles that 
are individually constant within each partition. The first should succeed; the 
second should error during merging. Verify the physical plan actually contains 
partial/final aggregation.
   > - Descending order with a non-median percentile. For example, 
percentile_cont(p) WITHIN GROUP (ORDER BY x DESC) over (10.0, 0.25), (20.0, 
0.25), (30.0, 0.25), (40.0, 0.25) should return 32.5. Also test the approximate 
variant against its literal-0.25 equivalent. Using 0.5 cannot catch a missing 1 
- p adjustment.
   > - Nonempty DISTINCT with a column percentile. percentile_cont(DISTINCT x, 
p) over (10.0, 0.5), (10.0, 0.5), (30.0, 0.5) should return 20. The new 
DISTINCT test currently only exercises empty input.
   > - Deferred resolution after null percentiles. Have an initial batch with 
only null percentile values, followed by a batch with p = 0.5. Also cover 
nonempty input where the percentile remains null throughout, to establish the 
intended behavior.
   > - Out-of-range column percentiles. Constant columns containing -0.1 or 1.1 
should fail, covering validation after deferred resolution rather than only 
literal validation.
   
   >  It would also be useful to exercise the partitioned cases with 
approx_percentile_cont_with_weight, since it has a separate update path.



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