diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 50b854091..2e909d07c 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -185,8 +185,6 @@ class Envs: SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = EnvInt(10000) SGLANG_DISABLE_OUTLINES_DISK_CACHE = EnvBool(False) - # CuTe DSL GDN Decode - SGLANG_USE_CUTEDSL_GDN_DECODE = EnvBool(False) # Test & Debug SGLANG_DETECT_SLOW_RANK = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py index 25ffc3f0b..6eabb10b5 100644 --- a/python/sglang/srt/layers/attention/attention_registry.py +++ b/python/sglang/srt/layers/attention/attention_registry.py @@ -189,15 +189,21 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac if cfg := runner.mambaish_config: from sglang.srt.layers.attention.fla.utils import check_environments from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( - GDNAttnBackend, HybridLinearAttnBackend, - KimiLinearAttnBackend, - LightningAttentionBackend, Mamba2AttnBackend, ) + from sglang.srt.layers.attention.linear.gdn_backend import GDNAttnBackend + from sglang.srt.layers.attention.linear.kda_backend import KDAAttnBackend + from sglang.srt.layers.attention.linear.lightning_backend import ( + LightningAttentionBackend, + ) + from sglang.srt.layers.attention.linear.utils import ( + initialize_linear_attn_config, + ) from sglang.srt.utils import is_blackwell, is_npu check_environments() + initialize_linear_attn_config(runner.server_args) if runner.hybrid_gdn_config is not None: if is_blackwell(): assert ( @@ -213,7 +219,7 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac elif runner.mamba2_config is not None: linear_attn_backend = Mamba2AttnBackend(runner) elif runner.kimi_linear_config is not None: - linear_attn_backend = KimiLinearAttnBackend(runner) + linear_attn_backend = KDAAttnBackend(runner) elif runner.hybrid_lightning_config is not None: linear_attn_backend = LightningAttentionBackend(runner) else: diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index d73190595..d700f588b 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -1,32 +1,12 @@ import logging -import math from typing import Optional, Union import torch import triton import triton.language as tl -from einops import rearrange -from sglang.srt.environ import Envs from sglang.srt.layers.attention.base_attn_backend import AttentionBackend -from sglang.srt.layers.attention.fla.fused_gdn_gating import fused_gdn_gating -from sglang.srt.layers.attention.fla.fused_recurrent import ( - fused_recurrent_gated_delta_rule_update, -) -from sglang.srt.layers.attention.fla.fused_sigmoid_gating_recurrent import ( - fused_sigmoid_gating_delta_rule_update, -) -from sglang.srt.layers.attention.linear.lightning_attn import ( - BailingLinearKernel, - linear_decode_forward_triton, -) -from sglang.srt.layers.attention.linear.linear_metadata import BailingLinearMetadata -from sglang.srt.layers.attention.linear.seg_la import SegLaMeta, seg_la_fwd -from sglang.srt.layers.attention.mamba.causal_conv1d_triton import ( - PAD_SLOT_ID, - causal_conv1d_fn, - causal_conv1d_update, -) +from sglang.srt.layers.attention.mamba.causal_conv1d_triton import PAD_SLOT_ID from sglang.srt.layers.attention.mamba.mamba import MambaMixer2 from sglang.srt.layers.attention.mamba.mamba2_metadata import ( ForwardMetadata, @@ -36,69 +16,18 @@ from sglang.srt.layers.attention.mamba.mamba_state_scatter_triton import ( fused_mamba_state_scatter_with_mask, ) from sglang.srt.layers.radix_attention import RadixAttention -from sglang.srt.layers.radix_linear_attention import RadixLinearAttention -from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, MambaPool +from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput from sglang.srt.speculative.spec_info import SpecInput -from sglang.srt.utils import cpu_has_amx_support, is_cpu, is_cuda, is_npu -from sglang.srt.utils.common import rank0_log +from sglang.srt.utils import is_cpu if not is_cpu(): - # fix import error on CPU device, no impacts when non-CPU path - try: - from sglang.jit_kernel.cutedsl_gdn import ( - cutedsl_fused_sigmoid_gating_delta_rule_update, - ) - except ModuleNotFoundError: - # CuTe DSL path requires cuda-python (cuda.bindings.*). Keep runtime usable - # by falling back to non-CuTe kernels when it's unavailable. - cutedsl_fused_sigmoid_gating_delta_rule_update = None - from sglang.srt.layers.attention.fla.chunk import chunk_gated_delta_rule from sglang.srt.layers.attention.fla.chunk_delta_h import ( CHUNK_SIZE as FLA_CHUNK_SIZE, ) - from sglang.srt.layers.attention.fla.kda import chunk_kda - -if is_cuda(): - from sglang.srt.layers.attention.mamba.causal_conv1d import ( - causal_conv1d_fn as causal_conv1d_fn_cuda, - ) - - causal_conv1d_fn = causal_conv1d_fn_cuda -elif is_npu(): - from sgl_kernel_npu.fla.chunk import chunk_gated_delta_rule_npu - from sgl_kernel_npu.fla.fused_sigmoid_gating_recurrent import ( - fused_sigmoid_gating_delta_rule_update_npu, - ) - from sgl_kernel_npu.mamba.causal_conv1d import ( - causal_conv1d_fn_npu, - causal_conv1d_update_npu, - ) - - chunk_gated_delta_rule = chunk_gated_delta_rule_npu - fused_sigmoid_gating_delta_rule_update = fused_sigmoid_gating_delta_rule_update_npu - causal_conv1d_fn = causal_conv1d_fn_npu - causal_conv1d_update = causal_conv1d_update_npu -elif is_cpu(): - assert ( - cpu_has_amx_support() - ), "CPU requires AMX support for hybrid linear attn backend" - from sgl_kernel.mamba import ( - causal_conv1d_fn_cpu, - causal_conv1d_update_cpu, - chunk_gated_delta_rule_cpu, - ) - - chunk_gated_delta_rule = chunk_gated_delta_rule_cpu - causal_conv1d_fn = causal_conv1d_fn_cpu - causal_conv1d_update = causal_conv1d_update_cpu - fused_sigmoid_gating_delta_rule_update = ( - torch.ops.sgl_kernel.fused_sigmoid_gating_delta_rule_update_cpu - ) - fused_gdn_gating = torch.ops.sgl_kernel.fused_gdn_gating_cpu logger = logging.getLogger(__name__) @@ -661,410 +590,6 @@ class MambaAttnBackendBase(AttentionBackend): ] -class KimiLinearAttnBackend(MambaAttnBackendBase): - """Attention backend using Mamba kernel.""" - - def forward_decode( - self, - layer: RadixLinearAttention, - mixed_qkv: torch.Tensor, - a: torch.Tensor, - b: torch.Tensor, - **kwargs, - ): - q_proj_states, k_proj_states, v_proj_states = torch.split( - mixed_qkv, - [layer.q_dim, layer.k_dim, layer.v_dim], - dim=-1, - ) - - q_conv_weights, k_conv_weights, v_conv_weights = layer.conv_weights - q_conv_bias, k_conv_bias, v_conv_bias = layer.bias - - layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id) - q_conv_state, k_conv_state, v_conv_state = layer_cache.conv - ssm_states = layer_cache.temporal - query_start_loc = self.forward_metadata.query_start_loc - cache_indices = self.forward_metadata.mamba_cache_indices - - q_conv_state = q_conv_state.transpose(-1, -2) - k_conv_state = k_conv_state.transpose(-1, -2) - v_conv_state = v_conv_state.transpose(-1, -2) - - q = causal_conv1d_update( - q_proj_states, - q_conv_state, - q_conv_weights, - q_conv_bias, - activation="silu", - conv_state_indices=cache_indices, - ) - k = causal_conv1d_update( - k_proj_states, - k_conv_state, - k_conv_weights, - k_conv_bias, - activation="silu", - conv_state_indices=cache_indices, - ) - v = causal_conv1d_update( - v_proj_states, - v_conv_state, - v_conv_weights, - v_conv_bias, - activation="silu", - conv_state_indices=cache_indices, - ) - - q = rearrange(q, "n (h d) -> 1 n h d", d=layer.head_q_dim) - k = rearrange(k, "n (h d) -> 1 n h d", d=layer.head_k_dim) - v = rearrange(v, "n (h d) -> 1 n h d", d=layer.head_v_dim) - - return fused_sigmoid_gating_delta_rule_update( - A_log=layer.A_log, - dt_bias=layer.dt_bias, - q=q, - k=k, - v=v, - a=a, - b=b, - initial_state_source=ssm_states, - initial_state_indices=cache_indices, - cu_seqlens=query_start_loc, - use_qk_l2norm_in_kernel=True, - softplus_beta=1.0, - softplus_threshold=20.0, - is_kda=True, - ) - - def forward_extend( - self, - layer: RadixLinearAttention, - forward_batch: ForwardBatch, - mixed_qkv: torch.Tensor, - a: torch.Tensor, - b: torch.Tensor, - **kwargs, # Unused, for compatibility with HybridLinearAttnBackend - ): - from sglang.srt.layers.attention.mamba.causal_conv1d_triton import ( - causal_conv1d_fn, - ) - - q_proj_states, k_proj_states, v_proj_states = torch.split( - mixed_qkv, - [layer.q_dim, layer.k_dim, layer.v_dim], - dim=-1, - ) - - q_conv_weights, k_conv_weights, v_conv_weights = layer.conv_weights - q_conv_bias, k_conv_bias, v_conv_bias = layer.bias - - query_start_loc = self.forward_metadata.query_start_loc - cache_indices = self.forward_metadata.mamba_cache_indices - - mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id) - conv_state_q, conv_state_k, conv_state_v = mamba_cache_params.conv - # deal with strides - conv_state_q = conv_state_q.transpose(-1, -2) - conv_state_k = conv_state_k.transpose(-1, -2) - conv_state_v = conv_state_v.transpose(-1, -2) - - ssm_states = mamba_cache_params.temporal - - has_initial_state = forward_batch.extend_prefix_lens > 0 - - q_proj_states = q_proj_states.transpose(0, 1) - k_proj_states = k_proj_states.transpose(0, 1) - v_proj_states = v_proj_states.transpose(0, 1) - - q = causal_conv1d_fn( - q_proj_states, - q_conv_weights, - q_conv_bias, - activation="silu", - conv_states=conv_state_q, - has_initial_state=has_initial_state, - cache_indices=cache_indices, - query_start_loc=query_start_loc, - seq_lens_cpu=forward_batch.extend_seq_lens_cpu, - ).transpose(0, 1) - - k = causal_conv1d_fn( - k_proj_states, - k_conv_weights, - k_conv_bias, - activation="silu", - conv_states=conv_state_k, - has_initial_state=has_initial_state, - cache_indices=cache_indices, - query_start_loc=query_start_loc, - seq_lens_cpu=forward_batch.extend_seq_lens_cpu, - ).transpose(0, 1) - - v = causal_conv1d_fn( - v_proj_states, - v_conv_weights, - v_conv_bias, - activation="silu", - conv_states=conv_state_v, - has_initial_state=has_initial_state, - cache_indices=cache_indices, - query_start_loc=query_start_loc, - seq_lens_cpu=forward_batch.extend_seq_lens_cpu, - ).transpose(0, 1) - - q = rearrange(q, "n (h d) -> 1 n h d", d=layer.head_q_dim) - k = rearrange(k, "n (h d) -> 1 n h d", d=layer.head_k_dim) - v = rearrange(v, "n (h d) -> 1 n h d", d=layer.head_v_dim) - - core_attn_out = chunk_kda( - q=q, - k=k, - v=v, - g=a, - beta=b, - initial_state=ssm_states, - initial_state_indices=cache_indices, - use_qk_l2norm_in_kernel=True, - cu_seqlens=query_start_loc, - ) - - return core_attn_out - - -class GDNAttnBackend(MambaAttnBackendBase): - """Attention backend using Mamba kernel.""" - - def __init__(self, model_runner: ModelRunner): - super().__init__(model_runner) - self.conv_states_shape = ( - model_runner.req_to_token_pool.mamba_pool.mamba_cache.conv[0].shape - ) - if not is_cpu() and not is_npu(): - assert ( - self.conv_states_shape[-1] < FLA_CHUNK_SIZE - ), f"{self.conv_states_shape[-1]=} should be less than {FLA_CHUNK_SIZE}" - - use_cutedsl = Envs.SGLANG_USE_CUTEDSL_GDN_DECODE.get() - if use_cutedsl and cutedsl_fused_sigmoid_gating_delta_rule_update is None: - rank0_log( - "CuTe DSL GDN decode requested but unavailable " - "(missing cuda.bindings). Falling back to FLA decode kernel." - ) - use_cutedsl = False - rank0_log(f"CuTe DSL GDN decode enabled: {use_cutedsl}") - self._kernel_func = ( - cutedsl_fused_sigmoid_gating_delta_rule_update - if use_cutedsl - else fused_sigmoid_gating_delta_rule_update - ) - - def forward_decode( - self, - layer: RadixLinearAttention, - forward_batch: ForwardBatch, - mixed_qkv: torch.Tensor, - a: torch.Tensor, - b: torch.Tensor, - **kwargs, # Unused, for compatibility with HybridLinearAttnBackend - ): - layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id) - conv_states = layer_cache.conv[0] - ssm_states = layer_cache.temporal - query_start_loc = self.forward_metadata.query_start_loc - cache_indices = self.forward_metadata.mamba_cache_indices - mixed_qkv = causal_conv1d_update( - mixed_qkv, - conv_states, - layer.conv_weights, - layer.bias, - layer.activation, - conv_state_indices=cache_indices, - ) - - query, key, value = torch.split( - mixed_qkv, - [layer.q_dim, layer.k_dim, layer.v_dim], - dim=-1, - ) - # Reshape from [bs, h*d] to [1, bs, h, d] - bs = forward_batch.batch_size - query = query.view(1, bs, layer.num_q_heads, layer.head_q_dim) - key = key.view(1, bs, layer.num_k_heads, layer.head_k_dim) - value = value.view(1, bs, layer.num_v_heads, layer.head_v_dim) - - core_attn_out = self._kernel_func( - A_log=layer.A_log, - dt_bias=layer.dt_bias, - q=query, - k=key, - v=value, - a=a, - b=b, - initial_state_source=ssm_states, - initial_state_indices=cache_indices, - cu_seqlens=query_start_loc, - use_qk_l2norm_in_kernel=True, - softplus_beta=1.0, - softplus_threshold=20.0, - ) - - self._track_mamba_state_decode( - forward_batch, conv_states, ssm_states, cache_indices - ) - - return core_attn_out - - def forward_extend( - self, - layer: RadixLinearAttention, - forward_batch: ForwardBatch, - mixed_qkv: torch.Tensor, - a: torch.Tensor, - b: torch.Tensor, - **kwargs, # Unused, for compatibility with HybridLinearAttnBackend - ): - seq_len = mixed_qkv.shape[0] - - is_target_verify = forward_batch.forward_mode.is_target_verify() - forward_metadata = self.forward_metadata - - query_start_loc = forward_metadata.query_start_loc - cache_indices = forward_metadata.mamba_cache_indices - retrieve_next_token = forward_metadata.retrieve_next_token - retrieve_next_sibling = forward_metadata.retrieve_next_sibling - retrieve_parent_token = forward_metadata.retrieve_parent_token - - mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id) - conv_states = mamba_cache_params.conv[0] - ssm_states = mamba_cache_params.temporal - if is_target_verify: - assert isinstance(mamba_cache_params, MambaPool.SpeculativeState) - intermediate_state_cache = mamba_cache_params.intermediate_ssm - intermediate_conv_window_cache = ( - mamba_cache_params.intermediate_conv_window[0] - ) - has_initial_states = torch.ones( - seq_len // forward_batch.spec_info.draft_token_num, - dtype=torch.bool, - device=forward_batch.input_ids.device, - ) - intermediate_state_indices = torch.arange( - cache_indices.shape[0], dtype=torch.int32, device=cache_indices.device - ) - else: - has_initial_states = forward_batch.extend_prefix_lens > 0 - - if is_target_verify: - batch_size = seq_len // forward_batch.spec_info.draft_token_num - draft_token_num = forward_batch.spec_info.draft_token_num - mixed_qkv_reshaped = mixed_qkv.view( - batch_size, draft_token_num, -1 - ).transpose(1, 2) - mixed_qkv_processed = causal_conv1d_update( - mixed_qkv_reshaped, - conv_states, - layer.conv_weights, - layer.bias, - layer.activation, - conv_state_indices=cache_indices[:batch_size], - intermediate_conv_window=intermediate_conv_window_cache, - intermediate_state_indices=intermediate_state_indices[:batch_size], - retrieve_next_token=retrieve_next_token, - retrieve_next_sibling=retrieve_next_sibling, - retrieve_parent_token=retrieve_parent_token, - ) - mixed_qkv = mixed_qkv_processed.transpose(1, 2).view(seq_len, -1) - else: - mixed_qkv = mixed_qkv.transpose(0, 1) - if ( - forward_batch.mamba_track_mask is not None - and forward_batch.mamba_track_mask.any() - ): - conv_dst = forward_batch.mamba_track_indices - # Gather all slices at once: [:, track_conv_indices] -> [d, num_masked, slice_len] - # track_conv_indices is already filtered and clamped in _init_track_conv_indices - mixed_qkv_to_track = mixed_qkv[ - :, forward_metadata.track_conv_indices - ].transpose(0, 1) - # Apply mask and assign to destinations - mask_indices = forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0] - conv_states[conv_dst[mask_indices]] = mixed_qkv_to_track - - mixed_qkv = causal_conv1d_fn( - mixed_qkv, - layer.conv_weights, - layer.bias, - activation=layer.activation, - conv_states=conv_states, - has_initial_state=has_initial_states, - cache_indices=cache_indices, - query_start_loc=query_start_loc, - seq_lens_cpu=forward_batch.extend_seq_lens_cpu, - ).transpose(0, 1)[:seq_len] - - query, key, value = torch.split( - mixed_qkv, - [layer.q_dim, layer.k_dim, layer.v_dim], - dim=-1, - ) - - actual_seq_len = query.shape[0] - query = query.view(1, actual_seq_len, layer.num_q_heads, layer.head_q_dim) - key = key.view(1, actual_seq_len, layer.num_k_heads, layer.head_k_dim) - value = value.view(1, actual_seq_len, layer.num_v_heads, layer.head_v_dim) - - g, beta = fused_gdn_gating(layer.A_log, a, b, layer.dt_bias) - - if is_target_verify: - core_attn_out = fused_recurrent_gated_delta_rule_update( - q=query, - k=key, - v=value, - g=g, - beta=beta, - initial_state_source=ssm_states, - initial_state_indices=cache_indices, - cu_seqlens=query_start_loc, - use_qk_l2norm_in_kernel=True, - disable_state_update=True, - intermediate_states_buffer=intermediate_state_cache, - intermediate_state_indices=intermediate_state_indices, - cache_steps=forward_batch.spec_info.draft_token_num, - retrieve_parent_token=retrieve_parent_token, - ) - else: - # Only cuda env uses fuse ssm_states update - recurrent_state = ssm_states - recurrent_state_indices_args = {"initial_state_indices": cache_indices} - if is_npu() or is_cpu(): - recurrent_state = ssm_states[cache_indices] - recurrent_state_indices_args = {} - core_attn_out, last_recurrent_state, h = chunk_gated_delta_rule( - q=query, - k=key, - v=value, - g=g, - beta=beta, - initial_state=recurrent_state, - cu_seqlens=query_start_loc, - head_first=False, - use_qk_l2norm_in_kernel=True, - **recurrent_state_indices_args, - ) - if is_npu() or is_cpu(): - last_recurrent_state = last_recurrent_state.to( - ssm_states.dtype, copy=False - ) - ssm_states[cache_indices] = last_recurrent_state - - self._track_mamba_state_extend( - forward_batch, h, ssm_states, forward_metadata - ) - - return core_attn_out - - class Mamba2AttnBackend(MambaAttnBackendBase): """Attention backend wrapper for Mamba2Mixer kernels.""" @@ -1154,365 +679,6 @@ class Mamba2AttnBackend(MambaAttnBackendBase): ) -class LightningAttentionBackend(MambaAttnBackendBase): - """ - Note about the init: - - If no spec decoding - - FlashAttentionBackend will be init once when the server starts. - - If spec decoding - - FlashAttentionBackend will be init once for the target worker - - FlashAttentionMultiStepBackend will be once for the draft worker - - It will spawn num_steps FlashAttentionBackend for the draft worker - - Note about CUDA Graph: - - We only support CUDA Graph for Decode (Normal Decode and Draft Decode) and Target Verify. - - We don't support CUDA Graph for Extend and Draft Extend. - - When server init, init_cuda_graph_state will be called first and then init_cuda_graph_capture will be called. - - For each forward batch, init_replay_cuda_graph will be called first and then replay the graph. - """ - - def __init__(self, model_runner: ModelRunner): - super().__init__(model_runner) - - assert not ( - model_runner.sliding_window_size is not None - and model_runner.model_config.is_encoder_decoder - ), "Sliding window and cross attention are not supported together" - - # extra metadata for handling speculative decoding topk > 1, extended draft decode and verify - self.max_context_len = model_runner.model_config.context_len - self.device = model_runner.device - self.decode_cuda_graph_metadata = {} - self.kv_cache_dtype = model_runner.kv_cache_dtype - self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype - self.BLOCK = ( - model_runner.model_config.block - if hasattr(model_runner.model_config, "block") - else 256 - ) - total_num_heads = model_runner.model_config.hf_config.num_attention_heads - num_hidden_layers = model_runner.model_config.hf_config.num_hidden_layers - self.tp_slope = LightningAttentionBackend._build_slope_tensor( - total_num_heads, num_hidden_layers, self.device - ) - self.linear_backend = getattr( - model_runner.model_config.hf_config, "linear_backend", "seg_la" - ) - logger.info( - f"linear_backend for linear attention in hybrid_linear_backend: {self.linear_backend}" - ) - - def init_forward_metadata(self, forward_batch: ForwardBatch): - metadata = self._forward_metadata(forward_batch) - self.forward_metadata = BailingLinearMetadata.prepare_mixed( - metadata.query_start_loc, - metadata.mamba_cache_indices, - forward_batch, - ) - - def init_forward_metadata_capture_cuda_graph( - self, - bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - ): - metadata = self._capture_metadata(bs, req_pool_indices, forward_mode, spec_info) - self.forward_metadata = BailingLinearMetadata.prepare_decode( - metadata.query_start_loc, metadata.mamba_cache_indices, bs, seq_lens - ) - - def init_forward_metadata_replay_cuda_graph( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_sum: int, - encoder_lens: Optional[torch.Tensor], - forward_mode: ForwardMode, - spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], - seq_lens_cpu: Optional[torch.Tensor], - ): - metadata = self._replay_metadata( - bs, req_pool_indices, forward_mode, spec_info, seq_lens_cpu - ) - self.forward_metadata = BailingLinearMetadata.prepare_decode( - metadata.query_start_loc, metadata.mamba_cache_indices, bs, seq_lens - ) - - @staticmethod - def _build_slope_tensor( - n_attention_heads: int, num_hidden_layers: int, device="cuda" - ): - def get_slopes(n): - def get_slopes_power_of_2(n): - start = 2 ** (-(2 ** -(math.log2(n) - 3))) - ratio = start - return [start * ratio**i for i in range(n)] - - if math.log2(n).is_integer(): - return get_slopes_power_of_2(n) - else: - closest_power_of_2 = 2 ** math.floor(math.log2(n)) - return ( - get_slopes_power_of_2(closest_power_of_2) - + get_slopes(2 * closest_power_of_2)[0::2][: n - closest_power_of_2] - ) - - slopes = torch.tensor( - get_slopes(n_attention_heads), dtype=torch.float32 - ).reshape(n_attention_heads, 1, 1) - from sglang.srt.layers.dp_attention import ( - get_attention_tp_rank, - get_attention_tp_size, - ) - - tp_heads = n_attention_heads // get_attention_tp_size() - tp_rank = get_attention_tp_rank() - if num_hidden_layers <= 1: - slope_rate_list = [slopes * (1 + 1e-5)] - else: - slope_rate_list = [ - slopes * (1 - layer_id / (num_hidden_layers - 1) + 1e-5) - for layer_id in range(num_hidden_layers) - ] - - tp_slope = [ - slope_rate_list[layer_id][tp_rank * tp_heads : (tp_rank + 1) * tp_heads] - .contiguous() - .to(device) - for layer_id in range(num_hidden_layers) - ] - - return tp_slope - - def _prefill_and_mix_infer( - self, - q, - k, - v, - kv_cache, - state_indices_tensor, - forward_batch, - layer, - metadata, - ): - hidden = [] - for _prefill_idx in range(metadata.num_prefills): - if _prefill_idx >= forward_batch.extend_start_loc.shape[0]: - break - if _prefill_idx >= state_indices_tensor.shape[0]: - break - - _start = forward_batch.extend_start_loc[_prefill_idx] - - if _prefill_idx + 1 < forward_batch.extend_start_loc.shape[0]: - _end = forward_batch.extend_start_loc[_prefill_idx + 1] - else: - if ( - forward_batch.extend_seq_lens is not None - and _prefill_idx < forward_batch.extend_seq_lens.shape[0] - and metadata.num_decodes > 0 - ): - seq_len = forward_batch.extend_seq_lens[_prefill_idx] - _end = _start + seq_len - else: - _end = q.shape[0] - - slot_id = state_indices_tensor[_prefill_idx] - qs = q[_start:_end].transpose(0, 1).contiguous() - ks = k[_start:_end].transpose(0, 1).contiguous() - vs = v[_start:_end].transpose(0, 1).contiguous() - slice_layer_cache = kv_cache[slot_id, ...] - out_slice = BailingLinearKernel.jit_linear_forward_prefix( - qs, - ks, - vs, - slice_layer_cache, - self.tp_slope[layer.layer_id], - self.BLOCK, - layer_idx=layer.layer_id, - ) - hidden.append(out_slice.contiguous()) - if metadata.num_decodes > 0: - hidden.append( - self._decode_infer( - q, k, v, kv_cache, state_indices_tensor, metadata, layer - ) - ) - - if not hidden: - return torch.empty((0, q.size(-1)), device=q.device, dtype=q.dtype) - - hidden = torch.concat(hidden, dim=0).contiguous() - return hidden - - def _decode_infer(self, q, k, v, kv_cache, state_indices_tensor, metadata, layer): - num_prefill_tokens = metadata.num_prefill_tokens - num_prefills = metadata.num_prefills - q = q[num_prefill_tokens:].unsqueeze(2).contiguous() - k = k[num_prefill_tokens:].unsqueeze(2).contiguous() - v = v[num_prefill_tokens:].unsqueeze(2).contiguous() - slot_id = state_indices_tensor[num_prefills:] - - assert slot_id.shape[0] == q.shape[0], ( - f"slot_id length {slot_id.shape[0]} does not match decode batch size {q.shape[0]}. " - "This indicates a bug in the upstream logic that should be investigated." - ) - hidden = linear_decode_forward_triton( - q, k, v, kv_cache, self.tp_slope[layer.layer_id], slot_id, 32 - ) - return hidden - - def _linear_attention_entry( - self, - q, - k, - v, - kv_cache, - state_indices_tensor, - metadata, - layer, - mask=None, - temp_cache=None, - intermediate_state_indices=None, - ): - q_offsets = metadata.query_start_loc - - seg_meta = SegLaMeta( - batch_size=metadata.batch_size, - q_offsets=metadata.query_start_loc, - s_offsets=state_indices_tensor, - q_lengths=q_offsets.diff(), - s_scales=metadata.has_initial_states, - max_q_length=None, - mask=mask, - ) - hidden = seg_la_fwd( - q=q, - k=k, - v=v, - s=kv_cache, - decay_scales=self.tp_slope[layer.layer_id], - meta=seg_meta, - caches=temp_cache, - cache_indices=intermediate_state_indices, - decouple=True, - ) - return hidden - - def forward_extend( - self, - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - layer: RadixAttention, - forward_batch: ForwardBatch, - save_kv_cache=True, - **kwargs, - ): - q_rope = kwargs["q_rope"] if "q_rope" in kwargs else None - k_rope = kwargs["k_rope"] if "k_rope" in kwargs else None - layer_id = layer.layer_id if layer else kwargs["layer_id"] - - metadata = self.forward_metadata - - if self.kv_cache_dtype_str != "auto" and layer.k_scale is not None: - q = q.to(self.kv_cache_dtype) - - query_start_loc = self.forward_metadata.query_start_loc - cache_indices = self.forward_metadata.mamba_cache_indices - mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id) - ssm_states = mamba_cache_params.temporal - # logger.warning( - # f"---mix {layer.layer_id=}, {query_start_loc=}, {cache_indices=}, {ssm_states.shape=}" - # ) - if self.linear_backend == "minimax": - o = self._prefill_and_mix_infer( - q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim), - k, - v, - ssm_states, - cache_indices, - forward_batch, - layer, - metadata, - ) - elif self.linear_backend == "seg_la": - intermediate_state_indices = ( - torch.arange( - cache_indices.shape[0], - dtype=torch.int32, - device=cache_indices.device, - ) - if forward_batch.forward_mode.is_target_verify() - else None - ) - o = self._linear_attention_entry( - q, - k, - v, - ssm_states, - cache_indices, - metadata, - layer, - temp_cache=( - mamba_cache_params.intermediate_ssm - if forward_batch.forward_mode.is_target_verify() - else None - ), - intermediate_state_indices=intermediate_state_indices, - ) - else: - raise ValueError( - f"linear backend: {self.linear_backend} is not support for now" - ) - return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) - - def forward_decode( - self, - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - layer: RadixAttention, - forward_batch: ForwardBatch, - save_kv_cache=True, - **kwargs, - ) -> torch.Tensor: - q_rope = kwargs["q_rope"] if "q_rope" in kwargs else None - k_rope = kwargs["k_rope"] if "k_rope" in kwargs else None - layer_id = layer.layer_id if layer else kwargs["layer_id"] - - # Use precomputed metadata across all layers - metadata = self.forward_metadata - - if self.kv_cache_dtype_str != "auto": - q = q.to(self.kv_cache_dtype) - - # Do linear attention - query_start_loc = self.forward_metadata.query_start_loc - cache_indices = self.forward_metadata.mamba_cache_indices - mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id) - ssm_states = mamba_cache_params.temporal - # logger.warning( - # f"---mix {layer.layer_id=}, {query_start_loc.shape=}, {cache_indices.shape=}, {ssm_states.shape=}" - # ) - if self.linear_backend == "minimax": - o = self._decode_infer(q, k, v, ssm_states, cache_indices, metadata, layer) - elif self.linear_backend == "seg_la": - o = self._linear_attention_entry( - q, k, v, ssm_states, cache_indices, metadata, layer - ) - else: - raise ValueError( - f"linear backend: {self.linear_backend} is not support for now" - ) - return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) - - class HybridLinearAttnBackend(AttentionBackend): """Manages a full and linear attention backend""" diff --git a/python/sglang/srt/layers/attention/linear/__init__.py b/python/sglang/srt/layers/attention/linear/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py new file mode 100644 index 000000000..4706f0153 --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py @@ -0,0 +1,379 @@ +from typing import Tuple, Union + +import torch + +from sglang.srt.layers.attention.fla.fused_gdn_gating import fused_gdn_gating +from sglang.srt.layers.attention.hybrid_linear_attn_backend import MambaAttnBackendBase +from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel +from sglang.srt.layers.attention.linear.utils import ( + LinearAttnKernelBackend, + get_linear_attn_decode_backend, + get_linear_attn_prefill_backend, +) +from sglang.srt.layers.attention.mamba.causal_conv1d_triton import ( + causal_conv1d_fn, + causal_conv1d_update, +) +from sglang.srt.layers.radix_linear_attention import RadixLinearAttention +from sglang.srt.mem_cache.memory_pool import MambaPool +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.model_executor.model_runner import ModelRunner +from sglang.srt.utils import is_cpu, is_cuda, is_npu +from sglang.srt.utils.common import rank0_log + +if not is_cpu(): + from sglang.srt.layers.attention.fla.chunk_delta_h import ( + CHUNK_SIZE as FLA_CHUNK_SIZE, + ) + +if is_cuda(): + from sglang.srt.layers.attention.mamba.causal_conv1d import ( + causal_conv1d_fn as causal_conv1d_fn_cuda, + ) + + causal_conv1d_fn = causal_conv1d_fn_cuda +elif is_npu(): + from sgl_kernel_npu.mamba.causal_conv1d import ( + causal_conv1d_fn_npu, + causal_conv1d_update_npu, + ) + + causal_conv1d_fn = causal_conv1d_fn_npu + causal_conv1d_update = causal_conv1d_update_npu +elif is_cpu(): + from sgl_kernel.mamba import causal_conv1d_fn_cpu, causal_conv1d_update_cpu + + causal_conv1d_fn = causal_conv1d_fn_cpu + causal_conv1d_update = causal_conv1d_update_cpu + fused_gdn_gating = torch.ops.sgl_kernel.fused_gdn_gating_cpu + + +class GDNKernelDispatcher: + """Dispatches GDN kernel calls to the appropriate backend per mode.""" + + def __init__( + self, + decode_backend: LinearAttnKernelBackend, + prefill_backend: LinearAttnKernelBackend, + ): + triton_kernel = TritonGDNKernel() + + if decode_backend.is_triton(): + self.decode_kernel = triton_kernel + elif decode_backend.is_cutedsl(): + if not is_cuda(): + raise ValueError("CuTe DSL backend requires CUDA") + from sglang.srt.layers.attention.linear.kernels.gdn_cutedsl import ( + CuteDSLGDNKernel, + ) + + self.decode_kernel = CuteDSLGDNKernel() + else: + raise ValueError(f"Unsupported GDN decode backend: {decode_backend}") + + if prefill_backend.is_triton(): + self.extend_kernel = triton_kernel + elif prefill_backend.is_cutedsl(): + raise ValueError( + "CuTe DSL backend only supports decode, not prefill. " + "Use --linear-attn-prefill-backend triton instead." + ) + else: + raise ValueError(f"Unsupported GDN prefill backend: {prefill_backend}") + + self.verify_kernel = triton_kernel + + rank0_log( + f"GDN kernel dispatcher: decode={self.decode_kernel.__class__.__name__}, " + f"extend={self.extend_kernel.__class__.__name__}, " + f"verify={self.verify_kernel.__class__.__name__}" + ) + + def decode( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + *, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + return self.decode_kernel.decode( + q, + k, + v, + a, + b, + A_log=A_log, + dt_bias=dt_bias, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + **kwargs, + ) + + def extend( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + *, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> tuple: + return self.extend_kernel.extend( + q, + k, + v, + g, + beta, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + **kwargs, + ) + + def target_verify( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + *, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + return self.verify_kernel.target_verify( + q, + k, + v, + g, + beta, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + **kwargs, + ) + + +class GDNAttnBackend(MambaAttnBackendBase): + """Attention backend for GDN (Gated Delta Network) linear attention.""" + + def __init__(self, model_runner: ModelRunner): + super().__init__(model_runner) + self.conv_states_shape = ( + model_runner.req_to_token_pool.mamba_pool.mamba_cache.conv[0].shape + ) + if not is_cpu() and not is_npu(): + assert ( + self.conv_states_shape[-1] < FLA_CHUNK_SIZE + ), f"{self.conv_states_shape[-1]=} should be less than {FLA_CHUNK_SIZE}" + + decode_backend = get_linear_attn_decode_backend() + prefill_backend = get_linear_attn_prefill_backend() + self.kernel_dispatcher = GDNKernelDispatcher(decode_backend, prefill_backend) + + def forward_decode( + self, + layer: RadixLinearAttention, + forward_batch: ForwardBatch, + mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]], + a: torch.Tensor, + b: torch.Tensor, + **kwargs, + ): + layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id) + conv_states = layer_cache.conv[0] + ssm_states = layer_cache.temporal + query_start_loc = self.forward_metadata.query_start_loc + cache_indices = self.forward_metadata.mamba_cache_indices + + assert isinstance(mixed_qkv, torch.Tensor) + mixed_qkv = causal_conv1d_update( + mixed_qkv, + conv_states, + layer.conv_weights, + layer.bias, + layer.activation, + conv_state_indices=cache_indices, + ) + + query, key, value = torch.split( + mixed_qkv, + [layer.q_dim, layer.k_dim, layer.v_dim], + dim=-1, + ) + # Reshape from [bs, h*d] to [1, bs, h, d] + bs = forward_batch.batch_size + query = query.view(1, bs, layer.num_q_heads, layer.head_q_dim) + key = key.view(1, bs, layer.num_k_heads, layer.head_k_dim) + value = value.view(1, bs, layer.num_v_heads, layer.head_v_dim) + + core_attn_out = self.kernel_dispatcher.decode( + q=query, + k=key, + v=value, + a=a, + b=b, + A_log=layer.A_log, + dt_bias=layer.dt_bias, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + ) + + self._track_mamba_state_decode( + forward_batch, conv_states, ssm_states, cache_indices + ) + + return core_attn_out + + def forward_extend( + self, + layer: RadixLinearAttention, + forward_batch: ForwardBatch, + mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]], + a: torch.Tensor, + b: torch.Tensor, + **kwargs, + ): + assert isinstance(mixed_qkv, torch.Tensor) + seq_len = mixed_qkv.shape[0] + + is_target_verify = forward_batch.forward_mode.is_target_verify() + forward_metadata = self.forward_metadata + + query_start_loc = forward_metadata.query_start_loc + cache_indices = forward_metadata.mamba_cache_indices + retrieve_next_token = forward_metadata.retrieve_next_token + retrieve_next_sibling = forward_metadata.retrieve_next_sibling + retrieve_parent_token = forward_metadata.retrieve_parent_token + + mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id) + conv_states = mamba_cache_params.conv[0] + ssm_states = mamba_cache_params.temporal + if is_target_verify: + assert isinstance(mamba_cache_params, MambaPool.SpeculativeState) + intermediate_state_cache = mamba_cache_params.intermediate_ssm + intermediate_conv_window_cache = ( + mamba_cache_params.intermediate_conv_window[0] + ) + has_initial_states = torch.ones( + seq_len // forward_batch.spec_info.draft_token_num, + dtype=torch.bool, + device=forward_batch.input_ids.device, + ) + intermediate_state_indices = torch.arange( + cache_indices.shape[0], dtype=torch.int32, device=cache_indices.device + ) + else: + has_initial_states = forward_batch.extend_prefix_lens > 0 + + if is_target_verify: + batch_size = seq_len // forward_batch.spec_info.draft_token_num + draft_token_num = forward_batch.spec_info.draft_token_num + mixed_qkv_reshaped = mixed_qkv.view( + batch_size, draft_token_num, -1 + ).transpose(1, 2) + mixed_qkv_processed = causal_conv1d_update( + mixed_qkv_reshaped, + conv_states, + layer.conv_weights, + layer.bias, + layer.activation, + conv_state_indices=cache_indices[:batch_size], + intermediate_conv_window=intermediate_conv_window_cache, + intermediate_state_indices=intermediate_state_indices[:batch_size], + retrieve_next_token=retrieve_next_token, + retrieve_next_sibling=retrieve_next_sibling, + retrieve_parent_token=retrieve_parent_token, + ) + mixed_qkv = mixed_qkv_processed.transpose(1, 2).view(seq_len, -1) + else: + mixed_qkv = mixed_qkv.transpose(0, 1) + if ( + forward_batch.mamba_track_mask is not None + and forward_batch.mamba_track_mask.any() + ): + conv_dst = forward_batch.mamba_track_indices + mixed_qkv_to_track = mixed_qkv[ + :, forward_metadata.track_conv_indices + ].transpose(0, 1) + mask_indices = forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0] + conv_states[conv_dst[mask_indices]] = mixed_qkv_to_track + + mixed_qkv = causal_conv1d_fn( + mixed_qkv, + layer.conv_weights, + layer.bias, + activation=layer.activation, + conv_states=conv_states, + has_initial_state=has_initial_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + seq_lens_cpu=forward_batch.extend_seq_lens_cpu, + ).transpose(0, 1)[:seq_len] + + query, key, value = torch.split( + mixed_qkv, + [layer.q_dim, layer.k_dim, layer.v_dim], + dim=-1, + ) + + actual_seq_len = query.shape[0] + query = query.view(1, actual_seq_len, layer.num_q_heads, layer.head_q_dim) + key = key.view(1, actual_seq_len, layer.num_k_heads, layer.head_k_dim) + value = value.view(1, actual_seq_len, layer.num_v_heads, layer.head_v_dim) + + g, beta = fused_gdn_gating(layer.A_log, a, b, layer.dt_bias) + + if is_target_verify: + core_attn_out = self.kernel_dispatcher.target_verify( + q=query, + k=key, + v=value, + g=g, + beta=beta, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + intermediate_states_buffer=intermediate_state_cache, + intermediate_state_indices=intermediate_state_indices, + cache_steps=forward_batch.spec_info.draft_token_num, + retrieve_parent_token=retrieve_parent_token, + ) + else: + core_attn_out, last_recurrent_state, h = self.kernel_dispatcher.extend( + q=query, + k=key, + v=value, + g=g, + beta=beta, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + ) + if is_npu() or is_cpu(): + last_recurrent_state = last_recurrent_state.to( + ssm_states.dtype, copy=False + ) + ssm_states[cache_indices] = last_recurrent_state + + self._track_mamba_state_extend( + forward_batch, h, ssm_states, forward_metadata + ) + + return core_attn_out diff --git a/python/sglang/srt/layers/attention/linear/kda_backend.py b/python/sglang/srt/layers/attention/linear/kda_backend.py new file mode 100644 index 000000000..d005262b3 --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kda_backend.py @@ -0,0 +1,285 @@ +from typing import Tuple, Union + +import torch +from einops import rearrange + +from sglang.srt.layers.attention.hybrid_linear_attn_backend import MambaAttnBackendBase +from sglang.srt.layers.attention.linear.kernels.kda_triton import TritonKDAKernel +from sglang.srt.layers.attention.linear.utils import ( + LinearAttnKernelBackend, + get_linear_attn_decode_backend, + get_linear_attn_prefill_backend, +) +from sglang.srt.layers.attention.mamba.causal_conv1d_triton import ( + causal_conv1d_fn, + causal_conv1d_update, +) +from sglang.srt.layers.radix_linear_attention import RadixLinearAttention +from sglang.srt.utils import is_cpu, is_npu +from sglang.srt.utils.common import rank0_log + +# KDA always uses the triton causal_conv1d_fn (no CUDA override). +# Only causal_conv1d_update needs platform-specific overrides for decode. +if is_npu(): + from sgl_kernel_npu.mamba.causal_conv1d import causal_conv1d_update_npu + + causal_conv1d_update = causal_conv1d_update_npu +elif is_cpu(): + from sgl_kernel.mamba import causal_conv1d_update_cpu + + causal_conv1d_update = causal_conv1d_update_cpu + +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.model_executor.model_runner import ModelRunner + + +class KDAKernelDispatcher: + """Dispatches KDA kernel calls to the appropriate backend per mode.""" + + def __init__( + self, + decode_backend: LinearAttnKernelBackend, + prefill_backend: LinearAttnKernelBackend, + ): + triton_kernel = TritonKDAKernel() + + if decode_backend.is_triton(): + self.decode_kernel = triton_kernel + else: + raise ValueError( + f"Unsupported KDA decode backend: {decode_backend}. " + "KDA currently only supports 'triton'." + ) + + if prefill_backend.is_triton(): + self.extend_kernel = triton_kernel + else: + raise ValueError( + f"Unsupported KDA prefill backend: {prefill_backend}. " + "KDA currently only supports 'triton'." + ) + + rank0_log( + f"KDA kernel dispatcher: decode={self.decode_kernel.__class__.__name__}, " + f"extend={self.extend_kernel.__class__.__name__}" + ) + + def decode( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + *, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + return self.decode_kernel.decode( + q, + k, + v, + a, + b, + A_log=A_log, + dt_bias=dt_bias, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + **kwargs, + ) + + def extend( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + *, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + return self.extend_kernel.extend( + q, + k, + v, + g, + beta, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + **kwargs, + ) + + +class KDAAttnBackend(MambaAttnBackendBase): + """Attention backend for KDA (Kimi Delta Attention) linear attention.""" + + def __init__(self, model_runner: ModelRunner): + super().__init__(model_runner) + decode_backend = get_linear_attn_decode_backend() + prefill_backend = get_linear_attn_prefill_backend() + self.kernel_dispatcher = KDAKernelDispatcher(decode_backend, prefill_backend) + + def forward_decode( + self, + layer: RadixLinearAttention, + mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]], + a: torch.Tensor, + b: torch.Tensor, + **kwargs, + ): + q_proj_states, k_proj_states, v_proj_states = torch.split( + mixed_qkv, + [layer.q_dim, layer.k_dim, layer.v_dim], + dim=-1, + ) + q_conv_weights, k_conv_weights, v_conv_weights = layer.conv_weights + q_conv_bias, k_conv_bias, v_conv_bias = layer.bias + + layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id) + q_conv_state, k_conv_state, v_conv_state = layer_cache.conv + ssm_states = layer_cache.temporal + query_start_loc = self.forward_metadata.query_start_loc + cache_indices = self.forward_metadata.mamba_cache_indices + + q_conv_state = q_conv_state.transpose(-1, -2) + k_conv_state = k_conv_state.transpose(-1, -2) + v_conv_state = v_conv_state.transpose(-1, -2) + + q = causal_conv1d_update( + q_proj_states, + q_conv_state, + q_conv_weights, + q_conv_bias, + activation="silu", + conv_state_indices=cache_indices, + ) + k = causal_conv1d_update( + k_proj_states, + k_conv_state, + k_conv_weights, + k_conv_bias, + activation="silu", + conv_state_indices=cache_indices, + ) + v = causal_conv1d_update( + v_proj_states, + v_conv_state, + v_conv_weights, + v_conv_bias, + activation="silu", + conv_state_indices=cache_indices, + ) + + q = rearrange(q, "n (h d) -> 1 n h d", d=layer.head_q_dim) + k = rearrange(k, "n (h d) -> 1 n h d", d=layer.head_k_dim) + v = rearrange(v, "n (h d) -> 1 n h d", d=layer.head_v_dim) + + return self.kernel_dispatcher.decode( + q=q, + k=k, + v=v, + a=a, + b=b, + A_log=layer.A_log, + dt_bias=layer.dt_bias, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + ) + + def forward_extend( + self, + layer: RadixLinearAttention, + forward_batch: ForwardBatch, + mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]], + a: torch.Tensor, + b: torch.Tensor, + **kwargs, + ): + q_proj_states, k_proj_states, v_proj_states = torch.split( + mixed_qkv, + [layer.q_dim, layer.k_dim, layer.v_dim], + dim=-1, + ) + q_conv_weights, k_conv_weights, v_conv_weights = layer.conv_weights + q_conv_bias, k_conv_bias, v_conv_bias = layer.bias + + query_start_loc = self.forward_metadata.query_start_loc + cache_indices = self.forward_metadata.mamba_cache_indices + + mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id) + conv_state_q, conv_state_k, conv_state_v = mamba_cache_params.conv + # deal with strides + conv_state_q = conv_state_q.transpose(-1, -2) + conv_state_k = conv_state_k.transpose(-1, -2) + conv_state_v = conv_state_v.transpose(-1, -2) + + ssm_states = mamba_cache_params.temporal + + has_initial_state = forward_batch.extend_prefix_lens > 0 + + q_proj_states = q_proj_states.transpose(0, 1) + k_proj_states = k_proj_states.transpose(0, 1) + v_proj_states = v_proj_states.transpose(0, 1) + + q = causal_conv1d_fn( + q_proj_states, + q_conv_weights, + q_conv_bias, + activation="silu", + conv_states=conv_state_q, + has_initial_state=has_initial_state, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + seq_lens_cpu=forward_batch.extend_seq_lens_cpu, + ).transpose(0, 1) + + k = causal_conv1d_fn( + k_proj_states, + k_conv_weights, + k_conv_bias, + activation="silu", + conv_states=conv_state_k, + has_initial_state=has_initial_state, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + seq_lens_cpu=forward_batch.extend_seq_lens_cpu, + ).transpose(0, 1) + + v = causal_conv1d_fn( + v_proj_states, + v_conv_weights, + v_conv_bias, + activation="silu", + conv_states=conv_state_v, + has_initial_state=has_initial_state, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + seq_lens_cpu=forward_batch.extend_seq_lens_cpu, + ).transpose(0, 1) + + q = rearrange(q, "n (h d) -> 1 n h d", d=layer.head_q_dim) + k = rearrange(k, "n (h d) -> 1 n h d", d=layer.head_k_dim) + v = rearrange(v, "n (h d) -> 1 n h d", d=layer.head_v_dim) + + core_attn_out = self.kernel_dispatcher.extend( + q=q, + k=k, + v=v, + g=a, + beta=b, + ssm_states=ssm_states, + cache_indices=cache_indices, + query_start_loc=query_start_loc, + ) + + return core_attn_out diff --git a/python/sglang/srt/layers/attention/linear/kernels/__init__.py b/python/sglang/srt/layers/attention/linear/kernels/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/layers/attention/linear/kernels/gdn_cutedsl.py b/python/sglang/srt/layers/attention/linear/kernels/gdn_cutedsl.py new file mode 100644 index 000000000..fff4ef901 --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kernels/gdn_cutedsl.py @@ -0,0 +1,47 @@ +import torch + +from sglang.jit_kernel.cutedsl_gdn import cutedsl_fused_sigmoid_gating_delta_rule_update +from sglang.srt.layers.attention.linear.kernels.kernel_backend import ( + LinearAttnKernelBase, +) + + +class CuteDSLGDNKernel(LinearAttnKernelBase): + """CuTe DSL kernel for GDN decode (CUDA only).""" + + def decode( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + *, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + return cutedsl_fused_sigmoid_gating_delta_rule_update( + A_log=A_log, + dt_bias=dt_bias, + q=q, + k=k, + v=v, + a=a, + b=b, + initial_state_source=ssm_states, + initial_state_indices=cache_indices, + cu_seqlens=query_start_loc, + use_qk_l2norm_in_kernel=True, + softplus_beta=1.0, + softplus_threshold=20.0, + ) + + def extend(self, *args, **kwargs): + raise NotImplementedError("CuteDSLGDNKernel only supports decode") + + def target_verify(self, *args, **kwargs): + raise NotImplementedError("CuteDSLGDNKernel only supports decode") diff --git a/python/sglang/srt/layers/attention/linear/kernels/gdn_triton.py b/python/sglang/srt/layers/attention/linear/kernels/gdn_triton.py new file mode 100644 index 000000000..a1cf16bba --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kernels/gdn_triton.py @@ -0,0 +1,131 @@ +import torch + +from sglang.srt.layers.attention.linear.kernels.kernel_backend import ( + LinearAttnKernelBase, +) +from sglang.srt.utils import is_cpu, is_npu + +if not is_cpu(): + from sglang.srt.layers.attention.fla.chunk import chunk_gated_delta_rule + from sglang.srt.layers.attention.fla.fused_recurrent import ( + fused_recurrent_gated_delta_rule_update, + ) + from sglang.srt.layers.attention.fla.fused_sigmoid_gating_recurrent import ( + fused_sigmoid_gating_delta_rule_update, + ) + +if is_npu(): + from sgl_kernel_npu.fla.chunk import chunk_gated_delta_rule_npu + from sgl_kernel_npu.fla.fused_sigmoid_gating_recurrent import ( + fused_sigmoid_gating_delta_rule_update_npu, + ) + + chunk_gated_delta_rule = chunk_gated_delta_rule_npu + fused_sigmoid_gating_delta_rule_update = fused_sigmoid_gating_delta_rule_update_npu +elif is_cpu(): + from sgl_kernel.mamba import chunk_gated_delta_rule_cpu + + chunk_gated_delta_rule = chunk_gated_delta_rule_cpu + fused_sigmoid_gating_delta_rule_update = ( + torch.ops.sgl_kernel.fused_sigmoid_gating_delta_rule_update_cpu + ) + + +class TritonGDNKernel(LinearAttnKernelBase): + """Triton-based kernel for GDN (Gated Delta Network) linear attention.""" + + def decode( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + *, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + return fused_sigmoid_gating_delta_rule_update( + A_log=A_log, + dt_bias=dt_bias, + q=q, + k=k, + v=v, + a=a, + b=b, + initial_state_source=ssm_states, + initial_state_indices=cache_indices, + cu_seqlens=query_start_loc, + use_qk_l2norm_in_kernel=True, + softplus_beta=1.0, + softplus_threshold=20.0, + ) + + def extend( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + *, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> tuple: + recurrent_state = ssm_states + recurrent_state_indices_args = {"initial_state_indices": cache_indices} + if is_npu() or is_cpu(): + recurrent_state = ssm_states[cache_indices] + recurrent_state_indices_args = {} + return chunk_gated_delta_rule( + q=q, + k=k, + v=v, + g=g, + beta=beta, + initial_state=recurrent_state, + cu_seqlens=query_start_loc, + head_first=False, + use_qk_l2norm_in_kernel=True, + **recurrent_state_indices_args, + ) + + def target_verify( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + *, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + intermediate_states_buffer: torch.Tensor, + intermediate_state_indices: torch.Tensor, + cache_steps: int, + retrieve_parent_token: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + return fused_recurrent_gated_delta_rule_update( + q=q, + k=k, + v=v, + g=g, + beta=beta, + initial_state_source=ssm_states, + initial_state_indices=cache_indices, + cu_seqlens=query_start_loc, + use_qk_l2norm_in_kernel=True, + disable_state_update=True, + intermediate_states_buffer=intermediate_states_buffer, + intermediate_state_indices=intermediate_state_indices, + cache_steps=cache_steps, + retrieve_parent_token=retrieve_parent_token, + ) diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_triton.py b/python/sglang/srt/layers/attention/linear/kernels/kda_triton.py new file mode 100644 index 000000000..2ab8a5e0b --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_triton.py @@ -0,0 +1,73 @@ +import torch + +from sglang.srt.layers.attention.linear.kernels.kernel_backend import ( + LinearAttnKernelBase, +) +from sglang.srt.utils import is_cpu + +if not is_cpu(): + from sglang.srt.layers.attention.fla.fused_sigmoid_gating_recurrent import ( + fused_sigmoid_gating_delta_rule_update, + ) + from sglang.srt.layers.attention.fla.kda import chunk_kda + + +class TritonKDAKernel(LinearAttnKernelBase): + """Triton-based kernel for KDA (Kimi Delta Attention) linear attention.""" + + def decode( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + *, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + return fused_sigmoid_gating_delta_rule_update( + A_log=A_log, + dt_bias=dt_bias, + q=q, + k=k, + v=v, + a=a, + b=b, + initial_state_source=ssm_states, + initial_state_indices=cache_indices, + cu_seqlens=query_start_loc, + use_qk_l2norm_in_kernel=True, + softplus_beta=1.0, + softplus_threshold=20.0, + is_kda=True, + ) + + def extend( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + *, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + return chunk_kda( + q=q, + k=k, + v=v, + g=g, + beta=beta, + initial_state=ssm_states, + initial_state_indices=cache_indices, + use_qk_l2norm_in_kernel=True, + cu_seqlens=query_start_loc, + ) diff --git a/python/sglang/srt/layers/attention/linear/kernels/kernel_backend.py b/python/sglang/srt/layers/attention/linear/kernels/kernel_backend.py new file mode 100644 index 000000000..83a7fe0e7 --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kernels/kernel_backend.py @@ -0,0 +1,60 @@ +from abc import ABC, abstractmethod + +import torch + + +class LinearAttnKernelBase(ABC): + """Abstract base class for linear attention kernel implementations. + + Each concrete implementation wraps a specific kernel (Triton, CuTe DSL, etc.) + and provides decode/extend/target_verify methods with a unified interface. + """ + + @abstractmethod + def decode( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + *, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: ... + + @abstractmethod + def extend( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + *, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> tuple: ... + + def target_verify( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + *, + ssm_states: torch.Tensor, + cache_indices: torch.Tensor, + query_start_loc: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + raise NotImplementedError( + f"{self.__class__.__name__} does not support target_verify" + ) diff --git a/python/sglang/srt/layers/attention/linear/lightning_backend.py b/python/sglang/srt/layers/attention/linear/lightning_backend.py new file mode 100644 index 000000000..b34fefbfd --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/lightning_backend.py @@ -0,0 +1,372 @@ +import logging +import math +from typing import Optional, Union + +import torch + +from sglang.srt.layers.attention.hybrid_linear_attn_backend import MambaAttnBackendBase +from sglang.srt.layers.attention.linear.lightning_attn import ( + BailingLinearKernel, + linear_decode_forward_triton, +) +from sglang.srt.layers.attention.linear.linear_metadata import BailingLinearMetadata +from sglang.srt.layers.attention.linear.seg_la import SegLaMeta, seg_la_fwd +from sglang.srt.layers.radix_attention import RadixAttention +from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.srt.model_executor.model_runner import ModelRunner +from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput + +logger = logging.getLogger(__name__) + + +class LightningAttentionBackend(MambaAttnBackendBase): + """ + Note about the init: + - If no spec decoding + - FlashAttentionBackend will be init once when the server starts. + - If spec decoding + - FlashAttentionBackend will be init once for the target worker + - FlashAttentionMultiStepBackend will be once for the draft worker + - It will spawn num_steps FlashAttentionBackend for the draft worker + + Note about CUDA Graph: + - We only support CUDA Graph for Decode (Normal Decode and Draft Decode) and Target Verify. + - We don't support CUDA Graph for Extend and Draft Extend. + - When server init, init_cuda_graph_state will be called first and then init_cuda_graph_capture will be called. + - For each forward batch, init_replay_cuda_graph will be called first and then replay the graph. + """ + + def __init__(self, model_runner: ModelRunner): + super().__init__(model_runner) + + assert not ( + model_runner.sliding_window_size is not None + and model_runner.model_config.is_encoder_decoder + ), "Sliding window and cross attention are not supported together" + + # extra metadata for handling speculative decoding topk > 1, extended draft decode and verify + self.max_context_len = model_runner.model_config.context_len + self.device = model_runner.device + self.decode_cuda_graph_metadata = {} + self.kv_cache_dtype = model_runner.kv_cache_dtype + self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype + self.BLOCK = ( + model_runner.model_config.block + if hasattr(model_runner.model_config, "block") + else 256 + ) + total_num_heads = model_runner.model_config.hf_config.num_attention_heads + num_hidden_layers = model_runner.model_config.hf_config.num_hidden_layers + self.tp_slope = LightningAttentionBackend._build_slope_tensor( + total_num_heads, num_hidden_layers, self.device + ) + self.linear_backend = getattr( + model_runner.model_config.hf_config, "linear_backend", "seg_la" + ) + logger.info( + f"linear_backend for linear attention in hybrid_linear_backend: {self.linear_backend}" + ) + + def init_forward_metadata(self, forward_batch: ForwardBatch): + metadata = self._forward_metadata(forward_batch) + self.forward_metadata = BailingLinearMetadata.prepare_mixed( + metadata.query_start_loc, + metadata.mamba_cache_indices, + forward_batch, + ) + + def init_forward_metadata_capture_cuda_graph( + self, + bs: int, + num_tokens: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + encoder_lens: Optional[torch.Tensor], + forward_mode: ForwardMode, + spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], + ): + metadata = self._capture_metadata(bs, req_pool_indices, forward_mode, spec_info) + self.forward_metadata = BailingLinearMetadata.prepare_decode( + metadata.query_start_loc, metadata.mamba_cache_indices, bs, seq_lens + ) + + def init_forward_metadata_replay_cuda_graph( + self, + bs: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + seq_lens_sum: int, + encoder_lens: Optional[torch.Tensor], + forward_mode: ForwardMode, + spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], + seq_lens_cpu: Optional[torch.Tensor], + ): + metadata = self._replay_metadata( + bs, req_pool_indices, forward_mode, spec_info, seq_lens_cpu + ) + self.forward_metadata = BailingLinearMetadata.prepare_decode( + metadata.query_start_loc, metadata.mamba_cache_indices, bs, seq_lens + ) + + @staticmethod + def _build_slope_tensor( + n_attention_heads: int, num_hidden_layers: int, device="cuda" + ): + def get_slopes(n): + def get_slopes_power_of_2(n): + start = 2 ** (-(2 ** -(math.log2(n) - 3))) + ratio = start + return [start * ratio**i for i in range(n)] + + if math.log2(n).is_integer(): + return get_slopes_power_of_2(n) + else: + closest_power_of_2 = 2 ** math.floor(math.log2(n)) + return ( + get_slopes_power_of_2(closest_power_of_2) + + get_slopes(2 * closest_power_of_2)[0::2][: n - closest_power_of_2] + ) + + slopes = torch.tensor( + get_slopes(n_attention_heads), dtype=torch.float32 + ).reshape(n_attention_heads, 1, 1) + from sglang.srt.layers.dp_attention import ( + get_attention_tp_rank, + get_attention_tp_size, + ) + + tp_heads = n_attention_heads // get_attention_tp_size() + tp_rank = get_attention_tp_rank() + if num_hidden_layers <= 1: + slope_rate_list = [slopes * (1 + 1e-5)] + else: + slope_rate_list = [ + slopes * (1 - layer_id / (num_hidden_layers - 1) + 1e-5) + for layer_id in range(num_hidden_layers) + ] + + tp_slope = [ + slope_rate_list[layer_id][tp_rank * tp_heads : (tp_rank + 1) * tp_heads] + .contiguous() + .to(device) + for layer_id in range(num_hidden_layers) + ] + + return tp_slope + + def _prefill_and_mix_infer( + self, + q, + k, + v, + kv_cache, + state_indices_tensor, + forward_batch, + layer, + metadata, + ): + hidden = [] + for _prefill_idx in range(metadata.num_prefills): + if _prefill_idx >= forward_batch.extend_start_loc.shape[0]: + break + if _prefill_idx >= state_indices_tensor.shape[0]: + break + + _start = forward_batch.extend_start_loc[_prefill_idx] + + if _prefill_idx + 1 < forward_batch.extend_start_loc.shape[0]: + _end = forward_batch.extend_start_loc[_prefill_idx + 1] + else: + if ( + forward_batch.extend_seq_lens is not None + and _prefill_idx < forward_batch.extend_seq_lens.shape[0] + and metadata.num_decodes > 0 + ): + seq_len = forward_batch.extend_seq_lens[_prefill_idx] + _end = _start + seq_len + else: + _end = q.shape[0] + + slot_id = state_indices_tensor[_prefill_idx] + qs = q[_start:_end].transpose(0, 1).contiguous() + ks = k[_start:_end].transpose(0, 1).contiguous() + vs = v[_start:_end].transpose(0, 1).contiguous() + slice_layer_cache = kv_cache[slot_id, ...] + out_slice = BailingLinearKernel.jit_linear_forward_prefix( + qs, + ks, + vs, + slice_layer_cache, + self.tp_slope[layer.layer_id], + self.BLOCK, + layer_idx=layer.layer_id, + ) + hidden.append(out_slice.contiguous()) + if metadata.num_decodes > 0: + hidden.append( + self._decode_infer( + q, k, v, kv_cache, state_indices_tensor, metadata, layer + ) + ) + + if not hidden: + return torch.empty((0, q.size(-1)), device=q.device, dtype=q.dtype) + + hidden = torch.concat(hidden, dim=0).contiguous() + return hidden + + def _decode_infer(self, q, k, v, kv_cache, state_indices_tensor, metadata, layer): + num_prefill_tokens = metadata.num_prefill_tokens + num_prefills = metadata.num_prefills + q = q[num_prefill_tokens:].unsqueeze(2).contiguous() + k = k[num_prefill_tokens:].unsqueeze(2).contiguous() + v = v[num_prefill_tokens:].unsqueeze(2).contiguous() + slot_id = state_indices_tensor[num_prefills:] + + assert slot_id.shape[0] == q.shape[0], ( + f"slot_id length {slot_id.shape[0]} does not match decode batch size {q.shape[0]}. " + "This indicates a bug in the upstream logic that should be investigated." + ) + hidden = linear_decode_forward_triton( + q, k, v, kv_cache, self.tp_slope[layer.layer_id], slot_id, 32 + ) + return hidden + + def _linear_attention_entry( + self, + q, + k, + v, + kv_cache, + state_indices_tensor, + metadata, + layer, + mask=None, + temp_cache=None, + intermediate_state_indices=None, + ): + q_offsets = metadata.query_start_loc + + seg_meta = SegLaMeta( + batch_size=metadata.batch_size, + q_offsets=metadata.query_start_loc, + s_offsets=state_indices_tensor, + q_lengths=q_offsets.diff(), + s_scales=metadata.has_initial_states, + max_q_length=None, + mask=mask, + ) + hidden = seg_la_fwd( + q=q, + k=k, + v=v, + s=kv_cache, + decay_scales=self.tp_slope[layer.layer_id], + meta=seg_meta, + caches=temp_cache, + cache_indices=intermediate_state_indices, + decouple=True, + ) + return hidden + + def forward_extend( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + save_kv_cache=True, + **kwargs, + ): + q_rope = kwargs["q_rope"] if "q_rope" in kwargs else None + k_rope = kwargs["k_rope"] if "k_rope" in kwargs else None + layer_id = layer.layer_id if layer else kwargs["layer_id"] + + metadata = self.forward_metadata + + if self.kv_cache_dtype_str != "auto" and layer.k_scale is not None: + q = q.to(self.kv_cache_dtype) + + query_start_loc = self.forward_metadata.query_start_loc + cache_indices = self.forward_metadata.mamba_cache_indices + mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id) + ssm_states = mamba_cache_params.temporal + if self.linear_backend == "minimax": + o = self._prefill_and_mix_infer( + q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim), + k, + v, + ssm_states, + cache_indices, + forward_batch, + layer, + metadata, + ) + elif self.linear_backend == "seg_la": + intermediate_state_indices = ( + torch.arange( + cache_indices.shape[0], + dtype=torch.int32, + device=cache_indices.device, + ) + if forward_batch.forward_mode.is_target_verify() + else None + ) + o = self._linear_attention_entry( + q, + k, + v, + ssm_states, + cache_indices, + metadata, + layer, + temp_cache=( + mamba_cache_params.intermediate_ssm + if forward_batch.forward_mode.is_target_verify() + else None + ), + intermediate_state_indices=intermediate_state_indices, + ) + else: + raise ValueError( + f"linear backend: {self.linear_backend} is not support for now" + ) + return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + + def forward_decode( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + save_kv_cache=True, + **kwargs, + ) -> torch.Tensor: + q_rope = kwargs["q_rope"] if "q_rope" in kwargs else None + k_rope = kwargs["k_rope"] if "k_rope" in kwargs else None + layer_id = layer.layer_id if layer else kwargs["layer_id"] + + # Use precomputed metadata across all layers + metadata = self.forward_metadata + + if self.kv_cache_dtype_str != "auto": + q = q.to(self.kv_cache_dtype) + + # Do linear attention + query_start_loc = self.forward_metadata.query_start_loc + cache_indices = self.forward_metadata.mamba_cache_indices + mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer_id) + ssm_states = mamba_cache_params.temporal + if self.linear_backend == "minimax": + o = self._decode_infer(q, k, v, ssm_states, cache_indices, metadata, layer) + elif self.linear_backend == "seg_la": + o = self._linear_attention_entry( + q, k, v, ssm_states, cache_indices, metadata, layer + ) + else: + raise ValueError( + f"linear backend: {self.linear_backend} is not support for now" + ) + return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) diff --git a/python/sglang/srt/layers/attention/linear/utils.py b/python/sglang/srt/layers/attention/linear/utils.py new file mode 100644 index 000000000..59215ef3d --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/utils.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +import logging +from enum import Enum +from typing import TYPE_CHECKING, Optional + +from sglang.srt.utils.common import rank0_log + +if TYPE_CHECKING: + from sglang.srt.server_args import ServerArgs + +logger = logging.getLogger(__name__) + + +class LinearAttnKernelBackend(Enum): + TRITON = "triton" + CUTEDSL = "cutedsl" + + def is_triton(self): + return self == LinearAttnKernelBackend.TRITON + + def is_cutedsl(self): + return self == LinearAttnKernelBackend.CUTEDSL + + +LINEAR_ATTN_DECODE_BACKEND: Optional[LinearAttnKernelBackend] = None +LINEAR_ATTN_PREFILL_BACKEND: Optional[LinearAttnKernelBackend] = None + + +def initialize_linear_attn_config(server_args: ServerArgs): + global LINEAR_ATTN_DECODE_BACKEND + global LINEAR_ATTN_PREFILL_BACKEND + + base = server_args.linear_attn_backend + decode = server_args.linear_attn_decode_backend or base + prefill = server_args.linear_attn_prefill_backend or base + + LINEAR_ATTN_DECODE_BACKEND = LinearAttnKernelBackend(decode) + LINEAR_ATTN_PREFILL_BACKEND = LinearAttnKernelBackend(prefill) + rank0_log( + f"Linear attention kernel backend: " + f"decode={LINEAR_ATTN_DECODE_BACKEND.value}, " + f"prefill={LINEAR_ATTN_PREFILL_BACKEND.value}" + ) + + +def get_linear_attn_decode_backend() -> LinearAttnKernelBackend: + global LINEAR_ATTN_DECODE_BACKEND + if LINEAR_ATTN_DECODE_BACKEND is None: + logger.warning( + "LINEAR_ATTN_DECODE_BACKEND is not initialized, using triton backend" + ) + LINEAR_ATTN_DECODE_BACKEND = LinearAttnKernelBackend.TRITON + return LINEAR_ATTN_DECODE_BACKEND + + +def get_linear_attn_prefill_backend() -> LinearAttnKernelBackend: + global LINEAR_ATTN_PREFILL_BACKEND + if LINEAR_ATTN_PREFILL_BACKEND is None: + logger.warning( + "LINEAR_ATTN_PREFILL_BACKEND is not initialized, using triton backend" + ) + LINEAR_ATTN_PREFILL_BACKEND = LinearAttnKernelBackend.TRITON + return LINEAR_ATTN_PREFILL_BACKEND diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 228f7ef4d..044463090 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -216,6 +216,7 @@ MAMBA_SSM_DTYPE_CHOICES = ["float32", "bfloat16", "float16"] MAMBA_SCHEDULER_STRATEGY_CHOICES = ["auto", "no_buffer", "extra_buffer"] MAMBA_BACKEND_CHOICES = ["triton", "flashinfer"] +LINEAR_ATTN_KERNEL_BACKEND_CHOICES = ["triton", "cutedsl"] # Allow external code to add more choices @@ -524,6 +525,9 @@ class ServerArgs: mamba_full_memory_ratio: float = 0.9 mamba_scheduler_strategy: str = "auto" mamba_track_interval: int = 256 + linear_attn_backend: str = "triton" + linear_attn_decode_backend: Optional[str] = None + linear_attn_prefill_backend: Optional[str] = None # Hierarchical cache enable_hierarchical_cache: bool = False @@ -4305,6 +4309,31 @@ class ServerArgs: help="Choose the kernel backend for Mamba SSM operations. Default is 'triton'. " "Options: 'triton' (default), 'flashinfer' (requires FlashInfer with Mamba support).", ) + parser.add_argument( + "--linear-attn-backend", + type=str, + choices=LINEAR_ATTN_KERNEL_BACKEND_CHOICES, + default=ServerArgs.linear_attn_backend, + help="The default kernel backend for linear attention (GDN/KDA). " + "Can be overridden per-mode by --linear-attn-decode-backend " + "and --linear-attn-prefill-backend.", + ) + parser.add_argument( + "--linear-attn-decode-backend", + type=str, + choices=LINEAR_ATTN_KERNEL_BACKEND_CHOICES, + default=ServerArgs.linear_attn_decode_backend, + help="Override the kernel backend for linear attention decode. " + "If not set, uses --linear-attn-backend.", + ) + parser.add_argument( + "--linear-attn-prefill-backend", + type=str, + choices=LINEAR_ATTN_KERNEL_BACKEND_CHOICES, + default=ServerArgs.linear_attn_prefill_backend, + help="Override the kernel backend for linear attention prefill/extend. " + "If not set, uses --linear-attn-backend.", + ) # Hierarchical cache parser.add_argument(