perf(cp): add narrow output collection for NSA prefill CP

Replace full hidden all-gather at prefill tail with per-request last-token
hidden collection, reducing communication from total_tokens x hidden_size
to cp_size x bs x hidden_size for both in-seq-split and round-robin modes.

- nsa/utils.py: add cp_collect_last_token_hidden() with mode-specific
  narrow collection helpers that only gather the last token hidden
- deepseek_v2.py: add _should_use_narrow_output_path() gate on
  DeepseekV2Model, fallback to full gather for EAGLE/return_logprob/
  capture_hidden batches
- logits_processor.py: add _is_compact_hidden_states() to bypass
  _get_pruned_states() when hidden is already compact
This commit is contained in:
laoyao0822
2026-04-22 04:55:06 +08:00
parent 7403f14511
commit bc9a9f128b
3 changed files with 141 additions and 48 deletions
+60 -34
View File
@@ -57,6 +57,7 @@ from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
from sglang.srt.layers.attention.nsa.utils import (
can_cp_split,
cp_all_gather_rerange_output,
cp_collect_last_token_hidden,
cp_split_and_rebuild_data,
cp_split_and_rebuild_position,
is_nsa_enable_prefill_cp,
@@ -114,7 +115,11 @@ from sglang.srt.layers.vocab_parallel_embedding import (
VocabParallelEmbedding,
)
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
PPProxyTensors,
)
from sglang.srt.models.deepseek_common.attention_backend_handler import (
AttentionBackendRegistry,
)
@@ -222,8 +227,7 @@ class DeepseekV2MLP(nn.Module):
self.down_proj.weight = self.down_proj.weight_packed
if hidden_act != "silu":
raise ValueError(
f"Unsupported activation: {hidden_act}. "
"Only silu is supported for now."
f"Unsupported activation: {hidden_act}. Only silu is supported for now."
)
self.act_fn = SiluAndMul()
@@ -321,7 +325,6 @@ class MoEGate(nn.Module):
and (self.weight.shape[0] == 256 or self.weight.shape[0] == 384)
and _device_sm >= 90
):
# router gemm output float32
logits = dsv3_router_gemm(
hidden_states, self.weight, out_dtype=torch.float32
@@ -341,7 +344,6 @@ class MoEGate(nn.Module):
class DeepseekV2MoE(nn.Module):
def __init__(
self,
config: PretrainedConfig,
@@ -458,13 +460,15 @@ class DeepseekV2MoE(nn.Module):
else {}
),
)
is_packed_weight = hasattr(
self.shared_experts.gate_up_proj.quant_method, "quant_config"
) and self.shared_experts.gate_up_proj.quant_method.quant_config.get_name() in {
"awq",
"awq_marlin",
"moe_wna16",
}
is_packed_weight = (
hasattr(self.shared_experts.gate_up_proj.quant_method, "quant_config")
and self.shared_experts.gate_up_proj.quant_method.quant_config.get_name()
in {
"awq",
"awq_marlin",
"moe_wna16",
}
)
self.shared_experts_is_int8 = (
not is_packed_weight
and self.shared_experts.gate_up_proj.weight.dtype == torch.int8
@@ -486,9 +490,7 @@ class DeepseekV2MoE(nn.Module):
self.shared_experts.gate_up_proj.quant_method.quant_config.weight_block_size
== self.shared_experts.down_proj.quant_method.quant_config.weight_block_size
)
self.shared_experts_weight_block_size = (
self.shared_experts.gate_up_proj.quant_method.quant_config.weight_block_size
)
self.shared_experts_weight_block_size = self.shared_experts.gate_up_proj.quant_method.quant_config.weight_block_size
self.top_k = config.num_experts_per_tok
@@ -647,7 +649,6 @@ class DeepseekV2MoE(nn.Module):
def _pre_combine_hook(
dispatcher: BaseDispatcher, combine_input: CombineInput
):
nonlocal shared_output
self.alt_stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(self.alt_stream):
@@ -837,7 +838,6 @@ class DeepseekV2MoE(nn.Module):
def _post_dispatch_hook(
dispatcher: BaseDispatcher, dispatch_output: DispatchOutput
):
combine_overlap_args, down_gemm_overlap_args, meta_overlap_args = (
compute_overlap_args(dispatch_output, self.alt_stream)
)
@@ -855,7 +855,6 @@ class DeepseekV2MoE(nn.Module):
def _pre_combine_hook(
dispatcher: BaseDispatcher, combine_input: CombineInput
):
nonlocal shared_output
if (
@@ -893,7 +892,6 @@ class DeepseekV2MoE(nn.Module):
def _post_dispatch_hook(
dispatcher: BaseDispatcher, dispatch_output: DispatchOutput
):
combine_overlap_args, down_gemm_overlap_args, meta_overlap_args = (
compute_overlap_args(dispatch_output, self.alt_stream)
)
@@ -1076,7 +1074,6 @@ class DeepseekV2AttentionMLA(
DeepseekMLARocmForwardMixin,
DeepseekMLACpuForwardMixin,
):
def __init__(
self,
config: PretrainedConfig,
@@ -1357,18 +1354,18 @@ class DeepseekV2AttentionMLA(
not get_attn_tp_context().input_scattered
and hidden_states[0].shape[0] == 0
):
assert (
not self.o_proj.reduce_results
), "short-circuiting allreduce will lead to hangs"
assert not self.o_proj.reduce_results, (
"short-circuiting allreduce will lead to hangs"
)
return hidden_states[0]
else:
if (
not get_attn_tp_context().input_scattered
and hidden_states.shape[0] == 0
):
assert (
not self.o_proj.reduce_results
), "short-circuiting allreduce will lead to hangs"
assert not self.o_proj.reduce_results, (
"short-circuiting allreduce will lead to hangs"
)
return hidden_states, None, forward_batch, None
attn_forward_method = self.dispatch_attn_forward_method(forward_batch)
@@ -1499,7 +1496,6 @@ class DeepseekV2AttentionMLA(
class DeepseekV2DecoderLayer(nn.Module):
def __init__(
self,
config: PretrainedConfig,
@@ -2044,17 +2040,34 @@ class DeepseekV2Model(nn.Module):
hidden_states, _ = self.norm(hidden_states, residual)
if self.pp_group.is_last_rank and nsa_use_prefill_cp(forward_batch):
# allgather + rerrange
hidden_states = cp_all_gather_rerange_output(
hidden_states,
self.cp_size,
forward_batch,
torch.cuda.current_stream(),
)
if self._should_use_narrow_output_path(forward_batch):
hidden_states = cp_collect_last_token_hidden(
hidden_states, forward_batch, self.cp_size
)
else:
hidden_states = cp_all_gather_rerange_output(
hidden_states,
self.cp_size,
forward_batch,
torch.cuda.current_stream(),
)
if len(aux_hidden_states) == 0:
return hidden_states
return hidden_states, aux_hidden_states
def _should_use_narrow_output_path(self, forward_batch):
if not nsa_use_prefill_cp(forward_batch):
return False
if not self.pp_group.is_last_rank:
return False
if not forward_batch.forward_mode.is_extend():
return False
if forward_batch.return_logprob:
return False
if forward_batch.capture_hidden_mode != CaptureHiddenMode.NULL:
return False
return True
class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
# for quark model load
@@ -2208,6 +2221,19 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
else:
return hidden_states
def _should_use_narrow_output_path(self, forward_batch: ForwardBatch) -> bool:
if not nsa_use_prefill_cp(forward_batch):
return False
if not self.pp_group.is_last_rank:
return False
if not forward_batch.forward_mode.is_extend():
return False
if forward_batch.return_logprob:
return False
if forward_batch.capture_hidden_mode != CaptureHiddenMode.NULL:
return False
return True
@property
def start_layer(self):
return self.model.start_layer