diff --git a/python/sglang/kernels/ops/layernorm/mhc.py b/python/sglang/kernels/ops/layernorm/mhc.py index 3195265e9..9b069fb3f 100644 --- a/python/sglang/kernels/ops/layernorm/mhc.py +++ b/python/sglang/kernels/ops/layernorm/mhc.py @@ -13,12 +13,11 @@ from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) -from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.environ import envs from sglang.srt.layers.attention.dsa.utils import is_dsa_prefill_cp_interleave from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.utils.common import strict_contiguous -from sglang.srt.runtime_context import get_platform +from sglang.srt.runtime_context import get_parallel, get_platform from sglang.srt.utils.common import is_gfx1250_supported logger = logging.getLogger(__name__) @@ -1035,7 +1034,9 @@ def mhc_pre( # NCCL symmetric path: the Triton inplace MoE runner writes the expert # output back into this buffer, so a symmetric input yields a symmetric # all-reduce input. - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): layer_input = torch.empty( num_tokens, hidden_size, dtype=torch.bfloat16, device=residual.device ) @@ -1697,7 +1698,9 @@ def mhc_fused_post_pre( # layer_input_cur is the post-norm activation fed into the MoE; allocate it # in the symmetric memory pool so the Triton inplace MoE runner yields a # symmetric all-reduce input (see _mhc_pre_impl). - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): layer_input_cur = torch.empty( num_tokens, hidden_size, diff --git a/python/sglang/srt/disaggregation/ascend/conn.py b/python/sglang/srt/disaggregation/ascend/conn.py index 9c28fd871..c05cbdeb1 100644 --- a/python/sglang/srt/disaggregation/ascend/conn.py +++ b/python/sglang/srt/disaggregation/ascend/conn.py @@ -15,7 +15,7 @@ from sglang.srt.disaggregation.mooncake.conn import ( MooncakeKVReceiver, MooncakeKVSender, ) -from sglang.srt.distributed import get_pp_group +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils.network import get_local_ip_auto logger = logging.getLogger(__name__) @@ -130,7 +130,7 @@ class AscendKVManager(MooncakeKVManager): else: sliced_dst_kv_ptrs = [] start_layer = self.kv_args.prefill_start_layer - transfer_draft_kv = get_pp_group().is_last_rank and draft_kv_layers + transfer_draft_kv = get_parallel().pp_group.is_last_rank and draft_kv_layers if transfer_draft_kv: end_layer = start_layer + src_layers - draft_kv_layers else: diff --git a/python/sglang/srt/disaggregation/ascend/transfer_engine.py b/python/sglang/srt/disaggregation/ascend/transfer_engine.py index cac0b9c6e..96daaa072 100644 --- a/python/sglang/srt/disaggregation/ascend/transfer_engine.py +++ b/python/sglang/srt/disaggregation/ascend/transfer_engine.py @@ -57,10 +57,7 @@ class AscendTransferEngine(MooncakeTransferEngine): self.session_id = NetworkAddress(self.hostname, rpc_port).to_host_port_str() def initialize(self) -> None: - from sglang.srt.distributed.parallel_state import ( - get_world_group, - get_world_size, - ) + from sglang.srt.runtime_context import get_parallel transfer_protocol = self._get_transfer_protocol() if transfer_protocol == "device_rdma": @@ -68,10 +65,12 @@ class AscendTransferEngine(MooncakeTransferEngine): # through all_gather to avoid conflicts with rdma initialization. tmp_tensor = torch.zeros(1, device="npu") output_tensor_list = [ - torch.empty_like(tmp_tensor) for _ in range(get_world_size()) + torch.empty_like(tmp_tensor) for _ in range(get_parallel().world_size) ] torch.distributed.all_gather( - output_tensor_list, tmp_tensor, group=get_world_group().device_group + output_tensor_list, + tmp_tensor, + group=get_parallel().world_group.device_group, ) trans_op_type = self._resolve_trans_op_type(transfer_protocol) diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 1255a2056..d6935939f 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -31,7 +31,6 @@ from sglang.srt.disaggregation.utils import ( filter_kv_indices_for_cp_rank, get_dsv41_spec_layout, ) -from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.environ import envs from sglang.srt.runtime_context import ( get_disagg, @@ -259,7 +258,7 @@ class CommonKVManager(BaseKVManager): self._deferred_ack_targets: Dict[int, Tuple[str, int]] = {} self.req_to_decode_prefix_len: Dict[int, int] = {} self.decode_kv_args_table = {} - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group # If a timeout happens on the prefill side, it means prefill instances # fail to receive the KV indices from the decode instance of this request. # These timeout requests should be aborted to release the tree cache. @@ -1059,7 +1058,7 @@ class CommonKVManager(BaseKVManager): "multi-node prefill mode." ) - world_group = get_world_group() + world_group = get_parallel().world_group synced_port = world_group.broadcast_object(local_port, src=0) if synced_port != local_port: logger.info( @@ -1111,7 +1110,7 @@ class CommonKVManager(BaseKVManager): } if envs.SGLANG_RUST_SERVER.get() and self.attn_dp_size > 1: - topology_rows = get_world_group().all_gather_object(payload) + topology_rows = get_parallel().world_group.all_gather_object(payload) # Every scheduler contributes a topology row. Only the scheduler # ranks that own a Rust listener populate their local registry. if self.kv_args.rust_http_port is None: diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py index 68877a0ed..150a76580 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -35,7 +35,6 @@ from sglang.srt.disaggregation.encoder.receiver import ( from sglang.srt.distributed.parallel_state import ( get_default_distributed_backend, get_mooncake_transfer_engine, - get_tp_group, init_distributed_environment, initialize_model_parallel, ) @@ -654,7 +653,7 @@ class MMEncoder: get_parallel().tp_size, embedding_store=embedding_store, hidden_dims=self._embedding_dims, - tp_group=get_tp_group().cpu_group, + tp_group=get_parallel().tp_group.cpu_group, all_rank_get=False, dtype=self._embedding_dtype, ) @@ -1275,7 +1274,7 @@ class MMEncoder: layout_digest: tuple[int, int], ) -> List[torch.Tensor]: """Raise the same preparation error on every TP rank.""" - tp_group = get_tp_group() + tp_group = get_parallel().tp_group error_code = ( int( local_error.code diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index e4c590e43..18d069408 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -62,7 +62,6 @@ from sglang.srt.disaggregation.utils import ( prepare_abort, setup_state_kv_args, ) -from sglang.srt.distributed import get_pp_group from sglang.srt.environ import envs from sglang.srt.managers.schedule_batch import ( FINISH_ABORT, @@ -89,6 +88,7 @@ from sglang.srt.observability.scheduler_stage_metrics import ( ) from sglang.srt.runtime_context import ( get_disagg, + get_parallel, get_schedule, ) from sglang.srt.utils import is_npu @@ -266,7 +266,8 @@ class PrefillBootstrapQueue: draft_kv_pool = ( self.draft_token_to_kv_pool - if transfer_draft_cache and (not _is_npu or get_pp_group().is_last_rank) + if transfer_draft_cache + and (not _is_npu or get_parallel().pp_group.is_last_rank) else None ) num_draft_entries = 0 diff --git a/python/sglang/srt/elastic_ep/elastic_ep.py b/python/sglang/srt/elastic_ep/elastic_ep.py index 7894e11c8..abda6cd9e 100644 --- a/python/sglang/srt/elastic_ep/elastic_ep.py +++ b/python/sglang/srt/elastic_ep/elastic_ep.py @@ -7,7 +7,7 @@ from typing import TYPE_CHECKING, Callable, Iterator, List, Optional import torch -from sglang.srt.distributed import get_world_group, parallel_state +from sglang.srt.distributed import parallel_state from sglang.srt.distributed.utils import get_global_tcp_store from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata from sglang.srt.runtime_context import ( @@ -446,7 +446,7 @@ def join_process_groups() -> None: def get_healthy_expert_location_src_rank( *, invoked_in_elastic_ep_rejoin_path: bool ) -> int: - world_group = get_world_group() + world_group = get_parallel().world_group # NOTE: do not key off `self.server_args.elastic_ep_rejoin` here. # A rank that was started as a rejoin rank may later act as a healthy # rank in a subsequent recovery cycle. diff --git a/python/sglang/srt/elastic_ep/expert_backup_client.py b/python/sglang/srt/elastic_ep/expert_backup_client.py index 0f217aca6..adc04153a 100644 --- a/python/sglang/srt/elastic_ep/expert_backup_client.py +++ b/python/sglang/srt/elastic_ep/expert_backup_client.py @@ -7,10 +7,6 @@ from typing import Any, Callable import torch import zmq -from sglang.srt.distributed.parallel_state import ( - get_world_group, - get_world_size, -) from sglang.srt.environ import envs from sglang.srt.eplb.expert_location import get_global_expert_location_metadata from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send @@ -54,23 +50,23 @@ class ExpertBackupClient: self.buffer_size = 0 self.use_backup = False local_ip = get_local_ip_auto() - all_ips = [None] * get_world_size() + all_ips = [None] * get_parallel().world_size torch.distributed.all_gather_object( - all_ips, local_ip, group=get_world_group().cpu_group + all_ips, local_ip, group=get_parallel().world_group.cpu_group ) logger.info(f"all_ips: {all_ips}") for i in range(self.engine_num): self.recv_list[i] = context.socket(zmq.SUB) self.recv_list[i].connect( - f"tcp://{all_ips[i * get_world_size() // get_parallel().nnodes]}:{PORT_BASE + i * 2 + 1}" + f"tcp://{all_ips[i * get_parallel().world_size // get_parallel().nnodes]}:{PORT_BASE + i * 2 + 1}" ) self.recv_list[i].setsockopt(zmq.SUBSCRIBE, b"") # Synchronization channel to notify the manager when this client is ready. self.ready_sockets[i] = context.socket(zmq.PUSH) self.ready_sockets[i].connect( - f"tcp://{all_ips[i * get_world_size() // get_parallel().nnodes]}:{PORT_BASE + i * 2}" + f"tcp://{all_ips[i * get_parallel().world_size // get_parallel().nnodes]}:{PORT_BASE + i * 2}" ) sock_send(self.ready_sockets[i], UpdateExpertBackupReq()) diff --git a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py index fc855da78..d6ed7b35b 100644 --- a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py +++ b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py @@ -8,7 +8,7 @@ from flash_attn_interface import flash_attn_varlen_func from flash_attn_interface import flash_attn_with_kvcache as mate_flash_attn_with_kvcache from flash_attn_interface import get_scheduler_metadata -from sglang.srt.distributed import get_pp_group, get_pp_indices +from sglang.srt.distributed import get_pp_indices from sglang.srt.environ import envs from sglang.srt.hardware_backend.musa.layers.utils.cp_utils import ( musa_cp_attn_forward_extend as cp_attn_forward_extend, @@ -23,7 +23,7 @@ from sglang.srt.layers.utils.cp_utils import ( cp_allgather_and_save_kv_cache, ) from sglang.srt.mem_cache.memory_pool import KVWriteLoc -from sglang.srt.runtime_context import get_schedule +from sglang.srt.runtime_context import get_parallel, get_schedule if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -61,7 +61,7 @@ def _compute_scheduler_metadata( # Determine if scheduler metadata should be updated should_update = True - pp_group = get_pp_group() + pp_group = get_parallel().pp_group pp_rank = pp_group.rank_in_group start_layer_id, _ = get_pp_indices( backend.num_hidden_layers, pp_group.rank_in_group, pp_group.world_size diff --git a/python/sglang/srt/hardware_backend/npu/moe/fuseep.py b/python/sglang/srt/hardware_backend/npu/moe/fuseep.py index 609f6b6ca..1634e6f74 100644 --- a/python/sglang/srt/hardware_backend/npu/moe/fuseep.py +++ b/python/sglang/srt/hardware_backend/npu/moe/fuseep.py @@ -12,12 +12,11 @@ from typing import TYPE_CHECKING import torch -from sglang.srt.distributed import get_moe_ep_group from sglang.srt.environ import envs from sglang.srt.hardware_backend.npu.utils import npu_format_cast from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer from sglang.srt.layers.moe.utils import DeepEPMode -from sglang.srt.runtime_context import get_exec +from sglang.srt.runtime_context import get_exec, get_parallel if TYPE_CHECKING: from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE @@ -30,7 +29,7 @@ _PARAMS_BYTES = 2 # bf16 — Ascend's Dispatch & Combine does not support fp16 def _get_fuseep_buffer(layer: FusedMoE): DeepEPBuffer.set_dispatch_mode_as_low_latency() return DeepEPBuffer.get_deepep_buffer( - get_moe_ep_group().device_group, + get_parallel().moe_ep_group.device_group, layer.hidden_size, _PARAMS_BYTES, DeepEPMode.LOW_LATENCY, diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 26a457f62..2f0c73443 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -112,10 +112,6 @@ if _is_xpu: if _use_aiter: from aiter.ops.cache import indexer_k_quant_and_cache -from sglang.srt.distributed import ( - get_attn_tp_group, -) -from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers.cp.base import get_cp_strategy from sglang.srt.layers.cp.utils import is_cp_active @@ -157,7 +153,7 @@ if _is_cuda: def _broadcast_indexer_topk_from_rank0_impl(topk_indices: torch.Tensor) -> None: - group = get_attn_tp_group() + group = get_parallel().attn_tp_group if group.world_size == 1: return @@ -260,7 +256,9 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp): self.sm_count = deep_gemm.get_num_sms() self.half_device_sm_count = ceil_align(self.sm_count // 2, 8) pp_size = get_parallel().pp_size - self.logits_with_pp_recv = pp_size > 1 and not get_pp_group().is_last_rank + self.logits_with_pp_recv = ( + pp_size > 1 and not get_parallel().pp_group.is_last_rank + ) else: self.logits_with_pp_recv = False diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index a60b70f22..c5e7878cf 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -23,7 +23,6 @@ import torch from sglang.srt.distributed import ( attention_tensor_model_parallel_all_reduce, attention_tensor_model_parallel_quant_all_reduce, - get_tp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -266,7 +265,7 @@ class AttentionInputs: def tp_all_gather_hidden_states(self, hidden_states, forward_batch): total_tokens = forward_batch.input_ids.shape[0] output = hidden_states.new_empty((total_tokens, hidden_states.shape[-1])) - get_tp_group().all_gather_into_tensor(output, hidden_states) + get_parallel().tp_group.all_gather_into_tensor(output, hidden_states) return output def fetch_qkv_latent(self): @@ -510,7 +509,7 @@ def tp_reduce_scatter( ) local_tokens = hidden_states.shape[0] // context.tp_size output = hidden_states.new_empty(local_tokens, *hidden_states.shape[1:]) - get_tp_group().reduce_scatter_tensor(output, hidden_states) + get_parallel().tp_group.reduce_scatter_tensor(output, hidden_states) if residual is not None: residual = residual.tensor_split(context.tp_size)[context.tp_rank] return output, residual @@ -1075,7 +1074,7 @@ class CommunicateSimpleFn: gathered_hidden_states = [] for local_hidden_states in hidden_states: with use_symmetric_memory( - get_tp_group(), + get_parallel().tp_group, disabled=not is_allocation_symmetric(), ): output = torch.empty( @@ -1270,7 +1269,7 @@ class CommunicateWithAllReduceAndLayerNormFn: hidden_states ) with use_symmetric_memory( - get_tp_group(), + get_parallel().tp_group, disabled=not is_allocation_symmetric(), ): hidden_states, residual = layernorm(hidden_states, residual) @@ -1278,7 +1277,7 @@ class CommunicateWithAllReduceAndLayerNormFn: hidden_states += residual hidden_states, local_hidden_states = ( - get_global_dp_buffer(get_tp_group()), + get_global_dp_buffer(get_parallel().tp_group), hidden_states, ) if use_layer_norm_before_gather: @@ -1510,7 +1509,7 @@ class CommunicateSummableTensorPairFn: allow_reduce_scatter: bool = False, ): if get_parallel().tp_size == get_parallel().attn_dp_size: - group = get_tp_group() + group = get_parallel().tp_group else: group = get_parallel().attn_tp_group hidden_states, global_hidden_states = ( @@ -1518,7 +1517,7 @@ class CommunicateSummableTensorPairFn: hidden_states, ) if should_use_dp_reduce_scatterv(): - get_tp_group().reduce_scatterv( + get_parallel().tp_group.reduce_scatterv( global_hidden_states, output=hidden_states, sizes=get_dp_global_num_tokens(), @@ -1600,7 +1599,7 @@ class CommunicateSummableTensorPairFn: # DP scatter (if DP attention is enabled) if context.attn_dp_size > 1: if get_parallel().tp_size == get_parallel().attn_dp_size: - group = get_tp_group() + group = get_parallel().tp_group else: group = get_parallel().attn_tp_group hidden_states_output, global_hidden_states = ( diff --git a/python/sglang/srt/layers/communicator_mhc.py b/python/sglang/srt/layers/communicator_mhc.py index 8a61902fe..57f128253 100644 --- a/python/sglang/srt/layers/communicator_mhc.py +++ b/python/sglang/srt/layers/communicator_mhc.py @@ -18,7 +18,6 @@ from typing import Callable, Optional import torch from sglang.kernels.ops.layernorm.mhc import hc_contract, hc_expand -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.communication_op import ( attention_tensor_model_parallel_all_reduce, ) @@ -50,6 +49,7 @@ from sglang.srt.layers.dp_attention import ( ) from sglang.srt.layers.moe import should_use_dp_reduce_scatterv from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.runtime_context import get_parallel def tp_all_gather_hidden_states(hidden_states, forward_batch): @@ -58,7 +58,7 @@ def tp_all_gather_hidden_states(hidden_states, forward_batch): ) total_tokens = forward_batch.input_ids.shape[0] output = hidden_states.new_empty((total_tokens, hidden_states.shape[-1])) - get_tp_group().all_gather_into_tensor(output, hidden_states) + get_parallel().tp_group.all_gather_into_tensor(output, hidden_states) return output @@ -199,7 +199,7 @@ class MHCCommunicateWithAllReduceAndLayerNormFn(CommunicateWithAllReduceAndLayer return hidden_states, hidden_states scatter_states = hidden_states.tensor_split(context.tp_size)[context.tp_rank] - get_tp_group().reduce_scatter_tensor(scatter_states, hidden_states) + get_parallel().tp_group.reduce_scatter_tensor(scatter_states, hidden_states) scatter_states, residual = mhc.attn_to_mlp( scatter_states, residual, out_norm=layernorm @@ -238,7 +238,7 @@ class MHCCommunicateWithAllReduceAndLayerNormFn(CommunicateWithAllReduceAndLayer if context.attn_dp_size != 1: if hidden_states.shape[0] != 0: with use_symmetric_memory( - get_tp_group(), + get_parallel().tp_group, disabled=not is_allocation_symmetric(), ): hidden_states, residual = mhc.attn_to_mlp( @@ -248,7 +248,7 @@ class MHCCommunicateWithAllReduceAndLayerNormFn(CommunicateWithAllReduceAndLayer hidden_states, residual = mhc.attn_to_mlp(hidden_states, residual) hidden_states, local_hidden_states = ( - get_global_dp_buffer(get_tp_group()), + get_global_dp_buffer(get_parallel().tp_group), hidden_states, ) dp_gather_replicate(hidden_states, local_hidden_states, forward_batch) @@ -305,7 +305,7 @@ class MHCCommunicateSummableTensorPairFn(CommunicateSummableTensorPairFn): hidden_states = local_states.new_empty( local_states.shape[0] * context.tp_size, *local_states.shape[1:] ) - get_tp_group().all_gather_into_tensor(hidden_states, local_states) + get_parallel().tp_group.all_gather_into_tensor(hidden_states, local_states) return hidden_states, None @@ -322,13 +322,13 @@ class MHCCommunicateSummableTensorPairFn(CommunicateSummableTensorPairFn): **kwargs, ): hidden_states, global_hidden_states = ( - get_local_dp_buffer_mhc(get_tp_group(), 1), + get_local_dp_buffer_mhc(get_parallel().tp_group, 1), hidden_states, ) # MoE skips its post-expert all-reduce with reduce_scatterv, so this # scatter must reduce while combining local-expert partial sums. if should_use_dp_reduce_scatterv(): - get_tp_group().reduce_scatterv( + get_parallel().tp_group.reduce_scatterv( global_hidden_states, output=hidden_states, sizes=get_dp_global_num_tokens(), @@ -362,7 +362,7 @@ class MHCCommunicateSummableTensorPairFn(CommunicateSummableTensorPairFn): hidden_states, local_hidden_states = ( get_local_dp_buffer_mhc( - get_tp_group(), 1 if is_last_layer else mhc.hc_mult + get_parallel().tp_group, 1 if is_last_layer else mhc.hc_mult ), hidden_states, ) diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index 961412b7c..d2bc67116 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -15,16 +15,11 @@ from sglang.srt.arg_groups.model_override_base import ( ) from sglang.srt.distributed import ( GroupCoordinator, - get_attn_cp_group, - get_attn_tensor_model_parallel_rank, get_attn_tensor_model_parallel_world_size, - get_attn_tp_group, ) from sglang.srt.distributed import get_moe_dp_group as _get_moe_dp_group from sglang.srt.distributed import ( - get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size, - get_tp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -380,7 +375,7 @@ def initialize_dp_attention( dp.enabled = enable_dp_attention - tp_rank = get_tensor_model_parallel_rank() + tp_rank = get_parallel().tp_rank tp_size = get_tensor_model_parallel_world_size() _, _, attn_dp_rank, attn_dp_size = compute_dp_attention_world_info( @@ -503,9 +498,7 @@ 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_attn_tensor_model_parallel_rank() == 0 - ): + if local_tokens.shape[0] > 0 and (is_partial or get_parallel().attn_tp_rank == 0): assert local_tokens.untyped_storage() is not global_tokens.untyped_storage(), ( "aliasing between global_tokens and local_tokens not allowed" ) @@ -527,7 +520,9 @@ def _dp_gather_via_all_reduce( ): from sglang.srt.distributed.parallel_state import inplace_all_reduce - inplace_all_reduce(global_tokens, group_name=get_tp_group().unique_name) + inplace_all_reduce( + global_tokens, group_name=get_parallel().tp_group.unique_name + ) else: global_tokens[:] = tensor_model_parallel_all_reduce(global_tokens) @@ -549,16 +544,18 @@ def _dp_gather_via_all_gather( group=torch.distributed.group.WORLD, ) else: - get_tp_group().all_gather_into_tensor(global_tokens, local_tokens) + get_parallel().tp_group.all_gather_into_tensor(global_tokens, local_tokens) return if not is_partial: - if get_attn_tensor_model_parallel_rank() != 0: + if get_parallel().attn_tp_rank != 0: local_tokens.fill_(0) 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_parallel().attn_tp_rank] + get_parallel().attn_tp_group.reduce_scatter_tensor( + scattered_local_tokens, local_tokens + ) if use_world: torch.distributed.all_gather_into_tensor( global_tokens, @@ -566,7 +563,9 @@ def _dp_gather_via_all_gather( group=torch.distributed.group.WORLD, ) else: - get_tp_group().all_gather_into_tensor(global_tokens, scattered_local_tokens) + get_parallel().tp_group.all_gather_into_tensor( + global_tokens, scattered_local_tokens + ) # Variable-length DP-MoE gather (reference https://github.com/ROCm/ATOM/pull/930): instead of padding every @@ -696,7 +695,7 @@ def _dp_gather_via_all_gatherv_fp8( local_real.contiguous(), _DP_GATHER_FP8_GROUP ) gq, gs = _get_dp_gather_fp8_bufs(rows, hidden, global_tokens.device) - tp_group = get_tp_group() + tp_group = get_parallel().tp_group tp_group.all_gatherv(q.view(torch.uint8), sizes=sizes, output=gq) tp_group.all_gatherv(s, sizes=sizes, output=gs) _dequant_per_token_group_fp8_kernel[(rows,)]( @@ -785,7 +784,7 @@ def _dp_gather_via_all_gatherv( ): _dp_gather_via_all_gatherv_fp8(global_tokens, local_real, sizes) return - get_tp_group().all_gatherv(local_real, sizes=sizes, output=global_tokens) + get_parallel().tp_group.all_gatherv(local_real, sizes=sizes, output=global_tokens) def _note_dp_gather_in_prefill_graph() -> None: @@ -882,16 +881,18 @@ def dp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor): # the default reduce-scatter path if per-rank sizes are unavailable. sizes = get_dp_global_num_tokens() if sizes is not None: - get_tp_group().reduce_scatterv(input, output=output, sizes=sizes) + get_parallel().tp_group.reduce_scatterv(input, output=output, sizes=sizes) return if get_tensor_model_parallel_world_size() == get_parallel().attn_dp_size: - get_tp_group().reduce_scatter_tensor(output, input) + get_parallel().tp_group.reduce_scatter_tensor(output, input) else: scattered_local_tokens = input.tensor_split( get_tensor_model_parallel_world_size() - )[get_tensor_model_parallel_rank()] - get_tp_group().reduce_scatter_tensor(scattered_local_tokens, input) - get_attn_tp_group().all_gather_into_tensor(output, scattered_local_tokens) + )[get_parallel().tp_rank] + get_parallel().tp_group.reduce_scatter_tensor(scattered_local_tokens, input) + get_parallel().attn_tp_group.all_gather_into_tensor( + output, scattered_local_tokens + ) # --------------------------------------------------------------------------- @@ -989,29 +990,31 @@ def dp_reduce_scatterv_async( ev = _tbo_event(event_key) with torch.cuda.stream(comm): comm.wait_stream(compute) - get_tp_group().reduce_scatterv(global_tokens, output=output_local, sizes=sizes) + get_parallel().tp_group.reduce_scatterv( + global_tokens, output=output_local, sizes=sizes + ) ev.record(comm) return ev def attn_tp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor): - return get_attn_tp_group().reduce_scatter_tensor(output, input) + return get_parallel().attn_tp_group.reduce_scatter_tensor(output, input) def attn_cp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor): - return get_attn_cp_group().reduce_scatter_tensor(output, input) + return get_parallel().attn_cp_group.reduce_scatter_tensor(output, input) def attn_tp_all_reduce(input: torch.Tensor): - return get_attn_tp_group().all_reduce(input) + return get_parallel().attn_tp_group.all_reduce(input) def attn_tp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor): - return get_attn_tp_group().all_gather_into_tensor(output, input) + return get_parallel().attn_tp_group.all_gather_into_tensor(output, input) def attn_cp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor): - return get_attn_cp_group().all_gather_into_tensor(output, input) + return get_parallel().attn_cp_group.all_gather_into_tensor(output, input) def get_moe_cp_group() -> GroupCoordinator: @@ -1041,4 +1044,6 @@ 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_attn_tp_group().all_gather(input, output_tensor_list=output_list) + return get_parallel().attn_tp_group.all_gather( + input, output_tensor_list=output_list + ) diff --git a/python/sglang/srt/layers/engram.py b/python/sglang/srt/layers/engram.py index a1d499dc6..0cb2e953c 100644 --- a/python/sglang/srt/layers/engram.py +++ b/python/sglang/srt/layers/engram.py @@ -33,7 +33,6 @@ from sglang.kernels.ops.embeddings.engram_hash import ( engram_hash_ids_and_commit, ) from sglang.srt.distributed import tensor_model_parallel_all_reduce -from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import ( attn_cp_all_gather_into_tensor, @@ -708,7 +707,7 @@ class EngramEmbedding(nn.Module): layout, max(1, w_bytes + s_bytes), # mmap requires storage even for an empty shard. f"sglang_engram_{layer_id}", - get_tp_group(), + get_parallel().tp_group, ) raw = self.host_table.bytes[: w_bytes + s_bytes] weight = raw[:w_bytes].view(torch.float8_e4m3fn).view(n, dim) diff --git a/python/sglang/srt/layers/flashinfer_comm_fusion.py b/python/sglang/srt/layers/flashinfer_comm_fusion.py index 9cc7b795e..cdea99e21 100644 --- a/python/sglang/srt/layers/flashinfer_comm_fusion.py +++ b/python/sglang/srt/layers/flashinfer_comm_fusion.py @@ -6,12 +6,6 @@ import torch import torch.distributed as dist from torch.distributed import ProcessGroup -from sglang.srt.distributed import ( - get_attn_tp_group, - get_moe_ep_group, - get_moe_tp_group, - get_tp_group, -) from sglang.srt.distributed.parallel_state import in_the_same_node_as from sglang.srt.runtime_context import ( get_exec, @@ -335,7 +329,7 @@ def _preflight_check_workspace_memory( group = cpu_group if group is None: - tp_group = get_tp_group() + tp_group = get_parallel().tp_group if tp_group.world_size <= 1: return True group = tp_group.cpu_group @@ -670,12 +664,12 @@ def resolve_fusion_group(*, use_attn_tp_group: bool): parallel = get_parallel() if use_attn_tp_group: - return parallel.attn_tp_size, parallel.attn_tp_rank, get_attn_tp_group() + return parallel.attn_tp_size, parallel.attn_tp_rank, parallel.attn_tp_group if can_merge_post_experts_all_reduce(): - return parallel.tp_size, parallel.tp_rank, get_tp_group() + return parallel.tp_size, parallel.tp_rank, parallel.tp_group if parallel.moe_ep_size > 1: - return parallel.moe_ep_size, parallel.moe_ep_rank, get_moe_ep_group() - return parallel.moe_tp_size, parallel.moe_tp_rank, get_moe_tp_group() + return parallel.moe_ep_size, parallel.moe_ep_rank, parallel.moe_ep_group + return parallel.moe_tp_size, parallel.moe_tp_rank, parallel.moe_tp_group def _sync_allreduce_unavailable_across_tp(): @@ -691,7 +685,7 @@ def _sync_allreduce_unavailable_across_tp(): try: import torch.distributed as dist - tp_group = get_tp_group() + tp_group = get_parallel().tp_group if tp_group.world_size <= 1: return flag = torch.tensor( diff --git a/python/sglang/srt/layers/flashinfer_mnnvl_cutedsl.py b/python/sglang/srt/layers/flashinfer_mnnvl_cutedsl.py index da085bb43..c824e4c0c 100644 --- a/python/sglang/srt/layers/flashinfer_mnnvl_cutedsl.py +++ b/python/sglang/srt/layers/flashinfer_mnnvl_cutedsl.py @@ -308,10 +308,10 @@ def get_flashinfer_mnnvl_cutedsl_ar_fusion( assert max_m is not None assert rms_epsilon is not None assert weight_bias is not None - from sglang.srt.distributed.parallel_state import get_tp_group + from sglang.srt.runtime_context import get_parallel device = torch.device("cuda", torch.cuda.current_device()) - process_group = get_tp_group().device_group + process_group = get_parallel().tp_group.device_group domain = ( int(hidden_size), int(top_k), diff --git a/python/sglang/srt/layers/k3_ar_fusion.py b/python/sglang/srt/layers/k3_ar_fusion.py index b91774a5c..111ebb98c 100644 --- a/python/sglang/srt/layers/k3_ar_fusion.py +++ b/python/sglang/srt/layers/k3_ar_fusion.py @@ -25,7 +25,6 @@ import torch import sglang.srt.runtime_context as ctx from sglang.kernels.jit.utils import cache_once from sglang.srt.environ import envs -from sglang.srt.runtime_context import get_parallel if TYPE_CHECKING: from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( @@ -84,11 +83,11 @@ def _get_state() -> Optional[_State]: from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( CustomAllReduceV2, ) - from sglang.srt.distributed.parallel_state import get_tp_group + from sglang.srt.runtime_context import get_parallel if get_parallel().tp_size <= 1: return None - group = get_tp_group() + group = get_parallel().tp_group comm = group.ca_comm if ( not isinstance(comm, CustomAllReduceV2) @@ -198,9 +197,9 @@ def symm_buffer( its own. Each name belongs to one group, so the name alone identifies it. """ if group_name is None: - from sglang.srt.distributed.parallel_state import get_tp_group + from sglang.srt.runtime_context import get_parallel - group_name = get_tp_group().cpu_group.group_name + group_name = get_parallel().tp_group.cpu_group.group_name buf: _Buffer = ctx.get_buffer( f"k3_symm:{name}", lambda: _create_buffer(name, width, dtype, group_name) ) diff --git a/python/sglang/srt/layers/k3_gemm_ar.py b/python/sglang/srt/layers/k3_gemm_ar.py index 9a36809a8..2bbc06813 100644 --- a/python/sglang/srt/layers/k3_gemm_ar.py +++ b/python/sglang/srt/layers/k3_gemm_ar.py @@ -56,7 +56,7 @@ def maybe_wrap_o_proj(o_proj: RowParallelLinear) -> None: if not _init(): return from sglang.kernels.ops.kimi_k3 import gemm_ar as mod - from sglang.srt.distributed.parallel_state import get_tp_group + from sglang.srt.runtime_context import get_parallel parallel = get_parallel() world_size = parallel.tp_size @@ -73,7 +73,7 @@ def maybe_wrap_o_proj(o_proj: RowParallelLinear) -> None: mod.init( world_size=world_size, rank=parallel.tp_rank, - group=get_tp_group().cpu_group, + group=get_parallel().tp_group.cpu_group, k=weight.shape[1], ) # per-K compile + base-address stash, pre-capture diff --git a/python/sglang/srt/layers/layernorm_sp.py b/python/sglang/srt/layers/layernorm_sp.py index e3205a535..6f3f146bc 100644 --- a/python/sglang/srt/layers/layernorm_sp.py +++ b/python/sglang/srt/layers/layernorm_sp.py @@ -39,7 +39,6 @@ from typing import Optional import torch -from sglang.srt.distributed import get_tp_group from sglang.srt.runtime_context import ( get_flags, get_forward, @@ -112,7 +111,7 @@ def sp_entry_scatter(hidden_states: torch.Tensor) -> torch.Tensor: """ num_tokens = hidden_states.shape[0] set_sp_num_tokens(num_tokens) - tp_group = get_tp_group() + tp_group = get_parallel().tp_group tp_size = tp_group.world_size if tp_size == 1: return hidden_states @@ -127,7 +126,7 @@ def sp_entry_scatter(hidden_states: torch.Tensor) -> torch.Tensor: def sp_exit_gather(hidden_states: torch.Tensor, num_tokens: int) -> torch.Tensor: """g: all-gather the per-rank shards back to the full sequence along dim 0, then narrow to ``num_tokens`` (dropping the entry-scatter padding).""" - tp_group = get_tp_group() + tp_group = get_parallel().tp_group tp_size = tp_group.world_size if tp_size == 1: return hidden_states[:num_tokens] @@ -207,7 +206,7 @@ def column_parallel_g_matmul( """ num_tokens = sp_num_tokens() if sp_fused_matmul_eligible(linear): - group_name = get_tp_group().device_group.group_name + group_name = get_parallel().tp_group.device_group.group_name _, mm_outputs = torch.ops.symm_mem.fused_all_gather_matmul( input_parallel.contiguous(), [linear.weight.t()], @@ -235,7 +234,7 @@ def row_parallel_gbar_matmul(linear, input_: torch.Tensor, bias) -> torch.Tensor if padded != num_tokens: x = torch.nn.functional.pad(x, (0, 0, 0, padded - num_tokens)) if sp_fused_matmul_eligible(linear): - group_name = get_tp_group().device_group.group_name + group_name = get_parallel().tp_group.device_group.group_name return torch.ops.symm_mem.fused_matmul_reduce_scatter( x, linear.weight.t(), @@ -245,5 +244,5 @@ def row_parallel_gbar_matmul(linear, input_: torch.Tensor, bias) -> torch.Tensor ) full = linear.quant_method.apply(linear, x, bias) output = full.new_empty((padded // tp_size, *full.shape[1:])) - get_tp_group().reduce_scatter_tensor(output, full) + get_parallel().tp_group.reduce_scatter_tensor(output, full) return output diff --git a/python/sglang/srt/layers/linear.py b/python/sglang/srt/layers/linear.py index 1f3ee1f23..4504bbcb7 100644 --- a/python/sglang/srt/layers/linear.py +++ b/python/sglang/srt/layers/linear.py @@ -15,7 +15,6 @@ from torch.nn.parameter import Parameter, UninitializedParameter from sglang.kernels.kernel_api_logging import wrap_method_with_debug_kernel_once from sglang.srt.distributed import ( divide, - get_tp_group, split_tensor_along_last_dim, tensor_model_parallel_all_gather, tensor_model_parallel_all_reduce, @@ -1648,7 +1647,7 @@ class RowParallelLinear(LinearBase): symm_ctx = use_symmetric_memory(get_parallel().attn_tp_group) else: symm_ctx = use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ) with symm_ctx: if output_tensor is None: diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 02cedbe4f..f51bbb93f 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -26,7 +26,6 @@ from sglang.kernels.ops.activation.softcap import ( softcap_inplace_logits as fused_softcap, ) from sglang.srt.beam_search.logits_capture import BeamLogitsCapture -from sglang.srt.distributed import get_attn_tp_group, get_tp_group from sglang.srt.distributed.device_communicators import triton_symm_mem_ag from sglang.srt.environ import envs from sglang.srt.layers import layernorm_sp @@ -463,7 +462,12 @@ class LogitsProcessor(nn.Module): self.do_tensor_parallel_all_gather and not self.do_tensor_parallel_all_gather_dp_attn ): - group = get_attn_tp_group() if self.use_attn_tp_group else get_tp_group() + parallel = get_parallel() + group = ( + parallel.attn_tp_group + if self.use_attn_tp_group + else parallel.tp_group + ) chunking_group = group.cpu_group self.input_logprob_processor = InputLogprobProcessor( self.vocab_size, chunking_group=chunking_group @@ -1061,7 +1065,9 @@ class LogitsProcessor(nn.Module): """Exchange only the row block owned by each destination DP rank.""" logits = logits.contiguous() all_to_all_output = torch.empty_like(logits) - get_tp_group().all_to_all_single(all_to_all_output.view(-1), logits.view(-1)) + get_parallel().tp_group.all_to_all_single( + all_to_all_output.view(-1), logits.view(-1) + ) return _reassemble_tp_lm_head_all_to_all_output( all_to_all_output, get_parallel().tp_size ) diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index cf003a19d..188964489 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -14,8 +14,6 @@ from sglang.srt.batch_overlap.single_batch_overlap import DownGemmOverlapArgs from sglang.srt.batch_overlap.two_batch_overlap import MaybeTboDeepEPDispatcher from sglang.srt.configs.moe_model_registry import model_requires_fp32_silu_mul from sglang.srt.distributed import ( - get_moe_ep_group, - get_tp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -150,13 +148,13 @@ def _maybe_copy_weight_view_before_h2d( def _get_deepep_comm_group(a2a_backend): - group = get_tp_group().device_group + group = get_parallel().tp_group.device_group if a2a_backend.is_mori(): - group = get_tp_group() + group = get_parallel().tp_group elif _is_npu: - group = get_moe_ep_group().device_group + group = get_parallel().moe_ep_group.device_group return group @@ -204,7 +202,7 @@ def create_moe_dispatcher( _deepep_v2_experts_are_fp8(quant_method) ) return DeepEPv2Dispatcher( - group=get_tp_group().device_group, + group=get_parallel().tp_group.device_group, router_topk=moe_runner_config.top_k, num_experts=moe_runner_config.num_experts, num_local_experts=moe_runner_config.num_local_experts, @@ -219,7 +217,7 @@ def create_moe_dispatcher( ) elif a2a_backend.is_flashinfer(): return FlashinferDispatcher( - group=get_tp_group().device_group, + group=get_parallel().tp_group.device_group, router_topk=moe_runner_config.top_k, num_experts=moe_runner_config.num_experts, num_local_experts=moe_runner_config.num_local_experts, @@ -1583,7 +1581,7 @@ class FusedMoE(torch.nn.Module): dwdp_mgr.record_compute_and_prefetch_next(self.layer_id) with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): final_hidden_states = self.dispatcher.combine(combine_input=combine_input) diff --git a/python/sglang/srt/layers/moe/mega_moe.py b/python/sglang/srt/layers/moe/mega_moe.py index 20a94b338..4adb7826d 100644 --- a/python/sglang/srt/layers/moe/mega_moe.py +++ b/python/sglang/srt/layers/moe/mega_moe.py @@ -190,7 +190,7 @@ def _run_mega_routed( ) -> torch.Tensor: import deep_gemm - from sglang.srt.distributed.parallel_state import get_moe_ep_group + from sglang.srt.runtime_context import get_parallel hidden_size = moe.config.hidden_size @@ -216,7 +216,7 @@ def _run_mega_routed( topk_ids = None topk_weights = None - ep_group = get_moe_ep_group().device_group + ep_group = get_parallel().moe_ep_group.device_group num_experts = moe.experts.num_experts top_k = moe.config.num_experts_per_tok + moe.num_fused_shared_experts intermediate_size = moe.config.moe_intermediate_size diff --git a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py index af79c696a..7fe5311f4 100644 --- a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py +++ b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py @@ -17,7 +17,6 @@ from sglang.kernels.ops.quantization.per_token_group_quant import per_token_grou logger = logging.getLogger(__name__) -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -596,7 +595,7 @@ class DeepGemmRunnerCore(MoeRunnerCore): # symmetric path. Only this final output enters the pool; intermediate # buffers stay on the default allocator to bound pool occupancy. with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): down_output = torch.empty( (all_tokens, K), @@ -678,7 +677,7 @@ class DeepGemmRunnerCore(MoeRunnerCore): # GroupGemm-2: (M, N/2) (E, K, N/2) -> (M, K) with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): down_output = torch.empty( (all_tokens, K), @@ -855,7 +854,7 @@ class DeepGemmRunnerCore(MoeRunnerCore): activation_scale_width=down_input_scale.shape[-1], ) with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): down_output = torch.empty( (num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16 @@ -948,7 +947,7 @@ class DeepGemmRunnerCore(MoeRunnerCore): n = w2_weight.shape[1] with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): down_output = torch.empty( (num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16 @@ -1259,7 +1258,9 @@ def post_permute_deep_gemm_to_standard( src2dst = running_state["src2dst"] - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): output = torch.empty( hidden_states_shape, dtype=hidden_states_dtype, device=hidden_states_device ) diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py index ff6aef5ad..26aee883b 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py @@ -14,7 +14,6 @@ from typing import TYPE_CHECKING, Optional import torch from sglang.kernels.ops.quantization.fp8_kernel import scaled_fp8_quant -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -25,6 +24,7 @@ from sglang.srt.layers.moe.moe_runner.base import ( MoeRunnerConfig, register_fused_func, ) +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import is_flashinfer_available from sglang.srt.utils.common import next_power_of_2 @@ -206,7 +206,7 @@ def _run_flashinfer_cutlass( if output is None: with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): output = torch.empty( x.shape[0], @@ -435,7 +435,9 @@ def _fused_experts_flashinfer_mxfp4_cutlass( # new keyword at all on the existing W4A16/MXFP8 paths, so those paths keep # working with SGLang's currently pinned release. humming_kwargs = {"use_wfp4afp8_humming": True} if use_wfp4afp8_humming else {} - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): out = torch.empty(x.shape[0], out_hidden, dtype=output_dtype, device=x.device) flashinfer_cutlass_fused_moe( diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py index cc25486b9..95ba2c4b3 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py @@ -15,7 +15,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( ) # Import to register custom ops for torch.compile compatibility -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( is_symmetric_memory_enabled, is_tensor_in_symmetric_mempool, @@ -34,6 +33,7 @@ from sglang.srt.layers.moe.moe_runner.base import ( register_fused_func, ) from sglang.srt.layers.utils import copy_or_rebind_param +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils.common import ( is_flashinfer_available, next_power_of_2, @@ -817,7 +817,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8( # The deferred path returns FlashInfer's permuted/padded GEMM2 # materialization and must not allocate the ordinary final output. with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): symm_output = torch.empty( hidden_states.shape[0], @@ -943,7 +943,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8( # Allocate output inside symmetric memory context with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): symm_output = torch.empty( hidden_states.shape[0], @@ -1083,7 +1083,7 @@ def _fused_experts_flashinfer_mxfp4_sm100_trtllm_gen( ) if symm_output is None: with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): symm_output = torch.empty( num_tokens, @@ -1401,7 +1401,9 @@ def fused_experts_none_to_flashinfer_trtllm_fp4( ): symm_output = _provided else: - with use_symmetric_memory(get_tp_group(), disabled=not _symm_required): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not _symm_required + ): symm_output = torch.empty( num_tokens, hidden_size, @@ -1567,7 +1569,9 @@ def fused_experts_none_to_flashinfer_trtllm_bf16( hidden_states = dispatch_output.hidden_states topk_output = dispatch_output.topk_output - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): if use_routed_topk: assert runner_config.top_k is not None, ( "runner_config.top_k is required for flashinfer_trtllm_routed." diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py index 07a7bb6b9..677565008 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py @@ -22,14 +22,13 @@ from sglang.kernels.ops.moe.fused_moe_triton_kernels import ( support_tensor_descriptor, ) from sglang.srt.batch_invariant_ops import is_batch_invariant_mode_enabled -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig from sglang.srt.layers.moe.utils import get_moe_padding_size, get_moe_runner_backend -from sglang.srt.runtime_context import get_exec +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -578,7 +577,7 @@ def _fused_moe_kernel_sequence( # symmetric path. Only this output enters the pool; the intermediate caches # below stay on the default allocator to bound pool occupancy. with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): out_hidden_states = torch.empty_like(hidden_states) diff --git a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py index 2e334c826..4502751cb 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py @@ -7,7 +7,6 @@ from contextlib import nullcontext from dataclasses import dataclass from typing import TYPE_CHECKING, List, NamedTuple, Optional, Tuple, Union -from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.layers import deep_gemm_wrapper @@ -29,6 +28,7 @@ from sglang.srt.layers.moe.utils import ( get_deepep_output_dtype, is_tbo_enabled, ) +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import ( get_bool_env_var, get_cuda_version, @@ -100,7 +100,7 @@ def _deepep_precompile_tp_barrier() -> None: # To avoid this, we use torch.distributed's barrier during the compile stage. # We apply this barrier only in the compile stage to prevent extra all-reduce overhead at runtime. if envs.SGLANG_IN_DEEPGEMM_PRECOMPILE_STAGE.get(): - get_tp_group().barrier() + get_parallel().tp_group.barrier() class DeepEPPDispatchHooks(DispatcherBaseHooks): diff --git a/python/sglang/srt/layers/moe/token_dispatcher/standard.py b/python/sglang/srt/layers/moe/token_dispatcher/standard.py index b712c8128..e2bf17681 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/standard.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/standard.py @@ -4,9 +4,6 @@ from typing import TYPE_CHECKING, NamedTuple, Optional, Tuple import torch -from sglang.srt.distributed import ( - get_tp_group, -) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -153,7 +150,7 @@ class StandardDispatcher(BaseDispatcher): # Quantize before comm, swizzle after. with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): if hidden_states.shape[0] > 0: x, x_sf = fp4_quantize_flashinfer( @@ -167,7 +164,7 @@ class StandardDispatcher(BaseDispatcher): x_sf = torch.zeros( 0, x_col // 16, dtype=torch.uint8, device=hidden_states.device ) - topk_weights, topk_ids, x, x_sf = get_tp_group().all_gatherv( + topk_weights, topk_ids, x, x_sf = get_parallel().tp_group.all_gatherv( [topk_weights, topk_ids, x, x_sf], sizes=get_dp_global_num_tokens() ) # TODO: fuse into cutlass moe @@ -251,10 +248,10 @@ class StandardDispatcher(BaseDispatcher): (hidden_states,) = combine_input if should_use_flashinfer_cutlass_moe_fp4_allgather(): hidden_states, global_hidden_states = ( - get_local_dp_buffer(get_tp_group()), + get_local_dp_buffer(get_parallel().tp_group), hidden_states, ) - get_tp_group().reduce_scatterv( + get_parallel().tp_group.reduce_scatterv( global_hidden_states, output=hidden_states, sizes=get_dp_global_num_tokens(), diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index a734ca560..21586f29a 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -87,9 +87,6 @@ except ImportError: from sglang.kernels.fused_op import BaseFusedOp from sglang.kernels.ops.attention.dsv4 import mask_topk_ids -from sglang.srt.distributed import ( - get_tp_group, -) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -693,7 +690,7 @@ class TopK(BaseFusedOp): else: self.topk_config.torch_native = False with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): topk_output = select_experts( hidden_states=hidden_states, @@ -784,7 +781,7 @@ class TopK(BaseFusedOp): ) topk = self.topk_config.top_k - self.topk_config.num_fused_shared_experts with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): topk_weights = torch.empty((0, topk), dtype=torch.float32, device=device) topk_ids = torch.full((0, topk), -1, dtype=torch.int32, device=device) diff --git a/python/sglang/srt/layers/moe/waterfill.py b/python/sglang/srt/layers/moe/waterfill.py index 66e61e397..aced3aa7b 100644 --- a/python/sglang/srt/layers/moe/waterfill.py +++ b/python/sglang/srt/layers/moe/waterfill.py @@ -167,12 +167,12 @@ class WaterfillBalancer: local_routed_counts: Tensor, num_tokens: int ) -> Tuple[Tensor, Tensor]: """Aggregate dynamic load with SGLang EP communication.""" - from sglang.srt.distributed import get_moe_ep_group from sglang.srt.distributed.communication_op import ( moe_expert_parallel_all_reduce, ) + from sglang.srt.runtime_context import get_parallel - group = get_moe_ep_group() + group = get_parallel().moe_ep_group world = group.world_size buf = torch.zeros( world * 2, dtype=torch.int64, device=local_routed_counts.device diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_mxint4_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_mxint4_moe.py index 50e3ea6cf..c29e01efe 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_mxint4_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_mxint4_moe.py @@ -6,7 +6,6 @@ from typing import TYPE_CHECKING import torch from compressed_tensors import CompressionFormat -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -325,7 +324,7 @@ class CompressedTensorsMxInt4MoE(CompressedTensorsMoEScheme): ) with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): num_tokens = x.shape[0] hidden_size = x.shape[-1] diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index 9edc1d962..69e823c68 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -18,7 +18,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( per_token_group_quant_fp8, scaled_fp8_quant, ) -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -2902,7 +2901,7 @@ class Fp8MoEMethod(FusedMoEMethodBase): from sglang.srt.layers.moe.cutlass_moe import cutlass_fused_experts_fp8 with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): symm_output = torch.empty_like(x) diff --git a/python/sglang/srt/layers/quantization/moe_wna16.py b/python/sglang/srt/layers/quantization/moe_wna16.py index 711f361bb..6370378e0 100644 --- a/python/sglang/srt/layers/quantization/moe_wna16.py +++ b/python/sglang/srt/layers/quantization/moe_wna16.py @@ -9,7 +9,6 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional import numpy as np import torch -from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo from sglang.srt.layers.quantization.awq import AWQConfig @@ -453,7 +452,8 @@ class MoeWNA16Method(FusedMoEMethodBase): if not layer.quant_config.has_zp and "qzeros" in weight_name: return - device = get_tp_group().device + tp_group = get_parallel().tp_group + device = tp_group.device tp_rank = get_parallel().tp_rank loaded_weight = loaded_weight.to(device) shard_size = layer.intermediate_size_per_partition diff --git a/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py b/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py index 1b948e017..f90160145 100644 --- a/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py +++ b/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py @@ -7,7 +7,6 @@ import torch from torch.nn import Module from torch.nn.parameter import Parameter -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -15,6 +14,7 @@ from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe.utils import RoutingMethodType from sglang.srt.runtime_context import ( get_exec, + get_parallel, get_platform, ) from sglang.srt.utils import ( @@ -418,7 +418,7 @@ class Mxfp4FlashinferTrtllmMoEMethod: symm_output = None if not defer_finalize: with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): out_hidden_size = ( x_quant.shape[-1] * 2 @@ -530,7 +530,7 @@ def _fused_finalize_all_reduce_comm_world_size() -> Optional[int]: CustomAllReduceV2, ) - ca_comm = get_tp_group().ca_comm + ca_comm = get_parallel().tp_group.ca_comm if isinstance(ca_comm, CustomAllReduceV2) and not ca_comm.disabled: all_reduce_fusion.register_comm(ca_comm.obj) _fused_finalize_all_reduce_world_size = ca_comm.world_size @@ -560,7 +560,7 @@ def should_use_fuse_finalize_all_reduce( if not all_reduce_fusion.valid_cluster_sizes(hidden_dim): return False - tp_group = get_tp_group() + tp_group = get_parallel().tp_group if _fused_finalize_all_reduce_comm_world_size() != tp_group.world_size: return False # one push phase counter per row (the plane has num_sm of them) diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index 197bf7dff..de26d1eb6 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -7,7 +7,6 @@ import torch.distributed as dist from torch import nn from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32 -from sglang.srt.distributed import get_tp_group from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, @@ -109,7 +108,7 @@ def _select_sampling_mask_rows( class Sampler(nn.Module): def __init__(self): super().__init__() - self.tp_sync_group = get_tp_group().device_group + self.tp_sync_group = get_parallel().tp_group.device_group self.cp_sync_group = None if is_dp_attention_enabled(): self.tp_sync_group = get_parallel().attn_tp_group.device_group diff --git a/python/sglang/srt/layers/vocab_parallel_embedding.py b/python/sglang/srt/layers/vocab_parallel_embedding.py index 761f8e554..18f4ac3c4 100644 --- a/python/sglang/srt/layers/vocab_parallel_embedding.py +++ b/python/sglang/srt/layers/vocab_parallel_embedding.py @@ -14,7 +14,6 @@ from sglang.kernels.ops.embeddings.vocab_parallel_embedding import ( ) from sglang.srt.distributed import ( divide, - get_tp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -537,7 +536,7 @@ class VocabParallelEmbedding(torch.nn.Module): in-place fill deliberately stay outside the pool. """ symm_alloc = use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ) if self.tp_size == 1: with symm_alloc: diff --git a/python/sglang/srt/lora/mem_pool.py b/python/sglang/srt/lora/mem_pool.py index 0270ef3e6..7c8d805b8 100644 --- a/python/sglang/srt/lora/mem_pool.py +++ b/python/sglang/srt/lora/mem_pool.py @@ -17,7 +17,6 @@ import torch from sglang.srt.distributed import ( divide, - get_pp_group, ) from sglang.srt.environ import envs from sglang.srt.lora.eviction_policy import get_eviction_policy @@ -1457,7 +1456,7 @@ class LoRAMemoryPool: # Non-last PP stages do not own lm_head, so adapters can # legitimately contain lm_head LoRA weights with no local # module to load them into, otherwise we should have been able to load this weight. - assert not get_pp_group().is_last_rank, ( + assert not get_parallel().pp_group.is_last_rank, ( f"Failed to load lm_head LoRA weight: {name}, this is only expected to happen on non-last PP stages." ) continue diff --git a/python/sglang/srt/lora/trtllm_lora_temp/lora_dispatch.py b/python/sglang/srt/lora/trtllm_lora_temp/lora_dispatch.py index 972b1e4ce..4b062858c 100644 --- a/python/sglang/srt/lora/trtllm_lora_temp/lora_dispatch.py +++ b/python/sglang/srt/lora/trtllm_lora_temp/lora_dispatch.py @@ -19,11 +19,11 @@ from typing import TYPE_CHECKING import torch from sglang.kernels.ops.quantization.fp8_kernel import per_token_group_quant_fp8 -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) from sglang.srt.layers.dp_attention import is_allocation_symmetric +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils.common import next_power_of_2 if TYPE_CHECKING: @@ -167,7 +167,7 @@ def fused_experts_none_to_experimental_sgl_trtllm_fp8_lora( direct_down_output = None if use_virtual_lora_store: with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): direct_down_output = torch.empty( hidden_states.shape[0], @@ -277,7 +277,9 @@ def fused_experts_none_to_experimental_sgl_trtllm_fp8_lora( topk_ids, ) - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): output = torch.empty( hidden_states.shape[0], hidden_states.shape[1], @@ -401,7 +403,9 @@ def fused_experts_none_to_experimental_sgl_trtllm_bf16_lora( elif routing_method_type == RoutingMethodType.DeepSeekV3: routing_method_type = RoutingMethodType.TopK - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): direct_down_output = torch.empty( hidden_states.shape[0], hidden_states.shape[1], @@ -557,7 +561,9 @@ def fused_experts_none_to_experimental_sgl_trtllm_fp4_lora( topk_weights=topk_weights, ) - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): direct_down_output = torch.empty( hidden_states.shape[0], hidden_states.shape[1], diff --git a/python/sglang/srt/lora/trtllm_lora_temp/moe_overlap.py b/python/sglang/srt/lora/trtllm_lora_temp/moe_overlap.py index 12a2d62a4..9d2ba8244 100644 --- a/python/sglang/srt/lora/trtllm_lora_temp/moe_overlap.py +++ b/python/sglang/srt/lora/trtllm_lora_temp/moe_overlap.py @@ -62,7 +62,6 @@ def fused_experts_none_to_experimental_sgl_trtllm_fp8_lora_two_stream( merged_experts_fused_moe_lora_add, ) from sglang.kernels.ops.quantization.fp8_kernel import per_token_group_quant_fp8 - from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -73,6 +72,7 @@ def fused_experts_none_to_experimental_sgl_trtllm_fp8_lora_two_stream( from sglang.srt.lora.trtllm_lora_temp.shared_add_overlap import ( maybe_overlap_staged_shared_add, ) + from sglang.srt.runtime_context import get_parallel from sglang.srt.utils.common import next_power_of_2 assert runner_config.activation == "silu" and runner_config.is_gated, ( @@ -195,7 +195,9 @@ def fused_experts_none_to_experimental_sgl_trtllm_fp8_lora_two_stream( topk_weights=topk_weights, ) - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): direct_down_output = torch.empty( hidden_states.shape[0], hidden_states.shape[1], @@ -365,13 +367,13 @@ def fused_experts_none_to_experimental_sgl_trtllm_fp4_lora_two_stream( from sglang.kernels.ops.moe.trtllm_lora_temp.virtual_experts import ( merged_experts_fused_moe_lora_add, ) - from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput from sglang.srt.layers.moe.topk import TopKOutputChecker + from sglang.srt.runtime_context import get_parallel assert runner_config.activation == "silu" and runner_config.is_gated, ( "experimental_sgl_trtllm NVFP4 LoRA currently supports the gated SwiGLU path only." @@ -464,7 +466,9 @@ def fused_experts_none_to_experimental_sgl_trtllm_fp4_lora_two_stream( topk_ids=topk_ids, topk_weights=topk_weights, ) - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): direct_down_output = torch.empty( hidden_states.shape[0], hidden_states.shape[1], @@ -611,7 +615,6 @@ def fused_experts_none_to_experimental_sgl_trtllm_bf16_lora_two_stream( from sglang.kernels.ops.moe.trtllm_lora_temp.virtual_experts import ( merged_experts_fused_moe_lora_add, ) - from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -622,6 +625,7 @@ def fused_experts_none_to_experimental_sgl_trtllm_bf16_lora_two_stream( from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput from sglang.srt.layers.moe.topk import TopKOutputChecker from sglang.srt.layers.moe.utils import RoutingMethodType + from sglang.srt.runtime_context import get_parallel assert runner_config.activation == "silu" and runner_config.is_gated, ( "experimental_sgl_trtllm BF16 LoRA currently supports the gated SwiGLU path only." @@ -720,7 +724,9 @@ def fused_experts_none_to_experimental_sgl_trtllm_bf16_lora_two_stream( elif routing_method_type == RoutingMethodType.DeepSeekV3: routing_method_type = RoutingMethodType.TopK - with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + with use_symmetric_memory( + get_parallel().tp_group, disabled=not is_allocation_symmetric() + ): direct_down_output = torch.empty( hidden_states.shape[0], hidden_states.shape[1], diff --git a/python/sglang/srt/lora/trtllm_lora_temp/sgl_fp8_moe.py b/python/sglang/srt/lora/trtllm_lora_temp/sgl_fp8_moe.py index 3bcc0f4a1..732f5a262 100644 --- a/python/sglang/srt/lora/trtllm_lora_temp/sgl_fp8_moe.py +++ b/python/sglang/srt/lora/trtllm_lora_temp/sgl_fp8_moe.py @@ -29,7 +29,6 @@ def fused_experts_fp8_sgl( from sglang.kernels.ops.moe.trtllm_lora_temp.topk_pack import fused_pack_topk from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import ( - get_tp_group, is_allocation_symmetric, next_power_of_2, per_token_group_quant_fp8, @@ -46,6 +45,7 @@ def fused_experts_fp8_sgl( from sglang.srt.lora.trtllm_lora_temp.experimental_sgl_trtllm_moe import ( sgl_trtllm_fp8_block_scale_routed_moe_wrapper as trtllm_fp8_block_scale_routed_moe_wrapper, ) + from sglang.srt.runtime_context import get_parallel _SUPPORTED_FP8_ACTIVATIONS = {"silu", "relu2"} assert runner_config.activation in _SUPPORTED_FP8_ACTIVATIONS, ( @@ -98,7 +98,7 @@ def fused_experts_fp8_sgl( # Allocate output inside symmetric memory context with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): symm_output = torch.empty( hidden_states.shape[0], @@ -198,7 +198,7 @@ def fused_experts_fp8_sgl( # Allocate output inside symmetric memory context with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): symm_output = torch.empty( hidden_states.shape[0], diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 5de477fb7..2e70165b7 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -99,10 +99,8 @@ from sglang.srt.disaggregation.utils import ( prepare_abort, unified_memory_disagg_move_gate, ) -from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.distributed.parallel_state import ( abort_distributed_environment, - get_tp_group, ) from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.dllm.mixin.scheduler import SchedulerDllmMixin @@ -1213,14 +1211,14 @@ class Scheduler( ), ) - self.tp_group = get_tp_group() + self.tp_group = get_parallel().tp_group self.tp_cpu_group = self.tp_group.cpu_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_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() + self.pp_group = get_parallel().pp_group + self.world_group = get_parallel().world_group # NOTE: dp_tp_* are request/data-plane coordination groups (not tensor collectives). # When DP attention is enabled, scope to the attention-TP group; otherwise use diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index fd14b8d04..3766a72b8 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -7,7 +7,6 @@ import torch from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer from sglang.srt.configs.model_config import ModelConfig -from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.layers.cp.utils import get_cp_strategy @@ -192,9 +191,9 @@ class MLPSyncBatchInfo: ) num_ranks_in_tp_info = tp_info.shape[0] if device == "cpu": - tp_active_ranks = get_tp_group().active_ranks_cpu + tp_active_ranks = get_parallel().tp_group.active_ranks_cpu else: - tp_active_ranks = get_tp_group().active_ranks + tp_active_ranks = get_parallel().tp_group.active_ranks if tp_active_ranks.shape[0] < num_ranks_in_tp_info: tp_active_ranks = torch.ones( num_ranks_in_tp_info, @@ -432,9 +431,9 @@ def prepare_mlp_sync_batch_raw( tbo_preparer = TboDPAttentionPreparer() use_world_group = world_dp_gather_enabled() if use_world_group: - from sglang.srt.distributed.parallel_state import get_world_group + from sglang.srt.runtime_context import get_parallel - world = get_world_group() + world = get_parallel().world_group group = torch.distributed.group.WORLD device = world.device elif len(offload_tags) == 0 and ( diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index cad2ec201..db7259755 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -22,7 +22,6 @@ from typing import TYPE_CHECKING, List, Optional, Tuple import torch from sglang.srt.beam_search.logits_capture import capture_pre_sample_logits -from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.managers.io_struct import ( @@ -388,8 +387,8 @@ class TpModelWorker(BaseTpWorker): self.device = self.model_runner.device # Init nccl groups - self.pp_group = get_pp_group() - self.world_group = get_world_group() + self.pp_group = get_parallel().pp_group + self.world_group = get_parallel().world_group # Sync random seed across TP workers. # Elastic joiners and last-stage-only draft workers cannot enter the WORLD diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index ebb59e8e4..ee843abd9 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -28,7 +28,6 @@ from sglang.srt.configs.model_config import ( is_deepseek_v4, is_minimax_sparse, ) -from sglang.srt.distributed.parallel_state import get_world_group from sglang.srt.distributed.utils import get_pp_indices from sglang.srt.environ import envs from sglang.srt.layers.quantization.fp4_kv_cache_quant_method import ( @@ -2196,8 +2195,8 @@ class KVCacheConfigurator: available_gpu_memory = get_available_gpu_memory( self.device, self.gpu_id, - distributed=get_world_group().world_size > 1, - cpu_group=get_world_group().cpu_group, + distributed=get_parallel().world_group.world_size > 1, + cpu_group=get_parallel().world_group.cpu_group, ) slack_gb = pre_model_load_memory * (1 - get_schedule().mem_fraction_static) @@ -2292,7 +2291,7 @@ class KVCacheConfigurator: torch.distributed.all_reduce( tensor, op=torch.distributed.ReduceOp.MIN, - group=get_world_group().cpu_group, + group=get_parallel().world_group.cpu_group, ) token_capacity = tensor.item() diff --git a/python/sglang/srt/mem_cache/pool_host/base.py b/python/sglang/srt/mem_cache/pool_host/base.py index 7be9f25c3..7373a5fe5 100644 --- a/python/sglang/srt/mem_cache/pool_host/base.py +++ b/python/sglang/srt/mem_cache/pool_host/base.py @@ -9,7 +9,6 @@ from typing import Optional import psutil import torch -from sglang.srt.distributed.parallel_state import get_world_group from sglang.srt.mem_cache.memory_pool import KVCache from sglang.srt.mem_cache.pool_host.common import ( _cuda_host_unregister, @@ -40,7 +39,7 @@ def ranks_per_host() -> int: if not (torch.distributed.is_available() and torch.distributed.is_initialized()): return 1 try: - world_group = get_world_group() + world_group = get_parallel().world_group except AssertionError: return 1 if world_group.world_size == 1: @@ -75,9 +74,9 @@ def sync_fixed_hicache_size(size: int, host_size: int) -> int: return size try: - from sglang.srt.distributed.parallel_state import get_pp_group + from sglang.srt.runtime_context import get_parallel - pp_group = get_pp_group() + pp_group = get_parallel().pp_group except AssertionError: return size diff --git a/python/sglang/srt/mem_cache/storage/flexkv/__init__.py b/python/sglang/srt/mem_cache/storage/flexkv/__init__.py index a2bd129b6..49a12a0f9 100644 --- a/python/sglang/srt/mem_cache/storage/flexkv/__init__.py +++ b/python/sglang/srt/mem_cache/storage/flexkv/__init__.py @@ -22,19 +22,14 @@ def _flexkv_factory(ctx): """Build a :class:`FlexKVRadixCache` from a ``TreeCacheBuildContext``. ``TreeCacheBuildContext`` carries TP rank/size and the TP group - coordinator, but not PP/CP. We pick those up from the global - accessors in :mod:`sglang.srt.distributed.parallel_state`; FlexKV - needs them to fan out lookup/store decisions across the full TP × CP - × PP topology. + coordinator, but not PP/CP. We pick those up from ``get_parallel()``; + FlexKV needs them to fan out lookup/store decisions across the full + TP × CP × PP topology. """ - from sglang.srt.distributed.parallel_state import ( - get_attn_cp_group, - get_attn_tp_group, - get_pp_group, - ) from sglang.srt.mem_cache.storage.flexkv.flexkv_radix_cache import ( FlexKVRadixCache, ) + from sglang.srt.runtime_context import get_parallel server_args = ctx.server_args @@ -42,15 +37,15 @@ def _flexkv_factory(ctx): # the regular TP group when attn DP is off — that's fine, the # connector treats size-1 groups as no-ops. try: - pp_group = get_pp_group() + pp_group = get_parallel().pp_group except (RuntimeError, AssertionError): pp_group = None try: - attn_tp_group = get_attn_tp_group() + attn_tp_group = get_parallel().attn_tp_group except (RuntimeError, AssertionError): attn_tp_group = ctx.tp_group try: - attn_cp_group = get_attn_cp_group() + attn_cp_group = get_parallel().attn_cp_group except (RuntimeError, AssertionError): attn_cp_group = None diff --git a/python/sglang/srt/mem_cache/storage/flexkv/flexkv_comm.py b/python/sglang/srt/mem_cache/storage/flexkv/flexkv_comm.py index 609ba8f66..fa0341212 100644 --- a/python/sglang/srt/mem_cache/storage/flexkv/flexkv_comm.py +++ b/python/sglang/srt/mem_cache/storage/flexkv/flexkv_comm.py @@ -36,7 +36,7 @@ from typing import Any, Dict, List import torch import torch.distributed as dist -from sglang.srt.distributed.parallel_state import get_world_group +from sglang.srt.runtime_context import get_parallel logger = logging.getLogger(__name__) @@ -174,7 +174,7 @@ class FlexKVComm: self.pp_size > 1 or self.attn_tp_size > 1 or self.attn_cp_size > 1 ) - self._world_cpu_group = get_world_group().cpu_group + self._world_cpu_group = get_parallel().world_group.cpu_group self.pp_group = ( self.pp_cpu_group diff --git a/python/sglang/srt/model_executor/model_runner_components/attention_backend_setup.py b/python/sglang/srt/model_executor/model_runner_components/attention_backend_setup.py index 0f1fc9459..35d28e1c8 100644 --- a/python/sglang/srt/model_executor/model_runner_components/attention_backend_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/attention_backend_setup.py @@ -5,13 +5,13 @@ from typing import TYPE_CHECKING, Optional import msgspec -from sglang.srt.distributed import get_world_group from sglang.srt.environ import envs from sglang.srt.layers.attention.attention_registry import ( ATTENTION_BACKENDS, attn_backend_wrapper, ) from sglang.srt.layers.attention.tbo_backend import TboAttnBackend +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import init_cublas if TYPE_CHECKING: @@ -145,9 +145,9 @@ def build_attention_backends(*, model_runner: ModelRunner) -> AttentionBackends: lazy_init_zbal_gva_mem( model_runner.device, model_runner.gpu_id, - get_world_group().rank_in_group, - get_world_group().world_size, - get_world_group().cpu_group, + get_parallel().world_group.rank_in_group, + get_parallel().world_group.world_size, + get_parallel().world_group.cpu_group, ) # Record resolved per-mode backends on the backend for model dispatch. diff --git a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py index eda5b1099..4db397b96 100644 --- a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py @@ -8,7 +8,6 @@ from typing import TYPE_CHECKING, Any, Optional import msgspec from sglang.srt.configs.model_config import ModelImpl -from sglang.srt.distributed import get_world_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( prealloc_symmetric_memory_pool, ) @@ -224,7 +223,7 @@ def refresh_deep_gemm_layout_memory_budget( set_masked_standard_layout_memory_budget, ) - world_group = get_world_group() + world_group = get_parallel().world_group available_memory_gb = get_available_gpu_memory( model_runner.device, model_runner.gpu_id, diff --git a/python/sglang/srt/model_executor/model_runner_components/kv_pool_runtime.py b/python/sglang/srt/model_executor/model_runner_components/kv_pool_runtime.py index 984f29d51..4585fbbcb 100644 --- a/python/sglang/srt/model_executor/model_runner_components/kv_pool_runtime.py +++ b/python/sglang/srt/model_executor/model_runner_components/kv_pool_runtime.py @@ -8,7 +8,6 @@ import torch from sglang.srt.arg_groups.overrides import post_capture_kv_sizing_planned from sglang.srt.configs.hybrid_arch import mambaish_config -from sglang.srt.distributed import get_world_group from sglang.srt.mem_cache.kv_cache_configurator import mm_runtime_reservation_gb from sglang.srt.model_executor.cuda_graph_config import Backend from sglang.srt.model_executor.runner_utils.pool import graph_pool_borrow_enabled @@ -17,6 +16,7 @@ from sglang.srt.runtime_context import ( get_disagg, get_exec, get_mm, + get_parallel, pre_capture_activation_reserve_mb, ) from sglang.srt.utils.common import get_available_gpu_memory, get_device_memory_capacity @@ -59,8 +59,8 @@ def compute_post_capture_kv_resize( free_gb = get_available_gpu_memory( model_runner.device, model_runner.gpu_id, - distributed=get_world_group().world_size > 1, - cpu_group=get_world_group().cpu_group, + distributed=get_parallel().world_group.world_size > 1, + cpu_group=get_parallel().world_group.cpu_group, ) headroom_gb = model_runner.pre_model_load_memory * ( 1 - model_runner.mem_fraction_static diff --git a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py index b43cfdb45..c978df0ee 100644 --- a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py @@ -21,7 +21,6 @@ from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS from sglang.srt.debug_utils.tensor_dump_forward_hook import ( register_forward_hook_for_model, ) -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state from sglang.srt.model_loader.loader import get_model_loader from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( @@ -34,6 +33,7 @@ from sglang.srt.runtime_context import ( get_exec, get_model, get_observability, + get_parallel, ) from sglang.srt.utils.common import is_npu from sglang.srt.utils.network import NetworkAddress @@ -373,12 +373,12 @@ def dist_barrier_after_load( if elastic_ep_backend == "mooncake": # Mooncake does not support `monitored_barrier` if not is_ep_joiner: - dist.barrier(group=get_tp_group().cpu_group) + dist.barrier(group=get_parallel().tp_group.cpu_group) else: # Handle the case where some ranks do not finish loading. try: dist.monitored_barrier( - group=get_tp_group().cpu_group, + group=get_parallel().tp_group.cpu_group, timeout=datetime.timedelta(seconds=UNBALANCED_MODEL_LOADING_TIMEOUT_S), wait_all_ranks=True, ) diff --git a/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py b/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py index 5f5b63835..016fd2a38 100644 --- a/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/moe_ep_setup.py @@ -75,7 +75,7 @@ def prepare_moe_topk( def init_lplb_solvers(*, model_config: ModelConfig) -> None: """Initialize per-layer LPLB solvers from current expert location metadata.""" - from sglang.srt.distributed import get_moe_ep_group + from sglang.srt.runtime_context import get_parallel # Gate: refuse LP for non-DeepSeek MoE families whose empty-token paths # don't participate in the EP all-reduce (would deadlock under DP- @@ -88,7 +88,7 @@ def init_lplb_solvers(*, model_config: ModelConfig) -> None: if metadata is None: return clear_global_lplb_solvers() - ep_group = get_moe_ep_group() + ep_group = get_parallel().moe_ep_group for lid in range(metadata.num_layers): solver = LPLBSolver( phy2log=metadata.physical_to_logical_map[lid], diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 1d1a35c37..71ecd9ac5 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -2046,9 +2046,9 @@ class PreshardedModelLoader(DefaultModelLoader): cls, local_sig: Optional[str] ) -> Optional[str]: try: - from sglang.srt.distributed import get_world_group + from sglang.srt.runtime_context import get_parallel - group = get_world_group() + group = get_parallel().world_group if group.world_size <= 1: return local_sig all_sigs = group.all_gather_object(local_sig) @@ -2068,20 +2068,20 @@ class PreshardedModelLoader(DefaultModelLoader): @staticmethod def _world_rank_and_size() -> Tuple[int, int]: - from sglang.srt.distributed import get_world_group + from sglang.srt.runtime_context import get_parallel try: - g = get_world_group() + g = get_parallel().world_group return g.rank_in_group, g.world_size except (AssertionError, AttributeError): return 0, 1 @staticmethod def _world_barrier() -> None: - from sglang.srt.distributed import get_world_group + from sglang.srt.runtime_context import get_parallel try: - get_world_group().barrier() + get_parallel().world_group.barrier() except (AssertionError, AttributeError): pass diff --git a/python/sglang/srt/model_loader/weight_utils.py b/python/sglang/srt/model_loader/weight_utils.py index bb6cbba5c..993bc91e1 100644 --- a/python/sglang/srt/model_loader/weight_utils.py +++ b/python/sglang/srt/model_loader/weight_utils.py @@ -44,7 +44,6 @@ from sglang.srt.configs.model_config import ( ModelConfig, is_qwen3_5_mtp_draft, ) -from sglang.srt.distributed import get_world_group from sglang.srt.layers.quantization import QuantizationConfig, get_quantization_config from sglang.srt.layers.quantization.fp8 import Fp8Config from sglang.srt.layers.quantization.modelopt_quant import ( @@ -1007,7 +1006,7 @@ def _prefetch_all_checkpoints( # full checkpoint into its own page cache. Global rank would split files # across nodes, but page cache is not shared across nodes. if torch.distributed.is_initialized(): - world_group = get_world_group() + world_group = get_parallel().world_group local_rank = world_group.local_rank local_world_size = world_group.local_size or world_group.world_size else: diff --git a/python/sglang/srt/models/apertus.py b/python/sglang/srt/models/apertus.py index 6cbc918df..90241e21c 100644 --- a/python/sglang/srt/models/apertus.py +++ b/python/sglang/srt/models/apertus.py @@ -26,9 +26,6 @@ import torch from torch import nn from transformers import ApertusConfig -from sglang.srt.distributed import ( - get_pp_group, -) from sglang.srt.layers.activation import XIELU from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( @@ -298,7 +295,7 @@ class ApertusModel(nn.Module): self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size self.org_vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( config.vocab_size, @@ -430,7 +427,7 @@ class ApertusForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = self._init_model(config, quant_config, add_prefix("model", prefix)) diff --git a/python/sglang/srt/models/arcee.py b/python/sglang/srt/models/arcee.py index 79a6beff4..7c0e8d03f 100644 --- a/python/sglang/srt/models/arcee.py +++ b/python/sglang/srt/models/arcee.py @@ -20,9 +20,6 @@ import torch from torch import nn from transformers import LlamaConfig -from sglang.srt.distributed import ( - get_pp_group, -) from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( @@ -272,7 +269,7 @@ class ArceeModel(nn.Module): self.config = config self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( config.vocab_size, @@ -395,7 +392,7 @@ class ArceeForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = self._init_model(config, quant_config, add_prefix("model", prefix)) diff --git a/python/sglang/srt/models/bailing_mm.py b/python/sglang/srt/models/bailing_mm.py index 59533f104..5f9d446e6 100644 --- a/python/sglang/srt/models/bailing_mm.py +++ b/python/sglang/srt/models/bailing_mm.py @@ -21,7 +21,6 @@ import torch.nn as nn import torch.nn.functional as F from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.utils import PPMissingLayer from sglang.srt.managers.mm_utils import ( @@ -37,7 +36,7 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.bailing_moe import BailingMoeV2ForCausalLM from sglang.srt.models.qwen2_5_vl import Qwen2_5_VisionTransformer from sglang.srt.multimodal.mm_utils import materialize_multimodal_features -from sglang.srt.runtime_context import get_mm +from sglang.srt.runtime_context import get_mm, get_parallel from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -53,7 +52,7 @@ class BailingMMNativeForConditionalGeneration(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.use_data_parallel = get_mm().mm_enable_dp_encoder diff --git a/python/sglang/srt/models/bailing_mm_v3.py b/python/sglang/srt/models/bailing_mm_v3.py index dd604514f..021760c48 100644 --- a/python/sglang/srt/models/bailing_mm_v3.py +++ b/python/sglang/srt/models/bailing_mm_v3.py @@ -22,7 +22,6 @@ import torch.nn.functional as F from transformers import PretrainedConfig from sglang.srt.configs.bailing_hybrid import is_bailing_multi_gate_enabled -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.utils import PPMissingLayer from sglang.srt.managers.mm_utils import ( @@ -40,7 +39,7 @@ from sglang.srt.models.bailing_moe_v3 import ( ) from sglang.srt.models.qwen3_vl import Qwen3VLMoeVisionModel from sglang.srt.multimodal.mm_utils import materialize_multimodal_features -from sglang.srt.runtime_context import get_mm +from sglang.srt.runtime_context import get_mm, get_parallel from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -75,7 +74,7 @@ class BailingMoeV3VLForConditionalGeneration(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.norm_query_embeds = getattr(config, "norm_query_embeds", False) diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py index 7868411d1..36577fd59 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -28,8 +28,6 @@ from torch import nn from transformers import PretrainedConfig from sglang.srt.distributed import ( - get_pp_group, - parallel_state, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -331,7 +329,7 @@ class BailingMoESparseMoeBlock(nn.Module): self.ep_size = get_parallel().tp_size self.deepep_dispatcher = DeepEPDispatcher( - group=parallel_state.get_tp_group().device_group, + group=get_parallel().tp_group.device_group, router_topk=self.top_k, permute_fusion=True, num_experts=self.num_experts, @@ -790,7 +788,7 @@ class BailingMoEModel(nn.Module): prefix: str = "", ): super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.vocab_size = config.vocab_size self.embed_dim = config.hidden_size @@ -895,7 +893,7 @@ class BailingMoEForCausalLM(nn.Module): prefix: str = "", ): super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config alt_stream = get_stream("alt") if _is_cuda else None diff --git a/python/sglang/srt/models/bailing_moe_linear.py b/python/sglang/srt/models/bailing_moe_linear.py index 43d880b76..1777dc397 100644 --- a/python/sglang/srt/models/bailing_moe_linear.py +++ b/python/sglang/srt/models/bailing_moe_linear.py @@ -13,7 +13,6 @@ from sglang.kernels.ops.attention.fla.layernorm_gated import RMSNorm as RMSNormG from sglang.kernels.ops.attention.fla.layernorm_gated import layernorm_fn from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -923,7 +922,7 @@ class BailingMoELinearModel(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.vocab_size = config.vocab_size self.embed_dim = config.hidden_size @@ -1069,7 +1068,7 @@ class BailingMoELinearForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = BailingMoELinearModel( diff --git a/python/sglang/srt/models/bailing_moe_v3.py b/python/sglang/srt/models/bailing_moe_v3.py index a2dfac58d..6d0450211 100644 --- a/python/sglang/srt/models/bailing_moe_v3.py +++ b/python/sglang/srt/models/bailing_moe_v3.py @@ -20,7 +20,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( from sglang.srt.configs import KimiLinearConfig from sglang.srt.configs.bailing_hybrid import is_bailing_multi_gate_enabled from sglang.srt.distributed import ( - get_pp_group, moe_expert_parallel_all_reduce, moe_tensor_model_parallel_all_reduce, ) @@ -1263,7 +1262,7 @@ class BailingMoELinearModel(nn.Module): num_fused_shared_experts: int = 0, ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.vocab_size = config.vocab_size self.embed_dim = config.hidden_size @@ -1436,7 +1435,7 @@ class BailingMoeV3ForCausalLM(nn.Module): ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.tp_size = get_parallel().tp_size diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 7d61f6f7c..f4dc6bd83 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -25,7 +25,6 @@ from torch import nn from transformers import PretrainedConfig from sglang.kernels.ops.layernorm.fused_eh_norm import fused_eh_norm -from sglang.srt.distributed import get_pp_group from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.layers.attention.index_topk_share import IndexTopKShareState @@ -281,7 +280,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM): self.tp_size = get_parallel().tp_size self.quant_config = quant_config # if not set, model load will be broken in DeepseekV3ForCausalLM load_weights() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.determine_num_fused_shared_experts() nextn_quant_config = self._resolve_nextn_quant_config(config, quant_config) diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 778fcc16e..67243d3d7 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -52,7 +52,7 @@ from sglang.srt.configs.model_config import ( is_deepseek_dsa, is_glm_moe_dsa, ) -from sglang.srt.distributed import divide, get_pp_group +from sglang.srt.distributed import divide from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation @@ -2775,7 +2775,7 @@ class DeepseekV2Model(nn.Module): self.padding_id = config.pad_token_id self.vocab_size = config.vocab_size self.first_k_dense_replace = config.first_k_dense_replace - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank or (_is_npu and self.pp_group.is_last_rank): self.embed_tokens = VocabParallelEmbedding( @@ -3097,7 +3097,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): if quant_config is not None: quant_config.update_packed_modules_mapping(self.packed_modules_mapping) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.tp_size = get_parallel().tp_size self.quant_config = quant_config diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 40a261c6a..d46f0f90e 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -42,10 +42,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( ) 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_tp_group, -) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -2872,7 +2868,7 @@ class DeepseekV4DecoderLayer(nn.Module): # all-reduce input. Gated by is_allocation_symmetric() (mirrors the # TileLang path in _mhc_pre_impl / mhc_fused_post_pre). with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): y = hc_combine(x_flat, pre.squeeze(1), self.hc_mult, dtype) return y, post.squeeze(1), comb.squeeze(1), False @@ -3611,7 +3607,7 @@ class DeepseekV4DecoderLayer(nn.Module): ) elif _use_tp_moe_gather: hidden_states, local_hidden_states = ( - get_global_dp_buffer(get_tp_group()), + get_global_dp_buffer(get_parallel().tp_group), hidden_states, ) if _do_shared_local and local_hidden_states.shape[0] > 0: @@ -3661,7 +3657,7 @@ class DeepseekV4DecoderLayer(nn.Module): hidden_states = dsa_cp_reduce_scatter_hidden_states(hidden_states) elif _use_tp_moe_gather: hidden_states, global_hidden_states = ( - get_local_dp_buffer(get_tp_group()), + get_local_dp_buffer(get_parallel().tp_group), hidden_states, ) if should_use_dp_reduce_scatterv() or _use_reduce_scatterv: @@ -3669,7 +3665,7 @@ class DeepseekV4DecoderLayer(nn.Module): # each rank its own token slice, in one op. Correct because the # MoE-internal all_reduce was skipped (mlp_reduce_scatter above). # This is the symmetric inverse of the all_gatherv gather. - get_tp_group().reduce_scatterv( + get_parallel().tp_group.reduce_scatterv( global_hidden_states, output=hidden_states, sizes=get_dp_global_num_tokens(), @@ -3987,7 +3983,7 @@ class DeepseekV4Model(nn.Module): ) -> None: super().__init__() self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.hidden_size = config.hidden_size if self.pp_group.is_first_rank: embedding_quant_config = ( @@ -4376,7 +4372,7 @@ class DeepseekV4Model(nn.Module): # once across DP ranks, then populate each child's global_num_tokens + # global_dp_buffer_len so the gatherv/reduce_scatterv buffers size correctly. if get_moe_a2a_backend().is_none() and get_parallel().attn_dp_size > 1: - tp_group = get_tp_group() + tp_group = get_parallel().tp_group world = tp_group.world_size children = forward_batch.tbo_children local_lens = torch.tensor( @@ -4620,7 +4616,7 @@ class DeepseekV4ForCausalLM(nn.Module): if config.model_type == "deepseek_v41" and config.vision_n_layers > 0: if ( get_parallel().attn_cp_size != 1 - or get_pp_group().world_size != 1 + or get_parallel().pp_group.world_size != 1 or not get_moe_a2a_backend().is_none() ): raise ValueError( @@ -4636,7 +4632,7 @@ class DeepseekV4ForCausalLM(nn.Module): self.model = DeepseekV4Model( config, quant_config, prefix=add_prefix("model", prefix) ) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_last_rank: if self.pp_group.world_size == 1 and config.tie_word_embeddings: self.lm_head = self.model.embed_tokens @@ -5108,7 +5104,7 @@ class DeepseekV4ForCausalLM(nn.Module): compile_secs = time.perf_counter() - tic # Runs before init_memory_pool(); don't let transients skew pool sizing. torch.cuda.empty_cache() - get_tp_group().barrier() + get_parallel().tp_group.barrier() logger.info( "DeepSeek V4 MHC prewarm at load: compile %.1fs, rank sync +%.1fs", compile_secs, diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index f85401e1e..6694abb60 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -6,7 +6,6 @@ import torch.nn.functional as F from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.hardware_backend.npu.dsv4.dsv4_rope import prime_rope_cos_sin from sglang.srt.layers.attention.dsa.utils import ( dsa_use_prefill_cp, @@ -219,7 +218,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM): nn.Module.__init__(self) self.config = config self.tp_size = get_parallel().tp_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.quant_config = quant_config self.wo_a_fp8 = wo_a_fp8_gemm_enabled(quant_config) self.determine_num_fused_shared_experts() diff --git a/python/sglang/srt/models/dots3_common/modeling.py b/python/sglang/srt/models/dots3_common/modeling.py index ad9ce1de6..a49e3e7b4 100644 --- a/python/sglang/srt/models/dots3_common/modeling.py +++ b/python/sglang/srt/models/dots3_common/modeling.py @@ -45,8 +45,6 @@ from sglang.srt.batch_overlap.two_batch_overlap import ( ) from sglang.srt.configs.dots3 import Dots3Config from sglang.srt.distributed import ( - get_pp_group, - parallel_state, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -391,7 +389,7 @@ class Dots3MoE(nn.Module): ) self.deepep_dispatcher = MaybeTboDeepEPDispatcher( - group=parallel_state.get_tp_group().device_group, + group=get_parallel().tp_group.device_group, router_topk=self.top_k, permute_fusion=True, num_experts=self.num_experts, @@ -459,7 +457,7 @@ class Dots3MoE(nn.Module): final_hidden_states = self.experts(hidden_states, topk_output) current_stream.wait_stream(self.alt_stream) - with use_symmetric_memory(parallel_state.get_tp_group()) as sm: + with use_symmetric_memory(get_parallel().tp_group) as sm: final_hidden_states_out = torch.empty_like(final_hidden_states) torch.add(final_hidden_states, shared_output, out=final_hidden_states_out) @@ -491,7 +489,7 @@ class Dots3MoE(nn.Module): final_hidden_states = self.experts(hidden_states, topk_output) if shared_output is not None: - with use_symmetric_memory(parallel_state.get_tp_group()) as sm: + with use_symmetric_memory(get_parallel().tp_group) as sm: final_hidden_states_out = torch.empty_like(final_hidden_states) torch.add(final_hidden_states, shared_output, out=final_hidden_states_out) final_hidden_states = final_hidden_states_out @@ -1731,7 +1729,7 @@ class Dots3Model(nn.Module): super().__init__() _require_cuda() self.first_k_dense_replace = config.first_k_dense_replace - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -1865,7 +1863,7 @@ class Dots3LanguageModelForCausalLM(nn.Module): "g_proj", ] - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.tp_size = get_parallel().tp_size self.quant_config = quant_config @@ -2605,7 +2603,7 @@ class DotsNoteOmniThinkerForConditionalGeneration(nn.Module): ) self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group model_dir = Path(config._name_or_path) self.language_model = Dots3LanguageModelForCausalLM( config, diff --git a/python/sglang/srt/models/dots3_common/nextn.py b/python/sglang/srt/models/dots3_common/nextn.py index 2ae34a728..d97853dbe 100644 --- a/python/sglang/srt/models/dots3_common/nextn.py +++ b/python/sglang/srt/models/dots3_common/nextn.py @@ -7,7 +7,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm @@ -148,7 +147,7 @@ class Dots3NoteForCausalLMNextN(Dots3LanguageModelForCausalLM): self.config = config self.tp_size = get_parallel().tp_size self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.fuse_qkv_a_g_proj = True self.packed_modules_mapping = { "fused_qkv_a_g_proj_with_mqa": [ diff --git a/python/sglang/srt/models/dots_vlm.py b/python/sglang/srt/models/dots_vlm.py index 7165b4c7c..730d6f8a6 100644 --- a/python/sglang/srt/models/dots_vlm.py +++ b/python/sglang/srt/models/dots_vlm.py @@ -23,7 +23,6 @@ import torch from torch import nn from sglang.srt.configs.dots_vlm import DotsVLMConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.managers.mm_utils import ( MultiModalityDataPaddingPatternMultimodalTokens, @@ -33,6 +32,7 @@ from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInp from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM +from sglang.srt.runtime_context import get_parallel from .dots_vlm_vit import DotsVisionTransformer @@ -56,7 +56,7 @@ class DotsVLMForCausalLM(nn.Module): self.config = config self.image_token_id = config.im_span_id self.video_token_id = config.video_span_id - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if not config.encoder_only: self.language_model = DeepseekV2ForCausalLM( diff --git a/python/sglang/srt/models/ernie45_moe_vl.py b/python/sglang/srt/models/ernie45_moe_vl.py index 3cd69b1ee..f15071207 100644 --- a/python/sglang/srt/models/ernie45_moe_vl.py +++ b/python/sglang/srt/models/ernie45_moe_vl.py @@ -23,7 +23,6 @@ from torch import nn from transformers import PretrainedConfig from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.layers.dp_attention import is_dp_attention_enabled @@ -472,7 +471,7 @@ class Ernie4_5_VLMoeModel(nn.Module): ) -> None: super().__init__() self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( diff --git a/python/sglang/srt/models/exaone4.py b/python/sglang/srt/models/exaone4.py index 63ee1b342..44563fed1 100644 --- a/python/sglang/srt/models/exaone4.py +++ b/python/sglang/srt/models/exaone4.py @@ -5,7 +5,6 @@ import torch from torch import nn from transformers import Exaone4Config -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( @@ -310,7 +309,7 @@ class Exaone4Model(nn.Module): self.config = config self.quant_config = quant_config self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( config.vocab_size, @@ -423,7 +422,7 @@ class Exaone4ForCausalLM(nn.Module): prefix: str = "", ): super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config diff --git a/python/sglang/srt/models/exaone_moe.py b/python/sglang/srt/models/exaone_moe.py index 61c6cd255..1747e17c6 100755 --- a/python/sglang/srt/models/exaone_moe.py +++ b/python/sglang/srt/models/exaone_moe.py @@ -25,7 +25,6 @@ from torch import nn from transformers import PretrainedConfig from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -536,7 +535,7 @@ class ExaoneMoEModel(nn.Module): self.config = config self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -624,7 +623,7 @@ class ExaoneMoEForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config alt_stream = get_stream("alt") if _is_cuda else None @@ -856,7 +855,7 @@ class ExaoneMoEForCausalLM(nn.Module): ) def set_eagle3_layers_to_capture(self, layer_ids: Optional[list[int]] = None): - if not get_pp_group().is_last_rank: + if not get_parallel().pp_group.is_last_rank: return self.capture_aux_hidden_states = True diff --git a/python/sglang/srt/models/exaone_moe_mtp.py b/python/sglang/srt/models/exaone_moe_mtp.py index 10dea5461..92bce030c 100644 --- a/python/sglang/srt/models/exaone_moe_mtp.py +++ b/python/sglang/srt/models/exaone_moe_mtp.py @@ -23,7 +23,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -48,7 +47,7 @@ class ExaoneMoEForCausalLMMTP(ExaoneMoEForCausalLM): config.num_hidden_layers = 1 self.tp_size = get_parallel().tp_size self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.fc = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False) self.pre_fc_norm_embedding = RMSNorm( diff --git a/python/sglang/srt/models/falcon_h1.py b/python/sglang/srt/models/falcon_h1.py index e49f964ed..21e9c3244 100644 --- a/python/sglang/srt/models/falcon_h1.py +++ b/python/sglang/srt/models/falcon_h1.py @@ -5,7 +5,6 @@ import torch from torch import nn from sglang.srt.configs.falcon_h1 import FalconH1Config -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( HybridLinearAttnBackend, @@ -457,7 +456,7 @@ class FalconH1ForCausalLM(nn.Module): ) -> None: super().__init__() self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group assert self.pp_group.is_first_rank and self.pp_group.is_last_rank self.quant_config = quant_config self.model = FalconH1Model( diff --git a/python/sglang/srt/models/gemma4_causal.py b/python/sglang/srt/models/gemma4_causal.py index 26917a2de..01dfae96d 100644 --- a/python/sglang/srt/models/gemma4_causal.py +++ b/python/sglang/srt/models/gemma4_causal.py @@ -31,9 +31,6 @@ from sglang.kernels.ops.layernorm.gemma4_fused_ops import ( gemma_rmsnorm_residual_scalar, gemma_routing_post_topk, ) -from sglang.srt.distributed import ( - get_pp_group, -) from sglang.srt.layers.layernorm import Gemma4RMSNorm, RMSNorm from sglang.srt.layers.linear import ( QKVParallelLinear, @@ -782,7 +779,7 @@ class Gemma4TextModel(PreTrainedModel): self.quant_config = quant_config self.vocab_size = config.vocab_size self.padding_idx = getattr(config, "pad_token_id", None) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group # Token / per-layer embedding tables and the per-layer projection only # produce activations consumed at the model entry, so they live on the @@ -1090,7 +1087,7 @@ class Gemma4ForCausalLM(PreTrainedModel): prefix: str = "", ) -> None: super().__init__(config=config) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config diff --git a/python/sglang/srt/models/gemma4_mm.py b/python/sglang/srt/models/gemma4_mm.py index afa9a0c09..113f9e30d 100644 --- a/python/sglang/srt/models/gemma4_mm.py +++ b/python/sglang/srt/models/gemma4_mm.py @@ -28,7 +28,6 @@ from transformers import ( PreTrainedModel, ) -from sglang.srt.distributed import get_pp_group from sglang.srt.environ import envs from sglang.srt.layers.attention.triton_backend import TritonAttnBackend from sglang.srt.layers.layernorm import Gemma4RMSNorm @@ -65,6 +64,7 @@ from sglang.srt.models.gemma4_causal import ( pp_filter_load_weight, ) from sglang.srt.models.gemma4_vision import Gemma4VisionEncoder +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import add_prefix, cpu_has_amx_support, is_cpu from sglang.srt.utils.hf_transformers_utils import get_processor @@ -187,7 +187,7 @@ class Gemma4ForConditionalGeneration(PreTrainedModel): prefix: str = "", ) -> None: super().__init__(config=config) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config diff --git a/python/sglang/srt/models/gemma4_mtp.py b/python/sglang/srt/models/gemma4_mtp.py index 3978a1be1..1a1ec65d6 100644 --- a/python/sglang/srt/models/gemma4_mtp.py +++ b/python/sglang/srt/models/gemma4_mtp.py @@ -21,7 +21,6 @@ import torch from torch import nn from transformers import PretrainedConfig, PreTrainedModel -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.logits_processor import ( LogitsMetadata, @@ -32,6 +31,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.mem_cache.memory_pool import KVCache from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.gemma4_causal import Gemma4ForCausalLM, Gemma4TextModel +from sglang.srt.runtime_context import get_parallel from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPContext from sglang.srt.utils import add_prefix @@ -73,7 +73,7 @@ class Gemma4AssistantForCausalLM(Gemma4ForCausalLM): self.assistant_config = config self.config = text_config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.vocab_size = text_config.vocab_size self.hidden_size = text_config.hidden_size diff --git a/python/sglang/srt/models/gemma4_unified.py b/python/sglang/srt/models/gemma4_unified.py index 610f0c20c..1d00d0cbc 100644 --- a/python/sglang/srt/models/gemma4_unified.py +++ b/python/sglang/srt/models/gemma4_unified.py @@ -40,7 +40,6 @@ import torch from torch import nn from transformers import PreTrainedModel -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.layernorm import Gemma4RMSNorm from sglang.srt.layers.logits_processor import LogitsProcessor, LogitsProcessorOutput from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -53,6 +52,7 @@ from sglang.srt.managers.schedule_batch import ( from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.gemma4_causal import Gemma4TextModel, pp_filter_load_weight from sglang.srt.models.gemma4_mm import Gemma4ForConditionalGeneration +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -144,7 +144,7 @@ class Gemma4UnifiedForConditionalGeneration(Gemma4ForConditionalGeneration): # Skip Gemma4ForConditionalGeneration.__init__ (it builds the SigLIP / # conformer towers we do not have) and initialise the HF base directly. PreTrainedModel.__init__(self, config=config) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config diff --git a/python/sglang/srt/models/glm4.py b/python/sglang/srt/models/glm4.py index 00a5057d0..f1feda9e9 100644 --- a/python/sglang/srt/models/glm4.py +++ b/python/sglang/srt/models/glm4.py @@ -23,9 +23,6 @@ from typing import Any, Dict, Iterable, Optional, Tuple, Union import torch from torch import nn -from sglang.srt.distributed import ( - get_pp_group, -) from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm @@ -298,7 +295,7 @@ class Glm4Model(nn.Module): self.config = config self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -420,7 +417,7 @@ class Glm4ForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = Glm4Model( diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index 91584d233..b24f069db 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -28,9 +28,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz from sglang.srt.batch_overlap.single_batch_overlap import SboFlags from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo from sglang.srt.distributed import ( - get_pp_group, get_pp_indices, - parallel_state, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -646,7 +644,7 @@ class Glm4MoeSparseMoeBlock(nn.Module): final_hidden_states *= self.routed_scaling_factor if shared_output is not None: with use_symmetric_memory( - parallel_state.get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): final_hidden_states_out = torch.empty_like(final_hidden_states) torch.add(final_hidden_states, shared_output, out=final_hidden_states_out) @@ -1059,7 +1057,7 @@ class Glm4MoeModel(nn.Module): prefix: str = "", ): super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.vocab_size = config.vocab_size self.first_k_dense_replace = config.first_k_dense_replace @@ -1185,7 +1183,7 @@ class Glm4MoeForCausalLM(nn.Module): prefix: str = "", ) -> None: nn.Module.__init__(self) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.tp_size = get_parallel().tp_size self.quant_config = quant_config diff --git a/python/sglang/srt/models/glm4_moe_lite.py b/python/sglang/srt/models/glm4_moe_lite.py index 4d74c41dc..e30c2387a 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -26,8 +26,6 @@ from transformers import PretrainedConfig from sglang.srt.batch_overlap.single_batch_overlap import SboFlags from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo from sglang.srt.distributed import ( - get_pp_group, - parallel_state, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -365,7 +363,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module): final_hidden_states *= self.routed_scaling_factor if shared_output is not None: with use_symmetric_memory( - parallel_state.get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): final_hidden_states_out = torch.empty_like(final_hidden_states) torch.add(final_hidden_states, shared_output, out=final_hidden_states_out) @@ -760,7 +758,7 @@ class Glm4MoeLiteModel(nn.Module): self.padding_id = config.pad_token_id self.vocab_size = config.vocab_size self.first_k_dense_replace = config.first_k_dense_replace - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -892,7 +890,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): self.config = config self.tp_size = get_parallel().tp_size self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.determine_num_fused_shared_experts() self.model = Glm4MoeLiteModel( config, quant_config, prefix=add_prefix("model", prefix) diff --git a/python/sglang/srt/models/glm4v.py b/python/sglang/srt/models/glm4v.py index 43a74f4ad..f4df5a2a2 100644 --- a/python/sglang/srt/models/glm4v.py +++ b/python/sglang/srt/models/glm4v.py @@ -27,7 +27,6 @@ import torch.nn.functional as F from einops import rearrange from transformers.models.glm4v.configuration_glm4v import Glm4vConfig, Glm4vVisionConfig -from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.attention import vision_utils from sglang.srt.layers.attention.vision import ( @@ -556,7 +555,7 @@ class Glm4vForConditionalGeneration(nn.Module): ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.use_data_parallel = get_mm().mm_enable_dp_encoder vision_utils.update_vit_attn_dummy_heads_config(self.config) diff --git a/python/sglang/srt/models/glm4v_moe.py b/python/sglang/srt/models/glm4v_moe.py index 5f9d2ed51..6acc2789a 100644 --- a/python/sglang/srt/models/glm4v_moe.py +++ b/python/sglang/srt/models/glm4v_moe.py @@ -6,7 +6,6 @@ import torch import torch.nn as nn from transformers.models.glm4v_moe.configuration_glm4v_moe import Glm4vMoeConfig -from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.layers.attention import vision_utils from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.moe import get_moe_a2a_backend @@ -40,7 +39,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): ) -> None: nn.Module.__init__(self) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.use_data_parallel = get_mm().mm_enable_dp_encoder vision_utils.update_vit_attn_dummy_heads_config(self.config) diff --git a/python/sglang/srt/models/glm5_next.py b/python/sglang/srt/models/glm5_next.py index 9a5f4d4d2..47d2b20c0 100644 --- a/python/sglang/srt/models/glm5_next.py +++ b/python/sglang/srt/models/glm5_next.py @@ -16,7 +16,6 @@ from sglang.srt.batch_overlap.two_batch_overlap import ( ) from sglang.srt.configs.glm5_next import Glm5NextConfig, Glm5NextTextConfig from sglang.srt.configs.model_config import is_deepseek_dsa -from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.distributed.utils import divide from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import ( @@ -861,7 +860,7 @@ class Glm5NextModel(nn.Module): self.padding_id = config.pad_token_id self.vocab_size = config.vocab_size self.first_k_dense_replace = config.first_k_dense_replace - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -1123,7 +1122,7 @@ class Glm5NextForConditionalGeneration(nn.Module): and getattr(text_config, "q_lora_rank", None) is not None ) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = text_config self.tp_size = get_parallel().tp_size self.quant_config = quant_config diff --git a/python/sglang/srt/models/glm_ocr.py b/python/sglang/srt/models/glm_ocr.py index 15fb7d6c9..e26385d6e 100644 --- a/python/sglang/srt/models/glm_ocr.py +++ b/python/sglang/srt/models/glm_ocr.py @@ -29,7 +29,6 @@ from transformers.models.glm_ocr.configuration_glm_ocr import ( GlmOcrVisionConfig, ) -from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.layers.attention import vision_utils from sglang.srt.layers.attention.vision import ( VisionAttention, @@ -53,7 +52,7 @@ from sglang.srt.models.glm4v import ( Glm4vVisionModel, Glm4vVisionPatchEmbed, ) -from sglang.srt.runtime_context import get_mm +from sglang.srt.runtime_context import get_mm, get_parallel from sglang.srt.utils import add_prefix from sglang.srt.utils.hf_transformers_utils import get_processor @@ -285,7 +284,7 @@ class GlmOcrForConditionalGeneration(Glm4vForConditionalGeneration): ) -> None: super().__init__(config, quant_config, prefix) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.use_data_parallel = get_mm().mm_enable_dp_encoder self.visual = GlmOcrVisionModel( diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index 301d7e659..22283080e 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -28,7 +28,6 @@ from transformers import PretrainedConfig from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -680,7 +679,7 @@ class GptOssModel(nn.Module): super().__init__() self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if _is_npu: config.hidden_act = "npu_swiglu_oai" @@ -790,7 +789,7 @@ class GptOssForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = GptOssModel( diff --git a/python/sglang/srt/models/granitemoehybrid.py b/python/sglang/srt/models/granitemoehybrid.py index b6345e7bf..01021ebba 100644 --- a/python/sglang/srt/models/granitemoehybrid.py +++ b/python/sglang/srt/models/granitemoehybrid.py @@ -4,7 +4,6 @@ import torch from torch import nn from sglang.srt.configs.granitemoehybrid import GraniteMoeHybridConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( HybridLinearAttnBackend, Mamba2AttnBackend, @@ -327,7 +326,7 @@ class GraniteMoeHybridModel(nn.Module): self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -443,7 +442,7 @@ class GraniteMoeHybridForCausalLM( super().__init__() self.capture_aux_hidden_states = False - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.quant_config = quant_config self.config = config diff --git a/python/sglang/srt/models/hunyuan_v4.py b/python/sglang/srt/models/hunyuan_v4.py index c7971dd89..8f59b56ad 100644 --- a/python/sglang/srt/models/hunyuan_v4.py +++ b/python/sglang/srt/models/hunyuan_v4.py @@ -6,7 +6,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.attention.index_topk_share import IndexTopKShareState from sglang.srt.layers.communicator import AttentionInputs, get_attn_tp_context from sglang.srt.layers.layernorm import RMSNorm @@ -614,7 +613,7 @@ class HYV4DecoderLayer(nn.Module): class HYV4Model(nn.Module): def __init__(self, config, quant_config=None, prefix=""): super().__init__() - if get_pp_group().world_size != 1: + if get_parallel().pp_group.world_size != 1: raise ValueError("HYV4 pipeline parallelism is not supported") self.config = config self.start_layer = 0 @@ -675,7 +674,7 @@ class HYV4ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): super().__init__() self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.model = HYV4Model(config, quant_config, f"{prefix}.model") self.num_fused_shared_experts = max( ( diff --git a/python/sglang/srt/models/hunyuan_v4_nextn.py b/python/sglang/srt/models/hunyuan_v4_nextn.py index fce1b9229..60790f4ee 100644 --- a/python/sglang/srt/models/hunyuan_v4_nextn.py +++ b/python/sglang/srt/models/hunyuan_v4_nextn.py @@ -4,7 +4,6 @@ from typing import Iterable, Tuple import torch from torch import nn -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.attention.index_topk_share import IndexTopKShareState from sglang.srt.layers.communicator import AttentionInputs, get_attn_tp_context from sglang.srt.layers.layernorm import RMSNorm @@ -170,7 +169,7 @@ class HYV4ForCausalLMNextN(nn.Module, DeepseekV2WeightLoaderMixin): super().__init__() self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group nextn_quant_config = _mtp_quant_config(quant_config) self.model = HYV4ModelNextN( config, nextn_quant_config, prefix=f"{prefix}.model" diff --git a/python/sglang/srt/models/interns2_mobius.py b/python/sglang/srt/models/interns2_mobius.py index 2e5aa8468..5705c7cc0 100644 --- a/python/sglang/srt/models/interns2_mobius.py +++ b/python/sglang/srt/models/interns2_mobius.py @@ -10,7 +10,7 @@ from sglang.srt.configs.interns2_mobius import ( InternS2MobiusConfig, InternS2MobiusTextConfig, ) -from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce +from sglang.srt.distributed import tensor_model_parallel_all_reduce from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import GemmaRMSNorm @@ -695,7 +695,7 @@ class InternS2MobiusForCausalLM(Qwen3_5ForCausalLM): nn.Module.__init__(self) self.config = config self.hidden_size = config.hidden_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.world_size != 1: raise ValueError( "Intern-S2-Mobius baseline does not support pipeline parallelism" diff --git a/python/sglang/srt/models/kimi_k25_eagle3.py b/python/sglang/srt/models/kimi_k25_eagle3.py index dff7478fa..7d6913d56 100644 --- a/python/sglang/srt/models/kimi_k25_eagle3.py +++ b/python/sglang/srt/models/kimi_k25_eagle3.py @@ -24,7 +24,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.distributed.device_communicators import triton_symm_mem_ag from sglang.srt.layers.communicator import AttentionInputs, get_attn_tp_context from sglang.srt.layers.layernorm import RMSNorm @@ -39,6 +38,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import BumpAllocator, add_prefix logger = logging.getLogger(__name__) @@ -336,7 +336,7 @@ class Eagle3DeepseekV2ForCausalLM(nn.Module): ) quant_config = None self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.model = Eagle3MLAModel( config, quant_config=quant_config, prefix=add_prefix("model", prefix) diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index d9ccf744a..578f443ee 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -21,9 +21,7 @@ from sglang.srt.configs.kimi_k3 import KimiK3Config from sglang.srt.configs.kimi_linear import KimiLinearConfig from sglang.srt.distributed import ( divide, - get_pp_group, get_shared_experts_tp_group, - get_tp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -257,7 +255,7 @@ def _dp_local_buffer_group(): CommunicateSummableTensorPairFn._scatter_hidden_states).""" parallel = get_parallel() if parallel.tp_size == parallel.attn_dp_size: - return get_tp_group() + return get_parallel().tp_group return parallel.attn_tp_group @@ -361,7 +359,7 @@ class KimiK3MLP(nn.Module): ) if use_dp: local_hidden_states = hidden_states - hidden_states = get_global_dp_buffer(get_tp_group()) + hidden_states = get_global_dp_buffer(get_parallel().tp_group) dp_gather_replicate(hidden_states, local_hidden_states, forward_batch) gate_up, _ = self.gate_up_proj(hidden_states) hidden_states = self.act_fn(gate_up) @@ -796,12 +794,12 @@ class KimiK3MoE(nn.Module): import deep_gemm from sglang.kernels.ops.attention.dsv4 import mega_moe_pre_dispatch - from sglang.srt.distributed.parallel_state import get_moe_ep_group from sglang.srt.environ import envs from sglang.srt.layers.moe.mega_moe import ( _configure_mega_moe_deep_gemm_num_sms, _get_mega_moe_symm_buffer, ) + from sglang.srt.runtime_context import get_parallel # In SP-MoE mode (KimiK3DecoderLayer reduce-scatters the o_proj # output) the incoming rows are already this rank's token shard, so @@ -819,7 +817,7 @@ class KimiK3MoE(nn.Module): f"the env var to cover the per-rank rows" ) buf = _get_mega_moe_symm_buffer( - get_moe_ep_group().device_group, + get_parallel().moe_ep_group.device_group, num_experts=self.experts.num_experts, num_max_tokens_per_rank=num_max_tokens_per_rank, num_topk=self._mega_top_k, @@ -1389,7 +1387,7 @@ class KimiK3MoE(nn.Module): ).view(-1) else: with use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ): buf = hidden_states.new_empty(latent_numel + num_tokens * hidden_size) @@ -1505,7 +1503,7 @@ class KimiK3MoE(nn.Module): use_dp = self._dp_attention and forward_batch is not None and not self._ep_a2a if use_dp: local_hidden_states = hidden_states - hidden_states = get_global_dp_buffer(get_tp_group()) + hidden_states = get_global_dp_buffer(get_parallel().tp_group) dp_gather_replicate(hidden_states, local_hidden_states, forward_batch) dp_prefix_sum, prefix_sum = prefix_sum, None if hidden_states.shape[0] > 0 and self._eligible_for_fused_front: @@ -2919,7 +2917,7 @@ class KimiK3LinearModel(nn.Module): ): super().__init__() self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.dspark_layers_to_capture: Optional[list[int]] = None self._dp_attention = is_dp_attention_enabled() self._trim_padded_attn = require_mlp_sync() @@ -2990,7 +2988,7 @@ class KimiK3LinearModel(nn.Module): inputs_embeds: torch.Tensor | None = None, pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> torch.Tensor: - if get_pp_group().is_first_rank: + if get_parallel().pp_group.is_first_rank: if inputs_embeds is not None: hidden_states = inputs_embeds else: @@ -3188,7 +3186,7 @@ class KimiK3LinearForCausalLM(nn.Module): self.model = KimiK3LinearModel( config, quant_config, prefix=maybe_prefix(prefix, "model") ) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_last_rank: self.lm_head = ParallelLMHead( config.vocab_size, diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index 1b010b42f..87fdb0ae4 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -12,7 +12,6 @@ from sglang.kernels.ops.attention.fla.fused_norm_gate import FusedRMSNormGated from sglang.srt.configs.kimi_linear import KimiLinearConfig from sglang.srt.distributed import ( divide, - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -653,7 +652,7 @@ class KimiLinearModel(nn.Module): self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.dspark_layers_to_capture: Optional[list[int]] = None if self.pp_group.is_first_rank: @@ -699,7 +698,7 @@ class KimiLinearModel(nn.Module): inputs_embeds: torch.Tensor | None = None, pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> torch.Tensor: - if get_pp_group().is_first_rank: + if get_parallel().pp_group.is_first_rank: if inputs_embeds is not None: hidden_states = inputs_embeds else: @@ -771,7 +770,7 @@ class KimiLinearForCausalLM(nn.Module): ) self.start_layer = self.model.start_layer self.end_layer = self.model.end_layer - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_last_rank: self.lm_head = ParallelLMHead( self.config.vocab_size, diff --git a/python/sglang/srt/models/laguna.py b/python/sglang/srt/models/laguna.py index 9aa8d4591..b5c763d5a 100644 --- a/python/sglang/srt/models/laguna.py +++ b/python/sglang/srt/models/laguna.py @@ -18,7 +18,6 @@ from torch import nn from sglang.srt.configs.laguna import LagunaConfig, normalize_gating from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.environ import envs @@ -534,7 +533,7 @@ class LagunaModel(nn.Module): self.config = config self.padding_idx = getattr(config, "pad_token_id", None) self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -656,7 +655,7 @@ class LagunaForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.model = LagunaModel( config, quant_config=quant_config, prefix=add_prefix("model", prefix) diff --git a/python/sglang/srt/models/lfm2.py b/python/sglang/srt/models/lfm2.py index 1be184d47..d55101498 100644 --- a/python/sglang/srt/models/lfm2.py +++ b/python/sglang/srt/models/lfm2.py @@ -22,7 +22,6 @@ from sglang.kernels.ops.mamba.causal_conv1d_triton import ( causal_conv1d_update as causal_conv1d_update_triton, ) from sglang.srt.configs.lfm2 import Lfm2Config -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.attention.mamba.causal_conv1d import ( causal_conv1d_fn, causal_conv1d_update, @@ -688,7 +687,7 @@ class Lfm2ForCausalLM(nn.Module): ) -> None: super().__init__() self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group assert self.pp_group.is_first_rank and self.pp_group.is_last_rank self.quant_config = quant_config diff --git a/python/sglang/srt/models/lfm2_moe.py b/python/sglang/srt/models/lfm2_moe.py index cea383e87..edfe57830 100644 --- a/python/sglang/srt/models/lfm2_moe.py +++ b/python/sglang/srt/models/lfm2_moe.py @@ -23,7 +23,6 @@ from sglang.kernels.ops.mamba.lfm_short_conv import ( fused_lfm_short_conv_prefill, ) from sglang.srt.configs.lfm2_moe import Lfm2MoeConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.attention.mamba.causal_conv1d import ( causal_conv1d_fn, @@ -598,7 +597,7 @@ class Lfm2MoeForCausalLM(nn.Module): ) -> None: super().__init__() self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group assert self.pp_group.is_first_rank and self.pp_group.is_last_rank self.quant_config = quant_config diff --git a/python/sglang/srt/models/llada2.py b/python/sglang/srt/models/llada2.py index 1e4e804c4..01bd9c3fe 100644 --- a/python/sglang/srt/models/llada2.py +++ b/python/sglang/srt/models/llada2.py @@ -28,8 +28,6 @@ from torch import nn from transformers import PretrainedConfig from sglang.srt.distributed import ( - get_pp_group, - parallel_state, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -309,7 +307,7 @@ class LLaDA2MoeSparseMoeBlock(nn.Module): self.ep_size = get_parallel().tp_size self.deepep_dispatcher = DeepEPDispatcher( - group=parallel_state.get_tp_group().device_group, + group=get_parallel().tp_group.device_group, router_topk=self.top_k, permute_fusion=True, num_experts=self.num_experts, @@ -717,7 +715,7 @@ class LLaDA2MoeModel(nn.Module): prefix: str = "", ): super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.vocab_size = config.vocab_size self.embed_dim = config.hidden_size @@ -804,7 +802,7 @@ class LLaDA2MoeModelLM(nn.Module): prefix: str = "", ): super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config alt_stream = get_stream("alt") if _is_cuda else None diff --git a/python/sglang/srt/models/llama.py b/python/sglang/srt/models/llama.py index ad6aeb8bc..0da5cb20a 100644 --- a/python/sglang/srt/models/llama.py +++ b/python/sglang/srt/models/llama.py @@ -26,7 +26,6 @@ from torch import nn from transformers import LlamaConfig from sglang.srt.distributed import ( - get_pp_group, get_pp_indices, ) from sglang.srt.layers.activation import SiluAndMul @@ -379,7 +378,7 @@ class LlamaModel(nn.Module): self.config = config self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( config.vocab_size, @@ -521,7 +520,7 @@ class LlamaForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = self._init_model(config, quant_config, add_prefix("model", prefix)) diff --git a/python/sglang/srt/models/llama_eagle.py b/python/sglang/srt/models/llama_eagle.py index 5b7b95d47..df64c3b51 100644 --- a/python/sglang/srt/models/llama_eagle.py +++ b/python/sglang/srt/models/llama_eagle.py @@ -25,7 +25,6 @@ import torch from torch import nn from transformers import LlamaConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.vocab_parallel_embedding import ( @@ -34,6 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.models.llama import LlamaDecoderLayer, LlamaForCausalLM +from sglang.srt.runtime_context import get_parallel class LlamaDecoderLayer(LlamaDecoderLayer): @@ -120,7 +120,7 @@ class LlamaForCausalLMEagle(LlamaForCausalLM): nn.Module.__init__(self) self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.model = LlamaModel( config, quant_config=quant_config, prefix=add_prefix("model", prefix) ) diff --git a/python/sglang/srt/models/llama_eagle3.py b/python/sglang/srt/models/llama_eagle3.py index 89f2ac4d9..9f5ddceef 100644 --- a/python/sglang/srt/models/llama_eagle3.py +++ b/python/sglang/srt/models/llama_eagle3.py @@ -27,7 +27,6 @@ import torch from torch import nn from transformers import LlamaConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import QKVParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor @@ -39,6 +38,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.llama import LlamaDecoderLayer, LlamaForCausalLM, LlamaMLP +from sglang.srt.runtime_context import get_parallel class LlamaDecoderLayer(LlamaDecoderLayer): @@ -270,7 +270,7 @@ class LlamaForCausalLMEagle3(LlamaForCausalLM): nn.Module.__init__(self) self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group # Cache draft SWA size from server args once; consumed both by the post-init # attention patch below and by `get_attention_sliding_window_size` later. diff --git a/python/sglang/srt/models/mellum.py b/python/sglang/srt/models/mellum.py index 625bef1cf..be64381f6 100644 --- a/python/sglang/srt/models/mellum.py +++ b/python/sglang/srt/models/mellum.py @@ -25,7 +25,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import QKVParallelLinear, RowParallelLinear @@ -494,7 +493,7 @@ class MellumForCausalLM(Qwen3MoeForCausalLM): from sglang.srt.layers.logits_processor import LogitsProcessor - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group cfg = cast(Any, config) self.config = cfg self.quant_config = quant_config diff --git a/python/sglang/srt/models/mimo_v2.py b/python/sglang/srt/models/mimo_v2.py index eb5e042ae..c09e54b02 100644 --- a/python/sglang/srt/models/mimo_v2.py +++ b/python/sglang/srt/models/mimo_v2.py @@ -24,7 +24,6 @@ from torch import nn from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo from sglang.srt.configs.model_config import get_mimo_v2_fused_qkv_expected_tp_size from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -1003,7 +1002,7 @@ class MiMoV2Model(nn.Module): self.config = config self.padding_idx = getattr(config, "pad_token_id", None) self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.layers_to_capture = [] if self.pp_group.is_first_rank: @@ -1197,7 +1196,7 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self._encoder_processor = None # lazy-created in preprocess_mm_for_encoder diff --git a/python/sglang/srt/models/minimax_m2.py b/python/sglang/srt/models/minimax_m2.py index 54037429c..b0d470d6d 100644 --- a/python/sglang/srt/models/minimax_m2.py +++ b/python/sglang/srt/models/minimax_m2.py @@ -29,7 +29,6 @@ from transformers import PretrainedConfig from sglang.kernels.kernel_api_logging import debug_kernel_api from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -1123,7 +1122,7 @@ class MiniMaxM2Model(nn.Module): self.padding_idx = getattr(config, "pad_token_id", 0) self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.embed_tokens = VocabParallelEmbedding( config.vocab_size, @@ -1252,7 +1251,7 @@ class MiniMaxM2ForCausalLM(nn.Module): config, quant_config, prefix=add_prefix("model", prefix) ) - if get_pp_group().is_last_rank: + if get_parallel().pp_group.is_last_rank: self.lm_head = ParallelLMHead( config.vocab_size, config.hidden_size, @@ -1263,7 +1262,7 @@ class MiniMaxM2ForCausalLM(nn.Module): self.lm_head = PPMissingLayer() self.logits_processor = LogitsProcessor(config) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group # For EAGLE3 self.capture_aux_hidden_states = False @@ -1272,7 +1271,7 @@ class MiniMaxM2ForCausalLM(nn.Module): return self.model.get_input_embeddings(input_ids) def set_eagle3_layers_to_capture(self, layer_ids: Optional[list[int]] = None): - if not get_pp_group().is_last_rank: + if not get_parallel().pp_group.is_last_rank: return self.capture_aux_hidden_states = True diff --git a/python/sglang/srt/models/minimax_m3.py b/python/sglang/srt/models/minimax_m3.py index 3cd2a0931..eef7c270a 100644 --- a/python/sglang/srt/models/minimax_m3.py +++ b/python/sglang/srt/models/minimax_m3.py @@ -29,7 +29,6 @@ from sglang.srt.configs.model_config import ( get_minimax_sparse_layer_ids, ) from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.environ import envs @@ -1441,7 +1440,7 @@ class MiniMaxM3Model(nn.Module): self.padding_idx = getattr(config, "pad_token_id", 0) self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.use_gemma_norm = getattr(config, "use_gemma_norm", False) if self.pp_group.is_first_rank: @@ -1574,7 +1573,7 @@ class MiniMaxM3SparseForCausalLM(nn.Module): self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.num_fused_shared_experts = 0 self.determine_num_fused_shared_experts() diff --git a/python/sglang/srt/models/minimax_m3_vl.py b/python/sglang/srt/models/minimax_m3_vl.py index 22249126d..96b1c6ab3 100644 --- a/python/sglang/srt/models/minimax_m3_vl.py +++ b/python/sglang/srt/models/minimax_m3_vl.py @@ -6,9 +6,6 @@ from typing import Iterable, List, Optional, Tuple import torch import torch.nn as nn -from sglang.srt.distributed import ( - get_pp_group, -) from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.moe.utils import ( get_moe_a2a_backend, @@ -86,7 +83,7 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module): super().__init__() self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.use_data_parallel = get_mm().mm_enable_dp_encoder diff --git a/python/sglang/srt/models/mistral_eagle.py b/python/sglang/srt/models/mistral_eagle.py index 1e22022f9..bdb8253cc 100644 --- a/python/sglang/srt/models/mistral_eagle.py +++ b/python/sglang/srt/models/mistral_eagle.py @@ -40,7 +40,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import RowParallelLinear from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -48,6 +47,7 @@ from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.models.llama import LlamaDecoderLayer, LlamaForCausalLM from sglang.srt.models.llama_eagle import LlamaForCausalLMEagle +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -65,10 +65,10 @@ class MistralEagleModel(nn.Module): super().__init__() self.config = config self.vocab_size = config.vocab_size - assert get_pp_group().world_size == 1, ( + assert get_parallel().pp_group.world_size == 1, ( "MistralForCausalLMEagle currently does not support pipeline parallelism" ) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.embed_tokens = VocabParallelEmbedding( config.vocab_size, config.hidden_size, diff --git a/python/sglang/srt/models/mixtral.py b/python/sglang/srt/models/mixtral.py index 5e34952cd..4b2b5c908 100644 --- a/python/sglang/srt/models/mixtral.py +++ b/python/sglang/srt/models/mixtral.py @@ -26,7 +26,6 @@ from torch import nn from transformers import MixtralConfig from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.layers.layernorm import RMSNorm @@ -270,7 +269,7 @@ class MixtralModel(nn.Module): super().__init__() self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -343,7 +342,7 @@ class MixtralForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = MixtralModel( diff --git a/python/sglang/srt/models/mllama.py b/python/sglang/srt/models/mllama.py index b76ff7b59..c03c59b1a 100644 --- a/python/sglang/srt/models/mllama.py +++ b/python/sglang/srt/models/mllama.py @@ -20,7 +20,6 @@ from transformers.models.mllama.modeling_mllama import ( _prepare_aspect_ratio_attention_mask, ) -import sglang.srt.distributed.parallel_state as ps from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.layers.layernorm import RMSNorm @@ -392,7 +391,7 @@ class MllamaVisionModel(nn.Module): pixel_values.to(self.layernorm_pre.weight.dtype) ) hidden_state = patch_embeds - hidden_state = ps.get_tp_group().all_gather(hidden_state) + hidden_state = get_parallel().tp_group.all_gather(hidden_state) # tile embeddings _, num_patches, dim = hidden_state.shape diff --git a/python/sglang/srt/models/nanbeige.py b/python/sglang/srt/models/nanbeige.py index be9d6d5de..b5fd34b16 100644 --- a/python/sglang/srt/models/nanbeige.py +++ b/python/sglang/srt/models/nanbeige.py @@ -5,7 +5,6 @@ import torch from torch import nn from sglang.srt.configs import NanbeigeConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm @@ -262,7 +261,7 @@ class NanbeigeModel(nn.Module): super().__init__() self.config = config self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group pp_size = self.pp_group.world_size assert pp_size == 1, ( "The NanbeigeModel only supports a pipeline parallelism (PP) value of 1." @@ -389,7 +388,7 @@ class NanbeigeForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = NanbeigeModel( diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index fd85df6bc..1abd6da50 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -26,8 +26,6 @@ from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.configs import NemotronHConfig from sglang.srt.configs.nemotron_h import ATTENTION, MAMBA, MLP, MOE from sglang.srt.distributed import ( - get_moe_ep_group, - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.layers.activation import ReLU2 @@ -189,7 +187,7 @@ class NemotronHMoE(nn.Module): self.routed_scaling_factor = config.routed_scaling_factor self.device_module = torch.get_device_module() - self.ep_group = get_moe_ep_group().device_group + self.ep_group = get_parallel().moe_ep_group.device_group self.ep_rank = self.ep_group.rank() self.ep_size = self.ep_group.size() self.n_routed_experts = config.n_routed_experts @@ -836,7 +834,7 @@ class NemotronHModel(nn.Module): ) self.vocab_size = config.vocab_size + lora_vocab self.org_vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -973,7 +971,7 @@ class NemotronHForCausalLM(nn.Module): self.model = self._init_model( config=config, quant_config=quant_config, prefix=prefix ) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_last_rank: if self.pp_group.world_size == 1 and self.config.tie_word_embeddings: diff --git a/python/sglang/srt/models/nemotron_h_mtp.py b/python/sglang/srt/models/nemotron_h_mtp.py index d64caae1f..10065f48e 100644 --- a/python/sglang/srt/models/nemotron_h_mtp.py +++ b/python/sglang/srt/models/nemotron_h_mtp.py @@ -18,7 +18,6 @@ import torch from torch import nn 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, is_dp_attention_enabled, @@ -341,7 +340,7 @@ class NemotronHForCausalLMMTP(NemotronHForCausalLM): self.quant_config = quant_config self._owns_lm_head = False # Required for parent's load_weights - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group # Override config for MTP pattern (which has no Mamba layers) config.num_hidden_layers = len(config.mtp_hybrid_override_pattern) diff --git a/python/sglang/srt/models/nemotron_nas.py b/python/sglang/srt/models/nemotron_nas.py index f45e16591..aa4e86eb2 100644 --- a/python/sglang/srt/models/nemotron_nas.py +++ b/python/sglang/srt/models/nemotron_nas.py @@ -23,7 +23,6 @@ import torch from torch import nn from transformers import LlamaConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.logits_processor import LogitsProcessor, LogitsProcessorOutput from sglang.srt.layers.pooler import Pooler, PoolingType @@ -40,6 +39,7 @@ from sglang.srt.model_loader.weight_utils import ( maybe_remap_kv_scale_name, ) from sglang.srt.models.llama import LlamaAttention, LlamaMLP +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import add_prefix, make_layers from sglang.utils import logger @@ -179,7 +179,7 @@ class DeciModel(nn.Module): else 0 ) vocab_size = config.vocab_size + lora_vocab - if get_pp_group().is_first_rank: + if get_parallel().pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( vocab_size, config.hidden_size, @@ -200,11 +200,11 @@ class DeciModel(nn.Module): self.layers, self.start_layer, self.end_layer = make_layers( config.num_hidden_layers, get_layer, - pp_rank=get_pp_group().rank_in_group, - pp_size=get_pp_group().world_size, + pp_rank=get_parallel().pp_group.rank_in_group, + pp_size=get_parallel().pp_group.world_size, prefix=add_prefix("layers", prefix), ) - if get_pp_group().is_last_rank: + if get_parallel().pp_group.is_last_rank: self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) else: self.norm = PPMissingLayer(return_tuple=True) @@ -220,7 +220,7 @@ class DeciModel(nn.Module): inputs_embeds: Optional[torch.Tensor] = None, pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> Union[torch.Tensor, PPProxyTensors]: - if get_pp_group().is_first_rank: + if get_parallel().pp_group.is_first_rank: if inputs_embeds is not None: hidden_states = inputs_embeds else: @@ -244,7 +244,7 @@ class DeciModel(nn.Module): positions, hidden_states, forward_batch, residual ) - if not get_pp_group().is_last_rank: + if not get_parallel().pp_group.is_last_rank: return PPProxyTensors( {"hidden_states": hidden_states, "residual": residual} ) @@ -360,7 +360,7 @@ class DeciLMForCausalLM(nn.Module): inputs_embeds, pp_proxy_tensors=pp_proxy_tensors, ) - if get_pp_group().is_last_rank: + if get_parallel().pp_group.is_last_rank: if not get_embedding: return self.logits_processor( input_ids, hidden_states, self.lm_head, forward_batch diff --git a/python/sglang/srt/models/opt.py b/python/sglang/srt/models/opt.py index 3ac98bfd2..ee09deec8 100644 --- a/python/sglang/srt/models/opt.py +++ b/python/sglang/srt/models/opt.py @@ -22,9 +22,6 @@ import torch from torch import nn from transformers import OPTConfig -from sglang.srt.distributed import ( - get_pp_group, -) from sglang.srt.layers.linear import ( ColumnParallelLinear, QKVParallelLinear, @@ -229,7 +226,7 @@ class OPTDecoder(nn.Module): self.max_target_positions = config.max_position_embeddings self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.embed_tokens = VocabParallelEmbedding( config.vocab_size, @@ -333,7 +330,7 @@ class OPTModel(nn.Module): self.config = config self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.decoder = OPTDecoder( config=config, @@ -408,7 +405,7 @@ class OPTForCausalLM(nn.Module): self.logits_processor = LogitsProcessor(config) self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) self.capture_aux_hidden_states = False - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.stacked_params_mapping = [ # (param_name, shard_name, shard_id) (".qkv_proj", ".q_proj", "q"), diff --git a/python/sglang/srt/models/orion.py b/python/sglang/srt/models/orion.py index 9d34fa921..95df14375 100644 --- a/python/sglang/srt/models/orion.py +++ b/python/sglang/srt/models/orion.py @@ -15,7 +15,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.linear import ( MergedColumnParallelLinear, @@ -223,7 +222,7 @@ class OrionModel(nn.Module): ): super().__init__() self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -285,7 +284,7 @@ class OrionForCausalLM(nn.Module): super().__init__() self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.model = OrionModel( config=config, quant_config=quant_config, prefix=add_prefix("model", prefix) ) diff --git a/python/sglang/srt/models/persimmon.py b/python/sglang/srt/models/persimmon.py index 28f4cf250..f62a6f281 100644 --- a/python/sglang/srt/models/persimmon.py +++ b/python/sglang/srt/models/persimmon.py @@ -5,7 +5,6 @@ import torch from torch import nn from transformers import PersimmonConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -201,7 +200,7 @@ class PersimmonModel(nn.Module): ): super().__init__() self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( diff --git a/python/sglang/srt/models/phi.py b/python/sglang/srt/models/phi.py index 1dda797cb..0fe6faa6c 100644 --- a/python/sglang/srt/models/phi.py +++ b/python/sglang/srt/models/phi.py @@ -7,7 +7,6 @@ import torch from torch import nn from transformers import PhiConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -176,7 +175,7 @@ class PhiModel(nn.Module): config.vocab_size, config.hidden_size ) - pp_group = get_pp_group() + pp_group = get_parallel().pp_group pp_size = pp_group.world_size pp_rank = pp_group.rank diff --git a/python/sglang/srt/models/phi3_small.py b/python/sglang/srt/models/phi3_small.py index 15b90ef7b..ca8b25894 100644 --- a/python/sglang/srt/models/phi3_small.py +++ b/python/sglang/srt/models/phi3_small.py @@ -6,7 +6,6 @@ from torch import nn from transformers import Phi3Config from transformers.configuration_utils import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.linear import ( MergedColumnParallelLinear, QKVParallelLinear, @@ -293,7 +292,7 @@ class Phi3SmallModel(nn.Module): self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( config.vocab_size, diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py index 49eaecf8b..9181b275c 100644 --- a/python/sglang/srt/models/qwen2.py +++ b/python/sglang/srt/models/qwen2.py @@ -23,7 +23,6 @@ import torch from torch import nn from sglang.srt.distributed import ( - get_pp_group, get_pp_indices, ) from sglang.srt.layers.activation import SiluAndMul @@ -324,7 +323,7 @@ class Qwen2Model(nn.Module): self.config = config self.padding_idx = getattr(config, "pad_token_id", None) self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -495,7 +494,7 @@ class Qwen2ForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = Qwen2Model( diff --git a/python/sglang/srt/models/qwen2_5_vl.py b/python/sglang/srt/models/qwen2_5_vl.py index 6b805a113..17e451c67 100644 --- a/python/sglang/srt/models/qwen2_5_vl.py +++ b/python/sglang/srt/models/qwen2_5_vl.py @@ -38,7 +38,6 @@ from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import ( Qwen2_5_VLVisionConfig, ) -from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.environ import envs from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.attention.vision import ( @@ -721,7 +720,7 @@ class Qwen2_5_VLForConditionalGeneration(nn.Module): ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.use_data_parallel = get_mm().mm_enable_dp_encoder diff --git a/python/sglang/srt/models/qwen2_eagle.py b/python/sglang/srt/models/qwen2_eagle.py index 20c4a152a..f1842a2af 100644 --- a/python/sglang/srt/models/qwen2_eagle.py +++ b/python/sglang/srt/models/qwen2_eagle.py @@ -24,7 +24,6 @@ from typing import Iterable, Optional, Tuple import torch from torch import nn -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.vocab_parallel_embedding import ( @@ -33,6 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.models.qwen2 import Qwen2DecoderLayer, Qwen2ForCausalLM +from sglang.srt.runtime_context import get_parallel Qwen2Config = None @@ -121,7 +121,7 @@ class Qwen2ForCausalLMEagle(Qwen2ForCausalLM): nn.Module.__init__(self) self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.model = Qwen2Model( config, quant_config=quant_config, prefix=add_prefix("model", prefix) ) diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index 6e7b6363c..8c99eba64 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -33,7 +33,6 @@ from sglang.kernels.ops.elementwise.elementwise import ( ) from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo from sglang.srt.distributed import ( - get_pp_group, get_pp_indices, moe_expert_parallel_all_reduce, moe_tensor_model_parallel_all_reduce, @@ -1056,7 +1055,7 @@ class Qwen2MoeModel(nn.Module): super().__init__() self.config = config self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.moe_dp_size = get_parallel().moe_dp_size @@ -1206,7 +1205,7 @@ class Qwen2MoeForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config alt_stream = get_stream("alt") if _is_cuda else None diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 9fbd04670..b4d88886c 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -5,9 +5,6 @@ from typing import Any, Dict, Iterable, List, Optional, Tuple import torch from torch import nn -from sglang.srt.distributed import ( - get_pp_group, -) from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import QKVParallelLinear, RowParallelLinear @@ -476,7 +473,7 @@ class Qwen3ForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = Qwen3Model( diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py index d3691ea5e..8c48dd68d 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -39,7 +39,6 @@ from sglang.srt.configs.qwen3_5 import ( ) # Distributed -from sglang.srt.distributed import get_pp_group from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation @@ -1628,7 +1627,7 @@ class Qwen3_5ForCausalLM(nn.Module): super().__init__() self.config = config self.hidden_size = config.hidden_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group alt_stream = get_stream("alt") if _is_cuda or _hip_use_alt_stream else None diff --git a/python/sglang/srt/models/qwen3_5_mtp.py b/python/sglang/srt/models/qwen3_5_mtp.py index db1f79fd6..c6b21f0c8 100644 --- a/python/sglang/srt/models/qwen3_5_mtp.py +++ b/python/sglang/srt/models/qwen3_5_mtp.py @@ -23,7 +23,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation @@ -112,7 +111,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module): self.config = config self.tp_size = get_parallel().tp_size self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.fc = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False) RMSNorm_cls = GemmaRMSNorm @@ -130,7 +129,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module): is_nextn=True, ) - if get_pp_group().is_last_rank: + if get_parallel().pp_group.is_last_rank: if config.tie_word_embeddings: self.lm_head = self.model.embed_tokens else: diff --git a/python/sglang/srt/models/qwen3_5_text.py b/python/sglang/srt/models/qwen3_5_text.py index 814d6ff8a..be8d67034 100644 --- a/python/sglang/srt/models/qwen3_5_text.py +++ b/python/sglang/srt/models/qwen3_5_text.py @@ -19,7 +19,6 @@ from typing import Iterable, Optional, Set, Tuple, Union import torch from torch import nn -from sglang.srt.distributed import get_pp_group from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -60,7 +59,7 @@ class Qwen3_5ForCausalLM(nn.Module): super().__init__() self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if quant_config is not None and hasattr(quant_config, "packed_modules_mapping"): quant_config.packed_modules_mapping = self.packed_modules_mapping diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 5e204ecf3..8eed5d729 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -26,9 +26,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import ( - get_pp_group, -) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo @@ -972,7 +969,7 @@ class Qwen3MoeForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = Qwen3MoeModel( diff --git a/python/sglang/srt/models/qwen3_moe_mtp.py b/python/sglang/srt/models/qwen3_moe_mtp.py index e351fb4d7..c401c2e1f 100644 --- a/python/sglang/srt/models/qwen3_moe_mtp.py +++ b/python/sglang/srt/models/qwen3_moe_mtp.py @@ -21,7 +21,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.logits_processor import LogitsProcessor @@ -48,7 +47,7 @@ class Qwen3MoeForCausalLMMTP(Qwen3MoeForCausalLM): config.num_hidden_layers = 1 self.tp_size = get_parallel().tp_size self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.fc = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False) self.pre_fc_norm_embedding = RMSNorm( diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index 43da47e00..4ba179874 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -12,7 +12,6 @@ from sglang.kernels.ops.attention.triton_gdn_fused_proj import ( fused_qkvzba_split_reshape_cat, ) from sglang.srt.configs.qwen3_next import Qwen3NextConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.layers.attention.mamba.mamba import mamba_v2_sharded_weight_loader @@ -1005,7 +1004,7 @@ class Qwen3NextForCausalLM(nn.Module): ) -> None: super().__init__() self.config = config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group assert self.pp_group.is_first_rank and self.pp_group.is_last_rank # The quant config's packed_modules_mapping may be None if it wasn't diff --git a/python/sglang/srt/models/qwen3_next_mtp.py b/python/sglang/srt/models/qwen3_next_mtp.py index f39b989ba..b7c2cb90a 100644 --- a/python/sglang/srt/models/qwen3_next_mtp.py +++ b/python/sglang/srt/models/qwen3_next_mtp.py @@ -23,7 +23,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.layers.layernorm import GemmaRMSNorm @@ -54,7 +53,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM): quant_config = None self.quant_config = quant_config # if not set, model load will be broken in Qwen3NextForCausalLM load_weights() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group # currently based on the provided ckpt, we: # (1) do not use_dedicated_mtp_embeddings provided in ckpt since not provided and directly use the target model embeddings diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index 53310d53b..e620d3e0f 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -27,7 +27,6 @@ from einops import rearrange from transformers.activations import ACT2FN from sglang.srt.configs.qwen3_vl import Qwen3VLConfig, Qwen3VLVisionConfig -from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.environ import envs from sglang.srt.layers.attention.vision import ( BATCH_BUCKETS, @@ -359,7 +358,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): use_data_parallel: bool = False, ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.hidden_size = vision_config.hidden_size self.num_heads = vision_config.num_heads self.num_position_embeddings = vision_config.num_position_embeddings @@ -1307,7 +1306,7 @@ class Qwen3VLForConditionalGeneration(nn.Module): language_model_cls=Qwen3LLMModel, ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.quant_config = quant_config self.use_data_parallel = get_mm().mm_enable_dp_encoder diff --git a/python/sglang/srt/models/qwen4_exp.py b/python/sglang/srt/models/qwen4_exp.py index 80bb46262..34ac3e22d 100644 --- a/python/sglang/srt/models/qwen4_exp.py +++ b/python/sglang/srt/models/qwen4_exp.py @@ -14,7 +14,7 @@ from torch import nn from sglang.kernels.ops.elementwise.elementwise import fused_sigmoid_mul from sglang.srt.configs.qwen4_exp import Qwen4ExpConfig, Qwen4ExpTextConfig -from sglang.srt.distributed import get_tp_group, tensor_model_parallel_all_reduce +from sglang.srt.distributed import tensor_model_parallel_all_reduce from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -853,7 +853,7 @@ class Qwen4ExpPinnedHostEmbedding(VocabParallelEmbedding): allocation_context = nullcontext() if self.tp_size > 1: allocation_context = use_symmetric_memory( - get_tp_group(), disabled=not is_allocation_symmetric() + get_parallel().tp_group, disabled=not is_allocation_symmetric() ) with allocation_context, torch.inference_mode(False): # The gather kernel emits bf16 rows regardless of the table dtype. @@ -1377,7 +1377,7 @@ class Qwen4ExpLayerExtensionMixin: if use_dp_moe_gather: hidden_states, local_hidden_states = ( - get_global_dp_buffer(get_tp_group()), + get_global_dp_buffer(get_parallel().tp_group), hidden_states, ) dp_gather_replicate(hidden_states, local_hidden_states, forward_batch) @@ -1396,11 +1396,11 @@ class Qwen4ExpLayerExtensionMixin: if use_dp_moe_gather: hidden_states, global_hidden_states = ( - get_local_dp_buffer(get_tp_group()), + get_local_dp_buffer(get_parallel().tp_group), hidden_states, ) if should_use_dp_reduce_scatterv(): - get_tp_group().reduce_scatterv( + get_parallel().tp_group.reduce_scatterv( global_hidden_states, output=hidden_states, sizes=get_dp_global_num_tokens(), diff --git a/python/sglang/srt/models/qwen4_exp_mtp.py b/python/sglang/srt/models/qwen4_exp_mtp.py index 31a873014..103e1600d 100644 --- a/python/sglang/srt/models/qwen4_exp_mtp.py +++ b/python/sglang/srt/models/qwen4_exp_mtp.py @@ -9,7 +9,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.layers.layernorm import GemmaRMSNorm @@ -50,7 +49,7 @@ class Qwen4ExpForCausalLMMTP(Qwen3_5ForCausalLMMTP): self.config = config self.tp_size = get_parallel().tp_size self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.hidden_size = config.hidden_size self.hc_count = config.hc_count self._mtp_input_fusion = self._init_mtp_input_fusion(config) diff --git a/python/sglang/srt/models/sarvam_moe.py b/python/sglang/srt/models/sarvam_moe.py index 90d877c9b..48b0fdb4b 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -14,7 +14,6 @@ from transformers import PretrainedConfig from sglang.kernels.ops.attention.utils import concat_and_cast_mha_k_triton from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -1130,7 +1129,7 @@ class SarvamMLAModel(nn.Module): self.config = config self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.alt_stream = get_stream("alt") if _is_cuda else None if self.pp_group.is_first_rank: @@ -1211,7 +1210,7 @@ class SarvamMLAForCausalLM(nn.Module): ) -> None: super().__init__() self._remap_config(config) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = SarvamMLAModel(config, quant_config, add_prefix("model", prefix)) diff --git a/python/sglang/srt/models/sdar.py b/python/sglang/srt/models/sdar.py index 3d5912d97..6c8c744a6 100644 --- a/python/sglang/srt/models/sdar.py +++ b/python/sglang/srt/models/sdar.py @@ -10,7 +10,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.dp_attention import ( @@ -356,7 +355,7 @@ class SDARModel(nn.Module): self.config = config self.vocab_size = config.vocab_size self.embed_dim = config.hidden_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -439,7 +438,7 @@ class SDARForCausalLM(nn.Module): prefix: str = "", ): super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group assert self.pp_group.world_size == 1, ( f"SDARMoeForCausalLM does not support pipeline parallel (pp_size={self.pp_group.world_size}). " "Please set pp_size=1." diff --git a/python/sglang/srt/models/sdar_moe.py b/python/sglang/srt/models/sdar_moe.py index 9ac15f3fe..c2ef19027 100644 --- a/python/sglang/srt/models/sdar_moe.py +++ b/python/sglang/srt/models/sdar_moe.py @@ -11,7 +11,6 @@ from torch import nn from transformers import PretrainedConfig from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -438,7 +437,7 @@ class SDARMoeModel(nn.Module): self.config = config self.vocab_size = config.vocab_size self.embed_dim = config.hidden_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -525,13 +524,13 @@ class SDARMoeForCausalLM(nn.Module): prefix: str = "", ): super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group assert self.pp_group.world_size == 1, ( f"SDARMoeForCausalLM does not support pipeline parallel (pp_size={self.pp_group.world_size}). " "Please set pp_size=1." ) - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config alt_stream = get_stream("alt") if _is_cuda else None diff --git a/python/sglang/srt/models/solar.py b/python/sglang/srt/models/solar.py index bf54128f5..309877d78 100644 --- a/python/sglang/srt/models/solar.py +++ b/python/sglang/srt/models/solar.py @@ -8,7 +8,6 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( @@ -270,7 +269,7 @@ class SolarModel(nn.Module): self.vocab_size = config.vocab_size self.org_vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( config.vocab_size, @@ -290,7 +289,7 @@ class SolarModel(nn.Module): ), prefix=f"{prefix}.layers", ) - if get_pp_group().is_last_rank: + if get_parallel().pp_group.is_last_rank: self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) else: self.norm = PPMissingLayer() @@ -417,7 +416,7 @@ class SolarForCausalLM(nn.Module): prefix: str = "", ): super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = SolarModel( diff --git a/python/sglang/srt/models/spark2_5.py b/python/sglang/srt/models/spark2_5.py index 2fa3ea70c..77a1c0e5d 100644 --- a/python/sglang/srt/models/spark2_5.py +++ b/python/sglang/srt/models/spark2_5.py @@ -4,7 +4,6 @@ from typing import Iterable, Optional, Tuple, Union import torch from torch import nn -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import GeluAndMul from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm @@ -300,7 +299,7 @@ class Spark2_5Model(nn.Module): super().__init__() self.config = config self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -385,7 +384,7 @@ class Spark2_5ForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = Spark2_5Model( diff --git a/python/sglang/srt/models/starcoder2.py b/python/sglang/srt/models/starcoder2.py index 5b33f704a..71c45fcdf 100644 --- a/python/sglang/srt/models/starcoder2.py +++ b/python/sglang/srt/models/starcoder2.py @@ -29,7 +29,6 @@ import torch from torch import nn from transformers import Starcoder2Config -from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -234,7 +233,7 @@ class Starcoder2Model(nn.Module): prefix=f"{prefix}.embed_tokens", ) - pp_group = get_pp_group() + pp_group = get_parallel().pp_group pp_size = pp_group.world_size pp_rank = pp_group.rank self.start_layer = pp_rank * config.num_hidden_layers // pp_size diff --git a/python/sglang/srt/models/step3p5.py b/python/sglang/srt/models/step3p5.py index c9220f2f3..f2d62ad33 100644 --- a/python/sglang/srt/models/step3p5.py +++ b/python/sglang/srt/models/step3p5.py @@ -5,7 +5,6 @@ import torch.nn.functional as F from torch import nn from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -658,7 +657,7 @@ class Step3p5Model(nn.Module): super().__init__() self.config = config self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group alt_stream = get_stream("alt") if _is_cuda else None @@ -797,7 +796,7 @@ class Step3p5ForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config self.model = Step3p5Model( diff --git a/python/sglang/srt/models/transformers.py b/python/sglang/srt/models/transformers.py index af5e65b48..d03c3a19a 100644 --- a/python/sglang/srt/models/transformers.py +++ b/python/sglang/srt/models/transformers.py @@ -34,7 +34,6 @@ from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS from sglang.srt.distributed import ( divide, - get_pp_group, get_pp_indices, tensor_model_parallel_all_reduce, ) @@ -576,7 +575,7 @@ class TransformersBase(nn.Module): self.config = config self.text_config = get_hf_text_config(config) self.weight_mapper = self.hf_to_sglang_mapper - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group # Weight loading attrs self.skip_prefixes: list[str] = [] diff --git a/python/sglang/srt/models/xllm.py b/python/sglang/srt/models/xllm.py index 69ec06355..c22da5565 100644 --- a/python/sglang/srt/models/xllm.py +++ b/python/sglang/srt/models/xllm.py @@ -35,7 +35,7 @@ import torch.nn.functional as F from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce +from sglang.srt.distributed import tensor_model_parallel_all_reduce from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo @@ -1604,7 +1604,7 @@ class XllmModel(nn.Module): self.config = config self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -1701,7 +1701,7 @@ class XllmForCausalLM(nn.Module): prefix: str = "", ) -> None: super().__init__() - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.config = config self.quant_config = quant_config _validate_mova_config(config, quant_config) diff --git a/python/sglang/srt/models/zaya.py b/python/sglang/srt/models/zaya.py index 8f9b19ba1..2d17576ce 100644 --- a/python/sglang/srt/models/zaya.py +++ b/python/sglang/srt/models/zaya.py @@ -53,7 +53,6 @@ from torch import nn from sglang.srt.configs.zaya import ZayaConfig from sglang.srt.distributed import ( - get_pp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.layers.layernorm import RMSNorm @@ -1389,7 +1388,7 @@ class ZayaModel(nn.Module): self.config = config self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group if self.pp_group.is_first_rank: self.embed_tokens = VocabParallelEmbedding( @@ -1489,7 +1488,7 @@ class ZayaForCausalLM(nn.Module): super().__init__() self.config = config self.quant_config = quant_config - self.pp_group = get_pp_group() + self.pp_group = get_parallel().pp_group self.model = ZayaModel( config=config, diff --git a/python/sglang/srt/multimodal/vit_cuda_graph_runner.py b/python/sglang/srt/multimodal/vit_cuda_graph_runner.py index 7a40a45ed..4388493f7 100644 --- a/python/sglang/srt/multimodal/vit_cuda_graph_runner.py +++ b/python/sglang/srt/multimodal/vit_cuda_graph_runner.py @@ -23,8 +23,8 @@ from typing import Dict, Hashable, List, Optional, Tuple import torch import torch.nn as nn -from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.layers.attention.vision import VisionAttention +from sglang.srt.runtime_context import get_parallel class ViTCudaGraphRunner: @@ -167,7 +167,7 @@ class ViTCudaGraphRunner: # graph, and all layers are local in DP mode, so capture locally. if getattr(self.vit, "use_data_parallel", False): return nullcontext() - ca_comm = get_tp_group().ca_comm + ca_comm = get_parallel().tp_group.ca_comm return ca_comm.capture() if ca_comm is not None else nullcontext() def _create_graph( diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 28c5e42cd..fafe1d64d 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -19,7 +19,6 @@ from sglang.kernels.ops.speculative.dspark.dspark_accept import ( accept_sampling, ) from sglang.srt.configs.hybrid_arch import mambaish_config -from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.layers.logits_processor import should_apply_lm_head_quant_method @@ -391,7 +390,7 @@ class DFlashWorkerV2(BaseSpecWorker): self._tp_sync = SpecTpSync( get_parallel().attn_tp_group if get_parallel().enable_dp_attention - else get_tp_group() + else get_parallel().tp_group ) # Under dp attention, the draft worker runs on the per-DP attn-TP @@ -443,7 +442,7 @@ class DFlashWorkerV2(BaseSpecWorker): ) validate_domino_runtime( device=torch.device(self.device), - tp_size=int(get_tp_group().world_size), + tp_size=int(get_parallel().tp_group.world_size), tp_rank=int(self.ps.tp_rank), target_vocab_size=int(self.model_runner.model_config.vocab_size), draft_vocab_size=int(self.draft_model_runner.model_config.vocab_size), @@ -503,7 +502,7 @@ class DFlashWorkerV2(BaseSpecWorker): if self._is_domino: logger.info( "DFLASH Domino rollout enabled (BF16, TP=%s, block-shared candidate pool size=%s).", - int(get_tp_group().world_size), + int(get_parallel().tp_group.world_size), self.domino_candidate_pool_size, ) logger.info( @@ -628,7 +627,7 @@ class DFlashWorkerV2(BaseSpecWorker): SpecTpSyncSite.DFLASH_MEM, self.device, self.gpu_id, - group=get_tp_group(), + group=get_parallel().tp_group, ) if available_mem < 1.0: capture_decode_cuda_graph = False @@ -805,7 +804,7 @@ class DFlashWorkerV2(BaseSpecWorker): if not is_dense_head_weight(lm_head.weight): # Quantized lm_head (FP8/INT) would break the static matmul. return _eager("quantized lm_head") - tp_group = get_tp_group() + tp_group = get_parallel().tp_group if self._is_domino: prefix_gru = self.draft_model.prefix_gru embed_proj = self.draft_model.embed_proj @@ -1253,7 +1252,7 @@ class DFlashWorkerV2(BaseSpecWorker): if not get_parallel().enable_dp_attention: return - tp_group = get_tp_group() + tp_group = get_parallel().tp_group tp_size = int(tp_group.world_size) if tp_size <= 1: return @@ -1461,7 +1460,7 @@ class DFlashWorkerV2(BaseSpecWorker): makes for GGUF models. Padding rows are excluded so argmax cannot return an id outside the real vocabulary. """ - tp_size = int(get_tp_group().world_size) + tp_size = int(get_parallel().tp_group.world_size) if tp_size != 1: raise RuntimeError( "DFLASH with a quantized target lm_head is only supported at " @@ -1523,7 +1522,7 @@ class DFlashWorkerV2(BaseSpecWorker): return out_tokens shard = lm_head.shard_indices - tp_group = get_tp_group() + tp_group = get_parallel().tp_group tp_size = int(tp_group.world_size) # Valid ranges in the local shard (excluding padding): @@ -2517,7 +2516,7 @@ class DFlashWorkerV2(BaseSpecWorker): embed_proj = self.draft_model.embed_proj if prefix_gru is None or embed_proj is None: raise RuntimeError("DFLASH Domino projector modules are unavailable.") - tp_group = get_tp_group() + tp_group = get_parallel().tp_group shard = getattr(lm_head, "shard_indices", None) draft_next = domino_greedy_rollout( draft_hidden=draft_hidden, diff --git a/python/sglang/srt/speculative/dspark_components/dspark_planner.py b/python/sglang/srt/speculative/dspark_components/dspark_planner.py index e56026572..a80767482 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_planner.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_planner.py @@ -11,7 +11,6 @@ from sglang.kernels.ops.speculative.dspark.dspark_schedule import ( ScheduleVerifyLensTopk, compute_sort_survival, ) -from sglang.srt.distributed import get_tp_group from sglang.srt.environ import InvariantCheckLevel, envs from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.managers.overlap_utils import ( @@ -317,7 +316,7 @@ class DSparkVerifyPlanner: if batch.is_extend_in_batch: batch.global_spec_verify_tier_num_tokens = None return - cpu_group = get_tp_group().cpu_group + cpu_group = get_parallel().tp_group.cpu_group local_tensor = torch.tensor([local_tier_num_tokens], dtype=torch.int64) gathered = torch.empty( (torch.distributed.get_world_size(group=cpu_group),), dtype=torch.int64 diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py index 29eebc046..91ee6212b 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -19,7 +19,7 @@ from sglang.srt.mem_cache.allocation_sizing import ( get_alloc_reserve_per_decode, page_aligned_decode_alloc_lens, ) -from sglang.srt.runtime_context import get_parallel, get_spec +from sglang.srt.runtime_context import get_spec from sglang.srt.utils import ( is_cpu, is_cuda, @@ -741,10 +741,10 @@ def eagle_sample( """ import torch.nn.functional as F - from sglang.srt.distributed import get_tp_group from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) + from sglang.srt.runtime_context import get_parallel from sglang.srt.sampling.penaltylib.repetition_penalty import ( apply_scaling_penalties, ) @@ -836,7 +836,7 @@ def eagle_sample( tp_group = ( get_parallel().attn_tp_group if is_dp_attention_enabled() - else get_tp_group() + else get_parallel().tp_group ) if tp_group.world_size > 1: tp_group.broadcast(predict, src=0) @@ -872,7 +872,7 @@ def eagle_sample( tp_group = ( get_parallel().attn_tp_group if is_dp_attention_enabled() - else get_tp_group() + else get_parallel().tp_group ) if tp_group.world_size > 1: tp_group.broadcast(predict, src=0) @@ -999,7 +999,7 @@ def eagle_sample( tp_group = ( get_parallel().attn_tp_group if is_dp_attention_enabled() - else get_tp_group() + else get_parallel().tp_group ) if tp_group.world_size > 1: tp_group.broadcast(predict, src=0) diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index dde9b5dca..e40c64186 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -8,7 +8,6 @@ import torch from sglang.kernels.ops.speculative.topk1 import draft_topk1_postprocess from sglang.srt.configs.model_config import get_dsa_mtp_topk_width -from sglang.srt.distributed import get_pp_group from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_runner import ( @@ -1300,7 +1299,7 @@ class EAGLEWorkerV2(BaseSpecWorker): # Only the last PP stage runs the draft; other EAGLEWorkerV2 instances # return proxies so scheduler dispatch remains rank-uniform. - self._hosts_draft = get_pp_group().is_last_rank + self._hosts_draft = get_parallel().pp_group.is_last_rank self._draft_worker = ( EagleDraftWorker( server_args, diff --git a/python/sglang/srt/utils/poll_based_barrier.py b/python/sglang/srt/utils/poll_based_barrier.py index db1d22763..26a902ff1 100644 --- a/python/sglang/srt/utils/poll_based_barrier.py +++ b/python/sglang/srt/utils/poll_based_barrier.py @@ -1,6 +1,6 @@ import torch -from sglang.srt.distributed import get_world_group +from sglang.srt.runtime_context import get_parallel class PollBasedBarrier: @@ -26,6 +26,6 @@ class PollBasedBarrier: torch.distributed.all_reduce( global_arrived, torch.distributed.ReduceOp.MIN, - group=get_world_group().cpu_group, + group=get_parallel().world_group.cpu_group, ) return global_arrived.item() diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py index b5fc99a51..46aac0c27 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py @@ -5,19 +5,20 @@ import torch import torch.nn.functional as F from torch import nn -# Patch TP world size / rank before importing modules that read them at __init__. -import sglang.srt.layers.linear as _linear_mod +# State the topology before importing modules that read it at __init__. The +# group is stated too: `RowParallelLinear.forward` asks for it to manage +# symmetric memory, and `world_size=1` short-circuits that. from sglang.srt.runtime_context import get_context, get_parallel _parallel_override = get_parallel().override( - tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0 + tp_size=1, + tp_rank=0, + attn_tp_size=1, + attn_tp_rank=0, + tp_group=SimpleNamespace(world_size=1), ) _parallel_override.__enter__() -# RowParallelLinear.forward calls get_tp_group() to manage symmetric memory. -# Provide a stub group with world_size=1 so use_symmetric_memory short-circuits. -_linear_mod.get_tp_group = lambda: SimpleNamespace(world_size=1) - from sglang.srt.configs.falcon_h1 import FalconH1Config # noqa: E402 from sglang.srt.configs.mamba_utils import ( # noqa: E402 Mamba2CacheParams, diff --git a/test/registered/kernels/ops/layernorm/test_mhc_kernels.py b/test/registered/kernels/ops/layernorm/test_mhc_kernels.py index 3311b34f2..099d2b4cc 100644 --- a/test/registered/kernels/ops/layernorm/test_mhc_kernels.py +++ b/test/registered/kernels/ops/layernorm/test_mhc_kernels.py @@ -10,11 +10,25 @@ from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-large") +@pytest.fixture +def stated_tp_group(): + """A TP group for a test that runs in a process without one. + + The production call passes the group *into* `use_symmetric_memory`, so + stubbing that context manager does not stop the read -- the argument is + evaluated first. Stating it on the context answers every spelling. + """ + from sglang.srt.runtime_context import get_parallel + + with get_parallel().override(tp_group=None): + yield + + @pytest.mark.parametrize("hidden_size", [4096, 7168]) @pytest.mark.parametrize("num_tokens", [0, 1, 8, 17, 32, 64]) @pytest.mark.parametrize("use_norm", [False, True]) def test_mhc_fused_post_pre_matches_unfused( - monkeypatch, hidden_size, num_tokens, use_norm + monkeypatch, hidden_size, num_tokens, use_norm, stated_tp_group ): if not torch.cuda.is_available(): pytest.skip("CUDA is required for TileLang mHC kernels") @@ -22,12 +36,10 @@ def test_mhc_fused_post_pre_matches_unfused( monkeypatch.setattr(mhc, "is_dsa_prefill_cp_interleave", lambda: False) # This is a single-process kernel unit test with no TP group initialized. # mhc_pre / mhc_fused_post_pre allocate the MoE input in the symmetric-memory - # pool via use_symmetric_memory(get_tp_group(), ...); bypass that path so the - # kernel runs with a plain torch.empty allocation. Mirrors the workaround in - # test_mxfp4_sm90_cutlass.py for the same TP-group-not-initialized case. + # pool, which asks for the TP group; bypassing the allocation is enough, and + # then nothing asks. Mirrors the workaround in test_mxfp4_sm90_cutlass.py. monkeypatch.setattr(mhc, "use_symmetric_memory", lambda *a, **kw: nullcontext()) monkeypatch.setattr(mhc, "is_allocation_symmetric", lambda: False) - monkeypatch.setattr(mhc, "get_tp_group", lambda: None) torch.manual_seed(0) device = torch.device("cuda") hc_mult = 4 diff --git a/test/registered/kernels/ops/moe/test_minimax_quant_scatter.py b/test/registered/kernels/ops/moe/test_minimax_quant_scatter.py index d88cbec27..e905cfe9f 100644 --- a/test/registered/kernels/ops/moe/test_minimax_quant_scatter.py +++ b/test/registered/kernels/ops/moe/test_minimax_quant_scatter.py @@ -35,6 +35,20 @@ register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="4-gpu-b dev = "cuda" +@pytest.fixture +def stated_tp_group(): + """A TP group for a test that runs in a process without one. + + The production call passes the group *into* `use_symmetric_memory`, so + stubbing that context manager does not stop the read -- the argument is + evaluated first. Stating it on the context answers every spelling. + """ + from sglang.srt.runtime_context import get_parallel + + with get_parallel().override(tp_group=None): + yield + + def test_sm120_mxfp8_dispatch_preserves_activation_scale_recipe(monkeypatch): """SM120 group-128 activations must not use the MXFP8 weight-scale recipe.""" from sglang.srt.layers import deep_gemm_wrapper @@ -257,7 +271,9 @@ def test_standard_layout_auto_memory_policy(monkeypatch): @pytest.mark.parametrize("weight_dtype", ["fp8", "bf16"]) -def test_standard_masked_runner_matches_compact_end_to_end(monkeypatch, weight_dtype): +def test_standard_masked_runner_matches_compact_end_to_end( + monkeypatch, weight_dtype, stated_tp_group +): """Exercise both production grouped GEMMs through the standard path.""" arch_major, _ = torch.cuda.get_device_capability(torch.cuda.current_device()) if arch_major <= 9: @@ -266,7 +282,6 @@ def test_standard_masked_runner_matches_compact_end_to_end(monkeypatch, weight_d # This kernel test runs outside a model-parallel process. Bypass only the # symmetric-allocation context; all pre-permute, DeepGEMM, activation, # quantization, down-GEMM, and post-permute kernels remain real. - monkeypatch.setattr(deep_gemm_runner, "get_tp_group", lambda: None) monkeypatch.setattr( deep_gemm_runner, "use_symmetric_memory", diff --git a/test/registered/ops/test_aiter_greedy_sample_amd.py b/test/registered/ops/test_aiter_greedy_sample_amd.py index e605d1d28..a479b497c 100644 --- a/test/registered/ops/test_aiter_greedy_sample_amd.py +++ b/test/registered/ops/test_aiter_greedy_sample_amd.py @@ -22,7 +22,7 @@ register_amd_ci(est_time=60, suite="stage-b-test-1-gpu-small-amd") def _mock_global_server_args(backend="pytorch"): - from sglang.srt.layers import sampler as sampler_mod + from sglang.srt.runtime_context import get_parallel from sglang.srt.server_args import ( ServerArgs, set_global_server_args_for_scheduler, @@ -37,7 +37,11 @@ def _mock_global_server_args(backend="pytorch"): class _DummyTPGroup: device_group = None - sampler_mod.get_tp_group = lambda: _DummyTPGroup() + # `Sampler.__init__` asks the context for the group; state one for the rest + # of the process, since this process has no distributed init. Not the scoped + # `override()`: its context manager would be collected here and take the + # value back down with it. + get_parallel().override_permanently(tp_group=_DummyTPGroup()) from sglang.srt.runtime_context import get_flags get_flags().dp.enabled = False diff --git a/test/registered/unit/disaggregation/test_encode_server.py b/test/registered/unit/disaggregation/test_encode_server.py index 94724679c..757dbfb39 100644 --- a/test/registered/unit/disaggregation/test_encode_server.py +++ b/test/registered/unit/disaggregation/test_encode_server.py @@ -1012,7 +1012,7 @@ class TestEncoderDelivery(CustomTestCase): with ( patch( - "sglang.srt.disaggregation.encoder.server.get_tp_group", + "sglang.srt.distributed.parallel_state.get_tp_group", return_value=TPGroup(), ), patch( @@ -1052,7 +1052,7 @@ class TestEncoderDelivery(CustomTestCase): with ( patch( - "sglang.srt.disaggregation.encoder.server.get_tp_group", + "sglang.srt.distributed.parallel_state.get_tp_group", return_value=TPGroup(), ), patch( @@ -1096,7 +1096,7 @@ class TestEncoderDelivery(CustomTestCase): with ( patch( - "sglang.srt.disaggregation.encoder.server.get_tp_group", + "sglang.srt.distributed.parallel_state.get_tp_group", return_value=TPGroup(), ), patch( diff --git a/test/registered/unit/disaggregation/test_register_to_bootstrap.py b/test/registered/unit/disaggregation/test_register_to_bootstrap.py index 285d3f425..196a5cafe 100644 --- a/test/registered/unit/disaggregation/test_register_to_bootstrap.py +++ b/test/registered/unit/disaggregation/test_register_to_bootstrap.py @@ -201,7 +201,9 @@ class TestRegisterToBootstrap(CustomTestCase): self.assertIn("10.0.0.1", url_used) @patch("sglang.srt.disaggregation.common.conn.requests.put") - @patch("sglang.srt.disaggregation.common.conn.get_world_group") + # The consumer reads the group through `get_parallel()`, which reads + # through to the canonical getter, so that is where the stub belongs. + @patch("sglang.srt.distributed.parallel_state.get_world_group") def test_rust_attention_dp_replicates_complete_topology_across_hosts( self, mock_world_group, mock_put ): diff --git a/test/registered/unit/layers/quantization/test_mxfp4_sm100_trtllm_gen.py b/test/registered/unit/layers/quantization/test_mxfp4_sm100_trtllm_gen.py index d6912c27b..a47428b38 100644 --- a/test/registered/unit/layers/quantization/test_mxfp4_sm100_trtllm_gen.py +++ b/test/registered/unit/layers/quantization/test_mxfp4_sm100_trtllm_gen.py @@ -31,6 +31,20 @@ if not is_sm100_supported(): GROUP_SIZE = 32 # MXFP4 block size +@pytest.fixture +def stated_tp_group(): + """A TP group for a test that runs in a process without one. + + The production call passes the group *into* `use_symmetric_memory`, so + stubbing that context manager does not stop the read -- the argument is + evaluated first. Stating it on the context answers every spelling. + """ + from sglang.srt.runtime_context import get_parallel + + with get_parallel().override(tp_group=None): + yield + + class _MockLayer: """Hand-built ``FusedMoE`` stand-in (avoids distributed init).""" @@ -283,7 +297,7 @@ def _ref_trtllm(x, layer, method, precision, top_k, router_logits): ], ) def test_apply_trtllm_gen_matches_flashinfer_direct( - tokens, num_experts, hidden, inter, top_k, precision, monkeypatch + tokens, num_experts, hidden, inter, top_k, precision, monkeypatch, stated_tp_group ): """``Mxfp4MoEMethod.apply`` (SM100 branch) must produce the same output as a direct ``trtllm_fp4_block_scale_moe`` call fed the same inputs. @@ -301,7 +315,6 @@ def test_apply_trtllm_gen_matches_flashinfer_direct( fi_trtllm_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext() ) monkeypatch.setattr(fi_trtllm_mod, "is_allocation_symmetric", lambda: False) - monkeypatch.setattr(fi_trtllm_mod, "get_tp_group", lambda: None) fixtures = _make_random_mxfp4(num_experts, hidden, inter) x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda") * 0.1 diff --git a/test/registered/unit/layers/quantization/test_mxfp4_sm120_cutlass.py b/test/registered/unit/layers/quantization/test_mxfp4_sm120_cutlass.py index 30b759085..043642587 100644 --- a/test/registered/unit/layers/quantization/test_mxfp4_sm120_cutlass.py +++ b/test/registered/unit/layers/quantization/test_mxfp4_sm120_cutlass.py @@ -17,6 +17,20 @@ from sglang.test.ci.ci_register import register_cuda_ci register_cuda_ci(est_time=14, stage="base-b", runner_config="1-gpu-small") +@pytest.fixture +def stated_tp_group(): + """A TP group for a test that runs in a process without one. + + The production call passes the group *into* `use_symmetric_memory`, so + stubbing that context manager does not stop the read -- the argument is + evaluated first. Stating it on the context answers every spelling. + """ + from sglang.srt.runtime_context import get_parallel + + with get_parallel().override(tp_group=None): + yield + + def _random_weights(num_experts: int, hidden: int, intermediate: int): generator = torch.Generator(device="cuda").manual_seed(0) w13 = torch.randint( @@ -110,7 +124,7 @@ def test_dsv4_sm120_load_contract(monkeypatch, request): assert captured["fp4_scale_dtype"] == torch.float8_e8m0fnu -def test_dsv4_sm120_matches_direct_flashinfer(monkeypatch): +def test_dsv4_sm120_matches_direct_flashinfer(monkeypatch, stated_tp_group): if not torch.cuda.is_available(): pytest.skip("CUDA required") if torch.cuda.get_device_capability()[0] != 12: @@ -134,7 +148,6 @@ def test_dsv4_sm120_matches_direct_flashinfer(monkeypatch): runner_module, "use_symmetric_memory", lambda *args, **kwargs: nullcontext() ) monkeypatch.setattr(runner_module, "is_allocation_symmetric", lambda: False) - monkeypatch.setattr(runner_module, "get_tp_group", lambda: None) num_experts, hidden, intermediate = 4, 256, 256 w13, w2, w13_scale, w2_scale = _random_weights(num_experts, hidden, intermediate) @@ -254,7 +267,7 @@ def test_dsv4_sm120_matches_direct_flashinfer(monkeypatch): assert torch.equal(actual, expected) -def test_gpt_oss_sm120_padding_layout_and_kernel(monkeypatch): +def test_gpt_oss_sm120_padding_layout_and_kernel(monkeypatch, stated_tp_group): if not torch.cuda.is_available(): pytest.skip("CUDA required") if torch.cuda.get_device_capability() != (12, 0): @@ -277,7 +290,6 @@ def test_gpt_oss_sm120_padding_layout_and_kernel(monkeypatch): runner_module, "use_symmetric_memory", lambda *args, **kwargs: nullcontext() ) monkeypatch.setattr(runner_module, "is_allocation_symmetric", lambda: False) - monkeypatch.setattr(runner_module, "get_tp_group", lambda: None) num_experts, hidden, intermediate = 4, 160, 160 padded_hidden = padded_intermediate = 256 diff --git a/test/registered/unit/layers/quantization/test_mxfp4_sm90_cutlass.py b/test/registered/unit/layers/quantization/test_mxfp4_sm90_cutlass.py index 29dfbcd9d..9fc43d277 100644 --- a/test/registered/unit/layers/quantization/test_mxfp4_sm90_cutlass.py +++ b/test/registered/unit/layers/quantization/test_mxfp4_sm90_cutlass.py @@ -63,6 +63,20 @@ from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig GROUP_SIZE = 32 # MXFP4 block size +@pytest.fixture +def stated_tp_group(): + """A TP group for a test that runs in a process without one. + + The production call passes the group *into* `use_symmetric_memory`, so + stubbing that context manager does not stop the read -- the argument is + evaluated first. Stating it on the context answers every spelling. + """ + from sglang.srt.runtime_context import get_parallel + + with get_parallel().override(tp_group=None): + yield + + class _MockLayer: """Stand-in for ``FusedMoE`` carrying the attributes the SM90 helpers read. @@ -336,7 +350,7 @@ def test_process_weights_matches_direct_interleave(num_experts, hidden, inter): ], ) def test_apply_sm90_cutlass_matches_flashinfer_direct( - tokens, num_experts, hidden, inter, top_k, monkeypatch + tokens, num_experts, hidden, inter, top_k, monkeypatch, stated_tp_group ): """End-to-end: SGLang's ``_apply_sm90_cutlass`` must produce the same output as a direct FlashInfer ``cutlass_fused_moe`` call fed with the @@ -352,7 +366,6 @@ def test_apply_sm90_cutlass_matches_flashinfer_direct( fi_cutlass_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext() ) monkeypatch.setattr(fi_cutlass_mod, "is_allocation_symmetric", lambda: False) - monkeypatch.setattr(fi_cutlass_mod, "get_tp_group", lambda: None) monkeypatch.setattr( fi_cutlass_mod.envs.SGLANG_FLASHINFER_MOE_FUSED_FINALIZE, "get", @@ -593,7 +606,7 @@ def test_humming_range_ignores_prerounded_hidden_tail(): [(8, 256, 256, 1, 0), (8, 192, 192, 1, 0), (8, 256, 256, 2, 1)], ) def test_apply_sm90_humming_matches_flashinfer_direct( - tokens, hidden, inter, ep_size, ep_rank, monkeypatch + tokens, hidden, inter, ep_size, ep_rank, monkeypatch, stated_tp_group ): """SGLang must forward the five Humming scales and enable the new kernel.""" import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass as fi_cutlass_mod @@ -602,7 +615,6 @@ def test_apply_sm90_humming_matches_flashinfer_direct( fi_cutlass_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext() ) monkeypatch.setattr(fi_cutlass_mod, "is_allocation_symmetric", lambda: False) - monkeypatch.setattr(fi_cutlass_mod, "get_tp_group", lambda: None) monkeypatch.setattr( fi_cutlass_mod.envs.SGLANG_FLASHINFER_MOE_FUSED_FINALIZE, "get", @@ -725,7 +737,7 @@ def _make_random_dsv4_mxfp4(num_experts, hidden, inter, seed=0): ], ) def test_dsv4_apply_matches_flashinfer_direct( - tokens, num_experts, hidden, inter, top_k, monkeypatch + tokens, num_experts, hidden, inter, top_k, monkeypatch, stated_tp_group ): """End-to-end: SGLang's DSv4 ``Mxfp4FlashinferCutlassMoEMethod.apply`` output must match a direct FlashInfer ``cutlass_fused_moe`` call with @@ -741,7 +753,6 @@ def test_dsv4_apply_matches_flashinfer_direct( fi_cutlass_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext() ) monkeypatch.setattr(fi_cutlass_mod, "is_allocation_symmetric", lambda: False) - monkeypatch.setattr(fi_cutlass_mod, "get_tp_group", lambda: None) w13, w2, w13_s, w2_s = _make_random_dsv4_mxfp4(num_experts, hidden, inter) w1, w3 = w13.chunk(2, dim=1) diff --git a/test/registered/unit/mem_cache/test_mem_pool_host.py b/test/registered/unit/mem_cache/test_mem_pool_host.py index a080edab4..4a7179536 100644 --- a/test/registered/unit/mem_cache/test_mem_pool_host.py +++ b/test/registered/unit/mem_cache/test_mem_pool_host.py @@ -270,8 +270,9 @@ class TestHostMemoryBudget(CustomTestCase): unittest.mock.patch.object( torch.distributed, "is_initialized", return_value=True ), - unittest.mock.patch.object( - base, "get_world_group", return_value=fake_group + unittest.mock.patch( + "sglang.srt.distributed.parallel_state.get_world_group", + return_value=fake_group, ), ): self.assertEqual(base.ranks_per_host(), 8) diff --git a/test/registered/unit/model_executor/runner_utils/test_graph_pool_borrow.py b/test/registered/unit/model_executor/runner_utils/test_graph_pool_borrow.py index 59a49400f..831bf434a 100644 --- a/test/registered/unit/model_executor/runner_utils/test_graph_pool_borrow.py +++ b/test/registered/unit/model_executor/runner_utils/test_graph_pool_borrow.py @@ -217,7 +217,12 @@ class TestGraphPoolBorrow(CustomTestCase): "sglang.srt.layers.dp_attention.is_dp_attention_enabled", return_value=False, ), - patch("sglang.srt.distributed.get_tp_group", return_value=tp_group), + # `parallel_state`, not the package re-export: a stub on the + # re-export is never consulted. + patch( + "sglang.srt.distributed.parallel_state.get_tp_group", + return_value=tp_group, + ), patch( "sglang.kernels.ops.speculative.sampling.tree_speculative_sampling_target_only", side_effect=fake_sampling, diff --git a/test/registered/unit/model_executor/test_kv_canary_headroom.py b/test/registered/unit/model_executor/test_kv_canary_headroom.py index a123b89f6..901e6057b 100644 --- a/test/registered/unit/model_executor/test_kv_canary_headroom.py +++ b/test/registered/unit/model_executor/test_kv_canary_headroom.py @@ -60,9 +60,8 @@ class TestCanaryHeadroom(CustomTestCase): ), ), patch.object(kv_pool_runtime.torch.cuda, "synchronize"), - patch.object( - kv_pool_runtime, - "get_world_group", + patch( + "sglang.srt.distributed.parallel_state.get_world_group", return_value=SimpleNamespace(world_size=1, cpu_group=None), ), patch.object(kv_pool_runtime, "get_available_gpu_memory", return_value=20), diff --git a/test/registered/unit/model_loader/test_prefetch_checkpoints.py b/test/registered/unit/model_loader/test_prefetch_checkpoints.py index 1685cd323..13e992be0 100644 --- a/test/registered/unit/model_loader/test_prefetch_checkpoints.py +++ b/test/registered/unit/model_loader/test_prefetch_checkpoints.py @@ -238,7 +238,7 @@ class TestPrefetchCheckpoints(CustomTestCase): patch("concurrent.futures.ThreadPoolExecutor", _InlineExecutor), patch("concurrent.futures.wait", side_effect=_wait_all), patch( - "sglang.srt.model_loader.weight_utils.get_world_group", + "sglang.srt.distributed.parallel_state.get_world_group", return_value=FakeWorldGroup(), ), patch( diff --git a/test/registered/unit/model_loader/test_presharded_loader.py b/test/registered/unit/model_loader/test_presharded_loader.py index f1b5fddef..bc74cd311 100644 --- a/test/registered/unit/model_loader/test_presharded_loader.py +++ b/test/registered/unit/model_loader/test_presharded_loader.py @@ -653,7 +653,8 @@ class TestStructuralSignature(unittest.TestCase): fake_group.all_gather_object.side_effect = lambda local: ["sig-pp0", "sig-pp1"] with mock.patch( - "sglang.srt.distributed.get_world_group", return_value=fake_group + "sglang.srt.distributed.parallel_state.get_world_group", + return_value=fake_group, ): agg_from_rank0 = ( PreshardedModelLoader._make_rank_invariant_structural_signature( @@ -673,7 +674,8 @@ class TestStructuralSignature(unittest.TestCase): "sig-pp1-changed", ] with mock.patch( - "sglang.srt.distributed.get_world_group", return_value=fake_group + "sglang.srt.distributed.parallel_state.get_world_group", + return_value=fake_group, ): agg_changed = ( PreshardedModelLoader._make_rank_invariant_structural_signature( diff --git a/test/registered/unit/model_loader/test_transformers_fallback.py b/test/registered/unit/model_loader/test_transformers_fallback.py index 671c20f05..9989c1569 100644 --- a/test/registered/unit/model_loader/test_transformers_fallback.py +++ b/test/registered/unit/model_loader/test_transformers_fallback.py @@ -55,7 +55,7 @@ class TestTransformersFallbackSkipSubstrs(CustomTestCase): with ( patch( - "sglang.srt.models.transformers.get_pp_group", + "sglang.srt.distributed.parallel_state.get_pp_group", return_value=SimpleNamespace(), ), patch( diff --git a/test/registered/unit/multimodal/test_vit_cuda_graph_runner.py b/test/registered/unit/multimodal/test_vit_cuda_graph_runner.py index 3858acb19..09a72e4fd 100644 --- a/test/registered/unit/multimodal/test_vit_cuda_graph_runner.py +++ b/test/registered/unit/multimodal/test_vit_cuda_graph_runner.py @@ -33,7 +33,7 @@ def _runner(*, use_data_parallel: bool) -> ViTCudaGraphRunner: def test_dp_vit_graph_capture_does_not_enter_tp_communication_capture(): runner = _runner(use_data_parallel=True) with patch( - "sglang.srt.multimodal.vit_cuda_graph_runner.get_tp_group", + "sglang.srt.distributed.parallel_state.get_tp_group", side_effect=AssertionError("DP capture must be rank-local"), ): with runner._capture_context(): @@ -53,7 +53,7 @@ def test_non_dp_vit_graph_capture_uses_tp_communication_capture(): group = SimpleNamespace(ca_comm=SimpleNamespace(capture=lambda: Capture())) runner = _runner(use_data_parallel=False) with patch( - "sglang.srt.multimodal.vit_cuda_graph_runner.get_tp_group", return_value=group + "sglang.srt.distributed.parallel_state.get_tp_group", return_value=group ): with runner._capture_context(): pass