Multiple flexibility fixes for DP attention (#33537)

This commit is contained in:
Lianmin Zheng
2026-08-04 15:40:40 -07:00
committed by GitHub
parent 6808c6d571
commit b0fd31ba07
3 changed files with 122 additions and 48 deletions
@@ -702,6 +702,39 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
if skip_attn_backend_init:
self.mark_forward_metadata_ready()
def init_mlp_sync_metadata(
self, batch: ScheduleBatch, device: Union[str, torch.device]
) -> None:
"""Populate per-rank token counts for DP-attention MLP synchronization."""
if batch.global_num_tokens is None:
return
assert batch.global_num_tokens_for_logprob is not None
if self.spec_info is not None:
from sglang.srt.speculative.spec_info import spec_scale_global_num_tokens
global_num_tokens, global_num_tokens_for_logprob = (
spec_scale_global_num_tokens(
self.spec_info,
batch.global_num_tokens,
batch.global_num_tokens_for_logprob,
)
)
else:
global_num_tokens = batch.global_num_tokens
global_num_tokens_for_logprob = batch.global_num_tokens_for_logprob
self.original_global_num_tokens_cpu = batch.global_num_tokens
self.global_num_tokens_cpu = global_num_tokens
self.global_num_tokens_gpu = torch.tensor(
global_num_tokens, dtype=torch.int64
).to(device, non_blocking=True)
self.global_num_tokens_for_logprob_cpu = global_num_tokens_for_logprob
self.global_num_tokens_for_logprob_gpu = torch.tensor(
global_num_tokens_for_logprob, dtype=torch.int64
).to(device, non_blocking=True)
self.can_run_dp_cuda_graph = batch.can_run_dp_cuda_graph
@classmethod
def init_new(
cls,
@@ -841,37 +874,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
)
ret.num_token_non_padded_cpu = num_tokens
# For MLP sync
if batch.global_num_tokens is not None:
assert batch.global_num_tokens_for_logprob is not None
# process global_num_tokens and global_num_tokens_for_logprob
if batch.spec_info is not None:
from sglang.srt.speculative.spec_info import (
spec_scale_global_num_tokens,
)
global_num_tokens, global_num_tokens_for_logprob = (
spec_scale_global_num_tokens(
batch.spec_info,
batch.global_num_tokens,
batch.global_num_tokens_for_logprob,
)
)
else:
global_num_tokens = batch.global_num_tokens
global_num_tokens_for_logprob = batch.global_num_tokens_for_logprob
ret.original_global_num_tokens_cpu = batch.global_num_tokens
ret.global_num_tokens_cpu = global_num_tokens
ret.global_num_tokens_gpu = torch.tensor(
global_num_tokens, dtype=torch.int64
).to(device, non_blocking=True)
ret.global_num_tokens_for_logprob_cpu = global_num_tokens_for_logprob
ret.global_num_tokens_for_logprob_gpu = torch.tensor(
global_num_tokens_for_logprob, dtype=torch.int64
).to(device, non_blocking=True)
ret.init_mlp_sync_metadata(batch, device)
if ret.forward_mode.is_idle():
ret.positions = torch.empty((0,), dtype=torch.int64, device=device)
@@ -1520,9 +1523,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
spec_info.num_accept_tokens = self._pad_tensor_to_size(
spec_info.num_accept_tokens, bs
)
spec_info.hidden_states = self._pad_tensor_to_size(
spec_info.hidden_states, num_tokens
)
if spec_info.hidden_states is not None:
spec_info.hidden_states = self._pad_tensor_to_size(
spec_info.hidden_states, num_tokens
)
def prepare_attn_tp_scatter_input(self, model_runner: ModelRunner):
from sglang.srt.layers.communicator import get_attn_tp_context
@@ -1569,12 +1573,15 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
]
logits_output.hidden_states = logits_output.hidden_states[:num_tokens]
elif self.forward_mode.is_target_verify(): # verify
num_tokens = bs * self.spec_info.draft_token_num
num_tokens = bs * self.spec_info.num_tokens_per_req
if logits_output.next_token_logits is not None:
logits_output.next_token_logits = logits_output.next_token_logits[
:num_tokens
]
logits_output.hidden_states = logits_output.hidden_states[:num_tokens]
if logits_output.hidden_states is not None:
logits_output.hidden_states = logits_output.hidden_states[
:num_tokens
]
elif self.forward_mode.is_draft_extend_v2(): # draft extend_v2
bs = bs * self.spec_info.num_tokens_per_req
if logits_output.next_token_logits is not None:
@@ -523,6 +523,15 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
cap_layout.verify_lens.copy_(live.verify_lens)
cap_layout.qo_indptr_device.copy_(live.qo_indptr_device)
@staticmethod
def _max_dp_batch_size(forward_batch: ForwardBatch) -> int:
request_counts = forward_batch.original_global_num_tokens_cpu
if request_counts is None:
raise RuntimeError(
"DP CUDA graph replay requires raw per-rank request counts"
)
return max(request_counts)
def can_run_graph(self, forward_batch: ForwardBatch):
# Disable for token embedding overrides (dynamic per-request)
if forward_batch.replace_embeds is not None:
@@ -550,9 +559,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
return False
if self.require_mlp_tp_gather:
# Raw sync values are per-rank request counts on decode-family
# rounds -- no width division, no per-algorithm enumeration.
cuda_graph_bs = max(forward_batch.original_global_num_tokens_cpu)
cuda_graph_bs = self._max_dp_batch_size(forward_batch)
else:
cuda_graph_bs = forward_batch.batch_size
@@ -1103,15 +1110,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
else:
raw_num_token = raw_bs * self.captured_req_width
if self.require_mlp_tp_gather:
max_num_tokens = max(forward_batch.global_num_tokens_cpu)
max_batch_size = (
max_num_tokens / self.captured_req_width
if self.model_runner.spec_algorithm.is_eagle()
or self.model_runner.spec_algorithm.is_standalone()
or self.model_runner.spec_algorithm.is_dflash_family()
else max_num_tokens
)
bs = self._pad_to_bucket(int(max_batch_size), self.capture_bs)
max_batch_size = self._max_dp_batch_size(forward_batch)
bs = self._pad_to_bucket(max_batch_size, self.capture_bs)
else:
bs = self._pad_to_bucket(raw_bs, self.capture_bs)
padded_num_tokens = bs * self.captured_req_width