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