From c5326d28a348dfba14e908cb055ea4c8ec2a92fa Mon Sep 17 00:00:00 2001 From: kk <43161300+kkHuang-amd@users.noreply.github.com> Date: Sun, 20 Sep 2026 02:53:54 +0800 Subject: [PATCH] [AMD] dsv4: pick kv_splits per index stream, not by occupancy alone (#39968) Co-authored-by: wunhuang --- .../dsv4/unified_kv_kernels/paged_decode.py | 48 +++++++++++++++++++ .../dsv4/unified_kv_kernels/runtime.py | 9 +++- .../deepseek_v4_backend_hip_radix.py | 7 +++ 3 files changed, 63 insertions(+), 1 deletion(-) diff --git a/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_decode.py b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_decode.py index 391fe9c2a..ae70c9eaf 100644 --- a/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_decode.py +++ b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/paged_decode.py @@ -52,6 +52,7 @@ the captured launch sequence. from __future__ import annotations import functools +import os import torch import triton @@ -89,6 +90,48 @@ _FP8_GROUP_SIZE = 64 _FP8_DTYPE = torch.float8_e4m3fnuz if is_fp8_fnuz() else torch.float8_e4m3fn +# --- layer-aware split-K (SGLANG_MLA_HCA_KV_SPLITS, 0 = keep the heuristic) -- +# `_kv_splits_heuristic` is an OCCUPANCY rule: it splits only when the base grid +# underfills the device. That is the wrong question for this kernel. At bs=14 +# the grid nearly saturates (196 CTAs against a 384 target) so it picks +# splits=1, while the actual cost is set by ONE straggler CTA walking ~5,000 KV +# entries next to CTAs walking 200. Splitting breaks up that CTA; occupancy +# never sees it. +# +# The discriminator is static and known at capture time: `compress_ratio`. +# CSA (ratio 4) is clamped to index_topk+128 = 1152, dispersion ~0.19, +# and split-K LOSES on it (+5.1 % measured at bs=14). +# HCA (ratio 128) is unclamped, kv_len ~ context/128 reaching ~5,000, +# dispersion 0.70-0.78, and split-K wins at every kv_len the +# run traverses. +# +# Measured on the HCA shape at bs=14 vs the heuristic's splits=1, swept over the +# kv_len range the run actually traverses: +# median 200: splits 4 -41.8 %, 8 -35.8 % +# median 1300: splits 4 -50.1 %, 8 -54.0 % <- steady state +# median 3000: splits 4 -23.8 %, 8 -27.0 % +# 4 wins outright at the low end, gives ~93 % of 8 at steady state, and halves +# the partial buffers: acc_partial is T x splits x h_padded x D x 4 B = 103 MB +# at splits=4 against 205 MB at 8 for bs=14, charged inside the cuda-graph pool. +# +# End to end on DSv4 / MI355X / c128, n-weighted at matched batch size against +# the same config with the plain heuristic: decode step 121.76 -> 114.12 ms +# (-7.50 ms, -6.2 %). At c256 the gain is smaller (-5.37 ms of 156.19). +_HCA_KV_SPLITS = int(os.environ.get("SGLANG_MLA_HCA_KV_SPLITS", "4")) + + +def _kv_splits_for_stream(compress_ratio: int) -> int | None: + """-> kv_splits override for a decode stream, or None to keep the heuristic. + + Only the unclamped stream is overridden. SWA (ratio 0) is a 128-entry + window and CSA (ratio 4) is clamped, so both are near-uniform and the + occupancy heuristic is already right for them. + """ + if _HCA_KV_SPLITS > 0 and compress_ratio == 128: + return _HCA_KV_SPLITS + return None + + @functools.lru_cache(maxsize=1) def _cu_count() -> int: """Compute-unit count of the active GPU, queried once via aiter. @@ -900,11 +943,15 @@ def sparse_attn_v4_paged_decode( attn_sink: torch.Tensor, softmax_scale: float, kv_scales: torch.Tensor | None = None, + kv_splits: int | None = None, ) -> torch.Tensor: """V4 decode sparse attention over a unified KV pool with paged indices. When ``kv_scales`` is provided, ``unified_kv`` must be fp8 (e4m3fnuz) and will be dequantized in-kernel using 1xGROUP_SIZE (default 64) block scales. + + ``kv_splits`` overrides `_kv_splits_heuristic`. The caller knows the layer's + ``compress_ratio`` and the heuristic does not; see `_kv_splits_for_stream`. """ if _is_gfx1250_supported: # aiter ships only on ROCm, and this module is imported by a CPU-registered @@ -930,4 +977,5 @@ def sparse_attn_v4_paged_decode( attn_sink, softmax_scale, kv_scales=kv_scales, + kv_splits=kv_splits, ) diff --git a/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/runtime.py b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/runtime.py index 39f02a9d7..6dfdf7fcc 100644 --- a/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/runtime.py +++ b/python/sglang/kernels/ops/attention/dsv4/unified_kv_kernels/runtime.py @@ -233,9 +233,16 @@ def decode( kv_indptr: torch.Tensor, attn_sink: torch.Tensor, # [H] fp32 softmax_scale: float, + kv_splits: int | None = None, ) -> torch.Tensor: return sparse_attn_v4_paged_decode( - q, unified_kv, kv_indices, kv_indptr, attn_sink, softmax_scale + q, + unified_kv, + kv_indices, + kv_indptr, + attn_sink, + softmax_scale, + kv_splits=kv_splits, ) diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index bd86fe749..12fbf76e4 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -1664,6 +1664,10 @@ class DeepseekV4HipRadixBackend( attn_sink=attn_sink, v_head_dim=layer.v_head_dim, ) + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.paged_decode import ( + _kv_splits_for_stream, + ) + return runtime.decode( q=q, unified_kv=unified, @@ -1671,6 +1675,9 @@ class DeepseekV4HipRadixBackend( kv_indptr=kv_indptr, attn_sink=attn_sink, softmax_scale=self.softmax_scale, + # Only this call site knows compress_ratio, and it is the one + # thing that separates the ragged stream from the clamped ones. + kv_splits=_kv_splits_for_stream(compress_ratio), ) # prefill / extend