Opt tp: tp attn support tp reduce scattered input (#10568)
This commit is contained in:
@@ -38,7 +38,10 @@ import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.distributed.parallel_state import get_moe_expert_parallel_world_size
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
get_moe_expert_parallel_world_size,
|
||||
get_tensor_model_parallel_world_size,
|
||||
)
|
||||
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
DpPaddingMode,
|
||||
@@ -766,6 +769,13 @@ class ForwardBatch:
|
||||
else:
|
||||
bs = self.batch_size = num_tokens
|
||||
|
||||
# padding
|
||||
self._pad_inputs_to_size(model_runner, num_tokens, bs)
|
||||
self.global_num_tokens_cpu = global_num_tokens
|
||||
global_num_tokens_pinned = torch.tensor(global_num_tokens, pin_memory=True)
|
||||
self.global_num_tokens_gpu.copy_(global_num_tokens_pinned, non_blocking=True)
|
||||
|
||||
def _pad_inputs_to_size(self, model_runner: ModelRunner, num_tokens, bs):
|
||||
# padding
|
||||
self.input_ids = self._pad_tensor_to_size(self.input_ids, num_tokens)
|
||||
self.req_pool_indices = self._pad_tensor_to_size(self.req_pool_indices, bs)
|
||||
@@ -788,9 +798,6 @@ class ForwardBatch:
|
||||
if self.encoder_lens is not None:
|
||||
self.encoder_lens = self._pad_tensor_to_size(self.encoder_lens, bs)
|
||||
self.positions = self._pad_tensor_to_size(self.positions, num_tokens)
|
||||
self.global_num_tokens_cpu = global_num_tokens
|
||||
global_num_tokens_pinned = torch.tensor(global_num_tokens, pin_memory=True)
|
||||
self.global_num_tokens_gpu.copy_(global_num_tokens_pinned, non_blocking=True)
|
||||
|
||||
if self.mrope_positions is not None:
|
||||
self.mrope_positions = self._pad_tensor_to_size(self.mrope_positions, bs)
|
||||
@@ -818,6 +825,19 @@ class ForwardBatch:
|
||||
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
|
||||
|
||||
attn_tp_context = get_attn_tp_context()
|
||||
input_scattered = attn_tp_context.use_input_scattered(self)
|
||||
if not input_scattered:
|
||||
return
|
||||
assert self.forward_mode.is_extend()
|
||||
tokens = self.input_ids.shape[0]
|
||||
rank_size = get_tensor_model_parallel_world_size()
|
||||
tokens_padded = (tokens + rank_size - 1) // rank_size * rank_size
|
||||
self._pad_inputs_to_size(model_runner, tokens_padded, self.batch_size)
|
||||
|
||||
def post_forward_mlp_sync_batch(self, logits_output: LogitsProcessorOutput):
|
||||
|
||||
self.forward_mode = getattr(self, "_original_forward_mode", self.forward_mode)
|
||||
|
||||
@@ -2221,6 +2221,8 @@ class ModelRunner:
|
||||
# For MLP sync
|
||||
if forward_batch.global_num_tokens_cpu is not None:
|
||||
forward_batch.prepare_mlp_sync_batch(self)
|
||||
else:
|
||||
forward_batch.prepare_attn_tp_scatter_input(self)
|
||||
|
||||
if forward_batch.forward_mode.is_decode():
|
||||
ret = self.forward_decode(
|
||||
|
||||
Reference in New Issue
Block a user