[DP] Fix FlashInfer dispatcher workspace sizing and set_dp_buffer_len (#26643)

This commit is contained in:
Hanming Lu
2026-06-02 13:25:20 -07:00
committed by GitHub
parent 9e717cae46
commit b603f08c0c
8 changed files with 90 additions and 161 deletions
@@ -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)