Opt tp: tp attn support tp reduce scattered input (#10568)

This commit is contained in:
Yongfei Xu
2025-11-15 18:08:12 +08:00
committed by GitHub
parent 4a10e37ba7
commit d91b16eb16
7 changed files with 275 additions and 36 deletions

View File

@@ -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)

View File

@@ -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(