[AMD] Optimize KIMI-K3 with Triton MLA decode kernel by tuning the stage-1 geometry for gfx950 (#34580)
Co-authored-by: Thomas Wang <thomawan@amd.com>
This commit is contained in:
co-authored by
Thomas Wang
parent
53621818e4
commit
d01812d89e
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user