Skip to content

Feature/quadrature tree shap gpu improvements - #12652

Merged
RAMitchell merged 1 commit into
dmlc:masterfrom
ron-wettenstein:feature/QuadratureTreeSHAP_GPU_improvements
Oct 7, 2026
Merged

RAMitchell merged 1 commit into
dmlc:masterfrom
ron-wettenstein:feature/QuadratureTreeSHAP_GPU_improvements

Conversation

@ron-wettenstein

@ron-wettenstein ron-wettenstein commented Oct 4, 2026 •

Copy link
Copy Markdown
Contributor

Simplify and speed up QuadratureTreeSHAP on GPU

This PR speeds up SHAP values and, especially, interaction values on the GPU, and makes shap.cu
more 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.

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

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

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

  4. Let every thread read its own values.
    Previously, one thread per row read values such as p_enter and q_prev and 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).

@coderabbitai

coderabbitai Bot commented Oct 4, 2026 •

Copy link
Copy Markdown

Review in Change Stack →

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration
  • Configuration used: Repository UI
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: 8f3609e3-c5e7-4f63-8665-279fb2c230c8
📥 Commits

Reviewing files that changed from the base of the PR and between 053d656 and 6a0cd83.

📒 Files selected for processing (1)
  • src/predictor/interpretability/shap.cu

Included review availability: This review used your included allowance. Your plan provides up to 10 included reviews per hour; 9 remain after this review.


Walkthrough

Both GPU quadrature SHAP task runners now read path probabilities and calculate q_prev and p_e per lane, replacing subgroup-leader calculations and broadcasts. Subgroup leaders still write shared path state. Interaction contribution calculations remain in place.

Suggested reviewers: ramitchell, trivialfis

Priority: ➖ Normal

Merge Risk: ⚪ Minimal · up to 6a0cd

No actionable regression was identified in the per-lane SHAP probability calculations. The change is mergeable after normal checks.

  • Autopilot · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

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.

❤️ Share

Comment @coderabbitai help to get the list of available commands.

@RAMitchell

Copy link
Copy Markdown
Member

Awesome work Ron. Same comment as the other PR. 4 changes = 4 PRs.

@ron-wettenstein

Copy link
Copy Markdown
Contributor Author

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.
Link to colab notebook: https://drive.google.com/file/d/1cSV57Bqz4FS6TaBIJmGS-pntO7qL17lA/view?usp=sharing
The PR makes Shapley interaction values between 1.2 to 6 times faster.
The runtime of first order Shapley values did not changed.

FYI @RAMitchell.

Model Features Depth Task Speedup
heloc 23 6 SHAP values 1.01×
interactions 1.19×
superconductivity 81 6 SHAP values 1.00×
interactions 1.25×
superconductivity 81 16 SHAP values 1.05×
interactions 2.29×
superconductivity 81 24 SHAP values 1.06×
interactions 3.56×
bioresponse 1776 6 SHAP values 0.98×
interactions 6.18×
synthetic, 3-class 20 6 SHAP values 1.00×
interactions 1.20×
Absolute times (ms per call, master → branch)
Model Task Rows Time
heloc d6 SHAP values 10000 119 → 117
heloc d6 interactions 10000 387 → 324
superconductivity d6 SHAP values 10000 217 → 216
superconductivity d6 interactions 5898 541 → 433
superconductivity d16 SHAP values 1155 444 → 425
superconductivity d16 interactions 254 991 → 433
superconductivity d24 SHAP values 545 450 → 423
superconductivity d24 interactions 115 1546 → 431
bioresponse d6 SHAP values 3751 65 → 66
bioresponse d6 interactions 53 2458 → 397
synthetic 3-class d6 SHAP values 10000 244 → 244
synthetic 3-class d6 interactions 6445 513 → 429

Setup.

  • GPU: Tesla T4, driver 580.82.07, CUDA 13.0, on Google Colab.
  • Builds: this branch (2696079f) vs master (b16b82e4), both Release builds for the T4's architecture.
  • Models: XGBoost with 100 trees at the given max depth (tree_method="hist", seed 0) on heloc,
    superconductivity and bioresponse from OpenML (ids 46932, 43174 and 4134), plus a synthetic 3-class
    multi:softprob model with 150 trees. Predictions use device="cuda".
  • Rows per call: chosen so the branch takes about 0.4 s per call, capped at 10,000 rows and at the dataset size. Interactions on bioresponse are also capped by the size of the output.
  • Method: each version runs in its own process, the order alternates every round, and there are 4 rounds.
    Each round takes the median of 3 calls after a warm-up, and the table shows the median speedup across rounds.

@ron-wettenstein

ron-wettenstein commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor Author

I split the contributions to their own branches

  1. Contribution 1 is in QuadratureTreeSHAP GPU: Remove the O(D²) partner scan in interaction values #12653 . It is stand alone
  2. Contribution 2 is in QuadratureTreeSHAP GPU: Parallelize the final step of SHAP interaction values #12654 . It is stand alone
  3. Contribution 3 is in QuadratureTreeSHAP GPU: Pad the last block of 4 rows #12655 . It depends on the code of contribution 1 and 2, so the PR includes their code too. After merging of the first two PRs this PR will include only the code of contributions 3.
  4. Contribution 4 depends on 1, 2 and 3. Instead of closing this PR, we can use it for contribution 4. After merging 1, 2 and 3 this PR will include only 4.

FYI @RAMitchell

@ron-wettenstein
ron-wettenstein force-pushed the feature/QuadratureTreeSHAP_GPU_improvements branch from 2696079 to 053d656 Compare October 6, 2026 08:23
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>
@ron-wettenstein
ron-wettenstein force-pushed the feature/QuadratureTreeSHAP_GPU_improvements branch from 053d656 to 6a0cd83 Compare October 6, 2026 19:25
@ron-wettenstein

ron-wettenstein commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor Author

@RAMitchell this branch is now updated, includes only contribution 4, and ready for review :)

@RAMitchell
RAMitchell merged commit 98eebd0 into dmlc:master Oct 7, 2026
84 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants