[Bugfix] Migrate retired parallel accessors (#30653)

This commit is contained in:
Mohammad Miadh Angkad
2026-07-09 11:22:34 -07:00
committed by GitHub
parent 26ba3458d3
commit b0ecbceed9
4 changed files with 12 additions and 22 deletions
+4 -10
View File
@@ -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
@@ -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
+2 -3
View File
@@ -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
@@ -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(