Clarify ModelRunner.dp_size into attn_dp_size (#31142)
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user