andygrove opened a new pull request, #6350:
URL: https://github.com/apache/datafusion-comet/pull/6350

   ## Which issue does this PR close?
   
   Closes #3025.
   
   ## Rationale for this change
   
   Native `CASE WHEN` and `IF` spent most of their time outside the expressions 
they evaluate.
   
   - DataFusion's `CaseExpr` filters the batch down to each branch's rows, 
evaluates the branch, and merges the partial results through 
`MutableArrayData`. That is one dynamically dispatched extend per run of 
consecutive rows taking the same branch, which on unsorted data is about one 
call for every two rows.
   - The planner wrapped every branch in a Spark `Cast` to the type it already 
had. The cast itself is cheap, but it hid literals and columns from 
DataFusion's own fast paths and made every CASE nullable.
   
   Filtering and merging is only needed when a branch could fail for a row 
Spark never evaluates it for, such as a division by a zero that the WHEN rules 
out. When nothing can fail, evaluating each branch over the whole batch gives 
the same values.
   
   ## What changes are included in this PR?
   
   - `CaseWhenExpr`, in `conditional_funcs/case_when.rs`. The planner now 
builds it for `CaseWhen`, and `IfExpr` evaluates through it.
     - It decides on the first batch whether every branch value, and every WHEN 
after the first, cannot fail or return something different for seeing more 
rows. That covers columns and literals, and comparisons, boolean logic, null 
checks, widening casts, and wrapping integer or IEEE floating point arithmetic 
over them. Spark evaluates the first WHEN for every row, so it is exempt.
     - If so, it evaluates them over the whole batch and takes each row's value 
from its branch, with typed merges for primitive, boolean, string and binary 
results. The merges copy whole 64-row words or whole runs where they can, so 
sorted data does not regress.
     - Otherwise, and for every other result type, it evaluates with 
DataFusion's `CaseExpr` as before. A branch that can fail is still evaluated 
only for the rows that choose it, and a function is always assumed to be able 
to fail.
     - An `IF` or `CASE WHEN` in the ELSE is merged in one pass with the outer 
branches, since `IF(a, x, IF(b, y, z))` is `CASE WHEN a THEN x WHEN b THEN y 
ELSE z END`.
     - Nullability follows Spark: the result can be NULL when a branch value 
can or there is no ELSE, which is what `IfExpr` already did.
   - `create_case_when` replaces the planner's `create_case_expr` and only 
casts a branch whose type differs from the common type.
   - `benches/conditional.rs` now builds each shape the way the planner does, 
on 8192-row batches of random data. It covers the queries in 
`CometConditionalExpressionBenchmark`, common TPC-H and TPC-DS shapes, Comet's 
divisor guard, coalesce, string and boolean results, a branch that can fail, 
sparse and dense nulls, and sorted data.
   
   ## How are these changes tested?
   
   - Unit tests in `case_when.rs`, including a differential test against 
DataFusion's `CaseExpr`. It covers 10 result types, 3 null densities, 5 
predicate layouts (with random value bits under NULL predicates) and 5 branch 
shapes, on whole and sliced batches. Seven mutations of the merges and the 
eligibility rules each fail at least one test.
   - New SQL tests:
   
     - `case_when_types.sql`: every result type the merges handle, with NULLs 
in both predicates and values, plus a 20000-row table in random and in 
ascending order.
     - `case_when_lazy.sql`: runs with ANSI on and off. A branch or a later 
WHEN would fail for rows that an earlier WHEN rules out.
     - `case_when_ansi.sql`: errors for rows that do choose a failing branch, 
and errors in the first WHEN, still surface.
   
     All three pass on main as well. Against a build with the primitive merge 
off by one and every expression treated as infallible, all three fail.
   
   - `CometSqlFileTestSuite`, `CometExpressionSuite`, `CometNativeCastSuite`, 
