diff --git a/op_tests/test_mha_d256_logits.py b/op_tests/test_mha_d256_logits.py new file mode 100644 index 0000000000..4ad28b11d5 --- /dev/null +++ b/op_tests/test_mha_d256_logits.py @@ -0,0 +1,328 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. + +"""flash_attn_func sweep: bf16, BSHD, Q/K/V head_dim=256. + +Dense MHA prefill (same as op_tests/test_mha.py / flash_attn_func). There is no +decode phase. Shape defaults follow the M x N serving axes in ROCm/aiter#5434: + + seqlen_q 1024, 2048, 4096, 8192, 16384 + seqlen_k 4096, 16384, 65664 (skips seqlen_k < seqlen_q) + nheads 32, 64 + batch 1, 16 + +Shapes that hit HIP illegal access on MI355X (b>=64, and b=16/h=64/sk=131072) +are not in the default product. + +Layout is BSHD [B, S, H, D]; D is fixed at 256. No dropout, bias, alibi, or +local window. Forward only. + +The full [B, H, Sq, Sk] fp32 score tensor does not fit at the large shapes, so +accuracy is checked on a query-row sample (grid edges first). + +Examples: + python op_tests/test_mha_d256_logits.py + python op_tests/test_mha_d256_logits.py -n 32 -b 1 -q 1024 -k 4096 + python op_tests/test_mha_d256_logits.py --ref +""" + +import argparse +import itertools + +import pandas as pd +import torch + +import aiter +from aiter import dtypes +from aiter.jit.utils.chip_info import get_gfx +from aiter.test_common import benchmark, checkAllclose, run_perftest +from aiter.test_mha_common import attention_ref + +torch.set_default_device("cuda") + +SUPPORTED_GFX = ["gfx942", "gfx950"] +HEAD_DIM = 256 +_REF_SCORE_BYTES = 1 << 30 +_REF_MAX_ROWS = 64 + + +def _ref_rows(seqlen_q, batch, nheads, seqlen_k): + budget = max(1, _REF_SCORE_BYTES // (max(batch, 1) * nheads * seqlen_k * 4)) + want = max(1, min(seqlen_q, _REF_MAX_ROWS, budget)) + spread = torch.linspace(0, seqlen_q - 1, steps=want).round().long().tolist() + rows = [] + for r in [ + 0, + 1, + seqlen_q // 2, + seqlen_q // 2 + 1, + seqlen_q - 2, + seqlen_q - 1, + ] + spread: + if 0 <= r < seqlen_q and r not in rows: + rows.append(r) + if len(rows) >= want: + break + return torch.tensor(sorted(rows), dtype=torch.long) + + +def _causal_bias(rows, seqlen_q, seqlen_k, device): + """Bottom-right causal visibility for sampled query rows [R, Sk].""" + q_pos = rows.to(device) + k_pos = torch.arange(seqlen_k, device=device) + visible = k_pos[None, :] <= (seqlen_k - seqlen_q + q_pos)[:, None] + bias = torch.zeros(1, 1, rows.numel(), seqlen_k, device=device) + bias.masked_fill_(~visible[None, None], float("-inf")) + return bias + + +def run_torch(q, k, v, causal=True, upcast=True, reorder_ops=False, attn_bias=None): + out, _, softmax_lse = attention_ref( + q, + k, + v, + attn_bias=attn_bias, + causal=causal if attn_bias is None else False, + upcast=upcast, + reorder_ops=reorder_ops, + ) + return out, softmax_lse + + +def run_flash(q, k, v, causal=True): + ret, us = run_perftest( + aiter.flash_attn_func, + q, + k, + v, + 0.0, + None, + causal, + (-1, -1, 0), + None, + None, + False, + return_lse=True, + return_attn_probs=False, + how_v3_bf16_cvt=2, + num_rotate_args=1, + ) + if not isinstance(ret, (tuple, list)): + return ret, None, us + out = ret[0] + softmax_lse = ret[1] if len(ret) > 1 else None + return out, softmax_lse, us + + +@benchmark() +def test_flash_attn_bshd_d256( + batch_size, + nheads, + seqlen_q, + seqlen_k, + causal, + check_ref=False, +): + torch.manual_seed(0) + torch.cuda.empty_cache() + dtype = dtypes.bf16 + d = HEAD_DIM + + q = torch.randn(batch_size, seqlen_q, nheads, d, dtype=dtype) + k = torch.randn(batch_size, seqlen_k, nheads, d, dtype=dtype) + v = torch.randn(batch_size, seqlen_k, nheads, d, dtype=dtype) + + out, softmax_lse, us = run_flash(q, k, v, causal=causal) + + err = None + ref_n = 0 + if check_ref: + rows = _ref_rows(seqlen_q, batch_size, nheads, seqlen_k) + ref_n = int(rows.numel()) + q_s = q[:, rows] + attn_bias = _causal_bias(rows, seqlen_q, seqlen_k, q.device) if causal else None + out_ref, lse_ref = run_torch(q_s, k, v, causal=causal, attn_bias=attn_bias) + out_pt, lse_pt = run_torch( + q_s, + k, + v, + causal=causal, + attn_bias=attn_bias, + upcast=False, + reorder_ops=True, + ) + + got = out[:, rows] + out_tol = max(2 * (out_pt - out_ref).abs().max().item(), 0.01) + err = checkAllclose( + out_ref, + got, + atol=out_tol, + rtol=0.0, + msg=f"flash_attn BSHD d={d} b={batch_size} h={nheads} " + f"sq={seqlen_q} sk={seqlen_k} causal={causal}", + ) + lse_got = softmax_lse[:, :, rows] if softmax_lse is not None else None + if lse_got is not None: + checkAllclose( + lse_ref.float(), + lse_got.float(), + atol=max( + 2 * (lse_pt.float() - lse_ref.float()).abs().max().item(), 0.01 + ), + rtol=0.0, + msg=f"flash_attn lse b={batch_size} h={nheads} sq={seqlen_q} sk={seqlen_k}", + ) + else: + aiter.logger.info( + "flash_attn BSHD d=%s b=%s h=%s sq=%s sk=%s [no-ref] %.2f us", + d, + batch_size, + nheads, + seqlen_q, + seqlen_k, + us, + ) + + flops = ( + batch_size + * nheads + * (seqlen_q * seqlen_k * d * 2 + seqlen_q * seqlen_k * d * 2) + ) + if causal: + flops = flops / 2 + nbytes = ( + batch_size + * nheads + * 2 + * (seqlen_q * d + seqlen_k * d + seqlen_k * d + seqlen_q * d) + ) + + return { + "gfx": get_gfx(), + "ref_rows": ref_n, + "fwd us": us, + "fwd TFLOPS": flops / us / 1e6, + "fwd TB/s": nbytes / us / 1e6, + "fwd err": err, + } + + +def _summarize(name, rows): + if not rows: + return + df = pd.DataFrame(rows) + keep = [ + c + for c in ( + "batch_size", + "nheads", + "seqlen_q", + "seqlen_k", + "causal", + "gfx", + "ref_rows", + "fwd us", + "fwd TFLOPS", + "fwd TB/s", + "fwd err", + ) + if c in df.columns + ] + aiter.logger.info( + "%s summary (markdown):\n%s", name, df[keep].to_markdown(index=False) + ) + + +def main(): + if get_gfx() not in SUPPORTED_GFX: + aiter.logger.warning("mha_d256_logits unsupported on %s; skipping", get_gfx()) + return + + parser = argparse.ArgumentParser( + formatter_class=argparse.RawTextHelpFormatter, + description="config input of test", + ) + parser.add_argument( + "-b", + "--batch_size", + type=int, + nargs="*", + default=[1, 16], + help="""Prefill batch size. + e.g.: -b 1 16""", + ) + parser.add_argument( + "-n", + "--nheads", + type=int, + nargs="*", + default=[32, 64], + help="""Q/K/V heads (MHA, nheads_k = nheads). + e.g.: -n 32""", + ) + parser.add_argument( + "-q", + "--seqlen_q", + type=int, + nargs="*", + default=[1024, 2048, 4096, 8192, 16384], + help="""Query length. + e.g.: -q 1024 4096""", + ) + parser.add_argument( + "-k", + "--seqlen_k", + type=int, + nargs="*", + default=[4096, 16384, 65664], + help="""KV length. Skips seqlen_k < seqlen_q. + e.g.: -k 4096 16384""", + ) + parser.add_argument( + "--causal", + action=argparse.BooleanOptionalAction, + default=True, + help="causal mask. Default: True.", + ) + parser.add_argument( + "--ref", + action=argparse.BooleanOptionalAction, + default=False, + help="Compare against torch golden. Default: False. Pass --ref to enable.", + ) + args = parser.parse_args() + + rows = [] + for batch_size, nheads, seqlen_q, seqlen_k in itertools.product( + args.batch_size, args.nheads, args.seqlen_q, args.seqlen_k + ): + if seqlen_k < seqlen_q: + continue + try: + rows.append( + test_flash_attn_bshd_d256( + batch_size, + nheads, + seqlen_q, + seqlen_k, + args.causal, + check_ref=args.ref, + ) + ) + except RuntimeError as e: + if "out of memory" not in str(e).lower(): + raise + aiter.logger.warning( + "OOM skip b=%s h=%s sq=%s sk=%s", + batch_size, + nheads, + seqlen_q, + seqlen_k, + ) + torch.cuda.empty_cache() + _summarize("mha_d256_logits", rows) + + +if __name__ == "__main__": + main() diff --git a/op_tests/test_mha_opus_d192_v128_logits.py b/op_tests/test_mha_opus_d192_v128_logits.py new file mode 100644 index 0000000000..7e700714b1 --- /dev/null +++ b/op_tests/test_mha_opus_d192_v128_logits.py @@ -0,0 +1,316 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. + +"""fmha_fwd_bf16_opus_fwd sweep: bf16, BSHD, D_QK=192 / D_V=128. + +Dense MHA prefill through the OPUS gfx950 kernel (same path as +test_flash_attn_func_opus_d192_v128 in op_tests/test_mha.py). There is no +decode phase. Shape defaults follow the M x N serving axes in ROCm/aiter#5434: + + seqlen_q 1024, 2048, 4096, 8192, 16384 + seqlen_k 4096, 16384, 65664 (skips seqlen_k < seqlen_q) + nheads 32, 64 + batch 1, 16, 64, 128 + +Shapes that OOM'd on MI355X (sk=131072; b=128/h=64/sk=65664; b=256) are +not in the default product. + +Layout is BSHD; Q/K last dim 192, V last dim 128. nheads_k = nheads. No dropout, +bias, alibi, or local window. Forward only. gfx950 only. + +The full [B, H, Sq, Sk] fp32 score tensor does not fit at the large shapes, so +output accuracy is checked on a query-row sample. LSE uses opus_ref_lse, which +already chunks over query rows. + +Examples: + python op_tests/test_mha_opus_d192_v128_logits.py + python op_tests/test_mha_opus_d192_v128_logits.py -n 32 -b 1 -q 1024 -k 4096 + python op_tests/test_mha_opus_d192_v128_logits.py --ref +""" + +import argparse +import itertools + +import pandas as pd +import torch + +import aiter +from aiter import dtypes +from aiter.jit.utils.chip_info import get_gfx +from aiter.ops.mha import fmha_fwd_bf16_opus_fwd +from aiter.test_common import benchmark, checkAllclose, run_perftest +from aiter.test_mha_common import attention_ref, opus_check_lse, opus_ref_lse + +torch.set_default_device("cuda") + +SUPPORTED_GFX = ["gfx950"] +D_QK = 192 +D_V = 128 +_REF_SCORE_BYTES = 1 << 30 +_REF_MAX_ROWS = 64 + + +def _ref_rows(seqlen_q, batch, nheads, seqlen_k): + budget = max(1, _REF_SCORE_BYTES // (max(batch, 1) * nheads * seqlen_k * 4)) + want = max(1, min(seqlen_q, _REF_MAX_ROWS, budget)) + spread = torch.linspace(0, seqlen_q - 1, steps=want).round().long().tolist() + rows = [] + for r in [ + 0, + 1, + seqlen_q // 2, + seqlen_q // 2 + 1, + seqlen_q - 2, + seqlen_q - 1, + ] + spread: + if 0 <= r < seqlen_q and r not in rows: + rows.append(r) + if len(rows) >= want: + break + return torch.tensor(sorted(rows), dtype=torch.long) + + +def _causal_bias(rows, seqlen_q, seqlen_k, device): + q_pos = rows.to(device) + k_pos = torch.arange(seqlen_k, device=device) + visible = k_pos[None, :] <= (seqlen_k - seqlen_q + q_pos)[:, None] + bias = torch.zeros(1, 1, rows.numel(), seqlen_k, device=device) + bias.masked_fill_(~visible[None, None], float("-inf")) + return bias + + +def run_torch(q, k, v, causal=True, upcast=True, reorder_ops=False, attn_bias=None): + out, _, softmax_lse = attention_ref( + q, + k, + v, + attn_bias=attn_bias, + causal=causal if attn_bias is None else False, + upcast=upcast, + reorder_ops=reorder_ops, + ) + return out, softmax_lse + + +@benchmark() +def test_mha_opus_d192_v128( + batch_size, nheads, seqlen_q, seqlen_k, causal, check_ref=False +): + torch.manual_seed(0) + torch.cuda.empty_cache() + dtype = dtypes.bf16 + + q = torch.randn(batch_size, seqlen_q, nheads, D_QK, dtype=dtype) + k = torch.randn(batch_size, seqlen_k, nheads, D_QK, dtype=dtype) + v = torch.randn(batch_size, seqlen_k, nheads, D_V, dtype=dtype) + + (out, lse), us = run_perftest( + fmha_fwd_bf16_opus_fwd, + q, + k, + v, + D_QK**-0.5, + causal, + return_lse=True, + num_rotate_args=1, + ) + aiter.logger.info( + "opus d192/v128 perf b=%s h=%s sq=%s sk=%s %.2f us", + batch_size, + nheads, + seqlen_q, + seqlen_k, + us, + ) + + err = None + ref_n = 0 + if check_ref: + rows = _ref_rows(seqlen_q, batch_size, nheads, seqlen_k) + ref_n = int(rows.numel()) + q_s = q[:, rows] + attn_bias = _causal_bias(rows, seqlen_q, seqlen_k, q.device) if causal else None + out_ref, _ = run_torch(q_s, k, v, causal=causal, attn_bias=attn_bias) + out_pt, _ = run_torch( + q_s, + k, + v, + causal=causal, + attn_bias=attn_bias, + upcast=False, + reorder_ops=True, + ) + + got = out[:, rows] + out_tol = max(2 * (out_pt - out_ref).abs().max().item(), 0.01) + err = checkAllclose( + out_ref, + got, + atol=out_tol, + rtol=0.0, + msg=f"opus d192/v128 BSHD b={batch_size} h={nheads} " + f"sq={seqlen_q} sk={seqlen_k} causal={causal} {us:.2f} us", + ) + + lse_ref = opus_ref_lse(q, k, causal) + assert tuple(lse.shape) == ( + batch_size, + nheads, + seqlen_q, + ), f"lse {tuple(lse.shape)}" + opus_check_lse("opus-d192", lse, lse_ref) + dead = torch.isneginf(lse_ref) + if dead.any(): + dead_o = dead.permute(0, 2, 1).unsqueeze(-1).expand_as(out) + assert (out[dead_o] == 0).all(), "fully-masked rows must produce O=0" + + flops = ( + batch_size + * nheads + * (seqlen_q * seqlen_k * D_QK * 2 + seqlen_q * seqlen_k * D_V * 2) + ) + if causal: + flops = flops / 2 + nbytes = ( + batch_size + * nheads + * 2 + * (seqlen_q * D_QK + seqlen_k * D_QK + seqlen_k * D_V + seqlen_q * D_V) + ) + + return { + "gfx": get_gfx(), + "ref_rows": ref_n, + "fwd us": us, + "fwd TFLOPS": flops / us / 1e6, + "fwd TB/s": nbytes / us / 1e6, + "fwd err": err, + } + + +def _summarize(name, rows): + if not rows: + return + df = pd.DataFrame(rows) + keep = [ + c + for c in ( + "batch_size", + "nheads", + "seqlen_q", + "seqlen_k", + "causal", + "gfx", + "ref_rows", + "fwd us", + "fwd TFLOPS", + "fwd TB/s", + "fwd err", + ) + if c in df.columns + ] + aiter.logger.info( + "%s summary (markdown):\n%s", name, df[keep].to_markdown(index=False) + ) + + +def main(): + if get_gfx() not in SUPPORTED_GFX: + aiter.logger.warning( + "mha_opus_d192_v128_logits unsupported on %s; skipping", get_gfx() + ) + return + + parser = argparse.ArgumentParser( + formatter_class=argparse.RawTextHelpFormatter, + description="config input of test", + ) + parser.add_argument( + "-b", + "--batch_size", + type=int, + nargs="*", + default=[1, 16, 64, 128], + help="""Prefill batch size. + e.g.: -b 1 16""", + ) + parser.add_argument( + "-n", + "--nheads", + type=int, + nargs="*", + default=[32, 64], + help="""Q/K/V heads (MHA, nheads_k = nheads). + e.g.: -n 32""", + ) + parser.add_argument( + "-q", + "--seqlen_q", + type=int, + nargs="*", + default=[1024, 2048, 4096, 8192, 16384], + help="""Query length. + e.g.: -q 1024 4096""", + ) + parser.add_argument( + "-k", + "--seqlen_k", + type=int, + nargs="*", + default=[4096, 16384, 65664], + help="""KV length. Skips seqlen_k < seqlen_q. + e.g.: -k 4096 16384""", + ) + parser.add_argument( + "--causal", + action=argparse.BooleanOptionalAction, + default=True, + help="causal mask. Default: True.", + ) + parser.add_argument( + "--ref", + action=argparse.BooleanOptionalAction, + default=False, + help="Compare against torch golden. Default: False. Pass --ref to enable.", + ) + parser.add_argument( + "--perf", + action="store_true", + help="Alias for --no-ref (kernel timing only).", + ) + args = parser.parse_args() + check_ref = bool(args.ref) and not args.perf + + rows = [] + for batch_size, nheads, seqlen_q, seqlen_k in itertools.product( + args.batch_size, args.nheads, args.seqlen_q, args.seqlen_k + ): + if seqlen_k < seqlen_q: + continue + try: + rows.append( + test_mha_opus_d192_v128( + batch_size, + nheads, + seqlen_q, + seqlen_k, + args.causal, + check_ref=check_ref, + ) + ) + except RuntimeError as e: + if "out of memory" not in str(e).lower(): + raise + aiter.logger.warning( + "OOM skip b=%s h=%s sq=%s sk=%s", + batch_size, + nheads, + seqlen_q, + seqlen_k, + ) + torch.cuda.empty_cache() + _summarize("mha_opus_d192_v128_logits", rows) + + +if __name__ == "__main__": + main() diff --git a/op_tests/test_mla_gqa_logits.py b/op_tests/test_mla_gqa_logits.py new file mode 100644 index 0000000000..714db23500 --- /dev/null +++ b/op_tests/test_mla_gqa_logits.py @@ -0,0 +1,1218 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. + +"""MLA decode GQA-head sweep for persistent LEGACY layout. + +Sweeps nhead in {32, 64, 96, 128}. KV is a single latent head (nhead_kv=1), so +GQA ratio == nhead. Page tables are fragmented (randperm); each sequence owns +its pages. Prefill, non-persistent decode, 3BUFFER and DS32_OPUS are out of +scope. KV length is uniform across the batch. + +Two phases (``-p``), matching the round-robin CP test in +op_tests/test_mla_persistent_round_robin.py: + + decode aiter.mla_decode_fwd vs torch_mla_extend + cp round-robin context-parallel: per-rank shard + online-softmax merge + +Examples: + python op_tests/test_mla_gqa_logits.py + python op_tests/test_mla_gqa_logits.py -p decode -n 32 64 -b 16 64 -c 4096 + python op_tests/test_mla_gqa_logits.py -p cp -cpw 4 -n 32 -mtp 1 -c 64 -b 1 + python op_tests/test_mla_gqa_logits.py -d fp8 -kvd fp8 -n 32 64 96 128 + python op_tests/test_mla_gqa_logits.py -p decode --ref +""" + +import argparse +import itertools + +import pandas as pd +import torch + +import aiter +from aiter import dtypes +from aiter.jit.utils.chip_info import get_gfx +from aiter.test_common import benchmark, checkAllclose, run_perftest + +torch.set_default_device("cuda") +torch.set_printoptions(sci_mode=False) + +SUPPORTED_GFX = ["gfx942", "gfx950"] + + +def check_support(dtype, kv_dtype, nhead): + if dtype != kv_dtype: + return False + return not (dtype == dtypes.bf16 and nhead == 32 and get_gfx() == "gfx942") + + +def cal_diff( + x: torch.Tensor, y: torch.Tensor, name: str, use_fp8: bool = False +) -> None: + x, y = x.double(), y.double() + # RMSE = ((x - y) * (x - y)).mean().sqrt().item() + cos_diff = 1 - 2 * (x * y).sum().item() / max((x * x + y * y).sum().item(), 1e-12) + if use_fp8: + assert cos_diff < 3e-2 + else: + assert cos_diff < 1e-5 + + +def ref_masked_attention( + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + scale: float, + dtype, + is_causal=True, + is_fp8_q=False, + is_fp8_kvc=False, + q_scale=None, + kv_scale=None, + causal_diagonal=None, + attn_mask=None, +): + if is_fp8_q and q_scale is not None: + scale *= q_scale + if is_fp8_kvc and kv_scale is not None: + scale *= kv_scale + attn_weights = torch.einsum("qhd,khd->hqk", query.float(), key.float()) * scale + + if attn_mask is not None: + attn_bias = torch.zeros_like(attn_weights) + attn_bias.masked_fill_(attn_mask[None].logical_not(), float("-inf")) + attn_weights = attn_weights + attn_bias + elif is_causal: + s_q = query.shape[0] + s_k = key.shape[0] + diagonal = causal_diagonal if causal_diagonal is not None else s_k - s_q + attn_bias = torch.zeros(s_q, s_k, dtype=query.dtype) + temp_mask = torch.ones(s_q, s_k, dtype=torch.bool).tril(diagonal=diagonal) + attn_bias.masked_fill_(temp_mask.logical_not(), float("-inf")) + attn_bias.to(query.dtype) + attn_weights += attn_bias + + lse = attn_weights.logsumexp(dim=-1) + m = attn_weights.max(-1).values + attn_weights_exp = torch.exp(attn_weights - m.unsqueeze(-1)) + l = attn_weights_exp.sum(-1) + if is_fp8_q: + attn_weights_fp8 = attn_weights_exp.to(dtypes.fp8) + attn_weights_exp = attn_weights_fp8.to(torch.float) + + out = torch.einsum("hqk,khd->qhd", attn_weights_exp.float(), value.float()) + out = out / l.transpose(0, 1).unsqueeze(-1) + if is_fp8_kvc and kv_scale is not None: + out *= kv_scale + + if attn_mask is not None: + invalid = attn_mask.any(dim=-1).logical_not() + if bool(invalid.any()): + out = out.clone() + out[invalid] = 0.0 + lse = lse.clone() + lse[:, invalid] = float("-inf") + + return out.to(dtype), lse + + +def torch_mla_extend( + q, # [total_q, nheads, headdim_q] + kvc_cache, # [num_page, page_size, nhead_kv, qk_head_dim] + qo_indptr, + kv_indptr, + kv_indices, + kv_last_page_lens, + sm_scale, + kv_lora_rank, + qk_rope_head_dim, + dtype, + is_causal=True, + q_scale=None, + kv_scale=None, +): + _num_page, page_size, _nhead_kv, _ = kvc_cache.shape + is_fp8_q = q.dtype == dtypes.fp8 + is_fp8_kvc = kvc_cache.dtype == dtypes.fp8 + + if is_fp8_q: + q = q.to(torch.float) + + if is_fp8_kvc: + kvc_cache = kvc_cache.to(torch.float) + + qs = torch.tensor_split(q, qo_indptr.tolist()[1:]) + kvc = torch.index_select(kvc_cache, 0, kv_indices) + kvs = torch.tensor_split(kvc, kv_indptr.tolist()[1:]) + bs = qo_indptr.shape[0] - 1 + + os = [] + lses = [] + for i in range(bs): + cur_num_page = kvs[i].shape[0] + real_kv_seq_len = (cur_num_page - 1) * page_size + kv_last_page_lens.tolist()[i] + kvc = kvs[i].flatten(0, 1)[:real_kv_seq_len,] + q = qs[i] + k = kvc + v, _ = torch.split(kvc, [kv_lora_rank, qk_rope_head_dim], dim=-1) + o, lse = ref_masked_attention( + q, + k, + v, + sm_scale, + dtype, + is_causal=is_causal, + is_fp8_q=is_fp8_q, + is_fp8_kvc=is_fp8_kvc, + q_scale=q_scale, + kv_scale=kv_scale, + ) + os.append(o) + lses.append(lse) + o = torch.concat(os) + # Each lse is (nheads, seq_q_i); concatenate query positions along dim=1, then (total_q, nheads). + lse = torch.concat(lses, dim=1).transpose(0, 1) + return o, lse + + +def torch_mla_extend_round_robin( + q, + kvc_cache, + qo_indptr, + kv_indptr_r, + kv_indices_r, + g_kv_indptr, + page_size, + sm_scale, + kv_lora_rank, + qk_rope_head_dim, + dtype, + cp_world_size, + cp_rank, + q_scale=None, + kv_scale=None, +): + """Round-robin CP reference for one rank. page_size must be 1.""" + dev = kvc_cache.device + is_fp8_q = q.dtype == dtypes.fp8 + is_fp8_kvc = kvc_cache.dtype == dtypes.fp8 + if is_fp8_q: + q = q.to(torch.float) + if is_fp8_kvc: + kvc_cache = kvc_cache.to(torch.float) + + qs = torch.tensor_split(q, qo_indptr.tolist()[1:]) + kvc = torch.index_select(kvc_cache, 0, kv_indices_r) + indptr_r = kv_indptr_r.tolist() + g_indptr = g_kv_indptr.tolist() + bs = qo_indptr.shape[0] - 1 + + os = [] + lses = [] + for i in range(bs): + q_i = qs[i] + s_q, nheads, _ = q_i.shape + p0, p1 = int(indptr_r[i]), int(indptr_r[i + 1]) + s_k = (p1 - p0) * page_size + + if s_k == 0: + os.append(torch.zeros(s_q, nheads, kv_lora_rank, dtype=dtype, device=dev)) + lses.append(torch.full((nheads, s_q), float("-inf"), device=dev)) + continue + + local_kv = kvc[p0:p1].flatten(0, 1)[:s_k] + k = local_kv + v, _ = torch.split(local_kv, [kv_lora_rank, qk_rope_head_dim], dim=-1) + + local_global_pos = torch.arange(s_k, device=dev) * cp_world_size + cp_rank + global_len = (int(g_indptr[i + 1]) - int(g_indptr[i])) * page_size + q_global = (global_len - s_q) + torch.arange(s_q, device=dev) + attn_mask = local_global_pos[None, :] <= q_global[:, None] + + o, lse = ref_masked_attention( + q_i, + k, + v, + sm_scale, + dtype, + is_fp8_q=is_fp8_q, + is_fp8_kvc=is_fp8_kvc, + q_scale=q_scale, + kv_scale=kv_scale, + attn_mask=attn_mask, + ) + os.append(o) + lses.append(lse) + + o = torch.concat(os) + lse = torch.concat(lses, dim=1).transpose(0, 1) + return o, lse + + +def merge_cp_ranks(cp_outs, cp_lses, out_dtype=torch.bfloat16): + LS = torch.stack([lse.float() for lse in cp_lses], 0) + glse = torch.logsumexp(LS, 0) + w = torch.exp(LS - glse).nan_to_num_(0.0) + out = sum(w[r][..., None] * cp_outs[r].float() for r in range(len(cp_outs))) + return out.to(out_dtype), glse + + +def aiter_cp_rank_decode( + q, + kv_buffer, + qo_indptr, + kv_indptr_r, + kv_indices_r, + g_kv_indptr, + kv_last_page_lens, + batch_size, + max_seqlen_q, + nhead, + nhead_kv, + kv_lora_rank, + qk_head_dim, + v_head_dim, + sm_scale, + dtype, + kvtype, + max_split_per_batch, + is_causal, + cp_world_size, + cp_rank, + q_scale=None, + kv_scale=None, +): + dev = q.device + total_q = q.shape[0] + o = torch.zeros(total_q, nhead, v_head_dim, dtype=torch.bfloat16, device=dev) + + info = aiter.get_mla_metadata_info_v1( + batch_size, + max_seqlen_q, + nhead, + dtype, + kvtype, + is_sparse=False, + fast_mode=True, + num_kv_splits=max_split_per_batch, + intra_batch_mode=False, + ) + + def _alloc(sz, ty): + return torch.empty(sz, dtype=ty, device=dev) + + work_meta_data = _alloc(*info[0]) + work_indptr = _alloc(*info[1]) + work_info_set = _alloc(*info[2]) + reduce_indptr = _alloc(*info[3]) + reduce_final_map = _alloc(*info[4]) + reduce_partial_map = _alloc(*info[5]) + + aiter.get_mla_metadata_v1( + qo_indptr, + kv_indptr_r, + kv_last_page_lens, + nhead // nhead_kv, + nhead_kv, + False, + work_meta_data, + work_info_set, + work_indptr, + reduce_indptr, + reduce_final_map, + reduce_partial_map, + page_size=1, + kv_granularity=16, + max_seqlen_qo=max_seqlen_q, + uni_seqlen_qo=max_seqlen_q, + fast_mode=True, + max_split_per_batch=max_split_per_batch, + intra_batch_mode=False, + dtype_q_nope=dtype, + dtype_kv_nope=kvtype, + is_cp_round_robin=True, + ) + + (_, final_lse), us = run_perftest( + aiter.mla.mla_decode_fwd, + q, + kv_buffer, + o, + qo_indptr, + kv_indptr_r, + kv_indices_r, + kv_last_page_lens, + max_seqlen_q=max_seqlen_q, + page_size=1, + nhead_kv=nhead_kv, + sm_scale=sm_scale, + num_kv_splits=max_split_per_batch, + q_scale=q_scale, + kv_scale=kv_scale, + work_meta_data=work_meta_data, + work_indptr=work_indptr, + work_info_set=work_info_set, + reduce_indptr=reduce_indptr, + reduce_final_map=reduce_final_map, + reduce_partial_map=reduce_partial_map, + intra_batch_mode=False, + return_lse=True, + g_kv_indptr=g_kv_indptr, + cp_world_size=cp_world_size, + cp_rank=cp_rank, + ) + lse = final_lse.float() if final_lse is not None else None + return o.float(), lse, us + + +@benchmark() +def test_mla_gqa_decode( + ctx_lens, + batch_size, + nhead, + kv_lora_rank, + qk_nope_head_dim, + qk_rope_head_dim, + v_head_dim, + dtype, + kvtype, + page_size, + decode_qlen, + max_split_per_batch, + return_lse, + causal, + check_ref=False, +): + """Persistent LEGACY paged MLA decode for one (B, N, nhead) shape.""" + ret = {} + out_dtype = torch.bfloat16 + + qo_indptr = torch.zeros(batch_size + 1, dtype=torch.int) + kv_indptr = torch.zeros(batch_size + 1, dtype=torch.int) + seq_lens_qo = torch.empty(batch_size, dtype=torch.int) + seq_lens_kv = torch.empty(batch_size, dtype=torch.int) + kv_block_nums = torch.empty(batch_size, dtype=torch.int) + kv_last_page_lens = torch.ones(batch_size, dtype=torch.int) + seq_lens_kv.fill_(ctx_lens) + kv_block_nums.fill_((ctx_lens + page_size - 1) // page_size) + if ctx_lens % page_size == 0: + kv_last_page_lens.fill_(page_size) + else: + kv_last_page_lens.fill_(ctx_lens % page_size) + + kv_indptr[1 : batch_size + 1] = torch.cumsum(kv_block_nums, dim=0) + num_page = kv_indptr[-1].item() + kv_indices = torch.randperm(num_page, dtype=torch.int) + total_kv = seq_lens_kv.sum().item() + + kv_buffer = torch.randn( + (num_page, page_size, 1, kv_lora_rank + qk_rope_head_dim), + dtype=torch.bfloat16, + ) + + qk_head_dim = kv_lora_rank + qk_rope_head_dim + sm_scale = 1.0 / (qk_head_dim**0.5) + torch.cuda.empty_cache() + nhead_kv = 1 + + seq_lens_qo.fill_(decode_qlen) + max_seqlen_qo = seq_lens_qo.max().item() + qo_indptr[1 : batch_size + 1] = torch.cumsum(seq_lens_qo, dim=0) + total_q = qo_indptr[-1].item() + q = torch.randn((total_q, nhead, qk_head_dim), dtype=torch.bfloat16) + + out_ref = lse_ref = None + if check_ref: + out_ref, lse_ref = torch_mla_extend( + q, + kv_buffer, + qo_indptr, + kv_indptr, + kv_indices, + kv_last_page_lens, + sm_scale, + kv_lora_rank, + qk_rope_head_dim, + is_causal=causal, + dtype=out_dtype, + ) + + if nhead >= 128: + gpu = torch.cuda.current_device() + device_properties = torch.cuda.get_device_properties(gpu) + cu_num = device_properties.multi_processor_count + max_split_per_batch = min( + (cu_num + batch_size - 1) // batch_size, max_split_per_batch + ) + + ( + (work_meta_data_size, work_meta_data_type), + (work_indptr_size, work_indptr_type), + (work_info_set_size, work_info_set_type), + (reduce_indptr_size, reduce_indptr_type), + (reduce_final_map_size, reduce_final_map_type), + (reduce_partial_map_size, reduce_partial_map_type), + ) = aiter.get_mla_metadata_info_v1( + batch_size, + max_seqlen_qo, + nhead, + dtype, + kvtype, + is_sparse=False, + fast_mode=True, + num_kv_splits=max_split_per_batch, + intra_batch_mode=False, + ) + + work_meta_data = torch.empty( + work_meta_data_size, dtype=work_meta_data_type, device="cuda" + ) + work_indptr = torch.empty(work_indptr_size, dtype=work_indptr_type, device="cuda") + work_info_set = torch.empty( + work_info_set_size, + dtype=work_info_set_type, + device="cuda", + ) + reduce_indptr = torch.empty( + reduce_indptr_size, dtype=reduce_indptr_type, device="cuda" + ) + reduce_final_map = torch.empty( + reduce_final_map_size, dtype=reduce_final_map_type, device="cuda" + ) + reduce_partial_map = torch.empty( + reduce_partial_map_size, dtype=reduce_partial_map_type, device="cuda" + ) + + aiter.get_mla_metadata_v1( + qo_indptr, + kv_indptr, + kv_last_page_lens, + nhead // nhead_kv, + nhead_kv, + causal, + work_meta_data, + work_info_set, + work_indptr, + reduce_indptr, + reduce_final_map, + reduce_partial_map, + page_size=page_size, + kv_granularity=max( + page_size, + ( + 32 + if (nhead == 64 and dtype == dtypes.fp8 and kvtype == dtypes.fp8) + else 16 + ), + ), + max_seqlen_qo=int(max_seqlen_qo), + uni_seqlen_qo=decode_qlen, + fast_mode=True, + max_split_per_batch=max_split_per_batch, + intra_batch_mode=False, + dtype_q_nope=dtype, + dtype_kv_nope=kvtype, + ) + + def test_absorb_decode_bf16(): + out_asm = torch.empty((total_q, nhead, v_head_dim), dtype=out_dtype).fill_(-1) + (_attn_logits, attn_lse), us_asm_decode = run_perftest( + aiter.mla.mla_decode_fwd, + q, + kv_buffer.view(num_page, page_size, nhead_kv, qk_head_dim), + out_asm, + qo_indptr, + kv_indptr, + kv_indices, + kv_last_page_lens, + max_seqlen_qo, + page_size, + nhead_kv, + sm_scale, + num_kv_splits=max_split_per_batch, + work_meta_data=work_meta_data, + work_indptr=work_indptr, + work_info_set=work_info_set, + reduce_indptr=reduce_indptr, + reduce_final_map=reduce_final_map, + reduce_partial_map=reduce_partial_map, + intra_batch_mode=False, + return_lse=return_lse, + causal=causal, + ) + + err = None + if check_ref: + err = checkAllclose( + out_ref, + out_asm, + msg=f"mla_decode-absorb [golden vs aiter_asm]: {us_asm_decode:>8.2f} us......", + ) + if return_lse: + checkAllclose( + lse_ref, + attn_lse.reshape(total_q, nhead), + msg=f"mla_decode-absorb [lse_ref vs attn_lse]: {us_asm_decode:>8.2f} us......", + ) + else: + aiter.logger.info("mla_decode-absorb [no-ref] %8.2f us", us_asm_decode) + return err, us_asm_decode + + def test_absorb_decode_fp8(): + out_asm = torch.empty((total_q, nhead, v_head_dim), dtype=out_dtype).fill_(-1) + q_fp8 = q.to(dtypes.fp8) + q_scale = torch.ones([1], dtype=torch.float, device="cuda") + kv_buffer_fp8 = kv_buffer.to(dtypes.fp8) + kv_scale = torch.ones([1], dtype=torch.float, device="cuda") + + out_ref_fp8 = None + if check_ref: + out_ref_fp8, _lse_ref_fp8 = torch_mla_extend( + q_fp8, + kv_buffer_fp8, + qo_indptr, + kv_indptr, + kv_indices, + kv_last_page_lens, + sm_scale, + kv_lora_rank, + qk_rope_head_dim, + dtype=out_dtype, + is_causal=causal, + q_scale=None, + kv_scale=kv_scale, + ) + + (_attn_logits, attn_lse), us_asm_decode = run_perftest( + aiter.mla.mla_decode_fwd, + q_fp8, + kv_buffer_fp8.view(num_page, page_size, nhead_kv, qk_head_dim), + out_asm, + qo_indptr, + kv_indptr, + kv_indices, + kv_last_page_lens, + max_seqlen_qo, + page_size, + nhead_kv, + sm_scale, + num_kv_splits=max_split_per_batch, + q_scale=q_scale, + kv_scale=kv_scale, + work_meta_data=work_meta_data, + work_indptr=work_indptr, + work_info_set=work_info_set, + reduce_indptr=reduce_indptr, + reduce_final_map=reduce_final_map, + reduce_partial_map=reduce_partial_map, + intra_batch_mode=False, + return_lse=return_lse, + causal=causal, + ) + + err = None + if check_ref: + err = checkAllclose( + out_ref, + out_asm, + msg=f"mla_decode-absorb_fp8 [golden vs aiter_asm]: {us_asm_decode:>8.2f} us......", + ) + if return_lse: + err = checkAllclose( + lse_ref, + attn_lse.reshape(total_q, nhead), + msg=f"mla_decode-absorb_fp8 [lse_ref vs attn_lse]: {us_asm_decode:>8.2f} us......", + ) + err = checkAllclose( + out_ref_fp8, + out_asm, + msg=f"mla_decode-absorb_fp8 [golden fp8 vs aiter_asm]: {us_asm_decode:>8.2f} us......", + ) + cal_diff(out_ref, out_asm, "out", True) + else: + aiter.logger.info( + "mla_decode-absorb_fp8 [no-ref] %8.2f us", us_asm_decode + ) + return err, us_asm_decode + + err = None + us_asm_decode = 1e12 + if dtype == torch.bfloat16: + err, us_asm_decode = test_absorb_decode_bf16() + elif kvtype == dtypes.fp8: + err, us_asm_decode = test_absorb_decode_fp8() + + ret["decode:err"] = err + ret["decode:asm_576"] = us_asm_decode + flops = decode_qlen * total_kv * nhead * (qk_head_dim + v_head_dim) * 2 + nbytes = ( + total_kv * nhead_kv * qk_head_dim * (torch.finfo(kvtype).bits // 8) + + total_q * nhead * qk_head_dim * (torch.finfo(dtype).bits // 8) + + total_q * nhead * v_head_dim * (torch.finfo(out_dtype).bits // 8) + ) + ret["decode:flops"] = flops + ret["decode:bytes"] = nbytes + ret["decode:TFLOPS"] = flops / us_asm_decode / 1e6 + ret["decode:TB/s"] = nbytes / us_asm_decode / 1e6 + return ret + + +@benchmark() +def test_mla_cp( + ctx_lens, + batch_size, + nhead, + kv_lora_rank, + qk_rope_head_dim, + v_head_dim, + dtype, + kvtype, + decode_qlen, + cp_world_size, + max_split_per_batch, + return_lse=False, + check_ref=False, +): + """Fixed-length round-robin CP decode. page_size is forced to 1.""" + ret = {} + W = cp_world_size + dev = "cuda" + out_dtype = torch.bfloat16 + nhead_kv = 1 + qlen = decode_qlen + qk_head_dim = kv_lora_rank + qk_rope_head_dim + sm_scale = 1.0 / (qk_head_dim**0.5) + is_causal = qlen > 1 + page_size = 1 + + kv_block_nums = torch.empty(batch_size, dtype=torch.int) + seq_lens_kv = torch.empty(batch_size, dtype=torch.int) + kv_last_page_lens = torch.ones(batch_size, dtype=torch.int) + seq_lens_kv.fill_(ctx_lens) + kv_block_nums.fill_((ctx_lens + page_size - 1) // page_size) + kv_last_page_lens.fill_( + page_size if ctx_lens % page_size == 0 else ctx_lens % page_size + ) + + assert ( + int(seq_lens_kv.min().item()) >= W + ), f"every request kv_len must be >= cp_world_size({W})" + + kv_indptr = torch.zeros(batch_size + 1, dtype=torch.int) + kv_indptr[1:] = torch.cumsum(kv_block_nums, dim=0) + num_page = int(kv_indptr[-1].item()) + kv_indices = torch.randperm(num_page, dtype=torch.int) + + seq_lens_qo = torch.full((batch_size,), qlen, dtype=torch.int) + qo_indptr = torch.zeros(batch_size + 1, dtype=torch.int) + qo_indptr[1:] = torch.cumsum(seq_lens_qo, dim=0) + total_q = int(qo_indptr[-1].item()) + + kv_buffer = torch.randn( + (num_page, page_size, 1, qk_head_dim), dtype=torch.bfloat16 + ).to(kvtype) + q = torch.randn((total_q, nhead, qk_head_dim), dtype=torch.bfloat16).to(dtype) + + q_scale = ( + torch.ones([1], dtype=torch.float, device=dev) if dtype == dtypes.fp8 else None + ) + kv_scale = ( + torch.ones([1], dtype=torch.float, device=dev) if kvtype == dtypes.fp8 else None + ) + + out_ref = lse_ref = None + if check_ref: + out_ref, lse_ref = torch_mla_extend( + q, + kv_buffer, + qo_indptr, + kv_indptr, + kv_indices, + kv_last_page_lens, + sm_scale, + kv_lora_rank, + qk_rope_head_dim, + dtype=out_dtype, + is_causal=True, + q_scale=q_scale, + kv_scale=kv_scale, + ) + + g_kv_indptr = kv_indptr.to(dev).to(torch.int32) + kv_indices_dev = kv_indices.to(dev).to(torch.int32) + qo_indptr_dev = qo_indptr.to(dev).to(torch.int32) + kv_buffer_dev = kv_buffer.to(dev) + + rank_kv_indptr_r, rank_kv_indices_r, rank_kv_last_r = [], [], [] + for r in range(W): + idx_r_list, local_lens = [], [] + for b in range(batch_size): + real_kv = int(seq_lens_kv[b].item()) + start = int(kv_indptr[b].item()) + pos = torch.arange(real_kv, device=dev) + pos = pos[pos % W == r] + idx_r_list.append(kv_indices_dev[start + pos]) + local_lens.append(int(pos.numel())) + kv_indices_r = ( + torch.cat(idx_r_list).to(torch.int32) + if sum(local_lens) > 0 + else torch.zeros(1, dtype=torch.int32, device=dev) + ) + kv_indptr_r = torch.zeros(batch_size + 1, dtype=torch.int32, device=dev) + kv_indptr_r[1:] = torch.cumsum( + torch.tensor(local_lens, dtype=torch.int32, device=dev), dim=0 + ) + rank_kv_indptr_r.append(kv_indptr_r) + rank_kv_indices_r.append(kv_indices_r) + rank_kv_last_r.append(torch.ones(batch_size, dtype=torch.int32, device=dev)) + + cp_outs, cp_lses = [], [] + aiter_outs, aiter_lses, rank_us = [], [], [] + for r in range(W): + kv_indptr_r = rank_kv_indptr_r[r] + kv_indices_r = rank_kv_indices_r[r] + kv_last_page_lens_r = rank_kv_last_r[r] + o_r = l_r = None + if check_ref: + o_r, l_r = torch_mla_extend_round_robin( + q, + kv_buffer_dev, + qo_indptr_dev, + kv_indptr_r, + kv_indices_r, + g_kv_indptr, + page_size, + sm_scale, + kv_lora_rank, + qk_rope_head_dim, + dtype=out_dtype, + cp_world_size=W, + cp_rank=r, + q_scale=q_scale, + kv_scale=kv_scale, + ) + cp_outs.append(o_r) + cp_lses.append(l_r) + + o_a, l_a, us = aiter_cp_rank_decode( + q, + kv_buffer_dev, + qo_indptr_dev, + kv_indptr_r, + kv_indices_r, + g_kv_indptr, + kv_last_page_lens_r, + batch_size, + qlen, + nhead, + nhead_kv, + kv_lora_rank, + qk_head_dim, + v_head_dim, + sm_scale, + dtype, + kvtype, + max_split_per_batch, + is_causal, + W, + r, + q_scale, + kv_scale, + ) + rank_us.append(us) + + if check_ref: + local_lens_r = (kv_indptr_r[1:] - kv_indptr_r[:-1]).tolist() + checkAllclose( + o_r.float(), + o_a, + msg=f"mla_cp_round_robin W={W} qlen={qlen} rank{r} " + f"local_len={local_lens_r} has_NaN={bool(torch.isnan(o_a).any())} " + f"[cp_ref vs aiter]:......", + ) + if return_lse: + checkAllclose( + l_r.float(), + l_a, + msg=f"mla_cp_round_robin W={W} qlen={qlen} rank{r} " + f"[cp_ref vs aiter lse]:......", + ) + aiter_outs.append(o_a) + aiter_lses.append(l_a) + + err_ref = err = None + if check_ref: + cp_merged_out, cp_merged_lse = merge_cp_ranks(cp_outs, cp_lses, out_dtype) + err_ref = checkAllclose( + out_ref, + cp_merged_out, + msg=f"mla_cp_round_robin W={W} qlen={qlen} [golden vs cp_ref_merge out]:......", + ) + checkAllclose( + lse_ref, + cp_merged_lse, + msg=f"mla_cp_round_robin W={W} qlen={qlen} [golden vs cp_ref_merge lse]:......", + ) + + aiter_merged_out, aiter_merged_lse = merge_cp_ranks( + aiter_outs, aiter_lses, out_dtype + ) + err = checkAllclose( + out_ref, + aiter_merged_out, + msg=f"mla_cp_round_robin W={W} qlen={qlen} [golden vs aiter_merge out]:......", + ) + if return_lse: + checkAllclose( + lse_ref, + aiter_merged_lse, + msg=f"mla_cp_round_robin W={W} qlen={qlen} [golden vs aiter_merge lse]:......", + ) + else: + aiter.logger.info( + "mla_cp_round_robin W=%s qlen=%s [no-ref] mean rank_us=%.2f", + W, + qlen, + sum(rank_us) / max(len(rank_us), 1), + ) + ret["cp:err_ref"] = err_ref + ret["cp:err_aiter"] = err + ret["cp:world_size"] = W + ret["cp:rank_us"] = sum(rank_us) / max(len(rank_us), 1) + return ret + + +def _summarize(name, rows): + if not rows: + return + df = pd.DataFrame(rows) + keep = [ + c + for c in ( + "nhead", + "decode_qlen", + "batch_size", + "ctx_lens", + "dtype", + "kvtype", + "gfx", + "decode:err", + "decode:asm_576", + "decode:TFLOPS", + "decode:TB/s", + "cp:world_size", + "cp:err_ref", + "cp:err_aiter", + "cp:rank_us", + ) + if c in df.columns + ] + aiter.logger.info( + "%s summary (markdown):\n%s", name, df[keep].to_markdown(index=False) + ) + + +def main(): + if get_gfx() not in SUPPORTED_GFX: + aiter.logger.warning("mla_gqa_logits unsupported on %s; skipping", get_gfx()) + return + + parser = argparse.ArgumentParser( + formatter_class=argparse.RawTextHelpFormatter, + description="config input of test", + ) + parser.add_argument( + "-n", + "--nhead", + type=int, + nargs="*", + default=[32, 64, 96, 128], + help="""Q heads (GQA ratio vs nhead_kv=1). Persistent LEGACY only. + e.g.: -n 32 96""", + ) + parser.add_argument( + "-c", + "--ctxLen", + type=int, + nargs="*", + default=[4096, 16384, 65536, 131072], + help="""KV length N. + e.g.: -c 8192 131072""", + ) + parser.add_argument( + "-b", + "--batchSize", + type=int, + nargs="*", + default=[16, 32, 64, 128], + help="""Decode batch. Decode M = batch * decode_qlen. + e.g.: -b 128 256""", + ) + parser.add_argument( + "-mtp", + "--decode_qlen", + type=int, + nargs="*", + default=[1, 2, 4, 8], + help="""Decode speculative rows per sequence. + e.g.: -mtp 1 2""", + ) + parser.add_argument( + "-d", + "--dtype", + type=dtypes.str2Dtype, + choices=[dtypes.d_dtypes["bf16"], dtypes.d_dtypes["fp8"]], + nargs="*", + default=[dtypes.d_dtypes["bf16"], dtypes.d_dtypes["fp8"]], + metavar="{bf16, fp8}", + help="""Data type of Q. + e.g.: -d bf16 fp8""", + ) + parser.add_argument( + "-kvd", + "--kv_dtype", + type=dtypes.str2Dtype, + choices=[dtypes.d_dtypes["bf16"], dtypes.d_dtypes["fp8"]], + nargs="*", + default=[dtypes.d_dtypes["bf16"], dtypes.d_dtypes["fp8"]], + metavar="{bf16, fp8}", + help="""Data type of KV. + e.g.: -kvd bf16 fp8""", + ) + parser.add_argument( + "-blk", + "--block_size", + type=int, + default=1, + help="""Paged KV page size (LEGACY layout). + e.g.: -blk 1""", + ) + parser.add_argument( + "-ms", + "--max_split_per_batch", + type=int, + default=32, + help="""kv seqlens max split num per batch. + e.g.: -ms 32""", + ) + parser.add_argument( + "-k", + "--kv_lora_rank", + type=int, + default=512, + help="kv lora rank.", + ) + parser.add_argument( + "-qn", + "--qk_nope_head_dim", + type=int, + default=512, + help="qk nope head dim.", + ) + parser.add_argument( + "-qr", + "--qk_rope_head_dim", + type=int, + default=64, + help="qk rope head dim.", + ) + parser.add_argument( + "-vh", + "--v_head_dim", + type=int, + default=512, + help="v head dim.", + ) + parser.add_argument( + "-lse", + "--return_lse", + action="store_true", + help="return lse.", + ) + parser.add_argument( + "--causal", + action=argparse.BooleanOptionalAction, + default=True, + help="causal mask across decode_qlen tokens. Default: True.", + ) + parser.add_argument( + "-p", + "--phase", + type=str, + nargs="*", + choices=["decode", "cp"], + default=["decode", "cp"], + help="""Which phases to run. + e.g.: -p decode + e.g.: -p cp""", + ) + parser.add_argument( + "-cpw", + "--cp_world_size", + type=int, + nargs="*", + default=[2, 3, 4, 7, 8], + help="""CP ranks for the round-robin phase. Skipped when ctx < W. + e.g.: -cpw 4""", + ) + parser.add_argument( + "--ref", + action=argparse.BooleanOptionalAction, + default=False, + help="Compare against torch golden. Default: False. Pass --ref to enable.", + ) + args = parser.parse_args() + + def _as_list(x): + if isinstance(x, (list, tuple)): + return list(x) + return [x] + + dtypes_q = _as_list(args.dtype) + dtypes_kv = _as_list(args.kv_dtype) + + gfx = get_gfx() + if "decode" in args.phase: + rows = [] + for nhead, decode_qlen, dtype, kvtype, ctx_len, batch_size in itertools.product( + args.nhead, + args.decode_qlen, + dtypes_q, + dtypes_kv, + args.ctxLen, + args.batchSize, + ): + if not check_support(dtype, kvtype, nhead): + aiter.logger.warning( + "skip unsupported combo nhead=%s dtype=%s kvtype=%s", + nhead, + dtype, + kvtype, + ) + continue + try: + ret = test_mla_gqa_decode( + ctx_len, + batch_size, + nhead, + args.kv_lora_rank, + args.qk_nope_head_dim, + args.qk_rope_head_dim, + args.v_head_dim, + dtype, + kvtype, + args.block_size, + decode_qlen=decode_qlen, + max_split_per_batch=args.max_split_per_batch, + return_lse=args.return_lse, + causal=args.causal, + check_ref=args.ref, + ) + except (RuntimeError, AttributeError) as e: + msg = str(e).lower() + if not any( + s in msg + for s in ( + "out of memory", + "hk_mla", + "heuristic", + "cannot get", + ) + ): + raise + aiter.logger.warning( + "skip decode nhead=%s mtp=%s dtype=%s ctx=%s b=%s (%s)", + nhead, + decode_qlen, + dtype, + ctx_len, + batch_size, + e, + ) + torch.cuda.empty_cache() + continue + ret["gfx"] = gfx + rows.append(ret) + torch.cuda.empty_cache() + _summarize("mla_gqa_logits decode", rows) + + if "cp" in args.phase: + if gfx != "gfx950": + aiter.logger.warning("mla_gqa_logits cp unsupported on %s; skipping", gfx) + else: + rows = [] + for ( + nhead, + decode_qlen, + dtype, + kvtype, + ctx_len, + batch_size, + cp_world_size, + ) in itertools.product( + args.nhead, + args.decode_qlen, + dtypes_q, + dtypes_kv, + args.ctxLen, + args.batchSize, + args.cp_world_size, + ): + if not check_support(dtype, kvtype, nhead): + continue + if ctx_len < cp_world_size: + continue + if dtype == dtypes.fp8 or kvtype == dtypes.fp8: + aiter.logger.warning( + "skip cp fp8: no cprr heuristic kernel (gqa=%s qseqlen=%s)", + nhead, + decode_qlen, + ) + continue + try: + ret = test_mla_cp( + ctx_len, + batch_size, + nhead, + args.kv_lora_rank, + args.qk_rope_head_dim, + args.v_head_dim, + dtype, + kvtype, + decode_qlen=decode_qlen, + cp_world_size=cp_world_size, + max_split_per_batch=args.max_split_per_batch, + return_lse=args.return_lse, + check_ref=args.ref, + ) + except (RuntimeError, AttributeError) as e: + msg = str(e).lower() + if not any( + s in msg + for s in ( + "out of memory", + "hk_mla", + "heuristic", + "cannot get", + ) + ): + raise + aiter.logger.warning( + "skip cp nhead=%s mtp=%s dtype=%s ctx=%s b=%s W=%s (%s)", + nhead, + decode_qlen, + dtype, + ctx_len, + batch_size, + cp_world_size, + e, + ) + torch.cuda.empty_cache() + continue + ret["gfx"] = gfx + rows.append(ret) + torch.cuda.empty_cache() + _summarize("mla_gqa_logits cp", rows) + + +if __name__ == "__main__": + main()