Give the attention-DP width and rank one home (#40067)
This commit is contained in:
@@ -1739,7 +1739,7 @@ class _SGLangPlugin(_FrameworkPlugin):
|
|||||||
info["moe_tp_size"] = parallel.moe_tp_size
|
info["moe_tp_size"] = parallel.moe_tp_size
|
||||||
info["moe_dp_rank"] = parallel.moe_dp_rank
|
info["moe_dp_rank"] = parallel.moe_dp_rank
|
||||||
info["moe_dp_size"] = self._dp_attn.get_moe_cp_size()
|
info["moe_dp_size"] = self._dp_attn.get_moe_cp_size()
|
||||||
except (AttributeError, AssertionError, ValueError):
|
except (AttributeError, AssertionError, ValueError, RuntimeError):
|
||||||
info["distributed_error"] = True
|
info["distributed_error"] = True
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -1747,11 +1747,12 @@ class _SGLangPlugin(_FrameworkPlugin):
|
|||||||
info["enable_dp_attention"] = self._dp_attn.is_dp_attention_enabled()
|
info["enable_dp_attention"] = self._dp_attn.is_dp_attention_enabled()
|
||||||
info["attn_tp_rank"] = parallel.attn_tp_rank
|
info["attn_tp_rank"] = parallel.attn_tp_rank
|
||||||
info["attn_tp_size"] = parallel.attn_tp_size
|
info["attn_tp_size"] = parallel.attn_tp_size
|
||||||
info["attn_dp_rank"] = self._dp_attn.get_attention_dp_rank()
|
info["attn_dp_rank"] = parallel.attn_dp_rank
|
||||||
info["attn_dp_size"] = self._dp_attn.get_attention_dp_size()
|
info["attn_dp_size"] = parallel.attn_dp_size
|
||||||
info["attn_cp_rank"] = parallel.attn_cp_rank
|
info["attn_cp_rank"] = parallel.attn_cp_rank
|
||||||
info["attn_cp_size"] = parallel.attn_cp_size
|
info["attn_cp_size"] = parallel.attn_cp_size
|
||||||
except (AttributeError, AssertionError, ValueError):
|
# An unstamped topology name raises RuntimeError.
|
||||||
|
except (AttributeError, AssertionError, ValueError, RuntimeError):
|
||||||
info["dp_attention_error"] = True
|
info["dp_attention_error"] = True
|
||||||
|
|
||||||
return info
|
return info
|
||||||
|
|||||||
@@ -33,10 +33,6 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.distributed import get_pp_group, get_world_group
|
from sglang.srt.distributed import get_pp_group, get_world_group
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import (
|
|
||||||
get_attention_dp_rank,
|
|
||||||
get_attention_dp_size,
|
|
||||||
)
|
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_disagg,
|
get_disagg,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
@@ -197,8 +193,8 @@ class CommonKVManager(BaseKVManager):
|
|||||||
self.attn_cp_rank = parallel.attn_cp_rank
|
self.attn_cp_rank = parallel.attn_cp_rank
|
||||||
self.dcp_size = parallel.attn_dcp_size
|
self.dcp_size = parallel.attn_dcp_size
|
||||||
self.dcp_rank = parallel.attn_dcp_rank
|
self.dcp_rank = parallel.attn_dcp_rank
|
||||||
self.attn_dp_size = get_attention_dp_size()
|
self.attn_dp_size = parallel.attn_dp_size
|
||||||
self.attn_dp_rank = get_attention_dp_rank()
|
self.attn_dp_rank = parallel.attn_dp_rank
|
||||||
self.system_dp_size = (
|
self.system_dp_size = (
|
||||||
1 if get_parallel().enable_dp_attention else get_parallel().dp_size
|
1 if get_parallel().enable_dp_attention else get_parallel().dp_size
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import functools
|
import functools
|
||||||
import logging
|
import logging
|
||||||
from contextlib import contextmanager
|
|
||||||
from enum import IntEnum, auto
|
from enum import IntEnum, auto
|
||||||
from typing import TYPE_CHECKING, List, Optional, Tuple
|
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||||
|
|
||||||
@@ -53,9 +52,6 @@ logger = logging.getLogger(__name__)
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
_ATTN_DP_RANK: Optional[int] = None
|
|
||||||
_ATTN_DP_SIZE: Optional[int] = None
|
|
||||||
|
|
||||||
|
|
||||||
def world_dp_gather_enabled() -> bool:
|
def world_dp_gather_enabled() -> bool:
|
||||||
"""Whether DP gathers should use expanded WORLD after joiner admission."""
|
"""Whether DP gathers should use expanded WORLD after joiner admission."""
|
||||||
@@ -68,10 +64,9 @@ def enable_joiner_all_gather():
|
|||||||
|
|
||||||
|
|
||||||
def update_dp_attention_post_scale(new_dp_size: int, new_dp_rank: int):
|
def update_dp_attention_post_scale(new_dp_size: int, new_dp_rank: int):
|
||||||
global _ATTN_DP_SIZE, _ATTN_DP_RANK
|
get_parallel().override_permanently(
|
||||||
_ATTN_DP_SIZE = new_dp_size
|
attn_dp_size=new_dp_size, attn_dp_rank=new_dp_rank
|
||||||
_ATTN_DP_RANK = new_dp_rank
|
)
|
||||||
get_parallel().override_permanently(attn_dp_size=new_dp_size)
|
|
||||||
get_flags().dp.use_world_group_for_gather = True
|
get_flags().dp.use_world_group_for_gather = True
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"[Elastic EP] dp_attention switched to WORLD: dp_size=%d dp_rank=%d",
|
"[Elastic EP] dp_attention switched to WORLD: dp_size=%d dp_rank=%d",
|
||||||
@@ -101,7 +96,7 @@ class DpPaddingMode(IntEnum):
|
|||||||
def get_dp_padding_mode(
|
def get_dp_padding_mode(
|
||||||
cls, is_extend_in_batch, global_num_tokens: List[int]
|
cls, is_extend_in_batch, global_num_tokens: List[int]
|
||||||
) -> DpPaddingMode:
|
) -> DpPaddingMode:
|
||||||
dp_size = get_attention_dp_size()
|
dp_size = get_parallel().attn_dp_size
|
||||||
|
|
||||||
# (trangdough) pplx-kernels a2a is a symmetric collective: every EP rank
|
# (trangdough) pplx-kernels a2a is a symmetric collective: every EP rank
|
||||||
# must dispatch the same number of tokens or the device-side handshake
|
# must dispatch the same number of tokens or the device-side handshake
|
||||||
@@ -375,7 +370,6 @@ def initialize_dp_attention(
|
|||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
):
|
):
|
||||||
global _ATTN_DP_RANK, _ATTN_DP_SIZE
|
|
||||||
dp = get_flags().dp
|
dp = get_flags().dp
|
||||||
dp.max_len_with_idle = (
|
dp.max_len_with_idle = (
|
||||||
getattr(model_config.hf_config, "hybrid_override_pattern", None) is not None
|
getattr(model_config.hf_config, "hybrid_override_pattern", None) is not None
|
||||||
@@ -389,13 +383,12 @@ def initialize_dp_attention(
|
|||||||
tp_rank = get_tensor_model_parallel_rank()
|
tp_rank = get_tensor_model_parallel_rank()
|
||||||
tp_size = get_tensor_model_parallel_world_size()
|
tp_size = get_tensor_model_parallel_world_size()
|
||||||
|
|
||||||
_, _, _ATTN_DP_RANK, _ATTN_DP_SIZE = compute_dp_attention_world_info(
|
_, _, attn_dp_rank, attn_dp_size = compute_dp_attention_world_info(
|
||||||
enable_dp_attention, tp_rank, tp_size, dp_size, attn_cp_size
|
enable_dp_attention, tp_rank, tp_size, dp_size, attn_cp_size
|
||||||
)
|
)
|
||||||
get_parallel().override_permanently(attn_dp_size=_ATTN_DP_SIZE)
|
|
||||||
|
|
||||||
if get_exec().moe.elastic_ep_backend is not None and get_parallel().max_ep_size:
|
if get_exec().moe.elastic_ep_backend is not None and get_parallel().max_ep_size:
|
||||||
_ATTN_DP_RANK = tp_rank + get_parallel().ep_join_rank_offset
|
attn_dp_rank = tp_rank + get_parallel().ep_join_rank_offset
|
||||||
# Reads the resolution, not a bag: this runs under
|
# Reads the resolution, not a bag: this runs under
|
||||||
# `initialize_dp_attention`, which the weight-cache daemon calls from
|
# `initialize_dp_attention`, which the weight-cache daemon calls from
|
||||||
# `_init_distributed` -- and other callers reach it from processes
|
# `_init_distributed` -- and other callers reach it from processes
|
||||||
@@ -404,6 +397,13 @@ def initialize_dp_attention(
|
|||||||
if ep_scale_joiner_of(resolving_view(server_args)):
|
if ep_scale_joiner_of(resolving_view(server_args)):
|
||||||
dp.joiner_skip_all_gather = True
|
dp.joiner_skip_all_gather = True
|
||||||
|
|
||||||
|
# Stamped together, after the elastic adjustment: the width and the rank
|
||||||
|
# describe one topology, and a reader that caught them mid-update would
|
||||||
|
# see this process placed in a group it is not in.
|
||||||
|
get_parallel().override_permanently(
|
||||||
|
attn_dp_size=attn_dp_size, attn_dp_rank=attn_dp_rank
|
||||||
|
)
|
||||||
|
|
||||||
_DpGatheredBufferWrapper.set_metadata(
|
_DpGatheredBufferWrapper.set_metadata(
|
||||||
hidden_size=model_config.hidden_size,
|
hidden_size=model_config.hidden_size,
|
||||||
dtype=model_config.dtype,
|
dtype=model_config.dtype,
|
||||||
@@ -419,42 +419,9 @@ def is_allocation_symmetric() -> bool:
|
|||||||
return not is_dp_attention_enabled() or is_dp_max_padding()
|
return not is_dp_attention_enabled() or is_dp_max_padding()
|
||||||
|
|
||||||
|
|
||||||
def get_attention_dp_rank() -> int:
|
|
||||||
assert _ATTN_DP_RANK is not None, "dp attention not initialized!"
|
|
||||||
return _ATTN_DP_RANK
|
|
||||||
|
|
||||||
|
|
||||||
def get_attention_dp_size() -> int:
|
|
||||||
assert _ATTN_DP_SIZE is not None, "dp attention not initialized!"
|
|
||||||
return _ATTN_DP_SIZE
|
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
|
||||||
def disable_dp_size():
|
|
||||||
"""Run without DP attention until this scope ends.
|
|
||||||
|
|
||||||
This is for draft workers of speculative decoding, which run the draft model
|
|
||||||
at a different width from the target model's workers.
|
|
||||||
|
|
||||||
The scope replaces both the module global that ``get_attention_dp_size()``
|
|
||||||
reads and the derived width the runtime context answers with, so the two
|
|
||||||
spellings of the name cannot disagree inside it.
|
|
||||||
"""
|
|
||||||
global _ATTN_DP_SIZE
|
|
||||||
assert _ATTN_DP_SIZE is not None, "dp attention not initialized!"
|
|
||||||
|
|
||||||
old_dp_size = _ATTN_DP_SIZE
|
|
||||||
_ATTN_DP_SIZE = 1
|
|
||||||
try:
|
|
||||||
with get_parallel().override(attn_dp_size=1):
|
|
||||||
yield
|
|
||||||
finally:
|
|
||||||
_ATTN_DP_SIZE = old_dp_size
|
|
||||||
|
|
||||||
|
|
||||||
def get_dp_local_info(forward_batch: ForwardBatch) -> Tuple[torch.Tensor, torch.Tensor]:
|
def get_dp_local_info(forward_batch: ForwardBatch) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
# `get_dp_local_info` is only called in global DP gather and scatter. We use global DP rank here.
|
# `get_dp_local_info` is only called in global DP gather and scatter. We use global DP rank here.
|
||||||
dp_rank = get_attention_dp_rank()
|
dp_rank = get_parallel().attn_dp_rank
|
||||||
|
|
||||||
if forward_batch.dp_local_start_pos is None:
|
if forward_batch.dp_local_start_pos is None:
|
||||||
cumtokens = torch.cumsum(forward_batch.global_num_tokens_gpu, dim=0)
|
cumtokens = torch.cumsum(forward_batch.global_num_tokens_gpu, dim=0)
|
||||||
@@ -478,7 +445,7 @@ def get_dp_local_slice_cpu(
|
|||||||
# CPU (start, length) slice for DP-local data in a rank-padded buffer.
|
# CPU (start, length) slice for DP-local data in a rank-padded buffer.
|
||||||
# Returns Python ints (no D2H sync) and handles the cuda-graph-padded layout.
|
# Returns Python ints (no D2H sync) and handles the cuda-graph-padded layout.
|
||||||
global_num_tokens = forward_batch.global_num_tokens_cpu
|
global_num_tokens = forward_batch.global_num_tokens_cpu
|
||||||
dp_rank = get_attention_dp_rank()
|
dp_rank = get_parallel().attn_dp_rank
|
||||||
local_num_tokens = global_num_tokens[dp_rank]
|
local_num_tokens = global_num_tokens[dp_rank]
|
||||||
if can_run_graph:
|
if can_run_graph:
|
||||||
local_start_pos = dp_rank * cuda_graph_batch
|
local_start_pos = dp_rank * cuda_graph_batch
|
||||||
@@ -758,7 +725,7 @@ def is_dp_gatherv_active() -> bool:
|
|||||||
_USE_DP_GATHERV
|
_USE_DP_GATHERV
|
||||||
and not world_dp_gather_enabled()
|
and not world_dp_gather_enabled()
|
||||||
and get_attn_tensor_model_parallel_world_size() == 1
|
and get_attn_tensor_model_parallel_world_size() == 1
|
||||||
and get_tensor_model_parallel_world_size() == get_attention_dp_size()
|
and get_tensor_model_parallel_world_size() == get_parallel().attn_dp_size
|
||||||
and not _DpGatheredBufferWrapper.is_dp_max_padding()
|
and not _DpGatheredBufferWrapper.is_dp_max_padding()
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -792,7 +759,7 @@ def _dp_gather_via_all_gatherv(
|
|||||||
# each rank's local tensor up to sizes[rank] with zeros (matching the
|
# each rank's local tensor up to sizes[rank] with zeros (matching the
|
||||||
# buffer's reserved per-rank slot) so sum(sizes) == buffer rows and there
|
# buffer's reserved per-rank slot) so sum(sizes) == buffer rows and there
|
||||||
# is no uninitialized tail for the MoE to read.
|
# is no uninitialized tail for the MoE to read.
|
||||||
rank = get_attention_dp_rank()
|
rank = get_parallel().attn_dp_rank
|
||||||
local_rows = sizes[rank]
|
local_rows = sizes[rank]
|
||||||
if local_tokens.shape[0] == local_rows:
|
if local_tokens.shape[0] == local_rows:
|
||||||
local_real = local_tokens
|
local_real = local_tokens
|
||||||
@@ -917,7 +884,7 @@ def dp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor):
|
|||||||
if sizes is not None:
|
if sizes is not None:
|
||||||
get_tp_group().reduce_scatterv(input, output=output, sizes=sizes)
|
get_tp_group().reduce_scatterv(input, output=output, sizes=sizes)
|
||||||
return
|
return
|
||||||
if get_tensor_model_parallel_world_size() == get_attention_dp_size():
|
if get_tensor_model_parallel_world_size() == get_parallel().attn_dp_size:
|
||||||
get_tp_group().reduce_scatter_tensor(output, input)
|
get_tp_group().reduce_scatter_tensor(output, input)
|
||||||
else:
|
else:
|
||||||
scattered_local_tokens = input.tensor_split(
|
scattered_local_tokens = input.tensor_split(
|
||||||
|
|||||||
@@ -40,7 +40,6 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
dp_gather_replicate,
|
dp_gather_replicate,
|
||||||
dp_reduce_scatter_tensor,
|
dp_reduce_scatter_tensor,
|
||||||
dp_scatter,
|
dp_scatter,
|
||||||
get_attention_dp_size,
|
|
||||||
get_global_dp_buffer_len,
|
get_global_dp_buffer_len,
|
||||||
is_dp_gatherv_active,
|
is_dp_gatherv_active,
|
||||||
)
|
)
|
||||||
@@ -766,7 +765,7 @@ class EngramEmbedding(nn.Module):
|
|||||||
attn_cp_all_gather_into_tensor(all_indices, indices.contiguous())
|
attn_cp_all_gather_into_tensor(all_indices, indices.contiguous())
|
||||||
start = parallel.attn_cp_rank * local_rows
|
start = parallel.attn_cp_rank * local_rows
|
||||||
return self._lookup(all_indices)[start : start + local_rows]
|
return self._lookup(all_indices)[start : start + local_rows]
|
||||||
if self.tp_size > 1 and get_attention_dp_size() > 1:
|
if self.tp_size > 1 and get_parallel().attn_dp_size > 1:
|
||||||
return self._dp_sharded_lookup(indices, forward_batch)
|
return self._dp_sharded_lookup(indices, forward_batch)
|
||||||
return self._lookup(indices)
|
return self._lookup(indices)
|
||||||
|
|
||||||
@@ -832,7 +831,7 @@ class EngramEmbedding(nn.Module):
|
|||||||
if (
|
if (
|
||||||
padding is not None
|
padding is not None
|
||||||
and padding.is_max_len()
|
and padding.is_max_len()
|
||||||
and self.tp_size == get_attention_dp_size()
|
and self.tp_size == get_parallel().attn_dp_size
|
||||||
and rows == self.tp_size * local.shape[0]
|
and rows == self.tp_size * local.shape[0]
|
||||||
) or is_dp_gatherv_active():
|
) or is_dp_gatherv_active():
|
||||||
dp_reduce_scatter_tensor(local, values)
|
dp_reduce_scatter_tensor(local, values)
|
||||||
|
|||||||
@@ -37,7 +37,6 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.mem_cache.pool_host import HostKVCache
|
from sglang.srt.mem_cache.pool_host import HostKVCache
|
||||||
|
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
get_attention_dp_rank,
|
|
||||||
is_dp_attention_enabled,
|
is_dp_attention_enabled,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine
|
from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine
|
||||||
@@ -696,7 +695,7 @@ class HiCacheController:
|
|||||||
if is_dp_attention_enabled():
|
if is_dp_attention_enabled():
|
||||||
self.tp_rank = get_parallel().attn_tp_rank
|
self.tp_rank = get_parallel().attn_tp_rank
|
||||||
self.tp_size = get_parallel().attn_tp_size
|
self.tp_size = get_parallel().attn_tp_size
|
||||||
self.dp_rank = get_attention_dp_rank()
|
self.dp_rank = get_parallel().attn_dp_rank
|
||||||
else:
|
else:
|
||||||
self.tp_rank = get_parallel().tp_rank
|
self.tp_rank = get_parallel().tp_rank
|
||||||
self.tp_size = get_parallel().tp_size
|
self.tp_size = get_parallel().tp_size
|
||||||
|
|||||||
@@ -58,9 +58,8 @@ def _resolve_elastic_world_dp_size(
|
|||||||
return dp_size
|
return dp_size
|
||||||
|
|
||||||
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||||
from sglang.srt.layers.dp_attention import get_attention_dp_size
|
|
||||||
|
|
||||||
live_dp_size = get_attention_dp_size()
|
live_dp_size = get_parallel().attn_dp_size
|
||||||
effective_ep_size = ElasticEPStateManager.get_effective_ep_size()
|
effective_ep_size = ElasticEPStateManager.get_effective_ep_size()
|
||||||
world_size = torch.distributed.get_world_size(group)
|
world_size = torch.distributed.get_world_size(group)
|
||||||
|
|
||||||
|
|||||||
@@ -12,14 +12,13 @@ from tqdm import tqdm
|
|||||||
|
|
||||||
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 (
|
||||||
get_attention_dp_rank,
|
|
||||||
get_attention_dp_size,
|
|
||||||
is_dp_attention_enabled,
|
is_dp_attention_enabled,
|
||||||
set_is_extend_in_batch,
|
set_is_extend_in_batch,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||||
from sglang.srt.mem_cache.common import release_kv_cache
|
from sglang.srt.mem_cache.common import release_kv_cache
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.utils import broadcast_pyobj
|
from sglang.srt.utils import broadcast_pyobj
|
||||||
from sglang.srt.utils.common import get_device_module
|
from sglang.srt.utils.common import get_device_module
|
||||||
@@ -196,9 +195,9 @@ class DynamicChunkSizer:
|
|||||||
|
|
||||||
if is_dp_attention_enabled():
|
if is_dp_attention_enabled():
|
||||||
# Profiling runs one request on this rank; other DP ranks report 0.
|
# Profiling runs one request on this rank; other DP ranks report 0.
|
||||||
dp_size = get_attention_dp_size()
|
dp_size = get_parallel().attn_dp_size
|
||||||
global_num_tokens = [0] * dp_size
|
global_num_tokens = [0] * dp_size
|
||||||
dp_rank = get_attention_dp_rank()
|
dp_rank = get_parallel().attn_dp_rank
|
||||||
global_num_tokens[dp_rank] = current_seq_len
|
global_num_tokens[dp_rank] = current_seq_len
|
||||||
batch.global_num_tokens = global_num_tokens
|
batch.global_num_tokens = global_num_tokens
|
||||||
batch.global_num_tokens_for_logprob = global_num_tokens
|
batch.global_num_tokens_for_logprob = global_num_tokens
|
||||||
|
|||||||
@@ -310,15 +310,14 @@ class UMBPStore(HiCacheStorage):
|
|||||||
if dp_rank_hint is None:
|
if dp_rank_hint is None:
|
||||||
try:
|
try:
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
get_attention_dp_rank,
|
|
||||||
get_attention_dp_size,
|
|
||||||
is_dp_attention_enabled,
|
is_dp_attention_enabled,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
if is_dp_attention_enabled():
|
if is_dp_attention_enabled():
|
||||||
dp_rank_hint = get_attention_dp_rank()
|
dp_rank_hint = get_parallel().attn_dp_rank
|
||||||
dp_size_hint = get_attention_dp_size()
|
dp_size_hint = get_parallel().attn_dp_size
|
||||||
except (ImportError, AssertionError):
|
except (ImportError, AssertionError, RuntimeError):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
if local_rank_hint is not None:
|
if local_rank_hint is not None:
|
||||||
@@ -807,14 +806,13 @@ class UMBPStore(HiCacheStorage):
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
get_attention_dp_rank,
|
|
||||||
get_attention_dp_size,
|
|
||||||
is_dp_attention_enabled,
|
is_dp_attention_enabled,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
if is_dp_attention_enabled():
|
if is_dp_attention_enabled():
|
||||||
dp_rank = get_attention_dp_rank()
|
dp_rank = get_parallel().attn_dp_rank
|
||||||
dp_size = get_attention_dp_size()
|
dp_size = get_parallel().attn_dp_size
|
||||||
dp_rank_hint = dp_rank
|
dp_rank_hint = dp_rank
|
||||||
dp_size_hint = dp_size
|
dp_size_hint = dp_size
|
||||||
if cfg.ssd.enabled:
|
if cfg.ssd.enabled:
|
||||||
@@ -883,7 +881,7 @@ class UMBPStore(HiCacheStorage):
|
|||||||
dp_size,
|
dp_size,
|
||||||
cfg.ssd.storage_dir,
|
cfg.ssd.storage_dir,
|
||||||
)
|
)
|
||||||
except (ImportError, AssertionError):
|
except (ImportError, AssertionError, RuntimeError):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -27,7 +27,6 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
attn_tp_all_reduce,
|
attn_tp_all_reduce,
|
||||||
dp_gather_replicate,
|
dp_gather_replicate,
|
||||||
dp_scatter,
|
dp_scatter,
|
||||||
get_attention_dp_size,
|
|
||||||
get_dp_global_num_tokens,
|
get_dp_global_num_tokens,
|
||||||
get_global_dp_buffer,
|
get_global_dp_buffer,
|
||||||
get_local_dp_buffer,
|
get_local_dp_buffer,
|
||||||
@@ -505,7 +504,7 @@ class Qwen4ExpNGramEmbedding(nn.Module):
|
|||||||
self.use_attn_tp_ngram = _use_attn_tp_ngram()
|
self.use_attn_tp_ngram = _use_attn_tp_ngram()
|
||||||
self.gather_dp_tokens = (
|
self.gather_dp_tokens = (
|
||||||
is_dp_attention_enabled()
|
is_dp_attention_enabled()
|
||||||
and get_attention_dp_size() > 1
|
and get_parallel().attn_dp_size > 1
|
||||||
and not self.use_attn_tp_ngram
|
and not self.use_attn_tp_ngram
|
||||||
)
|
)
|
||||||
ngram_prefix = f"{prefix}.ngram_embedding" if prefix else "ngram_embedding"
|
ngram_prefix = f"{prefix}.ngram_embedding" if prefix else "ngram_embedding"
|
||||||
@@ -1360,7 +1359,7 @@ class Qwen4ExpLayerExtensionMixin:
|
|||||||
return hidden_states, residual
|
return hidden_states, residual
|
||||||
|
|
||||||
def _qwen4_exp_use_dp_moe_gather(self) -> bool:
|
def _qwen4_exp_use_dp_moe_gather(self) -> bool:
|
||||||
return get_attention_dp_size() > 1 and get_moe_a2a_backend().is_none()
|
return get_parallel().attn_dp_size > 1 and get_moe_a2a_backend().is_none()
|
||||||
|
|
||||||
def _qwen4_exp_use_attn_tp_a2a_scatter(self) -> bool:
|
def _qwen4_exp_use_attn_tp_a2a_scatter(self) -> bool:
|
||||||
return get_parallel().attn_tp_size > 1 and not get_moe_a2a_backend().is_none()
|
return get_parallel().attn_tp_size > 1 and not get_moe_a2a_backend().is_none()
|
||||||
|
|||||||
@@ -290,22 +290,53 @@ class ParallelContext:
|
|||||||
raise AttributeError(f"ParallelContext has no {name!r}")
|
raise AttributeError(f"ParallelContext has no {name!r}")
|
||||||
|
|
||||||
def _v(self, name, getter):
|
def _v(self, name, getter):
|
||||||
overrides = self._overrides
|
"""Scoped override, else the permanent stamp, else the live group.
|
||||||
return overrides[name] if name in overrides else getter()
|
|
||||||
|
|
||||||
def override_permanently(self, **widths) -> None:
|
One priority order for ranks and widths alike (`_derived_width`), so a
|
||||||
"""Permanently correct a derived width the published bag can't answer
|
stamped value wins over the coordinator for both.
|
||||||
|
"""
|
||||||
|
overrides = self._overrides
|
||||||
|
if name in overrides:
|
||||||
|
return overrides[name]
|
||||||
|
derived = self._derived
|
||||||
|
if name in derived:
|
||||||
|
return derived[name]
|
||||||
|
return getter()
|
||||||
|
|
||||||
|
def _stamped(self, name, why):
|
||||||
|
"""A per-process fact no configuration implies: scoped override, else
|
||||||
|
the permanent stamp, else fail.
|
||||||
|
|
||||||
|
Unlike a width, this has nothing to fall back on -- the configuration
|
||||||
|
does not carry this process's rank, and there is no group to ask --
|
||||||
|
so an unstamped read is a missing initialization rather than a
|
||||||
|
missing override, and says so.
|
||||||
|
"""
|
||||||
|
overrides = self._overrides
|
||||||
|
if name in overrides:
|
||||||
|
return overrides[name]
|
||||||
|
derived = self._derived
|
||||||
|
if name in derived:
|
||||||
|
return derived[name]
|
||||||
|
raise RuntimeError(f"parallel rank {name!r} is not available: {why}")
|
||||||
|
|
||||||
|
def override_permanently(self, **values) -> None:
|
||||||
|
"""Permanently record a width or rank the published bag can't answer
|
||||||
or no longer answers correctly -- not `RuntimeContext.override`,
|
or no longer answers correctly -- not `RuntimeContext.override`,
|
||||||
because a derived width is not a resolved config leaf and this must
|
because neither is a resolved config leaf and this must work with no
|
||||||
work with no config published at all (`multimodal_gen` lends a TP
|
config published at all (`multimodal_gen` lends a TP group to `srt`
|
||||||
group to `srt` layers with no `srt` config to publish against).
|
layers with no `srt` config to publish against).
|
||||||
|
|
||||||
|
Widths are quotients of the configured leaves, so the bag can usually
|
||||||
|
answer and this only corrects it; a rank is a per-process fact the
|
||||||
|
configuration never carries, so for those this is the only source.
|
||||||
|
|
||||||
Lives beside, not inside, the `@contextmanager` `override` above -- a
|
Lives beside, not inside, the `@contextmanager` `override` above -- a
|
||||||
name it cannot also have on this class -- because these are permanent
|
name it cannot also have on this class -- because these are permanent
|
||||||
for the process, not scoped to a `with` block: none of the real
|
for the process, not scoped to a `with` block: none of the real
|
||||||
callers ever restore the value they set here.
|
callers ever restore the value they set here.
|
||||||
"""
|
"""
|
||||||
self._derived.update(widths)
|
self._derived.update(values)
|
||||||
|
|
||||||
def clear_derived_widths(self) -> None:
|
def clear_derived_widths(self) -> None:
|
||||||
self._derived.clear()
|
self._derived.clear()
|
||||||
@@ -401,7 +432,12 @@ class ParallelContext:
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def attn_dp_rank(self) -> int:
|
def attn_dp_rank(self) -> int:
|
||||||
return self._v("attn_dp_rank", _dp().get_attention_dp_rank)
|
return self._stamped(
|
||||||
|
"attn_dp_rank",
|
||||||
|
"it is computed from this process's `tp_rank` when the attention "
|
||||||
|
"topology is initialized, so a process that never ran "
|
||||||
|
"`initialize_dp_attention` has no answer to give",
|
||||||
|
)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def world_group(self) -> Any:
|
def world_group(self) -> Any:
|
||||||
@@ -560,9 +596,9 @@ class MoeFlags(_FlagGroupBase):
|
|||||||
|
|
||||||
class DpFlags(_FlagGroupBase):
|
class DpFlags(_FlagGroupBase):
|
||||||
"""DP-attention runtime flags, materialized by ``initialize_dp_attention``
|
"""DP-attention runtime flags, materialized by ``initialize_dp_attention``
|
||||||
(after distributed setup; reads the model config). Topology values
|
(after distributed setup; reads the model config). The topology values it
|
||||||
(sizes/ranks) stay on ``layers.dp_attention`` until the parallel vertical
|
also computes -- the attention-DP width and rank -- are stamped on
|
||||||
migrates them."""
|
``get_parallel()``, not kept here."""
|
||||||
|
|
||||||
enabled: bool = False
|
enabled: bool = False
|
||||||
use_world_group_for_gather: bool = False
|
use_world_group_for_gather: bool = False
|
||||||
|
|||||||
@@ -65,7 +65,9 @@ _DP = "sglang.srt.layers.dp_attention"
|
|||||||
# anything, so there is nothing to derive them from. The quotients used to be
|
# anything, so there is nothing to derive them from. The quotients used to be
|
||||||
# in this table and are not any more -- `attn_tp_size` and its siblings are
|
# in this table and are not any more -- `attn_tp_size` and its siblings are
|
||||||
# functions of the configured leaves, and `TestDerivedWidthsComeFromTheLeaves`
|
# functions of the configured leaves, and `TestDerivedWidthsComeFromTheLeaves`
|
||||||
# is what pins them.
|
# is what pins them. `attn_dp_rank` is not here either: no group coordinator
|
||||||
|
# knows it, so it is stamped when the attention topology is initialized and
|
||||||
|
# `TestStampedRanks` is what pins it.
|
||||||
SIZE_RANK_DELEGATIONS = [
|
SIZE_RANK_DELEGATIONS = [
|
||||||
("world_size", f"{_PS}.get_world_size"),
|
("world_size", f"{_PS}.get_world_size"),
|
||||||
("world_rank", f"{_PS}.get_world_rank"),
|
("world_rank", f"{_PS}.get_world_rank"),
|
||||||
@@ -77,7 +79,6 @@ SIZE_RANK_DELEGATIONS = [
|
|||||||
("moe_tp_rank", f"{_PS}.get_moe_tensor_parallel_rank"),
|
("moe_tp_rank", f"{_PS}.get_moe_tensor_parallel_rank"),
|
||||||
("attn_tp_rank", f"{_PS}.get_attn_tensor_model_parallel_rank"),
|
("attn_tp_rank", f"{_PS}.get_attn_tensor_model_parallel_rank"),
|
||||||
("attn_cp_rank", f"{_PS}.get_attn_context_model_parallel_rank"),
|
("attn_cp_rank", f"{_PS}.get_attn_context_model_parallel_rank"),
|
||||||
("attn_dp_rank", f"{_DP}.get_attention_dp_rank"),
|
|
||||||
]
|
]
|
||||||
|
|
||||||
GROUP_DELEGATIONS = [
|
GROUP_DELEGATIONS = [
|
||||||
@@ -150,6 +151,87 @@ class TestParallelDelegation(_IsolatedOverrides):
|
|||||||
self.assertFalse(hasattr(ParallelContext, "local_attn_dp_size"))
|
self.assertFalse(hasattr(ParallelContext, "local_attn_dp_size"))
|
||||||
|
|
||||||
|
|
||||||
|
class TestStampedRanks(_IsolatedOverrides):
|
||||||
|
"""`attn_dp_rank` comes from the stamp, and says so when there is none.
|
||||||
|
|
||||||
|
It is the one rank no group answers with: `initialize_dp_attention`
|
||||||
|
computes it from this process's `tp_rank`, and an elastic scale-up
|
||||||
|
replaces it with a rank in the expanded WORLD. Falling back to anything
|
||||||
|
would be inventing a placement for this process.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
super().setUp()
|
||||||
|
parallel = get_parallel()
|
||||||
|
self._saved_derived = dict(parallel._derived)
|
||||||
|
parallel.clear_derived_widths()
|
||||||
|
self.addCleanup(
|
||||||
|
lambda: (
|
||||||
|
parallel.clear_derived_widths(),
|
||||||
|
parallel.override_permanently(**self._saved_derived),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_the_stamp_is_the_answer(self):
|
||||||
|
parallel = get_parallel()
|
||||||
|
parallel.override_permanently(attn_dp_rank=3)
|
||||||
|
self.assertEqual(parallel.attn_dp_rank, 3)
|
||||||
|
# An elastic scale-up restamps it; the newest stamp wins.
|
||||||
|
parallel.override_permanently(attn_dp_rank=9)
|
||||||
|
self.assertEqual(parallel.attn_dp_rank, 9)
|
||||||
|
|
||||||
|
def test_a_scope_still_wins_over_the_stamp(self):
|
||||||
|
parallel = get_parallel()
|
||||||
|
parallel.override_permanently(attn_dp_rank=3)
|
||||||
|
with parallel.override(attn_dp_rank=0):
|
||||||
|
self.assertEqual(parallel.attn_dp_rank, 0)
|
||||||
|
self.assertEqual(parallel.attn_dp_rank, 3)
|
||||||
|
|
||||||
|
def test_unstamped_names_the_cause(self):
|
||||||
|
with self.assertRaises(RuntimeError) as caught:
|
||||||
|
get_parallel().attn_dp_rank
|
||||||
|
self.assertIn("initialize_dp_attention", str(caught.exception))
|
||||||
|
|
||||||
|
def test_a_stated_width_reaches_the_padding_mode(self):
|
||||||
|
"""The reason this PR exists, from a reader's side.
|
||||||
|
|
||||||
|
`get_dp_padding_mode` reads the attention-DP width. Before the width
|
||||||
|
had one home, a scoped `override` moved the context and left the
|
||||||
|
module global answering, so stating a topology moved only half the
|
||||||
|
runtime: this asserted `SUM_LEN` with the width stated as 1.
|
||||||
|
"""
|
||||||
|
from sglang.srt.layers.dp_attention import DpPaddingMode
|
||||||
|
|
||||||
|
with get_parallel().override(attn_dp_size=1):
|
||||||
|
mode = DpPaddingMode.get_dp_padding_mode(
|
||||||
|
is_extend_in_batch=True, global_num_tokens=[3, 5]
|
||||||
|
)
|
||||||
|
self.assertIs(mode, DpPaddingMode.MAX_LEN)
|
||||||
|
|
||||||
|
# And the branch it would have taken with the target's width.
|
||||||
|
with get_parallel().override(attn_dp_size=2):
|
||||||
|
mode = DpPaddingMode.get_dp_padding_mode(
|
||||||
|
is_extend_in_batch=True, global_num_tokens=[3, 5]
|
||||||
|
)
|
||||||
|
self.assertIs(mode, DpPaddingMode.SUM_LEN)
|
||||||
|
|
||||||
|
def test_a_scale_up_stamps_the_width_and_the_rank_together(self):
|
||||||
|
"""The two describe one topology; a reader that saw only one moved
|
||||||
|
would place this process in a group it is not in."""
|
||||||
|
from sglang.srt.layers.dp_attention import update_dp_attention_post_scale
|
||||||
|
|
||||||
|
# It also flips a process-wide gather flag; put it back, or every
|
||||||
|
# later test in this process runs as if a scale-up had happened.
|
||||||
|
dp_flags = get_flags().dp
|
||||||
|
saved_gather = dp_flags.use_world_group_for_gather
|
||||||
|
self.addCleanup(setattr, dp_flags, "use_world_group_for_gather", saved_gather)
|
||||||
|
|
||||||
|
parallel = get_parallel()
|
||||||
|
update_dp_attention_post_scale(new_dp_size=16, new_dp_rank=11)
|
||||||
|
self.assertEqual(parallel.attn_dp_size, 16)
|
||||||
|
self.assertEqual(parallel.attn_dp_rank, 11)
|
||||||
|
|
||||||
|
|
||||||
class TestParallelOverride(_IsolatedOverrides):
|
class TestParallelOverride(_IsolatedOverrides):
|
||||||
def test_override_takes_precedence(self):
|
def test_override_takes_precedence(self):
|
||||||
p = get_parallel()
|
p = get_parallel()
|
||||||
@@ -1573,22 +1655,6 @@ class TestDerivedWidths(_IsolatedOverrides):
|
|||||||
with self.assertRaisesRegex(RuntimeError, r"derived parallel width"):
|
with self.assertRaisesRegex(RuntimeError, r"derived parallel width"):
|
||||||
get_parallel().attn_tp_size
|
get_parallel().attn_tp_size
|
||||||
|
|
||||||
def test_a_temporary_disable_beats_the_permanent_override(self):
|
|
||||||
"""`disable_dp_size()` runs a draft scope without DP attention. It moves
|
|
||||||
the module global the legacy getter reads, so it has to move the derived
|
|
||||||
width too -- the scoped override wins over the permanent one, and a
|
|
||||||
scope that left it alone would answer with the target model's width
|
|
||||||
for its duration."""
|
|
||||||
from sglang.srt.layers import dp_attention
|
|
||||||
|
|
||||||
parallel = get_parallel()
|
|
||||||
parallel.override_permanently(attn_dp_size=4)
|
|
||||||
with patch.object(dp_attention, "_ATTN_DP_SIZE", 4):
|
|
||||||
with dp_attention.disable_dp_size():
|
|
||||||
self.assertEqual(dp_attention.get_attention_dp_size(), 1)
|
|
||||||
self.assertEqual(parallel.attn_dp_size, 1)
|
|
||||||
self.assertEqual(parallel.attn_dp_size, 4)
|
|
||||||
|
|
||||||
def test_the_permanent_override_is_cleared_and_reset(self):
|
def test_the_permanent_override_is_cleared_and_reset(self):
|
||||||
parallel = get_parallel()
|
parallel = get_parallel()
|
||||||
parallel.override_permanently(attn_dp_size=2)
|
parallel.override_permanently(attn_dp_size=2)
|
||||||
|
|||||||
Reference in New Issue
Block a user