Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@ kernel for non-attention ops). **AITER** = ROCm AITER backend.
| **POD attention** | ✅ `fa2` | — | HIP | MHA / GQA / MQA; single + batch variants (`PODWithPagedKVCacheWrapper`, `BatchPODWithPagedKVCacheWrapper`); JIT-only (excluded from AOT, same as upstream CUDA) |
| **RoPE (positional encoding)** | ✅ `native` | ✅ | **AITER** for the cos/sin-cache path on gfx942/gfx950; else **HIP `native`** | LLaMA-style + LLaMA 3.1 scaling; fused RoPE + fp8 quant + paged-KV append (E4M3FNUZ, E5M2FNUZ). AITER backend covers `apply_rope_with_cos_sin_cache` and its inplace variant via AITER's C++ `rope_cached_positions_2c_fwd_impl` (linked at the C++ level, no runtime `import aiter`); cos/sin passed as float32 |
| **Paged KV-cache append** | ✅ `native` | ✅ | **AITER** when `fp16/bf16` + `NHD` + gfx942/gfx950 + AITER importable; else **HIP `native`** | `append_paged_kv_cache`; fp8 KV-cache supported on the HIP path |
| **RMSNorm** | ✅ `native` | ✅ | **AITER** for 2-D inputs on gfx942/gfx950; else **HIP `native`** (3-D inputs or AITER unavailable) | `rmsnorm`; AITER's C++ `rms_norm` (linked at the C++ level, no runtime `import aiter`); fp16/bf16, 2-D only, slightly lower precision at `hidden_size >= 1024` |
| **RMSNorm** | ✅ `native` | ✅ | **AITER** for 2-D fp16/bf16 inputs on gfx942/gfx950; else **HIP `native`** (3-D inputs, fp32, or AITER unavailable) | `rmsnorm`; AITER's C++ CK `rmsnorm2d` (the `rmsnorm2d_fwd` entry point, linked at the C++ level, no runtime `import aiter`); fp16/bf16, 2-D only, weight dtype must match input, slightly lower precision at `hidden_size >= 1024` |
| **Fused add RMSNorm** | ✅ `native` | ✅ | **AITER** on gfx942/gfx950; else **HIP `native`** | `fused_add_rmsnorm`; AITER's C++ CK `rmsnorm2d_with_add` (linked at the C++ level, no runtime `import aiter`); 2-D only, slightly lower precision at `hidden_size >= 1024` |
| **LayerNorm / Gemma RMSNorm** | ✅ | — | HIP | |
| **Sampling** | ✅ | — | HIP | Top-K / Top-P / Min-P / OnlineSoftmax / SamplingFromLogits |
Expand Down Expand Up @@ -328,7 +328,7 @@ prefill/decode), `backend="native"` for non-attention ops
The `rmsnorm`, `fused_add_rmsnorm`, `silu_and_mul`, and `rope`
(cos/sin-cache) AITER backends are integrated at the **C++ level**:
FlashInfer's JIT compiles a small HIP shim that calls AITER's C++ kernels
(`rms_norm`, `rmsnorm2d_with_add`, `aiter::silu_and_mul`,
(`rmsnorm2d`, `rmsnorm2d_with_add`, `aiter::silu_and_mul`,
`rope_cached_positions_2c_fwd_impl`) directly and links a symbol-visible
AITER `.so` — there is no runtime `import aiter` on these paths. The first
JIT build of each op builds the corresponding AITER module once with
Expand All @@ -342,8 +342,9 @@ shape-based gating to the other ops.

Backend-specific exceptions to "auto picks AITER when supported":

* `rmsnorm`: `backend="auto"` picks the AITER C++ path (`rms_norm`) only
for 2-D inputs; 3-D inputs fall back to the HIP `native` kernel.
* `rmsnorm`: `backend="auto"` picks the AITER C++ path (CK `rmsnorm2d`) only
for 2-D fp16/bf16 inputs whose weight dtype matches; 3-D inputs, fp32, or a
mismatched weight dtype fall back to the HIP `native` kernel.
* `batch_decode`: `use_cuda_graph=True` or `use_tensor_cores=True`
force `auto` back to `fa2` (AITER decode does not support either),
and `pos_encoding_mode != "NONE"` raises under `backend="aiter"`.
Expand Down
17 changes: 11 additions & 6 deletions flashinfer/csrc_rocm/norm_aiter.cu
Original file line number Diff line number Diff line change
@@ -1,10 +1,10 @@
// SPDX-FileCopyrightText: 2026 Advanced Micro Devices, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// PyTorch entry point for AITER's CK fused-add RMSNorm (rmsnorm2d_with_add).
// FlashInfer links the symbol-visible AITER module (see
// flashinfer/jit/aiter_source.py) and calls the kernel directly — no Python
// `import aiter` at runtime.
// PyTorch entry points for AITER's CK RMSNorm kernels: plain rmsnorm
// (rmsnorm2d) and fused-add rmsnorm (rmsnorm2d_with_add). FlashInfer links the
// symbol-visible AITER module (see flashinfer/jit/aiter_source.py) and calls the
// kernels directly — no Python `import aiter` at runtime.
//
// FlashInfer's fused_add_rmsnorm is in-place:
// residual = input + residual; input = rmsnorm(residual) * weight
Expand All @@ -22,7 +22,10 @@
void rmsnorm2d_with_add(at::Tensor& out, at::Tensor& input, at::Tensor& residual_in,
at::Tensor& residual_out, at::Tensor& weight, double epsilon,
int use_model_sensitive_rmsnorm);
void rms_norm(at::Tensor& out, at::Tensor& input, at::Tensor& weight, double epsilon);
// CK 2D forward (the symbol the AITER `rmsnorm2d_fwd` pybind name binds to);
// the in-place overload writes `out` directly.
void rmsnorm2d(at::Tensor& out, at::Tensor& input, at::Tensor& weight, double epsilon,
int use_model_sensitive_rmsnorm);

void fused_add_rmsnorm_aiter(at::Tensor input, at::Tensor residual, at::Tensor weight, double eps) {
const c10::hip::OptionalHIPGuardMasqueradingAsCUDA device_guard(input.device());
Expand All @@ -36,5 +39,7 @@ void fused_add_rmsnorm_aiter(at::Tensor input, at::Tensor residual, at::Tensor w

void rmsnorm_aiter(at::Tensor out, at::Tensor input, at::Tensor weight, double eps) {
const c10::hip::OptionalHIPGuardMasqueradingAsCUDA device_guard(input.device());
rms_norm(out, input, weight, eps);
// CK expects weight as [1, n]; FlashInfer passes [n] (see fused-add note above).
at::Tensor weight2d = weight.reshape({1, -1});
rmsnorm2d(out, input, weight2d, eps, /*use_model_sensitive_rmsnorm=*/0);
}
44 changes: 33 additions & 11 deletions flashinfer/norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,15 +37,24 @@ def get_norm_aiter_module():

return gen_norm_aiter_module().build_and_load()

def _auto_select_norm_backend(input: torch.Tensor) -> str:
# auto routes plain rmsnorm to the C++ AITER kernel on supported devices
# (2D inputs only) and falls back to native everywhere else (3D inputs, or
# when AITER is not installed, so auto never raises). Note: AITER's rms_norm
# uses lower-precision reductions that exceed the flashinfer test tolerance
# at hidden_size >= 1024 (fp16 atol ~4e-3, bf16 ~7e-2).
def _auto_select_norm_backend(input: torch.Tensor, weight: torch.Tensor) -> str:
# auto routes plain rmsnorm to the C++ AITER kernel only when the input is
# 2D fp16/bf16 with a matching weight dtype (the CK rmsnorm2d kernel rejects
# other ranks/dtypes and reads weight with the input dtype) and AITER is
# available; everything else falls back to native. (fp32 is unsupported by
# both the AITER and native ROCm kernels, so it still raises downstream —
# routing it to native just keeps the error consistent with the default
# kernel rather than exposing a CK-specific message.) Note: AITER's CK
# rmsnorm2d uses lower-precision reductions that exceed the flashinfer test
# tolerance at hidden_size >= 1024 (fp16 atol ~4e-3, bf16 ~7e-2).
from .aiter_utils import is_aiter_available

if input.ndim == 2 and is_aiter_available(input.device):
if (
input.ndim == 2
and input.dtype in (torch.float16, torch.bfloat16)
and weight.dtype == input.dtype
and is_aiter_available(input.device)
):
return "aiter"
return "native"

Expand Down Expand Up @@ -86,10 +95,10 @@ def rmsnorm(
<https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html#programmatic-dependent-launch-and-synchronization>`_
backend: str
Kernel backend to use. ``"auto"`` (default) selects the best available backend:
the AITER C++ kernel for 2D inputs on supported ROCm devices, else native.
the AITER C++ kernel for 2D fp16/bf16 inputs on supported ROCm devices, else native.
``"native"`` uses the FlashInfer JIT kernel on all platforms.
``"aiter"`` uses AMD AITER's ``rms_norm`` C++ kernel — ROCm (gfx942/gfx950)
only, 2D inputs only. Precision is slightly lower than ``"native"`` at
``"aiter"`` uses AMD AITER's CK ``rmsnorm2d`` C++ kernel — ROCm (gfx942/gfx950)
only, 2D fp16/bf16 inputs only. Precision is slightly lower than ``"native"`` at
``hidden_size >= 1024`` (fp16 atol ~4e-3, bf16 ~7e-2).

Returns
Expand All @@ -98,7 +107,9 @@ def rmsnorm(
Normalized tensor, 2D shape (batch_size, hidden_size) or 3D shape (batch_size, num_heads, hidden_size).
"""
if IS_HIP:
_backend = backend if backend != "auto" else _auto_select_norm_backend(input)
_backend = (
backend if backend != "auto" else _auto_select_norm_backend(input, weight)
)
if _backend == "aiter":
from .aiter_utils import require_aiter

Expand All @@ -108,6 +119,17 @@ def rmsnorm(
f"AITER rmsnorm only supports 2D inputs; got {input.ndim}D. "
"Use backend='native' for 3D inputs."
)
if input.dtype not in (torch.float16, torch.bfloat16):
raise ValueError(
f"AITER rmsnorm only supports float16/bfloat16 inputs; got {input.dtype}."
)
if weight.dtype != input.dtype:
# CK rmsnorm2d derives a single dtype from input and reads weight
# bytes with it; a mismatched weight dtype silently yields NaN/garbage.
raise ValueError(
f"AITER rmsnorm requires weight.dtype == input.dtype; got "
f"weight {weight.dtype} vs input {input.dtype}."
)
if out is None:
out = torch.empty_like(input)
get_norm_aiter_module().rmsnorm_aiter(out, input, weight, eps)
Expand Down
43 changes: 38 additions & 5 deletions tests/rocm_tests/test_rmsnorm_aiter_hip.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
# SPDX-FileCopyrightText: 2026 Advanced Micro Devices, Inc.
# SPDX-License-Identifier: Apache-2.0
#
# Tests for the AITER rms_norm backend exposed via flashinfer.rmsnorm(backend="aiter").
# Tests for the AITER CK rmsnorm2d backend exposed via flashinfer.rmsnorm(backend="aiter").
#
# Note on tolerances: AITER rms_norm uses lower-precision reductions than the
# Note on tolerances: AITER's CK rmsnorm2d uses lower-precision reductions than the
# flashinfer native JIT kernel. For production inference these differences are
# negligible, but they exceed the native kernel's test tolerance (fp16 atol=1e-3,
# bf16 atol=1.6e-2). The tolerances below reflect AITER's actual precision.
Expand Down Expand Up @@ -43,14 +43,47 @@ def test_rmsnorm_aiter_vs_ref(dtype, hidden_size, batch_size):

@requires_aiter
def test_rmsnorm_auto_backend_selects_aiter_for_2d():
"""auto backend on gfx942/950 routes 2D inputs to AITER and 3D inputs to native."""
"""auto routes 2D fp16/bf16 (matching weight) to AITER; 3D, fp32, or a
mismatched weight dtype to native."""
from flashinfer.norm import _auto_select_norm_backend

device = torch.device("cuda:0")
x2d = torch.randn(8, 128, dtype=torch.float16, device=device)
w = torch.randn(128, dtype=torch.float16, device=device)
x3d = torch.randn(8, 4, 128, dtype=torch.float16, device=device)
assert _auto_select_norm_backend(x2d) == "aiter"
assert _auto_select_norm_backend(x3d) == "native"
x2d_fp32 = torch.randn(8, 128, dtype=torch.float32, device=device)
w_fp32 = torch.randn(128, dtype=torch.float32, device=device)
assert _auto_select_norm_backend(x2d, w) == "aiter"
assert _auto_select_norm_backend(x3d, w) == "native"
# CK rmsnorm2d rejects fp32, so auto must not select it (routes to native).
assert _auto_select_norm_backend(x2d_fp32, w_fp32) == "native"
# CK reads weight with the input dtype, so a mismatch must not select AITER.
assert _auto_select_norm_backend(x2d, w_fp32) == "native"


@requires_aiter
def test_rmsnorm_aiter_rejects_fp32():
"""backend='aiter' with fp32 raises a clear FlashInfer error, not the deep CK one.

fp32 is unsupported by both ROCm rmsnorm kernels; the Python-level guard
rejects it before reaching AITER so the message is actionable.
"""
device = torch.device("cuda:0")
x = torch.randn(8, 128, dtype=torch.float32, device=device)
w = torch.randn(128, dtype=torch.float32, device=device)
with pytest.raises(ValueError, match="float16/bfloat16"):
flashinfer.rmsnorm(x, w, backend="aiter")


@requires_aiter
def test_rmsnorm_aiter_rejects_weight_dtype_mismatch():
"""backend='aiter' with weight.dtype != input.dtype raises rather than silently
producing NaN/garbage (CK rmsnorm2d reads weight bytes with the input dtype)."""
device = torch.device("cuda:0")
x = torch.randn(8, 128, dtype=torch.float16, device=device)
w = torch.randn(128, dtype=torch.float32, device=device)
with pytest.raises(ValueError, match="weight.dtype == input.dtype"):
flashinfer.rmsnorm(x, w, backend="aiter")


@requires_aiter
Expand Down
Loading