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
137 changes: 127 additions & 10 deletions aiter/ops/triton/attention/sparse_mla.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
)
from aiter.ops.triton.utils._triton import arch_info
from aiter.ops.triton.utils.common_utils import max_addressable_bytes
from aiter.ops.triton.utils.config_utils import load_config_json, resolve_config_dir
from aiter.ops.triton.utils.device_info import get_num_sms
from aiter.ops.triton.utils.logger import AiterTritonLogger

Expand Down Expand Up @@ -61,6 +62,105 @@ def _cache_pointers(fmt, kv, d_qk, kv_scale):
return u8, u8.view(torch.bfloat16), u8.view(torch.float32), u8.shape[1]


SUPPORTED_ARCHS = ("gfx942", "gfx950")

# The kernel reads every fp8 byte, in q, the cache and the dot operands, as OCP
# e4m3, which is gfx950's native fp8. gfx942's is fnuz, and the kernel does not
# decode it yet, so gfx942 takes bf16 q and a bf16 cache only.
FP8_ARCHS = ("gfx950",)
Comment thread
Copilot marked this conversation as resolved.

# The packed caches (fp8_dsv4_mla, fp8_g64) and the SWA+top-k two-loop do not
# reach the kernel below: they route to pa_decode_sparse, whose packed driver is
# gfx950-only. Everywhere else they land in its fallback path, which reads a
# plain grouped fp8 pool rather than these records, so it rejects them on dtype.
PACKED_ARCHS = ("gfx950",)


def _get_config(arch: str | None = None) -> dict:
"""The _sparse_mla launch config published for arch, the running one by default.

BLOCK_K is per arch because gfx950's tile does not fit gfx942's 64 KB of
LDS. num_warps is its own entry rather than BLOCK_K // 16, so the smaller
tile does not halve the warps too.
"""
cfg_dir = resolve_config_dir("attention", "SPARSE_MLA", backend="gluon", arch=arch)
return dict(load_config_json(f"{cfg_dir}/DEFAULT.json")["_sparse_mla"])


def _check_packed_arch(arch: str) -> None:
"""Packed caches and the two-loop are gfx950-only, whatever SUPPORTED_ARCHS says.

Left alone these reach pa_decode_sparse and fail inside it on a cache dtype,
which says nothing about the arch being the reason.
"""
if arch not in PACKED_ARCHS:
flat = "bf16, fp8_scalar or fp8_dsv32_mla" if arch in FP8_ARCHS else "bf16"
raise ValueError(
f"the fp8_dsv4_mla and fp8_g64 caches and the SWA+top-k two-loop are "
f"{'/'.join(PACKED_ARCHS)}-only and have no implementation on {arch}. "
f"Use the flat {flat} cache there."
)


def _check_fp8_arch(arch: str, fmt: str, q_dtype: torch.dtype) -> None:
"""fp8 q and caches only where OCP e4m3 is the native fp8.

A cache arrives as bytes, usually behind a uint8 view, so a gfx942 cache in
its native fnuz cannot be told apart from OCP here. Read as OCP it comes out
2x too large, and NaN wherever the quantizer saturated at 240.
"""
if arch in FP8_ARCHS:
return
got = [
what
for what, is_fp8 in (
(f"q is {q_dtype}", q_dtype.itemsize == 1),
(f"the cache is {fmt}", fmt != "bf16"),
)
if is_fp8
]
if got:
raise ValueError(
f"sparse_mla_fwd takes bf16 q and a bf16 cache on {arch}, but "
f"{' and '.join(got)}. The kernel reads fp8 as OCP e4m3, and {arch}'s "
"native fp8 is fnuz, which it does not decode yet."
)


# Row pitch padding, and the scratch the kernel takes beyond the tiles. Both
# hold only for bf16 tiles with the async path off, which is every gfx942
# launch: fp8 dots are rejected there and its launch config keeps ASYNC_LDS off.
# A nonzero KV_LDS_PAD replaces the pad on the KV tile alone.
_LDS_PAD = 8
_LDS_SCRATCH_PER_BLOCK_K = 32
Comment thread
jin-amd marked this conversation as resolved.


def _check_lds_budget(arch, block_k, kv_lora_rank, qk_rope_head_dim, kv_lds_pad):
"""Reject a geometry whose tiles cannot fit, naming what would.

kv_lds_pad is the launch's KV_LDS_PAD, so prefill is checked at the wider
pitch it stages at. Left to the launcher this surfaces as an opaque
OutOfResources. gfx942 is the only arch checked: its config already takes
the smaller of the two published tiles, so a latent too wide to fit has
nowhere left to go. gfx950 is left to the launcher, as before.
"""
if arch != "gfx942":
return
budget = arch_info._LDS_CAP_BYTES[arch]
rope = block_k * (qk_rope_head_dim + _LDS_PAD) * 2 if qk_rope_head_dim else 0
need = (
block_k * (kv_lora_rank + (kv_lds_pad or _LDS_PAD)) * 2
+ rope
+ _LDS_SCRATCH_PER_BLOCK_K * block_k
)
if need > budget:
raise ValueError(
f"kv_lora_rank={kv_lora_rank} with qk_rope_head_dim="
f"{qk_rope_head_dim} needs {need} B of LDS at BLOCK_K={block_k}, "
f"over {arch}'s {budget} B. This geometry needs gfx950's 160 KB."
Comment thread
jin-amd marked this conversation as resolved.
)


