Skip to content

[FEAT][kernels]: integrate Composable Kernel / ROCm FlashAttention for AMD parity #39

Description

@Flink-ddd

Description:

Context:

Following the hardware abstraction layer refactoring (Issue #36), we need to ensure our AMD users have parity with CUDA users for attention operations. We should integrate AMD's Composable Kernel or the ROCm-ported FlashAttention as the primary backend for AMD GPUs.

Tasks:

  1. Add ROCm backend dispatch logic in rl_engine/kernels/ops/rocm/attention/.
  2. Wrap ROCm Flash-Attention-2 (from flash-attn ROCm fork) or natively call CK attention operators.
  3. Ensure the input/output signature perfectly matches the CUDA flash_attn.py wrapper for seamless higher-level routing.
  4. Update the CI/CD pipeline to skip ROCm specific tests if AMD hardware is not detected, but ensure linting passes.

Activity

  1. added
    component: kernelsTasks involving the development of CUDA and Triton underlying operators
    platform: rocmSpecific tasks specific to AMD graphics cards (such as CK, bpreshuffle/FA)
    type: performancePerformance optimization tasks aimed at increasing throughput and reducing latency etc.
    type: ci-cdModify GitHub Actions, automated tests, and packaging/deployment tasks.
    priority: highSevere congestion issues require the highest priority for resolution.
    on May 31, 2026
  2. FED4 commented on Jun 10, 2026

    @FED4
    Contributor

    Hi, I’d like to take this. I have AMD/NVIDIA GPU access and can implement ROCm attention backend dispatch, validate parity with the CUDA wrapper, and add correctness/perf results. Please assign it to me.

  3. Flink-ddd commented on Jun 10, 2026

    @Flink-ddd
    CollaboratorAuthor

    That's great, This task is primarily for AMD ROCm, so you have an AMD MI300X or 355, right?

  4. FED4 commented on Jun 10, 2026

    @FED4
    Contributor

    That's great, This task is primarily for AMD ROCm, so you have an AMD MI300X or 355, right?

    Yes, I can access MI300X.

  5. Flink-ddd commented on Jun 10, 2026

    @Flink-ddd
    CollaboratorAuthor

    okay, I will assign this task to you. Thanks.

  6. FED4 commented on Jun 11, 2026

    @FED4
    Contributor

    I’ll implement the ROCm attention wrapper with the same public call signature as the existing CUDA FlashAttentionOp.

    My current understanding:

    • Dispatch selection lives in rl_engine/kernels/registry.py.
    • The ROCm attention backend should live under rl_engine/kernels/ops/rocm/attention/.
    • I’ll start with ROCm FlashAttention-2 because it is closest to the existing CUDA FlashAttention wrapper; CK can be added as a follow-up backend if needed.

    For correctness, I plan to use PyTorch SDPA math backend as the numerical ground truth and put the shared checks in tests/test_attention_correctness.py, so the same cases can cover the existing CUDA wrapper and the new ROCm wrapper. I’ll start with fp16 atol=1e-2/rtol=1e-2, following the 1e-2 tolerance scale used in fused-attention validation such as Triton’s fused attention tests. For bf16, I’ll start slightly looser at atol=2e-2/rtol=2e-2. I’ll report MI300X max absolute and relative error so the threshold can be tightened if the observed results allow it.

    I also noticed RL-Kernel does not currently have a dedicated PyTorch SDPA attention fallback wrapper under rl_engine/kernels/ops/pytorch/attention/. I’ll use PyTorch SDPA directly in tests for now.

    I’ll focus on correctness and dispatch first. Once correctness and dispatch are stable, I’ll collect performance numbers. If those results should be added to the project’s benchmark docs, I can align that part with #20.

  7. Flink-ddd commented on Jun 11, 2026

    @Flink-ddd
    CollaboratorAuthor

    Hi @FED4,

    Thanks for the detailed plan. The approach is solid prioritizing the ROCm FA2 fork for API parity and using PyTorch SDPA math backend as the ground truth is exactly the right path.

    A few quick notes for production-readiness:

    1. Tolerances: 1e-2 is too loose for FP16, as FlashAttention uses FP32 accumulation internally. Let's enforce atol=1e-3, rtol=1e-3 for FP16 as the baseline. Your proposed 2e-2 for BF16 is acceptable.
    2. Edge Cases: Please ensure explicit validation for supported headdim raise NotImplementedError for unsupported sizes and ensure the shared tests cover both causal and non-causal paths, as dispatch behavior often differs.

    Agree with your priority: dispatch and correctness first. Looking forward to the PR. Let me know if you need any testing support on the MI300X.

  8. FED4 commented on Jun 11, 2026

    @FED4
    Contributor

    Thanks, I’ll tighten FP16 to atol=1e-3/rtol=1e-3 and keep BF16 at atol=2e-2/rtol=2e-2. I’ll also add explicit supported-head-dim validation and NotImplementedError coverage for unsupported sizes, with shared causal/non-causal correctness tests.

  9. FED4 commented on Jun 12, 2026

    @FED4
    Contributor

    FYI, I tested ROCm FlashAttention against PyTorch SDPA math backend on MI300/gfx942.

    Coverage:

    • FP16 and BF16
    • head dims 64, 128, 256
    • causal and non-causal
    • default scale and explicit 1 / sqrt(head_dim) scale

    Command:

    PYTEST_DISABLE_PLUGIN_AUTOLOAD=1 python -m pytest tests/test_attention_correctness.py -q -rs

    Result:

    • 24 ROCm cases passed
    • 24 CUDA-only cases skipped as expected on ROCm PyTorch

    Observed max absolute error:

    • FP16: 0.00101447105 with tolerance atol=1e-3, rtol=1e-3
    • BF16: 0.00883960724 with tolerance atol=2e-2, rtol=2e-2
     dtype    cases     max abs diff    max mean abs diff    max rel diff    result
    ━━━━━━━  ━━━━━━━  ━━━━━━━━━━━━━━━  ━━━━━━━━━━━━━━━━━━━  ━━━━━━━━━━━━━━  ━━━━━━━━
     FP16        12    0.00101447105       4.47822458e-05      195.472168    passed
    ───────  ───────  ───────────────  ───────────────────  ──────────────  ────────
     BF16        12    0.00883960724       3.57312470e-04      1585.35059    passed
    
  10. FED4 commented on Jun 16, 2026

    @FED4
    Contributor

    Quick scope check for #39:

    On the MI300/gfx942 machine, PyTorch SDPA is already using ROCm FlashAttention via AOTriton.

    I added RocmFlashAttentionOp for external FlashAttention 2 Triton AMD and verified correctness, but it did not consistently beat PyTorch SDPA/AOTriton in the tested shapes.

    One-off benchmark
    This is a quick sanity benchmark, not a formal performance claim. It compares the new default ROCm attention path (NativeAttentionOp, PyTorch SDPA/AOTriton) with external RocmFlashAttentionOp.

    With FLASH_ATTENTION_TRITON_AMD_AUTOTUNE=TRUE:

    dtype shape (B,S,H,D) causal NativeAttentionOp ms RocmFlashAttentionOp autotune ms native/external
    fp16 (1,128,4,64) false 0.020 0.101 0.19x
    fp16 (2,256,8,128) false 0.047 0.103 0.45x
    fp16 (1,512,8,256) false 0.099 0.116 0.85x
    fp16 (1,1024,16,64) false 0.118 0.111 1.07x
    fp16 (1,2048,16,64) false 0.311 0.344 0.90x
    bf16 (1,512,8,256) false 0.100 0.114 0.87x
    bf16 (1,1024,16,64) false 0.136 0.116 1.17x
    bf16 (1,2048,16,64) false 0.345 0.377 0.92x

    External FlashAttention wins on some mid-size cases, but not broadly enough to be the default ROCm dispatch target.

    Should this PR keep PyTorch SDPA/AOTriton as the default ROCm attention path and make external flash-attn opt-in, or should external RocmFlashAttentionOp be the default to match the original issue wording?

    If this needs to beat SDPA/AOTriton, I’ll look into CK next. I understand that’s the project goal, but not necessarily required for this issue.

  11. FED4 commented on Jun 18, 2026

    @FED4
    Contributor

    I will still use rocm flash attention as default to keep everything simple in #104.
    Permformance issue can be looked into in a separate PR.

  12. Flink-ddd commented on Jun 18, 2026

    @Flink-ddd
    CollaboratorAuthor

    Hi @FED4 Thanks for your update, I think we can select the external RocmFlashAttentionOp as the default backend for this PR.

    This is a solid engineering decision and aligns well with the architectural evolution in frameworks like vLLM. Prioritizing API parity with the CUDA FlashAttention wrapper and establishing a unified, stable dispatch layer ensures our system's correctness first.

    Let's keep things simple and get the baseline merged in #104. We can definitely track the performance edge of PyTorch SDPA/AOTriton for those specific mid-size shapes in a separate issue. Down the road, we might even consider introducing shape-based dynamic dispatch to get the best of both worlds.

    Fantastic job on the thorough benchmarking and validation!

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

component: kernelsTasks involving the development of CUDA and Triton underlying operatorsplatform: rocmSpecific tasks specific to AMD graphics cards (such as CK, bpreshuffle/FA)priority: highSevere congestion issues require the highest priority for resolution.type: ci-cdModify GitHub Actions, automated tests, and packaging/deployment tasks.type: performancePerformance optimization tasks aimed at increasing throughput and reducing latency etc.

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions