[scheduler] remove scheduler allgather for best throughout (#14294)
This commit is contained in:
@@ -195,6 +195,7 @@ class Envs:
|
|||||||
SGLANG_SCHEDULER_MAX_RECV_PER_POLL = EnvInt(-1)
|
SGLANG_SCHEDULER_MAX_RECV_PER_POLL = EnvInt(-1)
|
||||||
SGLANG_EXPERIMENTAL_CPP_RADIX_TREE = EnvBool(False)
|
SGLANG_EXPERIMENTAL_CPP_RADIX_TREE = EnvBool(False)
|
||||||
SGLANG_DYNAMIC_CHUNKING_SMOOTH_FACTOR = EnvFloat(0.75)
|
SGLANG_DYNAMIC_CHUNKING_SMOOTH_FACTOR = EnvFloat(0.75)
|
||||||
|
SGLANG_SCHEDULER_SKIP_ALL_GATHER = EnvBool(False)
|
||||||
SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE = EnvBool(False)
|
SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE = EnvBool(False)
|
||||||
|
|
||||||
# Test: pd-disaggregation
|
# Test: pd-disaggregation
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Callable
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer
|
from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.utils.common import require_mlp_tp_gather
|
from sglang.srt.utils.common import require_mlp_tp_gather
|
||||||
|
|
||||||
@@ -73,6 +74,7 @@ def _update_gather_batch(
|
|||||||
batch: ScheduleBatch,
|
batch: ScheduleBatch,
|
||||||
mlp_sync_info: MLPSyncBatchInfo,
|
mlp_sync_info: MLPSyncBatchInfo,
|
||||||
require_mlp_tp_gather: bool,
|
require_mlp_tp_gather: bool,
|
||||||
|
skip_all_gather=False,
|
||||||
):
|
):
|
||||||
# TODO: handle the case when moe_dense_tp_size != 1
|
# TODO: handle the case when moe_dense_tp_size != 1
|
||||||
if not require_mlp_tp_gather:
|
if not require_mlp_tp_gather:
|
||||||
@@ -83,9 +85,10 @@ def _update_gather_batch(
|
|||||||
batch.global_num_tokens_for_logprob = (
|
batch.global_num_tokens_for_logprob = (
|
||||||
mlp_sync_info.global_num_tokens_for_logprob
|
mlp_sync_info.global_num_tokens_for_logprob
|
||||||
)
|
)
|
||||||
batch.is_extend_in_batch = mlp_sync_info.is_extend_in_batch
|
if not skip_all_gather:
|
||||||
batch.tbo_split_seq_index = mlp_sync_info.tbo_split_seq_index
|
batch.is_extend_in_batch = mlp_sync_info.is_extend_in_batch
|
||||||
batch.global_forward_mode = mlp_sync_info.global_forward_mode
|
batch.tbo_split_seq_index = mlp_sync_info.tbo_split_seq_index
|
||||||
|
batch.global_forward_mode = mlp_sync_info.global_forward_mode
|
||||||
|
|
||||||
# Check forward mode for cuda graph
|
# Check forward mode for cuda graph
|
||||||
batch.can_run_dp_cuda_graph = mlp_sync_info.can_cuda_graph
|
batch.can_run_dp_cuda_graph = mlp_sync_info.can_cuda_graph
|
||||||
@@ -124,6 +127,7 @@ def prepare_mlp_sync_batch_raw(
|
|||||||
# When return_logprob = False, only need last token per request
|
# When return_logprob = False, only need last token per request
|
||||||
num_tokens_for_logprob = local_batch.batch_size()
|
num_tokens_for_logprob = local_batch.batch_size()
|
||||||
|
|
||||||
|
skip_all_gather = envs.SGLANG_SCHEDULER_SKIP_ALL_GATHER.get()
|
||||||
can_cuda_graph = (
|
can_cuda_graph = (
|
||||||
local_batch is None
|
local_batch is None
|
||||||
or local_batch.forward_mode.is_decode_or_idle()
|
or local_batch.forward_mode.is_decode_or_idle()
|
||||||
@@ -154,15 +158,17 @@ def prepare_mlp_sync_batch_raw(
|
|||||||
local_can_run_tbo=local_can_run_tbo,
|
local_can_run_tbo=local_can_run_tbo,
|
||||||
local_forward_mode=local_forward_mode,
|
local_forward_mode=local_forward_mode,
|
||||||
)
|
)
|
||||||
mlp_sync_info.all_gather(device=device, group=group)
|
|
||||||
|
|
||||||
mlp_sync_info.tbo_split_seq_index, mlp_sync_info.global_forward_mode = (
|
if not skip_all_gather:
|
||||||
tbo_preparer.compute_output(
|
mlp_sync_info.all_gather(device=device, group=group)
|
||||||
mlp_sync_info.tp0_info[:, 4:6],
|
|
||||||
|
mlp_sync_info.tbo_split_seq_index, mlp_sync_info.global_forward_mode = (
|
||||||
|
tbo_preparer.compute_output(
|
||||||
|
mlp_sync_info.tp0_info[:, 4:6],
|
||||||
|
)
|
||||||
)
|
)
|
||||||
)
|
|
||||||
|
|
||||||
need_idle_batch = max(mlp_sync_info.global_num_tokens) > 0
|
need_idle_batch = skip_all_gather or max(mlp_sync_info.global_num_tokens) > 0
|
||||||
if need_idle_batch:
|
if need_idle_batch:
|
||||||
batch_to_gather = local_batch
|
batch_to_gather = local_batch
|
||||||
if local_batch is None:
|
if local_batch is None:
|
||||||
@@ -170,7 +176,9 @@ def prepare_mlp_sync_batch_raw(
|
|||||||
elif local_batch.forward_mode.is_prebuilt():
|
elif local_batch.forward_mode.is_prebuilt():
|
||||||
# NOTE: for prebuilt batch, we add an inner idle batch to run MLP sync
|
# NOTE: for prebuilt batch, we add an inner idle batch to run MLP sync
|
||||||
batch_to_gather = local_batch.inner_idle_batch = get_idle_batch()
|
batch_to_gather = local_batch.inner_idle_batch = get_idle_batch()
|
||||||
_update_gather_batch(batch_to_gather, mlp_sync_info, require_mlp_tp_gather)
|
_update_gather_batch(
|
||||||
|
batch_to_gather, mlp_sync_info, require_mlp_tp_gather, skip_all_gather
|
||||||
|
)
|
||||||
|
|
||||||
return local_batch
|
return local_batch
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user