[DP] Fix FlashInfer dispatcher workspace sizing and set_dp_buffer_len (#26643)
This commit is contained in:
@@ -100,13 +100,23 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
# TODO: Can other moe runners use payload_in_workspace too?
|
||||
self.payload_in_workspace = get_moe_runner_backend().is_flashinfer_cutlass()
|
||||
|
||||
# TODO: Can this be a server arg and shared with deepep/mooncakeep?
|
||||
# FlashInfer sizes the workspace from the maximum dispatched tokens per
|
||||
# EP rank. See FlashInfer's moe_a2a_get_workspace_size_per_rank(),
|
||||
# which reserves ep_size * max_num_tokens * payload bytes, and the C++
|
||||
# dispatch op's epSize * runtimeMaxTokensPerRank payload buffer.
|
||||
#
|
||||
# The workspace must fit both:
|
||||
# (a) the fattest prefill batch (bounded by chunked_prefill_size), and
|
||||
# (b) the largest decode batch (bounded by max_running_requests, which
|
||||
# _resolve_max_num_reqs caps at 4096 per DP worker).
|
||||
# max_running_requests is not yet resolved at model-construction time,
|
||||
# so we use 4096 as a floor to cover decode batches and _dummy_run
|
||||
# (which warms up at batch_size = req_to_token_pool.size).
|
||||
cps = get_global_server_args().chunked_prefill_size
|
||||
default_max_tokens = max(cps if cps and cps > 0 else 4096, 4096)
|
||||
self.max_num_tokens = get_int_env_var(
|
||||
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK", 4096
|
||||
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK",
|
||||
default_max_tokens,
|
||||
)
|
||||
|
||||
# Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized.
|
||||
|
||||
@@ -1036,37 +1036,19 @@ class CudaGraphRunner:
|
||||
)
|
||||
|
||||
if self.require_mlp_tp_gather:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=input_ids.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens * self.dp_size
|
||||
global_num_tokens_cpu = [num_tokens] * self.dp_size
|
||||
elif self.require_attn_tp_gather:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=input_ids.device,
|
||||
)
|
||||
global_num_tokens_cpu = [num_tokens]
|
||||
else:
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
if global_num_tokens_cpu is not None:
|
||||
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
||||
num_tokens_tensor = torch.tensor(
|
||||
global_num_tokens_cpu, dtype=torch.int32, device=input_ids.device
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens
|
||||
buffers.global_num_tokens_gpu.copy_(num_tokens_tensor)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(num_tokens_tensor)
|
||||
else:
|
||||
global_dp_buffer_len = None
|
||||
|
||||
@@ -1171,6 +1153,7 @@ class CudaGraphRunner:
|
||||
global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
global_num_tokens_cpu,
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
|
||||
@@ -2606,39 +2606,19 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
)
|
||||
|
||||
if require_mlp_tp_gather_:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.server_args.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.server_args.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens * self.server_args.dp_size
|
||||
global_num_tokens_cpu = [num_tokens] * self.server_args.dp_size
|
||||
elif require_attn_tp_gather(self.server_args):
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens
|
||||
global_num_tokens_cpu = [num_tokens]
|
||||
else:
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
if global_num_tokens_cpu is not None:
|
||||
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
||||
num_tokens_tensor = torch.tensor(
|
||||
global_num_tokens_cpu, dtype=torch.int32, device=self.device
|
||||
)
|
||||
buffers.global_num_tokens_gpu.copy_(num_tokens_tensor)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(num_tokens_tensor)
|
||||
else:
|
||||
global_dp_buffer_len = None
|
||||
global_num_tokens_cpu = None
|
||||
@@ -2754,6 +2734,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
global_num_tokens_cpu,
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
|
||||
@@ -537,6 +537,7 @@ class PiecewiseCudaGraphRunner:
|
||||
)
|
||||
|
||||
global_dp_buffer_len = None
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
if self.model_runner.server_args.enable_lora:
|
||||
# It is safe to capture CUDA graph using empty LoRA id, as the LoRA kernels will always be launched whenever
|
||||
@@ -611,6 +612,7 @@ class PiecewiseCudaGraphRunner:
|
||||
global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
global_num_tokens_cpu,
|
||||
)
|
||||
# FIXME: the implementation is hacky. `is_extend_in_batch`` is for determining the deepep mode.
|
||||
# It is True in this context but we need to set it to use low latency deepep mode.
|
||||
|
||||
@@ -289,44 +289,26 @@ class EAGLEDraftCudaGraphRunner:
|
||||
topk_index = buffers.topk_index[:num_seqs]
|
||||
|
||||
if self.require_mlp_tp_gather:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_num_tokens = buffers.global_num_tokens_gpu
|
||||
global_dp_buffer_len = num_tokens * self.dp_size
|
||||
global_num_tokens_for_logprob = buffers.global_num_tokens_for_logprob_gpu
|
||||
global_num_tokens_cpu = [num_tokens] * self.dp_size
|
||||
elif self.require_attn_tp_gather:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
global_num_tokens_cpu = [num_tokens]
|
||||
else:
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
if global_num_tokens_cpu is not None:
|
||||
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
||||
num_tokens_tensor = torch.tensor(
|
||||
global_num_tokens_cpu,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
buffers.global_num_tokens_gpu.copy_(num_tokens_tensor)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(num_tokens_tensor)
|
||||
global_num_tokens = buffers.global_num_tokens_gpu
|
||||
global_dp_buffer_len = num_tokens
|
||||
global_num_tokens_for_logprob = buffers.global_num_tokens_for_logprob_gpu
|
||||
else:
|
||||
global_num_tokens = None
|
||||
global_dp_buffer_len = None
|
||||
global_num_tokens = None
|
||||
global_num_tokens_for_logprob = None
|
||||
|
||||
capture_mode = (
|
||||
@@ -380,6 +362,7 @@ class EAGLEDraftCudaGraphRunner:
|
||||
global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
global_num_tokens_cpu,
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
|
||||
@@ -316,37 +316,28 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
)
|
||||
|
||||
if self.require_mlp_tp_gather:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens_for_logprob] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens * self.dp_size
|
||||
global_num_tokens_cpu = [num_tokens] * self.dp_size
|
||||
elif self.require_attn_tp_gather:
|
||||
global_num_tokens_cpu = [num_tokens]
|
||||
else:
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
if global_num_tokens_cpu is not None:
|
||||
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
global_num_tokens_cpu,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens_for_logprob],
|
||||
[num_tokens_for_logprob] * len(global_num_tokens_cpu),
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens
|
||||
else:
|
||||
global_dp_buffer_len = None
|
||||
|
||||
@@ -392,6 +383,7 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
global_num_tokens_cpu,
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
|
||||
@@ -215,45 +215,27 @@ class FrozenKVMTPCudaGraphRunner:
|
||||
bonus_tokens = buffers.bonus_tokens[:request_bs]
|
||||
|
||||
if self.require_mlp_tp_gather:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[expanded_bs] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=buffers.positions.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[expanded_bs] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=buffers.positions.device,
|
||||
)
|
||||
)
|
||||
global_num_tokens = buffers.global_num_tokens_gpu
|
||||
global_num_tokens_for_logprob = buffers.global_num_tokens_for_logprob_gpu
|
||||
global_dp_buffer_len = expanded_bs * self.dp_size
|
||||
global_num_tokens_cpu = [expanded_bs] * self.dp_size
|
||||
elif self.require_attn_tp_gather:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[expanded_bs],
|
||||
dtype=torch.int32,
|
||||
device=buffers.positions.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[expanded_bs],
|
||||
dtype=torch.int32,
|
||||
device=buffers.positions.device,
|
||||
)
|
||||
global_num_tokens_cpu = [expanded_bs]
|
||||
else:
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
if global_num_tokens_cpu is not None:
|
||||
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
||||
num_tokens_tensor = torch.tensor(
|
||||
global_num_tokens_cpu,
|
||||
dtype=torch.int32,
|
||||
device=buffers.positions.device,
|
||||
)
|
||||
buffers.global_num_tokens_gpu.copy_(num_tokens_tensor)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(num_tokens_tensor)
|
||||
global_num_tokens = buffers.global_num_tokens_gpu
|
||||
global_num_tokens_for_logprob = buffers.global_num_tokens_for_logprob_gpu
|
||||
global_dp_buffer_len = expanded_bs
|
||||
else:
|
||||
global_dp_buffer_len = None
|
||||
global_num_tokens = None
|
||||
global_num_tokens_for_logprob = None
|
||||
global_dp_buffer_len = None
|
||||
|
||||
spec_info = FrozenKVMTPDraftInput(
|
||||
topk_p=topk_p,
|
||||
@@ -296,6 +278,7 @@ class FrozenKVMTPCudaGraphRunner:
|
||||
global_dp_buffer_len,
|
||||
expanded_bs,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
global_num_tokens_cpu,
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
|
||||
@@ -318,37 +318,30 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
]
|
||||
|
||||
if self.require_mlp_tp_gather:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens * self.dp_size
|
||||
global_num_tokens_cpu = [num_tokens] * self.dp_size
|
||||
global_num_tokens_for_logprob_cpu = [num_tokens] * self.dp_size
|
||||
elif self.require_attn_tp_gather:
|
||||
global_num_tokens_cpu = [num_tokens]
|
||||
global_num_tokens_for_logprob_cpu = [bs]
|
||||
else:
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
if global_num_tokens_cpu is not None:
|
||||
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
global_num_tokens_cpu,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[bs],
|
||||
global_num_tokens_for_logprob_cpu,
|
||||
dtype=torch.int32,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens
|
||||
else:
|
||||
global_dp_buffer_len = None
|
||||
|
||||
@@ -383,6 +376,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
|
||||
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
||||
global_dp_buffer_len=global_dp_buffer_len,
|
||||
global_num_tokens_cpu=global_num_tokens_cpu,
|
||||
spec_algorithm=self.model_runner.spec_algorithm,
|
||||
spec_info=spec_info,
|
||||
capture_hidden_mode=capture_mode,
|
||||
@@ -416,6 +410,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
forward_batch.global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
forward_batch.global_num_tokens_cpu,
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user