config: read resolved config via namespace accessors (#33013)

This commit is contained in:
Cheng Wan
2026-07-31 15:06:59 -07:00
committed by GitHub
parent 4862edc85f
commit 55b6769b0e
187 changed files with 1110 additions and 923 deletions
+6 -11
View File
@@ -261,19 +261,14 @@ 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) on the published ``server_args``: resolution has weight-resolved dtypes): validated against the resolvable whitelist, then
already materialized, so the declaration writes through, joining the written to the config bags via ``get_context().override``; ``server_args``
declaration stash for provenance and republish consistency.""" stays the pristine startup record."""
from sglang.srt.runtime_context import get_context from sglang.srt.runtime_context import get_context
server_args = get_context().server_args context = get_context()
validate_declarations(server_args, [(source, dict(declared))]) validate_declarations(context.server_args, [(source, dict(declared))])
override = getattr(server_args, "override", None) context.override(source, **declared)
if override is not None:
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(
@@ -40,7 +40,7 @@ from sglang.srt.model_executor.forward_batch_info import (
compute_position, compute_position,
) )
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_device, get_parallel, get_server_args
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
@@ -184,7 +184,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_server_args().device, non_blocking=True) ).to(device=get_device().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:
@@ -336,7 +336,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_server_args().device (2,), dtype=torch.int32, device=get_device().device
) )
def capture_one_batch_size(self, batch: ForwardBatch, num_tokens: int): def capture_one_batch_size(self, batch: ForwardBatch, num_tokens: int):
@@ -835,7 +835,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_server_args().device, non_blocking=True device=get_device().device, non_blocking=True
) )
@classmethod @classmethod
+2 -2
View File
@@ -8,6 +8,7 @@ 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):
@@ -224,9 +225,8 @@ 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_server_args().enable_scattered_sconv: if get_exec().comm.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 (
+2 -1
View File
@@ -23,6 +23,7 @@ import torch.distributed as dist
import zmq import zmq
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
from sglang.srt.runtime_context import get_serving
# -------------------------------------- config base ------------------------------------------ # -------------------------------------- config base ------------------------------------------
@@ -1798,7 +1799,7 @@ class _SGLangPlugin(_FrameworkPlugin):
if args is None: if args is None:
return None return None
return args.tokenizer_path return get_serving().tokenizer_path
except Exception: except Exception:
return None return None
@@ -41,6 +41,7 @@ class KVArgs:
kv_data_lens: List[int] kv_data_lens: List[int]
kv_item_lens: List[int] kv_item_lens: List[int]
kv_layer_ids: List[int] kv_layer_ids: List[int]
kv_cache_dtype_str: str
aux_data_ptrs: List[int] aux_data_ptrs: List[int]
aux_data_lens: List[int] aux_data_lens: List[int]
aux_item_lens: List[int] aux_item_lens: List[int]
@@ -36,7 +36,7 @@ from sglang.srt.layers.dp_attention import (
get_attention_dp_rank, get_attention_dp_rank,
get_attention_dp_size, get_attention_dp_size,
) )
from sglang.srt.runtime_context import get_model, get_parallel from sglang.srt.runtime_context import get_parallel, get_serving
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,
@@ -148,6 +148,7 @@ class CommonKVManager(BaseKVManager):
is_mla_backend: Optional[bool] = False, is_mla_backend: Optional[bool] = False,
): ):
self.kv_args = args self.kv_args = args
self.kv_cache_dtype_str = args.kv_cache_dtype_str
self.kv_item_lens_sum = sum(args.kv_item_lens) self.kv_item_lens_sum = sum(args.kv_item_lens)
self.state_item_lens_sum = sum(x for comp in args.state_item_lens for x in comp) self.state_item_lens_sum = sum(x for comp in args.state_item_lens for x in comp)
self.is_mla_backend = is_mla_backend self.is_mla_backend = is_mla_backend
@@ -533,11 +534,11 @@ class CommonKVManager(BaseKVManager):
if ( if (
info.kv_cache_dtype is not None info.kv_cache_dtype is not None
and info.kv_cache_dtype != get_model().kv_cache_dtype and info.kv_cache_dtype != self.kv_cache_dtype_str
): ):
raise RuntimeError( raise RuntimeError(
f"KV cache dtype mismatch: prefill server has kv_cache_dtype={info.kv_cache_dtype}, " f"KV cache dtype mismatch: prefill server has kv_cache_dtype={info.kv_cache_dtype}, "
f"but decode server has kv_cache_dtype={get_model().kv_cache_dtype}. " f"but decode server has kv_cache_dtype={self.kv_cache_dtype_str}. "
f"Both servers must use the same --kv-cache-dtype value." f"Both servers must use the same --kv-cache-dtype value."
) )
@@ -701,7 +702,7 @@ class CommonKVManager(BaseKVManager):
"rank_ip": self.local_ip, "rank_ip": self.local_ip,
"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": self.kv_cache_dtype_str,
"load_balance_method": self.server_args.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
@@ -709,7 +710,7 @@ class CommonKVManager(BaseKVManager):
# 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": self.server_args.port, "prefill_http_port": get_serving().port,
} }
max_retries, initial_delay, max_delay = 5, 1.0, 30.0 max_retries, initial_delay, max_delay = 5, 1.0, 30.0
+7 -6
View File
@@ -87,7 +87,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_parallel from sglang.srt.runtime_context import get_disagg, get_parallel
from sglang.srt.utils import get_num_new_pages, is_npu from sglang.srt.utils import get_num_new_pages, is_npu
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
@@ -423,6 +423,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
kv_args.pp_rank = self.pp_rank kv_args.pp_rank = self.pp_rank
kv_args.system_dp_rank = self.scheduler.ps.dp_rank kv_args.system_dp_rank = self.scheduler.ps.dp_rank
kv_args.kv_cache_dtype_str = (
self.scheduler.tp_worker.model_runner.kv_cache_dtype_str
)
transfer_kv_pool = ( transfer_kv_pool = (
self.scheduler.hisparse_coordinator.mem_pool_host self.scheduler.hisparse_coordinator.mem_pool_host
if self.scheduler.enable_hisparse if self.scheduler.enable_hisparse
@@ -2244,7 +2247,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 self.server_args.disaggregation_decode_enable_radix_cache: if get_disagg().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
@@ -2284,7 +2287,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 self.server_args.disaggregation_decode_enable_offload_kvcache: if get_disagg().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
@@ -2296,9 +2299,7 @@ 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 = ( self.polling_interval = get_disagg().disaggregation_decode_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,6 +28,7 @@ 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
@@ -117,13 +118,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 self.server_args.encoder_transfer_backend == "mooncake": if get_disagg().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 self.server_args.encoder_transfer_backend == "zmq_to_scheduler": elif get_disagg().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:
@@ -141,7 +142,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 self.server_args.encoder_transfer_backend == "zmq_to_tokenizer": elif get_disagg().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
) )
@@ -63,7 +63,7 @@ from sglang.srt.observability.trace import (
process_tracing_init, process_tracing_init,
trace_set_thread_info, trace_set_thread_info,
) )
from sglang.srt.runtime_context import publish from sglang.srt.runtime_context import get_disagg, get_exec, get_mm, publish
from sglang.srt.server_args import ( from sglang.srt.server_args import (
PortArgs, PortArgs,
ServerArgs, ServerArgs,
@@ -352,7 +352,7 @@ class MMEncoder:
[], dtype=self._embedding_dtype [], dtype=self._embedding_dtype
).element_size() ).element_size()
if self.server_args.enable_mm_global_cache: if get_mm().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,
) )
@@ -370,15 +370,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 self.server_args.encoder_transfer_backend == "mooncake": if get_disagg().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: {self.server_args.encoder_transfer_backend}" f"Using transfer backend: {get_disagg().encoder_transfer_backend}"
) )
if self.server_args.encoder_transfer_backend == "mooncake": if get_disagg().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()
@@ -391,8 +391,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=(
self.server_args.disaggregation_ib_device get_disagg().disaggregation_ib_device
or self.server_args.mooncake_ib_device or get_exec().moe.mooncake_ib_device
), ),
) )
@@ -401,7 +401,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 self.server_args.encoder_transfer_backend == "mooncake": if get_disagg().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
@@ -415,12 +415,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 self.server_args.encoder_transfer_backend == "mooncake": if get_disagg().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 self.server_args.encoder_transfer_backend == "mooncake": if get_disagg().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
@@ -1710,7 +1710,7 @@ class MMEncoder:
mm_item.set(k, _convert(v)) mm_item.set(k, _convert(v))
cache_hit = False cache_hit = False
use_mm_cache = self.server_args.enable_prefix_mm_cache and log_metrics use_mm_cache = get_mm().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])
@@ -1806,7 +1806,7 @@ class MMEncoder:
embedding_port=None, embedding_port=None,
url=None, url=None,
): ):
if self.server_args.encoder_transfer_backend == "mooncake": if get_disagg().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:
@@ -1878,7 +1878,7 @@ class MMEncoder:
logger.info(f"{endpoint = }") logger.info(f"{endpoint = }")
# Serialize data # Serialize data
if self.server_args.encoder_transfer_backend == "mooncake": if get_disagg().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)
@@ -1910,11 +1910,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 self.server_args.encoder_transfer_backend != "mooncake" and get_disagg().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=self.server_args.encoder_transfer_backend, backend=get_disagg().encoder_transfer_backend,
) )
async def encode( async def encode(
@@ -60,6 +60,7 @@ from sglang.srt.observability.trace import (
TraceReqContext, TraceReqContext,
trace_set_thread_info, trace_set_thread_info,
) )
from sglang.srt.runtime_context import 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
@@ -327,7 +328,7 @@ class MooncakeKVManager(CommonKVManager):
lambda ptr, size: self.engine.batch_register([ptr], [size]), lambda ptr, size: self.engine.batch_register([ptr], [size]),
self.kv_args, self.kv_args,
count, count,
self.server_args.chunked_prefill_size, get_schedule().chunked_prefill_size,
) )
self.kv_buffer_tensors = None self.kv_buffer_tensors = None
@@ -498,7 +499,7 @@ class MooncakeKVManager(CommonKVManager):
room, room,
self.transfer_infos, self.transfer_infos,
self.kv_buffer_tensors, self.kv_buffer_tensors,
self.server_args.chunked_prefill_size, get_schedule().chunked_prefill_size,
self._staging_ctx.prefetch_requested, self._staging_ctx.prefetch_requested,
self._staging_ctx.prefetch_sockets, self._staging_ctx.prefetch_sockets,
) )
@@ -44,6 +44,7 @@ from sglang.srt.disaggregation.utils import (
resolve_dcp_dst_entry_indices, resolve_dcp_dst_entry_indices,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_schedule
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
try: try:
@@ -538,7 +539,7 @@ class NixlKVManager(CommonKVManager):
lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id), lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id),
self.kv_args, self.kv_args,
count, count,
self.server_args.chunked_prefill_size, get_schedule().chunked_prefill_size,
) )
def _init_staging_allocator(self): def _init_staging_allocator(self):
@@ -670,7 +671,7 @@ class NixlKVManager(CommonKVManager):
room, room,
self.transfer_infos, self.transfer_infos,
self.kv_buffer_tensors, self.kv_buffer_tensors,
self.server_args.chunked_prefill_size, get_schedule().chunked_prefill_size,
self._staging_ctx.prefetch_requested, self._staging_ctx.prefetch_requested,
self._staging_ctx.prefetch_sockets, self._staging_ctx.prefetch_sockets,
) )
+5 -1
View File
@@ -65,6 +65,7 @@ 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 import is_npu from sglang.srt.utils import is_npu
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
@@ -183,6 +184,9 @@ class PrefillBootstrapQueue:
kv_args.engine_rank = self.tp_rank kv_args.engine_rank = self.tp_rank
kv_args.pp_rank = self.pp_rank kv_args.pp_rank = self.pp_rank
kv_args.system_dp_rank = self.scheduler.ps.dp_rank kv_args.system_dp_rank = self.scheduler.ps.dp_rank
kv_args.kv_cache_dtype_str = (
self.scheduler.tp_worker.model_runner.kv_cache_dtype_str
)
layer_shard_enabled = getattr( layer_shard_enabled = getattr(
self.token_to_kv_pool, "layer_shard_enabled", False self.token_to_kv_pool, "layer_shard_enabled", False
) )
@@ -1249,7 +1253,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 = self.server_args.optimistic_prefill_attempts max_attempts = get_disagg().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_server_args from sglang.srt.runtime_context import get_exec
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_server_args().enable_symm_mem return get_exec().comm.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_server_args from sglang.srt.runtime_context import get_exec
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_server_args().enable_symm_mem return get_exec().comm.enable_symm_mem
except ValueError: except ValueError:
return False return False
@@ -12,6 +12,7 @@ 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:
@@ -98,10 +99,9 @@ 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_server_args().enable_scattered_sconv get_exec().comm.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
+3 -2
View File
@@ -11,6 +11,7 @@ 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__)
@@ -22,7 +23,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 self.server_args.dllm_algorithm is not None if get_exec().dllm.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)
@@ -200,7 +201,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=self.server_args.prefill_max_requests, prefill_max_requests=get_schedule().prefill_max_requests,
dllm_config=self.dllm_config, dllm_config=self.dllm_config,
) )
@@ -14,6 +14,7 @@ from sglang.srt.distributed.parallel_state import (
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
@@ -111,7 +112,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
+ self.server_args.ep_num_redundant_experts + get_exec().moe.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):
+3 -3
View File
@@ -377,9 +377,9 @@ 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": self.server_args.tokenizer_path, "tokenizer_path": self.tokenizer_manager.server_args.tokenizer_path,
"is_generation": self.tokenizer_manager.is_generation, "is_generation": self.tokenizer_manager.is_generation,
"weight_version": self.server_args.weight_version, "weight_version": self.tokenizer_manager.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),
} }
@@ -432,7 +432,7 @@ class RuntimeHandle:
"max_model_len": self.tokenizer_manager.model_config.context_len, "max_model_len": self.tokenizer_manager.model_config.context_len,
} }
] ]
if self.server_args.enable_lora and hasattr( if self.tokenizer_manager.server_args.enable_lora and hasattr(
self.tokenizer_manager, "lora_registry" self.tokenizer_manager, "lora_registry"
): ):
lora_registry = self.tokenizer_manager.lora_registry lora_registry = self.tokenizer_manager.lora_registry
+3 -3
View File
@@ -18,7 +18,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_server_args from sglang.srt.runtime_context import get_model
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
@@ -343,8 +343,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_server_args().model_path, get_model().model_path,
get_server_args().load_format, get_model().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_server_args from sglang.srt.runtime_context import get_exec
@dataclass @dataclass
@@ -40,8 +40,7 @@ class ExpertLocationDispatchInfo:
@classmethod @classmethod
def init_new(cls, layer_id: int): def init_new(cls, layer_id: int):
server_args = get_server_args() ep_dispatch_algorithm = get_exec().moe.ep_dispatch_algorithm
ep_dispatch_algorithm = 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
@@ -50,7 +49,7 @@ class ExpertLocationDispatchInfo:
return cls( return cls(
ep_dispatch_algorithm=ep_dispatch_algorithm, ep_dispatch_algorithm=ep_dispatch_algorithm,
rank_invariant=server_args.moe_a2a_backend == "none", rank_invariant=get_exec().moe.moe_a2a_backend == "none",
partial_logical_to_rank_dispatch_physical_map=( partial_logical_to_rank_dispatch_physical_map=(
expert_location_metadata.logical_to_rank_dispatch_physical_map[ expert_location_metadata.logical_to_rank_dispatch_physical_map[
layer_id, : layer_id, :
@@ -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_server_args from sglang.srt.runtime_context import get_device
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__)
@@ -109,7 +109,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_server_args().device, non_blocking=True) .to(device=get_device().device, non_blocking=True)
) )
routed_experts_weights_of_layer[layer_id].append(canary_tensor) routed_experts_weights_of_layer[layer_id].append(canary_tensor)
@@ -19,6 +19,7 @@ from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.model_executor.model_runner_components.layer_setup import ( from sglang.srt.model_executor.model_runner_components.layer_setup import (
ModelLayerInfo, ModelLayerInfo,
) )
from sglang.srt.runtime_context import get_exec, get_memory, get_schedule
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -144,13 +145,13 @@ class MlxModelRunnerStub(ModelRunner):
(``MlxAuxiliaryStateComponent``) raises ``NotImplementedError`` for (``MlxAuxiliaryStateComponent``) raises ``NotImplementedError`` for
the mode. the mode.
""" """
if self.server_args.disable_radix_cache: if get_memory().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
def _explicit_aux_state_size_per_worker(self) -> int | None: def _explicit_aux_state_size_per_worker(self) -> int | None:
"""Return the explicit auxiliary-state cap for this attention-DP owner.""" """Return the explicit auxiliary-state cap for this attention-DP owner."""
aux_state_size = self.server_args.max_mamba_cache_size aux_state_size = get_schedule().max_mamba_cache_size
if aux_state_size is None: if aux_state_size is None:
return None return None
return aux_state_size // self.ps.attn_dp_size return aux_state_size // self.ps.attn_dp_size
@@ -173,7 +174,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 = self.server_args.max_running_requests requested = get_schedule().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)
@@ -189,7 +190,7 @@ class MlxModelRunnerStub(ModelRunner):
ratio = self._aux_state_slots_per_request() ratio = self._aux_state_slots_per_request()
resolved = min(resolved, aux_state_size // ratio) resolved = min(resolved, aux_state_size // ratio)
if resolved <= 0: if resolved <= 0:
global_aux_state_size = self.server_args.max_mamba_cache_size global_aux_state_size = get_schedule().max_mamba_cache_size
min_global_aux_state_size = ratio * self.ps.attn_dp_size min_global_aux_state_size = ratio * self.ps.attn_dp_size
raise RuntimeError( raise RuntimeError(
f"MLX auxiliary-state cache is too small to serve any " f"MLX auxiliary-state cache is too small to serve any "
@@ -221,7 +222,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=self.server_args.enable_memory_saver enable=get_exec().features.enable_memory_saver
) )
# Load model (sets metadata only) # Load model (sets metadata only)
@@ -267,7 +268,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=self.server_args.disable_radix_cache, owns_auxiliary_state_release=get_memory().disable_radix_cache,
) )
else: else:
self.req_to_token_pool = ReqToTokenPool( self.req_to_token_pool = ReqToTokenPool(
@@ -31,6 +31,7 @@ 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__)
@@ -53,19 +54,19 @@ class MlxTpModelWorker(TpModelWorker):
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=self.server_args.model_path, model_path=get_model().model_path,
trust_remote_code=self.server_args.trust_remote_code, trust_remote_code=get_model().trust_remote_code,
disable_radix_cache=self.server_args.disable_radix_cache, disable_radix_cache=get_memory().disable_radix_cache,
mem_fraction_static=self.server_args.mem_fraction_static, mem_fraction_static=get_schedule().mem_fraction_static,
quantization=self.server_args.quantization, quantization=get_model().quantization,
) )
if self.server_args.max_total_tokens is not None: if get_schedule().max_total_tokens is not None:
init_kwargs["pool_size"] = self.server_args.max_total_tokens init_kwargs["pool_size"] = get_schedule().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=self.server_args.mem_fraction_static, mem_fraction_static=get_schedule().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,
@@ -23,7 +23,7 @@ from sglang.srt.layers.utils.cp_utils import (
cp_allgather_and_save_kv_cache, 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_server_args from sglang.srt.runtime_context import get_schedule
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -515,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_server_args().disable_chunked_prefix_cache assert not get_schedule().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
@@ -26,7 +26,7 @@ from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_kv_cache
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_flags from sglang.srt.runtime_context import get_flags, get_spec
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
from sglang.srt.utils import get_bool_env_var, get_current_device_stream_fast from sglang.srt.utils import get_bool_env_var, get_current_device_stream_fast
@@ -336,9 +336,7 @@ class AscendAttnBackend(AttentionBackend):
self.use_fa = get_bool_env_var("ASCEND_USE_FA", "False") self.use_fa = get_bool_env_var("ASCEND_USE_FA", "False")
self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False") self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False")
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.speculative_num_draft_tokens = ( self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
model_runner.server_args.speculative_num_draft_tokens
)
self.ascend_attn_mask_builder = AscendAttnMaskBuilder( self.ascend_attn_mask_builder = AscendAttnMaskBuilder(
model_runner, self.device, self.use_fia, self.use_mla model_runner, self.device, self.use_fia, self.use_mla
) )
@@ -13,7 +13,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 from sglang.srt.runtime_context import get_parallel, get_spec
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -1493,9 +1493,8 @@ 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_server_args().speculative_num_draft_tokens or 1 n_draft = get_spec().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
) )
@@ -1540,9 +1539,8 @@ 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_server_args().speculative_num_draft_tokens or 1 max_seqlen_q = get_spec().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_server_args from sglang.srt.runtime_context import get_mm
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_server_args().mm_attention_backend override_backend = get_mm().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_server_args from sglang.srt.runtime_context import get_exec
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_server_args().fuseep_mode, fuse_mode=get_exec().moe.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_server_args().fuseep_mode == 1: if get_exec().moe.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_server_args().fuseep_mode == 2: elif get_exec().moe.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)
+2 -2
View File
@@ -33,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_parallel, get_server_args from sglang.srt.runtime_context import get_exec, 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,
@@ -130,7 +130,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_server_args().rl_on_policy_target is not None: if get_exec().deterministic.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
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_spec
""" """
end to end attention solution with aiter kernels end to end attention solution with aiter kernels
@@ -148,8 +148,8 @@ class AiterAttnBackend(AttentionBackend):
self.device = model_runner.device self.device = model_runner.device
self.is_multimodal = model_runner.model_config.is_multimodal self.is_multimodal = model_runner.model_config.is_multimodal
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.num_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_num_steps = model_runner.server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.topk = topk self.topk = topk
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
@@ -60,7 +60,7 @@ from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import (
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_spec
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
from sglang.srt.speculative.ragged_verify import ( from sglang.srt.speculative.ragged_verify import (
RaggedVerifyMode, RaggedVerifyMode,
@@ -537,9 +537,7 @@ class DeepseekV4AttnBackend(
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4" assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
self.mtp_enabled = self.topk > 0 self.mtp_enabled = self.topk > 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens: int = ( self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens
model_runner.server_args.speculative_num_draft_tokens
)
if self.speculative_num_draft_tokens is not None: if self.speculative_num_draft_tokens is not None:
# Persistent target-verify metadata buffers. Allocated here (not # Persistent target-verify metadata buffers. Allocated here (not
# lazily) so they are ordinary tensors: the first touch of a lazy # lazily) so they are ordinary tensors: the first touch of a lazy
@@ -39,7 +39,7 @@ from sglang.srt.layers.attention.dsv4.metadata 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.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_spec
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
from sglang.srt.utils import ceil_align from sglang.srt.utils import ceil_align
@@ -455,9 +455,7 @@ class DeepseekV4HipRadixBackend(
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4" assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
self.mtp_enabled = self.topk > 0 self.mtp_enabled = self.topk > 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens: int = ( self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens
model_runner.server_args.speculative_num_draft_tokens
)
self.speculative_step_id = speculative_step_id self.speculative_step_id = speculative_step_id
self.forward_metadata: Union[ self.forward_metadata: Union[
DSV4Metadata, DSV4Metadata,
@@ -37,7 +37,14 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
get_tc_piecewise_forward_context, get_tc_piecewise_forward_context,
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import (
get_device,
get_exec,
get_parallel,
get_schedule,
get_server_args,
get_spec,
)
from sglang.srt.state_capturer.indexer_topk import ( from sglang.srt.state_capturer.indexer_topk import (
maybe_capture_indexer_topk, maybe_capture_indexer_topk,
) )
@@ -163,7 +170,7 @@ def _uses_dsa_attention_backend(forward_batch: ForwardBatch) -> bool:
): ):
backend_name = ( backend_name = (
decode_backend decode_backend
if server_args.speculative_attention_mode == "decode" if get_spec().speculative_attention_mode == "decode"
else prefill_backend else prefill_backend
) )
else: else:
@@ -460,7 +467,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_server_args().device, device=get_device().device,
) )
self.block_size = block_size self.block_size = block_size
self.scale_fmt = scale_fmt self.scale_fmt = scale_fmt
@@ -471,7 +478,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_server_args().dsa_paged_mqa_logits_backend get_exec().kernel.dsa_paged_mqa_logits_backend
) )
@contextlib.contextmanager @contextlib.contextmanager
@@ -1066,7 +1073,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_server_args().mem_fraction_static mem_fraction_static = get_schedule().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:
@@ -15,7 +15,7 @@ from typing import (
import torch import torch
from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_spec
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from sglang.kernels.ops.attention.dsa.dequant_k_cache import ( from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
@@ -469,9 +469,7 @@ class DeepseekSparseAttnBackend(
# Speculative decoding # Speculative decoding
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = model_runner.server_args.speculative_eagle_topk or 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens = ( self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
model_runner.server_args.speculative_num_draft_tokens
)
self.speculative_step_id = speculative_step_id self.speculative_step_id = speculative_step_id
self.use_fused_topk = should_use_dsa_fused_topk( self.use_fused_topk = should_use_dsa_fused_topk(
model_runner.server_args, seed_dsa_topk_from_draft_extend model_runner.server_args, seed_dsa_topk_from_draft_extend
@@ -40,7 +40,7 @@ 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_parallel from sglang.srt.runtime_context import get_exec, 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
@@ -922,9 +922,8 @@ 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_server_args().enable_deepseek_v4_fp4_indexer self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer
self.alt_streams = alt_streams self.alt_streams = alt_streams
def compute_q( def compute_q(
@@ -30,7 +30,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_server_args from sglang.srt.runtime_context import get_schedule, get_spec
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
@@ -204,9 +204,7 @@ class FlashAttentionBackend(AttentionBackend):
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = model_runner.server_args.speculative_eagle_topk or 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens = ( self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
model_runner.server_args.speculative_num_draft_tokens
)
if ( if (
self.speculative_num_draft_tokens is not None self.speculative_num_draft_tokens is not None
and model_runner.is_draft_worker and model_runner.is_draft_worker
@@ -1513,7 +1511,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_server_args().disable_chunked_prefix_cache assert not get_schedule().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_parallel from sglang.srt.runtime_context import get_disagg, get_exec, get_parallel, get_schedule
""" """
Support attention backend for flashinfer MLA. Support attention backend for flashinfer MLA.
@@ -33,7 +33,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, get_server_args from sglang.srt.runtime_context import get_buffer
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,
@@ -224,9 +224,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_server_args().disaggregation_mode != "decode" and get_disagg().disaggregation_mode != "decode"
and not get_server_args().disable_chunked_prefix_cache and not get_schedule().disable_chunked_prefix_cache
and not get_server_args().flashinfer_mla_disable_ragged and not get_exec().kernel.flashinfer_mla_disable_ragged
) )
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
@@ -402,7 +402,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_server_args().flashinfer_mla_disable_ragged not get_exec().kernel.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()
@@ -20,7 +20,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import scaled_fp8_quant
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_spec
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -94,7 +94,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
torch.float8_e5m2, torch.float8_e5m2,
} }
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.num_draft_tokens = get_spec().speculative_num_draft_tokens
self.cuda_graph_kv_indices = None self.cuda_graph_kv_indices = None
self.cuda_graph_mla_metadata = None self.cuda_graph_mla_metadata = None
@@ -24,7 +24,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_server_args from sglang.srt.runtime_context import get_exec, get_memory, 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
@@ -392,7 +392,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_server_args().mamba_track_interval interval = get_exec().mamba.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)
@@ -823,7 +823,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_server_args().enable_page_major_kv_layout use_triton_causal_conv or get_memory().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(
@@ -8,6 +8,7 @@ from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
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 from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_spec
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -59,7 +60,7 @@ class IntelAMXAttnBackend(AttentionBackend):
self.num_kv_splits = 8 self.num_kv_splits = 8
# speculative decoding params # speculative decoding params
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.num_draft_tokens = get_spec().speculative_num_draft_tokens
def _build_extend_metadata(self, forward_batch: ForwardBatch): def _build_extend_metadata(self, forward_batch: ForwardBatch):
"""Resolve (seq_lens, extend_seq_lens, extend_start_loc, tree_mask) for """Resolve (seq_lens, extend_seq_lens, extend_start_loc, tree_mask) for
@@ -60,7 +60,7 @@ from sglang.srt.models.inkling_common.kernels.sconv import (
fused_extend_sconv_metadata, fused_extend_sconv_metadata,
precompute_helion_extend_metadata, precompute_helion_extend_metadata,
) )
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_exec, get_server_args, get_spec
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -117,7 +117,7 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend):
growing a buffer after a graph captured it moves the address that graph growing a buffer after a graph captured it moves the address that graph
reads, and prefill captures before the decode runner reports its bounds.""" reads, and prefill captures before the decode runner reports its bounds."""
server_args = get_server_args() server_args = get_server_args()
cuda_graph_config = server_args.cuda_graph_config cuda_graph_config = get_exec().graph.cuda_graph_config
decode_bs: list[int] = [] decode_bs: list[int] = []
prefill_tokens: list[int] = [] prefill_tokens: list[int] = []
decode_max_bs = 0 decode_max_bs = 0
@@ -125,7 +125,7 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend):
decode_bs = list(cuda_graph_config.decode.bs or []) decode_bs = list(cuda_graph_config.decode.bs or [])
prefill_tokens = list(cuda_graph_config.prefill.bs or []) prefill_tokens = list(cuda_graph_config.prefill.bs or [])
decode_max_bs = cuda_graph_config.decode.max_bs or 0 decode_max_bs = cuda_graph_config.decode.max_bs or 0
draft_token_num = server_args.speculative_num_draft_tokens or 1 draft_token_num = get_spec().speculative_num_draft_tokens or 1
# req_to_token_pool.size is the runner's max_bs for both graph phases. # req_to_token_pool.size is the runner's max_bs for both graph phases.
max_bs = max([self.req_to_token_pool.size, decode_max_bs, *decode_bs]) max_bs = max([self.req_to_token_pool.size, decode_max_bs, *decode_bs])
max_tokens = max([max_bs, *prefill_tokens, max_bs * draft_token_num]) max_tokens = max([max_bs, *prefill_tokens, max_bs * draft_token_num])
@@ -37,7 +37,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
cuda_graph_fully_disabled, cuda_graph_fully_disabled,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_spec
from sglang.srt.speculative.spec_utils import ( from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width, draft_kv_indices_buffer_width,
draft_kv_indices_used_len, draft_kv_indices_used_len,
@@ -168,9 +168,9 @@ class TritonAttnBackend(AttentionBackend):
self._translate_kv_loc = getattr( self._translate_kv_loc = getattr(
self.token_to_kv_pool_allocator, "translate_kv_loc_dense", None self.token_to_kv_pool_allocator, "translate_kv_loc_dense", None
) or getattr(self.token_to_kv_pool_allocator, "translate_kv_loc", None) ) or getattr(self.token_to_kv_pool_allocator, "translate_kv_loc", None)
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.num_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_num_steps = model_runner.server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = get_spec().speculative_eagle_topk or 0
# Split-KV verify is bit-equivalent only for a pure-causal chain (topk==1) # Split-KV verify is bit-equivalent only for a pure-causal chain (topk==1)
# and is gfx95-only; else fall back to extend_attention_fwd. # and is gfx95-only; else fall back to extend_attention_fwd.
self.use_verify_splitkv = ( self.use_verify_splitkv = (
@@ -34,7 +34,7 @@ from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_buffer from sglang.srt.runtime_context import get_buffer, get_spec
from sglang.srt.speculative.ragged_verify import ( from sglang.srt.speculative.ragged_verify import (
build_ragged_target_verify_geometry, build_ragged_target_verify_geometry,
resolve_ragged_verify_layout, resolve_ragged_verify_layout,
@@ -162,9 +162,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
self.speculative_step_id = speculative_step_id self.speculative_step_id = speculative_step_id
self.target_verify_metadata = {} self.target_verify_metadata = {}
self.speculative_num_draft_tokens = ( self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
model_runner.server_args.speculative_num_draft_tokens
)
# True iff the model declares ENCODER_ONLY (bidirectional) layers, which # True iff the model declares ENCODER_ONLY (bidirectional) layers, which
# need the expanded TARGET_VERIFY metadata (TRTLLMMHAMetadata.encoder_*). # need the expanded TARGET_VERIFY metadata (TRTLLMMHAMetadata.encoder_*).
self.expand_encoder_only_verify = any( self.expand_encoder_only_verify = any(
@@ -41,7 +41,12 @@ 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, get_parallel, get_server_args from sglang.srt.runtime_context import (
get_buffer,
get_parallel,
get_schedule,
get_spec,
)
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():
@@ -238,11 +243,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 = ( self.disable_chunked_prefix_cache = get_schedule().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 = get_spec().speculative_num_draft_tokens
self._verify_mask = None self._verify_mask = None
# Tree-mask scratch is fetched from the target backend only. # Tree-mask scratch is fetched from the target backend only.
self.is_draft_runner = model_runner.is_draft_worker self.is_draft_runner = model_runner.is_draft_worker
+5 -6
View File
@@ -17,7 +17,7 @@ from sglang.kernels.ops.layernorm.norm import (
) )
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_parallel from sglang.srt.runtime_context import get_exec, get_mm, 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,
@@ -88,7 +88,6 @@ 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
@@ -1047,7 +1046,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_server_args().mm_attention_backend is None and _passed_backend is None: if get_mm().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.")
@@ -1126,7 +1125,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_server_args().rl_on_policy_target is not None if get_exec().deterministic.rl_on_policy_target is not None
else {} else {}
) )
q_norm = RMSNorm( q_norm = RMSNorm(
@@ -1154,7 +1153,7 @@ class VisionAttention(nn.Module):
- CUDA (other): "triton_attn" - CUDA (other): "triton_attn"
- Non-CUDA: "sdpa" - Non-CUDA: "sdpa"
""" """
override_backend = get_server_args().mm_attention_backend override_backend = get_mm().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:
@@ -1259,7 +1258,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_server_args().rl_on_policy_target is not None get_exec().deterministic.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), (
@@ -13,7 +13,7 @@ from sglang.kernels.ops.kvcache.kv_indices import (
) )
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_spec
from sglang.srt.utils import get_bool_env_var, get_device_core_count from sglang.srt.utils import get_bool_env_var, get_device_core_count
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -92,7 +92,7 @@ class WaveAttnBackend(AttentionBackend):
(max_bs + 1,), dtype=torch.int64, device=model_runner.device (max_bs + 1,), dtype=torch.int64, device=model_runner.device
) )
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.num_draft_tokens = get_spec().speculative_num_draft_tokens
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
@@ -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_server_args from sglang.srt.runtime_context import get_schedule, get_spec
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -86,9 +86,7 @@ class XPUAttentionBackend(AttentionBackend):
) )
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = model_runner.server_args.speculative_eagle_topk or 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens = ( self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
model_runner.server_args.speculative_num_draft_tokens
)
self.speculative_step_id = speculative_step_id self.speculative_step_id = speculative_step_id
# Local attention settings # Local attention settings
@@ -638,7 +636,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_server_args().disable_chunked_prefix_cache assert not get_schedule().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
+14 -8
View File
@@ -73,7 +73,13 @@ from sglang.srt.model_executor.cuda_graph_config import (
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args from sglang.srt.runtime_context import (
get_exec,
get_forward,
get_parallel,
get_server_args,
get_spec,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm 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,
@@ -169,7 +175,7 @@ def apply_flashinfer_allreduce_fusion(batch_size: int):
(_is_sm90_supported or _is_sm100_supported) (_is_sm90_supported or _is_sm100_supported)
and _is_flashinfer_available and _is_flashinfer_available
and not is_dp_attention_enabled() and not is_dp_attention_enabled()
and get_server_args().flashinfer_allreduce_fusion_backend is not None and get_exec().comm.flashinfer_allreduce_fusion_backend is not None
and not is_flashinfer_allreduce_unavailable() and not is_flashinfer_allreduce_unavailable()
# Symbolic size checks stay last: under Dynamo tracing they guard on # Symbolic size checks stay last: under Dynamo tracing they guard on
# the dynamic token dim, so statically-off configs must short-circuit # the dynamic token dim, so statically-off configs must short-circuit
@@ -190,7 +196,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_server_args().enable_aiter_allreduce_fusion and get_exec().comm.enable_aiter_allreduce_fusion
) )
@@ -278,7 +284,7 @@ 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_server_args().speculative_algorithm != "EAGLE3" and get_spec().speculative_algorithm != "EAGLE3"
) )
if get_server_args().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:
@@ -411,7 +417,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_server_args().enable_two_batch_overlap and get_exec().overlap.enable_two_batch_overlap
) )
@classmethod @classmethod
@@ -475,7 +481,7 @@ class LayerCommunicator:
) )
self._post_init_communicate() self._post_init_communicate()
self._speculative_algo = SpeculativeAlgorithm.from_string( self._speculative_algo = SpeculativeAlgorithm.from_string(
get_server_args().speculative_algorithm get_spec().speculative_algorithm
) )
def _post_init_communicate(self): def _post_init_communicate(self):
@@ -846,7 +852,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_server_args().enable_aiter_allreduce_fusion and get_exec().comm.enable_aiter_allreduce_fusion
) )
) )
and (not self.is_last_layer) and (not self.is_last_layer)
@@ -1151,7 +1157,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_server_args().enable_quant_communications and get_exec().comm.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(
+2 -4
View File
@@ -53,7 +53,7 @@ from sglang.srt.layers.dp_attention import (
) )
from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_device, get_parallel
@dataclass @dataclass
@@ -208,10 +208,8 @@ 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_server_args().device) device = torch.device(get_device().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_parallel, get_server_args from sglang.srt.runtime_context import get_device, get_parallel
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_server_args().device, device=get_device().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_server_args().device, device=get_device().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_server_args().device, device=get_device().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_server_args().device, device=get_device().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_server_args().device, device=get_device().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_server_args().device, device=get_device().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 -6
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_parallel, get_server_args from sglang.srt.runtime_context import get_exec, 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,
@@ -223,7 +223,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_server_args().enable_aiter_allreduce_fusion: if _use_aiter and get_exec().comm.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)
@@ -425,7 +425,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_server_args().rl_on_policy_target == "fsdp" 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)
out = rms_norm_batch_invariant( out = rms_norm_batch_invariant(
@@ -532,7 +532,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_server_args().rl_on_policy_target == "fsdp" or get_exec().deterministic.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)
@@ -593,7 +593,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_server_args().rl_on_policy_target == "fsdp" 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(
@@ -720,7 +720,10 @@ 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 residual is not None or get_server_args().rl_on_policy_target == "fsdp": if (
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,
+2 -2
View File
@@ -39,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_parallel, get_server_args from sglang.srt.runtime_context import get_exec, get_parallel
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:
@@ -1597,7 +1597,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_server_args().enable_quant_communications and get_exec().comm.enable_quant_communications
) )
if forward_batch is not None if forward_batch is not None
else False else False
+4 -4
View File
@@ -51,7 +51,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_exec, 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,
@@ -350,7 +350,7 @@ class LogitsProcessor(nn.Module):
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_server_args().enable_dp_lm_head self.use_attn_tp_group = get_server_args().enable_dp_lm_head
self.use_fp32_lm_head = get_server_args().enable_fp32_lm_head self.use_fp32_lm_head = get_exec().features.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 = (
@@ -374,8 +374,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_server_args().enable_mis self.enable_mis = get_exec().features.enable_mis
self.rl_on_policy_target = get_server_args().rl_on_policy_target self.rl_on_policy_target = get_exec().deterministic.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(
@@ -71,6 +71,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
) )
from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec,
get_global_dwdp_manager, get_global_dwdp_manager,
get_parallel, get_parallel,
get_server_args, get_server_args,
@@ -260,7 +261,7 @@ class FusedMoE(torch.nn.Module):
self._num_global_routed = num_experts - num_shared_slots self._num_global_routed = num_experts - num_shared_slots
server_args = get_server_args() server_args = get_server_args()
if server_args.ep_join_mode == "scale": if get_exec().moe.ep_join_mode == "scale":
storage_ep_size = server_args.elastic_ep_initial_size storage_ep_size = server_args.elastic_ep_initial_size
assert storage_ep_size is not None assert storage_ep_size is not None
self._expert_storage_rank = ( self._expert_storage_rank = (
@@ -359,7 +360,7 @@ class FusedMoE(torch.nn.Module):
print_info_once( print_info_once(
"FlashInfer TRTLLM MoE deferred finalize is " "FlashInfer TRTLLM MoE deferred finalize is "
f"{'enabled' if self.supports_deferred_finalize else 'disabled'} " f"{'enabled' if self.supports_deferred_finalize else 'disabled'} "
f"(moe_runner_backend={server_args.moe_runner_backend}, " f"(moe_runner_backend={get_exec().moe.moe_runner_backend}, "
f"quant_method={type(self.quant_method).__name__})." f"quant_method={type(self.quant_method).__name__})."
) )
+2 -2
View File
@@ -22,6 +22,7 @@ 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__)
@@ -44,10 +45,9 @@ 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_server_args().enable_waterfill num_fused_shared_experts > 0 and get_exec().moe.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_server_args from sglang.srt.runtime_context import get_exec
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,
@@ -531,7 +531,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_server_args().enable_fused_moe_sum_all_reduce get_exec().moe.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_server_args from sglang.srt.runtime_context import get_exec
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_server_args().enable_deterministic_inference: if get_exec().deterministic.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_server_args().enable_deterministic_inference: if get_exec().deterministic.enable_deterministic_inference:
config = { config = {
"BLOCK_SIZE_M": 64, "BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64, "BLOCK_SIZE_N": 64,
@@ -27,7 +27,7 @@ from sglang.srt.layers.moe.topk import (
TopKOutputChecker, 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_server_args from sglang.srt.runtime_context import get_schedule, get_spec
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
@@ -123,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_server_args().chunked_prefill_size cps = get_schedule().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",
@@ -132,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_server_args().speculative_algorithm get_spec().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 = (
+6 -8
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_parallel from sglang.srt.runtime_context import get_exec, get_lora, get_parallel
try: try:
from triton_kernels.matmul_ogs import GatherIndx, RoutingData, ScatterIndx from triton_kernels.matmul_ogs import GatherIndx, RoutingData, ScatterIndx
@@ -430,10 +430,9 @@ 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_server_args().enable_waterfill num_fused_shared_experts > 0 and get_exec().moe.enable_waterfill
) )
self.waterfill_balancer = None self.waterfill_balancer = None
@@ -507,9 +506,8 @@ 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_server_args().enable_lora) use_standard_for_lora = bool(get_lora().enable_lora)
except ValueError: except ValueError:
use_standard_for_lora = False use_standard_for_lora = False
output_format = ( output_format = (
@@ -1362,9 +1360,9 @@ def _eplb_remap_enabled() -> bool:
# there is no EPLB mapping, so the remap must be skipped. # there is no EPLB mapping, so the remap must be skipped.
return False return False
return ( return (
server_args.enable_eplb get_exec().moe.enable_eplb
or server_args.init_expert_location != "trivial" or get_exec().moe.init_expert_location != "trivial"
or server_args.ep_num_redundant_experts > 0 or get_exec().moe.ep_num_redundant_experts > 0
) )
+3 -3
View File
@@ -12,7 +12,7 @@ from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.runtime_context import get_flags, get_forward, get_parallel from sglang.srt.runtime_context import get_exec, get_flags, get_forward, get_parallel
from sglang.srt.utils import is_cuda, is_npu from sglang.srt.utils import is_cuda, is_npu
_is_npu = is_npu() _is_npu = is_npu()
@@ -239,8 +239,8 @@ def get_deepep_output_dtype(self) -> DispatcherOutputDtype:
# 0. Parse server argument. # 0. Parse server argument.
server_args = get_server_args() server_args = get_server_args()
if server_args and server_args.deepep_dispatcher_output_dtype != "auto": if server_args and get_exec().moe.deepep_dispatcher_output_dtype != "auto":
return DispatcherOutputDtype(server_args.deepep_dispatcher_output_dtype) return DispatcherOutputDtype(get_exec().moe.deepep_dispatcher_output_dtype)
# 1. Parse deprecated environment variables. # 1. Parse deprecated environment variables.
if envs.SGLANG_DEEPEP_BF16_DISPATCH.get(): if envs.SGLANG_DEEPEP_BF16_DISPATCH.get():
@@ -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_parallel from sglang.srt.runtime_context import get_exec, 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,7 +34,6 @@ 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,
@@ -1844,7 +1843,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_server_args().cuda_graph_config.prefill.tc_compiler == "inductor" and get_exec().graph.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_server_args from sglang.srt.runtime_context import get_exec
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
is_cpu, is_cpu,
@@ -333,7 +333,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_server_args().flashinfer_mxfp4_moe_precision get_exec().moe.flashinfer_mxfp4_moe_precision
) )
# When `flashinfer_mxfp4` is enabled, dispatch to one of three FlashInfer # When `flashinfer_mxfp4` is enabled, dispatch to one of three 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_server_args from sglang.srt.runtime_context import get_exec
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_server_args().flashinfer_mxfp4_moe_precision get_exec().moe.flashinfer_mxfp4_moe_precision
) )
def create_moe_runner(self, layer, moe_runner_config): def create_moe_runner(self, layer, moe_runner_config):
@@ -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_server_args from sglang.srt.runtime_context import get_exec
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,
@@ -129,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_server_args().rl_on_policy_target is not None or _is_musa: if get_exec().deterministic.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,
@@ -153,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_server_args().rl_on_policy_target is not None else None "cpu" if get_exec().deterministic.rl_on_policy_target is not None else None
) )
inv_freq = 1.0 / ( inv_freq = 1.0 / (
base base
@@ -164,7 +164,7 @@ class RotaryEmbedding(MultiPlatformOp):
/ self.rotary_dim / self.rotary_dim
) )
) )
if get_server_args().rl_on_policy_target is not None: if get_exec().deterministic.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_server_args from sglang.srt.runtime_context import get_exec, 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,
@@ -132,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_server_args().rl_on_policy_target is not None: if get_exec().deterministic.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):
+8 -6
View File
@@ -15,7 +15,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.logprob_processor import ( from sglang.srt.layers.logprob_processor import (
OutputLogprobProcessor, OutputLogprobProcessor,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_exec, 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
@@ -74,12 +74,14 @@ 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_server_args().rl_on_policy_target self.rl_on_policy_target = get_exec().deterministic.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 = get_server_args().enable_deterministic_inference self.enable_deterministic = (
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_server_args().sampling_backend == "ascend" self.use_ascend_backend = get_exec().kernel.sampling_backend == "ascend"
self.output_logprob_processor = OutputLogprobProcessor() self.output_logprob_processor = OutputLogprobProcessor()
@@ -260,7 +262,7 @@ class Sampler(nn.Module):
positions=positions, positions=positions,
) )
else: else:
backend = get_server_args().sampling_backend backend = get_exec().kernel.sampling_backend
if backend == "flashinfer": if backend == "flashinfer":
assert ( assert (
sampling_info.sampling_seed is None sampling_info.sampling_seed is None
@@ -540,7 +542,7 @@ def create_sampler(backend: Optional[str] = None) -> "Sampler":
"""Create a sampler honoring custom backend registrations.""" """Create a sampler honoring custom backend registrations."""
server_args = get_server_args() server_args = get_server_args()
backend = backend or (server_args.sampling_backend if server_args else None) backend = backend or (get_exec().kernel.sampling_backend if server_args else None)
if backend in _CUSTOM_SAMPLER_FACTORIES: if backend in _CUSTOM_SAMPLER_FACTORIES:
sampler = _CUSTOM_SAMPLER_FACTORIES[backend]() sampler = _CUSTOM_SAMPLER_FACTORIES[backend]()
@@ -48,7 +48,7 @@ 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 publish from sglang.srt.runtime_context import get_exec, publish
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 +232,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 self.server_args.elastic_ep_backend is not None: if get_exec().moe.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 +485,7 @@ class DataParallelController:
logger.debug("Worker port broadcast completed") logger.debug("Worker port broadcast completed")
return worker_ports return worker_ports
finally: finally:
if self.server_args.elastic_ep_backend is None: if get_exec().moe.elastic_ep_backend is None:
rep_socket.close() rep_socket.close()
else: else:
threading.Thread( threading.Thread(
+9 -4
View File
@@ -33,7 +33,12 @@ 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 get_parallel, get_server_args from sglang.srt.runtime_context import (
get_disagg,
get_parallel,
get_schedule,
get_server_args,
)
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
@@ -931,7 +936,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_server_args().chunked_prefill_size chunked_prefill_size = get_schedule().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"
@@ -1295,7 +1300,7 @@ def general_mm_embed_routine(
# encoder/ViT execution and multimodal feature placement, while # encoder/ViT execution and multimodal feature placement, while
# the language model range below excludes both. # the language model range below excludes both.
with torch.profiler.record_function("sglang.vlm.mm_embedding"): with torch.profiler.record_function("sglang.vlm.mm_embedding"):
if server_args and server_args.enable_adaptive_dispatch_to_encoder: if server_args and get_disagg().enable_adaptive_dispatch_to_encoder:
# Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches # Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches
input_embeds, other_info = _embed_mm_inputs_with_split( input_embeds, other_info = _embed_mm_inputs_with_split(
mm_inputs_list=mm_inputs_list, mm_inputs_list=mm_inputs_list,
@@ -1340,7 +1345,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_server_args().language_only: if get_disagg().language_only:
precomputed_embeddings = getattr( precomputed_embeddings = getattr(
mm_item, "precomputed_embeddings", None mm_item, "precomputed_embeddings", None
) )
+6 -5
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
from sglang.srt.dllm.config import DllmConfig from sglang.srt.dllm.config import DllmConfig
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_exec, get_schedule, get_serving, get_spec
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
Range, Range,
ceil_align, ceil_align,
@@ -1097,7 +1098,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_server_args().speculative_algorithm spec_alg = get_spec().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
@@ -1118,7 +1119,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_server_args().strip_thinking_cache and self.reasoning_tokens > 0: if get_serving().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
@@ -2922,7 +2923,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
) )
if server_args.enable_mamba_extra_buffer(): if server_args.enable_mamba_extra_buffer():
mamba_track_interval = server_args.mamba_track_interval mamba_track_interval = get_exec().mamba.mamba_track_interval
if len(self.reqs) == 0: if len(self.reqs) == 0:
self.mamba_track_indices = torch.empty( self.mamba_track_indices = torch.empty(
@@ -3168,8 +3169,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
continue continue
else: else:
pre_len = ( pre_len = (
pre_len - server_args.chunked_prefill_size pre_len - get_schedule().chunked_prefill_size
if server_args.chunked_prefill_size > 0 if get_schedule().chunked_prefill_size > 0
else pre_len else pre_len
) )
self._evict_swa(req, pre_len) self._evict_swa(req, pre_len)
@@ -5,6 +5,7 @@ from array import array
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor
from sglang.srt.runtime_context import get_disagg
from sglang.srt.utils import get_bool_env_var from sglang.srt.utils import get_bool_env_var
_ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG") _ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG")
@@ -56,7 +57,6 @@ 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_server_args
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -195,7 +195,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_server_args().disaggregation_mode != "decode" and get_disagg().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)
+84 -75
View File
@@ -27,6 +27,20 @@ from functools import partial
from http import HTTPStatus from http import HTTPStatus
from typing import Any, Deque, Dict, List, Optional, Tuple, Union from typing import Any, Deque, Dict, List, Optional, Tuple, Union
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_exec,
get_lora,
get_memory,
get_mm,
get_model,
get_observability,
get_schedule,
get_serving,
get_spec,
)
from sglang.srt.utils.common import suppress_noisy_warnings # isort: skip from sglang.srt.utils.common import suppress_noisy_warnings # isort: skip
suppress_noisy_warnings() suppress_noisy_warnings()
@@ -482,9 +496,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=self.server_args.enable_metrics, enable_metrics=get_observability().enable_metrics,
enable_kv_cache_events=bool( enable_kv_cache_events=bool(
self.server_args.kv_events_config get_observability().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
@@ -526,8 +540,8 @@ class Scheduler(
self.init_hisparse_coordinator() self.init_hisparse_coordinator()
if ( if (
self.server_args.disaggregation_mode == "decode" get_disagg().disaggregation_mode == "decode"
and self.server_args.disaggregation_decode_enable_offload_kvcache and get_disagg().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,
@@ -642,7 +656,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 self.server_args.dllm_algorithm is not None if get_exec().dllm.dllm_algorithm is not None
else None else None
) )
@@ -671,10 +685,10 @@ class Scheduler(
port_args=port_args, port_args=port_args,
is_rank_zero=is_rank_zero, is_rank_zero=is_rank_zero,
skip_tokenizer_init=self.server_args.skip_tokenizer_init, skip_tokenizer_init=self.server_args.skip_tokenizer_init,
metrics_enabled=self.server_args.enable_metrics metrics_enabled=get_observability().enable_metrics
and ( and (
self.ps.attn_tp_rank == 0 self.ps.attn_tp_rank == 0
or self.server_args.enable_metrics_for_all_schedulers or get_observability().enable_metrics_for_all_schedulers
), ),
enable_scripted_runtime=envs.SGLANG_TEST_SCRIPTED_RUNTIME.get(), enable_scripted_runtime=envs.SGLANG_TEST_SCRIPTED_RUNTIME.get(),
) )
@@ -693,7 +707,7 @@ class Scheduler(
port_args, port_args,
self.ps.dp_size, self.ps.dp_size,
dp_rank, dp_rank,
publish_interval=self.server_args.load_snapshot_publish_interval, publish_interval=get_observability().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)
@@ -703,7 +717,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 self.server_args.sleep_on_idle and get_device().sleep_on_idle
): ):
self.idle_sleeper = IdleSleeper( self.idle_sleeper = IdleSleeper(
sockets=[ sockets=[
@@ -737,22 +751,22 @@ class Scheduler(
else: else:
if self.model_config.is_multimodal: if self.model_config.is_multimodal:
self.processor = get_processor( self.processor = get_processor(
server_args.tokenizer_path, get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode, tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=get_model().trust_remote_code,
revision=server_args.revision, revision=get_model().revision,
use_fast=not server_args.disable_fast_image_processor, use_fast=not get_mm().disable_fast_image_processor,
tokenizer_backend=server_args.tokenizer_backend, tokenizer_backend=get_serving().tokenizer_backend,
model_name=server_args.model_path, model_name=get_model().model_path,
) )
self.tokenizer = get_tokenizer_from_processor(self.processor) self.tokenizer = get_tokenizer_from_processor(self.processor)
else: else:
self.tokenizer = get_tokenizer( self.tokenizer = get_tokenizer(
server_args.tokenizer_path, get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode, tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=get_model().trust_remote_code,
revision=server_args.revision, revision=get_model().revision,
tokenizer_backend=server_args.tokenizer_backend, tokenizer_backend=get_serving().tokenizer_backend,
) )
# Load multimodal processor for M-RoPE fallback computation. # Load multimodal processor for M-RoPE fallback computation.
@@ -774,9 +788,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 self.server_args.reasoning_parser and self.tokenizer: if get_serving().reasoning_parser and self.tokenizer:
reasoning_parser = ReasoningParser( reasoning_parser = ReasoningParser(
model_type=self.server_args.reasoning_parser, model_type=get_serving().reasoning_parser,
stream_reasoning=False, stream_reasoning=False,
tokenizer=self.tokenizer, tokenizer=self.tokenizer,
) )
@@ -847,7 +861,7 @@ class Scheduler(
target_worker=self.tp_worker, target_worker=self.tp_worker,
) )
if self.server_args.speculative_draft_load_format is not None: if get_spec().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
@@ -855,10 +869,10 @@ class Scheduler(
# format. # format.
self.server_args.override( self.server_args.override(
"scheduler.draft_load_format", "scheduler.draft_load_format",
load_format=self.server_args.speculative_draft_load_format, load_format=get_spec().speculative_draft_load_format,
) )
logger.info( logger.info(
f"Using draft model load_format: '{self.server_args.speculative_draft_load_format}'" f"Using draft model load_format: '{get_spec().speculative_draft_load_format}'"
) )
DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args) DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args)
@@ -925,8 +939,8 @@ class Scheduler(
model_runner.post_capture_resize_kv_pool() model_runner.post_capture_resize_kv_pool()
if ( if (
self.server_args.elastic_ep_backend is not None get_exec().moe.elastic_ep_backend is not None
and self.server_args.ep_join_mode == "recover" and get_exec().moe.ep_join_mode == "recover"
): ):
model_runner.post_capture_elastic_ep_recover() model_runner.post_capture_elastic_ep_recover()
@@ -955,7 +969,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(
self.server_args.min_free_slots_delay, get_schedule().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(),
) )
@@ -1001,14 +1015,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={self.server_args.chunked_prefill_size}, " f"chunked_prefill_size={get_schedule().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 self.server_args.enable_metrics: if get_observability().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.
@@ -1055,7 +1069,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 = self.server_args.chunked_prefill_size self.chunked_prefill_size = get_schedule().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
) )
@@ -1075,13 +1089,12 @@ 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 self.chunked_prefill_size is not None and get_schedule().enable_mixed_chunk
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 = (
self.server_args.enable_dynamic_chunking and self.ps.pp_size > 1 get_schedule().enable_dynamic_chunking and self.ps.pp_size > 1
) )
if self.enable_dynamic_chunking: if self.enable_dynamic_chunking:
try: try:
@@ -1117,8 +1130,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 self.server_args.enable_prefill_delayer: if get_schedule().enable_prefill_delayer:
if self.server_args.disaggregation_mode == "decode": if get_disagg().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)."
@@ -1135,15 +1148,15 @@ class Scheduler(
if self.metrics_reporter.enable_metrics if self.metrics_reporter.enable_metrics
else None else None
), ),
max_delay_passes=self.server_args.prefill_delayer_max_delay_passes, max_delay_passes=get_schedule().prefill_delayer_max_delay_passes,
token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark, token_usage_low_watermark=get_schedule().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 self.server_args.disable_priority_preemption and not get_schedule().disable_priority_preemption
) )
self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args( self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args(
@@ -1159,12 +1172,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=self.server_args.watchdog_timeout self, watchdog_timeout=get_device().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=self.server_args.enable_memory_saver enable=get_exec().features.enable_memory_saver
) )
# Init recv skipper and input blocker # Init recv skipper and input blocker
@@ -1186,11 +1199,9 @@ 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( self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
self.server_args.disaggregation_mode
)
self.transfer_backend = TransferBackend( self.transfer_backend = TransferBackend(
self.server_args.disaggregation_transfer_backend get_disagg().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?
@@ -1260,10 +1271,10 @@ class Scheduler(
tp_size=self.ps.tp_size, tp_size=self.ps.tp_size,
dp_size=self.server_args.dp_size, dp_size=self.server_args.dp_size,
gpu_id=self.ps.gpu_id, gpu_id=self.ps.gpu_id,
bootstrap_port=self.server_args.disaggregation_bootstrap_port, bootstrap_port=get_disagg().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=self.server_args.num_reserved_decode_tokens, num_reserved_decode_tokens=get_disagg().num_reserved_decode_tokens,
transfer_backend=self.transfer_backend, transfer_backend=self.transfer_backend,
) )
@@ -1289,7 +1300,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=self.server_args.disaggregation_bootstrap_port, bootstrap_port=get_disagg().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,
@@ -1303,11 +1314,10 @@ 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 ( if get_disagg().language_only and get_disagg().encoder_transfer_backend in [
self.server_args.language_only "zmq_to_scheduler",
and self.server_args.encoder_transfer_backend "mooncake",
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,
@@ -1388,7 +1398,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 self.server_args.enable_deterministic_inference: if not get_exec().deterministic.enable_deterministic_inference:
self.truncation_align_size = None self.truncation_align_size = None
return return
@@ -1794,10 +1804,10 @@ class Scheduler(
) )
def init_lora_drainer(self) -> None: def init_lora_drainer(self) -> None:
if self.server_args.lora_drain_wait_threshold > 0.0: if get_lora().lora_drain_wait_threshold > 0.0:
self.lora_drainer = LoRADrainer( self.lora_drainer = LoRADrainer(
self.server_args.max_loras_per_batch, get_lora().max_loras_per_batch,
self.server_args.lora_drain_wait_threshold, get_lora().lora_drain_wait_threshold,
) )
else: else:
self.lora_drainer = None self.lora_drainer = None
@@ -1923,7 +1933,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=self.server_args.kv_events_config, kv_events_config=get_observability().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,
@@ -2107,7 +2117,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 self.server_args.enable_broadcast_mm_inputs_process: if get_mm().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)
@@ -2154,7 +2164,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 (
self.server_args.elastic_ep_backend is None get_exec().moe.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()
): ):
@@ -2200,8 +2210,7 @@ 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 recv_req.session_id is not None and get_memory().enable_session_radix_cache
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:
@@ -2213,7 +2222,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 = self.server_args.disaggregation_bootstrap_port recv_req.bootstrap_port = get_disagg().disaggregation_bootstrap_port
req = Req( req = Req(
recv_req.rid, recv_req.rid,
@@ -2366,7 +2375,7 @@ class Scheduler(
self._add_request_to_queue(req) self._add_request_to_queue(req)
return return
if req.return_sampling_mask and self.server_args.sampling_backend == "ascend": if req.return_sampling_mask and get_exec().kernel.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 = (
@@ -2415,7 +2424,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,
self.server_args.allow_auto_truncate, get_serving().allow_auto_truncate,
) )
if error_msg: if error_msg:
req.set_finish_with_abort(error_msg) req.set_finish_with_abort(error_msg)
@@ -2693,7 +2702,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,
self.server_args.allow_auto_truncate, get_serving().allow_auto_truncate,
) )
if error_msg: if error_msg:
self._add_request_to_queue(req) self._add_request_to_queue(req)
@@ -2905,7 +2914,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 self.server_args.speculative_skip_dp_mlp_sync and not get_spec().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:
@@ -2979,7 +2988,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 self.server_args.enable_flexkv: if self.enable_hierarchical_cache or get_memory().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:
@@ -3046,7 +3055,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=self.server_args.prefill_max_requests, prefill_max_requests=get_schedule().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),
@@ -3619,7 +3628,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 self.server_args.elastic_ep_backend is not None self.enable_dp_attention and get_exec().moe.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
@@ -3924,7 +3933,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=self.server_args.served_model_name, served_model_name=get_serving().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,
) )
@@ -4044,7 +4053,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 self.server_args.elastic_ep_backend is not None: if get_exec().moe.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()
@@ -4583,10 +4592,10 @@ class Scheduler(
return None return None
def close_session(self, recv_req: CloseSessionReqInput): def close_session(self, recv_req: CloseSessionReqInput):
if self.server_args.enable_session_radix_cache: if get_memory().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 (
self.server_args.enable_session_radix_cache get_memory().enable_session_radix_cache
): ):
self.session_controller.close(recv_req) self.session_controller.close(recv_req)
@@ -27,7 +27,13 @@ from sglang.srt.mem_cache.common import (
maybe_cache_unfinished_req, maybe_cache_unfinished_req,
release_kv_cache, release_kv_cache,
) )
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_memory,
get_observability,
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
@@ -84,7 +90,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 = self.server_args.disaggregation_decode_enable_radix_cache use_free_group = get_disagg().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:
@@ -92,7 +98,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 self.server_args.enable_hisparse: if get_memory().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)
@@ -243,7 +249,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 self.server_args.enable_hisparse: if get_memory().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)
@@ -756,7 +762,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 self.server_args.enable_metrics: if get_observability().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
) )
@@ -939,7 +945,7 @@ class SchedulerBatchResultProcessor:
self._mamba_prefix_cache_update(req, batch, result, i) self._mamba_prefix_cache_update(req, batch, result, i)
if ( if (
self.server_args.disaggregation_decode_enable_offload_kvcache get_disagg().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)
@@ -959,12 +965,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 self.server_args.disaggregation_decode_enable_offload_kvcache: if get_disagg().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 self.server_args.enable_hisparse: if get_memory().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
@@ -1063,7 +1069,7 @@ class SchedulerBatchResultProcessor:
other_idx other_idx
].item() == -1 and mamba_lazy_spec_in_window( ].item() == -1 and mamba_lazy_spec_in_window(
req, req,
server_args.mamba_track_interval, get_exec().mamba.mamba_track_interval,
server_args.max_speculative_num_draft_tokens, server_args.max_speculative_num_draft_tokens,
) )
if ( if (
@@ -1102,7 +1108,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_server_args().mamba_track_interval interval = get_exec().mamba.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:
@@ -26,6 +26,7 @@ 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_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
@@ -385,7 +386,7 @@ class SchedulerDPAttnAdapter:
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=self.server_args.disable_overlap_schedule, disable_overlap_schedule=get_schedule().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,6 +14,7 @@ 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
@@ -155,7 +156,7 @@ class SchedulerLoadInquirer:
) )
lora = None lora = None
if self.server_args.enable_lora: if get_lora().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,
@@ -11,6 +11,7 @@ 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, MIS_DELIMITER_TOKEN_ID,
ServerArgs, ServerArgs,
@@ -164,7 +165,7 @@ class SchedulerLogprobResultProcessor:
delimiter token receive logprobs. delimiter token receive logprobs.
""" """
return ( return (
self.server_args.enable_mis get_exec().features.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
) )
@@ -26,6 +26,7 @@ from sglang.srt.observability.metrics_collector import (
SchedulerStats, SchedulerStats,
compute_routing_key_stats, compute_routing_key_stats,
) )
from sglang.srt.runtime_context import get_spec
from sglang.srt.utils.device_timer import DeviceTimer from sglang.srt.utils.device_timer import DeviceTimer
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
@@ -764,12 +765,10 @@ class SchedulerMetricsReporter:
else: else:
spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct
num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct
if self.scheduler.server_args.speculative_num_draft_tokens: if get_spec().speculative_num_draft_tokens:
draft_per_round = ( draft_per_round = get_spec().speculative_num_draft_tokens - 1
self.scheduler.server_args.speculative_num_draft_tokens - 1
)
else: else:
draft_per_round = self.scheduler.server_args.speculative_num_steps or 0 draft_per_round = get_spec().speculative_num_steps or 0
total_draft_tokens = self.spec_num_forward_ct * draft_per_round total_draft_tokens = self.spec_num_forward_ct * draft_per_round
spec_accept_rate = ( spec_accept_rate = (
num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0 num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0
@@ -27,6 +27,7 @@ from sglang.srt.managers.schedule_batch import (
Req, 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
@@ -153,7 +154,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=self.server_args.stream_interval, default_stream_interval=get_serving().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,
rust_server_mode=self.rust_server is not None, rust_server_mode=self.rust_server is not None,
@@ -184,7 +185,7 @@ class SchedulerOutputStreamer:
if ( if (
req.finished() req.finished()
and self.ps.attn_tp_rank == 0 and self.ps.attn_tp_rank == 0
and self.server_args.enable_request_time_stats_logging and get_observability().enable_request_time_stats_logging
): ):
req.log_time_stats() req.log_time_stats()
@@ -19,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_server_args from sglang.srt.runtime_context import get_device
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
@@ -257,7 +257,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_server_args().base_gpu_id: if self.ps.gpu_id == get_device().base_gpu_id:
torch.cuda.cudart().cudaProfilerStart() torch.cuda.cudart().cudaProfilerStart()
self.profile_in_progress = True self.profile_in_progress = True
@@ -368,7 +368,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_server_args().base_gpu_id: if self.ps.gpu_id == get_device().base_gpu_id:
torch.cuda.cudart().cudaProfilerStop() torch.cuda.cudart().cudaProfilerStop()
merge_message = self._merge_profile_traces() merge_message = self._merge_profile_traces()
@@ -27,6 +27,7 @@ from sglang.srt.managers.mm_utils import (
has_shm_features, has_shm_features,
unwrap_shm_features, unwrap_shm_features,
) )
from sglang.srt.runtime_context import get_disagg
from sglang.srt.utils import ( from sglang.srt.utils import (
broadcast_pyobj, broadcast_pyobj,
point_to_point_pyobj, point_to_point_pyobj,
@@ -231,8 +232,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 self.server_args.language_only and get_disagg().language_only
and self.server_args.encoder_transfer_backend and get_disagg().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)
@@ -36,6 +36,7 @@ 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
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
@@ -479,7 +480,7 @@ class SchedulerPPMixin:
) )
) )
if self.server_args.disaggregation_decode_enable_offload_kvcache: if get_disagg().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:
@@ -549,7 +550,7 @@ 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 self.server_args.disaggregation_decode_enable_offload_kvcache: if get_disagg().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:
+11 -10
View File
@@ -47,6 +47,7 @@ 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 (
@@ -408,14 +409,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=(
self.server_args.model_path get_model().model_path
if not self.is_draft_worker if not self.is_draft_worker
else self.server_args.speculative_draft_model_path else get_spec().speculative_draft_model_path
), ),
model_revision=( model_revision=(
self.server_args.revision get_model().revision
if not self.is_draft_worker if not self.is_draft_worker
else self.server_args.speculative_draft_model_revision else get_spec().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,
@@ -426,7 +427,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=self.server_args.mem_fraction_static, mem_fraction_static=get_schedule().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,
@@ -442,11 +443,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, self.server_args.speculative_num_steps): for i in range(1, get_spec().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=self.server_args.mem_fraction_static, mem_fraction_static=get_schedule().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,
@@ -462,7 +463,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 self.server_args.dllm_algorithm is not None: if get_exec().dllm.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
@@ -488,9 +489,9 @@ class TpModelWorker(BaseTpWorker):
) )
return ( return (
self.model_runner.max_total_num_tokens, self.model_runner.max_total_num_tokens,
self.server_args.max_prefill_tokens, get_schedule().max_prefill_tokens,
self.model_runner.max_running_requests, self.model_runner.max_running_requests,
self.server_args.max_queued_requests, get_schedule().max_queued_requests,
max_req_len, max_req_len,
max_req_len - 5, max_req_len - 5,
self.random_seed, self.random_seed,
@@ -23,6 +23,7 @@ from sglang.srt.observability.metrics_collector import (
RadixCacheMetricsCollector, RadixCacheMetricsCollector,
resolve_collector_class, resolve_collector_class,
) )
from sglang.srt.runtime_context import get_observability
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
@@ -238,8 +239,8 @@ class BasePrefixCache(ABC, PrefixCacheTrait):
server_args = get_server_args() server_args = get_server_args()
labels = {"cache_type": self.__class__.__name__} labels = {"cache_type": self.__class__.__name__}
if server_args.extra_metric_labels: if get_observability().extra_metric_labels:
labels.update(server_args.extra_metric_labels) labels.update(get_observability().extra_metric_labels)
radix_cache_cls = resolve_collector_class( radix_cache_cls = resolve_collector_class(
server_args, server_args,
STAT_LOGGER_ROLE_RADIX_CACHE, STAT_LOGGER_ROLE_RADIX_CACHE,
+9 -4
View File
@@ -16,7 +16,12 @@ 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 from sglang.srt.runtime_context import (
get_schedule,
get_server_args,
get_serving,
get_spec,
)
from sglang.srt.utils.common import ceil_align from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -179,12 +184,12 @@ def _release_overallocated_kv_indices(
req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache
) -> None: ) -> None:
global_server_args = get_server_args() global_server_args = get_server_args()
page_size = global_server_args.page_size page_size = get_schedule().page_size
spec_algo = global_server_args.speculative_algorithm spec_algo = get_spec().speculative_algorithm
# 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 global_server_args.strip_thinking_cache: if spec_algo is None and not get_serving().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_server_args from sglang.srt.runtime_context import get_exec, get_server_args, get_spec
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_server_args().enable_deepseek_v4_fp4_indexer self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer
self._create_buffer() self._create_buffer()
@@ -577,8 +577,8 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self.c128_kv_pool = None self.c128_kv_pool = None
server_args = get_server_args() server_args = get_server_args()
spec_extra = ( spec_extra = (
(server_args.speculative_num_draft_tokens - 1) (get_spec().speculative_num_draft_tokens - 1)
if server_args.speculative_algorithm is not None if get_spec().speculative_algorithm is not None
else 0 else 0
) )
self.unified_kv_pool = DeepSeekV4UnifiedKVPool( self.unified_kv_pool = DeepSeekV4UnifiedKVPool(
@@ -659,7 +659,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
def get_ring_size(self, compress_ratio: int) -> int: def get_ring_size(self, compress_ratio: int) -> int:
server_args = get_server_args() server_args = get_server_args()
is_speculative = server_args.speculative_algorithm is not None is_speculative = get_spec().speculative_algorithm is not None
return get_compress_state_ring_size(compress_ratio, is_speculative) return get_compress_state_ring_size(compress_ratio, is_speculative)
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor): def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
@@ -58,7 +58,15 @@ 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 get_model, get_parallel from sglang.srt.runtime_context import (
get_context,
get_disagg,
get_exec,
get_memory,
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 (
@@ -184,6 +192,7 @@ class KVCacheConfigurator:
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator]
memory_pool_config: Optional[MemoryPoolConfig] memory_pool_config: Optional[MemoryPoolConfig]
draft_model_idx: Optional[int] = None draft_model_idx: Optional[int] = None
kv_cache_dtype_str: Optional[str] = None
mambaish_config: Optional[Any] = field(init=False) mambaish_config: Optional[Any] = field(init=False)
hybrid_gdn_config: Optional[Any] = field(init=False) hybrid_gdn_config: Optional[Any] = field(init=False)
is_inkling_mtp_draft: bool = field(init=False) is_inkling_mtp_draft: bool = field(init=False)
@@ -211,7 +220,7 @@ class KVCacheConfigurator:
def _build_fp4_quant_method(self, *, num_layers: int): def _build_fp4_quant_method(self, *, num_layers: int):
if not is_float4_e2m1fn_x2(self.kv_cache_dtype): if not is_float4_e2m1fn_x2(self.kv_cache_dtype):
return None return None
quant_name = resolve_kv_cache_quant(get_model().kv_cache_dtype) quant_name = resolve_kv_cache_quant(self.kv_cache_dtype_str)
if quant_name is None: if quant_name is None:
return None return None
quant_method = get_kv_cache_quant_method( quant_method = get_kv_cache_quant_method(
@@ -314,8 +323,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 (
self.server_args.enable_unified_memory get_memory().enable_unified_memory
and self.server_args.disaggregation_mode == "null" and get_disagg().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 +373,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=self.server_args.max_mamba_cache_size, mamba_size=get_schedule().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=self.server_args.speculative_eagle_topk, speculative_eagle_topk=get_spec().speculative_eagle_topk,
) )
# Initialize token_to_kv_pool # Initialize token_to_kv_pool
@@ -400,7 +409,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 (
self.server_args.prefill_only_disable_kv_cache get_schedule().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)
): ):
@@ -435,8 +444,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 self.server_args.speculative_num_draft_tokens is not None: if get_spec().speculative_num_draft_tokens is not None:
extra_max_context_len += self.server_args.speculative_num_draft_tokens extra_max_context_len += get_spec().speculative_num_draft_tokens
mamba_layer_ids = [ mamba_layer_ids = [
i i
@@ -471,14 +480,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=self.server_args.max_mamba_cache_size, max_mamba_cache_size=get_schedule().max_mamba_cache_size,
max_num_reqs=max_num_reqs, max_num_reqs=max_num_reqs,
enable_memory_saver=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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=self.server_args.speculative_num_draft_tokens, speculative_num_draft_tokens=get_spec().speculative_num_draft_tokens,
disable_overlap_schedule=self.server_args.disable_overlap_schedule, disable_overlap_schedule=get_schedule().disable_overlap_schedule,
need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"), need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"),
mamba_full_memory_ratio=self.server_args.mamba_full_memory_ratio, mamba_full_memory_ratio=get_schedule().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.
@@ -511,13 +520,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 self.server_args.speculative_num_draft_tokens is not None: if get_spec().speculative_num_draft_tokens is not None:
extra_max_context_len += self.server_args.speculative_num_draft_tokens extra_max_context_len += get_spec().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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)
@@ -567,8 +576,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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.enable_memory_saver,
need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"), need_sort=get_disagg().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,
@@ -588,7 +597,7 @@ class KVCacheConfigurator:
is_dsv4_model: bool, is_dsv4_model: bool,
current_platform, current_platform,
): ):
if not self.server_args.prefill_only_disable_kv_cache or self.is_draft_worker: if not get_schedule().prefill_only_disable_kv_cache or self.is_draft_worker:
return return
unsupported_pool_family = None unsupported_pool_family = None
@@ -623,9 +632,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 self.server_args.disaggregation_mode == "decode": if get_disagg().disaggregation_mode == "decode":
# Extra slots for pre-allocated requests # Extra slots for pre-allocated requests
pre_alloc_size = self.server_args.disaggregation_decode_extra_slots pre_alloc_size = get_disagg().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,
@@ -665,7 +674,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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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=(
[ [
@@ -675,16 +684,16 @@ 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=self.server_args.speculative_eagle_topk, speculative_eagle_topk=get_spec().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 self.server_args.disable_overlap_schedule, enable_overlap_schedule=not get_schedule().disable_overlap_schedule,
mamba_size=self.server_args.max_mamba_cache_size, mamba_size=get_schedule().max_mamba_cache_size,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len, linear_replayssm_cache_len=get_exec().mamba.linear_replayssm_cache_len,
mamba_envelope_layout=self.server_args.enable_page_major_kv_layout, mamba_envelope_layout=get_memory().enable_page_major_kv_layout,
enable_gdn_replayssm_spec=( enable_gdn_replayssm_spec=(
self.server_args.enable_gdn_replayssm_spec get_exec().mamba.enable_gdn_replayssm_spec
and self.hybrid_gdn_config is not None and self.hybrid_gdn_config is not None
), ),
) )
@@ -708,7 +717,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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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
@@ -721,11 +730,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=self.server_args.max_mamba_cache_size, mamba_size=get_schedule().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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=(
[ [
@@ -737,14 +746,14 @@ 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=self.server_args.speculative_eagle_topk, speculative_eagle_topk=get_spec().speculative_eagle_topk,
enable_overlap_schedule=not self.server_args.disable_overlap_schedule, enable_overlap_schedule=not get_schedule().disable_overlap_schedule,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
enable_linear_replayssm=self.server_args.enable_linear_replayssm, enable_linear_replayssm=get_exec().mamba.enable_linear_replayssm,
linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len, linear_replayssm_cache_len=get_exec().mamba.linear_replayssm_cache_len,
mamba_envelope_layout=self.server_args.enable_page_major_kv_layout, mamba_envelope_layout=get_memory().enable_page_major_kv_layout,
enable_gdn_replayssm_spec=( enable_gdn_replayssm_spec=(
self.server_args.enable_gdn_replayssm_spec get_exec().mamba.enable_gdn_replayssm_spec
and self.hybrid_gdn_config is not None and self.hybrid_gdn_config is not None
), ),
) )
@@ -770,7 +779,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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.enable_memory_saver,
) )
return req_to_token_pool return req_to_token_pool
@@ -786,7 +795,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 = self.server_args.enable_page_major_kv_layout enable_page_major = get_memory().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
) )
@@ -894,7 +903,7 @@ 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 = self.server_args.page_size swa_page_size = get_schedule().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."
@@ -928,12 +937,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=self.server_args.page_size, page_size=get_schedule().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=self.server_args.page_size, page_size=get_schedule().page_size,
max_num_reqs=max_running_requests, max_num_reqs=max_running_requests,
) )
else: else:
@@ -951,7 +960,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=self.server_args.page_size, page_size=get_schedule().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,
@@ -962,11 +971,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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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=self.server_args.enable_hisparse, enable_hisparse=get_memory().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
), ),
@@ -977,7 +986,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=self.server_args.page_size, page_size=get_schedule().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,
@@ -988,7 +997,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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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),
@@ -1001,14 +1010,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=self.server_args.page_size, page_size=get_schedule().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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,
) )
@@ -1018,13 +1027,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=self.server_args.page_size, page_size=get_schedule().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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,
) )
@@ -1055,7 +1064,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=self.server_args.page_size, page_size=get_schedule().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),
@@ -1077,14 +1086,14 @@ class KVCacheConfigurator:
token_to_kv_pool = NPUMLATokenToKVPool( token_to_kv_pool = NPUMLATokenToKVPool(
max_total_num_tokens, max_total_num_tokens,
page_size=self.server_args.page_size, page_size=get_schedule().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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,13 +1106,13 @@ class KVCacheConfigurator:
token_to_kv_pool = NPUMHATokenToKVPool( token_to_kv_pool = NPUMHATokenToKVPool(
max_total_num_tokens, max_total_num_tokens,
page_size=self.server_args.page_size, page_size=get_schedule().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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,
) )
@@ -1117,7 +1126,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 self.server_args.enable_hisparse: if get_memory().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
@@ -1137,7 +1146,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=self.server_args.page_size, page_size=get_schedule().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,
@@ -1148,7 +1157,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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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),
@@ -1159,13 +1168,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=self.server_args.page_size, page_size=get_schedule().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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,
) )
@@ -1174,13 +1183,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=self.server_args.page_size, page_size=get_schedule().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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,
) )
@@ -1207,7 +1216,7 @@ class KVCacheConfigurator:
} }
swa_pool_class = ( swa_pool_class = (
MHATokenToKVPoolMXFP8 MHATokenToKVPoolMXFP8
if get_model().kv_cache_dtype == "mxfp8" if self.kv_cache_dtype_str == "mxfp8"
else mha_pool_class else mha_pool_class
) )
swa_attention_layer_ids = self.model_config.swa_attention_layer_ids swa_attention_layer_ids = self.model_config.swa_attention_layer_ids
@@ -1237,7 +1246,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=self.server_args.page_size, page_size=get_schedule().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),
@@ -1245,7 +1254,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=(self.server_args.speculative_algorithm is not None), enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
token_to_kv_pool_class=swa_pool_class, token_to_kv_pool_class=swa_pool_class,
**kwargs, **kwargs,
) )
@@ -1260,7 +1269,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=self.server_args.page_size, page_size=get_schedule().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),
@@ -1270,7 +1279,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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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,
) )
@@ -1305,11 +1314,11 @@ class KVCacheConfigurator:
# buffers) for the full-attention layers, same as the SWA branch. # buffers) for the full-attention layers, same as the SWA branch.
full_pool_class = ( full_pool_class = (
MHATokenToKVPoolMXFP8 MHATokenToKVPoolMXFP8
if get_model().kv_cache_dtype == "mxfp8" and not self.use_mla_backend if self.kv_cache_dtype_str == "mxfp8" and not self.use_mla_backend
else mha_pool_class else mha_pool_class
) )
token_to_kv_pool = HybridLinearKVPool( token_to_kv_pool = HybridLinearKVPool(
page_size=self.server_args.page_size, page_size=get_schedule().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),
@@ -1318,8 +1327,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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.enable_memory_saver,
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), enable_kv_cache_copy=(get_spec().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,
@@ -1332,30 +1341,30 @@ 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=self.server_args.page_size, page_size=get_schedule().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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 self.server_args.enable_pdmux, enable_alt_stream=not get_disagg().enable_pdmux,
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
) )
return token_to_kv_pool return token_to_kv_pool
def _build_mha_kv_pool( def _build_mha_kv_pool(
self, *, max_total_num_tokens: int, mha_pool_class: type, quant_method=None self, *, max_total_num_tokens: int, mha_pool_class: type, quant_method=None
) -> KVCache: ) -> KVCache:
if get_model().kv_cache_dtype == "mxfp8": if self.kv_cache_dtype_str == "mxfp8":
pool_cls = MHATokenToKVPoolMXFP8 pool_cls = MHATokenToKVPoolMXFP8
else: else:
pool_cls = ( pool_cls = (
NoOpMHATokenToKVPool NoOpMHATokenToKVPool
if self.server_args.prefill_only_disable_kv_cache if get_schedule().prefill_only_disable_kv_cache
else mha_pool_class else mha_pool_class
) )
pool_kwargs = {} pool_kwargs = {}
@@ -1365,18 +1374,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=self.server_args.page_size, page_size=get_schedule().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=self.server_args.enable_memory_saver, enable_memory_saver=get_exec().features.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 self.server_args.enable_pdmux, enable_alt_stream=not get_disagg().enable_pdmux,
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
**pool_kwargs, **pool_kwargs,
) )
return token_to_kv_pool return token_to_kv_pool
@@ -1391,13 +1400,13 @@ 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 = self.server_args.disaggregation_mode in ("decode", "prefill") need_sort = get_disagg().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=self.server_args.page_size, page_size=get_schedule().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,
@@ -1422,7 +1431,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=self.server_args.page_size, page_size=get_schedule().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,7 +1444,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=self.server_args.page_size, page_size=get_schedule().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,
@@ -1445,7 +1454,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=self.server_args.page_size, page_size=get_schedule().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,
@@ -1455,14 +1464,14 @@ 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=self.server_args.page_size, page_size=get_schedule().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 self.server_args.enable_hisparse: if get_memory().enable_hisparse:
from sglang.srt.mem_cache.sparsity import ( from sglang.srt.mem_cache.sparsity import (
parse_hisparse_config, parse_hisparse_config,
) )
@@ -1470,7 +1479,7 @@ class KVCacheConfigurator:
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=self.server_args.page_size, page_size=get_schedule().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,
@@ -1478,8 +1487,7 @@ class KVCacheConfigurator:
host_to_device_ratio=hisparse_cfg.host_to_device_ratio, host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
) )
elif ( elif (
self.server_args.page_size == 1 get_schedule().page_size == 1 and self.server_args.dcp_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,
@@ -1491,7 +1499,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=self.server_args.page_size page_size=get_schedule().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,
@@ -1499,7 +1507,7 @@ class KVCacheConfigurator:
need_sort=need_sort, need_sort=need_sort,
) )
if self.server_args.enable_hisparse and is_dsv4_model: if get_memory().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
@@ -1551,7 +1559,7 @@ class KVCacheConfigurator:
cpu_group=get_world_group().cpu_group, cpu_group=get_world_group().cpu_group,
) )
slack_gb = pre_model_load_memory * (1 - self.server_args.mem_fraction_static) slack_gb = pre_model_load_memory * (1 - get_schedule().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(
@@ -1575,7 +1583,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={self.server_args.mem_fraction_static}. " f"--mem-fraction-static={get_schedule().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 = "
@@ -1586,7 +1594,7 @@ 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 self.server_args.disable_radix_cache: if get_memory().disable_radix_cache:
return 1 return 1
skip_decode_lock = envs.SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK.get() skip_decode_lock = envs.SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK.get()
@@ -1598,7 +1606,7 @@ class KVCacheConfigurator:
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 self.server_args.disable_overlap_schedule: if not get_schedule().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:
@@ -1622,7 +1630,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 = self.server_args.max_total_tokens user_limit = get_schedule().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:
@@ -1652,7 +1660,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 = self.server_args.max_running_requests max_num_reqs = get_schedule().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)
@@ -1663,13 +1671,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, self.server_args.max_mamba_cache_size // ratio max_num_reqs, get_schedule().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={self.server_args.max_mamba_cache_size}, " f"any requests. max_mamba_cache_size={get_schedule().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 "
@@ -1699,7 +1707,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 = self.server_args.mem_fraction_static config.mem_fraction_static = get_schedule().mem_fraction_static
return config return config
def config_from_budget( def config_from_budget(
@@ -1715,14 +1723,14 @@ 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, self.server_args.page_size budget_bytes, get_schedule().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, self.server_args.page_size max_tokens, get_schedule().page_size
) )
return config return config
@@ -1735,13 +1743,14 @@ class KVCacheConfigurator:
# The ring is allocated per slot but is not part of mamba_cache_per_req; # The ring is allocated per slot but is not part of mamba_cache_per_req;
# the solve must charge it too or num_slots is over-provisioned. # the solve must charge it too or num_slots is over-provisioned.
replayssm_active = ( replayssm_active = (
server_args.enable_gdn_replayssm_spec and self.hybrid_gdn_config is not None get_exec().mamba.enable_gdn_replayssm_spec
and self.hybrid_gdn_config is not None
) )
if replayssm_active: if replayssm_active:
record_len = ( record_len = (
server_args.max_speculative_num_draft_tokens server_args.max_speculative_num_draft_tokens
if server_args.max_speculative_num_draft_tokens is not None if server_args.max_speculative_num_draft_tokens is not None
else server_args.linear_replayssm_cache_len else get_exec().mamba.linear_replayssm_cache_len
) )
replayssm_ring_per_req = ( replayssm_ring_per_req = (
config.mamba2_cache_params.replayssm_ring_bytes_per_req( config.mamba2_cache_params.replayssm_ring_bytes_per_req(
@@ -1751,45 +1760,45 @@ class KVCacheConfigurator:
else: else:
replayssm_ring_per_req = 0 replayssm_ring_per_req = 0
if has_spec_dec: if has_spec_dec:
assert server_args.speculative_num_draft_tokens is not None assert get_spec().speculative_num_draft_tokens is not None
assert server_args.max_running_requests is not None assert get_schedule().max_running_requests is not None
if server_args.max_mamba_cache_size is not None: if get_schedule().max_mamba_cache_size is not None:
# Use explicitly set max_mamba_cache_size # Use explicitly set max_mamba_cache_size
server_args.override( get_context().override(
"mamba_pool.per_dp_shard", "mamba_pool.per_dp_shard",
max_mamba_cache_size=server_args.max_mamba_cache_size max_mamba_cache_size=get_schedule().max_mamba_cache_size
// self.ps.attn_dp_size, // self.ps.attn_dp_size,
) )
# Reserve intermediate memory based on capped max_num_reqs (+1 padding slot) # Reserve intermediate memory based on capped max_num_reqs (+1 padding slot)
if has_spec_dec and not replayssm_active: if has_spec_dec and not replayssm_active:
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, get_schedule().max_running_requests // self.ps.attn_dp_size,
server_args.max_mamba_cache_size // ratio, get_schedule().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
* (capped_reqs + 1) * (capped_reqs + 1)
* server_args.speculative_num_draft_tokens * get_spec().speculative_num_draft_tokens
) )
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
elif ( elif (
server_args.disable_radix_cache get_memory().disable_radix_cache
and server_args.max_running_requests is not None and get_schedule().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
server_args.override( get_context().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=get_schedule().max_running_requests
// self.ps.attn_dp_size, // self.ps.attn_dp_size,
) )
# Reserve intermediate memory based on capped max_num_reqs (+1 padding slot) # Reserve intermediate memory based on capped max_num_reqs (+1 padding slot)
if has_spec_dec and not replayssm_active: if has_spec_dec and not replayssm_active:
intermediate_size = ( intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req config.mamba2_cache_params.mamba_cache_per_req
* (server_args.max_mamba_cache_size + 1) * (get_schedule().max_mamba_cache_size + 1)
* server_args.speculative_num_draft_tokens * get_spec().speculative_num_draft_tokens
) )
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
else: else:
@@ -1802,16 +1811,16 @@ class KVCacheConfigurator:
# (K + 1) * per_req + (K / ratio + 1) * D * per_req = mamba_budget_bytes # (K + 1) * per_req + (K / ratio + 1) * D * per_req = mamba_budget_bytes
mamba_budget = ( mamba_budget = (
total_rest_memory total_rest_memory
* server_args.mamba_full_memory_ratio * get_schedule().mamba_full_memory_ratio
/ (1 + server_args.mamba_full_memory_ratio) / (1 + get_schedule().mamba_full_memory_ratio)
) )
mamba_budget_bytes = mamba_budget * (1 << 30) mamba_budget_bytes = mamba_budget * (1 << 30)
if has_spec_dec and not replayssm_active: if has_spec_dec and not replayssm_active:
ratio = self._calculate_mamba_ratio() ratio = self._calculate_mamba_ratio()
D = server_args.speculative_num_draft_tokens D = get_spec().speculative_num_draft_tokens
# Joint solve: main_state + intermediate = mamba_budget # Joint solve: main_state + intermediate = mamba_budget
server_args.override( get_context().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)) (mamba_budget_bytes - per_req * (1 + D))
@@ -1821,14 +1830,14 @@ class KVCacheConfigurator:
# Intermediate memory is included in mamba_budget, subtract it # Intermediate memory is included in mamba_budget, subtract it
# 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, get_schedule().max_running_requests // self.ps.attn_dp_size,
server_args.max_mamba_cache_size // ratio, get_schedule().max_mamba_cache_size // ratio,
) )
intermediate_size = per_req * (capped_reqs + 1) * D intermediate_size = per_req * (capped_reqs + 1) * D
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
else: else:
per_slot = per_req + replayssm_ring_per_req per_slot = per_req + replayssm_ring_per_req
server_args.override( get_context().override(
"mamba_pool.memory_budget", "mamba_pool.memory_budget",
max_mamba_cache_size=int( max_mamba_cache_size=int(
(mamba_budget_bytes - per_slot) // per_slot (mamba_budget_bytes - per_slot) // per_slot
@@ -1839,10 +1848,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 server_args.max_mamba_cache_size <= 0: if get_schedule().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={server_args.max_mamba_cache_size} " f"Computed max_mamba_cache_size={get_schedule().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, "
@@ -1853,7 +1862,7 @@ class KVCacheConfigurator:
# +1: the pool's padding slot # +1: the pool's padding slot
mamba_state_memory = ( mamba_state_memory = (
(server_args.max_mamba_cache_size + 1) (get_schedule().max_mamba_cache_size + 1)
* (config.mamba2_cache_params.mamba_cache_per_req + replayssm_ring_per_req) * (config.mamba2_cache_params.mamba_cache_per_req + replayssm_ring_per_req)
/ (1 << 30) / (1 << 30)
) )
@@ -42,6 +42,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
) )
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.mem_cache.storage.flexkv.flexkv_connector import FlexKVConnector from sglang.srt.mem_cache.storage.flexkv.flexkv_connector import FlexKVConnector
from sglang.srt.runtime_context import get_spec
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
@@ -393,7 +394,7 @@ class FlexKVRadixCache(RadixCache):
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_server_args
global_server_args = get_server_args() global_server_args = get_server_args()
topk = global_server_args.speculative_eagle_topk topk = get_spec().speculative_eagle_topk
enable_kv_committed_len = topk is None or topk == 1 enable_kv_committed_len = topk is None or topk == 1
if enable_kv_committed_len: if enable_kv_committed_len:
kv_committed_len = req.kv_committed_len kv_committed_len = req.kv_committed_len
@@ -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_server_args from sglang.srt.runtime_context import get_memory, get_server_args, get_spec
from sglang.srt.utils import create_device_stream, device_stream_context from sglang.srt.utils import create_device_stream, device_stream_context
try: try:
@@ -109,7 +109,7 @@ class LMCRadixCache(RadixCache):
): ):
super().__init__(params) super().__init__(params)
cli_lmc_cfg = get_server_args().lmcache_config_file or "" cli_lmc_cfg = get_memory().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(
@@ -448,7 +448,7 @@ class LMCRadixCache(RadixCache):
return return
global_server_args = get_server_args() global_server_args = get_server_args()
topk = global_server_args.speculative_eagle_topk topk = get_spec().speculative_eagle_topk
enable_kv_committed_len = topk is None or topk == 1 enable_kv_committed_len = topk is None or topk == 1
if enable_kv_committed_len: if enable_kv_committed_len:
kv_committed_len = req.kv_committed_len kv_committed_len = req.kv_committed_len
@@ -35,7 +35,7 @@ from sglang.srt.mem_cache.unified_cache.components.tree_component import (
TreeComponent, TreeComponent,
get_and_increase_time_counter, get_and_increase_time_counter,
) )
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_exec, get_server_args
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
@@ -66,7 +66,7 @@ class MambaComponent(TreeComponent):
), f"MambaComponent requires page_size=1 when mamba_extra_buffer is disabled, got {params.page_size}" ), f"MambaComponent requires page_size=1 when mamba_extra_buffer is disabled, got {params.page_size}"
super().__init__(cache, params) super().__init__(cache, params)
self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
self.mamba_max_states_per_path = get_server_args().mamba_max_states_per_path self.mamba_max_states_per_path = get_exec().mamba.mamba_max_states_per_path
# HiCache state # HiCache state
self._mamba_pool_host = None # set to host mamba pool when HiCache enabled self._mamba_pool_host = None # set to host mamba pool when HiCache enabled
@@ -37,7 +37,7 @@ from sglang.srt.model_executor.forward_batch_info import (
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode
from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.runtime_context import get_flags, get_parallel, get_spec
from sglang.srt.utils import ( from sglang.srt.utils import (
empty_context, empty_context,
log_info_on_rank0, log_info_on_rank0,
@@ -1027,9 +1027,9 @@ class CPUGraphRunner:
retrieve_next_token=None, retrieve_next_token=None,
retrieve_next_sibling=None, retrieve_next_sibling=None,
retrieve_cum_len=None, retrieve_cum_len=None,
spec_steps=self.model_runner.server_args.speculative_num_steps, spec_steps=get_spec().speculative_num_steps,
topk=self.model_runner.server_args.speculative_eagle_topk, topk=self.model_runner.server_args.speculative_eagle_topk,
draft_token_num=self.model_runner.server_args.speculative_num_draft_tokens, draft_token_num=get_spec().speculative_num_draft_tokens,
capture_hidden_mode=CaptureHiddenMode.FULL, capture_hidden_mode=CaptureHiddenMode.FULL,
seq_lens_sum=None, seq_lens_sum=None,
seq_lens_cpu=None, seq_lens_cpu=None,
@@ -16,7 +16,7 @@ cuda_graph_config, and the --cuda-graph-config JSON CLI parser.
Module-level imports are pure stdlib — no torch / sglang.srt deps — so Module-level imports are pure stdlib — no torch / sglang.srt deps — so
ServerArgs can import everything here without pulling in backend ServerArgs can import everything here without pulling in backend
classes. check_cuda_graph_backend lazy-imports get_server_args classes. check_cuda_graph_backend lazy-imports the config accessor
inside the function body to preserve that invariant. inside the function body to preserve that invariant.
""" """
@@ -179,15 +179,14 @@ def _diff_phase(actual: PhaseConfig, baseline: PhaseConfig) -> Dict[str, Any]:
def check_cuda_graph_backend(phase: str, backend: str) -> bool: def check_cuda_graph_backend(phase: str, backend: str) -> bool:
"""True if cuda_graph_config[phase].backend == backend on the """True if cuda_graph_config[phase].backend == backend on the
global server args. Returns False if the global server args have not published config. Returns False if the config has not been published
been initialized yet (e.g. unit tests, early startup).""" yet (e.g. unit tests, early startup)."""
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_exec
try: try:
server_args = get_server_args() cfg = get_exec().graph.cuda_graph_config
except ValueError: except ValueError:
return False return False
cfg = server_args.cuda_graph_config
if cfg is None or phase not in Phase.ALL: if cfg is None or phase not in Phase.ALL:
return False return False
return getattr(cfg, phase).backend == backend return getattr(cfg, phase).backend == backend
@@ -51,7 +51,7 @@ 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_parallel, get_server_args from sglang.srt.runtime_context import get_exec, get_parallel
from sglang.srt.utils import ( from sglang.srt.utils import (
is_cuda, is_cuda,
is_hip, is_hip,
@@ -965,7 +965,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_server_args().enable_mis and any( if get_exec().features.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(
@@ -1134,7 +1134,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_server_args().rl_on_policy_target rl_on_policy_target = get_exec().deterministic.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():
@@ -162,9 +162,13 @@ from sglang.srt.model_executor.runner import (
) )
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_context,
get_exec,
get_global_dwdp_manager, get_global_dwdp_manager,
get_lora,
get_model,
get_parallel, get_parallel,
get_server_args, 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
@@ -322,7 +326,7 @@ class ModelRunner:
self.init_threads_binding() self.init_threads_binding()
# Set float32 matmul precision # Set float32 matmul precision
if get_server_args().enable_tf32_matmul: if get_exec().features.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)
@@ -399,7 +403,7 @@ class ModelRunner:
def _initialize_elastic_ep_joiner(self) -> None: def _initialize_elastic_ep_joiner(self) -> None:
if not ( if not (
self.server_args.elastic_ep_backend is not None get_exec().moe.elastic_ep_backend is not None
and self.server_args.is_ep_scale_joiner and self.server_args.is_ep_scale_joiner
): ):
return return
@@ -473,7 +477,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=self.server_args.custom_weight_loader, custom_weight_loaders=get_model().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,
@@ -527,6 +531,7 @@ class ModelRunner:
model_config=self.model_config, model_config=self.model_config,
server_args=self.server_args, server_args=self.server_args,
kv_cache_dtype=self.kv_cache_dtype, kv_cache_dtype=self.kv_cache_dtype,
kv_cache_dtype_str=self.kv_cache_dtype_str,
model_dtype=self.dtype, model_dtype=self.dtype,
page_size=self.page_size, page_size=self.page_size,
sliding_window_size=self.sliding_window_size, sliding_window_size=self.sliding_window_size,
@@ -550,7 +555,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 self.server_args.model_impl.lower() == ModelImpl.MINDSPORE and _is_npu: if get_model().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 +612,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=self.server_args.enable_memory_saver enable=get_exec().features.enable_memory_saver
) )
def maybe_init_remote_instance_transfer_engine(self): def maybe_init_remote_instance_transfer_engine(self):
@@ -643,7 +648,7 @@ class ModelRunner:
) )
def maybe_init_lplb_solvers(self): def maybe_init_lplb_solvers(self):
if self.server_args.ep_dispatch_algorithm == "lp" and not self.is_draft_worker: if get_exec().moe.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 +662,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 self.server_args.enable_eplb and (not self.is_draft_worker) if get_exec().moe.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 self.server_args.elastic_ep_backend: if get_exec().moe.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 +686,8 @@ class ModelRunner:
get_model=lambda: self.model, get_model=lambda: self.model,
) )
if ( if (
self.server_args.enable_elastic_expert_backup get_exec().moe.enable_elastic_expert_backup
and self.server_args.elastic_ep_backend is not None and get_exec().moe.elastic_ep_backend is not None
) )
else None else None
) )
@@ -691,17 +696,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_server_args().torchao_config) apply_torchao_config_to_model(self.model, get_exec().graph.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 self.server_args.enable_lora: if get_lora().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 self.server_args.enable_deterministic_inference: if get_exec().deterministic.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()
@@ -996,8 +1001,8 @@ class ModelRunner:
remote_instance_weight_transporter_engine=self.remote_instance_weight_transporter.engine, remote_instance_weight_transporter_engine=self.remote_instance_weight_transporter.engine,
remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id, remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id,
draft_model_idx=self.draft_model_idx, draft_model_idx=self.draft_model_idx,
weight_cache_mode=self.server_args.weight_cache_mode, weight_cache_mode=get_model().weight_cache_mode,
weight_cache_socket=self.server_args.weight_cache_socket, weight_cache_socket=get_model().weight_cache_socket,
) )
# If the weight cache is enabled, override the load format to IPC_CACHE # If the weight cache is enabled, override the load format to IPC_CACHE
@@ -1038,7 +1043,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 self.server_args.enable_layerwise_nvtx_marker: if get_exec().comm.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")
@@ -1095,7 +1100,7 @@ class ModelRunner:
) )
dist_barrier_after_load( dist_barrier_after_load(
elastic_ep_backend=self.server_args.elastic_ep_backend, elastic_ep_backend=get_exec().moe.elastic_ep_backend,
tp_rank=self.ps.tp_rank, tp_rank=self.ps.tp_rank,
is_ep_joiner=self.server_args.is_ep_joiner, is_ep_joiner=self.server_args.is_ep_joiner,
) )
@@ -1115,16 +1120,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=self.server_args.max_loras_per_batch, max_loras_per_batch=get_lora().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=self.server_args.lora_backend, lora_backend=get_lora().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=self.server_args.max_lora_rank, max_lora_rank=get_lora().max_lora_rank,
target_modules=self.server_args.lora_target_modules, target_modules=get_lora().lora_target_modules,
lora_paths=self.server_args.lora_paths, lora_paths=get_lora().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(
@@ -1157,29 +1162,10 @@ class ModelRunner:
else: else:
return self.max_total_num_tokens return self.max_total_num_tokens
def _record_kv_cache_dtype(self, resolved: str) -> None:
# the weight-resolved kv-cache dtype is written to the config
# bags via get_context().override, so get_model().kv_cache_dtype readers
# see it. server_args stays the pristine RAW record -- configure_kv_cache
# _dtype reads it as the resolver INPUT. A draft / mock runner whose
# server_args is not the published object keeps the private-bag write.
from sglang.srt.runtime_context import get_context
if get_context()._server_args is self.server_args:
get_context().override(
"ModelRunner.configure_kv_cache_dtype", kv_cache_dtype=resolved
)
else:
self.server_args.override(
"ModelRunner.configure_kv_cache_dtype", kv_cache_dtype=resolved
)
def configure_kv_cache_dtype(self): def configure_kv_cache_dtype(self):
spec_algorithm = getattr(self, "spec_algorithm", None) spec_algorithm = getattr(self, "spec_algorithm", None)
resolved_kv_cache_dtype, self.kv_cache_dtype = ( resolved_kv_cache_dtype, self.kv_cache_dtype = (
kv_cache_dtype.configure_kv_cache_dtype( kv_cache_dtype.configure_kv_cache_dtype(
# RAW user intent = resolver INPUT; server_args stays pristine
# so read it here -- not the resolved get_model() bag.
server_args_kv_cache_dtype=self.server_args.kv_cache_dtype, server_args_kv_cache_dtype=self.server_args.kv_cache_dtype,
model=getattr(self, "model", None), model=getattr(self, "model", None),
model_dtype=getattr(self, "dtype", torch.bfloat16), model_dtype=getattr(self, "dtype", torch.bfloat16),
@@ -1201,8 +1187,6 @@ class ModelRunner:
if resolved_kv_cache_dtype is not None if resolved_kv_cache_dtype is not None
else self.server_args.kv_cache_dtype else self.server_args.kv_cache_dtype
) )
if resolved_kv_cache_dtype is not None:
self._record_kv_cache_dtype(resolved_kv_cache_dtype)
def _get_attention_backend(self, init_new_workspace: bool = False): def _get_attention_backend(self, init_new_workspace: bool = False):
return get_attention_backend( return get_attention_backend(
@@ -1391,7 +1375,7 @@ class ModelRunner:
) )
output.expert_distribution_metrics = recorder_outputs.get("metrics") output.expert_distribution_metrics = recorder_outputs.get("metrics")
no_copy_to_cpu = not self.server_args.disable_overlap_schedule no_copy_to_cpu = not get_schedule().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
@@ -1421,7 +1405,7 @@ class ModelRunner:
self.msprobe_debugger.stop() self.msprobe_debugger.stop()
self.msprobe_debugger.step() self.msprobe_debugger.step()
if self.server_args.elastic_ep_backend is not None: if get_exec().moe.elastic_ep_backend is not None:
self.maybe_join_ep_ranks() self.maybe_join_ep_ranks()
return output return output
@@ -1852,7 +1836,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
> self.server_args.elastic_ep_scale_timeout > get_exec().moe.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)
@@ -1922,7 +1906,7 @@ class ModelRunner:
load_config: LoadConfig, load_config: LoadConfig,
) -> None: ) -> None:
self.model = new_model self.model = new_model
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,
@@ -11,6 +11,7 @@ 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
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
@@ -58,7 +59,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 self.server_args.remote_instance_weight_loader_backend and get_model().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
@@ -84,7 +85,7 @@ class RemoteInstanceWeightTransporter:
else: else:
bootstrap_host = "127.0.0.1" bootstrap_host = "127.0.0.1"
bootstrap_port = self.server_args.engine_info_bootstrap_port bootstrap_port = get_model().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"
@@ -33,7 +33,7 @@ from sglang.srt.environ import envs
from sglang.srt.mem_cache.allocation_sizing import get_alloc_len_per_decode from sglang.srt.mem_cache.allocation_sizing import get_alloc_len_per_decode
from sglang.srt.mem_cache.deepseek_v4_memory_pool import get_compress_state_ring_size from sglang.srt.mem_cache.deepseek_v4_memory_pool import get_compress_state_ring_size
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
from sglang.srt.runtime_context import get_model, get_parallel from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
ceil_align, ceil_align,
ceil_div, ceil_div,
@@ -119,6 +119,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
""" """
def __init__(self, kvc: KVCacheConfigurator): def __init__(self, kvc: KVCacheConfigurator):
self.kv_cache_dtype_str = kvc.kv_cache_dtype_str
# Determine effective number of layers for KV cache # Determine effective number of layers for KV cache
if mambaish := mambaish_config(kvc.model_config): if mambaish := mambaish_config(kvc.model_config):
effective_layer_ids = [ effective_layer_ids = [
@@ -304,7 +305,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
) )
# FP4 prefill uses one shared FP8 dequant workspace across layers. # FP4 prefill uses one shared FP8 dequant workspace across layers.
cell_size += n * k * 2 * kv_size cell_size += n * k * 2 * kv_size
elif get_model().kv_cache_dtype == "mxfp8": elif self.kv_cache_dtype_str == "mxfp8":
scale_block_size = 32 scale_block_size = 32
n = model_config.get_num_kv_heads(tp_size) n = model_config.get_num_kv_heads(tp_size)
cell_size += ( cell_size += (
@@ -339,6 +340,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
""" """
def __init__(self, kvc: KVCacheConfigurator): def __init__(self, kvc: KVCacheConfigurator):
self.kv_cache_dtype_str = kvc.kv_cache_dtype_str
model_config = kvc.model_config model_config = kvc.model_config
kv_cache_dtype = kvc.kv_cache_dtype kv_cache_dtype = kvc.kv_cache_dtype
kv_size = torch._utils._element_size(kv_cache_dtype) kv_size = torch._utils._element_size(kv_cache_dtype)
@@ -368,7 +370,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
* kv_size * kv_size
) )
if get_model().kv_cache_dtype == "mxfp8": if self.kv_cache_dtype_str == "mxfp8":
scale_block_size = 32 scale_block_size = 32
self._full_per_token += ( self._full_per_token += (
model_config.get_num_kv_heads(tp_size) model_config.get_num_kv_heads(tp_size)
@@ -501,6 +503,7 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
""" """
def __init__(self, kvc: KVCacheConfigurator): def __init__(self, kvc: KVCacheConfigurator):
self.kv_cache_dtype_str = kvc.kv_cache_dtype_str
super().__init__(kvc) super().__init__(kvc)
assert self._full_layers_num > 0 assert self._full_layers_num > 0
@@ -613,6 +616,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
""" """
def __init__(self, kvc: KVCacheConfigurator): def __init__(self, kvc: KVCacheConfigurator):
self.kv_cache_dtype_str = kvc.kv_cache_dtype_str
cfg = kvc.model_config cfg = kvc.model_config
self.qk_nope_head_dim = cfg.qk_nope_head_dim self.qk_nope_head_dim = cfg.qk_nope_head_dim
self.qk_rope_head_dim = cfg.qk_rope_head_dim self.qk_rope_head_dim = cfg.qk_rope_head_dim
@@ -91,7 +91,7 @@ from sglang.srt.model_executor.runner_utils.deepep_adapter import (
DeepEPCudaGraphRunnerAdapter, DeepEPCudaGraphRunnerAdapter,
) )
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.runtime_context import get_flags, get_parallel, get_spec
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
from sglang.srt.utils import ( from sglang.srt.utils import (
empty_context, empty_context,
@@ -246,12 +246,12 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.is_dllm = self.dllm_config is not None self.is_dllm = self.dllm_config is not None
self.attn_backend = attn_backend or model_runner.attn_backend self.attn_backend = attn_backend or model_runner.attn_backend
self.speculative_num_steps = ( self.speculative_num_steps = (
model_runner.server_args.speculative_num_steps get_spec().speculative_num_steps
if speculative_num_steps is None if speculative_num_steps is None
else speculative_num_steps else speculative_num_steps
) )
self.speculative_num_draft_tokens = ( self.speculative_num_draft_tokens = (
model_runner.server_args.speculative_num_draft_tokens get_spec().speculative_num_draft_tokens
if speculative_num_draft_tokens is None if speculative_num_draft_tokens is None
else speculative_num_draft_tokens else speculative_num_draft_tokens
) )
+14 -15
View File
@@ -47,7 +47,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_server_args from sglang.srt.runtime_context import get_exec, get_model, 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)
@@ -495,10 +495,10 @@ class DefaultModelLoader(BaseModelLoader):
hf_folder = model_name_or_path hf_folder = model_name_or_path
server_args = get_server_args() server_args = get_server_args()
if server_args and server_args.model_checksum is not None: if server_args and get_model().model_checksum is not None:
from sglang.srt.utils.model_file_verifier import verify from sglang.srt.utils.model_file_verifier import verify
checksums_source = server_args.model_checksum or model_name_or_path checksums_source = get_model().model_checksum or model_name_or_path
verify(model_path=hf_folder, checksums_source=checksums_source) verify(model_path=hf_folder, checksums_source=checksums_source)
hf_weights_files: List[str] = [] hf_weights_files: List[str] = []
@@ -581,11 +581,11 @@ class DefaultModelLoader(BaseModelLoader):
) )
elif use_safetensors: elif use_safetensors:
server_args = get_server_args() server_args = get_server_args()
weight_loader_disable_mmap = server_args.weight_loader_disable_mmap weight_loader_disable_mmap = get_model().weight_loader_disable_mmap
weight_loader_prefetch = server_args.weight_loader_prefetch_checkpoints weight_loader_prefetch = get_model().weight_loader_prefetch_checkpoints
prefetch_num_threads = server_args.weight_loader_prefetch_num_threads prefetch_num_threads = get_model().weight_loader_prefetch_num_threads
weight_loader_drop_cache_after_load = ( weight_loader_drop_cache_after_load = (
server_args.weight_loader_drop_cache_after_load get_model().weight_loader_drop_cache_after_load
) )
# Prefetch and multi-threaded loading both read the same shards, # Prefetch and multi-threaded loading both read the same shards,
@@ -879,9 +879,8 @@ 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_server_args().torchao_config torchao_config = get_exec().graph.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)
@@ -1751,13 +1750,13 @@ class PreshardedModelLoader(DefaultModelLoader):
"moe_dense_tp_size": server_args.moe_dense_tp_size, "moe_dense_tp_size": server_args.moe_dense_tp_size,
"moe_dp_size": server_args.moe_dp_size, "moe_dp_size": server_args.moe_dp_size,
"enable_dp_lm_head": server_args.enable_dp_lm_head, "enable_dp_lm_head": server_args.enable_dp_lm_head,
"enable_fp32_lm_head": server_args.enable_fp32_lm_head, "enable_fp32_lm_head": get_exec().features.enable_fp32_lm_head,
"quantization": model_config.quantization, "quantization": model_config.quantization,
"model_dtype": str(model_config.dtype), "model_dtype": str(model_config.dtype),
"ep_num_redundant_experts": server_args.ep_num_redundant_experts, "ep_num_redundant_experts": get_exec().moe.ep_num_redundant_experts,
"enable_eplb": server_args.enable_eplb, "enable_eplb": get_exec().moe.enable_eplb,
"init_expert_location": self._normalize_init_expert_location( "init_expert_location": self._normalize_init_expert_location(
server_args.init_expert_location get_exec().moe.init_expert_location
), ),
"structural_signature": self._compute_structural_signature(model_config), "structural_signature": self._compute_structural_signature(model_config),
} }
@@ -3934,10 +3933,10 @@ class RunaiModelStreamerLoader(BaseModelLoader):
) )
server_args = get_server_args() server_args = get_server_args()
if server_args and server_args.model_checksum is not None: if server_args and get_model().model_checksum is not None:
from sglang.srt.utils.model_file_verifier import verify from sglang.srt.utils.model_file_verifier import verify
checksums_source = server_args.model_checksum or model_name_or_path checksums_source = get_model().model_checksum or model_name_or_path
verify(model_path=hf_folder, checksums_source=checksums_source) verify(model_path=hf_folder, checksums_source=checksums_source)
hf_weights_files = list_safetensors(path=hf_folder) hf_weights_files = list_safetensors(path=hf_folder)
+3 -4
View File
@@ -78,6 +78,7 @@ 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_server_args,
@@ -209,7 +210,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_server_args().ep_num_redundant_experts == 0 assert get_exec().moe.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)
@@ -223,9 +224,7 @@ 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 = ( self.num_experts = config.num_experts + get_exec().moe.ep_num_redundant_experts
config.num_experts + get_server_args().ep_num_redundant_experts
)
self.gate = BailingMoEGate( self.gate = BailingMoEGate(
config=config, config=config,
@@ -59,6 +59,7 @@ 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_server_args,
@@ -529,7 +530,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_server_args().device, device=get_device().device,
dtype=torch.float32, dtype=torch.float32,
) )
@@ -690,7 +691,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_server_args().device, device=get_device().device,
) )
self.attn = RadixAttention( self.attn = RadixAttention(
self.num_heads, self.num_heads,
+2 -4
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_parallel, get_server_args from sglang.srt.runtime_context import get_model, get_parallel
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
BertConfig = None BertConfig = None
@@ -365,9 +365,7 @@ class BertModel(nn.Module):
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("encoder", prefix), prefix=add_prefix("encoder", prefix),
) )
pooling_type = ( pooling_type = PoolingType.CLS if get_model().is_embedding else PoolingType.LAST
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, get_server_args from sglang.srt.runtime_context import get_exec
from sglang.srt.utils import is_sm100_or_sm110_supported, use_intel_amx_backend from sglang.srt.utils import is_sm100_or_sm110_supported, 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_server_args().enable_deterministic_inference: if get_exec().deterministic.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")
@@ -194,7 +194,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_server_args().enable_deterministic_inference: if get_exec().deterministic.enable_deterministic_inference:
return _dispatch_mla_subtype(attn, forward_batch) return _dispatch_mla_subtype(attn, forward_batch)
if ( if (

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