def _mla_num_splits(
num_queries: int, heads_blocks: int, avg_topk: float, block_k: int = 64
) -> int:
Expand Down Expand Up @@ -104,6 +204,8 @@ def _async_launch_config(
and uni_tile
and not has_invalid
and not has_extra
# The async tiles below are sized for gfx950's LDS.
and arch_info.get_arch() != "gfx942"
)
workgroups = num_queries * heads_blocks * max(1, num_splits)
num_sms = get_num_sms()
Expand All @@ -115,13 +217,19 @@ def _async_launch_config(
return enabled, (128 if enabled else block_k), waves_per_eu


def _resolve_dot_precision(dot_precision: str, fmt: str) -> bool:
def _resolve_dot_precision(dot_precision: str, fmt: str, arch: str) -> bool:
if dot_precision not in ("bf16", "fp8"):
raise ValueError(
f"dot_precision must be 'bf16' or 'fp8', got {dot_precision!r}"
)
if dot_precision == "bf16":
return False
if arch not in FP8_ARCHS:
raise ValueError(
f"dot_precision='fp8' is not supported on {arch}: the kernel feeds the "
f"matrix core OCP e4m3, but {arch}'s native fp8 is fnuz. Use "
"dot_precision='bf16'."
)
if fmt == "fp8_dsv32_mla":
raise ValueError(
"dot_precision='fp8' does not support the fp8_dsv32_mla cache."
Expand Down Expand Up @@ -291,6 +399,7 @@ def _forward_paged(
"""dsv4 and the SWA+top-k two-loop, until the two launchers merge."""
from aiter.ops.triton.attention.pa_decode_sparse import pa_decode_sparse

_check_packed_arch(arch_info.get_arch())
unsupported = [
name
for name, asked in (
Expand Down Expand Up @@ -352,7 +461,9 @@ def sparse_mla_fwd(

Supported KV cache formats. The format is inferred from kv_buffer's shape,
dtype and kv_scale; each row gives what the caller has to pass. R is the
QK width, kv_lora_rank + qk_rope_head_dim.
QK width, kv_lora_rank + qk_rope_head_dim. Every fp8 format, like fp8 q, is
gfx950-only: the kernel reads fp8 as OCP e4m3, and gfx942's native fp8 is
fnuz.

format kv_buffer kv_scale geometry args
bf16 [slots, R], [nb, block, R], None as the model
Expand Down Expand Up @@ -404,10 +515,11 @@ def sparse_mla_fwd(
dot_precision: what the QK and PV matrix-core ops run in.

"bf16" (default): the KV tile is dequantized to bf16 on its way
into LDS and both dots are bf16. Works with every cache format.
into LDS and both dots are bf16. Works with every cache format
the arch takes: all of them on gfx950, bf16 alone on gfx942.
"fp8": the cache's own code points go to the fp8 matrix core with no
dequant, and the per-tensor scale folds outside the tile loop.
tensor scale fp8 kv cache only.
tensor scale fp8 kv cache only; gfx950 only.

q is adapted to the choice. bf16 q is quantized in the kernel
prologue, one scale per (query, head-block) tile; fp8 q is passed
Expand Down Expand Up @@ -487,7 +599,9 @@ def sparse_mla_fwd(
extra_indptr,
extra_indices,
)
assert arch_info.get_arch() == "gfx950", "sparse_mla_fwd is gfx950-only"
arch = arch_info.get_arch()
assert arch in SUPPORTED_ARCHS, f"sparse_mla_fwd does not support {arch}"
Comment thread
jin-amd marked this conversation as resolved.
_check_fp8_arch(arch, fmt, q.dtype)
q_is_fp8 = q.dtype == torch.float8_e4m3fn
if q.dtype not in (torch.bfloat16, torch.float8_e4m3fn):
raise ValueError(
Expand All @@ -509,7 +623,7 @@ def sparse_mla_fwd(
)
q_scale = q_scale.reshape(1).to(torch.float32).contiguous()
cache, alt, scl, block_size = _cache_pointers(fmt, kv_buffer, d_qk, kv_scale)
fp8_dots = _resolve_dot_precision(dot_precision, fmt)
fp8_dots = _resolve_dot_precision(dot_precision, fmt, arch)
if not q_is_fp8:
# q_scale describes an fp8 q's encoding. With bf16 q the kernel quantizes
# per (query, head-block) tile when the dots are fp8, so a caller-supplied
Expand All @@ -527,8 +641,9 @@ def sparse_mla_fwd(
# Tuned launch config (gfx950 / MI355). H < 16 runs natively at
# BLOCK_M = next_pow2(H) instead of padding heads
block_m = 16 if num_heads >= 16 else max(8, 1 << (num_heads - 1).bit_length())
block_k = 64
num_warps = block_k // 16
cfg = _get_config()
block_k = cfg["BLOCK_K"]
num_warps = cfg["num_warps"]

num_rows = cache.shape[0] * block_size if cache.ndim >= 2 else cache.shape[0]
avg_topk = kv_indices.numel() / max(1, num_queries)
Expand Down Expand Up @@ -604,6 +719,9 @@ def sparse_mla_fwd(
has_extra=False,
block_k=block_k,
)
# bf16-staged tiles only; fp8 dots stage raw fp8 in their own layout
kv_lds_pad = 16 if num_queries >= _PREFILL_MIN_ROWS and not fp8_dots else 0
_check_lds_budget(arch, block_k, kv_lora_rank, qk_rope_head_dim, kv_lds_pad)

# Q is read once per query without split-K, and re-read by every split
q_cache = ".cg" if num_splits == 1 else ""
Expand Down Expand Up @@ -672,8 +790,7 @@ def sparse_mla_fwd(
FP8_MFMA=fp8_dots,
ASYNC_LDS=async_lds_on,
GATHER_CACHE="",
# bf16-staged tiles only; fp8 dots stage raw fp8 in their own layout
KV_LDS_PAD=16 if num_queries >= _PREFILL_MIN_ROWS and not fp8_dots else 0,
KV_LDS_PAD=kv_lds_pad,
q_scl_ptr=q_scale,
Q_FP8=q_is_fp8,
lse_ptr=lse,
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
{
"_sparse_mla": {
"BLOCK_K": 32,
"num_warps": 4
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
{
"_sparse_mla": {
"BLOCK_K": 64,
"num_warps": 4
}
}
33 changes: 26 additions & 7 deletions op_tests/op_benchmarks/triton/bench_sparse_mla.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: MIT
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.

"""Benchmark for sparse_mla_fwd (gfx950 gluon, separated-rope MLA).
"""Benchmark for sparse_mla_fwd (gluon, separated-rope MLA; gfx950 and gfx942).

The cache is flushed between iterations by default. Leaving it warm lets the loop
re-read its KV and flatters decode shapes badly.
Expand All @@ -19,7 +19,11 @@
from torch.autograd import DeviceType
from torch.profiler import ProfilerActivity, profile

from aiter.ops.triton.attention.sparse_mla import sparse_mla_fwd
from aiter.ops.triton.attention.sparse_mla import (
FP8_ARCHS,
SUPPORTED_ARCHS,
sparse_mla_fwd,
)
from aiter.ops.triton.utils._triton import arch_info
from aiter.ops.triton.utils.types import get_fp8_e4m3_dtype
from op_tests.op_benchmarks.triton.utils.benchmark_utils import (
Expand Down Expand Up @@ -152,13 +156,28 @@ def run_benchmark(args):
for tokens in args.num_tokens
]

# sparse_mla_fwd raises on fp8 where the arch's native fp8 is not OCP e4m3,
# and triton's harness does not catch it, so drop the series instead.
dot_vals = ["bf16"]
dot_names = ["bf16 dots"]
dot_styles = [("green", "-")]
if arch_info.get_arch() in FP8_ARCHS:
dot_vals.append("fp8")
dot_names.append("fp8 dots")
dot_styles.append(("blue", "-"))
else:
print(
f"note: skipping the fp8-dot series, {arch_info.get_arch()}'s native "
"fp8 is fnuz and the kernel reads OCP e4m3"
)

benchmark = triton.testing.Benchmark(
x_names=["phase", "num_seqs", "num_tokens", "num_heads", "context", "topk"],
x_vals=x_vals_list,
line_arg="dots",
line_vals=["bf16", "fp8"],
line_names=["bf16 dots", "fp8 dots"],
styles=[("green", "-"), ("blue", "-")],
line_vals=dot_vals,
line_names=dot_names,
styles=dot_styles,
ylabel=ylabel,
plot_name=get_caller_name_no_ext(),
args={"metric": args.metric},
Expand Down Expand Up @@ -263,8 +282,8 @@ def main():
f"--num_tokens {over} exceeds --context {args.context}: a sequence "
"cannot prefill more tokens than its context holds"
)
if arch_info.get_arch() != "gfx950":
raise SystemExit(f"sparse_mla_fwd is gfx950-only (got {arch_info.get_arch()})")
if arch_info.get_arch() not in SUPPORTED_ARCHS:
raise SystemExit(f"sparse_mla_fwd does not support {arch_info.get_arch()}")
run_benchmark(args)


Expand Down
Loading
Loading