[Fix] Apply the attention-CP broadcast result in PP dynamic-chunk profiling (#37669)

This commit is contained in:
Liangsheng Yin
2026-09-02 17:20:07 -07:00
committed by GitHub
parent 3421d4375b
commit 5c46ce37f5
4 changed files with 30 additions and 67 deletions
@@ -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(