[Bugfix] [DSA] [Hisparse] Broadcast TP Rank 0 Topk Indexes to other TPs (#24654)
Co-authored-by: xz-keg <xuzou_keg@outlook.com> Co-authored-by: xuzou <xu.zou@aminer.cn>
This commit is contained in:
@@ -446,6 +446,11 @@ SGLang supports various environment variables that can be used to configure its
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>When the maximum kv len in current prefill batch exceeds this value, the sparse mla kernel will be applied, else it falls back to dense MHA implementation. Default to the index topk of model (2048 for DeepSeek V3.2). <code>SGLANG_NSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD</code> is a deprecated alias.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>2048</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_DSA_TOPK_BROADCAST</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Experimental. When enabled, broadcast the finalized NSA/DSA indexer top-k result from attention TP rank 0 to the other attention TP ranks. This can mitigate top-k mismatches in TP attention runs at the cost of some speed.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
|
||||
@@ -482,6 +482,7 @@ class Envs:
|
||||
)
|
||||
SGLANG_DSA_MQA_LOGITS_FREE_MEM_FRACTION = EnvFloat(0.2)
|
||||
SGLANG_USE_FUSED_METADATA_COPY = EnvBool(True)
|
||||
SGLANG_DSA_TOPK_BROADCAST = EnvBool(False)
|
||||
|
||||
# sgl-kernel
|
||||
SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK = EnvBool(False)
|
||||
|
||||
@@ -74,6 +74,7 @@ if is_npu():
|
||||
from sglang.srt.distributed import (
|
||||
get_attn_context_model_parallel_rank,
|
||||
get_attn_context_model_parallel_world_size,
|
||||
get_attn_tp_group,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import get_pp_group
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
@@ -167,6 +168,42 @@ if _is_cuda:
|
||||
weights = weights.unsqueeze(-1) * q_scale * softmax_scale
|
||||
return weights
|
||||
|
||||
@register_custom_op(mutates_args=["topk_indices"])
|
||||
@register_split_op()
|
||||
def broadcast_indexer_topk_from_rank0_(topk_indices: torch.Tensor) -> None:
|
||||
_broadcast_indexer_topk_from_rank0_impl(topk_indices)
|
||||
|
||||
|
||||
def _broadcast_indexer_topk_from_rank0_impl(topk_indices: torch.Tensor) -> None:
|
||||
group = get_attn_tp_group()
|
||||
if group.world_size == 1:
|
||||
return
|
||||
|
||||
if topk_indices.device.type == "cuda" and torch.cuda.is_current_stream_capturing():
|
||||
if group.pynccl_comm is None:
|
||||
raise RuntimeError(
|
||||
"SGLANG_DSA_TOPK_BROADCAST requires PyNCCL during CUDA graph capture."
|
||||
)
|
||||
with group.pynccl_comm.change_state(enable=True):
|
||||
group.pynccl_comm.broadcast(topk_indices, src=0)
|
||||
else:
|
||||
group.broadcast(topk_indices, src=0)
|
||||
|
||||
|
||||
def _broadcast_indexer_topk_from_rank0(
|
||||
topk_indices: Optional[torch.Tensor],
|
||||
) -> Optional[torch.Tensor]:
|
||||
# Sync only the finalized indexer output. Internal topk_transform calls can
|
||||
# be chunked differently across ranks, which would make collectives diverge.
|
||||
if topk_indices is None or not envs.SGLANG_DSA_TOPK_BROADCAST.get():
|
||||
return topk_indices
|
||||
|
||||
if is_in_piecewise_cuda_graph():
|
||||
broadcast_indexer_topk_from_rank0_(topk_indices)
|
||||
else:
|
||||
_broadcast_indexer_topk_from_rank0_impl(topk_indices)
|
||||
return topk_indices
|
||||
|
||||
|
||||
class BaseIndexerMetadata(ABC):
|
||||
@abstractmethod
|
||||
@@ -1302,19 +1339,18 @@ class Indexer(MultiPlatformOp):
|
||||
|
||||
# Optimization: fast path when skipping topk computation
|
||||
if skip_logits_computation and (not self.dsa_enable_prefill_cp):
|
||||
return maybe_capture_indexer_topk(
|
||||
topk_result = self._forward_cuda_k_only(
|
||||
x,
|
||||
positions,
|
||||
forward_batch,
|
||||
layer_id,
|
||||
self._forward_cuda_k_only(
|
||||
x,
|
||||
positions,
|
||||
forward_batch,
|
||||
layer_id,
|
||||
act_quant,
|
||||
enable_dual_stream,
|
||||
metadata,
|
||||
return_indices,
|
||||
),
|
||||
act_quant,
|
||||
enable_dual_stream,
|
||||
metadata,
|
||||
return_indices,
|
||||
)
|
||||
topk_result = _broadcast_indexer_topk_from_rank0(topk_result)
|
||||
return maybe_capture_indexer_topk(layer_id, topk_result)
|
||||
|
||||
if enable_dual_stream and forward_batch.forward_mode.is_decode_or_idle():
|
||||
current_stream = torch.cuda.current_stream()
|
||||
@@ -1427,15 +1463,14 @@ class Indexer(MultiPlatformOp):
|
||||
# print(
|
||||
# "HACK: seq_lens empty but x not empty, hackily return all-invalid topk_result"
|
||||
# )
|
||||
return maybe_capture_indexer_topk(
|
||||
layer_id,
|
||||
torch.full(
|
||||
(x_meta.shape[0], self.index_topk),
|
||||
-1,
|
||||
dtype=torch.int,
|
||||
device=x_meta.device,
|
||||
),
|
||||
topk_result = torch.full(
|
||||
(x_meta.shape[0], self.index_topk),
|
||||
-1,
|
||||
dtype=torch.int,
|
||||
device=x_meta.device,
|
||||
)
|
||||
topk_result = _broadcast_indexer_topk_from_rank0(topk_result)
|
||||
return maybe_capture_indexer_topk(layer_id, topk_result)
|
||||
|
||||
if (
|
||||
forward_batch.forward_mode.is_decode_or_idle()
|
||||
@@ -1488,10 +1523,9 @@ class Indexer(MultiPlatformOp):
|
||||
kv_len_next,
|
||||
actual_seq_q_next,
|
||||
)
|
||||
return maybe_capture_indexer_topk(
|
||||
layer_id,
|
||||
torch.cat([topk_result_prev, topk_result_next], dim=0),
|
||||
)
|
||||
topk_result = torch.cat([topk_result_prev, topk_result_next], dim=0)
|
||||
topk_result = _broadcast_indexer_topk_from_rank0(topk_result)
|
||||
return maybe_capture_indexer_topk(layer_id, topk_result)
|
||||
elif is_in_piecewise_cuda_graph():
|
||||
assert (
|
||||
not enable_dual_stream
|
||||
@@ -1527,6 +1561,7 @@ class Indexer(MultiPlatformOp):
|
||||
topk=self.index_topk,
|
||||
layer_id=layer_id,
|
||||
)
|
||||
topk_result = _broadcast_indexer_topk_from_rank0(topk_result)
|
||||
return maybe_capture_indexer_topk(layer_id, topk_result)
|
||||
|
||||
def forward_npu(
|
||||
|
||||
Reference in New Issue
Block a user