Repository navigation
Feature/quadrature tree shap gpu improvements - #12652
RAMitchell merged 1 commit into
Conversation
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (1)
Included review availability: This review used your included allowance. Your plan provides up to 10 included reviews per hour; 9 remain after this review. WalkthroughBoth GPU quadrature SHAP task runners now read path probabilities and calculate Suggested reviewers: Priority: ➖ Normal Merge Risk: ⚪ Minimal · up to No actionable regression was identified in the per-lane SHAP probability calculations. The change is mergeable after normal checks.
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
|
Awesome work Ron. Same comment as the other PR. 4 changes = 4 PRs. |
Benchmark: speedup over master on the GPU (master time / branch time)To measure the speedups of the new approach, I used the datasets and model configurations also used in #12650. FYI @RAMitchell.
Absolute times (ms per call, master → branch)
Setup.
|
|
I split the contributions to their own branches
FYI @RAMitchell |
2696079 to
053d656
Compare
Values such as p_enter and q_prev were read by one thread per row and passed to the row's other threads with warp shuffles. Since padding guarantees that every thread works on a real row, each thread now reads these values, and evaluates its row's split, directly. This removes more than half of the shuffles in the tree walk and the now unused Broadcast helper. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
053d656 to
6a0cd83
Compare
|
@RAMitchell this branch is now updated, includes only contribution 4, and ready for review :) |
Simplify and speed up QuadratureTreeSHAP on GPU
This PR speeds up SHAP values and, especially, interaction values on the GPU, and makes
shap.cumore than 100 lines shorter. Only the GPU code changes, and the results stay the same
(up to floating-point rounding). See the speedup results in a comment below.
Keep the features of the current path in shared memory (interaction values).
To compute interactions, the kernel repeatedly needs the distinct features on the current tree path. It used to rebuild this list at every step by rescanning the path in global memory, which costs O(D²) per step for a tree of depth D. It now keeps the list up to date in shared memory as it walks the tree, so the lookup costs O(D). On a T4, interaction values are about 1.1–1.2× faster for depth-6 trees and about 2× faster for depth-16 trees.
Parallelize the final step of interaction values.
After the tree walk, each interaction matrix is made symmetric (the (i, j) and (j, i) entries are
averaged) and its diagonal (each feature's interaction with itself) is computed. This used to run
with a single GPU thread per matrix. It now runs as two passes: one thread per entry to make the matrix symmetric, then one thread per column to compute its diagonal. The results are bit-for-bit the same. The gain is largest for models with many features: with 1000 features, interaction values are 3.4× faster in batch and about 60× faster for a single row.
Pad the last block of rows instead of handling it separately.
The GPU processes rows in blocks of 4. When the number of rows was not a multiple of 4, the leftover rows ran in a second, specialized kernel that could only start after the first one finished. For deep trees this could almost double the run time. Now the empty slots of the last block simply repeat the last row and are not allowed to write their results. For example, with 34 rows the last block gets rows 32, 33, [33], [33]. This is up to about 1.7× faster for deep trees (e.g. 70 rows at depth 16) and removes about 100 lines of code.
Let every thread read its own values.
Previously, one thread per row read values such as
p_enterandq_prevand passed them to the other 7 threads of that row with warp shuffles. Since padding guarantees that every thread works on a real row, each thread can now read these values directly. This removes more than half of the shuffles in the tree walk, gives a 3–8% speedup on deep trees, and removes about 40 lines of code.Validation: the GPU tests pass, and the outputs were checked against the CPU implementation and an exact float64 reference on several models (up to depth 16 and 1000 features).