[AMD] dsv4: pick kv_splits per index stream, not by occupancy alone (#39968)
Co-authored-by: wunhuang <wunhuang@amd.com>
This commit is contained in:
@@ -52,6 +52,7 @@ the captured launch sequence.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import functools
|
import functools
|
||||||
|
import os
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton
|
import triton
|
||||||
@@ -89,6 +90,48 @@ _FP8_GROUP_SIZE = 64
|
|||||||
_FP8_DTYPE = torch.float8_e4m3fnuz if is_fp8_fnuz() else torch.float8_e4m3fn
|
_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)
|
@functools.lru_cache(maxsize=1)
|
||||||
def _cu_count() -> int:
|
def _cu_count() -> int:
|
||||||
"""Compute-unit count of the active GPU, queried once via aiter.
|
"""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,
|
attn_sink: torch.Tensor,
|
||||||
softmax_scale: float,
|
softmax_scale: float,
|
||||||
kv_scales: torch.Tensor | None = None,
|
kv_scales: torch.Tensor | None = None,
|
||||||
|
kv_splits: int | None = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""V4 decode sparse attention over a unified KV pool with paged indices.
|
"""V4 decode sparse attention over a unified KV pool with paged indices.
|
||||||
|
|
||||||
When ``kv_scales`` is provided, ``unified_kv`` must be fp8 (e4m3fnuz) and
|
When ``kv_scales`` is provided, ``unified_kv`` must be fp8 (e4m3fnuz) and
|
||||||
will be dequantized in-kernel using 1xGROUP_SIZE (default 64) block scales.
|
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:
|
if _is_gfx1250_supported:
|
||||||
# aiter ships only on ROCm, and this module is imported by a CPU-registered
|
# 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,
|
attn_sink,
|
||||||
softmax_scale,
|
softmax_scale,
|
||||||
kv_scales=kv_scales,
|
kv_scales=kv_scales,
|
||||||
|
kv_splits=kv_splits,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -233,9 +233,16 @@ def decode(
|
|||||||
kv_indptr: torch.Tensor,
|
kv_indptr: torch.Tensor,
|
||||||
attn_sink: torch.Tensor, # [H] fp32
|
attn_sink: torch.Tensor, # [H] fp32
|
||||||
softmax_scale: float,
|
softmax_scale: float,
|
||||||
|
kv_splits: int | None = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
return sparse_attn_v4_paged_decode(
|
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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1664,6 +1664,10 @@ class DeepseekV4HipRadixBackend(
|
|||||||
attn_sink=attn_sink,
|
attn_sink=attn_sink,
|
||||||
v_head_dim=layer.v_head_dim,
|
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(
|
return runtime.decode(
|
||||||
q=q,
|
q=q,
|
||||||
unified_kv=unified,
|
unified_kv=unified,
|
||||||
@@ -1671,6 +1675,9 @@ class DeepseekV4HipRadixBackend(
|
|||||||
kv_indptr=kv_indptr,
|
kv_indptr=kv_indptr,
|
||||||
attn_sink=attn_sink,
|
attn_sink=attn_sink,
|
||||||
softmax_scale=self.softmax_scale,
|
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
|
# prefill / extend
|
||||||
|
|||||||
Reference in New Issue
Block a user