[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
|
# 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
|
# 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
|
||||||
import torch.distributed
|
import torch.distributed
|
||||||
|
|
||||||
|
from sglang.srt.utils import broadcast_pyobj
|
||||||
|
|
||||||
from .parallel_state import (
|
from .parallel_state import (
|
||||||
|
get_attn_cp_group,
|
||||||
get_attn_tp_group,
|
get_attn_tp_group,
|
||||||
get_moe_ep_group,
|
get_moe_ep_group,
|
||||||
get_moe_tp_group,
|
get_moe_tp_group,
|
||||||
@@ -96,6 +99,23 @@ def broadcast_tensor_dict(
|
|||||||
return get_tp_group().broadcast_tensor_dict(tensor_dict, src)
|
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:
|
def attention_tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
||||||
"""All-reduce the input tensor across attention parallel group."""
|
"""All-reduce the input tensor across attention parallel group."""
|
||||||
return get_attn_tp_group().all_reduce(input_)
|
return get_attn_tp_group().all_reduce(input_)
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ import zmq
|
|||||||
from torch.distributed import ReduceOp, all_reduce, barrier
|
from torch.distributed import ReduceOp, all_reduce, barrier
|
||||||
|
|
||||||
from sglang.srt.disaggregation.utils import prepare_abort
|
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.environ import envs
|
||||||
from sglang.srt.managers.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
BatchTokenizedEmbeddingReqInput,
|
BatchTokenizedEmbeddingReqInput,
|
||||||
@@ -164,21 +165,7 @@ class SchedulerRequestReceiver:
|
|||||||
work_reqs = None
|
work_reqs = None
|
||||||
control_reqs = None
|
control_reqs = None
|
||||||
|
|
||||||
if self.ps.attn_tp_size != 1:
|
work_reqs = attn_cp_tp_broadcast_pyobj(work_reqs)
|
||||||
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],
|
|
||||||
)
|
|
||||||
|
|
||||||
# When dp_attention_local_control_broadcast is enabled, each DP
|
# When dp_attention_local_control_broadcast is enabled, each DP
|
||||||
# group leader already receives control messages from the DP
|
# group leader already receives control messages from the DP
|
||||||
@@ -190,20 +177,7 @@ class SchedulerRequestReceiver:
|
|||||||
or is_ep_scale_joiner()
|
or is_ep_scale_joiner()
|
||||||
)
|
)
|
||||||
if _local_ctrl:
|
if _local_ctrl:
|
||||||
if self.ps.attn_tp_size != 1:
|
control_reqs = attn_cp_tp_broadcast_pyobj(control_reqs)
|
||||||
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],
|
|
||||||
)
|
|
||||||
elif self.ps.tp_size != 1:
|
elif self.ps.tp_size != 1:
|
||||||
control_reqs = broadcast_pyobj(
|
control_reqs = broadcast_pyobj(
|
||||||
control_reqs,
|
control_reqs,
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from tqdm import tqdm
|
|||||||
|
|
||||||
from sglang.srt.disaggregation.base.conn import KVPoll
|
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.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.distributed.parallel_state import P2PWork
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import (
|
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,
|
pop_auxiliary_output_from_pp_tensors,
|
||||||
)
|
)
|
||||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
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
|
from sglang.srt.utils.common import get_device_module, is_xpu
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -735,24 +736,7 @@ class SchedulerPPMixin:
|
|||||||
f"seq_lens={seq_lens}, latencies_ms={latencies}"
|
f"seq_lens={seq_lens}, latencies_ms={latencies}"
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.ps.attn_tp_size > 1:
|
seq_lens, latencies = attn_cp_tp_broadcast_pyobj([seq_lens, latencies])
|
||||||
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],
|
|
||||||
)
|
|
||||||
|
|
||||||
# Broadcast data to all ranks
|
# Broadcast data to all ranks
|
||||||
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
||||||
@@ -997,22 +981,7 @@ class SchedulerPPMixin:
|
|||||||
else:
|
else:
|
||||||
data = None
|
data = None
|
||||||
|
|
||||||
if self.ps.attn_tp_size > 1:
|
data = attn_cp_tp_broadcast_pyobj(data)
|
||||||
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],
|
|
||||||
)
|
|
||||||
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
def _pp_prepare_tensor_dict(
|
def _pp_prepare_tensor_dict(
|
||||||
|
|||||||
@@ -103,8 +103,8 @@ class TestPPCPRankOffsets(unittest.TestCase):
|
|||||||
side_effect=fake_point_to_point_pyobj,
|
side_effect=fake_point_to_point_pyobj,
|
||||||
),
|
),
|
||||||
patch(
|
patch(
|
||||||
"sglang.srt.managers.scheduler_pp_mixin.broadcast_pyobj",
|
"sglang.srt.managers.scheduler_pp_mixin.attn_cp_tp_broadcast_pyobj",
|
||||||
side_effect=lambda data, *args, **kwargs: data,
|
side_effect=lambda data: data,
|
||||||
),
|
),
|
||||||
):
|
):
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
|
|||||||
Reference in New Issue
Block a user