feat(kda): add persistent SM90 bwd intra CUDA kernel - #122
Closed
zheyang0825 wants to merge 1 commit into
Closed
Conversation
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
📌 Description
Integrates the optimized persistent CUDA C++ KDA intra-chunk backward kernel under
csrc/kda/sm90/bwdand dispatches supportedsafe_gate=True, BF16, K=128, chunk-size-64 workloads to it. Unsupported inputs continue to use the existing Triton implementation.The kernel uses warp-level
mma.sync.m16n8k8. It is organized under SM90 to match the current repository layout and is compiled for SM90, SM100, and SM103.Also adds:
🔍 Related Issues
None.
🚀 Pull Request Checklist
✅ Pre-commit Checks
pre-commit.pre-commit run --all-filesand fixed all reported issues.🧪 Tests
Validation on NVIDIA L20X (SM90), CUDA 12.9, PyTorch 2.9.1+cu129:
⚡ Performance
Same-machine median latency, 25 ms warmup and 100 ms measurement windows:
The full eight-shape tables and reproduction steps are in
BENCHMARK_KDA_BWD_INTRA_SM90.md.Reviewer Notes
Three review/fix passes were completed before opening this PR:
The optimized path intentionally supports only
safe_gate=True, BF16 Q/K/beta, FP32 gate/gradient inputs, K=128, and chunk size 64. Other cases retain the Triton fallback.