[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, ZigzagContextParallelMetadata,
ZigzagCPStrategy, ZigzagCPStrategy,
) )
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner 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). ``(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): if not is_glm_dsa_cache_layer_split_enabled(model_runner):
return None, 1 return None, 1
shard_size = get_attention_cp_size() shard_size = get_parallel().attn_cp_size
if shard_size <= 1: if shard_size <= 1:
return None, 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( 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 layers, plus one extra layer for the remote scratch buffer used when reading
a layer owned by another CP rank. 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): if not is_glm_dsa_cache_layer_split_enabled(model_runner):
return num_layers return num_layers
shard_size = get_attention_cp_size() shard_size = get_parallel().attn_cp_size
if shard_size <= 1: if shard_size <= 1:
return num_layers return num_layers
owned_layers_upper_bound = (num_layers + shard_size - 1) // shard_size 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.attention.dsa import index_buf_accessor
from sglang.srt.layers.cp.utils import get_layer_owner, get_layer_shard_range 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 ( from sglang.srt.mem_cache.memory_pool import (
GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_KV_CACHE,
DSATokenToKVPool, DSATokenToKVPool,
@@ -46,6 +45,7 @@ from sglang.srt.mem_cache.memory_pool import (
maybe_detect_oob, maybe_detect_oob,
unwrap_write_loc, unwrap_write_loc,
) )
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.cache_controller import LayerDoneCounter from sglang.srt.managers.cache_controller import LayerDoneCounter
@@ -113,7 +113,7 @@ class LayerSplitDSATokenToKVPool(DSATokenToKVPool):
# ---- broadcast plumbing ----------------------------------------------- # ---- broadcast plumbing -----------------------------------------------
def _init_layer_broadcast_comm(self) -> None: 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: if cp_group.world_size <= 1 or cp_group.pynccl_comm is None:
return return
@@ -143,7 +143,7 @@ class LayerSplitDSATokenToKVPool(DSATokenToKVPool):
if tensor.data_ptr() != src_tensor.data_ptr(): if tensor.data_ptr() != src_tensor.data_ptr():
tensor.copy_(src_tensor) tensor.copy_(src_tensor)
cp_group = get_attention_cp_group() cp_group = get_parallel().attn_cp_group
comm = ( comm = (
self.layer_broadcast_comm self.layer_broadcast_comm
if use_layer_broadcast_comm and self.layer_broadcast_comm is not None 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 import torch.nn.functional as F
from einops import rearrange 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.attention.vision import VisionAttention
from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.dp_attention import is_dp_attention_enabled
from sglang.srt.layers.layernorm import RMSNorm 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.qwen3_vl import Qwen3_VisionMLP as GlmImageVisionMLP
from sglang.srt.models.utils import compute_cu_seqlens_from_grid_numpy 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.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 from sglang.srt.utils import add_prefix, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -593,7 +592,7 @@ class GlmImageTextAttention(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.layer_id = layer_id self.layer_id = layer_id
self.hidden_size = hidden_size self.hidden_size = hidden_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
@@ -45,10 +45,7 @@ def _run(rank: int, world: int, port: int):
init_distributed_environment, init_distributed_environment,
initialize_model_parallel, initialize_model_parallel,
) )
from sglang.srt.layers.dp_attention import ( from sglang.srt.runtime_context import get_parallel
get_attention_cp_rank,
get_attention_cp_size,
)
init_distributed_environment( init_distributed_environment(
world_size=world, world_size=world,
@@ -66,8 +63,8 @@ def _run(rank: int, world: int, port: int):
LayerSplitDSATokenToKVPool, LayerSplitDSATokenToKVPool,
) )
cp_rank = get_attention_cp_rank() cp_rank = get_parallel().attn_cp_rank
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
assert cp_size == world assert cp_size == world
pool = LayerSplitDSATokenToKVPool( pool = LayerSplitDSATokenToKVPool(