`CometMathExpressionSuite` and `CometAggregateSuite` pass locally on the 
default profile.
   
   Criterion, against main on an Apple M3 Max:
   
   | Shape                               |     main | This PR | Speedup |
   | ----------------------------------- | -------: | ------: | ------: |
   | case literal 3 branches             |  79.9 µs | 28.1 µs |   2.84x |
   | case literal 10 branches            | 151.6 µs | 47.5 µs |   3.19x |
   | case column 3 branches              |  62.0 µs |  7.9 µs |   7.82x |
   | case column 10 branches             | 185.2 µs | 82.6 µs |   2.24x |
   | if literal                          |  41.2 µs | 27.1 µs |   1.52x |
   | if column                           |  45.0 µs |  6.8 µs |   6.56x |
   | nested if literal                   | 100.3 µs | 32.2 µs |   3.11x |
   | nested if column                    | 104.1 µs | 17.5 µs |   5.94x |
   | case int literals                   |  12.4 µs |  3.1 µs |   4.05x |
   | case column or null                 |   9.9 µs |  2.7 µs |   3.74x |
   | divisor guard                       |  19.9 µs |  2.8 µs |   7.14x |
   | coalesce, sparse nulls              |  50.1 µs |  3.9 µs |  12.85x |
   | coalesce, dense nulls               |  44.0 µs |  3.7 µs |  11.83x |
   | if short strings                    | 146.3 µs | 31.1 µs |   4.71x |
   | if long strings                     | 228.8 µs | 68.8 µs |   3.33x |
   | if boolean                          |  50.2 µs |  4.3 µs |  11.64x |
   | case fallible branch                |  14.1 µs | 14.6 µs |   0.97x |
   | if column, sparse nulls             |  73.6 µs |  7.3 µs |  10.09x |
   | if column, dense nulls              |  30.2 µs |  5.7 µs |   5.26x |
   | if short strings, sparse nulls      | 181.9 µs | 30.9 µs |   5.89x |
   | if short strings, dense nulls       |  55.7 µs | 16.5 µs |   3.37x |
   | case column 3 branches, dense nulls |  40.5 µs |  6.6 µs |   6.14x |
   | if column, sorted                   |  14.0 µs |  5.1 µs |   2.76x |
   | if short strings, sorted            |  95.6 µs | 13.0 µs |   7.33x |
   | case literal 10 branches, sorted    |  68.1 µs | 29.5 µs |   2.31x |
   | case column 10 branches, sorted     |  91.3 µs | 78.7 µs |   1.16x |
   
   The one shape that is slower is a branch that can fail, which is still 
evaluated with `CaseExpr`. It pays about 0.3 µs per batch, measured with the 
three versions interleaved in one process, because `CaseExpr` derives the 
result type of its first branch for every batch, and that is now a `BinaryExpr` 
rather than a cast that stores its type. `IfExpr` has always given `CaseExpr` 
bare branches.
   
   `CometConditionalExpressionBenchmark`, 1M rows, best of three runs of each 
build interleaved, on a machine with other load:
   
   | Query                                 | Spark | Comet, main   | Comet, 
this PR |
   | ------------------------------------- | ----- | ------------- | 
-------------- |
   | Case When Literal (3 branches)        | 44 ms | 31 ms (1.42x) | 31 ms 
(1.42x)  |
   | Case When Literal (10 branches)       | 44 ms | 39 ms (1.13x) | 28 ms 
(1.57x)  |
   | Case When Column Result (3 branches)  | 32 ms | 25 ms (1.28x) | 25 ms 
(1.28x)  |
   | Case When Column Result (10 branches) | 45 ms | 55 ms (0.82x) | 26 ms 
(1.73x)  |
   | If Literal                            | 36 ms | 31 ms (1.16x) | 31 ms 
(1.16x)  |
   | If Column Result                      | 28 ms | 24 ms (1.17x) | 23 ms 
(1.22x)  |
   | Nested If Literal (4 outcomes)        | 38 ms | 27 ms (1.41x) | 26 ms 
(1.46x)  |
   | Nested If Column Result (4 outcomes)  | 40 ms | 29 ms (1.38x) | 23 ms 
(1.74x)  |
   
   Several queries do not change here. With a native scan, the native plan runs 
a batch ahead of the JVM, and for these queries converting every output row 
back to a Spark row takes longer than the native work does. With each query 
ending in an aggregate instead (`sum(...)` over numbers, `sum(hash(...))` over 
strings), Comet goes from slower than Spark on six of the eight queries to 
faster on all of them:
   
   | Query                                 | Spark | Comet, main   | Comet, 
this PR |
   | ------------------------------------- | ----- | ------------- | 
-------------- |
   | Case When Literal (3 branches)        | 33 ms | 34 ms (0.97x) | 28 ms 
(1.18x)  |
   | Case When Literal (10 branches)       | 35 ms | 47 ms (0.74x) | 26 ms 
(1.35x)  |
   | Case When Column Result (3 branches)  | 28 ms | 33 ms (0.85x) | 21 ms 
(1.33x)  |
   | Case When Column Result (10 branches) | 42 ms | 64 ms (0.66x) | 36 ms 
(1.17x)  |
   | If Literal                            | 31 ms | 29 ms (1.07x) | 28 ms 
(1.11x)  |
   | If Column Result                      | 24 ms | 26 ms (0.92x) | 20 ms 
(1.20x)  |
   | Nested If Literal (4 outcomes)        | 31 ms | 35 ms (0.89x) | 24 ms 
(1.29x)  |
   | Nested If Column Result (4 outcomes)  | 38 ms | 36 ms (1.06x) | 21 ms 
(1.81x)  |
   
   Found while testing: #6334. `IF` fails when its branches differ only in 
nested nullability. It fails the same way on main and this PR does not change 
it. `case_when_types.sql` carries it as an ignored query.
   


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