Repository navigation
QuadratureTreeSHAP GPU: Remove the O(D²) partner scan in interaction values - #12653
RAMitchell merged 2 commits into
Conversation
…ctions The interaction kernel needs the distinct features on the current tree path at every return edge. Keep each depth's split feature, and whether a deeper split on the same feature shadows it, in per-warp shared memory, updated as the walk enters and leaves nodes. This replaces the O(D^2) rescan of the path in global memory with an O(D) pass over shared memory. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (2)
Included review availability: This review used your included allowance. Your plan provides up to 10 included reviews per hour; 9 remain after this review. WalkthroughGPU SHAP interaction traversal now tracks split features and whether earlier path entries are shadowed by repeated splits. Partner enumeration walks the active path and uses only unshadowed entries. The interaction kernel uses the expanded shared-state type and adds synchronization around traversal state reuse and changes. A new test compares CPU and GPU SHAP outputs for partial row tiles. Suggested reviewers: Priority: ⬇️ Low Merge Risk: ⚪ Minimal · up to The investigated interaction traversal preserves the expected partner selection. No identified issue remains to resolve before normal merge 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 |
Finish shared path reads before reusing node and stage slots, and synchronize stage reads before the warp leader advances the traversal. Apply the ordering to SHAP values and interactions, and cover full and partial row tiles with a CPU/GPU regression test suitable for Racecheck.
|
I fixed a shared memory race condition that was there before this PR. |
This PR is split out of #12652. See that PR for more details, including benchmark results against master for this change and the three related ones.
This PR include contribution 1 of #12652 :
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.