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_rank = moe_ep_rank
|
||||||
self.moe_ep_size = moe_ep_size
|
self.moe_ep_size = moe_ep_size
|
||||||
self.dp_rank = dp_rank
|
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_rank = pp_rank
|
||||||
self.pp_size = pp_size
|
self.pp_size = pp_size
|
||||||
self.attn_cp_rank = attn_cp_rank
|
self.attn_cp_rank = attn_cp_rank
|
||||||
@@ -1280,7 +1282,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
initialize_model_parallel(
|
initialize_model_parallel(
|
||||||
tensor_model_parallel_size=self.tp_size,
|
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,
|
pipeline_model_parallel_size=self.pp_size,
|
||||||
expert_model_parallel_size=self.moe_ep_size,
|
expert_model_parallel_size=self.moe_ep_size,
|
||||||
attention_context_model_parallel_size=self.attn_cp_size,
|
attention_context_model_parallel_size=self.attn_cp_size,
|
||||||
@@ -3012,7 +3014,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
# Try msprob debugger
|
# Try msprob debugger
|
||||||
if self.msprobe_debugger is not None:
|
if self.msprobe_debugger is not None:
|
||||||
rank_id = (
|
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)
|
self.msprobe_debugger.start(model=self.model, rank_id=rank_id)
|
||||||
|
|
||||||
|
|||||||
@@ -168,7 +168,7 @@ class ModelRunnerKVCacheMixin:
|
|||||||
ratio = self._calculate_mamba_ratio()
|
ratio = self._calculate_mamba_ratio()
|
||||||
capped_reqs = min(
|
capped_reqs = min(
|
||||||
server_args.max_running_requests
|
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,
|
server_args.max_mamba_cache_size // ratio,
|
||||||
)
|
)
|
||||||
intermediate_size = (
|
intermediate_size = (
|
||||||
@@ -226,7 +226,7 @@ class ModelRunnerKVCacheMixin:
|
|||||||
# so the return value only has main_state subtracted from total
|
# so the return value only has main_state subtracted from total
|
||||||
capped_reqs = min(
|
capped_reqs = min(
|
||||||
server_args.max_running_requests
|
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,
|
server_args.max_mamba_cache_size // ratio,
|
||||||
)
|
)
|
||||||
intermediate_size = per_req * capped_reqs * D
|
intermediate_size = per_req * capped_reqs * D
|
||||||
@@ -1353,7 +1353,7 @@ class ModelRunnerKVCacheMixin:
|
|||||||
|
|
||||||
max_num_reqs = self.server_args.max_running_requests
|
max_num_reqs = self.server_args.max_running_requests
|
||||||
if max_num_reqs is not None:
|
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)
|
max_num_reqs = min(requested_per_worker, token_capacity // 2)
|
||||||
else:
|
else:
|
||||||
requested_per_worker = None
|
requested_per_worker = None
|
||||||
|
|||||||
@@ -107,7 +107,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.tp_size
|
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.pp_size = model_runner.server_args.pp_size
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
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_gathered_buffer:
|
||||||
if self.require_mlp_tp_gather:
|
if self.require_mlp_tp_gather:
|
||||||
global_num_tokens_gpu = torch.zeros(
|
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(
|
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||||
(self.dp_size,), dtype=torch.int32
|
(self.attn_dp_size,), dtype=torch.int32
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
assert self.require_attn_tp_gather
|
assert self.require_attn_tp_gather
|
||||||
@@ -350,7 +350,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.require_mlp_tp_gather:
|
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:
|
elif self.require_attn_tp_gather:
|
||||||
global_num_tokens_cpu = [num_tokens]
|
global_num_tokens_cpu = [num_tokens]
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -96,7 +96,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
self.tp_size = model_runner.tp_size
|
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.pp_size = model_runner.server_args.pp_size
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
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_gathered_buffer:
|
||||||
if self.require_mlp_tp_gather:
|
if self.require_mlp_tp_gather:
|
||||||
global_num_tokens_gpu = torch.zeros(
|
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(
|
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||||
(self.dp_size,), dtype=torch.int32
|
(self.attn_dp_size,), dtype=torch.int32
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
assert self.require_attn_tp_gather
|
assert self.require_attn_tp_gather
|
||||||
@@ -339,7 +339,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
num_tokens_for_logprob = num_tokens
|
num_tokens_for_logprob = num_tokens
|
||||||
|
|
||||||
if self.require_mlp_tp_gather:
|
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:
|
elif self.require_attn_tp_gather:
|
||||||
global_num_tokens_cpu = [num_tokens]
|
global_num_tokens_cpu = [num_tokens]
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -92,7 +92,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
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.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||||
self.tp_size = self.model_runner.tp_size
|
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.pp_size = model_runner.server_args.pp_size
|
||||||
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
|
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk
|
self.topk = model_runner.server_args.speculative_eagle_topk
|
||||||
@@ -149,10 +149,10 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
if self.require_gathered_buffer:
|
if self.require_gathered_buffer:
|
||||||
if self.require_mlp_tp_gather:
|
if self.require_mlp_tp_gather:
|
||||||
global_num_tokens_gpu = torch.zeros(
|
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(
|
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||||
(self.dp_size,), dtype=torch.int32
|
(self.attn_dp_size,), dtype=torch.int32
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
assert self.require_attn_tp_gather
|
assert self.require_attn_tp_gather
|
||||||
@@ -240,7 +240,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
bonus_tokens = buffers.bonus_tokens[:request_bs]
|
bonus_tokens = buffers.bonus_tokens[:request_bs]
|
||||||
|
|
||||||
if self.require_mlp_tp_gather:
|
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:
|
elif self.require_attn_tp_gather:
|
||||||
global_num_tokens_cpu = [expanded_bs]
|
global_num_tokens_cpu = [expanded_bs]
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user