Skip to content

support sm100 fwd hdim256 - #2432

Closed
lyppg wants to merge 1 commit into
Dao-AILab:mainfrom
lyppg:sm100_fwd_hdim256_support
Closed

lyppg wants to merge 1 commit into
Dao-AILab:mainfrom
lyppg:sm100_fwd_hdim256_support

Conversation

@lyppg

@lyppg lyppg commented Apr 3, 2026 •

Copy link
Copy Markdown

Enable head_dim=256 for forward kernels on SM100 GPUs.
Key changes: Force q_stage=1 and tile_n=64 to fit within the 512-column TMEM budget
This is an initial implementation with large room for optimization. We will continue to iterate and improve performance in follow-up PRs.
Test:

python benchmarks/benchmark_attn.py --backend fa4 --fwd --causal true --batch 16 --headdim 256
Benchmarking FA4 fwd, hdim=256, seqlen=8192, causal=True, nheads=8, nheads_kv=8

==================================================
  FWD (ms / TFLOPS / MFU%)
==================================================
     hdim causal batch seqlen                  FA4
--------------------------------------------------
      256   True    16   8192      4.15/1059/42.4%

pytest tests/cute/test_flash_attn.py::test_flash_attn_output[4096-4096-256-True-0-0.0-False-False-False-mha-dtype0]
===================================================================================== test session starts =====================================================================================
platform linux -- Python 3.12.3, pytest-9.0.2, pluggy-1.6.0
rootdir: /home/jingxinpan/git/jxp/flash-attention/tests
configfile: pyproject.toml
plugins: anyio-4.13.0, typeguard-4.5.1, xdist-3.8.0
collected 1 item

tests/cute/test_flash_attn.py .                                                                                                                                                         [100%]

====================================================================================== warnings summary =======================================================================================
cute/test_flash_attn.py: 62 warnings
  /usr/local/lib/python3.12/dist-packages/nvidia_cutlass_dsl/python_packages/cutlass/base_dsl/_mlir_helpers/op.py:63: DeprecationWarning: `make_fragment` is deprecated, use `make_rmem_tensor` instead
    res_or_list = opFunc(*args, **kwargs, loc=loc)

-- Docs: https://docs.pytest.org/en/stable/how-to/capture-warnings.html
=============================================================================== 1 passed, 62 warnings in 12.34s ===============================================================================

@tridao

tridao commented Apr 3, 2026

Copy link
Copy Markdown
Member

Thanks!
We also have an ongoing PR #2412 for hdim 256 fwd + bwd

@lyppg

lyppg commented Apr 3, 2026

Copy link
Copy Markdown
Author

Thanks! We also have an ongoing PR #2412 for hdim 256 fwd + bwd

Just saw it! That one looks much comprehensive, I will close this one, thanks.

@lyppg lyppg closed this Apr 3, 2026
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