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

This commit is contained in:
Cheng Wan
2026-07-22 01:18:05 -07:00
committed by GitHub
parent 1d0a6ee178
commit 11a4c2d057
162 changed files with 1103 additions and 1209 deletions
@@ -39,7 +39,12 @@ 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_exec,
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
@@ -183,7 +188,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:
@@ -335,7 +340,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):
@@ -633,7 +638,7 @@ class TboForwardBatchPreparer:
sum_field=None, sum_field=None,
) )
_, child_b.extend_start_loc = compute_position( _, child_b.extend_start_loc = compute_position(
get_server_args().attention_backend, get_exec().kernel.attention_backend,
child_b.extend_prefix_lens, child_b.extend_prefix_lens,
child_b.extend_seq_lens, child_b.extend_seq_lens,
child_b.extend_num_tokens, child_b.extend_num_tokens,
@@ -832,7 +837,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 (
@@ -14,6 +14,7 @@ from sglang.srt.constrained.base_grammar_backend import (
from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject
from sglang.srt.distributed.communication_tags import P2PTag from sglang.srt.distributed.communication_tags import P2PTag
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_serving
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.io_struct import AbortReq from sglang.srt.managers.io_struct import AbortReq
@@ -28,7 +29,7 @@ class GrammarManager:
self.scheduler = scheduler self.scheduler = scheduler
self.server_args = scheduler.server_args self.server_args = scheduler.server_args
self.grammar_queue: List[Req] = [] self.grammar_queue: List[Req] = []
if not self.server_args.skip_tokenizer_init: if not get_serving().skip_tokenizer_init:
self.grammar_backend = create_grammar_backend( self.grammar_backend = create_grammar_backend(
self.server_args, self.server_args,
scheduler.tokenizer, scheduler.tokenizer,
@@ -32,11 +32,8 @@ from sglang.srt.disaggregation.utils import (
) )
from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.distributed import get_pp_group, get_world_group
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import get_attention_dp_rank, get_attention_dp_size
get_attention_dp_rank, from sglang.srt.runtime_context import get_model, get_parallel, get_serving
get_attention_dp_size,
)
from sglang.srt.runtime_context import get_model, get_parallel
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import ( from sglang.srt.utils.network import (
NetworkAddress, NetworkAddress,
@@ -634,7 +631,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
+4 -6
View File
@@ -86,7 +86,7 @@ from sglang.srt.observability.req_time_stats import (
set_schedule_time_batch, set_schedule_time_batch,
set_time_batch, set_time_batch,
) )
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_disagg, get_parallel
from sglang.srt.utils import get_num_new_pages from sglang.srt.utils import get_num_new_pages
from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.network import NetworkAddress
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
@@ -2151,7 +2151,7 @@ class SchedulerDisaggregationDecodeMixin:
# Decode-radix path: new requests already matched in # Decode-radix path: new requests already matched in
# `pop_preallocated`. Retracted requests reset `last_node`, # `pop_preallocated`. Retracted requests reset `last_node`,
# so re-match only when that state is missing. # so re-match only when that state is missing.
if 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
@@ -2191,7 +2191,7 @@ class SchedulerDisaggregationDecodeMixin:
if self.enable_decode_hicache: if self.enable_decode_hicache:
self.tree_cache.check_hicache_events() self.tree_cache.check_hicache_events()
if 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
@@ -2203,9 +2203,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
) )
@@ -59,14 +59,9 @@ from sglang.srt.model_loader import get_model
from sglang.srt.multimodal.processors.qwen_vl import preprocess_video from sglang.srt.multimodal.processors.qwen_vl import preprocess_video
from sglang.srt.observability.metrics_collector import EncoderMetricsCollector from sglang.srt.observability.metrics_collector import EncoderMetricsCollector
from sglang.srt.observability.req_time_stats import EncoderReqTimeStats from sglang.srt.observability.req_time_stats import EncoderReqTimeStats
from sglang.srt.observability.trace import ( from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
process_tracing_init, from sglang.srt.runtime_context import get_disagg, get_exec, get_mm
trace_set_thread_info, from sglang.srt.server_args import PortArgs, ServerArgs
)
from sglang.srt.server_args import (
PortArgs,
ServerArgs,
)
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prometheus_middleware, add_prometheus_middleware,
configure_logger, configure_logger,
@@ -350,7 +345,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,
) )
@@ -368,15 +363,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()
@@ -389,8 +384,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
), ),
) )
@@ -399,7 +394,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
@@ -413,12 +408,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
@@ -1688,7 +1683,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])
@@ -1784,7 +1779,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:
@@ -1855,7 +1850,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)
@@ -1887,11 +1882,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(
@@ -55,6 +55,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
@@ -314,9 +315,7 @@ class MooncakeKVManager(CommonKVManager):
self.kv_buffer_tensors = None self.kv_buffer_tensors = None
def _handle_staging_req(self, msg): def _handle_staging_req(self, msg):
from sglang.srt.disaggregation.common.staging_handler import ( from sglang.srt.disaggregation.common.staging_handler import handle_staging_req
handle_staging_req,
)
room = int(msg[1].decode("ascii")) room = int(msg[1].decode("ascii"))
session_id = msg[4].decode("ascii") session_id = msg[4].decode("ascii")
@@ -350,9 +349,7 @@ class MooncakeKVManager(CommonKVManager):
def _is_watermark_ready( def _is_watermark_ready(
self, session_id: str, alloc_round: int, alloc_end: int self, session_id: str, alloc_round: int, alloc_end: int
) -> bool: ) -> bool:
from sglang.srt.disaggregation.common.staging_handler import ( from sglang.srt.disaggregation.common.staging_handler import is_watermark_ready
is_watermark_ready,
)
return is_watermark_ready(self._staging_ctx, session_id, alloc_round, alloc_end) return is_watermark_ready(self._staging_ctx, session_id, alloc_round, alloc_end)
@@ -469,7 +466,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,
) )
+6 -10
View File
@@ -13,6 +13,8 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
from sglang.srt.runtime_context import get_schedule
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo
@@ -535,9 +537,7 @@ class NixlKVManager(CommonKVManager):
def _is_watermark_ready( def _is_watermark_ready(
self, agent_name: str, alloc_round: int, alloc_end: int self, agent_name: str, alloc_round: int, alloc_end: int
) -> bool: ) -> bool:
from sglang.srt.disaggregation.common.staging_handler import ( from sglang.srt.disaggregation.common.staging_handler import is_watermark_ready
is_watermark_ready,
)
return is_watermark_ready(self._staging_ctx, agent_name, alloc_round, alloc_end) return is_watermark_ready(self._staging_ctx, agent_name, alloc_round, alloc_end)
@@ -558,9 +558,7 @@ class NixlKVManager(CommonKVManager):
threading.Thread(target=decode_staging_thread, daemon=True).start() threading.Thread(target=decode_staging_thread, daemon=True).start()
def _handle_staging_req(self, msg): def _handle_staging_req(self, msg):
from sglang.srt.disaggregation.common.staging_handler import ( from sglang.srt.disaggregation.common.staging_handler import handle_staging_req
handle_staging_req,
)
room = int(msg[1].decode("ascii")) room = int(msg[1].decode("ascii"))
session_id = msg[4].decode("ascii") session_id = msg[4].decode("ascii")
@@ -625,7 +623,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,
) )
@@ -1739,9 +1737,7 @@ class NixlKVManager(CommonKVManager):
req, page_start, num_pages, session_id=req.agent_name req, page_start, num_pages, session_id=req.agent_name
) )
if not ready: if not ready:
from sglang.srt.disaggregation.common.staging_buffer import ( from sglang.srt.disaggregation.common.staging_buffer import StagingAllocator
StagingAllocator,
)
if c_offset == StagingAllocator.ALLOC_OVERSIZED: if c_offset == StagingAllocator.ALLOC_OVERSIZED:
raise RuntimeError( raise RuntimeError(
+2 -1
View File
@@ -64,6 +64,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.nvtx_utils import scheduler_nvtx_method from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -1181,7 +1182,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,
) )
+1 -1
View File
@@ -442,7 +442,7 @@ def get_healthy_expert_location_src_rank(
*, invoked_in_elastic_ep_rejoin_path: bool *, invoked_in_elastic_ep_rejoin_path: bool
) -> int: ) -> int:
world_group = get_world_group() world_group = get_world_group()
# NOTE: do not key off `self.server_args.elastic_ep_rejoin` here. # NOTE: do not key off `get_exec().moe.elastic_ep_rejoin` here.
# A rank that was started as a rejoin rank may later act as a healthy # A rank that was started as a rejoin rank may later act as a healthy
# rank in a subsequent recovery cycle. # rank in a subsequent recovery cycle.
local_rejoin_flag = bool(invoked_in_elastic_ep_rejoin_path) local_rejoin_flag = bool(invoked_in_elastic_ep_rejoin_path)
@@ -7,13 +7,11 @@ from typing import Any, Callable
import torch import torch
import zmq import zmq
from sglang.srt.distributed.parallel_state import ( from sglang.srt.distributed.parallel_state import get_world_group, get_world_size
get_world_group,
get_world_size,
)
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send
from sglang.srt.runtime_context import get_exec
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import get_local_ip_auto from sglang.srt.utils.network import get_local_ip_auto
@@ -111,7 +109,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):
+15 -2
View File
@@ -877,7 +877,14 @@ class Engine(EngineScoreMixin, EngineBase):
server_args, port_args server_args, port_args
) )
else: else:
# Launch multi-tokenizer router # Launch multi-tokenizer router. Unlike TokenizerManager, the router
# does not publish; but it runs in this parent process and reads
# resolved config through the namespace accessors (e.g. get_parallel()
# for routed_dp_rank), so publish here. The child TokenizerWorkers
# publish independently in their own processes.
from sglang.srt.runtime_context import publish
publish(server_args, role="tokenizer")
tokenizer_manager = MultiTokenizerRouter(server_args, port_args) tokenizer_manager = MultiTokenizerRouter(server_args, port_args)
template_manager = None template_manager = None
@@ -997,12 +1004,18 @@ class Engine(EngineScoreMixin, EngineBase):
) )
def get_server_info(self): def get_server_info(self):
from sglang.srt.runtime_context import get_context
internal_states = self.loop.run_until_complete( internal_states = self.loop.run_until_complete(
self.tokenizer_manager.get_internal_state() self.tokenizer_manager.get_internal_state()
) )
return msgspec_to_builtins( return msgspec_to_builtins(
{ {
**dataclasses.asdict(self.tokenizer_manager.server_args), # Overlay post-publish overrides so the report reflects current
# config (weight version, model path, runtime tunables).
**get_context().resolved_server_args_dict(
base=dataclasses.asdict(self.tokenizer_manager.server_args)
),
**self._scheduler_init_result.scheduler_infos[0], **self._scheduler_init_result.scheduler_infos[0],
"internal_states": internal_states, "internal_states": internal_states,
"version": __version__, "version": __version__,
+10 -9
View File
@@ -15,6 +15,7 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional
from pydantic import ValidationError from pydantic import ValidationError
from sglang.srt.runtime_context import get_context, get_lora, get_serving
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -229,9 +230,7 @@ class RuntimeHandle:
return self._openai_serving_classes return self._openai_serving_classes
from sglang.srt.entrypoints.openai.serving_chat import OpenAIServingChat from sglang.srt.entrypoints.openai.serving_chat import OpenAIServingChat
from sglang.srt.entrypoints.openai.serving_classify import ( from sglang.srt.entrypoints.openai.serving_classify import OpenAIServingClassify
OpenAIServingClassify,
)
from sglang.srt.entrypoints.openai.serving_completions import ( from sglang.srt.entrypoints.openai.serving_completions import (
OpenAIServingCompletion, OpenAIServingCompletion,
) )
@@ -376,16 +375,20 @@ 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": get_serving().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": get_serving().weight_version,
"model_type": getattr(model_config.hf_config, "model_type", None), "model_type": getattr(model_config.hf_config, "model_type", None),
"architectures": getattr(model_config.hf_config, "architectures", None), "architectures": getattr(model_config.hf_config, "architectures", None),
} }
return json.dumps(result, default=str) return json.dumps(result, default=str)
def get_server_info(self) -> str: def get_server_info(self) -> str:
result: Dict[str, Any] = dataclasses.asdict(self.server_args) # Overlay post-publish overrides (weight version, model path, runtime
# tunables) so the report reflects current config, not the startup record.
result: Dict[str, Any] = get_context().resolved_server_args_dict(
base=dataclasses.asdict(self.server_args)
)
result.update(self.scheduler_info) result.update(self.scheduler_info)
return json.dumps(msgspec_to_builtins(result), default=str) return json.dumps(msgspec_to_builtins(result), default=str)
@@ -424,9 +427,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 get_lora().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
for _, lora_ref in lora_registry.get_all_adapters().items(): for _, lora_ref in lora_registry.get_all_adapters().items():
models.append( models.append(
+13 -4
View File
@@ -693,18 +693,19 @@ async def get_model_info():
@app.get("/model_info") @app.get("/model_info")
async def model_info(): async def model_info():
"""Get the model information.""" """Get the model information."""
from sglang.srt.runtime_context import get_serving
model_config = _global_state.tokenizer_manager.model_config model_config = _global_state.tokenizer_manager.model_config
result = { result = {
"model_path": _global_state.tokenizer_manager.model_path, "model_path": _global_state.tokenizer_manager.model_path,
"tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path, "tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path,
"is_generation": _global_state.tokenizer_manager.is_generation, "is_generation": _global_state.tokenizer_manager.is_generation,
"preferred_sampling_params": _global_state.tokenizer_manager.server_args.preferred_sampling_params, "preferred_sampling_params": _global_state.tokenizer_manager.server_args.preferred_sampling_params,
"weight_version": _global_state.tokenizer_manager.server_args.weight_version, "weight_version": get_serving().weight_version,
"has_image_understanding": model_config.is_image_understandable_model, "has_image_understanding": model_config.is_image_understandable_model,
"has_audio_understanding": model_config.is_audio_understandable_model, "has_audio_understanding": model_config.is_audio_understandable_model,
"model_type": getattr(model_config.hf_config, "model_type", None), "model_type": getattr(model_config.hf_config, "model_type", None),
"architectures": getattr(model_config.hf_config, "architectures", None), "architectures": getattr(model_config.hf_config, "architectures", None),
"weight_version": _global_state.tokenizer_manager.server_args.weight_version,
# "hf_config": model_config.hf_config.to_dict(), # "hf_config": model_config.hf_config.to_dict(),
} }
return result return result
@@ -738,12 +739,18 @@ async def server_info():
await _global_state.tokenizer_manager.get_internal_state() await _global_state.tokenizer_manager.get_internal_state()
) )
from sglang.srt.runtime_context import get_context
server_args = _global_state.tokenizer_manager.server_args server_args = _global_state.tokenizer_manager.server_args
# server_args.model_config is not serializable but should be excluded by asdict. # server_args.model_config is not serializable but should be excluded by asdict.
# Overlay post-publish overrides so runtime updates (weight version, model
# path/load format) are reported, not the startup record.
return msgspec_to_builtins( return msgspec_to_builtins(
{ {
**dataclasses.asdict(server_args), **get_context().resolved_server_args_dict(
base=dataclasses.asdict(server_args)
),
**_global_state.scheduler_info, **_global_state.scheduler_info,
"internal_states": internal_states, "internal_states": internal_states,
"version": __version__, "version": __version__,
@@ -1368,7 +1375,9 @@ async def update_weight_version(
# since weight_version update is a simple operation that doesn't affect model weights # since weight_version update is a simple operation that doesn't affect model weights
try: try:
# Update the weight version in server args (the single source of truth) # Update the weight version in server args (the single source of truth)
_global_state.tokenizer_manager.server_args.override( from sglang.srt.runtime_context import get_context
get_context().override(
"http.update_weight_version", weight_version=obj.new_version "http.update_weight_version", weight_version=obj.new_version
) )
@@ -55,6 +55,8 @@ class HttpServerEngineAdapter(EngineBase):
def __init__(self, **kwargs): def __init__(self, **kwargs):
self.server_args = ServerArgs(**kwargs) self.server_args = ServerArgs(**kwargs)
# Read host/port from the adapter's own args: no config is published yet
# in this process (publish happens in the child from launch_server_process).
print( print(
f"Launch HttpServerEngineAdapter at: {self.server_args.host}:{self.server_args.port}" f"Launch HttpServerEngineAdapter at: {self.server_args.host}:{self.server_args.port}"
) )
@@ -72,6 +72,7 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import (
TranscriptionAdapter, TranscriptionAdapter,
) )
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_serving
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import random_uuid from sglang.srt.utils import random_uuid
@@ -338,12 +339,12 @@ class RealtimeConnection:
if ( if (
transcription is not None transcription is not None
and transcription.model and transcription.model
and transcription.model != self.server_args.served_model_name and transcription.model != get_serving().served_model_name
): ):
await self._send_error( await self._send_error(
"not_supported", "not_supported",
f"Model {transcription.model!r} is not served by this endpoint " f"Model {transcription.model!r} is not served by this endpoint "
f"(serving {self.server_args.served_model_name!r}); set " f"(serving {get_serving().served_model_name!r}); set "
f"transcription.model to null or to the server's model name.", f"transcription.model to null or to the server's model name.",
param="session.audio.input.transcription.model", param="session.audio.input.transcription.model",
) )
+3 -3
View File
@@ -16,7 +16,7 @@ from sglang.srt.eplb.expert_location import (
get_global_expert_location_metadata, get_global_expert_location_metadata,
) )
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
from sglang.srt.runtime_context import get_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
@@ -274,8 +274,8 @@ def update_expert_location_with_recovery(
else: else:
# Load the missing weights from disk # Load the missing weights from disk
update_weights_from_disk_callable( update_weights_from_disk_callable(
get_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
@@ -34,7 +34,7 @@ class ExpertLocationDispatchInfo:
@classmethod @classmethod
def init_new(cls, layer_id: int): def init_new(cls, layer_id: int):
ep_dispatch_algorithm = get_server_args().ep_dispatch_algorithm ep_dispatch_algorithm = get_exec().moe.ep_dispatch_algorithm
expert_location_metadata = get_global_expert_location_metadata() expert_location_metadata = get_global_expert_location_metadata()
assert expert_location_metadata is not None assert expert_location_metadata is not None
@@ -26,7 +26,7 @@ from sglang.srt.eplb.expert_location import (
ExpertLocationMetadata, ExpertLocationMetadata,
get_global_expert_location_metadata, get_global_expert_location_metadata,
) )
from sglang.srt.runtime_context import get_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__)
@@ -107,7 +107,7 @@ def _update_expert_weights_with_canary(
canary_tensor = ( canary_tensor = (
_get_canary_value(old_expert_location_metadata, layer_id) _get_canary_value(old_expert_location_metadata, layer_id)
.clone() .clone()
.to(device=get_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)
@@ -16,9 +16,8 @@ from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import (
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.model_executor.model_runner_components.layer_setup import ( 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,7 +143,7 @@ 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
@@ -165,7 +164,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)
@@ -173,7 +172,7 @@ class MlxModelRunnerStub(ModelRunner):
requested_per_worker = requested // self.dp_size requested_per_worker = requested // self.dp_size
resolved = min(requested_per_worker, capacity_cap) resolved = min(requested_per_worker, capacity_cap)
aux_state_size = self.server_args.max_mamba_cache_size aux_state_size = get_schedule().max_mamba_cache_size
if ( if (
mambaish_config(self.model_config) is not None mambaish_config(self.model_config) is not None
and aux_state_size is not None and aux_state_size is not None
@@ -209,7 +208,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)
@@ -241,7 +240,7 @@ class MlxModelRunnerStub(ModelRunner):
# Create minimal pools # Create minimal pools
if mambaish_config(self.model_config) is not None: if mambaish_config(self.model_config) is not None:
auxiliary_state_size = self.server_args.max_mamba_cache_size auxiliary_state_size = get_schedule().max_mamba_cache_size
if auxiliary_state_size is None: if auxiliary_state_size is None:
auxiliary_state_size = ( auxiliary_state_size = (
self.max_running_requests * self._aux_state_slots_per_request() self.max_running_requests * self._aux_state_slots_per_request()
@@ -255,7 +254,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__)
@@ -47,25 +48,23 @@ class MlxTpModelWorker(TpModelWorker):
def _init_model_runner(self): def _init_model_runner(self):
"""Create MLX runner first (auto-sizes pool), then stub with matching size.""" """Create MLX runner first (auto-sizes pool), then stub with matching size."""
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
from sglang.srt.hardware_backend.mlx.model_runner_stub import ( from sglang.srt.hardware_backend.mlx.model_runner_stub import MlxModelRunnerStub
MlxModelRunnerStub,
)
logger.info("Initializing MlxModelRunner for end-to-end MLX inference") logger.info("Initializing MlxModelRunner for end-to-end MLX inference")
init_kwargs = dict( init_kwargs = dict(
model_path=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,
@@ -19,11 +19,9 @@ from sglang.srt.layers.attention.flashattention_backend import (
merge_state_v2_wrapper, merge_state_v2_wrapper,
) )
from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.layers.utils.cp_utils import ( 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 +513,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
@@ -12,7 +12,7 @@ from sglang.srt.layers.attention.dsv4.compressor import CompressorBackendMixin
from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin
from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, ForwardMode from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, ForwardMode
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.runtime_context import get_parallel 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
@@ -1362,9 +1362,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
) )
@@ -1409,9 +1408,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)
+3 -5
View File
@@ -22,9 +22,7 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import divide
divide,
)
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
@@ -33,7 +31,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
Phase, Phase,
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.runtime_context import get_parallel, 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,
@@ -89,7 +87,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
@@ -37,10 +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 (
from sglang.srt.state_capturer.indexer_topk import ( get_device,
maybe_capture_indexer_topk, get_exec,
get_parallel,
get_schedule,
get_server_args,
) )
from sglang.srt.state_capturer.indexer_topk import maybe_capture_indexer_topk
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
ceil_align, ceil_align,
@@ -105,9 +109,7 @@ if is_npu():
import torch_npu import torch_npu
from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream
from sglang.srt.distributed import ( from sglang.srt.distributed import get_attn_tp_group
get_attn_tp_group,
)
from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.distributed.parallel_state import get_pp_group
from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.communicator import ScatterMode from sglang.srt.layers.communicator import ScatterMode
@@ -458,7 +460,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
@@ -469,7 +471,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
@@ -1055,7 +1057,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:
@@ -28,16 +28,14 @@ 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
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.dsv4.compressor import ( from sglang.srt.layers.attention.dsv4.compressor import CompressorBackendMixin
CompressorBackendMixin,
)
from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -129,9 +127,7 @@ def _aiter_fp8_paged_mqa_logits(
clean_logits: bool = False, clean_logits: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
"""Wrapper adapting aiter's deepgemm_fp8_paged_mqa_logits to SGLang's interface.""" """Wrapper adapting aiter's deepgemm_fp8_paged_mqa_logits to SGLang's interface."""
from aiter.ops.triton.attention.pa_mqa_logits import ( from aiter.ops.triton.attention.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits
deepgemm_fp8_paged_mqa_logits,
)
batch_size = q_fp8.shape[0] batch_size = q_fp8.shape[0]
next_n = q_fp8.shape[1] next_n = q_fp8.shape[1]
@@ -838,9 +834,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(
@@ -13,9 +13,7 @@ from sglang.kernels.ops.attention.metadata import (
) )
from sglang.kernels.ops.attention.pa_page_table import _build_pa_page_table from sglang.kernels.ops.attention.pa_page_table import _build_pa_page_table
from sglang.kernels.ops.attention.utils import assert_buffer_fits from sglang.kernels.ops.attention.utils import assert_buffer_fits
from sglang.kernels.ops.kvcache.trtllm_mha_page_table import ( from sglang.kernels.ops.kvcache.trtllm_mha_page_table import build_trtllm_mha_page_table
build_trtllm_mha_page_table,
)
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy
@@ -28,7 +26,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
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
@@ -166,9 +164,12 @@ class FlashAttentionBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype = model_runner.kv_cache_dtype
from sglang.srt.runtime_context import get_model
self.kv_cache_dtype_str = get_model().kv_cache_dtype self.kv_cache_dtype_str = getattr(
model_runner,
"kv_cache_dtype_str",
model_runner.server_args.kv_cache_dtype,
)
self.kv_cache_is_mxfp8 = self.kv_cache_dtype_str == "mxfp8" self.kv_cache_is_mxfp8 = self.kv_cache_dtype_str == "mxfp8"
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
# Static page-table width (upper bound). The device-side page-table build # Static page-table width (upper bound). The device-side page-table build
@@ -1479,7 +1480,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.
@@ -32,7 +32,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph, is_in_tc_piecewise_cuda_graph,
) )
from sglang.srt.runtime_context import get_buffer, 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,
@@ -223,9 +223,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
self.enable_chunk_kv = ( self.enable_chunk_kv = (
not skip_prefill not skip_prefill
and get_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
@@ -401,7 +401,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
prefix_lens = forward_batch.extend_prefix_lens prefix_lens = forward_batch.extend_prefix_lens
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
use_ragged = ( use_ragged = (
not get_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()
@@ -19,7 +19,7 @@ from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_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
@@ -350,7 +350,7 @@ class MambaAttnBackendBase(AttentionBackend):
"""Per-row (length bs) bool flush mask = the radix track's seq_lens_cpu % """Per-row (length bs) bool flush mask = the radix track's seq_lens_cpu %
mamba_track_interval == 0, so force-flush and snapshot fire on the same mamba_track_interval == 0, so force-flush and snapshot fire on the same
steps (no off-by-one).""" steps (no off-by-one)."""
interval = get_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)
@@ -764,7 +764,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
# Page-major stores state strided; only the stride-aware Triton causal-conv # Page-major stores state strided; only the stride-aware Triton causal-conv
# reads it (CUDA causal_conv1d garbles it). A model may also force Triton. # reads it (CUDA causal_conv1d garbles it). A model may also force Triton.
use_triton_causal_conv = ( use_triton_causal_conv = (
use_triton_causal_conv or get_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(
@@ -38,7 +38,11 @@ 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,
)
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():
@@ -197,9 +201,7 @@ 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 = model_runner.server_args.speculative_num_draft_tokens
self.cuda_graph_custom_mask = None self.cuda_graph_custom_mask = None
+6 -9
View File
@@ -15,7 +15,7 @@ from einops import rearrange
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm as can_use_jit_qk_norm from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm as can_use_jit_qk_norm
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.models.utils import apply_qk_norm from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_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,
@@ -69,9 +69,7 @@ if _is_npu:
if _is_xpu: if _is_xpu:
from sgl_kernel.flash_attn import flash_attn_varlen_func from sgl_kernel.flash_attn import flash_attn_varlen_func
from sglang.kernels.ops.attention.prefill_attention import ( from sglang.kernels.ops.attention.prefill_attention import context_attention_fwd
context_attention_fwd,
)
from sglang.srt.distributed import ( from sglang.srt.distributed import (
split_tensor_along_last_dim, split_tensor_along_last_dim,
tensor_model_parallel_all_gather, tensor_model_parallel_all_gather,
@@ -86,7 +84,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
@@ -1045,7 +1042,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.")
@@ -1124,7 +1121,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(
@@ -1152,7 +1149,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:
@@ -1257,7 +1254,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), (
@@ -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
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -69,9 +69,12 @@ class XPUAttentionBackend(AttentionBackend):
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype = model_runner.kv_cache_dtype
from sglang.srt.runtime_context import get_model
self.kv_cache_dtype_str = get_model().kv_cache_dtype self.kv_cache_dtype_str = getattr(
model_runner,
"kv_cache_dtype_str",
model_runner.server_args.kv_cache_dtype,
)
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.skip_prefill = skip_prefill self.skip_prefill = skip_prefill
@@ -640,7 +643,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
@@ -72,7 +72,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,
@@ -170,7 +176,7 @@ def apply_flashinfer_allreduce_fusion(batch_size: int):
and batch_size > 0 and batch_size > 0
and batch_size <= FUSE_ALLREDUCE_MAX_BATCH_SIZE and batch_size <= FUSE_ALLREDUCE_MAX_BATCH_SIZE
and not is_dp_attention_enabled() and not is_dp_attention_enabled()
and get_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()
) )
@@ -186,7 +192,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
) )
@@ -274,7 +280,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:
@@ -407,7 +413,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
@@ -467,7 +473,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):
@@ -815,7 +821,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)
@@ -1120,7 +1126,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(
+3 -7
View File
@@ -48,12 +48,10 @@ from sglang.srt.layers.cp.base import (
CPAttentionBackendKind, CPAttentionBackendKind,
) )
from sglang.srt.layers.cp.padding import pad_local_rows from sglang.srt.layers.cp.padding import pad_local_rows
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import is_allocation_symmetric
is_allocation_symmetric,
)
from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_device, get_parallel
@dataclass @dataclass
@@ -208,10 +206,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]
+11 -12
View File
@@ -31,7 +31,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
Phase, Phase,
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.runtime_context import get_parallel, 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,9 +130,7 @@ if _is_cuda:
# BEFORE the weight multiply, so the multiply is done in the narrow dtype. # BEFORE the weight multiply, so the multiply is done in the narrow dtype.
_jit_rmsnorm_hf_available = False _jit_rmsnorm_hf_available = False
try: try:
from sglang.jit_kernel.rmsnorm_hf import ( from sglang.jit_kernel.rmsnorm_hf import is_supported_rmsnorm_hf_hidden_size
is_supported_rmsnorm_hf_hidden_size,
)
from sglang.jit_kernel.rmsnorm_hf import rmsnorm_hf as _jit_rmsnorm_hf from sglang.jit_kernel.rmsnorm_hf import rmsnorm_hf as _jit_rmsnorm_hf
_jit_rmsnorm_hf_available = True _jit_rmsnorm_hf_available = True
@@ -144,9 +142,7 @@ if _is_cuda:
_jit_rmsnorm_hf = None _jit_rmsnorm_hf = None
from sglang.jit_kernel.norm import fused_add_rmsnorm as _jit_fused_add_rmsnorm from sglang.jit_kernel.norm import fused_add_rmsnorm as _jit_fused_add_rmsnorm
from sglang.jit_kernel.norm import ( from sglang.jit_kernel.norm import is_supported_jit_fused_add_rmsnorm_hidden_size
is_supported_jit_fused_add_rmsnorm_hidden_size,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -206,7 +202,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)
@@ -284,7 +280,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(
@@ -391,7 +387,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)
@@ -452,7 +448,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(
@@ -579,7 +575,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,
+5 -11
View File
@@ -25,9 +25,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, use_symmetric_memory,
) )
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import is_allocation_symmetric
is_allocation_symmetric,
)
from sglang.srt.layers.moe.utils import should_skip_mlp_all_reduce from sglang.srt.layers.moe.utils import should_skip_mlp_all_reduce
from sglang.srt.layers.parameter import ( from sglang.srt.layers.parameter import (
BasevLLMParameter, BasevLLMParameter,
@@ -39,7 +37,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:
@@ -759,9 +757,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
shard_offsets.append((i, current_shard_offset, output_size)) shard_offsets.append((i, current_shard_offset, output_size))
current_shard_offset += output_size current_shard_offset += output_size
if _is_cpu: if _is_cpu:
from sglang.srt.model_loader.weight_utils import ( from sglang.srt.model_loader.weight_utils import pad_loaded_weight
pad_loaded_weight,
)
loaded_weight = pad_loaded_weight( loaded_weight = pad_loaded_weight(
loaded_weight, param.output_dim, output_sizes loaded_weight, param.output_dim, output_sizes
@@ -805,9 +801,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
current_block_offset += shard_block_size current_block_offset += shard_block_size
if _is_cpu: if _is_cpu:
from sglang.srt.model_loader.weight_utils import ( from sglang.srt.model_loader.weight_utils import pad_loaded_weight
pad_loaded_weight,
)
loaded_weight = pad_loaded_weight( loaded_weight = pad_loaded_weight(
loaded_weight, param.output_dim, shard_block_sizes loaded_weight, param.output_dim, shard_block_sizes
@@ -1596,7 +1590,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
@@ -47,7 +47,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,
@@ -346,7 +346,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 = (
@@ -370,8 +370,8 @@ class LogitsProcessor(nn.Module):
self.final_logit_softcapping = None self.final_logit_softcapping = None
self.return_full_logits = return_full_logits self.return_full_logits = return_full_logits
self.enable_mis = get_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(
+3 -5
View File
@@ -7,9 +7,7 @@ import torch
from torch import nn from torch import nn
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.eplb.expert_distribution import ( from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
get_global_expert_distribution_recorder,
)
from sglang.srt.eplb.expert_location_dispatch import ( from sglang.srt.eplb.expert_location_dispatch import (
ExpertLocationDispatchInfo, ExpertLocationDispatchInfo,
topk_ids_logical_to_physical, topk_ids_logical_to_physical,
@@ -22,6 +20,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 +43,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,
@@ -506,7 +506,7 @@ def _fused_moe_kernel_sequence(
out_hidden_states = torch.empty_like(hidden_states) out_hidden_states = torch.empty_like(hidden_states)
use_fused_moe_sum_all_reduce = ( use_fused_moe_sum_all_reduce = (
get_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,
@@ -21,13 +21,9 @@ from sglang.srt.layers.moe.token_dispatcher import (
from sglang.srt.layers.moe.token_dispatcher.flashinfer_utils import ( from sglang.srt.layers.moe.token_dispatcher.flashinfer_utils import (
TorchDistributedCommBackend, TorchDistributedCommBackend,
) )
from sglang.srt.layers.moe.topk import ( from sglang.srt.layers.moe.topk import StandardTopKOutput, TopKOutput, TopKOutputChecker
StandardTopKOutput,
TopKOutput,
TopKOutputChecker,
)
from sglang.srt.layers.moe.utils import get_moe_runner_backend from sglang.srt.layers.moe.utils import get_moe_runner_backend
from sglang.srt.runtime_context import get_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 +119,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 +128,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 = (
+5 -11
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
@@ -83,9 +83,7 @@ except ImportError:
pass pass
from sglang.jit_kernel.dsv4 import mask_topk_ids from sglang.jit_kernel.dsv4 import mask_topk_ids
from sglang.srt.distributed import ( from sglang.srt.distributed import get_tp_group
get_tp_group,
)
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, use_symmetric_memory,
) )
@@ -98,9 +96,7 @@ from sglang.srt.eplb.expert_location_dispatch import (
) )
from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe import get_moe_runner_backend from sglang.srt.layers.moe import get_moe_runner_backend
from sglang.srt.layers.moe.utils import ( from sglang.srt.layers.moe.utils import has_per_rank_fused_shared_slots
has_per_rank_fused_shared_slots,
)
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -419,10 +415,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
@@ -496,9 +491,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 = (
@@ -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,
@@ -1470,9 +1469,7 @@ def requant_block_scale_ue8m0_for_deepgemm(
scales are not already UE8M0, and DeepGEMM can run the layer (bf16 output, scales are not already UE8M0, and DeepGEMM can run the layer (bf16 output,
aligned shape). Returns True when it requantizes. aligned shape). Returns True when it requantizes.
""" """
from sglang.srt.model_loader.utils import ( from sglang.srt.model_loader.utils import should_deepgemm_weight_requant_ue8m0
should_deepgemm_weight_requant_ue8m0,
)
if ( if (
not use_deepgemm_runner not use_deepgemm_runner
@@ -1794,7 +1791,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,
@@ -77,9 +77,7 @@ if is_flashinfer_available():
nvfp4_block_scale_interleave, nvfp4_block_scale_interleave,
trtllm_fp4_block_scale_moe, trtllm_fp4_block_scale_moe,
) )
from flashinfer.fused_moe.core import ( from flashinfer.fused_moe.core import get_w2_permute_indices_with_cache
get_w2_permute_indices_with_cache,
)
# SM90 mixed-input helpers landed in FlashInfer #3084 (post-0.6.10). Older # SM90 mixed-input helpers landed in FlashInfer #3084 (post-0.6.10). Older
# versions don't ship them; gate at import so unrelated code paths still load. # versions don't ship them; gate at import so unrelated code paths still load.
@@ -334,7 +332,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 two FlashInfer # When `flashinfer_mxfp4` is enabled, dispatch to one of two FlashInfer
# entry points depending on the GPU: # entry points depending on the GPU:
@@ -14,7 +14,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
) )
from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe.utils import RoutingMethodType from sglang.srt.layers.moe.utils import RoutingMethodType
from sglang.srt.runtime_context import get_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):
@@ -376,9 +376,7 @@ def maybe_fuse_routed_scale_and_shared_add(
from sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe import ( from sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe import (
Mxfp4FlashinferCutlassMoEMethod, Mxfp4FlashinferCutlassMoEMethod,
) )
from sglang.srt.layers.quantization.mxfp4_marlin_moe import ( from sglang.srt.layers.quantization.mxfp4_marlin_moe import Mxfp4MarlinMoEMethod
Mxfp4MarlinMoEMethod,
)
fused = isinstance( fused = isinstance(
experts.quant_method, experts.quant_method,
@@ -11,7 +11,7 @@ from sglang.srt.environ import envs
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb
from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_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,
@@ -65,9 +65,7 @@ if _is_npu:
) )
if _is_hip: if _is_hip:
from sglang.kernels.ops.attention.utils import ( from sglang.kernels.ops.attention.utils import fused_qk_rope_reshape_and_cache
fused_qk_rope_reshape_and_cache,
)
if _is_xpu: if _is_xpu:
from sgl_kernel import fused_qk_rope_with_cos_sin_cache_inplace from sgl_kernel import fused_qk_rope_with_cos_sin_cache_inplace
@@ -127,7 +125,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,
@@ -151,7 +149,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
@@ -162,7 +160,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
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
is_cuda, is_cuda,
@@ -42,7 +42,6 @@ if _is_xpu:
from sgl_kernel import multimodal_rotary_embedding from sgl_kernel import multimodal_rotary_embedding
from sglang.kernels.ops.attention.mrope import apply_interleaved_rope_triton from sglang.kernels.ops.attention.mrope import apply_interleaved_rope_triton
from sglang.srt.runtime_context import get_server_args
def apply_interleaved_rope(x: torch.Tensor, mrope_section: list) -> torch.Tensor: def apply_interleaved_rope(x: torch.Tensor, mrope_section: list) -> torch.Tensor:
@@ -132,7 +131,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):
@@ -144,7 +143,7 @@ class MRotaryEmbedding(RotaryEmbedding):
last_dim = cos_sin.size()[-1] last_dim = cos_sin.size()[-1]
cos, sin = cos_sin.chunk(2, dim=-1) cos, sin = cos_sin.chunk(2, dim=-1)
if self.mrope_interleaved: if self.mrope_interleaved:
if support_triton(get_server_args().attention_backend): if support_triton(get_exec().kernel.attention_backend):
cos = apply_interleaved_rope_triton(cos, self.mrope_section) cos = apply_interleaved_rope_triton(cos, self.mrope_section)
sin = apply_interleaved_rope_triton(sin, self.mrope_section) sin = apply_interleaved_rope_triton(sin, self.mrope_section)
else: else:
+11 -22
View File
@@ -8,34 +8,21 @@ from torch import nn
from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32 from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32
from sglang.srt.distributed import get_tp_group from sglang.srt.distributed import get_tp_group
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import is_dp_attention_enabled
is_dp_attention_enabled,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.logprob_processor import ( from sglang.srt.layers.logprob_processor import OutputLogprobProcessor
OutputLogprobProcessor, from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
)
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.sampling.sampling_params import TOP_K_ALL
from sglang.srt.utils.async_probe import sanitize_nan_logits from sglang.srt.utils.async_probe import sanitize_nan_logits
from sglang.srt.utils.common import ( from sglang.srt.utils.common import get_bool_env_var, is_cuda, is_hip, is_musa, is_npu
get_bool_env_var,
is_cuda,
is_hip,
is_musa,
is_npu,
)
if is_cuda(): if is_cuda():
from flashinfer.sampling import ( from flashinfer.sampling import (
min_p_sampling_from_probs, min_p_sampling_from_probs,
top_k_top_p_sampling_from_probs, top_k_top_p_sampling_from_probs,
) )
from sgl_kernel import ( from sgl_kernel import top_k_renorm_prob, top_p_renorm_prob
top_k_renorm_prob,
top_p_renorm_prob,
)
if is_musa(): if is_musa():
from sgl_kernel import ( from sgl_kernel import (
@@ -74,12 +61,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()
@@ -245,7 +234,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
@@ -48,6 +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 get_exec
from sglang.srt.server_args import ( from sglang.srt.server_args import (
DP_ATTENTION_HANDSHAKE_PORT_DELTA, DP_ATTENTION_HANDSHAKE_PORT_DELTA,
PortArgs, PortArgs,
@@ -231,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; "
@@ -484,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(
+11 -5
View File
@@ -33,7 +33,13 @@ 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,
get_serving,
)
from sglang.srt.utils import flatten_nested_list, is_hip, is_npu, print_warning_once from sglang.srt.utils import flatten_nested_list, is_hip, is_npu, print_warning_once
from sglang.srt.utils.stale_shm_cleanup import make_shm_name from sglang.srt.utils.stale_shm_cleanup import make_shm_name
from sglang.utils import logger from sglang.utils import logger
@@ -878,7 +884,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"
@@ -1287,7 +1293,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
) )
@@ -1967,7 +1973,7 @@ def wrap_shm_features(obj):
""" """
Scan the object for multimodal tensors and wrap them in SHM pointers. Scan the object for multimodal tensors and wrap them in SHM pointers.
""" """
if _get_is_default_transport() or get_server_args().skip_tokenizer_init: if _get_is_default_transport() or get_serving().skip_tokenizer_init:
return obj return obj
if obj.mm_inputs: if obj.mm_inputs:
@@ -2028,7 +2034,7 @@ def unwrap_shm_features(obj):
Restore ShmPointerMMData wrappers back into standard torch.Tensors. Restore ShmPointerMMData wrappers back into standard torch.Tensors.
Handles both single requests and batch requests. Handles both single requests and batch requests.
""" """
if _get_is_default_transport() or get_server_args().skip_tokenizer_init: if _get_is_default_transport() or get_serving().skip_tokenizer_init:
return obj return obj
# Handle batch requests # Handle batch requests
if isinstance(obj, BaseBatchReq): if isinstance(obj, BaseBatchReq):
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_disagg
# Copyright 2023-2024 SGLang Team # Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License. # you may not use this file except in compliance with the License.
@@ -645,15 +647,15 @@ class TokenizerWorker(TokenizerManager):
self.tokenizer_ipc_name = port_args.tokenizer_ipc_name self.tokenizer_ipc_name = port_args.tokenizer_ipc_name
# For PD disaggregtion # For PD disaggregtion
self.server_args.override( from sglang.srt.runtime_context import get_context
get_context().override(
"tokenizer_worker.restore_disaggregation_mode", "tokenizer_worker.restore_disaggregation_mode",
disaggregation_mode=disaggregation_mode, disaggregation_mode=disaggregation_mode,
) )
self.disaggregation_mode = DisaggregationMode( self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
self.server_args.disaggregation_mode
)
self.disaggregation_transfer_backend = TransferBackend( self.disaggregation_transfer_backend = TransferBackend(
self.server_args.disaggregation_transfer_backend get_disagg().disaggregation_transfer_backend
) )
# Register this worker with the router for pause/continue broadcasting # Register this worker with the router for pause/continue broadcasting
+9 -7
View File
@@ -77,10 +77,7 @@ from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import ( from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import (
NewTokenRatioTracker, NewTokenRatioTracker,
) )
from sglang.srt.mem_cache.allocation import ( from sglang.srt.mem_cache.allocation import alloc_for_decode, alloc_for_extend
alloc_for_decode,
alloc_for_extend,
)
from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
@@ -105,7 +102,12 @@ from sglang.srt.observability.req_time_stats import (
DPControllerReqTimeStats, DPControllerReqTimeStats,
SchedulerReqTimeStats, SchedulerReqTimeStats,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import (
get_parallel,
get_server_args,
get_serving,
get_spec,
)
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -1094,7 +1096,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
@@ -1115,7 +1117,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
@@ -56,7 +56,7 @@ from sglang.srt.mem_cache.multi_ended_allocator import (
UnifiedMambaTokenToKVPoolAllocator, UnifiedMambaTokenToKVPoolAllocator,
) )
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_disagg
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -193,7 +193,7 @@ class SchedulePolicy:
if ( if (
not isinstance(policy, CacheAwarePolicy) not isinstance(policy, CacheAwarePolicy)
and self.tree_cache.supports_fast_match_prefix() and self.tree_cache.supports_fast_match_prefix()
and get_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)
+73 -67
View File
@@ -210,9 +210,7 @@ from sglang.srt.managers.scheduler_components.pool_stats_observer import (
from sglang.srt.managers.scheduler_components.profiler_manager import ( from sglang.srt.managers.scheduler_components.profiler_manager import (
SchedulerProfilerManager, SchedulerProfilerManager,
) )
from sglang.srt.managers.scheduler_components.recv_skipper import ( from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper
SchedulerRecvSkipper,
)
from sglang.srt.managers.scheduler_components.request_receiver import ( from sglang.srt.managers.scheduler_components.request_receiver import (
SchedulerRequestReceiver, SchedulerRequestReceiver,
) )
@@ -241,7 +239,20 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.plugins import load_plugins from sglang.srt.plugins import load_plugins
from sglang.srt.runtime_context import get_context, get_parallel from sglang.srt.runtime_context import (
get_context,
get_device,
get_disagg,
get_exec,
get_lora,
get_memory,
get_mm,
get_observability,
get_parallel,
get_schedule,
get_serving,
get_spec,
)
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.sampling.sampling_params import TOP_K_ALL
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs
@@ -443,9 +454,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
@@ -471,8 +482,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,
@@ -583,7 +594,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
) )
@@ -611,11 +622,11 @@ class Scheduler(
self.ipc_channels = SchedulerIpcChannels.create( self.ipc_channels = SchedulerIpcChannels.create(
port_args=port_args, port_args=port_args,
is_rank_zero=is_rank_zero, is_rank_zero=is_rank_zero,
skip_tokenizer_init=self.server_args.skip_tokenizer_init, skip_tokenizer_init=get_serving().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(),
) )
@@ -631,7 +642,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)
@@ -641,7 +652,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=[
@@ -712,9 +723,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,
) )
@@ -785,7 +796,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
@@ -793,10 +804,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)
@@ -887,7 +898,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(),
) )
@@ -933,14 +944,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.
@@ -987,7 +998,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
) )
@@ -1007,13 +1018,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:
@@ -1049,8 +1059,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)."
@@ -1067,15 +1077,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(
@@ -1091,12 +1101,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
@@ -1118,11 +1128,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?
@@ -1192,10 +1200,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,
) )
@@ -1221,7 +1229,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,
@@ -1235,11 +1243,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,
@@ -1320,7 +1327,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
@@ -1329,7 +1336,7 @@ class Scheduler(
"triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096), "triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096),
} }
env_var, default_size = backend_sizes.get( env_var, default_size = backend_sizes.get(
self.server_args.attention_backend, (None, None) get_exec().kernel.attention_backend, (None, None)
) )
self.truncation_align_size = ( self.truncation_align_size = (
get_int_env_var(env_var, default_size) if env_var else None get_int_env_var(env_var, default_size) if env_var else None
@@ -1725,10 +1732,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
@@ -1830,7 +1837,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,
@@ -2006,7 +2013,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)
@@ -2053,7 +2060,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()
): ):
@@ -2099,8 +2106,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:
@@ -2112,7 +2118,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,
@@ -2265,7 +2271,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 = (
@@ -2314,7 +2320,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)
@@ -2592,7 +2598,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)
@@ -2804,7 +2810,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:
@@ -2878,7 +2884,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:
@@ -2945,7 +2951,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),
@@ -3516,7 +3522,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
@@ -3792,7 +3798,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,
) )
@@ -3912,7 +3918,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()
@@ -4445,10 +4451,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)
@@ -2,14 +2,7 @@ from __future__ import annotations
import logging import logging
from dataclasses import dataclass from dataclasses import dataclass
from typing import ( from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union
TYPE_CHECKING,
Callable,
List,
Optional,
Tuple,
Union,
)
import torch import torch
@@ -23,11 +16,14 @@ from sglang.srt.managers.schedule_batch import (
ScheduleBatch, ScheduleBatch,
mamba_lazy_spec_in_window, mamba_lazy_spec_in_window,
) )
from sglang.srt.mem_cache.common import ( from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache
maybe_cache_unfinished_req, from sglang.srt.runtime_context import (
release_kv_cache, get_disagg,
get_exec,
get_memory,
get_observability,
get_server_args,
) )
from sglang.srt.runtime_context import get_server_args
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
@@ -48,10 +44,7 @@ if TYPE_CHECKING:
SchedulerOutputStreamer, SchedulerOutputStreamer,
) )
from sglang.srt.managers.tp_worker import BaseTpWorker from sglang.srt.managers.tp_worker import BaseTpWorker
from sglang.srt.managers.utils import ( from sglang.srt.managers.utils import EmbeddingBatchResult, GenerationBatchResult
EmbeddingBatchResult,
GenerationBatchResult,
)
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
@@ -84,7 +77,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 +85,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 +236,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 +749,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 +932,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 +952,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
@@ -1102,7 +1095,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:
@@ -12,9 +12,7 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import world_dp_gather_enabled from sglang.srt.layers.dp_attention import world_dp_gather_enabled
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler_components.recv_skipper import ( from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper
SchedulerRecvSkipper,
)
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
@@ -26,6 +24,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 +384,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
@@ -144,7 +145,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,
@@ -1,20 +1,15 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from typing import ( from typing import List, Tuple
List,
Tuple,
)
import torch import torch
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.server_args import ( from sglang.srt.runtime_context import get_exec
MIS_DELIMITER_TOKEN_ID, from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID, ServerArgs
ServerArgs,
)
@dataclass(kw_only=True, slots=True, frozen=True) @dataclass(kw_only=True, slots=True, frozen=True)
@@ -164,7 +159,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
) )
@@ -2,12 +2,7 @@ from __future__ import annotations
import logging import logging
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import ( from typing import Any, Callable, List, Optional
Any,
Callable,
List,
Optional,
)
import torch import torch
import zmq import zmq
@@ -21,11 +16,9 @@ from sglang.srt.managers.io_struct import (
CachedTokensDetails, CachedTokensDetails,
wrap_as_pickle, wrap_as_pickle,
) )
from sglang.srt.managers.schedule_batch import ( from sglang.srt.managers.schedule_batch import BaseFinishReason, Req
BaseFinishReason,
Req,
)
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.runtime_context import get_observability, get_serving
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -144,7 +137,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,
) )
@@ -171,7 +164,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()
@@ -5,13 +5,7 @@ import os
import time import time
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import ( from typing import TYPE_CHECKING, Any, Callable, List, Optional
TYPE_CHECKING,
Any,
Callable,
List,
Optional,
)
import torch import torch
@@ -19,7 +13,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
@@ -255,7 +249,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
@@ -365,7 +359,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()
@@ -2,14 +2,7 @@ from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import ( from typing import TYPE_CHECKING, Any, Callable, List, Optional, Union
TYPE_CHECKING,
Any,
Callable,
List,
Optional,
Union,
)
import zmq import zmq
from torch.distributed import barrier from torch.distributed import barrier
@@ -22,14 +15,9 @@ from sglang.srt.managers.io_struct import (
TokenizedGenerateReqInput, TokenizedGenerateReqInput,
sock_recv, sock_recv,
) )
from sglang.srt.managers.mm_utils import ( from sglang.srt.managers.mm_utils import has_shm_features, unwrap_shm_features
has_shm_features, from sglang.srt.runtime_context import get_disagg
unwrap_shm_features, from sglang.srt.utils import broadcast_pyobj, point_to_point_pyobj
)
from sglang.srt.utils import (
broadcast_pyobj,
point_to_point_pyobj,
)
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -220,8 +208,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:
@@ -74,6 +74,7 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromTensorReqOutput, UpdateWeightsFromTensorReqOutput,
) )
from sglang.srt.managers.load_snapshot import LoadSnapshot from sglang.srt.managers.load_snapshot import LoadSnapshot
from sglang.srt.runtime_context import get_lora
from sglang.srt.server_args import LoRARef, ServerArgs from sglang.srt.server_args import LoRARef, ServerArgs
from sglang.srt.utils import ( from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
@@ -569,7 +570,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop() self.auto_create_handle_loop()
try: try:
if not self.server_args.enable_lora: if not get_lora().enable_lora:
raise ValueError( raise ValueError(
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA." "LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
) )
@@ -602,10 +603,10 @@ class TokenizerControlMixin:
await self.lora_registry.register(new_adapter) await self.lora_registry.register(new_adapter)
self.lora_ref_cache[obj.lora_name] = new_adapter self.lora_ref_cache[obj.lora_name] = new_adapter
if self.server_args.max_loaded_loras is not None: if get_lora().max_loaded_loras is not None:
while ( while (
self.lora_registry.num_registered_loras self.lora_registry.num_registered_loras
> self.server_args.max_loaded_loras > get_lora().max_loaded_loras
): ):
lru_lora_name = await self.lora_registry.lru_lora_name( lru_lora_name = await self.lora_registry.lru_lora_name(
exclude_pinned=True exclude_pinned=True
@@ -619,7 +620,7 @@ class TokenizerControlMixin:
logger.info( logger.info(
f"Unloading least recently used LoRA adapter '{lru_lora_name}' " f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
f"(current number of adapters: {self.lora_registry.num_registered_loras}, " f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
f"max allowed: {self.server_args.max_loaded_loras})" f"max allowed: {get_lora().max_loaded_loras})"
) )
unload_result = await self._unload_lora_adapter_locked( unload_result = await self._unload_lora_adapter_locked(
@@ -647,7 +648,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop() self.auto_create_handle_loop()
try: try:
if not self.server_args.enable_lora: if not get_lora().enable_lora:
raise ValueError( raise ValueError(
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA." "LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
) )
@@ -672,10 +673,10 @@ class TokenizerControlMixin:
if result.success: if result.success:
await self.lora_registry.register(new_adapter) await self.lora_registry.register(new_adapter)
self.lora_ref_cache[obj.lora_name] = new_adapter self.lora_ref_cache[obj.lora_name] = new_adapter
if self.server_args.max_loaded_loras is not None: if get_lora().max_loaded_loras is not None:
while ( while (
self.lora_registry.num_registered_loras self.lora_registry.num_registered_loras
> self.server_args.max_loaded_loras > get_lora().max_loaded_loras
): ):
lru_lora_name = await self.lora_registry.lru_lora_name( lru_lora_name = await self.lora_registry.lru_lora_name(
exclude_pinned=True exclude_pinned=True
@@ -689,7 +690,7 @@ class TokenizerControlMixin:
logger.info( logger.info(
f"Unloading least recently used LoRA adapter '{lru_lora_name}' " f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
f"(current number of adapters: {self.lora_registry.num_registered_loras}, " f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
f"max allowed: {self.server_args.max_loaded_loras})" f"max allowed: {get_lora().max_loaded_loras})"
) )
unload_result = await self._unload_lora_adapter_locked( unload_result = await self._unload_lora_adapter_locked(
@@ -717,7 +718,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop() self.auto_create_handle_loop()
try: try:
if not self.server_args.enable_lora: if not get_lora().enable_lora:
raise ValueError( raise ValueError(
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA." "LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
) )
@@ -893,6 +894,8 @@ class TokenizerControlMixin:
) -> None: ) -> None:
"""Update weight version if provided.""" """Update weight version if provided."""
if weight_version is not None: if weight_version is not None:
self.server_args.override( from sglang.srt.runtime_context import get_context
get_context().override(
"tokenizer.weight_version", weight_version=weight_version "tokenizer.weight_version", weight_version=weight_version
) )
+41 -33
View File
@@ -110,6 +110,14 @@ from sglang.srt.observability.request_metrics_exporter import (
RequestMetricsExporterManager, RequestMetricsExporterManager,
) )
from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_lora,
get_model,
get_observability,
get_serving,
)
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import ( from sglang.srt.server_args import (
PortArgs, PortArgs,
@@ -463,10 +471,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# TODO: Refactor and organize the log export code. # TODO: Refactor and organize the log export code.
# Request logging # Request logging
self.request_logger = RequestLogger( self.request_logger = RequestLogger(
log_requests=self.server_args.log_requests, log_requests=get_observability().log_requests,
log_requests_level=self.server_args.log_requests_level, log_requests_level=get_observability().log_requests_level,
log_requests_format=self.server_args.log_requests_format, log_requests_format=get_observability().log_requests_format,
log_requests_target=self.server_args.log_requests_target, log_requests_target=get_observability().log_requests_target,
) )
# Dumping # Dumping
@@ -489,7 +497,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def init_weight_update(self): def init_weight_update(self):
# Initial weights status # Initial weights status
self.initial_weights_loaded = True self.initial_weights_loaded = True
if self.server_args.checkpoint_engine_wait_weights_before_ready: if get_model().checkpoint_engine_wait_weights_before_ready:
self.initial_weights_loaded = False self.initial_weights_loaded = False
# Weight updates # Weight updates
@@ -509,7 +517,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# The registry dynamically updates as adapters are loaded / unloaded during runtime. It # The registry dynamically updates as adapters are loaded / unloaded during runtime. It
# serves as the source of truth for available adapters and maps user-friendly LoRA names # serves as the source of truth for available adapters and maps user-friendly LoRA names
# to internally used unique LoRA IDs. # to internally used unique LoRA IDs.
self.lora_registry = LoRARegistry(self.server_args.lora_paths) self.lora_registry = LoRARegistry(get_lora().lora_paths)
# Lock to serialize LoRA update operations. # Lock to serialize LoRA update operations.
# Please note that, unlike `model_update_lock`, this does not block inference, allowing # Please note that, unlike `model_update_lock`, this does not block inference, allowing
# LoRA updates and inference to overlap. # LoRA updates and inference to overlap.
@@ -518,15 +526,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# point to their latest LoRARef objects, so that they can be # point to their latest LoRARef objects, so that they can be
# dynamically loaded if needed for inference # dynamically loaded if needed for inference
self.lora_ref_cache: Dict[str, LoRARef] = {} self.lora_ref_cache: Dict[str, LoRARef] = {}
if self.server_args.lora_paths is not None: if get_lora().lora_paths is not None:
for lora_ref in self.server_args.lora_paths: for lora_ref in get_lora().lora_paths:
self.lora_ref_cache[lora_ref.lora_name] = lora_ref self.lora_ref_cache[lora_ref.lora_name] = lora_ref
def init_disaggregation(self): def init_disaggregation(self):
# PD Disaggregation # PD Disaggregation
self.disaggregation_mode = DisaggregationMode( self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
self.server_args.disaggregation_mode
)
# Keep a reference so the bootstrap server is not garbage-collected. # Keep a reference so the bootstrap server is not garbage-collected.
self.bootstrap_server = start_disagg_service(self.server_args) self.bootstrap_server = start_disagg_service(self.server_args)
# Single-source counter for auto-assigning fake bootstrap_room. # Single-source counter for auto-assigning fake bootstrap_room.
@@ -535,18 +541,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Encoder Disaggregation # Encoder Disaggregation
self.encoder_bootstrap_server = None self.encoder_bootstrap_server = None
if self.server_args.language_only: if self.server_args.language_only:
from sglang.srt.disaggregation.encode_receiver import ( from sglang.srt.disaggregation.encode_receiver import EncoderBootstrapServer
EncoderBootstrapServer,
)
# Shared mutable URL list: the bootstrap server appends / removes # Shared mutable URL list: the bootstrap server appends / removes
# entries as encoders register, the receiver reads from the same # entries as encoders register, the receiver reads from the same
# list. Pre-populated with static --encoder-urls so the legacy # list. Pre-populated with static --encoder-urls so the legacy
# CLI flag still works (alongside dynamic registrations). # CLI flag still works (alongside dynamic registrations).
self.encoder_urls: List[str] = list(self.server_args.encoder_urls) self.encoder_urls: List[str] = list(get_disagg().encoder_urls)
self.encoder_bootstrap_server = EncoderBootstrapServer( self.encoder_bootstrap_server = EncoderBootstrapServer(
host=self.server_args.host, host=get_serving().host,
port=self.server_args.encoder_bootstrap_port, port=get_disagg().encoder_bootstrap_port,
urls=self.encoder_urls, urls=self.encoder_urls,
) )
self.mm_receiver = create_mm_receiver( self.mm_receiver = create_mm_receiver(
@@ -560,20 +564,22 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Metrics # Metrics
if self.enable_metrics: if self.enable_metrics:
engine_type = DisaggregationMode.to_engine_type( engine_type = DisaggregationMode.to_engine_type(
self.server_args.disaggregation_mode get_disagg().disaggregation_mode
) )
labels = { labels = {
"model_name": self.server_args.served_model_name, "model_name": get_serving().served_model_name,
"engine_type": engine_type, "engine_type": engine_type,
} }
if self.enable_priority_scheduling: if self.enable_priority_scheduling:
labels["priority"] = "" labels["priority"] = ""
if self.server_args.tokenizer_metrics_allowed_custom_labels: if get_observability().tokenizer_metrics_allowed_custom_labels:
for label in self.server_args.tokenizer_metrics_allowed_custom_labels: for (
label
) in get_observability().tokenizer_metrics_allowed_custom_labels:
labels[label] = "" labels[label] = ""
if self.server_args.extra_metric_labels: if get_observability().extra_metric_labels:
labels.update(self.server_args.extra_metric_labels) labels.update(get_observability().extra_metric_labels)
tokenizer_collector_cls = resolve_collector_class( tokenizer_collector_cls = resolve_collector_class(
self.server_args, self.server_args,
STAT_LOGGER_ROLE_TOKENIZER, STAT_LOGGER_ROLE_TOKENIZER,
@@ -582,18 +588,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.metrics_collector = tokenizer_collector_cls( self.metrics_collector = tokenizer_collector_cls(
server_args=self.server_args, server_args=self.server_args,
labels=labels, labels=labels,
bucket_time_to_first_token=self.server_args.bucket_time_to_first_token, bucket_time_to_first_token=get_observability().bucket_time_to_first_token,
bucket_e2e_request_latency=self.server_args.bucket_e2e_request_latency, bucket_e2e_request_latency=get_observability().bucket_e2e_request_latency,
bucket_inter_token_latency=self.server_args.bucket_inter_token_latency, bucket_inter_token_latency=get_observability().bucket_inter_token_latency,
) )
start_cpu_monitor_thread("tokenizer") start_cpu_monitor_thread("tokenizer")
if self.server_args.gc_warning_threshold_secs > 0.0: if get_observability().gc_warning_threshold_secs > 0.0:
configure_gc_warning(self.server_args.gc_warning_threshold_secs) configure_gc_warning(get_observability().gc_warning_threshold_secs)
self.soft_watchdog = Watchdog.create( self.soft_watchdog = Watchdog.create(
debug_name="TokenizerManager", debug_name="TokenizerManager",
watchdog_timeout=self.server_args.soft_watchdog_timeout, watchdog_timeout=get_device().soft_watchdog_timeout,
soft=True, soft=True,
test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(), test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(),
) )
@@ -1757,7 +1763,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# default the load format to the server_args # default the load format to the server_args
if obj.load_format is None: if obj.load_format is None:
obj.load_format = self.server_args.load_format obj.load_format = get_model().load_format
logger.info("Start update_weights. Load format=%s", obj.load_format) logger.info("Start update_weights. Load format=%s", obj.load_format)
if obj.abort_all_requests: if obj.abort_all_requests:
@@ -1783,7 +1789,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def _update_model_path_info(self, model_path: str, load_format: str): def _update_model_path_info(self, model_path: str, load_format: str):
self.served_model_name = model_path self.served_model_name = model_path
self.server_args.override( from sglang.srt.runtime_context import get_context
get_context().override(
"tokenizer.update_weights", model_path=model_path, load_format=load_format "tokenizer.update_weights", model_path=model_path, load_format=load_format
) )
self.model_path = model_path self.model_path = model_path
@@ -1927,7 +1935,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
"id": rid, "id": rid,
"finish_reason": recv_obj.finished_reasons[i], "finish_reason": recv_obj.finished_reasons[i],
"prompt_tokens": recv_obj.prompt_tokens[i], "prompt_tokens": recv_obj.prompt_tokens[i],
"weight_version": self.server_args.weight_version, "weight_version": get_serving().weight_version,
"num_retractions": recv_obj.retraction_counts[i], "num_retractions": recv_obj.retraction_counts[i],
} }
@@ -2801,7 +2809,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
meta_info = { meta_info = {
"id": recv_obj.rid, "id": recv_obj.rid,
"finish_reason": finish_reason, "finish_reason": finish_reason,
"weight_version": self.server_args.weight_version, "weight_version": get_serving().weight_version,
"e2e_latency": state.time_stats.get_e2e_latency(), "e2e_latency": state.time_stats.get_e2e_latency(),
} }
is_stream = getattr(state.obj, "stream", False) is_stream = getattr(state.obj, "stream", False)
@@ -597,7 +597,10 @@ class TokenizerManagerScoreMixin:
f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})" f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})"
) )
# Check if multi-item scoring is enabled # Check if multi-item scoring is enabled. enable_mis is a static startup
# feature flag (never overridden post-publish), and score_request is also
# exercised on a bare mixin without a published context, so read it off
# server_args rather than the resolved-config bag.
use_multi_item_scoring = self.server_args.enable_mis use_multi_item_scoring = self.server_args.enable_mis
input_ids = None input_ids = None
+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 (
@@ -405,14 +406,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,
@@ -423,7 +424,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,
@@ -439,11 +440,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,
@@ -459,7 +460,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
@@ -485,9 +486,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,
+3 -3
View File
@@ -26,7 +26,7 @@ from sglang.srt.mem_cache.common import (
evict_from_tree_cache, evict_from_tree_cache,
) )
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_exec, get_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
is_cpu, is_cpu,
is_cuda, is_cuda,
@@ -65,7 +65,7 @@ def write_cache_indices(
prefix_tensors: list[torch.Tensor], prefix_tensors: list[torch.Tensor],
req_to_token_pool: ReqToTokenPool, req_to_token_pool: ReqToTokenPool,
): ):
if support_triton(get_server_args().attention_backend): if support_triton(get_exec().kernel.attention_backend):
prefix_pointers = torch.tensor( prefix_pointers = torch.tensor(
[t.data_ptr() for t in prefix_tensors], [t.data_ptr() for t in prefix_tensors],
dtype=torch.uint64, dtype=torch.uint64,
@@ -106,7 +106,7 @@ def get_last_loc(
req_pool_indices_tensor: torch.Tensor, req_pool_indices_tensor: torch.Tensor,
prefix_lens_tensor: torch.Tensor, prefix_lens_tensor: torch.Tensor,
) -> torch.Tensor: ) -> torch.Tensor:
attn_backend = get_server_args().attention_backend attn_backend = get_exec().kernel.attention_backend
uses_triton_dispatch = attn_backend not in ("ascend", "torch_native") uses_triton_dispatch = attn_backend not in ("ascend", "torch_native")
if _is_hip and uses_triton_dispatch: if _is_hip and uses_triton_dispatch:
+2 -2
View File
@@ -16,7 +16,7 @@ from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_server_args, get_serving
from sglang.srt.utils.common import ceil_align from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -183,7 +183,7 @@ def _release_overallocated_kv_indices(
# strip_thinking_cache intentionally reports output tokens as overallocated # strip_thinking_cache intentionally reports output tokens as overallocated
# so they fall into the free path below (#22373). # so they fall into the free path below (#22373).
if spec_algo is None and not 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
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()
@@ -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_disagg,
get_exec,
get_memory,
get_model,
get_parallel,
get_schedule,
get_spec,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
@@ -115,9 +123,7 @@ if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import ( from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import (
SpecAuxHiddenStateConfig, SpecAuxHiddenStateConfig,
) )
from sglang.srt.model_executor.pool_configurator import ( from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
MemoryPoolConfig,
)
class KVCacheConfigResult(msgspec.Struct, frozen=True, kw_only=True): class KVCacheConfigResult(msgspec.Struct, frozen=True, kw_only=True):
@@ -308,8 +314,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:
@@ -358,13 +364,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
@@ -394,7 +400,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)
): ):
@@ -432,8 +438,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
@@ -462,14 +468,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.
@@ -502,13 +508,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)
@@ -558,8 +564,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,
@@ -579,7 +585,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
@@ -588,7 +594,7 @@ class KVCacheConfigurator:
elif current_platform.is_out_of_tree() and not self.mambaish_config: elif current_platform.is_out_of_tree() and not self.mambaish_config:
unsupported_pool_family = "out-of-tree platform KV pool" unsupported_pool_family = "out-of-tree platform KV pool"
elif ( elif (
self.server_args.attention_backend == "ascend" and not self.mambaish_config get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config
): ):
unsupported_pool_family = "NPU/Ascend KV pool" unsupported_pool_family = "NPU/Ascend KV pool"
elif self.use_mla_backend and is_dsa_model: elif self.use_mla_backend and is_dsa_model:
@@ -614,9 +620,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,
@@ -648,15 +654,13 @@ class KVCacheConfigurator:
extra_max_context_len: int, extra_max_context_len: int,
pre_alloc_size: int, pre_alloc_size: int,
) -> ReqToTokenPool: ) -> ReqToTokenPool:
from sglang.srt.disaggregation.decode import ( from sglang.srt.disaggregation.decode import HybridMambaDecodeReqToTokenPool
HybridMambaDecodeReqToTokenPool,
)
req_to_token_pool = HybridMambaDecodeReqToTokenPool( req_to_token_pool = HybridMambaDecodeReqToTokenPool(
size=max_num_reqs, size=max_num_reqs,
max_context_len=self.model_config.context_len + extra_max_context_len, max_context_len=self.model_config.context_len + extra_max_context_len,
device=self.device, device=self.device,
enable_memory_saver=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=(
[ [
@@ -666,11 +670,11 @@ class KVCacheConfigurator:
] ]
), ),
speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens, speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens,
speculative_eagle_topk=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,
) )
return req_to_token_pool return req_to_token_pool
@@ -688,7 +692,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
@@ -701,11 +705,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=(
[ [
@@ -717,18 +721,18 @@ class KVCacheConfigurator:
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(), enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(),
speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens, speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens,
speculative_eagle_topk=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,
# ReplaySSM spec-verify is GDN-only: activate the pool machinery # ReplaySSM spec-verify is GDN-only: activate the pool machinery
# (rings + cursors + the intermediate_ssm gate) only for GDN-hybrid # (rings + cursors + the intermediate_ssm gate) only for GDN-hybrid
# models, so any other mamba-ish model (Mamba2/Nemotron, lightning, # models, so any other mamba-ish model (Mamba2/Nemotron, lightning,
# ...) run with the flag set stays byte-identical to flag-off. # ...) run with the flag set stays byte-identical to flag-off.
enable_gdn_replayssm_spec=( enable_gdn_replayssm_spec=(
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
), ),
) )
@@ -754,7 +758,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
@@ -770,7 +774,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
) )
@@ -802,7 +806,7 @@ class KVCacheConfigurator:
max_total_num_tokens=sizes.max_total_num_tokens, max_total_num_tokens=sizes.max_total_num_tokens,
) )
elif ( elif (
self.server_args.attention_backend == "ascend" and not self.mambaish_config get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config
): ):
if self.is_hybrid_swa: if self.is_hybrid_swa:
token_to_kv_pool = self._build_ascend_swa_kv_pool( token_to_kv_pool = self._build_ascend_swa_kv_pool(
@@ -878,14 +882,12 @@ 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."
if self.is_draft_worker: if self.is_draft_worker:
from sglang.srt.models.deepseek_v4_nextn import ( from sglang.srt.models.deepseek_v4_nextn import COMPRESS_RATIO_NEXTN_LAYER
COMPRESS_RATIO_NEXTN_LAYER,
)
compression_ratios = [ compression_ratios = [
COMPRESS_RATIO_NEXTN_LAYER COMPRESS_RATIO_NEXTN_LAYER
@@ -912,12 +914,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:
@@ -935,7 +937,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,
@@ -946,11 +948,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
), ),
@@ -961,7 +963,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,
@@ -972,7 +974,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),
@@ -985,14 +987,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,
) )
@@ -1002,13 +1004,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,
) )
@@ -1020,9 +1022,7 @@ class KVCacheConfigurator:
full_max_total_num_tokens: Optional[int], full_max_total_num_tokens: Optional[int],
swa_max_total_num_tokens: Optional[int], swa_max_total_num_tokens: Optional[int],
) -> KVCache: ) -> KVCache:
from sglang.srt.hardware_backend.npu.memory_pool_npu import ( from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool
NPUMHATokenToKVPool,
)
kwargs = {} kwargs = {}
if self.is_hybrid_swa_compress: if self.is_hybrid_swa_compress:
@@ -1039,7 +1039,7 @@ class KVCacheConfigurator:
token_to_kv_pool = SWAKVPool( token_to_kv_pool = SWAKVPool(
size=full_max_total_num_tokens, size=full_max_total_num_tokens,
size_swa=swa_max_total_num_tokens, size_swa=swa_max_total_num_tokens,
page_size=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),
@@ -1055,39 +1055,35 @@ class KVCacheConfigurator:
def _build_ascend_mla_kv_pool( def _build_ascend_mla_kv_pool(
self, *, max_total_num_tokens: int, is_dsa_model: bool self, *, max_total_num_tokens: int, is_dsa_model: bool
) -> KVCache: ) -> KVCache:
from sglang.srt.hardware_backend.npu.memory_pool_npu import ( from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMLATokenToKVPool
NPUMLATokenToKVPool,
)
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,
) )
return token_to_kv_pool return token_to_kv_pool
def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
from sglang.srt.hardware_backend.npu.memory_pool_npu import ( from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool
NPUMHATokenToKVPool,
)
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,
) )
@@ -1101,7 +1097,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
@@ -1121,7 +1117,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,
@@ -1132,7 +1128,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),
@@ -1143,13 +1139,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,
) )
@@ -1158,13 +1154,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,
) )
@@ -1221,7 +1217,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),
@@ -1229,7 +1225,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,
) )
@@ -1244,7 +1240,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),
@@ -1254,7 +1250,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,
) )
@@ -1293,7 +1289,7 @@ class KVCacheConfigurator:
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),
@@ -1302,8 +1298,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,
@@ -1316,18 +1312,18 @@ class KVCacheConfigurator:
def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
token_to_kv_pool = MHATokenToKVPoolFP4( token_to_kv_pool = MHATokenToKVPoolFP4(
max_total_num_tokens, max_total_num_tokens,
page_size=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
@@ -1339,7 +1335,7 @@ class KVCacheConfigurator:
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 = {}
@@ -1349,18 +1345,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
@@ -1375,20 +1371,20 @@ class KVCacheConfigurator:
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator], token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator],
) -> BaseTokenToKVPoolAllocator: ) -> BaseTokenToKVPoolAllocator:
# Initialize token_to_kv_pool_allocator # Initialize token_to_kv_pool_allocator
need_sort = 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,
need_sort=need_sort, need_sort=need_sort,
) )
elif _is_npu and ( elif _is_npu and (
self.server_args.attention_backend == "ascend" get_exec().kernel.attention_backend == "ascend"
or is_dsv4_model or is_dsv4_model
or self.hybrid_gdn_config is not None or self.hybrid_gdn_config is not None
): ):
@@ -1406,7 +1402,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,
@@ -1419,7 +1415,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,
@@ -1429,7 +1425,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,
@@ -1439,22 +1435,20 @@ 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,
)
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,
@@ -1462,8 +1456,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,
@@ -1475,7 +1468,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,
@@ -1483,7 +1476,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
@@ -1535,7 +1528,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(
@@ -1559,7 +1552,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 = "
@@ -1570,14 +1563,14 @@ class KVCacheConfigurator:
return int(rest_memory * (1 << 30)) # return in bytes return int(rest_memory * (1 << 30)) # return in bytes
def _calculate_mamba_ratio(self) -> int: def _calculate_mamba_ratio(self) -> int:
if self.server_args.disable_radix_cache: if get_memory().disable_radix_cache:
return 1 return 1
additional_ratio = 0 additional_ratio = 0
if self.server_args.enable_mamba_extra_buffer(): if self.server_args.enable_mamba_extra_buffer():
# ping-pong buffer size is 2 when overlap schedule is on, 1 otherwise. # ping-pong buffer size is 2 when overlap schedule is on, 1 otherwise.
# Lazy mode saves 1 slot (2 → 1) for overlap; non-overlap already uses 1. # Lazy mode saves 1 slot (2 → 1) for overlap; non-overlap already uses 1.
if not 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:
@@ -1596,7 +1589,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:
@@ -1626,7 +1619,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)
@@ -1637,13 +1630,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 "
@@ -1673,7 +1666,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(
@@ -1689,18 +1682,20 @@ 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
def _handle_max_mamba_cache(self, total_rest_memory): def _handle_max_mamba_cache(self, total_rest_memory):
from sglang.srt.runtime_context import get_context
config = self.mambaish_config config = self.mambaish_config
server_args = self.server_args server_args = self.server_args
assert config is not None assert config is not None
@@ -1710,11 +1705,11 @@ class KVCacheConfigurator:
assert server_args.speculative_num_draft_tokens is not None assert server_args.speculative_num_draft_tokens is not None
assert server_args.max_running_requests is not None assert server_args.max_running_requests is not None
if 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 # Reserve intermediate memory based on capped max_num_reqs
@@ -1722,7 +1717,7 @@ class KVCacheConfigurator:
ratio = self._calculate_mamba_ratio() ratio = self._calculate_mamba_ratio()
capped_reqs = min( capped_reqs = min(
server_args.max_running_requests // self.ps.attn_dp_size, server_args.max_running_requests // self.ps.attn_dp_size,
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
@@ -1735,7 +1730,7 @@ class KVCacheConfigurator:
and server_args.max_running_requests is not None and server_args.max_running_requests is not None
): ):
# Use explicitly set max_running_requests when radix cache is disabled # Use explicitly set max_running_requests when radix cache is disabled
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=server_args.max_running_requests
// self.ps.attn_dp_size, // self.ps.attn_dp_size,
@@ -1744,7 +1739,7 @@ class KVCacheConfigurator:
if has_spec_dec: if has_spec_dec:
intermediate_size = ( intermediate_size = (
config.mamba2_cache_params.mamba_cache_per_req config.mamba2_cache_params.mamba_cache_per_req
* server_args.max_mamba_cache_size * get_schedule().max_mamba_cache_size
* server_args.speculative_num_draft_tokens * server_args.speculative_num_draft_tokens
) )
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
@@ -1769,7 +1764,7 @@ class KVCacheConfigurator:
ratio = self._calculate_mamba_ratio() ratio = self._calculate_mamba_ratio()
D = server_args.speculative_num_draft_tokens D = server_args.speculative_num_draft_tokens
# Joint solve: main_state + intermediate = mamba_budget # Joint solve: main_state + intermediate = mamba_budget
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 / ratio)) mamba_budget_bytes // (per_req * (1 + D / ratio))
@@ -1779,12 +1774,12 @@ class KVCacheConfigurator:
# so the return value only has main_state subtracted from total # so the return value only has main_state subtracted from total
capped_reqs = min( capped_reqs = min(
server_args.max_running_requests // self.ps.attn_dp_size, server_args.max_running_requests // self.ps.attn_dp_size,
server_args.max_mamba_cache_size // ratio, get_schedule().max_mamba_cache_size // ratio,
) )
intermediate_size = per_req * capped_reqs * D intermediate_size = per_req * capped_reqs * D
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
else: else:
server_args.override( get_context().override(
"mamba_pool.memory_budget", "mamba_pool.memory_budget",
max_mamba_cache_size=int(mamba_budget_bytes // per_req), max_mamba_cache_size=int(mamba_budget_bytes // per_req),
) )
@@ -1793,10 +1788,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, "
@@ -1806,7 +1801,7 @@ class KVCacheConfigurator:
) )
mamba_state_memory = ( mamba_state_memory = (
server_args.max_mamba_cache_size get_schedule().max_mamba_cache_size
* config.mamba2_cache_params.mamba_cache_per_req * config.mamba2_cache_params.mamba_cache_per_req
/ (1 << 30) / (1 << 30)
) )
@@ -16,7 +16,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchResult, MatchResult,
) )
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_memory, get_server_args
try: try:
from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector
@@ -108,7 +108,7 @@ class LMCRadixCache(RadixCache):
): ):
super().__init__(params) super().__init__(params)
cli_lmc_cfg = get_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(
@@ -51,13 +51,8 @@ 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_hip, is_npu, support_triton
is_cuda,
is_hip,
is_npu,
support_triton,
)
from sglang.srt.utils.common import ceil_align, is_pin_memory_available from sglang.srt.utils.common import ceil_align, is_pin_memory_available
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -941,7 +936,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(
@@ -1110,7 +1105,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():
@@ -26,11 +26,7 @@ import torch
import torch.distributed as dist import torch.distributed as dist
from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.load_config import LoadConfig
from sglang.srt.configs.model_config import ( from sglang.srt.configs.model_config import AttentionArch, ModelConfig, ModelImpl
AttentionArch,
ModelConfig,
ModelImpl,
)
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
from sglang.srt.debug_utils.dumper import dumper from sglang.srt.debug_utils.dumper import dumper
from sglang.srt.distributed import bootstrap from sglang.srt.distributed import bootstrap
@@ -74,9 +70,7 @@ from sglang.srt.kv_canary.runner.canary_manager import context_tuple
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
from sglang.srt.layers import deep_gemm_wrapper, model_parallel from sglang.srt.layers import deep_gemm_wrapper, model_parallel
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.cp.utils import ( from sglang.srt.layers.cp.utils import get_cp_strategy
get_cp_strategy,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.sampler import create_sampler from sglang.srt.layers.sampler import create_sampler
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
@@ -86,17 +80,10 @@ from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
from sglang.srt.mem_cache import kv_cache_dtype from sglang.srt.mem_cache import kv_cache_dtype
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.kv_cache_configurator import ( from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator
KVCacheConfigurator,
)
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.model_executor.cuda_graph_config import ( 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, PPProxyTensors
)
from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch,
PPProxyTensors,
)
from sglang.srt.model_executor.forward_context import ( from sglang.srt.model_executor.forward_context import (
ForwardContext, ForwardContext,
forward_context, forward_context,
@@ -155,14 +142,15 @@ from sglang.srt.model_executor.model_runner_components.weight_updater import (
WeightUpdater, WeightUpdater,
) )
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.model_executor.runner import ( from sglang.srt.model_executor.runner import EagerRunner, get_batch_sizes_to_capture
EagerRunner,
get_batch_sizes_to_capture,
)
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_device,
get_exec,
get_global_dwdp_manager, get_global_dwdp_manager,
get_server_args, get_lora,
get_model,
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
@@ -319,7 +307,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)
@@ -396,12 +384,12 @@ 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_joiner and self.server_args.is_ep_joiner
): ):
return return
is_scale_join = self.server_args.ep_join_mode == "scale" is_scale_join = get_exec().moe.ep_join_mode == "scale"
if is_scale_join: if is_scale_join:
join_effective_ep_size = ( join_effective_ep_size = (
self.server_args.ep_join_rank_offset + self.ps.tp_size self.server_args.ep_join_rank_offset + self.ps.tp_size
@@ -484,7 +472,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,
@@ -561,7 +549,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(
@@ -618,7 +606,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):
@@ -654,7 +642,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):
@@ -668,12 +656,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):
@@ -692,8 +680,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
) )
@@ -702,17 +690,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()
@@ -973,7 +961,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")
@@ -1030,7 +1018,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_scale_joiner=self.server_args.is_ep_scale_joiner, is_ep_scale_joiner=self.server_args.is_ep_scale_joiner,
) )
@@ -1050,16 +1038,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(
@@ -1331,7 +1319,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
@@ -1361,7 +1349,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
@@ -1765,7 +1753,7 @@ class ModelRunner:
recovered = maybe_recover_ep_ranks( recovered = maybe_recover_ep_ranks(
tp_group=self.tp_group, tp_group=self.tp_group,
eplb_manager=self.eplb_manager, eplb_manager=self.eplb_manager,
random_seed=self.server_args.random_seed, random_seed=get_device().random_seed,
) )
if recovered: if recovered:
self.forward_pass_id = 0 self.forward_pass_id = 0
@@ -1774,7 +1762,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)
@@ -1842,7 +1830,9 @@ class ModelRunner:
load_config: LoadConfig, load_config: LoadConfig,
) -> None: ) -> None:
self.model = new_model self.model = new_model
self.server_args.override( from sglang.srt.runtime_context import get_context
get_context().override(
"model_runner.update_model_fields", "model_runner.update_model_fields",
model_path=model_path, model_path=model_path,
load_format=load_format, load_format=load_format,
@@ -24,6 +24,12 @@ def maybe_disable_chunked_prefix_cache(
# model's (often non-MLA) config must not flip the shared setting. # model's (often non-MLA) config must not flip the shared setting.
if is_draft_worker: if is_draft_worker:
return return
# This is a load-time gate that runs in ModelRunner.__init__ BEFORE the
# runner publishes its config (and direct/benchmark construction never
# publishes earlier), so read/write the supplied server_args. The runner's
# subsequent publish snapshots this into the schedule bag for get_schedule()
# readers.
if ( if (
not use_mla_backend not use_mla_backend
or server_args.attention_backend or server_args.attention_backend
@@ -11,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"
+3 -6
View File
@@ -44,7 +44,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
get_remote_instance_transfer_engine_info_per_rank, get_remote_instance_transfer_engine_info_per_rank,
register_memory_region, register_memory_region,
) )
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_exec, get_server_args
from sglang.srt.utils import get_available_gpu_memory from sglang.srt.utils import get_available_gpu_memory
# Try to import accelerate (optional dependency) # Try to import accelerate (optional dependency)
@@ -71,9 +71,7 @@ from sglang.srt.connector import (
get_connector_type, get_connector_type,
) )
from sglang.srt.connector.utils import parse_model_name from sglang.srt.connector.utils import parse_model_name
from sglang.srt.distributed import ( from sglang.srt.distributed import model_parallel_is_initialized
model_parallel_is_initialized,
)
from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
@@ -865,9 +863,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)
+4 -7
View File
@@ -41,9 +41,7 @@ from sglang.srt.layers.communicator import (
LayerScatterModes, LayerScatterModes,
enable_moe_dense_fully_dp, enable_moe_dense_fully_dp,
) )
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import is_dp_attention_enabled
is_dp_attention_enabled,
)
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
MergedColumnParallelLinear, MergedColumnParallelLinear,
@@ -78,6 +76,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 +208,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 +222,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,
@@ -12,17 +12,12 @@ from transformers import PretrainedConfig
from sglang.kernels.ops.attention.fla.layernorm_gated import RMSNorm as RMSNormGated from sglang.kernels.ops.attention.fla.layernorm_gated import RMSNorm as RMSNormGated
from sglang.kernels.ops.attention.fla.layernorm_gated import layernorm_fn from sglang.kernels.ops.attention.fla.layernorm_gated import layernorm_fn
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
from sglang.srt.distributed import ( from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce
get_pp_group,
tensor_model_parallel_all_reduce,
)
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import is_dp_attention_enabled
is_dp_attention_enabled,
)
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
@@ -59,6 +54,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 +525,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 +686,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_server_args from sglang.srt.runtime_context import get_exec
from sglang.srt.utils import use_intel_amx_backend from sglang.srt.utils import use_intel_amx_backend
MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"] MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"]
@@ -118,7 +118,7 @@ def handle_attention_flashinfer(attn, forward_batch):
def handle_attention_fa3(attn, forward_batch): def handle_attention_fa3(attn, forward_batch):
# when deterministic inference is enabled, use MLA # when deterministic inference is enabled, use MLA
if get_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")
@@ -187,7 +187,7 @@ def handle_attention_triton(attn, forward_batch):
return AttnForwardMethod.MLA return AttnForwardMethod.MLA
# when deterministic inference is enabled, use MLA # when deterministic inference is enabled, use MLA
if get_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 (
@@ -30,7 +30,11 @@ from sglang.srt.models.deepseek_common.utils import (
_use_aiter_bpreshuffle_gfx95, _use_aiter_bpreshuffle_gfx95,
_use_aiter_gfx95, _use_aiter_gfx95,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_schedule,
)
from sglang.srt.utils import BumpAllocator, get_bool_env_var, next_power_of_2 from sglang.srt.utils import BumpAllocator, get_bool_env_var, next_power_of_2
_use_fp8_prefill_attn = ( _use_fp8_prefill_attn = (
@@ -142,9 +146,7 @@ def _forward_dsa_indexer_for_mha(
class DeepseekMHAForwardMixin: class DeepseekMHAForwardMixin:
def init_mha_forward(self: DeepseekV2AttentionMLA): def init_mha_forward(self: DeepseekV2AttentionMLA):
self.disable_chunked_prefix_cache = ( self.disable_chunked_prefix_cache = get_schedule().disable_chunked_prefix_cache
get_server_args().disable_chunked_prefix_cache
)
# TODO: Design a finer way to determine the threshold # TODO: Design a finer way to determine the threshold
self.chunked_prefix_cache_threshold = ( self.chunked_prefix_cache_threshold = (
@@ -305,8 +307,8 @@ class DeepseekMHAForwardMixin:
self.use_dsa self.use_dsa
and self.kv_cache_dtype == "fp8_e4m3" and self.kv_cache_dtype == "fp8_e4m3"
and ( and (
not get_server_args().dsa_decode_backend == "trtllm" not get_exec().kernel.dsa_decode_backend == "trtllm"
or not get_server_args().dsa_prefill_backend == "trtllm" or not get_exec().kernel.dsa_prefill_backend == "trtllm"
) )
): ):
# FP8 path: dequantize DSA-specific FP8 format to BF16 # FP8 path: dequantize DSA-specific FP8 format to BF16
@@ -65,10 +65,8 @@ from sglang.srt.models.deepseek_common.utils import (
_use_aiter_bpreshuffle_gfx95, _use_aiter_bpreshuffle_gfx95,
_use_aiter_gfx95, _use_aiter_gfx95,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
from sglang.srt.state_capturer.indexer_topk import ( from sglang.srt.state_capturer.indexer_topk import maybe_capture_indexer_topk
maybe_capture_indexer_topk,
)
from sglang.srt.utils import BumpAllocator from sglang.srt.utils import BumpAllocator
from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.custom_op import register_custom_op
@@ -153,7 +151,7 @@ def _should_defer_dsa_cp_kv_gather(
class DeepseekMLAForwardMixin: class DeepseekMLAForwardMixin:
def init_mla_forward(self: DeepseekV2AttentionMLA): def init_mla_forward(self: DeepseekV2AttentionMLA):
self.flashinfer_mla_disable_ragged = ( self.flashinfer_mla_disable_ragged = (
get_server_args().flashinfer_mla_disable_ragged get_exec().kernel.flashinfer_mla_disable_ragged
) )
def should_run_indexer( def should_run_indexer(
@@ -990,8 +988,8 @@ class DeepseekMLAForwardMixin:
""" """
if self.current_attention_backend in ("dsa", "nsa"): if self.current_attention_backend in ("dsa", "nsa"):
return ( return (
get_server_args().dsa_decode_backend == "trtllm" get_exec().kernel.dsa_decode_backend == "trtllm"
or get_server_args().dsa_prefill_backend == "trtllm" or get_exec().kernel.dsa_prefill_backend == "trtllm"
) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn ) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn
return ( return (
+9 -5
View File
@@ -59,7 +59,12 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8 from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8
from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
from sglang.srt.models.utils import WeightsMapper from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import (
get_model,
get_parallel,
get_server_args,
get_spec,
)
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
@@ -148,7 +153,7 @@ class DeepseekModelNextN(nn.Module):
self.rot_weight = None self.rot_weight = None
if _is_npu: if _is_npu:
rot_weight_path = get_server_args().model_path + "/rot.safetensors" rot_weight_path = get_model().model_path + "/rot.safetensors"
if os.path.isfile(rot_weight_path): if os.path.isfile(rot_weight_path):
self.rot_weight = load_file(rot_weight_path) self.rot_weight = load_file(rot_weight_path)
self.rot_weight = self.rot_weight["rot.weight"].npu() self.rot_weight = self.rot_weight["rot.weight"].npu()
@@ -161,8 +166,7 @@ class DeepseekModelNextN(nn.Module):
layer_name = "decoder" layer_name = "decoder"
if _is_npu and ( if _is_npu and (
get_server_args().speculative_draft_model_path get_spec().speculative_draft_model_path == get_model().model_path
== get_server_args().model_path
): ):
layer_name = "layers." + str(config.num_hidden_layers) layer_name = "layers." + str(config.num_hidden_layers)
@@ -201,7 +205,7 @@ class DeepseekModelNextN(nn.Module):
if ( if (
_is_npu _is_npu
and self.quant_config is None and self.quant_config is None
and get_server_args().quantization is not None and get_model().quantization is not None
): ):
# ascend mtp unquant # ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
+14 -21
View File
@@ -29,10 +29,7 @@ import torch.nn.functional as F
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.jit_kernel.dsv4 import ( from sglang.jit_kernel.dsv4 import silu_and_mul_clamp, silu_and_mul_contig_post_quant
silu_and_mul_clamp,
silu_and_mul_contig_post_quant,
)
from sglang.kernels.ops.quantization.fp8_kernel import ( from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale, create_per_token_group_quant_fp8_output_scale,
) )
@@ -78,9 +75,7 @@ from sglang.srt.layers.communicator_dsa_cp import (
maybe_prefetch_next_full_attention_kv, maybe_prefetch_next_full_attention_kv,
) )
from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.layers.cp.utils import is_cp_v2_active
from sglang.srt.layers.dcp.planner import ( from sglang.srt.layers.dcp.planner import prepare_decode_context_parallel_metadata
prepare_decode_context_parallel_metadata,
)
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear, ColumnParallelLinear,
@@ -115,9 +110,7 @@ from sglang.srt.layers.moe.utils import (
) )
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.quantization.fp8 import Fp8Config from sglang.srt.layers.quantization.fp8 import Fp8Config
from sglang.srt.layers.quantization.fp8_utils import ( from sglang.srt.layers.quantization.fp8_utils import materialize_bpreshuffle_fp8_scale
materialize_bpreshuffle_fp8_scale,
)
from sglang.srt.layers.quantization.mxfp4_flashinfer_trtllm_moe import ( from sglang.srt.layers.quantization.mxfp4_flashinfer_trtllm_moe import (
maybe_fuse_routed_scale_and_shared_add, maybe_fuse_routed_scale_and_shared_add,
) )
@@ -183,11 +176,14 @@ from sglang.srt.models.deepseek_common.utils import (
is_wint4afp8_or_wint4a16_config, is_wint4afp8_or_wint4a16_config,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_device,
get_exec,
get_flags, get_flags,
get_forward, get_forward,
get_model, get_model,
get_parallel, get_parallel,
get_server_args, 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 (
@@ -381,9 +377,7 @@ class DeepseekV2MLP(nn.Module):
return down_output return down_output
if self.use_fused_clamp_act_mul and self.swiglu_limit is not None: if self.use_fused_clamp_act_mul and self.swiglu_limit is not None:
from aiter.ops.triton.fusions.fused_clamp_act_mul import ( from aiter.ops.triton.fusions.fused_clamp_act_mul import fused_clamp_act_mul
fused_clamp_act_mul,
)
if not self._fused_clamp_fp8_checked: if not self._fused_clamp_fp8_checked:
from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod
@@ -494,7 +488,7 @@ class MoEGate(nn.Module):
True, # is_vnni True, # is_vnni
) )
if get_server_args().enable_deterministic_inference: if get_exec().deterministic.enable_deterministic_inference:
return F.linear(hidden_states, self.weight, None) return F.linear(hidden_states, self.weight, None)
if ( if (
@@ -560,7 +554,7 @@ class DeepseekV2MoE(nn.Module):
n_shared_experts = ( n_shared_experts = (
0 if config.n_shared_experts is None else int(config.n_shared_experts) 0 if config.n_shared_experts is None else int(config.n_shared_experts)
) )
_fusion_disabled = get_server_args().disable_shared_experts_fusion _fusion_disabled = get_exec().moe.disable_shared_experts_fusion
# num_fused_shared_experts drives weight remapping in deepseek_weight_loader: # num_fused_shared_experts drives weight remapping in deepseek_weight_loader:
# mlp.shared_experts → mlp.experts.256 when > 0. # mlp.shared_experts → mlp.experts.256 when > 0.
@@ -630,8 +624,7 @@ class DeepseekV2MoE(nn.Module):
fused_shared_experts_scaling_factor = 1.0 / float(self.moe_ep_size) fused_shared_experts_scaling_factor = 1.0 / float(self.moe_ep_size)
self.experts = get_moe_impl_class(quant_config)( self.experts = get_moe_impl_class(quant_config)(
num_experts=num_experts_for_moe num_experts=num_experts_for_moe + get_exec().moe.ep_num_redundant_experts,
+ get_server_args().ep_num_redundant_experts,
num_fused_shared_experts=self.num_fused_shared_experts, num_fused_shared_experts=self.num_fused_shared_experts,
top_k=top_k_for_moe, top_k=top_k_for_moe,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
@@ -804,7 +797,7 @@ class DeepseekV2MoE(nn.Module):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts + get_server_args().ep_num_redundant_experts config.n_routed_experts + get_exec().moe.ep_num_redundant_experts
) )
self.renormalize = config.norm_topk_prob self.renormalize = config.norm_topk_prob
self.topk_group = config.topk_group self.topk_group = config.topk_group
@@ -1718,7 +1711,7 @@ class DeepseekV2AttentionMLA(
base=rope_theta, base=rope_theta,
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,
) )
if rope_scaling and rope_scaling.get("apply_yarn_scaling", True): if rope_scaling and rope_scaling.get("apply_yarn_scaling", True):
@@ -2069,7 +2062,7 @@ class DeepseekV2DecoderLayer(nn.Module):
rope_scaling = config.rope_scaling rope_scaling = config.rope_scaling
max_position_embeddings = config.max_position_embeddings max_position_embeddings = config.max_position_embeddings
self.speculative_algorithm = SpeculativeAlgorithm.from_string( self.speculative_algorithm = SpeculativeAlgorithm.from_string(
get_server_args().speculative_algorithm get_spec().speculative_algorithm
) )
self.dsa_enable_prefill_cp = dsa_enable_prefill_cp self.dsa_enable_prefill_cp = dsa_enable_prefill_cp
self.mla_enable_prefill_cp = mla_enable_prefill_cp self.mla_enable_prefill_cp = mla_enable_prefill_cp
@@ -2765,7 +2758,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
self.num_fused_shared_experts = 0 self.num_fused_shared_experts = 0
server_args = get_server_args() server_args = get_server_args()
if get_server_args().disable_shared_experts_fusion: if get_exec().moe.disable_shared_experts_fusion:
return return
disable_reason = None disable_reason = None
+14 -17
View File
@@ -29,18 +29,11 @@ from sglang.jit_kernel.dsv4 import (
fused_rope_inplace, fused_rope_inplace,
sglang_per_token_group_quant_fp8_dsv4_wo_a, sglang_per_token_group_quant_fp8_dsv4_wo_a,
) )
from sglang.kernels.ops.attention.deepseek_v4_rope import ( from sglang.kernels.ops.attention.deepseek_v4_rope import v4_rope_inplace_npu
v4_rope_inplace_npu, from sglang.kernels.ops.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
)
from sglang.kernels.ops.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8,
)
from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.compilation.compilation_config import register_split_op
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
from sglang.srt.distributed import ( from sglang.srt.distributed import get_pp_group, get_tp_group
get_pp_group,
get_tp_group,
)
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, use_symmetric_memory,
) )
@@ -136,7 +129,13 @@ from sglang.srt.models.deepseek_v2 import (
_is_npu, _is_npu,
_is_xpu, _is_xpu,
) )
from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args from sglang.srt.runtime_context import (
get_device,
get_exec,
get_forward,
get_parallel,
get_server_args,
)
if not _is_hip: if not _is_hip:
from sglang.srt.layers.utils.cp_utils import ( from sglang.srt.layers.utils.cp_utils import (
@@ -311,9 +310,7 @@ def _freqs_cis_to_cos_sin(
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.attention.deepseek_v4_backend import ( from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend
DeepseekV4AttnBackend,
)
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
DeepseekV4HipRadixBackend, DeepseekV4HipRadixBackend,
) )
@@ -575,7 +572,7 @@ class MQALayer(MqaAttentionBase):
base=self.rope_base, base=self.rope_base,
rope_scaling=self.rope_scaling, rope_scaling=self.rope_scaling,
is_neox_style=False, is_neox_style=False,
device=get_server_args().device, device=get_device().device,
) )
if _is_hip: if _is_hip:
@@ -2458,11 +2455,11 @@ class DeepseekV4ForCausalLM(nn.Module):
def determine_num_fused_shared_experts(self): def determine_num_fused_shared_experts(self):
self.num_fused_shared_experts = 0 self.num_fused_shared_experts = 0
if get_server_args().disable_shared_experts_fusion: if get_exec().moe.disable_shared_experts_fusion:
return return
disable_reason = None disable_reason = None
if get_server_args().enforce_shared_experts_fusion: if get_exec().moe.enforce_shared_experts_fusion:
if self.config.n_shared_experts != 1: if self.config.n_shared_experts != 1:
raise ValueError( raise ValueError(
"DeepSeek V4 shared-experts fusion expects exactly one shared " "DeepSeek V4 shared-experts fusion expects exactly one shared "
+10 -10
View File
@@ -24,17 +24,12 @@ import torch
from torch import nn from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import ( from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce
get_pp_group,
tensor_model_parallel_all_reduce,
)
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import is_dp_attention_enabled
is_dp_attention_enabled,
)
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
MergedColumnParallelLinear, MergedColumnParallelLinear,
@@ -62,7 +57,12 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_server_args,
get_stream,
)
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -165,7 +165,7 @@ class ExaoneMoESparseMoEBlock(nn.Module):
) )
self.experts = get_moe_impl_class(quant_config)( self.experts = get_moe_impl_class(quant_config)(
num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts,
top_k=config.num_experts_per_tok, top_k=config.num_experts_per_tok,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
intermediate_size=config.moe_intermediate_size, intermediate_size=config.moe_intermediate_size,
@@ -206,7 +206,7 @@ class ExaoneMoESparseMoEBlock(nn.Module):
if get_moe_a2a_backend().is_deepep(): if get_moe_a2a_backend().is_deepep():
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.num_experts + get_server_args().ep_num_redundant_experts config.num_experts + get_exec().moe.ep_num_redundant_experts
) )
self.top_k = config.num_experts_per_tok self.top_k = config.num_experts_per_tok
+5 -13
View File
@@ -18,11 +18,7 @@ from typing import Iterable, List, Optional, Set, Tuple, Union
import torch import torch
from torch import nn from torch import nn
from transformers import ( from transformers import Gemma4TextConfig, PretrainedConfig, PreTrainedModel
Gemma4TextConfig,
PretrainedConfig,
PreTrainedModel,
)
from sglang.kernels.ops.layernorm.gemma4_fused_ops import ( from sglang.kernels.ops.layernorm.gemma4_fused_ops import (
gemma4_fused_routing, gemma4_fused_routing,
@@ -31,9 +27,7 @@ from sglang.kernels.ops.layernorm.gemma4_fused_ops import (
gemma_rmsnorm_residual_scalar, gemma_rmsnorm_residual_scalar,
gemma_routing_post_topk, gemma_routing_post_topk,
) )
from sglang.srt.distributed import ( from sglang.srt.distributed import get_pp_group
get_pp_group,
)
from sglang.srt.layers.layernorm import Gemma4RMSNorm, RMSNorm from sglang.srt.layers.layernorm import Gemma4RMSNorm, RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
QKVParallelLinear, QKVParallelLinear,
@@ -55,10 +49,8 @@ from sglang.srt.model_loader.weight_utils import (
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
) )
from sglang.srt.models.gemma3_causal import Gemma3MLP, Gemma3TextScaledWordEmbedding from sglang.srt.models.gemma3_causal import Gemma3MLP, Gemma3TextScaledWordEmbedding
from sglang.srt.models.utils import ( from sglang.srt.models.utils import create_fused_set_kv_buffer_arg
create_fused_set_kv_buffer_arg, from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
)
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils import add_prefix, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -254,7 +246,7 @@ class Gemma4MoE(nn.Module):
experts_type = get_moe_impl_class(quant_config) experts_type = get_moe_impl_class(quant_config)
self.experts = experts_type( self.experts = experts_type(
num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
intermediate_size=config.moe_intermediate_size, intermediate_size=config.moe_intermediate_size,
layer_id=layer_id, layer_id=layer_id,
+2 -3
View File
@@ -29,7 +29,7 @@ from sglang.srt.layers.clippable_linear import (
) )
from sglang.srt.layers.layernorm import Gemma4RMSNorm from sglang.srt.layers.layernorm import Gemma4RMSNorm
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_mm, get_parallel
from sglang.srt.utils import add_prefix, get_device_capability, is_cuda, is_hip from sglang.srt.utils import add_prefix, get_device_capability, is_cuda, is_hip
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -181,9 +181,8 @@ class Gemma4VisionAttention(nn.Module):
@staticmethod @staticmethod
def _select_backend() -> str: def _select_backend() -> str:
"""Mirror VisionAttention._determine_attention_backend for consistency.""" """Mirror VisionAttention._determine_attention_backend for consistency."""
from sglang.srt.runtime_context import get_server_args
override = get_server_args().mm_attention_backend override = get_mm().mm_attention_backend
if override is not None: if override is not None:
return override return override
if is_cuda(): if is_cuda():
+4 -3
View File
@@ -84,6 +84,7 @@ from sglang.srt.models.deepseek_nextn import DeepseekV3ForCausalLMNextN
from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM
from sglang.srt.models.utils import WeightsMapper, apply_qk_norm from sglang.srt.models.utils import WeightsMapper, apply_qk_norm
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,
@@ -406,7 +407,7 @@ class Glm4MoeSparseMoeBlock(nn.Module):
self.n_shared_experts = config.n_shared_experts self.n_shared_experts = config.n_shared_experts
self.num_fused_shared_experts = ( self.num_fused_shared_experts = (
0 0
if get_server_args().disable_shared_experts_fusion if get_exec().moe.disable_shared_experts_fusion
else config.n_shared_experts else config.n_shared_experts
) )
@@ -526,7 +527,7 @@ class Glm4MoeSparseMoeBlock(nn.Module):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts + get_server_args().ep_num_redundant_experts config.n_routed_experts + get_exec().moe.ep_num_redundant_experts
) )
self.renormalize = config.norm_topk_prob self.renormalize = config.norm_topk_prob
self.topk_group = config.topk_group self.topk_group = config.topk_group
@@ -1178,7 +1179,7 @@ class Glm4MoeForCausalLM(nn.Module):
self.capture_aux_hidden_states = False self.capture_aux_hidden_states = False
def determine_num_fused_shared_experts(self): def determine_num_fused_shared_experts(self):
if get_server_args().disable_shared_experts_fusion: if get_exec().moe.disable_shared_experts_fusion:
return return
disable_reason = None disable_reason = None
+5 -4
View File
@@ -75,6 +75,7 @@ from sglang.srt.models.deepseek_common.deepseek_weight_loader import (
from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec,
get_forward, get_forward,
get_parallel, get_parallel,
get_server_args, get_server_args,
@@ -189,7 +190,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
self.n_shared_experts = config.n_shared_experts self.n_shared_experts = config.n_shared_experts
self.num_fused_shared_experts = ( self.num_fused_shared_experts = (
0 0
if get_server_args().disable_shared_experts_fusion if get_exec().moe.disable_shared_experts_fusion
else config.n_shared_experts else config.n_shared_experts
) )
self.config = config self.config = config
@@ -216,7 +217,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
self.experts = get_moe_impl_class(quant_config)( self.experts = get_moe_impl_class(quant_config)(
num_experts=config.n_routed_experts num_experts=config.n_routed_experts
+ self.num_fused_shared_experts + self.num_fused_shared_experts
+ get_server_args().ep_num_redundant_experts, + get_exec().moe.ep_num_redundant_experts,
num_fused_shared_experts=self.num_fused_shared_experts, num_fused_shared_experts=self.num_fused_shared_experts,
top_k=config.num_experts_per_tok + self.num_fused_shared_experts, top_k=config.num_experts_per_tok + self.num_fused_shared_experts,
hidden_size=config.hidden_size, hidden_size=config.hidden_size,
@@ -284,7 +285,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
# TODO: we will support tp < ep in the future # TODO: we will support tp < ep in the future
self.ep_size = get_parallel().moe_ep_size self.ep_size = get_parallel().moe_ep_size
self.num_experts = ( self.num_experts = (
config.n_routed_experts + get_server_args().ep_num_redundant_experts config.n_routed_experts + get_exec().moe.ep_num_redundant_experts
) )
self.renormalize = config.norm_topk_prob self.renormalize = config.norm_topk_prob
self.topk_group = config.topk_group self.topk_group = config.topk_group
@@ -928,7 +929,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
self, architecture: str = "Glm4MoeLiteForCausalLM" self, architecture: str = "Glm4MoeLiteForCausalLM"
): ):
self.num_fused_shared_experts = 0 self.num_fused_shared_experts = 0
if get_server_args().disable_shared_experts_fusion: if get_exec().moe.disable_shared_experts_fusion:
return return
disable_reason = None disable_reason = None
@@ -35,7 +35,7 @@ from sglang.srt.models.glm4_moe_lite import (
Glm4MoeLiteDecoderLayer, Glm4MoeLiteDecoderLayer,
Glm4MoeLiteForCausalLM, Glm4MoeLiteForCausalLM,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec
from sglang.srt.utils import BumpAllocator, add_prefix, is_npu from sglang.srt.utils import BumpAllocator, add_prefix, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -139,7 +139,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_parallel().tp_size self.tp_size = get_parallel().tp_size
if is_npu() and get_server_args().speculative_draft_model_quantization is None: if is_npu() and get_spec().speculative_draft_model_quantization is None:
quant_config = None quant_config = None
self.quant_config = quant_config self.quant_config = quant_config
@@ -156,7 +156,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
self.logits_processor = LogitsProcessor(config) self.logits_processor = LogitsProcessor(config)
self.num_fused_shared_experts = ( self.num_fused_shared_experts = (
0 if get_server_args().disable_shared_experts_fusion else 1 0 if get_exec().moe.disable_shared_experts_fusion else 1
) )
@torch.no_grad() @torch.no_grad()
+3 -3
View File
@@ -32,7 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -125,7 +125,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
nn.Module.__init__(self) nn.Module.__init__(self)
self.config = config self.config = config
self.tp_size = get_parallel().tp_size self.tp_size = get_parallel().tp_size
if is_npu() and get_server_args().speculative_draft_model_quantization is None: if is_npu() and get_spec().speculative_draft_model_quantization is None:
quant_config = None quant_config = None
self.quant_config = quant_config self.quant_config = quant_config
@@ -142,7 +142,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
self.logits_processor = LogitsProcessor(config) self.logits_processor = LogitsProcessor(config)
self.num_fused_shared_experts = ( self.num_fused_shared_experts = (
0 if get_server_args().disable_shared_experts_fusion else 1 0 if get_exec().moe.disable_shared_experts_fusion else 1
) )
@torch.no_grad() @torch.no_grad()

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