[refactor] Adopt get_parallel() everywhere and close out the parallel wrapper surface (#30492)

This commit is contained in:
Cheng Wan
2026-07-09 02:09:39 -07:00
committed by GitHub
parent 06eb1b1838
commit e703f9e566
71 changed files with 341 additions and 415 deletions
+2 -2
View File
@@ -70,7 +70,6 @@ from sglang.srt.distributed.parallel_state import (
destroy_model_parallel,
)
from sglang.srt.entrypoints.engine import _set_envs_and_config
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.moe import initialize_moe_config
from sglang.srt.layers.quantization.fp4_utils import initialize_fp4_gemm_config
from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
@@ -80,6 +79,7 @@ from sglang.srt.mem_cache.base_prefix_cache import EvictParams
from sglang.srt.model_executor.cuda_graph_config import Phase
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_parallel
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -503,7 +503,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
prepare_mlp_sync_batch_raw(
batch,
dp_size=model_runner.server_args.dp_size,
attn_tp_size=get_attention_tp_size(),
attn_tp_size=get_parallel().attn_tp_size,
attn_cp_size=model_runner.attn_cp_size,
tp_group=model_runner.tp_group,
get_idle_batch=None,
+2 -2
View File
@@ -20,6 +20,7 @@ from transformers.configuration_utils import PretrainedConfig
from transformers.utils import logging
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.runtime_context import get_parallel
logger = logging.get_logger(__name__)
@@ -173,10 +174,9 @@ class BailingHybridConfig(PretrainedConfig):
@property
def mamba2_cache_params(self) -> Mamba2CacheParams:
from sglang.srt.layers.dp_attention import get_attention_tp_size
shape = Mamba2StateShape.create(
tp_world_size=get_attention_tp_size(),
tp_world_size=get_parallel().attn_tp_size,
intermediate_size=0,
n_groups=0,
num_heads=self.num_linear_key_value_heads,
+2 -2
View File
@@ -22,6 +22,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape,
mamba2_state_dtype,
)
from sglang.srt.runtime_context import get_parallel
logger = logging.get_logger(__name__)
@@ -299,10 +300,9 @@ class FalconH1Config(PretrainedConfig):
@property
def mamba2_cache_params(self):
from sglang.srt.layers.dp_attention import get_attention_tp_size
shape = Mamba2StateShape.create(
tp_world_size=get_attention_tp_size(),
tp_world_size=get_parallel().attn_tp_size,
intermediate_size=self.mamba_intermediate,
n_groups=self.mamba_n_groups,
num_heads=self.mamba_n_heads,
@@ -18,6 +18,7 @@ from transformers.configuration_utils import PretrainedConfig
from transformers.utils import logging
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.runtime_context import get_parallel
logger = logging.get_logger(__name__)
@@ -287,10 +288,9 @@ class GraniteMoeHybridConfig(PretrainedConfig):
@property
def mamba2_cache_params(self):
"""Returns the Mamba2 cache parameters for this configuration."""
from sglang.srt.layers.dp_attention import get_attention_tp_size
shape = Mamba2StateShape.create(
tp_world_size=get_attention_tp_size(),
tp_world_size=get_parallel().attn_tp_size,
intermediate_size=self.mamba_intermediate_size,
n_groups=self.mamba_n_groups,
num_heads=self.mamba_n_heads,
+2 -2
View File
@@ -8,6 +8,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape,
mamba2_state_dtype,
)
from sglang.srt.runtime_context import get_parallel
@dataclass
@@ -56,7 +57,6 @@ class JetNemotronConfig(PretrainedConfig):
@property
def mamba2_cache_params(self) -> Mamba2CacheParams:
from sglang.srt.layers.dp_attention import get_attention_tp_size
jet_block_config = JetBlockConfig(**self.efficient_attention_config["jet"])
@@ -66,7 +66,7 @@ class JetNemotronConfig(PretrainedConfig):
total_v_dim = num_heads * head_v_dim
shape = Mamba2StateShape.create(
tp_world_size=get_attention_tp_size(),
tp_world_size=get_parallel().attn_tp_size,
intermediate_size=total_v_dim,
n_groups=num_heads,
num_heads=num_heads,
+2 -2
View File
@@ -4,6 +4,7 @@
from transformers.configuration_utils import PretrainedConfig
from sglang.srt.configs.mamba_utils import KimiLinearCacheParams, KimiLinearStateShape
from sglang.srt.runtime_context import get_parallel
class KimiLinearConfig(PretrainedConfig):
@@ -151,10 +152,9 @@ class KimiLinearConfig(PretrainedConfig):
@property
def mamba2_cache_params(self) -> KimiLinearCacheParams:
from sglang.srt.layers.dp_attention import get_attention_tp_size
shape = KimiLinearStateShape.create(
tp_world_size=get_attention_tp_size(),
tp_world_size=get_parallel().attn_tp_size,
num_heads=self.linear_attn_config["num_heads"],
head_dim=self.linear_attn_config["head_dim"],
conv_kernel_size=self.linear_attn_config["short_conv_kernel_size"],
+3 -3
View File
@@ -25,6 +25,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape,
mamba2_state_dtype,
)
from sglang.srt.runtime_context import get_parallel
logger = logging.get_logger(__name__)
@@ -62,7 +63,6 @@ class Lfm2Config(HFLfm2Config):
LFM2 uses ShortConv layers with a small fixed-size cache (kernel_size - 1).
Unlike full Mamba2 models, LFM2 only uses the conv state, not SSM temporal state.
"""
from sglang.srt.layers.dp_attention import get_attention_tp_size
conv_layer_ids = self.linear_layer_ids
if not conv_layer_ids:
@@ -71,9 +71,9 @@ class Lfm2Config(HFLfm2Config):
hidden_size = self.hidden_size
conv_kernel = int(self.conv_L_cache)
# get_attention_tp_size() requires initialization, default to 1 if not available
# get_parallel().attn_tp_size requires initialization, default to 1 if not available
try:
tp_size = get_attention_tp_size()
tp_size = get_parallel().attn_tp_size
except (AssertionError, RuntimeError):
tp_size = 1
+2 -2
View File
@@ -23,6 +23,7 @@ from transformers import CONFIG_MAPPING
from transformers.configuration_utils import PretrainedConfig
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.runtime_context import get_parallel
class Lfm2MoeConfig(PretrainedConfig):
@@ -149,7 +150,6 @@ class Lfm2MoeConfig(PretrainedConfig):
LFM2-MoE uses ShortConv layers with a small fixed-size cache.
"""
from sglang.srt.layers.dp_attention import get_attention_tp_size
conv_layer_ids = self.linear_layer_ids
if not conv_layer_ids:
@@ -161,7 +161,7 @@ class Lfm2MoeConfig(PretrainedConfig):
# actual cache size is kernel_size - 1 (e.g., 2 for kernel=3)
try:
tp_size = get_attention_tp_size()
tp_size = get_parallel().attn_tp_size
except (AssertionError, RuntimeError):
tp_size = 1
+3 -3
View File
@@ -19,6 +19,7 @@ from transformers import Lfm2VlConfig as HFLfm2VlConfig
from transformers.utils import logging
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.runtime_context import get_parallel
logger = logging.get_logger(__name__)
@@ -65,7 +66,6 @@ class Lfm2VlConfig(HFLfm2VlConfig):
LFM2 uses ShortConv layers with a small fixed-size cache (kernel_size - 1).
Unlike full Mamba2 models, LFM2 only uses the conv state, not SSM temporal state.
"""
from sglang.srt.layers.dp_attention import get_attention_tp_size
conv_layer_ids = self.linear_layer_ids
if not conv_layer_ids:
@@ -75,9 +75,9 @@ class Lfm2VlConfig(HFLfm2VlConfig):
# conv_L_cache in config is kernel_size (e.g., 3)
conv_kernel = int(self.text_config.conv_L_cache)
# get_attention_tp_size() requires initialization, default to 1 if not available
# get_parallel().attn_tp_size requires initialization, default to 1 if not available
try:
tp_size = get_attention_tp_size()
tp_size = get_parallel().attn_tp_size
except (AssertionError, RuntimeError):
tp_size = 1
+2 -2
View File
@@ -28,6 +28,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape,
mamba2_state_dtype,
)
from sglang.srt.runtime_context import get_parallel
logger = logging.get_logger(__name__)
@@ -422,10 +423,9 @@ class NemotronHConfig(PretrainedConfig):
@property
def mamba2_cache_params(self) -> Mamba2CacheParams:
from sglang.srt.layers.dp_attention import get_attention_tp_size
shape = Mamba2StateShape.create(
tp_world_size=get_attention_tp_size(),
tp_world_size=get_parallel().attn_tp_size,
intermediate_size=self.mamba_num_heads * self.mamba_head_dim,
n_groups=self.n_groups,
num_heads=self.mamba_num_heads,
+3 -3
View File
@@ -25,6 +25,7 @@ from sglang.srt.configs.mamba_utils import (
mamba2_state_dtype,
)
from sglang.srt.configs.update_config import adjust_tp_num_heads_if_necessary
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_cpu
logger = logging.get_logger(__name__)
@@ -285,14 +286,13 @@ class Qwen3NextConfig(PretrainedConfig):
@property
def mamba2_cache_params(self) -> Mamba2CacheParams:
from sglang.srt.layers.dp_attention import get_attention_tp_size
if _is_cpu:
world_size = get_attention_tp_size()
world_size = get_parallel().attn_tp_size
adjust_tp_num_heads_if_necessary(self, world_size, False)
shape = Mamba2StateShape.create(
tp_world_size=get_attention_tp_size(),
tp_world_size=get_parallel().attn_tp_size,
intermediate_size=self.linear_value_head_dim * self.linear_num_value_heads,
n_groups=self.linear_num_key_heads,
num_heads=self.linear_num_value_heads,
+3 -4
View File
@@ -20,6 +20,8 @@ from typing import TYPE_CHECKING, List, Optional
from transformers.configuration_utils import PretrainedConfig
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING:
from sglang.srt.configs.mamba_utils import Mamba2CacheParams
@@ -264,11 +266,8 @@ class ZayaConfig(PretrainedConfig):
# equals the global TP group (DP attention is unsupported), so the two
# are always identical in practice.
try:
from sglang.srt.distributed import (
get_tensor_model_parallel_world_size,
)
tp_size = get_tensor_model_parallel_world_size()
tp_size = get_parallel().tp_size
except (AssertionError, RuntimeError):
tp_size = 1
+18 -16
View File
@@ -1712,30 +1712,32 @@ class _SGLangPlugin(_FrameworkPlugin):
info = {}
from sglang.srt.runtime_context import get_parallel
try:
info["tp_rank"] = self._dist.get_tensor_model_parallel_rank()
info["tp_size"] = self._dist.get_tensor_model_parallel_world_size()
info["pp_rank"] = self._dist.get_pipeline_model_parallel_rank()
info["pp_size"] = self._dist.get_pipeline_model_parallel_world_size()
info["moe_ep_rank"] = self._dist.get_moe_expert_parallel_rank()
info["moe_ep_size"] = self._dist.get_moe_expert_parallel_world_size()
info["moe_tp_rank"] = self._dist.get_moe_tensor_parallel_rank()
info["moe_tp_size"] = self._dist.get_moe_tensor_parallel_world_size()
info["moe_dp_rank"] = self._dist.get_moe_data_parallel_rank()
info["moe_dp_size"] = self._dist.get_moe_data_parallel_world_size()
parallel = get_parallel()
info["tp_rank"] = parallel.tp_rank
info["tp_size"] = parallel.tp_size
info["pp_rank"] = parallel.pp_rank
info["pp_size"] = parallel.pp_size
info["moe_ep_rank"] = parallel.moe_ep_rank
info["moe_ep_size"] = parallel.moe_ep_size
info["moe_tp_rank"] = parallel.moe_tp_rank
info["moe_tp_size"] = parallel.moe_tp_size
info["moe_dp_rank"] = parallel.moe_dp_rank
info["moe_dp_size"] = parallel.moe_dp_size
except (AttributeError, AssertionError):
info["distributed_error"] = True
try:
parallel = get_parallel()
info["enable_dp_attention"] = self._dp_attn.is_dp_attention_enabled()
info["attn_tp_rank"] = self._dp_attn.get_attention_tp_rank()
info["attn_tp_size"] = self._dp_attn.get_attention_tp_size()
info["attn_tp_rank"] = parallel.attn_tp_rank
info["attn_tp_size"] = parallel.attn_tp_size
info["attn_dp_rank"] = self._dp_attn.get_attention_dp_rank()
info["attn_dp_size"] = self._dp_attn.get_attention_dp_size()
info["local_attn_dp_rank"] = self._dp_attn.get_local_attention_dp_rank()
info["local_attn_dp_size"] = self._dp_attn.get_local_attention_dp_size()
info["attn_cp_rank"] = self._dp_attn.get_attention_cp_rank()
info["attn_cp_size"] = self._dp_attn.get_attention_cp_size()
info["attn_cp_rank"] = parallel.attn_cp_rank
info["attn_cp_size"] = parallel.attn_cp_size
except (AttributeError, AssertionError):
info["dp_attention_error"] = True
@@ -33,13 +33,10 @@ from sglang.srt.disaggregation.utils import (
from sglang.srt.distributed import get_pp_group, get_world_group
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import (
get_attention_cp_rank,
get_attention_cp_size,
get_attention_dp_rank,
get_attention_dp_size,
get_attention_tp_rank,
get_attention_tp_size,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import (
NetworkAddress,
@@ -134,10 +131,10 @@ class CommonKVManager(BaseKVManager):
self.bootstrap_host = server_args.host
self.bootstrap_port = server_args.disaggregation_bootstrap_port
self.dist_init_addr = server_args.dist_init_addr
self.attn_tp_size = get_attention_tp_size()
self.attn_tp_rank = get_attention_tp_rank()
self.attn_cp_size = get_attention_cp_size()
self.attn_cp_rank = get_attention_cp_rank()
self.attn_tp_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_parallel().attn_tp_rank
self.attn_cp_size = get_parallel().attn_cp_size
self.attn_cp_rank = get_parallel().attn_cp_rank
self.attn_dp_size = get_attention_dp_size()
self.attn_dp_rank = get_attention_dp_rank()
self.system_dp_size = (
+2 -2
View File
@@ -60,7 +60,6 @@ from sglang.srt.disaggregation.utils import (
setup_state_kv_args,
)
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.managers.schedule_batch import FINISH_ABORT, ScheduleBatch
from sglang.srt.managers.schedule_policy import match_prefix_for_req
from sglang.srt.managers.utils import GenerationBatchResult
@@ -85,6 +84,7 @@ from sglang.srt.observability.req_time_stats import (
set_schedule_time_batch,
set_time_batch,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import get_num_new_pages
from sglang.srt.utils.network import NetworkAddress
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
@@ -400,7 +400,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
kv_args_class = get_kv_class(self.transfer_backend, KVClassType.KVARGS)
kv_args = kv_args_class()
attn_tp_size = get_attention_tp_size()
attn_tp_size = get_parallel().attn_tp_size
kv_args.engine_rank = self.tp_rank % (attn_tp_size)
kv_args.pp_rank = self.pp_rank
@@ -21,7 +21,6 @@ from sglang.srt.hardware_backend.npu.attention.mla_preprocess import (
)
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_kv_cache
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
@@ -39,6 +38,8 @@ import logging
import numpy as np
from sglang.srt.runtime_context import get_parallel
logger = logging.getLogger(__name__)
FULL_ATTENTION_WINDOW = 2147483647
@@ -351,7 +352,8 @@ class AscendAttnBackend(AttentionBackend):
self.q_head_num_padding = None
if hasattr(model_runner.model_config, "num_attention_heads") and self.use_mla:
self.tp_q_head_num = (
model_runner.model_config.num_attention_heads // get_attention_tp_size()
model_runner.model_config.num_attention_heads
// get_parallel().attn_tp_size
)
for num in self.padding_size_list:
if num >= self.tp_q_head_num:
@@ -10,9 +10,9 @@ import torch.nn.functional as F
from sglang.srt.hardware_backend.npu.attention.ascend_backend import AscendAttnBackend
from sglang.srt.layers.attention.dsv4.compressor import CompressorBackendMixin
from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, ForwardMode
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention
@@ -781,7 +781,6 @@ class C4IndexerAscendBackendMixin(C4IndexerBackendMixin):
assert (
not skip_compressor
), "skip_compressor=True is not supported by forward_c4_indexer_npu"
from sglang.srt.layers.dp_attention import get_attention_tp_group
ratio = c4_indexer.compressor.ratio
device = x.device
@@ -825,7 +824,7 @@ class C4IndexerAscendBackendMixin(C4IndexerBackendMixin):
seqlens_cpu = forward_batch.seq_lens_cpu
end_pos = forward_batch.seq_lens.cumsum(dim=0)
page_table = self.forward_metadata.c4_page_table
attn_tp_size = get_attention_tp_size()
attn_tp_size = get_parallel().attn_tp_size
topk_idxs: list[torch.Tensor] = []
for i, _end_token in enumerate(end_pos):
seq_i = int(seqlens_cpu[i])
@@ -847,7 +846,7 @@ class C4IndexerAscendBackendMixin(C4IndexerBackendMixin):
index_score.relu_() * weights.unsqueeze(-1)[start:end, ...]
).sum(dim=1)
if attn_tp_size > 1 and getattr(c4_indexer, "enable_indexer_tp", False):
get_attention_tp_group().all_reduce(index_score)
get_parallel().attn_tp_group.all_reduce(index_score)
arange_kv = torch.arange(seq_i // ratio, device=device)
arange_q = torch.arange(1, seq_i + 1, device=device).unsqueeze(1)
causal = arange_kv.repeat(seq_i, 1) >= (arange_q // ratio)
@@ -972,7 +971,7 @@ class DeepseekV4AscendAttnBackend(
self.use_graph_swa_mask = False
cfg = model_runner.model_config
self._dsv4_config = cfg
tp_size = get_attention_tp_size()
tp_size = get_parallel().attn_tp_size
self._dsv4_q_head_num = cfg.num_attention_heads // tp_size
self._dsv4_kv_head_num = 1 # V4 MQA / latent
self._dsv4_head_dim = cfg.head_dim
@@ -19,7 +19,6 @@ from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp
from sglang.srt.layers.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton,
)
from sglang.srt.layers.dp_attention import get_attention_cp_size
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ReplicatedLinear
from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output
@@ -473,7 +472,7 @@ class Compressor(MultiPlatformOp):
if dsa_use_prefill_cp(forward_batch):
x = cp_all_gather_rerange_output(
x,
get_attention_cp_size(),
get_parallel().attn_cp_size,
forward_batch,
torch.cuda.current_stream(),
)
@@ -21,7 +21,6 @@ from sglang.srt.layers.attention.dsv4.metadata import (
NonPagedIndexerPlan,
PagedIndexerMetadata,
)
from sglang.srt.layers.dp_attention import get_attention_cp_size
from sglang.srt.layers.linear import ReplicatedLinear
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import (
@@ -30,6 +29,7 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
from sglang.srt.utils import add_prefix, is_cuda, is_hip
from sglang.srt.utils.common import is_sm120_supported
@@ -464,7 +464,7 @@ class C4IndexerBackendMixin:
):
return False
if (
get_attention_cp_size() != 1
get_parallel().attn_cp_size != 1
or self.hisparse_coordinator is not None
or is_in_tc_piecewise_cuda_graph()
or is_in_breakable_cuda_graph()
@@ -9,7 +9,6 @@ from sglang.srt.distributed.communication_op import (
from sglang.srt.layers.attention.fla.layernorm_gated import rms_norm_gated
from sglang.srt.layers.dp_attention import (
attn_tp_all_reduce,
get_attention_tp_group,
is_dp_attention_enabled,
)
from sglang.srt.layers.utils import MultiPlatformOp
@@ -92,7 +91,7 @@ class Mixer2RMSNormGated(MultiPlatformOp):
# To handle the general case, redundantly apply the variance
if self.use_attn_tp_group:
parts = [torch.empty_like(x) for _ in range(self.tp_size)]
get_attention_tp_group().all_gather(x, output_tensor_list=parts)
get_parallel().attn_tp_group.all_gather(x, output_tensor_list=parts)
x = torch.cat(parts, dim=-1)
else:
x = tensor_model_parallel_all_gather(x, -1)
+6 -7
View File
@@ -42,7 +42,6 @@ from sglang.srt.layers.dp_attention import (
dp_gather_replicate,
dp_reduce_scatter_tensor,
dp_scatter,
get_attention_tp_group,
get_dp_global_num_tokens,
get_global_dp_buffer,
get_local_dp_buffer,
@@ -923,7 +922,7 @@ class CommunicateSimpleFn:
return tuple(gathered_hidden_states)
hidden_states, local_hidden_states = (
get_local_dp_buffer(get_attention_tp_group()),
get_local_dp_buffer(get_parallel().attn_tp_group),
hidden_states,
)
attn_tp_all_gather_into_tensor(
@@ -1044,7 +1043,7 @@ class CommunicateWithAllReduceAndLayerNormFn:
(``moe_dense_tp_size > 1``): both hidden states and residual stay in
``TP_ATTN_FULL`` across the boundary.
"""
hidden_states = get_attention_tp_group().all_reduce(hidden_states)
hidden_states = get_parallel().attn_tp_group.all_reduce(hidden_states)
if hidden_states.shape[0] != 0:
hidden_states, residual = layernorm(hidden_states, residual)
return hidden_states, residual
@@ -1069,7 +1068,7 @@ class CommunicateWithAllReduceAndLayerNormFn:
if residual_input_mode == ScatterMode.SCATTERED and context.attn_tp_size > 1:
residual, local_residual = (
get_local_dp_buffer(get_attention_tp_group()),
get_local_dp_buffer(get_parallel().attn_tp_group),
residual,
)
attn_tp_all_gather_into_tensor(residual, local_residual)
@@ -1325,7 +1324,7 @@ class CommunicateSummableTensorPairFn:
if get_parallel().tp_size == get_parallel().attn_dp_size:
group = get_tp_group()
else:
group = get_attention_tp_group()
group = get_parallel().attn_tp_group
hidden_states, global_hidden_states = (
get_local_dp_buffer(group),
hidden_states,
@@ -1353,7 +1352,7 @@ class CommunicateSummableTensorPairFn:
hidden_states += residual
residual = None
hidden_states, local_hidden_states = (
get_local_dp_buffer(get_attention_tp_group()),
get_local_dp_buffer(get_parallel().attn_tp_group),
hidden_states,
)
attn_tp_all_gather_into_tensor(
@@ -1415,7 +1414,7 @@ class CommunicateSummableTensorPairFn:
if get_parallel().tp_size == get_parallel().attn_dp_size:
group = get_tp_group()
else:
group = get_attention_tp_group()
group = get_parallel().attn_tp_group
hidden_states_output, global_hidden_states = (
get_local_dp_buffer(group),
hidden_states,
@@ -34,7 +34,6 @@ from sglang.srt.layers.communicator import (
from sglang.srt.layers.dp_attention import (
attn_cp_all_gather_into_tensor,
attn_cp_reduce_scatter_tensor,
get_attention_cp_group,
get_local_dp_buffer,
)
from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp
@@ -54,7 +53,7 @@ def dsa_cp_gather_hidden_states(hidden_states: torch.Tensor):
attn_tp_size = get_parallel().attn_tp_size
assert attn_dp_size == 1 and attn_tp_size == 1
hidden_states, local_hidden_states = (
get_local_dp_buffer(get_attention_cp_group()),
get_local_dp_buffer(get_parallel().attn_cp_group),
hidden_states,
)
attn_cp_all_gather_into_tensor(hidden_states, local_hidden_states)
+2 -2
View File
@@ -48,11 +48,11 @@ from sglang.srt.layers.cp.base import (
CPAttentionBackendKind,
)
from sglang.srt.layers.dp_attention import (
get_attention_cp_group,
is_allocation_symmetric,
)
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
from sglang.srt.runtime_context import get_parallel
@dataclass
@@ -362,7 +362,7 @@ class ZigzagCPStrategy(ContextParallelStrategy):
padding = [0, 0] * (x.ndim - 1) + [0, pad_size]
x = F.pad(x, padding, mode="constant", value=0)
group = get_attention_cp_group()
group = get_parallel().attn_cp_group
ctx = (
use_symmetric_memory(group, disabled=not is_allocation_symmetric())
if x.is_cuda
+7 -8
View File
@@ -31,10 +31,9 @@ from sglang.srt.distributed.parallel_state import (
GroupCoordinator,
get_dcp_group,
get_dcp_group_no_assert,
get_dcp_rank,
get_dcp_world_size,
)
from sglang.srt.layers.dcp.kernels import CPTritonContext, correct_attn_out
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_cuda
@@ -46,19 +45,19 @@ def dcp_enabled() -> bool:
return False
if not is_cuda():
return False
return get_dcp_world_size() > 1
return get_parallel().dcp_size > 1
def get_attention_dcp_world_size() -> int:
if not dcp_enabled():
return 1
return get_dcp_world_size()
return get_parallel().dcp_size
def get_attention_dcp_rank() -> int:
if not dcp_enabled():
return 0
return get_dcp_rank()
return get_parallel().dcp_rank
def _ag_lse(cp_attn_lse: torch.Tensor, cp_group: GroupCoordinator) -> torch.Tensor:
@@ -133,7 +132,7 @@ def cp_lse_ag_out_rs_mla(
def _all_gather_dcp_kv_cache(kv_a: torch.Tensor):
dcp_world_size = get_dcp_world_size()
dcp_world_size = get_parallel().dcp_size
# not use symmetric_memory unless torch mem_pool updated, see https://github.com/pytorch/pytorch/issues/178138
gathered_kv_a = kv_a.new_empty(
(kv_a.shape[0] * dcp_world_size, *kv_a.shape[1:]),
@@ -282,8 +281,8 @@ def all_gather_kv_cache_for_dcp(
if not dcp_enabled():
return torch.cat([prefix_kv_a, prefix_k_pe], dim=-1)
# 1. compute max kv_lens for each seq
dcp_world_size = get_dcp_world_size()
dcp_rank = get_dcp_rank()
dcp_world_size = get_parallel().dcp_size
dcp_rank = get_parallel().dcp_rank
if prefix_starts_cpu is None:
prefix_starts_cpu = torch.zeros_like(prefix_kv_lens_cpu)
+6 -4
View File
@@ -17,8 +17,8 @@ the owner-rule local-index filter."""
import torch
from sglang.srt.distributed.parallel_state import get_dcp_rank, get_dcp_world_size
from sglang.srt.layers.dcp.comm import dcp_enabled
from sglang.srt.runtime_context import get_parallel
def get_dcp_lens(
@@ -45,8 +45,8 @@ def get_dcp_lens(
def filter_dcp_local_kv_indices(kv_indices: torch.Tensor):
if dcp_enabled():
kv_indices = (
kv_indices[kv_indices % get_dcp_world_size() == get_dcp_rank()]
// get_dcp_world_size()
kv_indices[kv_indices % get_parallel().dcp_size == get_parallel().dcp_rank]
// get_parallel().dcp_size
)
return kv_indices
@@ -61,4 +61,6 @@ def update_local_kv_lens_for_dcp(kv_len_arr):
"""
if not dcp_enabled():
return
kv_len_arr.copy_(get_dcp_lens(kv_len_arr, get_dcp_world_size(), get_dcp_rank()))
kv_len_arr.copy_(
get_dcp_lens(kv_len_arr, get_parallel().dcp_size, get_parallel().dcp_rank)
)
+6 -6
View File
@@ -20,7 +20,6 @@ from typing import Optional
import torch
from sglang.srt.distributed.parallel_state import get_dcp_rank, get_dcp_world_size
from sglang.srt.layers.dcp.comm import dcp_enabled
from sglang.srt.layers.dcp.kernels import (
create_dcp_kv_indices,
@@ -28,6 +27,7 @@ from sglang.srt.layers.dcp.kernels import (
)
from sglang.srt.layers.dcp.layout import update_local_kv_lens_for_dcp
from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args
@@ -108,13 +108,13 @@ def prepare_decode_context_parallel_metadata(
extend_cu_prefix_lens,
dcp_kv_indices,
extend_prefix_lens_sum,
get_dcp_world_size(),
get_parallel().dcp_size,
)
dcp_local_prefix_kv_indices = (
dcp_prefix_kv_indices[
dcp_prefix_kv_indices % get_dcp_world_size() == get_dcp_rank()
dcp_prefix_kv_indices % get_parallel().dcp_size == get_parallel().dcp_rank
]
// get_dcp_world_size()
// get_parallel().dcp_size
)
dcp_kv_buffer = torch.empty(
(
@@ -179,8 +179,8 @@ def plan_dcp_decode_metadata(
local_kv_lens,
local_kv_lens_cumsum,
local_kv_indices,
dcp_rank=get_dcp_rank(),
dcp_world_size=get_dcp_world_size(),
dcp_rank=get_parallel().dcp_rank,
dcp_world_size=get_parallel().dcp_size,
BLOCK_SIZE=BLOCK_SIZE,
)
kv_indices[:total_local_len] = local_kv_indices[:total_local_len]
@@ -16,6 +16,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
from sglang.srt.environ import envs
from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ceil_align, ceil_div, get_available_gpu_memory, is_musa
@@ -435,7 +436,6 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
# in-seq-split). _dummy_run does not pad q/hidden like the real flow, so
# an unaligned bs makes DSA's padded num_splits longer than the q tokens
# and trips FlashMLA's "num_splits must have shape (b+1)" check.
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.utils.cp_utils import get_cp_padding_align_size
from sglang.srt.utils.common import require_mlp_sync
@@ -443,7 +443,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
block_m = 64
cp = max(get_cp_padding_align_size(), 1)
attn_tp_size = get_attention_tp_size()
attn_tp_size = get_parallel().attn_tp_size
mlp_sync = require_mlp_sync(model_runner.server_args)
def _align(bs: int) -> int:
+18 -84
View File
@@ -12,8 +12,6 @@ import triton.language as tl
from sglang.srt.distributed import (
GroupCoordinator,
get_attn_context_model_parallel_rank,
get_attn_context_model_parallel_world_size,
get_attn_cp_group,
get_attn_tensor_model_parallel_rank,
get_attn_tensor_model_parallel_world_size,
@@ -43,8 +41,6 @@ if TYPE_CHECKING:
_ATTN_DP_RANK: Optional[int] = None
_ATTN_DP_SIZE: Optional[int] = None
_LOCAL_ATTN_DP_SIZE: Optional[int] = None
_LOCAL_ATTN_DP_RANK: Optional[int] = None
_is_hip = is_hip()
_USE_ROCM700A_WA = _is_hip and get_bool_env_var("SGLANG_USE_ROCM700A")
@@ -275,29 +271,11 @@ def compute_dp_attention_world_info(
return attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size
def compute_dp_attention_local_info(
enable_dp_attention, tp_rank, tp_size, dp_size, moe_dense_tp_size
):
if not enable_dp_attention:
return tp_rank, tp_size, 0
local_tp_size = moe_dense_tp_size if moe_dense_tp_size else tp_size
local_tp_rank = tp_rank % local_tp_size
local_dp_size = max(1, dp_size // (tp_size // local_tp_size))
local_attn_tp_size = local_tp_size // local_dp_size
local_attn_dp_rank = local_tp_rank // local_attn_tp_size
local_attn_tp_rank = local_tp_rank % local_attn_tp_size
return local_attn_tp_rank, local_attn_tp_size, local_attn_dp_rank
def initialize_dp_attention(
server_args: ServerArgs,
model_config: ModelConfig,
):
global _ATTN_DP_RANK, _ATTN_DP_SIZE
global _LOCAL_ATTN_DP_SIZE, _LOCAL_ATTN_DP_RANK
dp = get_flags().dp
dp.max_len_with_idle = (
getattr(model_config.hf_config, "hybrid_override_pattern", None) is not None
@@ -315,19 +293,7 @@ def initialize_dp_attention(
_, _, _ATTN_DP_RANK, _ = compute_dp_attention_world_info(
enable_dp_attention, tp_rank, tp_size, dp_size, attn_cp_size
)
_, _, _LOCAL_ATTN_DP_RANK = compute_dp_attention_local_info(
enable_dp_attention, tp_rank, tp_size, dp_size, moe_dense_tp_size
)
if enable_dp_attention:
_ATTN_DP_SIZE = dp_size
if moe_dense_tp_size is None:
_LOCAL_ATTN_DP_SIZE = _ATTN_DP_SIZE
else:
_LOCAL_ATTN_DP_SIZE = max(1, dp_size // (tp_size // moe_dense_tp_size))
else:
_ATTN_DP_SIZE = 1
_LOCAL_ATTN_DP_SIZE = 1
_ATTN_DP_SIZE = dp_size if enable_dp_attention else 1
_DpGatheredBufferWrapper.set_metadata(
hidden_size=model_config.hidden_size,
@@ -344,30 +310,6 @@ def is_allocation_symmetric() -> bool:
return not is_dp_attention_enabled() or is_dp_max_padding()
def get_attention_tp_group() -> GroupCoordinator:
return get_attn_tp_group()
def get_attention_tp_rank() -> int:
return get_attn_tensor_model_parallel_rank()
def get_attention_tp_size() -> int:
return get_attn_tensor_model_parallel_world_size()
def get_attention_cp_group() -> GroupCoordinator:
return get_attn_cp_group()
def get_attention_cp_rank() -> int:
return get_attn_context_model_parallel_rank()
def get_attention_cp_size() -> int:
return get_attn_context_model_parallel_world_size()
def get_attention_dp_rank() -> int:
assert _ATTN_DP_RANK is not None, "dp attention not initialized!"
return _ATTN_DP_RANK
@@ -378,16 +320,6 @@ def get_attention_dp_size() -> int:
return _ATTN_DP_SIZE
def get_local_attention_dp_rank() -> int:
assert _LOCAL_ATTN_DP_RANK is not None, "dp attention not initialized!"
return _LOCAL_ATTN_DP_RANK
def get_local_attention_dp_size() -> int:
assert _LOCAL_ATTN_DP_SIZE is not None, "dp attention not initialized!"
return _LOCAL_ATTN_DP_SIZE
@contextmanager
def disable_dp_size():
"""Patch the tp group temporarily until this function ends.
@@ -497,7 +429,9 @@ def _dp_gather_via_all_reduce(
assert local_tokens.is_contiguous()
assert global_tokens.is_contiguous()
if local_tokens.shape[0] > 0 and (is_partial or get_attention_tp_rank() == 0):
if local_tokens.shape[0] > 0 and (
is_partial or get_attn_tensor_model_parallel_rank() == 0
):
assert (
local_tokens.untyped_storage() is not global_tokens.untyped_storage()
), "aliasing between global_tokens and local_tokens not allowed"
@@ -526,17 +460,17 @@ def _dp_gather_via_all_gather(
forward_batch: ForwardBatch,
is_partial: bool,
):
if get_attention_tp_size() == 1:
if get_attn_tensor_model_parallel_world_size() == 1:
get_tp_group().all_gather_into_tensor(global_tokens, local_tokens)
return
if not is_partial:
if get_attention_tp_rank() != 0:
if get_attn_tensor_model_parallel_rank() != 0:
local_tokens.fill_(0)
scattered_local_tokens = local_tokens.tensor_split(get_attention_tp_size())[
get_attention_tp_rank()
]
get_attention_tp_group().reduce_scatter_tensor(scattered_local_tokens, local_tokens)
scattered_local_tokens = local_tokens.tensor_split(
get_attn_tensor_model_parallel_world_size()
)[get_attn_tensor_model_parallel_rank()]
get_attn_tp_group().reduce_scatter_tensor(scattered_local_tokens, local_tokens)
get_tp_group().all_gather_into_tensor(global_tokens, scattered_local_tokens)
@@ -560,7 +494,7 @@ def is_dp_gatherv_active() -> bool:
dp_reduce_scatter_tensor) consistent."""
return (
_USE_DP_GATHERV
and get_attention_tp_size() == 1
and get_attn_tensor_model_parallel_world_size() == 1
and get_tensor_model_parallel_world_size() == get_attention_dp_size()
and not _DpGatheredBufferWrapper.is_dp_max_padding()
)
@@ -704,7 +638,7 @@ def dp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor):
get_tensor_model_parallel_world_size()
)[get_tensor_model_parallel_rank()]
get_tp_group().reduce_scatter_tensor(scattered_local_tokens, input)
get_attention_tp_group().all_gather_into_tensor(output, scattered_local_tokens)
get_attn_tp_group().all_gather_into_tensor(output, scattered_local_tokens)
# ---------------------------------------------------------------------------
@@ -810,23 +744,23 @@ def dp_reduce_scatterv_async(
def attn_tp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor):
return get_attention_tp_group().reduce_scatter_tensor(output, input)
return get_attn_tp_group().reduce_scatter_tensor(output, input)
def attn_cp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor):
return get_attention_cp_group().reduce_scatter_tensor(output, input)
return get_attn_cp_group().reduce_scatter_tensor(output, input)
def attn_tp_all_reduce(input: torch.Tensor):
return get_attention_tp_group().all_reduce(input)
return get_attn_tp_group().all_reduce(input)
def attn_tp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor):
return get_attention_tp_group().all_gather_into_tensor(output, input)
return get_attn_tp_group().all_gather_into_tensor(output, input)
def attn_cp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor):
return get_attention_cp_group().all_gather_into_tensor(output, input)
return get_attn_cp_group().all_gather_into_tensor(output, input)
def get_moe_cp_group() -> GroupCoordinator:
@@ -855,4 +789,4 @@ def moe_cp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor):
def attn_tp_all_gather(output_list: List[torch.Tensor], input: torch.Tensor):
return get_attention_tp_group().all_gather(input, output_tensor_list=output_list)
return get_attn_tp_group().all_gather(input, output_tensor_list=output_list)
+2 -3
View File
@@ -25,7 +25,6 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
from sglang.srt.layers.dp_attention import (
get_attention_tp_group,
is_allocation_symmetric,
)
from sglang.srt.layers.parameter import (
@@ -1531,7 +1530,7 @@ class RowParallelLinear(LinearBase):
# bias will not get added more than once in TP>1 case)
bias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias
if self.use_dp_attention_reduce:
symm_ctx = use_symmetric_memory(get_attention_tp_group())
symm_ctx = use_symmetric_memory(get_parallel().attn_tp_group)
else:
symm_ctx = use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()
@@ -1541,7 +1540,7 @@ class RowParallelLinear(LinearBase):
if self.reduce_results and self.tp_size > 1 and not skip_all_reduce:
if self.use_dp_attention_reduce:
output = get_attention_tp_group().all_reduce(output_parallel)
output = get_parallel().attn_tp_group.all_reduce(output_parallel)
else:
quantize_communications = (
(
@@ -5,12 +5,7 @@ import re
import torch
from sglang.srt.distributed import (
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_rank,
get_moe_tensor_parallel_world_size,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_cuda
_is_cuda = is_cuda()
@@ -66,10 +61,10 @@ def _load_gptoss_quark_expert_weights(model, weights, quark_expert_pat):
loaded_params: set[str] = set()
mxfp4_block = 32
moe_tp_rank = get_moe_tensor_parallel_rank()
moe_tp_size = get_moe_tensor_parallel_world_size()
moe_ep_rank = get_moe_expert_parallel_rank()
moe_ep_size = get_moe_expert_parallel_world_size()
moe_tp_rank = get_parallel().moe_tp_rank
moe_tp_size = get_parallel().moe_tp_size
moe_ep_rank = get_parallel().moe_ep_rank
moe_ep_size = get_parallel().moe_ep_size
intermediate_size = model.config.intermediate_size
assert (
+2 -3
View File
@@ -7,13 +7,12 @@ from torch import nn
from sglang.srt.distributed import get_tp_group
from sglang.srt.layers.dp_attention import (
get_attention_tp_group,
is_dp_attention_enabled,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.utils.hash import murmur_hash32
from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs
from sglang.srt.runtime_context import get_server_args
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import TOP_K_ALL
from sglang.srt.server_args import get_global_server_args
@@ -71,7 +70,7 @@ class Sampler(nn.Module):
super().__init__()
self.tp_sync_group = get_tp_group().device_group
if is_dp_attention_enabled():
self.tp_sync_group = get_attention_tp_group().device_group
self.tp_sync_group = get_parallel().attn_tp_group.device_group
self.rl_on_policy_target = get_global_server_args().rl_on_policy_target
# In RL on-policy mode, deterministic inference is automatically enabled.
+5 -6
View File
@@ -10,7 +10,6 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
)
from sglang.srt.layers.dp_attention import (
attn_cp_all_gather_into_tensor,
get_attention_cp_group,
is_allocation_symmetric,
)
from sglang.srt.layers.moe import get_moe_a2a_backend
@@ -227,7 +226,7 @@ def cp_all_gather_reorganized_into_tensor(input_tensor, cp_size, forward_batch,
input_tensor, (0, 0, 0, pad_size), mode="constant", value=0
)
with use_symmetric_memory(
get_attention_cp_group(), disabled=not is_allocation_symmetric()
get_parallel().attn_cp_group, disabled=not is_allocation_symmetric()
):
input_tensor_full = torch.empty(
max_len * cp_size,
@@ -236,7 +235,7 @@ def cp_all_gather_reorganized_into_tensor(input_tensor, cp_size, forward_batch,
dtype=input_tensor.dtype,
)
get_attention_cp_group().cp_all_gather_into_tensor_async(
get_parallel().attn_cp_group.cp_all_gather_into_tensor_async(
input_tensor_full, input_tensor, stream
)
@@ -276,7 +275,7 @@ def cp_all_gather_reorganized_into_tensor_kv_cache(
# Create output tensor with proper shape for all dimensions
with use_symmetric_memory(
get_attention_cp_group(), disabled=not is_allocation_symmetric()
get_parallel().attn_cp_group, disabled=not is_allocation_symmetric()
):
input_tensor_full = torch.empty(
max_len * cp_size,
@@ -285,7 +284,7 @@ def cp_all_gather_reorganized_into_tensor_kv_cache(
dtype=input_tensor.dtype,
)
get_attention_cp_group().cp_all_gather_into_tensor_async(
get_parallel().attn_cp_group.cp_all_gather_into_tensor_async(
input_tensor_full, input_tensor, stream
)
@@ -340,7 +339,7 @@ def cp_all_gather_rerange_output(input_tensor, cp_size, forward_batch, stream):
if is_dsa_prefill_cp_round_robin_split():
with use_symmetric_memory(
get_attention_cp_group(), disabled=not is_allocation_symmetric()
get_parallel().attn_cp_group, disabled=not is_allocation_symmetric()
):
output_tensor = input_tensor.new_empty(
(input_tensor.shape[0] * cp_size, *input_tensor.shape[1:]),
+2 -2
View File
@@ -5,7 +5,6 @@ import torch.nn.functional as F
from torch import nn
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
split_tensor_along_last_dim,
tensor_model_parallel_all_gather,
tensor_model_parallel_all_reduce,
@@ -26,6 +25,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
)
from sglang.srt.lora.backend.base_backend import BaseLoRABackend
from sglang.srt.lora.utils import LoRABatchInfo, get_lm_head_lora_b_shard_size
from sglang.srt.runtime_context import get_parallel
_SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get()
@@ -713,7 +713,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
if self.base_layer.input_is_parallel:
input_parallel = input_
else:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_parallel().tp_rank
splitted_input = split_tensor_along_last_dim(
input_, num_partitions=self.base_layer.tp_size
)
+3 -6
View File
@@ -17,10 +17,6 @@ import torch
from sglang.srt.distributed import (
divide,
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_rank,
get_moe_tensor_parallel_world_size,
get_pp_group,
)
from sglang.srt.environ import envs
@@ -41,6 +37,7 @@ from sglang.srt.lora.utils import (
get_stacked_multiply,
get_target_module_name,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_pin_memory_available
from sglang.srt.utils.hf_transformers_utils import AutoConfig
@@ -95,7 +92,7 @@ def _get_moe_ep_context() -> Tuple[int, int]:
"""Return `(moe_ep_size, moe_ep_rank)`, or `(1, 0)` if the MoE EP group
is not initialized (hermetic tests or pure-TP launches)."""
try:
return get_moe_expert_parallel_world_size(), get_moe_expert_parallel_rank()
return get_parallel().moe_ep_size, get_parallel().moe_ep_rank
except Exception: # pragma: no cover - MoE EP group not initialized
return 1, 0
@@ -107,7 +104,7 @@ def _get_moe_tp_context() -> Tuple[int, int]:
MoE weights are NOT sharded along their inner dim even though attention
weights are."""
try:
return get_moe_tensor_parallel_world_size(), get_moe_tensor_parallel_rank()
return get_parallel().moe_tp_size, get_parallel().moe_tp_rank
except Exception: # pragma: no cover - MoE TP group not initialized
return 1, 0
@@ -10,7 +10,6 @@ preserved and called for batches where two-stream isn't active.
import torch
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
split_tensor_along_last_dim,
tensor_model_parallel_all_gather,
tensor_model_parallel_all_reduce,
@@ -24,6 +23,7 @@ from sglang.srt.lora.trtllm_lora_temp import (
is_two_stream_active,
lora_overlap_alloc_stream,
)
from sglang.srt.runtime_context import get_parallel
def qkv_proj_lora_forward(self, input_: torch.Tensor):
@@ -93,7 +93,7 @@ def row_parallel_lora_forward(
if self.base_layer.input_is_parallel:
input_parallel = input_
else:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_parallel().tp_rank
splitted_input = split_tensor_along_last_dim(
input_, num_partitions=self.base_layer.tp_size
)
+7 -14
View File
@@ -33,19 +33,12 @@ if TYPE_CHECKING:
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.pool_host import HostKVCache
from sglang.srt.distributed import (
get_pipeline_model_parallel_rank,
get_pipeline_model_parallel_world_size,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.layers.dp_attention import (
get_attention_dp_rank,
get_attention_tp_rank,
get_attention_tp_size,
is_dp_attention_enabled,
)
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import get_device_module
logger = logging.getLogger(__name__)
@@ -569,16 +562,16 @@ class HiCacheController:
storage_backend_extra_config = {}
if is_dp_attention_enabled():
self.tp_rank = get_attention_tp_rank()
self.tp_size = get_attention_tp_size()
self.tp_rank = get_parallel().attn_tp_rank
self.tp_size = get_parallel().attn_tp_size
self.dp_rank = get_attention_dp_rank()
else:
self.tp_rank = get_tensor_model_parallel_rank()
self.tp_size = get_tensor_model_parallel_world_size()
self.tp_rank = get_parallel().tp_rank
self.tp_size = get_parallel().tp_size
self.dp_rank = 0
self.pp_rank = get_pipeline_model_parallel_rank()
self.pp_size = get_pipeline_model_parallel_world_size()
self.pp_rank = get_parallel().pp_rank
self.pp_size = get_parallel().pp_size
# Currently, NPUMLATokenToKVPool is the subclass of MLATokenToKVPool.
# DeepSeekV4TokenToKVPool has compressed MLA-style rank-replicated cache
+2 -2
View File
@@ -66,7 +66,6 @@ from sglang.srt.disaggregation.decode_schedule_batch_mixin import (
ScheduleBatchDisaggregationDecodeMixin,
)
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationMode
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank
from sglang.srt.dllm.mixin.req import ReqDllmMixin
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
@@ -107,6 +106,7 @@ from sglang.srt.observability.req_time_stats import (
DPControllerReqTimeStats,
SchedulerReqTimeStats,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import ServerArgs, get_global_server_args
@@ -1577,7 +1577,7 @@ class Req(ReqDllmMixin):
self.has_log_time_stats = True
def set_finish_with_abort(self, error_msg: str):
if get_tensor_model_parallel_rank() == 0:
if get_parallel().tp_rank == 0:
logger.error(f"{error_msg}, {self.rid=}")
self.multimodal_inputs = None
self.grammar = None
+3 -4
View File
@@ -74,8 +74,6 @@ from sglang.srt.layers.attention.mamba.ops import (
)
from sglang.srt.layers.dp_attention import (
compute_dp_attention_world_info,
get_attention_cp_group,
get_attention_tp_group,
)
from sglang.srt.layers.moe import initialize_moe_config
from sglang.srt.layers.quantization.fp4_utils import initialize_fp4_gemm_config
@@ -239,6 +237,7 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.platforms import current_platform
from sglang.srt.plugins import load_plugins
from sglang.srt.runtime_context import get_parallel
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.server_args import PortArgs, ServerArgs, get_global_server_args
from sglang.srt.session.session_controller import SessionController
@@ -903,9 +902,9 @@ class Scheduler(
self.tp_group = get_tp_group()
self.tp_cpu_group = self.tp_group.cpu_group
self.attn_tp_group = get_attention_tp_group()
self.attn_tp_group = get_parallel().attn_tp_group
self.attn_tp_cpu_group = self.attn_tp_group.cpu_group
self.attn_cp_group = get_attention_cp_group()
self.attn_cp_group = get_parallel().attn_cp_group
self.attn_cp_cpu_group = self.attn_cp_group.cpu_group
self.pp_group = get_pp_group()
self.world_group = get_world_group()
@@ -27,7 +27,6 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
import torch
from numpy import float64
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.mem_cache.allocator import (
PagedTokenToKVPoolAllocator,
TokenToKVPoolAllocator,
@@ -59,6 +58,8 @@ if TYPE_CHECKING:
import logging
from sglang.srt.runtime_context import get_parallel
logger = logging.getLogger(__name__)
@@ -412,7 +413,7 @@ class LRUList:
evictable_size == lru_list_evictable_size
), f"{self.mamba=}, total nodes: {total_nodes}, total lru: {total_lru}, evictable size: {evictable_size} != lru list evictable size: {lru_list_evictable_size}"
except Exception as e:
if get_tensor_model_parallel_rank() == 0:
if get_parallel().tp_rank == 0:
msg = f"Mamba Radix tree sanity check failed, ping @yizhang2077: {e}"
logger.error(msg)
tree_cache.pretty_print()
@@ -26,7 +26,6 @@ import psutil
import torch
import tqdm
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.model_executor.forward_batch_info import (
@@ -38,7 +37,7 @@ from sglang.srt.model_executor.forward_batch_info import (
)
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode
from sglang.srt.runtime_context import get_flags
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.utils import (
empty_context,
log_info_on_rank0,
@@ -710,11 +709,11 @@ class CPUGraphRunner:
def capture(self) -> None:
capture_range = (
tqdm.tqdm(list(reversed(self.capture_bs)))
if get_tensor_model_parallel_rank() == 0
if get_parallel().tp_rank == 0
else reversed(self.capture_bs)
)
for bs in capture_range:
if get_tensor_model_parallel_rank() == 0:
if get_parallel().tp_rank == 0:
avail_mem = psutil.virtual_memory().available / (1 << 30)
capture_range.set_description(
f"Capturing batches ({bs=} {avail_mem=:.2f} GB)"
@@ -127,7 +127,6 @@ from sglang.srt.layers.cp.utils import (
get_cp_strategy,
)
from sglang.srt.layers.dp_attention import (
get_attention_tp_group,
initialize_dp_attention,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
@@ -182,7 +181,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
from sglang.srt.model_loader.utils import set_default_torch_dtype
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_flags, get_server_args
from sglang.srt.runtime_context import get_flags, get_parallel, get_server_args
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.server_args import ( # noqa: F401 (re-export)
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
@@ -1293,7 +1292,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
)
self.tp_group = get_tp_group()
self.pp_group = get_pp_group()
self.attention_tp_group = get_attention_tp_group()
self.attention_tp_group = get_parallel().attn_tp_group
# Check memory for tensor parallelism
local_gpu_memory = get_available_gpu_memory(self.device, self.gpu_id)
@@ -17,7 +17,6 @@ from sglang.srt.configs.model_config import (
)
from sglang.srt.distributed.parallel_state import get_world_group
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.mem_cache.allocator import (
PagedTokenToKVPoolAllocator,
TokenToKVPoolAllocator,
@@ -49,6 +48,7 @@ from sglang.srt.mem_cache.memory_pool import (
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.common import (
get_available_gpu_memory,
get_device_memory_capacity,
@@ -487,7 +487,7 @@ class ModelRunnerKVCacheMixin:
bundle = init_unified_mamba_pools(
device=self.device,
kv_cache_dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_attention_tp_size()),
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim,
page_size=self.page_size,
start_layer=self.start_layer,
@@ -544,7 +544,7 @@ class ModelRunnerKVCacheMixin:
enable_memory_saver=self.server_args.enable_memory_saver,
)
head_num = self.model_config.get_num_kv_heads(get_attention_tp_size())
head_num = self.model_config.get_num_kv_heads(get_parallel().attn_tp_size)
head_dim = self.model_config.head_dim
if self.is_hybrid_swa_compress:
# Asymmetric head dims between full and SWA (NPU compress path):
@@ -553,7 +553,7 @@ class ModelRunnerKVCacheMixin:
swa_head_num = max(
1,
self.model_config.hf_text_config.swa_num_key_value_heads
// get_attention_tp_size(),
// get_parallel().attn_tp_size,
)
swa_head_dim = self.model_config.hf_text_config.swa_head_dim
swa_v_head_dim = self.model_config.hf_text_config.swa_v_head_dim
@@ -862,7 +862,7 @@ class ModelRunnerKVCacheMixin:
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
get_parallel().attn_tp_size
),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
@@ -885,7 +885,7 @@ class ModelRunnerKVCacheMixin:
"swa_head_num": max(
1,
self.model_config.hf_text_config.swa_num_key_value_heads
// get_attention_tp_size(),
// get_parallel().attn_tp_size,
),
"swa_head_dim": self.model_config.swa_head_dim,
"swa_v_head_dim": self.model_config.swa_v_head_dim,
@@ -898,7 +898,7 @@ class ModelRunnerKVCacheMixin:
dtype=self.kv_cache_dtype,
post_capture_active=self.post_capture_kv_active,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
get_parallel().attn_tp_size
),
head_dim=self.model_config.head_dim,
swa_attention_layer_ids=self.model_config.swa_attention_layer_ids,
@@ -937,7 +937,7 @@ class ModelRunnerKVCacheMixin:
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
get_parallel().attn_tp_size
),
head_dim=self.model_config.head_dim,
layer_num=self.num_effective_layers,
@@ -1008,7 +1008,7 @@ class ModelRunnerKVCacheMixin:
"swa_head_num": max(
1,
self.model_config.hf_text_config.swa_num_key_value_heads
// get_attention_tp_size(),
// get_parallel().attn_tp_size,
),
"swa_head_dim": self.model_config.swa_head_dim,
"swa_v_head_dim": self.model_config.swa_v_head_dim,
@@ -1020,7 +1020,7 @@ class ModelRunnerKVCacheMixin:
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
get_parallel().attn_tp_size
),
head_dim=self.model_config.head_dim,
swa_attention_layer_ids=self.model_config.swa_attention_layer_ids,
@@ -1047,7 +1047,7 @@ class ModelRunnerKVCacheMixin:
dtype=self.kv_cache_dtype,
index_dtype=self.dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
get_parallel().attn_tp_size
),
head_dim=self.model_config.head_dim,
idx_head_dim=sparse_cfg["sparse_index_dim"],
@@ -1071,7 +1071,7 @@ class ModelRunnerKVCacheMixin:
size=self.max_total_num_tokens,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
get_parallel().attn_tp_size
),
head_dim=self.model_config.head_dim,
# if draft worker, we only need 1 attention layer's kv pool
@@ -1106,7 +1106,7 @@ class ModelRunnerKVCacheMixin:
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
get_parallel().attn_tp_size
),
head_dim=self.model_config.head_dim,
v_head_dim=self.model_config.v_head_dim,
@@ -1131,7 +1131,7 @@ class ModelRunnerKVCacheMixin:
page_size=self.page_size,
dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(
get_attention_tp_size()
get_parallel().attn_tp_size
),
head_dim=self.model_config.head_dim,
v_head_dim=self.model_config.v_head_dim,
@@ -29,10 +29,10 @@ from sglang.srt.configs.model_config import (
is_minimax_sparse,
)
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.mem_cache.common import get_alloc_len_per_decode
from sglang.srt.mem_cache.deepseek_v4_memory_pool import get_compress_state_ring_size
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.common import (
ceil_align,
ceil_div,
@@ -179,7 +179,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
kv_cache_dtype = mr.kv_cache_dtype
kv_size = torch._utils._element_size(kv_cache_dtype)
tp_size = get_attention_tp_size()
tp_size = get_parallel().attn_tp_size
if mr.use_mla_backend:
cell_size = (
@@ -232,7 +232,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
)
num_indexer_kv = num_sparse - num_indexer_k_only
kv_heads = model_config.get_num_kv_heads(get_attention_tp_size())
kv_heads = model_config.get_num_kv_heads(get_parallel().attn_tp_size)
head_dim = model_config.head_dim
indexer_head_dim = sparse_cfg["sparse_index_dim"]
indexer_dtype_size = torch._utils._element_size(mr.dtype)
@@ -292,7 +292,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
model_config = mr.model_config
kv_cache_dtype = mr.kv_cache_dtype
kv_size = torch._utils._element_size(kv_cache_dtype)
tp_size = get_attention_tp_size()
tp_size = get_parallel().attn_tp_size
self._full_layers_num = len(model_config.full_attention_layer_ids)
self._swa_layers_num = len(model_config.swa_attention_layer_ids)
@@ -37,7 +37,6 @@ from torch.profiler import ProfilerActivity, profile
from sglang.srt.compilation import torch_compile_decoration
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.distributed.parallel_state import (
graph_capture,
set_pdmux_status,
@@ -46,8 +45,6 @@ from sglang.srt.dllm.config import DllmConfig
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
get_attention_tp_rank,
get_attention_tp_size,
set_dp_buffer_len,
set_is_extend_in_batch,
)
@@ -93,7 +90,7 @@ from sglang.srt.model_executor.runner_utils.deepep_adapter import (
DeepEPCudaGraphRunnerAdapter,
)
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
from sglang.srt.runtime_context import get_flags
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.utils import (
empty_context,
get_available_gpu_memory,
@@ -209,8 +206,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
)
self.enable_pdmux = model_runner.server_args.enable_pdmux
self.attn_tp_size = get_attention_tp_size()
self.attn_tp_rank = get_attention_tp_rank()
self.attn_tp_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_parallel().attn_tp_rank
# True if a DSACPLayerCommunicator-style prefill-CP flavor is active
# (DSA or MLA). These flavors feed a zigzag-split rank-local layout
# into the runner; MHA-arch prefill CP (Qwen3/Qwen2 MoE via PR
@@ -507,7 +504,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
export_cuda_graph_capture_trace(
prof_context,
runner_name=type(self).__name__,
tp_rank=get_tensor_model_parallel_rank(),
tp_rank=get_parallel().tp_rank,
)
def capture_prepare(
@@ -722,7 +719,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# Reverse so cuda graphs share memory better.
capture_range = (
tqdm.tqdm(list(reversed(self.capture_bs)))
if get_tensor_model_parallel_rank() == 0
if get_parallel().tp_rank == 0
else reversed(self.capture_bs)
)
lora_variants = (
@@ -731,7 +728,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
else [(None, None)]
)
for bs in capture_range:
if get_tensor_model_parallel_rank() == 0:
if get_parallel().tp_rank == 0:
avail_mem = get_available_gpu_memory(
self.model_runner.device,
self.model_runner.gpu_id,
@@ -43,7 +43,6 @@ from typing import TYPE_CHECKING, Dict, Optional, Union
import torch
import tqdm
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.distributed.parallel_state import graph_capture
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
@@ -91,6 +90,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
from sglang.srt.model_executor.runner_utils.buffers import (
PrefillInputBuffers,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
from sglang.srt.utils import (
get_available_gpu_memory,
@@ -761,11 +761,11 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
)
capture_range = (
tqdm.tqdm(list(reversed(self.capture_num_tokens)))
if get_tensor_model_parallel_rank() == 0
if get_parallel().tp_rank == 0
else reversed(self.capture_num_tokens)
)
for num_tokens in capture_range:
if get_tensor_model_parallel_rank() == 0:
if get_parallel().tp_rank == 0:
avail_mem = get_available_gpu_memory(
self.model_runner.device,
self.model_runner.gpu_id,
@@ -34,7 +34,6 @@ from sglang.srt.compilation.compile_phase import (
enable_torch_compile_warmup,
set_pcg_capture_stream,
)
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
set_graph_pool_id,
)
@@ -49,6 +48,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
from sglang.srt.model_executor.runner_utils.pool import (
get_or_create_global_graph_memory_pool,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_hip
if TYPE_CHECKING:
@@ -191,11 +191,11 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
tqdm.tqdm(
list(reversed(cuda_graph_runner.capture_num_tokens))
)
if get_tensor_model_parallel_rank() == 0
if get_parallel().tp_rank == 0
else reversed(cuda_graph_runner.capture_num_tokens)
)
for num_tokens in compile_range:
if get_tensor_model_parallel_rank() == 0:
if get_parallel().tp_rank == 0:
compile_range.set_description(
f"Compiling num tokens ({num_tokens=})"
)
+13 -21
View File
@@ -72,8 +72,6 @@ from sglang.srt.connector import (
)
from sglang.srt.connector.utils import parse_model_name
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
model_parallel_is_initialized,
)
from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES
@@ -111,6 +109,7 @@ from sglang.srt.model_loader.weight_utils import (
set_runai_streamer_env,
)
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import (
get_bool_env_var,
get_device_capability,
@@ -527,9 +526,9 @@ class DefaultModelLoader(BaseModelLoader):
if k >= 0:
hf_weights_files.sort()
if k > 0:
tp_size = get_tensor_model_parallel_world_size()
tp_size = get_parallel().tp_size
if tp_size > 1:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_parallel().tp_rank
group_size = tp_size * k
staggered: List[str] = []
for i in range(0, len(hf_weights_files), group_size):
@@ -1120,8 +1119,8 @@ class QuantizedRLModelLoader(DefaultModelLoader):
if scale_info is None:
return
# Get tp rank and size
tp_rank = get_tensor_model_parallel_rank()
tp_size = get_tensor_model_parallel_world_size()
tp_rank = get_parallel().tp_rank
tp_size = get_parallel().tp_size
def _get_tp_sharded_scale(full_scale_tensor):
"""Get tp sharded scale from full scale tensor"""
@@ -1539,8 +1538,6 @@ class ShardedStateLoader(BaseModelLoader):
) -> nn.Module:
from safetensors.torch import safe_open
from sglang.srt.distributed import get_tensor_model_parallel_rank
local_model_path = self._prepare_weights(
model_config.model_path, model_config.revision
)
@@ -1554,7 +1551,7 @@ class ShardedStateLoader(BaseModelLoader):
quant_method = getattr(module, "quant_method", None)
if quant_method is not None:
quant_method.process_weights_after_loading(module)
rank = get_tensor_model_parallel_rank()
rank = get_parallel().tp_rank
pattern = os.path.join(
local_model_path,
self.pattern.format(rank=rank, part="*"),
@@ -1605,11 +1602,9 @@ class ShardedStateLoader(BaseModelLoader):
) -> None:
from safetensors.torch import save_file
from sglang.srt.distributed import get_tensor_model_parallel_rank
if pattern is None:
pattern = ShardedStateLoader.DEFAULT_PATTERN
rank = get_tensor_model_parallel_rank()
rank = get_parallel().tp_rank
part_idx = 0
total_size = 0
state_dict = ShardedStateLoader._filter_subtensors(model.state_dict())
@@ -1908,8 +1903,8 @@ class BitsAndBytesModelLoader(BaseModelLoader):
) -> Generator:
from bitsandbytes.functional import quantize_4bit
tp_size = get_tensor_model_parallel_world_size()
tp_rank = get_tensor_model_parallel_rank()
tp_size = get_parallel().tp_size
tp_rank = get_parallel().tp_rank
for weight_name, weight_tensor in self._hf_weight_iter(
hf_weights_files, use_safetensors
@@ -2003,7 +1998,7 @@ class BitsAndBytesModelLoader(BaseModelLoader):
# The quant_states in pre_quantized models cannot work with a split
# weight tensor. So TP does not work with pre_quantized bnb models.
if pre_quant and get_tensor_model_parallel_world_size() > 1:
if pre_quant and get_parallel().tp_size > 1:
raise ValueError(
"Prequant BitsAndBytes models with TP is not supported."
"Please try with PP."
@@ -2445,7 +2440,7 @@ class RemoteModelLoader(BaseModelLoader):
) -> Generator[Tuple[str, torch.Tensor], None, None]:
"""Get an iterator for the model weights from remote storage."""
assert get_connector_type(client) == ConnectorType.KV
rank = get_tensor_model_parallel_rank()
rank = get_parallel().tp_rank
return client.weight_iterator(rank)
def _get_weights_iterator_fs(
@@ -2468,7 +2463,7 @@ class RemoteModelLoader(BaseModelLoader):
with create_remote_connector(url) as client:
assert get_connector_type(client) == ConnectorType.KV
model_name = parse_model_name(url)
rank = get_tensor_model_parallel_rank()
rank = get_parallel().tp_rank
state_dict = ShardedStateLoader._filter_subtensors(model.state_dict())
for key, tensor in state_dict.items():
r_key = f"{model_name}/keys/rank_{rank}/{key}"
@@ -2804,10 +2799,7 @@ class ModelOptModelLoader(DefaultModelLoader):
# Apply quantization
mtq.quantize(model, quant_cfg, forward_loop=calibrate_loop)
if (
not model_parallel_is_initialized()
or get_tensor_model_parallel_rank() == 0
):
if not model_parallel_is_initialized() or get_parallel().tp_rank == 0:
mtq.print_quant_summary(model)
# Save checkpoint if path provided
+2 -4
View File
@@ -30,7 +30,6 @@ from sglang.srt.compilation.compilation_config import register_split_op
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
from sglang.srt.distributed import (
get_pp_group,
get_tensor_model_parallel_world_size,
get_tp_group,
)
from sglang.srt.environ import envs
@@ -469,8 +468,7 @@ class MQALayer(nn.Module):
self.hidden_size,
bias=False,
quant_config=quant_config,
reduce_results=attn_tp_size == get_tensor_model_parallel_world_size()
and attn_tp_size > 1,
reduce_results=attn_tp_size == get_parallel().tp_size and attn_tp_size > 1,
prefix=add_prefix("wo_b", prefix),
tp_rank=attn_tp_rank,
tp_size=attn_tp_size,
@@ -1104,7 +1102,7 @@ class MQALayer(nn.Module):
o = torch.einsum("tgd,grd->tgr", o, wo_a)
o, _ = self.wo_b(o.flatten(1))
if self.tp_size > 1 and self.tp_size < get_tensor_model_parallel_world_size():
if self.tp_size > 1 and self.tp_size < get_parallel().tp_size:
o = attn_tp_all_reduce(o)
return o
-4
View File
@@ -7,9 +7,6 @@ from transformers import Exaone4Config
from sglang.srt.distributed import get_pp_group
from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.dp_attention import (
get_local_attention_dp_size,
)
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import (
MergedColumnParallelLinear,
@@ -238,7 +235,6 @@ class Exaone4DecoderLayer(nn.Module):
max_position_embeddings = getattr(config, "max_position_embeddings", 8192)
self.local_dp_size = get_local_attention_dp_size()
self.attn_tp_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_parallel().attn_tp_rank
+2 -2
View File
@@ -40,7 +40,6 @@ import torch
from torch import nn
from transformers import PretrainedConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import (
@@ -57,6 +56,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__)
@@ -144,7 +144,7 @@ class HrmTextAttention(nn.Module):
) -> None:
super().__init__()
self.hidden_size = config.hidden_size
tp_size = get_tensor_model_parallel_world_size()
tp_size = get_parallel().tp_size
self.total_num_heads = config.num_attention_heads
assert self.total_num_heads % tp_size == 0, (
+2 -3
View File
@@ -45,7 +45,6 @@ from sglang.srt.layers.communicator import (
)
from sglang.srt.layers.dp_attention import (
attn_tp_all_reduce,
get_attention_tp_group,
is_dp_attention_enabled,
)
from sglang.srt.layers.layernorm import RMSNorm
@@ -437,7 +436,7 @@ class MiniMaxM2QKRMSNorm:
# typically, this should not exceed 1M, since max_tokens is usually less than 16384
max_size = ((8 * max_tokens + ALIGN - 1) // ALIGN) * ALIGN
comm = CustomAllReduceV2(
group=get_attention_tp_group().cpu_group,
group=get_parallel().attn_tp_group.cpu_group,
device=device,
max_pull_size=0,
max_pull_blocks=0,
@@ -877,7 +876,7 @@ class MiniMaxM2Attention(nn.Module):
rotary_dim=self.rotary_dim,
eps=self.q_norm.variance_epsilon,
tp_world=self.q_norm.attn_tp_size,
tp_group=get_attention_tp_group().device_group,
tp_group=get_parallel().attn_tp_group.device_group,
)
else:
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
+2 -3
View File
@@ -21,7 +21,6 @@ from sglang.srt.configs import NemotronHConfig
from sglang.srt.distributed import get_pp_group
from sglang.srt.layers.dp_attention import (
attn_tp_all_reduce,
get_attention_tp_group,
is_dp_attention_enabled,
)
from sglang.srt.layers.layernorm import RMSNorm
@@ -104,7 +103,7 @@ class NemotronHMTPAttentionDecoderLayer(NemotronHAttentionDecoderLayer):
)
hidden_states, _ = self.eh_proj(fused)
if is_dp_attention_enabled():
hidden_states = get_attention_tp_group().all_gather(
hidden_states = get_parallel().attn_tp_group.all_gather(
hidden_states, dim=-1
)
@@ -190,7 +189,7 @@ class NemotronHMTPMoEDecoderLayer(NemotronHMoEDecoderLayer):
)
hidden_states, _ = self.eh_proj(fused)
if is_dp_attention_enabled():
hidden_states = get_attention_tp_group().all_gather(
hidden_states = get_parallel().attn_tp_group.all_gather(
hidden_states, dim=-1
)
+2 -4
View File
@@ -29,8 +29,6 @@ from transformers import PretrainedConfig
from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
from sglang.srt.distributed import (
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_world_size,
get_pp_group,
get_pp_indices,
moe_expert_parallel_all_reduce,
@@ -955,9 +953,9 @@ class Qwen2MoeModel(nn.Module):
and hasattr(hidden_states, "_sglang_needs_allreduce_fusion")
and hidden_states._sglang_needs_allreduce_fusion
):
if get_moe_expert_parallel_world_size() > 1:
if get_parallel().moe_ep_size > 1:
hidden_states = moe_expert_parallel_all_reduce(hidden_states)
if get_moe_tensor_parallel_world_size() > 1:
if get_parallel().moe_tp_size > 1:
hidden_states = moe_tensor_model_parallel_all_reduce(hidden_states)
hidden_states._sglang_needs_allreduce_fusion = False
return PPProxyTensors(
+6 -15
View File
@@ -40,11 +40,8 @@ import pybase64
import torch
from PIL import Image
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.distributed.communication_op import tensor_model_parallel_all_gather
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import flatten_nested_list
@@ -452,12 +449,12 @@ def run_dp_sharded_vision_model(
"""
num_chunks = image_input.shape[0]
mp_world_size = get_tensor_model_parallel_world_size()
mp_world_size = get_parallel().tp_size
num_chunks_per_rank = (num_chunks + mp_world_size - 1) // mp_world_size
num_padded_chunks = num_chunks_per_rank * mp_world_size - num_chunks
pad = (0,) * (2 * (image_input.dim() - 1)) + (0, num_padded_chunks)
image_input_padded = torch.nn.functional.pad(image_input, pad)
rank = get_tensor_model_parallel_rank()
rank = get_parallel().tp_rank
image_input_per_rank = image_input_padded[
rank * num_chunks_per_rank : (rank + 1) * num_chunks_per_rank, ...
]
@@ -504,19 +501,13 @@ def run_dp_sharded_mrope_vision_model(
```
"""
from sglang.srt.layers.dp_attention import (
get_attention_tp_group,
get_attention_tp_rank,
get_attention_tp_size,
)
tp_size = get_attention_tp_size()
tp_size = get_parallel().attn_tp_size
if tp_size == 1:
return vision_model(pixel_values, grid_thw=torch.tensor(grid_thw_list))
# GPU_0 tp_rank_local = 0
# GPU_1 tp_rank_local = 1
tp_rank_local = get_attention_tp_rank()
tp_rank_local = get_parallel().attn_tp_rank
# patches_per_image = [1000, 100, 200, 50]
patches_per_image = [math.prod(grid_thw) for grid_thw in grid_thw_list]
@@ -628,7 +619,7 @@ def run_dp_sharded_mrope_vision_model(
image_embeds_local_padded = image_embeds_local
# Do all_gather to collect embeddings from all ranks
gathered_embeds = get_attention_tp_group().all_gather(
gathered_embeds = get_parallel().attn_tp_group.all_gather(
image_embeds_local_padded, dim=0
)
+15
View File
@@ -81,6 +81,8 @@ _PARALLEL_FIELDS = frozenset(
"attn_cp_rank",
"attn_dp_size",
"attn_dp_rank",
"dcp_size",
"dcp_rank",
"world_group",
"tp_group",
"pp_group",
@@ -89,6 +91,7 @@ _PARALLEL_FIELDS = frozenset(
"moe_tp_group",
"attn_tp_group",
"attn_cp_group",
"dcp_group",
}
)
@@ -183,6 +186,14 @@ class ParallelContext:
def attn_cp_rank(self) -> int:
return self._v("attn_cp_rank", _ps().get_attn_context_model_parallel_rank)
@property
def dcp_size(self) -> int:
return self._v("dcp_size", _ps().get_dcp_world_size)
@property
def dcp_rank(self) -> int:
return self._v("dcp_rank", _ps().get_dcp_rank)
@property
def attn_dp_size(self) -> int:
return self._v("attn_dp_size", _dp().get_attention_dp_size)
@@ -223,6 +234,10 @@ class ParallelContext:
def attn_cp_group(self) -> Any:
return self._v("attn_cp_group", _ps().get_attn_cp_group)
@property
def dcp_group(self) -> Any:
return self._v("dcp_group", _ps().get_dcp_group)
class _FlagGroupBase:
"""Shared flag-group behavior: typo-safe writes + transactional ``override()``.
+4 -2
View File
@@ -20,6 +20,7 @@ from sglang.srt.mem_cache.common import (
get_alloc_reserve_per_decode,
get_last_loc,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.speculative.triton_ops.spec_tree import (
sgl_build_tree_kernel_efficient_triton,
verify_tree_greedy_kernel_triton,
@@ -591,7 +592,6 @@ def eagle_sample(
from sglang.srt.distributed import get_tp_group
from sglang.srt.layers.dp_attention import (
get_attention_tp_group,
is_dp_attention_enabled,
)
from sglang.srt.sampling.penaltylib.repetition_penalty import (
@@ -762,7 +762,9 @@ def eagle_sample(
# non-determinism in softmax/top_k/top_p, causing different
# sampled tokens. Broadcast from rank 0 to ensure consistency.
tp_group = (
get_attention_tp_group() if is_dp_attention_enabled() else get_tp_group()
get_parallel().attn_tp_group
if is_dp_attention_enabled()
else get_tp_group()
)
if tp_group.world_size > 1:
tp_group.broadcast(predict, src=0)
@@ -22,7 +22,6 @@ from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend
from sglang.srt.layers.attention.trtllm_mla_backend import (
TRTLLMMLABackend,
)
from sglang.srt.layers.dp_attention import get_attention_tp_group
from sglang.srt.layers.moe.utils import (
speculative_moe_a2a_backend_context,
speculative_moe_backend_context,
@@ -47,6 +46,7 @@ from sglang.srt.model_executor.runner import (
DecodeCudaGraphRunner,
get_batch_sizes_to_capture,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.adaptive_runtime_state import (
AdaptiveController,
@@ -175,7 +175,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
# Load draft model weights only.
if server_args.enable_dp_attention and self.speculative_algorithm.is_eagle3():
ctx = draft_tp_context(get_attention_tp_group())
ctx = draft_tp_context(get_parallel().attn_tp_group)
else:
ctx = empty_context()
with (
@@ -5,7 +5,7 @@ import numpy as np
import pybase64
import torch
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.runtime_context import get_parallel
from sglang.srt.state_capturer.base import BaseTopkCapturer
logger = logging.getLogger(__name__)
@@ -25,7 +25,7 @@ class IndexerTopkCapturer(BaseTopkCapturer):
self.num_indexer_layers = num_indexer_layers
self.index_topk = index_topk
attn_tp_size = get_attention_tp_size()
attn_tp_size = get_parallel().attn_tp_size
assert attn_tp_size == 1, "IndexerTopkCapturer now only supports DP attention"
# DP-attention capture is per-rank-local: each rank writes [:local_batch, ...]
@@ -7,12 +7,12 @@ import torch
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.layers.dp_attention import (
attn_tp_all_gather_into_tensor,
get_attention_tp_size,
get_dp_local_slice_cpu,
is_dp_attention_enabled,
)
from sglang.srt.layers.moe import get_moe_a2a_backend
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args
from sglang.srt.state_capturer.base import BaseTopkCapturer
@@ -84,7 +84,9 @@ class RoutedExpertsCapturer(BaseTopkCapturer):
# holds the full batch and the existing _get_local_slice / D2H sync
# paths work unchanged. Pre-allocate the gather target.
if get_moe_a2a_backend().is_deepep():
attn_tp_size = get_attention_tp_size() if is_dp_attention_enabled() else 1
attn_tp_size = (
get_parallel().attn_tp_size if is_dp_attention_enabled() else 1
)
self.gather_buffer = torch.empty(
(
self.device_cache.buffer.shape[0] * attn_tp_size,
@@ -98,7 +100,7 @@ class RoutedExpertsCapturer(BaseTopkCapturer):
if get_moe_a2a_backend().is_deepep():
local_topk = topk_indices
topk_indices = self.gather_buffer[
: local_topk.size(0) * get_attention_tp_size()
: local_topk.size(0) * get_parallel().attn_tp_size
]
attn_tp_all_gather_into_tensor(topk_indices, local_topk)
super().capture(layer_id, topk_indices)
+4 -6
View File
@@ -97,6 +97,7 @@ from typing_extensions import Literal
from sglang.srt.environ import envs
from sglang.srt.observability.func_timer import enable_func_timer
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.video_decoder import _BACKEND, VideoDecoderWrapper
if TYPE_CHECKING:
@@ -2058,11 +2059,10 @@ def set_ulimit(target_soft_limit=65535):
def rank0_log(msg: str):
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
model_parallel_is_initialized,
)
if not model_parallel_is_initialized() or get_tensor_model_parallel_rank() == 0:
if not model_parallel_is_initialized() or get_parallel().tp_rank == 0:
logger.info(msg)
@@ -3364,10 +3364,9 @@ class BumpAllocator:
def log_info_on_rank0(logger, msg):
from sglang.srt.distributed import get_tensor_model_parallel_rank
try:
if torch.distributed.is_initialized() and get_tensor_model_parallel_rank() == 0:
if torch.distributed.is_initialized() and get_parallel().tp_rank == 0:
logger.info(msg)
except Exception as e:
if torch.distributed.is_initialized():
@@ -3382,10 +3381,9 @@ def log_debug_on_rank0(logger, msg):
Log a debug message only on tensor model parallel rank 0.
Falls back to logging if distributed is not initialized or error occurs.
"""
from sglang.srt.distributed import get_tensor_model_parallel_rank
try:
if torch.distributed.is_initialized() and get_tensor_model_parallel_rank() == 0:
if torch.distributed.is_initialized() and get_parallel().tp_rank == 0:
logger.debug(msg)
except Exception as e:
if torch.distributed.is_initialized():
+4 -11
View File
@@ -12,7 +12,7 @@ from sglang.srt.distributed.naive_distributed import (
set_naive_distributed,
)
from sglang.srt.layers.parameter import ModelWeightParameter
from sglang.srt.runtime_context import get_stream
from sglang.srt.runtime_context import get_parallel, get_stream
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available
from sglang.srt.utils.host_shared_memory import (
@@ -170,11 +170,8 @@ class OffloaderV2(BaseOffloader):
# Temporarily init inside Offloader, can move if other modules also need this
if self.mode in {"sharded_gpu", "shm_cpu"}:
from sglang.srt.distributed import get_tensor_model_parallel_world_size
assert (
get_tensor_model_parallel_world_size() == 1
), "not yet support tp_size!=1"
assert get_parallel().tp_size == 1, "not yet support tp_size!=1"
set_naive_distributed(
NaiveDistributed(
rank=dp_rank,
@@ -390,9 +387,7 @@ class _ShmCpuParamOffloader(_BaseParamOffloader):
self._rank = get_naive_distributed().get_rank()
self._world_size = get_naive_distributed().get_world_size()
from sglang.srt.distributed import get_tensor_model_parallel_world_size
assert get_tensor_model_parallel_world_size() == 1, "not yet support tp_size!=1"
assert get_parallel().tp_size == 1, "not yet support tp_size!=1"
assert (
self._param.data.is_contiguous()
), f"not yet support non-contiguous tensor {self._param.shape=} {self._param.stride()=}"
@@ -497,9 +492,7 @@ class _ShardedGpuParamOffloader(_BaseParamOffloader):
self._rank = get_naive_distributed().get_rank()
self._world_size = get_naive_distributed().get_world_size()
from sglang.srt.distributed import get_tensor_model_parallel_world_size
assert get_tensor_model_parallel_world_size() == 1, "not yet support tp_size!=1"
assert get_parallel().tp_size == 1, "not yet support tp_size!=1"
assert (
self._param.data.is_contiguous()
), f"not yet support non-contiguous tensor {self._param.shape=} {self._param.stride()=}"
@@ -555,15 +555,11 @@ def _capture_eagle_draft_extend_graph_runner(
"sglang.srt.model_executor.runner.decode_cuda_graph_runner.graph_capture",
_single_rank_graph_capture,
),
patch(
"sglang.srt.model_executor.runner.decode_cuda_graph_runner.get_tensor_model_parallel_rank",
lambda: 0,
),
patch(
"sglang.srt.model_executor.runner.decode_cuda_graph_runner.get_available_gpu_memory",
lambda *args, **kwargs: 0.0,
),
get_parallel().override(attn_cp_size=1),
get_parallel().override(attn_cp_size=1, tp_rank=0),
):
_reset_cuda_graph_test_buffers()
return EAGLEDraftExtendCudaGraphRunner(
@@ -437,15 +437,11 @@ def _capture_eagle_draft_graph_runner(
"sglang.srt.model_executor.runner.decode_cuda_graph_runner.graph_capture",
_single_rank_graph_capture,
),
patch(
"sglang.srt.model_executor.runner.decode_cuda_graph_runner.get_tensor_model_parallel_rank",
lambda: 0,
),
patch(
"sglang.srt.model_executor.runner.decode_cuda_graph_runner.get_available_gpu_memory",
lambda *args, **kwargs: 0.0,
),
get_parallel().override(attn_cp_size=1),
get_parallel().override(attn_cp_size=1, tp_rank=0),
):
_reset_cuda_graph_test_buffers()
return EAGLEDraftCudaGraphRunner(
@@ -463,15 +459,11 @@ def _capture_frozen_kv_mtp_graph_runner(
"sglang.srt.model_executor.runner.decode_cuda_graph_runner.graph_capture",
_single_rank_graph_capture,
),
patch(
"sglang.srt.model_executor.runner.decode_cuda_graph_runner.get_tensor_model_parallel_rank",
lambda: 0,
),
patch(
"sglang.srt.model_executor.runner.decode_cuda_graph_runner.get_available_gpu_memory",
lambda *args, **kwargs: 0.0,
),
get_parallel().override(attn_cp_size=1),
get_parallel().override(attn_cp_size=1, tp_rank=0),
):
_reset_cuda_graph_test_buffers()
return FrozenKVMTPCudaGraphRunner(worker)