diff --git a/python/sglang/srt/layers/cp/utils.py b/python/sglang/srt/layers/cp/utils.py index 9ac1dc3ad..fbadc0c74 100644 --- a/python/sglang/srt/layers/cp/utils.py +++ b/python/sglang/srt/layers/cp/utils.py @@ -32,6 +32,7 @@ from sglang.srt.layers.cp.zigzag import ( ZigzagContextParallelMetadata, ZigzagCPStrategy, ) +from sglang.srt.runtime_context import get_parallel if TYPE_CHECKING: from sglang.srt.model_executor.model_runner import ModelRunner @@ -66,17 +67,12 @@ def get_glm_dsa_cp_layer_shard_info( ``(None, 1)`` disables sharding (feature off or only one CP rank). """ - from sglang.srt.layers.dp_attention import ( - get_attention_cp_rank, - get_attention_cp_size, - ) - if not is_glm_dsa_cache_layer_split_enabled(model_runner): return None, 1 - shard_size = get_attention_cp_size() + shard_size = get_parallel().attn_cp_size if shard_size <= 1: return None, 1 - return get_attention_cp_rank(), shard_size + return get_parallel().attn_cp_rank, shard_size def get_glm_dsa_layer_split_effective_num_layers( @@ -88,11 +84,9 @@ def get_glm_dsa_layer_split_effective_num_layers( layers, plus one extra layer for the remote scratch buffer used when reading a layer owned by another CP rank. """ - from sglang.srt.layers.dp_attention import get_attention_cp_size - if not is_glm_dsa_cache_layer_split_enabled(model_runner): return num_layers - shard_size = get_attention_cp_size() + shard_size = get_parallel().attn_cp_size if shard_size <= 1: return num_layers owned_layers_upper_bound = (num_layers + shard_size - 1) // shard_size diff --git a/python/sglang/srt/mem_cache/dsa_cache_layer_split.py b/python/sglang/srt/mem_cache/dsa_cache_layer_split.py index 9f44ce0be..9df0b2061 100644 --- a/python/sglang/srt/mem_cache/dsa_cache_layer_split.py +++ b/python/sglang/srt/mem_cache/dsa_cache_layer_split.py @@ -37,7 +37,6 @@ import torch from sglang.srt.layers.attention.dsa import index_buf_accessor from sglang.srt.layers.cp.utils import get_layer_owner, get_layer_shard_range -from sglang.srt.layers.dp_attention import get_attention_cp_group from sglang.srt.mem_cache.memory_pool import ( GPU_MEMORY_TYPE_KV_CACHE, DSATokenToKVPool, @@ -46,6 +45,7 @@ from sglang.srt.mem_cache.memory_pool import ( maybe_detect_oob, unwrap_write_loc, ) +from sglang.srt.runtime_context import get_parallel if TYPE_CHECKING: from sglang.srt.managers.cache_controller import LayerDoneCounter @@ -113,7 +113,7 @@ class LayerSplitDSATokenToKVPool(DSATokenToKVPool): # ---- broadcast plumbing ----------------------------------------------- def _init_layer_broadcast_comm(self) -> None: - cp_group = get_attention_cp_group() + cp_group = get_parallel().attn_cp_group if cp_group.world_size <= 1 or cp_group.pynccl_comm is None: return @@ -143,7 +143,7 @@ class LayerSplitDSATokenToKVPool(DSATokenToKVPool): if tensor.data_ptr() != src_tensor.data_ptr(): tensor.copy_(src_tensor) - cp_group = get_attention_cp_group() + cp_group = get_parallel().attn_cp_group comm = ( self.layer_broadcast_comm if use_layer_broadcast_comm and self.layer_broadcast_comm is not None diff --git a/python/sglang/srt/models/glm_image_vl.py b/python/sglang/srt/models/glm_image_vl.py index 91cd80ad8..c3b0d1812 100644 --- a/python/sglang/srt/models/glm_image_vl.py +++ b/python/sglang/srt/models/glm_image_vl.py @@ -27,7 +27,6 @@ import torch.nn as nn import torch.nn.functional as F from einops import rearrange -from sglang.srt.distributed import get_tensor_model_parallel_world_size from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm @@ -54,7 +53,7 @@ from sglang.srt.models.qwen2 import Qwen2MLP as GlmImageTextMLP from sglang.srt.models.qwen3_vl import Qwen3_VisionMLP as GlmImageVisionMLP from sglang.srt.models.utils import compute_cu_seqlens_from_grid_numpy from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -593,7 +592,7 @@ class GlmImageTextAttention(nn.Module): prefix: str = "", ): super().__init__() - tp_size = get_tensor_model_parallel_world_size() + tp_size = get_parallel().tp_size self.layer_id = layer_id self.hidden_size = hidden_size self.total_num_heads = num_heads diff --git a/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py b/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py index a56e01e4a..0e2c1c7cc 100644 --- a/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py +++ b/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py @@ -45,10 +45,7 @@ def _run(rank: int, world: int, port: int): init_distributed_environment, initialize_model_parallel, ) - from sglang.srt.layers.dp_attention import ( - get_attention_cp_rank, - get_attention_cp_size, - ) + from sglang.srt.runtime_context import get_parallel init_distributed_environment( world_size=world, @@ -66,8 +63,8 @@ def _run(rank: int, world: int, port: int): LayerSplitDSATokenToKVPool, ) - cp_rank = get_attention_cp_rank() - cp_size = get_attention_cp_size() + cp_rank = get_parallel().attn_cp_rank + cp_size = get_parallel().attn_cp_size assert cp_size == world pool = LayerSplitDSATokenToKVPool(