[Fix] Apply the attention-CP broadcast result in PP dynamic-chunk profiling (#37669)
This commit is contained in:
@@ -2,12 +2,15 @@
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.6.4.post1/vllm/distributed/communication_op.py
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
|
||||
from sglang.srt.utils import broadcast_pyobj
|
||||
|
||||
from .parallel_state import (
|
||||
get_attn_cp_group,
|
||||
get_attn_tp_group,
|
||||
get_moe_ep_group,
|
||||
get_moe_tp_group,
|
||||
@@ -96,6 +99,23 @@ def broadcast_tensor_dict(
|
||||
return get_tp_group().broadcast_tensor_dict(tensor_dict, src)
|
||||
|
||||
|
||||
def attn_cp_tp_broadcast_pyobj(data: List[Any]) -> List[Any]:
|
||||
"""Broadcast from the (attn-TP 0, attn-CP 0) rank to every rank of its DP shard."""
|
||||
# attn-TP and attn-CP are orthogonal factors of the shard and no single group
|
||||
# covers both, so broadcast along TP first and then along CP.
|
||||
tp_group = get_attn_tp_group()
|
||||
if tp_group.world_size > 1:
|
||||
data = broadcast_pyobj(
|
||||
data, tp_group.rank, tp_group.cpu_group, src=tp_group.ranks[0]
|
||||
)
|
||||
cp_group = get_attn_cp_group()
|
||||
if cp_group.world_size > 1:
|
||||
data = broadcast_pyobj(
|
||||
data, cp_group.rank, cp_group.cpu_group, src=cp_group.ranks[0]
|
||||
)
|
||||
return data
|
||||
|
||||
|
||||
def attention_tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
||||
"""All-reduce the input tensor across attention parallel group."""
|
||||
return get_attn_tp_group().all_reduce(input_)
|
||||
|
||||
@@ -17,6 +17,7 @@ import zmq
|
||||
from torch.distributed import ReduceOp, all_reduce, barrier
|
||||
|
||||
from sglang.srt.disaggregation.utils import prepare_abort
|
||||
from sglang.srt.distributed.communication_op import attn_cp_tp_broadcast_pyobj
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.io_struct import (
|
||||
BatchTokenizedEmbeddingReqInput,
|
||||
@@ -164,21 +165,7 @@ class SchedulerRequestReceiver:
|
||||
work_reqs = None
|
||||
control_reqs = None
|
||||
|
||||
if self.ps.attn_tp_size != 1:
|
||||
work_reqs = broadcast_pyobj(
|
||||
work_reqs,
|
||||
self.attn_tp_group.rank,
|
||||
self.attn_tp_cpu_group,
|
||||
src=self.attn_tp_group.ranks[0],
|
||||
)
|
||||
|
||||
if self.ps.attn_cp_size != 1:
|
||||
work_reqs = broadcast_pyobj(
|
||||
work_reqs,
|
||||
self.attn_cp_group.rank,
|
||||
self.attn_cp_cpu_group,
|
||||
src=self.attn_cp_group.ranks[0],
|
||||
)
|
||||
work_reqs = attn_cp_tp_broadcast_pyobj(work_reqs)
|
||||
|
||||
# When dp_attention_local_control_broadcast is enabled, each DP
|
||||
# group leader already receives control messages from the DP
|
||||
@@ -190,20 +177,7 @@ class SchedulerRequestReceiver:
|
||||
or is_ep_scale_joiner()
|
||||
)
|
||||
if _local_ctrl:
|
||||
if self.ps.attn_tp_size != 1:
|
||||
control_reqs = broadcast_pyobj(
|
||||
control_reqs,
|
||||
self.attn_tp_group.rank,
|
||||
self.attn_tp_cpu_group,
|
||||
src=self.attn_tp_group.ranks[0],
|
||||
)
|
||||
if self.ps.attn_cp_size != 1:
|
||||
control_reqs = broadcast_pyobj(
|
||||
control_reqs,
|
||||
self.attn_cp_group.rank,
|
||||
self.attn_cp_cpu_group,
|
||||
src=self.attn_cp_group.ranks[0],
|
||||
)
|
||||
control_reqs = attn_cp_tp_broadcast_pyobj(control_reqs)
|
||||
elif self.ps.tp_size != 1:
|
||||
control_reqs = broadcast_pyobj(
|
||||
control_reqs,
|
||||
|
||||
@@ -15,6 +15,7 @@ from tqdm import tqdm
|
||||
|
||||
from sglang.srt.disaggregation.base.conn import KVPoll
|
||||
from sglang.srt.disaggregation.utils import poll_and_all_reduce_attn_cp_tp_group
|
||||
from sglang.srt.distributed.communication_op import attn_cp_tp_broadcast_pyobj
|
||||
from sglang.srt.distributed.parallel_state import P2PWork
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
@@ -44,7 +45,7 @@ from sglang.srt.sampling.sampling_observer_pp import (
|
||||
pop_auxiliary_output_from_pp_tensors,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
|
||||
from sglang.srt.utils import DynamicGradMode, point_to_point_pyobj
|
||||
from sglang.srt.utils.common import get_device_module, is_xpu
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -735,24 +736,7 @@ class SchedulerPPMixin:
|
||||
f"seq_lens={seq_lens}, latencies_ms={latencies}"
|
||||
)
|
||||
|
||||
if self.ps.attn_tp_size > 1:
|
||||
data_to_sync_tp = [seq_lens, latencies]
|
||||
data_to_sync_tp = broadcast_pyobj(
|
||||
data_to_sync_tp,
|
||||
self.attn_tp_group.rank,
|
||||
self.attn_tp_cpu_group,
|
||||
src=self.attn_tp_group.ranks[0],
|
||||
)
|
||||
seq_lens, latencies = data_to_sync_tp
|
||||
|
||||
if self.ps.attn_cp_size > 1:
|
||||
data_to_sync_tp = [seq_lens, latencies]
|
||||
data_to_sync_tp = broadcast_pyobj(
|
||||
data_to_sync_tp,
|
||||
self.attn_cp_group.rank,
|
||||
self.attn_cp_cpu_group,
|
||||
src=self.attn_cp_group.ranks[0],
|
||||
)
|
||||
seq_lens, latencies = attn_cp_tp_broadcast_pyobj([seq_lens, latencies])
|
||||
|
||||
# Broadcast data to all ranks
|
||||
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
||||
@@ -997,22 +981,7 @@ class SchedulerPPMixin:
|
||||
else:
|
||||
data = None
|
||||
|
||||
if self.ps.attn_tp_size > 1:
|
||||
data = broadcast_pyobj(
|
||||
data,
|
||||
self.attn_tp_group.rank,
|
||||
self.attn_tp_cpu_group,
|
||||
src=self.attn_tp_group.ranks[0],
|
||||
)
|
||||
|
||||
if self.ps.attn_cp_size > 1:
|
||||
data = broadcast_pyobj(
|
||||
data,
|
||||
self.attn_cp_group.rank,
|
||||
self.attn_cp_cpu_group,
|
||||
src=self.attn_cp_group.ranks[0],
|
||||
)
|
||||
|
||||
data = attn_cp_tp_broadcast_pyobj(data)
|
||||
return data
|
||||
|
||||
def _pp_prepare_tensor_dict(
|
||||
|
||||
Reference in New Issue
Block a user