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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user