[Refactor] Clean up and split DSA indexer (#33443)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user