refactor context parallel state (#17213)
Co-authored-by: Shunkang <182541032+Shunkangz@users.noreply.github.co>
This commit is contained in:
co-authored by
Shunkang
parent
0012d6a4eb
commit
8b4c364960
@@ -43,6 +43,7 @@ from sglang.srt.dllm.config import DllmConfig
|
||||
from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
DpPaddingMode,
|
||||
get_attention_cp_size,
|
||||
get_attention_tp_rank,
|
||||
get_attention_tp_size,
|
||||
set_dp_buffer_len,
|
||||
@@ -204,6 +205,9 @@ def get_batch_sizes_to_capture(model_runner: ModelRunner, num_tokens_per_bs=1):
|
||||
if require_gathered_buffer(server_args):
|
||||
mul_base *= get_attention_tp_size()
|
||||
|
||||
if mul_base % get_attention_cp_size() != 0:
|
||||
mul_base *= get_attention_cp_size()
|
||||
|
||||
# Model input token count = bs * num_tokens_per_bs; must be a multiple of attn_tp_size.
|
||||
capture_bs = [bs for bs in capture_bs if bs * num_tokens_per_bs % mul_base == 0]
|
||||
|
||||
|
||||
@@ -45,6 +45,7 @@ from sglang.srt.distributed.parallel_state import (
|
||||
from sglang.srt.layers.attention.nsa.utils import NSAContextParallelMetadata
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
DpPaddingMode,
|
||||
get_attention_cp_size,
|
||||
get_attention_dp_rank,
|
||||
get_attention_tp_rank,
|
||||
get_attention_tp_size,
|
||||
@@ -749,6 +750,11 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
# there is no reduce-scatter in LM logprob, so we do not need to adjust the padded length for logprob
|
||||
global_num_tokens[i] = ceil_align(global_num_tokens[i], attn_tp_size)
|
||||
|
||||
# make sure that each rank has the same number of tokens to do collective communication.
|
||||
attn_cp_size = get_attention_cp_size()
|
||||
for i in range(sync_group_size):
|
||||
global_num_tokens[i] = ceil_align(global_num_tokens[i], attn_cp_size)
|
||||
|
||||
dp_padding_mode = DpPaddingMode.get_dp_padding_mode(
|
||||
self.is_extend_in_batch, global_num_tokens
|
||||
)
|
||||
|
||||
@@ -290,6 +290,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
nccl_port: int,
|
||||
server_args: ServerArgs,
|
||||
dp_rank: Optional[int] = None,
|
||||
attn_cp_rank: Optional[int] = None,
|
||||
moe_dp_rank: Optional[int] = None,
|
||||
is_draft_worker: bool = False,
|
||||
req_to_token_pool: Optional[ReqToTokenPool] = None,
|
||||
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
|
||||
@@ -303,9 +305,13 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
self.tp_size = tp_size
|
||||
self.moe_ep_rank = moe_ep_rank
|
||||
self.moe_ep_size = moe_ep_size
|
||||
self.dp_size = server_args.dp_size
|
||||
self.dp_size = server_args.dp_size if server_args.enable_dp_attention else 1
|
||||
self.pp_rank = pp_rank
|
||||
self.pp_size = pp_size
|
||||
self.attn_cp_rank = attn_cp_rank
|
||||
self.attn_cp_size = server_args.attn_cp_size
|
||||
self.moe_dp_rank = moe_dp_rank
|
||||
self.moe_dp_size = server_args.moe_dp_size
|
||||
self.model_config = model_config
|
||||
self.dist_port = nccl_port
|
||||
self.server_args = server_args
|
||||
@@ -586,8 +592,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
(
|
||||
self.max_total_num_tokens // 2
|
||||
if server_args.max_running_requests is None
|
||||
else server_args.max_running_requests
|
||||
// (server_args.dp_size if server_args.enable_dp_attention else 1)
|
||||
else server_args.max_running_requests // (self.dp_size)
|
||||
),
|
||||
self.req_to_token_pool.size,
|
||||
)
|
||||
@@ -797,8 +802,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
)
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=self.tp_size,
|
||||
attention_data_parallel_size=self.dp_size,
|
||||
pipeline_model_parallel_size=self.pp_size,
|
||||
expert_model_parallel_size=self.moe_ep_size,
|
||||
attention_context_model_parallel_size=self.attn_cp_size,
|
||||
moe_data_model_parallel_size=self.moe_dp_size,
|
||||
duplicate_tp_group=self.server_args.enable_pdmux,
|
||||
)
|
||||
initialize_dp_attention(
|
||||
|
||||
Reference in New Issue
Block a user