[DP Attention] Refactor: adding some utility functions (#9136)

This commit is contained in:
Cheng Wan
2025-08-13 21:08:06 -07:00
committed by GitHub
parent b3363cc1aa
commit b87aacb5c5
21 changed files with 216 additions and 159 deletions

View File

@@ -27,7 +27,7 @@ from sglang.srt.distributed import (
tensor_model_parallel_all_gather,
)
from sglang.srt.layers.dp_attention import (
DPPaddingMode,
DpPaddingMode,
attn_tp_all_gather,
attn_tp_all_gather_into_tensor,
dp_gather_replicate,
@@ -35,7 +35,9 @@ from sglang.srt.layers.dp_attention import (
get_attention_dp_rank,
get_attention_dp_size,
get_attention_tp_size,
get_global_dp_buffer,
get_local_attention_dp_size,
set_dp_buffer_len,
)
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.managers.schedule_batch import global_server_args_dict
@@ -108,14 +110,12 @@ class LogitsMetadata:
# The start position of local hidden states.
dp_local_start_pos: Optional[torch.Tensor] = None
dp_local_num_tokens: Optional[torch.Tensor] = None
gathered_buffer: Optional[torch.Tensor] = None
# Buffer to gather logits from all ranks.
forward_batch_gathered_buffer: Optional[torch.Tensor] = None
global_dp_buffer_len: Optional[int] = None
# Number of tokens to sample per DP rank
global_num_tokens_for_logprob_cpu: Optional[torch.Tensor] = None
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor] = None
# The gather mode for DP attention
dp_padding_mode: Optional[DPPaddingMode] = None
dp_padding_mode: Optional[DpPaddingMode] = None
# for padding
padded_static_len: int = -1
@@ -164,11 +164,10 @@ class LogitsMetadata:
global_num_tokens_gpu=forward_batch.global_num_tokens_gpu,
dp_local_start_pos=forward_batch.dp_local_start_pos,
dp_local_num_tokens=forward_batch.dp_local_num_tokens,
gathered_buffer=forward_batch.gathered_buffer,
forward_batch_gathered_buffer=forward_batch.gathered_buffer,
global_dp_buffer_len=forward_batch.global_dp_buffer_len,
global_num_tokens_for_logprob_cpu=forward_batch.global_num_tokens_for_logprob_cpu,
global_num_tokens_for_logprob_gpu=forward_batch.global_num_tokens_for_logprob_gpu,
dp_padding_mode=DPPaddingMode.SUM_LEN,
dp_padding_mode=DpPaddingMode.SUM_LEN,
)
def compute_dp_attention_metadata(self):
@@ -188,16 +187,11 @@ class LogitsMetadata:
if self.global_num_tokens_for_logprob_cpu is not None:
# create a smaller buffer to reduce peak memory usage
self.gathered_buffer = torch.empty(
(
sum(self.global_num_tokens_for_logprob_cpu),
self.gathered_buffer.shape[1],
),
dtype=self.gathered_buffer.dtype,
device=self.gathered_buffer.device,
)
self.global_dp_buffer_len = sum(self.global_num_tokens_for_logprob_cpu)
else:
self.gathered_buffer = torch.empty_like(self.gathered_buffer)
self.global_dp_buffer_len = self.global_dp_buffer_len
set_dp_buffer_len(self.global_dp_buffer_len, self.dp_local_num_tokens)
class LogitsProcessor(nn.Module):
@@ -443,7 +437,7 @@ class LogitsProcessor(nn.Module):
if self.do_tensor_parallel_all_gather_dp_attn:
logits_metadata.compute_dp_attention_metadata()
hidden_states, local_hidden_states = (
logits_metadata.gathered_buffer,
get_global_dp_buffer(),
hidden_states,
)
dp_gather_replicate(hidden_states, local_hidden_states, logits_metadata)