[Bugfix] Migrate retired parallel accessors (#30653)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user