Revert RuntimeContext config-namespace reads/roles (#31813–#31817) (#32100)

This commit is contained in:
Cheng Wan
2026-07-22 11:52:41 -07:00
committed by GitHub
parent 0bdd4730af
commit f5dcbe8f14
187 changed files with 1432 additions and 1441 deletions
+11 -9
View File
@@ -261,17 +261,19 @@ def mamba_extra_buffer_of(cfg: Any) -> bool:
def declare_load_time_override(source: str, declared: Dict[str, Any]) -> None: def declare_load_time_override(source: str, declared: Dict[str, Any]) -> None:
"""Declare a load-time resolved field (model-file config overrides, """Declare a load-time resolved field (model-file config overrides,
weight-resolved dtypes) after publish. it is written to the config weight-resolved dtypes) on the published ``server_args``: resolution has
bags via ``get_context().override`` (namespace readers see it); server_args already materialized, so the declaration writes through, joining the
stays the pristine startup record. Validated against the resolvable declaration stash for provenance and republish consistency."""
whitelist first."""
from sglang.srt.runtime_context import get_context from sglang.srt.runtime_context import get_context
context = get_context() server_args = get_context().server_args
validate_declarations(context.server_args, [(source, dict(declared))]) validate_declarations(server_args, [(source, dict(declared))])
# write the config bags (namespace readers see it); server_args override = getattr(server_args, "override", None)
# stays the pristine startup record. if override is not None:
context.override(source, **declared) override(source, **declared)
else:
# Config-shaped fixtures without the mutation entry point.
_apply_fields(server_args, declared)
def collect_model_override_declarations( def collect_model_override_declarations(
@@ -39,11 +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 ( from sglang.srt.runtime_context import get_parallel, get_server_args
get_device,
get_exec,
get_parallel,
)
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
@@ -187,7 +183,7 @@ def _update_device_and_sum_field_from_cpu_field(
cpu_value cpu_value
if isinstance(cpu_value, torch.Tensor) if isinstance(cpu_value, torch.Tensor)
else torch.tensor(cpu_value, dtype=old_device_value.dtype) else torch.tensor(cpu_value, dtype=old_device_value.dtype)
).to(device=get_device().device, non_blocking=True) ).to(device=get_server_args().device, non_blocking=True)
setattr(batch, device_field, new_device_value) setattr(batch, device_field, new_device_value)
if sum_field is not None: if sum_field is not None:
@@ -339,7 +335,7 @@ def compute_split_indices_for_cuda_graph_replay(
class TboCudaGraphRunnerPlugin: class TboCudaGraphRunnerPlugin:
def __init__(self): def __init__(self):
self._tbo_children_num_token_non_padded = torch.zeros( self._tbo_children_num_token_non_padded = torch.zeros(
(2,), dtype=torch.int32, device=get_device().device (2,), dtype=torch.int32, device=get_server_args().device
) )
def capture_one_batch_size(self, batch: ForwardBatch, num_tokens: int): def capture_one_batch_size(self, batch: ForwardBatch, num_tokens: int):
@@ -637,7 +633,7 @@ class TboForwardBatchPreparer:
sum_field=None, sum_field=None,
) )
_, child_b.extend_start_loc = compute_position( _, child_b.extend_start_loc = compute_position(
get_exec().kernel.attention_backend, get_server_args().attention_backend,
child_b.extend_prefix_lens, child_b.extend_prefix_lens,
child_b.extend_seq_lens, child_b.extend_seq_lens,
child_b.extend_num_tokens, child_b.extend_num_tokens,
@@ -761,7 +757,7 @@ class TboForwardBatchPreparer:
# TODO improve, e.g. unify w/ `init_raw` # TODO improve, e.g. unify w/ `init_raw`
if ( if (
get_parallel().moe_dense_tp_size == 1 get_server_args().moe_dense_tp_size == 1
and batch.global_dp_buffer_len is not None and batch.global_dp_buffer_len is not None
): ):
sum_len = end_token_index - start_token_index sum_len = end_token_index - start_token_index
@@ -836,7 +832,7 @@ class TboForwardBatchPreparer:
value_a = min(tbo_split_token_index, num_token_non_padded) value_a = min(tbo_split_token_index, num_token_non_padded)
value_b = max(0, num_token_non_padded - tbo_split_token_index) value_b = max(0, num_token_non_padded - tbo_split_token_index)
return torch.tensor([value_a, value_b], dtype=torch.int32).to( return torch.tensor([value_a, value_b], dtype=torch.int32).to(
device=get_device().device, non_blocking=True device=get_server_args().device, non_blocking=True
) )
@classmethod @classmethod
+2 -2
View File
@@ -8,7 +8,6 @@ from transformers import CONFIG_MAPPING
from transformers.configuration_utils import PretrainedConfig from transformers.configuration_utils import PretrainedConfig
from sglang.srt.configs.mamba_utils import BaseLinearStateParams from sglang.srt.configs.mamba_utils import BaseLinearStateParams
from sglang.srt.runtime_context import get_exec
class InklingModelConfig(PretrainedConfig): class InklingModelConfig(PretrainedConfig):
@@ -225,8 +224,9 @@ class InklingModelConfig(PretrainedConfig):
self.swa_num_key_value_heads, self.swa_head_dim self.swa_num_key_value_heads, self.swa_head_dim
) )
stream_dim = self.hidden_size stream_dim = self.hidden_size
from sglang.srt.runtime_context import get_server_args
if get_exec().comm.enable_scattered_sconv: if get_server_args().enable_scattered_sconv:
# Scattered sconv: the attn/mlp output sconvs run on the [T, H/P] # Scattered sconv: the attn/mlp output sconvs run on the [T, H/P]
# hidden shard, so their conv-state caches shard with them. # hidden shard, so their conv-state caches shard with them.
assert ( assert (
@@ -14,7 +14,6 @@ from sglang.srt.constrained.base_grammar_backend import (
from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject
from sglang.srt.distributed.communication_tags import P2PTag from sglang.srt.distributed.communication_tags import P2PTag
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_serving
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.io_struct import AbortReq from sglang.srt.managers.io_struct import AbortReq
@@ -29,7 +28,7 @@ class GrammarManager:
self.scheduler = scheduler self.scheduler = scheduler
self.server_args = scheduler.server_args self.server_args = scheduler.server_args
self.grammar_queue: List[Req] = [] self.grammar_queue: List[Req] = []
if not get_serving().skip_tokenizer_init: if not self.server_args.skip_tokenizer_init:
self.grammar_backend = create_grammar_backend( self.grammar_backend = create_grammar_backend(
self.server_args, self.server_args,
scheduler.tokenizer, scheduler.tokenizer,
@@ -32,8 +32,11 @@ from sglang.srt.disaggregation.utils import (
) )
from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.distributed import get_pp_group, get_world_group
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import get_attention_dp_rank, get_attention_dp_size from sglang.srt.layers.dp_attention import (
from sglang.srt.runtime_context import get_model, get_parallel, get_serving get_attention_dp_rank,
get_attention_dp_size,
)
from sglang.srt.runtime_context import get_model, get_parallel
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import ( from sglang.srt.utils.network import (
NetworkAddress, NetworkAddress,
@@ -572,7 +575,7 @@ class CommonKVManager(BaseKVManager):
`Connection refused`, and the leader's `prefill_port_table` ends `Connection refused`, and the leader's `prefill_port_table` ends
up missing rows. up missing rows.
""" """
if not self.dist_init_addr or get_parallel().nnodes == 1: if not self.dist_init_addr or self.server_args.nnodes == 1:
return local_port return local_port
if not (dist.is_available() and dist.is_initialized()): if not (dist.is_available() and dist.is_initialized()):
@@ -624,14 +627,14 @@ class CommonKVManager(BaseKVManager):
"rank_port": self.rank_port, "rank_port": self.rank_port,
"page_size": self.kv_args.page_size, "page_size": self.kv_args.page_size,
"kv_cache_dtype": get_model().kv_cache_dtype, "kv_cache_dtype": get_model().kv_cache_dtype,
"load_balance_method": get_parallel().load_balance_method, "load_balance_method": self.server_args.load_balance_method,
"enable_dsa_cache_layer_split": getattr( "enable_dsa_cache_layer_split": getattr(
self.server_args, "enable_dsa_cache_layer_split", False self.server_args, "enable_dsa_cache_layer_split", False
), ),
# Self-register the HTTP API port so the decode can derive the PD # Self-register the HTTP API port so the decode can derive the PD
# retract rebootstrap /generate URL from bootstrap info instead of a # retract rebootstrap /generate URL from bootstrap info instead of a
# router-injected pd_rebootstrap_prefill_url. # router-injected pd_rebootstrap_prefill_url.
"prefill_http_port": get_serving().port, "prefill_http_port": self.server_args.port,
} }
max_retries, initial_delay, max_delay = 5, 1.0, 30.0 max_retries, initial_delay, max_delay = 5, 1.0, 30.0
+6 -4
View File
@@ -86,7 +86,7 @@ from sglang.srt.observability.req_time_stats import (
set_schedule_time_batch, set_schedule_time_batch,
set_time_batch, set_time_batch,
) )
from sglang.srt.runtime_context import get_disagg, get_parallel from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import get_num_new_pages from sglang.srt.utils import get_num_new_pages
from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.network import NetworkAddress
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
@@ -2151,7 +2151,7 @@ class SchedulerDisaggregationDecodeMixin:
# Decode-radix path: new requests already matched in # Decode-radix path: new requests already matched in
# `pop_preallocated`. Retracted requests reset `last_node`, # `pop_preallocated`. Retracted requests reset `last_node`,
# so re-match only when that state is missing. # so re-match only when that state is missing.
if get_disagg().disaggregation_decode_enable_radix_cache: if self.server_args.disaggregation_decode_enable_radix_cache:
tree_cache = self.tree_cache if req.last_node is None else None tree_cache = self.tree_cache if req.last_node is None else None
else: else:
tree_cache = self.tree_cache tree_cache = self.tree_cache
@@ -2191,7 +2191,7 @@ class SchedulerDisaggregationDecodeMixin:
if self.enable_decode_hicache: if self.enable_decode_hicache:
self.tree_cache.check_hicache_events() self.tree_cache.check_hicache_events()
if get_disagg().disaggregation_decode_enable_offload_kvcache: if self.server_args.disaggregation_decode_enable_offload_kvcache:
self.decode_offload_manager.check_offload_progress() self.decode_offload_manager.check_offload_progress()
# try to resume retracted requests if there are enough space for another `num_reserved_decode_tokens` decode steps # try to resume retracted requests if there are enough space for another `num_reserved_decode_tokens` decode steps
@@ -2203,7 +2203,9 @@ class SchedulerDisaggregationDecodeMixin:
if not hasattr(self, "polling_count"): if not hasattr(self, "polling_count"):
self.polling_count = 0 self.polling_count = 0
self.polling_interval = get_disagg().disaggregation_decode_polling_interval self.polling_interval = (
self.server_args.disaggregation_decode_polling_interval
)
self.polling_count = (self.polling_count + 1) % self.polling_interval self.polling_count = (self.polling_count + 1) % self.polling_interval
@@ -28,7 +28,6 @@ from sglang.srt.disaggregation.encode_server import (
) )
from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle
from sglang.srt.managers.schedule_batch import Modality from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.runtime_context import get_disagg
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import random_uuid from sglang.srt.utils import random_uuid
from sglang.srt.utils.network import NetworkAddress, get_zmq_socket from sglang.srt.utils.network import NetworkAddress, get_zmq_socket
@@ -118,13 +117,13 @@ class SGLangEncoderServer(SGLangEncoderServicer):
context.set_details(error_msg) context.set_details(error_msg)
return sglang_encoder_pb2.EncodeResponse() return sglang_encoder_pb2.EncodeResponse()
if get_disagg().encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
return sglang_encoder_pb2.EncodeResponse( return sglang_encoder_pb2.EncodeResponse(
embedding_size=nbytes, embedding_size=nbytes,
embedding_len=embedding_len, embedding_len=embedding_len,
embedding_dim=embedding_dim, embedding_dim=embedding_dim,
) )
elif get_disagg().encoder_transfer_backend == "zmq_to_scheduler": elif self.server_args.encoder_transfer_backend == "zmq_to_scheduler":
embedding_ports = list(request.embedding_port) embedding_ports = list(request.embedding_port)
logger.info(f"embedding_port = {embedding_ports}") logger.info(f"embedding_port = {embedding_ports}")
if not embedding_ports: if not embedding_ports:
@@ -142,7 +141,7 @@ class SGLangEncoderServer(SGLangEncoderServicer):
await asyncio.gather(*tasks) await asyncio.gather(*tasks)
self.encoder.embedding_to_send.pop(request.req_id, None) self.encoder.embedding_to_send.pop(request.req_id, None)
return sglang_encoder_pb2.EncodeResponse() return sglang_encoder_pb2.EncodeResponse()
elif get_disagg().encoder_transfer_backend == "zmq_to_tokenizer": elif self.server_args.encoder_transfer_backend == "zmq_to_tokenizer":
embedding_port = ( embedding_port = (
request.embedding_port[0] if request.embedding_port else 0 request.embedding_port[0] if request.embedding_port else 0
) )
@@ -59,9 +59,15 @@ from sglang.srt.model_loader import get_model
from sglang.srt.multimodal.processors.qwen_vl import preprocess_video from sglang.srt.multimodal.processors.qwen_vl import preprocess_video
from sglang.srt.observability.metrics_collector import EncoderMetricsCollector from sglang.srt.observability.metrics_collector import EncoderMetricsCollector
from sglang.srt.observability.req_time_stats import EncoderReqTimeStats from sglang.srt.observability.req_time_stats import EncoderReqTimeStats
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info from sglang.srt.observability.trace import (
from sglang.srt.runtime_context import get_disagg, get_exec, get_mm process_tracing_init,
from sglang.srt.server_args import PortArgs, ServerArgs trace_set_thread_info,
)
from sglang.srt.server_args import (
PortArgs,
ServerArgs,
set_global_server_args_for_scheduler,
)
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prometheus_middleware, add_prometheus_middleware,
configure_logger, configure_logger,
@@ -256,9 +262,7 @@ class MMEncoder:
): ):
logger.info(f"init MMEncoder {rank}/{server_args.tp_size}") logger.info(f"init MMEncoder {rank}/{server_args.tp_size}")
self.server_args = server_args self.server_args = server_args
from sglang.srt.runtime_context import publish set_global_server_args_for_scheduler(server_args)
publish(server_args, role="encoder")
self.rank = rank self.rank = rank
# DP rank for metric labels; overridden by run_dp_worker in DP mode. # DP rank for metric labels; overridden by run_dp_worker in DP mode.
# 0 in the single-instance (non-DP) path. # 0 in the single-instance (non-DP) path.
@@ -345,7 +349,7 @@ class MMEncoder:
[], dtype=self._embedding_dtype [], dtype=self._embedding_dtype
).element_size() ).element_size()
if get_mm().enable_mm_global_cache: if self.server_args.enable_mm_global_cache:
from sglang.srt.mem_cache.storage.mooncake_store.embedding_cache_controller import ( from sglang.srt.mem_cache.storage.mooncake_store.embedding_cache_controller import (
EmbeddingCacheController, EmbeddingCacheController,
) )
@@ -363,15 +367,15 @@ class MMEncoder:
self.mm_global_cache = None self.mm_global_cache = None
# Pre-compute embedding metadata (needed by all ranks for mooncake) # Pre-compute embedding metadata (needed by all ranks for mooncake)
if get_disagg().encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
self._embedding_dims = self._infer_embedding_dims() self._embedding_dims = self._infer_embedding_dims()
if self.rank == 0: if self.rank == 0:
logger.info( logger.info(
f"Using transfer backend: {get_disagg().encoder_transfer_backend}" f"Using transfer backend: {self.server_args.encoder_transfer_backend}"
) )
if get_disagg().encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
self.local_ip = get_local_ip_auto() self.local_ip = get_local_ip_auto()
self.engine = get_mooncake_transfer_engine() self.engine = get_mooncake_transfer_engine()
@@ -384,8 +388,8 @@ class MMEncoder:
hostname=self.local_ip, hostname=self.local_ip,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
ib_device=( ib_device=(
get_disagg().disaggregation_ib_device self.server_args.disaggregation_ib_device
or get_exec().moe.mooncake_ib_device or self.server_args.mooncake_ib_device
), ),
) )
@@ -394,7 +398,7 @@ class MMEncoder:
self.encode_dispatch_lock = asyncio.Lock() self.encode_dispatch_lock = asyncio.Lock()
# Async mooncake state: track background VIT forward completion # Async mooncake state: track background VIT forward completion
if get_disagg().encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
self._forward_ready_events: Dict[str, asyncio.Event] = {} self._forward_ready_events: Dict[str, asyncio.Event] = {}
self._forward_results: Dict[str, dict] = {} self._forward_results: Dict[str, dict] = {}
# when multiple decoder TP ranks call # when multiple decoder TP ranks call
@@ -408,12 +412,12 @@ class MMEncoder:
# Bind unified encode entry point based on backend and cache config # Bind unified encode entry point based on backend and cache config
if self.mm_global_cache is not None: if self.mm_global_cache is not None:
if get_disagg().encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
self._encode_fn = self.encode_with_global_cache_mooncake self._encode_fn = self.encode_with_global_cache_mooncake
else: else:
self._encode_fn = self.encode_with_global_cache self._encode_fn = self.encode_with_global_cache
else: else:
if get_disagg().encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
self._encode_fn = self.encode_with_mooncake self._encode_fn = self.encode_with_mooncake
else: else:
self._encode_fn = self.encode self._encode_fn = self.encode
@@ -1683,7 +1687,7 @@ class MMEncoder:
mm_item.set(k, _convert(v)) mm_item.set(k, _convert(v))
cache_hit = False cache_hit = False
use_mm_cache = get_mm().enable_prefix_mm_cache and log_metrics use_mm_cache = self.server_args.enable_prefix_mm_cache and log_metrics
if use_mm_cache: if use_mm_cache:
mm_item.set_pad_value() mm_item.set_pad_value()
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash]) mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
@@ -1779,7 +1783,7 @@ class MMEncoder:
embedding_port=None, embedding_port=None,
url=None, url=None,
): ):
if get_disagg().encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
# Wait for async VIT forward completion if needed # Wait for async VIT forward completion if needed
req_id = mm_data.req_id req_id = mm_data.req_id
if req_id in self._forward_ready_events: if req_id in self._forward_ready_events:
@@ -1850,7 +1854,7 @@ class MMEncoder:
logger.info(f"{endpoint = }") logger.info(f"{endpoint = }")
# Serialize data # Serialize data
if get_disagg().encoder_transfer_backend == "mooncake": if self.server_args.encoder_transfer_backend == "mooncake":
# Mooncake already pushed the embedding via RDMA; # Mooncake already pushed the embedding via RDMA;
new_mm_data = mm_data.copy_without_embedding() new_mm_data = mm_data.copy_without_embedding()
serialized_data = pickle.dumps(new_mm_data) serialized_data = pickle.dumps(new_mm_data)
@@ -1882,11 +1886,11 @@ class MMEncoder:
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket) await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
if ( if (
encoder_metrics_collector is not None encoder_metrics_collector is not None
and get_disagg().encoder_transfer_backend != "mooncake" and self.server_args.encoder_transfer_backend != "mooncake"
): ):
encoder_metrics_collector.observe_transfer( encoder_metrics_collector.observe_transfer(
time.perf_counter() - _zmq_xfer_start, time.perf_counter() - _zmq_xfer_start,
backend=get_disagg().encoder_transfer_backend, backend=self.server_args.encoder_transfer_backend,
) )
async def encode( async def encode(
@@ -55,7 +55,6 @@ from sglang.srt.observability.trace import (
TraceReqContext, TraceReqContext,
trace_set_thread_info, trace_set_thread_info,
) )
from sglang.srt.runtime_context import get_parallel, get_schedule
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.network import NetworkAddress
@@ -315,7 +314,9 @@ class MooncakeKVManager(CommonKVManager):
self.kv_buffer_tensors = None self.kv_buffer_tensors = None
def _handle_staging_req(self, msg): def _handle_staging_req(self, msg):
from sglang.srt.disaggregation.common.staging_handler import handle_staging_req from sglang.srt.disaggregation.common.staging_handler import (
handle_staging_req,
)
room = int(msg[1].decode("ascii")) room = int(msg[1].decode("ascii"))
session_id = msg[4].decode("ascii") session_id = msg[4].decode("ascii")
@@ -349,7 +350,9 @@ class MooncakeKVManager(CommonKVManager):
def _is_watermark_ready( def _is_watermark_ready(
self, session_id: str, alloc_round: int, alloc_end: int self, session_id: str, alloc_round: int, alloc_end: int
) -> bool: ) -> bool:
from sglang.srt.disaggregation.common.staging_handler import is_watermark_ready from sglang.srt.disaggregation.common.staging_handler import (
is_watermark_ready,
)
return is_watermark_ready(self._staging_ctx, session_id, alloc_round, alloc_end) return is_watermark_ready(self._staging_ctx, session_id, alloc_round, alloc_end)
@@ -466,7 +469,7 @@ class MooncakeKVManager(CommonKVManager):
room, room,
self.transfer_infos, self.transfer_infos,
self.kv_buffer_tensors, self.kv_buffer_tensors,
get_schedule().chunked_prefill_size, self.server_args.chunked_prefill_size,
self._staging_ctx.prefetch_requested, self._staging_ctx.prefetch_requested,
self._staging_ctx.prefetch_sockets, self._staging_ctx.prefetch_sockets,
) )
@@ -948,7 +951,7 @@ class MooncakeKVManager(CommonKVManager):
if ( if (
self.attn_cp_size > 1 self.attn_cp_size > 1
and self.attn_cp_rank != 0 and self.attn_cp_rank != 0
and not get_parallel().enable_dsa_cache_layer_split and not self.server_args.enable_dsa_cache_layer_split
): ):
skip_state = True skip_state = True
+10 -6
View File
@@ -13,8 +13,6 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
from sglang.srt.runtime_context import get_schedule
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo
@@ -537,7 +535,9 @@ class NixlKVManager(CommonKVManager):
def _is_watermark_ready( def _is_watermark_ready(
self, agent_name: str, alloc_round: int, alloc_end: int self, agent_name: str, alloc_round: int, alloc_end: int
) -> bool: ) -> bool:
from sglang.srt.disaggregation.common.staging_handler import is_watermark_ready from sglang.srt.disaggregation.common.staging_handler import (
is_watermark_ready,
)
return is_watermark_ready(self._staging_ctx, agent_name, alloc_round, alloc_end) return is_watermark_ready(self._staging_ctx, agent_name, alloc_round, alloc_end)
@@ -558,7 +558,9 @@ class NixlKVManager(CommonKVManager):
threading.Thread(target=decode_staging_thread, daemon=True).start() threading.Thread(target=decode_staging_thread, daemon=True).start()
def _handle_staging_req(self, msg): def _handle_staging_req(self, msg):
from sglang.srt.disaggregation.common.staging_handler import handle_staging_req from sglang.srt.disaggregation.common.staging_handler import (
handle_staging_req,
)
room = int(msg[1].decode("ascii")) room = int(msg[1].decode("ascii"))
session_id = msg[4].decode("ascii") session_id = msg[4].decode("ascii")
@@ -623,7 +625,7 @@ class NixlKVManager(CommonKVManager):
room, room,
self.transfer_infos, self.transfer_infos,
self.kv_buffer_tensors, self.kv_buffer_tensors,
get_schedule().chunked_prefill_size, self.server_args.chunked_prefill_size,
self._staging_ctx.prefetch_requested, self._staging_ctx.prefetch_requested,
self._staging_ctx.prefetch_sockets, self._staging_ctx.prefetch_sockets,
) )
@@ -1737,7 +1739,9 @@ class NixlKVManager(CommonKVManager):
req, page_start, num_pages, session_id=req.agent_name req, page_start, num_pages, session_id=req.agent_name
) )
if not ready: if not ready:
from sglang.srt.disaggregation.common.staging_buffer import StagingAllocator from sglang.srt.disaggregation.common.staging_buffer import (
StagingAllocator,
)
if c_offset == StagingAllocator.ALLOC_OVERSIZED: if c_offset == StagingAllocator.ALLOC_OVERSIZED:
raise RuntimeError( raise RuntimeError(
+1 -2
View File
@@ -64,7 +64,6 @@ from sglang.srt.mem_cache.common 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.observability.req_time_stats import set_schedule_time_batch from sglang.srt.observability.req_time_stats import set_schedule_time_batch
from sglang.srt.runtime_context import get_disagg
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -1182,7 +1181,7 @@ class SchedulerDisaggregationPrefillMixin:
def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None: def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None:
"""Release KV cache and requeue an optimistic prefill request.""" """Release KV cache and requeue an optimistic prefill request."""
max_attempts = get_disagg().optimistic_prefill_attempts max_attempts = self.server_args.optimistic_prefill_attempts
maybe_cache_unfinished_req(req, self.tree_cache) maybe_cache_unfinished_req(req, self.tree_cache)
release_kv_cache(req, self.tree_cache) release_kv_cache(req, self.tree_cache)
req.reset_for_retract() req.reset_for_retract()
@@ -14,7 +14,7 @@ from sglang.srt.compilation.compile_phase import (
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_exec from sglang.srt.runtime_context import get_server_args
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -25,7 +25,7 @@ class PyMscclppCommunicator:
def _is_symm_mem_enabled(self) -> bool: def _is_symm_mem_enabled(self) -> bool:
try: try:
return get_exec().comm.enable_symm_mem return get_server_args().enable_symm_mem
except ValueError: except ValueError:
return False return False
@@ -15,7 +15,7 @@ from torch.cuda.memory import (
from sglang.srt.distributed.parallel_state import GroupCoordinator from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils.common import torch_release from sglang.srt.utils.common import torch_release
after_2_8_0 = torch_release >= (2, 8) after_2_8_0 = torch_release >= (2, 8)
@@ -159,7 +159,7 @@ _register_func = None
def is_symmetric_memory_enabled(): def is_symmetric_memory_enabled():
try: try:
return get_exec().comm.enable_symm_mem return get_server_args().enable_symm_mem
except ValueError: except ValueError:
return False return False
@@ -12,7 +12,6 @@ from sglang.srt.distributed.device_communicators.all_reduce_utils import (
TORCH_SYMM_MEM_ALL_REDUCE_MAX_SIZES, TORCH_SYMM_MEM_ALL_REDUCE_MAX_SIZES,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_exec
from sglang.srt.utils import is_cuda, is_hip from sglang.srt.utils import is_cuda, is_hip
try: try:
@@ -99,9 +98,10 @@ class TorchSymmMemCommunicator:
# ([16384, 6144] bf16 = 192 MiB), including room for tail regions. # ([16384, 6144] bf16 = 192 MiB), including room for tail regions.
if envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR.get(): if envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR.get():
self.max_size = max(self.max_size, 256 * 1024 * 1024) self.max_size = max(self.max_size, 256 * 1024 * 1024)
from sglang.srt.runtime_context import get_server_args
if ( if (
get_exec().comm.enable_scattered_sconv get_server_args().enable_scattered_sconv
or envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV.get() or envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV.get()
): ):
# Fused extend kernels are out-of-place, so OUT must hold the # Fused extend kernels are out-of-place, so OUT must hold the
@@ -16,8 +16,6 @@ import torch.distributed._symmetric_memory as symm_mem
import triton import triton
import triton.language as tl import triton.language as tl
from sglang.srt.runtime_context import get_parallel
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Each thread moves _NUMEL_PER_THREAD bf16 via one 128-bit multimem op; the # Each thread moves _NUMEL_PER_THREAD bf16 via one 128-bit multimem op; the
@@ -468,6 +466,7 @@ class MultimemAllGatherer:
# Lazy import avoids a module-load dependency on the distributed facade. # Lazy import avoids a module-load dependency on the distributed facade.
from sglang.srt.distributed import get_tp_group from sglang.srt.distributed import 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_server_args
tp_group = get_tp_group() tp_group = get_tp_group()
# Only probe node topology when the deployment can actually span # Only probe node topology when the deployment can actually span
@@ -478,7 +477,7 @@ class MultimemAllGatherer:
# EP/mooncake setups, and keep multimem enabled. # EP/mooncake setups, and keep multimem enabled.
if ( if (
tp_group.world_size > 1 tp_group.world_size > 1
and get_parallel().nnodes > 1 and get_server_args().nnodes > 1
and not all(in_the_same_node_as(tp_group.cpu_group, source_rank=0)) and not all(in_the_same_node_as(tp_group.cpu_group, source_rank=0))
): ):
logger.warning( logger.warning(
+2 -3
View File
@@ -11,7 +11,6 @@ from sglang.srt.managers.schedule_policy import AddReqResult, PrefillAdder
from sglang.srt.mem_cache.common import release_kv_cache from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.req_time_stats import set_time_batch from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.runtime_context import get_exec, get_schedule
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -23,7 +22,7 @@ class SchedulerDllmMixin:
def init_diffusion_llm(self: Scheduler): def init_diffusion_llm(self: Scheduler):
self.dllm_config = ( self.dllm_config = (
DllmConfig.from_server_args(self.server_args) DllmConfig.from_server_args(self.server_args)
if get_exec().dllm.dllm_algorithm is not None if self.server_args.dllm_algorithm is not None
else None else None
) )
self.dllm_manager = DllmManager(dllm_config=self.dllm_config) self.dllm_manager = DllmManager(dllm_config=self.dllm_config)
@@ -201,7 +200,7 @@ class SchedulerDllmMixin:
self.chunked_prefill_size, self.chunked_prefill_size,
running_bs if self.is_mixed_chunk else 0, running_bs if self.is_mixed_chunk else 0,
self.priority_scheduling_preemption_threshold, self.priority_scheduling_preemption_threshold,
prefill_max_requests=get_schedule().prefill_max_requests, prefill_max_requests=self.server_args.prefill_max_requests,
dllm_config=self.dllm_config, dllm_config=self.dllm_config,
) )
+1 -1
View File
@@ -442,7 +442,7 @@ def get_healthy_expert_location_src_rank(
*, invoked_in_elastic_ep_rejoin_path: bool *, invoked_in_elastic_ep_rejoin_path: bool
) -> int: ) -> int:
world_group = get_world_group() world_group = get_world_group()
# NOTE: do not key off `get_exec().moe.elastic_ep_rejoin` here. # NOTE: do not key off `self.server_args.elastic_ep_rejoin` here.
# A rank that was started as a rejoin rank may later act as a healthy # A rank that was started as a rejoin rank may later act as a healthy
# rank in a subsequent recovery cycle. # rank in a subsequent recovery cycle.
local_rejoin_flag = bool(invoked_in_elastic_ep_rejoin_path) local_rejoin_flag = bool(invoked_in_elastic_ep_rejoin_path)
@@ -7,11 +7,13 @@ from typing import Any, Callable
import torch import torch
import zmq import zmq
from sglang.srt.distributed.parallel_state import get_world_group, get_world_size from sglang.srt.distributed.parallel_state import (
get_world_group,
get_world_size,
)
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send
from sglang.srt.runtime_context import get_exec
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import get_local_ip_auto from sglang.srt.utils.network import get_local_ip_auto
@@ -109,7 +111,7 @@ class ExpertBackupClient:
global_expert_location_metadata = get_global_expert_location_metadata() global_expert_location_metadata = get_global_expert_location_metadata()
num_experts = ( num_experts = (
self.model_config.hf_config.n_routed_experts self.model_config.hf_config.n_routed_experts
+ get_exec().moe.ep_num_redundant_experts + self.server_args.ep_num_redundant_experts
) )
num_local_experts = num_experts // self.moe_ep_size num_local_experts = num_experts // self.moe_ep_size
for i in range(self.engine_num): for i in range(self.engine_num):
@@ -20,6 +20,7 @@ from sglang.srt.model_loader.utils import set_default_torch_dtype
from sglang.srt.server_args import ( from sglang.srt.server_args import (
PortArgs, PortArgs,
ServerArgs, ServerArgs,
set_global_server_args_for_scheduler,
) )
from sglang.srt.utils.network import get_local_ip_auto from sglang.srt.utils.network import get_local_ip_auto
@@ -158,9 +159,7 @@ def run_expert_backup_manager_process(
server_args: ServerArgs, server_args: ServerArgs,
port_args: PortArgs, port_args: PortArgs,
): ):
from sglang.srt.runtime_context import publish set_global_server_args_for_scheduler(server_args)
publish(server_args, role="expert_backup")
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
init_mooncake_transfer_engine, init_mooncake_transfer_engine,
) )
+4 -18
View File
@@ -93,7 +93,6 @@ 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_detection import resolve_auto_parsers
from sglang.srt.parser.template_manager import TemplateManager from sglang.srt.parser.template_manager import TemplateManager
from sglang.srt.plugins import load_plugins 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.server_args import PortArgs, ServerArgs
from sglang.srt.utils import ( from sglang.srt.utils import (
MultiprocessingSerializer, MultiprocessingSerializer,
@@ -254,7 +253,7 @@ class Engine(EngineScoreMixin, EngineBase):
# Initialize ZMQ sockets # Initialize ZMQ sockets
context = zmq.Context(2) context = zmq.Context(2)
if server_args.node_rank == 0: if self.server_args.node_rank == 0:
self.send_to_rpc = get_zmq_socket( self.send_to_rpc = get_zmq_socket(
context, zmq.DEALER, self.port_args.rpc_ipc_name, True context, zmq.DEALER, self.port_args.rpc_ipc_name, True
) )
@@ -302,7 +301,7 @@ class Engine(EngineScoreMixin, EngineBase):
routed_dp_rank = data_parallel_rank routed_dp_rank = data_parallel_rank
if routed_dp_rank is not None: if routed_dp_rank is not None:
dp_size = get_parallel().dp_size dp_size = self.server_args.dp_size
if dp_size <= 1 and routed_dp_rank == 0: if dp_size <= 1 and routed_dp_rank == 0:
logger.debug( logger.debug(
f"routed_dp_rank={routed_dp_rank} is ignored because dp_size={dp_size}" f"routed_dp_rank={routed_dp_rank} is ignored because dp_size={dp_size}"
@@ -878,14 +877,7 @@ class Engine(EngineScoreMixin, EngineBase):
server_args, port_args server_args, port_args
) )
else: else:
# Launch multi-tokenizer router. Unlike TokenizerManager, the router # Launch multi-tokenizer router
# does not publish; but it runs in this parent process and reads
# resolved config through the namespace accessors (e.g. get_parallel()
# for routed_dp_rank), so publish here. The child TokenizerWorkers
# publish independently in their own processes.
from sglang.srt.runtime_context import publish
publish(server_args, role="tokenizer")
tokenizer_manager = MultiTokenizerRouter(server_args, port_args) tokenizer_manager = MultiTokenizerRouter(server_args, port_args)
template_manager = None template_manager = None
@@ -1005,18 +997,12 @@ class Engine(EngineScoreMixin, EngineBase):
) )
def get_server_info(self): def get_server_info(self):
from sglang.srt.runtime_context import get_context
internal_states = self.loop.run_until_complete( internal_states = self.loop.run_until_complete(
self.tokenizer_manager.get_internal_state() self.tokenizer_manager.get_internal_state()
) )
return msgspec_to_builtins( return msgspec_to_builtins(
{ {
# Overlay post-publish overrides so the report reflects current **dataclasses.asdict(self.tokenizer_manager.server_args),
# config (weight version, model path, runtime tunables).
**get_context().resolved_server_args_dict(
base=dataclasses.asdict(self.tokenizer_manager.server_args)
),
**self._scheduler_init_result.scheduler_infos[0], **self._scheduler_init_result.scheduler_infos[0],
"internal_states": internal_states, "internal_states": internal_states,
"version": __version__, "version": __version__,
+9 -10
View File
@@ -15,7 +15,6 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional
from pydantic import ValidationError from pydantic import ValidationError
from sglang.srt.runtime_context import get_context, get_lora, get_serving
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -230,7 +229,9 @@ class RuntimeHandle:
return self._openai_serving_classes return self._openai_serving_classes
from sglang.srt.entrypoints.openai.serving_chat import OpenAIServingChat from sglang.srt.entrypoints.openai.serving_chat import OpenAIServingChat
from sglang.srt.entrypoints.openai.serving_classify import OpenAIServingClassify from sglang.srt.entrypoints.openai.serving_classify import (
OpenAIServingClassify,
)
from sglang.srt.entrypoints.openai.serving_completions import ( from sglang.srt.entrypoints.openai.serving_completions import (
OpenAIServingCompletion, OpenAIServingCompletion,
) )
@@ -375,20 +376,16 @@ class RuntimeHandle:
model_config = self.tokenizer_manager.model_config model_config = self.tokenizer_manager.model_config
result = { result = {
"model_path": self.tokenizer_manager.model_path, "model_path": self.tokenizer_manager.model_path,
"tokenizer_path": get_serving().tokenizer_path, "tokenizer_path": self.server_args.tokenizer_path,
"is_generation": self.tokenizer_manager.is_generation, "is_generation": self.tokenizer_manager.is_generation,
"weight_version": get_serving().weight_version, "weight_version": self.server_args.weight_version,
"model_type": getattr(model_config.hf_config, "model_type", None), "model_type": getattr(model_config.hf_config, "model_type", None),
"architectures": getattr(model_config.hf_config, "architectures", None), "architectures": getattr(model_config.hf_config, "architectures", None),
} }
return json.dumps(result, default=str) return json.dumps(result, default=str)
def get_server_info(self) -> str: def get_server_info(self) -> str:
# Overlay post-publish overrides (weight version, model path, runtime result: Dict[str, Any] = dataclasses.asdict(self.server_args)
# tunables) so the report reflects current config, not the startup record.
result: Dict[str, Any] = get_context().resolved_server_args_dict(
base=dataclasses.asdict(self.server_args)
)
result.update(self.scheduler_info) result.update(self.scheduler_info)
return json.dumps(msgspec_to_builtins(result), default=str) return json.dumps(msgspec_to_builtins(result), default=str)
@@ -427,7 +424,9 @@ class RuntimeHandle:
"max_model_len": self.tokenizer_manager.model_config.context_len, "max_model_len": self.tokenizer_manager.model_config.context_len,
} }
] ]
if get_lora().enable_lora and hasattr(self.tokenizer_manager, "lora_registry"): if self.server_args.enable_lora and hasattr(
self.tokenizer_manager, "lora_registry"
):
lora_registry = self.tokenizer_manager.lora_registry lora_registry = self.tokenizer_manager.lora_registry
for _, lora_ref in lora_registry.get_all_adapters().items(): for _, lora_ref in lora_registry.get_all_adapters().items():
models.append( models.append(
+4 -13
View File
@@ -703,19 +703,18 @@ async def get_model_info():
@app.get("/model_info") @app.get("/model_info")
async def model_info(): async def model_info():
"""Get the model information.""" """Get the model information."""
from sglang.srt.runtime_context import get_serving
model_config = _global_state.tokenizer_manager.model_config model_config = _global_state.tokenizer_manager.model_config
result = { result = {
"model_path": _global_state.tokenizer_manager.model_path, "model_path": _global_state.tokenizer_manager.model_path,
"tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path, "tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path,
"is_generation": _global_state.tokenizer_manager.is_generation, "is_generation": _global_state.tokenizer_manager.is_generation,
"preferred_sampling_params": _global_state.tokenizer_manager.server_args.preferred_sampling_params, "preferred_sampling_params": _global_state.tokenizer_manager.server_args.preferred_sampling_params,
"weight_version": get_serving().weight_version, "weight_version": _global_state.tokenizer_manager.server_args.weight_version,
"has_image_understanding": model_config.is_image_understandable_model, "has_image_understanding": model_config.is_image_understandable_model,
"has_audio_understanding": model_config.is_audio_understandable_model, "has_audio_understanding": model_config.is_audio_understandable_model,
"model_type": getattr(model_config.hf_config, "model_type", None), "model_type": getattr(model_config.hf_config, "model_type", None),
"architectures": getattr(model_config.hf_config, "architectures", None), "architectures": getattr(model_config.hf_config, "architectures", None),
"weight_version": _global_state.tokenizer_manager.server_args.weight_version,
# "hf_config": model_config.hf_config.to_dict(), # "hf_config": model_config.hf_config.to_dict(),
} }
return result return result
@@ -749,18 +748,12 @@ async def server_info():
await _global_state.tokenizer_manager.get_internal_state() await _global_state.tokenizer_manager.get_internal_state()
) )
from sglang.srt.runtime_context import get_context
server_args = _global_state.tokenizer_manager.server_args server_args = _global_state.tokenizer_manager.server_args
# server_args.model_config is not serializable but should be excluded by asdict. # server_args.model_config is not serializable but should be excluded by asdict.
# Overlay post-publish overrides so runtime updates (weight version, model
# path/load format) are reported, not the startup record.
return msgspec_to_builtins( return msgspec_to_builtins(
{ {
**get_context().resolved_server_args_dict( **dataclasses.asdict(server_args),
base=dataclasses.asdict(server_args)
),
**_global_state.scheduler_info, **_global_state.scheduler_info,
"internal_states": internal_states, "internal_states": internal_states,
"version": __version__, "version": __version__,
@@ -1385,9 +1378,7 @@ async def update_weight_version(
# since weight_version update is a simple operation that doesn't affect model weights # since weight_version update is a simple operation that doesn't affect model weights
try: try:
# Update the weight version in server args (the single source of truth) # Update the weight version in server args (the single source of truth)
from sglang.srt.runtime_context import get_context _global_state.tokenizer_manager.server_args.override(
get_context().override(
"http.update_weight_version", weight_version=obj.new_version "http.update_weight_version", weight_version=obj.new_version
) )
@@ -55,8 +55,6 @@ class HttpServerEngineAdapter(EngineBase):
def __init__(self, **kwargs): def __init__(self, **kwargs):
self.server_args = ServerArgs(**kwargs) self.server_args = ServerArgs(**kwargs)
# Read host/port from the adapter's own args: no config is published yet
# in this process (publish happens in the child from launch_server_process).
print( print(
f"Launch HttpServerEngineAdapter at: {self.server_args.host}:{self.server_args.port}" f"Launch HttpServerEngineAdapter at: {self.server_args.host}:{self.server_args.port}"
) )
@@ -72,7 +72,6 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import (
TranscriptionAdapter, TranscriptionAdapter,
) )
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_serving
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import random_uuid from sglang.srt.utils import random_uuid
@@ -339,12 +338,12 @@ class RealtimeConnection:
if ( if (
transcription is not None transcription is not None
and transcription.model and transcription.model
and transcription.model != get_serving().served_model_name and transcription.model != self.server_args.served_model_name
): ):
await self._send_error( await self._send_error(
"not_supported", "not_supported",
f"Model {transcription.model!r} is not served by this endpoint " f"Model {transcription.model!r} is not served by this endpoint "
f"(serving {get_serving().served_model_name!r}); set " f"(serving {self.server_args.served_model_name!r}); set "
f"transcription.model to null or to the server's model name.", f"transcription.model to null or to the server's model name.",
param="session.audio.input.transcription.model", param="session.audio.input.transcription.model",
) )
+3 -3
View File
@@ -16,7 +16,7 @@ from sglang.srt.eplb.expert_location import (
get_global_expert_location_metadata, get_global_expert_location_metadata,
) )
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
from sglang.srt.runtime_context import get_model from sglang.srt.runtime_context import get_server_args
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
@@ -274,8 +274,8 @@ def update_expert_location_with_recovery(
else: else:
# Load the missing weights from disk # Load the missing weights from disk
update_weights_from_disk_callable( update_weights_from_disk_callable(
get_model().model_path, get_server_args().model_path,
get_model().load_format, get_server_args().load_format,
weight_name_filter=weight_name_filter, weight_name_filter=weight_name_filter,
) )
@@ -18,7 +18,7 @@ from typing import Literal, Optional
import torch import torch
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_server_args
@dataclass @dataclass
@@ -34,7 +34,7 @@ class ExpertLocationDispatchInfo:
@classmethod @classmethod
def init_new(cls, layer_id: int): def init_new(cls, layer_id: int):
ep_dispatch_algorithm = get_exec().moe.ep_dispatch_algorithm ep_dispatch_algorithm = get_server_args().ep_dispatch_algorithm
expert_location_metadata = get_global_expert_location_metadata() expert_location_metadata = get_global_expert_location_metadata()
assert expert_location_metadata is not None assert expert_location_metadata is not None
@@ -26,7 +26,7 @@ from sglang.srt.eplb.expert_location import (
ExpertLocationMetadata, ExpertLocationMetadata,
get_global_expert_location_metadata, get_global_expert_location_metadata,
) )
from sglang.srt.runtime_context import get_device from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import get_bool_env_var from sglang.srt.utils import get_bool_env_var
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -107,7 +107,7 @@ def _update_expert_weights_with_canary(
canary_tensor = ( canary_tensor = (
_get_canary_value(old_expert_location_metadata, layer_id) _get_canary_value(old_expert_location_metadata, layer_id)
.clone() .clone()
.to(device=get_device().device, non_blocking=True) .to(device=get_server_args().device, non_blocking=True)
) )
routed_experts_weights_of_layer[layer_id].append(canary_tensor) routed_experts_weights_of_layer[layer_id].append(canary_tensor)
@@ -16,8 +16,9 @@ from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import (
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.model_executor.model_runner_components.layer_setup import ModelLayerInfo from sglang.srt.model_executor.model_runner_components.layer_setup import (
from sglang.srt.runtime_context import get_exec, get_memory, get_schedule ModelLayerInfo,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -143,7 +144,7 @@ class MlxModelRunnerStub(ModelRunner):
(``MlxAuxiliaryStateComponent``) raises ``NotImplementedError`` for (``MlxAuxiliaryStateComponent``) raises ``NotImplementedError`` for
the mode. the mode.
""" """
if get_memory().disable_radix_cache: if self.server_args.disable_radix_cache:
return 1 return 1
return MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO return MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO
@@ -164,7 +165,7 @@ class MlxModelRunnerStub(ModelRunner):
Requires ``self.max_total_num_tokens`` to already be set. Requires ``self.max_total_num_tokens`` to already be set.
""" """
capacity_cap = self.max_total_num_tokens // 2 capacity_cap = self.max_total_num_tokens // 2
requested = get_schedule().max_running_requests requested = self.server_args.max_running_requests
if requested is None: if requested is None:
requested_per_worker = None requested_per_worker = None
resolved = min(capacity_cap, 4096) resolved = min(capacity_cap, 4096)
@@ -172,7 +173,7 @@ class MlxModelRunnerStub(ModelRunner):
requested_per_worker = requested // self.dp_size requested_per_worker = requested // self.dp_size
resolved = min(requested_per_worker, capacity_cap) resolved = min(requested_per_worker, capacity_cap)
aux_state_size = get_schedule().max_mamba_cache_size aux_state_size = self.server_args.max_mamba_cache_size
if ( if (
mambaish_config(self.model_config) is not None mambaish_config(self.model_config) is not None
and aux_state_size is not None and aux_state_size is not None
@@ -208,7 +209,7 @@ class MlxModelRunnerStub(ModelRunner):
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
self.memory_saver_adapter = TorchMemorySaverAdapter.create( self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=get_exec().features.enable_memory_saver enable=self.server_args.enable_memory_saver
) )
# Load model (sets metadata only) # Load model (sets metadata only)
@@ -240,7 +241,7 @@ class MlxModelRunnerStub(ModelRunner):
# Create minimal pools # Create minimal pools
if mambaish_config(self.model_config) is not None: if mambaish_config(self.model_config) is not None:
auxiliary_state_size = get_schedule().max_mamba_cache_size auxiliary_state_size = self.server_args.max_mamba_cache_size
if auxiliary_state_size is None: if auxiliary_state_size is None:
auxiliary_state_size = ( auxiliary_state_size = (
self.max_running_requests * self._aux_state_slots_per_request() self.max_running_requests * self._aux_state_slots_per_request()
@@ -254,7 +255,7 @@ class MlxModelRunnerStub(ModelRunner):
# With the radix cache disabled no tree component exists to # With the radix cache disabled no tree component exists to
# release auxiliary slots, so the pool owns their release # release auxiliary slots, so the pool owns their release
# (see MlxAuxiliaryStateReqToTokenPool docstring). # (see MlxAuxiliaryStateReqToTokenPool docstring).
owns_auxiliary_state_release=get_memory().disable_radix_cache, owns_auxiliary_state_release=self.server_args.disable_radix_cache,
) )
else: else:
self.req_to_token_pool = ReqToTokenPool( self.req_to_token_pool = ReqToTokenPool(
@@ -31,7 +31,6 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.runtime_context import get_memory, get_model, get_schedule
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -48,23 +47,25 @@ class MlxTpModelWorker(TpModelWorker):
def _init_model_runner(self): def _init_model_runner(self):
"""Create MLX runner first (auto-sizes pool), then stub with matching size.""" """Create MLX runner first (auto-sizes pool), then stub with matching size."""
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
from sglang.srt.hardware_backend.mlx.model_runner_stub import MlxModelRunnerStub from sglang.srt.hardware_backend.mlx.model_runner_stub import (
MlxModelRunnerStub,
)
logger.info("Initializing MlxModelRunner for end-to-end MLX inference") logger.info("Initializing MlxModelRunner for end-to-end MLX inference")
init_kwargs = dict( init_kwargs = dict(
model_path=get_model().model_path, model_path=self.server_args.model_path,
trust_remote_code=get_model().trust_remote_code, trust_remote_code=self.server_args.trust_remote_code,
disable_radix_cache=get_memory().disable_radix_cache, disable_radix_cache=self.server_args.disable_radix_cache,
mem_fraction_static=get_schedule().mem_fraction_static, mem_fraction_static=self.server_args.mem_fraction_static,
quantization=get_model().quantization, quantization=self.server_args.quantization,
) )
if get_schedule().max_total_tokens is not None: if self.server_args.max_total_tokens is not None:
init_kwargs["pool_size"] = get_schedule().max_total_tokens init_kwargs["pool_size"] = self.server_args.max_total_tokens
self._mlx_runner = MlxModelRunner(**init_kwargs) self._mlx_runner = MlxModelRunner(**init_kwargs)
self._model_runner = MlxModelRunnerStub( self._model_runner = MlxModelRunnerStub(
model_config=self.model_config, model_config=self.model_config,
mem_fraction_static=get_schedule().mem_fraction_static, mem_fraction_static=self.server_args.mem_fraction_static,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
ps=self.ps, ps=self.ps,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
@@ -19,9 +19,11 @@ from sglang.srt.layers.attention.flashattention_backend import (
merge_state_v2_wrapper, merge_state_v2_wrapper,
) )
from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.layers.utils.cp_utils import cp_allgather_and_save_kv_cache from sglang.srt.layers.utils.cp_utils import (
cp_allgather_and_save_kv_cache,
)
from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.runtime_context import get_schedule from sglang.srt.runtime_context import get_server_args
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -513,7 +515,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
and not forward_batch.forward_mode.is_draft_extend_v2() and not forward_batch.forward_mode.is_draft_extend_v2()
): ):
if forward_batch.attn_attend_prefix_cache: if forward_batch.attn_attend_prefix_cache:
assert not get_schedule().disable_chunked_prefix_cache assert not get_server_args().disable_chunked_prefix_cache
assert forward_batch.prefix_chunk_idx is not None assert forward_batch.prefix_chunk_idx is not None
assert forward_batch.prefix_chunk_cu_seq_lens is not None assert forward_batch.prefix_chunk_cu_seq_lens is not None
assert forward_batch.prefix_chunk_max_seq_lens is not None assert forward_batch.prefix_chunk_max_seq_lens is not None
@@ -12,7 +12,7 @@ from sglang.srt.layers.attention.dsv4.compressor import CompressorBackendMixin
from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin
from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, ForwardMode from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, ForwardMode
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.runtime_context import get_parallel, get_spec 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
@@ -1362,8 +1362,9 @@ class DeepseekV4AscendAttnBackend(
or forward_batch.forward_mode.is_draft_extend_v2() or forward_batch.forward_mode.is_draft_extend_v2()
): ):
B = forward_batch.batch_size B = forward_batch.batch_size
from sglang.srt.runtime_context import get_server_args
n_draft = get_spec().speculative_num_draft_tokens or 1 n_draft = get_server_args().speculative_num_draft_tokens or 1
actual_q = torch.arange( actual_q = torch.arange(
n_draft, B * n_draft + 1, n_draft, dtype=torch.int32, device=device n_draft, B * n_draft + 1, n_draft, dtype=torch.int32, device=device
) )
@@ -1408,8 +1409,9 @@ class DeepseekV4AscendAttnBackend(
forward_batch.forward_mode.is_target_verify() forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2() or forward_batch.forward_mode.is_draft_extend_v2()
): ):
from sglang.srt.runtime_context import get_server_args
max_seqlen_q = get_spec().speculative_num_draft_tokens or 1 max_seqlen_q = get_server_args().speculative_num_draft_tokens or 1
else: else:
max_seqlen_q = 1 max_seqlen_q = 1
return self._kernel_metadata_from_parts( return self._kernel_metadata_from_parts(
@@ -27,7 +27,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
) )
from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
from sglang.srt.runtime_context import get_mm from sglang.srt.runtime_context import get_server_args
class ViTNpuGraphRunner(ViTCudaGraphRunner): class ViTNpuGraphRunner(ViTCudaGraphRunner):
@@ -70,7 +70,7 @@ class ViTNpuGraphRunner(ViTCudaGraphRunner):
graph = torch_npu.npu.NPUGraph() graph = torch_npu.npu.NPUGraph()
vit = self.vit vit = self.vit
override_backend = get_mm().mm_attention_backend override_backend = get_server_args().mm_attention_backend
with torch_npu.npu.graph(graph, pool=ViTNpuGraphRunner._graph_memory_pool): with torch_npu.npu.graph(graph, pool=ViTNpuGraphRunner._graph_memory_pool):
y = None y = None
deepstack_outs: List[torch.Tensor] = [] deepstack_outs: List[torch.Tensor] = []
@@ -17,7 +17,7 @@ from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.utils import npu_format_cast from sglang.srt.hardware_backend.npu.utils import npu_format_cast
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer
from sglang.srt.layers.moe.utils import DeepEPMode from sglang.srt.layers.moe.utils import DeepEPMode
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_server_args
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
@@ -57,7 +57,7 @@ def forward_fuseep(
envs.SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get() envs.SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get()
), ),
num_experts=layer.num_experts, num_experts=layer.num_experts,
fuse_mode=get_exec().moe.fuseep_mode, fuse_mode=get_server_args().fuseep_mode,
) )
return hidden_states return hidden_states
@@ -126,7 +126,7 @@ def process_fuseep_weights(layer: torch.nn.Module, weight_prefix: str) -> None:
Invoked by ``maybe_apply_fuseep_weights`` for both ``"w13"`` and ``"w2"``. Invoked by ``maybe_apply_fuseep_weights`` for both ``"w13"`` and ``"w2"``.
""" """
if get_exec().moe.fuseep_mode == 1: if get_server_args().fuseep_mode == 1:
# -- The fused MoE optimization mode "1": dispatch_gmm_combine_decode -- # -- The fused MoE optimization mode "1": dispatch_gmm_combine_decode --
if weight_prefix == "w13": if weight_prefix == "w13":
cpu_w13 = layer.w13_weight.data.transpose(1, 2).cpu() cpu_w13 = layer.w13_weight.data.transpose(1, 2).cpu()
@@ -143,7 +143,7 @@ def process_fuseep_weights(layer: torch.nn.Module, weight_prefix: str) -> None:
layer.w2_weight_scale = torch.nn.Parameter( layer.w2_weight_scale = torch.nn.Parameter(
w2_scale.to(torch.float32), requires_grad=False w2_scale.to(torch.float32), requires_grad=False
) )
elif get_exec().moe.fuseep_mode == 2: elif get_server_args().fuseep_mode == 2:
# -- The fused MoE optimization mode "2": dispatch_ffn_combine -- # -- The fused MoE optimization mode "2": dispatch_ffn_combine --
if weight_prefix == "w13": if weight_prefix == "w13":
w13_weight = _release_weight_cache(layer.w13_weight) w13_weight = _release_weight_cache(layer.w13_weight)
+5 -3
View File
@@ -22,7 +22,9 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import divide from sglang.srt.distributed import (
divide,
)
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
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
@@ -31,7 +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_exec, get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
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,
@@ -87,7 +89,7 @@ logger = logging.getLogger(__name__)
class SiluAndMul(MultiPlatformOp): class SiluAndMul(MultiPlatformOp):
def __init__(self, *args, **kwargs): def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs) super().__init__(*args, **kwargs)
if get_exec().deterministic.rl_on_policy_target is not None: if get_server_args().rl_on_policy_target is not None:
self._forward_method = self.forward_native self._forward_method = self.forward_native
elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get(): elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get():
self._forward_method = self.forward_aiter self._forward_method = self.forward_aiter
@@ -37,14 +37,10 @@ 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 ( from sglang.srt.runtime_context import get_parallel, get_server_args
get_device, from sglang.srt.state_capturer.indexer_topk import (
get_exec, maybe_capture_indexer_topk,
get_parallel,
get_schedule,
get_server_args,
) )
from sglang.srt.state_capturer.indexer_topk import maybe_capture_indexer_topk
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
ceil_align, ceil_align,
@@ -109,7 +105,9 @@ if is_npu():
import torch_npu import torch_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 get_attn_tp_group from sglang.srt.distributed import (
get_attn_tp_group,
)
from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.distributed.parallel_state import get_pp_group
from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.communicator import ScatterMode from sglang.srt.layers.communicator import ScatterMode
@@ -460,7 +458,7 @@ class Indexer(MultiPlatformOp):
base=rope_theta, # type: ignore base=rope_theta, # type: ignore
rope_scaling=rope_scaling, rope_scaling=rope_scaling,
is_neox_style=is_neox_style, is_neox_style=is_neox_style,
device=get_device().device, device=get_server_args().device,
) )
self.block_size = block_size self.block_size = block_size
self.scale_fmt = scale_fmt self.scale_fmt = scale_fmt
@@ -471,7 +469,7 @@ class Indexer(MultiPlatformOp):
self.num_local_tokens = getattr(config, "index_local_tokens", 0) self.num_local_tokens = getattr(config, "index_local_tokens", 0)
self.paged_mqa_logits_backend = DSAPagedMQALogitsBackend.resolve( self.paged_mqa_logits_backend = DSAPagedMQALogitsBackend.resolve(
get_exec().kernel.dsa_paged_mqa_logits_backend get_server_args().dsa_paged_mqa_logits_backend
) )
@contextlib.contextmanager @contextlib.contextmanager
@@ -1057,7 +1055,7 @@ class Indexer(MultiPlatformOp):
total_mem = torch.cuda.get_device_properties(device_index).total_memory total_mem = torch.cuda.get_device_properties(device_index).total_memory
total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION) total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION)
mem_fraction_static = get_schedule().mem_fraction_static mem_fraction_static = get_server_args().mem_fraction_static
if mem_fraction_static is None: if mem_fraction_static is None:
static_budget = total_mem_budget static_budget = total_mem_budget
else: else:
@@ -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 ( 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.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip
from sglang.srt.utils.common import ceil_align, ceil_div 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(): def is_dsa_enable_prefill_cp():
return get_parallel().enable_dsa_prefill_context_parallel return get_server_args().enable_dsa_prefill_context_parallel
def is_dsa_prefill_cp_in_seq_split(): def is_dsa_prefill_cp_in_seq_split():
return ( return (
is_dsa_enable_prefill_cp() is_dsa_enable_prefill_cp()
and get_parallel().dsa_prefill_cp_mode == "in-seq-split" and get_server_args().dsa_prefill_cp_mode == "in-seq-split"
) )
def is_dsa_prefill_cp_round_robin_split(): def is_dsa_prefill_cp_round_robin_split():
return ( return (
is_dsa_enable_prefill_cp() is_dsa_enable_prefill_cp()
and get_parallel().dsa_prefill_cp_mode == "round-robin-split" and get_server_args().dsa_prefill_cp_mode == "round-robin-split"
) )
@@ -28,14 +28,16 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( 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_exec, get_parallel from sglang.srt.runtime_context import get_parallel
from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
from sglang.srt.utils import add_prefix, is_cuda, is_hip, is_xpu from sglang.srt.utils import add_prefix, is_cuda, is_hip, is_xpu
from sglang.srt.utils.common import is_sm120_supported from sglang.srt.utils.common import is_sm120_supported
if TYPE_CHECKING: if TYPE_CHECKING:
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.dsv4.compressor import CompressorBackendMixin from sglang.srt.layers.attention.dsv4.compressor import (
CompressorBackendMixin,
)
from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.quantization import QuantizationConfig
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 from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -127,7 +129,9 @@ def _aiter_fp8_paged_mqa_logits(
clean_logits: bool = False, clean_logits: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
"""Wrapper adapting aiter's deepgemm_fp8_paged_mqa_logits to SGLang's interface.""" """Wrapper adapting aiter's deepgemm_fp8_paged_mqa_logits to SGLang's interface."""
from aiter.ops.triton.attention.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits from aiter.ops.triton.attention.pa_mqa_logits import (
deepgemm_fp8_paged_mqa_logits,
)
batch_size = q_fp8.shape[0] batch_size = q_fp8.shape[0]
next_n = q_fp8.shape[1] next_n = q_fp8.shape[1]
@@ -834,8 +838,9 @@ class C4Indexer(nn.Module):
self.rotary_emb = rotary_emb self.rotary_emb = rotary_emb
self.freqs_cis = freqs_cis self.freqs_cis = freqs_cis
self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5 self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5
from sglang.srt.runtime_context import get_server_args
self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer self.use_fp4_indexer = get_server_args().enable_deepseek_v4_fp4_indexer
self.alt_streams = alt_streams self.alt_streams = alt_streams
def compute_q( def compute_q(
@@ -13,7 +13,9 @@ from sglang.kernels.ops.attention.metadata import (
) )
from sglang.kernels.ops.attention.pa_page_table import _build_pa_page_table from sglang.kernels.ops.attention.pa_page_table import _build_pa_page_table
from sglang.kernels.ops.attention.utils import assert_buffer_fits from sglang.kernels.ops.attention.utils import assert_buffer_fits
from sglang.kernels.ops.kvcache.trtllm_mha_page_table import build_trtllm_mha_page_table from sglang.kernels.ops.kvcache.trtllm_mha_page_table import (
build_trtllm_mha_page_table,
)
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy
@@ -26,7 +28,7 @@ from sglang.srt.layers.utils.cp_utils import (
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.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_schedule from sglang.srt.runtime_context import get_server_args
from sglang.srt.speculative.ragged_verify import build_ragged_target_verify_geometry from sglang.srt.speculative.ragged_verify import build_ragged_target_verify_geometry
from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
@@ -164,12 +166,9 @@ class FlashAttentionBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
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.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype = model_runner.kv_cache_dtype
from sglang.srt.runtime_context import get_model
self.kv_cache_dtype_str = getattr( self.kv_cache_dtype_str = get_model().kv_cache_dtype
model_runner,
"kv_cache_dtype_str",
model_runner.server_args.kv_cache_dtype,
)
self.kv_cache_is_mxfp8 = self.kv_cache_dtype_str == "mxfp8" self.kv_cache_is_mxfp8 = self.kv_cache_dtype_str == "mxfp8"
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
# Static page-table width (upper bound). The device-side page-table build # Static page-table width (upper bound). The device-side page-table build
@@ -1480,7 +1479,7 @@ class FlashAttentionBackend(AttentionBackend):
): ):
# Do multi-head attention with chunked prefix cache # Do multi-head attention with chunked prefix cache
if forward_batch.attn_attend_prefix_cache: if forward_batch.attn_attend_prefix_cache:
assert not get_schedule().disable_chunked_prefix_cache assert not get_server_args().disable_chunked_prefix_cache
# MHA for chunked prefix kv cache when running model with MLA # MHA for chunked prefix kv cache when running model with MLA
assert forward_batch.prefix_chunk_idx is not None assert forward_batch.prefix_chunk_idx is not None
assert forward_batch.prefix_chunk_cu_seq_lens is not None assert forward_batch.prefix_chunk_cu_seq_lens is not None
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_disagg, get_exec, get_parallel, get_schedule from sglang.srt.runtime_context import get_parallel
""" """
Support attention backend for flashinfer MLA. Support attention backend for flashinfer MLA.
@@ -32,7 +32,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
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_buffer from sglang.srt.runtime_context import get_buffer, get_server_args
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.speculative.spec_utils import ( from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width, draft_kv_indices_buffer_width,
@@ -223,9 +223,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
self.enable_chunk_kv = ( self.enable_chunk_kv = (
not skip_prefill not skip_prefill
and get_disagg().disaggregation_mode != "decode" and get_server_args().disaggregation_mode != "decode"
and not get_schedule().disable_chunked_prefix_cache and not get_server_args().disable_chunked_prefix_cache
and not get_exec().kernel.flashinfer_mla_disable_ragged and not get_server_args().flashinfer_mla_disable_ragged
) )
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
@@ -401,7 +401,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
prefix_lens = forward_batch.extend_prefix_lens prefix_lens = forward_batch.extend_prefix_lens
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
use_ragged = ( use_ragged = (
not get_exec().kernel.flashinfer_mla_disable_ragged not get_server_args().flashinfer_mla_disable_ragged
and extend_no_prefix and extend_no_prefix
# Piecewise cuda graph should use paged prefill to be compatible with prefix cache # Piecewise cuda graph should use paged prefill to be compatible with prefix cache
and not is_in_tc_piecewise_cuda_graph() and not is_in_tc_piecewise_cuda_graph()
@@ -19,7 +19,7 @@ from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
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.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_exec, get_memory, get_server_args from sglang.srt.runtime_context import get_server_args
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
@@ -350,7 +350,7 @@ class MambaAttnBackendBase(AttentionBackend):
"""Per-row (length bs) bool flush mask = the radix track's seq_lens_cpu % """Per-row (length bs) bool flush mask = the radix track's seq_lens_cpu %
mamba_track_interval == 0, so force-flush and snapshot fire on the same mamba_track_interval == 0, so force-flush and snapshot fire on the same
steps (no off-by-one).""" steps (no off-by-one)."""
interval = get_exec().mamba.mamba_track_interval interval = get_server_args().mamba_track_interval
if seq_lens_cpu is None: if seq_lens_cpu is None:
# Should not happen for the supported config; stay safe and never flush. # Should not happen for the supported config; stay safe and never flush.
return torch.zeros((bs,), dtype=torch.bool) return torch.zeros((bs,), dtype=torch.bool)
@@ -764,7 +764,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
# Page-major stores state strided; only the stride-aware Triton causal-conv # Page-major stores state strided; only the stride-aware Triton causal-conv
# reads it (CUDA causal_conv1d garbles it). A model may also force Triton. # reads it (CUDA causal_conv1d garbles it). A model may also force Triton.
use_triton_causal_conv = ( use_triton_causal_conv = (
use_triton_causal_conv or get_memory().enable_page_major_kv_layout use_triton_causal_conv or get_server_args().enable_page_major_kv_layout
) )
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id) layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
mixer_out, intermediate_states = mixer.forward( mixer_out, intermediate_states = mixer.forward(
@@ -38,11 +38,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
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 ( from sglang.srt.runtime_context import get_buffer, get_parallel, get_server_args
get_buffer,
get_parallel,
get_schedule,
)
from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2 from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2
if is_flashinfer_available(): if is_flashinfer_available():
@@ -201,7 +197,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
self.forward_prefill_metadata: Optional[TRTLLMMLAPrefillMetadata] = None self.forward_prefill_metadata: Optional[TRTLLMMLAPrefillMetadata] = None
self.forward_decode_metadata: Union[TRTLLMMLADecodeMetadata, None] = None self.forward_decode_metadata: Union[TRTLLMMLADecodeMetadata, None] = None
self.disable_chunked_prefix_cache = get_schedule().disable_chunked_prefix_cache self.disable_chunked_prefix_cache = (
get_server_args().disable_chunked_prefix_cache
)
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.cuda_graph_custom_mask = None self.cuda_graph_custom_mask = None
+9 -6
View File
@@ -15,7 +15,7 @@ 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.models.utils import apply_qk_norm from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel 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,
@@ -69,7 +69,9 @@ if _is_npu:
if _is_xpu: if _is_xpu:
from sgl_kernel.flash_attn import flash_attn_varlen_func from sgl_kernel.flash_attn import flash_attn_varlen_func
from sglang.kernels.ops.attention.prefill_attention import context_attention_fwd from sglang.kernels.ops.attention.prefill_attention import (
context_attention_fwd,
)
from sglang.srt.distributed import ( from sglang.srt.distributed import (
split_tensor_along_last_dim, split_tensor_along_last_dim,
tensor_model_parallel_all_gather, tensor_model_parallel_all_gather,
@@ -84,6 +86,7 @@ from sglang.srt.layers.linear import (
from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.layers.rotary_embedding import apply_rotary_pos_emb from sglang.srt.layers.rotary_embedding import apply_rotary_pos_emb
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_pos_emb_native_eager from sglang.srt.layers.rotary_embedding.utils import apply_rotary_pos_emb_native_eager
from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
@@ -1042,7 +1045,7 @@ class VisionAttention(nn.Module):
# Select attention backend via a unified method # Select attention backend via a unified method
_passed_backend = qkv_backend _passed_backend = qkv_backend
qkv_backend = self._determine_attention_backend(_passed_backend) qkv_backend = self._determine_attention_backend(_passed_backend)
if get_mm().mm_attention_backend is None and _passed_backend is None: if get_server_args().mm_attention_backend is None and _passed_backend is None:
print_info_once(f"Multimodal attention backend not set. Use {qkv_backend}.") print_info_once(f"Multimodal attention backend not set. Use {qkv_backend}.")
print_info_once(f"Using {qkv_backend} as multimodal attention backend.") print_info_once(f"Using {qkv_backend} as multimodal attention backend.")
@@ -1121,7 +1124,7 @@ class VisionAttention(nn.Module):
weight_dtype=torch.float32, weight_dtype=torch.float32,
cast_x_before_out_mul=True, cast_x_before_out_mul=True,
) )
if get_exec().deterministic.rl_on_policy_target is not None if get_server_args().rl_on_policy_target is not None
else {} else {}
) )
q_norm = RMSNorm( q_norm = RMSNorm(
@@ -1149,7 +1152,7 @@ class VisionAttention(nn.Module):
- CUDA (other): "triton_attn" - CUDA (other): "triton_attn"
- Non-CUDA: "sdpa" - Non-CUDA: "sdpa"
""" """
override_backend = get_mm().mm_attention_backend override_backend = get_server_args().mm_attention_backend
if override_backend is not None: if override_backend is not None:
backend = override_backend backend = override_backend
elif passed_backend is not None: elif passed_backend is not None:
@@ -1254,7 +1257,7 @@ class VisionAttention(nn.Module):
x = x.unsqueeze(0) x = x.unsqueeze(0)
assert x.dim() == 3, x.shape assert x.dim() == 3, x.shape
if ( if (
get_exec().deterministic.rl_on_policy_target is not None get_server_args().rl_on_policy_target is not None
and position_embeddings is not None and position_embeddings is not None
): ):
assert isinstance(position_embeddings, tuple), ( assert isinstance(position_embeddings, tuple), (
@@ -15,7 +15,7 @@ from sglang.srt.layers.attention.flashattention_backend import (
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.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_schedule from sglang.srt.runtime_context import get_server_args
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -69,12 +69,9 @@ class XPUAttentionBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
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.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype = model_runner.kv_cache_dtype
from sglang.srt.runtime_context import get_model
self.kv_cache_dtype_str = getattr( self.kv_cache_dtype_str = get_model().kv_cache_dtype
model_runner,
"kv_cache_dtype_str",
model_runner.server_args.kv_cache_dtype,
)
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.skip_prefill = skip_prefill self.skip_prefill = skip_prefill
@@ -643,7 +640,7 @@ class XPUAttentionBackend(AttentionBackend):
): ):
# Do multi-head attention with chunked prefix cache # Do multi-head attention with chunked prefix cache
if forward_batch.attn_attend_prefix_cache: if forward_batch.attn_attend_prefix_cache:
assert not get_schedule().disable_chunked_prefix_cache assert not get_server_args().disable_chunked_prefix_cache
# MHA for chunked prefix kv cache when running model with MLA # MHA for chunked prefix kv cache when running model with MLA
assert forward_batch.prefix_chunk_idx is not None assert forward_batch.prefix_chunk_idx is not None
assert forward_batch.prefix_chunk_cu_seq_lens is not None assert forward_batch.prefix_chunk_cu_seq_lens is not None
+12 -17
View File
@@ -72,12 +72,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 ( from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args
get_exec,
get_forward,
get_parallel,
get_spec,
)
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 (
get_bool_env_var, get_bool_env_var,
@@ -175,7 +170,7 @@ def apply_flashinfer_allreduce_fusion(batch_size: int):
and batch_size > 0 and batch_size > 0
and batch_size <= FUSE_ALLREDUCE_MAX_BATCH_SIZE and batch_size <= FUSE_ALLREDUCE_MAX_BATCH_SIZE
and not is_dp_attention_enabled() and not is_dp_attention_enabled()
and get_exec().comm.flashinfer_allreduce_fusion_backend is not None and get_server_args().flashinfer_allreduce_fusion_backend is not None
and not is_flashinfer_allreduce_unavailable() and not is_flashinfer_allreduce_unavailable()
) )
@@ -191,7 +186,7 @@ def apply_aiter_all_reduce_fusion(input_tensor: torch.Tensor):
and total_bytes <= 8 * 1024 * 8192 and total_bytes <= 8 * 1024 * 8192
and get_parallel().tp_size != 6 and get_parallel().tp_size != 6
and not is_dp_attention_enabled() and not is_dp_attention_enabled()
and get_exec().comm.enable_aiter_allreduce_fusion and get_server_args().enable_aiter_allreduce_fusion
) )
@@ -270,7 +265,7 @@ class AttnTpContext:
def init_context(self, q_lora_rank, is_dsa): def init_context(self, q_lora_rank, is_dsa):
self.is_dsa = is_dsa self.is_dsa = is_dsa
self.allow_input_scattered = ( self.allow_input_scattered = (
get_parallel().enable_attn_tp_input_scattered get_server_args().enable_attn_tp_input_scattered
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
@@ -279,9 +274,9 @@ class AttnTpContext:
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()
and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE)
and get_spec().speculative_algorithm != "EAGLE3" and get_server_args().speculative_algorithm != "EAGLE3"
) )
if get_parallel().enable_attn_tp_input_scattered: if get_server_args().enable_attn_tp_input_scattered:
if not self.allow_input_scattered: if not self.allow_input_scattered:
logging.info( logging.info(
"attn_tp_input_scattered is not enabled while other conditions are not met" "attn_tp_input_scattered is not enabled while other conditions are not met"
@@ -412,7 +407,7 @@ class LayerScatterModes:
not context.is_layer_sparse not context.is_layer_sparse
and context.is_next_layer_sparse and context.is_next_layer_sparse
and enable_moe_dense_fully_dp() and enable_moe_dense_fully_dp()
and get_exec().overlap.enable_two_batch_overlap and get_server_args().enable_two_batch_overlap
) )
@classmethod @classmethod
@@ -439,11 +434,11 @@ class LayerScatterModes:
def enable_moe_dense_fully_dp(): def enable_moe_dense_fully_dp():
return get_parallel().moe_dense_tp_size == 1 return get_server_args().moe_dense_tp_size == 1
def enable_dwdp(): def enable_dwdp():
return get_parallel().dwdp_size > 1 return get_server_args().dwdp_size > 1
class LayerCommunicator: class LayerCommunicator:
@@ -476,7 +471,7 @@ class LayerCommunicator:
) )
self._post_init_communicate() self._post_init_communicate()
self._speculative_algo = SpeculativeAlgorithm.from_string( self._speculative_algo = SpeculativeAlgorithm.from_string(
get_spec().speculative_algorithm get_server_args().speculative_algorithm
) )
def _post_init_communicate(self): def _post_init_communicate(self):
@@ -845,7 +840,7 @@ class LayerCommunicator:
and get_parallel().tp_size != 6 and get_parallel().tp_size != 6
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 get_exec().comm.enable_aiter_allreduce_fusion and get_server_args().enable_aiter_allreduce_fusion
) )
) )
and (not self.is_last_layer) and (not self.is_last_layer)
@@ -1150,7 +1145,7 @@ class CommunicateWithAllReduceAndLayerNormFn:
if not handled: if not handled:
quantize_communications = ( quantize_communications = (
not forward_batch.forward_mode.is_decode_or_idle() not forward_batch.forward_mode.is_decode_or_idle()
and get_exec().comm.enable_quant_communications and get_server_args().enable_quant_communications
) )
if quantize_communications: if quantize_communications:
hidden_states = attention_tensor_model_parallel_quant_all_reduce( hidden_states = attention_tensor_model_parallel_quant_all_reduce(
+7 -3
View File
@@ -48,10 +48,12 @@ from sglang.srt.layers.cp.base import (
CPAttentionBackendKind, CPAttentionBackendKind,
) )
from sglang.srt.layers.cp.padding import pad_local_rows from sglang.srt.layers.cp.padding import pad_local_rows
from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.dp_attention import (
is_allocation_symmetric,
)
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_device, get_parallel from sglang.srt.runtime_context import get_parallel
@dataclass @dataclass
@@ -206,8 +208,10 @@ class ZigzagCPStrategy(ContextParallelStrategy):
actual_seq_q_prev_list.append(block_sizes[cp_rank]) actual_seq_q_prev_list.append(block_sizes[cp_rank])
actual_seq_q_next_list.append(block_sizes[cp_segment_num - cp_rank - 1]) actual_seq_q_next_list.append(block_sizes[cp_segment_num - cp_rank - 1])
from sglang.srt.runtime_context import get_server_args
try: try:
device = torch.device(get_device().device) device = torch.device(get_server_args().device)
except Exception: except Exception:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
cu_prev = [0] + list(accumulate(actual_seq_q_prev_list)) cu_prev = [0] + list(accumulate(actual_seq_q_prev_list))
+7 -7
View File
@@ -26,7 +26,7 @@ from sglang.kernels.ops.attention.dcp_kernels import (
) )
from sglang.srt.layers.dcp.layout import update_local_kv_lens_for_dcp from sglang.srt.layers.dcp.layout import update_local_kv_lens_for_dcp
from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
from sglang.srt.runtime_context import get_device, get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
def prepare_decode_context_parallel_metadata( def prepare_decode_context_parallel_metadata(
@@ -53,12 +53,12 @@ def prepare_decode_context_parallel_metadata(
extend_prefix_starts = torch.zeros( extend_prefix_starts = torch.zeros(
len(seq_lens), len(seq_lens),
dtype=torch.int32, dtype=torch.int32,
device=get_device().device, device=get_server_args().device,
) )
extend_cu_prefix_lens = torch.zeros( extend_cu_prefix_lens = torch.zeros(
len(seq_lens) + 1, len(seq_lens) + 1,
dtype=torch.int32, dtype=torch.int32,
device=get_device().device, device=get_server_args().device,
) )
extend_cu_prefix_lens[1:] = torch.cumsum(extend_prefix_lens, dim=0) extend_cu_prefix_lens[1:] = torch.cumsum(extend_prefix_lens, dim=0)
extend_cu_prefix_lens = extend_cu_prefix_lens[:-1] extend_cu_prefix_lens = extend_cu_prefix_lens[:-1]
@@ -67,7 +67,7 @@ def prepare_decode_context_parallel_metadata(
dcp_prefix_kv_indices = torch.empty( dcp_prefix_kv_indices = torch.empty(
sum(extend_prefix_lens_cpu), sum(extend_prefix_lens_cpu),
dtype=torch.int32, dtype=torch.int32,
device=get_device().device, device=get_server_args().device,
) )
create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)]( create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)](
req_to_token, req_to_token,
@@ -81,20 +81,20 @@ def prepare_decode_context_parallel_metadata(
dcp_kv_indptr = torch.zeros( dcp_kv_indptr = torch.zeros(
len(seq_lens) + 1, len(seq_lens) + 1,
dtype=torch.int32, dtype=torch.int32,
device=get_device().device, device=get_server_args().device,
) )
dcp_kv_indptr[1:] = seq_lens.cumsum(dim=0) dcp_kv_indptr[1:] = seq_lens.cumsum(dim=0)
dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)] dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)]
dcp_kv_indices = torch.zeros( dcp_kv_indices = torch.zeros(
seq_lens_sum, seq_lens_sum,
dtype=torch.int32, dtype=torch.int32,
device=get_device().device, device=get_server_args().device,
) )
extend_cu_lens = torch.zeros( extend_cu_lens = torch.zeros(
len(seq_lens) + 1, len(seq_lens) + 1,
dtype=torch.int32, dtype=torch.int32,
device=get_device().device, device=get_server_args().device,
) )
extend_cu_lens[1:] = torch.cumsum(extend_seq_lens, dim=0) extend_cu_lens[1:] = torch.cumsum(extend_seq_lens, dim=0)
extend_cu_lens = extend_cu_lens[:-1] extend_cu_lens = extend_cu_lens[:-1]
+9 -10
View File
@@ -32,7 +32,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_exec, get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
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,
@@ -148,7 +148,9 @@ if _is_cuda:
_jit_rmsnorm_hf = None _jit_rmsnorm_hf = None
from sglang.jit_kernel.norm import fused_add_rmsnorm as _jit_fused_add_rmsnorm from sglang.jit_kernel.norm import fused_add_rmsnorm as _jit_fused_add_rmsnorm
from sglang.jit_kernel.norm import is_supported_jit_fused_add_rmsnorm_hidden_size from sglang.jit_kernel.norm import (
is_supported_jit_fused_add_rmsnorm_hidden_size,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -219,7 +221,7 @@ def _forward_with_allreduce_fusion(
return fused_result return fused_result
# For AITER route, preserve correctness when fused path is unavailable. # For AITER route, preserve correctness when fused path is unavailable.
if _use_aiter and get_exec().comm.enable_aiter_allreduce_fusion: if _use_aiter and get_server_args().enable_aiter_allreduce_fusion:
x = tensor_model_parallel_all_reduce(x) x = tensor_model_parallel_all_reduce(x)
return norm_module.forward(x, residual, None) return norm_module.forward(x, residual, None)
@@ -421,7 +423,7 @@ class RMSNorm(MultiPlatformOp):
if ( if (
residual is not None residual is not None
or self.cast_x_before_out_mul or self.cast_x_before_out_mul
or get_exec().deterministic.rl_on_policy_target == "fsdp" or get_server_args().rl_on_policy_target == "fsdp"
): ):
return self.forward_native(x, residual, post_residual_addition) return self.forward_native(x, residual, post_residual_addition)
out = rms_norm_batch_invariant( out = rms_norm_batch_invariant(
@@ -528,7 +530,7 @@ class RMSNorm(MultiPlatformOp):
if ( if (
residual is not None residual is not None
or self.cast_x_before_out_mul or self.cast_x_before_out_mul
or get_exec().deterministic.rl_on_policy_target == "fsdp" or get_server_args().rl_on_policy_target == "fsdp"
or (self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0) or (self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0)
): ):
return self.forward_native(x, residual, post_residual_addition) return self.forward_native(x, residual, post_residual_addition)
@@ -589,7 +591,7 @@ class RMSNorm(MultiPlatformOp):
if ( if (
residual is not None residual is not None
or self.cast_x_before_out_mul or self.cast_x_before_out_mul
or get_exec().deterministic.rl_on_policy_target == "fsdp" or get_server_args().rl_on_policy_target == "fsdp"
): ):
return self.forward_native(x, residual, post_residual_addition) return self.forward_native(x, residual, post_residual_addition)
return rms_norm_batch_invariant( return rms_norm_batch_invariant(
@@ -716,10 +718,7 @@ class RMSNorm(MultiPlatformOp):
if self.variance_size_override is not None: if self.variance_size_override is not None:
return self.forward_native(x, residual, post_residual_addition) return self.forward_native(x, residual, post_residual_addition)
if is_batch_invariant_mode_enabled(): if is_batch_invariant_mode_enabled():
if ( if residual is not None or get_server_args().rl_on_policy_target == "fsdp":
residual is not None
or get_exec().deterministic.rl_on_policy_target == "fsdp"
):
return self.forward_native(x, residual, post_residual_addition) return self.forward_native(x, residual, post_residual_addition)
return rms_norm_batch_invariant( return rms_norm_batch_invariant(
x, x,
+11 -5
View File
@@ -25,7 +25,9 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, use_symmetric_memory,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.dp_attention import (
is_allocation_symmetric,
)
from sglang.srt.layers.moe.utils import should_skip_mlp_all_reduce from sglang.srt.layers.moe.utils import should_skip_mlp_all_reduce
from sglang.srt.layers.parameter import ( from sglang.srt.layers.parameter import (
BasevLLMParameter, BasevLLMParameter,
@@ -37,7 +39,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_exec, get_parallel from sglang.srt.runtime_context import get_parallel, get_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
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -757,7 +759,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
shard_offsets.append((i, current_shard_offset, output_size)) shard_offsets.append((i, current_shard_offset, output_size))
current_shard_offset += output_size current_shard_offset += output_size
if _is_cpu: if _is_cpu:
from sglang.srt.model_loader.weight_utils import pad_loaded_weight from sglang.srt.model_loader.weight_utils import (
pad_loaded_weight,
)
loaded_weight = pad_loaded_weight( loaded_weight = pad_loaded_weight(
loaded_weight, param.output_dim, output_sizes loaded_weight, param.output_dim, output_sizes
@@ -801,7 +805,9 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
current_block_offset += shard_block_size current_block_offset += shard_block_size
if _is_cpu: if _is_cpu:
from sglang.srt.model_loader.weight_utils import pad_loaded_weight from sglang.srt.model_loader.weight_utils import (
pad_loaded_weight,
)
loaded_weight = pad_loaded_weight( loaded_weight = pad_loaded_weight(
loaded_weight, param.output_dim, shard_block_sizes loaded_weight, param.output_dim, shard_block_sizes
@@ -1590,7 +1596,7 @@ class RowParallelLinear(LinearBase):
quantize_communications = ( quantize_communications = (
( (
not forward_batch.forward_mode.is_decode_or_idle() not forward_batch.forward_mode.is_decode_or_idle()
and get_exec().comm.enable_quant_communications and get_server_args().enable_quant_communications
) )
if forward_batch is not None if forward_batch is not None
else False else False
+5 -5
View File
@@ -47,7 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
is_cpu, is_cpu,
is_npu, is_npu,
@@ -345,8 +345,8 @@ class LogitsProcessor(nn.Module):
self.config = config self.config = config
self.vocab_size = config.vocab_size self.vocab_size = config.vocab_size
self.logit_scale = logit_scale self.logit_scale = logit_scale
self.use_attn_tp_group = get_parallel().enable_dp_lm_head self.use_attn_tp_group = get_server_args().enable_dp_lm_head
self.use_fp32_lm_head = get_exec().features.enable_fp32_lm_head self.use_fp32_lm_head = get_server_args().enable_fp32_lm_head
if self.use_attn_tp_group: if self.use_attn_tp_group:
self.attn_tp_size = get_parallel().attn_tp_size self.attn_tp_size = get_parallel().attn_tp_size
self.do_tensor_parallel_all_gather = ( self.do_tensor_parallel_all_gather = (
@@ -370,8 +370,8 @@ class LogitsProcessor(nn.Module):
self.final_logit_softcapping = None self.final_logit_softcapping = None
self.return_full_logits = return_full_logits self.return_full_logits = return_full_logits
self.enable_mis = get_exec().features.enable_mis self.enable_mis = get_server_args().enable_mis
self.rl_on_policy_target = get_exec().deterministic.rl_on_policy_target self.rl_on_policy_target = get_server_args().rl_on_policy_target
self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer( self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer(
max_tokens=triton_symm_mem_ag.recommended_max_tokens( max_tokens=triton_symm_mem_ag.recommended_max_tokens(
+5 -3
View File
@@ -7,7 +7,9 @@ import torch
from torch import nn from torch import nn
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.eplb.expert_location_dispatch import ( from sglang.srt.eplb.expert_location_dispatch import (
ExpertLocationDispatchInfo, ExpertLocationDispatchInfo,
topk_ids_logical_to_physical, topk_ids_logical_to_physical,
@@ -20,7 +22,6 @@ from sglang.srt.layers.moe.topk import (
remap_topk_for_per_rank_shared_slots, remap_topk_for_per_rank_shared_slots,
) )
from sglang.srt.layers.moe.utils import has_per_rank_fused_shared_slots from sglang.srt.layers.moe.utils import has_per_rank_fused_shared_slots
from sglang.srt.runtime_context import get_exec
from sglang.srt.utils import is_hip, is_npu from sglang.srt.utils import is_hip, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -43,9 +44,10 @@ class HashTopK(nn.Module):
): ):
super().__init__() super().__init__()
self.layer_id = layer_id self.layer_id = layer_id
from sglang.srt.runtime_context import get_server_args
self.enable_waterfill = ( self.enable_waterfill = (
num_fused_shared_experts > 0 and get_exec().moe.enable_waterfill num_fused_shared_experts > 0 and get_server_args().enable_waterfill
) )
self.waterfill_balancer = None self.waterfill_balancer = None
@@ -28,7 +28,7 @@ from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
from sglang.srt.layers.moe.utils import get_moe_padding_size from sglang.srt.layers.moe.utils import get_moe_padding_size
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_server_args
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,
@@ -506,7 +506,7 @@ def _fused_moe_kernel_sequence(
out_hidden_states = torch.empty_like(hidden_states) out_hidden_states = torch.empty_like(hidden_states)
use_fused_moe_sum_all_reduce = ( use_fused_moe_sum_all_reduce = (
get_exec().moe.enable_fused_moe_sum_all_reduce get_server_args().enable_fused_moe_sum_all_reduce
and (not no_combine) and (not no_combine)
and (topk > 2) and (topk > 2)
and (not use_int8_w8a16) and (not use_int8_w8a16)
@@ -9,7 +9,7 @@ from typing import Any, Dict, List, Optional, Tuple
import torch import torch
import triton import triton
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import get_device_name, is_hip from sglang.srt.utils import get_device_name, is_hip
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -69,7 +69,7 @@ def get_moe_configs(
kernel on a given batch size bs, the closest batch size in the grid should kernel on a given batch size bs, the closest batch size in the grid should
be picked and the associated configuration chosen to invoke the kernel. be picked and the associated configuration chosen to invoke the kernel.
""" """
if get_exec().deterministic.enable_deterministic_inference: if get_server_args().enable_deterministic_inference:
logger.warning( logger.warning(
"Deterministic inference is enabled, using default MoE kernel config." "Deterministic inference is enabled, using default MoE kernel config."
) )
@@ -187,7 +187,7 @@ def get_default_config(
is_marlin: bool, is_marlin: bool,
block_shape: Optional[List[int]] = None, block_shape: Optional[List[int]] = None,
) -> Dict[str, int]: ) -> Dict[str, int]:
if get_exec().deterministic.enable_deterministic_inference: if get_server_args().enable_deterministic_inference:
config = { config = {
"BLOCK_SIZE_M": 64, "BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64, "BLOCK_SIZE_N": 64,
@@ -21,9 +21,13 @@ from sglang.srt.layers.moe.token_dispatcher import (
from sglang.srt.layers.moe.token_dispatcher.flashinfer_utils import ( from sglang.srt.layers.moe.token_dispatcher.flashinfer_utils import (
TorchDistributedCommBackend, TorchDistributedCommBackend,
) )
from sglang.srt.layers.moe.topk import StandardTopKOutput, TopKOutput, TopKOutputChecker from sglang.srt.layers.moe.topk import (
StandardTopKOutput,
TopKOutput,
TopKOutputChecker,
)
from sglang.srt.layers.moe.utils import get_moe_runner_backend from sglang.srt.layers.moe.utils import get_moe_runner_backend
from sglang.srt.runtime_context import get_schedule, get_spec from sglang.srt.runtime_context import get_server_args
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils import get_int_env_var from sglang.srt.utils import get_int_env_var
@@ -119,7 +123,7 @@ class FlashinferDispatcher(BaseDispatcher):
# max_running_requests is not yet resolved at model-construction time, # max_running_requests is not yet resolved at model-construction time,
# so we use 4096 as a floor to cover decode batches and _dummy_run # so we use 4096 as a floor to cover decode batches and _dummy_run
# (which warms up at batch_size = req_to_token_pool.size). # (which warms up at batch_size = req_to_token_pool.size).
cps = get_schedule().chunked_prefill_size cps = get_server_args().chunked_prefill_size
default_max_tokens = max(cps if cps and cps > 0 else 4096, 4096) default_max_tokens = max(cps if cps and cps > 0 else 4096, 4096)
self.max_num_tokens = get_int_env_var( self.max_num_tokens = get_int_env_var(
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK", "SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK",
@@ -128,7 +132,7 @@ class FlashinferDispatcher(BaseDispatcher):
# Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized. # Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized.
speculative_algo = SpeculativeAlgorithm.from_string( speculative_algo = SpeculativeAlgorithm.from_string(
get_spec().speculative_algorithm get_server_args().speculative_algorithm
) )
if MOE_NVFP4_DISPATCH and not speculative_algo.is_eagle(): if MOE_NVFP4_DISPATCH and not speculative_algo.is_eagle():
total_dispatch_payload_size_per_token = ( total_dispatch_payload_size_per_token = (
@@ -23,7 +23,6 @@ from sglang.srt.layers.moe.token_dispatcher.deepep import (
) )
from sglang.srt.layers.moe.topk import TopKOutput from sglang.srt.layers.moe.topk import TopKOutput
from sglang.srt.layers.moe.utils import DeepEPMode from sglang.srt.layers.moe.utils import DeepEPMode
from sglang.srt.runtime_context import get_parallel
try: try:
from nixl_ep import Buffer from nixl_ep import Buffer
@@ -128,7 +127,9 @@ class NixlEPBuffer:
offset = ElasticEPStateManager.get_ep_join_rank_offset() offset = ElasticEPStateManager.get_ep_join_rank_offset()
global_rank = rank + offset global_rank = rank + offset
max_ep_size = get_parallel().max_ep_size or world_size from sglang.srt.runtime_context import get_server_args
max_ep_size = get_server_args().max_ep_size or world_size
nixl_max_ranks = max_ep_size nixl_max_ranks = max_ep_size
num_rdma_bytes = 0 num_rdma_bytes = 0
@@ -225,8 +226,9 @@ class _NixlEPDispatcherImplBase:
elastic_state.active_ranks if elastic_state is not None else None elastic_state.active_ranks if elastic_state is not None else None
) )
self._active_world_size = dist.get_world_size(group) self._active_world_size = dist.get_world_size(group)
from sglang.srt.runtime_context import get_server_args
_max_ep = get_parallel().max_ep_size or self._active_world_size _max_ep = get_server_args().max_ep_size or self._active_world_size
self._mask_buffer = ( self._mask_buffer = (
torch.zeros(_max_ep, dtype=torch.int32, device="cuda") torch.zeros(_max_ep, dtype=torch.int32, device="cuda")
if self.active_ranks is not None if self.active_ranks is not None
+11 -5
View File
@@ -32,7 +32,7 @@ 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_exec, get_lora, get_parallel 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
@@ -83,7 +83,9 @@ except ImportError:
pass pass
from sglang.kernels.ops.attention.dsv4 import mask_topk_ids from sglang.kernels.ops.attention.dsv4 import mask_topk_ids
from sglang.srt.distributed import 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,
) )
@@ -96,7 +98,9 @@ from sglang.srt.eplb.expert_location_dispatch import (
) )
from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe import get_moe_runner_backend from sglang.srt.layers.moe import get_moe_runner_backend
from sglang.srt.layers.moe.utils import has_per_rank_fused_shared_slots from sglang.srt.layers.moe.utils import (
has_per_rank_fused_shared_slots,
)
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -415,9 +419,10 @@ class TopK(MultiPlatformOp):
assert num_expert_group is not None and topk_group is not None assert num_expert_group is not None and topk_group is not None
self.layer_id = layer_id self.layer_id = layer_id
from sglang.srt.runtime_context import get_server_args
self.enable_waterfill = ( self.enable_waterfill = (
num_fused_shared_experts > 0 and get_exec().moe.enable_waterfill num_fused_shared_experts > 0 and get_server_args().enable_waterfill
) )
self.waterfill_balancer = None self.waterfill_balancer = None
@@ -491,8 +496,9 @@ class TopK(MultiPlatformOp):
# ===== TO BE REFACTORED ==== # ===== TO BE REFACTORED ====
elif get_moe_runner_backend().is_experimental_sgl_trtllm(): elif get_moe_runner_backend().is_experimental_sgl_trtllm():
try: try:
from sglang.srt.runtime_context import get_server_args
use_standard_for_lora = bool(get_lora().enable_lora) use_standard_for_lora = bool(get_server_args().enable_lora)
except ValueError: except ValueError:
use_standard_for_lora = False use_standard_for_lora = False
output_format = ( output_format = (
@@ -13,7 +13,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
) )
from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers import deep_gemm_wrapper
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_exec, get_parallel 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:
@@ -34,6 +34,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
w8a8_block_fp8_matmul_deepgemm, w8a8_block_fp8_matmul_deepgemm,
w8a8_block_fp8_matmul_triton, w8a8_block_fp8_matmul_triton,
) )
from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
ceil_align, ceil_align,
ceil_div, ceil_div,
@@ -1469,7 +1470,9 @@ def requant_block_scale_ue8m0_for_deepgemm(
scales are not already UE8M0, and DeepGEMM can run the layer (bf16 output, scales are not already UE8M0, and DeepGEMM can run the layer (bf16 output,
aligned shape). Returns True when it requantizes. aligned shape). Returns True when it requantizes.
""" """
from sglang.srt.model_loader.utils import should_deepgemm_weight_requant_ue8m0 from sglang.srt.model_loader.utils import (
should_deepgemm_weight_requant_ue8m0,
)
if ( if (
not use_deepgemm_runner not use_deepgemm_runner
@@ -1791,7 +1794,7 @@ def apply_fp8_linear(
if ( if (
input_scale is not None input_scale is not None
and input_scale.numel() == 1 and input_scale.numel() == 1
and get_exec().graph.cuda_graph_config.prefill.tc_compiler == "inductor" and get_server_args().cuda_graph_config.prefill.tc_compiler == "inductor"
): ):
qinput = ( qinput = (
(input_2d * input_scale.reciprocal()) (input_2d * input_scale.reciprocal())
@@ -48,7 +48,7 @@ from sglang.srt.layers.quantization.base_config import (
QuantizeMethodBase, QuantizeMethodBase,
) )
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_exec from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
is_cpu, is_cpu,
@@ -77,7 +77,9 @@ if is_flashinfer_available():
nvfp4_block_scale_interleave, nvfp4_block_scale_interleave,
trtllm_fp4_block_scale_moe, trtllm_fp4_block_scale_moe,
) )
from flashinfer.fused_moe.core import get_w2_permute_indices_with_cache from flashinfer.fused_moe.core import (
get_w2_permute_indices_with_cache,
)
# SM90 mixed-input helpers landed in FlashInfer #3084 (post-0.6.10). Older # SM90 mixed-input helpers landed in FlashInfer #3084 (post-0.6.10). Older
# versions don't ship them; gate at import so unrelated code paths still load. # versions don't ship them; gate at import so unrelated code paths still load.
@@ -332,7 +334,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
self.use_flashinfer = get_moe_runner_backend().is_flashinfer_mxfp4() self.use_flashinfer = get_moe_runner_backend().is_flashinfer_mxfp4()
self.use_marlin = get_moe_runner_backend().is_marlin() self.use_marlin = get_moe_runner_backend().is_marlin()
self.flashinfer_mxfp4_moe_precision = ( self.flashinfer_mxfp4_moe_precision = (
get_exec().moe.flashinfer_mxfp4_moe_precision get_server_args().flashinfer_mxfp4_moe_precision
) )
# When `flashinfer_mxfp4` is enabled, dispatch to one of two FlashInfer # When `flashinfer_mxfp4` is enabled, dispatch to one of two FlashInfer
# entry points depending on the GPU: # entry points depending on the GPU:
@@ -14,7 +14,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
) )
from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe.utils import RoutingMethodType from sglang.srt.layers.moe.utils import RoutingMethodType
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
is_flashinfer_available, is_flashinfer_available,
log_info_on_rank0, log_info_on_rank0,
@@ -51,7 +51,7 @@ class Mxfp4FlashinferTrtllmMoEMethod:
self._fp8 = fp8_method self._fp8 = fp8_method
self.prefix = prefix self.prefix = prefix
self.flashinfer_mxfp4_moe_precision = ( self.flashinfer_mxfp4_moe_precision = (
get_exec().moe.flashinfer_mxfp4_moe_precision get_server_args().flashinfer_mxfp4_moe_precision
) )
def create_moe_runner(self, layer, moe_runner_config): def create_moe_runner(self, layer, moe_runner_config):
@@ -376,7 +376,9 @@ def maybe_fuse_routed_scale_and_shared_add(
from sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe import ( from sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe import (
Mxfp4FlashinferCutlassMoEMethod, Mxfp4FlashinferCutlassMoEMethod,
) )
from sglang.srt.layers.quantization.mxfp4_marlin_moe import Mxfp4MarlinMoEMethod from sglang.srt.layers.quantization.mxfp4_marlin_moe import (
Mxfp4MarlinMoEMethod,
)
fused = isinstance( fused = isinstance(
experts.quant_method, experts.quant_method,
@@ -11,7 +11,7 @@ from sglang.srt.environ import envs
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_server_args
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,
@@ -67,7 +67,9 @@ if _is_npu:
) )
if _is_hip: if _is_hip:
from sglang.kernels.ops.attention.utils import fused_qk_rope_reshape_and_cache from sglang.kernels.ops.attention.utils import (
fused_qk_rope_reshape_and_cache,
)
if _is_xpu: if _is_xpu:
from sgl_kernel import fused_qk_rope_with_cos_sin_cache_inplace from sgl_kernel import fused_qk_rope_with_cos_sin_cache_inplace
@@ -127,7 +129,7 @@ class RotaryEmbedding(MultiPlatformOp):
self._apply_rotary_emb_wrapped = apply_rotary_emb self._apply_rotary_emb_wrapped = apply_rotary_emb
# XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend # XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend
if get_exec().deterministic.rl_on_policy_target is not None or _is_musa: if get_server_args().rl_on_policy_target is not None or _is_musa:
self._forward_method = self.forward_native self._forward_method = self.forward_native
self._apply_rotary_emb_wrapped = torch.compile( self._apply_rotary_emb_wrapped = torch.compile(
dynamic=True, dynamic=True,
@@ -151,7 +153,7 @@ class RotaryEmbedding(MultiPlatformOp):
# create the cache on GPU for faster initialization. This may cause # create the cache on GPU for faster initialization. This may cause
# a slight numerical difference between the HF implementation and ours. # a slight numerical difference between the HF implementation and ours.
init_device = ( init_device = (
"cpu" if get_exec().deterministic.rl_on_policy_target is not None else None "cpu" if get_server_args().rl_on_policy_target is not None else None
) )
inv_freq = 1.0 / ( inv_freq = 1.0 / (
base base
@@ -162,7 +164,7 @@ class RotaryEmbedding(MultiPlatformOp):
/ self.rotary_dim / self.rotary_dim
) )
) )
if get_exec().deterministic.rl_on_policy_target is not None: if get_server_args().rl_on_policy_target is not None:
inv_freq = inv_freq.cuda() inv_freq = inv_freq.cuda()
return inv_freq return inv_freq
@@ -18,7 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import (
yarn_get_mscale_simple, yarn_get_mscale_simple,
yarn_linear_ramp_mask, yarn_linear_ramp_mask,
) )
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
is_cuda, is_cuda,
@@ -42,6 +42,7 @@ if _is_xpu:
from sgl_kernel import multimodal_rotary_embedding from sgl_kernel import multimodal_rotary_embedding
from sglang.kernels.ops.attention.mrope import apply_interleaved_rope_triton from sglang.kernels.ops.attention.mrope import apply_interleaved_rope_triton
from sglang.srt.runtime_context import get_server_args
def apply_interleaved_rope(x: torch.Tensor, mrope_section: list) -> torch.Tensor: def apply_interleaved_rope(x: torch.Tensor, mrope_section: list) -> torch.Tensor:
@@ -131,7 +132,7 @@ class MRotaryEmbedding(RotaryEmbedding):
self.register_buffer("axis_map", axis_map, persistent=False) self.register_buffer("axis_map", axis_map, persistent=False)
else: else:
self.axis_map = None self.axis_map = None
if get_exec().deterministic.rl_on_policy_target is not None: if get_server_args().rl_on_policy_target is not None:
self._forward_method = self.forward_native self._forward_method = self.forward_native
def get_cos_sin_with_position(self, positions): def get_cos_sin_with_position(self, positions):
@@ -143,7 +144,7 @@ class MRotaryEmbedding(RotaryEmbedding):
last_dim = cos_sin.size()[-1] last_dim = cos_sin.size()[-1]
cos, sin = cos_sin.chunk(2, dim=-1) cos, sin = cos_sin.chunk(2, dim=-1)
if self.mrope_interleaved: if self.mrope_interleaved:
if support_triton(get_exec().kernel.attention_backend): if support_triton(get_server_args().attention_backend):
cos = apply_interleaved_rope_triton(cos, self.mrope_section) cos = apply_interleaved_rope_triton(cos, self.mrope_section)
sin = apply_interleaved_rope_triton(sin, self.mrope_section) sin = apply_interleaved_rope_triton(sin, self.mrope_section)
else: else:
+22 -11
View File
@@ -8,21 +8,34 @@ from torch import nn
from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32 from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32
from sglang.srt.distributed import get_tp_group from sglang.srt.distributed import get_tp_group
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.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.logprob_processor import OutputLogprobProcessor from sglang.srt.layers.logprob_processor import (
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args OutputLogprobProcessor,
)
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.sampling.sampling_params import TOP_K_ALL
from sglang.srt.utils.async_probe import sanitize_nan_logits from sglang.srt.utils.async_probe import sanitize_nan_logits
from sglang.srt.utils.common import get_bool_env_var, is_cuda, is_hip, is_musa, is_npu from sglang.srt.utils.common import (
get_bool_env_var,
is_cuda,
is_hip,
is_musa,
is_npu,
)
if is_cuda(): if is_cuda():
from flashinfer.sampling import ( from flashinfer.sampling import (
min_p_sampling_from_probs, min_p_sampling_from_probs,
top_k_top_p_sampling_from_probs, top_k_top_p_sampling_from_probs,
) )
from sgl_kernel import top_k_renorm_prob, top_p_renorm_prob from sgl_kernel import (
top_k_renorm_prob,
top_p_renorm_prob,
)
if is_musa(): if is_musa():
from sgl_kernel import ( from sgl_kernel import (
@@ -61,14 +74,12 @@ class Sampler(nn.Module):
if is_dp_attention_enabled(): if is_dp_attention_enabled():
self.tp_sync_group = get_parallel().attn_tp_group.device_group self.tp_sync_group = get_parallel().attn_tp_group.device_group
self.rl_on_policy_target = get_exec().deterministic.rl_on_policy_target self.rl_on_policy_target = get_server_args().rl_on_policy_target
# In RL on-policy mode, deterministic inference is automatically enabled. # In RL on-policy mode, deterministic inference is automatically enabled.
self.enable_deterministic = ( self.enable_deterministic = get_server_args().enable_deterministic_inference
get_exec().deterministic.enable_deterministic_inference
)
# In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer. # In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer.
self.use_log_softmax_logprob = self.rl_on_policy_target is not None self.use_log_softmax_logprob = self.rl_on_policy_target is not None
self.use_ascend_backend = get_exec().kernel.sampling_backend == "ascend" self.use_ascend_backend = get_server_args().sampling_backend == "ascend"
self.output_logprob_processor = OutputLogprobProcessor() self.output_logprob_processor = OutputLogprobProcessor()
@@ -234,7 +245,7 @@ class Sampler(nn.Module):
positions=positions, positions=positions,
) )
else: else:
backend = get_exec().kernel.sampling_backend backend = get_server_args().sampling_backend
if backend == "flashinfer": if backend == "flashinfer":
assert ( assert (
sampling_info.sampling_seed is None sampling_info.sampling_seed is None
+2 -2
View File
@@ -58,13 +58,13 @@ class ContextParallelMetadata:
def is_prefill_context_parallel_enabled(): def is_prefill_context_parallel_enabled():
return get_parallel().enable_prefill_context_parallel return get_server_args().enable_prefill_context_parallel
def is_prefill_cp_in_seq_split(): def is_prefill_cp_in_seq_split():
return ( return (
is_prefill_context_parallel_enabled() is_prefill_context_parallel_enabled()
and get_parallel().prefill_cp_mode == "in-seq-split" and get_server_args().prefill_cp_mode == "in-seq-split"
) )
@@ -48,7 +48,6 @@ from sglang.srt.managers.scheduler import run_scheduler_process
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
from sglang.srt.runtime_context import get_exec
from sglang.srt.server_args import ( from sglang.srt.server_args import (
DP_ATTENTION_HANDSHAKE_PORT_DELTA, DP_ATTENTION_HANDSHAKE_PORT_DELTA,
PortArgs, PortArgs,
@@ -232,7 +231,7 @@ class DataParallelController:
sock_send(worker, obj) sock_send(worker, obj)
def update_active_ranks(self, ranks: ActiveRanksOutput): def update_active_ranks(self, ranks: ActiveRanksOutput):
if get_exec().moe.elastic_ep_backend is not None: if self.server_args.elastic_ep_backend is not None:
if len(ranks.status) != self.max_dp_size: if len(ranks.status) != self.max_dp_size:
logger.warning( logger.warning(
"[Elastic EP][DPC] active rank status len=%d != max_dp_size=%d; " "[Elastic EP][DPC] active rank status len=%d != max_dp_size=%d; "
@@ -485,7 +484,7 @@ class DataParallelController:
logger.debug("Worker port broadcast completed") logger.debug("Worker port broadcast completed")
return worker_ports return worker_ports
finally: finally:
if get_exec().moe.elastic_ep_backend is None: if self.server_args.elastic_ep_backend is None:
rep_socket.close() rep_socket.close()
else: else:
threading.Thread( threading.Thread(
@@ -816,12 +815,6 @@ def run_data_parallel_controller_process(
kill_itself_when_parent_died() kill_itself_when_parent_died()
parent_process = psutil.Process().parent() parent_process = psutil.Process().parent()
# Publish the resolved config at DP-controller process entry: this process
# reads config namespaces (e.g. get_exec().moe.*) in its own address space
# before spawning schedulers.
from sglang.srt.runtime_context import publish
publish(server_args, role="scheduler")
configure_logger(server_args) configure_logger(server_args)
if server_args.enable_trace: if server_args.enable_trace:
process_tracing_init( process_tracing_init(
+5 -11
View File
@@ -33,13 +33,7 @@ from sglang.srt.managers.schedule_batch import (
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.multimodal.evs import EVSEmbeddingResult from sglang.srt.multimodal.evs import EVSEmbeddingResult
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import get_parallel, get_server_args
get_disagg,
get_parallel,
get_schedule,
get_server_args,
get_serving,
)
from sglang.srt.utils import flatten_nested_list, is_hip, is_npu, print_warning_once from sglang.srt.utils import flatten_nested_list, is_hip, is_npu, print_warning_once
from sglang.srt.utils.stale_shm_cleanup import make_shm_name from sglang.srt.utils.stale_shm_cleanup import make_shm_name
from sglang.utils import logger from sglang.utils import logger
@@ -884,7 +878,7 @@ def _adjust_embedding_length(
f"tokens from multimodal embeddings." f"tokens from multimodal embeddings."
) )
if num_mm_tokens_in_input_ids < num_mm_tokens_in_embedding: if num_mm_tokens_in_input_ids < num_mm_tokens_in_embedding:
chunked_prefill_size = get_schedule().chunked_prefill_size chunked_prefill_size = get_server_args().chunked_prefill_size
if chunked_prefill_size != -1: if chunked_prefill_size != -1:
logger.warning( logger.warning(
"You may want to avoid this issue by raising `chunked_prefill_size`, or disabling chunked prefill" "You may want to avoid this issue by raising `chunked_prefill_size`, or disabling chunked prefill"
@@ -1293,7 +1287,7 @@ def general_mm_embed_routine(
feature = getattr(mm_item, "feature", None) feature = getattr(mm_item, "feature", None)
if isinstance(feature, torch.Tensor) and feature.is_cuda: if isinstance(feature, torch.Tensor) and feature.is_cuda:
mm_item.feature = feature.to("cpu", non_blocking=True) mm_item.feature = feature.to("cpu", non_blocking=True)
if get_disagg().language_only: if get_server_args().language_only:
precomputed_embeddings = getattr( precomputed_embeddings = getattr(
mm_item, "precomputed_embeddings", None mm_item, "precomputed_embeddings", None
) )
@@ -1973,7 +1967,7 @@ def wrap_shm_features(obj):
""" """
Scan the object for multimodal tensors and wrap them in SHM pointers. Scan the object for multimodal tensors and wrap them in SHM pointers.
""" """
if _get_is_default_transport() or get_serving().skip_tokenizer_init: if _get_is_default_transport() or get_server_args().skip_tokenizer_init:
return obj return obj
if obj.mm_inputs: if obj.mm_inputs:
@@ -2034,7 +2028,7 @@ def unwrap_shm_features(obj):
Restore ShmPointerMMData wrappers back into standard torch.Tensors. Restore ShmPointerMMData wrappers back into standard torch.Tensors.
Handles both single requests and batch requests. Handles both single requests and batch requests.
""" """
if _get_is_default_transport() or get_serving().skip_tokenizer_init: if _get_is_default_transport() or get_server_args().skip_tokenizer_init:
return obj return obj
# Handle batch requests # Handle batch requests
if isinstance(obj, BaseBatchReq): if isinstance(obj, BaseBatchReq):
@@ -1,7 +1,5 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_disagg
# Copyright 2023-2024 SGLang Team # Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License. # you may not use this file except in compliance with the License.
@@ -647,15 +645,15 @@ class TokenizerWorker(TokenizerManager):
self.tokenizer_ipc_name = port_args.tokenizer_ipc_name self.tokenizer_ipc_name = port_args.tokenizer_ipc_name
# For PD disaggregtion # For PD disaggregtion
from sglang.srt.runtime_context import get_context self.server_args.override(
get_context().override(
"tokenizer_worker.restore_disaggregation_mode", "tokenizer_worker.restore_disaggregation_mode",
disaggregation_mode=disaggregation_mode, disaggregation_mode=disaggregation_mode,
) )
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode) self.disaggregation_mode = DisaggregationMode(
self.server_args.disaggregation_mode
)
self.disaggregation_transfer_backend = TransferBackend( self.disaggregation_transfer_backend = TransferBackend(
get_disagg().disaggregation_transfer_backend self.server_args.disaggregation_transfer_backend
) )
# Register this worker with the router for pause/continue broadcasting # Register this worker with the router for pause/continue broadcasting
+7 -9
View File
@@ -77,7 +77,10 @@ from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import ( from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import (
NewTokenRatioTracker, NewTokenRatioTracker,
) )
from sglang.srt.mem_cache.allocation import alloc_for_decode, alloc_for_extend from sglang.srt.mem_cache.allocation import (
alloc_for_decode,
alloc_for_extend,
)
from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
@@ -102,12 +105,7 @@ from sglang.srt.observability.req_time_stats import (
DPControllerReqTimeStats, DPControllerReqTimeStats,
SchedulerReqTimeStats, SchedulerReqTimeStats,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import get_parallel, get_server_args
get_parallel,
get_server_args,
get_serving,
get_spec,
)
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -1096,7 +1094,7 @@ class Req(ReqDllmMixin):
"""Check if this request is prefill-only (no token generation needed).""" """Check if this request is prefill-only (no token generation needed)."""
# NOTE: when spec is enabled, prefill_only optimizations are disabled # NOTE: when spec is enabled, prefill_only optimizations are disabled
spec_alg = get_spec().speculative_algorithm spec_alg = get_server_args().speculative_algorithm
return self.sampling_params.max_new_tokens == 0 and spec_alg is None return self.sampling_params.max_new_tokens == 0 and spec_alg is None
@property @property
@@ -1117,7 +1115,7 @@ class Req(ReqDllmMixin):
def effective_kv_committed_len(self) -> int: def effective_kv_committed_len(self) -> int:
# Report only the prompt prefix so thinking + answer fall into the # Report only the prompt prefix so thinking + answer fall into the
# overallocated range and are reclaimed by release_kv_cache. #22373. # overallocated range and are reclaimed by release_kv_cache. #22373.
if get_serving().strip_thinking_cache and self.reasoning_tokens > 0: if get_server_args().strip_thinking_cache and self.reasoning_tokens > 0:
return min(self.kv_committed_len, len(self.origin_input_ids)) return min(self.kv_committed_len, len(self.origin_input_ids))
return self.kv_committed_len return self.kv_committed_len
@@ -56,7 +56,7 @@ from sglang.srt.mem_cache.multi_ended_allocator import (
UnifiedMambaTokenToKVPoolAllocator, UnifiedMambaTokenToKVPoolAllocator,
) )
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
from sglang.srt.runtime_context import get_disagg from sglang.srt.runtime_context import get_server_args
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -193,7 +193,7 @@ class SchedulePolicy:
if ( if (
not isinstance(policy, CacheAwarePolicy) not isinstance(policy, CacheAwarePolicy)
and self.tree_cache.supports_fast_match_prefix() and self.tree_cache.supports_fast_match_prefix()
and get_disagg().disaggregation_mode != "decode" and get_server_args().disaggregation_mode != "decode"
): ):
for r in waiting_queue: for r in waiting_queue:
match_prefix_for_req(self.tree_cache, r, include_req=True) match_prefix_for_req(self.tree_cache, r, include_req=True)
+69 -82
View File
@@ -210,7 +210,9 @@ from sglang.srt.managers.scheduler_components.pool_stats_observer import (
from sglang.srt.managers.scheduler_components.profiler_manager import ( from sglang.srt.managers.scheduler_components.profiler_manager import (
SchedulerProfilerManager, SchedulerProfilerManager,
) )
from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper from sglang.srt.managers.scheduler_components.recv_skipper import (
SchedulerRecvSkipper,
)
from sglang.srt.managers.scheduler_components.request_receiver import ( from sglang.srt.managers.scheduler_components.request_receiver import (
SchedulerRequestReceiver, SchedulerRequestReceiver,
) )
@@ -239,20 +241,7 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.plugins import load_plugins from sglang.srt.plugins import load_plugins
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import get_context, get_parallel
get_context,
get_device,
get_disagg,
get_exec,
get_lora,
get_memory,
get_mm,
get_observability,
get_parallel,
get_schedule,
get_serving,
get_spec,
)
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.sampling.sampling_params import TOP_K_ALL
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs
@@ -454,9 +443,9 @@ class Scheduler(
attn_tp_cpu_group=self.attn_tp_cpu_group, attn_tp_cpu_group=self.attn_tp_cpu_group,
tp_cpu_group=self.tp_cpu_group, tp_cpu_group=self.tp_cpu_group,
attn_cp_cpu_group=self.attn_cp_cpu_group, attn_cp_cpu_group=self.attn_cp_cpu_group,
enable_metrics=get_observability().enable_metrics, enable_metrics=self.server_args.enable_metrics,
enable_kv_cache_events=bool( enable_kv_cache_events=bool(
get_observability().kv_events_config self.server_args.kv_events_config
and self.ps.pp_rank == 0 and self.ps.pp_rank == 0
and self.ps.attn_tp_rank == 0 and self.ps.attn_tp_rank == 0
and self.ps.attn_cp_rank == 0 and self.ps.attn_cp_rank == 0
@@ -482,8 +471,8 @@ class Scheduler(
self.init_hisparse_coordinator() self.init_hisparse_coordinator()
if ( if (
get_disagg().disaggregation_mode == "decode" self.server_args.disaggregation_mode == "decode"
and get_disagg().disaggregation_decode_enable_offload_kvcache and self.server_args.disaggregation_decode_enable_offload_kvcache
): ):
self.decode_offload_manager = DecodeKVCacheOffloadManager( self.decode_offload_manager = DecodeKVCacheOffloadManager(
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
@@ -594,7 +583,7 @@ class Scheduler(
self.dllm_config = ( # For diffusion LLM self.dllm_config = ( # For diffusion LLM
DllmConfig.from_server_args(self.server_args) DllmConfig.from_server_args(self.server_args)
if get_exec().dllm.dllm_algorithm is not None if self.server_args.dllm_algorithm is not None
else None else None
) )
@@ -622,11 +611,11 @@ class Scheduler(
self.ipc_channels = SchedulerIpcChannels.create( self.ipc_channels = SchedulerIpcChannels.create(
port_args=port_args, port_args=port_args,
is_rank_zero=is_rank_zero, is_rank_zero=is_rank_zero,
skip_tokenizer_init=get_serving().skip_tokenizer_init, skip_tokenizer_init=self.server_args.skip_tokenizer_init,
metrics_enabled=get_observability().enable_metrics metrics_enabled=self.server_args.enable_metrics
and ( and (
self.ps.attn_tp_rank == 0 self.ps.attn_tp_rank == 0
or get_observability().enable_metrics_for_all_schedulers or self.server_args.enable_metrics_for_all_schedulers
), ),
enable_scripted_runtime=envs.SGLANG_TEST_SCRIPTED_RUNTIME.get(), enable_scripted_runtime=envs.SGLANG_TEST_SCRIPTED_RUNTIME.get(),
) )
@@ -642,7 +631,7 @@ class Scheduler(
port_args, port_args,
self.ps.dp_size, self.ps.dp_size,
dp_rank, dp_rank,
publish_interval=get_observability().load_snapshot_publish_interval, publish_interval=self.server_args.load_snapshot_publish_interval,
) )
except Exception as e: except Exception as e:
logger.warning("load snapshot writer init failed: %s", e) logger.warning("load snapshot writer init failed: %s", e)
@@ -652,7 +641,7 @@ class Scheduler(
self.ps.pp_rank == 0 self.ps.pp_rank == 0
and self.ps.attn_tp_rank == 0 and self.ps.attn_tp_rank == 0
and self.ps.attn_cp_rank == 0 and self.ps.attn_cp_rank == 0
and get_device().sleep_on_idle and self.server_args.sleep_on_idle
): ):
self.idle_sleeper = IdleSleeper( self.idle_sleeper = IdleSleeper(
sockets=[ sockets=[
@@ -723,9 +712,9 @@ class Scheduler(
) )
# Set reasoning_parser and think_end_id if --reasoning_parser is enabled # Set reasoning_parser and think_end_id if --reasoning_parser is enabled
if get_serving().reasoning_parser and self.tokenizer: if self.server_args.reasoning_parser and self.tokenizer:
reasoning_parser = ReasoningParser( reasoning_parser = ReasoningParser(
model_type=get_serving().reasoning_parser, model_type=self.server_args.reasoning_parser,
stream_reasoning=False, stream_reasoning=False,
tokenizer=self.tokenizer, tokenizer=self.tokenizer,
) )
@@ -796,7 +785,7 @@ class Scheduler(
target_worker=self.tp_worker, target_worker=self.tp_worker,
) )
if get_spec().speculative_draft_load_format is not None: if self.server_args.speculative_draft_load_format is not None:
# Write the draft load_format onto server_args (not just the bag): # Write the draft load_format onto server_args (not just the bag):
# the draft worker is built from a copy of self.server_args and # the draft worker is built from a copy of self.server_args and
# build_load_config reads server_args.load_format, so a bag-only # build_load_config reads server_args.load_format, so a bag-only
@@ -804,10 +793,10 @@ class Scheduler(
# format. # format.
self.server_args.override( self.server_args.override(
"scheduler.draft_load_format", "scheduler.draft_load_format",
load_format=get_spec().speculative_draft_load_format, load_format=self.server_args.speculative_draft_load_format,
) )
logger.info( logger.info(
f"Using draft model load_format: '{get_spec().speculative_draft_load_format}'" f"Using draft model load_format: '{self.server_args.speculative_draft_load_format}'"
) )
DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args) DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args)
@@ -898,7 +887,7 @@ class Scheduler(
# --min-free-slots-delay. Built independently of the prefill delayer. # --min-free-slots-delay. Built independently of the prefill delayer.
self.min_free_slots_delayer: Optional[MinFreeSlotsDelayer] = None self.min_free_slots_delayer: Optional[MinFreeSlotsDelayer] = None
min_free_slots = resolve_min_free_slots( min_free_slots = resolve_min_free_slots(
get_schedule().min_free_slots_delay, self.server_args.min_free_slots_delay,
self.max_running_requests, self.max_running_requests,
is_dflash_family=self.spec_algorithm.is_dflash_family(), is_dflash_family=self.spec_algorithm.is_dflash_family(),
) )
@@ -944,14 +933,14 @@ class Scheduler(
if self.ps.tp_rank == 0: if self.ps.tp_rank == 0:
logger.info( logger.info(
f"max_total_num_tokens={self.max_total_num_tokens}, " f"max_total_num_tokens={self.max_total_num_tokens}, "
f"chunked_prefill_size={get_schedule().chunked_prefill_size}, " f"chunked_prefill_size={self.server_args.chunked_prefill_size}, "
f"max_prefill_tokens={self.max_prefill_tokens}, " f"max_prefill_tokens={self.max_prefill_tokens}, "
f"max_running_requests={self.max_running_requests}, " f"max_running_requests={self.max_running_requests}, "
f"context_len={self.model_config.context_len}, " f"context_len={self.model_config.context_len}, "
f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}={avail_mem:.2f} GB" f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}={avail_mem:.2f} GB"
) )
if get_observability().enable_metrics: if self.server_args.enable_metrics:
self.metrics_collector.emit_constants( self.metrics_collector.emit_constants(
max_total_num_tokens=self.max_total_num_tokens, max_total_num_tokens=self.max_total_num_tokens,
# TODO: max_running_requests_under_SLO has no setter — dead chain. # TODO: max_running_requests_under_SLO has no setter — dead chain.
@@ -998,7 +987,7 @@ class Scheduler(
self._engine_paused = False self._engine_paused = False
def init_chunked_prefill(self): def init_chunked_prefill(self):
self.chunked_prefill_size = get_schedule().chunked_prefill_size self.chunked_prefill_size = self.server_args.chunked_prefill_size
uses_transformers_backend = ( uses_transformers_backend = (
get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS
) )
@@ -1018,12 +1007,13 @@ class Scheduler(
self.chunked_req = None self.chunked_req = None
self._pending_chunked_abort_req = None self._pending_chunked_abort_req = None
self.is_mixed_chunk = ( self.is_mixed_chunk = (
self.chunked_prefill_size is not None and get_schedule().enable_mixed_chunk self.chunked_prefill_size is not None
and self.server_args.enable_mixed_chunk
) )
# Init the dynamic chunking predictor for PP # Init the dynamic chunking predictor for PP
self.enable_dynamic_chunking = ( self.enable_dynamic_chunking = (
get_schedule().enable_dynamic_chunking and self.ps.pp_size > 1 self.server_args.enable_dynamic_chunking and self.ps.pp_size > 1
) )
if self.enable_dynamic_chunking: if self.enable_dynamic_chunking:
try: try:
@@ -1059,8 +1049,8 @@ class Scheduler(
) )
self.prefill_delayer: Optional[PrefillDelayer] = None self.prefill_delayer: Optional[PrefillDelayer] = None
self.max_prefill_bs: int = 0 self.max_prefill_bs: int = 0
if get_schedule().enable_prefill_delayer: if self.server_args.enable_prefill_delayer:
if get_disagg().disaggregation_mode == "decode": if self.server_args.disaggregation_mode == "decode":
logger.info( logger.info(
"Ignoring --enable-prefill-delayer on decode engine " "Ignoring --enable-prefill-delayer on decode engine "
"(no prefill scheduling path; delayer would be a no-op)." "(no prefill scheduling path; delayer would be a no-op)."
@@ -1077,15 +1067,15 @@ class Scheduler(
if self.metrics_reporter.enable_metrics if self.metrics_reporter.enable_metrics
else None else None
), ),
max_delay_passes=get_schedule().prefill_delayer_max_delay_passes, max_delay_passes=self.server_args.prefill_delayer_max_delay_passes,
token_usage_low_watermark=get_schedule().prefill_delayer_token_usage_low_watermark, token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark,
device=self.tp_group.device, device=self.tp_group.device,
) )
# NOTE: preemption is enabled by default for priority scheduling. # NOTE: preemption is enabled by default for priority scheduling.
self.enable_priority_preemption = ( self.enable_priority_preemption = (
self.enable_priority_scheduling self.enable_priority_scheduling
and not get_schedule().disable_priority_preemption and not self.server_args.disable_priority_preemption
) )
self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args( self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args(
@@ -1101,12 +1091,12 @@ class Scheduler(
def init_watch_dog_memory_saver_input_blocker(self): def init_watch_dog_memory_saver_input_blocker(self):
# Start watchdog thread # Start watchdog thread
self.watchdog = create_scheduler_watchdog( self.watchdog = create_scheduler_watchdog(
self, watchdog_timeout=get_device().watchdog_timeout self, watchdog_timeout=self.server_args.watchdog_timeout
) )
# Init memory saver, profiler and metric stats # Init memory saver, profiler and metric stats
self.memory_saver_adapter = TorchMemorySaverAdapter.create( self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=get_exec().features.enable_memory_saver enable=self.server_args.enable_memory_saver
) )
# Init recv skipper and input blocker # Init recv skipper and input blocker
@@ -1128,9 +1118,11 @@ class Scheduler(
self.disagg_decode_prealloc_queue = None self.disagg_decode_prealloc_queue = None
self.disagg_decode_transfer_queue = None self.disagg_decode_transfer_queue = None
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode) self.disaggregation_mode = DisaggregationMode(
self.server_args.disaggregation_mode
)
self.transfer_backend = TransferBackend( self.transfer_backend = TransferBackend(
get_disagg().disaggregation_transfer_backend self.server_args.disaggregation_transfer_backend
) )
# todo: should we fix this when enabling mtp or it doesn't matter since we only enable mtp in decode node thus we don't transfer draft kvs between P and D? # todo: should we fix this when enabling mtp or it doesn't matter since we only enable mtp in decode node thus we don't transfer draft kvs between P and D?
@@ -1198,12 +1190,12 @@ class Scheduler(
gloo_group=self.attn_tp_cpu_group, gloo_group=self.attn_tp_cpu_group,
tp_rank=self.ps.tp_rank, tp_rank=self.ps.tp_rank,
tp_size=self.ps.tp_size, tp_size=self.ps.tp_size,
dp_size=get_parallel().dp_size, dp_size=self.server_args.dp_size,
gpu_id=self.ps.gpu_id, gpu_id=self.ps.gpu_id,
bootstrap_port=get_disagg().disaggregation_bootstrap_port, bootstrap_port=self.server_args.disaggregation_bootstrap_port,
max_total_num_tokens=self.max_total_num_tokens, max_total_num_tokens=self.max_total_num_tokens,
pp_rank=self.ps.pp_rank, pp_rank=self.ps.pp_rank,
num_reserved_decode_tokens=get_disagg().num_reserved_decode_tokens, num_reserved_decode_tokens=self.server_args.num_reserved_decode_tokens,
transfer_backend=self.transfer_backend, transfer_backend=self.transfer_backend,
) )
@@ -1229,7 +1221,7 @@ class Scheduler(
tp_rank=self.ps.tp_rank, tp_rank=self.ps.tp_rank,
tp_size=self.ps.tp_size, tp_size=self.ps.tp_size,
gpu_id=self.ps.gpu_id, gpu_id=self.ps.gpu_id,
bootstrap_port=get_disagg().disaggregation_bootstrap_port, bootstrap_port=self.server_args.disaggregation_bootstrap_port,
gloo_group=self.attn_tp_cpu_group, gloo_group=self.attn_tp_cpu_group,
max_total_num_tokens=self.max_total_num_tokens, max_total_num_tokens=self.max_total_num_tokens,
scheduler=self, scheduler=self,
@@ -1243,10 +1235,11 @@ class Scheduler(
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get() self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
# Init mm receiver for EPD disaggregation mode # Init mm receiver for EPD disaggregation mode
if get_disagg().language_only and get_disagg().encoder_transfer_backend in [ if (
"zmq_to_scheduler", self.server_args.language_only
"mooncake", and self.server_args.encoder_transfer_backend
]: in ["zmq_to_scheduler", "mooncake"]
):
self.mm_receiver = create_mm_receiver( self.mm_receiver = create_mm_receiver(
self.server_args, self.server_args,
dtype=self.model_config.dtype, dtype=self.model_config.dtype,
@@ -1327,7 +1320,7 @@ class Scheduler(
def init_deterministic_inference_config(self): def init_deterministic_inference_config(self):
"""Initialize deterministic inference configuration for different attention backends.""" """Initialize deterministic inference configuration for different attention backends."""
if not get_exec().deterministic.enable_deterministic_inference: if not self.server_args.enable_deterministic_inference:
self.truncation_align_size = None self.truncation_align_size = None
return return
@@ -1336,7 +1329,7 @@ class Scheduler(
"triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096), "triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096),
} }
env_var, default_size = backend_sizes.get( env_var, default_size = backend_sizes.get(
get_exec().kernel.attention_backend, (None, None) self.server_args.attention_backend, (None, None)
) )
self.truncation_align_size = ( self.truncation_align_size = (
get_int_env_var(env_var, default_size) if env_var else None get_int_env_var(env_var, default_size) if env_var else None
@@ -1732,10 +1725,10 @@ class Scheduler(
) )
def init_lora_drainer(self) -> None: def init_lora_drainer(self) -> None:
if get_lora().lora_drain_wait_threshold > 0.0: if self.server_args.lora_drain_wait_threshold > 0.0:
self.lora_drainer = LoRADrainer( self.lora_drainer = LoRADrainer(
get_lora().max_loras_per_batch, self.server_args.max_loras_per_batch,
get_lora().lora_drain_wait_threshold, self.server_args.lora_drain_wait_threshold,
) )
else: else:
self.lora_drainer = None self.lora_drainer = None
@@ -1837,7 +1830,7 @@ class Scheduler(
def init_kv_events_publisher(self) -> None: def init_kv_events_publisher(self) -> None:
self.kv_events_publisher = SchedulerKvEventsPublisher( self.kv_events_publisher = SchedulerKvEventsPublisher(
kv_events_config=get_observability().kv_events_config, kv_events_config=self.server_args.kv_events_config,
ps=self.ps, ps=self.ps,
attn_tp_rank=self.ps.attn_tp_rank, attn_tp_rank=self.ps.attn_tp_rank,
attn_cp_rank=self.ps.attn_cp_rank, attn_cp_rank=self.ps.attn_cp_rank,
@@ -2013,7 +2006,7 @@ class Scheduler(
return image_inputs return image_inputs
def _get_multimodal_inputs(self, mm_inputs_dict): def _get_multimodal_inputs(self, mm_inputs_dict):
if get_mm().enable_broadcast_mm_inputs_process: if self.server_args.enable_broadcast_mm_inputs_process:
return self._process_and_broadcast_mm_inputs(mm_inputs_dict) return self._process_and_broadcast_mm_inputs(mm_inputs_dict)
else: else:
return MultimodalInputs.from_processor_output(mm_inputs_dict) return MultimodalInputs.from_processor_output(mm_inputs_dict)
@@ -2060,7 +2053,7 @@ class Scheduler(
def _maybe_namespace_elastic_radix_cache(self, req: Req) -> None: def _maybe_namespace_elastic_radix_cache(self, req: Req) -> None:
if ( if (
get_exec().moe.elastic_ep_backend is None self.server_args.elastic_ep_backend is None
or self.disable_radix_cache or self.disable_radix_cache
or not self.tree_cache.is_tree_cache() or not self.tree_cache.is_tree_cache()
): ):
@@ -2106,7 +2099,8 @@ class Scheduler(
) )
# Radix-native sessions use only the top-level session_id. # Radix-native sessions use only the top-level session_id.
radix_native_session = ( radix_native_session = (
recv_req.session_id is not None and get_memory().enable_session_radix_cache recv_req.session_id is not None
and self.server_args.enable_session_radix_cache
) )
if session_id is None or radix_native_session: if session_id is None or radix_native_session:
@@ -2118,7 +2112,7 @@ class Scheduler(
if recv_req.bootstrap_port is None: if recv_req.bootstrap_port is None:
# Use default bootstrap port # Use default bootstrap port
recv_req.bootstrap_port = get_disagg().disaggregation_bootstrap_port recv_req.bootstrap_port = self.server_args.disaggregation_bootstrap_port
req = Req( req = Req(
recv_req.rid, recv_req.rid,
@@ -2271,7 +2265,7 @@ class Scheduler(
self._add_request_to_queue(req) self._add_request_to_queue(req)
return return
if req.return_sampling_mask and get_exec().kernel.sampling_backend == "ascend": if req.return_sampling_mask and self.server_args.sampling_backend == "ascend":
# The ascend backend samples from logits directly and never builds the # The ascend backend samples from logits directly and never builds the
# top-k/top-p support, so it cannot produce a sampling mask. # top-k/top-p support, so it cannot produce a sampling mask.
error_msg = ( error_msg = (
@@ -2320,7 +2314,7 @@ class Scheduler(
error_msg = validate_input_length( error_msg = validate_input_length(
req, req,
self.max_req_input_len, self.max_req_input_len,
get_serving().allow_auto_truncate, self.server_args.allow_auto_truncate,
) )
if error_msg: if error_msg:
req.set_finish_with_abort(error_msg) req.set_finish_with_abort(error_msg)
@@ -2598,7 +2592,7 @@ class Scheduler(
error_msg = validate_input_length( error_msg = validate_input_length(
req, req,
self.max_req_input_len, self.max_req_input_len,
get_serving().allow_auto_truncate, self.server_args.allow_auto_truncate,
) )
if error_msg: if error_msg:
self._add_request_to_queue(req) self._add_request_to_queue(req)
@@ -2810,7 +2804,7 @@ class Scheduler(
if ( if (
need_mlp_sync need_mlp_sync
and not self.spec_algorithm.is_none() and not self.spec_algorithm.is_none()
and not get_spec().speculative_skip_dp_mlp_sync and not self.server_args.speculative_skip_dp_mlp_sync
): ):
# NOTE: This branch makes sure prefill and decode batches will not be mixed when spec and dp-attn is enabled. # NOTE: This branch makes sure prefill and decode batches will not be mixed when spec and dp-attn is enabled.
# Before merging the new batch into running batch: # Before merging the new batch into running batch:
@@ -2884,7 +2878,7 @@ class Scheduler(
for req in ready_grammar_requests: for req in ready_grammar_requests:
self._add_request_to_queue(req) self._add_request_to_queue(req)
if self.enable_hierarchical_cache or get_memory().enable_flexkv: if self.enable_hierarchical_cache or self.server_args.enable_flexkv:
self.tree_cache.check_hicache_events() self.tree_cache.check_hicache_events()
if self.enable_priority_preemption or self.is_hybrid_swa: if self.enable_priority_preemption or self.is_hybrid_swa:
@@ -2951,7 +2945,7 @@ class Scheduler(
self.priority_scheduling_preemption_threshold, self.priority_scheduling_preemption_threshold,
max_prefill_bs=self.max_prefill_bs, max_prefill_bs=self.max_prefill_bs,
max_running_requests=self.max_running_requests, max_running_requests=self.max_running_requests,
prefill_max_requests=get_schedule().prefill_max_requests, prefill_max_requests=self.server_args.prefill_max_requests,
prefill_delayer_single_pass=prefill_delayer_single_pass, prefill_delayer_single_pass=prefill_delayer_single_pass,
dllm_config=self.dllm_config, dllm_config=self.dllm_config,
waiting_queue_len=len(self.waiting_queue), waiting_queue_len=len(self.waiting_queue),
@@ -3522,7 +3516,7 @@ class Scheduler(
def _maybe_report_active_ranks(self) -> None: def _maybe_report_active_ranks(self) -> None:
if not ( if not (
self.enable_dp_attention and get_exec().moe.elastic_ep_backend is not None self.enable_dp_attention and self.server_args.elastic_ep_backend is not None
): ):
return return
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
@@ -3798,7 +3792,7 @@ class Scheduler(
ok, msg = self.tree_cache.attach_storage_backend( ok, msg = self.tree_cache.attach_storage_backend(
storage_backend=recv_req.hicache_storage_backend, storage_backend=recv_req.hicache_storage_backend,
storage_backend_extra_config_json=recv_req.hicache_storage_backend_extra_config_json, storage_backend_extra_config_json=recv_req.hicache_storage_backend_extra_config_json,
served_model_name=get_serving().served_model_name, served_model_name=self.server_args.served_model_name,
hicache_storage_prefetch_policy=recv_req.hicache_storage_prefetch_policy, hicache_storage_prefetch_policy=recv_req.hicache_storage_prefetch_policy,
hicache_write_policy=recv_req.hicache_write_policy, hicache_write_policy=recv_req.hicache_write_policy,
) )
@@ -3918,7 +3912,7 @@ class Scheduler(
} }
ret["effective_max_running_requests_per_dp"] = self.max_running_requests ret["effective_max_running_requests_per_dp"] = self.max_running_requests
if get_exec().moe.elastic_ep_backend is not None: if self.server_args.elastic_ep_backend is not None:
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
ret["is_scaling_elastic_ep"] = ElasticEPStateManager.is_scaling() ret["is_scaling_elastic_ep"] = ElasticEPStateManager.is_scaling()
@@ -4313,7 +4307,7 @@ class Scheduler(
old_ep_size = ElasticEPStateManager.get_effective_ep_size() old_ep_size = ElasticEPStateManager.get_effective_ep_size()
new_ep_size = recv_req.new_ep_size new_ep_size = recv_req.new_ep_size
max_ep_size = get_parallel().max_ep_size or old_ep_size max_ep_size = self.server_args.max_ep_size or old_ep_size
logger.debug( logger.debug(
"[Elastic EP][scale] request received: new_ep_size=%d " "[Elastic EP][scale] request received: new_ep_size=%d "
@@ -4451,10 +4445,10 @@ class Scheduler(
return None return None
def close_session(self, recv_req: CloseSessionReqInput): def close_session(self, recv_req: CloseSessionReqInput):
if get_memory().enable_session_radix_cache: if self.server_args.enable_session_radix_cache:
self.tree_cache.release_radix_session(recv_req.session_id) self.tree_cache.release_radix_session(recv_req.session_id)
if recv_req.session_id in self.session_controller or not ( if recv_req.session_id in self.session_controller or not (
get_memory().enable_session_radix_cache self.server_args.enable_session_radix_cache
): ):
self.session_controller.close(recv_req) self.session_controller.close(recv_req)
@@ -4633,13 +4627,6 @@ def run_scheduler_process(
display_dp_rank=display_dp_rank, display_dp_rank=display_dp_rank,
display_moe_ep_rank=display_moe_ep_rank, display_moe_ep_rank=display_moe_ep_rank,
) )
# Publish the resolved config at scheduler process entry so the config
# namespaces (get_serving()/get_device()/get_exec()/...) are available to
# Scheduler.__init__ and its init_* helpers, which read them before the
# model worker's own publish. ModelRunner re-publishes idempotently.
from sglang.srt.runtime_context import publish
publish(server_args, role="scheduler")
parent_process = psutil.Process().parent() parent_process = psutil.Process().parent()
# Set up tracing # Set up tracing
@@ -2,7 +2,14 @@ from __future__ import annotations
import logging import logging
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union from typing import (
TYPE_CHECKING,
Callable,
List,
Optional,
Tuple,
Union,
)
import torch import torch
@@ -16,14 +23,11 @@ from sglang.srt.managers.schedule_batch import (
ScheduleBatch, ScheduleBatch,
mamba_lazy_spec_in_window, mamba_lazy_spec_in_window,
) )
from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache from sglang.srt.mem_cache.common import (
from sglang.srt.runtime_context import ( maybe_cache_unfinished_req,
get_disagg, release_kv_cache,
get_exec,
get_memory,
get_observability,
get_server_args,
) )
from sglang.srt.runtime_context import get_server_args
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
@@ -44,7 +48,10 @@ if TYPE_CHECKING:
SchedulerOutputStreamer, SchedulerOutputStreamer,
) )
from sglang.srt.managers.tp_worker import BaseTpWorker from sglang.srt.managers.tp_worker import BaseTpWorker
from sglang.srt.managers.utils import EmbeddingBatchResult, GenerationBatchResult from sglang.srt.managers.utils import (
EmbeddingBatchResult,
GenerationBatchResult,
)
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
@@ -77,7 +84,7 @@ class SchedulerBatchResultProcessor:
def process_batch_result_prebuilt(self, batch: ScheduleBatch): def process_batch_result_prebuilt(self, batch: ScheduleBatch):
assert self.disaggregation_mode == DisaggregationMode.DECODE assert self.disaggregation_mode == DisaggregationMode.DECODE
use_free_group = get_disagg().disaggregation_decode_enable_radix_cache use_free_group = self.server_args.disaggregation_decode_enable_radix_cache
if use_free_group: if use_free_group:
self.token_to_kv_pool_allocator.free_group_begin() self.token_to_kv_pool_allocator.free_group_begin()
for req in batch.reqs: for req in batch.reqs:
@@ -85,7 +92,7 @@ class SchedulerBatchResultProcessor:
req.update_finish_state() req.update_finish_state()
if req.finished(): if req.finished():
req.time_stats.set_quick_finish_time() req.time_stats.set_quick_finish_time()
if get_memory().enable_hisparse: if self.server_args.enable_hisparse:
self.hisparse_coordinator.request_finished(req) self.hisparse_coordinator.request_finished(req)
release_kv_cache(req, self.tree_cache) release_kv_cache(req, self.tree_cache)
@@ -236,7 +243,7 @@ class SchedulerBatchResultProcessor:
req.time_stats.set_completion_time() req.time_stats.set_completion_time()
elif not batch.decoding_reqs or req not in batch.decoding_reqs: elif not batch.decoding_reqs or req not in batch.decoding_reqs:
maybe_cache_unfinished_req(req, self.tree_cache) maybe_cache_unfinished_req(req, self.tree_cache)
if get_memory().enable_hisparse: if self.server_args.enable_hisparse:
self.hisparse_coordinator.admit_request_into_staging(req) self.hisparse_coordinator.admit_request_into_staging(req)
self._maybe_collect_customized_info(i, req, logits_output) self._maybe_collect_customized_info(i, req, logits_output)
@@ -749,7 +756,7 @@ class SchedulerBatchResultProcessor:
num_block_accept_tokens=result.num_block_accept_tokens, num_block_accept_tokens=result.num_block_accept_tokens,
num_cap_tokens=result.num_cap_tokens, num_cap_tokens=result.num_cap_tokens,
) )
if get_observability().enable_metrics: if self.server_args.enable_metrics:
self.metrics_collector.increment_decode_cuda_graph_pass( self.metrics_collector.increment_decode_cuda_graph_pass(
value=can_run_cuda_graph value=can_run_cuda_graph
) )
@@ -932,7 +939,7 @@ class SchedulerBatchResultProcessor:
self._mamba_prefix_cache_update(req, batch, result, i) self._mamba_prefix_cache_update(req, batch, result, i)
if ( if (
get_disagg().disaggregation_decode_enable_offload_kvcache self.server_args.disaggregation_decode_enable_offload_kvcache
and not req.finished() and not req.finished()
): ):
self.decode_offload_manager.offload_kv_cache(req) self.decode_offload_manager.offload_kv_cache(req)
@@ -952,12 +959,12 @@ class SchedulerBatchResultProcessor:
self._maybe_collect_routed_experts(req) self._maybe_collect_routed_experts(req)
self._maybe_collect_indexer_topk(req) self._maybe_collect_indexer_topk(req)
if get_disagg().disaggregation_decode_enable_offload_kvcache: if self.server_args.disaggregation_decode_enable_offload_kvcache:
# Asynchronously offload KV cache; release_kv_cache will be called after Device->Host transfer completes # Asynchronously offload KV cache; release_kv_cache will be called after Device->Host transfer completes
if not self.decode_offload_manager.offload_kv_cache(req): if not self.decode_offload_manager.offload_kv_cache(req):
self.decode_offload_manager.finalize_release_on_finish(req) self.decode_offload_manager.finalize_release_on_finish(req)
else: else:
if get_memory().enable_hisparse: if self.server_args.enable_hisparse:
self.hisparse_coordinator.request_finished(req) self.hisparse_coordinator.request_finished(req)
prepare_release = getattr( prepare_release = getattr(
self.model_worker, "prepare_for_kv_cache_release", None self.model_worker, "prepare_for_kv_cache_release", None
@@ -1095,7 +1102,7 @@ class SchedulerBatchResultProcessor:
For spec decode, the boundary is detected by comparing the For spec decode, the boundary is detected by comparing the
accepted seq_len range against interval boundaries. accepted seq_len range against interval boundaries.
""" """
interval = get_exec().mamba.mamba_track_interval interval = get_server_args().mamba_track_interval
if batch.spec_algorithm.is_none(): if batch.spec_algorithm.is_none():
if req.kv_committed_len % interval == 0: if req.kv_committed_len % interval == 0:
@@ -12,7 +12,9 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import world_dp_gather_enabled from sglang.srt.layers.dp_attention import world_dp_gather_enabled
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper from sglang.srt.managers.scheduler_components.recv_skipper import (
SchedulerRecvSkipper,
)
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
@@ -24,7 +26,6 @@ from sglang.srt.model_executor.cuda_graph_config import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.metrics_collector import DPCooperationInfo from sglang.srt.observability.metrics_collector import DPCooperationInfo
from sglang.srt.runtime_context import get_parallel, get_schedule
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.common import require_mlp_tp_gather from sglang.srt.utils.common import require_mlp_tp_gather
@@ -377,14 +378,14 @@ class SchedulerDPAttnAdapter:
def prepare_mlp_sync_batch(self, local_batch: ScheduleBatch): def prepare_mlp_sync_batch(self, local_batch: ScheduleBatch):
return prepare_mlp_sync_batch_raw( return prepare_mlp_sync_batch_raw(
local_batch, local_batch,
dp_size=get_parallel().dp_size, dp_size=self.server_args.dp_size,
attn_tp_size=self.ps.attn_tp_size, attn_tp_size=self.ps.attn_tp_size,
attn_cp_size=self.ps.attn_cp_size, attn_cp_size=self.ps.attn_cp_size,
tp_group=self.tp_group, tp_group=self.tp_group,
get_idle_batch=self.get_idle_batch, get_idle_batch=self.get_idle_batch,
disable_cuda_graph=cuda_graph_fully_disabled(), disable_cuda_graph=cuda_graph_fully_disabled(),
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args), require_mlp_tp_gather=require_mlp_tp_gather(self.server_args),
disable_overlap_schedule=get_schedule().disable_overlap_schedule, disable_overlap_schedule=self.server_args.disable_overlap_schedule,
offload_tags=self.offload_tags, offload_tags=self.offload_tags,
dwdp=self.server_args.dwdp_size > 1, dwdp=self.server_args.dwdp_size > 1,
) )
@@ -14,7 +14,6 @@ from sglang.srt.managers.load_snapshot import (
QueueMetrics, QueueMetrics,
SpeculativeMetrics, SpeculativeMetrics,
) )
from sglang.srt.runtime_context import get_lora
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.distributed.parallel_state_wrapper import ParallelState
@@ -145,7 +144,7 @@ class SchedulerLoadInquirer:
) )
lora = None lora = None
if get_lora().enable_lora: if self.server_args.enable_lora:
lora = LoRAMetrics( lora = LoRAMetrics(
slots_used=stats.lora_pool_slots_used, slots_used=stats.lora_pool_slots_used,
slots_total=stats.lora_pool_slots_total, slots_total=stats.lora_pool_slots_total,
@@ -1,15 +1,20 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from typing import List, Tuple from typing import (
List,
Tuple,
)
import torch import torch
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.runtime_context import get_exec from sglang.srt.server_args import (
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID, ServerArgs MIS_DELIMITER_TOKEN_ID,
ServerArgs,
)
@dataclass(kw_only=True, slots=True, frozen=True) @dataclass(kw_only=True, slots=True, frozen=True)
@@ -159,7 +164,7 @@ class SchedulerLogprobResultProcessor:
delimiter token receive logprobs. delimiter token receive logprobs.
""" """
return ( return (
get_exec().features.enable_mis self.server_args.enable_mis
and req.is_prefill_only and req.is_prefill_only
and req.multi_item_delimiter_indices is not None and req.multi_item_delimiter_indices is not None
) )
@@ -2,7 +2,12 @@ from __future__ import annotations
import logging import logging
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Callable, List, Optional from typing import (
Any,
Callable,
List,
Optional,
)
import torch import torch
import zmq import zmq
@@ -16,9 +21,11 @@ from sglang.srt.managers.io_struct import (
CachedTokensDetails, CachedTokensDetails,
wrap_as_pickle, wrap_as_pickle,
) )
from sglang.srt.managers.schedule_batch import BaseFinishReason, Req from sglang.srt.managers.schedule_batch import (
BaseFinishReason,
Req,
)
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.runtime_context import get_observability, get_serving
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -137,7 +144,7 @@ class SchedulerOutputStreamer:
return_sampling_mask=return_sampling_mask, return_sampling_mask=return_sampling_mask,
spec_algorithm=self.spec_algorithm, spec_algorithm=self.spec_algorithm,
disaggregation_mode=self.disaggregation_mode, disaggregation_mode=self.disaggregation_mode,
default_stream_interval=get_serving().stream_interval, default_stream_interval=self.server_args.stream_interval,
default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL, default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL,
get_cached_tokens_details=self.get_cached_tokens_details, get_cached_tokens_details=self.get_cached_tokens_details,
) )
@@ -164,7 +171,7 @@ class SchedulerOutputStreamer:
if ( if (
req.finished() req.finished()
and self.ps.attn_tp_rank == 0 and self.ps.attn_tp_rank == 0
and get_observability().enable_request_time_stats_logging and self.server_args.enable_request_time_stats_logging
): ):
req.log_time_stats() req.log_time_stats()
@@ -5,7 +5,13 @@ import os
import time import time
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any, Callable, List, Optional from typing import (
TYPE_CHECKING,
Any,
Callable,
List,
Optional,
)
import torch import torch
@@ -13,7 +19,7 @@ from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_device from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import is_mps, is_npu from sglang.srt.utils import is_mps, is_npu
from sglang.srt.utils.profile_merger import ProfileMerger from sglang.srt.utils.profile_merger import ProfileMerger
from sglang.srt.utils.profile_utils import ProfileManager from sglang.srt.utils.profile_utils import ProfileManager
@@ -249,7 +255,7 @@ class SchedulerProfilerManager:
self.profile_in_progress = True self.profile_in_progress = True
if "CUDA_PROFILER" in activities: if "CUDA_PROFILER" in activities:
if self.ps.gpu_id == get_device().base_gpu_id: if self.ps.gpu_id == get_server_args().base_gpu_id:
torch.cuda.cudart().cudaProfilerStart() torch.cuda.cudart().cudaProfilerStart()
self.profile_in_progress = True self.profile_in_progress = True
@@ -359,7 +365,7 @@ class SchedulerProfilerManager:
torch.cuda.memory._record_memory_history(enabled=None) torch.cuda.memory._record_memory_history(enabled=None)
if "CUDA_PROFILER" in self.profiler_activities: if "CUDA_PROFILER" in self.profiler_activities:
if self.ps.gpu_id == get_device().base_gpu_id: if self.ps.gpu_id == get_server_args().base_gpu_id:
torch.cuda.cudart().cudaProfilerStop() torch.cuda.cudart().cudaProfilerStop()
merge_message = self._merge_profile_traces() merge_message = self._merge_profile_traces()
@@ -2,7 +2,14 @@ from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, Any, Callable, List, Optional, Union from typing import (
TYPE_CHECKING,
Any,
Callable,
List,
Optional,
Union,
)
import zmq import zmq
from torch.distributed import barrier from torch.distributed import barrier
@@ -15,9 +22,14 @@ from sglang.srt.managers.io_struct import (
TokenizedGenerateReqInput, TokenizedGenerateReqInput,
sock_recv, sock_recv,
) )
from sglang.srt.managers.mm_utils import has_shm_features, unwrap_shm_features from sglang.srt.managers.mm_utils import (
from sglang.srt.runtime_context import get_disagg, get_parallel has_shm_features,
from sglang.srt.utils import broadcast_pyobj, point_to_point_pyobj unwrap_shm_features,
)
from sglang.srt.utils import (
broadcast_pyobj,
point_to_point_pyobj,
)
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -127,7 +139,7 @@ class SchedulerRequestReceiver:
return recv_reqs return recv_reqs
def _broadcast_reqs_across_ranks(self, recv_reqs: Optional[List]) -> List: def _broadcast_reqs_across_ranks(self, recv_reqs: Optional[List]) -> List:
if get_parallel().enable_dp_attention: if self.server_args.enable_dp_attention:
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0: 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) work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs)
else: else:
@@ -156,7 +168,7 @@ class SchedulerRequestReceiver:
# instead of the full tp_group. This avoids an expensive # instead of the full tp_group. This avoids an expensive
# all-ranks gloo sync. # all-ranks gloo sync.
_local_ctrl = ( _local_ctrl = (
get_parallel().enable_dp_attention_local_control_broadcast self.server_args.enable_dp_attention_local_control_broadcast
or self.server_args.is_ep_scale_joiner or self.server_args.is_ep_scale_joiner
) )
if _local_ctrl: if _local_ctrl:
@@ -208,8 +220,8 @@ class SchedulerRequestReceiver:
# Process MM requests under EPD-disaggregation mode # Process MM requests under EPD-disaggregation mode
if ( if (
self.ps.pp_rank == 0 self.ps.pp_rank == 0
and get_disagg().language_only and self.server_args.language_only
and get_disagg().encoder_transfer_backend and self.server_args.encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"] in ["zmq_to_scheduler", "mooncake"]
): ):
recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs) recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
@@ -233,7 +245,7 @@ class SchedulerRequestReceiver:
# peer ranks may still be unpickling ShmPointerMMData # peer ranks may still be unpickling ShmPointerMMData
# (-> shm_open). Synchronize the same CPU groups that carried # (-> shm_open). Synchronize the same CPU groups that carried
# SHM-backed work requests before materialize() unlinks them. # SHM-backed work requests before materialize() unlinks them.
if get_parallel().enable_dp_attention: if self.server_args.enable_dp_attention:
if self.ps.attn_tp_size > 1: if self.ps.attn_tp_size > 1:
barrier(group=self.attn_tp_cpu_group) barrier(group=self.attn_tp_cpu_group)
if self.ps.attn_cp_size > 1: if self.ps.attn_cp_size > 1:
@@ -36,7 +36,6 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.observability.req_time_stats import set_time_batch from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.runtime_context import get_disagg, get_parallel
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
from sglang.srt.utils.common import get_device_module, is_xpu from sglang.srt.utils.common import get_device_module, is_xpu
@@ -123,7 +122,7 @@ class SchedulerPPMixin:
next_pp_outputs = None next_pp_outputs = None
next_batch_result = None next_batch_result = None
d2h_event = None d2h_event = None
if get_parallel().pp_async_batch_depth > 0: if self.server_args.pp_async_batch_depth > 0:
next_pp_outputs, next_batch_result, d2h_event = ( next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors( self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id, next_first_rank_mb_id,
@@ -139,7 +138,7 @@ class SchedulerPPMixin:
self.mb_metadata, self.mb_metadata,
self.last_rank_comm_queue, self.last_rank_comm_queue,
) )
if get_parallel().pp_async_batch_depth == 0: if self.server_args.pp_async_batch_depth == 0:
next_pp_outputs, next_batch_result, d2h_event = ( next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors( self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id, next_first_rank_mb_id,
@@ -269,7 +268,7 @@ class SchedulerPPMixin:
server_is_idle = False server_is_idle = False
pp_proxy_tensors = self._pp_recv_proxy_tensors() pp_proxy_tensors = self._pp_recv_proxy_tensors()
if get_parallel().pp_async_batch_depth > 0: if self.server_args.pp_async_batch_depth > 0:
next_pp_outputs, next_batch_result, d2h_event = ( next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors( self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id, next_first_rank_mb_id,
@@ -285,7 +284,7 @@ class SchedulerPPMixin:
self.mb_metadata, self.mb_metadata,
self.last_rank_comm_queue, self.last_rank_comm_queue,
) )
if get_parallel().pp_async_batch_depth == 0: if self.server_args.pp_async_batch_depth == 0:
next_pp_outputs, next_batch_result, d2h_event = ( next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors( self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id, next_first_rank_mb_id,
@@ -428,7 +427,7 @@ class SchedulerPPMixin:
pp_proxy_tensors = self._pp_recv_proxy_tensors() pp_proxy_tensors = self._pp_recv_proxy_tensors()
# early send output if possible # early send output if possible
if get_parallel().pp_async_batch_depth > 0: if self.server_args.pp_async_batch_depth > 0:
next_pp_outputs, next_batch_result, d2h_event = ( next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors( self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id, next_first_rank_mb_id,
@@ -446,7 +445,7 @@ class SchedulerPPMixin:
self.last_rank_comm_queue, self.last_rank_comm_queue,
) )
if get_parallel().pp_async_batch_depth == 0: if self.server_args.pp_async_batch_depth == 0:
next_pp_outputs, next_batch_result, d2h_event = ( next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors( self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id, next_first_rank_mb_id,
@@ -480,7 +479,7 @@ class SchedulerPPMixin:
) )
) )
if get_disagg().disaggregation_decode_enable_offload_kvcache: if self.server_args.disaggregation_decode_enable_offload_kvcache:
self.decode_offload_manager.check_offload_progress() self.decode_offload_manager.check_offload_progress()
if rmbs[next_mb_id] is not None: if rmbs[next_mb_id] is not None:
@@ -550,17 +549,17 @@ class SchedulerPPMixin:
+ len(self.disagg_decode_transfer_queue.queue) + len(self.disagg_decode_transfer_queue.queue)
+ len(self.disagg_decode_prealloc_queue.queue) + len(self.disagg_decode_prealloc_queue.queue)
) )
if get_disagg().disaggregation_decode_enable_offload_kvcache: if self.server_args.disaggregation_decode_enable_offload_kvcache:
queue_size += len(self.decode_offload_manager.ongoing_offload) queue_size += len(self.decode_offload_manager.ongoing_offload)
if server_is_idle and queue_size == 0: if server_is_idle and queue_size == 0:
self.on_idle() self.on_idle()
def init_pp_loop_state(self: Scheduler): def init_pp_loop_state(self: Scheduler):
self.pp_loop_size: int = self.ps.pp_size + get_parallel().pp_async_batch_depth self.pp_loop_size: int = self.ps.pp_size + self.server_args.pp_async_batch_depth
# In CP mode, attention weights are duplicated, eliminating the need for the attention TP all-gather operation. # In CP mode, attention weights are duplicated, eliminating the need for the attention TP all-gather operation.
self.require_attn_tp_allgather = ( self.require_attn_tp_allgather = (
not get_parallel().enable_dsa_prefill_context_parallel not self.server_args.enable_dsa_prefill_context_parallel
) )
self.mbs = [None] * self.pp_loop_size self.mbs = [None] * self.pp_loop_size
self.last_mbs = [None] * self.pp_loop_size self.last_mbs = [None] * self.pp_loop_size
@@ -74,7 +74,6 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromTensorReqOutput, UpdateWeightsFromTensorReqOutput,
) )
from sglang.srt.managers.load_snapshot import LoadSnapshot from sglang.srt.managers.load_snapshot import LoadSnapshot
from sglang.srt.runtime_context import get_lora, get_parallel
from sglang.srt.server_args import LoRARef, ServerArgs from sglang.srt.server_args import LoRARef, ServerArgs
from sglang.srt.utils import ( from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
@@ -146,8 +145,8 @@ class TokenizerControlMixin:
def update_control_communicator_fan_out(self: TokenizerManager, worker_count: int): def update_control_communicator_fan_out(self: TokenizerManager, worker_count: int):
primary_group_control = ( primary_group_control = (
get_parallel().enable_dp_attention self.server_args.enable_dp_attention
and not get_parallel().enable_dp_attention_local_control_broadcast and not self.server_args.enable_dp_attention_local_control_broadcast
) )
if primary_group_control: if primary_group_control:
control_fan_out = ( control_fan_out = (
@@ -397,7 +396,7 @@ class TokenizerControlMixin:
) -> Tuple[bool, str]: ) -> Tuple[bool, str]:
self.auto_create_handle_loop() self.auto_create_handle_loop()
assert ( assert (
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for update weights from distributed" ), "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) results = await self.init_weights_update_group_communicator(obj)
@@ -410,7 +409,7 @@ class TokenizerControlMixin:
) -> Tuple[bool, str]: ) -> Tuple[bool, str]:
self.auto_create_handle_loop() self.auto_create_handle_loop()
assert ( assert (
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for destroy parameter update group" ), "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) results = await self.destroy_weights_update_group_communicator(obj)
@@ -423,7 +422,7 @@ class TokenizerControlMixin:
) -> Tuple[bool, str]: ) -> Tuple[bool, str]:
self.auto_create_handle_loop() self.auto_create_handle_loop()
assert ( assert (
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for update weights from distributed" ), "dp_size must be 1 or dp attention must be enabled for update weights from distributed"
if obj.abort_all_requests: if obj.abort_all_requests:
@@ -454,7 +453,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop() self.auto_create_handle_loop()
# TODO: support DP # TODO: support DP
assert ( assert (
get_parallel().dp_size == 1 self.server_args.dp_size == 1
), "dp_size must be 1 for init_weights_send_group_for_remote_instance" ), "dp_size must be 1 for init_weights_send_group_for_remote_instance"
result = ( result = (
await self.init_weights_send_group_for_remote_instance_communicator(obj) await self.init_weights_send_group_for_remote_instance_communicator(obj)
@@ -469,7 +468,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop() self.auto_create_handle_loop()
# TODO: support DP # TODO: support DP
assert ( assert (
get_parallel().dp_size == 1 self.server_args.dp_size == 1
), "dp_size must be 1 for send_weights_to_remote_instance" ), "dp_size must be 1 for send_weights_to_remote_instance"
result = (await self.send_weights_to_remote_instance_communicator(obj))[0] result = (await self.send_weights_to_remote_instance_communicator(obj))[0]
return result.success, result.message return result.success, result.message
@@ -481,7 +480,7 @@ class TokenizerControlMixin:
) -> Tuple[bool, str]: ) -> Tuple[bool, str]:
self.auto_create_handle_loop() self.auto_create_handle_loop()
assert ( assert (
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for update weights from tensor" ), "dp_size must be 1 or dp attention must be enabled for update weights from tensor"
if obj.abort_all_requests: if obj.abort_all_requests:
@@ -517,7 +516,7 @@ class TokenizerControlMixin:
try: try:
# For now, we only support single data parallel instance # For now, we only support single data parallel instance
assert ( assert (
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for update weights from IPC" ), "dp_size must be 1 or dp attention must be enabled for update weights from IPC"
logger.info("Starting IPC weight update") logger.info("Starting IPC weight update")
@@ -570,7 +569,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop() self.auto_create_handle_loop()
try: try:
if not get_lora().enable_lora: if not self.server_args.enable_lora:
raise ValueError( raise ValueError(
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA." "LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
) )
@@ -578,7 +577,7 @@ class TokenizerControlMixin:
# TODO (lifuhuang): Remove this after we verify that dynamic lora loading works # TODO (lifuhuang): Remove this after we verify that dynamic lora loading works
# with dp_size > 1. # with dp_size > 1.
assert ( assert (
get_parallel().dp_size == 1 self.server_args.dp_size == 1
), "dp_size must be 1 for dynamic lora loading" ), "dp_size must be 1 for dynamic lora loading"
logger.info( logger.info(
"Start load Lora adapter. Lora name=%s, path=%s", "Start load Lora adapter. Lora name=%s, path=%s",
@@ -603,10 +602,10 @@ class TokenizerControlMixin:
await self.lora_registry.register(new_adapter) await self.lora_registry.register(new_adapter)
self.lora_ref_cache[obj.lora_name] = new_adapter self.lora_ref_cache[obj.lora_name] = new_adapter
if get_lora().max_loaded_loras is not None: if self.server_args.max_loaded_loras is not None:
while ( while (
self.lora_registry.num_registered_loras self.lora_registry.num_registered_loras
> get_lora().max_loaded_loras > self.server_args.max_loaded_loras
): ):
lru_lora_name = await self.lora_registry.lru_lora_name( lru_lora_name = await self.lora_registry.lru_lora_name(
exclude_pinned=True exclude_pinned=True
@@ -620,7 +619,7 @@ class TokenizerControlMixin:
logger.info( logger.info(
f"Unloading least recently used LoRA adapter '{lru_lora_name}' " f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
f"(current number of adapters: {self.lora_registry.num_registered_loras}, " f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
f"max allowed: {get_lora().max_loaded_loras})" f"max allowed: {self.server_args.max_loaded_loras})"
) )
unload_result = await self._unload_lora_adapter_locked( unload_result = await self._unload_lora_adapter_locked(
@@ -648,13 +647,13 @@ class TokenizerControlMixin:
self.auto_create_handle_loop() self.auto_create_handle_loop()
try: try:
if not get_lora().enable_lora: if not self.server_args.enable_lora:
raise ValueError( raise ValueError(
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA." "LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
) )
assert ( assert (
get_parallel().dp_size == 1 self.server_args.dp_size == 1
), "dp_size must be 1 for dynamic lora loading" ), "dp_size must be 1 for dynamic lora loading"
logger.info( logger.info(
"Start load Lora adapter from tensors. Lora name=%s", "Start load Lora adapter from tensors. Lora name=%s",
@@ -673,10 +672,10 @@ class TokenizerControlMixin:
if result.success: if result.success:
await self.lora_registry.register(new_adapter) await self.lora_registry.register(new_adapter)
self.lora_ref_cache[obj.lora_name] = new_adapter self.lora_ref_cache[obj.lora_name] = new_adapter
if get_lora().max_loaded_loras is not None: if self.server_args.max_loaded_loras is not None:
while ( while (
self.lora_registry.num_registered_loras self.lora_registry.num_registered_loras
> get_lora().max_loaded_loras > self.server_args.max_loaded_loras
): ):
lru_lora_name = await self.lora_registry.lru_lora_name( lru_lora_name = await self.lora_registry.lru_lora_name(
exclude_pinned=True exclude_pinned=True
@@ -690,7 +689,7 @@ class TokenizerControlMixin:
logger.info( logger.info(
f"Unloading least recently used LoRA adapter '{lru_lora_name}' " f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
f"(current number of adapters: {self.lora_registry.num_registered_loras}, " f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
f"max allowed: {get_lora().max_loaded_loras})" f"max allowed: {self.server_args.max_loaded_loras})"
) )
unload_result = await self._unload_lora_adapter_locked( unload_result = await self._unload_lora_adapter_locked(
@@ -718,7 +717,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop() self.auto_create_handle_loop()
try: try:
if not get_lora().enable_lora: if not self.server_args.enable_lora:
raise ValueError( raise ValueError(
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA." "LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
) )
@@ -730,7 +729,7 @@ class TokenizerControlMixin:
# TODO (lifuhuang): Remove this after we verify that dynamic lora loading works # TODO (lifuhuang): Remove this after we verify that dynamic lora loading works
# with dp_size > 1. # with dp_size > 1.
assert ( assert (
get_parallel().dp_size == 1 self.server_args.dp_size == 1
), "dp_size must be 1 for dynamic lora loading" ), "dp_size must be 1 for dynamic lora loading"
logger.info( logger.info(
"Start unload Lora adapter. Lora name=%s", "Start unload Lora adapter. Lora name=%s",
@@ -750,7 +749,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop() self.auto_create_handle_loop()
results = await self.get_weights_by_name_communicator(obj) results = await self.get_weights_by_name_communicator(obj)
all_parameters = [r.parameter for r in results] all_parameters = [r.parameter for r in results]
if get_parallel().dp_size == 1: if self.server_args.dp_size == 1:
return all_parameters[0] return all_parameters[0]
else: else:
return all_parameters return all_parameters
@@ -894,8 +893,6 @@ class TokenizerControlMixin:
) -> None: ) -> None:
"""Update weight version if provided.""" """Update weight version if provided."""
if weight_version is not None: if weight_version is not None:
from sglang.srt.runtime_context import get_context self.server_args.override(
get_context().override(
"tokenizer.weight_version", weight_version=weight_version "tokenizer.weight_version", weight_version=weight_version
) )
+34 -43
View File
@@ -110,15 +110,6 @@ from sglang.srt.observability.request_metrics_exporter import (
RequestMetricsExporterManager, RequestMetricsExporterManager,
) )
from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_lora,
get_model,
get_observability,
get_parallel,
get_serving,
)
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import ( from sglang.srt.server_args import (
PortArgs, PortArgs,
@@ -472,10 +463,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# TODO: Refactor and organize the log export code. # TODO: Refactor and organize the log export code.
# Request logging # Request logging
self.request_logger = RequestLogger( self.request_logger = RequestLogger(
log_requests=get_observability().log_requests, log_requests=self.server_args.log_requests,
log_requests_level=get_observability().log_requests_level, log_requests_level=self.server_args.log_requests_level,
log_requests_format=get_observability().log_requests_format, log_requests_format=self.server_args.log_requests_format,
log_requests_target=get_observability().log_requests_target, log_requests_target=self.server_args.log_requests_target,
) )
# Dumping # Dumping
@@ -498,7 +489,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def init_weight_update(self): def init_weight_update(self):
# Initial weights status # Initial weights status
self.initial_weights_loaded = True self.initial_weights_loaded = True
if get_model().checkpoint_engine_wait_weights_before_ready: if self.server_args.checkpoint_engine_wait_weights_before_ready:
self.initial_weights_loaded = False self.initial_weights_loaded = False
# Weight updates # Weight updates
@@ -518,7 +509,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# The registry dynamically updates as adapters are loaded / unloaded during runtime. It # The registry dynamically updates as adapters are loaded / unloaded during runtime. It
# serves as the source of truth for available adapters and maps user-friendly LoRA names # serves as the source of truth for available adapters and maps user-friendly LoRA names
# to internally used unique LoRA IDs. # to internally used unique LoRA IDs.
self.lora_registry = LoRARegistry(get_lora().lora_paths) self.lora_registry = LoRARegistry(self.server_args.lora_paths)
# Lock to serialize LoRA update operations. # Lock to serialize LoRA update operations.
# Please note that, unlike `model_update_lock`, this does not block inference, allowing # Please note that, unlike `model_update_lock`, this does not block inference, allowing
# LoRA updates and inference to overlap. # LoRA updates and inference to overlap.
@@ -527,13 +518,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# point to their latest LoRARef objects, so that they can be # point to their latest LoRARef objects, so that they can be
# dynamically loaded if needed for inference # dynamically loaded if needed for inference
self.lora_ref_cache: Dict[str, LoRARef] = {} self.lora_ref_cache: Dict[str, LoRARef] = {}
if get_lora().lora_paths is not None: if self.server_args.lora_paths is not None:
for lora_ref in get_lora().lora_paths: for lora_ref in self.server_args.lora_paths:
self.lora_ref_cache[lora_ref.lora_name] = lora_ref self.lora_ref_cache[lora_ref.lora_name] = lora_ref
def init_disaggregation(self): def init_disaggregation(self):
# PD Disaggregation # PD Disaggregation
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode) self.disaggregation_mode = DisaggregationMode(
self.server_args.disaggregation_mode
)
# Keep a reference so the bootstrap server is not garbage-collected. # Keep a reference so the bootstrap server is not garbage-collected.
self.bootstrap_server = start_disagg_service(self.server_args) self.bootstrap_server = start_disagg_service(self.server_args)
# Single-source counter for auto-assigning fake bootstrap_room. # Single-source counter for auto-assigning fake bootstrap_room.
@@ -542,16 +535,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Encoder Disaggregation # Encoder Disaggregation
self.encoder_bootstrap_server = None self.encoder_bootstrap_server = None
if self.server_args.language_only: if self.server_args.language_only:
from sglang.srt.disaggregation.encode_receiver import EncoderBootstrapServer from sglang.srt.disaggregation.encode_receiver import (
EncoderBootstrapServer,
)
# Shared mutable URL list: the bootstrap server appends / removes # Shared mutable URL list: the bootstrap server appends / removes
# entries as encoders register, the receiver reads from the same # entries as encoders register, the receiver reads from the same
# list. Pre-populated with static --encoder-urls so the legacy # list. Pre-populated with static --encoder-urls so the legacy
# CLI flag still works (alongside dynamic registrations). # CLI flag still works (alongside dynamic registrations).
self.encoder_urls: List[str] = list(get_disagg().encoder_urls) self.encoder_urls: List[str] = list(self.server_args.encoder_urls)
self.encoder_bootstrap_server = EncoderBootstrapServer( self.encoder_bootstrap_server = EncoderBootstrapServer(
host=get_serving().host, host=self.server_args.host,
port=get_disagg().encoder_bootstrap_port, port=self.server_args.encoder_bootstrap_port,
urls=self.encoder_urls, urls=self.encoder_urls,
) )
self.mm_receiver = create_mm_receiver( self.mm_receiver = create_mm_receiver(
@@ -565,22 +560,20 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Metrics # Metrics
if self.enable_metrics: if self.enable_metrics:
engine_type = DisaggregationMode.to_engine_type( engine_type = DisaggregationMode.to_engine_type(
get_disagg().disaggregation_mode self.server_args.disaggregation_mode
) )
labels = { labels = {
"model_name": get_serving().served_model_name, "model_name": self.server_args.served_model_name,
"engine_type": engine_type, "engine_type": engine_type,
} }
if self.enable_priority_scheduling: if self.enable_priority_scheduling:
labels["priority"] = "" labels["priority"] = ""
if get_observability().tokenizer_metrics_allowed_custom_labels: if self.server_args.tokenizer_metrics_allowed_custom_labels:
for ( for label in self.server_args.tokenizer_metrics_allowed_custom_labels:
label
) in get_observability().tokenizer_metrics_allowed_custom_labels:
labels[label] = "" labels[label] = ""
if get_observability().extra_metric_labels: if self.server_args.extra_metric_labels:
labels.update(get_observability().extra_metric_labels) labels.update(self.server_args.extra_metric_labels)
tokenizer_collector_cls = resolve_collector_class( tokenizer_collector_cls = resolve_collector_class(
self.server_args, self.server_args,
STAT_LOGGER_ROLE_TOKENIZER, STAT_LOGGER_ROLE_TOKENIZER,
@@ -589,18 +582,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.metrics_collector = tokenizer_collector_cls( self.metrics_collector = tokenizer_collector_cls(
server_args=self.server_args, server_args=self.server_args,
labels=labels, labels=labels,
bucket_time_to_first_token=get_observability().bucket_time_to_first_token, bucket_time_to_first_token=self.server_args.bucket_time_to_first_token,
bucket_e2e_request_latency=get_observability().bucket_e2e_request_latency, bucket_e2e_request_latency=self.server_args.bucket_e2e_request_latency,
bucket_inter_token_latency=get_observability().bucket_inter_token_latency, bucket_inter_token_latency=self.server_args.bucket_inter_token_latency,
) )
start_cpu_monitor_thread("tokenizer") start_cpu_monitor_thread("tokenizer")
if get_observability().gc_warning_threshold_secs > 0.0: if self.server_args.gc_warning_threshold_secs > 0.0:
configure_gc_warning(get_observability().gc_warning_threshold_secs) configure_gc_warning(self.server_args.gc_warning_threshold_secs)
self.soft_watchdog = Watchdog.create( self.soft_watchdog = Watchdog.create(
debug_name="TokenizerManager", debug_name="TokenizerManager",
watchdog_timeout=get_device().soft_watchdog_timeout, watchdog_timeout=self.server_args.soft_watchdog_timeout,
soft=True, soft=True,
test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(), test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(),
) )
@@ -1366,7 +1359,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
return batch_size > 0 and ( return batch_size > 0 and (
self.server_args.enable_tokenizer_batch_encode self.server_args.enable_tokenizer_batch_encode
or ( or (
(not get_parallel().enable_dp_attention) (not self.server_args.enable_dp_attention)
and (not self._batch_has_text(batch_size, requests)) and (not self._batch_has_text(batch_size, requests))
) )
) )
@@ -1764,7 +1757,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# default the load format to the server_args # default the load format to the server_args
if obj.load_format is None: if obj.load_format is None:
obj.load_format = get_model().load_format obj.load_format = self.server_args.load_format
logger.info("Start update_weights. Load format=%s", obj.load_format) logger.info("Start update_weights. Load format=%s", obj.load_format)
if obj.abort_all_requests: if obj.abort_all_requests:
@@ -1790,9 +1783,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def _update_model_path_info(self, model_path: str, load_format: str): def _update_model_path_info(self, model_path: str, load_format: str):
self.served_model_name = model_path self.served_model_name = model_path
from sglang.srt.runtime_context import get_context self.server_args.override(
get_context().override(
"tokenizer.update_weights", model_path=model_path, load_format=load_format "tokenizer.update_weights", model_path=model_path, load_format=load_format
) )
self.model_path = model_path self.model_path = model_path
@@ -1936,7 +1927,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
"id": rid, "id": rid,
"finish_reason": recv_obj.finished_reasons[i], "finish_reason": recv_obj.finished_reasons[i],
"prompt_tokens": recv_obj.prompt_tokens[i], "prompt_tokens": recv_obj.prompt_tokens[i],
"weight_version": get_serving().weight_version, "weight_version": self.server_args.weight_version,
"num_retractions": recv_obj.retraction_counts[i], "num_retractions": recv_obj.retraction_counts[i],
} }
@@ -2810,7 +2801,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
meta_info = { meta_info = {
"id": recv_obj.rid, "id": recv_obj.rid,
"finish_reason": finish_reason, "finish_reason": finish_reason,
"weight_version": get_serving().weight_version, "weight_version": self.server_args.weight_version,
"e2e_latency": state.time_stats.get_e2e_latency(), "e2e_latency": state.time_stats.get_e2e_latency(),
} }
is_stream = getattr(state.obj, "stream", False) is_stream = getattr(state.obj, "stream", False)
@@ -597,10 +597,7 @@ class TokenizerManagerScoreMixin:
f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})" f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})"
) )
# Check if multi-item scoring is enabled. enable_mis is a static startup # Check if multi-item scoring is enabled
# feature flag (never overridden post-publish), and score_request is also
# exercised on a bare mixin without a published context, so read it off
# server_args rather than the resolved-config bag.
use_multi_item_scoring = self.server_args.enable_mis use_multi_item_scoring = self.server_args.enable_mis
input_ids = None input_ids = None
+10 -11
View File
@@ -47,7 +47,6 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.runtime_context import get_exec, get_model, get_schedule, get_spec
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed
from sglang.srt.utils.hf_transformers_utils import ( from sglang.srt.utils.hf_transformers_utils import (
@@ -406,14 +405,14 @@ class TpModelWorker(BaseTpWorker):
self.model_config = ModelConfig.from_server_args( self.model_config = ModelConfig.from_server_args(
self.server_args, self.server_args,
model_path=( model_path=(
get_model().model_path self.server_args.model_path
if not self.is_draft_worker if not self.is_draft_worker
else get_spec().speculative_draft_model_path else self.server_args.speculative_draft_model_path
), ),
model_revision=( model_revision=(
get_model().revision self.server_args.revision
if not self.is_draft_worker if not self.is_draft_worker
else get_spec().speculative_draft_model_revision else self.server_args.speculative_draft_model_revision
), ),
is_draft_model=self.is_draft_worker, is_draft_model=self.is_draft_worker,
context_length=self.context_length, context_length=self.context_length,
@@ -424,7 +423,7 @@ class TpModelWorker(BaseTpWorker):
self._model_runner = ModelRunner( self._model_runner = ModelRunner(
model_config=self.model_config, model_config=self.model_config,
mem_fraction_static=get_schedule().mem_fraction_static, mem_fraction_static=self.server_args.mem_fraction_static,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
ps=self.ps, ps=self.ps,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
@@ -440,11 +439,11 @@ class TpModelWorker(BaseTpWorker):
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
self.model_runner_list.append(self.model_runner) self.model_runner_list.append(self.model_runner)
for i in range(1, get_spec().speculative_num_steps): for i in range(1, self.server_args.speculative_num_steps):
self.model_runner_list.append( self.model_runner_list.append(
ModelRunner( ModelRunner(
model_config=self.model_config, model_config=self.model_config,
mem_fraction_static=get_schedule().mem_fraction_static, mem_fraction_static=self.server_args.mem_fraction_static,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
ps=self.ps, ps=self.ps,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
@@ -460,7 +459,7 @@ class TpModelWorker(BaseTpWorker):
def _init_dllm_algorithm(self): def _init_dllm_algorithm(self):
from sglang.srt.dllm.algorithm.base import DllmAlgorithm from sglang.srt.dllm.algorithm.base import DllmAlgorithm
if get_exec().dllm.dllm_algorithm is not None: if self.server_args.dllm_algorithm is not None:
self.dllm_algorithm = DllmAlgorithm.from_server_args(self.server_args) self.dllm_algorithm = DllmAlgorithm.from_server_args(self.server_args)
else: else:
self.dllm_algorithm = None self.dllm_algorithm = None
@@ -486,9 +485,9 @@ class TpModelWorker(BaseTpWorker):
) )
return ( return (
self.model_runner.max_total_num_tokens, self.model_runner.max_total_num_tokens,
get_schedule().max_prefill_tokens, self.server_args.max_prefill_tokens,
self.model_runner.max_running_requests, self.model_runner.max_running_requests,
get_schedule().max_queued_requests, self.server_args.max_queued_requests,
max_req_len, max_req_len,
max_req_len - 5, max_req_len - 5,
self.random_seed, self.random_seed,
+3 -3
View File
@@ -26,7 +26,7 @@ from sglang.srt.mem_cache.common import (
evict_from_tree_cache, evict_from_tree_cache,
) )
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.runtime_context import get_exec, get_server_args from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
is_cpu, is_cpu,
is_cuda, is_cuda,
@@ -65,7 +65,7 @@ def write_cache_indices(
prefix_tensors: list[torch.Tensor], prefix_tensors: list[torch.Tensor],
req_to_token_pool: ReqToTokenPool, req_to_token_pool: ReqToTokenPool,
): ):
if support_triton(get_exec().kernel.attention_backend): if support_triton(get_server_args().attention_backend):
prefix_pointers = torch.tensor( prefix_pointers = torch.tensor(
[t.data_ptr() for t in prefix_tensors], [t.data_ptr() for t in prefix_tensors],
dtype=torch.uint64, dtype=torch.uint64,
@@ -106,7 +106,7 @@ def get_last_loc(
req_pool_indices_tensor: torch.Tensor, req_pool_indices_tensor: torch.Tensor,
prefix_lens_tensor: torch.Tensor, prefix_lens_tensor: torch.Tensor,
) -> torch.Tensor: ) -> torch.Tensor:
attn_backend = get_exec().kernel.attention_backend attn_backend = get_server_args().attention_backend
uses_triton_dispatch = attn_backend not in ("ascend", "torch_native") uses_triton_dispatch = attn_backend not in ("ascend", "torch_native")
if _is_hip and uses_triton_dispatch: if _is_hip and uses_triton_dispatch:
+2 -2
View File
@@ -16,7 +16,7 @@ from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.runtime_context import get_server_args, get_serving from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils.common import ceil_align from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -183,7 +183,7 @@ def _release_overallocated_kv_indices(
# strip_thinking_cache intentionally reports output tokens as overallocated # strip_thinking_cache intentionally reports output tokens as overallocated
# so they fall into the free path below (#22373). # so they fall into the free path below (#22373).
if spec_algo is None and not get_serving().strip_thinking_cache: if spec_algo is None and not global_server_args.strip_thinking_cache:
assert ( assert (
start_p == end_p start_p == end_p
), f"Unexpected overallocated KV cache, {req.kv_committed_len=}, {req.kv.kv_allocated_len=}" ), f"Unexpected overallocated KV cache, {req.kv_committed_len=}, {req.kv.kv_allocated_len=}"
@@ -21,7 +21,7 @@ from sglang.srt.environ import envs
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.deepseek_v4_compress_state import CompressStatePool from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
from sglang.srt.mem_cache.memory_pool import KVCache from sglang.srt.mem_cache.memory_pool import KVCache
from sglang.srt.runtime_context import get_exec, get_server_args from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import ceil_div, is_hip from sglang.srt.utils import ceil_div, is_hip
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -276,7 +276,7 @@ class DeepSeekV4IndexerPool(KVCache):
end_layer, end_layer,
) )
self.index_head_dim = index_head_dim self.index_head_dim = index_head_dim
self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer self.use_fp4_indexer = get_server_args().enable_deepseek_v4_fp4_indexer
self._create_buffer() self._create_buffer()
@@ -58,15 +58,7 @@ from sglang.srt.mem_cache.memory_pool import (
) )
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import get_model, get_parallel
get_disagg,
get_exec,
get_memory,
get_model,
get_parallel,
get_schedule,
get_spec,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
@@ -123,7 +115,9 @@ if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import ( from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import (
SpecAuxHiddenStateConfig, SpecAuxHiddenStateConfig,
) )
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import (
MemoryPoolConfig,
)
class KVCacheConfigResult(msgspec.Struct, frozen=True, kw_only=True): class KVCacheConfigResult(msgspec.Struct, frozen=True, kw_only=True):
@@ -314,8 +308,8 @@ class KVCacheConfigurator:
# from one byte buffer, then return. Gated to the target worker # from one byte buffer, then return. Gated to the target worker
# (req_to_token_pool is None); supports hybrid Mamba and hybrid SWA (not DSV4). # (req_to_token_pool is None); supports hybrid Mamba and hybrid SWA (not DSV4).
if ( if (
get_memory().enable_unified_memory self.server_args.enable_unified_memory
and get_disagg().disaggregation_mode == "null" and self.server_args.disaggregation_mode == "null"
and req_to_token_pool is None and req_to_token_pool is None
): ):
if self.mambaish_config is not None: if self.mambaish_config is not None:
@@ -364,13 +358,13 @@ class KVCacheConfigurator:
# TARGET_VERIFY, so their pools skip the per-step intermediate # TARGET_VERIFY, so their pools skip the per-step intermediate
# (SpeculativeState) buffers only the target pool consumes. # (SpeculativeState) buffers only the target pool consumes.
req_to_token_pool = req_to_token_pool.clone_with_new_mamba( req_to_token_pool = req_to_token_pool.clone_with_new_mamba(
mamba_size=get_schedule().max_mamba_cache_size, mamba_size=self.server_args.max_mamba_cache_size,
mamba_spec_state_size=sizes.max_running_requests, mamba_spec_state_size=sizes.max_running_requests,
cache_params=self.mambaish_config.mamba2_cache_params, cache_params=self.mambaish_config.mamba2_cache_params,
device=self.device, device=self.device,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
draft_model_idx=self.draft_model_idx, draft_model_idx=self.draft_model_idx,
speculative_eagle_topk=get_spec().speculative_eagle_topk, speculative_eagle_topk=self.server_args.speculative_eagle_topk,
) )
# Initialize token_to_kv_pool # Initialize token_to_kv_pool
@@ -400,7 +394,7 @@ class KVCacheConfigurator:
# unsupported pool families before allocation. Keep this guard here so # unsupported pool families before allocation. Keep this guard here so
# future pool-selection refactors fail at boot instead of on first use. # future pool-selection refactors fail at boot instead of on first use.
if ( if (
get_schedule().prefill_only_disable_kv_cache self.server_args.prefill_only_disable_kv_cache
and not self.is_draft_worker and not self.is_draft_worker
and not isinstance(token_to_kv_pool, NoOpMHATokenToKVPool) and not isinstance(token_to_kv_pool, NoOpMHATokenToKVPool)
): ):
@@ -438,8 +432,8 @@ class KVCacheConfigurator:
assert self.page_size >= 1, f"page_size must be >= 1, got {self.page_size}" assert self.page_size >= 1, f"page_size must be >= 1, got {self.page_size}"
# Mirror the non-shared path's extra_max_context_len computation. # Mirror the non-shared path's extra_max_context_len computation.
extra_max_context_len = 4 extra_max_context_len = 4
if get_spec().speculative_num_draft_tokens is not None: if self.server_args.speculative_num_draft_tokens is not None:
extra_max_context_len += get_spec().speculative_num_draft_tokens extra_max_context_len += self.server_args.speculative_num_draft_tokens
mamba_layer_ids = [ mamba_layer_ids = [
i i
@@ -468,14 +462,14 @@ class KVCacheConfigurator:
model_context_len=self.model_config.context_len, model_context_len=self.model_config.context_len,
extra_max_context_len=extra_max_context_len, extra_max_context_len=extra_max_context_len,
max_total_num_tokens=max_total_num_tokens, max_total_num_tokens=max_total_num_tokens,
max_mamba_cache_size=get_schedule().max_mamba_cache_size, max_mamba_cache_size=self.server_args.max_mamba_cache_size,
max_num_reqs=max_num_reqs, max_num_reqs=max_num_reqs,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
speculative_num_draft_tokens=get_spec().speculative_num_draft_tokens, speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
disable_overlap_schedule=get_schedule().disable_overlap_schedule, disable_overlap_schedule=self.server_args.disable_overlap_schedule,
need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"), need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"),
mamba_full_memory_ratio=get_schedule().mamba_full_memory_ratio, mamba_full_memory_ratio=self.server_args.mamba_full_memory_ratio,
# Overlap mode: the allocator's `free` drops a wait_stream(forward_stream) # Overlap mode: the allocator's `free` drops a wait_stream(forward_stream)
# barrier so eager compaction serializes after the in-flight forward's # barrier so eager compaction serializes after the in-flight forward's
# v2p/KV reads. Near-no-op in normal mode. # v2p/KV reads. Near-no-op in normal mode.
@@ -508,13 +502,13 @@ class KVCacheConfigurator:
), "unified memory pool does not support MLA-SWA hybrid yet" ), "unified memory pool does not support MLA-SWA hybrid yet"
# Mirror the non-shared path's extra_max_context_len computation. # Mirror the non-shared path's extra_max_context_len computation.
extra_max_context_len = 4 extra_max_context_len = 4
if get_spec().speculative_num_draft_tokens is not None: if self.server_args.speculative_num_draft_tokens is not None:
extra_max_context_len += get_spec().speculative_num_draft_tokens extra_max_context_len += self.server_args.speculative_num_draft_tokens
req_to_token_pool = ReqToTokenPool( req_to_token_pool = ReqToTokenPool(
size=max_num_reqs, size=max_num_reqs,
max_context_len=self.model_config.context_len + extra_max_context_len, max_context_len=self.model_config.context_len + extra_max_context_len,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
) )
head_num = self.model_config.get_num_kv_heads(get_parallel().attn_tp_size) head_num = self.model_config.get_num_kv_heads(get_parallel().attn_tp_size)
@@ -564,8 +558,8 @@ class KVCacheConfigurator:
full_attention_layer_ids=full_attention_layer_ids, full_attention_layer_ids=full_attention_layer_ids,
full_max_total_num_tokens=full_max_total_num_tokens, full_max_total_num_tokens=full_max_total_num_tokens,
swa_max_total_num_tokens=swa_max_total_num_tokens, swa_max_total_num_tokens=swa_max_total_num_tokens,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"), need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"),
# Overlap mode: same wait_stream(forward_stream) rationale as # Overlap mode: same wait_stream(forward_stream) rationale as
# `_init_unified_mamba_pools`. # `_init_unified_mamba_pools`.
forward_stream=self.forward_stream, forward_stream=self.forward_stream,
@@ -585,7 +579,7 @@ class KVCacheConfigurator:
is_dsv4_model: bool, is_dsv4_model: bool,
current_platform, current_platform,
): ):
if not get_schedule().prefill_only_disable_kv_cache or self.is_draft_worker: if not self.server_args.prefill_only_disable_kv_cache or self.is_draft_worker:
return return
unsupported_pool_family = None unsupported_pool_family = None
@@ -594,7 +588,7 @@ class KVCacheConfigurator:
elif current_platform.is_out_of_tree() and not self.mambaish_config: elif current_platform.is_out_of_tree() and not self.mambaish_config:
unsupported_pool_family = "out-of-tree platform KV pool" unsupported_pool_family = "out-of-tree platform KV pool"
elif ( elif (
get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config self.server_args.attention_backend == "ascend" and not self.mambaish_config
): ):
unsupported_pool_family = "NPU/Ascend KV pool" unsupported_pool_family = "NPU/Ascend KV pool"
elif self.use_mla_backend and is_dsa_model: elif self.use_mla_backend and is_dsa_model:
@@ -620,9 +614,9 @@ class KVCacheConfigurator:
def _build_req_to_token_pool(self, *, max_num_reqs: int) -> ReqToTokenPool: def _build_req_to_token_pool(self, *, max_num_reqs: int) -> ReqToTokenPool:
extra_max_context_len = get_req_to_token_extra_context_len(self.server_args) extra_max_context_len = get_req_to_token_extra_context_len(self.server_args)
if get_disagg().disaggregation_mode == "decode": if self.server_args.disaggregation_mode == "decode":
# Extra slots for pre-allocated requests # Extra slots for pre-allocated requests
pre_alloc_size = get_disagg().disaggregation_decode_extra_slots pre_alloc_size = self.server_args.disaggregation_decode_extra_slots
if self.mambaish_config: if self.mambaish_config:
req_to_token_pool = self._build_hybrid_mamba_decode_req_pool( req_to_token_pool = self._build_hybrid_mamba_decode_req_pool(
max_num_reqs=max_num_reqs, max_num_reqs=max_num_reqs,
@@ -654,13 +648,15 @@ class KVCacheConfigurator:
extra_max_context_len: int, extra_max_context_len: int,
pre_alloc_size: int, pre_alloc_size: int,
) -> ReqToTokenPool: ) -> ReqToTokenPool:
from sglang.srt.disaggregation.decode import HybridMambaDecodeReqToTokenPool from sglang.srt.disaggregation.decode import (
HybridMambaDecodeReqToTokenPool,
)
req_to_token_pool = HybridMambaDecodeReqToTokenPool( req_to_token_pool = HybridMambaDecodeReqToTokenPool(
size=max_num_reqs, size=max_num_reqs,
max_context_len=self.model_config.context_len + extra_max_context_len, max_context_len=self.model_config.context_len + extra_max_context_len,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
cache_params=self.mambaish_config.mamba2_cache_params, cache_params=self.mambaish_config.mamba2_cache_params,
mamba_layer_ids=( mamba_layer_ids=(
[ [
@@ -670,11 +666,11 @@ class KVCacheConfigurator:
] ]
), ),
speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens, speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens,
speculative_eagle_topk=get_spec().speculative_eagle_topk, speculative_eagle_topk=self.server_args.speculative_eagle_topk,
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
pre_alloc_size=pre_alloc_size, pre_alloc_size=pre_alloc_size,
enable_overlap_schedule=not get_schedule().disable_overlap_schedule, enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
mamba_size=get_schedule().max_mamba_cache_size, mamba_size=self.server_args.max_mamba_cache_size,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
) )
return req_to_token_pool return req_to_token_pool
@@ -692,7 +688,7 @@ class KVCacheConfigurator:
size=max_num_reqs, size=max_num_reqs,
max_context_len=self.model_config.context_len + extra_max_context_len, max_context_len=self.model_config.context_len + extra_max_context_len,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
pre_alloc_size=pre_alloc_size, pre_alloc_size=pre_alloc_size,
) )
return req_to_token_pool return req_to_token_pool
@@ -705,11 +701,11 @@ class KVCacheConfigurator:
) -> ReqToTokenPool: ) -> ReqToTokenPool:
req_to_token_pool = HybridReqToTokenPool( req_to_token_pool = HybridReqToTokenPool(
size=max_num_reqs, size=max_num_reqs,
mamba_size=get_schedule().max_mamba_cache_size, mamba_size=self.server_args.max_mamba_cache_size,
mamba_spec_state_size=max_num_reqs, mamba_spec_state_size=max_num_reqs,
max_context_len=self.model_config.context_len + extra_max_context_len, max_context_len=self.model_config.context_len + extra_max_context_len,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
cache_params=self.mambaish_config.mamba2_cache_params, cache_params=self.mambaish_config.mamba2_cache_params,
mamba_layer_ids=( mamba_layer_ids=(
[ [
@@ -721,18 +717,18 @@ class KVCacheConfigurator:
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(), enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(),
speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens, speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens,
speculative_eagle_topk=get_spec().speculative_eagle_topk, speculative_eagle_topk=self.server_args.speculative_eagle_topk,
enable_overlap_schedule=not get_schedule().disable_overlap_schedule, enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
enable_linear_replayssm=get_exec().mamba.enable_linear_replayssm, enable_linear_replayssm=self.server_args.enable_linear_replayssm,
linear_replayssm_cache_len=get_exec().mamba.linear_replayssm_cache_len, linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len,
mamba_envelope_layout=get_memory().enable_page_major_kv_layout, mamba_envelope_layout=self.server_args.enable_page_major_kv_layout,
# ReplaySSM spec-verify is GDN-only: activate the pool machinery # ReplaySSM spec-verify is GDN-only: activate the pool machinery
# (rings + cursors + the intermediate_ssm gate) only for GDN-hybrid # (rings + cursors + the intermediate_ssm gate) only for GDN-hybrid
# models, so any other mamba-ish model (Mamba2/Nemotron, lightning, # models, so any other mamba-ish model (Mamba2/Nemotron, lightning,
# ...) run with the flag set stays byte-identical to flag-off. # ...) run with the flag set stays byte-identical to flag-off.
enable_gdn_replayssm_spec=( enable_gdn_replayssm_spec=(
get_exec().mamba.enable_gdn_replayssm_spec self.server_args.enable_gdn_replayssm_spec
and self.hybrid_gdn_config is not None and self.hybrid_gdn_config is not None
), ),
) )
@@ -758,7 +754,7 @@ class KVCacheConfigurator:
size=max_num_reqs, size=max_num_reqs,
max_context_len=self.model_config.context_len + extra_max_context_len, max_context_len=self.model_config.context_len + extra_max_context_len,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
) )
return req_to_token_pool return req_to_token_pool
@@ -774,7 +770,7 @@ class KVCacheConfigurator:
# selected by swapping in the PageMajorMHATokenToKVPool subclass. The # selected by swapping in the PageMajorMHATokenToKVPool subclass. The
# default keeps upstream's per-layer layout. The Mamba state pool is routed # default keeps upstream's per-layer layout. The Mamba state pool is routed
# separately via `mamba_envelope_layout` on the req-to-token pool above. # separately via `mamba_envelope_layout` on the req-to-token pool above.
enable_page_major = get_memory().enable_page_major_kv_layout enable_page_major = self.server_args.enable_page_major_kv_layout
mha_pool_class = ( mha_pool_class = (
PageMajorMHATokenToKVPool if enable_page_major else MHATokenToKVPool PageMajorMHATokenToKVPool if enable_page_major else MHATokenToKVPool
) )
@@ -806,7 +802,7 @@ class KVCacheConfigurator:
max_total_num_tokens=sizes.max_total_num_tokens, max_total_num_tokens=sizes.max_total_num_tokens,
) )
elif ( elif (
get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config self.server_args.attention_backend == "ascend" and not self.mambaish_config
): ):
if self.is_hybrid_swa: if self.is_hybrid_swa:
token_to_kv_pool = self._build_ascend_swa_kv_pool( token_to_kv_pool = self._build_ascend_swa_kv_pool(
@@ -882,12 +878,14 @@ class KVCacheConfigurator:
c128_state_dtype: Optional[torch.dtype], c128_state_dtype: Optional[torch.dtype],
req_to_token_pool: ReqToTokenPool, req_to_token_pool: ReqToTokenPool,
) -> KVCache: ) -> KVCache:
swa_page_size = get_schedule().page_size swa_page_size = self.server_args.page_size
if not _is_npu: if not _is_npu:
assert swa_page_size == 256, "In paged swa mode, page_size must be 256." assert swa_page_size == 256, "In paged swa mode, page_size must be 256."
if self.is_draft_worker: if self.is_draft_worker:
from sglang.srt.models.deepseek_v4_nextn import COMPRESS_RATIO_NEXTN_LAYER from sglang.srt.models.deepseek_v4_nextn import (
COMPRESS_RATIO_NEXTN_LAYER,
)
compression_ratios = [ compression_ratios = [
COMPRESS_RATIO_NEXTN_LAYER COMPRESS_RATIO_NEXTN_LAYER
@@ -914,12 +912,12 @@ class KVCacheConfigurator:
# sliding eviction in ``ScheduleBatch._evict_swa``. # sliding eviction in ``ScheduleBatch._evict_swa``.
c4_state_pool_size = npu_state_pool_size( c4_state_pool_size = npu_state_pool_size(
ratio=4, ratio=4,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
max_num_reqs=max_running_requests, max_num_reqs=max_running_requests,
) )
c128_state_pool_size = npu_state_pool_size( c128_state_pool_size = npu_state_pool_size(
ratio=128, ratio=128,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
max_num_reqs=max_running_requests, max_num_reqs=max_running_requests,
) )
else: else:
@@ -937,7 +935,7 @@ class KVCacheConfigurator:
c128_size=c128_max_total_num_tokens, c128_size=c128_max_total_num_tokens,
c4_state_pool_size=c4_state_pool_size, c4_state_pool_size=c4_state_pool_size,
c128_state_pool_size=c128_state_pool_size, c128_state_pool_size=c128_state_pool_size,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
swa_page_size=swa_page_size, swa_page_size=swa_page_size,
sliding_window=self.model_config.window_size, sliding_window=self.model_config.window_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
@@ -948,11 +946,11 @@ class KVCacheConfigurator:
indexer_head_dim=self.model_config.index_head_dim, indexer_head_dim=self.model_config.index_head_dim,
layer_num=self.layer_info.num_effective_layers, layer_num=self.layer_info.num_effective_layers,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
compression_ratios=compression_ratios, compression_ratios=compression_ratios,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
enable_hisparse=get_memory().enable_hisparse, enable_hisparse=self.server_args.enable_hisparse,
online_mtp_max_draft_tokens=( online_mtp_max_draft_tokens=(
self.server_args.max_speculative_num_draft_tokens or 0 self.server_args.max_speculative_num_draft_tokens or 0
), ),
@@ -963,7 +961,7 @@ class KVCacheConfigurator:
PoolCls = current_platform.get_dsa_kv_pool_cls() PoolCls = current_platform.get_dsa_kv_pool_cls()
token_to_kv_pool = PoolCls( token_to_kv_pool = PoolCls(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank, kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim, qk_rope_head_dim=self.model_config.qk_rope_head_dim,
@@ -974,7 +972,7 @@ class KVCacheConfigurator:
kv_cache_dtype=self.kv_cache_dtype, kv_cache_dtype=self.kv_cache_dtype,
server_args=self.server_args, server_args=self.server_args,
), ),
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config), index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
@@ -987,14 +985,14 @@ class KVCacheConfigurator:
PoolCls = current_platform.get_mla_kv_pool_cls() PoolCls = current_platform.get_mla_kv_pool_cls()
token_to_kv_pool = PoolCls( token_to_kv_pool = PoolCls(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank, kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim, qk_rope_head_dim=self.model_config.qk_rope_head_dim,
index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None), index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None),
layer_num=self.layer_info.num_effective_layers, layer_num=self.layer_info.num_effective_layers,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
) )
@@ -1004,13 +1002,13 @@ class KVCacheConfigurator:
PoolCls = current_platform.get_mha_kv_pool_cls() PoolCls = current_platform.get_mha_kv_pool_cls()
token_to_kv_pool = PoolCls( token_to_kv_pool = PoolCls(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim, head_dim=self.model_config.head_dim,
layer_num=self.layer_info.num_effective_layers, layer_num=self.layer_info.num_effective_layers,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
) )
@@ -1022,7 +1020,9 @@ class KVCacheConfigurator:
full_max_total_num_tokens: Optional[int], full_max_total_num_tokens: Optional[int],
swa_max_total_num_tokens: Optional[int], swa_max_total_num_tokens: Optional[int],
) -> KVCache: ) -> KVCache:
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMHATokenToKVPool,
)
kwargs = {} kwargs = {}
if self.is_hybrid_swa_compress: if self.is_hybrid_swa_compress:
@@ -1039,7 +1039,7 @@ class KVCacheConfigurator:
token_to_kv_pool = SWAKVPool( token_to_kv_pool = SWAKVPool(
size=full_max_total_num_tokens, size=full_max_total_num_tokens,
size_swa=swa_max_total_num_tokens, size_swa=swa_max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
post_capture_active=self.post_capture_kv_active, post_capture_active=self.post_capture_kv_active,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
@@ -1055,35 +1055,39 @@ class KVCacheConfigurator:
def _build_ascend_mla_kv_pool( def _build_ascend_mla_kv_pool(
self, *, max_total_num_tokens: int, is_dsa_model: bool self, *, max_total_num_tokens: int, is_dsa_model: bool
) -> KVCache: ) -> KVCache:
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMLATokenToKVPool from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMLATokenToKVPool,
)
token_to_kv_pool = NPUMLATokenToKVPool( token_to_kv_pool = NPUMLATokenToKVPool(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank, kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim, qk_rope_head_dim=self.model_config.qk_rope_head_dim,
index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None), index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None),
layer_num=self.layer_info.num_effective_layers, layer_num=self.layer_info.num_effective_layers,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
) )
return token_to_kv_pool return token_to_kv_pool
def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMHATokenToKVPool,
)
token_to_kv_pool = NPUMHATokenToKVPool( token_to_kv_pool = NPUMHATokenToKVPool(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim, head_dim=self.model_config.head_dim,
layer_num=self.layer_info.num_effective_layers, layer_num=self.layer_info.num_effective_layers,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
) )
@@ -1097,7 +1101,7 @@ class KVCacheConfigurator:
dsa_cp_layer_shard_size, dsa_cp_layer_shard_size,
) = get_glm_dsa_cp_layer_shard_info(self) ) = get_glm_dsa_cp_layer_shard_info(self)
pool_kwargs = {} pool_kwargs = {}
if get_memory().enable_hisparse: if self.server_args.enable_hisparse:
PoolCls = HiSparseDSATokenToKVPool PoolCls = HiSparseDSATokenToKVPool
from sglang.srt.mem_cache.sparsity import parse_hisparse_config from sglang.srt.mem_cache.sparsity import parse_hisparse_config
@@ -1117,7 +1121,7 @@ class KVCacheConfigurator:
PoolCls = DSATokenToKVPool PoolCls = DSATokenToKVPool
token_to_kv_pool = PoolCls( token_to_kv_pool = PoolCls(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank, kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim, qk_rope_head_dim=self.model_config.qk_rope_head_dim,
@@ -1128,7 +1132,7 @@ class KVCacheConfigurator:
kv_cache_dtype=self.kv_cache_dtype, kv_cache_dtype=self.kv_cache_dtype,
server_args=self.server_args, server_args=self.server_args,
), ),
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config), index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
@@ -1139,13 +1143,13 @@ class KVCacheConfigurator:
def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
token_to_kv_pool = MLATokenToKVPoolFP4( token_to_kv_pool = MLATokenToKVPoolFP4(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank, kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim, qk_rope_head_dim=self.model_config.qk_rope_head_dim,
layer_num=self.layer_info.num_effective_layers, layer_num=self.layer_info.num_effective_layers,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
) )
@@ -1154,13 +1158,13 @@ class KVCacheConfigurator:
def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
token_to_kv_pool = MLATokenToKVPool( token_to_kv_pool = MLATokenToKVPool(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank, kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim, qk_rope_head_dim=self.model_config.qk_rope_head_dim,
layer_num=self.layer_info.num_effective_layers, layer_num=self.layer_info.num_effective_layers,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
) )
@@ -1217,7 +1221,7 @@ class KVCacheConfigurator:
token_to_kv_pool = SWAKVPool( token_to_kv_pool = SWAKVPool(
size=full_max_total_num_tokens, size=full_max_total_num_tokens,
size_swa=size_swa, size_swa=size_swa,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
post_capture_active=self.post_capture_kv_active, post_capture_active=self.post_capture_kv_active,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
@@ -1225,7 +1229,7 @@ class KVCacheConfigurator:
swa_attention_layer_ids=swa_attention_layer_ids, swa_attention_layer_ids=swa_attention_layer_ids,
full_attention_layer_ids=full_attention_layer_ids, full_attention_layer_ids=full_attention_layer_ids,
device=self.device, device=self.device,
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None), enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
token_to_kv_pool_class=swa_pool_class, token_to_kv_pool_class=swa_pool_class,
**kwargs, **kwargs,
) )
@@ -1240,7 +1244,7 @@ class KVCacheConfigurator:
) )
token_to_kv_pool = MiniMaxSparseKVPool( token_to_kv_pool = MiniMaxSparseKVPool(
size=max_total_num_tokens, size=max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
index_dtype=self.model_dtype, index_dtype=self.model_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
@@ -1250,7 +1254,7 @@ class KVCacheConfigurator:
sparse_layer_ids=sparse_layer_ids, sparse_layer_ids=sparse_layer_ids,
disable_value_sparse_layer_ids=disable_value_sparse_layer_ids, disable_value_sparse_layer_ids=disable_value_sparse_layer_ids,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
) )
@@ -1289,7 +1293,7 @@ class KVCacheConfigurator:
else mha_pool_class else mha_pool_class
) )
token_to_kv_pool = HybridLinearKVPool( token_to_kv_pool = HybridLinearKVPool(
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
size=max_total_num_tokens, size=max_total_num_tokens,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
@@ -1298,8 +1302,8 @@ class KVCacheConfigurator:
full_attention_layer_ids=full_attention_layer_ids, full_attention_layer_ids=full_attention_layer_ids,
device=self.device, device=self.device,
mamba_pool=req_to_token_pool.mamba_pool, mamba_pool=req_to_token_pool.mamba_pool,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None), enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
use_mla=self.use_mla_backend, use_mla=self.use_mla_backend,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
full_kv_pool_class=full_pool_class, full_kv_pool_class=full_pool_class,
@@ -1312,18 +1316,18 @@ class KVCacheConfigurator:
def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
token_to_kv_pool = MHATokenToKVPoolFP4( token_to_kv_pool = MHATokenToKVPoolFP4(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim, head_dim=self.model_config.head_dim,
v_head_dim=self.model_config.v_head_dim, v_head_dim=self.model_config.v_head_dim,
layer_num=self.layer_info.num_effective_layers, layer_num=self.layer_info.num_effective_layers,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
enable_alt_stream=not get_disagg().enable_pdmux, enable_alt_stream=not self.server_args.enable_pdmux,
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None), enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
) )
return token_to_kv_pool return token_to_kv_pool
@@ -1335,7 +1339,7 @@ class KVCacheConfigurator:
else: else:
pool_cls = ( pool_cls = (
NoOpMHATokenToKVPool NoOpMHATokenToKVPool
if get_schedule().prefill_only_disable_kv_cache if self.server_args.prefill_only_disable_kv_cache
else mha_pool_class else mha_pool_class
) )
pool_kwargs = {} pool_kwargs = {}
@@ -1345,18 +1349,18 @@ class KVCacheConfigurator:
pool_kwargs["post_capture_active"] = self.post_capture_kv_active pool_kwargs["post_capture_active"] = self.post_capture_kv_active
token_to_kv_pool = pool_cls( token_to_kv_pool = pool_cls(
max_total_num_tokens, max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim, head_dim=self.model_config.head_dim,
v_head_dim=self.model_config.v_head_dim, v_head_dim=self.model_config.v_head_dim,
layer_num=self.layer_info.num_effective_layers, layer_num=self.layer_info.num_effective_layers,
device=self.device, device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer, end_layer=self.layer_info.end_layer,
enable_alt_stream=not get_disagg().enable_pdmux, enable_alt_stream=not self.server_args.enable_pdmux,
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None), enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
**pool_kwargs, **pool_kwargs,
) )
return token_to_kv_pool return token_to_kv_pool
@@ -1371,20 +1375,20 @@ class KVCacheConfigurator:
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator], token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator],
) -> BaseTokenToKVPoolAllocator: ) -> BaseTokenToKVPoolAllocator:
# Initialize token_to_kv_pool_allocator # Initialize token_to_kv_pool_allocator
need_sort = get_disagg().disaggregation_mode in ("decode", "prefill") need_sort = self.server_args.disaggregation_mode in ("decode", "prefill")
if token_to_kv_pool_allocator is None: if token_to_kv_pool_allocator is None:
if current_platform.is_out_of_tree(): if current_platform.is_out_of_tree():
AllocatorCls = current_platform.get_paged_allocator_cls() AllocatorCls = current_platform.get_paged_allocator_cls()
token_to_kv_pool_allocator = AllocatorCls( token_to_kv_pool_allocator = AllocatorCls(
sizes.max_total_num_tokens, sizes.max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
device=self.device, device=self.device,
kvcache=token_to_kv_pool, kvcache=token_to_kv_pool,
need_sort=need_sort, need_sort=need_sort,
) )
elif _is_npu and ( elif _is_npu and (
get_exec().kernel.attention_backend == "ascend" self.server_args.attention_backend == "ascend"
or is_dsv4_model or is_dsv4_model
or self.hybrid_gdn_config is not None or self.hybrid_gdn_config is not None
): ):
@@ -1402,7 +1406,7 @@ class KVCacheConfigurator:
token_to_kv_pool_allocator = swa_allocator_cls( token_to_kv_pool_allocator = swa_allocator_cls(
sizes.full_max_total_num_tokens, sizes.full_max_total_num_tokens,
sizes.swa_max_total_num_tokens, sizes.swa_max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
device=self.device, device=self.device,
kvcache=token_to_kv_pool, kvcache=token_to_kv_pool,
@@ -1415,7 +1419,7 @@ class KVCacheConfigurator:
token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator( token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator(
sizes.max_total_num_tokens, sizes.max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
device=self.device, device=self.device,
kvcache=token_to_kv_pool, kvcache=token_to_kv_pool,
@@ -1425,7 +1429,7 @@ class KVCacheConfigurator:
if self.is_hybrid_swa and sizes.full_max_total_num_tokens == 0: if self.is_hybrid_swa and sizes.full_max_total_num_tokens == 0:
token_to_kv_pool_allocator = PureSWATokenToKVPoolAllocator( token_to_kv_pool_allocator = PureSWATokenToKVPoolAllocator(
sizes.swa_max_total_num_tokens, sizes.swa_max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
device=self.device, device=self.device,
kvcache=token_to_kv_pool, kvcache=token_to_kv_pool,
@@ -1435,20 +1439,22 @@ class KVCacheConfigurator:
token_to_kv_pool_allocator = SWATokenToKVPoolAllocator( token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
sizes.full_max_total_num_tokens, sizes.full_max_total_num_tokens,
sizes.swa_max_total_num_tokens, sizes.swa_max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
device=self.device, device=self.device,
kvcache=token_to_kv_pool, kvcache=token_to_kv_pool,
need_sort=need_sort, need_sort=need_sort,
) )
else: else:
if get_memory().enable_hisparse: if self.server_args.enable_hisparse:
from sglang.srt.mem_cache.sparsity import parse_hisparse_config from sglang.srt.mem_cache.sparsity import (
parse_hisparse_config,
)
hisparse_cfg = parse_hisparse_config(self.server_args) hisparse_cfg = parse_hisparse_config(self.server_args)
token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator( token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator(
sizes.max_total_num_tokens, sizes.max_total_num_tokens,
page_size=get_schedule().page_size, page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
device=self.device, device=self.device,
kvcache=token_to_kv_pool, kvcache=token_to_kv_pool,
@@ -1456,7 +1462,8 @@ class KVCacheConfigurator:
host_to_device_ratio=hisparse_cfg.host_to_device_ratio, host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
) )
elif ( elif (
get_schedule().page_size == 1 and self.server_args.dcp_size == 1 self.server_args.page_size == 1
and self.server_args.dcp_size == 1
): ):
token_to_kv_pool_allocator = TokenToKVPoolAllocator( token_to_kv_pool_allocator = TokenToKVPoolAllocator(
sizes.max_total_num_tokens, sizes.max_total_num_tokens,
@@ -1468,7 +1475,7 @@ class KVCacheConfigurator:
else: else:
token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator( token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
sizes.max_total_num_tokens * self.server_args.dcp_size, sizes.max_total_num_tokens * self.server_args.dcp_size,
page_size=get_schedule().page_size page_size=self.server_args.page_size
* self.server_args.dcp_size, * self.server_args.dcp_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
device=self.device, device=self.device,
@@ -1476,7 +1483,7 @@ class KVCacheConfigurator:
need_sort=need_sort, need_sort=need_sort,
) )
if get_memory().enable_hisparse and is_dsv4_model: if self.server_args.enable_hisparse and is_dsv4_model:
assert self.is_hybrid_swa, "DeepSeek V4 HiSparse requires SWA mode." assert self.is_hybrid_swa, "DeepSeek V4 HiSparse requires SWA mode."
token_to_kv_pool_allocator = DeepSeekV4HiSparseTokenToKVPoolAllocator( token_to_kv_pool_allocator = DeepSeekV4HiSparseTokenToKVPoolAllocator(
token_to_kv_pool_allocator token_to_kv_pool_allocator
@@ -1528,7 +1535,7 @@ class KVCacheConfigurator:
cpu_group=get_world_group().cpu_group, cpu_group=get_world_group().cpu_group,
) )
slack_gb = pre_model_load_memory * (1 - get_schedule().mem_fraction_static) slack_gb = pre_model_load_memory * (1 - self.server_args.mem_fraction_static)
if self.mambaish_config is not None and self.post_capture_kv_active: if self.mambaish_config is not None and self.post_capture_kv_active:
# Mamba state is a fixed pre-capture allocation, so it can't ride the ~0 post-capture slack. # Mamba state is a fixed pre-capture allocation, so it can't ride the ~0 post-capture slack.
slack_gb = max( slack_gb = max(
@@ -1552,7 +1559,7 @@ class KVCacheConfigurator:
) )
raise ValueError( raise ValueError(
f"Loaded weights leave no GPU memory for the KV cache under " f"Loaded weights leave no GPU memory for the KV cache under "
f"--mem-fraction-static={get_schedule().mem_fraction_static}. " f"--mem-fraction-static={self.server_args.mem_fraction_static}. "
f"Raise --mem-fraction-static above " f"Raise --mem-fraction-static above "
f"{suggested_mem_fraction_static:.3f} " f"{suggested_mem_fraction_static:.3f} "
f"(minimum viable = 1 - available/pre = " f"(minimum viable = 1 - available/pre = "
@@ -1563,14 +1570,14 @@ class KVCacheConfigurator:
return int(rest_memory * (1 << 30)) # return in bytes return int(rest_memory * (1 << 30)) # return in bytes
def _calculate_mamba_ratio(self) -> int: def _calculate_mamba_ratio(self) -> int:
if get_memory().disable_radix_cache: if self.server_args.disable_radix_cache:
return 1 return 1
additional_ratio = 0 additional_ratio = 0
if self.server_args.enable_mamba_extra_buffer(): if self.server_args.enable_mamba_extra_buffer():
# ping-pong buffer size is 2 when overlap schedule is on, 1 otherwise. # ping-pong buffer size is 2 when overlap schedule is on, 1 otherwise.
# Lazy mode saves 1 slot (2 → 1) for overlap; non-overlap already uses 1. # Lazy mode saves 1 slot (2 → 1) for overlap; non-overlap already uses 1.
if not get_schedule().disable_overlap_schedule: if not self.server_args.disable_overlap_schedule:
if self.server_args.enable_mamba_extra_buffer_lazy(): if self.server_args.enable_mamba_extra_buffer_lazy():
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY
else: else:
@@ -1589,7 +1596,7 @@ class KVCacheConfigurator:
Page alignment is handled by the configurator, not here. Page alignment is handled by the configurator, not here.
If constraints change the value, the configurator re-runs and re-aligns. If constraints change the value, the configurator re-runs and re-aligns.
""" """
user_limit = get_schedule().max_total_tokens user_limit = self.server_args.max_total_tokens
# Apply user-specified upper bound # Apply user-specified upper bound
if user_limit is not None: if user_limit is not None:
@@ -1619,7 +1626,7 @@ class KVCacheConfigurator:
estimated = int(token_capacity / self.model_config.context_len * 512) estimated = int(token_capacity / self.model_config.context_len * 512)
estimated = max(min(estimated, 4096), 2048) estimated = max(min(estimated, 4096), 2048)
max_num_reqs = get_schedule().max_running_requests max_num_reqs = self.server_args.max_running_requests
if max_num_reqs is not None: if max_num_reqs is not None:
requested_per_worker = max_num_reqs // self.ps.attn_dp_size requested_per_worker = max_num_reqs // self.ps.attn_dp_size
max_num_reqs = min(requested_per_worker, token_capacity // 2) max_num_reqs = min(requested_per_worker, token_capacity // 2)
@@ -1630,13 +1637,13 @@ class KVCacheConfigurator:
if self.mambaish_config is not None: if self.mambaish_config is not None:
ratio = self._calculate_mamba_ratio() ratio = self._calculate_mamba_ratio()
max_num_reqs = min( max_num_reqs = min(
max_num_reqs, get_schedule().max_mamba_cache_size // ratio max_num_reqs, self.server_args.max_mamba_cache_size // ratio
) )
if max_num_reqs <= 0: if max_num_reqs <= 0:
raise RuntimeError( raise RuntimeError(
f"Hybrid (mamba/linear-attention) state cache is too small to serve " f"Hybrid (mamba/linear-attention) state cache is too small to serve "
f"any requests. max_mamba_cache_size={get_schedule().max_mamba_cache_size}, " f"any requests. max_mamba_cache_size={self.server_args.max_mamba_cache_size}, "
f"mamba_ratio={ratio}, resulting max_num_reqs={max_num_reqs}. " f"mamba_ratio={ratio}, resulting max_num_reqs={max_num_reqs}. "
f"Try: (1) reduce --max-running-requests, " f"Try: (1) reduce --max-running-requests, "
f"(2) increase --mem-fraction-static, or " f"(2) increase --mem-fraction-static, or "
@@ -1666,7 +1673,7 @@ class KVCacheConfigurator:
) )
configurator = create_memory_pool_configurator(self) configurator = create_memory_pool_configurator(self)
config = configurator.finalize_with_max_running_requests(config) config = configurator.finalize_with_max_running_requests(config)
config.mem_fraction_static = get_schedule().mem_fraction_static config.mem_fraction_static = self.server_args.mem_fraction_static
return config return config
def config_from_budget( def config_from_budget(
@@ -1682,20 +1689,18 @@ class KVCacheConfigurator:
configurator = create_memory_pool_configurator(self) configurator = create_memory_pool_configurator(self)
config = configurator.calculate_pool_sizes( config = configurator.calculate_pool_sizes(
budget_bytes, get_schedule().page_size budget_bytes, self.server_args.page_size
) )
max_tokens = self._apply_token_constraints(config.max_total_num_tokens) max_tokens = self._apply_token_constraints(config.max_total_num_tokens)
if cap_tokens is not None: if cap_tokens is not None:
max_tokens = min(max_tokens, cap_tokens) max_tokens = min(max_tokens, cap_tokens)
if max_tokens != config.max_total_num_tokens: if max_tokens != config.max_total_num_tokens:
config = configurator.calculate_pool_sizes_from_max_tokens( config = configurator.calculate_pool_sizes_from_max_tokens(
max_tokens, get_schedule().page_size max_tokens, self.server_args.page_size
) )
return config return config
def _handle_max_mamba_cache(self, total_rest_memory): def _handle_max_mamba_cache(self, total_rest_memory):
from sglang.srt.runtime_context import get_context
config = self.mambaish_config config = self.mambaish_config
server_args = self.server_args server_args = self.server_args
assert config is not None assert config is not None
@@ -1705,11 +1710,11 @@ class KVCacheConfigurator:
assert server_args.speculative_num_draft_tokens is not None assert server_args.speculative_num_draft_tokens is not None
assert server_args.max_running_requests is not None assert server_args.max_running_requests is not None
if get_schedule().max_mamba_cache_size is not None: if server_args.max_mamba_cache_size is not None:
# Use explicitly set max_mamba_cache_size # Use explicitly set max_mamba_cache_size
get_context().override( server_args.override(
"mamba_pool.per_dp_shard", "mamba_pool.per_dp_shard",
max_mamba_cache_size=get_schedule().max_mamba_cache_size max_mamba_cache_size=server_args.max_mamba_cache_size
// self.ps.attn_dp_size, // self.ps.attn_dp_size,
) )
# Reserve intermediate memory based on capped max_num_reqs # Reserve intermediate memory based on capped max_num_reqs
@@ -1717,7 +1722,7 @@ class KVCacheConfigurator:
ratio = self._calculate_mamba_ratio() ratio = self._calculate_mamba_ratio()
capped_reqs = min( capped_reqs = min(
server_args.max_running_requests // self.ps.attn_dp_size, server_args.max_running_requests // self.ps.attn_dp_size,
get_schedule().max_mamba_cache_size // ratio, server_args.max_mamba_cache_size // ratio,
) )
intermediate_size = ( intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req config.mamba2_cache_params.mamba_cache_per_req
@@ -1730,7 +1735,7 @@ class KVCacheConfigurator:
and server_args.max_running_requests is not None and server_args.max_running_requests is not None
): ):
# Use explicitly set max_running_requests when radix cache is disabled # Use explicitly set max_running_requests when radix cache is disabled
get_context().override( server_args.override(
"mamba_pool.from_max_running_requests", "mamba_pool.from_max_running_requests",
max_mamba_cache_size=server_args.max_running_requests max_mamba_cache_size=server_args.max_running_requests
// self.ps.attn_dp_size, // self.ps.attn_dp_size,
@@ -1739,7 +1744,7 @@ class KVCacheConfigurator:
if has_spec_dec: if has_spec_dec:
intermediate_size = ( intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req config.mamba2_cache_params.mamba_cache_per_req
* get_schedule().max_mamba_cache_size * server_args.max_mamba_cache_size
* server_args.speculative_num_draft_tokens * server_args.speculative_num_draft_tokens
) )
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
@@ -1764,7 +1769,7 @@ class KVCacheConfigurator:
ratio = self._calculate_mamba_ratio() ratio = self._calculate_mamba_ratio()
D = server_args.speculative_num_draft_tokens D = server_args.speculative_num_draft_tokens
# Joint solve: main_state + intermediate = mamba_budget # Joint solve: main_state + intermediate = mamba_budget
get_context().override( server_args.override(
"mamba_pool.memory_budget_spec", "mamba_pool.memory_budget_spec",
max_mamba_cache_size=int( max_mamba_cache_size=int(
mamba_budget_bytes // (per_req * (1 + D / ratio)) mamba_budget_bytes // (per_req * (1 + D / ratio))
@@ -1774,12 +1779,12 @@ class KVCacheConfigurator:
# so the return value only has main_state subtracted from total # so the return value only has main_state subtracted from total
capped_reqs = min( capped_reqs = min(
server_args.max_running_requests // self.ps.attn_dp_size, server_args.max_running_requests // self.ps.attn_dp_size,
get_schedule().max_mamba_cache_size // ratio, server_args.max_mamba_cache_size // ratio,
) )
intermediate_size = per_req * capped_reqs * D intermediate_size = per_req * capped_reqs * D
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
else: else:
get_context().override( server_args.override(
"mamba_pool.memory_budget", "mamba_pool.memory_budget",
max_mamba_cache_size=int(mamba_budget_bytes // per_req), max_mamba_cache_size=int(mamba_budget_bytes // per_req),
) )
@@ -1788,10 +1793,10 @@ class KVCacheConfigurator:
# A non-positive value means GPU memory is insufficient for the requested # A non-positive value means GPU memory is insufficient for the requested
# configuration. Fail fast with actionable advice instead of silently # configuration. Fail fast with actionable advice instead of silently
# producing garbled output at runtime. # producing garbled output at runtime.
if get_schedule().max_mamba_cache_size <= 0: if server_args.max_mamba_cache_size <= 0:
raise RuntimeError( raise RuntimeError(
f"Not enough GPU memory for hybrid (mamba/linear-attention) state cache. " f"Not enough GPU memory for hybrid (mamba/linear-attention) state cache. "
f"Computed max_mamba_cache_size={get_schedule().max_mamba_cache_size} " f"Computed max_mamba_cache_size={server_args.max_mamba_cache_size} "
f"(total_rest_memory={total_rest_memory:.2f} GB, " f"(total_rest_memory={total_rest_memory:.2f} GB, "
f"mamba_cache_per_req={config.mamba2_cache_params.mamba_cache_per_req / (1 << 20):.2f} MB). " f"mamba_cache_per_req={config.mamba2_cache_params.mamba_cache_per_req / (1 << 20):.2f} MB). "
f"Try: (1) reduce --max-running-requests, " f"Try: (1) reduce --max-running-requests, "
@@ -1801,7 +1806,7 @@ class KVCacheConfigurator:
) )
mamba_state_memory = ( mamba_state_memory = (
get_schedule().max_mamba_cache_size server_args.max_mamba_cache_size
* config.mamba2_cache_params.mamba_cache_per_req * config.mamba2_cache_params.mamba_cache_per_req
/ (1 << 30) / (1 << 30)
) )
@@ -16,7 +16,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchResult, MatchResult,
) )
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
from sglang.srt.runtime_context import get_memory, get_server_args from sglang.srt.runtime_context import get_server_args
try: try:
from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector
@@ -108,7 +108,7 @@ class LMCRadixCache(RadixCache):
): ):
super().__init__(params) super().__init__(params)
cli_lmc_cfg = get_memory().lmcache_config_file or "" cli_lmc_cfg = get_server_args().lmcache_config_file or ""
kvcache = self.token_to_kv_pool_allocator.get_kvcache() kvcache = self.token_to_kv_pool_allocator.get_kvcache()
connector_kwargs = dict( connector_kwargs = dict(
@@ -51,8 +51,13 @@ from sglang.srt.layers.dp_attention import (
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import ( from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
ForwardBatchDeepSeekMHAMixin, ForwardBatchDeepSeekMHAMixin,
) )
from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import is_cuda, is_hip, is_npu, support_triton from sglang.srt.utils import (
is_cuda,
is_hip,
is_npu,
support_triton,
)
from sglang.srt.utils.common import ceil_align, is_pin_memory_available from sglang.srt.utils.common import ceil_align, is_pin_memory_available
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -936,7 +941,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# --enable-mis: every request must carry delimiter indices (the score # --enable-mis: every request must carry delimiter indices (the score
# endpoint always produces MIS-structured requests; consumers index # endpoint always produces MIS-structured requests; consumers index
# without None-checking). # without None-checking).
if get_exec().features.enable_mis and any( if get_server_args().enable_mis and any(
r.multi_item_delimiter_indices is not None for r in batch.reqs r.multi_item_delimiter_indices is not None for r in batch.reqs
): ):
assert all( assert all(
@@ -1105,7 +1110,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# batch_size * [3 * seq_len] # batch_size * [3 * seq_len]
batch_size = self.seq_lens_cpu.shape[0] batch_size = self.seq_lens_cpu.shape[0]
mrope_positions_list = [[]] * batch_size mrope_positions_list = [[]] * batch_size
rl_on_policy_target = get_exec().deterministic.rl_on_policy_target rl_on_policy_target = get_server_args().rl_on_policy_target
for batch_idx in range(batch_size): for batch_idx in range(batch_size):
mm_input = batch.multimodal_inputs[batch_idx] mm_input = batch.multimodal_inputs[batch_idx]
if self.forward_mode.is_decode(): if self.forward_mode.is_decode():
@@ -26,7 +26,11 @@ import torch
import torch.distributed as dist import torch.distributed as dist
from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.load_config import LoadConfig
from sglang.srt.configs.model_config import AttentionArch, ModelConfig, ModelImpl from sglang.srt.configs.model_config import (
AttentionArch,
ModelConfig,
ModelImpl,
)
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
from sglang.srt.debug_utils.dumper import dumper from sglang.srt.debug_utils.dumper import dumper
from sglang.srt.distributed import bootstrap from sglang.srt.distributed import bootstrap
@@ -70,7 +74,9 @@ from sglang.srt.kv_canary.runner.canary_manager import context_tuple
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
from sglang.srt.layers import deep_gemm_wrapper, model_parallel from sglang.srt.layers import deep_gemm_wrapper, model_parallel
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.cp.utils import get_cp_strategy from sglang.srt.layers.cp.utils import (
get_cp_strategy,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.sampler import create_sampler from sglang.srt.layers.sampler import create_sampler
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
@@ -80,10 +86,17 @@ from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
from sglang.srt.mem_cache import kv_cache_dtype from sglang.srt.mem_cache import kv_cache_dtype
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator from sglang.srt.mem_cache.kv_cache_configurator import (
KVCacheConfigurator,
)
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled from sglang.srt.model_executor.cuda_graph_config import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors cuda_graph_fully_disabled,
)
from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch,
PPProxyTensors,
)
from sglang.srt.model_executor.forward_context import ( from sglang.srt.model_executor.forward_context import (
ForwardContext, ForwardContext,
forward_context, forward_context,
@@ -142,16 +155,14 @@ from sglang.srt.model_executor.model_runner_components.weight_updater import (
WeightUpdater, WeightUpdater,
) )
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.model_executor.runner import EagerRunner, get_batch_sizes_to_capture from sglang.srt.model_executor.runner import (
EagerRunner,
get_batch_sizes_to_capture,
)
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_device,
get_exec,
get_global_dwdp_manager, get_global_dwdp_manager,
get_lora, get_server_args,
get_model,
get_parallel,
get_schedule,
set_global_dwdp_manager, set_global_dwdp_manager,
) )
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
@@ -308,7 +319,7 @@ class ModelRunner:
self.init_threads_binding() self.init_threads_binding()
# Set float32 matmul precision # Set float32 matmul precision
if get_exec().features.enable_tf32_matmul: if get_server_args().enable_tf32_matmul:
torch.set_float32_matmul_precision("high") torch.set_float32_matmul_precision("high")
# Set device early so that TransferEngine init (e.g. Ascend NPU) # Set device early so that TransferEngine init (e.g. Ascend NPU)
@@ -385,20 +396,20 @@ class ModelRunner:
def _initialize_elastic_ep_joiner(self) -> None: def _initialize_elastic_ep_joiner(self) -> None:
if not ( if not (
get_exec().moe.elastic_ep_backend is not None self.server_args.elastic_ep_backend is not None
and self.server_args.is_ep_joiner and self.server_args.is_ep_joiner
): ):
return return
is_scale_join = get_exec().moe.ep_join_mode == "scale" is_scale_join = self.server_args.ep_join_mode == "scale"
if is_scale_join: if is_scale_join:
join_effective_ep_size = ( join_effective_ep_size = (
get_parallel().ep_join_rank_offset + self.ps.tp_size self.server_args.ep_join_rank_offset + self.ps.tp_size
) )
dist.barrier(group=self.tp_group.cpu_group) dist.barrier(group=self.tp_group.cpu_group)
if self.ps.tp_rank == 0: if self.ps.tp_rank == 0:
register_scale_cohort( register_scale_cohort(
get_parallel().ep_join_rank_offset, self.server_args.ep_join_rank_offset,
join_effective_ep_size, join_effective_ep_size,
) )
join_scale_process_group() join_scale_process_group()
@@ -408,7 +419,7 @@ class ModelRunner:
else: else:
join_process_groups() join_process_groups()
global_ep_rank = self.ps.tp_rank + get_parallel().ep_join_rank_offset global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset
broadcast_global_expert_location_metadata( broadcast_global_expert_location_metadata(
model_config=self.model_config, model_config=self.model_config,
moe_ep_rank=global_ep_rank, moe_ep_rank=global_ep_rank,
@@ -442,9 +453,9 @@ class ModelRunner:
new_dp_size=join_effective_ep_size, new_dp_size=join_effective_ep_size,
new_dp_rank=global_ep_rank, new_dp_rank=global_ep_rank,
) )
from sglang.srt.runtime_context import get_context self.server_args.override(
"elastic_ep.scale_join", dp_size=join_effective_ep_size
get_context().override("elastic_ep.scale_join", dp_size=join_effective_ep_size) )
if self.eplb_manager is not None: if self.eplb_manager is not None:
self.eplb_manager.disable_rebalance( self.eplb_manager.disable_rebalance(
"EPLB rebalance is disabled after elastic EP scale-up" "EPLB rebalance is disabled after elastic EP scale-up"
@@ -473,7 +484,7 @@ class ModelRunner:
device=self.device, device=self.device,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
model_config=self.model_config, model_config=self.model_config,
custom_weight_loaders=get_model().custom_weight_loader, custom_weight_loaders=self.server_args.custom_weight_loader,
get_model=lambda: self.model, get_model=lambda: self.model,
update_model_fields=self.update_model_fields, update_model_fields=self.update_model_fields,
recapture_cuda_graph=self.init_decode_cuda_graph, recapture_cuda_graph=self.init_decode_cuda_graph,
@@ -550,7 +561,7 @@ class ModelRunner:
def init_mindspore_runner(self): def init_mindspore_runner(self):
# Init the mindspore runner # Init the mindspore runner
# for now, there is only some communication initialization work # for now, there is only some communication initialization work
if get_model().model_impl.lower() == ModelImpl.MINDSPORE and _is_npu: if self.server_args.model_impl.lower() == ModelImpl.MINDSPORE and _is_npu:
from sglang.srt.model_executor.mindspore_runner import init_ms_distributed from sglang.srt.model_executor.mindspore_runner import init_ms_distributed
init_ms_distributed( init_ms_distributed(
@@ -607,7 +618,7 @@ class ModelRunner:
def init_memory_saver_adapter(self): def init_memory_saver_adapter(self):
self.memory_saver_adapter = TorchMemorySaverAdapter.create( self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=get_exec().features.enable_memory_saver enable=self.server_args.enable_memory_saver
) )
def maybe_init_remote_instance_transfer_engine(self): def maybe_init_remote_instance_transfer_engine(self):
@@ -618,7 +629,7 @@ class ModelRunner:
if self.is_draft_worker: if self.is_draft_worker:
return return
expert_rank = self.ps.moe_ep_rank + ( expert_rank = self.ps.moe_ep_rank + (
get_parallel().ep_join_rank_offset self.server_args.ep_join_rank_offset
if self.server_args.is_ep_scale_joiner if self.server_args.is_ep_scale_joiner
else 0 else 0
) )
@@ -643,7 +654,7 @@ class ModelRunner:
) )
def maybe_init_lplb_solvers(self): def maybe_init_lplb_solvers(self):
if get_exec().moe.ep_dispatch_algorithm == "lp" and not self.is_draft_worker: if self.server_args.ep_dispatch_algorithm == "lp" and not self.is_draft_worker:
init_lplb_solvers(model_config=self.model_config) init_lplb_solvers(model_config=self.model_config)
def maybe_init_eplb_manager(self): def maybe_init_eplb_manager(self):
@@ -657,12 +668,12 @@ class ModelRunner:
get_expert_backup_client=lambda: self.expert_backup_client, get_expert_backup_client=lambda: self.expert_backup_client,
get_weight_updater=lambda: self.weight_updater, get_weight_updater=lambda: self.weight_updater,
) )
if get_exec().moe.enable_eplb and (not self.is_draft_worker) if self.server_args.enable_eplb and (not self.is_draft_worker)
else None else None
) )
def maybe_init_elastic_ep(self): def maybe_init_elastic_ep(self):
if get_exec().moe.elastic_ep_backend: if self.server_args.elastic_ep_backend:
ElasticEPStateManager.init(self.server_args) ElasticEPStateManager.init(self.server_args)
def init_token_oracle(self): def init_token_oracle(self):
@@ -681,8 +692,8 @@ class ModelRunner:
get_model=lambda: self.model, get_model=lambda: self.model,
) )
if ( if (
get_exec().moe.enable_elastic_expert_backup self.server_args.enable_elastic_expert_backup
and get_exec().moe.elastic_ep_backend is not None and self.server_args.elastic_ep_backend is not None
) )
else None else None
) )
@@ -691,17 +702,17 @@ class ModelRunner:
# In layered loading, torchao may have been applied # In layered loading, torchao may have been applied
torchao_applied = getattr(self.model, "torchao_applied", False) torchao_applied = getattr(self.model, "torchao_applied", False)
if not torchao_applied: if not torchao_applied:
apply_torchao_config_to_model(self.model, get_exec().graph.torchao_config) apply_torchao_config_to_model(self.model, get_server_args().torchao_config)
supports_torch_tp = getattr(self.model, "supports_torch_tp", False) supports_torch_tp = getattr(self.model, "supports_torch_tp", False)
if self.ps.tp_size > 1 and supports_torch_tp: if self.ps.tp_size > 1 and supports_torch_tp:
self.apply_torch_tp() self.apply_torch_tp()
def maybe_init_lora_manager(self): def maybe_init_lora_manager(self):
if get_lora().enable_lora: if self.server_args.enable_lora:
self.init_lora_manager() self.init_lora_manager()
def maybe_enable_batch_invariant_mode(self): def maybe_enable_batch_invariant_mode(self):
if get_exec().deterministic.enable_deterministic_inference: if self.server_args.enable_deterministic_inference:
from sglang.srt.batch_invariant_ops import enable_batch_invariant_mode from sglang.srt.batch_invariant_ops import enable_batch_invariant_mode
enable_batch_invariant_mode() enable_batch_invariant_mode()
@@ -798,7 +809,7 @@ class ModelRunner:
device=self.device, device=self.device,
tp_group=( tp_group=(
self.attention_tp_group.cpu_group self.attention_tp_group.cpu_group
if get_parallel().enable_dp_attention if self.server_args.enable_dp_attention
else self.tp_group.cpu_group else self.tp_group.cpu_group
), ),
host_to_device_ratio=hisparse_cfg.host_to_device_ratio, host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
@@ -962,7 +973,7 @@ class ModelRunner:
get_offloader().post_init() get_offloader().post_init()
# Register model for layerwise NVTX profiling if enabled # Register model for layerwise NVTX profiling if enabled
if get_exec().comm.enable_layerwise_nvtx_marker: if self.server_args.enable_layerwise_nvtx_marker:
pyt_hooks = PytHooks() pyt_hooks = PytHooks()
pyt_hooks.register_hooks(self.model, module_prefix="model") pyt_hooks.register_hooks(self.model, module_prefix="model")
@@ -1019,7 +1030,7 @@ class ModelRunner:
) )
dist_barrier_after_load( dist_barrier_after_load(
elastic_ep_backend=get_exec().moe.elastic_ep_backend, elastic_ep_backend=self.server_args.elastic_ep_backend,
tp_rank=self.ps.tp_rank, tp_rank=self.ps.tp_rank,
is_ep_scale_joiner=self.server_args.is_ep_scale_joiner, is_ep_scale_joiner=self.server_args.is_ep_scale_joiner,
) )
@@ -1039,16 +1050,16 @@ class ModelRunner:
self.lora_manager = LoRAManager( self.lora_manager = LoRAManager(
base_model=self.model, base_model=self.model,
base_hf_config=self.model_config.hf_config, base_hf_config=self.model_config.hf_config,
max_loras_per_batch=get_lora().max_loras_per_batch, max_loras_per_batch=self.server_args.max_loras_per_batch,
load_config=self.load_config, load_config=self.load_config,
dtype=self.dtype, dtype=self.dtype,
server_args=self.server_args, server_args=self.server_args,
lora_backend=get_lora().lora_backend, lora_backend=self.server_args.lora_backend,
tp_size=self.ps.tp_size, tp_size=self.ps.tp_size,
tp_rank=self.ps.tp_rank, tp_rank=self.ps.tp_rank,
max_lora_rank=get_lora().max_lora_rank, max_lora_rank=self.server_args.max_lora_rank,
target_modules=get_lora().lora_target_modules, target_modules=self.server_args.lora_target_modules,
lora_paths=get_lora().lora_paths, lora_paths=self.server_args.lora_paths,
) )
if not cuda_graph_fully_disabled(): if not cuda_graph_fully_disabled():
init_lora_cuda_graph_moe_buffers( init_lora_cuda_graph_moe_buffers(
@@ -1320,7 +1331,7 @@ class ModelRunner:
) )
output.expert_distribution_metrics = recorder_outputs.get("metrics") output.expert_distribution_metrics = recorder_outputs.get("metrics")
no_copy_to_cpu = not get_schedule().disable_overlap_schedule no_copy_to_cpu = not self.server_args.disable_overlap_schedule
if ( if (
not self.is_draft_worker not self.is_draft_worker
and (experts_capturer := get_global_experts_capturer()) is not None and (experts_capturer := get_global_experts_capturer()) is not None
@@ -1350,7 +1361,7 @@ class ModelRunner:
self.msprobe_debugger.stop() self.msprobe_debugger.stop()
self.msprobe_debugger.step() self.msprobe_debugger.step()
if get_exec().moe.elastic_ep_backend is not None: if self.server_args.elastic_ep_backend is not None:
self.maybe_join_ep_ranks() self.maybe_join_ep_ranks()
return output return output
@@ -1609,7 +1620,7 @@ class ModelRunner:
if added <= 0: if added <= 0:
return return
initial_ep_size = get_parallel().elastic_ep_initial_size initial_ep_size = self.server_args.elastic_ep_initial_size
assert initial_ep_size is not None assert initial_ep_size is not None
self.server_args.override("elastic_ep.scale", ep_size=effective_size) self.server_args.override("elastic_ep.scale", ep_size=effective_size)
@@ -1628,7 +1639,7 @@ class ModelRunner:
set_global_expert_location_metadata(new_metadata, allow_overwrite=True) set_global_expert_location_metadata(new_metadata, allow_overwrite=True)
def _elastic_global_rank(self) -> int: def _elastic_global_rank(self) -> int:
return self.ps.tp_rank + get_parallel().ep_join_rank_offset return self.ps.tp_rank + self.server_args.ep_join_rank_offset
def _report_elastic_scale_failure(self, error: str, effective_size: int) -> None: 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: if self.ps.tp_rank != 0 or self.server_args.is_ep_scale_joiner:
@@ -1705,9 +1716,7 @@ class ModelRunner:
new_dp_size=target_size, new_dp_size=target_size,
new_dp_rank=self._elastic_global_rank(), new_dp_rank=self._elastic_global_rank(),
) )
from sglang.srt.runtime_context import get_context self.server_args.override("elastic_ep.scale", dp_size=target_size)
get_context().override("elastic_ep.scale", dp_size=target_size)
ElasticEPStateManager.mark_syncing_new_world() ElasticEPStateManager.mark_syncing_new_world()
self._elastic_scale_ready_barrier( self._elastic_scale_ready_barrier(
@@ -1756,7 +1765,7 @@ class ModelRunner:
recovered = maybe_recover_ep_ranks( recovered = maybe_recover_ep_ranks(
tp_group=self.tp_group, tp_group=self.tp_group,
eplb_manager=self.eplb_manager, eplb_manager=self.eplb_manager,
random_seed=get_device().random_seed, random_seed=self.server_args.random_seed,
) )
if recovered: if recovered:
self.forward_pass_id = 0 self.forward_pass_id = 0
@@ -1765,7 +1774,7 @@ class ModelRunner:
local_timeout = ( local_timeout = (
state.pending_since is not None state.pending_since is not None
and time.monotonic() - state.pending_since and time.monotonic() - state.pending_since
> get_exec().moe.elastic_ep_scale_timeout > self.server_args.elastic_ep_scale_timeout
) )
timeout = state.active_ranks.new_tensor(int(local_timeout)) timeout = state.active_ranks.new_tensor(int(local_timeout))
dist.all_reduce(timeout, op=dist.ReduceOp.MAX, group=dist.group.WORLD) dist.all_reduce(timeout, op=dist.ReduceOp.MAX, group=dist.group.WORLD)
@@ -1833,9 +1842,7 @@ class ModelRunner:
load_config: LoadConfig, load_config: LoadConfig,
) -> None: ) -> None:
self.model = new_model self.model = new_model
from sglang.srt.runtime_context import get_context self.server_args.override(
get_context().override(
"model_runner.update_model_fields", "model_runner.update_model_fields",
model_path=model_path, model_path=model_path,
load_format=load_format, load_format=load_format,
@@ -24,12 +24,6 @@ def maybe_disable_chunked_prefix_cache(
# model's (often non-MLA) config must not flip the shared setting. # model's (often non-MLA) config must not flip the shared setting.
if is_draft_worker: if is_draft_worker:
return return
# This is a load-time gate that runs in ModelRunner.__init__ BEFORE the
# runner publishes its config (and direct/benchmark construction never
# publishes earlier), so read/write the supplied server_args. The runner's
# subsequent publish snapshots this into the schedule bag for get_schedule()
# readers.
if ( if (
not use_mla_backend not use_mla_backend
or server_args.attention_backend or server_args.attention_backend
@@ -11,7 +11,6 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend, RemoteInstanceWeightLoaderBackend,
register_memory_region, register_memory_region,
) )
from sglang.srt.runtime_context import get_model, get_parallel
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
@@ -59,7 +58,7 @@ class RemoteInstanceWeightTransporter:
# ModelExpress owns TransferEngine memory registration and metadata # ModelExpress owns TransferEngine memory registration and metadata
# publishing for backend=modelexpress. Re-registering here would # publishing for backend=modelexpress. Re-registering here would
# overlap the same weight buffers. # overlap the same weight buffers.
and get_model().remote_instance_weight_loader_backend and self.server_args.remote_instance_weight_loader_backend
!= RemoteInstanceWeightLoaderBackend.MODELEXPRESS != RemoteInstanceWeightLoaderBackend.MODELEXPRESS
and self.engine is not None and self.engine is not None
and self.weight_info is None and self.weight_info is None
@@ -76,16 +75,16 @@ class RemoteInstanceWeightTransporter:
""" """
import requests as http_requests import requests as http_requests
if get_parallel().dist_init_addr: if self.server_args.dist_init_addr:
# Multi-node: bootstrap server is on the head node (node_rank==0). # Multi-node: bootstrap server is on the head node (node_rank==0).
# Derive host from dist_init_addr (shared across all nodes). # Derive host from dist_init_addr (shared across all nodes).
bootstrap_host = ( bootstrap_host = (
NetworkAddress.parse(get_parallel().dist_init_addr).resolved().host NetworkAddress.parse(self.server_args.dist_init_addr).resolved().host
) )
else: else:
bootstrap_host = "127.0.0.1" bootstrap_host = "127.0.0.1"
bootstrap_port = get_model().engine_info_bootstrap_port bootstrap_port = self.server_args.engine_info_bootstrap_port
bootstrap_na = NetworkAddress(bootstrap_host, bootstrap_port) bootstrap_na = NetworkAddress(bootstrap_host, bootstrap_port)
url = f"{bootstrap_na.to_url()}/register_transfer_engine_info" url = f"{bootstrap_na.to_url()}/register_transfer_engine_info"
+6 -3
View File
@@ -44,7 +44,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
get_remote_instance_transfer_engine_info_per_rank, get_remote_instance_transfer_engine_info_per_rank,
register_memory_region, register_memory_region,
) )
from sglang.srt.runtime_context import get_exec, get_server_args from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import get_available_gpu_memory from sglang.srt.utils import get_available_gpu_memory
# Try to import accelerate (optional dependency) # Try to import accelerate (optional dependency)
@@ -71,7 +71,9 @@ from sglang.srt.connector import (
get_connector_type, get_connector_type,
) )
from sglang.srt.connector.utils import parse_model_name from sglang.srt.connector.utils import parse_model_name
from sglang.srt.distributed import model_parallel_is_initialized from sglang.srt.distributed import (
model_parallel_is_initialized,
)
from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
@@ -863,8 +865,9 @@ class LayeredModelLoader(DefaultModelLoader):
device_config: DeviceConfig, device_config: DeviceConfig,
) -> nn.Module: ) -> nn.Module:
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
from sglang.srt.runtime_context import get_server_args
torchao_config = get_exec().graph.torchao_config torchao_config = get_server_args().torchao_config
target_device = torch.device(device_config.device) target_device = torch.device(device_config.device)
quant_config = _get_quantization_config(model_config, self.load_config) quant_config = _get_quantization_config(model_config, self.load_config)
+5 -3
View File
@@ -26,7 +26,9 @@ import torch
from torch import nn from torch import nn
from transformers import ApertusConfig 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.activation import XIELU
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 (
@@ -50,7 +52,7 @@ from sglang.srt.model_loader.weight_utils import (
kv_cache_scales_loader, kv_cache_scales_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
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__)
@@ -440,7 +442,7 @@ class ApertusForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_parallel().enable_dp_lm_head, use_attn_tp_group=get_server_args().enable_dp_lm_head,
) )
self.logits_processor = LogitsProcessor(config) self.logits_processor = LogitsProcessor(config)
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
+5 -3
View File
@@ -20,7 +20,9 @@ import torch
from torch import nn from torch import nn
from transformers import LlamaConfig 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.activation import get_act_fn
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 (
@@ -44,7 +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.runtime_context import get_parallel, get_server_args
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__)
@@ -403,7 +405,7 @@ class ArceeForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_parallel().enable_dp_lm_head, use_attn_tp_group=get_server_args().enable_dp_lm_head,
) )
self.logits_processor = LogitsProcessor(config) self.logits_processor = LogitsProcessor(config)
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
+9 -5
View File
@@ -41,7 +41,9 @@ from sglang.srt.layers.communicator import (
LayerScatterModes, LayerScatterModes,
enable_moe_dense_fully_dp, 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.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
MergedColumnParallelLinear, MergedColumnParallelLinear,
@@ -76,9 +78,9 @@ from sglang.srt.models.utils import (
enable_fused_set_kv_buffer, enable_fused_set_kv_buffer,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec,
get_forward, get_forward,
get_parallel, get_parallel,
get_server_args,
get_stream, get_stream,
) )
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
@@ -207,7 +209,7 @@ class BailingMoESparseMoeBlock(nn.Module):
self.router_dtype = torch.bfloat16 self.router_dtype = torch.bfloat16
# TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now # TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now
assert get_exec().moe.ep_num_redundant_experts == 0 assert get_server_args().ep_num_redundant_experts == 0
# check group topk # check group topk
self.num_expert_group = getattr(config, "n_group", 0) self.num_expert_group = getattr(config, "n_group", 0)
self.topk_group = getattr(config, "topk_group", 0) self.topk_group = getattr(config, "topk_group", 0)
@@ -221,7 +223,9 @@ class BailingMoESparseMoeBlock(nn.Module):
self.num_expert_group = self.topk_group = None self.num_expert_group = self.topk_group = None
self.use_grouped_topk = False self.use_grouped_topk = False
self.num_experts = config.num_experts + get_exec().moe.ep_num_redundant_experts self.num_experts = (
config.num_experts + get_server_args().ep_num_redundant_experts
)
self.gate = BailingMoEGate( self.gate = BailingMoEGate(
config=config, config=config,
@@ -820,7 +824,7 @@ class BailingMoEForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_parallel().enable_dp_lm_head, use_attn_tp_group=get_server_args().enable_dp_lm_head,
) )
self.logits_processor = LogitsProcessor(config) self.logits_processor = LogitsProcessor(config)
+11 -6
View File
@@ -12,12 +12,17 @@ from transformers import PretrainedConfig
from sglang.kernels.ops.attention.fla.layernorm_gated import RMSNorm as RMSNormGated from sglang.kernels.ops.attention.fla.layernorm_gated import RMSNorm as RMSNormGated
from sglang.kernels.ops.attention.fla.layernorm_gated import layernorm_fn from sglang.kernels.ops.attention.fla.layernorm_gated import layernorm_fn
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.distributed import (
get_pp_group,
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.layers import deep_gemm_wrapper from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes 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.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
@@ -54,9 +59,9 @@ 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 ( from sglang.srt.runtime_context import (
get_device,
get_forward, get_forward,
get_parallel, get_parallel,
get_server_args,
get_stream, get_stream,
) )
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -524,7 +529,7 @@ class BailingMoELinearAttention(nn.Module):
base=self.rope_theta, base=self.rope_theta,
rope_scaling=config.rope_scaling, rope_scaling=config.rope_scaling,
is_neox_style=True, is_neox_style=True,
device=get_device().device, device=get_server_args().device,
dtype=torch.float32, dtype=torch.float32,
) )
@@ -685,7 +690,7 @@ class BailingMoEAttention(nn.Module):
max_position=self.max_position_embeddings, max_position=self.max_position_embeddings,
base=self.rope_theta, base=self.rope_theta,
rope_scaling=config.rope_scaling, rope_scaling=config.rope_scaling,
device=get_device().device, device=get_server_args().device,
) )
self.attn = RadixAttention( self.attn = RadixAttention(
self.num_heads, self.num_heads,
@@ -1084,7 +1089,7 @@ class BailingMoELinearForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
params_dtype=torch.float32, params_dtype=torch.float32,
quant_config=quant_config, quant_config=quant_config,
use_attn_tp_group=get_parallel().enable_dp_lm_head, use_attn_tp_group=get_server_args().enable_dp_lm_head,
) )
) )
self.logits_processor = LogitsProcessor(config) self.logits_processor = LogitsProcessor(config)
@@ -42,7 +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.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import BumpAllocator, add_prefix from sglang.srt.utils import BumpAllocator, add_prefix
LoraConfig = None LoraConfig = None
@@ -208,7 +208,7 @@ class BailingMoeForCausalLMNextN(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix), prefix=add_prefix("model.shared_head.head", prefix),
use_attn_tp_group=get_parallel().enable_dp_lm_head, use_attn_tp_group=get_server_args().enable_dp_lm_head,
) )
self.logits_processor = LogitsProcessor(config) self.logits_processor = LogitsProcessor(config)
if hasattr(self.config, "model_type") and config.model_type == "bailing_hybrid": if hasattr(self.config, "model_type") and config.model_type == "bailing_hybrid":
+4 -2
View File
@@ -16,7 +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_model, get_parallel from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
BertConfig = None BertConfig = None
@@ -365,7 +365,9 @@ class BertModel(nn.Module):
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("encoder", prefix), prefix=add_prefix("encoder", prefix),
) )
pooling_type = PoolingType.CLS if get_model().is_embedding else PoolingType.LAST pooling_type = (
PoolingType.CLS if get_server_args().is_embedding else PoolingType.LAST
)
self.pooler = ( self.pooler = (
BertPooler(config) BertPooler(config)
if self.use_bert_pooler if self.use_bert_pooler
@@ -11,7 +11,7 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods
AttnForwardMethod, AttnForwardMethod,
) )
from sglang.srt.models.deepseek_common.utils import _is_hip from sglang.srt.models.deepseek_common.utils import _is_hip
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_server_args
from sglang.srt.utils import use_intel_amx_backend from sglang.srt.utils import use_intel_amx_backend
MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"] MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"]
@@ -118,7 +118,7 @@ def handle_attention_flashinfer(attn, forward_batch):
def handle_attention_fa3(attn, forward_batch): def handle_attention_fa3(attn, forward_batch):
# when deterministic inference is enabled, use MLA # when deterministic inference is enabled, use MLA
if get_exec().deterministic.enable_deterministic_inference: if get_server_args().enable_deterministic_inference:
return _dispatch_mla_subtype(attn, forward_batch) return _dispatch_mla_subtype(attn, forward_batch)
else: else:
return _handle_attention_backend(attn, forward_batch, "fa3") return _handle_attention_backend(attn, forward_batch, "fa3")
@@ -187,7 +187,7 @@ def handle_attention_triton(attn, forward_batch):
return AttnForwardMethod.MLA return AttnForwardMethod.MLA
# when deterministic inference is enabled, use MLA # when deterministic inference is enabled, use MLA
if get_exec().deterministic.enable_deterministic_inference: if get_server_args().enable_deterministic_inference:
return _dispatch_mla_subtype(attn, forward_batch) return _dispatch_mla_subtype(attn, forward_batch)
if ( if (
@@ -30,11 +30,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 ( from sglang.srt.runtime_context import get_parallel, get_server_args
get_exec,
get_parallel,
get_schedule,
)
from sglang.srt.utils import BumpAllocator, get_bool_env_var, next_power_of_2 from sglang.srt.utils import BumpAllocator, get_bool_env_var, next_power_of_2
_use_fp8_prefill_attn = ( _use_fp8_prefill_attn = (
@@ -146,7 +142,9 @@ def _forward_dsa_indexer_for_mha(
class DeepseekMHAForwardMixin: class DeepseekMHAForwardMixin:
def init_mha_forward(self: DeepseekV2AttentionMLA): def init_mha_forward(self: DeepseekV2AttentionMLA):
self.disable_chunked_prefix_cache = get_schedule().disable_chunked_prefix_cache self.disable_chunked_prefix_cache = (
get_server_args().disable_chunked_prefix_cache
)
# TODO: Design a finer way to determine the threshold # TODO: Design a finer way to determine the threshold
self.chunked_prefix_cache_threshold = ( self.chunked_prefix_cache_threshold = (
@@ -307,8 +305,8 @@ class DeepseekMHAForwardMixin:
self.use_dsa self.use_dsa
and self.kv_cache_dtype == "fp8_e4m3" and self.kv_cache_dtype == "fp8_e4m3"
and ( and (
not get_exec().kernel.dsa_decode_backend == "trtllm" not get_server_args().dsa_decode_backend == "trtllm"
or not get_exec().kernel.dsa_prefill_backend == "trtllm" or not get_server_args().dsa_prefill_backend == "trtllm"
) )
): ):
# FP8 path: dequantize DSA-specific FP8 format to BF16 # FP8 path: dequantize DSA-specific FP8 format to BF16
@@ -65,8 +65,10 @@ 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_exec, get_parallel, get_server_args from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.state_capturer.indexer_topk import maybe_capture_indexer_topk from sglang.srt.state_capturer.indexer_topk import (
maybe_capture_indexer_topk,
)
from sglang.srt.utils import BumpAllocator from sglang.srt.utils import BumpAllocator
from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.custom_op import register_custom_op
@@ -151,7 +153,7 @@ def _should_defer_dsa_cp_kv_gather(
class DeepseekMLAForwardMixin: class DeepseekMLAForwardMixin:
def init_mla_forward(self: DeepseekV2AttentionMLA): def init_mla_forward(self: DeepseekV2AttentionMLA):
self.flashinfer_mla_disable_ragged = ( self.flashinfer_mla_disable_ragged = (
get_exec().kernel.flashinfer_mla_disable_ragged get_server_args().flashinfer_mla_disable_ragged
) )
def should_run_indexer( def should_run_indexer(
@@ -988,8 +990,8 @@ class DeepseekMLAForwardMixin:
""" """
if self.current_attention_backend in ("dsa", "nsa"): if self.current_attention_backend in ("dsa", "nsa"):
return ( return (
get_exec().kernel.dsa_decode_backend == "trtllm" get_server_args().dsa_decode_backend == "trtllm"
or get_exec().kernel.dsa_prefill_backend == "trtllm" or get_server_args().dsa_prefill_backend == "trtllm"
) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn ) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn
return ( return (
+6 -9
View File
@@ -59,11 +59,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 ( from sglang.srt.runtime_context import get_parallel, get_server_args
get_model,
get_parallel,
get_spec,
)
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
@@ -152,7 +148,7 @@ class DeepseekModelNextN(nn.Module):
self.rot_weight = None self.rot_weight = None
if _is_npu: if _is_npu:
rot_weight_path = get_model().model_path + "/rot.safetensors" rot_weight_path = get_server_args().model_path + "/rot.safetensors"
if os.path.isfile(rot_weight_path): if os.path.isfile(rot_weight_path):
self.rot_weight = load_file(rot_weight_path) self.rot_weight = load_file(rot_weight_path)
self.rot_weight = self.rot_weight["rot.weight"].npu() self.rot_weight = self.rot_weight["rot.weight"].npu()
@@ -165,7 +161,8 @@ class DeepseekModelNextN(nn.Module):
layer_name = "decoder" layer_name = "decoder"
if _is_npu and ( if _is_npu and (
get_spec().speculative_draft_model_path == get_model().model_path get_server_args().speculative_draft_model_path
== get_server_args().model_path
): ):
layer_name = "layers." + str(config.num_hidden_layers) layer_name = "layers." + str(config.num_hidden_layers)
@@ -204,7 +201,7 @@ class DeepseekModelNextN(nn.Module):
if ( if (
_is_npu _is_npu
and self.quant_config is None and self.quant_config is None
and get_model().quantization is not None and get_server_args().quantization is not None
): ):
# ascend mtp unquant # ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
@@ -380,7 +377,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix), prefix=add_prefix("model.shared_head.head", prefix),
use_attn_tp_group=get_parallel().enable_dp_lm_head, use_attn_tp_group=get_server_args().enable_dp_lm_head,
) )
self.logits_processor = LogitsProcessor(config) self.logits_processor = LogitsProcessor(config)

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