diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index cd1492ad9..9747787c1 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, List, Optional import torch import triton import triton.language as tl +from sgl_kernel.utils import is_arch_support_pdl from sglang.srt.configs.model_config import AttentionArch from sglang.srt.layers.attention.base_attn_backend import AttentionBackend @@ -28,6 +29,19 @@ if TYPE_CHECKING: from sglang.srt.speculative.spec_info import SpecInput +_MLA_DECODE_MIN_BLOCK_KV = 32 + + +def _mla_decode_kv_splits_cap( + base_max_kv_splits: int, sm_count: int, max_context_len: int +) -> int: + if sm_count <= 0: + return base_max_kv_splits + sm_cap = next_power_of_2(sm_count) + ctx_cap = next_power_of_2(triton.cdiv(max_context_len, _MLA_DECODE_MIN_BLOCK_KV)) + return max(base_max_kv_splits, min(sm_cap, ctx_cap)) + + def logit_capping_mod(logit_capping_method, logit_cap): # positive logit_cap -> tanh cap if logit_capping_method == "tanh": @@ -128,6 +142,13 @@ class TritonAttnBackend(AttentionBackend): "SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false" ) self.max_kv_splits = model_runner.server_args.triton_attention_num_kv_splits + if self.use_mla: + self.max_kv_splits = _mla_decode_kv_splits_cap( + self.max_kv_splits, + self.device_core_count, + self.max_context_len, + ) + self.use_pdl = is_arch_support_pdl() self.allow_bidirectional_attention_in_extend = ( model_runner.server_args.disable_cuda_graph @@ -887,15 +908,24 @@ class TritonAttnBackend(AttentionBackend): else: # Save KV cache first (must do this before unified kernel) if save_kv_cache: - if ( - self.use_mla or layer.k_scale is None - ): # Triton MLA currently doesn't support quantized kv cache + if layer.k_scale is None: forward_batch.token_to_kv_pool.set_kv_buffer( layer, forward_batch.out_cache_loc, k, v, ) + elif self.use_mla: + # For MLA, scale K manually before storing since MLATokenToKVPool + # doesn't accept scale parameters. Clone to protect k from mutation + # since it's used later in the attention kernel. + k_scaled = k.clone().div_(layer.k_scale) + forward_batch.token_to_kv_pool.set_kv_buffer( + layer, + forward_batch.out_cache_loc, + k_scaled, + v, + ) else: forward_batch.token_to_kv_pool.set_kv_buffer( layer, @@ -1139,7 +1169,11 @@ class TritonAttnBackend(AttentionBackend): logits_soft_cap = logit_capping_mod(layer.logit_capping_method, layer.logit_cap) if save_kv_cache: - if self.use_mla: # Triton MLA currently doesn't support quantized kv cache + if self.use_mla: + if layer.k_scale is not None: + # MLATokenToKVPool doesn't accept scale parameters; k is unused + # after this point in decode, so scale in place. + k.div_(layer.k_scale) forward_batch.token_to_kv_pool.set_kv_buffer( layer, forward_batch.out_cache_loc, @@ -1197,6 +1231,8 @@ class TritonAttnBackend(AttentionBackend): logit_cap=logits_soft_cap, sinks=sinks, xai_temperature_len=layer.xai_temperature_len, + has_mla=self.use_mla, + use_pdl=self.use_pdl, ) return o diff --git a/python/sglang/srt/layers/attention/triton_ops/decode_attention.py b/python/sglang/srt/layers/attention/triton_ops/decode_attention.py index 2b166f3b0..b42ffa433 100644 --- a/python/sglang/srt/layers/attention/triton_ops/decode_attention.py +++ b/python/sglang/srt/layers/attention/triton_ops/decode_attention.py @@ -281,6 +281,8 @@ def _fwd_grouped_kernel_stage1( xai_temperature_len: tl.constexpr, Lk: tl.constexpr, Lv: tl.constexpr, + HAS_MLA: tl.constexpr = False, + USE_PDL: tl.constexpr = False, ): cur_batch = tl.program_id(0) cur_head_id = tl.program_id(1) @@ -329,36 +331,36 @@ def _fwd_grouped_kernel_stage1( e_sum = tl.zeros([BLOCK_H], dtype=tl.float32) acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32) + # Hoist loop-invariant base offsets + base_offs_k = cur_kv_head * stride_buf_kh + offs_d[:, None] + if BLOCK_DPE > 0: + base_offs_kpe = cur_kv_head * stride_buf_kh + offs_dpe[:, None] + if not HAS_MLA: + base_offs_v = cur_kv_head * stride_buf_vh + offs_dv[None, :] + if split_kv_end > split_kv_start: q = tl.load(Q + offs_q, mask=(mask_h[:, None]) & (mask_d[None, :]), other=0.0) + q_k = q.to(K_Buffer.dtype.element_ty) if BLOCK_DPE > 0: qpe = tl.load( Q + off_qpe, mask=(mask_h[:, None]) & (mask_dpe[None, :]), other=0.0 ) - for start_n in range(split_kv_start, split_kv_end, BLOCK_N): + for start_n in tl.range(split_kv_start, split_kv_end, BLOCK_N): offs_n = start_n + tl.arange(0, BLOCK_N) kv_loc = tl.load( kv_indices + cur_batch_kv_start_idx + offs_n, mask=offs_n < split_kv_end, other=0, ) - offs_buf_k = ( - kv_loc[None, :] * stride_buf_kbs - + cur_kv_head * stride_buf_kh - + offs_d[:, None] - ) + offs_buf_k = kv_loc[None, :] * stride_buf_kbs + base_offs_k k = tl.load( K_Buffer + offs_buf_k, mask=(offs_n[None, :] < split_kv_end) & (mask_d[:, None]), other=0.0, ) - qk = tl.dot(q, k.to(q.dtype)) + qk = tl.dot(q_k, k) if BLOCK_DPE > 0: - offs_buf_kpe = ( - kv_loc[None, :] * stride_buf_kbs - + cur_kv_head * stride_buf_kh - + offs_dpe[:, None] - ) + offs_buf_kpe = kv_loc[None, :] * stride_buf_kbs + base_offs_kpe kpe = tl.load( K_Buffer + offs_buf_kpe, mask=(offs_n[None, :] < split_kv_end) & (mask_dpe[:, None]), @@ -376,17 +378,15 @@ def _fwd_grouped_kernel_stage1( qk = tl.where( mask_h[:, None] & (offs_n[None, :] < split_kv_end), qk, float("-inf") ) - - offs_buf_v = ( - kv_loc[:, None] * stride_buf_vbs - + cur_kv_head * stride_buf_vh - + offs_dv[None, :] - ) - v = tl.load( - V_Buffer + offs_buf_v, - mask=(offs_n[:, None] < split_kv_end) & (mask_dv[None, :]), - other=0.0, - ) + if HAS_MLA: + v = tl.trans(k) + else: + offs_buf_v = kv_loc[:, None] * stride_buf_vbs + base_offs_v + v = tl.load( + V_Buffer + offs_buf_v, + mask=(offs_n[:, None] < split_kv_end) & (mask_dv[None, :]), + other=0.0, + ) n_e_max = tl.maximum(tl.max(qk, 1), e_max) re_scale = tl.exp(e_max - n_e_max) @@ -422,6 +422,9 @@ def _fwd_grouped_kernel_stage1( mask=mask_h, ) + if USE_PDL: + tl.extra.cuda.gdc_launch_dependents() + def _decode_grouped_att_m_fwd( q, @@ -436,6 +439,8 @@ def _decode_grouped_att_m_fwd( sm_scale_withk, logit_cap, xai_temperature_len=-1, + has_mla=False, + use_pdl=False, ): BLOCK = 32 Lk = k_buffer.shape[-1] @@ -508,6 +513,8 @@ def _decode_grouped_att_m_fwd( num_stages=num_stages, Lk=Lk, Lv=Lv, + HAS_MLA=has_mla, + USE_PDL=use_pdl, **extra_kargs, ) @@ -531,10 +538,14 @@ def _fwd_kernel_stage2( BLOCK_DV: tl.constexpr, Lv: tl.constexpr, HAS_SINK: tl.constexpr, + USE_PDL: tl.constexpr = False, ): cur_batch = tl.program_id(0) cur_head = tl.program_id(1) + if USE_PDL: + tl.extra.cuda.gdc_wait() + cur_batch_seq_len = tl.load(kv_indptr + cur_batch + 1) - tl.load( kv_indptr + cur_batch ) @@ -553,7 +564,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 range(0, MAX_KV_SPLITS): + for split_kv_id in tl.range(0, MAX_KV_SPLITS, 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) @@ -594,6 +605,7 @@ def _decode_softmax_reducev_fwd( num_kv_splits, max_kv_splits, sinks=None, + use_pdl=False, ): batch, head_num = q.shape[0], q.shape[1] Lv = v_buffer.shape[-1] @@ -627,8 +639,10 @@ def _decode_softmax_reducev_fwd( BLOCK_DV=BLOCK_DV, Lv=Lv, HAS_SINK=HAS_SINK, + USE_PDL=use_pdl, num_warps=4, num_stages=2, + **({"launch_pdl": True} if use_pdl else {}), **extra_kargs, ) @@ -694,6 +708,8 @@ def decode_attention_fwd_grouped( logit_cap=0.0, sinks=None, xai_temperature_len=-1, + has_mla=False, + use_pdl=False, ): _decode_grouped_att_m_fwd( q, @@ -708,6 +724,8 @@ def decode_attention_fwd_grouped( sm_scale_withk, logit_cap, xai_temperature_len, + has_mla=has_mla, + use_pdl=use_pdl, ) _decode_softmax_reducev_fwd( attn_logits, @@ -720,6 +738,7 @@ def decode_attention_fwd_grouped( num_kv_splits, max_kv_splits, sinks, + use_pdl=use_pdl, ) @@ -740,6 +759,8 @@ def decode_attention_fwd( logit_cap=0.0, sinks=None, xai_temperature_len=-1, + has_mla=False, + use_pdl=False, ): assert max_kv_splits == attn_logits.shape[2] assert q.shape[0] <= kv_indptr.shape[0] - 1 @@ -784,4 +805,6 @@ def decode_attention_fwd( logit_cap=logit_cap, sinks=sinks, xai_temperature_len=xai_temperature_len, + has_mla=has_mla, + use_pdl=use_pdl, ) diff --git a/python/sglang/srt/layers/attention/triton_ops/extend_attention.py b/python/sglang/srt/layers/attention/triton_ops/extend_attention.py index fb57df6ed..e6a353e9b 100644 --- a/python/sglang/srt/layers/attention/triton_ops/extend_attention.py +++ b/python/sglang/srt/layers/attention/triton_ops/extend_attention.py @@ -382,7 +382,6 @@ def _fwd_kernel( mask=(mask_n[None, :]) & (mask_d[:, None]), other=0.0, ) - qk = tl.dot(q.to(k.dtype), k) if BLOCK_DPE > 0: offs_kpe = ( @@ -887,7 +886,6 @@ def _fwd_kernel_unified( other=0.0, ) - # Compute QK qk = tl.dot(q.to(k.dtype), k) if BLOCK_DPE > 0: offs_kpe = (