From 7f6a2e2b5038fad550f718461a9837eb13edc60d Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Mon, 3 Aug 2026 18:05:01 -0700 Subject: [PATCH] [Refactor] Clean up and split DSA indexer (#33443) --- .../modules/deepseek_v2_attention_mla_npu.py | 2 +- .../srt/layers/attention/base_attn_backend.py | 4 +- .../srt/layers/attention/dsa/dsa_indexer.py | 759 +----------------- .../attention/dsa/dsa_indexer_metadata.py | 168 ++++ .../layers/attention/dsa/dsa_npu_indexer.py | 369 +++++++++ .../attention/dsa/dsa_prefill_cuda_graph.py | 191 +++++ .../srt/layers/attention/dsa_backend.py | 84 +- .../layers/attention/hybrid_attn_backend.py | 2 +- .../kernels/ops/attention/test_dsa_indexer.py | 10 +- 9 files changed, 753 insertions(+), 836 deletions(-) create mode 100644 python/sglang/srt/layers/attention/dsa/dsa_indexer_metadata.py create mode 100644 python/sglang/srt/layers/attention/dsa/dsa_npu_indexer.py create mode 100644 python/sglang/srt/layers/attention/dsa/dsa_prefill_cuda_graph.py diff --git a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py index 60dc42e8c..b3f60a561 100644 --- a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py +++ b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py @@ -11,7 +11,7 @@ from sglang.srt.hardware_backend.npu.attention.mla_preprocess import ( is_fia_nz, is_mla_preprocess_enabled, ) -from sglang.srt.layers.attention.dsa.dsa_indexer import scattered_to_tp_attn_full +from sglang.srt.layers.attention.dsa.dsa_npu_indexer import scattered_to_tp_attn_full from sglang.srt.layers.attention.dsa.utils import ( dsa_use_prefill_cp, ) diff --git a/python/sglang/srt/layers/attention/base_attn_backend.py b/python/sglang/srt/layers/attention/base_attn_backend.py index fb8ced91a..9b7cf373a 100644 --- a/python/sglang/srt/layers/attention/base_attn_backend.py +++ b/python/sglang/srt/layers/attention/base_attn_backend.py @@ -9,7 +9,9 @@ from sglang.kernel_api_logging import debug_kernel_api from sglang.srt.utils.common import is_npu if TYPE_CHECKING: - from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata + from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import ( + BaseIndexerMetadata, + ) from sglang.srt.layers.attention.verify_mask import VerifyMask from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.model_executor.forward_batch_info import ForwardBatch diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 2ba465ba3..79136182b 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -2,7 +2,6 @@ from __future__ import annotations import contextlib import logging -from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union import torch @@ -15,6 +14,14 @@ from sglang.kernels.ops.attention.fused_store_index_cache import ( from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype, is_fp8_fnuz from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.environ import envs +from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import BaseIndexerMetadata +from sglang.srt.layers.attention.dsa.dsa_npu_indexer import DSANPUIndexerMixin +from sglang.srt.layers.attention.dsa.dsa_prefill_cuda_graph import ( + GRAPH_WEIGHTS_PROJ_LORA_ERROR, + _is_in_piecewise_or_breakable_cuda_graph, + bcg_dsa_indexer_prefill_split, + pcg_dsa_indexer_prefill_split, +) from sglang.srt.layers.attention.dsa.paged_mqa_logits_backend import ( DSAPagedMQALogitsBackend, ) @@ -24,17 +31,12 @@ from sglang.srt.layers.attention.dsa.utils import ( is_dsa_prefill_cp_in_seq_split, is_graph_dsa_split_op_surface, ) -from sglang.srt.layers.dp_attention import attn_tp_all_gather_into_tensor from sglang.srt.layers.layernorm import LayerNorm, RMSNorm from sglang.srt.layers.utils import MultiPlatformOp -from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( - eager_on_graph, -) from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( is_in_breakable_cuda_graph, ) from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( - get_tc_piecewise_forward_context, is_in_tc_piecewise_cuda_graph, ) from sglang.srt.runtime_context import ( @@ -43,7 +45,6 @@ from sglang.srt.runtime_context import ( get_parallel, get_schedule, get_server_args, - get_spec, ) from sglang.srt.state_capturer.indexer_topk import ( maybe_capture_indexer_topk, @@ -62,7 +63,6 @@ from sglang.srt.utils.custom_op import register_custom_op logger = logging.getLogger(__name__) -global _use_multi_stream _is_cuda = is_cuda() _is_hip = is_hip() _is_npu = is_npu() @@ -108,16 +108,11 @@ if _is_cuda: if _use_aiter: from aiter.ops.cache import indexer_k_quant_and_cache -if is_npu(): - import torch_npu - from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream - from sglang.srt.distributed import ( get_attn_tp_group, ) from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.layers import deep_gemm_wrapper -from sglang.srt.layers.communicator import ScatterMode from sglang.srt.layers.cp.base import get_cp_strategy from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.layers.linear import ReplicatedLinear @@ -127,56 +122,15 @@ from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_context import ( get_attn_backend, - get_req_to_token_pool, get_token_to_kv_pool, ) from sglang.srt.model_executor.runner import get_is_capture_mode -from sglang.srt.runtime_context import get_server_args -_use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get() if TYPE_CHECKING: from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool DUAL_STREAM_TOKEN_THRESHOLD = 1024 if _is_cuda else 0 -GRAPH_WEIGHTS_PROJ_LORA_ERROR = ( - "DSA indexer weights_proj LoRA is incompatible with " - "piecewise/breakable CUDA graph; remove the explicit " - "prefill cuda-graph backend override or drop " - "indexer.weights_proj from the LoRA target modules." -) - - -def _is_in_piecewise_or_breakable_cuda_graph() -> bool: - return is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph() - - -def _uses_dsa_attention_backend(forward_batch: ForwardBatch) -> bool: - attn_backend = get_attn_backend() - server_args = get_server_args() - prefill_backend, decode_backend = server_args.get_attention_backends() - prefill_backend = ( - getattr(attn_backend, "prefill_attention_backend_str", None) or prefill_backend - ) - decode_backend = ( - getattr(attn_backend, "decode_attention_backend_str", None) or decode_backend - ) - - if forward_batch.forward_mode.is_decode_or_idle(): - backend_name = decode_backend - elif ( - forward_batch.forward_mode.is_target_verify() - or forward_batch.forward_mode.is_draft_extend_v2() - ): - backend_name = ( - decode_backend - if get_spec().speculative_attention_mode == "decode" - else prefill_backend - ) - else: - backend_name = prefill_backend - - return backend_name in ("dsa", "nsa") if _is_cuda: @@ -185,57 +139,10 @@ if _is_cuda: fused_k_indexer_norm_rope, fused_k_indexer_norm_rope_store, ) - - def _scale_head_gate_graph_fake_impl( - weights_raw: torch.Tensor, - n_heads_inv_sqrt: float, - softmax_scale: float, - q_scale: torch.Tensor, - ) -> torch.Tensor: - return torch.empty( - (weights_raw.shape[0], weights_raw.shape[1], q_scale.shape[-1]), - dtype=torch.float32, - device=weights_raw.device, - ) - - # In-graph (PCG/BCG) head gate for the fused path: weights_proj is folded - # into wk_weights_proj, so weights_raw is precomputed and there is no GEMM. - @register_custom_op(fake_impl=_scale_head_gate_graph_fake_impl) - def scale_head_gate_graph( - weights_raw: torch.Tensor, - n_heads_inv_sqrt: float, - softmax_scale: float, - q_scale: torch.Tensor, - ) -> torch.Tensor: - weights = weights_raw * n_heads_inv_sqrt - return weights.unsqueeze(-1) * q_scale * softmax_scale - - def _logits_head_gate_graph_fake_impl( - x: torch.Tensor, - weight: torch.Tensor, - n_heads_inv_sqrt: float, - softmax_scale: float, - q_scale: torch.Tensor, - ) -> torch.Tensor: - return torch.empty( - (x.shape[0], weight.shape[0], q_scale.shape[-1]), - dtype=torch.float32, - device=x.device, - ) - - # In-graph (PCG/BCG) head gate for the NON-prefill path - @register_custom_op(fake_impl=_logits_head_gate_graph_fake_impl) - def logits_head_gate_graph( - x: torch.Tensor, - weight: torch.Tensor, - n_heads_inv_sqrt: float, - softmax_scale: float, - q_scale: torch.Tensor, - ) -> torch.Tensor: - out = torch.mm(x, weight.t(), out_dtype=torch.float32) - weights = out * n_heads_inv_sqrt - weights = weights.unsqueeze(-1) * q_scale * softmax_scale - return weights + from sglang.srt.layers.attention.dsa.dsa_prefill_cuda_graph import ( + logits_head_gate_graph, + scale_head_gate_graph, + ) @register_custom_op(mutates_args=["topk_indices"]) @register_split_op() @@ -274,77 +181,6 @@ def _broadcast_indexer_topk_from_rank0( return topk_indices -class BaseIndexerMetadata(ABC): - @abstractmethod - def get_seqlens_int32(self) -> torch.Tensor: - """ - Return: (batch_size,) int32 tensor - """ - - @abstractmethod - def get_page_table_64(self) -> torch.Tensor: - """ - Return: (batch_size, num_blocks) int32, page table. - The page size of the table is 64. - """ - - @abstractmethod - def get_page_table_1(self) -> torch.Tensor: - """ - Return: (batch_size, num_blocks) int32, page table. - The page size of the table is 1. - """ - - @abstractmethod - def get_seqlens_expanded(self) -> torch.Tensor: - """ - Return: (sum_extend_seq_len,) int32 tensor - """ - - def get_indexer_kvcache_range(self) -> Tuple[torch.Tensor, torch.Tensor]: - """ - Return: (tokens, ), (tokens, ) int32, k_start and k_end in kv cache(token,xxx) for each token. - """ - - def get_indexer_seq_len_cpu(self) -> torch.Tensor: - """ - Return: seq lens for each batch. - """ - - def get_indexer_seq_len(self) -> torch.Tensor: - """ - Return: seq lens for each batch. - """ - - def get_dsa_extend_len_cpu(self) -> List[int]: - """ - Return: extend seq lens for each batch. - """ - - def get_token_to_batch_idx(self) -> torch.Tensor: - """ - Return: batch idx for each token. - """ - - @abstractmethod - def topk_transform( - self, - logits: torch.Tensor, - topk: int, - ) -> torch.Tensor: - """ - Perform topk selection on the logits and possibly transform the result. - - NOTE that attention backend may override this function to do some - transformation, which means the result of this topk_transform may not - be the topk indices of the input logits. - - Return: Anything, since it will be passed to the attention backend - for further processing on sparse attention computation. - Don't assume it is the topk indices of the input logits. - """ - - def rotate_activation(x: torch.Tensor) -> torch.Tensor: # from sgl_kernel import hadamard_transform if _is_hip: @@ -361,7 +197,7 @@ def rotate_activation(x: torch.Tensor) -> torch.Tensor: return hadamard_transform(x, scale=hidden_size**-0.5) -class Indexer(MultiPlatformOp): +class Indexer(DSANPUIndexerMixin, MultiPlatformOp): _MQA_LOGITS_BYTES_PER_ELEM = 4 _MQA_LOGITS_STATIC_SKIP_ELEMS = 8_000_000 _MQA_LOGITS_TOTAL_MEM_FRACTION = 0.3 @@ -408,10 +244,8 @@ class Indexer(MultiPlatformOp): self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() if self.dsa_enable_prefill_cp: self.cp_size = get_parallel().attn_cp_size - self.cp_rank = get_parallel().attn_cp_rank else: self.cp_size = None - self.cp_rank = None if _is_cuda: self.sm_count = deep_gemm.get_num_sms() self.half_device_sm_count = ceil_align(self.sm_count // 2, 8) @@ -675,7 +509,6 @@ class Indexer(MultiPlatformOp): self, x: torch.Tensor, positions: torch.Tensor, - enable_dual_stream: bool, ): # Non-fusion path only; self.wk does not exist when fusion is on. key, _ = self.wk(x) @@ -1329,7 +1162,6 @@ class Indexer(MultiPlatformOp): forward_batch: ForwardBatch, layer_id: int, act_quant, - enable_dual_stream: bool, metadata: BaseIndexerMetadata, return_indices: bool = True, *, @@ -1372,7 +1204,7 @@ class Indexer(MultiPlatformOp): out_cache_loc=out_cache_loc, ) else: - key = self._get_k_bf16(x, positions, enable_dual_stream) + key = self._get_k_bf16(x, positions) if num_tokens is not None: assert num_tokens <= key.shape[0] key = key[:num_tokens] @@ -1560,93 +1392,6 @@ class Indexer(MultiPlatformOp): return topk_result - def forward_indexer( - self, - q_fp8: torch.Tensor, - weights: torch.Tensor, - forward_batch: ForwardBatch, - topk: int, - layer_id: int, - ) -> Optional[torch.Tensor]: - assert not _is_in_piecewise_or_breakable_cuda_graph(), ( - "DSA forward_indexer (non-CUDA loop path) not supported under " - "piecewise/breakable CUDA graph" - ) - if not _is_npu: - from sglang.kernels.ops.attention.dsa.tilelang_kernel import fp8_index - - page_size = get_token_to_kv_pool().page_size - assert page_size == 64, "only support page size 64" - - assert len(weights.shape) == 3 - weights = weights.squeeze(-1) - - # logits = deep_gemm.fp8_mqa_logits(q_fp8, kv_fp8, weights, ks, ke) - k_fp8_list = [] - k_scale_list = [] - - topk_indices_list = [] - - block_tables = get_req_to_token_pool().req_to_token[ - forward_batch.req_pool_indices, : - ] - strided_indices = torch.arange( - 0, block_tables.shape[-1], page_size, device="cuda" - ) - block_tables = block_tables[:, strided_indices] // page_size - - q_len_start = 0 - - for i in range(forward_batch.batch_size): - seq_len = forward_batch.seq_lens[i].item() - q_len = ( - forward_batch.extend_seq_lens_cpu[i] - if forward_batch.forward_mode.is_extend() - else 1 - ) - q_len_end = q_len_start + q_len - - q_fp8_partial = q_fp8[q_len_start:q_len_end] - q_fp8_partial = q_fp8_partial.unsqueeze(0).contiguous() - - weights_partial = weights[q_len_start:q_len_end] - weights_partial = weights_partial.squeeze(-1).unsqueeze(0).contiguous() - - k_fp8 = get_token_to_kv_pool().get_index_k_continuous( - layer_id, - seq_len, - block_tables[i], - ) - k_scale = get_token_to_kv_pool().get_index_k_scale_continuous( - layer_id, - seq_len, - block_tables[i], - ) - - k_fp8 = k_fp8.view(torch.float8_e4m3fn).unsqueeze(0).contiguous() - k_scale = k_scale.view(torch.float32).squeeze(-1).unsqueeze(0).contiguous() - - index_score = fp8_index( - q_fp8_partial, - weights_partial, - k_fp8, - k_scale, - ) - end_pos = seq_len - topk_indices = index_score.topk(min(topk, end_pos), dim=-1)[1].squeeze(0) - - pad_len = ceil_align(topk_indices.shape[-1], 2048) - topk_indices.shape[-1] - topk_indices = torch.nn.functional.pad( - topk_indices, (0, pad_len), "constant", -1 - ) - - topk_indices_list.append(topk_indices) - - q_len_start = q_len_end - - topk_indices = torch.cat(topk_indices_list, dim=0) - return topk_indices - def _store_index_k_cache( self, forward_batch: ForwardBatch, @@ -1729,19 +1474,6 @@ class Indexer(MultiPlatformOp): index_k_scale=k_scale, ) - def forward_xpu( - self, - x: torch.Tensor, - q_lora: torch.Tensor, - positions: torch.Tensor, - forward_batch: ForwardBatch, - layer_id: int, - return_indices: bool = True, - ) -> Optional[torch.Tensor]: - return self.forward_cuda( - x, q_lora, positions, forward_batch, layer_id, return_indices - ) - def forward_cuda( self, x: torch.Tensor, @@ -1801,7 +1533,6 @@ class Indexer(MultiPlatformOp): forward_batch, layer_id, act_quant, - enable_dual_stream, metadata, return_indices, ) @@ -2081,466 +1812,6 @@ class Indexer(MultiPlatformOp): metadata, ) else: - topk_result = self.forward_indexer( - q_fp8.contiguous(), - weights, - forward_batch, - topk=self.index_topk, - layer_id=layer_id, - ) + raise NotImplementedError("DSA indexer only supports CUDA, HIP, and NPU") topk_result = _broadcast_indexer_topk_from_rank0(topk_result) return maybe_capture_indexer_topk(layer_id, topk_result) - - def forward_npu( - self, - x: torch.Tensor, - q_lora: torch.Tensor, - positions: torch.Tensor, - forward_batch: ForwardBatch, - layer_id: int, - layer_scatter_modes=None, - dynamic_scale: torch.Tensor = None, - ) -> torch.Tensor: - if get_attn_backend().forward_metadata.seq_lens_cpu_int is None: - actual_seq_lengths_kv = get_attn_backend().forward_metadata.seq_lens - else: - actual_seq_lengths_kv = get_attn_backend().forward_metadata.seq_lens_cpu_int - is_prefill = ( - forward_batch.forward_mode.is_extend() - and not forward_batch.forward_mode.is_draft_extend_v2() - and not forward_batch.forward_mode.is_target_verify() - ) - - bs = q_lora.shape[0] - - if self.rotary_emb.is_neox_style: - if not hasattr(forward_batch, "npu_indexer_sin_cos_cache"): - cos_sin = self.rotary_emb.cos_sin_cache[positions] - cos, sin = cos_sin.chunk(2, dim=-1) - cos = cos.repeat(1, 2).view(-1, 1, 1, self.rope_head_dim) - sin = sin.repeat(1, 2).view(-1, 1, 1, self.rope_head_dim) - forward_batch.npu_indexer_sin_cos_cache = (sin, cos) - else: - sin, cos = forward_batch.npu_indexer_sin_cos_cache - - if self.alt_stream is not None: - self.alt_stream.wait_stream(torch.npu.current_stream()) - with torch.npu.stream(self.alt_stream): - q_lora = ( - (q_lora, dynamic_scale) if dynamic_scale is not None else q_lora - ) - q = self.wq_b(q_lora)[ - 0 - ] # [bs, 1536] @ [1536, 64 * 128] = [bs, 64 * 128] - wq_b_event = self.alt_stream.record_event() - q = q.view(bs, self.n_heads, self.head_dim) # [bs, 64, 128] - q_pe, q_nope = torch.split( - q, - [self.rope_head_dim, self.head_dim - self.rope_head_dim], - dim=-1, - ) # [bs, 64, 64 + 64] - q_pe = q_pe.view(bs, self.n_heads, 1, self.rope_head_dim) - q_pe = torch_npu.npu_rotary_mul(q_pe, cos, sin).view( - bs, self.n_heads, self.rope_head_dim - ) # [bs, n, d] - q = torch.cat([q_pe, q_nope], dim=-1) - q.record_stream(self.alt_stream) - q_rope_event = self.alt_stream.record_event() - else: - q_lora = ( - (q_lora, dynamic_scale) if dynamic_scale is not None else q_lora - ) - q = self.wq_b(q_lora)[ - 0 - ] # [bs, 1536] @ [1536, 64 * 128] = [bs, 64 * 128] - q = q.view(bs, self.n_heads, self.head_dim) # [bs, 64, 128] - q_pe, q_nope = torch.split( - q, - [self.rope_head_dim, self.head_dim - self.rope_head_dim], - dim=-1, - ) # [bs, 64, 64 + 64] - q_pe = q_pe.view(bs, self.n_heads, 1, self.rope_head_dim) - q_pe = torch_npu.npu_rotary_mul(q_pe, cos, sin).view( - bs, self.n_heads, self.rope_head_dim - ) # [bs, n, d] - q = torch.cat([q_pe, q_nope], dim=-1) - - if envs.SGLANG_NPU_USE_MULTI_STREAM.get(): - indexer_weight_stream = get_indexer_weight_stream() - indexer_weight_stream.wait_stream(torch.npu.current_stream()) - with torch.npu.stream(indexer_weight_stream): - x = x.view(-1, self.hidden_size) - weights = self.weights_proj(x.float())[0].to(torch.bfloat16) - weights.record_stream(indexer_weight_stream) - weights_event = indexer_weight_stream.record_event() - else: - x = x.view(-1, self.hidden_size) - weights = self.weights_proj(x.float())[0].to(torch.bfloat16) - - k_proj = self.wk(x)[0] # [b, s, 7168] @ [7168, 128] = [b, s, 128] - k = self.k_norm(k_proj) - if ( - _use_ag_after_qlora - and layer_scatter_modes.layer_input_mode == ScatterMode.SCATTERED - and layer_scatter_modes.attn_mode == ScatterMode.TP_ATTN_FULL - ): - k = scattered_to_tp_attn_full(k, forward_batch) - k_pe, k_nope = torch.split( - k, - [self.rope_head_dim, self.head_dim - self.rope_head_dim], - dim=-1, - ) # [bs, 64 + 64] - - k_pe = k_pe.view(-1, 1, 1, self.rope_head_dim) - k_pe = torch.ops.npu.npu_rotary_mul(k_pe, cos, sin).view( - bs, 1, self.rope_head_dim - ) # [bs, 1, d] - k = torch.cat([k_pe, k_nope.unsqueeze(1)], dim=-1) # [bs, 1, 128] - - else: - if envs.SGLANG_NPU_USE_MULTI_STREAM.get(): - indexer_weight_stream = get_indexer_weight_stream() - indexer_weight_stream.wait_stream(torch.npu.current_stream()) - with torch.npu.stream(indexer_weight_stream): - x = x.view(-1, self.hidden_size) - weights = self.weights_proj(x.float())[0].to(torch.bfloat16) - weights.record_stream(indexer_weight_stream) - weights_event = indexer_weight_stream.record_event() - else: - x = x.view(-1, self.hidden_size) - weights = self.weights_proj(x.float())[0].to(torch.bfloat16) - - q_lora = (q_lora, dynamic_scale) if dynamic_scale is not None else q_lora - q = self.wq_b(q_lora)[0] # [bs, 1536] @ [1536, 64 * 128] = [bs, 64 * 128] - q = q.view(bs, self.n_heads, self.head_dim) # [bs, 64, 128] - q_pe, q_nope = torch.split( - q, - [self.rope_head_dim, self.head_dim - self.rope_head_dim], - dim=-1, - ) # [bs, 64, 64 + 64] - - k_proj = self.wk(x)[0] # [b, s, 7168] @ [7168, 128] = [b, s, 128] - k = self.k_norm(k_proj) - k_pe, k_nope = torch.split( - k, - [self.rope_head_dim, self.head_dim - self.rope_head_dim], - dim=-1, - ) # [bs, 64 + 64] - - k_pe = k_pe.unsqueeze(1) - - if layer_id == 0: - self.rotary_emb.sin_cos_cache = ( - self.rotary_emb.cos_sin_cache.index_select(0, positions) - ) - - q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe) - k_pe = k_pe.squeeze(1) - q = torch.cat([q_pe, q_nope], dim=-1) - k = torch.cat([k_pe, k_nope], dim=-1) - - if ( - is_prefill - and self.dsa_enable_prefill_cp - and forward_batch.attn_cp_metadata is not None - ): - k = cp_all_gather_rerange_output( - k.contiguous().view(-1, self.head_dim), - self.cp_size, - forward_batch, - torch.npu.current_stream(), - ) - - get_token_to_kv_pool().set_index_k_buffer( - layer_id, forward_batch.out_cache_loc, k - ) - if is_prefill: - if ( - self.dsa_enable_prefill_cp - and forward_batch.attn_cp_metadata is not None - ): - get_attn_backend().forward_metadata.actual_seq_lengths_q = ( - forward_batch.attn_cp_metadata.actual_seq_q_prev_tensor, - forward_batch.attn_cp_metadata.actual_seq_q_next_tensor, - ) - if sum(forward_batch.extend_prefix_lens_cpu) > 0: - total_kv_len_prev_tensor = ( - forward_batch.attn_cp_metadata.kv_len_prev_tensor - + forward_batch.extend_prefix_lens.squeeze() - ) - total_kv_len_next_tensor = ( - forward_batch.attn_cp_metadata.kv_len_next_tensor - + forward_batch.extend_prefix_lens.squeeze() - ) - get_attn_backend().forward_metadata.actual_seq_lengths_kv = ( - total_kv_len_prev_tensor, - total_kv_len_next_tensor, - ) - else: - get_attn_backend().forward_metadata.actual_seq_lengths_kv = ( - forward_batch.attn_cp_metadata.kv_len_prev_tensor, - forward_batch.attn_cp_metadata.kv_len_next_tensor, - ) - actual_seq_lengths_q = ( - get_attn_backend().forward_metadata.actual_seq_lengths_q - ) - actual_seq_lengths_kv = ( - get_attn_backend().forward_metadata.actual_seq_lengths_kv - ) - else: - actual_seq_lengths_kv = forward_batch.seq_lens - actual_seq_lengths_q = forward_batch.extend_seq_lens.cumsum(dim=0) - else: - if get_attn_backend().forward_metadata.actual_seq_lengths_q is None: - if ( - forward_batch.forward_mode.is_draft_extend_v2() - or forward_batch.forward_mode.is_target_verify() - ): - num_draft_tokens = get_attn_backend().speculative_num_draft_tokens - actual_seq_lengths_q = torch.arange( - num_draft_tokens, - num_draft_tokens + bs, - num_draft_tokens, - dtype=torch.int32, - device=k.device, - ) - else: - actual_seq_lengths_q = torch.tensor( - [1 + i * 1 for i in range(bs)], - dtype=torch.int32, - device=k.device, - ) - else: - actual_seq_lengths_q = ( - get_attn_backend().forward_metadata.actual_seq_lengths_q - ) - - past_key_states = get_token_to_kv_pool().get_index_k_buffer(layer_id) - - if self.rotary_emb.is_neox_style and self.alt_stream is not None: - torch.npu.current_stream().wait_event(q_rope_event) - if envs.SGLANG_NPU_USE_MULTI_STREAM.get(): - torch.npu.current_stream().wait_event(weights_event) - if ( - _use_ag_after_qlora - and layer_scatter_modes.layer_input_mode == ScatterMode.SCATTERED - and layer_scatter_modes.attn_mode == ScatterMode.TP_ATTN_FULL - ): - weights = scattered_to_tp_attn_full(weights, forward_batch) - block_table = get_attn_backend().forward_metadata.block_tables - if ( - is_prefill - and self.dsa_enable_prefill_cp - and forward_batch.attn_cp_metadata is not None - ): - block_table = block_table[: actual_seq_lengths_q[0].numel()] - topk_indices = self.do_npu_cp_balance_indexer( - q.view(-1, self.n_heads, self.head_dim), - past_key_states, - weights, - actual_seq_lengths_q, - actual_seq_lengths_kv, - block_table, - ) - return topk_indices - else: - block_table = ( - block_table[: actual_seq_lengths_q.size()[0]] - if is_prefill - else block_table - ) - - topk_indices = torch_npu.npu_lightning_indexer( - query=q.view(-1, self.n_heads, self.head_dim), - key=past_key_states, - weights=weights, - actual_seq_lengths_query=actual_seq_lengths_q.to(torch.int32), - actual_seq_lengths_key=actual_seq_lengths_kv.to(k.device).to( - torch.int32 - ), - block_table=block_table, - layout_query="TND", - layout_key="PA_BSND", - sparse_count=self.index_topk, - sparse_mode=3, - ) - # Keep DSA top-k as [T, K]; NPU attention expands it when needed. - return topk_indices[0].squeeze(1) - - def do_npu_cp_balance_indexer( - self, - q, - past_key_states, - indexer_weights, - actual_seq_lengths_q, - actual_seq_lengths_kv, - block_table, - ): - q_prev, q_next = torch.split(q, (q.size(0) + 1) // 2, dim=0) - weights_prev, weights_next = None, None - if indexer_weights is not None: - weights_prev, weights_next = torch.split( - indexer_weights, (indexer_weights.size(0) + 1) // 2, dim=0 - ) - weights_prev = weights_prev.contiguous().view(-1, weights_prev.shape[-1]) - weights_next = weights_next.contiguous().view(-1, weights_next.shape[-1]) - - actual_seq_lengths_q_prev, actual_seq_lengths_q_next = actual_seq_lengths_q - actual_seq_lengths_kv_prev, actual_seq_lengths_kv_next = actual_seq_lengths_kv - - topk_indices_prev = torch_npu.npu_lightning_indexer( - query=q_prev, - key=past_key_states, - weights=weights_prev, - actual_seq_lengths_query=actual_seq_lengths_q_prev.to( - device=q.device, dtype=torch.int32 - ), - actual_seq_lengths_key=actual_seq_lengths_kv_prev.to( - device=q.device, dtype=torch.int32 - ), - block_table=block_table, - layout_query="TND", - layout_key="PA_BSND", - sparse_count=self.index_topk, - sparse_mode=3, - ) - topk_indices_next = torch_npu.npu_lightning_indexer( - query=q_next, - key=past_key_states, - weights=weights_next, - actual_seq_lengths_query=actual_seq_lengths_q_next.to( - device=q.device, dtype=torch.int32 - ), - actual_seq_lengths_key=actual_seq_lengths_kv_next.to( - device=q.device, dtype=torch.int32 - ), - block_table=block_table, - layout_query="TND", - layout_key="PA_BSND", - sparse_count=self.index_topk, - sparse_mode=3, - ) - return torch.cat([topk_indices_prev[0], topk_indices_next[0]], dim=0).squeeze(1) - - -@register_custom_op(mutates_args=["topk_result"]) -@register_split_op() -def pcg_dsa_indexer_prefill_split( - layer_id: int, - x: torch.Tensor, - q_lora: torch.Tensor, - positions: torch.Tensor, - topk_result: torch.Tensor, -) -> None: - # Default in-graph indexer path for non-CP prefill: runs the whole indexer - # (q/k proj, head gate, k-cache store, topk) as one eager split op. PCG calls - # this as a split op; BCG uses the explicit eager wrapper below. - # - # Output contract (differs from the eager `forward` path): a split op returns - # None, so results are delivered only by mutating `topk_result` in place. The - # call site pre-allocates it at a static, padded shape and a downstream - # captured graph reads it at a fixed address; eager code instead allocates - # and returns a fresh, naturally-sized tensor each call. - assert _is_cuda, "Internal error: DSA graph dispatch is only supported on CUDA" - from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant - - forward_context = get_tc_piecewise_forward_context() - forward_batch = forward_context.forward_batch - indexer = forward_context.dsa_indexers[layer_id] - metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch) - - extend_num_tokens = forward_batch.extend_num_tokens - # Empty buffer encodes return_indices=False for graph dispatch. - return_indices = topk_result.numel() != 0 - k_only = not return_indices or ( - indexer._should_skip_logits_computation(forward_batch) - and not indexer.dsa_enable_prefill_cp - ) - if k_only: - indexer._forward_cuda_k_only( - x, - positions, - forward_batch, - layer_id, - act_quant, - enable_dual_stream=False, - metadata=metadata, - return_indices=return_indices, - num_tokens=extend_num_tokens, - topk_result=topk_result, - ) - return - - # Fused path stores K (no-Hadamard) and computes q_fp8 + head gate in the - # fused kernels, sliced to the unpadded count. Single stream: the split op is - # captured, so the dual-stream overlap is disabled. - if indexer.use_dsa_indexer_fusion: - q_fp8, weights = indexer._fused_q_prepare_and_store( - x, - q_lora, - positions, - forward_batch, - layer_id, - act_quant, - num_tokens=extend_num_tokens, - enable_dual_stream=False, - ) - indexer._get_topk_ragged( - False, - forward_batch, - layer_id, - q_fp8, - weights, - metadata, - topk_result, - ) - return - - query, key, _ = indexer._get_q_k_bf16( - q_lora, - x, - positions, - enable_dual_stream=False, - forward_batch=forward_batch, - ) - q_fp8, q_scale = act_quant(query, indexer.block_size, indexer.scale_fmt) - # Reuse the compiled head-gate util shared with the eager path. - weights = indexer._get_logits_head_gate(x, q_scale) - # Store K cache + ragged top-k, sliced to the unpadded count and writing into - # the static padded topk_result buffer (the graph contract). Mirrors the eager - # path's store + _get_topk_ragged. - indexer._store_index_k_cache( - forward_batch=forward_batch, - layer_id=layer_id, - key=key[:extend_num_tokens], - act_quant=act_quant, - out_cache_loc=forward_batch.out_cache_loc[:extend_num_tokens], - ) - indexer._get_topk_ragged( - False, - forward_batch, - layer_id, - q_fp8[:extend_num_tokens], - weights, - metadata, - topk_result, - ) - - -bcg_dsa_indexer_prefill_split = eager_on_graph(True)(pcg_dsa_indexer_prefill_split) - - -def scattered_to_tp_attn_full( - hidden_states: torch.Tensor, - forward_batch, -) -> torch.Tensor: - hidden_states, local_hidden_states = ( - torch.empty( - (forward_batch.input_ids.shape[0], hidden_states.shape[1]), - dtype=hidden_states.dtype, - device=hidden_states.device, - ), - hidden_states, - ) - attn_tp_all_gather_into_tensor(hidden_states, local_hidden_states.contiguous()) - return hidden_states diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer_metadata.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer_metadata.py new file mode 100644 index 000000000..a35118263 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer_metadata.py @@ -0,0 +1,168 @@ +from __future__ import annotations + +from abc import ABC, abstractmethod +from dataclasses import dataclass +from typing import TYPE_CHECKING, List, Optional, Tuple + +import torch + +from sglang.srt.layers.attention.dsa.dsa_backend_mtp_precompute import ( + compute_cu_seqlens, +) +from sglang.srt.layers.attention.dsa.dsa_topk_backend import ( + DSATopKBackend, + TopkTransformMethod, +) + +if TYPE_CHECKING: + from sglang.srt.layers.attention.dsa_backend import DSAMetadata + + +class BaseIndexerMetadata(ABC): + @abstractmethod + def get_seqlens_int32(self) -> torch.Tensor: + """ + Return: (batch_size,) int32 tensor + """ + + @abstractmethod + def get_page_table_64(self) -> torch.Tensor: + """ + Return: (batch_size, num_blocks) int32, page table. + The page size of the table is 64. + """ + + @abstractmethod + def get_page_table_1(self) -> torch.Tensor: + """ + Return: (batch_size, num_blocks) int32, page table. + The page size of the table is 1. + """ + + @abstractmethod + def get_seqlens_expanded(self) -> torch.Tensor: + """ + Return: (sum_extend_seq_len,) int32 tensor + """ + + def get_indexer_kvcache_range(self) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Return: (tokens, ), (tokens, ) int32, k_start and k_end in kv cache(token,xxx) for each token. + """ + + def get_indexer_seq_len_cpu(self) -> torch.Tensor: + """ + Return: seq lens for each batch. + """ + + def get_indexer_seq_len(self) -> torch.Tensor: + """ + Return: seq lens for each batch. + """ + + def get_dsa_extend_len_cpu(self) -> List[int]: + """ + Return: extend seq lens for each batch. + """ + + def get_token_to_batch_idx(self) -> torch.Tensor: + """ + Return: batch idx for each token. + """ + + @abstractmethod + def topk_transform( + self, + logits: torch.Tensor, + topk: int, + ) -> torch.Tensor: + """ + Perform topk selection on the logits and possibly transform the result. + + NOTE that attention backend may override this function to do some + transformation, which means the result of this topk_transform may not + be the topk indices of the input logits. + + Return: Anything, since it will be passed to the attention backend + for further processing on sparse attention computation. + Don't assume it is the topk indices of the input logits. + """ + + +@dataclass(frozen=True) +class DSAIndexerMetadata(BaseIndexerMetadata): + attn_metadata: DSAMetadata + topk_transform_method: TopkTransformMethod + topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL + paged_mqa_schedule_metadata: Optional[torch.Tensor] = None + paged_mqa_ctx_lens_2d: Optional[torch.Tensor] = None + force_unfused_topk: bool = False + + def get_seqlens_int32(self) -> torch.Tensor: + return self.attn_metadata.cache_seqlens_int32 + + def get_page_table_64(self) -> torch.Tensor: + return self.attn_metadata.real_page_table + + def get_page_table_1(self) -> torch.Tensor: + return self.attn_metadata.page_table_1 + + def get_seqlens_expanded(self) -> torch.Tensor: + return self.attn_metadata.dsa_seqlens_expanded + + def get_indexer_kvcache_range(self) -> Tuple[torch.Tensor, torch.Tensor]: + return self.attn_metadata.indexer_k_start_end + + def get_indexer_seq_len(self) -> torch.Tensor: + return self.attn_metadata.indexer_seq_lens + + def get_indexer_seq_len_cpu(self) -> torch.Tensor: + return self.attn_metadata.indexer_seq_lens_cpu + + def get_dsa_extend_len_cpu(self) -> List[int]: + return self.attn_metadata.dsa_extend_seq_lens_list + + def get_token_to_batch_idx(self) -> torch.Tensor: + return self.attn_metadata.token_to_batch_idx + + def topk_transform( + self, + logits: torch.Tensor, + topk: int, + ks: Optional[torch.Tensor] = None, + cu_seqlens_q: Optional[torch.Tensor] = None, + ke_offset: Optional[torch.Tensor] = None, + batch_idx_list: Optional[List[int]] = None, + topk_indices_offset_override: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + if topk_indices_offset_override is not None: + cu_topk_indices_offset = topk_indices_offset_override + cu_seqlens_q_topk = None + elif cu_seqlens_q is not None: + cu_seqlens_q = cu_seqlens_q.to(torch.int32) + cu_seqlens_q_topk = compute_cu_seqlens(cu_seqlens_q) + cu_topk_indices_offset = torch.repeat_interleave( + cu_seqlens_q_topk[:-1], + cu_seqlens_q, + # Avoid reading sum(cu_seqlens_q) back to the host. + output_size=logits.shape[0], + ) + else: + cu_seqlens_q_topk = self.attn_metadata.cu_seqlens_q + cu_topk_indices_offset = self.attn_metadata.topk_indices_offset + if ke_offset is not None: + seq_lens_topk = ke_offset + else: + seq_lens_topk = self.get_seqlens_expanded() + return self.topk_backend.topk_transform( + logits=logits, + lengths=seq_lens_topk, + topk=topk, + topk_transform_method=self.topk_transform_method, + attn_metadata=self.attn_metadata, + cu_seqlens_q_topk=cu_seqlens_q_topk, + topk_indices_offset=cu_topk_indices_offset, + row_starts=ks, + batch_idx_list=batch_idx_list, + force_unfused_topk=self.force_unfused_topk, + ) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_npu_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_npu_indexer.py new file mode 100644 index 000000000..9329dad7a --- /dev/null +++ b/python/sglang/srt/layers/attention/dsa/dsa_npu_indexer.py @@ -0,0 +1,369 @@ +from __future__ import annotations + +import torch + +from sglang.srt.environ import envs +from sglang.srt.layers.communicator import ScatterMode +from sglang.srt.layers.dp_attention import attn_tp_all_gather_into_tensor +from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.model_executor.forward_context import ( + get_attn_backend, + get_token_to_kv_pool, +) +from sglang.srt.utils import is_npu + +if is_npu(): + import torch_npu + + from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream + +_use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get() + + +class DSANPUIndexerMixin: + def forward_npu( + self, + x: torch.Tensor, + q_lora: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + layer_id: int, + layer_scatter_modes=None, + dynamic_scale: torch.Tensor = None, + ) -> torch.Tensor: + if get_attn_backend().forward_metadata.seq_lens_cpu_int is None: + actual_seq_lengths_kv = get_attn_backend().forward_metadata.seq_lens + else: + actual_seq_lengths_kv = get_attn_backend().forward_metadata.seq_lens_cpu_int + is_prefill = ( + forward_batch.forward_mode.is_extend() + and not forward_batch.forward_mode.is_draft_extend_v2() + and not forward_batch.forward_mode.is_target_verify() + ) + + bs = q_lora.shape[0] + + if self.rotary_emb.is_neox_style: + if not hasattr(forward_batch, "npu_indexer_sin_cos_cache"): + cos_sin = self.rotary_emb.cos_sin_cache[positions] + cos, sin = cos_sin.chunk(2, dim=-1) + cos = cos.repeat(1, 2).view(-1, 1, 1, self.rope_head_dim) + sin = sin.repeat(1, 2).view(-1, 1, 1, self.rope_head_dim) + forward_batch.npu_indexer_sin_cos_cache = (sin, cos) + else: + sin, cos = forward_batch.npu_indexer_sin_cos_cache + + if self.alt_stream is not None: + self.alt_stream.wait_stream(torch.npu.current_stream()) + with torch.npu.stream(self.alt_stream): + q_lora = ( + (q_lora, dynamic_scale) if dynamic_scale is not None else q_lora + ) + q = self.wq_b(q_lora)[ + 0 + ] # [bs, 1536] @ [1536, 64 * 128] = [bs, 64 * 128] + q = q.view(bs, self.n_heads, self.head_dim) # [bs, 64, 128] + q_pe, q_nope = torch.split( + q, + [self.rope_head_dim, self.head_dim - self.rope_head_dim], + dim=-1, + ) # [bs, 64, 64 + 64] + q_pe = q_pe.view(bs, self.n_heads, 1, self.rope_head_dim) + q_pe = torch_npu.npu_rotary_mul(q_pe, cos, sin).view( + bs, self.n_heads, self.rope_head_dim + ) # [bs, n, d] + q = torch.cat([q_pe, q_nope], dim=-1) + q.record_stream(self.alt_stream) + q_rope_event = self.alt_stream.record_event() + else: + q_lora = ( + (q_lora, dynamic_scale) if dynamic_scale is not None else q_lora + ) + q = self.wq_b(q_lora)[ + 0 + ] # [bs, 1536] @ [1536, 64 * 128] = [bs, 64 * 128] + q = q.view(bs, self.n_heads, self.head_dim) # [bs, 64, 128] + q_pe, q_nope = torch.split( + q, + [self.rope_head_dim, self.head_dim - self.rope_head_dim], + dim=-1, + ) # [bs, 64, 64 + 64] + q_pe = q_pe.view(bs, self.n_heads, 1, self.rope_head_dim) + q_pe = torch_npu.npu_rotary_mul(q_pe, cos, sin).view( + bs, self.n_heads, self.rope_head_dim + ) # [bs, n, d] + q = torch.cat([q_pe, q_nope], dim=-1) + + if envs.SGLANG_NPU_USE_MULTI_STREAM.get(): + indexer_weight_stream = get_indexer_weight_stream() + indexer_weight_stream.wait_stream(torch.npu.current_stream()) + with torch.npu.stream(indexer_weight_stream): + x = x.view(-1, self.hidden_size) + weights = self.weights_proj(x.float())[0].to(torch.bfloat16) + weights.record_stream(indexer_weight_stream) + weights_event = indexer_weight_stream.record_event() + else: + x = x.view(-1, self.hidden_size) + weights = self.weights_proj(x.float())[0].to(torch.bfloat16) + + k_proj = self.wk(x)[0] # [b, s, 7168] @ [7168, 128] = [b, s, 128] + k = self.k_norm(k_proj) + if ( + _use_ag_after_qlora + and layer_scatter_modes.layer_input_mode == ScatterMode.SCATTERED + and layer_scatter_modes.attn_mode == ScatterMode.TP_ATTN_FULL + ): + k = scattered_to_tp_attn_full(k, forward_batch) + k_pe, k_nope = torch.split( + k, + [self.rope_head_dim, self.head_dim - self.rope_head_dim], + dim=-1, + ) # [bs, 64 + 64] + + k_pe = k_pe.view(-1, 1, 1, self.rope_head_dim) + k_pe = torch.ops.npu.npu_rotary_mul(k_pe, cos, sin).view( + bs, 1, self.rope_head_dim + ) # [bs, 1, d] + k = torch.cat([k_pe, k_nope.unsqueeze(1)], dim=-1) # [bs, 1, 128] + + else: + if envs.SGLANG_NPU_USE_MULTI_STREAM.get(): + indexer_weight_stream = get_indexer_weight_stream() + indexer_weight_stream.wait_stream(torch.npu.current_stream()) + with torch.npu.stream(indexer_weight_stream): + x = x.view(-1, self.hidden_size) + weights = self.weights_proj(x.float())[0].to(torch.bfloat16) + weights.record_stream(indexer_weight_stream) + weights_event = indexer_weight_stream.record_event() + else: + x = x.view(-1, self.hidden_size) + weights = self.weights_proj(x.float())[0].to(torch.bfloat16) + + q_lora = (q_lora, dynamic_scale) if dynamic_scale is not None else q_lora + q = self.wq_b(q_lora)[0] # [bs, 1536] @ [1536, 64 * 128] = [bs, 64 * 128] + q = q.view(bs, self.n_heads, self.head_dim) # [bs, 64, 128] + q_pe, q_nope = torch.split( + q, + [self.rope_head_dim, self.head_dim - self.rope_head_dim], + dim=-1, + ) # [bs, 64, 64 + 64] + + k_proj = self.wk(x)[0] # [b, s, 7168] @ [7168, 128] = [b, s, 128] + k = self.k_norm(k_proj) + k_pe, k_nope = torch.split( + k, + [self.rope_head_dim, self.head_dim - self.rope_head_dim], + dim=-1, + ) # [bs, 64 + 64] + + k_pe = k_pe.unsqueeze(1) + + if layer_id == 0: + self.rotary_emb.sin_cos_cache = ( + self.rotary_emb.cos_sin_cache.index_select(0, positions) + ) + + q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe) + k_pe = k_pe.squeeze(1) + q = torch.cat([q_pe, q_nope], dim=-1) + k = torch.cat([k_pe, k_nope], dim=-1) + + if ( + is_prefill + and self.dsa_enable_prefill_cp + and forward_batch.attn_cp_metadata is not None + ): + k = cp_all_gather_rerange_output( + k.contiguous().view(-1, self.head_dim), + self.cp_size, + forward_batch, + torch.npu.current_stream(), + ) + + get_token_to_kv_pool().set_index_k_buffer( + layer_id, forward_batch.out_cache_loc, k + ) + if is_prefill: + if ( + self.dsa_enable_prefill_cp + and forward_batch.attn_cp_metadata is not None + ): + get_attn_backend().forward_metadata.actual_seq_lengths_q = ( + forward_batch.attn_cp_metadata.actual_seq_q_prev_tensor, + forward_batch.attn_cp_metadata.actual_seq_q_next_tensor, + ) + if sum(forward_batch.extend_prefix_lens_cpu) > 0: + total_kv_len_prev_tensor = ( + forward_batch.attn_cp_metadata.kv_len_prev_tensor + + forward_batch.extend_prefix_lens.squeeze() + ) + total_kv_len_next_tensor = ( + forward_batch.attn_cp_metadata.kv_len_next_tensor + + forward_batch.extend_prefix_lens.squeeze() + ) + get_attn_backend().forward_metadata.actual_seq_lengths_kv = ( + total_kv_len_prev_tensor, + total_kv_len_next_tensor, + ) + else: + get_attn_backend().forward_metadata.actual_seq_lengths_kv = ( + forward_batch.attn_cp_metadata.kv_len_prev_tensor, + forward_batch.attn_cp_metadata.kv_len_next_tensor, + ) + actual_seq_lengths_q = ( + get_attn_backend().forward_metadata.actual_seq_lengths_q + ) + actual_seq_lengths_kv = ( + get_attn_backend().forward_metadata.actual_seq_lengths_kv + ) + else: + actual_seq_lengths_kv = forward_batch.seq_lens + actual_seq_lengths_q = forward_batch.extend_seq_lens.cumsum(dim=0) + else: + if get_attn_backend().forward_metadata.actual_seq_lengths_q is None: + if ( + forward_batch.forward_mode.is_draft_extend_v2() + or forward_batch.forward_mode.is_target_verify() + ): + num_draft_tokens = get_attn_backend().speculative_num_draft_tokens + actual_seq_lengths_q = torch.arange( + num_draft_tokens, + num_draft_tokens + bs, + num_draft_tokens, + dtype=torch.int32, + device=k.device, + ) + else: + actual_seq_lengths_q = torch.tensor( + [1 + i * 1 for i in range(bs)], + dtype=torch.int32, + device=k.device, + ) + else: + actual_seq_lengths_q = ( + get_attn_backend().forward_metadata.actual_seq_lengths_q + ) + + past_key_states = get_token_to_kv_pool().get_index_k_buffer(layer_id) + + if self.rotary_emb.is_neox_style and self.alt_stream is not None: + torch.npu.current_stream().wait_event(q_rope_event) + if envs.SGLANG_NPU_USE_MULTI_STREAM.get(): + torch.npu.current_stream().wait_event(weights_event) + if ( + _use_ag_after_qlora + and layer_scatter_modes.layer_input_mode == ScatterMode.SCATTERED + and layer_scatter_modes.attn_mode == ScatterMode.TP_ATTN_FULL + ): + weights = scattered_to_tp_attn_full(weights, forward_batch) + block_table = get_attn_backend().forward_metadata.block_tables + if ( + is_prefill + and self.dsa_enable_prefill_cp + and forward_batch.attn_cp_metadata is not None + ): + block_table = block_table[: actual_seq_lengths_q[0].numel()] + topk_indices = self.do_npu_cp_balance_indexer( + q.view(-1, self.n_heads, self.head_dim), + past_key_states, + weights, + actual_seq_lengths_q, + actual_seq_lengths_kv, + block_table, + ) + return topk_indices + else: + block_table = ( + block_table[: actual_seq_lengths_q.size()[0]] + if is_prefill + else block_table + ) + + topk_indices = torch_npu.npu_lightning_indexer( + query=q.view(-1, self.n_heads, self.head_dim), + key=past_key_states, + weights=weights, + actual_seq_lengths_query=actual_seq_lengths_q.to(torch.int32), + actual_seq_lengths_key=actual_seq_lengths_kv.to(k.device).to( + torch.int32 + ), + block_table=block_table, + layout_query="TND", + layout_key="PA_BSND", + sparse_count=self.index_topk, + sparse_mode=3, + ) + # Keep DSA top-k as [T, K]; NPU attention expands it when needed. + return topk_indices[0].squeeze(1) + + def do_npu_cp_balance_indexer( + self, + q, + past_key_states, + indexer_weights, + actual_seq_lengths_q, + actual_seq_lengths_kv, + block_table, + ): + q_prev, q_next = torch.split(q, (q.size(0) + 1) // 2, dim=0) + weights_prev, weights_next = None, None + if indexer_weights is not None: + weights_prev, weights_next = torch.split( + indexer_weights, (indexer_weights.size(0) + 1) // 2, dim=0 + ) + weights_prev = weights_prev.contiguous().view(-1, weights_prev.shape[-1]) + weights_next = weights_next.contiguous().view(-1, weights_next.shape[-1]) + + actual_seq_lengths_q_prev, actual_seq_lengths_q_next = actual_seq_lengths_q + actual_seq_lengths_kv_prev, actual_seq_lengths_kv_next = actual_seq_lengths_kv + + topk_indices_prev = torch_npu.npu_lightning_indexer( + query=q_prev, + key=past_key_states, + weights=weights_prev, + actual_seq_lengths_query=actual_seq_lengths_q_prev.to( + device=q.device, dtype=torch.int32 + ), + actual_seq_lengths_key=actual_seq_lengths_kv_prev.to( + device=q.device, dtype=torch.int32 + ), + block_table=block_table, + layout_query="TND", + layout_key="PA_BSND", + sparse_count=self.index_topk, + sparse_mode=3, + ) + topk_indices_next = torch_npu.npu_lightning_indexer( + query=q_next, + key=past_key_states, + weights=weights_next, + actual_seq_lengths_query=actual_seq_lengths_q_next.to( + device=q.device, dtype=torch.int32 + ), + actual_seq_lengths_key=actual_seq_lengths_kv_next.to( + device=q.device, dtype=torch.int32 + ), + block_table=block_table, + layout_query="TND", + layout_key="PA_BSND", + sparse_count=self.index_topk, + sparse_mode=3, + ) + return torch.cat([topk_indices_prev[0], topk_indices_next[0]], dim=0).squeeze(1) + + +def scattered_to_tp_attn_full( + hidden_states: torch.Tensor, + forward_batch, +) -> torch.Tensor: + hidden_states, local_hidden_states = ( + torch.empty( + (forward_batch.input_ids.shape[0], hidden_states.shape[1]), + dtype=hidden_states.dtype, + device=hidden_states.device, + ), + hidden_states, + ) + attn_tp_all_gather_into_tensor(hidden_states, local_hidden_states.contiguous()) + return hidden_states diff --git a/python/sglang/srt/layers/attention/dsa/dsa_prefill_cuda_graph.py b/python/sglang/srt/layers/attention/dsa/dsa_prefill_cuda_graph.py new file mode 100644 index 000000000..f90e7ba95 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsa/dsa_prefill_cuda_graph.py @@ -0,0 +1,191 @@ +from __future__ import annotations + +import torch + +from sglang.srt.compilation.compilation_config import register_split_op +from sglang.srt.model_executor.forward_context import get_attn_backend +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( + eager_on_graph, +) +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( + is_in_breakable_cuda_graph, +) +from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( + get_tc_piecewise_forward_context, + is_in_tc_piecewise_cuda_graph, +) +from sglang.srt.utils import is_cuda +from sglang.srt.utils.custom_op import register_custom_op + +_is_cuda = is_cuda() + +GRAPH_WEIGHTS_PROJ_LORA_ERROR = ( + "DSA indexer weights_proj LoRA is incompatible with " + "piecewise/breakable CUDA graph; remove the explicit " + "prefill cuda-graph backend override or drop " + "indexer.weights_proj from the LoRA target modules." +) + + +def _is_in_piecewise_or_breakable_cuda_graph() -> bool: + return is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph() + + +if _is_cuda: + + def _scale_head_gate_graph_fake_impl( + weights_raw: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + return torch.empty( + (weights_raw.shape[0], weights_raw.shape[1], q_scale.shape[-1]), + dtype=torch.float32, + device=weights_raw.device, + ) + + # In-graph (PCG/BCG) head gate for the fused path: weights_proj is folded + # into wk_weights_proj, so weights_raw is precomputed and there is no GEMM. + @register_custom_op(fake_impl=_scale_head_gate_graph_fake_impl) + def scale_head_gate_graph( + weights_raw: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + weights = weights_raw * n_heads_inv_sqrt + return weights.unsqueeze(-1) * q_scale * softmax_scale + + def _logits_head_gate_graph_fake_impl( + x: torch.Tensor, + weight: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + return torch.empty( + (x.shape[0], weight.shape[0], q_scale.shape[-1]), + dtype=torch.float32, + device=x.device, + ) + + # In-graph (PCG/BCG) head gate for the NON-prefill path + @register_custom_op(fake_impl=_logits_head_gate_graph_fake_impl) + def logits_head_gate_graph( + x: torch.Tensor, + weight: torch.Tensor, + n_heads_inv_sqrt: float, + softmax_scale: float, + q_scale: torch.Tensor, + ) -> torch.Tensor: + out = torch.mm(x, weight.t(), out_dtype=torch.float32) + weights = out * n_heads_inv_sqrt + weights = weights.unsqueeze(-1) * q_scale * softmax_scale + return weights + + +@register_custom_op(mutates_args=["topk_result"]) +@register_split_op() +def pcg_dsa_indexer_prefill_split( + layer_id: int, + x: torch.Tensor, + q_lora: torch.Tensor, + positions: torch.Tensor, + topk_result: torch.Tensor, +) -> None: + # Default in-graph indexer path for non-CP prefill: runs the whole indexer + # (q/k proj, head gate, k-cache store, topk) as one eager split op. PCG calls + # this as a split op; BCG uses the explicit eager wrapper below. + # + # Output contract (differs from the eager `forward` path): a split op returns + # None, so results are delivered only by mutating `topk_result` in place. The + # call site pre-allocates it at a static, padded shape and a downstream + # captured graph reads it at a fixed address; eager code instead allocates + # and returns a fresh, naturally-sized tensor each call. + assert _is_cuda, "Internal error: DSA graph dispatch is only supported on CUDA" + from sglang.kernels.ops.attention.dsa.triton_kernel import act_quant + + forward_context = get_tc_piecewise_forward_context() + forward_batch = forward_context.forward_batch + indexer = forward_context.dsa_indexers[layer_id] + metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch) + + extend_num_tokens = forward_batch.extend_num_tokens + # Empty buffer encodes return_indices=False for graph dispatch. + return_indices = topk_result.numel() != 0 + k_only = not return_indices or ( + indexer._should_skip_logits_computation(forward_batch) + and not indexer.dsa_enable_prefill_cp + ) + if k_only: + indexer._forward_cuda_k_only( + x, + positions, + forward_batch, + layer_id, + act_quant, + metadata=metadata, + return_indices=return_indices, + num_tokens=extend_num_tokens, + topk_result=topk_result, + ) + return + + # Fused path stores K (no-Hadamard) and computes q_fp8 + head gate in the + # fused kernels, sliced to the unpadded count. Single stream: the split op is + # captured, so the dual-stream overlap is disabled. + if indexer.use_dsa_indexer_fusion: + q_fp8, weights = indexer._fused_q_prepare_and_store( + x, + q_lora, + positions, + forward_batch, + layer_id, + act_quant, + num_tokens=extend_num_tokens, + enable_dual_stream=False, + ) + indexer._get_topk_ragged( + False, + forward_batch, + layer_id, + q_fp8, + weights, + metadata, + topk_result, + ) + return + + query, key, _ = indexer._get_q_k_bf16( + q_lora, + x, + positions, + enable_dual_stream=False, + forward_batch=forward_batch, + ) + q_fp8, q_scale = act_quant(query, indexer.block_size, indexer.scale_fmt) + # Reuse the compiled head-gate util shared with the eager path. + weights = indexer._get_logits_head_gate(x, q_scale) + # Store K cache + ragged top-k, sliced to the unpadded count and writing into + # the static padded topk_result buffer (the graph contract). Mirrors the eager + # path's store + _get_topk_ragged. + indexer._store_index_k_cache( + forward_batch=forward_batch, + layer_id=layer_id, + key=key[:extend_num_tokens], + act_quant=act_quant, + out_cache_loc=forward_batch.out_cache_loc[:extend_num_tokens], + ) + indexer._get_topk_ragged( + False, + forward_batch, + layer_id, + q_fp8[:extend_num_tokens], + weights, + metadata, + topk_result, + ) + + +bcg_dsa_indexer_prefill_split = eager_on_graph(True)(pcg_dsa_indexer_prefill_split) diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index c029f4496..37fcbe404 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -42,7 +42,7 @@ from sglang.srt.layers.attention.dsa.dsa_backend_mtp_precompute import ( PrecomputedMetadata, compute_cu_seqlens, ) -from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata +from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import DSAIndexerMetadata from sglang.srt.layers.attention.dsa.dsa_topk_backend import ( DSATopKBackend, TopkTransformMethod, @@ -271,88 +271,6 @@ def _cat(tensors: list[torch.Tensor], dim: int = -1) -> torch.Tensor: return _compiled_cat([qk_nope, qk_rope], dim=dim) -@dataclass(frozen=True) -class DSAIndexerMetadata(BaseIndexerMetadata): - attn_metadata: DSAMetadata - topk_transform_method: TopkTransformMethod - topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL - paged_mqa_schedule_metadata: Optional[torch.Tensor] = None - paged_mqa_ctx_lens_2d: Optional[torch.Tensor] = None - force_unfused_topk: bool = False - - def get_seqlens_int32(self) -> torch.Tensor: - return self.attn_metadata.cache_seqlens_int32 - - def get_page_table_64(self) -> torch.Tensor: - return self.attn_metadata.real_page_table - - def get_page_table_1(self) -> torch.Tensor: - return self.attn_metadata.page_table_1 - - def get_seqlens_expanded(self) -> torch.Tensor: - return self.attn_metadata.dsa_seqlens_expanded - - def get_cu_seqlens_k(self) -> torch.Tensor: - return self.attn_metadata.cu_seqlens_k - - def get_indexer_kvcache_range(self) -> Tuple[torch.Tensor, torch.Tensor]: - return self.attn_metadata.indexer_k_start_end - - def get_indexer_seq_len(self) -> torch.Tensor: - return self.attn_metadata.indexer_seq_lens - - def get_indexer_seq_len_cpu(self) -> torch.Tensor: - return self.attn_metadata.indexer_seq_lens_cpu - - def get_dsa_extend_len_cpu(self) -> List[int]: - return self.attn_metadata.dsa_extend_seq_lens_list - - def get_token_to_batch_idx(self) -> torch.Tensor: - return self.attn_metadata.token_to_batch_idx - - def topk_transform( - self, - logits: torch.Tensor, - topk: int, - ks: Optional[torch.Tensor] = None, - cu_seqlens_q: Optional[torch.Tensor] = None, - ke_offset: Optional[torch.Tensor] = None, - batch_idx_list: Optional[List[int]] = None, - topk_indices_offset_override: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - if topk_indices_offset_override is not None: - cu_topk_indices_offset = topk_indices_offset_override - cu_seqlens_q_topk = None - elif cu_seqlens_q is not None: - cu_seqlens_q = cu_seqlens_q.to(torch.int32) - cu_seqlens_q_topk = compute_cu_seqlens(cu_seqlens_q) - cu_topk_indices_offset = torch.repeat_interleave( - cu_seqlens_q_topk[:-1], - cu_seqlens_q, - # Avoid reading sum(cu_seqlens_q) back to the host. - output_size=logits.shape[0], - ) - else: - cu_seqlens_q_topk = self.attn_metadata.cu_seqlens_q - cu_topk_indices_offset = self.attn_metadata.topk_indices_offset - if ke_offset is not None: - seq_lens_topk = ke_offset - else: - seq_lens_topk = self.get_seqlens_expanded() - return self.topk_backend.topk_transform( - logits=logits, - lengths=seq_lens_topk, - topk=topk, - topk_transform_method=self.topk_transform_method, - attn_metadata=self.attn_metadata, - cu_seqlens_q_topk=cu_seqlens_q_topk, - topk_indices_offset=cu_topk_indices_offset, - row_starts=ks, - batch_idx_list=batch_idx_list, - force_unfused_topk=self.force_unfused_topk, - ) - - _DSA_IMPL_T: TypeAlias = Literal[ "flashmla_sparse", "flashmla_sparse_q8", diff --git a/python/sglang/srt/layers/attention/hybrid_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_attn_backend.py index 9a1ebe000..cd350b6d8 100644 --- a/python/sglang/srt/layers/attention/hybrid_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_attn_backend.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Optional import torch from sglang.srt.layers.attention.base_attn_backend import AttentionBackend -from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata +from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import BaseIndexerMetadata 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 diff --git a/test/registered/kernels/ops/attention/test_dsa_indexer.py b/test/registered/kernels/ops/attention/test_dsa_indexer.py index 2cbd4dd62..93d53b382 100644 --- a/test/registered/kernels/ops/attention/test_dsa_indexer.py +++ b/test/registered/kernels/ops/attention/test_dsa_indexer.py @@ -12,10 +12,10 @@ _parallel_override = get_parallel().override(attn_tp_size=1) _parallel_override.__enter__() from sglang.srt.configs.model_config import AttentionArch -from sglang.srt.layers.attention.dsa.dsa_indexer import ( +from sglang.srt.layers.attention.dsa.dsa_indexer import Indexer, rotate_activation +from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import ( BaseIndexerMetadata, - Indexer, - rotate_activation, + DSAIndexerMetadata, ) from sglang.srt.layers.attention.dsa.dsa_topk_backend import ( DSATopKBackend, @@ -23,7 +23,6 @@ from sglang.srt.layers.attention.dsa.dsa_topk_backend import ( ) from sglang.srt.layers.attention.dsa_backend import ( DeepseekSparseAttnBackend, - DSAIndexerMetadata, DSAMetadata, ) from sglang.srt.layers.layernorm import LayerNorm @@ -394,8 +393,7 @@ class TestDSAIndexer(CustomTestCase): # Pool refs + attn_backend are now resolved via the ForwardContext; # publish ``self.backend`` for the duration of this fixture call so - # ``get_attn_backend()`` / ``get_token_to_kv_pool()`` / - # ``get_req_to_token_pool()`` resolve correctly. + # ``get_attn_backend()`` / ``get_token_to_kv_pool()`` resolve correctly. from sglang.srt.model_executor.forward_context import ( ForwardContext, set_forward_context,