Clarify ModelRunner.dp_size into attn_dp_size (#31142)

This commit is contained in:
fzyzcjy
2026-07-14 15:49:05 +08:00
committed by GitHub
parent afa3c06d1f
commit 2cf753c4fe
5 changed files with 22 additions and 18 deletions
@@ -370,7 +370,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.moe_ep_rank = moe_ep_rank
self.moe_ep_size = moe_ep_size
self.dp_rank = dp_rank
self.dp_size = server_args.dp_size if server_args.enable_dp_attention else 1
self.attn_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
@@ -1280,7 +1282,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
)
initialize_model_parallel(
tensor_model_parallel_size=self.tp_size,
attention_data_parallel_size=self.dp_size,
attention_data_parallel_size=self.attn_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,
@@ -3012,7 +3014,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
# Try msprob debugger
if self.msprobe_debugger is not None:
rank_id = (
self.gpu_id if self.dp_size is not None and self.dp_size > 1 else None
self.gpu_id
if self.attn_dp_size is not None and self.attn_dp_size > 1
else None
)
self.msprobe_debugger.start(model=self.model, rank_id=rank_id)
@@ -168,7 +168,7 @@ class ModelRunnerKVCacheMixin:
ratio = self._calculate_mamba_ratio()
capped_reqs = min(
server_args.max_running_requests
// (self.dp_size if server_args.enable_dp_attention else 1),
// (self.attn_dp_size if server_args.enable_dp_attention else 1),
server_args.max_mamba_cache_size // ratio,
)
intermediate_size = (
@@ -226,7 +226,7 @@ class ModelRunnerKVCacheMixin:
# so the return value only has main_state subtracted from total
capped_reqs = min(
server_args.max_running_requests
// (self.dp_size if server_args.enable_dp_attention else 1),
// (self.attn_dp_size if server_args.enable_dp_attention else 1),
server_args.max_mamba_cache_size // ratio,
)
intermediate_size = per_req * capped_reqs * D
@@ -1353,7 +1353,7 @@ class ModelRunnerKVCacheMixin:
max_num_reqs = self.server_args.max_running_requests
if max_num_reqs is not None:
requested_per_worker = max_num_reqs // self.dp_size
requested_per_worker = max_num_reqs // self.attn_dp_size
max_num_reqs = min(requested_per_worker, token_capacity // 2)
else:
requested_per_worker = None
@@ -107,7 +107,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
self.device = model_runner.device
self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.tp_size
self.dp_size = model_runner.dp_size
self.attn_dp_size = model_runner.attn_dp_size
self.pp_size = model_runner.server_args.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
@@ -210,10 +210,10 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
if self.require_gathered_buffer:
if self.require_mlp_tp_gather:
global_num_tokens_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
(self.attn_dp_size,), dtype=torch.int32
)
global_num_tokens_for_logprob_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
(self.attn_dp_size,), dtype=torch.int32
)
else:
assert self.require_attn_tp_gather
@@ -350,7 +350,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
)
if self.require_mlp_tp_gather:
global_num_tokens_cpu = [num_tokens] * self.dp_size
global_num_tokens_cpu = [num_tokens] * self.attn_dp_size
elif self.require_attn_tp_gather:
global_num_tokens_cpu = [num_tokens]
else:
@@ -96,7 +96,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.device = model_runner.device
self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.tp_size
self.dp_size = model_runner.dp_size
self.attn_dp_size = model_runner.attn_dp_size
self.pp_size = model_runner.server_args.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
@@ -193,10 +193,10 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
if self.require_gathered_buffer:
if self.require_mlp_tp_gather:
global_num_tokens_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
(self.attn_dp_size,), dtype=torch.int32
)
global_num_tokens_for_logprob_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
(self.attn_dp_size,), dtype=torch.int32
)
else:
assert self.require_attn_tp_gather
@@ -339,7 +339,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
num_tokens_for_logprob = num_tokens
if self.require_mlp_tp_gather:
global_num_tokens_cpu = [num_tokens] * self.dp_size
global_num_tokens_cpu = [num_tokens] * self.attn_dp_size
elif self.require_attn_tp_gather:
global_num_tokens_cpu = [num_tokens]
else:
@@ -92,7 +92,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.tp_size = self.model_runner.tp_size
self.dp_size = self.model_runner.dp_size
self.attn_dp_size = self.model_runner.attn_dp_size
self.pp_size = model_runner.server_args.pp_size
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.topk = model_runner.server_args.speculative_eagle_topk
@@ -149,10 +149,10 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
if self.require_gathered_buffer:
if self.require_mlp_tp_gather:
global_num_tokens_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
(self.attn_dp_size,), dtype=torch.int32
)
global_num_tokens_for_logprob_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
(self.attn_dp_size,), dtype=torch.int32
)
else:
assert self.require_attn_tp_gather
@@ -240,7 +240,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
bonus_tokens = buffers.bonus_tokens[:request_bs]
if self.require_mlp_tp_gather:
global_num_tokens_cpu = [expanded_bs] * self.dp_size
global_num_tokens_cpu = [expanded_bs] * self.attn_dp_size
elif self.require_attn_tp_gather:
global_num_tokens_cpu = [expanded_bs]
else: