[Perf] gemm_bf16 gfx1250: cross-tile carry to hide K-tile prologue ds… - #1125
Open
amd-hhashemi wants to merge 1 commit into
Open
amd-hhashemi wants to merge 1 commit into
amd-hhashemi wants to merge 1 commit into
Conversation
…_load latency Defer each K-tile's last WMMA K-step across the tile boundary so the next tile's step-0 LDS reads hide behind the previous tile's WMMAs, instead of stalling on the exposed post-barrier s_wait_dscnt. The last K-step loads directly into persistent (zeroed) carry fragment tensors -- no copy, and the first tile's carry-WMMA adds 0 so no prologue peel is needed. compute_head runs _mma_carry() right after issuing step-0 reads, then computes steps 0..K_WS-2; an epilogue drains the final tile's carry. TDM/nb staging is unchanged; only the ds_load->WMMA schedule is restructured, and accumulation stays in natural K order so results are bit-consistent. Measured on M=64 N=65536 K=16384 (tiles 64,256,128, cluster 1,4): +7.5-8.0% vs baseline, PASS on both random and const inputs (rel_err matches baseline). Gain is realized in the compute-bound clock regime (high fclk / lower gfxclk); ~neutral when memory-bound.
Contributor
There was a problem hiding this comment.
🔵 Needs a closer look
Address the wmma_m_rep == 1 scheduling concern before approval.
Pull request overview
Optimizes gfx1250 BF16 GEMM by carrying the final WMMA K-step across tile boundaries to overlap LDS reads.
Changes:
- Adds persistent carry fragments.
- Restructures WMMA scheduling and drains the final carry.
- Preserves TDM staging and accumulation order.
File summaries
| File | Summary |
|---|---|
kernels/gemm/gemm_bf16_gfx1250.py |
Implements cross-tile carry scheduling; the single-WMMA-row prefetch timing requires retaining the prior split or targeted benchmarking. |
Review details
Suppressed comments (1)
kernels/gemm/gemm_bf16_gfx1250.py:256
- This removes the existing
wmma_m_repsplit and now issues the next TDM prefetch before_mma_kseven when there is only one WMMA row repeat. The old path deliberately issued it after that WMMA forwmma_m_rep == 1; moving the descriptor issue in front of the only compute instruction can expose TDM issue latency in the skinny-M configurations, while the reported benchmark only coverswmma_m_rep > 1. Please retain the prior before/after split for this case or add a targeted benchmark before accepting the schedule change.
if const_expr(ks == 0 and prefetch_kt is not None):
rocdl.sched_barrier(0)
issue(prefetch_kt % num_buffers, prefetch_kt)
rocdl.sched_barrier(0)
_mma_ks(cur)
- Files reviewed: 1/1 changed files
- Comments generated: 0
- Review effort level: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
This branch has not been deployed
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.
Defer each K-tile's last WMMA K-step across the tile boundary so the next tile's step-0 LDS reads hide behind the previous tile's WMMAs, instead of stalling on the exposed post-barrier s_wait_dscnt.
The last K-step loads directly into persistent (zeroed) carry fragment tensors -- no copy, and the first tile's carry-WMMA adds 0 so no prologue peel is needed. compute_head runs _mma_carry() right after issuing step-0 reads, then computes steps 0..K_WS-2; an epilogue drains the final tile's carry. TDM/nb staging is unchanged; only the ds_load->WMMA schedule is restructured, and accumulation stays in natural K order so results are bit-consistent.
Measured on M=64 N=65536 K=16384 (tiles 64,256,128, cluster 1,4): +7.5-8.0% vs baseline, PASS on both random and const inputs (rel_err matches baseline). Gain is realized in the compute-bound clock regime (high fclk / lower gfxclk); ~neutral when memory-bound.
Motivation
Technical Details
Test Plan
Test Result
Submission Checklist