perf(pwmj):settle never-matching streamed rows against the buffered extreme… - #25840
SubhamSinghal wants to merge 6 commits into
Conversation
… key before sorting
|
Benchmark PR: #25839 |
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## main #25840 +/- ##
==========================================
+ Coverage 82.55% 82.67% +0.11%
==========================================
Files 1141 1147 +6
Lines 440209 446769 +6560
Branches 440209 446769 +6560
==========================================
+ Hits 363409 369349 +5940
- Misses 54829 55006 +177
- Partials 21971 22414 +443 ☔ View full report in Codecov by Harness. 🚀 New features to boost your workflow:
|
|
@comphead @jayzhan211 can you help in reviewing this PR. |
jayzhan211
left a comment
There was a problem hiding this comment.
Thanks @SubhamSinghal, LGTM!
| buffer_idx = first_match(buffer_idx, buffered_len, |idx| { | ||
| is_match(cmp.compare(row_idx, idx), match_on_equal) | ||
| }); | ||
| debug_assert!(buffer_idx < buffered_len); |
There was a problem hiding this comment.
Maybe internal error here?
There was a problem hiding this comment.
+1. In release builds nothing fails here. With buffer_idx == buffered_len, count is 0 and an empty batch is pushed, so for Right/Full the streamed row is neither matched nor emitted as unmatched and is silently lost. The debug_assert_eq! on streamed NULLs a few lines up has the same shape. A NULL key that reached the scan compares Less than every non-NULL buffered key (nulls_first), so it would match all of them. Both checks are cheap enough for internal_err!.
comphead
left a comment
There was a problem hiding this comment.
Checked locally against NestedLoopJoinExec with a differential fuzz over random 1 to 7 row, unsorted streamed batches (Inner/Left/Right/Full, all four operators, batch sizes 2, 3 and 8192, one or two streamed partitions). No mismatches on this head for i32, f64 with -0.0/NaN, Utf8, lists without NULL elements and the other flat key types in matchable_rows_agrees_with_scan. PWMJ output for list keys with NULL elements is identical to main on the same inputs, as the description says.
fuzz_pwmj_matches_nested_loopbuilds the streamed side as one-row batches (pwmj_parts_exec), so a batch is either fully matchable or not at all. The new mixed path (rows split off, the rest sorted and mapped back throughpositions) is only reached by the small hand-written cases. Chunking the streamed side into random multi-row batches there is a cheap guard for it.- The description mentions a new benchmark and
cargo bench -p datafusion --bench pwmj_classic_sql, but neither is in this diff. The suite is in #25839 and runs throughbenchmark_runner -- pwmj. Please update the reproduction command and link it. - The classic join section of the
PiecewiseMergeJoinExecdocs inexec.rsstill describes the linear walk. A line on the pre-sort check against the last buffered key and the binary search would keep it accurate. - Nits:
null_padded_streamed_batchkeeps the oldnew_stream_batchparameter name. The "settled before the sort, unmatched forRight/Full, dropped otherwise" explanation is repeated infetch_stream_batch,split_off_never_matching,matchable_rowsand the loop inresolve_classic_join, so a pointer is enough in three of them. Once the nested-key issue is filed, please reference it next to the nested branch inmatchable_rowsso the branch can go when the comparator is fixed.
| // NULL keys never match and sort first, so the scan starts past the buffered ones. | ||
| // Streamed NULL keys never get here: `matchable_rows` settled them before the sort. | ||
| debug_assert_eq!(stream_values[0].null_count(), 0); | ||
| if !batch_process_state.processed_null_count { |
There was a problem hiding this comment.
With streamed NULLs gone from the scan, processed_null_count only seeds buffer_idx with the buffered NULL count once per stream batch. Seeding start_buffer_idx right after batch_process_state.reset() in fetch_stream_batch removes the flag, its reset and this first-call branch. I tried it locally and the PWMJ unit tests and the fuzz still pass.
| /// holds, so a check built on that would let the `+0.0` row through as a candidate | ||
| /// that the scan then rejects; with SQL semantics both rows are unmatched. | ||
| #[tokio::test] | ||
| async fn join_right_less_than_signed_zero_prefilter_agrees_with_scan() -> Result<()> { |
There was a problem hiding this comment.
This overlaps the float64 case of matchable_rows_agrees_with_scan and the f64 pass of fuzz_pwmj_matches_nested_loop. A classic RIGHT/FULL case in piecewise_merge_join_matrix.slt Part 1 would replace it and be checked against NestedLoopJoin at batch sizes 1, 2, 100 and 8192, instead of a hard-coded snapshot. The matrix only has -0.0 for existence joins today. The same file could take one unsorted streamed batch mixing matching, non-matching and NULL keys to cover the positions remap. I tried both cases across the matrix locally and they pass on this head.
| &sort_to_indices(buffered.as_ref(), Some(sort_options), None)?, | ||
| None, | ||
| )?; | ||
| let extreme = match sorted.len().checked_sub(1) { |
There was a problem hiding this comment.
This copies the extreme computation from collect_buffered_side, so the test would not notice a bug in the production version. A small shared function, for example fn buffered_extreme(values: &ArrayRef) -> Result<Option<ColumnarValue>>, would remove the duplicate and make the test exercise the real code.
jayzhan211
left a comment
There was a problem hiding this comment.
There is one issue from the new commits
| fn buffered_extreme(values: &ArrayRef) -> Result<Option<ColumnarValue>> { | ||
| Ok(match values.len().checked_sub(1) { | ||
| Some(last) if values.is_valid(last) => Some(ColumnarValue::Scalar( | ||
| ScalarValue::try_from_array(values, last)?, |
There was a problem hiding this comment.
For Dictionary(_, Float*) keys the pre-filter disagrees with the scan on -0.0. apply_cmp normalizes the streamed array (dictionary values included), but normalize_float_zero_scalar skips ScalarValue::Dictionary, so the extreme keeps -0.0 under total order while JoinKeyComparator treats it as 0.0. This gives an internal error when the filter lets a row through, and wrong rows when it drops a real match. The base commit returns the right results.
Repro:
set datafusion.optimizer.enable_piecewise_merge_join = true;
CREATE TABLE dl AS SELECT * FROM (VALUES (1, arrow_cast(-0.0, 'Dictionary(Int32, Float64)')), (2, arrow_cast(0.5, 'Dictionary(Int32, Float64)'))) t(lid, lv);
CREATE TABLE dr AS SELECT * FROM (VALUES (1, arrow_cast(0.0, 'Dictionary(Int32, Float64)')), (2, arrow_cast(0.25, 'Dictionary(Int32, Float64)'))) t(rid, rv);
SELECT * FROM dl JOIN dr ON dl.lv < dr.rv;
-- Internal error: PiecewiseMergeJoin: streamed row 1 reached the scan without a match
CREATE TABLE dl2 AS SELECT * FROM (VALUES (1, arrow_cast(-1.0, 'Dictionary(Int32, Float64)')), (2, arrow_cast(-0.0, 'Dictionary(Int32, Float64)'))) t(lid, lv);
CREATE TABLE dr2 AS SELECT * FROM (VALUES (1, arrow_cast(0.0, 'Dictionary(Int32, Float64)'))) t(rid, rv);
SELECT * FROM dl2 RIGHT JOIN dr2 ON dl2.lv >= dr2.rv;
-- got `NULL NULL 1 0`, expected `2 0 1 0` (NLJ and the base commit agree)Fix (normalize the extreme before taking it as a scalar):
+use datafusion_common::utils::normalize_float_zero; fn buffered_extreme(values: &ArrayRef) -> Result<Option<ColumnarValue>> {
Ok(match values.len().checked_sub(1) {
- Some(last) if values.is_valid(last) => Some(ColumnarValue::Scalar(
- ScalarValue::try_from_array(values, last)?,
- )),
+ // `apply_cmp` normalizes `-0.0` only in flat float scalars, not inside a
+ // `ScalarValue::Dictionary`, so normalize the key before taking it.
+ Some(last) if values.is_valid(last) => {
+ let extreme = normalize_float_zero(&values.slice(last, 1));
+ Some(ColumnarValue::Scalar(ScalarValue::try_from_array(
+ &extreme, 0,
+ )?))
+ }
_ => None,
})
}Cases for matchable_rows_agrees_with_scan. They fail on the current head (dictionary_float64_neg_zero_min/full <: streamed row 0) and pass with the fix:
(
// `-0.0` as the buffered extreme inside a dictionary: the smallest
// key for `<`/`<=`, the largest for `>`/`>=`.
"dictionary_float64_neg_zero_min",
Arc::new(DictionaryArray::<Int32Type>::new(
Int32Array::from(vec![0, 1]),
Arc::new(Float64Array::from(vec![-0.0, 0.5])),
)),
Arc::new(DictionaryArray::<Int32Type>::new(
Int32Array::from(vec![0, 1, 2]),
Arc::new(Float64Array::from(vec![0.0, -0.0, 0.25])),
)),
),
(
"dictionary_float64_neg_zero_max",
Arc::new(DictionaryArray::<Int32Type>::new(
Int32Array::from(vec![0, 1]),
Arc::new(Float64Array::from(vec![-1.0, -0.0])),
)),
Arc::new(DictionaryArray::<Int32Type>::new(
Int32Array::from(vec![0, 1, 2]),
Arc::new(Float64Array::from(vec![0.0, -0.0, -0.5])),
)),
),The root cause, normalize_float_zero_scalar not recursing into ScalarValue::Dictionary, is already on main: SELECT arrow_cast(0.0, 'Dictionary(Int32, Float64)') > arrow_cast(-0.0, 'Dictionary(Int32, Float64)') returns true. Worth filing on its own; fixing it there would also cover this.
There was a problem hiding this comment.
Thanks @jayzhan211. Addressed in 036502f
| )?; | ||
| let last = buffered_values.len() - 1; | ||
| let matchable = BooleanBuffer::collect_bool(num_rows, |row| { | ||
| stream_values.is_valid(row) |
There was a problem hiding this comment.
is_valid only checks the physical null buffer, and a RunArray has none. So a run-end encoded NULL key passes as matchable, gets past the null_count() > 0 guard at :492, and the scan joins it with every buffered row (nulls_first → Less). main gives the same rows, so this is fine as a follow-up. But join_right_run_end_encoded_keys_empty_sliced_batch already feeds such a NULL (a2 = 0) and asserts nothing.
let last = buffered_values.len() - 1;
+ // A run-end encoded NULL has no physical null buffer.
+ let nulls = stream_values.logical_nulls();
let matchable = BooleanBuffer::collect_bool(num_rows, |row| {
- stream_values.is_valid(row)
+ nulls.as_ref().is_none_or(|n| n.is_valid(row))
&& is_match(cmp.compare(row, last), match_on_equal)
});- if stream_values[0].null_count() > 0 {
+ if stream_values[0].logical_null_count() > 0 {In the test, after importing batches_to_sort_string. This fails on this head and passes with the fix:
let (_, batches, _) =
join_collect(left, right, on, Operator::Lt, JoinType::Right).await?;
// The NULL key (a2 = 0) and 0 (a2 = 3) match nothing.
assert_snapshot!(batches_to_sort_string(&batches), @r"
+----+----+----+----+
| a1 | b1 | a2 | b2 |
+----+----+----+----+
| | | 0 | |
| | | 3 | 0 |
| 1 | 5 | 4 | 6 |
| 2 | 5 | 4 | 6 |
| 3 | 1 | 1 | 2 |
| 3 | 1 | 2 | 2 |
| 3 | 1 | 4 | 6 |
| 4 | 1 | 1 | 2 |
| 4 | 1 | 2 | 2 |
| 4 | 1 | 4 | 6 |
+----+----+----+----+
");| &[sort_options], | ||
| NullEquality::NullEqualsNothing, | ||
| )?; | ||
| for row in 0..streamed.len() { |
There was a problem hiding this comment.
This PR also fixes dictionary keys whose values are NULL (a valid key pointing at a NULL value). apply_cmp sees the logical NULL and sets the row aside, while main joins it with every buffered row. Nothing pins this, and the oracle here uses physical is_valid, so it would call such a row matchable. Fine as a follow-up; this passes on the current head:
+ let nulls = streamed.logical_nulls();
for row in 0..streamed.len() {
- let expected = streamed.is_valid(row)
+ let expected = nulls.as_ref().is_none_or(|n| n.is_valid(row))(
// A valid key pointing at a NULL value: logically NULL, physically valid.
"dictionary_null_values",
Arc::new(DictionaryArray::<Int32Type>::new(
Int32Array::from(vec![0, 1, 2]),
Arc::new(Int32Array::from(vec![Some(5), None, Some(1)])),
)),
Arc::new(DictionaryArray::<Int32Type>::new(
Int32Array::from(vec![0, 1, 2, 3]),
Arc::new(Int32Array::from(vec![Some(2), None, Some(6), Some(0)])),
)),
),
Which issue does this PR close?
PiecewiseMergeJoinperformance for large tables #18221 (and the PWMJ epic [EPIC]: MakePiecewiseMergeJoinwork in Datafusion #17427).It doesn't close #18221. On the issue's 100K × 100K shape, the join's time goes to its huge output, so this change leaves it about the same. See the benchmarks below.
Rationale for this change
The classic
PiecewiseMergeJoinstream (INNER/LEFT/RIGHT/FULL) sorts each streamed batch, then walks the sorted buffered side row by row to find each streamed row's first match. Two costs are avoidable:<, and NULL keys. When most streamed rows can't match, which is common for a selective range predicate, the operator does all this work and outputs nothing for them.[k, len)of the sorted buffered side, sokcan be found with a binary search instead of stepping over every non-matching buffered row.What changes are included in this PR?
All changes are in
joins/piecewise_merge_join/, plus a new benchmark.classic_join.rs). Because every match set is a suffix, a streamed row matches anything at all only if it matches the last buffered key.Right/Fulland dropped forInner/Left. Only the remaining rows are sorted, gathered and scanned.apply_cmp, which normalizes-0.0/+0.0just as the scan'sJoinKeyComparatordoes.apply_cmporders NULL elements inside a key ascending, while the comparator applies the sort options, descending for</<=, at every nesting level.</<=gives different results fromNestedLoopJoinExecwhen keys contain NULL elements. That is a pre-existing bug onmain, which I'll file and fix separately. This PR keeps PWMJ's current results unchanged.utils.rs).first_match(lo, hi, matches)replaces the linear walk inresolve_classic_join.[previous row's match, len). The streamed batch is sorted, so each row's first match is at or after the previous row's.matches_on_equalandis_match, used by both streams.debug_assert. The buffered-side NULL skip stays.Benchmarks
The
pwmjsuite from #25839 (100K buffered keys, 2M streamed rows, five match regimes ×INNER/LEFT/RIGHT/FULL). Reproduce with:main(median)no_match: no streamed row matchesnull_heavy: half NULL, restno_matchhalf_match: halfno_match, halfselectiveselective: every row matches the 1–4 smallest buffered keysall_match: every row matches every buffered keyWhat is the testing strategy for this PR?
first_match_agrees_with_linear_scan(utils.rs): exhaustive over every range, answer and start position up to 40. It checks the result against a linear scan and bounds the number of comparisons by⌈log2(hi − lo + 1)⌉.matchable_rows_agrees_with_scan(classic_join.rs): the pre-sort check against a brute-force version of the scan's own definition of a match. It covers 12 key types, 4 operators, and a normal, all-NULL and empty buffered side.±0.0/NaN/±inf, Utf8, Utf8View, Binary, Dictionary, Decimal128, Date32, Timestamp, Boolean, and two List cases with NULL elements.join_right_less_than_signed_zero_prefilter_agrees_with_scan: end to end with-0.0/+0.0keys.fuzz_pwmj_matches_nested_loop(--features extended_tests), covering Inner/Left/Right/Full, NULL keys, small batches and multiple partitions;Are there any user-facing changes?
No API or configuration changes.