config: borrowed-record reads follow the config bags (#35908)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-08-23 01:19:20 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 64aa859da2
commit 362c2ee849
65 changed files with 617 additions and 377 deletions
+6 -6
View File
@@ -79,10 +79,10 @@ from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw
from sglang.srt.mem_cache.base_prefix_cache import EvictParams from sglang.srt.mem_cache.base_prefix_cache import EvictParams
from sglang.srt.model_executor.cuda_graph_config import Phase from sglang.srt.model_executor.cuda_graph_config import Phase, cuda_graph_fully_disabled
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel, get_schedule
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -487,7 +487,7 @@ class TreeCacheNamespace(SimpleNamespace):
def extend(reqs, model_runner): def extend(reqs, model_runner):
# Create dummy tree_cache for benchmarks (no prefix caching, just allocation) # Create dummy tree_cache for benchmarks (no prefix caching, just allocation)
dummy_tree_cache = TreeCacheNamespace( dummy_tree_cache = TreeCacheNamespace(
page_size=model_runner.server_args.page_size, page_size=get_schedule().page_size,
device=model_runner.device, device=model_runner.device,
token_to_kv_pool_allocator=model_runner.token_to_kv_pool_allocator, token_to_kv_pool_allocator=model_runner.token_to_kv_pool_allocator,
) )
@@ -542,14 +542,14 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
prepare_mlp_sync_batch_raw( prepare_mlp_sync_batch_raw(
batch, batch,
model_runner=model_runner, model_runner=model_runner,
dp_size=model_runner.server_args.dp_size, dp_size=get_parallel().dp_size,
attn_tp_size=get_parallel().attn_tp_size, attn_tp_size=get_parallel().attn_tp_size,
attn_cp_size=model_runner.ps.attn_cp_size, attn_cp_size=model_runner.ps.attn_cp_size,
tp_group=model_runner.tp_group, tp_group=model_runner.tp_group,
get_idle_batch=None, get_idle_batch=None,
disable_cuda_graph=model_runner.server_args.disable_cuda_graph, disable_cuda_graph=cuda_graph_fully_disabled(),
require_mlp_tp_gather=require_mlp_tp_gather(model_runner.server_args), require_mlp_tp_gather=require_mlp_tp_gather(model_runner.server_args),
disable_overlap_schedule=model_runner.server_args.disable_overlap_schedule, disable_overlap_schedule=get_schedule().disable_overlap_schedule,
offload_tags=set(), offload_tags=set(),
) )
@@ -16,6 +16,10 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
import torch import torch
from sglang.srt.runtime_context import (
get_schedule,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Bounded wait for a watermark advance before re-enqueueing a deferred staging # Bounded wait for a watermark advance before re-enqueueing a deferred staging
@@ -548,7 +552,7 @@ class PrefillStagingStrategy:
self.staging_buffer = staging_buffer self.staging_buffer = staging_buffer
page_size = kv_manager.kv_buffer_tensors["page_size"] page_size = kv_manager.kv_buffer_tensors["page_size"]
self.full_chunk_pages = ( self.full_chunk_pages = (
staging_grid_tokens(kv_manager.server_args.chunked_prefill_size, page_size) staging_grid_tokens(get_schedule().chunked_prefill_size, page_size)
// page_size // page_size
) )
+11 -11
View File
@@ -94,7 +94,11 @@ from sglang.srt.observability.req_time_stats import (
set_schedule_time_batch, set_schedule_time_batch,
set_time_batch, set_time_batch,
) )
from sglang.srt.runtime_context import get_disagg, get_parallel from sglang.srt.runtime_context import (
get_disagg,
get_memory,
get_parallel,
)
from sglang.srt.utils import ceil_align, get_num_new_pages, is_npu from sglang.srt.utils import ceil_align, get_num_new_pages, is_npu
from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.network import NetworkAddress
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
@@ -477,7 +481,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
full_len = ceil_align(full_len, page_size) full_len = ceil_align(full_len, page_size)
swa_len = ceil_align(swa_len, page_size) swa_len = ceil_align(swa_len, page_size)
swa_reserved = self.num_reserved_decode_tokens swa_reserved = self.num_reserved_decode_tokens
if self.scheduler.server_args.disable_radix_cache: if get_memory().disable_radix_cache:
swa_reserved = 0 swa_reserved = 0
return ( return (
full_len + self.num_reserved_decode_tokens, full_len + self.num_reserved_decode_tokens,
@@ -556,7 +560,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
req_to_token_pool=getattr(self, "req_to_token_pool", None), req_to_token_pool=getattr(self, "req_to_token_pool", None),
) )
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device kv_args.ib_device = get_disagg().disaggregation_ib_device
kv_args.gpu_id = self.scheduler.ps.gpu_id kv_args.gpu_id = self.scheduler.ps.gpu_id
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER) kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class( kv_manager = kv_manager_class(
@@ -608,7 +612,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
) )
# NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue # NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue
if _is_fake_transfer(req, self.scheduler.server_args): if _is_fake_transfer(req):
decode_req.kv_receiver.init(0) decode_req.kv_receiver.init(0)
return return
@@ -663,9 +667,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
self, req: Req, is_rebootstrap: bool = False self, req: Req, is_rebootstrap: bool = False
) -> DecodeRequest: ) -> DecodeRequest:
backend = ( backend = (
TransferBackend.FAKE TransferBackend.FAKE if _is_fake_transfer(req) else self.transfer_backend
if _is_fake_transfer(req, self.scheduler.server_args)
else self.transfer_backend
) )
kv_receiver_class = get_kv_class(backend, KVClassType.RECEIVER) kv_receiver_class = get_kv_class(backend, KVClassType.RECEIVER)
@@ -1455,7 +1457,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
if ( if (
self.scheduler.enable_hisparse self.scheduler.enable_hisparse
and isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool) and isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool)
and not _is_fake_transfer(decode_req.req, self.scheduler.server_args) and not _is_fake_transfer(decode_req.req)
): ):
# alloc_logical_only() already allocated the shared logical pages # alloc_logical_only() already allocated the shared logical pages
# used by C4 indexer and C128 KV. These device buffers do not use # used by C4 indexer and C128 KV. These device buffers do not use
@@ -2059,7 +2061,7 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
else 0 else 0
) )
if _is_fake_transfer(decode_req.req, self.scheduler.server_args): if _is_fake_transfer(decode_req.req):
pass pass
elif actual_room == 0: elif actual_room == 0:
# Should never happen: _poll_with_metadata_gate already confirmed # Should never happen: _poll_with_metadata_gate already confirmed
@@ -2197,7 +2199,6 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
self.gloo_group, self.gloo_group,
decode_reqs=self.queue, decode_reqs=self.queue,
metadata_buffers=self.metadata_buffers, metadata_buffers=self.metadata_buffers,
server_args=self.scheduler.server_args,
) )
def _poll_with_staging(self) -> list: def _poll_with_staging(self) -> list:
@@ -2206,7 +2207,6 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
self.staging_handler, self.staging_handler,
self.gloo_group, self.gloo_group,
metadata_buffers=self.metadata_buffers, metadata_buffers=self.metadata_buffers,
server_args=self.scheduler.server_args,
) )
def _init_staging_handler(self, kv_manager): def _init_staging_handler(self, kv_manager):
+2 -3
View File
@@ -253,7 +253,7 @@ class PrefillBootstrapQueue:
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = ( kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
self.metadata_buffers.get_buf_infos() self.metadata_buffers.get_buf_infos()
) )
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device kv_args.ib_device = get_disagg().disaggregation_ib_device
kv_args.gpu_id = self.scheduler.ps.gpu_id kv_args.gpu_id = self.scheduler.ps.gpu_id
req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None) req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
@@ -438,8 +438,7 @@ class PrefillBootstrapQueue:
failed_reqs.append(req) failed_reqs.append(req)
elif poll == KVPoll.Bootstrapping: elif poll == KVPoll.Bootstrapping:
if ( if (
req.prefill_attempt_count req.prefill_attempt_count < get_disagg().optimistic_prefill_attempts
< self.scheduler.server_args.optimistic_prefill_attempts
and not req.is_retracted # engine paused and not req.is_retracted # engine paused
): ):
if not self.ensure_metadata_buffer(req): if not self.ensure_metadata_buffer(req):
+11 -15
View File
@@ -22,6 +22,9 @@ import torch.distributed as dist
from sglang.srt.configs.model_config import get_dsa_index_topk from sglang.srt.configs.model_config import get_dsa_index_topk
from sglang.srt.disaggregation.base import KVPoll from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import (
get_disagg,
)
from sglang.srt.utils import is_hip, is_npu from sglang.srt.utils import is_hip, is_npu
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -33,7 +36,6 @@ if TYPE_CHECKING:
CommonKVSender, CommonKVSender,
) )
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.server_args import ServerArgs
if is_npu(): if is_npu():
from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import ( from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import (
@@ -168,14 +170,14 @@ def _poll_with_failure_injection(pollers) -> List[int]:
return [int(poller.poll()) for poller in pollers] return [int(poller.poll()) for poller in pollers]
def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool: def _is_fake_transfer(req: Req) -> bool:
return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or ( return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
req.bootstrap_host is None req.bootstrap_host is None
and server_args.disaggregation_transfer_backend == "fake" and get_disagg().disaggregation_transfer_backend == "fake"
) )
def _apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args) -> None: def _apply_metadata_gate(polls, decode_reqs, metadata_buffers) -> None:
"""Downgrade Success → Transferring for requests whose metadata hasn't landed. """Downgrade Success → Transferring for requests whose metadata hasn't landed.
Mutates `polls` in-place. Called before all-reduce so that MIN across TP Mutates `polls` in-place. Called before all-reduce so that MIN across TP
@@ -184,7 +186,7 @@ def _apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args) -> N
for i, poll_val in enumerate(polls): for i, poll_val in enumerate(polls):
if poll_val == int(KVPoll.Success): if poll_val == int(KVPoll.Success):
decode_req = decode_reqs[i] decode_req = decode_reqs[i]
if _is_fake_transfer(decode_req.req, server_args): if _is_fake_transfer(decode_req.req):
continue continue
actual_room = metadata_buffers.bootstrap_room[ actual_room = metadata_buffers.bootstrap_room[
decode_req.metadata_buffer_index, 0 decode_req.metadata_buffer_index, 0
@@ -205,18 +207,13 @@ def poll_and_all_reduce(
gloo_group: dist.ProcessGroup, gloo_group: dist.ProcessGroup,
decode_reqs=None, decode_reqs=None,
metadata_buffers: Optional[MetadataBuffers] = None, metadata_buffers: Optional[MetadataBuffers] = None,
server_args: Optional[ServerArgs] = None,
): ):
# at a certain prob, the poll is failed to simulate failure # at a certain prob, the poll is failed to simulate failure
polls = _poll_with_failure_injection(pollers) polls = _poll_with_failure_injection(pollers)
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed. # Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
if ( if decode_reqs is not None and metadata_buffers is not None:
decode_reqs is not None _apply_metadata_gate(polls, decode_reqs, metadata_buffers)
and metadata_buffers is not None
and server_args is not None
):
_apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args)
return _all_reduce_polls(polls, gloo_group) return _all_reduce_polls(polls, gloo_group)
@@ -239,7 +236,6 @@ def poll_and_all_reduce_with_staging(
staging_handler, staging_handler,
gloo_group: dist.ProcessGroup, gloo_group: dist.ProcessGroup,
metadata_buffers: Optional[MetadataBuffers] = None, metadata_buffers: Optional[MetadataBuffers] = None,
server_args: Optional[ServerArgs] = None,
): ):
"""Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce.""" """Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce."""
for decode_req in decode_reqs: for decode_req in decode_reqs:
@@ -265,8 +261,8 @@ def poll_and_all_reduce_with_staging(
): ):
raw_polls[i] = int(KVPoll.Transferring) raw_polls[i] = int(KVPoll.Transferring)
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed. # Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
if metadata_buffers is not None and server_args is not None: if metadata_buffers is not None:
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers, server_args) _apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers)
return _all_reduce_polls(raw_polls, gloo_group) return _all_reduce_polls(raw_polls, gloo_group)
+10 -7
View File
@@ -27,7 +27,10 @@ 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 initialize_dp_attention from sglang.srt.layers.dp_attention import initialize_dp_attention
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import (
get_exec,
get_parallel,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
@@ -111,7 +114,7 @@ def init_torch_distributed(
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request # Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs) # Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
if server_args.pre_warm_nccl and ( if get_exec().comm.pre_warm_nccl and (
ps.tp_size > 1 or ps.pp_size > 1 or ps.moe_ep_size > 1 ps.tp_size > 1 or ps.pp_size > 1 or ps.moe_ep_size > 1
): ):
_prewarm_nccl( _prewarm_nccl(
@@ -185,11 +188,11 @@ def _resolve_dist_init_method(*, server_args: ServerArgs, dist_port: int) -> str
def _set_all_reduce_flags(*, server_args: ServerArgs) -> None: def _set_all_reduce_flags(*, server_args: ServerArgs) -> None:
set_custom_all_reduce(not server_args.disable_custom_all_reduce) set_custom_all_reduce(not get_exec().comm.disable_custom_all_reduce)
set_mscclpp_all_reduce(server_args.enable_mscclpp) set_mscclpp_all_reduce(server_args.enable_mscclpp)
set_torch_symm_mem_all_reduce(server_args.enable_torch_symm_mem) set_torch_symm_mem_all_reduce(get_exec().comm.enable_torch_symm_mem)
set_flashinfer_allreduce_only( set_flashinfer_allreduce_only(
server_args.flashinfer_allreduce_fusion_backend is not None get_exec().comm.flashinfer_allreduce_fusion_backend is not None
) )
@@ -242,7 +245,7 @@ def _init_parallel_groups(
local_rank=gpu_id, local_rank=gpu_id,
distributed_init_method=dist_init_method, distributed_init_method=dist_init_method,
timeout=server_args.dist_timeout, timeout=server_args.dist_timeout,
moe_a2a_backend=server_args.moe_a2a_backend, moe_a2a_backend=get_exec().moe.moe_a2a_backend,
recovered_rank=is_ep_joiner, recovered_rank=is_ep_joiner,
max_world_size=server_args.max_ep_size, max_world_size=server_args.max_ep_size,
) )
@@ -260,7 +263,7 @@ def _init_parallel_groups(
and server_args.enable_two_batch_overlap and server_args.enable_two_batch_overlap
and get_parallel().enable_dsa_prefill_context_parallel and get_parallel().enable_dsa_prefill_context_parallel
), ),
enable_symm_mem=server_args.enable_symm_mem, enable_symm_mem=get_exec().comm.enable_symm_mem,
recovered_rank=is_ep_joiner, recovered_rank=is_ep_joiner,
rank_offset=rank_offset, rank_offset=rank_offset,
max_world_size=server_args.max_ep_size, max_world_size=server_args.max_ep_size,
+6 -4
View File
@@ -16,6 +16,10 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional
from pydantic import ValidationError from pydantic import ValidationError
from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan
from sglang.srt.runtime_context import (
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__)
@@ -400,7 +404,7 @@ class RuntimeHandle:
result = { result = {
"model_path": self.tokenizer_manager.model_path, "model_path": self.tokenizer_manager.model_path,
"served_model_name": self.tokenizer_manager.served_model_name, "served_model_name": self.tokenizer_manager.served_model_name,
"tokenizer_path": self.tokenizer_manager.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.tokenizer_manager.config_value("weight_version"), "weight_version": self.tokenizer_manager.config_value("weight_version"),
"load_format": self.tokenizer_manager.config_value("load_format"), "load_format": self.tokenizer_manager.config_value("load_format"),
@@ -461,9 +465,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.tokenizer_manager.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(
+5 -4
View File
@@ -394,7 +394,7 @@ async def lifespan(fast_api_app: FastAPI):
if ( if (
getattr(fast_api_app, "is_single_tokenizer_mode", False) getattr(fast_api_app, "is_single_tokenizer_mode", False)
and get_serving().grpc_port is not None and get_serving().grpc_port is not None
and not (server_args.smg_grpc_mode or server_args.grpc_mode) and not (get_serving().smg_grpc_mode or server_args.grpc_mode)
): ):
grpc_handle = _start_native_grpc_server_for_runtime( grpc_handle = _start_native_grpc_server_for_runtime(
server_args=server_args, server_args=server_args,
@@ -484,6 +484,7 @@ from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_router
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_disagg, get_disagg,
get_exec, get_exec,
get_lora,
get_model, get_model,
get_parallel, get_parallel,
get_serving, get_serving,
@@ -744,9 +745,9 @@ async def model_info():
# Manager-owned, and moved by a weight update alongside `model_path`: # Manager-owned, and moved by a weight update alongside `model_path`:
# this is where a client reads the identity the server answers under. # this is where a client reads the identity the server answers under.
"served_model_name": _global_state.tokenizer_manager.served_model_name, "served_model_name": _global_state.tokenizer_manager.served_model_name,
"tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path, "tokenizer_path": get_serving().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": get_serving().preferred_sampling_params,
"weight_version": _global_state.tokenizer_manager.config_value( "weight_version": _global_state.tokenizer_manager.config_value(
"weight_version" "weight_version"
), ),
@@ -1857,7 +1858,7 @@ async def available_models():
) )
# Add loaded LoRA adapters # Add loaded LoRA adapters
if _global_state.tokenizer_manager.server_args.enable_lora: if get_lora().enable_lora:
lora_registry = _global_state.tokenizer_manager.lora_registry lora_registry = _global_state.tokenizer_manager.lora_registry
for _, lora_ref in lora_registry.get_all_adapters().items(): for _, lora_ref in lora_registry.get_all_adapters().items():
model_cards.append( model_cards.append(
+8 -7
View File
@@ -18,6 +18,9 @@ import logging
import multiprocessing as mp import multiprocessing as mp
import os import os
from sglang.srt.runtime_context import (
get_serving,
)
from sglang.srt.utils.common import kill_itself_when_parent_died, kill_process_tree from sglang.srt.utils.common import kill_itself_when_parent_died, kill_process_tree
from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.network import NetworkAddress
from sglang.srt.utils.watchdog import SubprocessWatchdog from sglang.srt.utils.watchdog import SubprocessWatchdog
@@ -36,12 +39,10 @@ def _loopback_host(host: str) -> str:
return host return host
def build_sidecar_endpoint(server_args) -> str: def build_sidecar_endpoint(host: str, grpc_port: int) -> str:
"""Both halves of the endpoint come from the argument: this is a helper """Both halves are passed in: this is a string helper, and the caller is
over a config object, callable before anything is published.""" the one that knows where the effective values live."""
return NetworkAddress( return NetworkAddress(_loopback_host(host), grpc_port).to_url()
_loopback_host(server_args.host), server_args.grpc_port
).to_url()
def _parse_sidecar_args(args: list[str] | None) -> tuple[list[str], float]: def _parse_sidecar_args(args: list[str] | None) -> tuple[list[str], float]:
@@ -117,7 +118,7 @@ def start_sidecar(server_args) -> Sidecar:
module_name = server_args.sidecar module_name = server_args.sidecar
assert module_name is not None assert module_name is not None
sidecar_args, shutdown_timeout = _parse_sidecar_args(server_args.sidecar_args) sidecar_args, shutdown_timeout = _parse_sidecar_args(server_args.sidecar_args)
endpoint = build_sidecar_endpoint(server_args) endpoint = build_sidecar_endpoint(server_args.host, get_serving().grpc_port)
proc = mp.get_context("spawn").Process( proc = mp.get_context("spawn").Process(
name=f"sglang_sidecar_{module_name}", name=f"sglang_sidecar_{module_name}",
target=_run_sidecar, target=_run_sidecar,
+7 -3
View File
@@ -26,6 +26,10 @@ from typing import Optional
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import Response from fastapi.responses import Response
from sglang.srt.runtime_context import (
configured_pp_size,
get_parallel,
)
from sglang.srt.utils import get_device_name from sglang.srt.utils import get_device_name
from sglang.version import __version__ from sglang.version import __version__
@@ -144,9 +148,9 @@ async def get_loads(
"accelerator": _accelerator_name(), "accelerator": _accelerator_name(),
"num_accelerators": _num_accelerators_per_dp_rank( "num_accelerators": _num_accelerators_per_dp_rank(
tokenizer_manager.server_args.tp_size, tokenizer_manager.server_args.tp_size,
tokenizer_manager.server_args.pp_size, configured_pp_size(),
tokenizer_manager.server_args.dp_size, get_parallel().dp_size,
tokenizer_manager.server_args.enable_dp_attention, get_parallel().enable_dp_attention,
), ),
"loads": loads, "loads": loads,
} }
@@ -1,6 +1,10 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.runtime_context import (
get_parallel,
get_schedule,
get_spec,
)
""" """
end to end attention solution with aiter kernels end to end attention solution with aiter kernels
@@ -152,7 +156,7 @@ class AiterAttnBackend(AttentionBackend):
self.input_dtype = model_runner.model_config.dtype self.input_dtype = model_runner.model_config.dtype
self.page_size = model_runner.server_args.page_size self.page_size = get_schedule().page_size
self.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd) self.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd)
@@ -2957,7 +2961,7 @@ class AiterMultiStepDraftBackend:
# Cached variables for generate_draft_decode_kv_indices # Cached variables for generate_draft_decode_kv_indices
self.req_to_token_pool = model_runner.req_to_token_pool self.req_to_token_pool = model_runner.req_to_token_pool
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1] self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
self.page_size = model_runner.server_args.page_size self.page_size = get_schedule().page_size
def common_template( def common_template(
self, forward_batch: ForwardBatch, kv_indices_buffer: torch.Tensor, call_fn: int self, forward_batch: ForwardBatch, kv_indices_buffer: torch.Tensor, call_fn: int
@@ -13,7 +13,10 @@ from sglang.srt.configs.linear_attn_model_registry import (
get_linear_attn_config, get_linear_attn_config,
import_backend_class, import_backend_class,
) )
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import (
get_parallel,
get_spec,
)
from sglang.srt.utils import get_device_capability, is_hip, is_musa, is_npu from sglang.srt.utils import get_device_capability, is_hip, is_musa, is_npu
_is_musa = is_musa() _is_musa = is_musa()
@@ -49,7 +52,7 @@ def create_flashinfer_backend(runner):
) )
# Init streams # Init streams
if runner.server_args.speculative_algorithm == "EAGLE": if get_spec().speculative_algorithm == "EAGLE":
if ( if (
not hasattr(runner, "plan_stream_for_flashinfer") not hasattr(runner, "plan_stream_for_flashinfer")
or not runner.plan_stream_for_flashinfer or not runner.plan_stream_for_flashinfer
@@ -70,10 +73,7 @@ def create_flashinfer_backend(runner):
def create_trtllm_mla_backend(runner): def create_trtllm_mla_backend(runner):
if not runner.use_mla_backend: if not runner.use_mla_backend:
raise ValueError("trtllm_mla backend can only be used with MLA models.") raise ValueError("trtllm_mla backend can only be used with MLA models.")
if ( if get_parallel().dcp_enabled and get_spec().speculative_algorithm is not None:
get_parallel().dcp_enabled
and runner.server_args.speculative_algorithm is not None
):
_, decode_backend = runner.server_args.get_attention_backends() _, decode_backend = runner.server_args.get_attention_backends()
if decode_backend == "trtllm_mla": if decode_backend == "trtllm_mla":
raise ValueError( raise ValueError(
@@ -265,7 +265,7 @@ def create_hpc_ops_backend(runner):
raise ValueError( raise ValueError(
"Cross attention is not supported in the hpc_ops attention backend." "Cross attention is not supported in the hpc_ops attention backend."
) )
if runner.server_args.speculative_algorithm is not None: if get_spec().speculative_algorithm is not None:
raise ValueError( raise ValueError(
"hpc_ops backend does not support speculative decoding for now." "hpc_ops backend does not support speculative decoding for now."
) )
@@ -69,7 +69,10 @@ from sglang.srt.layers.attention.verify_mask import (
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.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.runtime_context import (
get_parallel,
get_spec,
)
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
from sglang.srt.speculative.ragged_verify import ( from sglang.srt.speculative.ragged_verify import (
RaggedVerifyMode, RaggedVerifyMode,
@@ -576,7 +579,7 @@ class DeepseekV4AttnBackend(
self._q8kv8_qpad_buf = None self._q8kv8_qpad_buf = None
self._q8kv8_attn_sink_pad = None self._q8kv8_attn_sink_pad = None
self._q8kv8_identity_scale = None self._q8kv8_identity_scale = None
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = get_spec().speculative_eagle_topk or 0
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4" assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
self.mtp_enabled = self.topk > 0 self.mtp_enabled = self.topk > 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
@@ -36,7 +36,10 @@ from sglang.srt.layers.attention.dsv4.metadata import (
) )
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.runtime_context import (
get_parallel,
get_spec,
)
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
from sglang.srt.utils import ceil_align from sglang.srt.utils import ceil_align
@@ -456,7 +459,7 @@ class DeepseekV4HipRadixBackend(
self.enable_deepseek_v4_fp4_indexer: bool = ( self.enable_deepseek_v4_fp4_indexer: bool = (
model_runner.server_args.enable_deepseek_v4_fp4_indexer model_runner.server_args.enable_deepseek_v4_fp4_indexer
) )
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = get_spec().speculative_eagle_topk or 0
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4" assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
self.mtp_enabled = self.topk > 0 self.mtp_enabled = self.topk > 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
@@ -15,7 +15,11 @@ from typing import (
import torch import torch
from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_spec,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from sglang.kernels.ops.attention.dsa.dequant_k_cache import ( from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
@@ -309,7 +313,7 @@ class DeepseekSparseAttnBackend(
assert isinstance(model_runner.page_size, int) assert isinstance(model_runner.page_size, int)
self.real_page_size = model_runner.page_size self.real_page_size = model_runner.page_size
self.num_splits = ( self.num_splits = (
1 if model_runner.server_args.enable_deterministic_inference else 0 1 if get_exec().deterministic.enable_deterministic_inference else 0
) )
self.use_dsa = is_deepseek_dsa(model_runner.model_config.hf_config) self.use_dsa = is_deepseek_dsa(model_runner.model_config.hf_config)
assert self.use_dsa, "DSA backend only supports DeepSeek DSA" assert self.use_dsa, "DSA backend only supports DeepSeek DSA"
@@ -334,10 +338,8 @@ class DeepseekSparseAttnBackend(
self.use_mha: bool = False self.use_mha: bool = False
self.supports_mha_one_shot: bool = True self.supports_mha_one_shot: bool = True
self.dsa_prefill_impl: _DSA_IMPL_T = ( self.dsa_prefill_impl: _DSA_IMPL_T = get_exec().kernel.dsa_prefill_backend
model_runner.server_args.dsa_prefill_backend self.dsa_decode_impl: _DSA_IMPL_T = get_exec().kernel.dsa_decode_backend
)
self.dsa_decode_impl: _DSA_IMPL_T = model_runner.server_args.dsa_decode_backend
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend( self.dsa_topk_backend: DSATopKBackend = DSATopKBackend(
model_runner.server_args.dsa_topk_backend model_runner.server_args.dsa_topk_backend
) )
@@ -390,7 +392,7 @@ class DeepseekSparseAttnBackend(
) )
# Speculative decoding # Speculative decoding
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = get_spec().speculative_eagle_topk or 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_step_id = speculative_step_id self.speculative_step_id = speculative_step_id
@@ -1,6 +1,9 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import (
get_exec,
get_parallel,
)
""" """
Support different attention backends. Support different attention backends.
@@ -403,7 +406,7 @@ class FlashInferAttnBackend(AttentionBackend):
# Also set split tile sizes for prefill and decode from environment variables, and disable kv split for cuda graph # Also set split tile sizes for prefill and decode from environment variables, and disable kv split for cuda graph
# More information can be found here: https://github.com/flashinfer-ai/flashinfer/pull/1675 # More information can be found here: https://github.com/flashinfer-ai/flashinfer/pull/1675
self.enable_deterministic = ( self.enable_deterministic = (
model_runner.server_args.enable_deterministic_inference get_exec().deterministic.enable_deterministic_inference
) )
self.prefill_split_tile_size = None self.prefill_split_tile_size = None
self.decode_split_tile_size = None self.decode_split_tile_size = None
@@ -1,6 +1,11 @@
from __future__ import annotations from __future__ import annotations
from sglang.srt.runtime_context import get_disagg, get_exec, get_parallel, get_schedule from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_parallel,
get_schedule,
)
""" """
Support attention backend for flashinfer MLA. Support attention backend for flashinfer MLA.
@@ -1101,7 +1106,7 @@ class FlashInferMLAMultiStepDraftBackend:
# Cached variables for generate_draft_decode_kv_indices # Cached variables for generate_draft_decode_kv_indices
self.req_to_token_pool = model_runner.req_to_token_pool self.req_to_token_pool = model_runner.req_to_token_pool
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1] self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
self.page_size = model_runner.server_args.page_size self.page_size = get_schedule().page_size
def common_template( def common_template(
self, self,
@@ -12,6 +12,9 @@ from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import BaseIndexerMeta
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, 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_spec,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.attention.verify_mask import VerifyMask from sglang.srt.layers.attention.verify_mask import VerifyMask
@@ -33,12 +36,8 @@ class HybridAttnBackend(AttentionBackend):
self.data_type = model_runner.kv_cache_dtype self.data_type = model_runner.kv_cache_dtype
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_pool = model_runner.req_to_token_pool self.req_to_token_pool = model_runner.req_to_token_pool
self.spec_attn_is_decode = ( self.spec_attn_is_decode = get_spec().speculative_attention_mode == "decode"
model_runner.server_args.speculative_attention_mode == "decode" self.spec_attn_is_prefill = get_spec().speculative_attention_mode == "prefill"
)
self.spec_attn_is_prefill = (
model_runner.server_args.speculative_attention_mode == "prefill"
)
# Gates the FutureMap's per-step seq_lens D2H (decide_needs_cpu_seq_lens # Gates the FutureMap's per-step seq_lens D2H (decide_needs_cpu_seq_lens
# ORs it across backends). Count only what runs in the spec decode loop: # ORs it across backends). Count only what runs in the spec decode loop:
# decode always, prefill only when mode=prefill routes verify to it -- # decode always, prefill only when mode=prefill routes verify to it --
@@ -121,10 +120,7 @@ class HybridAttnBackend(AttentionBackend):
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens) self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens)
if ( if get_spec().speculative_algorithm is not None and self.spec_attn_is_prefill:
self.model_runner.server_args.speculative_algorithm is not None
and self.spec_attn_is_prefill
):
# When speculative decoding is enabled, we need to initialize the backend # When speculative decoding is enabled, we need to initialize the backend
# that will be used for target_verify. # that will be used for target_verify.
self.prefill_backend.init_cuda_graph_state(max_bs, max_num_tokens) self.prefill_backend.init_cuda_graph_state(max_bs, max_num_tokens)
@@ -31,6 +31,7 @@ from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec, get_exec,
get_memory, get_memory,
get_spec,
mamba_cache_chunk_size, mamba_cache_chunk_size,
) )
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
@@ -47,7 +48,7 @@ class MambaAttnBackendBase(AttentionBackend):
super().__init__() super().__init__()
self.pad_slot_id = PAD_SLOT_ID self.pad_slot_id = PAD_SLOT_ID
self.device = model_runner.device self.device = model_runner.device
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = get_spec().speculative_eagle_topk or 0
self.is_draft_worker = model_runner.is_draft_worker self.is_draft_worker = model_runner.is_draft_worker
self.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_pool self.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
@@ -31,6 +31,9 @@ elif is_cpu():
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import (
get_spec,
)
class KDAKernelDispatcher: class KDAKernelDispatcher:
@@ -384,7 +387,7 @@ class KDAAttnBackend(MambaAttnBackendBase):
# traversal). Reject EAGLE tree verify (topk > 1) early at setup, keyed on the # traversal). Reject EAGLE tree verify (topk > 1) early at setup, keyed on the
# verify backend (not decode). The kernel keeps a per-call # verify backend (not decode). The kernel keeps a per-call
# retrieve_parent_token backstop that also covers ngram tree. # retrieve_parent_token backstop that also covers ngram tree.
speculative_topk = model_runner.server_args.speculative_eagle_topk or 1 speculative_topk = get_spec().speculative_eagle_topk or 1
if verify_backend.is_flashinfer() and speculative_topk > 1: if verify_backend.is_flashinfer() and speculative_topk > 1:
raise ValueError( raise ValueError(
"KDA FlashInfer speculative decoding only supports topk=1 " "KDA FlashInfer speculative decoding only supports topk=1 "
@@ -38,7 +38,12 @@ from sglang.srt.model_executor.cuda_graph_config import (
cuda_graph_fully_disabled, cuda_graph_fully_disabled,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_schedule,
get_spec,
)
from sglang.srt.speculative.spec_utils import ( from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width, draft_kv_indices_buffer_width,
draft_kv_indices_used_len, draft_kv_indices_used_len,
@@ -254,7 +259,7 @@ class TritonAttnBackend(AttentionBackend):
self.static_kv_splits = get_bool_env_var( self.static_kv_splits = get_bool_env_var(
"SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false" "SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false"
) )
self.max_kv_splits = model_runner.server_args.triton_attention_num_kv_splits self.max_kv_splits = get_exec().kernel.triton_attention_num_kv_splits
if self.use_mla and not _is_xpu: if self.use_mla and not _is_xpu:
self.max_kv_splits = _mla_decode_kv_splits_cap( self.max_kv_splits = _mla_decode_kv_splits_cap(
self.max_kv_splits, self.max_kv_splits,
@@ -280,11 +285,11 @@ class TritonAttnBackend(AttentionBackend):
cuda_graph_fully_disabled() cuda_graph_fully_disabled()
or check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE) or check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
) )
and model_runner.server_args.chunked_prefill_size == -1 and get_schedule().chunked_prefill_size == -1
) )
self.enable_deterministic = ( self.enable_deterministic = (
model_runner.server_args.enable_deterministic_inference get_exec().deterministic.enable_deterministic_inference
) )
if self.enable_deterministic: if self.enable_deterministic:
@@ -1987,7 +1992,7 @@ class TritonMultiStepDraftBackend:
# Cached variables for generate_draft_decode_kv_indices # Cached variables for generate_draft_decode_kv_indices
self.req_to_token_pool = model_runner.req_to_token_pool self.req_to_token_pool = model_runner.req_to_token_pool
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1] self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
self.page_size = model_runner.server_args.page_size self.page_size = get_schedule().page_size
def common_template( def common_template(
self, self,
@@ -13,7 +13,11 @@ from sglang.kernels.ops.kvcache.kv_indices import (
) )
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_spec,
)
from sglang.srt.utils import get_bool_env_var, get_device_core_count from sglang.srt.utils import get_bool_env_var, get_device_core_count
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -109,7 +113,7 @@ class WaveAttnBackend(AttentionBackend):
self.static_kv_splits = get_bool_env_var( self.static_kv_splits = get_bool_env_var(
"SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false" "SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false"
) )
self.max_kv_splits = model_runner.server_args.triton_attention_num_kv_splits self.max_kv_splits = get_exec().kernel.triton_attention_num_kv_splits
self.v_head_dim = model_runner.token_to_kv_pool.get_value_buffer(0).shape[-1] self.v_head_dim = model_runner.token_to_kv_pool.get_value_buffer(0).shape[-1]
self.forward_metadata: ForwardMetadata = None self.forward_metadata: ForwardMetadata = None
@@ -15,7 +15,10 @@ from sglang.srt.layers.attention.flashattention_backend import (
from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_schedule, get_spec from sglang.srt.runtime_context import (
get_schedule,
get_spec,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -78,7 +81,7 @@ class XPUAttentionBackend(AttentionBackend):
isinstance(model_runner.token_to_kv_pool, SWAKVPool) isinstance(model_runner.token_to_kv_pool, SWAKVPool)
and model_runner.token_to_kv_pool.swa_layer_nums > 0 and model_runner.token_to_kv_pool.swa_layer_nums > 0
) )
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = get_spec().speculative_eagle_topk or 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_step_id = speculative_step_id self.speculative_step_id = speculative_step_id
@@ -16,7 +16,11 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.runtime_context import get_parallel, get_schedule from sglang.srt.runtime_context import (
get_disagg,
get_parallel,
get_schedule,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ceil_align, ceil_div, get_available_gpu_memory, is_musa from sglang.srt.utils import ceil_align, ceil_div, get_available_gpu_memory, is_musa
@@ -472,7 +476,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
# In PD, prefill-only nodes never decode (indexer would OOM at large # In PD, prefill-only nodes never decode (indexer would OOM at large
# bs) and decode-only nodes never extend. # bs) and decode-only nodes never extend.
disagg_mode = model_runner.server_args.disaggregation_mode disagg_mode = get_disagg().disaggregation_mode
run_decode = model_runner.is_generation and disagg_mode != "prefill" run_decode = model_runner.is_generation and disagg_mode != "prefill"
run_extend = disagg_mode != "decode" run_extend = disagg_mode != "decode"
@@ -13,7 +13,10 @@ from typing import TYPE_CHECKING, Optional
import torch import torch
from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import (
get_parallel,
get_schedule,
)
from sglang.srt.utils import get_compiler_backend from sglang.srt.utils import get_compiler_backend
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -88,7 +91,7 @@ def create_kt_config_from_server_args(
cpuinfer_threads=server_args.kt_cpuinfer, cpuinfer_threads=server_args.kt_cpuinfer,
threadpool_count=server_args.kt_threadpool_count, threadpool_count=server_args.kt_threadpool_count,
weight_path=server_args.kt_weight_path, weight_path=server_args.kt_weight_path,
chunked_prefill_size=server_args.chunked_prefill_size, chunked_prefill_size=get_schedule().chunked_prefill_size,
method=server_args.kt_method, method=server_args.kt_method,
max_deferred_experts_per_token=server_args.kt_max_deferred_experts_per_token, max_deferred_experts_per_token=server_args.kt_max_deferred_experts_per_token,
num_layers=num_layers, num_layers=num_layers,
@@ -31,7 +31,10 @@ from sglang.srt.layers.quantization.base_config import (
QuantizeMethodBase, QuantizeMethodBase,
) )
from sglang.srt.layers.utils import copy_or_rebind_param from sglang.srt.layers.utils import copy_or_rebind_param
from sglang.srt.runtime_context import get_exec, get_lora from sglang.srt.runtime_context import (
get_exec,
get_lora,
)
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,
@@ -100,7 +103,9 @@ def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
backend_str = server_args.bf16_gemm_backend backend_str = server_args.bf16_gemm_backend
if backend_str == "auto" and is_sm100_supported(): if backend_str == "auto" and is_sm100_supported():
backend_str = ( backend_str = (
"torch" if server_args.enable_deterministic_inference else "cutedsl" "torch"
if get_exec().deterministic.enable_deterministic_inference
else "cutedsl"
) )
backend = Bf16GemmBackend(backend_str) backend = Bf16GemmBackend(backend_str)
@@ -118,7 +123,7 @@ def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
_hopper_bf16_gemv = hopper_bf16_gemv _hopper_bf16_gemv = hopper_bf16_gemv
_use_hopper_bf16_gemv = use_hopper_bf16_gemv _use_hopper_bf16_gemv = use_hopper_bf16_gemv
elif backend.is_cutedsl(): elif backend.is_cutedsl():
if server_args.enable_deterministic_inference: if get_exec().deterministic.enable_deterministic_inference:
raise ValueError( raise ValueError(
"--bf16-gemm-backend cutedsl is batch-size dependent and cannot " "--bf16-gemm-backend cutedsl is batch-size dependent and cannot "
"be combined with --enable-deterministic-inference" "be combined with --enable-deterministic-inference"
@@ -20,8 +20,8 @@ def validate_experimental_sgl_marlin_server_args(
# A provided adapter path implicitly enables LoRA later unless it was # A provided adapter path implicitly enables LoRA later unless it was
# explicitly disabled. No-LoRA delegates to the stock Marlin fused path. # explicitly disabled. No-LoRA delegates to the stock Marlin fused path.
lora_enabled = bool(server_args.enable_lora) or ( lora_enabled = bool(resolved_args.enable_lora) or (
server_args.enable_lora is None and bool(server_args.lora_paths) resolved_args.enable_lora is None and bool(server_args.lora_paths)
) )
if not lora_enabled: if not lora_enabled:
return return
@@ -61,6 +61,9 @@ from sglang.srt.managers.load_snapshot import (
zmq_reader_owner, zmq_reader_owner,
) )
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
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 ( from sglang.srt.utils import (
configure_logger, configure_logger,
@@ -667,7 +670,7 @@ class TokenizerWorker(TokenizerManager):
# For PD disaggregation # For PD disaggregation
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
@@ -8,7 +8,10 @@ from typing import TYPE_CHECKING, NamedTuple, Optional
import torch import torch
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import (
get_parallel,
get_schedule,
)
from sglang.srt.utils import get_bool_env_var from sglang.srt.utils import get_bool_env_var
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -112,7 +115,7 @@ class PrefillDelayer:
# env flag is on (or overlap scheduling is disabled), ride the NCCL # env flag is on (or overlap scheduling is disabled), ride the NCCL
# device group on `device` instead of gloo on CPU. # device group on `device` instead of gloo on CPU.
use_nccl = ( use_nccl = (
server_args.disable_overlap_schedule get_schedule().disable_overlap_schedule
or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get() or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get()
) )
if use_nccl: if use_nccl:
@@ -140,7 +143,7 @@ class PrefillDelayer:
self.skip_first_delayer = True self.skip_first_delayer = True
assert ( assert (
not server_args.disable_overlap_schedule not get_schedule().disable_overlap_schedule
), "To use PrefillDelayer, disable_overlap_schedule must be False." ), "To use PrefillDelayer, disable_overlap_schedule must be False."
def _negotiate_should_allow_prefill( def _negotiate_should_allow_prefill(
+6 -4
View File
@@ -27,6 +27,10 @@ from sglang.srt.managers.utils import (
compute_num_reserved_tokens, compute_num_reserved_tokens,
msgpack_decode_explained, msgpack_decode_explained,
) )
from sglang.srt.runtime_context import (
get_mm,
get_serving,
)
from sglang.srt.utils.flatten import ( from sglang.srt.utils.flatten import (
FlatPairColumns, FlatPairColumns,
NestedRowColumns, NestedRowColumns,
@@ -216,9 +220,7 @@ class NativeMmHost:
# `--mm-process-config {"image": {...}}`: only pixel-limit overrides are # `--mm-process-config {"image": {...}}`: only pixel-limit overrides are
# mirrored natively, anything else disables the pipeline. # mirrored natively, anything else disables the pipeline.
image_overrides = dict( image_overrides = dict((get_mm().mm_process_config or {}).get("image", {}))
(self.server_args.mm_process_config or {}).get("image", {})
)
if not set(image_overrides) <= {"min_pixels", "max_pixels"}: if not set(image_overrides) <= {"min_pixels", "max_pixels"}:
return None return None
@@ -386,7 +388,7 @@ class RustServer:
# Refuse rather than run: silently dropping it means generating with # Refuse rather than run: silently dropping it means generating with
# sampling the operator did not configure, and `/get_model_info` would go on # sampling the operator did not configure, and `/get_model_info` would go on
# advertising values no request ever receives. # advertising values no request ever receives.
if server_args.preferred_sampling_params: if get_serving().preferred_sampling_params:
raise ValueError( raise ValueError(
"SGLANG_RUST_SERVER does not yet apply --preferred-sampling-params " "SGLANG_RUST_SERVER does not yet apply --preferred-sampling-params "
"(the Python TokenizerManager merges it into every request; the rust " "(the Python TokenizerManager merges it into every request; the rust "
@@ -22,7 +22,13 @@ from sglang.srt.observability.metrics_collector import (
SchedulerStats, SchedulerStats,
compute_routing_key_stats, compute_routing_key_stats,
) )
from sglang.srt.runtime_context import get_context, get_observability, get_spec from sglang.srt.runtime_context import (
configured_pp_size,
get_context,
get_disagg,
get_observability,
get_spec,
)
from sglang.srt.utils.device_timer import DeviceTimer from sglang.srt.utils.device_timer import DeviceTimer
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
@@ -621,9 +627,8 @@ class SchedulerMetricsReporter:
msg += f"#optimistic-req: {num_optimistic}, " msg += f"#optimistic-req: {num_optimistic}, "
if ( if (
self.scheduler.server_args.language_only get_disagg().language_only
and self.scheduler.server_args.encoder_transfer_backend and get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
== "zmq_to_scheduler"
): ):
msg += ( msg += (
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, " f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
@@ -875,9 +880,8 @@ class SchedulerMetricsReporter:
msg += f"#retracted-req: {len(self.scheduler.disagg_decode_prealloc_queue.retracted_queue)}, " msg += f"#retracted-req: {len(self.scheduler.disagg_decode_prealloc_queue.retracted_queue)}, "
if ( if (
self.scheduler.server_args.language_only get_disagg().language_only
and self.scheduler.server_args.encoder_transfer_backend and get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
== "zmq_to_scheduler"
): ):
msg += ( msg += (
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, " f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
@@ -1113,7 +1117,7 @@ class SchedulerMetricsReporter:
active_lora_ids = set() active_lora_ids = set()
# For PP mode, check all running micro batches # For PP mode, check all running micro batches
if self.scheduler.server_args.pp_size > 1: if configured_pp_size() > 1:
for batch in self.scheduler.running_mbs: for batch in self.scheduler.running_mbs:
if batch and hasattr(batch, "reqs"): if batch and hasattr(batch, "reqs"):
for req in batch.reqs: for req in batch.reqs:
@@ -74,7 +74,11 @@ 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_parallel from sglang.srt.runtime_context import (
get_lora,
get_parallel,
get_spec,
)
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,
@@ -186,7 +190,7 @@ class TokenizerControlMixin:
self: TokenizerManager, obj: AddExternalCorpusReqInput self: TokenizerManager, obj: AddExternalCorpusReqInput
) -> AddExternalCorpusReqOutput: ) -> AddExternalCorpusReqOutput:
self.auto_create_handle_loop() self.auto_create_handle_loop()
if self.server_args.speculative_algorithm != "NGRAM": if get_spec().speculative_algorithm != "NGRAM":
return AddExternalCorpusReqOutput( return AddExternalCorpusReqOutput(
success=False, success=False,
message="Ngram speculative decoding is not enabled.", message="Ngram speculative decoding is not enabled.",
@@ -262,7 +266,7 @@ class TokenizerControlMixin:
self: TokenizerManager, corpus_id: str self: TokenizerManager, corpus_id: str
) -> RemoveExternalCorpusReqOutput: ) -> RemoveExternalCorpusReqOutput:
self.auto_create_handle_loop() self.auto_create_handle_loop()
if self.server_args.speculative_algorithm != "NGRAM": if get_spec().speculative_algorithm != "NGRAM":
return RemoveExternalCorpusReqOutput( return RemoveExternalCorpusReqOutput(
success=False, success=False,
message="Ngram speculative decoding is not enabled.", message="Ngram speculative decoding is not enabled.",
@@ -277,7 +281,7 @@ class TokenizerControlMixin:
self: TokenizerManager, self: TokenizerManager,
) -> ListExternalCorporaReqOutput: ) -> ListExternalCorporaReqOutput:
self.auto_create_handle_loop() self.auto_create_handle_loop()
if self.server_args.speculative_algorithm != "NGRAM": if get_spec().speculative_algorithm != "NGRAM":
return ListExternalCorporaReqOutput( return ListExternalCorporaReqOutput(
success=False, success=False,
message="Ngram speculative decoding is not enabled.", message="Ngram speculative decoding is not enabled.",
@@ -602,7 +606,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."
) )
@@ -680,7 +684,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."
) )
@@ -756,7 +760,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."
) )
@@ -121,7 +121,7 @@ def _register_legacy_hicache_draft(
host_to_device_ratio=primary_host_pool.logical_size / pool.size, host_to_device_ratio=primary_host_pool.logical_size / pool.size,
host_size=0, host_size=0,
page_size=page_size, page_size=page_size,
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend, allocator_type=server_args.hicache_storage_backend,
pool_label="draft", pool_label="draft",
) )
@@ -70,6 +70,7 @@ from sglang.srt.runtime_context import (
get_disagg, get_disagg,
get_exec, get_exec,
get_memory, get_memory,
get_mm,
get_parallel, get_parallel,
get_schedule, get_schedule,
get_spec, get_spec,
@@ -1253,7 +1254,7 @@ class KVCacheConfigurator:
token_to_kv_pool = NPUMiniMaxSparseKVPool( token_to_kv_pool = NPUMiniMaxSparseKVPool(
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),
@@ -1846,7 +1847,7 @@ class KVCacheConfigurator:
) )
mm_reservation_gb = mm_runtime_reservation_gb( mm_reservation_gb = mm_runtime_reservation_gb(
is_multimodal=self.model_config.is_multimodal, is_multimodal=self.model_config.is_multimodal,
mm_feature_transport=self.server_args.mm_feature_transport, mm_feature_transport=get_mm().mm_feature_transport,
) )
rest_memory = available_gpu_memory - slack_gb - mm_reservation_gb rest_memory = available_gpu_memory - slack_gb - mm_reservation_gb
if self.mambaish_config is not None: if self.mambaish_config is not None:
@@ -2219,16 +2220,16 @@ def calculate_mla_kv_cache_dim(
# since it is not compatible for trtllm and other mla attn backend due to the different # since it is not compatible for trtllm and other mla attn backend due to the different
# kv cache layout. # kv cache layout.
if ( if (
server_args.dsa_prefill_backend == "trtllm" get_exec().kernel.dsa_prefill_backend == "trtllm"
or server_args.dsa_decode_backend == "trtllm" or get_exec().kernel.dsa_decode_backend == "trtllm"
): ):
return kv_cache_dim return kv_cache_dim
# On HIP, TileLang and AITER DSA kernels consume the raw MLA KV layout: # On HIP, TileLang and AITER DSA kernels consume the raw MLA KV layout:
# nope(512 fp8) + rope(64 fp8), without extra per-block scales. # nope(512 fp8) + rope(64 fp8), without extra per-block scales.
if _is_hip and ( if _is_hip and (
server_args.dsa_prefill_backend in ("tilelang", "aiter") get_exec().kernel.dsa_prefill_backend in ("tilelang", "aiter")
or server_args.dsa_decode_backend in ("tilelang", "aiter") or get_exec().kernel.dsa_decode_backend in ("tilelang", "aiter")
): ):
return kv_cache_dim return kv_cache_dim
+1 -1
View File
@@ -214,7 +214,7 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache:
source = "default" source = "default"
if ( if (
ctx.server_args.enable_hierarchical_cache get_memory().enable_hierarchical_cache
and ctx.server_args.hicache_host_memory_mode == "buffer_only" and ctx.server_args.hicache_host_memory_mode == "buffer_only"
): ):
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
@@ -39,7 +39,14 @@ from sglang.srt.model_executor.forward_batch_info import (
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode
from sglang.srt.runtime_context import get_exec, get_flags, get_parallel, get_spec from sglang.srt.runtime_context import (
configured_pp_size,
get_exec,
get_flags,
get_lora,
get_parallel,
get_spec,
)
from sglang.srt.utils import ( from sglang.srt.utils import (
empty_context, empty_context,
log_info_on_rank0, log_info_on_rank0,
@@ -585,22 +592,20 @@ class CPUGraphRunner:
self.enable_two_batch_overlap = ( self.enable_two_batch_overlap = (
model_runner.server_args.enable_two_batch_overlap model_runner.server_args.enable_two_batch_overlap
) )
self.speculative_algorithm = model_runner.server_args.speculative_algorithm self.speculative_algorithm = get_spec().speculative_algorithm
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph model_runner.server_args.enable_profile_cuda_graph
) )
self.tp_size = model_runner.server_args.tp_size self.tp_size = model_runner.server_args.tp_size
self.dp_size = get_parallel().dp_size self.dp_size = get_parallel().dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = configured_pp_size()
self.capture_forward_mode = ForwardMode.DECODE self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = self.return_hidden_states_mode self.capture_hidden_mode = self.return_hidden_states_mode
# Static capture width: CPU graphs are decode-only. # Static capture width: CPU graphs are decode-only.
self.captured_req_width = 1 self.captured_req_width = 1
assert ( assert not get_lora().enable_lora, "CPUGraphRunner does not support LoRA yet."
not self.model_runner.server_args.enable_lora
), "CPUGraphRunner does not support LoRA yet."
assert ( assert (
not self.enable_two_batch_overlap not self.enable_two_batch_overlap
), "CPUGraphRunner does not support two batch overlap yet." ), "CPUGraphRunner does not support two batch overlap yet."
@@ -991,7 +996,7 @@ class CPUGraphRunner:
retrieve_next_sibling=None, retrieve_next_sibling=None,
retrieve_cum_len=None, retrieve_cum_len=None,
spec_steps=get_spec().speculative_num_steps, spec_steps=get_spec().speculative_num_steps,
topk=self.model_runner.server_args.speculative_eagle_topk, topk=get_spec().speculative_eagle_topk,
draft_token_num=get_spec().speculative_num_draft_tokens, draft_token_num=get_spec().speculative_num_draft_tokens,
capture_hidden_mode=CaptureHiddenMode.FULL, capture_hidden_mode=CaptureHiddenMode.FULL,
seq_lens_sum=None, seq_lens_sum=None,
@@ -51,7 +51,11 @@ from sglang.srt.layers.dp_attention import (
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import ( from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
ForwardBatchDeepSeekMHAMixin, ForwardBatchDeepSeekMHAMixin,
) )
from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.runtime_context import (
get_exec,
get_lora,
get_parallel,
)
from sglang.srt.utils import ( from sglang.srt.utils import (
is_cpu, is_cpu,
is_cuda, is_cuda,
@@ -922,7 +926,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
if model_runner.lora_manager is not None: if model_runner.lora_manager is not None:
# In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool # In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool
# as a batch, right before running the batch # as a batch, right before running the batch
if not model_runner.server_args.enable_lora_overlap_loading: if not get_lora().enable_lora_overlap_loading:
model_runner.lora_manager.fetch_new_loras(set(ret.lora_ids)) model_runner.lora_manager.fetch_new_loras(set(ret.lora_ids))
model_runner.lora_manager.prepare_lora_batch(ret) model_runner.lora_manager.prepare_lora_batch(ret)
@@ -1320,7 +1324,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# graph; larger prefills fall back to eager and keep the # graph; larger prefills fall back to eager and keep the
# memory-efficient SUM_LEN. global_num_tokens is identical across ranks # memory-efficient SUM_LEN. global_num_tokens is identical across ranks
# (all-gathered), so the decision is consistent cluster-wide. # (all-gathered), so the decision is consistent cluster-wide.
prefill_cg = model_runner.server_args.cuda_graph_config.prefill prefill_cg = get_exec().graph.cuda_graph_config.prefill
if ( if (
self.can_run_dp_breakable_cuda_graph self.can_run_dp_breakable_cuda_graph
and self.is_extend_in_batch and self.is_extend_in_batch
@@ -5,6 +5,9 @@ from typing import TYPE_CHECKING, Dict, Optional
import torch import torch
from sglang.srt.model_executor.cuda_graph_config import Backend from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.runtime_context import (
get_exec,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
@@ -29,7 +32,7 @@ class GraphSharedOutput:
def create_for_model_runner( def create_for_model_runner(
cls, model_runner: ModelRunner cls, model_runner: ModelRunner
) -> Optional[GraphSharedOutput]: ) -> Optional[GraphSharedOutput]:
cuda_graph_config = model_runner.server_args.cuda_graph_config cuda_graph_config = get_exec().graph.cuda_graph_config
if cuda_graph_config is None: if cuda_graph_config is None:
return None return None
@@ -41,7 +41,13 @@ from sglang.srt.model_executor.runner import (
) )
from sglang.srt.model_loader.utils import resolve_language_model from sglang.srt.model_loader.utils import resolve_language_model
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_flags from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_flags,
get_schedule,
get_spec,
)
from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0 from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -111,16 +117,15 @@ def capture_cuda_graphs(
if model_runner.is_draft_worker: if model_runner.is_draft_worker:
moe_runner_backend = ( moe_runner_backend = (
model_runner.server_args.speculative_moe_runner_backend get_spec().speculative_moe_runner_backend
or model_runner.server_args.moe_runner_backend or get_exec().moe.moe_runner_backend
) )
moe_a2a_backend = ( moe_a2a_backend = (
model_runner.server_args.speculative_moe_a2a_backend get_spec().speculative_moe_a2a_backend or get_exec().moe.moe_a2a_backend
or model_runner.server_args.moe_a2a_backend
) )
else: else:
moe_runner_backend = model_runner.server_args.moe_runner_backend moe_runner_backend = get_exec().moe.moe_runner_backend
moe_a2a_backend = model_runner.server_args.moe_a2a_backend moe_a2a_backend = get_exec().moe.moe_a2a_backend
uses_deep_gemm_moe_runner = moe_runner_backend == "deep_gemm" uses_deep_gemm_moe_runner = moe_runner_backend == "deep_gemm"
if moe_runner_backend == "auto" and model_runner.model_config.quantization in ( if moe_runner_backend == "auto" and model_runner.model_config.quantization in (
@@ -200,7 +205,7 @@ def capture_cuda_graphs(
prealloc_symmetric_memory_pool( prealloc_symmetric_memory_pool(
is_draft_worker=model_runner.is_draft_worker, is_draft_worker=model_runner.is_draft_worker,
enable_symm_mem=model_runner.server_args.enable_symm_mem, enable_symm_mem=get_exec().comm.enable_symm_mem,
device=model_runner.device, device=model_runner.device,
forward_stream=model_runner.forward_stream, forward_stream=model_runner.forward_stream,
) )
@@ -292,19 +297,17 @@ def capture_prefill_graph(
return result(None) return result(None)
# Disable prefill CUDA graph for non capture size # Disable prefill CUDA graph for non capture size
if not model_runner.server_args.cuda_graph_config.prefill.bs: if not get_exec().graph.cuda_graph_config.prefill.bs:
logger.warning("Disable prefill CUDA graph because the capture size is not set") logger.warning("Disable prefill CUDA graph because the capture size is not set")
return result(None) return result(None)
prefill_config = model_runner.server_args.cuda_graph_config.prefill prefill_config = get_exec().graph.cuda_graph_config.prefill
prefill_backend = prefill_config.backend prefill_backend = prefill_config.backend
context_length = model_runner.model_config.context_len context_length = model_runner.model_config.context_len
if prefill_backend == Backend.FULL: if prefill_backend == Backend.FULL:
max_capture_requests = prefill_config.full_prefill_max_req max_capture_requests = prefill_config.full_prefill_max_req
if max_capture_requests is None: if max_capture_requests is None:
max_capture_requests = max( max_capture_requests = max(get_schedule().chunked_prefill_size // 512, 1)
model_runner.server_args.chunked_prefill_size // 512, 1
)
max_capture_requests = min( max_capture_requests = min(
max_capture_requests, model_runner.req_to_token_pool.size max_capture_requests, model_runner.req_to_token_pool.size
) )
@@ -418,7 +421,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
if ( if (
model_runner.spec_algorithm.is_speculative() model_runner.spec_algorithm.is_speculative()
and not model_runner.is_draft_worker and not model_runner.is_draft_worker
and model_runner.server_args.disaggregation_mode == "prefill" and get_disagg().disaggregation_mode == "prefill"
): ):
return no_capture return no_capture
if not model_runner.is_generation: if not model_runner.is_generation:
@@ -453,7 +456,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
capture_name = f"{role} decode" capture_name = f"{role} decode"
num_tokens_per_req = 1 num_tokens_per_req = 1
capture_bs, _ = get_batch_sizes_to_capture(model_runner, num_tokens_per_req) capture_bs, _ = get_batch_sizes_to_capture(model_runner, num_tokens_per_req)
decode_backend = model_runner.server_args.cuda_graph_config.decode.backend decode_backend = get_exec().graph.cuda_graph_config.decode.backend
logger.info( logger.info(
f"Capture {capture_name} {graph_backend[model_runner.device]} begin. " f"Capture {capture_name} {graph_backend[model_runner.device]} begin. "
f"backend={decode_backend}, num_tokens_per_req={num_tokens_per_req}, " f"backend={decode_backend}, num_tokens_per_req={num_tokens_per_req}, "
@@ -11,7 +11,12 @@ from sglang.srt.distributed import get_world_group
from sglang.srt.mem_cache.kv_cache_configurator import mm_runtime_reservation_gb from sglang.srt.mem_cache.kv_cache_configurator import mm_runtime_reservation_gb
from sglang.srt.model_executor.cuda_graph_config import Backend from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import pre_capture_activation_reserve_mb from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_mm,
pre_capture_activation_reserve_mb,
)
from sglang.srt.utils.common import get_available_gpu_memory, get_device_memory_capacity from sglang.srt.utils.common import get_available_gpu_memory, get_device_memory_capacity
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -55,11 +60,11 @@ def compute_post_capture_kv_resize(
headroom_gb = model_runner.pre_model_load_memory * ( headroom_gb = model_runner.pre_model_load_memory * (
1 - model_runner.mem_fraction_static 1 - model_runner.mem_fraction_static
) )
decode_cuda_graph_config = model_runner.server_args.cuda_graph_config.decode decode_cuda_graph_config = get_exec().graph.cuda_graph_config.decode
decode_max_bs = int(decode_cuda_graph_config.max_bs or 0) decode_max_bs = int(decode_cuda_graph_config.max_bs or 0)
running_requests = int(model_runner.max_running_requests or decode_max_bs or 1) running_requests = int(model_runner.max_running_requests or decode_max_bs or 1)
eager_decode_gap = ( eager_decode_gap = (
model_runner.server_args.disaggregation_mode != "prefill" get_disagg().disaggregation_mode != "prefill"
and decode_cuda_graph_config.backend != Backend.DISABLED and decode_cuda_graph_config.backend != Backend.DISABLED
and decode_max_bs < running_requests and decode_max_bs < running_requests
) )
@@ -80,7 +85,7 @@ def compute_post_capture_kv_resize(
) )
mm_reservation_gb = mm_runtime_reservation_gb( mm_reservation_gb = mm_runtime_reservation_gb(
is_multimodal=model_runner.model_config.is_multimodal, is_multimodal=model_runner.model_config.is_multimodal,
mm_feature_transport=model_runner.server_args.mm_feature_transport, mm_feature_transport=get_mm().mm_feature_transport,
) )
budget_bytes = ( budget_bytes = (
int(max(0.0, free_gb - headroom_gb - mm_reservation_gb) * (1 << 30)) int(max(0.0, free_gb - headroom_gb - mm_reservation_gb) * (1 << 30))
@@ -451,7 +451,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
self._swa_layers_num > 0 self._swa_layers_num > 0
), "Hybrid SWA model must have at least one SWA layer" ), "Hybrid SWA model must have at least one SWA layer"
self._swa_full_tokens_ratio = kvc.server_args.swa_full_tokens_ratio self._swa_full_tokens_ratio = get_schedule().swa_full_tokens_ratio
self._sliding_window_size = kvc.sliding_window_size self._sliding_window_size = kvc.sliding_window_size
self._page_size = kvc.page_size self._page_size = kvc.page_size
@@ -779,19 +779,19 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
f"local={len(self.compression_ratios)}/{len(cfg.compress_ratios)}" f"local={len(self.compression_ratios)}/{len(cfg.compress_ratios)}"
) )
self.swa_page_size = cfg.window_size self.swa_page_size = cfg.window_size
self.swa_ratio = kvc.server_args.swa_full_tokens_ratio self.swa_ratio = get_schedule().swa_full_tokens_ratio
self.is_speculative = kvc.server_args.speculative_algorithm is not None self.is_speculative = get_spec().speculative_algorithm is not None
self.online_c128_mtp_max_draft_tokens = ( self.online_c128_mtp_max_draft_tokens = (
kvc.server_args.max_speculative_num_draft_tokens or 0 kvc.server_args.max_speculative_num_draft_tokens or 0
) )
self.requested_max_running_requests_per_worker = ( self.requested_max_running_requests_per_worker = (
kvc.server_args.max_running_requests // kvc.ps.attn_dp_size get_schedule().max_running_requests // kvc.ps.attn_dp_size
if kvc.server_args.max_running_requests is not None if get_schedule().max_running_requests is not None
else None else None
) )
self.disaggregation_mode = kvc.server_args.disaggregation_mode self.disaggregation_mode = get_disagg().disaggregation_mode
self.disaggregation_decode_extra_slots = ( self.disaggregation_decode_extra_slots = (
kvc.server_args.disaggregation_decode_extra_slots or 0 get_disagg().disaggregation_decode_extra_slots or 0
) )
if kvc.server_args.enable_hisparse: if kvc.server_args.enable_hisparse:
from sglang.srt.mem_cache.sparsity import parse_hisparse_config from sglang.srt.mem_cache.sparsity import parse_hisparse_config
@@ -46,7 +46,13 @@ from sglang.srt.model_executor.runner.flashinfer_autotune import (
run_flashinfer_autotune_forward, run_flashinfer_autotune_forward,
should_run_flashinfer_autotune, should_run_flashinfer_autotune,
) )
from sglang.srt.runtime_context import get_flags, get_parallel from sglang.srt.runtime_context import (
configured_pp_size,
get_disagg,
get_exec,
get_flags,
get_parallel,
)
from sglang.srt.speculative.spec_info import create_dummy_verify_input from sglang.srt.speculative.spec_info import create_dummy_verify_input
from sglang.srt.utils import ( from sglang.srt.utils import (
empty_context, empty_context,
@@ -214,7 +220,7 @@ class BaseRunner(ABC):
self.tp_size = model_runner.server_args.tp_size self.tp_size = model_runner.server_args.tp_size
# elastic-EP scale-up rewrites dp_size on the published config # elastic-EP scale-up rewrites dp_size on the published config
self.dp_size = get_parallel().dp_size self.dp_size = get_parallel().dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = configured_pp_size()
self.enable_pdmux = model_runner.server_args.enable_pdmux self.enable_pdmux = model_runner.server_args.enable_pdmux
self.return_hidden_states_mode = ( self.return_hidden_states_mode = (
CaptureHiddenMode.NULL CaptureHiddenMode.NULL
@@ -265,7 +271,7 @@ class BaseRunner(ABC):
with custom_all_reduce.register_graph_buffers). with custom_all_reduce.register_graph_buffers).
""" """
mr = self.model_runner mr = self.model_runner
if mr.server_args.flashinfer_allreduce_fusion_backend is None: if get_exec().comm.flashinfer_allreduce_fusion_backend is None:
return return
from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE
@@ -344,7 +350,7 @@ class BaseRunner(ABC):
vocab_size=mr.model_config.vocab_size, vocab_size=mr.model_config.vocab_size,
dtype=mr.model_config.dtype, dtype=mr.model_config.dtype,
dp_size=get_parallel().dp_size, dp_size=get_parallel().dp_size,
pp_size=mr.server_args.pp_size, pp_size=configured_pp_size(),
is_encoder_decoder=mr.model_config.is_encoder_decoder, is_encoder_decoder=mr.model_config.is_encoder_decoder,
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args), require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args),
seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(), seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(),
@@ -410,7 +416,7 @@ class BaseRunner(ABC):
# TARGET_VERIFY dummy forward would trip the linear-attn backend's # TARGET_VERIFY dummy forward would trip the linear-attn backend's
# pool-type assert. Warm up in plain DECODE instead. # pool-type assert. Warm up in plain DECODE instead.
_is_pd_prefill_target = ( _is_pd_prefill_target = (
mr.server_args.disaggregation_mode == "prefill" and not mr.is_draft_worker get_disagg().disaggregation_mode == "prefill" and not mr.is_draft_worker
) )
if mr.spec_algorithm.is_speculative() and not _is_pd_prefill_target: if mr.spec_algorithm.is_speculative() and not _is_pd_prefill_target:
if mr.is_draft_worker: if mr.is_draft_worker:
@@ -514,7 +520,7 @@ class BaseRunner(ABC):
extend_prefix_lens = None extend_prefix_lens = None
extend_start_loc = None extend_start_loc = None
if mr.server_args.pp_size > 1: if configured_pp_size() > 1:
# PP0 already cp-split hidden_states before send. # PP0 already cp-split hidden_states before send.
pp_hidden_tokens = num_tokens pp_hidden_tokens = num_tokens
if ( if (
@@ -638,7 +644,7 @@ class BaseRunner(ABC):
kwargs = {} kwargs = {}
if ( if (
mr.server_args.pp_size > 1 configured_pp_size() > 1
and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters
): ):
kwargs["pp_proxy_tensors"] = PPProxyTensors( kwargs["pp_proxy_tensors"] = PPProxyTensors(
@@ -99,7 +99,12 @@ from sglang.srt.model_executor.runner_utils.deepep_adapter import (
) )
from sglang.srt.model_executor.runner_utils.shared_read_event import make_external_event from sglang.srt.model_executor.runner_utils.shared_read_event import make_external_event
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
from sglang.srt.runtime_context import get_flags, get_parallel, get_spec from sglang.srt.runtime_context import (
get_exec,
get_flags,
get_parallel,
get_spec,
)
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
from sglang.srt.utils import ( from sglang.srt.utils import (
empty_context, empty_context,
@@ -233,7 +238,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.require_mlp_tp_gather or self.require_attn_tp_gather self.require_mlp_tp_gather or self.require_attn_tp_gather
) )
self.require_mlp_sync = ( self.require_mlp_sync = (
model_runner.server_args.enable_dp_attention or self.require_gathered_buffer get_parallel().enable_dp_attention or self.require_gathered_buffer
) )
self.enable_two_batch_overlap = ( self.enable_two_batch_overlap = (
model_runner.server_args.enable_two_batch_overlap model_runner.server_args.enable_two_batch_overlap
@@ -243,7 +248,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
hf_config = model_runner.model_config.hf_config hf_config = model_runner.model_config.hf_config
self.ngram_embedding_n = hf_config.ngram_embedding_n self.ngram_embedding_n = hf_config.ngram_embedding_n
self.ngram_embedding_k = hf_config.ngram_embedding_k self.ngram_embedding_k = hf_config.ngram_embedding_k
self.speculative_algorithm = model_runner.server_args.speculative_algorithm self.speculative_algorithm = get_spec().speculative_algorithm
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph model_runner.server_args.enable_profile_cuda_graph
) )
@@ -1127,7 +1132,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size
# Sanity-check: --debug-cuda-graph requires breakable backend. # Sanity-check: --debug-cuda-graph requires breakable backend.
if self.model_runner.server_args.debug_cuda_graph: if get_exec().graph.debug_cuda_graph:
assert isinstance( assert isinstance(
self.backend, BreakableCudaGraphBackend self.backend, BreakableCudaGraphBackend
), "Breakable CUDA graph is required for --debug-cuda-graph" ), "Breakable CUDA graph is required for --debug-cuda-graph"
@@ -1477,7 +1482,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
retrieve_next_sibling=None, retrieve_next_sibling=None,
retrieve_cum_len=None, retrieve_cum_len=None,
spec_steps=self.speculative_num_steps, spec_steps=self.speculative_num_steps,
topk=self.model_runner.server_args.speculative_eagle_topk, topk=get_spec().speculative_eagle_topk,
draft_token_num=self.speculative_num_draft_tokens, draft_token_num=self.speculative_num_draft_tokens,
capture_hidden_mode=capture_mode, capture_hidden_mode=capture_mode,
seq_lens_sum=None, seq_lens_sum=None,
@@ -116,7 +116,12 @@ from sglang.srt.model_executor.runner_utils.buffers import (
PrefillInputBuffers, PrefillInputBuffers,
) )
from sglang.srt.model_loader.utils import resolve_language_model from sglang.srt.model_loader.utils import resolve_language_model
from sglang.srt.runtime_context import get_parallel, get_schedule from sglang.srt.runtime_context import (
get_exec,
get_memory,
get_parallel,
get_schedule,
)
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
from sglang.srt.utils import ( from sglang.srt.utils import (
get_available_gpu_memory, get_available_gpu_memory,
@@ -262,7 +267,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
self.capture_return_pooled_hidden_states = not model_runner.is_generation self.capture_return_pooled_hidden_states = not model_runner.is_generation
# --- prefill graph config ------------------------------------- # --- prefill graph config -------------------------------------
prefill_config = model_runner.server_args.cuda_graph_config.prefill prefill_config = get_exec().graph.cuda_graph_config.prefill
self.prefill_backend_name = prefill_config.backend self.prefill_backend_name = prefill_config.backend
# bs in prefill carries the captured shape (token count for # bs in prefill carries the captured shape (token count for
# tc_piecewise) — one shape knob per phase. # tc_piecewise) — one shape knob per phase.
@@ -424,7 +429,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
f"{type(attn_backend).__name__} does not support chunked-prefix " f"{type(attn_backend).__name__} does not support chunked-prefix "
"Full prefill CUDA graphs" "Full prefill CUDA graphs"
) )
prefix_config = model_runner.server_args.cuda_graph_config.prefill prefix_config = get_exec().graph.cuda_graph_config.prefill
( (
self._prefix_chunk_len, self._prefix_chunk_len,
self._prefix_chunk_capacity, self._prefix_chunk_capacity,
@@ -531,7 +536,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
def _is_mamba_track_enabled(self) -> bool: def _is_mamba_track_enabled(self) -> bool:
return self.model_runner.server_args.enable_mamba_extra_buffer() and ( return self.model_runner.server_args.enable_mamba_extra_buffer() and (
not self.model_runner.server_args.disable_radix_cache not get_memory().disable_radix_cache
) )
def _cache_loc_dtype(self): def _cache_loc_dtype(self):
@@ -802,10 +807,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
model_runner, capture_req_slots: int model_runner, capture_req_slots: int
) -> tuple[int, int]: ) -> tuple[int, int]:
"""Resolve per-request length and aggregate capacity of one chunk.""" """Resolve per-request length and aggregate capacity of one chunk."""
prefix_config = model_runner.server_args.cuda_graph_config.prefill prefix_config = get_exec().graph.cuda_graph_config.prefill
requested_capacity = prefix_config.full_prefill_prefix_chunk_tokens requested_capacity = prefix_config.full_prefill_prefix_chunk_tokens
if requested_capacity is None: if requested_capacity is None:
requested_capacity = model_runner.server_args.chunked_prefill_size requested_capacity = get_schedule().chunked_prefill_size
if requested_capacity is None or requested_capacity <= 0: if requested_capacity is None or requested_capacity <= 0:
requested_capacity = prefix_config.max_bs requested_capacity = prefix_config.max_bs
if requested_capacity is None or requested_capacity <= 0: if requested_capacity is None or requested_capacity <= 0:
@@ -49,7 +49,10 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
from sglang.srt.model_executor.runner_utils.pool import ( from sglang.srt.model_executor.runner_utils.pool import (
get_or_create_global_graph_memory_pool, get_or_create_global_graph_memory_pool,
) )
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import (
get_exec,
get_parallel,
)
from sglang.srt.utils import is_hip from sglang.srt.utils import is_hip
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -110,7 +113,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
def build_compilation_config(server_args: ServerArgs) -> CompilationConfig: def build_compilation_config(server_args: ServerArgs) -> CompilationConfig:
"""Construct a CompilationConfig from ServerArgs and """Construct a CompilationConfig from ServerArgs and
register the MoE A2A split-op when DeepEP / Mooncake is in use.""" register the MoE A2A split-op when DeepEP / Mooncake is in use."""
prefill = server_args.cuda_graph_config.prefill prefill = get_exec().graph.cuda_graph_config.prefill
num_tokens = prefill.bs num_tokens = prefill.bs
compiler = prefill.tc_compiler compiler = prefill.tc_compiler
assert num_tokens is not None, "cuda_graph_config[prefill].bs is not set" assert num_tokens is not None, "cuda_graph_config[prefill].bs is not set"
@@ -37,6 +37,9 @@ from sglang.srt.model_executor.runner_backend.full_cuda_graph_backend import (
from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import ( from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import (
TcPiecewiseCudaGraphBackend, TcPiecewiseCudaGraphBackend,
) )
from sglang.srt.runtime_context import (
get_exec,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.runner.base_cuda_graph_runner import ( from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
@@ -58,7 +61,7 @@ def resolve_decode_backend(
the Full-style backend is wired for NPU today). the Full-style backend is wired for NPU today).
""" """
model_runner = cuda_graph_runner.model_runner model_runner = cuda_graph_runner.model_runner
cfg = model_runner.server_args.cuda_graph_config cfg = get_exec().graph.cuda_graph_config
backend_name = cfg.decode.backend if cfg is not None else Backend.FULL backend_name = cfg.decode.backend if cfg is not None else Backend.FULL
enable_memory_saver = model_runner.server_args.enable_memory_saver enable_memory_saver = model_runner.server_args.enable_memory_saver
@@ -86,7 +89,7 @@ def resolve_decode_backend(
return BreakableCudaGraphBackend( return BreakableCudaGraphBackend(
cuda_graph_runner, cuda_graph_runner,
enable_memory_saver=enable_memory_saver, enable_memory_saver=enable_memory_saver,
debug_eager=model_runner.server_args.debug_cuda_graph, debug_eager=get_exec().graph.debug_cuda_graph,
) )
if backend_name == Backend.TC_PIECEWISE: if backend_name == Backend.TC_PIECEWISE:
global _TC_PIECEWISE_DECODE_FALLBACK_LOGGED global _TC_PIECEWISE_DECODE_FALLBACK_LOGGED
@@ -106,14 +109,14 @@ def resolve_prefill_backend(
) -> BaseCudaGraphBackend: ) -> BaseCudaGraphBackend:
"""Pick a backend instance from cuda_graph_config['prefill']['backend'].""" """Pick a backend instance from cuda_graph_config['prefill']['backend']."""
model_runner = cuda_graph_runner.model_runner model_runner = cuda_graph_runner.model_runner
cfg = model_runner.server_args.cuda_graph_config cfg = get_exec().graph.cuda_graph_config
backend_name = cfg.prefill.backend if cfg is not None else Backend.TC_PIECEWISE backend_name = cfg.prefill.backend if cfg is not None else Backend.TC_PIECEWISE
if backend_name == Backend.BREAKABLE: if backend_name == Backend.BREAKABLE:
return BreakableCudaGraphBackend( return BreakableCudaGraphBackend(
cuda_graph_runner, cuda_graph_runner,
enable_memory_saver=model_runner.server_args.enable_memory_saver, enable_memory_saver=model_runner.server_args.enable_memory_saver,
debug_eager=model_runner.server_args.debug_cuda_graph, debug_eager=get_exec().graph.debug_cuda_graph,
) )
if backend_name == Backend.FULL: if backend_name == Backend.FULL:
return FullCudaGraphBackend( return FullCudaGraphBackend(
@@ -34,7 +34,11 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
from sglang.srt.model_executor.runner_backend_utils import ( from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG, CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
from sglang.srt.runtime_context import get_flags, get_spec from sglang.srt.runtime_context import (
configured_pp_size,
get_flags,
get_spec,
)
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec
@@ -109,7 +113,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.ps.tp_size self.tp_size = model_runner.ps.tp_size
self.attn_dp_size = model_runner.ps.attn_dp_size self.attn_dp_size = model_runner.ps.attn_dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = configured_pp_size()
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
@@ -124,7 +128,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
if speculative_num_steps is None if speculative_num_steps is None
else speculative_num_steps else speculative_num_steps
) )
self.topk = model_runner.server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.draft_attn_backend = draft_attn_backend or model_runner.draft_attn_backend self.draft_attn_backend = draft_attn_backend or model_runner.draft_attn_backend
# Patch_model in parent's capture() needs an attn_backend reference. # Patch_model in parent's capture() needs an attn_backend reference.
@@ -35,7 +35,11 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
from sglang.srt.model_executor.runner_backend_utils import ( from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG, CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
from sglang.srt.runtime_context import get_flags, get_spec from sglang.srt.runtime_context import (
configured_pp_size,
get_flags,
get_spec,
)
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
@@ -95,7 +99,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.ps.tp_size self.tp_size = model_runner.ps.tp_size
self.attn_dp_size = model_runner.ps.attn_dp_size self.attn_dp_size = model_runner.ps.attn_dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = configured_pp_size()
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
@@ -32,7 +32,11 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
from sglang.srt.model_executor.runner_backend_utils import ( from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG, CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
from sglang.srt.runtime_context import get_flags, get_spec from sglang.srt.runtime_context import (
configured_pp_size,
get_flags,
get_spec,
)
from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -95,9 +99,9 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.tp_size = self.model_runner.ps.tp_size self.tp_size = self.model_runner.ps.tp_size
self.attn_dp_size = self.model_runner.ps.attn_dp_size self.attn_dp_size = self.model_runner.ps.attn_dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = configured_pp_size()
self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.topk = model_runner.server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph model_runner.server_args.enable_profile_cuda_graph
@@ -59,7 +59,12 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
from sglang.srt.model_executor.runner_backend_utils import ( from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG, CUDA_GRAPH_CAPTURE_FAILED_MSG,
) )
from sglang.srt.runtime_context import get_flags, get_parallel, get_spec from sglang.srt.runtime_context import (
configured_pp_size,
get_flags,
get_parallel,
get_spec,
)
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
from sglang.srt.speculative.multi_layer_eagle_utils import ( from sglang.srt.speculative.multi_layer_eagle_utils import (
@@ -151,7 +156,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.ps.tp_size self.tp_size = model_runner.ps.tp_size
self.dp_size = get_parallel().dp_size self.dp_size = get_parallel().dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = configured_pp_size()
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
@@ -161,7 +166,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.enable_pdmux = model_runner.server_args.enable_pdmux self.enable_pdmux = model_runner.server_args.enable_pdmux
self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.topk = model_runner.server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph model_runner.server_args.enable_profile_cuda_graph
) )
@@ -33,6 +33,7 @@ from sglang.srt.managers.schedule_batch import (
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
configured_tp_size, configured_tp_size,
get_mm,
get_parallel, get_parallel,
) )
from sglang.srt.utils.cuda_ipc_transport_utils import ( from sglang.srt.utils.cuda_ipc_transport_utils import (
@@ -933,7 +934,7 @@ class CudaVmmFeatureTransport:
def __init__(self, server_args, mm_processor) -> None: def __init__(self, server_args, mm_processor) -> None:
self.pool: CudaVmmMemoryPool | None = None self.pool: CudaVmmMemoryPool | None = None
if server_args.mm_feature_transport != "cuda_vmm": if get_mm().mm_feature_transport != "cuda_vmm":
return return
if mm_processor is None: if mm_processor is None:
raise RuntimeError( raise RuntimeError(
@@ -121,47 +121,55 @@ def test_draft_extend_in_graph_uses_captured_static_q_stride(monkeypatch):
def test_hybrid_wrappers_forward_in_graph_hook(): def test_hybrid_wrappers_forward_in_graph_hook():
"""Hybrid wrappers must forward init_forward_metadata_in_graph to the # The hybrid backend reads the mode from the published configuration.
wrapped backend(s) — the inherited no-op would leave the fused metadata from sglang.srt.runtime_context import get_context
rebuild out of the captured graph (stale page table on every replay)."""
from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend,
)
def make_fake(name, calls): override = get_context().override_server_args(speculative_attention_mode="decode")
return SimpleNamespace( override.install()
token_to_kv_pool=None, try:
req_to_token_pool=None, """Hybrid wrappers must forward init_forward_metadata_in_graph to the
needs_cpu_seq_lens=False, wrapped backend(s) — the inherited no-op would leave the fused metadata
init_forward_metadata_in_graph=lambda fb: calls.append(name), rebuild out of the captured graph (stale page table on every replay)."""
from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend,
) )
fb = SimpleNamespace(forward_mode=ForwardMode.DECODE) def make_fake(name, calls):
return SimpleNamespace(
token_to_kv_pool=None,
req_to_token_pool=None,
needs_cpu_seq_lens=False,
init_forward_metadata_in_graph=lambda fb: calls.append(name),
)
calls = [] fb = SimpleNamespace(forward_mode=ForwardMode.DECODE)
hybrid = HybridAttnBackend(
SimpleNamespace(
kv_cache_dtype=torch.bfloat16,
token_to_kv_pool=None,
req_to_token_pool=None,
server_args=SimpleNamespace(speculative_attention_mode="decode"),
model_config=SimpleNamespace(context_len=2048),
),
prefill_backend=make_fake("prefill", calls),
decode_backend=make_fake("decode", calls),
)
hybrid.init_forward_metadata_in_graph(fb)
assert calls == ["decode"]
calls = [] calls = []
hybrid_linear = HybridLinearAttnBackend( hybrid = HybridAttnBackend(
full_attn_backend=make_fake("full", calls), SimpleNamespace(
linear_attn_backend=make_fake("linear", calls), kv_cache_dtype=torch.bfloat16,
full_attn_layers=[0], token_to_kv_pool=None,
) req_to_token_pool=None,
hybrid_linear.init_forward_metadata_in_graph(fb) server_args=SimpleNamespace(speculative_attention_mode="decode"),
assert calls == ["full", "linear"] model_config=SimpleNamespace(context_len=2048),
),
prefill_backend=make_fake("prefill", calls),
decode_backend=make_fake("decode", calls),
)
hybrid.init_forward_metadata_in_graph(fb)
assert calls == ["decode"]
calls = []
hybrid_linear = HybridLinearAttnBackend(
full_attn_backend=make_fake("full", calls),
linear_attn_backend=make_fake("linear", calls),
full_attn_layers=[0],
)
hybrid_linear.init_forward_metadata_in_graph(fb)
assert calls == ["full", "linear"]
finally:
override.restore()
def test_metadata_update_records_inside_cuda_graph(): def test_metadata_update_records_inside_cuda_graph():
@@ -66,13 +66,24 @@ class _FakeHttpTokenizerManager:
pp_size=1, pp_size=1,
enable_dp_attention=False, enable_dp_attention=False,
): ):
from sglang.srt.runtime_context import get_context
self.loads = loads self.loads = loads
self.server_args = SimpleNamespace( # `tp_size` is raw input and still read off the record; the leaves
# resolution writes come from the bags.
self.server_args = SimpleNamespace(tp_size=tp_size)
# The accelerator arithmetic answers "what will this server do", so it
# reads the resolved topology out of the bags; publish the shape under test.
self._override = get_context().override_server_args(
tp_size=tp_size, tp_size=tp_size,
dp_size=dp_size, dp_size=dp_size,
pp_size=pp_size, pp_size=pp_size,
enable_dp_attention=enable_dp_attention, enable_dp_attention=enable_dp_attention,
) )
self._override.install()
def restore(self):
self._override.restore()
async def get_loads(self, include=None, dp_rank=None): async def get_loads(self, include=None, dp_rank=None):
results = [] results = []
@@ -95,6 +106,7 @@ class TestLoadsResponse(CustomTestCase):
) )
] ]
) )
self.addCleanup(manager.restore)
response = asyncio.run(get_loads(tokenizer_manager=manager)) response = asyncio.run(get_loads(tokenizer_manager=manager))
@@ -111,6 +123,7 @@ class TestLoadsAcceleratorField(CustomTestCase):
"""Guards the response contract: the JSON envelope carries an """Guards the response contract: the JSON envelope carries an
accelerator name and the accelerator count for each DP rank.""" accelerator name and the accelerator count for each DP rank."""
manager = _FakeHttpTokenizerManager([LoadSnapshot(dp_rank=0)], tp_size=16) manager = _FakeHttpTokenizerManager([LoadSnapshot(dp_rank=0)], tp_size=16)
self.addCleanup(manager.restore)
with mock.patch.object( with mock.patch.object(
v1_loads, "_accelerator_name", return_value="NVIDIA GB300" v1_loads, "_accelerator_name", return_value="NVIDIA GB300"
@@ -126,6 +139,7 @@ class TestLoadsAcceleratorField(CustomTestCase):
dp_size=8, dp_size=8,
enable_dp_attention=True, enable_dp_attention=True,
) )
self.addCleanup(manager.restore)
response = asyncio.run(get_loads(tokenizer_manager=manager)) response = asyncio.run(get_loads(tokenizer_manager=manager))
@@ -120,11 +120,11 @@ class TestLinearAttnBackends(CustomTestCase):
from sglang.srt.layers.attention.linear.gdn_backend import GDNAttnBackend from sglang.srt.layers.attention.linear.gdn_backend import GDNAttnBackend
# The draft-token width is a bag leaf read before the stamp.
self._publish(speculative_eagle_topk=0)
runner = SimpleNamespace( runner = SimpleNamespace(
device="cpu", device="cpu",
server_args=SimpleNamespace( server_args=SimpleNamespace(enable_unified_memory=False),
speculative_eagle_topk=0, enable_unified_memory=False
),
is_draft_worker=False, is_draft_worker=False,
req_to_token_pool=SimpleNamespace( req_to_token_pool=SimpleNamespace(
mamba_pool=SimpleNamespace( mamba_pool=SimpleNamespace(
@@ -1,3 +1,4 @@
import contextlib
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
@@ -9,6 +10,7 @@ from sglang.srt.layers.attention.verify_mask import (
maybe_create_verify_mask, maybe_create_verify_mask,
tree_mask_numel, tree_mask_numel,
) )
from sglang.srt.runtime_context import get_context
from sglang.srt.speculative.eagle_utils import TreeMaskMode, default_tree_mask_mode from sglang.srt.speculative.eagle_utils import TreeMaskMode, default_tree_mask_mode
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -123,6 +125,24 @@ def _mask(numel, **kwargs):
) )
@contextlib.contextmanager
def _published(speculative_attention_mode):
"""The backend reads the mode from the config bags, so publish one.
A stand-in on the model runner stopped being read when the mode became a
published leaf -- the record it would come from is not the one this process
resolved.
"""
override = get_context().override_server_args(
speculative_attention_mode=speculative_attention_mode
)
override.install()
try:
yield
finally:
override.restore()
def _make_hybrid_backend(speculative_attention_mode, prefill_mask, decode_mask): def _make_hybrid_backend(speculative_attention_mode, prefill_mask, decode_mask):
model_runner = SimpleNamespace( model_runner = SimpleNamespace(
kv_cache_dtype=None, kv_cache_dtype=None,
@@ -133,11 +153,12 @@ def _make_hybrid_backend(speculative_attention_mode, prefill_mask, decode_mask):
), ),
model_config=SimpleNamespace(context_len=_MAX_CONTEXT_LEN), model_config=SimpleNamespace(context_len=_MAX_CONTEXT_LEN),
) )
return HybridAttnBackend( with _published(speculative_attention_mode):
model_runner, return HybridAttnBackend(
prefill_backend=_FakeAttnBackend(prefill_mask), model_runner,
decode_backend=_FakeAttnBackend(decode_mask), prefill_backend=_FakeAttnBackend(prefill_mask),
) decode_backend=_FakeAttnBackend(decode_mask),
)
class TestHybridAttnBackendHandsOutSelectedChildMask(CustomTestCase): class TestHybridAttnBackendHandsOutSelectedChildMask(CustomTestCase):
@@ -40,7 +40,11 @@ def _validate_server(**overrides):
server_args.update(overrides) server_args.update(overrides)
return validate_experimental_sgl_marlin_server_args( return validate_experimental_sgl_marlin_server_args(
types.SimpleNamespace(**server_args), types.SimpleNamespace(**server_args),
types.SimpleNamespace(ep_size=4, moe_a2a_backend="none"), types.SimpleNamespace(
ep_size=4,
moe_a2a_backend="none",
enable_lora=server_args["enable_lora"],
),
) )
@@ -64,7 +64,13 @@ class TestDraftSidecarPoolDispatch(CustomTestCase):
page_size=512, page_size=512,
) )
) )
server_args = SimpleNamespace(hicache_mem_layout="page_first") # The layout comes from the published configuration.
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs
server_args = ServerArgs(model_path="dummy", hicache_mem_layout="page_first")
publish(server_args, role="scheduler")
self.addCleanup(reset_context)
with ( with (
patch( patch(
@@ -24,49 +24,60 @@ class _FakeBackend:
def test_split_full_attention_applies_model_wrapper_once(): def test_split_full_attention_applies_model_wrapper_once():
runner = SimpleNamespace( # The hybrid backend takes the speculative attention mode from the
server_args=SimpleNamespace(speculative_attention_mode="prefill"), # published configuration.
model_config=SimpleNamespace(context_len=2048), from sglang.srt.runtime_context import get_context
kv_cache_dtype=None,
token_to_kv_pool=object(),
req_to_token_pool=object(),
init_new_workspace=None,
)
wrapper_inputs = []
wrapped_backend = object()
def wrap_once(model_runner, backend): override = get_context().override_server_args(speculative_attention_mode="prefill")
assert model_runner is runner override.install()
wrapper_inputs.append(backend) try:
return wrapped_backend runner = SimpleNamespace(
server_args=SimpleNamespace(speculative_attention_mode="prefill"),
model_config=SimpleNamespace(context_len=2048),
kv_cache_dtype=None,
token_to_kv_pool=object(),
req_to_token_pool=object(),
init_new_workspace=None,
)
wrapper_inputs = []
wrapped_backend = object()
constructors = { def wrap_once(model_runner, backend):
"decode-test": lambda model_runner: _FakeBackend("decode"), assert model_runner is runner
"prefill-test": lambda model_runner: _FakeBackend("prefill"), wrapper_inputs.append(backend)
} return wrapped_backend
resolved = ResolvedAttentionBackendStr(decode="decode-test", prefill="prefill-test")
with ( constructors = {
patch.dict(attention_backend_setup.ATTENTION_BACKENDS, constructors), "decode-test": lambda model_runner: _FakeBackend("decode"),
patch.object( "prefill-test": lambda model_runner: _FakeBackend("prefill"),
attention_backend_setup, }
"attn_backend_wrapper", resolved = ResolvedAttentionBackendStr(
side_effect=wrap_once, decode="decode-test", prefill="prefill-test"
),
):
result = attention_backend_setup._build_resolved_backend(
model_runner=runner,
resolved=resolved,
init_new_workspace=True,
) )
assert result is wrapped_backend with (
assert len(wrapper_inputs) == 1 patch.dict(attention_backend_setup.ATTENTION_BACKENDS, constructors),
split_backend = wrapper_inputs[0] patch.object(
assert isinstance(split_backend, HybridAttnBackend) attention_backend_setup,
assert split_backend.decode_backend.name == "decode" "attn_backend_wrapper",
assert split_backend.prefill_backend.name == "prefill" side_effect=wrap_once,
assert runner.init_new_workspace is True ),
):
result = attention_backend_setup._build_resolved_backend(
model_runner=runner,
resolved=resolved,
init_new_workspace=True,
)
assert result is wrapped_backend
assert len(wrapper_inputs) == 1
split_backend = wrapper_inputs[0]
assert isinstance(split_backend, HybridAttnBackend)
assert split_backend.decode_backend.name == "decode"
assert split_backend.prefill_backend.name == "prefill"
assert runner.init_new_workspace is True
finally:
override.restore()
if __name__ == "__main__": if __name__ == "__main__":
@@ -13,6 +13,15 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def test_model_runner_can_override_decode_graph_runner(monkeypatch): def test_model_runner_can_override_decode_graph_runner(monkeypatch):
from sglang.srt.runtime_context import get_context
# The capture decision reads the graph configuration and the MoE backends
# out of the bags.
override = get_context().override_server_args(
cuda_graph_config=SimpleNamespace(decode=SimpleNamespace(backend="default")),
)
override.install()
class CustomGraphRunner: class CustomGraphRunner:
def __init__(self, model_runner): def __init__(self, model_runner):
self.model_runner = model_runner self.model_runner = model_runner
@@ -23,12 +32,7 @@ def test_model_runner_can_override_decode_graph_runner(monkeypatch):
gpu_id = 0 gpu_id = 0
is_draft_worker = False is_draft_worker = False
spec_algorithm = SimpleNamespace(is_speculative=lambda: False) spec_algorithm = SimpleNamespace(is_speculative=lambda: False)
server_args = SimpleNamespace( server_args = SimpleNamespace(model_impl="auto")
model_impl="auto",
cuda_graph_config=SimpleNamespace(
decode=SimpleNamespace(backend="default")
),
)
def _decode_cuda_graph_runner_cls(self): def _decode_cuda_graph_runner_cls(self):
return CustomGraphRunner return CustomGraphRunner
@@ -43,10 +47,13 @@ def test_model_runner_can_override_decode_graph_runner(monkeypatch):
cuda_graph_setup.current_platform, "is_out_of_tree", lambda: False cuda_graph_setup.current_platform, "is_out_of_tree", lambda: False
) )
capture = capture_decode_graph(model_runner=model_runner) try:
capture = capture_decode_graph(model_runner=model_runner)
assert isinstance(capture.runner, CustomGraphRunner) assert isinstance(capture.runner, CustomGraphRunner)
assert capture.runner.model_runner is model_runner assert capture.runner.model_runner is model_runner
finally:
override.restore()
if __name__ == "__main__": if __name__ == "__main__":
@@ -65,6 +65,16 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
def test_low_free_memory_still_captures_prefill_graph(self): def test_low_free_memory_still_captures_prefill_graph(self):
eager_runner = object() eager_runner = object()
prefill_runner = object() prefill_runner = object()
# The capture decision reads the graph configuration and the LoRA flag
# out of the bags.
override = get_context().override_server_args(
enable_lora=False,
cuda_graph_config=SimpleNamespace(
prefill=SimpleNamespace(bs=[1], backend=Backend.BREAKABLE)
),
)
override.install()
self.addCleanup(override.restore)
model_runner = SimpleNamespace( model_runner = SimpleNamespace(
device="cuda", device="cuda",
gpu_id=0, gpu_id=0,
@@ -73,11 +83,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
# reads it rather than the process-wide LoRA config. # reads it rather than the process-wide LoRA config.
lora_manager=None, lora_manager=None,
spec_algorithm=SimpleNamespace(is_eagle=lambda: False), spec_algorithm=SimpleNamespace(is_eagle=lambda: False),
server_args=SimpleNamespace( server_args=SimpleNamespace(),
cuda_graph_config=SimpleNamespace(
prefill=SimpleNamespace(bs=[1], backend=Backend.BREAKABLE)
),
),
model=SimpleNamespace(), model=SimpleNamespace(),
model_config=SimpleNamespace(context_len=8192, num_hidden_layers=1), model_config=SimpleNamespace(context_len=8192, num_hidden_layers=1),
req_to_token_pool=SimpleNamespace(size=1), req_to_token_pool=SimpleNamespace(size=1),
@@ -140,15 +146,18 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
self.assertIs(capture.runner, eager_runner) self.assertIs(capture.runner, eager_runner)
def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self): def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self):
graph_config = SimpleNamespace(
prefill=SimpleNamespace(full_prefill_prefix_chunk_tokens=None, max_bs=8)
)
# Both leaves come from the bags; the published object is this one, so
# the cases below still drive them by mutating it.
override = get_context().override_server_args(
chunked_prefill_size=16, cuda_graph_config=graph_config
)
published = override.install()
self.addCleanup(override.restore)
model_runner = SimpleNamespace( model_runner = SimpleNamespace(
server_args=SimpleNamespace( server_args=SimpleNamespace(),
chunked_prefill_size=16,
cuda_graph_config=SimpleNamespace(
prefill=SimpleNamespace(
full_prefill_prefix_chunk_tokens=None, max_bs=8
)
),
),
# Wider than the token table, so the table is the binding limit. # Wider than the token table, so the table is the binding limit.
model_config=SimpleNamespace(context_len=4096), model_config=SimpleNamespace(context_len=4096),
req_to_token_pool=SimpleNamespace( req_to_token_pool=SimpleNamespace(
@@ -161,24 +170,20 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
(4, 16), (4, 16),
) )
model_runner.server_args.chunked_prefill_size = -1 get_context().override("test", chunked_prefill_size=-1)
self.assertEqual( self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4), PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(2, 8), (2, 8),
) )
model_runner.server_args.chunked_prefill_size = 16 get_context().override("test", chunked_prefill_size=16)
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = ( graph_config.prefill.full_prefill_prefix_chunk_tokens = 24
24
)
self.assertEqual( self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4), PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(6, 24), (6, 24),
) )
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = ( graph_config.prefill.full_prefill_prefix_chunk_tokens = 256
256
)
self.assertEqual( self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4), PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(32, 128), (32, 128),
@@ -186,9 +191,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
# At least one token is reserved per request lane even if the requested # At least one token is reserved per request lane even if the requested
# aggregate capacity is smaller than the fixed request-slot count. # aggregate capacity is smaller than the fixed request-slot count.
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = ( graph_config.prefill.full_prefill_prefix_chunk_tokens = 2
2
)
self.assertEqual( self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4), PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(1, 4), (1, 4),
@@ -197,17 +200,13 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
# A context shorter than the token table binds instead: a draft runner # A context shorter than the token table binds instead: a draft runner
# capped at the target's context, or a short --context-length. # capped at the target's context, or a short --context-length.
model_runner.model_config.context_len = 8 model_runner.model_config.context_len = 8
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = ( graph_config.prefill.full_prefill_prefix_chunk_tokens = 256
256
)
self.assertEqual( self.assertEqual(
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4), PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4),
(8, 32), (8, 32),
) )
model_runner.server_args.cuda_graph_config.prefill.full_prefill_prefix_chunk_tokens = ( graph_config.prefill.full_prefill_prefix_chunk_tokens = 0
0
)
with self.assertRaisesRegex(ValueError, "must be positive"): with self.assertRaisesRegex(ValueError, "must be positive"):
PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4) PrefillCudaGraphRunner._resolve_prefix_chunk_shape(model_runner, 4)
@@ -101,7 +101,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
) )
# The consumer count comes from the published topology. # The consumer count comes from the published topology.
override = get_context().override_server_args( override = get_context().override_server_args(
enable_dp_attention=False, tp_size=4 enable_dp_attention=False, tp_size=4, mm_feature_transport="cuda_vmm"
) )
override.install() override.install()
self.addCleanup(override.restore) self.addCleanup(override.restore)
@@ -122,13 +122,16 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
) )
def test_disabled_transport_is_a_noop(self): def test_disabled_transport_is_a_noop(self):
from sglang.srt.runtime_context import get_context
from sglang.srt.utils.cuda_vmm_transport_utils import ( from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport, CudaVmmFeatureTransport,
) )
transport = CudaVmmFeatureTransport( # The transport choice is a bag leaf.
SimpleNamespace(mm_feature_transport="cpu"), None override = get_context().override_server_args(mm_feature_transport="cpu")
) override.install()
self.addCleanup(override.restore)
transport = CudaVmmFeatureTransport(SimpleNamespace(), None)
self.assertEqual(transport.prepare_for_dispatch([None]), []) self.assertEqual(transport.prepare_for_dispatch([None]), [])
transport.cancel_for_dispatch([]) transport.cancel_for_dispatch([])
@@ -136,14 +139,16 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
self.assertIsNone(transport.pool) self.assertIsNone(transport.pool)
def test_vmm_transport_requires_processor(self): def test_vmm_transport_requires_processor(self):
from sglang.srt.runtime_context import get_context
from sglang.srt.utils.cuda_vmm_transport_utils import ( from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport, CudaVmmFeatureTransport,
) )
override = get_context().override_server_args(mm_feature_transport="cuda_vmm")
override.install()
self.addCleanup(override.restore)
with self.assertRaisesRegex(RuntimeError, "multimodal processor"): with self.assertRaisesRegex(RuntimeError, "multimodal processor"):
CudaVmmFeatureTransport( CudaVmmFeatureTransport(SimpleNamespace(), None)
SimpleNamespace(mm_feature_transport="cuda_vmm"), None
)
def test_image_features_are_packed_per_request(self): def test_image_features_are_packed_per_request(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
@@ -1683,16 +1683,22 @@ class TestCudaGraphConfigDataclassAccess(CustomTestCase):
mock_backend = mock_get_moe_a2a_backend.return_value mock_backend = mock_get_moe_a2a_backend.return_value
mock_backend.is_deepep.return_value = False mock_backend.is_deepep.return_value = False
mock_backend.is_mooncake.return_value = False mock_backend.is_mooncake.return_value = False
server_args = SimpleNamespace( from sglang.srt.runtime_context import get_context
# The graph configuration is a bag leaf; the debug switch is raw input
# and stays on the argument.
override = get_context().override_server_args(
cuda_graph_config=CudaGraphConfig( cuda_graph_config=CudaGraphConfig(
prefill=PhaseConfig( prefill=PhaseConfig(
backend=Backend.TC_PIECEWISE, backend=Backend.TC_PIECEWISE,
bs=[32, 64], bs=[32, 64],
tc_compiler="eager", tc_compiler="eager",
) )
), )
enable_torch_compile_debug_mode=False,
) )
override.install()
self.addCleanup(override.restore)
server_args = SimpleNamespace(enable_torch_compile_debug_mode=False)
config = TcPiecewiseCudaGraphBackend.build_compilation_config(server_args) config = TcPiecewiseCudaGraphBackend.build_compilation_config(server_args)
@@ -2063,15 +2069,15 @@ class TestGrpcServerArgs(CustomTestCase):
def test_sidecar_builds_loopback_grpc_endpoints(self): def test_sidecar_builds_loopback_grpc_endpoints(self):
self.assertEqual( self.assertEqual(
build_sidecar_endpoint(SimpleNamespace(host="0.0.0.0", grpc_port=50051)), build_sidecar_endpoint("0.0.0.0", 50051),
"http://127.0.0.1:50051", "http://127.0.0.1:50051",
) )
self.assertEqual( self.assertEqual(
build_sidecar_endpoint(SimpleNamespace(host="::", grpc_port=50051)), build_sidecar_endpoint("::", 50051),
"http://[::1]:50051", "http://[::1]:50051",
) )
self.assertEqual( self.assertEqual(
build_sidecar_endpoint(SimpleNamespace(host="[::]", grpc_port=50051)), build_sidecar_endpoint("[::]", 50051),
"http://[::1]:50051", "http://[::1]:50051",
) )
@@ -2083,6 +2089,8 @@ class TestGrpcServerArgs(CustomTestCase):
self.assertEqual(parsed.sidecar_args, argv) self.assertEqual(parsed.sidecar_args, argv)
def test_start_sidecar_passes_endpoint_and_provider_argv_separately(self): def test_start_sidecar_passes_endpoint_and_provider_argv_separately(self):
from sglang.srt.runtime_context import get_context as get_context_for_config
server_args = SimpleNamespace( server_args = SimpleNamespace(
sidecar="example.sidecar", sidecar="example.sidecar",
sidecar_args=[ sidecar_args=[
@@ -2092,8 +2100,11 @@ class TestGrpcServerArgs(CustomTestCase):
"2", "2",
], ],
host="127.0.0.1", host="127.0.0.1",
grpc_port=50051,
) )
# The port the sidecar dials is the resolved one, off the bag.
override = get_context_for_config().override_server_args(grpc_port=50051)
override.install()
self.addCleanup(override.restore)
with ( with (
patch("sglang.srt.entrypoints.sidecar.mp.get_context") as get_context, patch("sglang.srt.entrypoints.sidecar.mp.get_context") as get_context,
patch("sglang.srt.entrypoints.sidecar.Sidecar") as sidecar_class, patch("sglang.srt.entrypoints.sidecar.Sidecar") as sidecar_class,
@@ -8,6 +8,7 @@ from types import SimpleNamespace
import torch import torch
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -236,6 +237,13 @@ class TestHybridNeedsCpuSeqLens(CustomTestCase):
req_to_token_pool=None, req_to_token_pool=None,
model_config=SimpleNamespace(context_len=2048), model_config=SimpleNamespace(context_len=2048),
) )
# The backend takes the mode from the published configuration, not from
# the runner it is handed.
override = get_context().override_server_args(
speculative_attention_mode=spec_mode
)
override.install()
self.addCleanup(override.restore)
return HybridAttnBackend(runner, backend(prefill_flag), backend(decode_flag)) return HybridAttnBackend(runner, backend(prefill_flag), backend(decode_flag))
def test_delegation(self): def test_delegation(self):
@@ -107,6 +107,40 @@ _CONFIGURED_SIZE_CALL_SITES = {
("srt/managers/scheduler.py", "configured_dcp_size"): ( ("srt/managers/scheduler.py", "configured_dcp_size"): (
"same pre-distributed-init arithmetic in configure_scheduler_process" "same pre-distributed-init arithmetic in configure_scheduler_process"
), ),
("srt/model_executor/runner/base_runner.py", "configured_pp_size"): (
"the runner's layer window is arithmetic over the configured stage "
"count; a draft runner shares the target's groups, so the live "
"property would answer for the wrong runner"
),
("srt/model_executor/cpu_graph_runner.py", "configured_pp_size"): (
"the same window, on the CPU graph path"
),
(
"srt/managers/scheduler_components/metrics_reporter.py",
"configured_pp_size",
): (
"the reporter labels its metrics with the stage count it was launched "
"with, which is configuration; the live group answers per process"
),
("srt/speculative/eagle_draft_cuda_graph_runner.py", "configured_pp_size"): (
"the draft runner's window over the target's stages: its own groups are "
"the target's, so the configured count is the one that describes it"
),
(
"srt/speculative/eagle_draft_extend_cuda_graph_runner.py",
"configured_pp_size",
): ("the same draft window, on the extend path"),
(
"srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py",
"configured_pp_size",
): ("the same draft window, multi-layer extend"),
("srt/speculative/frozen_kv_mtp_cuda_graph_runner.py", "configured_pp_size"): (
"the same draft window, frozen-KV MTP"
),
("srt/entrypoints/v1_loads.py", "configured_pp_size"): (
"the /v1/loads accelerator count is arithmetic over the launch shape, "
"reported from the tokenizer process, which holds no model groups"
),
("srt/disaggregation/common/conn.py", "configured_pp_size"): ( ("srt/disaggregation/common/conn.py", "configured_pp_size"): (
"the bootstrap connection is built by the KV manager on the transfer " "the bootstrap connection is built by the KV manager on the transfer "
"path, which the CPU-only conn tests exercise without ever starting " "path, which the CPU-only conn tests exercise without ever starting "
@@ -144,7 +144,6 @@ _EXPOSED = {
("layers/moe/utils.py", "moe_runner_backend"), ("layers/moe/utils.py", "moe_runner_backend"),
("layers/moe/utils.py", "quantization"), ("layers/moe/utils.py", "quantization"),
("layers/moe/utils.py", "speculative_moe_runner_backend"), ("layers/moe/utils.py", "speculative_moe_runner_backend"),
("entrypoints/sidecar.py", "grpc_port"),
("configs/embedding_model_spec.py", "chunked_prefill_size"), ("configs/embedding_model_spec.py", "chunked_prefill_size"),
("configs/embedding_model_spec.py", "cuda_graph_config"), ("configs/embedding_model_spec.py", "cuda_graph_config"),
("configs/embedding_model_spec.py", "disable_radix_cache"), ("configs/embedding_model_spec.py", "disable_radix_cache"),
@@ -159,13 +158,6 @@ _EXPOSED = {
("configs/model_config.py", "quantization"), ("configs/model_config.py", "quantization"),
("configs/model_config.py", "speculative_algorithm"), ("configs/model_config.py", "speculative_algorithm"),
("configs/model_config.py", "speculative_draft_model_quantization"), ("configs/model_config.py", "speculative_draft_model_quantization"),
("disaggregation/utils.py", "disaggregation_transfer_backend"),
("distributed/bootstrap.py", "disable_custom_all_reduce"),
("distributed/bootstrap.py", "enable_symm_mem"),
("distributed/bootstrap.py", "enable_torch_symm_mem"),
("distributed/bootstrap.py", "flashinfer_allreduce_fusion_backend"),
("distributed/bootstrap.py", "moe_a2a_backend"),
("distributed/bootstrap.py", "pre_warm_nccl"),
("entrypoints/engine.py", "attn_cp_size"), ("entrypoints/engine.py", "attn_cp_size"),
("entrypoints/engine.py", "enable_symm_mem"), ("entrypoints/engine.py", "enable_symm_mem"),
("entrypoints/engine.py", "moe_dp_size"), ("entrypoints/engine.py", "moe_dp_size"),
@@ -183,10 +175,7 @@ _EXPOSED = {
("layers/cp/bcg.py", "cp_strategy"), ("layers/cp/bcg.py", "cp_strategy"),
("layers/cp/bcg.py", "enable_prefill_cp"), ("layers/cp/bcg.py", "enable_prefill_cp"),
("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"), ("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"),
("layers/moe/kt_ep_wrapper.py", "chunked_prefill_size"),
("layers/quantization/unquant.py", "enable_deterministic_inference"),
("lora/lora_manager.py", "enable_lora_overlap_loading"), ("lora/lora_manager.py", "enable_lora_overlap_loading"),
("lora/marlin_lora_temp/policy.py", "enable_lora"),
("lora/marlin_lora_temp/policy.py", "lora_paths"), ("lora/marlin_lora_temp/policy.py", "lora_paths"),
("managers/data_parallel_controller.py", "attn_cp_size"), ("managers/data_parallel_controller.py", "attn_cp_size"),
("managers/data_parallel_controller.py", "disaggregation_mode"), ("managers/data_parallel_controller.py", "disaggregation_mode"),
@@ -194,13 +183,6 @@ _EXPOSED = {
("managers/data_parallel_controller.py", "moe_dp_size"), ("managers/data_parallel_controller.py", "moe_dp_size"),
("managers/data_parallel_controller.py", "pp_size"), ("managers/data_parallel_controller.py", "pp_size"),
("managers/data_parallel_controller.py", "soft_watchdog_timeout"), ("managers/data_parallel_controller.py", "soft_watchdog_timeout"),
("managers/prefill_delayer.py", "disable_overlap_schedule"),
("managers/rust_server.py", "mm_process_config"),
("mem_cache/kv_cache_builder.py", "hicache_mem_layout"),
(
"model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py",
"cuda_graph_config",
),
("parser/template_detection.py", "model_path"), ("parser/template_detection.py", "model_path"),
("speculative/adaptive_spec_params.py", "speculative_algorithm"), ("speculative/adaptive_spec_params.py", "speculative_algorithm"),
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"), ("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
@@ -208,7 +190,6 @@ _EXPOSED = {
("speculative/spec_info.py", "enable_multi_layer_eagle"), ("speculative/spec_info.py", "enable_multi_layer_eagle"),
("utils/common.py", "speculative_num_draft_tokens"), ("utils/common.py", "speculative_num_draft_tokens"),
("utils/common.py", "speculative_num_steps"), ("utils/common.py", "speculative_num_steps"),
("utils/cuda_vmm_transport_utils.py", "mm_feature_transport"),
("utils/hf_transformers/processor.py", "image_processor_backend"), ("utils/hf_transformers/processor.py", "image_processor_backend"),
} }