From d01812d89e21a7972e3cad33898830181f05be06 Mon Sep 17 00:00:00 2001 From: amd-danli103 Date: Tue, 18 Aug 2026 12:16:14 +0800 Subject: [PATCH] [AMD] Optimize KIMI-K3 with Triton MLA decode kernel by tuning the stage-1 geometry for gfx950 (#34580) Co-authored-by: Thomas Wang --- .../kernels/ops/attention/decode_attention.py | 217 ++++++++++++++- python/sglang/srt/environ.py | 3 + .../test_mla_decode_forced_splits.py | 259 ++++++++++++++++++ .../attention/test_mla_decode_geometry.py | 224 +++++++++++++++ 4 files changed, 692 insertions(+), 11 deletions(-) create mode 100644 test/registered/unit/layers/attention/test_mla_decode_forced_splits.py create mode 100644 test/registered/unit/layers/attention/test_mla_decode_geometry.py diff --git a/python/sglang/kernels/ops/attention/decode_attention.py b/python/sglang/kernels/ops/attention/decode_attention.py index 6ab7cd272..696990895 100644 --- a/python/sglang/kernels/ops/attention/decode_attention.py +++ b/python/sglang/kernels/ops/attention/decode_attention.py @@ -21,12 +21,14 @@ It supports page size = 1. # https://github.com/ModelTC/lightllm/blob/96353e868a840db4d103138caf15ed9dbea8c186/lightllm/models/deepseek2/triton_kernel/gqa_flash_decoding_stage2.py import logging +from typing import NamedTuple, Optional, Tuple import triton import triton.language as tl from sglang.kernels.ops.attention.score_mod import unpack_aux_tensors -from sglang.srt.utils import is_hip +from sglang.srt.environ import envs +from sglang.srt.utils import get_device_core_count, is_gfx95_supported, is_hip _is_hip = is_hip() @@ -35,6 +37,160 @@ logger = logging.getLogger(__name__) _MIN_BLOCK_KV = 32 +# heads per stage-1 tile, shared so the budget's head_tiles cannot drift from the launch +_GROUPED_BLOCK_H = 16 + + +# gfx950 wants 32 where the HIP path otherwise takes 16. That is the model it was picked +# against, not something a sweep isolated: at 16 the first dot is a single 16x16 MFMA +# tile, so the warps only have K=576 to split along and pay a cross-warp reduction every +# KV step, where 32 gives two of them an N tile each. 64 was timed at the batches the +# 4-warp bucket covers and never came out ahead: 3-5% behind at batch 1-3, noise at 4-5. +_MLA_BLOCK_N = 32 + + +class _MlaBucket(NamedTuple): + """Stage-1 geometry for a batch range. ``batch_max=None`` is the catch-all.""" + + num_warps: int + num_stages: int + max_splits: int + batch_max: Optional[int] = None + + +# gfx950 MLA decode, from a split-count sweep at every captured batch size, +# head_tiles == 1, 68k context (K3 at tp 8). max_splits is where more splits stopped +# paying at small batch, and dividing by batch * head_tiles keeps a smaller tp sane, +# though tuned at tp 8. +_MLA_BUCKETS = ( + _MlaBucket(num_warps=4, num_stages=2, max_splits=112, batch_max=5), + _MlaBucket(num_warps=2, num_stages=2, max_splits=256, batch_max=24), + _MlaBucket(num_warps=1, num_stages=1, max_splits=256), +) + +# For the paths that must not depend on the batch; the mid bucket sits between the +# other two geometries. Retuning it moves what deterministic inference produces, which +# test_batch_free_geometry_is_pinned guards. max_splits goes unused there. +_MLA_BUCKET_BATCH_FREE = _MLA_BUCKETS[1] + +_KEEP_SCHEDULER_SPLITS = None +_CORE_COUNT = {} +_LOGGED_TUNE = False + + +def _keep_scheduler_splits() -> bool: + """Whether the caller asked for a specific per-sequence num_kv_splits. + + ``--enable-deterministic-inference`` derives it from a fixed tile size so a + request's reduction tree cannot depend on its batch mates; a batch-wide count puts + that back. An explicit tile size or the static-splits env asks for the same thing. + """ + global _KEEP_SCHEDULER_SPLITS + if _KEEP_SCHEDULER_SPLITS is None: + from sglang.srt.runtime_context import get_exec + + try: + exec_cfg = get_exec() + except ValueError: + return False # not published yet, ask again on the next call + _KEEP_SCHEDULER_SPLITS = bool( + exec_cfg.deterministic.enable_deterministic_inference + or exec_cfg.kernel.triton_attention_split_tile_size + or envs.SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS.get() + ) + if _KEEP_SCHEDULER_SPLITS: + logger.info("MLA decode: keeping the scheduler's num_kv_splits") + return _KEEP_SCHEDULER_SPLITS + + +def _grouped_head_tiles(head_num: int, kv_group_num: int) -> int: + """Stage-1's grid extent along heads.""" + return triton.cdiv(head_num, min(_GROUPED_BLOCK_H, kv_group_num)) + + +def _mla_bucket(batch: int) -> _MlaBucket: + for bucket in _MLA_BUCKETS[:-1]: + if batch <= bucket.batch_max: + return bucket + return _MLA_BUCKETS[-1] + + +def _mla_split_budget(num_warps: int, core_count: int) -> int: + # about one wave of stage-1 workgroups, taking 4 warps to get one per CU and + # halving the warps to double how many fit. core_count, not a whole MI355X: a CPX + # partition exposes 32 of the 256 + return core_count * 4 // num_warps + + +def _mla_core_count(device_index: Optional[int]) -> int: + count = _CORE_COUNT.get(device_index) + if count is None: + count = get_device_core_count(device_index if device_index is not None else 0) + _CORE_COUNT[device_index] = count + return count + + +def _mla_kv_splits( + batch: int, head_tiles: int, max_kv_splits: int, core_count: int +) -> int: + """Batch-wide split count for stage-1, or 0 with no device to size it against. + + The budget is a ceiling, not a rounding target: crossing it costs a step, not a + proportional slice (batch 24, 68k: 21 splits / 504 blocks 358 us, 22 splits / + 528 blocks 528 us). Below it the count stays exact, since each split + shortens the KV every workgroup walks (batch 136: 7 splits 1628 us, 4 at 2734 us). + """ + if core_count <= 0: + return 0 + bucket = _mla_bucket(batch) + budget = _mla_split_budget(bucket.num_warps, core_count) + splits = min(max_kv_splits, bucket.max_splits, budget // max(1, batch * head_tiles)) + return max(1, splits) + + +def _mla_tuning_applies(has_mla: bool, head_dim: int) -> bool: + # both gates matter: tuned on gfx950 and on Lk=576. Cheapest term first since this + # runs per layer per decode step, and the env read stays uncached so a test + # override lands + return ( + _is_hip + and has_mla + and head_dim == 576 + and is_gfx95_supported() + and envs.SGLANG_MLA_DECODE_TUNE.get() + ) + + +def _mla_launch_plan( + q, k_buffer, max_kv_splits: int, has_mla: bool +) -> Tuple[bool, int]: + """``(take the tuned geometry, batch-wide split count)`` for one decode call. + + Both launches get one decision: stage-2 must merge exactly as many partials as + stage-1 wrote and a mismatch is silent, so neither the count nor the gate is + re-derived per launcher. 0 leaves both stages on the scheduler's per-sequence + counts, their default. + """ + if not _mla_tuning_applies(has_mla, k_buffer.shape[-1]): + return False, 0 + if _keep_scheduler_splits(): + return True, 0 + head_num = q.shape[1] + head_tiles = _grouped_head_tiles(head_num, head_num // k_buffer.shape[-2]) + splits = _mla_kv_splits( + q.shape[0], head_tiles, max_kv_splits, _mla_core_count(q.device.index) + ) + + global _LOGGED_TUNE + if splits and not _LOGGED_TUNE: + _LOGGED_TUNE = True + logger.info( + "MLA decode: gfx950 tuned stage-1 geometry, replacing the scheduler's " + "num_kv_splits and capped by --triton-attention-num-kv-splits " + "(SGLANG_MLA_DECODE_TUNE=0 to disable)" + ) + return True, splits + def _extract_kv_strides(buf, page_size: int): """Extract (slot_stride, head_stride, page_stride, tok_stride) for a @@ -425,6 +581,8 @@ def _fwd_grouped_kernel_stage1( aux0_stride_t=0, aux0_stride_h=0, aux0_len=0, + forced_kv_splits=0, + USE_FORCED: tl.constexpr = False, ): # int64 to avoid overflow of flat offsets into Mid_O when # batch * num_head * max_kv_splits * head_dim exceeds 2**31. @@ -448,7 +606,14 @@ def _fwd_grouped_kernel_stage1( cur_batch_kv_start_idx = tl.load(kv_indptr + cur_batch) cur_batch_seq_len = tl.load(kv_indptr + cur_batch + 1) - cur_batch_kv_start_idx - kv_splits = tl.load(num_kv_splits + cur_batch) + # runtime, not constexpr: it only feeds the kv_len_per_split arithmetic below, so + # a constexpr buys nothing and costs one stage-1 variant per cuda-graph ladder + # rung (stage-2 does need it at compile time). Any count covers any length since + # kv_len_per_split rounds cdiv(L, S) up; short sequences leave the tail empty. + if USE_FORCED: + kv_splits = forced_kv_splits + else: + kv_splits = tl.load(num_kv_splits + cur_batch) if xai_temperature_len > 0: offs_qidx = cur_batch_seq_len - 1 @@ -626,6 +791,8 @@ def _decode_grouped_att_m_fwd( page_size: int = 1, score_mod=None, aux_tensors=None, + tune_mla: bool = False, + forced_kv_splits: int = 0, ): BLOCK = 32 Lk = k_buffer.shape[-1] @@ -652,22 +819,32 @@ def _decode_grouped_att_m_fwd( batch, head_num = q.shape[0], q.shape[1] kv_group_num = q.shape[1] // kv_head_num - BLOCK_H = 16 + BLOCK_H = _GROUPED_BLOCK_H MAX_KV_SPLITS = max_kv_splits - grid = ( - batch, - triton.cdiv(head_num, min(BLOCK_H, kv_group_num)), - MAX_KV_SPLITS, - ) + head_tiles = _grouped_head_tiles(head_num, kv_group_num) extra_kargs = {} num_stages = 2 + num_warps = 4 if _is_hip: # https://rocm.docs.amd.com/en/docs-6.2.0/how-to/llm-fine-tuning-optimization/optimizing-triton-kernel.html # https://github.com/triton-lang/triton/blob/main/third_party/amd/backend/compiler.py extra_kargs = {"waves_per_eu": 1, "matrix_instr_nonkdim": 16, "kpack": 2} num_stages = 1 + if tune_mla: + # num_warps reorders the fp32 accumulation, so whoever declined the batch-wide + # count gets a batch-independent geometry too + bucket = _mla_bucket(batch) if forced_kv_splits else _MLA_BUCKET_BATCH_FREE + BLOCK, num_warps, num_stages = ( + _MLA_BLOCK_N, + bucket.num_warps, + bucket.num_stages, + ) + + # Blocks at or above the split count return immediately, so the grid shrinks too. + grid = (batch, head_tiles, forced_kv_splits or MAX_KV_SPLITS) + k_slot_stride, k_head_stride, k_page_stride, k_tok_stride = _extract_kv_strides( k_buffer, page_size ) @@ -712,7 +889,7 @@ def _decode_grouped_att_m_fwd( MIN_BLOCK_KV=_MIN_BLOCK_KV, logit_cap=logit_cap, xai_temperature_len=xai_temperature_len, - num_warps=4, + num_warps=num_warps, num_stages=num_stages, Lk=Lk, Lv=Lv, @@ -724,6 +901,8 @@ def _decode_grouped_att_m_fwd( aux0_stride_t=aux0_stride_t, aux0_stride_h=aux0_stride_h, aux0_len=aux0_len, + forced_kv_splits=forced_kv_splits, + USE_FORCED=forced_kv_splits > 0, **extra_kargs, ) @@ -748,6 +927,7 @@ def _fwd_kernel_stage2( Lv: tl.constexpr, HAS_SINK: tl.constexpr, USE_PDL: tl.constexpr = False, + FORCED_KV_SPLITS: tl.constexpr = 0, ): # int64 to avoid overflow of flat offsets into Mid_O when # batch * num_head * max_kv_splits * head_dim exceeds 2**31. @@ -760,7 +940,16 @@ def _fwd_kernel_stage2( cur_batch_seq_len = tl.load(kv_indptr + cur_batch + 1) - tl.load( kv_indptr + cur_batch ) - kv_splits = tl.load(num_kv_splits + cur_batch) + # Same count stage-1 used, or the two disagree about where split i starts. SPLIT_END + # is a constexpr in both branches: a dynamic bound would merge the same partials + # (stage-1 leaves the surplus splits masked out) but stops the unrolling, and + # reassociating the fp32 reduction moves the result a few ULP off stock. + if FORCED_KV_SPLITS > 0: + kv_splits = FORCED_KV_SPLITS + SPLIT_END: tl.constexpr = FORCED_KV_SPLITS + else: + kv_splits = tl.load(num_kv_splits + cur_batch) + SPLIT_END: tl.constexpr = MAX_KV_SPLITS offs_d = tl.arange(0, BLOCK_DV) mask_d = offs_d < Lv @@ -775,7 +964,7 @@ def _fwd_kernel_stage2( tl.cdiv(tl.cdiv(cur_batch_seq_len, kv_splits), MIN_BLOCK_KV) * MIN_BLOCK_KV ) - for split_kv_id in tl.range(0, MAX_KV_SPLITS, num_stages=2): + for split_kv_id in tl.range(0, SPLIT_END, num_stages=2): split_kv_start = kv_len_per_split * split_kv_id split_kv_end = tl.minimum(split_kv_start + kv_len_per_split, cur_batch_seq_len) @@ -817,6 +1006,7 @@ def _decode_softmax_reducev_fwd( max_kv_splits, sinks=None, use_pdl=False, + forced_kv_splits: int = 0, ): batch, head_num = q.shape[0], q.shape[1] Lv = v_buffer.shape[-1] @@ -851,6 +1041,7 @@ def _decode_softmax_reducev_fwd( Lv=Lv, HAS_SINK=HAS_SINK, USE_PDL=use_pdl, + FORCED_KV_SPLITS=forced_kv_splits, num_warps=4, num_stages=2, **({"launch_pdl": True} if use_pdl else {}), @@ -931,6 +1122,7 @@ def decode_attention_fwd_grouped( score_mod=None, aux_tensors=None, ): + tune_mla, forced_kv_splits = _mla_launch_plan(q, k_buffer, max_kv_splits, has_mla) _decode_grouped_att_m_fwd( q, k_buffer, @@ -949,6 +1141,8 @@ def decode_attention_fwd_grouped( page_size=page_size, score_mod=score_mod, aux_tensors=aux_tensors, + tune_mla=tune_mla, + forced_kv_splits=forced_kv_splits, ) _decode_softmax_reducev_fwd( attn_logits, @@ -962,6 +1156,7 @@ def decode_attention_fwd_grouped( max_kv_splits, sinks, use_pdl=use_pdl, + forced_kv_splits=forced_kv_splits, ) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 16007ece1..e3ad080f3 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -933,6 +933,9 @@ class Envs: SGLANG_CRASH_ON_TRITON_LOAD_AFTER_READY = EnvBool(False) SGLANG_TRITON_SLOW_COMPILE_THRESHOLD_SECS = EnvFloat(1.0) SGLANG_TRITON_LOAD_WARNING_THRESHOLD_GB = EnvFloat(1.0) + # gfx950 MLA decode stage-1: pick the launch geometry and split count per batch. + # Reorders the fp32 accumulation, so off by default. + SGLANG_MLA_DECODE_TUNE = EnvBool(False) SGLANG_ENABLE_TORCH_COMPILE = EnvBool(False) SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE = EnvInt(4096) SGLANG_TRITON_DECODE_SPLIT_TILE_SIZE = EnvInt(256) diff --git a/test/registered/unit/layers/attention/test_mla_decode_forced_splits.py b/test/registered/unit/layers/attention/test_mla_decode_forced_splits.py new file mode 100644 index 000000000..b7d1c7c01 --- /dev/null +++ b/test/registered/unit/layers/attention/test_mla_decode_forced_splits.py @@ -0,0 +1,259 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""The gfx950 MLA decode tuning has to stay numerically equivalent to the stock path. + +SGLANG_MLA_DECODE_TUNE replaces the scheduler's per-sequence num_kv_splits with one +batch-wide count. Its failure mode is silent: stage-2 derives where split `i` starts +from the same count stage-1 used, so if only one of the two launches gets the count, +the merge reads partials that were never written and the output is quietly wrong +rather than an error. Only running both stages for real catches that. + + python -m pytest test/registered/unit/layers/attention/test_mla_decode_forced_splits.py -v +""" + +import unittest + +import torch + +from sglang.kernels.ops.attention import decode_attention as da +from sglang.srt.environ import envs +from sglang.srt.runtime_context import get_context +from sglang.srt.utils import is_gfx95_supported, is_hip +from sglang.test.ci.ci_register import register_amd_ci +from sglang.test.test_utils import CustomTestCase + +# The gfx95 runner: stage-b-test-1-gpu-large-amd is MI300, where every case below skips. +register_amd_ci(est_time=6, suite="stage-b-test-1-gpu-small-amd-mi35x") + +LK, LV = 576, 512 + +# A different split count reassociates the fp32 softmax reduction, so the result can +# land a rounding step away once it is stored back to bf16. Measured at exactly one +# bf16 ULP on every case below, hence rtol of two. The atol floor is for outputs small +# enough that the error is set by the terms being summed rather than by the result: +# 1.34e-3 was the most any case needed, on elements around 0.04. +BF16_ULP = 2**-8 +ATOL = 3e-3 + +# Batches straddling every bucket edge, plus the head counts that give head_tiles 1 +# and 8, plus the sequence shapes a batch-wide split count has to survive. +# +# The two head_tiles=8 cases (h128) are the only ones that catch stage-1 being handed +# the wrong count while stage-2 keeps the right one; do not drop both. +CASES = ( + ("b1_h16", [4096], 16, 1), + ("b5_h128", [4096] * 5, 128, 1), + ("b6_h16", [4096] * 6, 16, 1), + ("b24_h16", [4096] * 24, 16, 1), + ("b25_h16", [4096] * 25, 16, 1), + ("b136_h16", [1024] * 136, 16, 1), + ("mixed_lengths", [1, 31, 33, 257, 1024, 4095, 4096, 16384], 16, 1), + ("mixed_skew", [16384, 1, 1, 1, 1, 1, 1, 1], 16, 1), + ("all_length_1", [1, 1, 1, 1], 16, 1), + ("mixed_h128", [7, 512, 1023, 1025, 2048, 4095], 128, 1), + ("page64", [4096] * 4, 16, 64), + ("page64_mixed", [1 + (257 * i) % 2048 for i in range(32)], 16, 64), +) + + +def _inputs(seq_lens, head_num, page_size, max_kv_splits, seed): + """Stock decode inputs, including the num_kv_splits the scheduler would write.""" + gen = torch.Generator(device="cuda").manual_seed(seed) + dev, batch = "cuda", len(seq_lens) + total = sum(seq_lens) + n_slots = total + 64 + + if page_size == 1: + pool = torch.randn( + n_slots, 1, LK, dtype=torch.bfloat16, device=dev, generator=gen + ) + else: + n_pages = (n_slots + page_size - 1) // page_size + pool = torch.randn( + n_pages, page_size, 1, LK, dtype=torch.bfloat16, device=dev, generator=gen + ) + n_slots = n_pages * page_size + + kv_indptr = torch.zeros(batch + 1, dtype=torch.int32, device=dev) + kv_indptr[1:] = torch.cumsum( + torch.tensor(seq_lens, dtype=torch.int32, device=dev), dim=0 + ) + # scattered slots, like a pool that has been recycled + kv_indices = torch.randperm(n_slots, device=dev, generator=gen)[:total].to( + torch.int32 + ) + + lens = torch.tensor(seq_lens, dtype=torch.int32, device=dev) + num_kv_splits = torch.clamp( + torch.div(lens, 256, rounding_mode="floor") + 1, 1, max_kv_splits + ).to(torch.int32) + + return { + "q": torch.randn( + batch, head_num, LK, dtype=torch.bfloat16, device=dev, generator=gen + ), + "k_buffer": pool, + "v_buffer": pool[..., :LV], + "kv_indptr": kv_indptr, + "kv_indices": kv_indices, + "num_kv_splits": num_kv_splits, + "max_kv_splits": max_kv_splits, + "page_size": page_size, + } + + +def _forced(inp): + return da._mla_launch_plan(inp["q"], inp["k_buffer"], inp["max_kv_splits"], True)[1] + + +def _run(inp): + batch, head_num = inp["q"].shape[0], inp["q"].shape[1] + dev, mks = "cuda", inp["max_kv_splits"] + o = torch.zeros(batch, head_num, LV, dtype=torch.bfloat16, device=dev) + logits = torch.empty(batch, head_num, mks, LV, dtype=torch.float32, device=dev) + # the DCP path relies on untouched entries staying -inf + lse = torch.full( + (batch, head_num, mks), -float("inf"), dtype=torch.float32, device=dev + ) + da.decode_attention_fwd_grouped( + inp["q"], + inp["k_buffer"], + inp["v_buffer"], + o, + inp["kv_indptr"], + inp["kv_indices"], + logits, + lse, + inp["num_kv_splits"], + mks, + 1.0 / LK**0.5, + 1.0, + has_mla=True, + page_size=inp["page_size"], + ) + torch.cuda.synchronize() + return o + + +@unittest.skipUnless( + is_hip() and is_gfx95_supported(), "the tuning only engages on gfx95" +) +class TestMlaDecodeForcedSplits(CustomTestCase): + def setUp(self): + # a plain config, so the count comes out batch-wide whatever ran before this + self._publish() + da._LOGGED_TUNE = False + + def tearDown(self): + # None, not False: the real resolution has to run again for anything later in + # this process. + da._KEEP_SCHEDULER_SPLITS = None + + def _publish(self, **fields): + override = get_context().override_server_args(**fields) + override.install() + self.addCleanup(override.restore) + da._KEEP_SCHEDULER_SPLITS = None # resolve it from what was just published + + def test_tuned_matches_stock(self): + # Not bit-for-bit: a different split count reassociates the fp32 softmax + # reduction. Bounded at two bf16 ULP, which is one rounding step of headroom + # over what this actually measures. + for seed, (name, seq_lens, head_num, page_size) in enumerate(CASES): + with self.subTest(case=name): + inp = _inputs(seq_lens, head_num, page_size, 256, seed=seed) + with envs.SGLANG_MLA_DECODE_TUNE.override(False): + stock = _run(inp) + self.assertEqual(_forced(inp), 0) + with envs.SGLANG_MLA_DECODE_TUNE.override(True): + # without this the comparison passes by comparing stock to stock + self.assertGreater(_forced(inp), 0, "tuning did not engage") + tuned = _run(inp) + self.assertFalse(torch.isnan(tuned).any()) + torch.testing.assert_close( + tuned.float(), + stock.float(), + rtol=2 * BF16_ULP, + atol=ATOL, + ) + + def test_both_stages_get_the_same_count(self): + # The mismatch the tolerance above can only catch indirectly: assert the entry + # point hands one count to both launches instead of letting them disagree. + # Dropping tune_mla is invisible to the numerics, it only costs the geometry. + seen = {} + + def record(key, real): + def wrapper(*args, forced_kv_splits=0, **kwargs): + seen[key] = (forced_kv_splits, kwargs.get("tune_mla")) + return real(*args, forced_kv_splits=forced_kv_splits, **kwargs) + + return wrapper + + stage1, stage2 = da._decode_grouped_att_m_fwd, da._decode_softmax_reducev_fwd + da._decode_grouped_att_m_fwd = record("stage1", stage1) + da._decode_softmax_reducev_fwd = record("stage2", stage2) + try: + with envs.SGLANG_MLA_DECODE_TUNE.override(True): + _run(_inputs([4096] * 24, 16, 1, 256, seed=0)) + finally: + da._decode_grouped_att_m_fwd = stage1 + da._decode_softmax_reducev_fwd = stage2 + + self.assertEqual(seen["stage1"][0], seen["stage2"][0]) + self.assertGreater( + seen["stage1"][0], 0, "tuning did not engage, nothing tested" + ) + self.assertTrue(seen["stage1"][1], "stage-1 was left on the stock geometry") + + def test_scheduler_splits_are_kept_when_asked(self): + # the geometry still changes under --enable-deterministic-inference, but it is + # batch-independent (see the next test); only the count has to survive verbatim + inp = _inputs([1, 4095, 4096, 16384], 16, 1, 256, seed=7) + with envs.SGLANG_MLA_DECODE_TUNE.override(False): + stock = _run(inp) + self._publish(enable_deterministic_inference=True) + with envs.SGLANG_MLA_DECODE_TUNE.override(True): + self.assertEqual(_forced(inp), 0) + kept = _run(inp) + torch.testing.assert_close( + kept.float(), stock.float(), rtol=2 * BF16_ULP, atol=ATOL + ) + + def test_deterministic_mode_is_batch_invariant(self): + # The point of keeping the scheduler's count: a request's reduction tree must + # not depend on who it shares the batch with. + self._publish(enable_deterministic_inference=True) + lens = [4096, 1, 777, 16384] + with envs.SGLANG_MLA_DECODE_TUNE.override(True): + batched = _run(_inputs(lens, 16, 1, 256, seed=3)) + for i in range(len(lens)): + inp = _inputs(lens, 16, 1, 256, seed=3) + # same q and same KV slots for row i, on its own + start, end = inp["kv_indptr"][i].item(), inp["kv_indptr"][i + 1].item() + alone = dict(inp) + alone["q"] = inp["q"][i : i + 1].clone() + alone["kv_indices"] = inp["kv_indices"][start:end].clone() + alone["kv_indptr"] = torch.tensor( + [0, end - start], dtype=torch.int32, device="cuda" + ) + alone["num_kv_splits"] = inp["num_kv_splits"][i : i + 1].clone() + with self.subTest(row=i, seq_len=lens[i]): + torch.testing.assert_close( + _run(alone)[0], batched[i], rtol=0, atol=0 + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/layers/attention/test_mla_decode_geometry.py b/test/registered/unit/layers/attention/test_mla_decode_geometry.py new file mode 100644 index 000000000..f33ae2ab8 --- /dev/null +++ b/test/registered/unit/layers/attention/test_mla_decode_geometry.py @@ -0,0 +1,224 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Stage-1 split budget for the gfx950 MLA decode geometry. + +Stage-1 launches `batch * head_tiles * kv_splits` workgroups, and crossing the +bucket's budget costs a step rather than a proportional slice: at batch 24 on a 68k +context, 21 splits measured 358 us against 528 us for 22. Which is also why the +budget is divided down with floor and not round -- `round(512/12) = 43` would put +batch 12 at 516 blocks, just over. + +The budget lives with the geometry rather than as a global constant, since halving +`num_warps` moved the cliff from 512 blocks to 1024. These tests pin the two +together, so retuning one without the other, or rounding the division up, fails here +instead of costing 50% at one batch size. + + python -m pytest test/registered/unit/layers/attention/test_mla_decode_geometry.py -v +""" + +import unittest + +from sglang.kernels.ops.attention import decode_attention as da +from sglang.kernels.ops.attention.decode_attention import ( + _MLA_BLOCK_N, + _MLA_BUCKET_BATCH_FREE, + _MLA_BUCKETS, + _fwd_grouped_kernel_stage1, + _grouped_head_tiles, + _keep_scheduler_splits, + _mla_bucket, + _mla_kv_splits, + _mla_split_budget, +) +from sglang.srt.environ import envs +from sglang.srt.runtime_context import get_context +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + +# gfx950 full-GPU. Passed in rather than read from the device so this stays a CPU test. +CORE_COUNT = 256 + +# The batch sizes cuda-graph capture walks, and the split count separately +# measured as optimal at each on a 68k context (head_tiles == 1, i.e. K3 at tp 8). +MEASURED_OPTIMUM = { + 1: 112, + 2: 112, + 3: 85, + 4: 64, + 5: 51, + 6: 85, + 7: 73, + 8: 64, + 10: 51, + 12: 42, + 14: 36, + 16: 32, + 20: 25, + 24: 21, + 28: 36, + 32: 32, +} + +MAX_KV_SPLITS = 256 + + +class TestMlaDecodeGeometry(unittest.TestCase): + def test_rule_reproduces_measured_optimum(self): + # a budget plus two constants, not a fit, so it hits the measured optimum + for batch, want in MEASURED_OPTIMUM.items(): + with self.subTest(batch=batch): + self.assertEqual( + _mla_kv_splits(batch, 1, MAX_KV_SPLITS, CORE_COUNT), want + ) + + def test_stage1_split_count_stays_runtime(self): + # A constexpr here would compile one stage-1 variant per rung of the capture + # ladder (21 for the default one), and the count only feeds kv_len_per_split. + params = {p.name: p for p in _fwd_grouped_kernel_stage1.params} + self.assertFalse(params["forced_kv_splits"].is_constexpr) + self.assertTrue(params["USE_FORCED"].is_constexpr) + + def test_batch_free_geometry_is_pinned(self): + # Deterministic inference runs on this geometry and BLOCK_N/num_warps reorder + # the fp32 accumulation, so retuning either moves those numbers. + self.assertEqual( + ( + _MLA_BLOCK_N, + _MLA_BUCKET_BATCH_FREE.num_warps, + _MLA_BUCKET_BATCH_FREE.num_stages, + ), + (32, 2, 2), + ) + + def test_head_tiles_matches_the_grid(self): + # The budget is divided by the same head_tiles the grid is launched with; a + # BLOCK_H that drifts between the two silently mis-sizes the budget. + for head_num, kv_group_num, want in ((16, 16, 1), (128, 128, 8), (8, 8, 1)): + with self.subTest(head_num=head_num): + self.assertEqual(_grouped_head_tiles(head_num, kv_group_num), want) + + def test_never_crosses_the_budget(self): + for head_tiles in (1, 2): + for batch in range(1, 1025): + splits = _mla_kv_splits(batch, head_tiles, MAX_KV_SPLITS, CORE_COUNT) + budget = _mla_split_budget(_mla_bucket(batch).num_warps, CORE_COUNT) + with self.subTest(batch=batch, head_tiles=head_tiles): + if batch * head_tiles <= budget: + self.assertLessEqual(batch * head_tiles * splits, budget) + else: + # already past the budget, so 1 is the floor + self.assertEqual(splits, 1) + + def test_low_batch_is_capped_not_scaled(self): + # below 6 the budget stops binding, and more splits stopped paying at 112 + # whatever the batch, so the cap sits on top of the budget instead of scaling + self.assertEqual(_mla_kv_splits(1, 1, MAX_KV_SPLITS, CORE_COUNT), 112) + self.assertEqual(_mla_kv_splits(2, 1, MAX_KV_SPLITS, CORE_COUNT), 112) + self.assertLess( + 1 * 112, _mla_split_budget(_mla_bucket(1).num_warps, CORE_COUNT) + ) + + def test_caller_max_kv_splits_wins(self): + for cap in (1, 4, 16): + with self.subTest(max_kv_splits=cap): + self.assertLessEqual(_mla_kv_splits(1, 1, cap, CORE_COUNT), cap) + + def test_at_least_one_split(self): + for batch in (1, 4096): + with self.subTest(batch=batch): + self.assertGreaterEqual( + _mla_kv_splits(batch, 1, MAX_KV_SPLITS, CORE_COUNT), 1 + ) + + def test_buckets_are_ordered_and_total(self): + bounds = [b.batch_max for b in _MLA_BUCKETS] + self.assertIsNone(bounds[-1], "the last bucket has to be the catch-all") + finite = bounds[:-1] + self.assertEqual(finite, sorted(finite)) + self.assertTrue(all(b is not None for b in finite)) + + def test_wider_workgroup_gets_a_tighter_budget(self): + # the budget divides by num_warps because that is where the cliff moved: + # halving the warps took it from 512 blocks to 1024 + budgets = [ + _mla_split_budget(w, CORE_COUNT) + for w in sorted({b.num_warps for b in _MLA_BUCKETS}, reverse=True) + ] + self.assertEqual(budgets, sorted(budgets)) + + def test_budget_follows_the_partition_size(self): + # A CPX partition exposes 32 of the 256 CUs while is_gfx95_supported() still + # says yes, so a budget pinned to the whole GPU would oversubscribe it 8x. + for warps in (1, 2, 4): + with self.subTest(num_warps=warps): + self.assertEqual( + _mla_split_budget(warps, 32) * 8, + _mla_split_budget(warps, 256), + ) + self.assertEqual( + _mla_kv_splits(8, 1, MAX_KV_SPLITS, 0), 0, "no core count, no budget" + ) + + def test_splits_shrink_on_a_partition(self): + # Same batch, smaller partition -> fewer splits, never more. + for batch in (8, 16, 32, 64): + with self.subTest(batch=batch): + self.assertLessEqual( + _mla_kv_splits(batch, 1, MAX_KV_SPLITS, 32), + _mla_kv_splits(batch, 1, MAX_KV_SPLITS, 256), + ) + + +class TestKeepSchedulerSplits(unittest.TestCase): + """Which configs decline the batch-wide count. + + Through override_server_args, so the flags resolve the way a launched server + resolves them; poking the cached decision keeps passing after they move namespace. + """ + + def _publish(self, **fields): + override = get_context().override_server_args(**fields) + override.install() + self.addCleanup(override.restore) + # resolved once per process, so clear it at both ends + da._KEEP_SCHEDULER_SPLITS = None + self.addCleanup(setattr, da, "_KEEP_SCHEDULER_SPLITS", None) + + def test_a_plain_config_takes_the_batch_wide_count(self): + self._publish() + self.assertFalse(_keep_scheduler_splits()) + + def test_deterministic_inference_keeps_the_scheduler_splits(self): + self._publish(enable_deterministic_inference=True) + self.assertTrue(_keep_scheduler_splits()) + + def test_an_explicit_split_tile_size_keeps_them(self): + self._publish(triton_attention_split_tile_size=256) + self.assertTrue(_keep_scheduler_splits()) + + def test_the_static_kv_splits_env_keeps_them(self): + self._publish() + with envs.SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS.override(True): + self.assertTrue(_keep_scheduler_splits()) + + def test_the_geometry_rule_does_not_read_the_config(self): + # _mla_kv_splits answers for a device, not for a config; the decline lives one + # level up, so a deterministic config elsewhere cannot rewrite the pins above + self._publish(enable_deterministic_inference=True) + self.assertEqual(_mla_kv_splits(24, 1, MAX_KV_SPLITS, CORE_COUNT), 21) + + +if __name__ == "__main__": + unittest.main()