Add get_parallel(): a structured accessor for parallel-topology state (#28567)

This commit is contained in:
Cheng Wan
2026-06-17 20:23:43 -07:00
committed by GitHub
parent d27d8b24de
commit 53318911ca
184 changed files with 1871 additions and 1733 deletions
@@ -19,7 +19,6 @@ from sglang.srt.layers.communicator import (
CommunicateSummableTensorPairFn, CommunicateSummableTensorPairFn,
ScatterMode, ScatterMode,
) )
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.moe import ( from sglang.srt.layers.moe import (
get_deepep_mode, get_deepep_mode,
get_moe_a2a_backend, get_moe_a2a_backend,
@@ -40,6 +39,7 @@ from sglang.srt.model_executor.forward_batch_info import (
compute_position, compute_position,
) )
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
@@ -649,7 +649,7 @@ class TboForwardBatchPreparer:
), f"{key=} {old_value=} {num_tokens=} {batch=}" ), f"{key=} {old_value=} {num_tokens=} {batch=}"
output_dict[key] = old_value[start_token_index:end_token_index] output_dict[key] = old_value[start_token_index:end_token_index]
attention_tp_size = get_attention_tp_size() attention_tp_size = get_parallel().attn_tp_size
output_dict["tbo_padded_len"] = ( output_dict["tbo_padded_len"] = (
(end_token_index - start_token_index - 1) // attention_tp_size + 1 (end_token_index - start_token_index - 1) // attention_tp_size + 1
) * attention_tp_size ) * attention_tp_size
+3 -4
View File
@@ -24,8 +24,6 @@ from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
divide, divide,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
@@ -35,6 +33,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
Phase, Phase,
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
@@ -356,7 +355,7 @@ class ScaledActivation(nn.Module):
self.act = act_module self.act = act_module
self.input_is_parallel = input_is_parallel self.input_is_parallel = input_is_parallel
if input_is_parallel: if input_is_parallel:
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
intermediate_size_per_partition = divide(intermediate_size, tp_size) intermediate_size_per_partition = divide(intermediate_size, tp_size)
else: else:
intermediate_size_per_partition = intermediate_size intermediate_size_per_partition = intermediate_size
@@ -373,7 +372,7 @@ class ScaledActivation(nn.Module):
def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor): def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
param_data = param.data param_data = param.data
if self.input_is_parallel: if self.input_is_parallel:
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
shard_size = param_data.shape[0] shard_size = param_data.shape[0]
start_idx = tp_rank * shard_size start_idx = tp_rank * shard_size
loaded_weight = loaded_weight.narrow(0, start_idx, shard_size) loaded_weight = loaded_weight.narrow(0, start_idx, shard_size)
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_parallel
""" """
end to end attention solution with aiter kernels end to end attention solution with aiter kernels
""" """
@@ -24,7 +26,6 @@ from sglang.srt.layers.attention.utils import (
get_num_kv_index_blocks_flashmla, get_num_kv_index_blocks_flashmla,
) )
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_tp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
@@ -154,11 +155,11 @@ class AiterAttnBackend(AttentionBackend):
self.speculative_num_steps = model_runner.server_args.speculative_num_steps self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.topk = topk self.topk = topk
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.head_dim = model_runner.model_config.head_dim self.head_dim = model_runner.model_config.head_dim
self.num_kv_head = model_runner.model_config.get_num_kv_heads( self.num_kv_head = model_runner.model_config.get_num_kv_heads(
get_attention_tp_size() get_parallel().attn_tp_size
) )
self.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype = model_runner.kv_cache_dtype
@@ -2565,10 +2566,10 @@ class AiterIndicesUpdaterPrefill:
def __init__(self, model_runner: ModelRunner, attn_backend: AttentionBackend): def __init__(self, model_runner: ModelRunner, attn_backend: AttentionBackend):
# Parse Constants # Parse Constants
self.num_qo_heads = ( self.num_qo_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.num_kv_heads = model_runner.model_config.get_num_kv_heads( self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
get_attention_tp_size() get_parallel().attn_tp_size
) )
self.head_dim = model_runner.model_config.head_dim self.head_dim = model_runner.model_config.head_dim
self.data_type = model_runner.kv_cache_dtype self.data_type = model_runner.kv_cache_dtype
@@ -2781,7 +2782,7 @@ class AiterMultiStepDraftBackend:
) )
self.max_context_len = self.attn_backends[0].max_context_len self.max_context_len = self.attn_backends[0].max_context_len
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.device = model_runner.device self.device = model_runner.device
# Cached variables for generate_draft_decode_kv_indices # Cached variables for generate_draft_decode_kv_indices
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_parallel
""" """
Support attention backend for Cutlass MLA. Support attention backend for Cutlass MLA.
@@ -16,7 +18,6 @@ from sglang.srt.layers.attention.utils import (
create_flashmla_kv_indices_triton, create_flashmla_kv_indices_triton,
get_num_kv_index_blocks_flashmla, get_num_kv_index_blocks_flashmla,
) )
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.utils import is_cuda from sglang.srt.utils import is_cuda
@@ -62,14 +63,14 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
) )
self.num_q_heads = ( self.num_q_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.num_kv_heads = model_runner.model_config.get_num_kv_heads( self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
get_attention_tp_size() get_parallel().attn_tp_size
) )
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.num_local_heads = ( self.num_local_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.forward_metadata: Union[CutlassMLADecodeMetadata] = None self.forward_metadata: Union[CutlassMLADecodeMetadata] = None
self.kv_lora_rank = model_runner.model_config.kv_lora_rank self.kv_lora_rank = model_runner.model_config.kv_lora_rank
@@ -20,6 +20,7 @@ import torch.nn.functional as F
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.runtime_context import get_parallel
if envs.SGLANG_OPT_USE_COMPRESSOR_V2.get(): if envs.SGLANG_OPT_USE_COMPRESSOR_V2.get():
# NOTE: should eventually be the only compressor backend # NOTE: should eventually be the only compressor backend
@@ -55,10 +56,6 @@ from sglang.srt.layers.attention.dsv4.quant_k_cache import (
from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import ( from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import (
SparsePrefillChunkCache, SparsePrefillChunkCache,
) )
from sglang.srt.layers.dp_attention import (
get_attention_cp_rank,
get_attention_cp_size,
)
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
@@ -311,8 +308,8 @@ class DSV4AttnMetadata:
] ]
def apply_cp_reindex(self) -> None: def apply_cp_reindex(self) -> None:
cp_rank = get_attention_cp_rank() cp_rank = get_parallel().attn_cp_rank
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
idx = slice(cp_rank, None, cp_size) idx = slice(cp_rank, None, cp_size)
pre_global_len = self.seq_lens_casual.shape[0] pre_global_len = self.seq_lens_casual.shape[0]
assert pre_global_len % cp_size == 0, ( assert pre_global_len % cp_size == 0, (
@@ -20,6 +20,7 @@ import torch.nn.functional as F
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.runtime_context import get_parallel
if envs.SGLANG_OPT_USE_COMPRESSOR_V2.get(): if envs.SGLANG_OPT_USE_COMPRESSOR_V2.get():
from sglang.srt.layers.attention.dsv4.compressor_v2 import ( from sglang.srt.layers.attention.dsv4.compressor_v2 import (
@@ -33,6 +34,7 @@ else:
FusedCompressMetadata, FusedCompressMetadata,
create_paged_compressor_data, create_paged_compressor_data,
) )
from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin
from sglang.srt.layers.attention.dsv4.metadata import ( from sglang.srt.layers.attention.dsv4.metadata import (
PagedIndexerMetadata, PagedIndexerMetadata,
@@ -45,10 +47,6 @@ from sglang.srt.layers.attention.dsv4.metadata_kernel import (
from sglang.srt.layers.attention.dsv4.quant_k_cache import ( from sglang.srt.layers.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton, quant_to_nope_fp8_rope_bf16_pack_triton,
) )
from sglang.srt.layers.dp_attention import (
get_attention_cp_rank,
get_attention_cp_size,
)
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
@@ -289,8 +287,8 @@ class DSV4AttnMetadata:
] ]
def apply_cp_reindex(self) -> None: def apply_cp_reindex(self) -> None:
cp_rank = get_attention_cp_rank() cp_rank = get_parallel().attn_cp_rank
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
idx = slice(cp_rank, None, cp_size) idx = slice(cp_rank, None, cp_size)
pre_global_len = self.seq_lens_casual.shape[0] pre_global_len = self.seq_lens_casual.shape[0]
assert pre_global_len % cp_size == 0, ( assert pre_global_len % cp_size == 0, (
@@ -1203,7 +1201,7 @@ class DeepseekV4HipRadixBackend(
# HIP backend (DeepseekV4HipRadixBackend, selected only when is_hip()). # HIP backend (DeepseekV4HipRadixBackend, selected only when is_hip()).
# The NVIDIA path uses DeepseekV4AttnBackend and never reaches here, so # The NVIDIA path uses DeepseekV4AttnBackend and never reaches here, so
# these CP changes do not affect B200/H200 execution. # these CP changes do not affect B200/H200 execution.
_cp_size = get_attention_cp_size() _cp_size = get_parallel().attn_cp_size
_cp_active = ( _cp_active = (
_cp_size > 1 _cp_size > 1
and is_dsa_prefill_cp_round_robin_split() and is_dsa_prefill_cp_round_robin_split()
@@ -1214,7 +1212,7 @@ class DeepseekV4HipRadixBackend(
final_pos_full = final_pos final_pos_full = final_pos
positions_full = positions positions_full = positions
if _cp_active: if _cp_active:
_sl = slice(get_attention_cp_rank(), None, _cp_size) _sl = slice(get_parallel().attn_cp_rank, None, _cp_size)
state_slot = state_slot[_sl].contiguous() state_slot = state_slot[_sl].contiguous()
chunk_start = chunk_start[_sl].contiguous() chunk_start = chunk_start[_sl].contiguous()
cu_q = cu_q[_sl].contiguous() cu_q = cu_q[_sl].contiguous()
@@ -26,6 +26,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
get_tc_piecewise_forward_context, get_tc_piecewise_forward_context,
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.state_capturer.indexer_topk import ( from sglang.srt.state_capturer.indexer_topk import (
maybe_capture_indexer_topk, maybe_capture_indexer_topk,
) )
@@ -72,8 +73,6 @@ if is_npu():
from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_attn_context_model_parallel_rank,
get_attn_context_model_parallel_world_size,
get_attn_tp_group, get_attn_tp_group,
) )
from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.distributed.parallel_state import get_pp_group
@@ -330,8 +329,8 @@ class Indexer(MultiPlatformOp):
self.alt_stream = alt_stream self.alt_stream = alt_stream
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
if self.dsa_enable_prefill_cp: if self.dsa_enable_prefill_cp:
self.cp_size = get_attn_context_model_parallel_world_size() self.cp_size = get_parallel().attn_cp_size
self.cp_rank = get_attn_context_model_parallel_rank() self.cp_rank = get_parallel().attn_cp_rank
else: else:
self.cp_size = None self.cp_size = None
self.cp_rank = None self.cp_rank = None
+11 -13
View File
@@ -8,10 +8,8 @@ import triton.language as tl
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
DpPaddingMode, DpPaddingMode,
get_attention_cp_rank,
get_attention_cp_size,
get_attention_dp_rank,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import get_bool_env_var, is_hip from sglang.srt.utils import get_bool_env_var, is_hip
from sglang.srt.utils.common import ceil_align, ceil_div from sglang.srt.utils.common import ceil_align, ceil_div
@@ -85,7 +83,7 @@ def is_dsa_prefill_cp_round_robin_split():
def can_dsa_prefill_cp_round_robin_split(forward_batch: "ForwardBatch"): def can_dsa_prefill_cp_round_robin_split(forward_batch: "ForwardBatch"):
if not forward_batch.forward_mode.is_context_parallel_extend(): if not forward_batch.forward_mode.is_context_parallel_extend():
return False return False
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
seq_len = sum(forward_batch.extend_seq_lens_cpu) seq_len = sum(forward_batch.extend_seq_lens_cpu)
return ( return (
is_dsa_prefill_cp_round_robin_split() is_dsa_prefill_cp_round_robin_split()
@@ -108,8 +106,8 @@ def dsa_cp_round_robin_split_data(input_: Union[torch.Tensor, List]):
| dp_atten_tp3: token3, token7, token11, token15, token19, ... | | dp_atten_tp3: token3, token7, token11, token15, token19, ... |
| +-------------------------+ | +-------------------------+
""" """
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
cp_rank = get_attention_cp_rank() cp_rank = get_parallel().attn_cp_rank
if isinstance(input_, (tuple, list)): if isinstance(input_, (tuple, list)):
indices = range(cp_rank, len(input_), cp_size) indices = range(cp_rank, len(input_), cp_size)
return input_[indices] return input_[indices]
@@ -133,7 +131,7 @@ def cal_padded_tokens(forward_batch: "ForwardBatch"):
global_num_tokens = forward_batch.global_num_tokens_cpu.copy() global_num_tokens = forward_batch.global_num_tokens_cpu.copy()
sync_group_size = len(global_num_tokens) sync_group_size = len(global_num_tokens)
attn_cp_size = get_attention_cp_size() attn_cp_size = get_parallel().attn_cp_size
# Must match the CP padding in ForwardBatch.prepare_mlp_sync_batch. # Must match the CP padding in ForwardBatch.prepare_mlp_sync_batch.
cp_align_size = get_cp_padding_align_size() cp_align_size = get_cp_padding_align_size()
for i in range(sync_group_size): for i in range(sync_group_size):
@@ -144,7 +142,7 @@ def cal_padded_tokens(forward_batch: "ForwardBatch"):
if dp_padding_mode.is_max_len(): if dp_padding_mode.is_max_len():
tokens = max(global_num_tokens) tokens = max(global_num_tokens)
elif len(global_num_tokens) > 1: elif len(global_num_tokens) > 1:
tokens = global_num_tokens[get_attention_dp_rank()] tokens = global_num_tokens[get_parallel().attn_dp_rank]
else: else:
tokens = global_num_tokens[0] tokens = global_num_tokens[0]
if can_dsa_prefill_cp_round_robin_split(forward_batch): if can_dsa_prefill_cp_round_robin_split(forward_batch):
@@ -153,7 +151,7 @@ def cal_padded_tokens(forward_batch: "ForwardBatch"):
def pad_dsa_cache_seqlens(forward_batch: "ForwardBatch", dsa_cache_seqlens): def pad_dsa_cache_seqlens(forward_batch: "ForwardBatch", dsa_cache_seqlens):
attn_cp_size = get_attention_cp_size() attn_cp_size = get_parallel().attn_cp_size
needs_cp_pad = attn_cp_size > 1 and can_dsa_prefill_cp_round_robin_split( needs_cp_pad = attn_cp_size > 1 and can_dsa_prefill_cp_round_robin_split(
forward_batch forward_batch
) )
@@ -219,8 +217,8 @@ def dsa_cp_round_robin_split_q_seqs_kernel(
def dsa_cp_round_robin_split_q_seqs_cpu(extend_seqs): def dsa_cp_round_robin_split_q_seqs_cpu(extend_seqs):
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
cp_rank = get_attention_cp_rank() cp_rank = get_parallel().attn_cp_rank
extra_seq = 0 extra_seq = 0
q_seqs = [] q_seqs = []
for bs, cur_len in enumerate(extend_seqs): for bs, cur_len in enumerate(extend_seqs):
@@ -245,8 +243,8 @@ def dsa_cp_round_robin_split_q_seqs(
bs_idx_cpu(List) and bs_idx(torch.Tensor): marks which sequences are ultimately selected, bs_idx_cpu(List) and bs_idx(torch.Tensor): marks which sequences are ultimately selected,
i.e., those with a partitioned length greater than zero. i.e., those with a partitioned length greater than zero.
""" """
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
cp_rank = get_attention_cp_rank() cp_rank = get_parallel().attn_cp_rank
# len(ret_q_lens_cpu) == len(bs_idx_cpu) # len(ret_q_lens_cpu) == len(bs_idx_cpu)
ret_q_lens_cpu, bs_idx_cpu = dsa_cp_round_robin_split_q_seqs_cpu(extend_seqs_cpu) ret_q_lens_cpu, bs_idx_cpu = dsa_cp_round_robin_split_q_seqs_cpu(extend_seqs_cpu)
ret_q_lens = torch.empty( ret_q_lens = torch.empty(
@@ -15,6 +15,7 @@ from typing import (
import torch import torch
from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
from sglang.srt.runtime_context import get_parallel
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -48,7 +49,6 @@ from sglang.srt.layers.attention.utils import (
mla_quantize_and_rope_for_fp8, mla_quantize_and_rope_for_fp8,
seqlens_expand_triton, seqlens_expand_triton,
) )
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.utils import is_cuda, is_hip, is_sm100_supported from sglang.srt.utils import is_cuda, is_hip, is_sm100_supported
@@ -311,7 +311,7 @@ class DeepseekSparseAttnBackend(
self.dsa_index_topk = get_dsa_index_topk(model_runner.model_config.hf_config) self.dsa_index_topk = get_dsa_index_topk(model_runner.model_config.hf_config)
self.max_context_len = model_runner.model_config.context_len self.max_context_len = model_runner.model_config.context_len
self.num_q_heads = ( self.num_q_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.kv_cache_dim = model_runner.token_to_kv_pool.kv_cache_dim self.kv_cache_dim = model_runner.token_to_kv_pool.kv_cache_dim
self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim
@@ -19,7 +19,6 @@ from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp
from sglang.srt.layers.attention.dsv4.quant_k_cache import ( from sglang.srt.layers.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton, quant_to_nope_fp8_rope_bf16_pack_triton,
) )
from sglang.srt.layers.dp_attention import get_attention_cp_size
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.linear import ReplicatedLinear
from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output
@@ -28,6 +27,7 @@ from sglang.srt.mem_cache.deepseek_v4_compress_state import (
) )
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.models.deepseek_v2 import _is_hip from sglang.srt.models.deepseek_v2 import _is_hip
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, get_bool_env_var, set_weight_attrs from sglang.srt.utils import add_prefix, get_bool_env_var, set_weight_attrs
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
@@ -427,7 +427,7 @@ class Compressor(nn.Module):
if dsa_use_prefill_cp(forward_batch): if dsa_use_prefill_cp(forward_batch):
kv_score = cp_all_gather_rerange_output( kv_score = cp_all_gather_rerange_output(
kv_score, kv_score,
get_attention_cp_size(), get_parallel().attn_cp_size,
forward_batch, forward_batch,
torch.cuda.current_stream(), torch.cuda.current_stream(),
) )
@@ -19,12 +19,12 @@ from sglang.jit_kernel.flash_attention import (
flash_attn_varlen_func, flash_attn_varlen_func,
flash_attn_with_kvcache, flash_attn_with_kvcache,
) )
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.flashattention_backend import ( from sglang.srt.layers.attention.flashattention_backend import (
FlashAttentionMetadata, FlashAttentionMetadata,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -170,7 +170,7 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
layer_sparse_attention_config = { layer_sparse_attention_config = {
int(i): j for i, j in self.sparse_attention_config[layer_idx].items() int(i): j for i, j in self.sparse_attention_config[layer_idx].items()
} }
start_head = self.num_heads * get_tensor_model_parallel_rank() start_head = self.num_heads * get_parallel().tp_rank
end_head = start_head + self.num_heads end_head = start_head + self.num_heads
return [layer_sparse_attention_config[i] for i in range(start_head, end_head)] return [layer_sparse_attention_config[i] for i in range(start_head, end_head)]
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_parallel
""" """
Support different attention backends. Support different attention backends.
Now there are two backends: FlashInfer and Triton. Now there are two backends: FlashInfer and Triton.
@@ -24,10 +26,6 @@ from sglang.srt.layers.attention.utils import (
assert_buffer_fits, assert_buffer_fits,
create_flashinfer_kv_indices_triton, create_flashinfer_kv_indices_triton,
) )
from sglang.srt.layers.dp_attention import (
get_attention_cp_size,
get_attention_tp_size,
)
from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.memory_pool import KVWriteLoc
@@ -67,9 +65,9 @@ def _cuda_graph_capture_max_bs(server_args, max_bs: int) -> int:
if server_args.enable_two_batch_overlap: if server_args.enable_two_batch_overlap:
mul_base *= 2 mul_base *= 2
if require_gathered_buffer(server_args): if require_gathered_buffer(server_args):
mul_base *= get_attention_tp_size() mul_base *= get_parallel().attn_tp_size
if mul_base % get_attention_cp_size() != 0: if mul_base % get_parallel().attn_cp_size != 0:
mul_base *= get_attention_cp_size() mul_base *= get_parallel().attn_cp_size
return (max_bs + mul_base - 1) // mul_base * mul_base return (max_bs + mul_base - 1) // mul_base * mul_base
@@ -208,9 +206,9 @@ class FlashInferAttnBackend(AttentionBackend):
self.decode_use_tensor_cores = should_use_tensor_core( self.decode_use_tensor_cores = should_use_tensor_core(
kv_cache_dtype=model_runner.kv_cache_dtype, kv_cache_dtype=model_runner.kv_cache_dtype,
num_attention_heads=model_runner.model_config.num_attention_heads num_attention_heads=model_runner.model_config.num_attention_heads
// get_attention_tp_size(), // get_parallel().attn_tp_size,
num_kv_heads=model_runner.model_config.get_num_kv_heads( num_kv_heads=model_runner.model_config.get_num_kv_heads(
get_attention_tp_size() get_parallel().attn_tp_size
), ),
) )
self.max_context_len = model_runner.model_config.context_len self.max_context_len = model_runner.model_config.context_len
@@ -1005,10 +1003,10 @@ class FlashInferIndicesUpdaterDecode:
def __init__(self, model_runner: ModelRunner, attn_backend: FlashInferAttnBackend): def __init__(self, model_runner: ModelRunner, attn_backend: FlashInferAttnBackend):
# Parse Constants # Parse Constants
self.num_qo_heads = ( self.num_qo_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.num_kv_heads = model_runner.model_config.get_num_kv_heads( self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
get_attention_tp_size() get_parallel().attn_tp_size
) )
self.head_dim = model_runner.model_config.head_dim self.head_dim = model_runner.model_config.head_dim
self.data_type = model_runner.kv_cache_dtype self.data_type = model_runner.kv_cache_dtype
@@ -1273,10 +1271,10 @@ class FlashInferIndicesUpdaterPrefill:
def __init__(self, model_runner: ModelRunner, attn_backend: FlashInferAttnBackend): def __init__(self, model_runner: ModelRunner, attn_backend: FlashInferAttnBackend):
# Parse Constants # Parse Constants
self.num_qo_heads = ( self.num_qo_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.num_kv_heads = model_runner.model_config.get_num_kv_heads( self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
get_attention_tp_size() get_parallel().attn_tp_size
) )
self.head_dim = model_runner.model_config.head_dim self.head_dim = model_runner.model_config.head_dim
self.data_type = model_runner.kv_cache_dtype self.data_type = model_runner.kv_cache_dtype
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_parallel
""" """
Support attention backend for flashinfer MLA. Support attention backend for flashinfer MLA.
The flashinfer_mla_disable_ragged flag controls whether to use ragged prefill wrapper and defaults to be false. The flashinfer_mla_disable_ragged flag controls whether to use ragged prefill wrapper and defaults to be false.
@@ -21,7 +23,6 @@ from sglang.srt.layers.attention.flashinfer_backend import (
create_flashinfer_kv_indices_triton, create_flashinfer_kv_indices_triton,
) )
from sglang.srt.layers.attention.utils import assert_buffer_fits from sglang.srt.layers.attention.utils import assert_buffer_fits
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
@@ -81,7 +82,7 @@ class FlashInferMhaChunkKVRunner:
): ):
# Parse Constants # Parse Constants
self.num_local_heads = ( self.num_local_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim
self.qk_rope_head_dim = model_runner.model_config.qk_rope_head_dim self.qk_rope_head_dim = model_runner.model_config.qk_rope_head_dim
@@ -639,7 +640,7 @@ class FlashInferMLAIndicesUpdaterDecode:
def __init__(self, model_runner: ModelRunner, attn_backend: AttentionBackend): def __init__(self, model_runner: ModelRunner, attn_backend: AttentionBackend):
# Parse Constants # Parse Constants
self.num_local_heads = ( self.num_local_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.kv_lora_rank = model_runner.model_config.kv_lora_rank self.kv_lora_rank = model_runner.model_config.kv_lora_rank
self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim
@@ -748,7 +749,7 @@ class FlashInferMLAIndicesUpdaterPrefill:
def __init__(self, model_runner: ModelRunner, attn_backend: AttentionBackend): def __init__(self, model_runner: ModelRunner, attn_backend: AttentionBackend):
# Parse Constants # Parse Constants
self.num_local_heads = ( self.num_local_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.kv_lora_rank = model_runner.model_config.kv_lora_rank self.kv_lora_rank = model_runner.model_config.kv_lora_rank
self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim self.qk_nope_head_dim = model_runner.model_config.qk_nope_head_dim
@@ -17,9 +17,9 @@ from sglang.srt.layers.attention.utils import (
create_flashmla_kv_indices_triton, create_flashmla_kv_indices_triton,
get_num_kv_index_blocks_flashmla, get_num_kv_index_blocks_flashmla,
) )
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -60,11 +60,11 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
) )
self.num_q_heads = ( self.num_q_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.num_local_heads = ( self.num_local_heads = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.forward_metadata: Union[FlashMLADecodeMetadata] = None self.forward_metadata: Union[FlashMLADecodeMetadata] = None
self.kv_lora_rank = model_runner.model_config.kv_lora_rank self.kv_lora_rank = model_runner.model_config.kv_lora_rank
@@ -15,6 +15,7 @@ from sglang.srt.layers.attention.linear.seg_la import SegLaMeta, seg_la_fwd
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_parallel
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -133,13 +134,9 @@ class LightningAttentionBackend(MambaAttnBackendBase):
slopes = torch.tensor( slopes = torch.tensor(
get_slopes(n_attention_heads), dtype=torch.float32 get_slopes(n_attention_heads), dtype=torch.float32
).reshape(n_attention_heads, 1, 1) ).reshape(n_attention_heads, 1, 1)
from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
)
tp_heads = n_attention_heads // get_attention_tp_size() tp_heads = n_attention_heads // get_parallel().attn_tp_size
tp_rank = get_attention_tp_rank() tp_rank = get_parallel().attn_tp_rank
if num_hidden_layers <= 1: if num_hidden_layers <= 1:
slope_rate_list = [slopes * (1 + 1e-5)] slope_rate_list = [slopes * (1 + 1e-5)]
else: else:
@@ -10,8 +10,6 @@ from sglang.srt.configs.mamba_utils import (
) )
from sglang.srt.distributed import ( from sglang.srt.distributed import (
divide, divide,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
) )
from sglang.srt.layers.attention.mamba.mamba2_metadata import Mamba2Metadata from sglang.srt.layers.attention.mamba.mamba2_metadata import Mamba2Metadata
from sglang.srt.layers.attention.mamba.mixer2_rms_norm_gated import Mixer2RMSNormGated from sglang.srt.layers.attention.mamba.mixer2_rms_norm_gated import Mixer2RMSNormGated
@@ -20,8 +18,6 @@ from sglang.srt.layers.attention.mamba.ops import (
selective_state_update, selective_state_update,
) )
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -36,6 +32,7 @@ from sglang.srt.model_loader.weight_utils import (
composed_weight_loader, composed_weight_loader,
sharded_weight_loader, sharded_weight_loader,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import ( from sglang.srt.utils import (
is_cpu, is_cpu,
is_cuda, is_cuda,
@@ -232,11 +229,11 @@ class MambaMixer2(torch.nn.Module):
# - NOTE: currently for the world size DOES NOT divide groups # - NOTE: currently for the world size DOES NOT divide groups
# case, we only support the case when n_groups == 1 # case, we only support the case when n_groups == 1
if is_dp_attention_enabled(): if is_dp_attention_enabled():
self.tp_size = get_attention_tp_size() self.tp_size = get_parallel().attn_tp_size
self.tp_rank = get_attention_tp_rank() self.tp_rank = get_parallel().attn_tp_rank
else: else:
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_parallel().tp_rank
self.num_heads = num_heads = cache_params.shape.num_heads self.num_heads = num_heads = cache_params.shape.num_heads
self.head_dim = cache_params.shape.head_dim self.head_dim = cache_params.shape.head_dim
@@ -6,20 +6,15 @@ from sglang.srt.distributed.communication_op import (
tensor_model_parallel_all_gather, tensor_model_parallel_all_gather,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.distributed.parallel_state import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.layers.attention.fla.layernorm_gated import rms_norm_gated from sglang.srt.layers.attention.fla.layernorm_gated import rms_norm_gated
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
attn_tp_all_reduce, attn_tp_all_reduce,
get_attention_tp_group, get_attention_tp_group,
get_attention_tp_rank,
get_attention_tp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.model_loader.weight_utils import sharded_weight_loader from sglang.srt.model_loader.weight_utils import sharded_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.common import set_weight_attrs from sglang.srt.utils.common import set_weight_attrs
@@ -34,11 +29,11 @@ class Mixer2RMSNormGated(MultiPlatformOp):
super().__init__() super().__init__()
self.use_attn_tp_group = is_dp_attention_enabled() self.use_attn_tp_group = is_dp_attention_enabled()
if self.use_attn_tp_group: if self.use_attn_tp_group:
self.tp_size = get_attention_tp_size() self.tp_size = get_parallel().attn_tp_size
self.tp_rank = get_attention_tp_rank() self.tp_rank = get_parallel().attn_tp_rank
else: else:
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_parallel().tp_rank
self.full_hidden_size = full_hidden_size self.full_hidden_size = full_hidden_size
self.group_size = full_hidden_size // full_n_groups self.group_size = full_hidden_size // full_n_groups
self.per_rank_hidden_size = full_hidden_size // self.tp_size self.per_rank_hidden_size = full_hidden_size // self.tp_size
@@ -12,12 +12,12 @@ from sglang.srt.layers.attention.triton_ops.kv_indices import (
create_flashinfer_kv_indices_triton, create_flashinfer_kv_indices_triton,
) )
from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel
from sglang.srt.speculative.spec_utils import ( from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width, draft_kv_indices_buffer_width,
draft_kv_indices_used_len, draft_kv_indices_used_len,
@@ -132,10 +132,10 @@ class TritonAttnBackend(AttentionBackend):
self.speculative_num_steps = model_runner.server_args.speculative_num_steps self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.num_kv_head = model_runner.model_config.get_num_kv_heads( self.num_kv_head = model_runner.model_config.get_num_kv_heads(
get_attention_tp_size() get_parallel().attn_tp_size
) )
# The decode triton kernel derives attn_lse offsets from attn_logits # The decode triton kernel derives attn_lse offsets from attn_logits
# strides via integer division by v_head_dim (the "// Lv" trick in # strides via integer division by v_head_dim (the "// Lv" trick in
@@ -1386,7 +1386,7 @@ class TritonMultiStepDraftBackend:
) )
self.max_context_len = self.attn_backends[0].max_context_len self.max_context_len = self.attn_backends[0].max_context_len
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.device = model_runner.device self.device = model_runner.device
# Cached variables for generate_draft_decode_kv_indices # Cached variables for generate_draft_decode_kv_indices
@@ -33,12 +33,12 @@ from sglang.srt.layers.attention.utils import (
concat_mla_absorb_q_general, concat_mla_absorb_q_general,
mla_quantize_and_rope_for_fp8, mla_quantize_and_rope_for_fp8,
) )
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2 from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2
@@ -149,9 +149,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
config = model_runner.model_config config = model_runner.model_config
# Model parameters # Model parameters
self.num_q_heads = config.num_attention_heads // get_attention_tp_size() self.num_q_heads = config.num_attention_heads // get_parallel().attn_tp_size
self.num_kv_heads = config.get_num_kv_heads(get_attention_tp_size()) self.num_kv_heads = config.get_num_kv_heads(get_parallel().attn_tp_size)
self.num_local_heads = config.num_attention_heads // get_attention_tp_size() self.num_local_heads = config.num_attention_heads // get_parallel().attn_tp_size
# MLA-specific dimensions # MLA-specific dimensions
self.kv_lora_rank = config.kv_lora_rank self.kv_lora_rank = config.kv_lora_rank
+5 -5
View File
@@ -14,8 +14,8 @@ from einops import rearrange
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm as can_use_jit_qk_norm from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm as can_use_jit_qk_norm
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import get_attention_tp_rank, get_attention_tp_size
from sglang.srt.models.utils import apply_qk_norm from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_bool_env_var, get_bool_env_var,
@@ -347,7 +347,7 @@ class VisionTritonAttention(nn.Module):
use_data_parallel = ( use_data_parallel = (
kwargs["use_data_parallel"] if "use_data_parallel" in kwargs else False kwargs["use_data_parallel"] if "use_data_parallel" in kwargs else False
) )
self.tp_size = 1 if use_data_parallel else get_attention_tp_size() self.tp_size = 1 if use_data_parallel else get_parallel().attn_tp_size
def forward( def forward(
self, self,
@@ -420,7 +420,7 @@ class VisionFlash3Attention(nn.Module):
use_data_parallel = ( use_data_parallel = (
kwargs["use_data_parallel"] if "use_data_parallel" in kwargs else False kwargs["use_data_parallel"] if "use_data_parallel" in kwargs else False
) )
self.tp_size = 1 if use_data_parallel else get_attention_tp_size() self.tp_size = 1 if use_data_parallel else get_parallel().attn_tp_size
def forward( def forward(
self, self,
@@ -925,8 +925,8 @@ class VisionAttention(nn.Module):
DeprecationWarning, DeprecationWarning,
stacklevel=2, stacklevel=2,
) )
self.tp_size = 1 if use_data_parallel else get_attention_tp_size() self.tp_size = 1 if use_data_parallel else get_parallel().attn_tp_size
self.tp_rank = 0 if use_data_parallel else get_attention_tp_rank() self.tp_rank = 0 if use_data_parallel else get_parallel().attn_tp_rank
self.dropout = dropout self.dropout = dropout
num_kv_heads = num_kv_heads if num_kv_heads is not None else num_heads num_kv_heads = num_kv_heads if num_kv_heads is not None else num_heads
self.head_size = head_dim if head_dim is not None else embed_dim // num_heads self.head_size = head_dim if head_dim is not None else embed_dim // num_heads
@@ -2,12 +2,12 @@
import torch import torch
from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.runtime_context import get_parallel
def update_vit_attn_dummy_heads_config(config): def update_vit_attn_dummy_heads_config(config):
"""Update HF config to ensure vision attention num_attention_heads is divisible by tp_size""" """Update HF config to ensure vision attention num_attention_heads is divisible by tp_size"""
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
num_heads = getattr( num_heads = getattr(
config.vision_config, config.vision_config,
"num_heads", "num_heads",
@@ -12,8 +12,8 @@ from sglang.srt.layers.attention.triton_ops.kv_indices import (
create_flashinfer_kv_indices_triton, create_flashinfer_kv_indices_triton,
) )
from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton from sglang.srt.layers.attention.triton_ops.metadata import get_num_kv_splits_triton
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import get_bool_env_var, get_device_core_count from sglang.srt.utils import get_bool_env_var, get_device_core_count
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -95,10 +95,10 @@ class WaveAttnBackend(AttentionBackend):
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
) )
self.num_kv_head = model_runner.model_config.get_num_kv_heads( self.num_kv_head = model_runner.model_config.get_num_kv_heads(
get_attention_tp_size() get_parallel().attn_tp_size
) )
self.static_kv_splits = get_bool_env_var( self.static_kv_splits = get_bool_env_var(
+4 -4
View File
@@ -27,7 +27,6 @@ from typing import Optional, Tuple
import torch import torch
import torch.nn as nn import torch.nn as nn
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
MergedColumnParallelLinear, MergedColumnParallelLinear,
@@ -35,6 +34,7 @@ from sglang.srt.layers.linear import (
RowParallelLinear, RowParallelLinear,
) )
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
_INF = float("inf") _INF = float("inf")
@@ -128,7 +128,7 @@ class ClippableQKVParallelLinear(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
self.q_size = (total_num_heads // tp_size) * head_size self.q_size = (total_num_heads // tp_size) * head_size
self.kv_size = (total_num_kv_heads // tp_size) * head_size self.kv_size = (total_num_kv_heads // tp_size) * head_size
@@ -192,7 +192,7 @@ class ClippableGLUParallelLinear(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
self.proj_size = hidden_size // tp_size self.proj_size = hidden_size // tp_size
self.linear = MergedColumnParallelLinear( self.linear = MergedColumnParallelLinear(
@@ -255,7 +255,7 @@ class ClippableGateUpParallelLinear(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
self.proj_size = intermediate_size // tp_size self.proj_size = intermediate_size // tp_size
self.gate_up_proj = MergedColumnParallelLinear( self.gate_up_proj = MergedColumnParallelLinear(
+13 -19
View File
@@ -23,8 +23,6 @@ import torch
from sglang.srt.distributed import ( from sglang.srt.distributed import (
attention_tensor_model_parallel_all_reduce, attention_tensor_model_parallel_all_reduce,
attention_tensor_model_parallel_quant_all_reduce, attention_tensor_model_parallel_quant_all_reduce,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
get_tp_group, get_tp_group,
moe_tensor_model_parallel_all_reduce, moe_tensor_model_parallel_all_reduce,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
@@ -44,12 +42,7 @@ from sglang.srt.layers.dp_attention import (
dp_gather_replicate, dp_gather_replicate,
dp_reduce_scatter_tensor, dp_reduce_scatter_tensor,
dp_scatter, dp_scatter,
get_attention_cp_rank,
get_attention_cp_size,
get_attention_dp_size,
get_attention_tp_group, get_attention_tp_group,
get_attention_tp_rank,
get_attention_tp_size,
get_dp_global_num_tokens, get_dp_global_num_tokens,
get_global_dp_buffer, get_global_dp_buffer,
get_local_dp_buffer, get_local_dp_buffer,
@@ -77,6 +70,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -189,7 +183,7 @@ def apply_aiter_all_reduce_fusion(input_tensor: torch.Tensor):
and total_bytes > 0 and total_bytes > 0
and n <= 16384 and n <= 16384
and total_bytes <= 8 * 1024 * 8192 and total_bytes <= 8 * 1024 * 8192
and get_tensor_model_parallel_world_size() != 6 and get_parallel().tp_size != 6
and not is_dp_attention_enabled() and not is_dp_attention_enabled()
and get_global_server_args().enable_aiter_allreduce_fusion and get_global_server_args().enable_aiter_allreduce_fusion
) )
@@ -276,7 +270,7 @@ class AttnTpContext:
and (_is_cuda or _is_npu) and (_is_cuda or _is_npu)
and q_lora_rank is not None and q_lora_rank is not None
and not is_dsa and not is_dsa
and get_tensor_model_parallel_world_size() > 1 and get_parallel().tp_size > 1
and not is_dp_attention_enabled() and not is_dp_attention_enabled()
and get_moe_a2a_backend().is_none() and get_moe_a2a_backend().is_none()
and not enable_moe_dense_fully_dp() and not enable_moe_dense_fully_dp()
@@ -801,7 +795,7 @@ class LayerCommunicator:
or ( or (
_use_aiter _use_aiter
and batch_size > 0 and batch_size > 0
and get_tensor_model_parallel_world_size() != 6 and get_parallel().tp_size != 6
and get_global_server_args().enable_aiter_allreduce_fusion and get_global_server_args().enable_aiter_allreduce_fusion
) )
) )
@@ -828,13 +822,13 @@ class CommunicateContext:
@classmethod @classmethod
def init_new(cls): def init_new(cls):
attn_tp_rank = get_attention_tp_rank() attn_tp_rank = get_parallel().attn_tp_rank
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
attn_dp_size = get_attention_dp_size() attn_dp_size = get_parallel().attn_dp_size
attn_cp_size = get_attention_cp_size() attn_cp_size = get_parallel().attn_cp_size
attn_cp_rank = get_attention_cp_rank() attn_cp_rank = get_parallel().attn_cp_rank
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
moe_cp_size = get_moe_cp_size() moe_cp_size = get_moe_cp_size()
process_group_sizes = { process_group_sizes = {
ScatterMode.SCATTERED: 1, ScatterMode.SCATTERED: 1,
@@ -1312,7 +1306,7 @@ class CommunicateSummableTensorPairFn:
context: CommunicateContext, context: CommunicateContext,
allow_reduce_scatter: bool = False, allow_reduce_scatter: bool = False,
): ):
if get_tensor_model_parallel_world_size() == get_attention_dp_size(): if get_parallel().tp_size == get_parallel().attn_dp_size:
group = get_tp_group() group = get_tp_group()
else: else:
group = get_attention_tp_group() group = get_attention_tp_group()
@@ -1402,7 +1396,7 @@ class CommunicateSummableTensorPairFn:
# DP scatter (if DP attention is enabled) # DP scatter (if DP attention is enabled)
if context.attn_dp_size > 1: if context.attn_dp_size > 1:
if get_tensor_model_parallel_world_size() == get_attention_dp_size(): if get_parallel().tp_size == get_parallel().attn_dp_size:
group = get_tp_group() group = get_tp_group()
else: else:
group = get_attention_tp_group() group = get_attention_tp_group()
@@ -35,14 +35,11 @@ from sglang.srt.layers.dp_attention import (
attn_cp_all_gather_into_tensor, attn_cp_all_gather_into_tensor,
attn_cp_reduce_scatter_tensor, attn_cp_reduce_scatter_tensor,
get_attention_cp_group, get_attention_cp_group,
get_attention_cp_rank,
get_attention_cp_size,
get_attention_dp_size,
get_attention_tp_size,
get_local_dp_buffer, get_local_dp_buffer,
) )
from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_parallel
def dsa_enable_prefill_cp(): def dsa_enable_prefill_cp():
@@ -53,8 +50,8 @@ def dsa_enable_prefill_cp():
def dsa_cp_gather_hidden_states(hidden_states: torch.Tensor): def dsa_cp_gather_hidden_states(hidden_states: torch.Tensor):
attn_dp_size = get_attention_dp_size() attn_dp_size = get_parallel().attn_dp_size
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
assert attn_dp_size == 1 and attn_tp_size == 1 assert attn_dp_size == 1 and attn_tp_size == 1
hidden_states, local_hidden_states = ( hidden_states, local_hidden_states = (
get_local_dp_buffer(get_attention_cp_group()), get_local_dp_buffer(get_attention_cp_group()),
@@ -65,11 +62,11 @@ def dsa_cp_gather_hidden_states(hidden_states: torch.Tensor):
def dsa_cp_reduce_scatter_hidden_states(hidden_states: torch.Tensor): def dsa_cp_reduce_scatter_hidden_states(hidden_states: torch.Tensor):
attn_dp_size = get_attention_dp_size() attn_dp_size = get_parallel().attn_dp_size
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
assert attn_dp_size == 1 and attn_tp_size == 1 assert attn_dp_size == 1 and attn_tp_size == 1
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
cp_rank = get_attention_cp_rank() cp_rank = get_parallel().attn_cp_rank
input_hidden_states = hidden_states input_hidden_states = hidden_states
hidden_states = hidden_states.tensor_split(cp_size)[cp_rank] hidden_states = hidden_states.tensor_split(cp_size)[cp_rank]
attn_cp_reduce_scatter_tensor(hidden_states, input_hidden_states) attn_cp_reduce_scatter_tensor(hidden_states, input_hidden_states)
+3 -2
View File
@@ -29,6 +29,8 @@ from dataclasses import dataclass
from enum import IntEnum from enum import IntEnum
from typing import TYPE_CHECKING, Any, Callable, List, Optional, Tuple from typing import TYPE_CHECKING, Any, Callable, List, Optional, Tuple
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -93,9 +95,8 @@ class ContextParallelStrategy(ABC):
@property @property
def cp_rank(self) -> int: def cp_rank(self) -> int:
from sglang.srt.layers.dp_attention import get_attention_cp_rank
return get_attention_cp_rank() return get_parallel().attn_cp_rank
@property @property
def per_layer_attn_cp_comm(self) -> bool: def per_layer_attn_cp_comm(self) -> bool:
@@ -7,18 +7,13 @@ import torch.distributed as dist
from torch.distributed import ProcessGroup from torch.distributed import ProcessGroup
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_attn_tensor_model_parallel_rank,
get_attn_tensor_model_parallel_world_size,
get_attn_tp_group, get_attn_tp_group,
get_moe_ep_group, get_moe_ep_group,
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_rank,
get_moe_tensor_parallel_world_size,
get_moe_tp_group, get_moe_tp_group,
get_tp_group, get_tp_group,
) )
from sglang.srt.distributed.parallel_state import in_the_same_node_as from sglang.srt.distributed.parallel_state import in_the_same_node_as
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
ceil_align, ceil_align,
@@ -637,17 +632,17 @@ def ensure_workspace_initialized(
return False return False
if use_attn_tp_group: if use_attn_tp_group:
world_size = get_attn_tensor_model_parallel_world_size() world_size = get_parallel().attn_tp_size
rank = get_attn_tensor_model_parallel_rank() rank = get_parallel().attn_tp_rank
coordinator = get_attn_tp_group() coordinator = get_attn_tp_group()
else: else:
if get_moe_expert_parallel_world_size() > 1: if get_parallel().moe_ep_size > 1:
world_size = get_moe_expert_parallel_world_size() world_size = get_parallel().moe_ep_size
rank = get_moe_expert_parallel_rank() rank = get_parallel().moe_ep_rank
coordinator = get_moe_ep_group() coordinator = get_moe_ep_group()
else: else:
world_size = get_moe_tensor_parallel_world_size() world_size = get_parallel().moe_tp_size
rank = get_moe_tensor_parallel_rank() rank = get_parallel().moe_tp_rank
coordinator = get_moe_tp_group() coordinator = get_moe_tp_group()
# Always pass the coordinator's groups: flashinfer >=0.6.10 reads the # Always pass the coordinator's groups: flashinfer >=0.6.10 reads the
@@ -757,12 +752,12 @@ def flashinfer_allreduce_residual_rmsnorm(
return None, None return None, None
if use_attn_tp_group: if use_attn_tp_group:
world_size = get_attn_tensor_model_parallel_world_size() world_size = get_parallel().attn_tp_size
else: else:
if get_moe_expert_parallel_world_size() > 1: if get_parallel().moe_ep_size > 1:
world_size = get_moe_expert_parallel_world_size() world_size = get_parallel().moe_ep_size
else: else:
world_size = get_moe_tensor_parallel_world_size() world_size = get_parallel().moe_tp_size
if world_size <= 1: if world_size <= 1:
logger.debug("Single GPU, no need for allreduce fusion") logger.debug("Single GPU, no need for allreduce fusion")
+5 -7
View File
@@ -31,6 +31,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
Phase, Phase,
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
@@ -153,9 +154,6 @@ def _forward_with_allreduce_fusion(
"""Shared allreduce-fused RMSNorm logic usable by any norm.""" """Shared allreduce-fused RMSNorm logic usable by any norm."""
if residual is not None: if residual is not None:
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_attn_tensor_model_parallel_world_size,
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
tensor_model_parallel_fused_allreduce_rmsnorm, tensor_model_parallel_fused_allreduce_rmsnorm,
) )
@@ -164,12 +162,12 @@ def _forward_with_allreduce_fusion(
) )
if use_attn_tp_group: if use_attn_tp_group:
world_size = get_attn_tensor_model_parallel_world_size() world_size = get_parallel().attn_tp_size
else: else:
if get_moe_expert_parallel_world_size() > 1: if get_parallel().moe_ep_size > 1:
world_size = get_moe_expert_parallel_world_size() world_size = get_parallel().moe_ep_size
else: else:
world_size = get_moe_tensor_parallel_world_size() world_size = get_parallel().moe_tp_size
if world_size > 1: if world_size > 1:
if post_residual_addition is not None: if post_residual_addition is not None:
+13 -14
View File
@@ -15,8 +15,6 @@ from torch.nn.parameter import Parameter, UninitializedParameter
from sglang.kernel_api_logging import wrap_method_with_debug_kernel_once from sglang.kernel_api_logging import wrap_method_with_debug_kernel_once
from sglang.srt.distributed import ( from sglang.srt.distributed import (
divide, divide,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
get_tp_group, get_tp_group,
split_tensor_along_last_dim, split_tensor_along_last_dim,
tensor_model_parallel_all_gather, tensor_model_parallel_all_gather,
@@ -40,6 +38,7 @@ from sglang.srt.layers.parameter import (
_ColumnvLLMParameter, _ColumnvLLMParameter,
) )
from sglang.srt.layers.utils import pad_or_narrow_weight from sglang.srt.layers.utils import pad_or_narrow_weight
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs
@@ -338,9 +337,9 @@ class ColumnParallelLinear(LinearBase):
# Divide the weight matrix along the last dimension. # Divide the weight matrix along the last dimension.
if tp_rank is None: if tp_rank is None:
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
if tp_size is None: if tp_size is None:
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.tp_rank, self.tp_size = tp_rank, tp_size self.tp_rank, self.tp_size = tp_rank, tp_size
assert self.quant_method is not None assert self.quant_method is not None
self.output_size_per_partition = divide(self.output_size, tp_size) self.output_size_per_partition = divide(self.output_size, tp_size)
@@ -526,9 +525,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
): ):
self.output_sizes = output_sizes self.output_sizes = output_sizes
if tp_rank is None: if tp_rank is None:
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
if tp_size is None: if tp_size is None:
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.tp_rank, self.tp_size = tp_rank, tp_size self.tp_rank, self.tp_size = tp_rank, tp_size
assert all(output_size % tp_size == 0 for output_size in output_sizes) assert all(output_size % tp_size == 0 for output_size in output_sizes)
self.use_presharded_weights = use_presharded_weights self.use_presharded_weights = use_presharded_weights
@@ -943,9 +942,9 @@ class QKVParallelLinear(ColumnParallelLinear):
self.total_num_kv_heads = total_num_kv_heads self.total_num_kv_heads = total_num_kv_heads
# Divide the weight matrix along the last dimension. # Divide the weight matrix along the last dimension.
if tp_rank is None: if tp_rank is None:
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
if tp_size is None: if tp_size is None:
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.tp_rank, self.tp_size = tp_rank, tp_size self.tp_rank, self.tp_size = tp_rank, tp_size
self.num_heads = divide(self.total_num_heads, tp_size) self.num_heads = divide(self.total_num_heads, tp_size)
if tp_size >= self.total_num_kv_heads: if tp_size >= self.total_num_kv_heads:
@@ -1390,9 +1389,9 @@ class RowParallelLinear(LinearBase):
# Divide the weight matrix along the last dimension. # Divide the weight matrix along the last dimension.
if tp_rank is None: if tp_rank is None:
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
if tp_size is None: if tp_size is None:
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.tp_rank, self.tp_size = tp_rank, tp_size self.tp_rank, self.tp_size = tp_rank, tp_size
self.input_size_per_partition = divide(input_size, self.tp_size) self.input_size_per_partition = divide(input_size, self.tp_size)
assert self.quant_method is not None assert self.quant_method is not None
@@ -1605,8 +1604,8 @@ class MergedColumnParallelRepeatedLinear(LinearBase):
prefix=prefix, prefix=prefix,
) )
self.num_column_parallel = len(column_output_sizes) self.num_column_parallel = len(column_output_sizes)
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_parallel().tp_rank
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.output_partition_sizes = [ self.output_partition_sizes = [
divide(x, self.tp_size) for x in column_output_sizes divide(x, self.tp_size) for x in column_output_sizes
@@ -1657,8 +1656,8 @@ class ColumnParallelBatchedLinear(nn.Module):
self, batch: int, input_size: int, output_size: int, dtype: torch.dtype self, batch: int, input_size: int, output_size: int, dtype: torch.dtype
): ):
super().__init__() super().__init__()
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_parallel().tp_rank
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.weight = nn.Parameter( self.weight = nn.Parameter(
torch.empty(batch, output_size // self.tp_size, input_size, dtype=dtype), torch.empty(batch, output_size // self.tp_size, input_size, dtype=dtype),
requires_grad=False, requires_grad=False,
+5 -8
View File
@@ -22,7 +22,6 @@ import torch
from torch import nn from torch import nn
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_gather, tensor_model_parallel_all_gather,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -32,9 +31,6 @@ from sglang.srt.layers.dp_attention import (
attn_tp_all_gather_into_tensor, attn_tp_all_gather_into_tensor,
dp_gather_replicate, dp_gather_replicate,
dp_scatter, dp_scatter,
get_attention_dp_rank,
get_attention_dp_size,
get_attention_tp_size,
get_dp_device, get_dp_device,
get_dp_dtype, get_dp_dtype,
get_dp_hidden_size, get_dp_hidden_size,
@@ -53,6 +49,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
is_cpu, is_cpu,
@@ -229,7 +226,7 @@ class LogitsMetadata:
def compute_dp_attention_metadata(self): def compute_dp_attention_metadata(self):
cumtokens = torch.cumsum(self.global_num_tokens_for_logprob_gpu, dim=0) cumtokens = torch.cumsum(self.global_num_tokens_for_logprob_gpu, dim=0)
dp_rank = get_attention_dp_rank() dp_rank = get_parallel().attn_dp_rank
if dp_rank == 0: if dp_rank == 0:
dp_local_start_pos = torch.zeros_like( dp_local_start_pos = torch.zeros_like(
self.global_num_tokens_for_logprob_gpu[0] self.global_num_tokens_for_logprob_gpu[0]
@@ -275,17 +272,17 @@ class LogitsProcessor(nn.Module):
self.use_attn_tp_group = get_global_server_args().enable_dp_lm_head self.use_attn_tp_group = get_global_server_args().enable_dp_lm_head
self.use_fp32_lm_head = get_global_server_args().enable_fp32_lm_head self.use_fp32_lm_head = get_global_server_args().enable_fp32_lm_head
if self.use_attn_tp_group: if self.use_attn_tp_group:
self.attn_tp_size = get_attention_tp_size() self.attn_tp_size = get_parallel().attn_tp_size
self.do_tensor_parallel_all_gather = ( self.do_tensor_parallel_all_gather = (
not skip_all_gather and self.attn_tp_size > 1 not skip_all_gather and self.attn_tp_size > 1
) )
self.do_tensor_parallel_all_gather_dp_attn = False self.do_tensor_parallel_all_gather_dp_attn = False
else: else:
self.do_tensor_parallel_all_gather = ( self.do_tensor_parallel_all_gather = (
not skip_all_gather and get_tensor_model_parallel_world_size() > 1 not skip_all_gather and get_parallel().tp_size > 1
) )
self.do_tensor_parallel_all_gather_dp_attn = ( self.do_tensor_parallel_all_gather_dp_attn = (
self.do_tensor_parallel_all_gather and get_attention_dp_size() != 1 self.do_tensor_parallel_all_gather and get_parallel().attn_dp_size != 1
) )
self.final_logit_softcapping = getattr( self.final_logit_softcapping = getattr(
self.config, "final_logit_softcapping", None self.config, "final_logit_softcapping", None
@@ -5,6 +5,7 @@ from typing import Optional
import torch import torch
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_cuda, is_cuda_alike from sglang.srt.utils import is_cuda, is_cuda_alike
_is_cuda = is_cuda() _is_cuda = is_cuda()
@@ -22,7 +23,6 @@ else:
from sgl_kernel import silu_and_mul from sgl_kernel import silu_and_mul
from sglang.jit_kernel.per_tensor_quant_fp8 import per_tensor_quant_fp8 from sglang.jit_kernel.per_tensor_quant_fp8 import per_tensor_quant_fp8
from sglang.srt.distributed import get_moe_expert_parallel_world_size
from sglang.srt.layers.moe.ep_moe.kernels import ( from sglang.srt.layers.moe.ep_moe.kernels import (
cutlass_w4_run_moe_ep_preproess, cutlass_w4_run_moe_ep_preproess,
deepep_ll_get_cutlass_w4a8_moe_mm_data, deepep_ll_get_cutlass_w4a8_moe_mm_data,
@@ -124,7 +124,7 @@ def cutlass_w4a8_moe(
assert topk == 1, "apply_router_weight_on_input is only implemented for topk=1" assert topk == 1, "apply_router_weight_on_input is only implemented for topk=1"
device = a.device device = a.device
if get_moe_expert_parallel_world_size() > 1: if get_parallel().moe_ep_size > 1:
topk_ids = torch.where(topk_ids == -1, num_local_experts, topk_ids) topk_ids = torch.where(topk_ids == -1, num_local_experts, topk_ids)
src2dst = cutlass_w4_run_moe_ep_preproess( src2dst = cutlass_w4_run_moe_ep_preproess(
@@ -13,10 +13,6 @@ from torch.nn.parameter import UninitializedParameter
from sglang.srt.batch_overlap.single_batch_overlap import DownGemmOverlapArgs from sglang.srt.batch_overlap.single_batch_overlap import DownGemmOverlapArgs
from sglang.srt.batch_overlap.two_batch_overlap import MaybeTboDeepEPDispatcher from sglang.srt.batch_overlap.two_batch_overlap import MaybeTboDeepEPDispatcher
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_rank,
get_moe_tensor_parallel_world_size,
get_tp_group, get_tp_group,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
@@ -65,6 +61,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
@@ -196,10 +193,10 @@ class FusedMoE(torch.nn.Module):
self.enable_flashinfer_cutlass_moe = ( self.enable_flashinfer_cutlass_moe = (
get_moe_runner_backend().is_flashinfer_cutlass() get_moe_runner_backend().is_flashinfer_cutlass()
) )
self.moe_ep_size = get_moe_expert_parallel_world_size() self.moe_ep_size = get_parallel().moe_ep_size
self.moe_ep_rank = get_moe_expert_parallel_rank() self.moe_ep_rank = get_parallel().moe_ep_rank
self.moe_tp_size = get_moe_tensor_parallel_world_size() self.moe_tp_size = get_parallel().moe_tp_size
self.moe_tp_rank = get_moe_tensor_parallel_rank() self.moe_tp_rank = get_parallel().moe_tp_rank
# DeepEP: each rank has its own shared expert slot, so total shared # DeepEP: each rank has its own shared expert slot, so total shared
# weight slots = num_fused_shared_experts * ep_size. # weight slots = num_fused_shared_experts * ep_size.
@@ -12,8 +12,8 @@ from typing import TYPE_CHECKING, Optional
import torch import torch
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import get_compiler_backend from sglang.srt.utils import get_compiler_backend
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -154,7 +154,7 @@ class KTEPWrapperMethod(FusedMoEMethodBase):
self.num_gpu_experts = kt_config.num_gpu_experts self.num_gpu_experts = kt_config.num_gpu_experts
self.override_num_local_experts = True self.override_num_local_experts = True
self.gpu_method.num_gpu_experts = self.num_gpu_experts self.gpu_method.num_gpu_experts = self.num_gpu_experts
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_parallel().tp_rank
# KT wrapper will be initialized in create_weights # KT wrapper will be initialized in create_weights
self.wrapper: Optional[KTMoEWrapper] = None self.wrapper: Optional[KTMoEWrapper] = None
@@ -20,6 +20,7 @@ from sglang.srt.layers.moe.utils import (
DeepEPMode, DeepEPMode,
is_tbo_enabled, is_tbo_enabled,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import ( from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
get_int_env_var, get_int_env_var,
@@ -35,10 +36,6 @@ from functools import lru_cache
import torch import torch
from sglang.srt.distributed import (
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
)
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
# Blockwise quantization group sizes: number of elements sharing one scale factor # Blockwise quantization group sizes: number of elements sharing one scale factor
@@ -217,8 +214,8 @@ def init_mori_op(
import mori import mori
world_size = get_moe_expert_parallel_world_size() world_size = get_parallel().moe_ep_size
rank = get_moe_expert_parallel_rank() rank = get_parallel().moe_ep_rank
gpu_per_node = 8 if world_size >= 8 else world_size gpu_per_node = 8 if world_size >= 8 else world_size
@@ -1048,7 +1045,7 @@ class MoriEPDispatcher(BaseDispatcher):
# experts that are not local to this rank. # experts that are not local to this rank.
self.expert_mask_gpu = None self.expert_mask_gpu = None
if _use_aiter and num_experts is not None and num_local_experts is not None: if _use_aiter and num_experts is not None and num_local_experts is not None:
ep_rank = get_moe_expert_parallel_rank() ep_rank = get_parallel().moe_ep_rank
expert_mask = torch.zeros( expert_mask = torch.zeros(
num_experts, num_experts,
device=torch.cuda.current_device(), device=torch.cuda.current_device(),
@@ -5,8 +5,6 @@ from typing import TYPE_CHECKING, NamedTuple, Optional
import torch import torch
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
get_tp_group, get_tp_group,
) )
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
@@ -30,6 +28,7 @@ from sglang.srt.layers.moe.utils import (
get_moe_runner_backend, get_moe_runner_backend,
should_use_flashinfer_cutlass_moe_fp4_allgather, should_use_flashinfer_cutlass_moe_fp4_allgather,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
get_bool_env_var, get_bool_env_var,
get_device, get_device,
@@ -88,7 +87,7 @@ class StandardDispatcher(BaseDispatcher):
def __init__(self, moe_runner_config: MoeRunnerConfig): def __init__(self, moe_runner_config: MoeRunnerConfig):
super().__init__() super().__init__()
self.moe_ep_size = get_moe_expert_parallel_world_size() self.moe_ep_size = get_parallel().moe_ep_size
backend = get_moe_runner_backend() backend = get_moe_runner_backend()
self.enable_flashinfer_cutlass_moe = backend.is_flashinfer_cutlass() self.enable_flashinfer_cutlass_moe = backend.is_flashinfer_cutlass()
self.enable_flashinfer_mxfp4_moe = backend.is_flashinfer_mxfp4() self.enable_flashinfer_mxfp4_moe = backend.is_flashinfer_mxfp4()
@@ -110,7 +109,7 @@ class StandardDispatcher(BaseDispatcher):
self.num_local_routed_experts = ( self.num_local_routed_experts = (
self.num_local_experts - self.num_local_shared_experts self.num_local_experts - self.num_local_shared_experts
) )
self.moe_ep_rank = get_moe_expert_parallel_rank() self.moe_ep_rank = get_parallel().moe_ep_rank
self.local_expert_mapping = None self.local_expert_mapping = None
self.expert_mask_gpu = None self.expert_mask_gpu = None
+4 -4
View File
@@ -31,6 +31,8 @@ from typing import (
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
from sglang.srt.runtime_context import get_parallel
try: try:
from triton_kernels.matmul_ogs import GatherIndx, RoutingData, ScatterIndx from triton_kernels.matmul_ogs import GatherIndx, RoutingData, ScatterIndx
from triton_kernels.tensor import make_ragged_tensor_metadata from triton_kernels.tensor import make_ragged_tensor_metadata
@@ -81,8 +83,6 @@ except ImportError:
from sglang.jit_kernel.dsv4 import mask_topk_ids from sglang.jit_kernel.dsv4 import mask_topk_ids
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
get_tp_group, get_tp_group,
) )
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
@@ -1484,8 +1484,8 @@ def _remap_topk_for_deepep(
if topk_ids.shape[0] == 0: if topk_ids.shape[0] == 0:
return topk_ids, topk_weights return topk_ids, topk_weights
ep_size = get_moe_expert_parallel_world_size() ep_size = get_parallel().moe_ep_size
ep_rank = get_moe_expert_parallel_rank() ep_rank = get_parallel().moe_ep_rank
# Static EPLB may add redundant physical experts. At this point routed # Static EPLB may add redundant physical experts. At this point routed
# topk_ids have already been remapped from logical to physical ids, so the # topk_ids have already been remapped from logical to physical ids, so the
# DeepEP interleaved layout must use the physical routed count. # DeepEP interleaved layout must use the physical routed count.
+4 -5
View File
@@ -8,12 +8,11 @@ from typing import TYPE_CHECKING, Optional
import torch import torch
from sglang.srt.distributed.parallel_state import get_moe_expert_parallel_world_size
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_dp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_cuda, is_npu from sglang.srt.utils import is_cuda, is_npu
_is_npu = is_npu() _is_npu = is_npu()
@@ -409,7 +408,7 @@ def should_use_flashinfer_cutlass_moe_fp4_allgather():
and get_moe_runner_backend().is_flashinfer_cutlass() and get_moe_runner_backend().is_flashinfer_cutlass()
and is_dp_attention_enabled() and is_dp_attention_enabled()
and MOE_QUANTIZATION == "modelopt_fp4" and MOE_QUANTIZATION == "modelopt_fp4"
and get_moe_expert_parallel_world_size() == get_attention_dp_size() and get_parallel().moe_ep_size == get_parallel().attn_dp_size
) )
@@ -423,8 +422,8 @@ def should_use_dp_reduce_scatterv():
not should_use_flashinfer_cutlass_moe_fp4_allgather() not should_use_flashinfer_cutlass_moe_fp4_allgather()
and get_moe_a2a_backend().is_none() and get_moe_a2a_backend().is_none()
and is_dp_attention_enabled() and is_dp_attention_enabled()
and get_attention_dp_size() > 1 and get_parallel().attn_dp_size > 1
and get_moe_expert_parallel_world_size() == get_attention_dp_size() and get_parallel().moe_ep_size == get_parallel().attn_dp_size
) )
@@ -10,7 +10,6 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional
import torch import torch
from torch.nn import Module from torch.nn import Module
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo
from sglang.srt.layers.parameter import BlockQuantScaleParameter, ModelWeightParameter from sglang.srt.layers.parameter import BlockQuantScaleParameter, ModelWeightParameter
@@ -23,6 +22,7 @@ from sglang.srt.layers.quantization.base_config import (
from sglang.srt.layers.quantization.int8_utils import apply_w8a8_block_int8_linear from sglang.srt.layers.quantization.int8_utils import apply_w8a8_block_int8_linear
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
from sglang.srt.layers.quantization.utils import is_layer_skipped from sglang.srt.layers.quantization.utils import is_layer_skipped
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import set_weight_attrs from sglang.srt.utils import set_weight_attrs
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -149,7 +149,7 @@ class BlockInt8LinearMethod(LinearMethodBase):
output_size_per_partition = sum(output_partition_sizes) output_size_per_partition = sum(output_partition_sizes)
weight_loader = extra_weight_attrs.get("weight_loader") weight_loader = extra_weight_attrs.get("weight_loader")
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
block_n, block_k = ( block_n, block_k = (
self.quant_config.weight_block_size[0], self.quant_config.weight_block_size[0],
@@ -271,7 +271,7 @@ class BlockInt8MoEMethod(FusedMoEMethodBase):
if self.quant_config.is_checkpoint_int8_serialized: if self.quant_config.is_checkpoint_int8_serialized:
params_dtype = torch.int8 params_dtype = torch.int8
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
block_n, block_k = ( block_n, block_k = (
self.quant_config.weight_block_size[0], self.quant_config.weight_block_size[0],
@@ -6,7 +6,7 @@ from typing import TYPE_CHECKING
import torch import torch
from compressed_tensors import CompressionFormat from compressed_tensors import CompressionFormat
from sglang.srt.distributed import get_moe_expert_parallel_rank, get_tp_group from sglang.srt.distributed import get_tp_group
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, use_symmetric_memory,
) )
@@ -17,6 +17,7 @@ from sglang.srt.layers.quantization.compressed_tensors.schemes import (
CompressedTensorsMoEScheme, CompressedTensorsMoEScheme,
) )
from sglang.srt.layers.quantization.utils import replace_parameter from sglang.srt.layers.quantization.utils import replace_parameter
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_flashinfer_available, next_power_of_2, set_weight_attrs from sglang.srt.utils import is_flashinfer_available, next_power_of_2, set_weight_attrs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -65,7 +66,7 @@ class CompressedTensorsMxInt4MoE(CompressedTensorsMoEScheme):
assert ( assert (
not config.actorder not config.actorder
), "Actorder is not supported by flashinfer_trtllm backend" ), "Actorder is not supported by flashinfer_trtllm backend"
self.moe_ep_rank = get_moe_expert_parallel_rank() self.moe_ep_rank = get_parallel().moe_ep_rank
if self.quant_config.quant_format != CompressionFormat.pack_quantized.value: if self.quant_config.quant_format != CompressionFormat.pack_quantized.value:
raise ValueError( raise ValueError(
@@ -6,7 +6,6 @@ from typing import TYPE_CHECKING
import torch import torch
from compressed_tensors.quantization import QuantizationStrategy from compressed_tensors.quantization import QuantizationStrategy
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import ( from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
FlashInferTrtllmFp8MoeQuantInfo, FlashInferTrtllmFp8MoeQuantInfo,
@@ -27,6 +26,7 @@ from sglang.srt.layers.quantization.utils import (
per_tensor_dequantize, per_tensor_dequantize,
swap_w13_to_w31, swap_w13_to_w31,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import get_bool_env_var, is_hip, set_weight_attrs from sglang.srt.utils import get_bool_env_var, is_hip, set_weight_attrs
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -99,7 +99,7 @@ class CompressedTensorsW8A8Fp8MoE(CompressedTensorsMoEScheme):
if self.block_quant: if self.block_quant:
assert self.weight_block_size is not None assert self.weight_block_size is not None
layer.weight_block_size = self.weight_block_size layer.weight_block_size = self.weight_block_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
block_n, block_k = ( block_n, block_k = (
self.weight_block_size[0], self.weight_block_size[0],
self.weight_block_size[1], self.weight_block_size[1],
+4 -3
View File
@@ -12,7 +12,7 @@ import torch.nn.functional as F
from torch.nn import Module from torch.nn import Module
from torch.nn.parameter import Parameter from torch.nn.parameter import Parameter
from sglang.srt.distributed import get_tensor_model_parallel_world_size, get_tp_group from sglang.srt.distributed import get_tp_group
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, use_symmetric_memory,
) )
@@ -79,6 +79,7 @@ from sglang.srt.layers.quantization.utils import (
requantize_with_max_scale, requantize_with_max_scale,
) )
from sglang.srt.layers.utils import copy_or_rebind_param from sglang.srt.layers.utils import copy_or_rebind_param
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_bool_env_var, get_bool_env_var,
@@ -374,7 +375,7 @@ class Fp8LinearMethod(LinearMethodBase):
output_partition_sizes: List[int], output_partition_sizes: List[int],
skip_block_quant_check: bool = False, skip_block_quant_check: bool = False,
): ):
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
block_n, block_k = ( block_n, block_k = (
self.quant_config.weight_block_size[0], self.quant_config.weight_block_size[0],
self.quant_config.weight_block_size[1], self.quant_config.weight_block_size[1],
@@ -916,7 +917,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
if self.quant_config.is_checkpoint_fp8_serialized: if self.quant_config.is_checkpoint_fp8_serialized:
params_dtype = torch.uint32 if _use_hip_int4 else torch.float8_e4m3fn params_dtype = torch.uint32 if _use_hip_int4 else torch.float8_e4m3fn
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
w13_up_dim, w2_up_dim, weight_padded = get_moe_weight_sizes( w13_up_dim, w2_up_dim, weight_padded = get_moe_weight_sizes(
intermediate_size_per_partition, intermediate_size_per_partition,
@@ -13,6 +13,7 @@ from sglang.srt.layers.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8_row_padded, sglang_per_token_group_quant_fp8_row_padded,
) )
from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.common import torch_release from sglang.srt.utils.common import torch_release
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -1779,9 +1780,8 @@ def validate_fp8_block_shape(
block_size: list[int], block_size: list[int],
) -> None: ) -> None:
"""Validate block quantization shapes for tensor parallelism.""" """Validate block quantization shapes for tensor parallelism."""
from sglang.srt.distributed import get_tensor_model_parallel_world_size
tp_size = getattr(layer, "tp_size", get_tensor_model_parallel_world_size()) tp_size = getattr(layer, "tp_size", get_parallel().tp_size)
block_n, block_k = block_size[0], block_size[1] block_n, block_k = block_size[0], block_size[1]
# Required by row parallel # Required by row parallel
@@ -9,7 +9,6 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional
import numpy as np import numpy as np
import torch import torch
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.distributed.parallel_state import get_tp_group
from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo
@@ -24,6 +23,7 @@ from sglang.srt.layers.quantization.unquant import (
UnquantizedFusedMoEMethod, UnquantizedFusedMoEMethod,
UnquantizedLinearMethod, UnquantizedLinearMethod,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import get_device_capability, set_weight_attrs from sglang.srt.utils import get_device_capability, set_weight_attrs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -455,7 +455,7 @@ class MoeWNA16Method(FusedMoEMethodBase):
return return
device = get_tp_group().device device = get_tp_group().device
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
loaded_weight = loaded_weight.to(device) loaded_weight = loaded_weight.to(device)
shard_size = layer.intermediate_size_per_partition shard_size = layer.intermediate_size_per_partition
@@ -5,7 +5,6 @@ import torch
from tqdm import tqdm from tqdm import tqdm
from tqdm.std import EMA from tqdm.std import EMA
from sglang.srt.distributed import get_tensor_model_parallel_rank
from sglang.srt.layers.int4fp8_utils import ( from sglang.srt.layers.int4fp8_utils import (
pack_int4_to_int32, pack_int4_to_int32,
quantize_fp8_scale_tensorwise, quantize_fp8_scale_tensorwise,
@@ -18,6 +17,7 @@ from sglang.srt.layers.quantization.base_config import (
QuantizeMethodBase, QuantizeMethodBase,
) )
from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import BAR_FORMAT, is_hip, set_weight_attrs from sglang.srt.utils import BAR_FORMAT, is_hip, set_weight_attrs
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -71,7 +71,7 @@ class QuarkInt4Fp8Config(QuantizationConfig):
self.num_quant_layers = 0 self.num_quant_layers = 0
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
# The weight iterator already has a progress bar on rank=0, account for that. # The weight iterator already has a progress bar on rank=0, account for that.
position = 1 + tqdm._get_free_pos() position = 1 + tqdm._get_free_pos()
@@ -138,7 +138,7 @@ class QuarkInt4Fp8MoEMethod(FusedMoEMethodBase):
self.online_quant_progress_bar = self.quant_config.online_quant_progress_bar self.online_quant_progress_bar = self.quant_config.online_quant_progress_bar
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_parallel().tp_rank
if not _is_hip: if not _is_hip:
raise NotImplementedError( raise NotImplementedError(
@@ -7,7 +7,6 @@ from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional, cast
import torch import torch
from torch.nn.parameter import Parameter from torch.nn.parameter import Parameter
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.amx_utils import ( from sglang.srt.layers.amx_utils import (
CPUQuantMethod, CPUQuantMethod,
_amx_process_weight_after_loading, _amx_process_weight_after_loading,
@@ -24,6 +23,7 @@ from sglang.srt.layers.quantization.base_config import (
from sglang.srt.layers.quantization.compressed_tensors.utils import should_ignore_layer from sglang.srt.layers.quantization.compressed_tensors.utils import should_ignore_layer
from sglang.srt.layers.quantization.int8_kernel import per_token_quant_int8 from sglang.srt.layers.quantization.int8_kernel import per_token_quant_int8
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
is_cpu, is_cpu,
@@ -260,7 +260,7 @@ class W8A8Int8MoEMethod(FusedMoEMethodBase):
): ):
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
# WEIGHTS # WEIGHTS
w13_weight = torch.nn.Parameter( w13_weight = torch.nn.Parameter(
+6 -7
View File
@@ -11,13 +11,12 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
attn_cp_all_gather_into_tensor, attn_cp_all_gather_into_tensor,
get_attention_cp_group, get_attention_cp_group,
get_attention_cp_rank,
get_attention_cp_size,
is_allocation_symmetric, is_allocation_symmetric,
) )
from sglang.srt.layers.moe import get_moe_a2a_backend from sglang.srt.layers.moe import get_moe_a2a_backend
from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
@@ -80,7 +79,7 @@ def get_cp_padding_align_size() -> int:
""" """
from sglang.srt.layers.attention.dsa.utils import is_dsa_prefill_cp_in_seq_split from sglang.srt.layers.attention.dsa.utils import is_dsa_prefill_cp_in_seq_split
attn_cp_size = get_attention_cp_size() attn_cp_size = get_parallel().attn_cp_size
if is_prefill_cp_in_seq_split() or is_dsa_prefill_cp_in_seq_split(): if is_prefill_cp_in_seq_split() or is_dsa_prefill_cp_in_seq_split():
return attn_cp_size * 2 return attn_cp_size * 2
return attn_cp_size return attn_cp_size
@@ -150,7 +149,7 @@ def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor):
) )
if is_dsa_prefill_cp_round_robin_split(): if is_dsa_prefill_cp_round_robin_split():
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
assert ( assert (
input_.shape[0] % cp_size == 0 input_.shape[0] % cp_size == 0
), f"Expect input shape 0 can divided by cp size, but got input shape {input_.shape}, cp size {cp_size}" ), f"Expect input shape 0 can divided by cp size, but got input shape {input_.shape}, cp size {cp_size}"
@@ -172,7 +171,7 @@ def cp_split_and_rebuild_position(forward_batch, positions: torch.Tensor):
) )
if is_dsa_prefill_cp_round_robin_split(): if is_dsa_prefill_cp_round_robin_split():
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
assert positions.shape[0] % cp_size == 0, ( assert positions.shape[0] % cp_size == 0, (
f"Expect positions shape 0 can divided by cp size, but got positions shape {positions.shape}, " f"Expect positions shape 0 can divided by cp size, but got positions shape {positions.shape}, "
f"cp size {cp_size}" f"cp size {cp_size}"
@@ -204,8 +203,8 @@ def cp_round_robin_input_ids(input_ids):
rank2: 2,10,18,... rank2: 2,10,18,...
... ...
""" """
cp_size = get_attention_cp_size() cp_size = get_parallel().attn_cp_size
cp_rank = get_attention_cp_rank() cp_rank = get_parallel().attn_cp_rank
if get_moe_a2a_backend().is_none(): if get_moe_a2a_backend().is_none():
input_ids = input_ids.reshape(-1, cp_size).T.flatten() input_ids = input_ids.reshape(-1, cp_size).T.flatten()
else: else:
@@ -11,8 +11,6 @@ from torch.nn.parameter import Parameter, UninitializedParameter
from sglang.srt.distributed import ( from sglang.srt.distributed import (
divide, divide,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
get_tp_group, get_tp_group,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
@@ -24,8 +22,6 @@ from sglang.srt.layers.amx_utils import PackWeightMethod
from sglang.srt.layers.communicator import get_attn_tp_context from sglang.srt.layers.communicator import get_attn_tp_context
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
attn_tp_all_reduce, attn_tp_all_reduce,
get_attention_tp_rank,
get_attention_tp_size,
is_allocation_symmetric, is_allocation_symmetric,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
@@ -36,6 +32,7 @@ from sglang.srt.layers.quantization.base_config import (
method_has_implemented_embedding, method_has_implemented_embedding,
) )
from sglang.srt.layers.quantization.unquant import UnquantizedEmbeddingMethod from sglang.srt.layers.quantization.unquant import UnquantizedEmbeddingMethod
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_compiler_backend, get_compiler_backend,
@@ -245,11 +242,11 @@ class VocabParallelEmbedding(torch.nn.Module):
self.use_attn_tp_group = use_attn_tp_group self.use_attn_tp_group = use_attn_tp_group
if self.enable_tp: if self.enable_tp:
if use_attn_tp_group: if use_attn_tp_group:
tp_rank = get_attention_tp_rank() tp_rank = get_parallel().attn_tp_rank
self.tp_size = get_attention_tp_size() self.tp_size = get_parallel().attn_tp_size
else: else:
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
else: else:
assert use_attn_tp_group is False assert use_attn_tp_group is False
tp_rank = 0 tp_rank = 0
@@ -36,19 +36,12 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Tuple, Union
import torch import torch
from sglang.srt.distributed.parallel_state import (
get_moe_expert_parallel_world_size,
get_tensor_model_parallel_world_size,
)
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.kv_canary.req_to_expected_token_ids_manager import ( from sglang.srt.kv_canary.req_to_expected_token_ids_manager import (
compute_req_all_ids_info, compute_req_all_ids_info,
) )
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
DpPaddingMode, DpPaddingMode,
get_attention_dp_rank,
get_attention_tp_rank,
get_attention_tp_size,
set_dp_buffer_len, set_dp_buffer_len,
set_is_extend_in_batch, set_is_extend_in_batch,
) )
@@ -56,6 +49,7 @@ from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
ForwardBatchDeepSeekMHAMixin, ForwardBatchDeepSeekMHAMixin,
) )
from sglang.srt.model_executor.triton_ops.position import compute_position_triton from sglang.srt.model_executor.triton_ops.position import compute_position_triton
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
is_cuda, is_cuda,
@@ -212,8 +206,8 @@ def compute_local_num_token_non_padded(
Converts a global count (across all TP ranks) to a local count for this rank. Converts a global count (across all TP ranks) to a local count for this rank.
The "global" scope is within the current DP rank; DP is handled via num_tokens_per_dp. The "global" scope is within the current DP rank; DP is handled via num_tokens_per_dp.
""" """
attn_tp_rank = get_attention_tp_rank() attn_tp_rank = get_parallel().attn_tp_rank
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
tokens_per_rank = num_tokens_per_dp // attn_tp_size tokens_per_rank = num_tokens_per_dp // attn_tp_size
return torch.clamp( return torch.clamp(
@@ -849,7 +843,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
"""Make num_token_non_padded local to this attention-TP rank.""" """Make num_token_non_padded local to this attention-TP rank."""
from sglang.srt.utils.common import require_mlp_tp_gather from sglang.srt.utils.common import require_mlp_tp_gather
dp_rank = get_attention_dp_rank() dp_rank = get_parallel().attn_dp_rank
assert self.global_num_tokens_cpu is not None assert self.global_num_tokens_cpu is not None
if require_mlp_tp_gather(server_args): if require_mlp_tp_gather(server_args):
@@ -1063,7 +1057,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
self._original_batch_size = self.batch_size self._original_batch_size = self.batch_size
global_num_tokens = self.global_num_tokens_cpu global_num_tokens = self.global_num_tokens_cpu
sync_group_size = len(global_num_tokens) sync_group_size = len(global_num_tokens)
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
for i in range(sync_group_size): for i in range(sync_group_size):
# make sure that the padded length is divisible by attn_tp_size because we may need reduce-scatter across attn_tp dim. # make sure that the padded length is divisible by attn_tp_size because we may need reduce-scatter across attn_tp dim.
@@ -1096,7 +1090,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
buffer_len = sum(global_num_tokens) buffer_len = sum(global_num_tokens)
if len(global_num_tokens) > 1: if len(global_num_tokens) > 1:
num_tokens = global_num_tokens[get_attention_dp_rank()] num_tokens = global_num_tokens[get_parallel().attn_dp_rank]
else: else:
num_tokens = global_num_tokens[0] num_tokens = global_num_tokens[0]
@@ -1299,7 +1293,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
return return
assert self.forward_mode.is_extend() assert self.forward_mode.is_extend()
tokens = self.input_ids.shape[0] tokens = self.input_ids.shape[0]
rank_size = get_tensor_model_parallel_world_size() rank_size = get_parallel().tp_size
tokens_padded = (tokens + rank_size - 1) // rank_size * rank_size tokens_padded = (tokens + rank_size - 1) // rank_size * rank_size
self._pad_inputs_to_size(model_runner, tokens_padded, self.batch_size) self._pad_inputs_to_size(model_runner, tokens_padded, self.batch_size)
@@ -1359,7 +1353,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
def enable_num_token_non_padded(): def enable_num_token_non_padded():
return get_moe_expert_parallel_world_size() > 1 return get_parallel().moe_ep_size > 1
def build_inner_fb_view( def build_inner_fb_view(
@@ -25,11 +25,7 @@ from typing import TYPE_CHECKING, Any, List, Sequence, Tuple
import torch import torch
from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin
from sglang.srt.layers.dp_attention import ( from sglang.srt.runtime_context import get_parallel
get_attention_cp_size,
get_attention_tp_rank,
get_attention_tp_size,
)
from sglang.srt.utils import require_gathered_buffer from sglang.srt.utils import require_gathered_buffer
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -81,10 +77,10 @@ def get_batch_sizes_to_capture(
num_tokens_per_bs = 1 num_tokens_per_bs = 1
if require_gathered_buffer(server_args): if require_gathered_buffer(server_args):
mul_base *= get_attention_tp_size() mul_base *= get_parallel().attn_tp_size
if mul_base % get_attention_cp_size() != 0: if mul_base % get_parallel().attn_cp_size != 0:
mul_base *= get_attention_cp_size() mul_base *= get_parallel().attn_cp_size
# pad `num_max_requests` to avoid being filtered out # pad `num_max_requests` to avoid being filtered out
num_max_requests = (num_max_requests + mul_base - 1) // mul_base * mul_base num_max_requests = (num_max_requests + mul_base - 1) // mul_base * mul_base
@@ -146,8 +142,8 @@ class BaseCudaGraphRunner(ABC):
self.tp_size = model_runner.server_args.tp_size self.tp_size = model_runner.server_args.tp_size
self.dp_size = model_runner.server_args.dp_size self.dp_size = model_runner.server_args.dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = model_runner.server_args.pp_size
self.attn_tp_size = get_attention_tp_size() self.attn_tp_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_attention_tp_rank() self.attn_tp_rank = get_parallel().attn_tp_rank
self.tbo_plugin = TboCudaGraphRunnerPlugin() self.tbo_plugin = TboCudaGraphRunnerPlugin()
@staticmethod @staticmethod
@@ -42,11 +42,8 @@ from tqdm.auto import tqdm
from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.load_config import LoadConfig
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
get_world_group, get_world_group,
) )
from sglang.srt.layers.dp_attention import get_attention_tp_rank
from sglang.srt.layers.quantization import QuantizationConfig, get_quantization_config from sglang.srt.layers.quantization import QuantizationConfig, get_quantization_config
from sglang.srt.layers.quantization.fp8 import Fp8Config from sglang.srt.layers.quantization.fp8 import Fp8Config
from sglang.srt.layers.quantization.modelopt_quant import ( from sglang.srt.layers.quantization.modelopt_quant import (
@@ -57,6 +54,7 @@ from sglang.srt.model_loader.ci_weight_validation import (
ci_download_with_validation_and_retry, ci_download_with_validation_and_retry,
ci_validate_and_cleanup_local_snapshot, ci_validate_and_cleanup_local_snapshot,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import ( from sglang.srt.utils import (
BAR_FORMAT, BAR_FORMAT,
find_local_repo_dir, find_local_repo_dir,
@@ -1327,7 +1325,7 @@ def row_parallel_weight_loader(
param: torch.Tensor, loaded_weight: torch.Tensor param: torch.Tensor, loaded_weight: torch.Tensor
) -> None: ) -> None:
"""Load weights that are row-parallelized.""" """Load weights that are row-parallelized."""
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
shard_dim = 0 if param.dim() != 1 else None shard_dim = 0 if param.dim() != 1 else None
if shard_dim is not None: if shard_dim is not None:
@@ -1345,7 +1343,7 @@ def sharded_weight_loader(shard_axis: int) -> LoaderFunction:
"""Create a weight loader that shards the weights along the given axis""" """Create a weight loader that shards the weights along the given axis"""
def loader(param: torch.Tensor, loaded_weight: torch.Tensor) -> None: def loader(param: torch.Tensor, loaded_weight: torch.Tensor) -> None:
tp_rank = get_attention_tp_rank() tp_rank = get_parallel().attn_tp_rank
shard_size = param.data.shape[shard_axis] shard_size = param.data.shape[shard_axis]
start_idx = tp_rank * shard_size start_idx = tp_rank * shard_size
@@ -1353,9 +1351,8 @@ def sharded_weight_loader(shard_axis: int) -> LoaderFunction:
if ( if (
is_cpu() is_cpu()
and ( and (
loaded_weight.size(0) % get_tensor_model_parallel_world_size() != 0 loaded_weight.size(0) % get_parallel().tp_size != 0
or loaded_weight.size(0) or loaded_weight.size(0) < get_parallel().tp_size * shard_size
< get_tensor_model_parallel_world_size() * shard_size
) )
and loaded_weight.dim() == 1 and loaded_weight.dim() == 1
): ):
+4 -5
View File
@@ -32,8 +32,6 @@ from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
@@ -58,6 +56,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
_is_npu = is_npu() _is_npu = is_npu()
@@ -160,8 +159,8 @@ class AfmoeMoE(nn.Module):
): ):
super().__init__() super().__init__()
self.config = config self.config = config
self.rank = get_tensor_model_parallel_rank() self.rank = get_parallel().tp_rank
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.n_routed_experts = getattr(config, "num_experts", None) self.n_routed_experts = getattr(config, "num_experts", None)
if self.n_routed_experts is None: if self.n_routed_experts is None:
@@ -309,7 +308,7 @@ class AfmoeAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
+4 -5
View File
@@ -28,8 +28,6 @@ from transformers import ApertusConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_pp_group, get_pp_group,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
) )
from sglang.srt.layers.activation import XIELU from sglang.srt.layers.activation import XIELU
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -54,6 +52,7 @@ from sglang.srt.model_loader.weight_utils import (
kv_cache_scales_loader, kv_cache_scales_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils import add_prefix, make_layers
@@ -130,7 +129,7 @@ class ApertusAttention(nn.Module):
super().__init__() super().__init__()
self.layer_id = layer_id self.layer_id = layer_id
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
@@ -383,8 +382,8 @@ class ApertusModel(nn.Module):
# factors (or else raise an exception). Thus, handled exceptions should # factors (or else raise an exception). Thus, handled exceptions should
# make sure to leave KV cache scale factors in a known good (dummy) state # make sure to leave KV cache scale factors in a known good (dummy) state
def load_kv_cache_scales(self, quantization_param_path: str) -> None: def load_kv_cache_scales(self, quantization_param_path: str) -> None:
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
for layer_idx, scaling_factor in kv_cache_scales_loader( for layer_idx, scaling_factor in kv_cache_scales_loader(
quantization_param_path, quantization_param_path,
tp_rank, tp_rank,
+4 -5
View File
@@ -22,8 +22,6 @@ from transformers import LlamaConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_pp_group, get_pp_group,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
) )
from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.activation import get_act_fn
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -48,6 +46,7 @@ from sglang.srt.model_loader.weight_utils import (
kv_cache_scales_loader, kv_cache_scales_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils import add_prefix, make_layers
@@ -119,7 +118,7 @@ class ArceeAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
@@ -348,8 +347,8 @@ class ArceeModel(nn.Module):
return hidden_states, aux_hidden_states return hidden_states, aux_hidden_states
def load_kv_cache_scales(self, quantization_param_path: str) -> None: def load_kv_cache_scales(self, quantization_param_path: str) -> None:
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
for layer_idx, scaling_factor in kv_cache_scales_loader( for layer_idx, scaling_factor in kv_cache_scales_loader(
quantization_param_path, quantization_param_path,
tp_rank, tp_rank,
+3 -6
View File
@@ -30,10 +30,6 @@ import torch
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -51,6 +47,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
from sglang.srt.utils.hf_transformers_utils import get_rope_config from sglang.srt.utils.hf_transformers_utils import get_rope_config
@@ -137,7 +134,7 @@ class BaiChuanAttention(nn.Module):
): ):
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
self.total_num_kv_heads = self.total_num_heads self.total_num_kv_heads = self.total_num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
@@ -187,7 +184,7 @@ class BaiChuanAttention(nn.Module):
# Create the alibi slopes and slice them. # Create the alibi slopes and slice them.
if self.position_embedding == "ALIBI": if self.position_embedding == "ALIBI":
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
head_start = tp_rank * self.num_heads head_start = tp_rank * self.num_heads
head_end = (tp_rank + 1) * self.num_heads head_end = (tp_rank + 1) * self.num_heads
alibi_slopes = _get_alibi_slopes(self.total_num_heads) alibi_slopes = _get_alibi_slopes(self.total_num_heads)
+9 -12
View File
@@ -29,7 +29,6 @@ from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_pp_group, get_pp_group,
get_tensor_model_parallel_world_size,
parallel_state, parallel_state,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
@@ -43,9 +42,6 @@ from sglang.srt.layers.communicator import (
enable_moe_dense_fully_dp, enable_moe_dense_fully_dp,
) )
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_dp_size,
get_attention_tp_rank,
get_attention_tp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -81,6 +77,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
@@ -188,7 +185,7 @@ class BailingMoESparseMoeBlock(nn.Module):
super().__init__() super().__init__()
self.layer_id = layer_id self.layer_id = layer_id
self.alt_stream = alt_stream self.alt_stream = alt_stream
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.top_k = config.num_experts_per_tok self.top_k = config.num_experts_per_tok
self.norm_topk_prob = config.norm_topk_prob self.norm_topk_prob = config.norm_topk_prob
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
@@ -290,7 +287,7 @@ class BailingMoESparseMoeBlock(nn.Module):
# dispatcher # dispatcher
if get_moe_a2a_backend().is_deepep(): if get_moe_a2a_backend().is_deepep():
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_tensor_model_parallel_world_size() self.ep_size = get_parallel().tp_size
self.deepep_dispatcher = DeepEPDispatcher( self.deepep_dispatcher = DeepEPDispatcher(
group=parallel_state.get_tp_group().device_group, group=parallel_state.get_tp_group().device_group,
@@ -434,9 +431,9 @@ class BailingMoEAttention(nn.Module):
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.total_num_heads = config.num_attention_heads self.total_num_heads = config.num_attention_heads
self.total_kv_heads = config.num_key_value_heads self.total_kv_heads = config.num_key_value_heads
self.dp_size = get_attention_dp_size() self.dp_size = get_parallel().attn_dp_size
attn_tp_rank = get_attention_tp_rank() attn_tp_rank = get_parallel().attn_tp_rank
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
assert self.total_num_heads % attn_tp_size == 0 assert self.total_num_heads % attn_tp_size == 0
if self.total_kv_heads >= attn_tp_size: if self.total_kv_heads >= attn_tp_size:
@@ -574,7 +571,7 @@ class BailingMoEBlock(nn.Module):
hidden_size = config.hidden_size hidden_size = config.hidden_size
self.input_layernorm = RMSNorm(hidden_size, eps=config.rms_norm_eps) self.input_layernorm = RMSNorm(hidden_size, eps=config.rms_norm_eps)
self.dp_size = get_attention_dp_size() self.dp_size = get_parallel().attn_dp_size
self.attention = BailingMoEAttention( self.attention = BailingMoEAttention(
config, config,
layer_id, layer_id,
@@ -584,8 +581,8 @@ class BailingMoEBlock(nn.Module):
alt_stream=alt_stream, alt_stream=alt_stream,
) )
self.layer_id = layer_id self.layer_id = layer_id
self.attn_tp_size = get_attention_tp_size() self.attn_tp_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_attention_tp_rank() self.attn_tp_rank = get_parallel().attn_tp_rank
self.is_layer_sparse = self._is_layer_sparse( self.is_layer_sparse = self._is_layer_sparse(
config, layer_id=layer_id, is_nextn=False config, layer_id=layer_id, is_nextn=False
+8 -11
View File
@@ -11,8 +11,6 @@ from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_pp_group, get_pp_group,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
@@ -22,8 +20,6 @@ from sglang.srt.layers.attention.fla.layernorm_gated import RMSNorm as RMSNormGa
from sglang.srt.layers.attention.fla.layernorm_gated import layernorm_fn from sglang.srt.layers.attention.fla.layernorm_gated import layernorm_fn
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -62,6 +58,7 @@ from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
from sglang.srt.models.utils import WeightsMapper from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
@@ -251,8 +248,8 @@ class BailingMoE(nn.Module):
self.alt_stream = alt_stream self.alt_stream = alt_stream
self.layer_id = layer_id self.layer_id = layer_id
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_parallel().tp_rank
self.top_k = config.num_experts_per_tok self.top_k = config.num_experts_per_tok
self.norm_expert_prob = getattr(config, "norm_topk_prob", False) self.norm_expert_prob = getattr(config, "norm_topk_prob", False)
@@ -406,8 +403,8 @@ class BailingGroupRMSNormGate(RMSNormGated):
param: torch.nn.Parameter, param: torch.nn.Parameter,
loaded_weight: torch.Tensor, loaded_weight: torch.Tensor,
) -> None: ) -> None:
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
tp_rank = get_attention_tp_rank() tp_rank = get_parallel().attn_tp_rank
shard_size = loaded_weight.shape[0] // tp_size shard_size = loaded_weight.shape[0] // tp_size
shard = slice(tp_rank * shard_size, (tp_rank + 1) * shard_size) shard = slice(tp_rank * shard_size, (tp_rank + 1) * shard_size)
param.data.copy_(loaded_weight[shard].contiguous()) param.data.copy_(loaded_weight[shard].contiguous())
@@ -437,8 +434,8 @@ class BailingMoELinearAttention(nn.Module):
self.hidden_inner_size = self.head_dim * self.total_num_heads self.hidden_inner_size = self.head_dim * self.total_num_heads
self.scaling = self.head_dim**-0.5 self.scaling = self.head_dim**-0.5
self.tp_size = get_attention_tp_size() self.tp_size = get_parallel().attn_tp_size
self.tp_rank = get_attention_tp_rank() self.tp_rank = get_parallel().attn_tp_rank
assert self.total_num_heads % self.tp_size == 0 assert self.total_num_heads % self.tp_size == 0
self.tp_heads = self.total_num_heads // self.tp_size self.tp_heads = self.total_num_heads // self.tp_size
@@ -642,7 +639,7 @@ class BailingMoEAttention(nn.Module):
self.layer_id = layer_id self.layer_id = layer_id
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
self.total_num_heads = config.num_attention_heads self.total_num_heads = config.num_attention_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
@@ -26,7 +26,6 @@ import torch
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.dp_attention import is_dp_attention_enabled
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.linear import ReplicatedLinear
@@ -43,6 +42,7 @@ from sglang.srt.models.bailing_moe_linear import (
BailingMoeV2_5ForCausalLM, BailingMoeV2_5ForCausalLM,
) )
from sglang.srt.models.utils import WeightsMapper from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import BumpAllocator, add_prefix from sglang.srt.utils import BumpAllocator, add_prefix
@@ -195,7 +195,7 @@ class BailingMoeForCausalLMNextN(nn.Module):
) -> None: ) -> None:
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
if hasattr(self, "determine_num_fused_shared_experts"): if hasattr(self, "determine_num_fused_shared_experts"):
# Asystem has determine_num_fused_shared_experts but theta does not. # Asystem has determine_num_fused_shared_experts but theta does not.
+2 -2
View File
@@ -4,7 +4,6 @@ from typing import Iterable, Optional, Set, Tuple
import torch import torch
from torch import nn from torch import nn
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.activation import get_act_fn
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
@@ -17,6 +16,7 @@ from sglang.srt.layers.radix_attention import AttentionType, RadixAttention
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -220,7 +220,7 @@ class BertSelfAttention(nn.Module):
): ):
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_attention_heads self.total_num_heads = num_attention_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
+2 -2
View File
@@ -23,7 +23,6 @@ from torch import nn
from torch.nn import LayerNorm from torch.nn import LayerNorm
from sglang.srt.configs import ChatGLMConfig from sglang.srt.configs import ChatGLMConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -41,6 +40,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
LoraConfig = None LoraConfig = None
@@ -56,7 +56,7 @@ class GLMAttention(nn.Module):
): ):
super().__init__() super().__init__()
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = config.num_attention_heads self.total_num_heads = config.num_attention_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
+4 -4
View File
@@ -11,7 +11,6 @@ from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
@@ -32,6 +31,7 @@ from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, get_compiler_backend, is_cuda, make_layers from sglang.srt.utils import add_prefix, get_compiler_backend, is_cuda, make_layers
@@ -117,7 +117,7 @@ class Cohere2MoeAttention(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.config = config self.config = config
self.layer_id = layer_id self.layer_id = layer_id
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
@@ -232,7 +232,7 @@ class Cohere2MoeSparseMoeBlock(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.num_experts = config.num_experts self.num_experts = config.num_experts
self.top_k = config.num_experts_per_tok self.top_k = config.num_experts_per_tok
@@ -403,7 +403,7 @@ class Cohere2MoeDecoderLayer(nn.Module):
norm_eps = getattr(config, "layer_norm_eps", 1e-5) norm_eps = getattr(config, "layer_norm_eps", 1e-5)
self.input_layernorm = Cohere2MoeLayerNorm(config.hidden_size, eps=norm_eps) self.input_layernorm = Cohere2MoeLayerNorm(config.hidden_size, eps=norm_eps)
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
def forward( def forward(
self, self,
+3 -6
View File
@@ -49,10 +49,6 @@ from torch import nn
from torch.nn.parameter import Parameter from torch.nn.parameter import Parameter
from transformers import Cohere2Config, CohereConfig, PretrainedConfig from transformers import Cohere2Config, CohereConfig, PretrainedConfig
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
MergedColumnParallelLinear, MergedColumnParallelLinear,
@@ -69,6 +65,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, get_compiler_backend, set_weight_attrs from sglang.srt.utils import add_prefix, get_compiler_backend, set_weight_attrs
@@ -97,7 +94,7 @@ class LayerNorm(nn.Module):
return hidden_states, residuals return hidden_states, residuals
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor): def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
shard_dim = 0 if param.dim() != 1 else None shard_dim = 0 if param.dim() != 1 else None
param_data = param.data param_data = param.data
if shard_dim is not None: if shard_dim is not None:
@@ -152,7 +149,7 @@ class CohereAttention(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.config = config self.config = config
self.attention_dropout = config.attention_dropout self.attention_dropout = config.attention_dropout
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
+5 -6
View File
@@ -25,8 +25,6 @@ import torch.nn as nn
from sglang.srt.configs import DbrxConfig from sglang.srt.configs import DbrxConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import ( from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import (
@@ -54,6 +52,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, is_npu, set_weight_attrs from sglang.srt.utils import add_prefix, is_npu, set_weight_attrs
_is_npu = is_npu() _is_npu = is_npu()
@@ -71,7 +70,7 @@ class DbrxRouter(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.num_total_experts = config.ffn_config.moe_num_experts self.num_total_experts = config.ffn_config.moe_num_experts
self.d_model = config.d_model self.d_model = config.d_model
self.layer = ReplicatedLinear( self.layer = ReplicatedLinear(
@@ -103,7 +102,7 @@ class DbrxExperts(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.num_total_experts = config.ffn_config.moe_num_experts self.num_total_experts = config.ffn_config.moe_num_experts
self.top_k = config.ffn_config.moe_top_k self.top_k = config.ffn_config.moe_top_k
self.d_model = config.d_model self.d_model = config.d_model
@@ -155,7 +154,7 @@ class DbrxExperts(nn.Module):
def weight_loader( def weight_loader(
self, param: nn.Parameter, loaded_weight: torch.Tensor, weight_name: str self, param: nn.Parameter, loaded_weight: torch.Tensor, weight_name: str
): ):
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
param_data = param.data param_data = param.data
shard_size = self.intermediate_size shard_size = self.intermediate_size
shard = slice(tp_rank * shard_size, (tp_rank + 1) * shard_size) shard = slice(tp_rank * shard_size, (tp_rank + 1) * shard_size)
@@ -242,7 +241,7 @@ class DbrxAttention(nn.Module):
is_neox_style=True, is_neox_style=True,
) )
tp_world_size = get_tensor_model_parallel_world_size() tp_world_size = get_parallel().tp_size
self.tp_size = tp_world_size self.tp_size = tp_world_size
assert self.total_num_heads % tp_world_size == 0 assert self.total_num_heads % tp_world_size == 0
self.num_heads = self.total_num_heads // tp_world_size self.num_heads = self.total_num_heads // tp_world_size
+4 -5
View File
@@ -25,8 +25,6 @@ from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
@@ -49,6 +47,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, cpu_has_amx_support, is_cpu, is_npu from sglang.srt.utils import add_prefix, cpu_has_amx_support, is_cpu, is_npu
from sglang.srt.utils.hf_transformers_utils import get_rope_config from sglang.srt.utils.hf_transformers_utils import get_rope_config
@@ -118,8 +117,8 @@ class DeepseekMoE(nn.Module):
): ):
super().__init__() super().__init__()
self.config = config self.config = config
self.rank = get_tensor_model_parallel_rank() self.rank = get_parallel().tp_rank
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.n_routed_experts = config.n_routed_experts self.n_routed_experts = config.n_routed_experts
self.top_k = config.num_experts_per_tok self.top_k = config.num_experts_per_tok
if self.tp_size > self.n_routed_experts: if self.tp_size > self.n_routed_experts:
@@ -244,7 +243,7 @@ class DeepseekAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
+6 -9
View File
@@ -25,7 +25,7 @@ from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.configs.model_config import is_deepseek_dsa from sglang.srt.configs.model_config import is_deepseek_dsa
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size from sglang.srt.distributed import get_pp_group
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.layers.attention.dsa.utils import ( from sglang.srt.layers.attention.dsa.utils import (
@@ -33,10 +33,6 @@ from sglang.srt.layers.attention.dsa.utils import (
dsa_use_prefill_cp, dsa_use_prefill_cp,
is_dsa_enable_prefill_cp, is_dsa_enable_prefill_cp,
) )
from sglang.srt.layers.dp_attention import (
get_attention_cp_rank,
get_attention_cp_size,
)
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.linear import ReplicatedLinear
from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.logits_processor import LogitsProcessor
@@ -60,6 +56,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8 from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8
from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
from sglang.srt.models.utils import WeightsMapper from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
@@ -145,7 +142,7 @@ class DeepseekModelNextN(nn.Module):
is_mla_prefill_cp_enabled() and not is_deepseek_dsa(config) is_mla_prefill_cp_enabled() and not is_deepseek_dsa(config)
) )
if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp:
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
else: else:
self.cp_size = None self.cp_size = None
self.decoder = DeepseekV2DecoderLayer( self.decoder = DeepseekV2DecoderLayer(
@@ -280,7 +277,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
) -> None: ) -> None:
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
# if not set, model load will be broken in DeepseekV3ForCausalLM load_weights() # if not set, model load will be broken in DeepseekV3ForCausalLM load_weights()
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
@@ -289,8 +286,8 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
self.mla_enable_prefill_cp = is_mla_prefill_cp_enabled() and not self.use_dsa self.mla_enable_prefill_cp = is_mla_prefill_cp_enabled() and not self.use_dsa
if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp:
self.cp_rank = get_attention_cp_rank() self.cp_rank = get_parallel().attn_cp_rank
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
else: else:
self.cp_rank = None self.cp_rank = None
self.cp_size = None self.cp_size = None
+14 -23
View File
@@ -47,9 +47,7 @@ from sglang.srt.configs.model_config import (
) )
from sglang.srt.distributed import ( from sglang.srt.distributed import (
divide, divide,
get_moe_expert_parallel_world_size,
get_pp_group, get_pp_group,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -72,12 +70,6 @@ from sglang.srt.layers.communicator import (
get_attn_tp_context, get_attn_tp_context,
) )
from sglang.srt.layers.communicator_dsa_cp import DSACPLayerCommunicator from sglang.srt.layers.communicator_dsa_cp import DSACPLayerCommunicator
from sglang.srt.layers.dp_attention import (
get_attention_cp_rank,
get_attention_cp_size,
get_attention_tp_rank,
get_attention_tp_size,
)
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
@@ -169,6 +161,7 @@ from sglang.srt.models.deepseek_common.utils import (
_use_aiter_bpreshuffle_gfx95, _use_aiter_bpreshuffle_gfx95,
_use_aiter_gfx95, _use_aiter_gfx95,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -527,8 +520,8 @@ class DeepseekV2MoE(nn.Module):
mla_enable_prefill_cp: bool = False, mla_enable_prefill_cp: bool = False,
): ):
super().__init__() super().__init__()
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.moe_ep_size = get_moe_expert_parallel_world_size() self.moe_ep_size = get_parallel().moe_ep_size
self.routed_scaling_factor = config.routed_scaling_factor self.routed_scaling_factor = config.routed_scaling_factor
self.n_shared_experts = config.n_shared_experts self.n_shared_experts = config.n_shared_experts
@@ -776,7 +769,7 @@ class DeepseekV2MoE(nn.Module):
or get_moe_a2a_backend().is_ascend_fuseep() or get_moe_a2a_backend().is_ascend_fuseep()
): ):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_moe_expert_parallel_world_size() self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts config.n_routed_experts
+ get_global_server_args().ep_num_redundant_experts + get_global_server_args().ep_num_redundant_experts
@@ -1510,8 +1503,8 @@ class DeepseekV2AttentionMLA(
self.kv_lora_rank = kv_lora_rank self.kv_lora_rank = kv_lora_rank
self.quant_config = quant_config self.quant_config = quant_config
self.is_nextn = is_nextn self.is_nextn = is_nextn
attn_tp_rank = get_attention_tp_rank() attn_tp_rank = get_parallel().attn_tp_rank
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
self.use_dsa = is_deepseek_dsa(config) self.use_dsa = is_deepseek_dsa(config)
self.dsa_enable_prefill_cp = dsa_enable_prefill_cp self.dsa_enable_prefill_cp = dsa_enable_prefill_cp
self.mla_enable_prefill_cp = mla_enable_prefill_cp self.mla_enable_prefill_cp = mla_enable_prefill_cp
@@ -1521,7 +1514,7 @@ class DeepseekV2AttentionMLA(
# store cp_size whenever either CP flavor is active so rebuild_cp_kv_cache # store cp_size whenever either CP flavor is active so rebuild_cp_kv_cache
# and the FA3 MLA wrapper can reach it on the dense MLA path too. # and the FA3 MLA wrapper can reach it on the dense MLA path too.
if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp:
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
self.num_heads = num_heads self.num_heads = num_heads
assert num_heads % attn_tp_size == 0 assert num_heads % attn_tp_size == 0
self.num_local_heads = num_heads // attn_tp_size self.num_local_heads = num_heads // attn_tp_size
@@ -2287,7 +2280,7 @@ class DeepseekV2Model(nn.Module):
is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config) is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config)
) )
if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp:
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
else: else:
self.cp_size = None self.cp_size = None
@@ -2372,11 +2365,9 @@ class DeepseekV2Model(nn.Module):
allocate_size = 0 allocate_size = 0
for i in range(len(self.layers)): for i in range(len(self.layers)):
if isinstance(self.layers[i].mlp, DeepseekV2MoE): if isinstance(self.layers[i].mlp, DeepseekV2MoE):
# tp_size = get_tensor_model_parallel_world_size() # tp_size = get_parallel().tp_size
is_a2a_moe = is_deepep_class_backend() is_a2a_moe = is_deepep_class_backend()
tp_size = ( tp_size = 1 if is_a2a_moe else get_parallel().tp_size
1 if is_a2a_moe else get_tensor_model_parallel_world_size()
)
intermediate_size = ( intermediate_size = (
config.moe_intermediate_size * config.n_shared_experts config.moe_intermediate_size * config.n_shared_experts
) )
@@ -2576,7 +2567,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
self.determine_num_fused_shared_experts() self.determine_num_fused_shared_experts()
self.use_dsa = is_deepseek_dsa(config) self.use_dsa = is_deepseek_dsa(config)
@@ -2614,8 +2605,8 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config) is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config)
) )
if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp:
self.cp_rank = get_attention_cp_rank() self.cp_rank = get_parallel().attn_cp_rank
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
else: else:
self.cp_rank = self.cp_size = None self.cp_rank = self.cp_size = None
@@ -2672,7 +2663,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
"or AMD-platform with capability >= gfx942(MI30x) can use shared experts fusion optimization." "or AMD-platform with capability >= gfx942(MI30x) can use shared experts fusion optimization."
"or MT-platform with capability >= 31 can use shared experts fusion optimization." "or MT-platform with capability >= 31 can use shared experts fusion optimization."
) )
elif get_moe_expert_parallel_world_size() > 1 and ( elif get_parallel().moe_ep_size > 1 and (
not _is_hip or torch.cuda.get_device_capability("cuda") < (9, 4) not _is_hip or torch.cuda.get_device_capability("cuda") < (9, 4)
): ):
disable_reason = ( disable_reason = (
+12 -17
View File
@@ -29,7 +29,6 @@ from sglang.srt.compilation.compilation_config import register_split_op
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_pp_group, get_pp_group,
get_tensor_model_parallel_world_size,
get_tp_group, get_tp_group,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -53,11 +52,6 @@ from sglang.srt.layers.dp_attention import (
attn_tp_all_gather, attn_tp_all_gather,
dp_gather_partial, dp_gather_partial,
dp_scatter, dp_scatter,
get_attention_cp_rank,
get_attention_cp_size,
get_attention_dp_size,
get_attention_tp_rank,
get_attention_tp_size,
get_dp_global_num_tokens, get_dp_global_num_tokens,
get_global_dp_buffer, get_global_dp_buffer,
get_local_dp_buffer, get_local_dp_buffer,
@@ -116,6 +110,7 @@ from sglang.srt.models.deepseek_v2 import ParallelLMHead, _is_cuda, _is_hip, _is
from sglang.srt.models.triton_ops.deepseek_v4 import ( from sglang.srt.models.triton_ops.deepseek_v4 import (
rms_normalize_triton as rms_normalize_triton, rms_normalize_triton as rms_normalize_triton,
) )
from sglang.srt.runtime_context import get_parallel
if not _is_hip: if not _is_hip:
from sglang.srt.layers.utils.cp_utils import ( from sglang.srt.layers.utils.cp_utils import (
@@ -271,11 +266,11 @@ class MQALayer(nn.Module):
compress_ratio_override: Optional[int] = None, compress_ratio_override: Optional[int] = None,
) -> None: ) -> None:
super().__init__() super().__init__()
self.tp_rank = attn_tp_rank = get_attention_tp_rank() self.tp_rank = attn_tp_rank = get_parallel().attn_tp_rank
self.tp_size = attn_tp_size = get_attention_tp_size() self.tp_size = attn_tp_size = get_parallel().attn_tp_size
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
if self.dsa_enable_prefill_cp: if self.dsa_enable_prefill_cp:
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
self.tp_rank = attn_tp_rank = 0 self.tp_rank = attn_tp_rank = 0
self.tp_size = attn_tp_size = 1 self.tp_size = attn_tp_size = 1
self.layer_id = layer_id self.layer_id = layer_id
@@ -1494,13 +1489,13 @@ class DeepseekV4DecoderLayer(nn.Module):
_use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch) _use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch)
_use_tp_moe_gather = ( _use_tp_moe_gather = (
not _use_cp not _use_cp
and get_attention_dp_size() > 1 and get_parallel().attn_dp_size > 1
and get_moe_a2a_backend().is_none() and get_moe_a2a_backend().is_none()
) )
_use_tp_attn_a2a_scatter = ( _use_tp_attn_a2a_scatter = (
not _use_cp not _use_cp
and envs.SGLANG_DSV4_FIX_TP_ATTN_A2A_SCATTER.get() and envs.SGLANG_DSV4_FIX_TP_ATTN_A2A_SCATTER.get()
and get_attention_tp_size() > 1 and get_parallel().attn_tp_size > 1
and not get_moe_a2a_backend().is_none() and not get_moe_a2a_backend().is_none()
) )
# symmetric gather+scatter for the no-EP TP-MoE dp-attn path: # symmetric gather+scatter for the no-EP TP-MoE dp-attn path:
@@ -1532,7 +1527,7 @@ class DeepseekV4DecoderLayer(nn.Module):
dp_gather_partial(hidden_states, local_hidden_states, forward_batch) dp_gather_partial(hidden_states, local_hidden_states, forward_batch)
_a2a_scatter_chunks: Optional[List[torch.Tensor]] = None _a2a_scatter_chunks: Optional[List[torch.Tensor]] = None
if _use_tp_attn_a2a_scatter: if _use_tp_attn_a2a_scatter:
s, r = get_attention_tp_size(), get_attention_tp_rank() s, r = get_parallel().attn_tp_size, get_parallel().attn_tp_rank
_a2a_scatter_chunks = list(hidden_states.tensor_split(s)) _a2a_scatter_chunks = list(hidden_states.tensor_split(s))
hidden_states = _a2a_scatter_chunks[r].contiguous() hidden_states = _a2a_scatter_chunks[r].contiguous()
input_ids = input_ids.tensor_split(s)[r].contiguous() input_ids = input_ids.tensor_split(s)[r].contiguous()
@@ -1646,7 +1641,7 @@ class DeepseekV4Model(nn.Module):
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
self.use_fused_mhc_post_pre = _is_fused_mhc_post_pre_enabled() self.use_fused_mhc_post_pre = _is_fused_mhc_post_pre_enabled()
if self.dsa_enable_prefill_cp: if self.dsa_enable_prefill_cp:
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
def hc_head( def hc_head(
self, self,
@@ -1694,7 +1689,7 @@ class DeepseekV4Model(nn.Module):
hidden_states.shape[0], self.hc_mult, self.hidden_size hidden_states.shape[0], self.hc_mult, self.hidden_size
) )
if get_attention_dp_size() > 1 and get_moe_a2a_backend().is_none(): if get_parallel().attn_dp_size > 1 and get_moe_a2a_backend().is_none():
input_ids_global = torch.empty( input_ids_global = torch.empty(
(_DpGatheredBufferWrapper._global_dp_buffer_len, 1), (_DpGatheredBufferWrapper._global_dp_buffer_len, 1),
dtype=input_ids.dtype, dtype=input_ids.dtype,
@@ -1776,7 +1771,7 @@ class DeepseekV4ForCausalLM(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
self.determine_num_fused_shared_experts() self.determine_num_fused_shared_experts()
self.model = DeepseekV4Model( self.model = DeepseekV4Model(
@@ -1816,8 +1811,8 @@ class DeepseekV4ForCausalLM(nn.Module):
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
if self.dsa_enable_prefill_cp: if self.dsa_enable_prefill_cp:
self.cp_rank = get_attention_cp_rank() self.cp_rank = get_parallel().attn_cp_rank
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
@property @property
def routed_experts_weights_of_layer(self): def routed_experts_weights_of_layer(self):
@@ -6,7 +6,7 @@ import torch.nn.functional as F
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size from sglang.srt.distributed import get_pp_group
from sglang.srt.layers.attention.dsa.utils import ( from sglang.srt.layers.attention.dsa.utils import (
can_dsa_cp_split, can_dsa_cp_split,
dsa_use_prefill_cp, dsa_use_prefill_cp,
@@ -16,9 +16,6 @@ from sglang.srt.layers.attention.dsa.utils import (
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
_DpGatheredBufferWrapper, _DpGatheredBufferWrapper,
dp_gather_partial, dp_gather_partial,
get_attention_cp_rank,
get_attention_cp_size,
get_attention_dp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -40,6 +37,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -108,7 +106,7 @@ class DeepseekV4ModelNextN(nn.Module):
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
if self.dsa_enable_prefill_cp: if self.dsa_enable_prefill_cp:
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
else: else:
self.cp_size = None self.cp_size = None
@@ -156,7 +154,7 @@ class DeepseekV4ModelNextN(nn.Module):
else: else:
hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1) hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1)
if get_attention_dp_size() > 1 and get_moe_a2a_backend().is_none(): if get_parallel().attn_dp_size > 1 and get_moe_a2a_backend().is_none():
input_ids_global = torch.empty( input_ids_global = torch.empty(
(_DpGatheredBufferWrapper._global_dp_buffer_len, 1), (_DpGatheredBufferWrapper._global_dp_buffer_len, 1),
dtype=input_ids.dtype, dtype=input_ids.dtype,
@@ -213,14 +211,14 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
) -> None: ) -> None:
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.quant_config = quant_config self.quant_config = quant_config
self.determine_num_fused_shared_experts() self.determine_num_fused_shared_experts()
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
if self.dsa_enable_prefill_cp: if self.dsa_enable_prefill_cp:
self.cp_rank = get_attention_cp_rank() self.cp_rank = get_parallel().attn_cp_rank
self.cp_size = get_attention_cp_size() self.cp_size = get_parallel().attn_cp_size
else: else:
self.cp_rank = None self.cp_rank = None
self.cp_size = None self.cp_size = None
+2 -2
View File
@@ -12,7 +12,6 @@ import torch
import torch.nn.functional as F import torch.nn.functional as F
from torch import nn from torch import nn
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -26,6 +25,7 @@ from sglang.srt.layers.rotary_embedding import get_rope
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.utils import apply_qk_norm from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_parallel
from sglang.srt.speculative.dflash_utils import ( from sglang.srt.speculative.dflash_utils import (
can_dflash_slice_qkv_weight, can_dflash_slice_qkv_weight,
get_dflash_attention_sliding_window_size, get_dflash_attention_sliding_window_size,
@@ -70,7 +70,7 @@ class DFlashAttention(nn.Module):
def __init__(self, config, layer_id: int) -> None: def __init__(self, config, layer_id: int) -> None:
super().__init__() super().__init__()
hidden_size = int(config.hidden_size) hidden_size = int(config.hidden_size)
tp_size = int(get_tensor_model_parallel_world_size()) tp_size = int(get_parallel().tp_size)
total_num_heads = int(config.num_attention_heads) total_num_heads = int(config.num_attention_heads)
total_num_kv_heads = int( total_num_kv_heads = int(
getattr(config, "num_key_value_heads", total_num_heads) getattr(config, "num_key_value_heads", total_num_heads)
+2 -2
View File
@@ -9,10 +9,10 @@ from torch.nn import LayerNorm
from transformers.modeling_utils import PreTrainedModel from transformers.modeling_utils import PreTrainedModel
from sglang.srt.configs.dots_vlm import DotsVisionConfig from sglang.srt.configs.dots_vlm import DotsVisionConfig
from sglang.srt.distributed import parallel_state
from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.conv import Conv2dLayer from sglang.srt.layers.conv import Conv2dLayer
from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -224,7 +224,7 @@ class DotsVisionTransformer(PreTrainedModel):
def _update_vision_config(self): def _update_vision_config(self):
"""update vision config to support tp""" """update vision config to support tp"""
world_size = parallel_state.get_tensor_model_parallel_world_size() world_size = get_parallel().tp_size
num_heads = self.config.num_attention_heads num_heads = self.config.num_attention_heads
head_dim = self.config.embed_dim // num_heads head_dim = self.config.embed_dim // num_heads
num_dummy_heads = 0 num_dummy_heads = 0
+2 -2
View File
@@ -24,7 +24,6 @@ from transformers.models.ernie4_5_moe.configuration_ernie4_5_moe import (
) )
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.layers.communicator import enable_moe_dense_fully_dp from sglang.srt.layers.communicator import enable_moe_dense_fully_dp
@@ -42,6 +41,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2MLP as Ernie4MLP from sglang.srt.models.deepseek_v2 import DeepseekV2MLP as Ernie4MLP
from sglang.srt.models.llama import LlamaAttention as Ernie4Attention from sglang.srt.models.llama import LlamaAttention as Ernie4Attention
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, is_npu, make_layers from sglang.srt.utils import add_prefix, is_npu, make_layers
from sglang.srt.utils.hf_transformers_utils import get_rope_config from sglang.srt.utils.hf_transformers_utils import get_rope_config
@@ -77,7 +77,7 @@ class Ernie4Moe(nn.Module):
): ):
super().__init__() super().__init__()
self.layer_id = layer_id self.layer_id = layer_id
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.moe_num_shared_experts = getattr(config, "moe_num_shared_experts", 0) self.moe_num_shared_experts = getattr(config, "moe_num_shared_experts", 0)
if config.hidden_act != "silu": if config.hidden_act != "silu":
+3 -3
View File
@@ -24,7 +24,6 @@ from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_pp_group, get_pp_group,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.dp_attention import is_dp_attention_enabled
@@ -43,6 +42,7 @@ from sglang.srt.layers.utils import PPMissingLayer
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.models.deepseek_v2 import DeepseekV2MLP as Ernie4_5_VLMoeMLP from sglang.srt.models.deepseek_v2 import DeepseekV2MLP as Ernie4_5_VLMoeMLP
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils import add_prefix, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -67,7 +67,7 @@ class Ernie4_5_VLMoeAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
@@ -158,7 +158,7 @@ class Ernie4_5_VLMoeMoE(nn.Module):
): ):
super().__init__() super().__init__()
self.layer_id = layer_id self.layer_id = layer_id
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.moe_num_shared_experts = getattr(config, "moe_num_shared_experts", 0) self.moe_num_shared_experts = getattr(config, "moe_num_shared_experts", 0)
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
+2 -2
View File
@@ -21,7 +21,6 @@ from typing import Any, Dict, Iterable, Optional, Tuple
import torch import torch
from torch import nn from torch import nn
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -39,6 +38,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
from sglang.srt.utils.hf_transformers_utils import get_rope_config from sglang.srt.utils.hf_transformers_utils import get_rope_config
@@ -98,7 +98,7 @@ class ExaoneAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
+7 -8
View File
@@ -5,11 +5,9 @@ import torch
from torch import nn from torch import nn
from transformers import Exaone4Config from transformers import Exaone4Config
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size from sglang.srt.distributed import get_pp_group
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
get_local_attention_dp_size, get_local_attention_dp_size,
) )
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -33,6 +31,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils import add_prefix, make_layers
from sglang.utils import get_exception_traceback, logger from sglang.utils import get_exception_traceback, logger
@@ -106,10 +105,10 @@ class Exaone4Attention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
attn_tp_rank = get_attention_tp_rank() attn_tp_rank = get_parallel().attn_tp_rank
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
@@ -241,8 +240,8 @@ class Exaone4DecoderLayer(nn.Module):
max_position_embeddings = getattr(config, "max_position_embeddings", 8192) max_position_embeddings = getattr(config, "max_position_embeddings", 8192)
self.local_dp_size = get_local_attention_dp_size() self.local_dp_size = get_local_attention_dp_size()
self.attn_tp_size = get_attention_tp_size() self.attn_tp_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_attention_tp_rank() self.attn_tp_rank = get_parallel().attn_tp_rank
self.self_attn = Exaone4Attention( self.self_attn = Exaone4Attention(
config=config, config=config,
+8 -11
View File
@@ -25,9 +25,7 @@ from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_moe_expert_parallel_world_size,
get_pp_group, get_pp_group,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
@@ -35,8 +33,6 @@ from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -66,6 +62,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
@@ -147,8 +144,8 @@ class ExaoneMoESparseMoEBlock(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.moe_ep_size = get_moe_expert_parallel_world_size() self.moe_ep_size = get_parallel().moe_ep_size
self.layer_id = layer_id self.layer_id = layer_id
self.routed_scaling_factor = config.routed_scaling_factor self.routed_scaling_factor = config.routed_scaling_factor
self.alt_stream = alt_stream self.alt_stream = alt_stream
@@ -214,7 +211,7 @@ class ExaoneMoESparseMoEBlock(nn.Module):
) )
if get_moe_a2a_backend().is_deepep(): if get_moe_a2a_backend().is_deepep():
self.ep_size = get_moe_expert_parallel_world_size() self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.num_experts + get_global_server_args().ep_num_redundant_experts config.num_experts + get_global_server_args().ep_num_redundant_experts
) )
@@ -330,8 +327,8 @@ class ExaoneMoEAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
attn_tp_rank = get_attention_tp_rank() attn_tp_rank = get_parallel().attn_tp_rank
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % attn_tp_size == 0 assert self.total_num_heads % attn_tp_size == 0
@@ -469,8 +466,8 @@ class ExaoneMoEDecoderLayer(nn.Module):
attention_bias = getattr(config, "attention_bias", False) or getattr( attention_bias = getattr(config, "attention_bias", False) or getattr(
config, "bias", False config, "bias", False
) )
self.attn_tp_size = get_attention_tp_size() self.attn_tp_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_attention_tp_rank() self.attn_tp_rank = get_parallel().attn_tp_rank
self.self_attn = ExaoneMoEAttention( self.self_attn = ExaoneMoEAttention(
config=config, config=config,
+3 -2
View File
@@ -23,13 +23,14 @@ import torch
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size from sglang.srt.distributed import get_pp_group
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.exaone_moe import ExaoneMoEForCausalLM, ExaoneMoEModel from sglang.srt.models.exaone_moe import ExaoneMoEForCausalLM, ExaoneMoEModel
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -46,7 +47,7 @@ class ExaoneMoEForCausalLMMTP(ExaoneMoEForCausalLM):
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
config.num_hidden_layers = 1 config.num_hidden_layers = 1
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
+6 -7
View File
@@ -5,7 +5,7 @@ import torch
from torch import nn from torch import nn
from sglang.srt.configs.falcon_h1 import FalconH1Config from sglang.srt.configs.falcon_h1 import FalconH1Config
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size from sglang.srt.distributed import get_pp_group
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend, HybridLinearAttnBackend,
@@ -14,8 +14,6 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
from sglang.srt.layers.attention.mamba.mamba import MambaMixer2 from sglang.srt.layers.attention.mamba.mamba import MambaMixer2
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -35,6 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, make_layers from sglang.srt.utils import add_prefix, is_cuda, make_layers
@@ -79,7 +78,7 @@ class FalconH1MLP(nn.Module):
self.layer_id = layer_id self.layer_id = layer_id
self.intermediate_size = intermediate_size self.intermediate_size = intermediate_size
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.gate_multiplier, self.down_multiplier = mlp_multipliers self.gate_multiplier, self.down_multiplier = mlp_multipliers
@@ -114,9 +113,9 @@ class FalconH1HybridAttentionDecoderLayer(nn.Module):
super().__init__() super().__init__()
self.config = config self.config = config
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.attn_tp_rank = get_attention_tp_rank() self.attn_tp_rank = get_parallel().attn_tp_rank
self.attn_tp_size = get_attention_tp_size() self.attn_tp_size = get_parallel().attn_tp_size
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.total_num_heads = config.num_attention_heads self.total_num_heads = config.num_attention_heads
assert self.total_num_heads % self.attn_tp_size == 0 assert self.total_num_heads % self.attn_tp_size == 0
self.num_heads = self.total_num_heads // self.attn_tp_size self.num_heads = self.total_num_heads // self.attn_tp_size
+2 -2
View File
@@ -25,7 +25,6 @@ import torch
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import GeluAndMul from sglang.srt.layers.activation import GeluAndMul
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -40,6 +39,7 @@ from sglang.srt.layers.rotary_embedding import get_rope
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -90,7 +90,7 @@ class GemmaAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
+2 -2
View File
@@ -24,7 +24,6 @@ import torch
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import GeluAndMul from sglang.srt.layers.activation import GeluAndMul
from sglang.srt.layers.layernorm import GemmaRMSNorm from sglang.srt.layers.layernorm import GemmaRMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -42,6 +41,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, is_npu, make_layers from sglang.srt.utils import add_prefix, is_npu, make_layers
_is_npu = is_npu() _is_npu = is_npu()
@@ -111,7 +111,7 @@ class Gemma2Attention(nn.Module):
self.layer_id = layer_id self.layer_id = layer_id
self.config = config self.config = config
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
+4 -7
View File
@@ -26,10 +26,6 @@ from transformers import (
PreTrainedModel, PreTrainedModel,
) )
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.layers.activation import GeluAndMul from sglang.srt.layers.activation import GeluAndMul
from sglang.srt.layers.layernorm import Gemma3RMSNorm from sglang.srt.layers.layernorm import Gemma3RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -47,6 +43,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, cpu_has_amx_support, is_cpu, make_layers from sglang.srt.utils import add_prefix, cpu_has_amx_support, is_cpu, make_layers
_is_cpu = is_cpu() _is_cpu = is_cpu()
@@ -126,7 +123,7 @@ class Gemma3Attention(nn.Module):
super().__init__() super().__init__()
self.layer_id = layer_id self.layer_id = layer_id
self.config = config self.config = config
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = config.num_attention_heads self.total_num_heads = config.num_attention_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
@@ -922,10 +919,10 @@ class Gemma3ForCausalLM(PreTrainedModel):
VocabParallelEmbedding (sharded). This method extracts the correct VocabParallelEmbedding (sharded). This method extracts the correct
shard so the weights can be shared. shard so the weights can be shared.
""" """
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
if tp_size <= 1: if tp_size <= 1:
return weight return weight
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
shard_size = (weight.shape[0] + tp_size - 1) // tp_size shard_size = (weight.shape[0] + tp_size - 1) // tp_size
return weight[tp_rank * shard_size : (tp_rank + 1) * shard_size] return weight[tp_rank * shard_size : (tp_rank + 1) * shard_size]
+2 -2
View File
@@ -5,7 +5,6 @@ import torch.nn.functional as F
from torch import nn from torch import nn
from transformers import AutoModel, Gemma3nTextConfig, PretrainedConfig, PreTrainedModel from transformers import AutoModel, Gemma3nTextConfig, PretrainedConfig, PreTrainedModel
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import GeluAndMul from sglang.srt.layers.activation import GeluAndMul
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -26,6 +25,7 @@ from sglang.srt.model_loader.weight_utils import (
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.models.gemma3_causal import Gemma3TextScaledWordEmbedding from sglang.srt.models.gemma3_causal import Gemma3TextScaledWordEmbedding
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils import add_prefix, make_layers
@@ -325,7 +325,7 @@ class Gemma3nAttention(nn.Module):
super().__init__() super().__init__()
self.layer_id = layer_id self.layer_id = layer_id
self.config = config self.config = config
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = config.num_attention_heads self.total_num_heads = config.num_attention_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
+5 -8
View File
@@ -37,16 +37,13 @@ from sglang.srt.layers.clippable_linear import (
ClippableQKVParallelLinear, ClippableQKVParallelLinear,
ClippableRowParallelLinear, ClippableRowParallelLinear,
) )
from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
)
from sglang.srt.layers.layernorm import Gemma4RMSNorm from sglang.srt.layers.layernorm import Gemma4RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
RowParallelLinear, RowParallelLinear,
) )
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, make_layers, set_weight_attrs from sglang.srt.utils import add_prefix, make_layers, set_weight_attrs
# SSCP convolution constants (no longer in config.json, never varied across models) # SSCP convolution constants (no longer in config.json, never varied across models)
@@ -69,7 +66,7 @@ class Gemma4AudioRelativePositionEmbedding(nn.Module):
super().__init__() super().__init__()
self.config = config self.config = config
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
total_num_heads = config.num_attention_heads total_num_heads = config.num_attention_heads
self.channels = config.hidden_size self.channels = config.hidden_size
self.head_dim = self.channels // total_num_heads self.head_dim = self.channels // total_num_heads
@@ -219,7 +216,7 @@ class Gemma4AudioAttention(nn.Module):
super().__init__() super().__init__()
self.config = config self.config = config
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
total_num_heads = config.num_attention_heads total_num_heads = config.num_attention_heads
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.head_dim = self.hidden_size // total_num_heads self.head_dim = self.hidden_size // total_num_heads
@@ -641,7 +638,7 @@ class Gemma4AudioConformerLightConv1d(nn.Module):
super().__init__() super().__init__()
self.config = config self.config = config
self.causal_padding = config.conv_kernel_size - 1 self.causal_padding = config.conv_kernel_size - 1
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
hidden_per_tp = config.hidden_size // tp_size hidden_per_tp = config.hidden_size // tp_size
self.register_buffer( self.register_buffer(
@@ -673,7 +670,7 @@ class Gemma4AudioConformerLightConv1d(nn.Module):
hidden_per_tp, eps=config.rms_norm_eps, scale_shift=0.0 hidden_per_tp, eps=config.rms_norm_eps, scale_shift=0.0
) )
tp_rank = get_attention_tp_rank() tp_rank = get_parallel().attn_tp_rank
def _shard_dim0(param, loaded_weight, _rank=tp_rank, _tp=tp_size): def _shard_dim0(param, loaded_weight, _rank=tp_rank, _tp=tp_size):
shard = param.shape[0] shard = param.shape[0]
+5 -6
View File
@@ -26,8 +26,6 @@ from transformers import (
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_pp_group, get_pp_group,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
) )
from sglang.srt.layers.gemma4_fused_ops import ( from sglang.srt.layers.gemma4_fused_ops import (
gemma4_fused_routing, gemma4_fused_routing,
@@ -60,6 +58,7 @@ from sglang.srt.models.gemma3_causal import Gemma3MLP, Gemma3TextScaledWordEmbed
from sglang.srt.models.utils import ( from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils import add_prefix, make_layers
@@ -210,7 +209,7 @@ class Gemma4MoE(nn.Module):
self.layer_id = layer_id self.layer_id = layer_id
self.hidden_size = hidden_size self.hidden_size = hidden_size
self.num_experts = config.num_experts self.num_experts = config.num_experts
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
# Per-expert output scale folded into routing weights so that # Per-expert output scale folded into routing weights so that
# MoE's fused kernel computes: Σ_e (expert_e * w_e * scale_e) # MoE's fused kernel computes: Σ_e (expert_e * w_e * scale_e)
@@ -291,7 +290,7 @@ class Gemma4Attention(nn.Module):
self.layer_id = layer_id self.layer_id = layer_id
self.config = config self.config = config
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
layer_type = config.layer_types[layer_id] layer_type = config.layer_types[layer_id]
self.sliding_window = ( self.sliding_window = (
@@ -1379,10 +1378,10 @@ class Gemma4ForCausalLM(PreTrainedModel):
VocabParallelEmbedding (sharded). This method extracts the correct VocabParallelEmbedding (sharded). This method extracts the correct
shard so the weights can be shared. shard so the weights can be shared.
""" """
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
if tp_size <= 1: if tp_size <= 1:
return weight return weight
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
shard_size = (weight.shape[0] + tp_size - 1) // tp_size shard_size = (weight.shape[0] + tp_size - 1) // tp_size
return weight[tp_rank * shard_size : (tp_rank + 1) * shard_size] return weight[tp_rank * shard_size : (tp_rank + 1) * shard_size]
+2 -2
View File
@@ -27,9 +27,9 @@ from sglang.srt.layers.clippable_linear import (
ClippableQKVParallelLinear, ClippableQKVParallelLinear,
ClippableRowParallelLinear, ClippableRowParallelLinear,
) )
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.layernorm import Gemma4RMSNorm from sglang.srt.layers.layernorm import Gemma4RMSNorm
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, get_device_capability, is_cuda, is_hip from sglang.srt.utils import add_prefix, get_device_capability, is_cuda, is_hip
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -140,7 +140,7 @@ class Gemma4VisionAttention(nn.Module):
super().__init__() super().__init__()
self.head_dim = config.head_dim self.head_dim = config.head_dim
tp_size = get_attention_tp_size() tp_size = get_parallel().attn_tp_size
self.num_heads_per_partition = config.num_attention_heads // tp_size self.num_heads_per_partition = config.num_attention_heads // tp_size
self.num_kv_heads_per_partition = config.num_key_value_heads // tp_size self.num_kv_heads_per_partition = config.num_key_value_heads // tp_size
+4 -5
View File
@@ -25,8 +25,6 @@ from torch import nn
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_pp_group, get_pp_group,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
) )
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.dp_attention import is_dp_attention_enabled
@@ -51,6 +49,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
kv_cache_scales_loader, kv_cache_scales_loader,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils import add_prefix, make_layers
from sglang.srt.utils.hf_transformers_utils import get_rope_config from sglang.srt.utils.hf_transformers_utils import get_rope_config
@@ -125,7 +124,7 @@ class Glm4Attention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
@@ -397,8 +396,8 @@ class Glm4Model(nn.Module):
# factors (or else raise an exception). Thus, handled exceptions should # factors (or else raise an exception). Thus, handled exceptions should
# make sure to leave KV cache scale factors in a known good (dummy) state # make sure to leave KV cache scale factors in a known good (dummy) state
def load_kv_cache_scales(self, quantization_param_path: str) -> None: def load_kv_cache_scales(self, quantization_param_path: str) -> None:
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
for layer_idx, scaling_factor in kv_cache_scales_loader( for layer_idx, scaling_factor in kv_cache_scales_loader(
quantization_param_path, quantization_param_path,
tp_rank, tp_rank,
+9 -13
View File
@@ -26,11 +26,8 @@ from transformers import PretrainedConfig
from sglang.srt.batch_overlap.single_batch_overlap import SboFlags 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.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_moe_expert_parallel_world_size,
get_pp_group, get_pp_group,
get_pp_indices, get_pp_indices,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
parallel_state, parallel_state,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
@@ -48,8 +45,6 @@ from sglang.srt.layers.communicator import (
enable_moe_dense_fully_dp, enable_moe_dense_fully_dp,
) )
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
is_allocation_symmetric, is_allocation_symmetric,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
@@ -87,6 +82,7 @@ from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM
from sglang.srt.models.utils import apply_qk_norm from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
@@ -205,8 +201,8 @@ class Glm4MoeAttention(nn.Module):
self.hidden_size = hidden_size self.hidden_size = hidden_size
self.start_layer = start_layer self.start_layer = start_layer
attn_tp_rank = get_attention_tp_rank() attn_tp_rank = get_parallel().attn_tp_rank
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % attn_tp_size == 0 assert self.total_num_heads % attn_tp_size == 0
@@ -228,7 +224,7 @@ class Glm4MoeAttention(nn.Module):
self.rope_theta = rope_theta self.rope_theta = rope_theta
self.use_qk_norm = use_qk_norm self.use_qk_norm = use_qk_norm
self.max_position_embeddings = max_position_embeddings self.max_position_embeddings = max_position_embeddings
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_parallel().tp_rank
self.qkv_proj = QKVParallelLinear( self.qkv_proj = QKVParallelLinear(
hidden_size, hidden_size,
@@ -403,8 +399,8 @@ class Glm4MoeSparseMoeBlock(nn.Module):
): ):
nn.Module.__init__(self) nn.Module.__init__(self)
self.top_k = config.num_experts_per_tok self.top_k = config.num_experts_per_tok
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.moe_ep_size = get_moe_expert_parallel_world_size() self.moe_ep_size = get_parallel().moe_ep_size
self.routed_scaling_factor = config.routed_scaling_factor self.routed_scaling_factor = config.routed_scaling_factor
self.n_shared_experts = config.n_shared_experts self.n_shared_experts = config.n_shared_experts
self.num_fused_shared_experts = ( self.num_fused_shared_experts = (
@@ -527,7 +523,7 @@ class Glm4MoeSparseMoeBlock(nn.Module):
or get_moe_a2a_backend().is_ascend_fuseep() or get_moe_a2a_backend().is_ascend_fuseep()
): ):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_moe_expert_parallel_world_size() self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts config.n_routed_experts
+ get_global_server_args().ep_num_redundant_experts + get_global_server_args().ep_num_redundant_experts
@@ -1178,7 +1174,7 @@ class Glm4MoeForCausalLM(nn.Module):
nn.Module.__init__(self) nn.Module.__init__(self)
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
self.num_fused_shared_experts = 0 self.num_fused_shared_experts = 0
self.determine_num_fused_shared_experts() self.determine_num_fused_shared_experts()
@@ -1209,7 +1205,7 @@ class Glm4MoeForCausalLM(nn.Module):
"Only GLM-4.5 on NV-platform with capability >= 80 " "Only GLM-4.5 on NV-platform with capability >= 80 "
"or AMD-platform with capability >= gfx942(MI30x) can use shared experts fusion optimization." "or AMD-platform with capability >= gfx942(MI30x) can use shared experts fusion optimization."
) )
elif get_moe_expert_parallel_world_size() > 1 and ( elif get_parallel().moe_ep_size > 1 and (
not _is_hip or torch.cuda.get_device_capability("cuda") < (9, 4) not _is_hip or torch.cuda.get_device_capability("cuda") < (9, 4)
): ):
disable_reason = "Only GLM-4.5 on AMD-platform with capability >= gfx942(MI30x) can use shared experts fusion optimization under expert parallelism." disable_reason = "Only GLM-4.5 on AMD-platform with capability >= gfx942(MI30x) can use shared experts fusion optimization under expert parallelism."
+5 -6
View File
@@ -26,9 +26,7 @@ from transformers import PretrainedConfig
from sglang.srt.batch_overlap.single_batch_overlap import SboFlags 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.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_moe_expert_parallel_world_size,
get_pp_group, get_pp_group,
get_tensor_model_parallel_world_size,
parallel_state, parallel_state,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
@@ -76,6 +74,7 @@ from sglang.srt.models.deepseek_common.deepseek_weight_loader import (
) )
from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
@@ -185,7 +184,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
is_nextn: bool = False, is_nextn: bool = False,
): ):
super().__init__() super().__init__()
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.routed_scaling_factor = config.routed_scaling_factor self.routed_scaling_factor = config.routed_scaling_factor
self.n_shared_experts = config.n_shared_experts self.n_shared_experts = config.n_shared_experts
self.num_fused_shared_experts = ( self.num_fused_shared_experts = (
@@ -283,7 +282,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
if get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake(): if get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake():
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_moe_expert_parallel_world_size() self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts config.n_routed_experts
+ get_global_server_args().ep_num_redundant_experts + get_global_server_args().ep_num_redundant_experts
@@ -907,7 +906,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
super().__init__() super().__init__()
config.moe_layer_freq = 1 config.moe_layer_freq = 1
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.determine_num_fused_shared_experts("Glm4MoeLiteForCausalLM") self.determine_num_fused_shared_experts("Glm4MoeLiteForCausalLM")
@@ -951,7 +950,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
or self.config.n_shared_experts != 1 or self.config.n_shared_experts != 1
): ):
disable_reason = "Only GLM-4.5 or GLM-4.6 on NV-platform with capability >= 80 can use shared experts fusion optimization." disable_reason = "Only GLM-4.5 or GLM-4.6 on NV-platform with capability >= 80 can use shared experts fusion optimization."
elif get_moe_expert_parallel_world_size() > 1: elif get_parallel().moe_ep_size > 1:
disable_reason = "GLM-4.5 or GLM-4.6 cannot use shared experts fusion optimization under expert parallelism." disable_reason = "GLM-4.5 or GLM-4.6 cannot use shared experts fusion optimization under expert parallelism."
if disable_reason is not None: if disable_reason is not None:
@@ -21,7 +21,6 @@ import torch
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder 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.dp_attention import is_dp_attention_enabled
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -36,6 +35,7 @@ from sglang.srt.models.glm4_moe_lite import (
Glm4MoeLiteDecoderLayer, Glm4MoeLiteDecoderLayer,
Glm4MoeLiteForCausalLM, Glm4MoeLiteForCausalLM,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import BumpAllocator, add_prefix, is_npu from sglang.srt.utils import BumpAllocator, add_prefix, is_npu
@@ -139,7 +139,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
) -> None: ) -> None:
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
if ( if (
is_npu() is_npu()
and get_global_server_args().speculative_draft_model_quantization is None and get_global_server_args().speculative_draft_model_quantization is None
+2 -2
View File
@@ -21,7 +21,6 @@ import torch
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder 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.dp_attention import is_dp_attention_enabled
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -33,6 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
@@ -125,7 +125,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
) -> None: ) -> None:
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
if ( if (
is_npu() is_npu()
and get_global_server_args().speculative_draft_model_quantization is None and get_global_server_args().speculative_draft_model_quantization is None
+5 -10
View File
@@ -27,10 +27,6 @@ import torch.nn.functional as F
from einops import rearrange from einops import rearrange
from transformers.models.glm4v.configuration_glm4v import Glm4vConfig, Glm4vVisionConfig from transformers.models.glm4v.configuration_glm4v import Glm4vConfig, Glm4vVisionConfig
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.distributed.parallel_state import get_pp_group
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.attention import vision_utils from sglang.srt.layers.attention import vision_utils
@@ -57,6 +53,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTe
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.glm4 import Glm4Model from sglang.srt.models.glm4 import Glm4Model
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
@@ -86,10 +83,8 @@ class Glm4vVisionMLP(nn.Module):
use_data_parallel: bool = False, use_data_parallel: bool = False,
): ):
super().__init__() super().__init__()
self.tp_size = ( self.tp_size = 1 if use_data_parallel else get_parallel().tp_size
1 if use_data_parallel else get_tensor_model_parallel_world_size() self.tp_rank = 0 if use_data_parallel else get_parallel().tp_rank
)
self.tp_rank = 0 if use_data_parallel else get_tensor_model_parallel_rank()
self.gate_up_proj = MergedColumnParallelLinear( self.gate_up_proj = MergedColumnParallelLinear(
input_size=in_features, input_size=in_features,
output_sizes=[hidden_features] * 2, # [gate_proj, up_proj] output_sizes=[hidden_features] * 2, # [gate_proj, up_proj]
@@ -237,8 +232,8 @@ class Glm4vPatchMerger(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = d_model self.hidden_size = d_model
tp_size = 1 if use_data_parallel else get_tensor_model_parallel_world_size() tp_size = 1 if use_data_parallel else get_parallel().tp_size
tp_rank = 0 if use_data_parallel else get_tensor_model_parallel_rank() tp_rank = 0 if use_data_parallel else get_parallel().tp_rank
self.proj = ReplicatedLinear( self.proj = ReplicatedLinear(
self.hidden_size, self.hidden_size,
self.hidden_size, self.hidden_size,
+3 -6
View File
@@ -6,10 +6,6 @@ import torch
import torch.nn as nn import torch.nn as nn
from transformers.models.glm4v_moe.configuration_glm4v_moe import Glm4vMoeConfig from transformers.models.glm4v_moe.configuration_glm4v_moe import Glm4vMoeConfig
from sglang.srt.distributed import (
get_moe_expert_parallel_world_size,
get_tensor_model_parallel_world_size,
)
from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.distributed.parallel_state import get_pp_group
from sglang.srt.layers.attention import vision_utils from sglang.srt.layers.attention import vision_utils
from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.logits_processor import LogitsProcessor
@@ -22,6 +18,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.glm4_moe import Glm4MoeModel from sglang.srt.models.glm4_moe import Glm4MoeModel
from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, get_device_sm, is_cuda, log_info_on_rank0 from sglang.srt.utils import add_prefix, get_device_sm, is_cuda, log_info_on_rank0
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
@@ -47,7 +44,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
self.config = config self.config = config
self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder
vision_utils.update_vit_attn_dummy_heads_config(self.config) vision_utils.update_vit_attn_dummy_heads_config(self.config)
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
self.num_fused_shared_experts = 0 self.num_fused_shared_experts = 0
self.determine_num_fused_shared_experts() self.determine_num_fused_shared_experts()
@@ -97,7 +94,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
disable_reason = "Shared experts fusion currently requires CUDA devices." disable_reason = "Shared experts fusion currently requires CUDA devices."
elif _is_cuda and (_device_sm is not None) and (_device_sm < 80): elif _is_cuda and (_device_sm is not None) and (_device_sm < 80):
disable_reason = "Shared experts fusion requires SM80 or newer GPUs." disable_reason = "Shared experts fusion requires SM80 or newer GPUs."
elif get_moe_expert_parallel_world_size() > 1: elif get_parallel().moe_ep_size > 1:
disable_reason = "Shared experts fusion is not supported together with expert parallelism yet." disable_reason = "Shared experts fusion is not supported together with expert parallelism yet."
elif get_moe_a2a_backend().is_deepep(): elif get_moe_a2a_backend().is_deepep():
disable_reason = "Shared experts fusion is not supported when Deepep MoE backend is enabled." disable_reason = "Shared experts fusion is not supported when Deepep MoE backend is enabled."
+2 -2
View File
@@ -21,7 +21,6 @@ import torch
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder 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.dp_attention import is_dp_attention_enabled
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -34,6 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.glm4 import Glm4DecoderLayer from sglang.srt.models.glm4 import Glm4DecoderLayer
from sglang.srt.models.glm_ocr import GlmOcrForConditionalGeneration from sglang.srt.models.glm_ocr import GlmOcrForConditionalGeneration
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -125,7 +125,7 @@ class GlmOcrForConditionalGenerationNextN(GlmOcrForConditionalGeneration):
) -> None: ) -> None:
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.quant_config = quant_config self.quant_config = quant_config
self.model = GlmOcrModelNextN( self.model = GlmOcrModelNextN(
config, quant_config, prefix=add_prefix("model", prefix) config, quant_config, prefix=add_prefix("model", prefix)
+2 -2
View File
@@ -24,7 +24,6 @@ import torch
from torch import nn from torch import nn
from transformers import GPT2Config from transformers import GPT2Config
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import NewGELU from sglang.srt.layers.activation import NewGELU
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
@@ -37,6 +36,7 @@ from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -52,7 +52,7 @@ class GPT2Attention(nn.Module):
super().__init__() super().__init__()
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
total_num_heads = config.num_attention_heads total_num_heads = config.num_attention_heads
tensor_model_parallel_world_size = get_tensor_model_parallel_world_size() tensor_model_parallel_world_size = get_parallel().tp_size
assert total_num_heads % tensor_model_parallel_world_size == 0 assert total_num_heads % tensor_model_parallel_world_size == 0
self.num_heads = total_num_heads // tensor_model_parallel_world_size self.num_heads = total_num_heads // tensor_model_parallel_world_size
self.head_dim = self.hidden_size // total_num_heads self.head_dim = self.hidden_size // total_num_heads
+2 -2
View File
@@ -25,7 +25,6 @@ import torch
from torch import nn from torch import nn
from transformers import GPTBigCodeConfig from transformers import GPTBigCodeConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.activation import get_act_fn
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
@@ -38,6 +37,7 @@ from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -53,7 +53,7 @@ class GPTBigCodeAttention(nn.Module):
super().__init__() super().__init__()
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
total_num_heads = config.num_attention_heads total_num_heads = config.num_attention_heads
self.tensor_model_parallel_world_size = get_tensor_model_parallel_world_size() self.tensor_model_parallel_world_size = get_parallel().tp_size
assert total_num_heads % self.tensor_model_parallel_world_size == 0 assert total_num_heads % self.tensor_model_parallel_world_size == 0
self.num_heads = total_num_heads // self.tensor_model_parallel_world_size self.num_heads = total_num_heads // self.tensor_model_parallel_world_size
self.head_dim = self.hidden_size // total_num_heads self.head_dim = self.hidden_size // total_num_heads
+2 -2
View File
@@ -25,7 +25,6 @@ import torch
from torch import nn from torch import nn
from transformers import GPTJConfig from transformers import GPTJConfig
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import get_act_fn from sglang.srt.layers.activation import get_act_fn
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
@@ -45,6 +44,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -78,7 +78,7 @@ class GPTJAttention(nn.Module):
prefix=add_prefix("out_proj", prefix), prefix=add_prefix("out_proj", prefix),
) )
tensor_model_parallel_world_size = get_tensor_model_parallel_world_size() tensor_model_parallel_world_size = get_parallel().tp_size
assert total_num_heads % tensor_model_parallel_world_size == 0 assert total_num_heads % tensor_model_parallel_world_size == 0
num_heads = total_num_heads // tensor_model_parallel_world_size num_heads = total_num_heads // tensor_model_parallel_world_size
+17 -24
View File
@@ -28,21 +28,13 @@ from transformers import PretrainedConfig
from sglang.jit_kernel.utils import is_arch_support_pdl from sglang.jit_kernel.utils import is_arch_support_pdl
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_rank,
get_moe_tensor_parallel_world_size,
get_pp_group, get_pp_group,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder 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 import ModelConfigForExpertLocation
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
@@ -76,6 +68,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
LazyValue, LazyValue,
@@ -188,7 +181,7 @@ def _resolve_moe_input_pad_multiple(
# output directly. # output directly.
if quant_config.get_name() != "mxfp4": if quant_config.get_name() != "mxfp4":
return 0 return 0
if get_tensor_model_parallel_world_size() != 1: if get_parallel().tp_size != 1:
# Mid-layer hidden_states still flow through CommunicateWith... # Mid-layer hidden_states still flow through CommunicateWith...
# AllReduceAndLayerNormFn helpers other than `_simple` when # AllReduceAndLayerNormFn helpers other than `_simple` when
# attn_tp_size > 1; those helpers haven't been updated to handle # attn_tp_size > 1; those helpers haven't been updated to handle
@@ -207,7 +200,7 @@ class GptOssSparseMoeBlock(nn.Module):
prefix: str = "", prefix: str = "",
): ):
super().__init__() super().__init__()
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
self.layer_id = layer_id self.layer_id = layer_id
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.activation = config.hidden_act self.activation = config.hidden_act
@@ -358,8 +351,8 @@ class GptOssAttention(nn.Module):
self.hidden_size = hidden_size self.hidden_size = hidden_size
self.sliding_window_size = sliding_window_size self.sliding_window_size = sliding_window_size
attn_tp_rank = get_attention_tp_rank() attn_tp_rank = get_parallel().attn_tp_rank
attn_tp_size = get_attention_tp_size() attn_tp_size = get_parallel().attn_tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % attn_tp_size == 0 assert self.total_num_heads % attn_tp_size == 0
@@ -380,7 +373,7 @@ class GptOssAttention(nn.Module):
self.scaling = self.head_dim**-0.5 self.scaling = self.head_dim**-0.5
self.rope_theta = rope_theta self.rope_theta = rope_theta
self.max_position_embeddings = max_position_embeddings self.max_position_embeddings = max_position_embeddings
self.tp_rank = get_tensor_model_parallel_rank() self.tp_rank = get_parallel().tp_rank
self.qkv_proj = QKVParallelLinear( self.qkv_proj = QKVParallelLinear(
hidden_size, hidden_size,
@@ -535,8 +528,8 @@ class GptOssDecoderLayer(nn.Module):
self.layer_id = layer_id self.layer_id = layer_id
self.attn_tp_size = get_attention_tp_size() self.attn_tp_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_attention_tp_rank() self.attn_tp_rank = get_parallel().attn_tp_rank
# GptOss all layers are sparse and have no nextn now # GptOss all layers are sparse and have no nextn now
self.is_layer_sparse = True self.is_layer_sparse = True
@@ -923,10 +916,10 @@ class GptOssForCausalLM(nn.Module):
loaded_params: set[str] = set() loaded_params: set[str] = set()
mxfp4_block = 32 mxfp4_block = 32
moe_tp_rank = get_moe_tensor_parallel_rank() moe_tp_rank = get_parallel().moe_tp_rank
moe_tp_size = get_moe_tensor_parallel_world_size() moe_tp_size = get_parallel().moe_tp_size
moe_ep_rank = get_moe_expert_parallel_rank() moe_ep_rank = get_parallel().moe_ep_rank
moe_ep_size = get_moe_expert_parallel_world_size() moe_ep_size = get_parallel().moe_ep_size
intermediate_size = self.config.intermediate_size intermediate_size = self.config.intermediate_size
assert ( assert (
@@ -1217,7 +1210,7 @@ class GptOssForCausalLM(nn.Module):
weight_loader = param.weight_loader weight_loader = param.weight_loader
if "bias" not in name: if "bias" not in name:
loaded_weight = loaded_weight.transpose(-2, -1) loaded_weight = loaded_weight.transpose(-2, -1)
if "w2_weight_bias" in name and get_moe_tensor_parallel_rank() != 0: if "w2_weight_bias" in name and get_parallel().moe_tp_rank != 0:
loaded_weight = loaded_weight.zero_() loaded_weight = loaded_weight.zero_()
weight_loader( weight_loader(
@@ -1235,8 +1228,8 @@ class GptOssForCausalLM(nn.Module):
if name in params_dict.keys(): if name in params_dict.keys():
param = params_dict[name] param = params_dict[name]
if "sinks" in name: if "sinks" in name:
start = get_attention_tp_rank() * param.numel() start = get_parallel().attn_tp_rank * param.numel()
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
full_shard_size = param.numel() * tp_size full_shard_size = param.numel() * tp_size
# This handles TP padding: if the checkpoint dim is not divisible by tp_size, # This handles TP padding: if the checkpoint dim is not divisible by tp_size,
# the last TP shard extends beyond `loaded_weight`, pad with zeros before slicing. # the last TP shard extends beyond `loaded_weight`, pad with zeros before slicing.
@@ -1337,7 +1330,7 @@ def _canonicalize_weights(config, weights_in: Iterable[Tuple[str, torch.Tensor]]
def _dequant_mlp_weight(debug_name, w_blocks, w_scales): def _dequant_mlp_weight(debug_name, w_blocks, w_scales):
if get_tensor_model_parallel_rank() == 0: if get_parallel().tp_rank == 0:
logger.info(f"Dequantize {debug_name} start") logger.info(f"Dequantize {debug_name} start")
original_device = w_blocks.device original_device = w_blocks.device
@@ -1348,7 +1341,7 @@ def _dequant_mlp_weight(debug_name, w_blocks, w_scales):
w_bf16 = dequant_mxfp4(w_block=w_blocks, w_scale=w_scales, out_dtype=torch.bfloat16) w_bf16 = dequant_mxfp4(w_block=w_blocks, w_scale=w_scales, out_dtype=torch.bfloat16)
w_bf16 = w_bf16.transpose(-2, -1).contiguous() w_bf16 = w_bf16.transpose(-2, -1).contiguous()
if get_tensor_model_parallel_rank() == 0: if get_parallel().tp_rank == 0:
logger.info( logger.info(
f"Dequantize {debug_name} end {w_blocks.shape=} {w_scales.shape=} {w_bf16.shape=}" f"Dequantize {debug_name} end {w_blocks.shape=} {w_scales.shape=} {w_bf16.shape=}"
) )
+2 -2
View File
@@ -26,7 +26,6 @@ import torch
from torch import nn from torch import nn
from transformers import GraniteConfig from transformers import GraniteConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
@@ -45,6 +44,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
from sglang.utils import get_exception_traceback from sglang.utils import get_exception_traceback
@@ -106,7 +106,7 @@ class GraniteAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
+2 -2
View File
@@ -6,7 +6,6 @@ import torch
from torch import nn from torch import nn
from transformers import GraniteConfig from transformers import GraniteConfig
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
QKVParallelLinear, QKVParallelLinear,
@@ -26,6 +25,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models import mixtral from sglang.srt.models import mixtral
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -105,7 +105,7 @@ class GraniteMoeAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
+5 -4
View File
@@ -5,7 +5,7 @@ from torch import nn
from transformers.models.granitemoeshared import GraniteMoeSharedConfig from transformers.models.granitemoeshared import GraniteMoeSharedConfig
from sglang.srt.configs.granitemoehybrid import GraniteMoeHybridConfig from sglang.srt.configs.granitemoehybrid import GraniteMoeHybridConfig
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size from sglang.srt.distributed import get_pp_group
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend, HybridLinearAttnBackend,
@@ -32,6 +32,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTe
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.transformers import maybe_prefix from sglang.srt.models.transformers import maybe_prefix
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import make_layers from sglang.srt.utils import make_layers
from .granitemoe import GraniteMoeMoE from .granitemoe import GraniteMoeMoE
@@ -112,7 +113,7 @@ class GraniteMoeHybridMambaDecoderLayer(nn.Module):
intermediate_size=config.intermediate_size, intermediate_size=config.intermediate_size,
layer_id=layer_idx, layer_id=layer_idx,
quant_config=quant_config, quant_config=quant_config,
tp_size=get_tensor_model_parallel_world_size(), tp_size=get_parallel().tp_size,
prefix=f"{prefix}.block_sparse_moe", prefix=f"{prefix}.block_sparse_moe",
) )
@@ -192,7 +193,7 @@ class GraniteMoeHybridAttention(nn.Module):
self.total_num_kv_heads = config.num_key_value_heads self.total_num_kv_heads = config.num_key_value_heads
# TensorParallel logic # TensorParallel logic
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
if self.total_num_kv_heads >= tp_size: if self.total_num_kv_heads >= tp_size:
@@ -299,7 +300,7 @@ class GraniteMoeHybridAttentionDecoderLayer(nn.Module):
intermediate_size=config.intermediate_size, intermediate_size=config.intermediate_size,
layer_id=layer_idx, layer_id=layer_idx,
quant_config=quant_config, quant_config=quant_config,
tp_size=get_tensor_model_parallel_world_size(), tp_size=get_parallel().tp_size,
prefix=f"{prefix}.block_sparse_moe", prefix=f"{prefix}.block_sparse_moe",
) )
+10 -14
View File
@@ -23,8 +23,6 @@ from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.layers.activation import GeluAndMul from sglang.srt.layers.activation import GeluAndMul
@@ -60,6 +58,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.loader import DefaultModelLoader from sglang.srt.model_loader.loader import DefaultModelLoader
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
_is_npu = is_npu() _is_npu = is_npu()
@@ -333,8 +332,8 @@ class Grok1Attention(nn.Module):
self.config = config self.config = config
self.layer_id = layer_id self.layer_id = layer_id
self.hidden_size = hidden_size self.hidden_size = hidden_size
attn_tp_rank = get_tensor_model_parallel_rank() attn_tp_rank = get_parallel().tp_rank
attn_tp_size = get_tensor_model_parallel_world_size() attn_tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % attn_tp_size == 0 assert self.total_num_heads % attn_tp_size == 0
self.num_heads = self.total_num_heads // attn_tp_size self.num_heads = self.total_num_heads // attn_tp_size
@@ -542,7 +541,7 @@ class Grok1DecoderLayer(nn.Module):
if self.residual_moe: if self.residual_moe:
# NOTE: self.block_sparse_moe modifies the input in-place, # NOTE: self.block_sparse_moe modifies the input in-place,
# so we have to call it later. Be aware of any possible related errors. # so we have to call it later. Be aware of any possible related errors.
if get_tensor_model_parallel_world_size() > 1: if get_parallel().tp_size > 1:
self.ffn = lambda x: tensor_model_parallel_all_reduce( self.ffn = lambda x: tensor_model_parallel_all_reduce(
self.moe_with_rmoe(x) self.moe_with_rmoe(x)
) )
@@ -593,7 +592,7 @@ class Grok1DecoderLayer(nn.Module):
forward_batch=forward_batch, forward_batch=forward_batch,
) )
if get_tensor_model_parallel_world_size() > 1: if get_parallel().tp_size > 1:
hidden_states = tensor_model_parallel_all_reduce(hidden_states) hidden_states = tensor_model_parallel_all_reduce(hidden_states)
hidden_states, residual = fused_dual_residual_rmsnorm( hidden_states, residual = fused_dual_residual_rmsnorm(
@@ -710,7 +709,7 @@ class Grok1ForCausalLM(nn.Module):
self.load_presharded_moe = ( self.load_presharded_moe = (
getattr(config, "load_presharded_moe", True) getattr(config, "load_presharded_moe", True)
and self.config.num_local_experts > 0 and self.config.num_local_experts > 0
and get_tensor_model_parallel_world_size() > 1 and get_parallel().tp_size > 1
) )
self.load_presharded_attn = getattr(config, "load_presharded_attn", False) self.load_presharded_attn = getattr(config, "load_presharded_attn", False)
self.load_presharded_embedding = getattr( self.load_presharded_embedding = getattr(
@@ -722,7 +721,7 @@ class Grok1ForCausalLM(nn.Module):
config, "replicate_lm_head", default_replicate_lm_head config, "replicate_lm_head", default_replicate_lm_head
) )
if get_tensor_model_parallel_world_size() > 1: if get_parallel().tp_size > 1:
setattr(DefaultModelLoader, "_prepare_weights", _prepare_presharded_weights) setattr(DefaultModelLoader, "_prepare_weights", _prepare_presharded_weights)
self.replicate_embedding = getattr(config, "replicate_embedding", False) self.replicate_embedding = getattr(config, "replicate_embedding", False)
@@ -939,10 +938,7 @@ class Grok1ForCausalLM(nn.Module):
return wq + wkv + out + ffn1 + ffn2 + embed return wq + wkv + out + ffn1 + ffn2 + embed
def get_num_params_torch(self): def get_num_params_torch(self):
return ( return sum(p.numel() for p in self.parameters()) * get_parallel().tp_size
sum(p.numel() for p in self.parameters())
* get_tensor_model_parallel_world_size()
)
old_prepare_weights = getattr(DefaultModelLoader, "_prepare_weights") old_prepare_weights = getattr(DefaultModelLoader, "_prepare_weights")
@@ -954,7 +950,7 @@ def _prepare_presharded_weights(
import glob import glob
import os import os
if get_tensor_model_parallel_world_size() == 1: if get_parallel().tp_size == 1:
return old_prepare_weights(self, model_name_or_path, revision, fall_back_to_pt) return old_prepare_weights(self, model_name_or_path, revision, fall_back_to_pt)
if not os.path.isdir(model_name_or_path): if not os.path.isdir(model_name_or_path):
@@ -971,7 +967,7 @@ def _prepare_presharded_weights(
else: else:
hf_folder = model_name_or_path hf_folder = model_name_or_path
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
# The old format # The old format
allow_patterns = [f"*-{tp_rank:03d}.bin"] allow_patterns = [f"*-{tp_rank:03d}.bin"]
+5 -6
View File
@@ -21,8 +21,6 @@ from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce, tensor_model_parallel_all_reduce,
) )
from sglang.srt.eplb.expert_distribution import ExpertDistributionRecorder from sglang.srt.eplb.expert_distribution import ExpertDistributionRecorder
@@ -52,6 +50,7 @@ from sglang.srt.model_loader.weight_utils import (
kv_cache_scales_loader, kv_cache_scales_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_hip from sglang.srt.utils import is_hip
from sglang.srt.utils.hf_transformers_utils import get_rope_config from sglang.srt.utils.hf_transformers_utils import get_rope_config
@@ -125,7 +124,7 @@ class HunYuanSparseMoeBlock(nn.Module):
layer_id: int = -1, layer_id: int = -1,
): ):
super().__init__() super().__init__()
self.tp_size = get_tensor_model_parallel_world_size() self.tp_size = get_parallel().tp_size
if self.tp_size > config.num_experts: if self.tp_size > config.num_experts:
raise ValueError( raise ValueError(
@@ -263,7 +262,7 @@ class HunYuanAttention(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = hidden_size self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
self.total_num_heads = num_heads self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0 assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size self.num_heads = self.total_num_heads // tp_size
@@ -783,8 +782,8 @@ class HunYuanMoEV1ForCausalLM(nn.Module):
# factors (or else raise an exception). Thus, handled exceptions should # factors (or else raise an exception). Thus, handled exceptions should
# make sure to leave KV cache scale factors in a known good (dummy) state # make sure to leave KV cache scale factors in a known good (dummy) state
def load_kv_cache_scales(self, quantization_param_path: str) -> None: def load_kv_cache_scales(self, quantization_param_path: str) -> None:
tp_size = get_tensor_model_parallel_world_size() tp_size = get_parallel().tp_size
tp_rank = get_tensor_model_parallel_rank() tp_rank = get_parallel().tp_rank
for layer_idx, scaling_factor in kv_cache_scales_loader( for layer_idx, scaling_factor in kv_cache_scales_loader(
quantization_param_path, quantization_param_path,
tp_rank, tp_rank,

Some files were not shown because too many files have changed in this diff Show More