[DP Attention] Refactor: adding some utility functions (#9136)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user