support GLM-5.2 MTP index sharing with prefill CP (#30992)
Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
This commit is contained in:
co-authored by
Baizhou Zhang
parent
1b4176cc46
commit
7e229e2a81
@@ -1,3 +1,4 @@
|
|||||||
|
from itertools import accumulate
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -6,11 +7,39 @@ import triton.language as tl
|
|||||||
|
|
||||||
|
|
||||||
def transform_index_page_table_prefill(**kwargs):
|
def transform_index_page_table_prefill(**kwargs):
|
||||||
return transform_index_page_table_prefill_ref(**kwargs)
|
return transform_index_page_table_prefill_fast(**kwargs)
|
||||||
|
|
||||||
|
|
||||||
def transform_index_page_table_decode(**kwargs):
|
def transform_index_page_table_decode(**kwargs):
|
||||||
return transform_index_page_table_decode_ref(**kwargs)
|
return transform_index_page_table_decode_fast(**kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def _allocate_prefill_result(
|
||||||
|
topk_indices: torch.Tensor,
|
||||||
|
real_num_tokens: int,
|
||||||
|
output_num_tokens: Optional[int],
|
||||||
|
) -> torch.Tensor:
|
||||||
|
topk_num_tokens = topk_indices.shape[0]
|
||||||
|
if output_num_tokens is None:
|
||||||
|
output_num_tokens = topk_num_tokens
|
||||||
|
|
||||||
|
assert real_num_tokens <= topk_num_tokens, (
|
||||||
|
f"sum(extend_lens_cpu) ({real_num_tokens}) exceeds "
|
||||||
|
f"topk_indices rows ({topk_num_tokens})"
|
||||||
|
)
|
||||||
|
assert topk_num_tokens <= output_num_tokens, (
|
||||||
|
f"topk_indices rows ({topk_num_tokens}) exceeds "
|
||||||
|
f"output_num_tokens ({output_num_tokens})"
|
||||||
|
)
|
||||||
|
|
||||||
|
result = torch.empty(
|
||||||
|
(output_num_tokens, topk_indices.shape[1]),
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=topk_indices.device,
|
||||||
|
)
|
||||||
|
if real_num_tokens < output_num_tokens:
|
||||||
|
result[real_num_tokens:].fill_(-1)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
@@ -19,11 +48,11 @@ def transform_index_page_table_decode_kernel(
|
|||||||
topk_indices_ptr: torch.Tensor,
|
topk_indices_ptr: torch.Tensor,
|
||||||
result_ptr: torch.Tensor,
|
result_ptr: torch.Tensor,
|
||||||
page_size: tl.constexpr,
|
page_size: tl.constexpr,
|
||||||
max_seqlen_k: tl.constexpr,
|
page_table_row_stride: tl.constexpr,
|
||||||
):
|
):
|
||||||
TOPK: tl.constexpr = 2048
|
TOPK: tl.constexpr = 2048
|
||||||
req_id = tl.program_id(0)
|
req_id = tl.program_id(0)
|
||||||
page_table_ptr = page_table_ptr + req_id * max_seqlen_k
|
page_table_ptr = page_table_ptr + req_id * page_table_row_stride
|
||||||
topk_indices_ptr = topk_indices_ptr + req_id * TOPK
|
topk_indices_ptr = topk_indices_ptr + req_id * TOPK
|
||||||
result_ptr = result_ptr + req_id * TOPK
|
result_ptr = result_ptr + req_id * TOPK
|
||||||
|
|
||||||
@@ -35,6 +64,61 @@ def transform_index_page_table_decode_kernel(
|
|||||||
tl.store(result_ptr + offset, -1, mask=~mask)
|
tl.store(result_ptr + offset, -1, mask=~mask)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def transform_index_page_table_prefill_kernel(
|
||||||
|
page_table_ptr: torch.Tensor,
|
||||||
|
topk_indices_ptr: torch.Tensor,
|
||||||
|
cu_seqlens_q_ptr: torch.Tensor,
|
||||||
|
result_ptr: torch.Tensor,
|
||||||
|
page_table_stride_0: tl.constexpr,
|
||||||
|
page_table_stride_1: tl.constexpr,
|
||||||
|
topk_indices_stride_0: tl.constexpr,
|
||||||
|
topk_indices_stride_1: tl.constexpr,
|
||||||
|
result_stride_0: tl.constexpr,
|
||||||
|
result_stride_1: tl.constexpr,
|
||||||
|
PAGE_TABLE_IS_EXPANDED: tl.constexpr,
|
||||||
|
TOPK: tl.constexpr,
|
||||||
|
BLOCK_Q: tl.constexpr,
|
||||||
|
BLOCK_TOPK: tl.constexpr,
|
||||||
|
):
|
||||||
|
request_id = tl.program_id(0)
|
||||||
|
query_offsets = tl.program_id(1) * BLOCK_Q + tl.arange(0, BLOCK_Q)
|
||||||
|
topk_offsets = tl.program_id(2) * BLOCK_TOPK + tl.arange(0, BLOCK_TOPK)
|
||||||
|
|
||||||
|
query_start = tl.load(cu_seqlens_q_ptr + request_id)
|
||||||
|
query_end = tl.load(cu_seqlens_q_ptr + request_id + 1)
|
||||||
|
token_indices = query_start + query_offsets
|
||||||
|
mask = (token_indices[:, None] < query_end) & (topk_offsets[None, :] < TOPK)
|
||||||
|
|
||||||
|
loaded_topk_indices = tl.load(
|
||||||
|
topk_indices_ptr
|
||||||
|
+ token_indices[:, None] * topk_indices_stride_0
|
||||||
|
+ topk_offsets[None, :] * topk_indices_stride_1,
|
||||||
|
mask=mask,
|
||||||
|
other=-1,
|
||||||
|
)
|
||||||
|
valid_topk_mask = mask & (loaded_topk_indices >= 0)
|
||||||
|
|
||||||
|
if PAGE_TABLE_IS_EXPANDED:
|
||||||
|
page_table_rows = token_indices
|
||||||
|
else:
|
||||||
|
page_table_rows = token_indices * 0 + request_id
|
||||||
|
loaded_kv_indices = tl.load(
|
||||||
|
page_table_ptr
|
||||||
|
+ page_table_rows[:, None] * page_table_stride_0
|
||||||
|
+ loaded_topk_indices * page_table_stride_1,
|
||||||
|
mask=valid_topk_mask,
|
||||||
|
other=-1,
|
||||||
|
)
|
||||||
|
tl.store(
|
||||||
|
result_ptr
|
||||||
|
+ token_indices[:, None] * result_stride_0
|
||||||
|
+ topk_offsets[None, :] * result_stride_1,
|
||||||
|
loaded_kv_indices,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def transform_index_page_table_decode_fast(
|
def transform_index_page_table_decode_fast(
|
||||||
page_table: torch.Tensor,
|
page_table: torch.Tensor,
|
||||||
topk_indices: torch.Tensor,
|
topk_indices: torch.Tensor,
|
||||||
@@ -54,7 +138,6 @@ def transform_index_page_table_decode_fast(
|
|||||||
assert page_table.shape[0] == topk_indices.shape[0]
|
assert page_table.shape[0] == topk_indices.shape[0]
|
||||||
assert topk_indices.shape[1] == 2048
|
assert topk_indices.shape[1] == 2048
|
||||||
qo_len = topk_indices.shape[0]
|
qo_len = topk_indices.shape[0]
|
||||||
max_seqlen_k = page_table.shape[1]
|
|
||||||
if result is None:
|
if result is None:
|
||||||
result = torch.empty_like(topk_indices, dtype=torch.int32)
|
result = torch.empty_like(topk_indices, dtype=torch.int32)
|
||||||
# Launch triton kernel
|
# Launch triton kernel
|
||||||
@@ -64,7 +147,7 @@ def transform_index_page_table_decode_fast(
|
|||||||
topk_indices,
|
topk_indices,
|
||||||
result,
|
result,
|
||||||
page_size,
|
page_size,
|
||||||
max_seqlen_k=max_seqlen_k,
|
page_table_row_stride=page_table.stride(0),
|
||||||
)
|
)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
@@ -74,20 +157,48 @@ def transform_index_page_table_prefill_fast(
|
|||||||
topk_indices: torch.Tensor,
|
topk_indices: torch.Tensor,
|
||||||
extend_lens_cpu: List[int],
|
extend_lens_cpu: List[int],
|
||||||
page_size: int = 1,
|
page_size: int = 1,
|
||||||
|
output_num_tokens: Optional[int] = None,
|
||||||
|
page_table_is_expanded: bool = False,
|
||||||
|
cu_seqlens_q: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
# TODO(baizhou): can be implemented with another triton kernel
|
|
||||||
assert page_size == 1
|
assert page_size == 1
|
||||||
result = torch.empty_like(topk_indices, dtype=torch.int32)
|
assert topk_indices.shape[1] == 2048
|
||||||
assert len(extend_lens_cpu) == page_table.shape[0]
|
real_num_tokens = sum(extend_lens_cpu)
|
||||||
offset = 0
|
result = _allocate_prefill_result(topk_indices, real_num_tokens, output_num_tokens)
|
||||||
for i, l in enumerate(extend_lens_cpu):
|
if real_num_tokens == 0:
|
||||||
transform_index_page_table_decode_fast(
|
return result
|
||||||
page_table[i].unsqueeze(0).expand(l, -1),
|
|
||||||
topk_indices[offset : offset + l],
|
max_extend_len = max(extend_lens_cpu)
|
||||||
result=result[offset : offset + l],
|
block_q = 1 if max_extend_len == 1 else 2 if max_extend_len == 2 else 4
|
||||||
|
block_topk = 256
|
||||||
|
if cu_seqlens_q is None:
|
||||||
|
cu_seqlens_q = torch.tensor(
|
||||||
|
[0, *accumulate(extend_lens_cpu)],
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=topk_indices.device,
|
||||||
|
)
|
||||||
|
grid = (
|
||||||
|
cu_seqlens_q.shape[0] - 1,
|
||||||
|
triton.cdiv(max_extend_len, block_q),
|
||||||
|
triton.cdiv(topk_indices.shape[1], block_topk),
|
||||||
|
)
|
||||||
|
transform_index_page_table_prefill_kernel[grid](
|
||||||
|
page_table,
|
||||||
|
topk_indices,
|
||||||
|
cu_seqlens_q,
|
||||||
|
result,
|
||||||
|
page_table.stride(0),
|
||||||
|
page_table.stride(1),
|
||||||
|
topk_indices.stride(0),
|
||||||
|
topk_indices.stride(1),
|
||||||
|
result.stride(0),
|
||||||
|
result.stride(1),
|
||||||
|
PAGE_TABLE_IS_EXPANDED=page_table_is_expanded,
|
||||||
|
TOPK=topk_indices.shape[1],
|
||||||
|
BLOCK_Q=block_q,
|
||||||
|
BLOCK_TOPK=block_topk,
|
||||||
|
num_warps=4,
|
||||||
)
|
)
|
||||||
offset += l
|
|
||||||
assert offset == topk_indices.shape[0]
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
@@ -117,10 +228,22 @@ def transform_index_page_table_prefill_ref(
|
|||||||
topk_indices: torch.Tensor,
|
topk_indices: torch.Tensor,
|
||||||
extend_lens_cpu: List[int],
|
extend_lens_cpu: List[int],
|
||||||
page_size: int = 1,
|
page_size: int = 1,
|
||||||
|
output_num_tokens: Optional[int] = None,
|
||||||
|
page_table_is_expanded: bool = False,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
assert page_size == 1
|
assert page_size == 1
|
||||||
result = torch.empty_like(topk_indices, dtype=torch.int32)
|
real_num_tokens = sum(extend_lens_cpu)
|
||||||
assert len(extend_lens_cpu) == page_table.shape[0]
|
result = _allocate_prefill_result(topk_indices, real_num_tokens, output_num_tokens)
|
||||||
|
|
||||||
|
if page_table_is_expanded:
|
||||||
|
if real_num_tokens > 0:
|
||||||
|
transform_index_page_table_decode_ref(
|
||||||
|
page_table[:real_num_tokens],
|
||||||
|
topk_indices[:real_num_tokens],
|
||||||
|
result=result[:real_num_tokens],
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
offset = 0
|
offset = 0
|
||||||
for i, l in enumerate(extend_lens_cpu):
|
for i, l in enumerate(extend_lens_cpu):
|
||||||
transform_index_page_table_decode_ref(
|
transform_index_page_table_decode_ref(
|
||||||
@@ -129,7 +252,6 @@ def transform_index_page_table_prefill_ref(
|
|||||||
result=result[offset : offset + l],
|
result=result[offset : offset + l],
|
||||||
)
|
)
|
||||||
offset += l
|
offset += l
|
||||||
assert offset == topk_indices.shape[0]
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1898,19 +1898,20 @@ class DeepseekSparseAttnBackend(
|
|||||||
q_nope = q_all[:, :, : layer.v_head_dim]
|
q_nope = q_all[:, :, : layer.v_head_dim]
|
||||||
q_rope = q_all[:, :, layer.v_head_dim :]
|
q_rope = q_all[:, :, layer.v_head_dim :]
|
||||||
|
|
||||||
# Align topk_indices with q dimensions
|
|
||||||
# This handles cases where q is padded (TP + partial DP attention)
|
|
||||||
if topk_indices is not None:
|
|
||||||
topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0])
|
|
||||||
|
|
||||||
# NOTE(dark): here, we use page size = 1
|
# NOTE(dark): here, we use page size = 1
|
||||||
topk_transform_method = self.get_topk_transform_method(
|
topk_transform_method = self.get_topk_transform_method(
|
||||||
forward_batch.forward_mode
|
forward_batch.forward_mode
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.use_fused_topk:
|
if self.use_fused_topk:
|
||||||
|
if topk_indices is not None:
|
||||||
|
topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0])
|
||||||
|
|
||||||
page_table_1 = self._get_fused_topk_page_table(topk_indices)
|
page_table_1 = self._get_fused_topk_page_table(topk_indices)
|
||||||
else:
|
else:
|
||||||
if topk_transform_method == TopkTransformMethod.RAGGED:
|
if topk_transform_method == TopkTransformMethod.RAGGED:
|
||||||
|
if topk_indices is not None:
|
||||||
|
topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0])
|
||||||
topk_indices_offset = metadata.topk_indices_offset
|
topk_indices_offset = metadata.topk_indices_offset
|
||||||
assert topk_indices_offset is not None
|
assert topk_indices_offset is not None
|
||||||
mask = topk_indices != -1
|
mask = topk_indices != -1
|
||||||
@@ -1929,6 +1930,12 @@ class DeepseekSparseAttnBackend(
|
|||||||
topk_indices=topk_indices,
|
topk_indices=topk_indices,
|
||||||
extend_lens_cpu=metadata.dsa_extend_seq_lens_list,
|
extend_lens_cpu=metadata.dsa_extend_seq_lens_list,
|
||||||
page_size=1,
|
page_size=1,
|
||||||
|
output_num_tokens=q_nope.shape[0],
|
||||||
|
page_table_is_expanded=(
|
||||||
|
forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
),
|
||||||
|
cu_seqlens_q=metadata.cu_seqlens_q,
|
||||||
)
|
)
|
||||||
|
|
||||||
# todo hisparse: to cover more backends
|
# todo hisparse: to cover more backends
|
||||||
@@ -2675,11 +2682,9 @@ class DeepseekSparseAttnBackend(
|
|||||||
else:
|
else:
|
||||||
q_all = q.view(-1, layer.tp_q_head_num, layer.head_dim)
|
q_all = q.view(-1, layer.tp_q_head_num, layer.head_dim)
|
||||||
|
|
||||||
# Align topk_indices with q dimensions
|
if self.use_fused_topk:
|
||||||
if topk_indices is not None:
|
if topk_indices is not None:
|
||||||
topk_indices = self._pad_topk_indices(topk_indices, q.shape[0])
|
topk_indices = self._pad_topk_indices(topk_indices, q.shape[0])
|
||||||
|
|
||||||
if self.use_fused_topk:
|
|
||||||
page_table_1 = self._get_fused_topk_page_table(topk_indices)
|
page_table_1 = self._get_fused_topk_page_table(topk_indices)
|
||||||
elif is_prefill:
|
elif is_prefill:
|
||||||
page_table_1 = transform_index_page_table_prefill(
|
page_table_1 = transform_index_page_table_prefill(
|
||||||
@@ -2687,8 +2692,16 @@ class DeepseekSparseAttnBackend(
|
|||||||
topk_indices=topk_indices,
|
topk_indices=topk_indices,
|
||||||
extend_lens_cpu=metadata.dsa_extend_seq_lens_list,
|
extend_lens_cpu=metadata.dsa_extend_seq_lens_list,
|
||||||
page_size=1,
|
page_size=1,
|
||||||
|
output_num_tokens=q.shape[0],
|
||||||
|
page_table_is_expanded=(
|
||||||
|
forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
),
|
||||||
|
cu_seqlens_q=metadata.cu_seqlens_q,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
if topk_indices is not None:
|
||||||
|
topk_indices = self._pad_topk_indices(topk_indices, q.shape[0])
|
||||||
page_table_1 = transform_index_page_table_decode(
|
page_table_1 = transform_index_page_table_decode(
|
||||||
page_table=metadata.page_table_1,
|
page_table=metadata.page_table_1,
|
||||||
topk_indices=topk_indices,
|
topk_indices=topk_indices,
|
||||||
|
|||||||
@@ -33,6 +33,7 @@ from sglang.srt.layers.attention.dsa.utils import (
|
|||||||
can_dsa_cp_split,
|
can_dsa_cp_split,
|
||||||
dsa_use_prefill_cp,
|
dsa_use_prefill_cp,
|
||||||
is_dsa_enable_prefill_cp,
|
is_dsa_enable_prefill_cp,
|
||||||
|
is_dsa_prefill_cp_round_robin_split,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
from sglang.srt.layers.linear import ReplicatedLinear
|
from sglang.srt.layers.linear import ReplicatedLinear
|
||||||
@@ -60,6 +61,34 @@ from sglang.srt.models.utils import WeightsMapper
|
|||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||||
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
|
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
|
||||||
|
|
||||||
|
|
||||||
|
def _gather_dsa_topk_indices_for_cp(
|
||||||
|
topk_indices: torch.Tensor,
|
||||||
|
local_num_tokens: int,
|
||||||
|
cp_size: int,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
stream,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if (
|
||||||
|
is_dsa_prefill_cp_round_robin_split()
|
||||||
|
and topk_indices.shape[0] < local_num_tokens
|
||||||
|
):
|
||||||
|
pad_rows = local_num_tokens - topk_indices.shape[0]
|
||||||
|
topk_indices = torch.cat(
|
||||||
|
[
|
||||||
|
topk_indices,
|
||||||
|
topk_indices.new_full((pad_rows, topk_indices.shape[1]), -1),
|
||||||
|
],
|
||||||
|
dim=0,
|
||||||
|
)
|
||||||
|
return cp_all_gather_rerange_output(
|
||||||
|
topk_indices,
|
||||||
|
cp_size,
|
||||||
|
forward_batch,
|
||||||
|
stream,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -222,12 +251,21 @@ class DeepseekModelNextN(nn.Module):
|
|||||||
else:
|
else:
|
||||||
hidden_states = self.eh_proj(eh_input)
|
hidden_states = self.eh_proj(eh_input)
|
||||||
|
|
||||||
if dsa_use_prefill_cp(
|
use_cp = dsa_use_prefill_cp(
|
||||||
forward_batch, self.dsa_enable_prefill_cp
|
forward_batch, self.dsa_enable_prefill_cp
|
||||||
) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp):
|
) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp)
|
||||||
|
if use_cp:
|
||||||
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
|
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
|
||||||
positions = cp_split_and_rebuild_position(forward_batch, positions)
|
positions = cp_split_and_rebuild_position(forward_batch, positions)
|
||||||
residual = None
|
residual = None
|
||||||
|
seed_buf = (
|
||||||
|
forward_batch.spec_info.dsa_seed_topk_capture
|
||||||
|
if forward_batch.forward_mode.is_extend(include_draft_extend_v2=True)
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
should_update_dsa_topk_indices = (
|
||||||
|
forward_batch.reuse_dsa_topk_indices or seed_buf is not None
|
||||||
|
)
|
||||||
with get_global_expert_distribution_recorder().disable_this_region():
|
with get_global_expert_distribution_recorder().disable_this_region():
|
||||||
hidden_states, residual, topk_indices = self.decoder(
|
hidden_states, residual, topk_indices = self.decoder(
|
||||||
positions,
|
positions,
|
||||||
@@ -241,34 +279,39 @@ class DeepseekModelNextN(nn.Module):
|
|||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
if forward_batch.reuse_dsa_topk_indices:
|
|
||||||
forward_batch.spec_info.dsa_topk_indices = topk_indices
|
|
||||||
|
|
||||||
# MTP IndexShare: on draft-extend, publish the last-token DSA
|
|
||||||
# indexer top-k to seed (avoid recomputing in) the draft-decode loop.
|
|
||||||
if forward_batch.forward_mode.is_extend(include_draft_extend_v2=True):
|
|
||||||
seed_buf = forward_batch.spec_info.dsa_seed_topk_capture
|
|
||||||
if seed_buf is not None and topk_indices is not None:
|
|
||||||
sel = forward_batch.spec_info.dsa_seed_topk_select
|
|
||||||
src = topk_indices if sel is None else topk_indices[sel]
|
|
||||||
seed_buf[: src.shape[0]].copy_(src)
|
|
||||||
|
|
||||||
if not forward_batch.forward_mode.is_idle():
|
if not forward_batch.forward_mode.is_idle():
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
hidden_states, _ = self.shared_head.norm(hidden_states, residual)
|
hidden_states, _ = self.shared_head.norm(hidden_states, residual)
|
||||||
else:
|
else:
|
||||||
hidden_states = self.shared_head.norm(hidden_states)
|
hidden_states = self.shared_head.norm(hidden_states)
|
||||||
|
|
||||||
if dsa_use_prefill_cp(
|
if use_cp:
|
||||||
forward_batch, self.dsa_enable_prefill_cp
|
local_num_tokens = hidden_states.shape[0]
|
||||||
) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp):
|
|
||||||
# allgather + rerrange
|
|
||||||
hidden_states = cp_all_gather_rerange_output(
|
hidden_states = cp_all_gather_rerange_output(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
self.cp_size,
|
self.cp_size,
|
||||||
forward_batch,
|
forward_batch,
|
||||||
torch.cuda.current_stream(),
|
torch.cuda.current_stream(),
|
||||||
)
|
)
|
||||||
|
if should_update_dsa_topk_indices and topk_indices is not None:
|
||||||
|
topk_indices = _gather_dsa_topk_indices_for_cp(
|
||||||
|
topk_indices,
|
||||||
|
local_num_tokens,
|
||||||
|
self.cp_size,
|
||||||
|
forward_batch,
|
||||||
|
torch.cuda.current_stream(),
|
||||||
|
)
|
||||||
|
if should_update_dsa_topk_indices and topk_indices is not None:
|
||||||
|
if forward_batch.reuse_dsa_topk_indices:
|
||||||
|
forward_batch.spec_info.dsa_topk_indices = topk_indices
|
||||||
|
if seed_buf is not None:
|
||||||
|
sel = forward_batch.spec_info.dsa_seed_topk_select
|
||||||
|
src = (
|
||||||
|
topk_indices[: seed_buf.shape[0]]
|
||||||
|
if sel is None
|
||||||
|
else topk_indices[sel]
|
||||||
|
)
|
||||||
|
seed_buf[: src.shape[0]].copy_(src)
|
||||||
finally:
|
finally:
|
||||||
exit_stack.close()
|
exit_stack.close()
|
||||||
|
|
||||||
|
|||||||
@@ -5204,7 +5204,7 @@ class ServerArgs:
|
|||||||
|
|
||||||
mode = strategy_to_legacy_mode[self.cp_strategy]
|
mode = strategy_to_legacy_mode[self.cp_strategy]
|
||||||
use_dsa_legacy_aliases = self.enable_dsa_prefill_context_parallel or getattr(
|
use_dsa_legacy_aliases = self.enable_dsa_prefill_context_parallel or getattr(
|
||||||
self, "attention_backend", None
|
self._resolved(), "attention_backend", None
|
||||||
) in ("dsa", "dsv4")
|
) in ("dsa", "dsv4")
|
||||||
if use_dsa_legacy_aliases:
|
if use_dsa_legacy_aliases:
|
||||||
self.enable_dsa_prefill_context_parallel = True
|
self.enable_dsa_prefill_context_parallel = True
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner i
|
|||||||
)
|
)
|
||||||
from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner
|
from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner
|
||||||
from sglang.srt.kv_canary.runner.canary_manager import context_tuple
|
from sglang.srt.kv_canary.runner.canary_manager import context_tuple
|
||||||
from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp
|
|
||||||
from sglang.srt.layers.attention.flashinfer_backend import FlashInferAttnBackend
|
from sglang.srt.layers.attention.flashinfer_backend import FlashInferAttnBackend
|
||||||
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
|
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
|
||||||
from sglang.srt.layers.attention.triton_backend import TritonAttnBackend
|
from sglang.srt.layers.attention.triton_backend import TritonAttnBackend
|
||||||
@@ -825,11 +824,9 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
|
|
||||||
# Seed the first draft-decode loop from each request's last prefill
|
# Seed the first draft-decode loop from each request's last prefill
|
||||||
# position. Gather last-per-req before the copy (prefill can be long).
|
# position. Gather last-per-req before the copy (prefill can be long).
|
||||||
# Skipped under context-parallel prefill (token layout wouldn't match).
|
|
||||||
seed_from_extend = (
|
seed_from_extend = (
|
||||||
self.seed_dsa_topk_from_draft_extend
|
self.seed_dsa_topk_from_draft_extend
|
||||||
and not forward_batch.forward_mode.is_idle()
|
and not forward_batch.forward_mode.is_idle()
|
||||||
and not dsa_use_prefill_cp(forward_batch)
|
|
||||||
)
|
)
|
||||||
if seed_from_extend:
|
if seed_from_extend:
|
||||||
bs = forward_batch.batch_size
|
bs = forward_batch.batch_size
|
||||||
|
|||||||
@@ -0,0 +1,63 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
popen_launch_server,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=400, stage="extra-b", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
|
GLM52_NVFP4_MODEL_PATH = "nvidia/GLM-5.2-NVFP4"
|
||||||
|
|
||||||
|
|
||||||
|
class TestGLM52CPInterleave(GSM8KMixin, CustomTestCase):
|
||||||
|
gsm8k_accuracy_thres = 0.935
|
||||||
|
gsm8k_num_examples = 500
|
||||||
|
gsm8k_num_threads = 32
|
||||||
|
gsm8k_num_shots = 20
|
||||||
|
gsm8k_accept_length_thres = 3
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = GLM52_NVFP4_MODEL_PATH
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--tp",
|
||||||
|
"4",
|
||||||
|
"--attn-cp-size",
|
||||||
|
"4",
|
||||||
|
"--enable-prefill-cp",
|
||||||
|
"--cp-strategy",
|
||||||
|
"interleave",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"EAGLE",
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"3",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"1",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"4",
|
||||||
|
"--mem-frac",
|
||||||
|
"0.85",
|
||||||
|
"--model-loader-extra-config",
|
||||||
|
'{"enable_multithread_load": true, "num_threads": 64}',
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "process") and cls.process:
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,240 @@
|
|||||||
|
import unittest
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
import sglang.srt.layers.attention.dsa.transform_index as transform_index_module
|
||||||
|
from sglang.srt.layers.attention.dsa.transform_index import (
|
||||||
|
transform_index_page_table_decode_fast,
|
||||||
|
transform_index_page_table_prefill_fast,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=60, stage="base-b", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
TOPK = 2048
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "CUDA is required for this test.")
|
||||||
|
class TestDSATransformIndex(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
super().setUp()
|
||||||
|
self.device = torch.device("cuda")
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
super().tearDown()
|
||||||
|
|
||||||
|
def _make_page_table(self, rows: int, context_length: int) -> torch.Tensor:
|
||||||
|
columns = torch.arange(context_length, dtype=torch.int32, device=self.device)
|
||||||
|
row_bias = (
|
||||||
|
torch.arange(rows, dtype=torch.int32, device=self.device).unsqueeze(1) * 17
|
||||||
|
)
|
||||||
|
return columns.unsqueeze(0) + row_bias
|
||||||
|
|
||||||
|
def _make_topk(self, rows: int, context_length: int) -> torch.Tensor:
|
||||||
|
topk = (
|
||||||
|
torch.arange(TOPK, dtype=torch.int64, device=self.device)
|
||||||
|
.remainder(context_length)
|
||||||
|
.repeat(rows, 1)
|
||||||
|
)
|
||||||
|
if rows > 0:
|
||||||
|
topk[:, 0] = 0
|
||||||
|
topk[:, 1] = context_length - 1
|
||||||
|
topk[:, 257::257] = -1
|
||||||
|
return topk
|
||||||
|
|
||||||
|
def _expected(
|
||||||
|
self,
|
||||||
|
page_table: torch.Tensor,
|
||||||
|
topk_indices: torch.Tensor,
|
||||||
|
extend_lens_cpu: list[int],
|
||||||
|
output_num_tokens: int,
|
||||||
|
page_table_is_expanded: bool,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
real_num_tokens = sum(extend_lens_cpu)
|
||||||
|
expected = torch.full(
|
||||||
|
(output_num_tokens, TOPK),
|
||||||
|
-1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
if real_num_tokens == 0:
|
||||||
|
return expected
|
||||||
|
|
||||||
|
if page_table_is_expanded:
|
||||||
|
source_rows = page_table[:real_num_tokens]
|
||||||
|
else:
|
||||||
|
request_ids = torch.repeat_interleave(
|
||||||
|
torch.arange(
|
||||||
|
len(extend_lens_cpu), dtype=torch.int64, device=self.device
|
||||||
|
),
|
||||||
|
torch.tensor(extend_lens_cpu, dtype=torch.int64, device=self.device),
|
||||||
|
)
|
||||||
|
source_rows = page_table[request_ids]
|
||||||
|
|
||||||
|
real_topk = topk_indices[:real_num_tokens]
|
||||||
|
torch.gather(
|
||||||
|
source_rows,
|
||||||
|
dim=1,
|
||||||
|
index=real_topk.clamp(min=0),
|
||||||
|
out=expected[:real_num_tokens],
|
||||||
|
)
|
||||||
|
expected[:real_num_tokens][real_topk < 0] = -1
|
||||||
|
return expected
|
||||||
|
|
||||||
|
def _check_decode_case(
|
||||||
|
self,
|
||||||
|
batch_size: int,
|
||||||
|
context_length: int,
|
||||||
|
*,
|
||||||
|
zero_row_stride: bool = False,
|
||||||
|
provide_result: bool = False,
|
||||||
|
) -> None:
|
||||||
|
if zero_row_stride:
|
||||||
|
page_table = self._make_page_table(1, context_length).expand(batch_size, -1)
|
||||||
|
else:
|
||||||
|
page_table = self._make_page_table(batch_size, context_length)
|
||||||
|
topk_indices = self._make_topk(batch_size, context_length)
|
||||||
|
expected = torch.empty(
|
||||||
|
(batch_size, TOPK), dtype=torch.int32, device=self.device
|
||||||
|
)
|
||||||
|
torch.gather(
|
||||||
|
page_table,
|
||||||
|
dim=1,
|
||||||
|
index=topk_indices.clamp(min=0),
|
||||||
|
out=expected,
|
||||||
|
)
|
||||||
|
expected[topk_indices < 0] = -1
|
||||||
|
result = torch.empty_like(expected) if provide_result else None
|
||||||
|
|
||||||
|
actual = transform_index_page_table_decode_fast(
|
||||||
|
page_table=page_table,
|
||||||
|
topk_indices=topk_indices,
|
||||||
|
result=result,
|
||||||
|
)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
if result is not None:
|
||||||
|
self.assertIs(actual, result)
|
||||||
|
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
|
||||||
|
|
||||||
|
def _check_case(
|
||||||
|
self,
|
||||||
|
extend_lens_cpu: list[int],
|
||||||
|
context_length: int,
|
||||||
|
*,
|
||||||
|
page_table_is_expanded: bool,
|
||||||
|
topk_padding: int = 0,
|
||||||
|
output_padding: int = 0,
|
||||||
|
) -> None:
|
||||||
|
real_num_tokens = sum(extend_lens_cpu)
|
||||||
|
page_table_rows = (
|
||||||
|
real_num_tokens if page_table_is_expanded else len(extend_lens_cpu)
|
||||||
|
)
|
||||||
|
topk_num_tokens = real_num_tokens + topk_padding
|
||||||
|
output_num_tokens = topk_num_tokens + output_padding
|
||||||
|
page_table = self._make_page_table(page_table_rows, context_length)
|
||||||
|
topk_indices = self._make_topk(topk_num_tokens, context_length)
|
||||||
|
expected = self._expected(
|
||||||
|
page_table,
|
||||||
|
topk_indices,
|
||||||
|
extend_lens_cpu,
|
||||||
|
output_num_tokens,
|
||||||
|
page_table_is_expanded,
|
||||||
|
)
|
||||||
|
|
||||||
|
actual = transform_index_page_table_prefill_fast(
|
||||||
|
page_table=page_table,
|
||||||
|
topk_indices=topk_indices,
|
||||||
|
extend_lens_cpu=extend_lens_cpu,
|
||||||
|
output_num_tokens=output_num_tokens,
|
||||||
|
page_table_is_expanded=page_table_is_expanded,
|
||||||
|
)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
|
||||||
|
|
||||||
|
def test_prefill_uses_dedicated_kernel(self):
|
||||||
|
extend_lens_cpu = [2, 1]
|
||||||
|
context_length = 4096
|
||||||
|
page_table = self._make_page_table(len(extend_lens_cpu), context_length)
|
||||||
|
topk_indices = self._make_topk(sum(extend_lens_cpu), context_length)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
transform_index_module,
|
||||||
|
"transform_index_page_table_decode_fast",
|
||||||
|
side_effect=AssertionError("prefill must not launch decode per request"),
|
||||||
|
):
|
||||||
|
transform_index_page_table_prefill_fast(
|
||||||
|
page_table=page_table,
|
||||||
|
topk_indices=topk_indices,
|
||||||
|
extend_lens_cpu=extend_lens_cpu,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_prefill_uses_device_cu_seqlens(self):
|
||||||
|
extend_lens_cpu = [2, 1]
|
||||||
|
context_length = 4096
|
||||||
|
page_table = self._make_page_table(len(extend_lens_cpu), context_length)
|
||||||
|
topk_indices = self._make_topk(sum(extend_lens_cpu), context_length)
|
||||||
|
cu_seqlens_q = torch.tensor([0, 2, 3], dtype=torch.int32, device=self.device)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
transform_index_module.torch,
|
||||||
|
"tensor",
|
||||||
|
side_effect=AssertionError("must reuse device-side metadata"),
|
||||||
|
):
|
||||||
|
transform_index_page_table_prefill_fast(
|
||||||
|
page_table=page_table,
|
||||||
|
topk_indices=topk_indices,
|
||||||
|
extend_lens_cpu=extend_lens_cpu,
|
||||||
|
cu_seqlens_q=cu_seqlens_q,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_mixed_lengths_padding_and_empty_batch(self):
|
||||||
|
self._check_case(
|
||||||
|
[0, 3, 1, 0, 4],
|
||||||
|
8192,
|
||||||
|
page_table_is_expanded=False,
|
||||||
|
topk_padding=5,
|
||||||
|
output_padding=7,
|
||||||
|
)
|
||||||
|
self._check_case(
|
||||||
|
[0, 0],
|
||||||
|
16,
|
||||||
|
page_table_is_expanded=False,
|
||||||
|
output_padding=8,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_large_batch_size(self):
|
||||||
|
self._check_case(
|
||||||
|
[1] * 8192,
|
||||||
|
4096,
|
||||||
|
page_table_is_expanded=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_large_context_lengths(self):
|
||||||
|
for context_length, page_table_is_expanded in (
|
||||||
|
(640_000, True),
|
||||||
|
(1_000_000, False),
|
||||||
|
):
|
||||||
|
with self.subTest(
|
||||||
|
context_length=context_length,
|
||||||
|
page_table_is_expanded=page_table_is_expanded,
|
||||||
|
):
|
||||||
|
self._check_case(
|
||||||
|
[2, 1],
|
||||||
|
context_length,
|
||||||
|
page_table_is_expanded=page_table_is_expanded,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_decode_fast_correctness_and_strides(self):
|
||||||
|
self._check_decode_case(17, 8192, provide_result=True)
|
||||||
|
self._check_decode_case(17, 8192, zero_row_stride=True)
|
||||||
|
|
||||||
|
def test_decode_fast_extreme_shapes(self):
|
||||||
|
self._check_decode_case(8192, 4096)
|
||||||
|
self._check_decode_case(2, 1_000_000)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user