Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
184 changes: 150 additions & 34 deletions datafusion/physical-plan/src/windows/bounded_window_agg_exec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ use crate::{

use arrow::compute::take_record_batch;
use arrow::{
array::{Array, ArrayRef, RecordBatchOptions, UInt32Builder},
array::{Array, ArrayRef, RecordBatchOptions, UInt32Array, UInt32Builder},
compute::{concat, concat_batches, sort_to_indices, take_arrays},
datatypes::SchemaRef,
record_batch::RecordBatch,
Expand Down Expand Up @@ -682,17 +682,25 @@ impl PartitionSearcher for LinearSearch {
evaluate_partition_by_column_values(record_batch, window_expr)?;
// NOTE: In Linear or PartiallySorted modes, we are sure that
// `partition_bys` are not empty.
// Calculate indices for each partition and construct a new record
// batch from the rows at these indices for each partition:
self.get_per_partition_indices(&partition_bys, record_batch)?
let (mut keys, permutation, bounds) =
self.compute_partition_permutation(&partition_bys, record_batch)?;
if keys.len() == 1 {
// The batch contains a single partition, so the gather below
// would be an identity permutation; use the batch as-is.
let key = keys.remove(0);
return Ok(vec![(key, record_batch.clone())]);
}
// Reorder the batch with a single `take` so that each partition's
// rows become contiguous, then hand each partition a zero-copy slice
// of the result. The slices share the gathered batch's buffers;
// `PartitionBatchState::extend` copies out of them the next time the
// partition receives rows.
let gathered = take_record_batch(record_batch, &UInt32Array::from(permutation))?;
Ok(keys
.into_iter()
.map(|(row, indices)| {
let mut new_indices = UInt32Builder::with_capacity(indices.len());
new_indices.append_slice(&indices);
let indices = new_indices.finish();
Ok((row, take_record_batch(record_batch, &indices)?))
})
.collect()
.zip(bounds.windows(2))
.map(|(key, bound)| (key, gathered.slice(bound[0], bound[1] - bound[0])))
.collect())
}

fn prune(&mut self, n_out: usize) {
Expand Down Expand Up @@ -746,42 +754,66 @@ impl LinearSearch {
}
}

/// Calculate indices of each partition (according to PARTITION BY expression)
/// `columns` contain partition by expression results.
fn get_per_partition_indices(
/// Splits the rows of `batch` by partition, according to the PARTITION BY
/// expression results in `columns`. Returns the distinct partition keys
/// in first-appearance order, a permutation of the row indices of
/// `batch` that groups each partition's rows together, and the
/// boundaries of each partition's run of rows within that permutation:
/// partition `p` occupies `permutation[bounds[p]..bounds[p + 1]]`, and
/// its indices are in ascending (stream) order.
fn compute_partition_permutation(
&mut self,
columns: &[ArrayRef],
batch: &RecordBatch,
) -> Result<Vec<(PartitionKey, Vec<u32>)>> {
let mut batch_hashes = vec![0; batch.num_rows()];
) -> Result<(Vec<PartitionKey>, Vec<u32>, Vec<usize>)> {
let num_rows = batch.num_rows();
let mut batch_hashes = vec![0; num_rows];
create_hashes(columns, &self.random_state, &mut batch_hashes)?;
self.input_buffer_hashes.extend(&batch_hashes);
// reset row_map for new calculation
self.row_map_batch.clear();
// res stores PartitionKey and row indices (indices where these partition occurs in the `batch`) for each partition.
let mut result: Vec<(PartitionKey, Vec<u32>)> = vec![];
let mut keys: Vec<PartitionKey> = vec![];
// Partition id of each row, in row order:
let mut row_partition_ids = Vec::with_capacity(num_rows);
// Number of rows in each partition:
let mut counts: Vec<usize> = vec![];
for (hash, row_idx) in batch_hashes.into_iter().zip(0u32..) {
let entry = self.row_map_batch.find_mut(hash, |(_, group_idx)| {
// We can safely get the first index of the partition indices
// since partition indices has one element during initialization.
let row = get_row_at_idx(columns, row_idx as usize).unwrap();
// Handle hash collusions with an equality check:
row.eq(&result[*group_idx].0)
// Handle hash collisions with an equality check:
row.eq(&keys[*group_idx])

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can't == be used here?

});
if let Some((_, group_idx)) = entry {
result[*group_idx].1.push(row_idx)
let group_idx = if let Some((_, group_idx)) = entry {
*group_idx
} else {
self.row_map_batch.insert_unique(
hash,
(hash, result.len()),
|(hash, _)| *hash,
);
let row = get_row_at_idx(columns, row_idx as usize)?;
// This is a new partition its only index is row_idx for now.
result.push((row, vec![row_idx]));
}
let group_idx = keys.len();
self.row_map_batch
.insert_unique(hash, (hash, group_idx), |(hash, _)| *hash);
keys.push(get_row_at_idx(columns, row_idx as usize)?);
counts.push(0);
group_idx
};
row_partition_ids.push(group_idx);
counts[group_idx] += 1;
}
// A prefix sum over the counts gives each partition's run boundaries
// in the permutation.
let mut bounds = Vec::with_capacity(counts.len() + 1);
let mut total = 0;
bounds.push(0);
for count in counts {
total += count;
bounds.push(total);
}
// Scatter each row's index into its partition's run. Visiting rows
// in ascending order keeps each run in ascending row order.
let mut cursors: Vec<usize> = bounds[..bounds.len() - 1].to_vec();
let mut permutation = vec![0u32; num_rows];
for (row_idx, group_idx) in row_partition_ids.into_iter().enumerate() {
permutation[cursors[group_idx]] = row_idx as u32;
cursors[group_idx] += 1;
}
Ok(result)
Ok((keys, permutation, bounds))
}

/// Calculates partition keys and result indices for each partition.
Expand Down Expand Up @@ -1924,4 +1956,88 @@ mod tests {
));
Ok(())
}

/// Checks the per-partition batches that `LinearSearch` splits an input
/// batch into: partitions appear in first-appearance order, rows within a
/// partition keep their stream order, NULL keys form their own partition,
/// and a single-partition batch is passed through without copying.
#[test]
fn test_linear_search_evaluate_partition_batches() -> Result<()> {
use super::{LinearSearch, PartitionSearcher};
use arrow::array::{Int32Array, Int64Array};

let schema = Arc::new(Schema::new(vec![
Field::new("a", DataType::Int32, true),
Field::new("b", DataType::Int64, false),
]));
let window_expr = create_window_expr(
&WindowFunctionDefinition::AggregateUDF(count_udaf()),
"count".to_string(),
&[col("b", &schema)?],
&[col("a", &schema)?],
&[],
Arc::new(WindowFrame::new(None)),
Arc::clone(&schema),
false,
false,
None,
)?;
let mut searcher = LinearSearch::new(vec![], Arc::clone(&schema));

let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(Int32Array::from(vec![
Some(1),
Some(2),
Some(1),
None,
Some(2),
Some(1),
])),
Arc::new(Int64Array::from(vec![10, 20, 11, 30, 21, 12])),
],
)?;
let result =
searcher.evaluate_partition_batches(&batch, &[Arc::clone(&window_expr)])?;
assert_eq!(result.len(), 3);
let expected = [
(
ScalarValue::Int32(Some(1)),
vec![Some(1); 3],
vec![10i64, 11, 12],
),
(ScalarValue::Int32(Some(2)), vec![Some(2); 2], vec![20, 21]),
(ScalarValue::Int32(None), vec![None], vec![30]),
];
for ((key, partition_batch), (exp_key, exp_a, exp_b)) in
result.iter().zip(expected)
{
assert_eq!(key, &vec![exp_key]);
let exp_batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(Int32Array::from(exp_a)),
Arc::new(Int64Array::from(exp_b)),
],
)?;
assert_eq!(partition_batch, &exp_batch);
}

let single = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(Int32Array::from(vec![Some(7), Some(7)])),
Arc::new(Int64Array::from(vec![70, 71])),
],
)?;
let result = searcher.evaluate_partition_batches(&single, &[window_expr])?;
assert_eq!(result.len(), 1);
assert_eq!(result[0].0, vec![ScalarValue::Int32(Some(7))]);
assert_eq!(result[0].1, single);
// The whole batch belongs to one partition, so its columns are reused
// rather than gathered into a new batch.
assert!(Arc::ptr_eq(result[0].1.column(0), single.column(0)));
Ok(())
}
}
Loading