[Refactor] Clean up and split DSA indexer (#33443)

This commit is contained in:
Baizhou Zhang
2026-08-03 18:05:01 -07:00
committed by GitHub
parent 3960983753
commit 7f6a2e2b50
9 changed files with 753 additions and 836 deletions
@@ -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,
)
@@ -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
@@ -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
@@ -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,
)
@@ -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
@@ -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)
@@ -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",
@@ -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