Add get_parallel(): a structured accessor for parallel-topology state (#28567)
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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],
|
||||||
|
|||||||
+3
-2
@@ -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(
|
||||||
|
|||||||
+2
-2
@@ -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],
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+686
-687
File diff suppressed because it is too large
Load Diff
@@ -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,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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 = (
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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":
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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."
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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."
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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=}"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,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",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
@@ -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
Reference in New Issue
Block a user