config: read parallel config leaves via get_parallel() (#31816)
This commit is contained in:
@@ -43,7 +43,6 @@ from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import SpecInput
|
||||
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
|
||||
@@ -762,7 +761,7 @@ class TboForwardBatchPreparer:
|
||||
|
||||
# TODO improve, e.g. unify w/ `init_raw`
|
||||
if (
|
||||
get_server_args().moe_dense_tp_size == 1
|
||||
get_parallel().moe_dense_tp_size == 1
|
||||
and batch.global_dp_buffer_len is not None
|
||||
):
|
||||
sum_len = end_token_index - start_token_index
|
||||
|
||||
@@ -572,7 +572,7 @@ class CommonKVManager(BaseKVManager):
|
||||
`Connection refused`, and the leader's `prefill_port_table` ends
|
||||
up missing rows.
|
||||
"""
|
||||
if not self.dist_init_addr or self.server_args.nnodes == 1:
|
||||
if not self.dist_init_addr or get_parallel().nnodes == 1:
|
||||
return local_port
|
||||
|
||||
if not (dist.is_available() and dist.is_initialized()):
|
||||
@@ -624,7 +624,7 @@ class CommonKVManager(BaseKVManager):
|
||||
"rank_port": self.rank_port,
|
||||
"page_size": self.kv_args.page_size,
|
||||
"kv_cache_dtype": get_model().kv_cache_dtype,
|
||||
"load_balance_method": self.server_args.load_balance_method,
|
||||
"load_balance_method": get_parallel().load_balance_method,
|
||||
"enable_dsa_cache_layer_split": getattr(
|
||||
self.server_args, "enable_dsa_cache_layer_split", False
|
||||
),
|
||||
|
||||
@@ -55,7 +55,7 @@ from sglang.srt.observability.trace import (
|
||||
TraceReqContext,
|
||||
trace_set_thread_info,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
|
||||
@@ -948,7 +948,7 @@ class MooncakeKVManager(CommonKVManager):
|
||||
if (
|
||||
self.attn_cp_size > 1
|
||||
and self.attn_cp_rank != 0
|
||||
and not self.server_args.enable_dsa_cache_layer_split
|
||||
and not get_parallel().enable_dsa_cache_layer_split
|
||||
):
|
||||
skip_state = True
|
||||
|
||||
|
||||
@@ -16,6 +16,8 @@ import torch.distributed._symmetric_memory as symm_mem
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Each thread moves _NUMEL_PER_THREAD bf16 via one 128-bit multimem op; the
|
||||
@@ -466,7 +468,6 @@ class MultimemAllGatherer:
|
||||
# Lazy import avoids a module-load dependency on the distributed facade.
|
||||
from sglang.srt.distributed import get_tp_group
|
||||
from sglang.srt.distributed.parallel_state import in_the_same_node_as
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
tp_group = get_tp_group()
|
||||
# Only probe node topology when the deployment can actually span
|
||||
@@ -477,7 +478,7 @@ class MultimemAllGatherer:
|
||||
# EP/mooncake setups, and keep multimem enabled.
|
||||
if (
|
||||
tp_group.world_size > 1
|
||||
and get_server_args().nnodes > 1
|
||||
and get_parallel().nnodes > 1
|
||||
and not all(in_the_same_node_as(tp_group.cpu_group, source_rank=0))
|
||||
):
|
||||
logger.warning(
|
||||
|
||||
@@ -93,6 +93,7 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
|
||||
from sglang.srt.parser.template_detection import resolve_auto_parsers
|
||||
from sglang.srt.parser.template_manager import TemplateManager
|
||||
from sglang.srt.plugins import load_plugins
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
MultiprocessingSerializer,
|
||||
@@ -253,7 +254,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
|
||||
# Initialize ZMQ sockets
|
||||
context = zmq.Context(2)
|
||||
if self.server_args.node_rank == 0:
|
||||
if server_args.node_rank == 0:
|
||||
self.send_to_rpc = get_zmq_socket(
|
||||
context, zmq.DEALER, self.port_args.rpc_ipc_name, True
|
||||
)
|
||||
@@ -301,7 +302,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
routed_dp_rank = data_parallel_rank
|
||||
|
||||
if routed_dp_rank is not None:
|
||||
dp_size = self.server_args.dp_size
|
||||
dp_size = get_parallel().dp_size
|
||||
if dp_size <= 1 and routed_dp_rank == 0:
|
||||
logger.debug(
|
||||
f"routed_dp_rank={routed_dp_rank} is ignored because dp_size={dp_size}"
|
||||
|
||||
@@ -12,7 +12,7 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import
|
||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip
|
||||
from sglang.srt.utils.common import ceil_align, ceil_div
|
||||
|
||||
@@ -76,20 +76,20 @@ def should_use_dsa_fused_topk(
|
||||
|
||||
|
||||
def is_dsa_enable_prefill_cp():
|
||||
return get_server_args().enable_dsa_prefill_context_parallel
|
||||
return get_parallel().enable_dsa_prefill_context_parallel
|
||||
|
||||
|
||||
def is_dsa_prefill_cp_in_seq_split():
|
||||
return (
|
||||
is_dsa_enable_prefill_cp()
|
||||
and get_server_args().dsa_prefill_cp_mode == "in-seq-split"
|
||||
and get_parallel().dsa_prefill_cp_mode == "in-seq-split"
|
||||
)
|
||||
|
||||
|
||||
def is_dsa_prefill_cp_round_robin_split():
|
||||
return (
|
||||
is_dsa_enable_prefill_cp()
|
||||
and get_server_args().dsa_prefill_cp_mode == "round-robin-split"
|
||||
and get_parallel().dsa_prefill_cp_mode == "round-robin-split"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -76,7 +76,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
@@ -271,7 +270,7 @@ class AttnTpContext:
|
||||
def init_context(self, q_lora_rank, is_dsa):
|
||||
self.is_dsa = is_dsa
|
||||
self.allow_input_scattered = (
|
||||
get_server_args().enable_attn_tp_input_scattered
|
||||
get_parallel().enable_attn_tp_input_scattered
|
||||
and (_is_cuda or _is_npu)
|
||||
and q_lora_rank is not None
|
||||
and not is_dsa
|
||||
@@ -282,7 +281,7 @@ class AttnTpContext:
|
||||
and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE)
|
||||
and get_spec().speculative_algorithm != "EAGLE3"
|
||||
)
|
||||
if get_server_args().enable_attn_tp_input_scattered:
|
||||
if get_parallel().enable_attn_tp_input_scattered:
|
||||
if not self.allow_input_scattered:
|
||||
logging.info(
|
||||
"attn_tp_input_scattered is not enabled while other conditions are not met"
|
||||
@@ -440,11 +439,11 @@ class LayerScatterModes:
|
||||
|
||||
|
||||
def enable_moe_dense_fully_dp():
|
||||
return get_server_args().moe_dense_tp_size == 1
|
||||
return get_parallel().moe_dense_tp_size == 1
|
||||
|
||||
|
||||
def enable_dwdp():
|
||||
return get_server_args().dwdp_size > 1
|
||||
return get_parallel().dwdp_size > 1
|
||||
|
||||
|
||||
class LayerCommunicator:
|
||||
|
||||
@@ -47,7 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils.common import (
|
||||
is_cpu,
|
||||
is_npu,
|
||||
@@ -345,7 +345,7 @@ class LogitsProcessor(nn.Module):
|
||||
self.config = config
|
||||
self.vocab_size = config.vocab_size
|
||||
self.logit_scale = logit_scale
|
||||
self.use_attn_tp_group = get_server_args().enable_dp_lm_head
|
||||
self.use_attn_tp_group = get_parallel().enable_dp_lm_head
|
||||
self.use_fp32_lm_head = get_exec().features.enable_fp32_lm_head
|
||||
if self.use_attn_tp_group:
|
||||
self.attn_tp_size = get_parallel().attn_tp_size
|
||||
|
||||
@@ -23,6 +23,7 @@ from sglang.srt.layers.moe.token_dispatcher.deepep import (
|
||||
)
|
||||
from sglang.srt.layers.moe.topk import TopKOutput
|
||||
from sglang.srt.layers.moe.utils import DeepEPMode
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
try:
|
||||
from nixl_ep import Buffer
|
||||
@@ -127,9 +128,7 @@ class NixlEPBuffer:
|
||||
offset = ElasticEPStateManager.get_ep_join_rank_offset()
|
||||
global_rank = rank + offset
|
||||
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
max_ep_size = get_server_args().max_ep_size or world_size
|
||||
max_ep_size = get_parallel().max_ep_size or world_size
|
||||
nixl_max_ranks = max_ep_size
|
||||
|
||||
num_rdma_bytes = 0
|
||||
@@ -226,9 +225,8 @@ class _NixlEPDispatcherImplBase:
|
||||
elastic_state.active_ranks if elastic_state is not None else None
|
||||
)
|
||||
self._active_world_size = dist.get_world_size(group)
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
_max_ep = get_server_args().max_ep_size or self._active_world_size
|
||||
_max_ep = get_parallel().max_ep_size or self._active_world_size
|
||||
self._mask_buffer = (
|
||||
torch.zeros(_max_ep, dtype=torch.int32, device="cuda")
|
||||
if self.active_ranks is not None
|
||||
|
||||
@@ -58,13 +58,13 @@ class ContextParallelMetadata:
|
||||
|
||||
|
||||
def is_prefill_context_parallel_enabled():
|
||||
return get_server_args().enable_prefill_context_parallel
|
||||
return get_parallel().enable_prefill_context_parallel
|
||||
|
||||
|
||||
def is_prefill_cp_in_seq_split():
|
||||
return (
|
||||
is_prefill_context_parallel_enabled()
|
||||
and get_server_args().prefill_cp_mode == "in-seq-split"
|
||||
and get_parallel().prefill_cp_mode == "in-seq-split"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1198,7 +1198,7 @@ class Scheduler(
|
||||
gloo_group=self.attn_tp_cpu_group,
|
||||
tp_rank=self.ps.tp_rank,
|
||||
tp_size=self.ps.tp_size,
|
||||
dp_size=self.server_args.dp_size,
|
||||
dp_size=get_parallel().dp_size,
|
||||
gpu_id=self.ps.gpu_id,
|
||||
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
||||
max_total_num_tokens=self.max_total_num_tokens,
|
||||
@@ -4313,7 +4313,7 @@ class Scheduler(
|
||||
|
||||
old_ep_size = ElasticEPStateManager.get_effective_ep_size()
|
||||
new_ep_size = recv_req.new_ep_size
|
||||
max_ep_size = self.server_args.max_ep_size or old_ep_size
|
||||
max_ep_size = get_parallel().max_ep_size or old_ep_size
|
||||
|
||||
logger.debug(
|
||||
"[Elastic EP][scale] request received: new_ep_size=%d "
|
||||
|
||||
@@ -24,7 +24,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.observability.metrics_collector import DPCooperationInfo
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.utils.common import require_mlp_tp_gather
|
||||
@@ -377,7 +377,7 @@ class SchedulerDPAttnAdapter:
|
||||
def prepare_mlp_sync_batch(self, local_batch: ScheduleBatch):
|
||||
return prepare_mlp_sync_batch_raw(
|
||||
local_batch,
|
||||
dp_size=self.server_args.dp_size,
|
||||
dp_size=get_parallel().dp_size,
|
||||
attn_tp_size=self.ps.attn_tp_size,
|
||||
attn_cp_size=self.ps.attn_cp_size,
|
||||
tp_group=self.tp_group,
|
||||
|
||||
@@ -16,7 +16,7 @@ from sglang.srt.managers.io_struct import (
|
||||
sock_recv,
|
||||
)
|
||||
from sglang.srt.managers.mm_utils import has_shm_features, unwrap_shm_features
|
||||
from sglang.srt.runtime_context import get_disagg
|
||||
from sglang.srt.runtime_context import get_disagg, get_parallel
|
||||
from sglang.srt.utils import broadcast_pyobj, point_to_point_pyobj
|
||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||
|
||||
@@ -127,7 +127,7 @@ class SchedulerRequestReceiver:
|
||||
return recv_reqs
|
||||
|
||||
def _broadcast_reqs_across_ranks(self, recv_reqs: Optional[List]) -> List:
|
||||
if self.server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
|
||||
work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs)
|
||||
else:
|
||||
@@ -156,7 +156,7 @@ class SchedulerRequestReceiver:
|
||||
# instead of the full tp_group. This avoids an expensive
|
||||
# all-ranks gloo sync.
|
||||
_local_ctrl = (
|
||||
self.server_args.enable_dp_attention_local_control_broadcast
|
||||
get_parallel().enable_dp_attention_local_control_broadcast
|
||||
or self.server_args.is_ep_scale_joiner
|
||||
)
|
||||
if _local_ctrl:
|
||||
@@ -233,7 +233,7 @@ class SchedulerRequestReceiver:
|
||||
# peer ranks may still be unpickling ShmPointerMMData
|
||||
# (-> shm_open). Synchronize the same CPU groups that carried
|
||||
# SHM-backed work requests before materialize() unlinks them.
|
||||
if self.server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
if self.ps.attn_tp_size > 1:
|
||||
barrier(group=self.attn_tp_cpu_group)
|
||||
if self.ps.attn_cp_size > 1:
|
||||
|
||||
@@ -36,7 +36,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.observability.req_time_stats import set_time_batch
|
||||
from sglang.srt.runtime_context import get_disagg
|
||||
from sglang.srt.runtime_context import get_disagg, get_parallel
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
|
||||
from sglang.srt.utils.common import get_device_module, is_xpu
|
||||
@@ -123,7 +123,7 @@ class SchedulerPPMixin:
|
||||
next_pp_outputs = None
|
||||
next_batch_result = None
|
||||
d2h_event = None
|
||||
if self.server_args.pp_async_batch_depth > 0:
|
||||
if get_parallel().pp_async_batch_depth > 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -139,7 +139,7 @@ class SchedulerPPMixin:
|
||||
self.mb_metadata,
|
||||
self.last_rank_comm_queue,
|
||||
)
|
||||
if self.server_args.pp_async_batch_depth == 0:
|
||||
if get_parallel().pp_async_batch_depth == 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -269,7 +269,7 @@ class SchedulerPPMixin:
|
||||
server_is_idle = False
|
||||
pp_proxy_tensors = self._pp_recv_proxy_tensors()
|
||||
|
||||
if self.server_args.pp_async_batch_depth > 0:
|
||||
if get_parallel().pp_async_batch_depth > 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -285,7 +285,7 @@ class SchedulerPPMixin:
|
||||
self.mb_metadata,
|
||||
self.last_rank_comm_queue,
|
||||
)
|
||||
if self.server_args.pp_async_batch_depth == 0:
|
||||
if get_parallel().pp_async_batch_depth == 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -428,7 +428,7 @@ class SchedulerPPMixin:
|
||||
pp_proxy_tensors = self._pp_recv_proxy_tensors()
|
||||
|
||||
# early send output if possible
|
||||
if self.server_args.pp_async_batch_depth > 0:
|
||||
if get_parallel().pp_async_batch_depth > 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -446,7 +446,7 @@ class SchedulerPPMixin:
|
||||
self.last_rank_comm_queue,
|
||||
)
|
||||
|
||||
if self.server_args.pp_async_batch_depth == 0:
|
||||
if get_parallel().pp_async_batch_depth == 0:
|
||||
next_pp_outputs, next_batch_result, d2h_event = (
|
||||
self._pp_commit_send_output_work_and_preprocess_output_tensors(
|
||||
next_first_rank_mb_id,
|
||||
@@ -557,10 +557,10 @@ class SchedulerPPMixin:
|
||||
self.on_idle()
|
||||
|
||||
def init_pp_loop_state(self: Scheduler):
|
||||
self.pp_loop_size: int = self.ps.pp_size + self.server_args.pp_async_batch_depth
|
||||
self.pp_loop_size: int = self.ps.pp_size + get_parallel().pp_async_batch_depth
|
||||
# In CP mode, attention weights are duplicated, eliminating the need for the attention TP all-gather operation.
|
||||
self.require_attn_tp_allgather = (
|
||||
not self.server_args.enable_dsa_prefill_context_parallel
|
||||
not get_parallel().enable_dsa_prefill_context_parallel
|
||||
)
|
||||
self.mbs = [None] * self.pp_loop_size
|
||||
self.last_mbs = [None] * self.pp_loop_size
|
||||
|
||||
@@ -74,7 +74,7 @@ from sglang.srt.managers.io_struct import (
|
||||
UpdateWeightsFromTensorReqOutput,
|
||||
)
|
||||
from sglang.srt.managers.load_snapshot import LoadSnapshot
|
||||
from sglang.srt.runtime_context import get_lora
|
||||
from sglang.srt.runtime_context import get_lora, get_parallel
|
||||
from sglang.srt.server_args import LoRARef, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
@@ -146,8 +146,8 @@ class TokenizerControlMixin:
|
||||
|
||||
def update_control_communicator_fan_out(self: TokenizerManager, worker_count: int):
|
||||
primary_group_control = (
|
||||
self.server_args.enable_dp_attention
|
||||
and not self.server_args.enable_dp_attention_local_control_broadcast
|
||||
get_parallel().enable_dp_attention
|
||||
and not get_parallel().enable_dp_attention_local_control_broadcast
|
||||
)
|
||||
if primary_group_control:
|
||||
control_fan_out = (
|
||||
@@ -397,7 +397,7 @@ class TokenizerControlMixin:
|
||||
) -> Tuple[bool, str]:
|
||||
self.auto_create_handle_loop()
|
||||
assert (
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for update weights from distributed"
|
||||
|
||||
results = await self.init_weights_update_group_communicator(obj)
|
||||
@@ -410,7 +410,7 @@ class TokenizerControlMixin:
|
||||
) -> Tuple[bool, str]:
|
||||
self.auto_create_handle_loop()
|
||||
assert (
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for destroy parameter update group"
|
||||
|
||||
results = await self.destroy_weights_update_group_communicator(obj)
|
||||
@@ -423,7 +423,7 @@ class TokenizerControlMixin:
|
||||
) -> Tuple[bool, str]:
|
||||
self.auto_create_handle_loop()
|
||||
assert (
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for update weights from distributed"
|
||||
|
||||
if obj.abort_all_requests:
|
||||
@@ -454,7 +454,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
# TODO: support DP
|
||||
assert (
|
||||
self.server_args.dp_size == 1
|
||||
get_parallel().dp_size == 1
|
||||
), "dp_size must be 1 for init_weights_send_group_for_remote_instance"
|
||||
result = (
|
||||
await self.init_weights_send_group_for_remote_instance_communicator(obj)
|
||||
@@ -469,7 +469,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
# TODO: support DP
|
||||
assert (
|
||||
self.server_args.dp_size == 1
|
||||
get_parallel().dp_size == 1
|
||||
), "dp_size must be 1 for send_weights_to_remote_instance"
|
||||
result = (await self.send_weights_to_remote_instance_communicator(obj))[0]
|
||||
return result.success, result.message
|
||||
@@ -481,7 +481,7 @@ class TokenizerControlMixin:
|
||||
) -> Tuple[bool, str]:
|
||||
self.auto_create_handle_loop()
|
||||
assert (
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for update weights from tensor"
|
||||
|
||||
if obj.abort_all_requests:
|
||||
@@ -517,7 +517,7 @@ class TokenizerControlMixin:
|
||||
try:
|
||||
# For now, we only support single data parallel instance
|
||||
assert (
|
||||
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
|
||||
), "dp_size must be 1 or dp attention must be enabled for update weights from IPC"
|
||||
logger.info("Starting IPC weight update")
|
||||
|
||||
@@ -578,7 +578,7 @@ class TokenizerControlMixin:
|
||||
# TODO (lifuhuang): Remove this after we verify that dynamic lora loading works
|
||||
# with dp_size > 1.
|
||||
assert (
|
||||
self.server_args.dp_size == 1
|
||||
get_parallel().dp_size == 1
|
||||
), "dp_size must be 1 for dynamic lora loading"
|
||||
logger.info(
|
||||
"Start load Lora adapter. Lora name=%s, path=%s",
|
||||
@@ -654,7 +654,7 @@ class TokenizerControlMixin:
|
||||
)
|
||||
|
||||
assert (
|
||||
self.server_args.dp_size == 1
|
||||
get_parallel().dp_size == 1
|
||||
), "dp_size must be 1 for dynamic lora loading"
|
||||
logger.info(
|
||||
"Start load Lora adapter from tensors. Lora name=%s",
|
||||
@@ -730,7 +730,7 @@ class TokenizerControlMixin:
|
||||
# TODO (lifuhuang): Remove this after we verify that dynamic lora loading works
|
||||
# with dp_size > 1.
|
||||
assert (
|
||||
self.server_args.dp_size == 1
|
||||
get_parallel().dp_size == 1
|
||||
), "dp_size must be 1 for dynamic lora loading"
|
||||
logger.info(
|
||||
"Start unload Lora adapter. Lora name=%s",
|
||||
@@ -750,7 +750,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
results = await self.get_weights_by_name_communicator(obj)
|
||||
all_parameters = [r.parameter for r in results]
|
||||
if self.server_args.dp_size == 1:
|
||||
if get_parallel().dp_size == 1:
|
||||
return all_parameters[0]
|
||||
else:
|
||||
return all_parameters
|
||||
|
||||
@@ -116,6 +116,7 @@ from sglang.srt.runtime_context import (
|
||||
get_lora,
|
||||
get_model,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
@@ -1365,7 +1366,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
return batch_size > 0 and (
|
||||
self.server_args.enable_tokenizer_batch_encode
|
||||
or (
|
||||
(not self.server_args.enable_dp_attention)
|
||||
(not get_parallel().enable_dp_attention)
|
||||
and (not self._batch_has_text(batch_size, requests))
|
||||
)
|
||||
)
|
||||
|
||||
@@ -150,6 +150,7 @@ from sglang.srt.runtime_context import (
|
||||
get_global_dwdp_manager,
|
||||
get_lora,
|
||||
get_model,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
set_global_dwdp_manager,
|
||||
)
|
||||
@@ -392,12 +393,12 @@ class ModelRunner:
|
||||
is_scale_join = get_exec().moe.ep_join_mode == "scale"
|
||||
if is_scale_join:
|
||||
join_effective_ep_size = (
|
||||
self.server_args.ep_join_rank_offset + self.ps.tp_size
|
||||
get_parallel().ep_join_rank_offset + self.ps.tp_size
|
||||
)
|
||||
dist.barrier(group=self.tp_group.cpu_group)
|
||||
if self.ps.tp_rank == 0:
|
||||
register_scale_cohort(
|
||||
self.server_args.ep_join_rank_offset,
|
||||
get_parallel().ep_join_rank_offset,
|
||||
join_effective_ep_size,
|
||||
)
|
||||
join_scale_process_group()
|
||||
@@ -407,7 +408,7 @@ class ModelRunner:
|
||||
else:
|
||||
join_process_groups()
|
||||
|
||||
global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
||||
global_ep_rank = self.ps.tp_rank + get_parallel().ep_join_rank_offset
|
||||
broadcast_global_expert_location_metadata(
|
||||
model_config=self.model_config,
|
||||
moe_ep_rank=global_ep_rank,
|
||||
@@ -441,9 +442,9 @@ class ModelRunner:
|
||||
new_dp_size=join_effective_ep_size,
|
||||
new_dp_rank=global_ep_rank,
|
||||
)
|
||||
self.server_args.override(
|
||||
"elastic_ep.scale_join", dp_size=join_effective_ep_size
|
||||
)
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override("elastic_ep.scale_join", dp_size=join_effective_ep_size)
|
||||
if self.eplb_manager is not None:
|
||||
self.eplb_manager.disable_rebalance(
|
||||
"EPLB rebalance is disabled after elastic EP scale-up"
|
||||
@@ -617,7 +618,7 @@ class ModelRunner:
|
||||
if self.is_draft_worker:
|
||||
return
|
||||
expert_rank = self.ps.moe_ep_rank + (
|
||||
self.server_args.ep_join_rank_offset
|
||||
get_parallel().ep_join_rank_offset
|
||||
if self.server_args.is_ep_scale_joiner
|
||||
else 0
|
||||
)
|
||||
@@ -797,7 +798,7 @@ class ModelRunner:
|
||||
device=self.device,
|
||||
tp_group=(
|
||||
self.attention_tp_group.cpu_group
|
||||
if self.server_args.enable_dp_attention
|
||||
if get_parallel().enable_dp_attention
|
||||
else self.tp_group.cpu_group
|
||||
),
|
||||
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
||||
@@ -1608,7 +1609,7 @@ class ModelRunner:
|
||||
if added <= 0:
|
||||
return
|
||||
|
||||
initial_ep_size = self.server_args.elastic_ep_initial_size
|
||||
initial_ep_size = get_parallel().elastic_ep_initial_size
|
||||
assert initial_ep_size is not None
|
||||
self.server_args.override("elastic_ep.scale", ep_size=effective_size)
|
||||
|
||||
@@ -1627,7 +1628,7 @@ class ModelRunner:
|
||||
set_global_expert_location_metadata(new_metadata, allow_overwrite=True)
|
||||
|
||||
def _elastic_global_rank(self) -> int:
|
||||
return self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
||||
return self.ps.tp_rank + get_parallel().ep_join_rank_offset
|
||||
|
||||
def _report_elastic_scale_failure(self, error: str, effective_size: int) -> None:
|
||||
if self.ps.tp_rank != 0 or self.server_args.is_ep_scale_joiner:
|
||||
@@ -1704,7 +1705,9 @@ class ModelRunner:
|
||||
new_dp_size=target_size,
|
||||
new_dp_rank=self._elastic_global_rank(),
|
||||
)
|
||||
self.server_args.override("elastic_ep.scale", dp_size=target_size)
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override("elastic_ep.scale", dp_size=target_size)
|
||||
|
||||
ElasticEPStateManager.mark_syncing_new_world()
|
||||
self._elastic_scale_ready_barrier(
|
||||
|
||||
+3
-3
@@ -11,7 +11,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||
RemoteInstanceWeightLoaderBackend,
|
||||
register_memory_region,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_model
|
||||
from sglang.srt.runtime_context import get_model, get_parallel
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
||||
|
||||
@@ -76,11 +76,11 @@ class RemoteInstanceWeightTransporter:
|
||||
"""
|
||||
import requests as http_requests
|
||||
|
||||
if self.server_args.dist_init_addr:
|
||||
if get_parallel().dist_init_addr:
|
||||
# Multi-node: bootstrap server is on the head node (node_rank==0).
|
||||
# Derive host from dist_init_addr (shared across all nodes).
|
||||
bootstrap_host = (
|
||||
NetworkAddress.parse(self.server_args.dist_init_addr).resolved().host
|
||||
NetworkAddress.parse(get_parallel().dist_init_addr).resolved().host
|
||||
)
|
||||
else:
|
||||
bootstrap_host = "127.0.0.1"
|
||||
|
||||
@@ -26,9 +26,7 @@ import torch
|
||||
from torch import nn
|
||||
from transformers import ApertusConfig
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
)
|
||||
from sglang.srt.distributed import get_pp_group
|
||||
from sglang.srt.layers.activation import XIELU
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
@@ -52,7 +50,7 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
kv_cache_scales_loader,
|
||||
maybe_remap_kv_scale_name,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix, make_layers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -442,7 +440,7 @@ class ApertusForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
||||
|
||||
@@ -20,9 +20,7 @@ import torch
|
||||
from torch import nn
|
||||
from transformers import LlamaConfig
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
)
|
||||
from sglang.srt.distributed import get_pp_group
|
||||
from sglang.srt.layers.activation import get_act_fn
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
@@ -46,7 +44,7 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
kv_cache_scales_loader,
|
||||
maybe_remap_kv_scale_name,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix, make_layers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -405,7 +403,7 @@ class ArceeForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
||||
|
||||
@@ -79,7 +79,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
|
||||
@@ -821,7 +820,7 @@ class BailingMoEForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -57,7 +57,6 @@ from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
@@ -1085,7 +1084,7 @@ class BailingMoELinearForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
params_dtype=torch.float32,
|
||||
quant_config=quant_config,
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
@@ -42,7 +42,7 @@ from sglang.srt.models.bailing_moe_linear import (
|
||||
BailingMoeV2_5ForCausalLM,
|
||||
)
|
||||
from sglang.srt.models.utils import WeightsMapper
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import BumpAllocator, add_prefix
|
||||
|
||||
LoraConfig = None
|
||||
@@ -208,7 +208,7 @@ class BailingMoeForCausalLMNextN(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("model.shared_head.head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
if hasattr(self.config, "model_type") and config.model_type == "bailing_hybrid":
|
||||
|
||||
@@ -62,7 +62,6 @@ from sglang.srt.models.utils import WeightsMapper
|
||||
from sglang.srt.runtime_context import (
|
||||
get_model,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
|
||||
@@ -381,7 +380,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("model.shared_head.head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -2719,7 +2719,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
else:
|
||||
# ranks other than the last rank will have a placeholder layer
|
||||
|
||||
@@ -134,7 +134,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
|
||||
if not _is_hip:
|
||||
@@ -2404,7 +2403,7 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
else:
|
||||
self.lm_head = PPMissingLayer()
|
||||
|
||||
@@ -38,7 +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_context import get_attn_backend
|
||||
from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -233,7 +233,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("model.shared_head.head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -28,7 +28,7 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
default_weight_loader,
|
||||
maybe_remap_kv_scale_name,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix, make_layers
|
||||
from sglang.utils import get_exception_traceback, logger
|
||||
|
||||
@@ -439,7 +439,7 @@ class Exaone4ForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
@@ -60,7 +60,6 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
|
||||
@@ -643,7 +642,7 @@ class ExaoneMoEForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
# For EAGLE3 support
|
||||
|
||||
@@ -30,7 +30,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.models.exaone_moe import ExaoneMoEForCausalLM, ExaoneMoEModel
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -63,7 +63,7 @@ class ExaoneMoEForCausalLMMTP(ExaoneMoEForCausalLM):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -13,9 +13,7 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
|
||||
)
|
||||
from sglang.srt.layers.attention.mamba.mamba import MambaMixer2
|
||||
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
||||
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.linear import (
|
||||
MergedColumnParallelLinear,
|
||||
@@ -36,7 +34,6 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.runtime_context import (
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import add_prefix, is_cuda, make_layers
|
||||
@@ -477,7 +474,7 @@ class FalconH1ForCausalLM(nn.Module):
|
||||
quant_config=quant_config,
|
||||
org_num_embeddings=config.vocab_size,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.lm_head = self.lm_head.float()
|
||||
self.lm_head_multiplier = config.lm_head_multiplier
|
||||
|
||||
@@ -87,7 +87,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
@@ -1171,7 +1170,7 @@ class Glm4MoeForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -78,7 +78,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
@@ -908,7 +907,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -35,7 +35,7 @@ from sglang.srt.models.glm4_moe_lite import (
|
||||
Glm4MoeLiteDecoderLayer,
|
||||
Glm4MoeLiteForCausalLM,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_spec
|
||||
from sglang.srt.utils import BumpAllocator, add_prefix, is_npu
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -151,7 +151,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("model.shared_head.head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_spec
|
||||
from sglang.srt.utils import add_prefix, is_npu
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -137,7 +137,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("model.shared_head.head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -18,7 +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.models.glm4_moe import Glm4MoeModel
|
||||
from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel
|
||||
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel
|
||||
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
|
||||
|
||||
@@ -69,7 +69,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
else:
|
||||
# ranks other than the last rank will have a placeholder layer
|
||||
|
||||
@@ -33,7 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.models.glm4 import Glm4DecoderLayer
|
||||
from sglang.srt.models.glm_ocr import GlmOcrForConditionalGeneration
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -134,7 +134,7 @@ class GlmOcrForConditionalGenerationNextN(GlmOcrForConditionalGeneration):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("model.shared_head.head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -70,7 +70,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
LazyValue,
|
||||
@@ -257,7 +256,7 @@ class GptOssSparseMoeBlock(nn.Module):
|
||||
hidden_states: torch.Tensor,
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
) -> torch.Tensor:
|
||||
if get_server_args().dwdp_size > 1:
|
||||
if get_parallel().dwdp_size > 1:
|
||||
return self.forward_dwdp(hidden_states)
|
||||
|
||||
if not get_moe_a2a_backend().is_deepep():
|
||||
@@ -775,7 +774,7 @@ class GptOssForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
# quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
self.capture_aux_hidden_states = False
|
||||
|
||||
@@ -49,7 +49,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.utils import LazyValue, add_prefix, make_layers
|
||||
|
||||
@@ -637,7 +636,7 @@ class LagunaForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
else:
|
||||
self.lm_head = PPMissingLayer()
|
||||
|
||||
@@ -78,7 +78,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
@@ -828,7 +827,7 @@ class LLaDA2MoeModelLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config, return_full_logits=True)
|
||||
|
||||
|
||||
@@ -25,10 +25,7 @@ import torch
|
||||
from torch import nn
|
||||
from transformers import LlamaConfig
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
get_pp_indices,
|
||||
)
|
||||
from sglang.srt.distributed import get_pp_group, get_pp_indices
|
||||
from sglang.srt.layers.activation import SiluAndMul
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
@@ -53,7 +50,7 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
maybe_remap_kv_scale_name,
|
||||
)
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix, is_cuda, is_npu, is_xpu, make_layers
|
||||
from sglang.utils import get_exception_traceback
|
||||
|
||||
@@ -530,7 +527,7 @@ class LlamaForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
|
||||
|
||||
@@ -41,17 +41,13 @@ from sglang.jit_kernel.dsv4 import linear_bf16_fp32
|
||||
from sglang.kernels.ops.moe.ep_moe_kernels import zero_experts_compute_triton
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
|
||||
from sglang.srt.configs import LongcatFlashConfig
|
||||
from sglang.srt.distributed import (
|
||||
tensor_model_parallel_all_reduce,
|
||||
)
|
||||
from sglang.srt.distributed import tensor_model_parallel_all_reduce
|
||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.activation import SiluAndMul
|
||||
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
||||
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.linear import (
|
||||
MergedColumnParallelLinear,
|
||||
@@ -87,7 +83,7 @@ from sglang.srt.model_loader.utils import (
|
||||
)
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
|
||||
from sglang.srt.runtime_context import get_parallel, get_stream
|
||||
from sglang.srt.utils import (
|
||||
BumpAllocator,
|
||||
add_prefix,
|
||||
@@ -714,7 +710,7 @@ class LongcatFlashForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
self.capture_aux_hidden_states = False
|
||||
|
||||
@@ -51,7 +51,7 @@ from sglang.srt.models.utils import (
|
||||
create_fused_set_kv_buffer_arg,
|
||||
enable_fused_set_kv_buffer,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import add_prefix, is_cuda
|
||||
|
||||
_is_cuda = is_cuda()
|
||||
@@ -520,7 +520,7 @@ class MellumForCausalLM(Qwen3MoeForCausalLM):
|
||||
cfg.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(cfg)
|
||||
self.capture_aux_hidden_states = False
|
||||
|
||||
@@ -77,7 +77,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
LazyValue,
|
||||
@@ -1187,7 +1186,7 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
else:
|
||||
self.lm_head = PPMissingLayer()
|
||||
|
||||
@@ -26,9 +26,7 @@ from sglang.srt.layers.communicator import (
|
||||
LayerScatterModes,
|
||||
enable_moe_dense_fully_dp,
|
||||
)
|
||||
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.logits_processor import LogitsProcessor
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
@@ -44,7 +42,7 @@ from sglang.srt.models.mimo_v2 import (
|
||||
MiMoV2MLP,
|
||||
load_mimo_v2_qkv_proj_weight,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
MiMoV2Config = None
|
||||
@@ -259,7 +257,7 @@ class MiMoV2MTP(MiMoV2ForCausalLM):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -76,7 +76,7 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
maybe_remap_kv_scale_name,
|
||||
)
|
||||
from sglang.srt.models.minimax_m2 import MiniMaxM2RMSNormTP
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import (
|
||||
add_prefix,
|
||||
get_device_sm,
|
||||
@@ -1438,7 +1438,7 @@ class MiniMaxM3SparseForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
@@ -105,7 +105,7 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module):
|
||||
text_config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("language_model.lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
else:
|
||||
self.lm_head = PPMissingLayer()
|
||||
|
||||
@@ -90,7 +90,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
add_prefix,
|
||||
@@ -939,7 +938,7 @@ class NemotronHForCausalLM(nn.Module):
|
||||
else lora_config.lora_vocab_padding_size
|
||||
),
|
||||
quant_config=quant_config,
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -19,10 +19,7 @@ from torch import nn
|
||||
|
||||
from sglang.srt.configs import NemotronHConfig
|
||||
from sglang.srt.distributed import get_pp_group
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
attn_tp_all_reduce,
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import attn_tp_all_reduce, is_dp_attention_enabled
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import ColumnParallelLinear
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
@@ -38,7 +35,7 @@ from sglang.srt.models.nemotron_h import (
|
||||
NemotronHMoEDecoderLayer,
|
||||
)
|
||||
from sglang.srt.models.nemotron_h_utils import is_attn_layer
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
|
||||
@@ -338,7 +335,7 @@ class NemotronHForCausalLMMTP(NemotronHForCausalLM):
|
||||
self.config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
@@ -93,7 +93,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
add_prefix,
|
||||
@@ -1011,7 +1010,7 @@ class Qwen2MoeForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
# For EAGLE3 support
|
||||
|
||||
@@ -34,7 +34,6 @@ from sglang.srt.models.utils import apply_qk_norm
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu
|
||||
@@ -495,7 +494,7 @@ class Qwen3ForCausalLM(nn.Module):
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -76,7 +76,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
@@ -956,7 +955,7 @@ class Qwen3MoeForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
self.capture_aux_hidden_states = False
|
||||
|
||||
@@ -30,7 +30,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.models.qwen3_moe import Qwen3MoeForCausalLM, Qwen3MoeModel
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -63,7 +63,7 @@ class Qwen3MoeForCausalLMMTP(Qwen3MoeForCausalLM):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -15,9 +15,7 @@ from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_r
|
||||
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
|
||||
from sglang.srt.layers.attention.mamba.mamba import mamba_v2_sharded_weight_loader
|
||||
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
||||
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 GemmaRMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
@@ -50,7 +48,6 @@ from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock
|
||||
from sglang.srt.runtime_context import (
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
@@ -831,9 +828,7 @@ class Qwen3HybridAttentionDecoderLayer(nn.Module):
|
||||
|
||||
if self.attn_output_gate:
|
||||
if _is_hip:
|
||||
from sglang.jit_kernel.triton.sigmoid_gate_mul import (
|
||||
sigmoid_gate_mul,
|
||||
)
|
||||
from sglang.jit_kernel.triton.sigmoid_gate_mul import sigmoid_gate_mul
|
||||
|
||||
attn_output = sigmoid_gate_mul(attn_output, gate)
|
||||
else:
|
||||
@@ -1030,7 +1025,7 @@ class Qwen3NextForCausalLM(nn.Module):
|
||||
quant_config=quant_config,
|
||||
org_num_embeddings=config.vocab_size,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
# For EAGLE3 support
|
||||
|
||||
@@ -35,7 +35,6 @@ from sglang.srt.models.qwen3_next import Qwen3NextForCausalLM, Qwen3NextModel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_model,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.utils import add_prefix, is_npu
|
||||
@@ -85,7 +84,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("model.shared_head.head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
# Mirror Qwen3NextForCausalLM.__init__'s shared-expert fusion setup so
|
||||
|
||||
@@ -68,7 +68,7 @@ from sglang.srt.models.utils import (
|
||||
)
|
||||
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
|
||||
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
||||
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel
|
||||
from sglang.srt.utils import add_prefix, cpu_has_amx_support, is_cpu, is_npu, round_up
|
||||
from sglang.srt.utils.hf_transformers_utils import get_processor
|
||||
|
||||
@@ -1269,7 +1269,7 @@ class Qwen3VLForConditionalGeneration(nn.Module):
|
||||
self.config.vocab_size,
|
||||
self.config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -1226,7 +1226,7 @@ class SarvamMLAForCausalLM(nn.Module):
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
)
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
|
||||
@@ -43,7 +43,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import add_prefix, is_cuda, make_layers
|
||||
@@ -473,7 +472,7 @@ class SDARForCausalLM(nn.Module):
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -56,7 +56,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
|
||||
@@ -557,7 +556,7 @@ class SDARMoeForCausalLM(nn.Module):
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -45,7 +45,6 @@ from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
|
||||
@@ -817,7 +816,7 @@ class Step3p5ForCausalLM(nn.Module):
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
use_attn_tp_group=get_parallel().enable_dp_lm_head,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -375,7 +375,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
self, batch: ScheduleBatch, on_publish
|
||||
) -> GenerationBatchResult:
|
||||
if batch.forward_mode.is_idle():
|
||||
if self.server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
self.target_worker.forward_batch_generation(
|
||||
batch, capture_hidden_mode=CaptureHiddenMode.FULL
|
||||
)
|
||||
@@ -443,7 +443,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
def _dp_verify_tier_num_tokens(self, batch: ScheduleBatch) -> Optional[int]:
|
||||
if not (
|
||||
self._draft_is_moe
|
||||
and self.server_args.enable_dp_attention
|
||||
and get_parallel().enable_dp_attention
|
||||
and batch.global_num_tokens is not None
|
||||
and self._verify_planner.is_compact_mode
|
||||
):
|
||||
@@ -487,7 +487,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
|
||||
if batch.forward_mode.is_idle():
|
||||
self._observers.note_idle_decode_step()
|
||||
if self.server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
if self._draft_is_moe:
|
||||
self._proposer.run_idle_participation(batch)
|
||||
self._verify_executor.run_idle_participation(
|
||||
@@ -549,7 +549,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
global_num_reqs = (
|
||||
max(batch.global_num_tokens)
|
||||
if self._draft_is_moe
|
||||
and self.server_args.enable_dp_attention
|
||||
and get_parallel().enable_dp_attention
|
||||
and batch.global_num_tokens is not None
|
||||
else None
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user