Skip to content

aiter-flash-attn: add AITER Flash Attention kernel for AMD ROCm - #903

Merged
danieldk merged 7 commits into
mainfrom
aiter-flash-attn-add
Jun 2, 2026
Merged

aiter-flash-attn: add AITER Flash Attention kernel for AMD ROCm#903
danieldk merged 7 commits into
mainfrom
aiter-flash-attn-add

Conversation

@danieldk

@danieldk danieldk commented Jun 1, 2026

Copy link
Copy Markdown
Member

This adds a Triton FlashAttention kernel for AMD ROCm, repackaged from the MHA implementation in AMD’s AITER project (https://github.com/ROCm/aiter).

The motivation is on the Transformers side: the ROCm FlashAttention fallback currently depends on installing the full aiter pip package, which reviewers (rightly) pushed back on. Having an equivalent kernel here means Transformers can route through get_kernel like every other FA backend instead of carrying a direct AITER dependency. Beyond just unblocking the dependency story, this is also an FA3-style kernel with native support for learnable attention sinks (sink=), which is required by models like gpt-oss on ROCm; something flash-attn2 does not provide.

It’s a slim copy; only the parts actually reachable from flash_attn_func and flash_attn_varlen_func, with the unused dao_ai implementation path removed and all absolute imports rewritten to be Hub-compliant. It has been tested locally on MI300X against an eager SDPA reference; numerics match within fp16 tolerance for dense, causal, and varlen cases.

Original PR: #890

@danieldk
danieldk requested a review from drbh as a code owner June 1, 2026 10:59
@danieldk
danieldk merged commit 27f5cac into main Jun 2, 2026
7 of 9 checks passed
@danieldk
danieldk deleted the aiter-flash-attn-add branch June 2, 2026 08:22
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