Skip to content

[Kernel][Perf] Optimize gfx950 head64 prefill attention - #1126

Open
michael604work wants to merge 1 commit into
ROCm:mainfrom
michael604work:michael/head64-page64-attention-pr
Open

michael604work wants to merge 1 commit into
ROCm:mainfrom
michael604work:michael/head64-page64-attention-pr

Conversation

@michael604work

Copy link
Copy Markdown

Summary

Optimize gfx950 BF16 prefill attention for the 42B MXFP4 serving shape:

  • head dimension 64 with 32 query heads and 4 KV heads
  • dense contiguous and page-64 vectorized-5D KV layouts
  • single-request and packed-varlen inputs
  • bottom-right causal global and causal sliding-window attention

The change adds workload-specific routing and improves the dual-wave/generic
pipelines through causal work reordering, mask specialization, page-table and
page-address reuse, overlapped K/V address preparation, and softmax/MFMA
scheduling.

Motivation

The existing gfx950 attention paths did not cover this head-dim-64/page-64
serving workload with the required combination of long-context paged KV,
packed varlen requests, and sliding-window attention. This is needed for the
42B MXFP4 model configuration.

Workload

  • GPU: MI350X (gfx950)
  • dtype: BF16
  • batch: 1 for dense and single-request paged cases
  • heads: Hq/Hkv = 32/4
  • head dimension: 64
  • page size: 64
  • dense global Q=K: 1K, 2K, 4K, 8K
  • paged global/SWA Q: 1K, 2K, 4K, 8K; KV: 8K, 32K, 128K
  • SWA windows: 2K, 4K, 8K, 16K
  • packed-varlen mixtures include non-page-aligned lengths and KV up to 128K

Performance

Representative dense-global throughput for the accepted campaign incumbent,
using causal-triangle QK+PV FLOPs and the MI350X 2.3 PFLOP/s BF16 peak:

Q=K Latency Useful TFLOP/s MFU
1K 29.07 us 147.9 6.43%
2K 53.13 us 323.5 14.07%
4K 128.93 us 533.1 23.18%
8K 347.70 us 790.7 34.38%

Each campaign promoted changes only after same-GPU paired measurements showed
at least 0.5% incremental geometric-mean improvement, no workload regression
above 3%, candidate MAD at most 3%, and all correctness gates passing.

Test plan

  • scripts/check_python_style.sh --include-local
  • Dense-global sampled-reference correctness at Q=K 1K, 2K, 4K, and 8K
  • pytest tests/kernels/test_swa_gfx950.py -q (5 passed)
  • Paged-global correctness with page-64 identity, reversed, and irregular block
    tables, including sampled-reference comparison
  • Paged-SWA correctness with page-64 identity, reversed, and irregular block
    tables, including sampled-reference comparison
  • Packed-varlen global: all short-heavy, balanced, and boundary workloads,
    including non-page-aligned lengths and KV up to 128K; zero mismatches
  • Packed-varlen SWA at window 4K across all three workload mixtures; zero
    mismatches

Kernel-level tests and microbenchmarks only; no model/e2e tests were run.

Add specialized dense, paged, packed-varlen, global, and sliding-window paths for BF16 GQA with head dimension 64 and page size 64. Improve causal scheduling, page-table address handling, and the dual-wave softmax/MFMA pipeline for long-context workloads.
@coderfeli

Copy link
Copy Markdown
Collaborator

Conflict. and perf compare with other hdim and existing solution?

This branch has not been deployed

No deployments
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