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.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw
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.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.server_args import PortArgs, ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -487,7 +487,7 @@ class TreeCacheNamespace(SimpleNamespace):
def extend(reqs, model_runner):
# Create dummy tree_cache for benchmarks (no prefix caching, just allocation)
dummy_tree_cache = TreeCacheNamespace(
page_size=model_runner.server_args.page_size,
page_size=get_schedule().page_size,
device=model_runner.device,
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(
batch,
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_cp_size=model_runner.ps.attn_cp_size,
tp_group=model_runner.tp_group,
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),
disable_overlap_schedule=model_runner.server_args.disable_overlap_schedule,
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
offload_tags=set(),
)
@@ -16,6 +16,10 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
import torch
from sglang.srt.runtime_context import (
get_schedule,
)
logger = logging.getLogger(__name__)
# Bounded wait for a watermark advance before re-enqueueing a deferred staging
@@ -548,7 +552,7 @@ class PrefillStagingStrategy:
self.staging_buffer = staging_buffer
page_size = kv_manager.kv_buffer_tensors["page_size"]
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
)
+11 -11
View File
@@ -94,7 +94,11 @@ from sglang.srt.observability.req_time_stats import (
set_schedule_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.network import NetworkAddress
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)
swa_len = ceil_align(swa_len, page_size)
swa_reserved = self.num_reserved_decode_tokens
if self.scheduler.server_args.disable_radix_cache:
if get_memory().disable_radix_cache:
swa_reserved = 0
return (
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),
)
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_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
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
if _is_fake_transfer(req, self.scheduler.server_args):
if _is_fake_transfer(req):
decode_req.kv_receiver.init(0)
return
@@ -663,9 +667,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
self, req: Req, is_rebootstrap: bool = False
) -> DecodeRequest:
backend = (
TransferBackend.FAKE
if _is_fake_transfer(req, self.scheduler.server_args)
else self.transfer_backend
TransferBackend.FAKE if _is_fake_transfer(req) else self.transfer_backend
)
kv_receiver_class = get_kv_class(backend, KVClassType.RECEIVER)
@@ -1455,7 +1457,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
if (
self.scheduler.enable_hisparse
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
# used by C4 indexer and C128 KV. These device buffers do not use
@@ -2059,7 +2061,7 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
else 0
)
if _is_fake_transfer(decode_req.req, self.scheduler.server_args):
if _is_fake_transfer(decode_req.req):
pass
elif actual_room == 0:
# Should never happen: _poll_with_metadata_gate already confirmed
@@ -2197,7 +2199,6 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
self.gloo_group,
decode_reqs=self.queue,
metadata_buffers=self.metadata_buffers,
server_args=self.scheduler.server_args,
)
def _poll_with_staging(self) -> list:
@@ -2206,7 +2207,6 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
self.staging_handler,
self.gloo_group,
metadata_buffers=self.metadata_buffers,
server_args=self.scheduler.server_args,
)
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 = (
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
req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
@@ -438,8 +438,7 @@ class PrefillBootstrapQueue:
failed_reqs.append(req)
elif poll == KVPoll.Bootstrapping:
if (
req.prefill_attempt_count
< self.scheduler.server_args.optimistic_prefill_attempts
req.prefill_attempt_count < get_disagg().optimistic_prefill_attempts
and not req.is_retracted # engine paused
):
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.disaggregation.base import KVPoll
from sglang.srt.environ import envs
from sglang.srt.runtime_context import (
get_disagg,
)
from sglang.srt.utils import is_hip, is_npu
if TYPE_CHECKING:
@@ -33,7 +36,6 @@ if TYPE_CHECKING:
CommonKVSender,
)
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.server_args import ServerArgs
if is_npu():
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]
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 (
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.
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):
if poll_val == int(KVPoll.Success):
decode_req = decode_reqs[i]
if _is_fake_transfer(decode_req.req, server_args):
if _is_fake_transfer(decode_req.req):
continue
actual_room = metadata_buffers.bootstrap_room[
decode_req.metadata_buffer_index, 0
@@ -205,18 +207,13 @@ def poll_and_all_reduce(
gloo_group: dist.ProcessGroup,
decode_reqs=None,
metadata_buffers: Optional[MetadataBuffers] = None,
server_args: Optional[ServerArgs] = None,
):
# at a certain prob, the poll is failed to simulate failure
polls = _poll_with_failure_injection(pollers)
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
if (
decode_reqs is not None
and metadata_buffers is not None
and server_args is not None
):
_apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args)
if decode_reqs is not None and metadata_buffers is not None:
_apply_metadata_gate(polls, decode_reqs, metadata_buffers)
return _all_reduce_polls(polls, gloo_group)
@@ -239,7 +236,6 @@ def poll_and_all_reduce_with_staging(
staging_handler,
gloo_group: dist.ProcessGroup,
metadata_buffers: Optional[MetadataBuffers] = None,
server_args: Optional[ServerArgs] = None,
):
"""Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce."""
for decode_req in decode_reqs:
@@ -265,8 +261,8 @@ def poll_and_all_reduce_with_staging(
):
raw_polls[i] = int(KVPoll.Transferring)
# 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:
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers, server_args)
if metadata_buffers is not None:
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers)
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.layers.dp_attention import initialize_dp_attention
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.utils import (
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
# 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
):
_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:
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_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(
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,
distributed_init_method=dist_init_method,
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,
max_world_size=server_args.max_ep_size,
)
@@ -260,7 +263,7 @@ def _init_parallel_groups(
and server_args.enable_two_batch_overlap
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,
rank_offset=rank_offset,
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 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
logger = logging.getLogger(__name__)
@@ -400,7 +404,7 @@ class RuntimeHandle:
result = {
"model_path": self.tokenizer_manager.model_path,
"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,
"weight_version": self.tokenizer_manager.config_value("weight_version"),
"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,
}
]
if self.tokenizer_manager.server_args.enable_lora and hasattr(
self.tokenizer_manager, "lora_registry"
):
if get_lora().enable_lora and hasattr(self.tokenizer_manager, "lora_registry"):
lora_registry = self.tokenizer_manager.lora_registry
for _, lora_ref in lora_registry.get_all_adapters().items():
models.append(
+5 -4
View File
@@ -394,7 +394,7 @@ async def lifespan(fast_api_app: FastAPI):
if (
getattr(fast_api_app, "is_single_tokenizer_mode", False)
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(
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 (
get_disagg,
get_exec,
get_lora,
get_model,
get_parallel,
get_serving,
@@ -744,9 +745,9 @@ async def model_info():
# Manager-owned, and moved by a weight update alongside `model_path`:
# this is where a client reads the identity the server answers under.
"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,
"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"
),
@@ -1857,7 +1858,7 @@ async def available_models():
)
# 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
for _, lora_ref in lora_registry.get_all_adapters().items():
model_cards.append(
+8 -7
View File
@@ -18,6 +18,9 @@ import logging
import multiprocessing as mp
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.network import NetworkAddress
from sglang.srt.utils.watchdog import SubprocessWatchdog
@@ -36,12 +39,10 @@ def _loopback_host(host: str) -> str:
return host
def build_sidecar_endpoint(server_args) -> str:
"""Both halves of the endpoint come from the argument: this is a helper
over a config object, callable before anything is published."""
return NetworkAddress(
_loopback_host(server_args.host), server_args.grpc_port
).to_url()
def build_sidecar_endpoint(host: str, grpc_port: int) -> str:
"""Both halves are passed in: this is a string helper, and the caller is
the one that knows where the effective values live."""
return NetworkAddress(_loopback_host(host), grpc_port).to_url()
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
assert module_name is not None
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(
name=f"sglang_sidecar_{module_name}",
target=_run_sidecar,
+7 -3
View File
@@ -26,6 +26,10 @@ from typing import Optional
from fastapi import APIRouter, Depends, HTTPException
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.version import __version__
@@ -144,9 +148,9 @@ async def get_loads(
"accelerator": _accelerator_name(),
"num_accelerators": _num_accelerators_per_dp_rank(
tokenizer_manager.server_args.tp_size,
tokenizer_manager.server_args.pp_size,
tokenizer_manager.server_args.dp_size,
tokenizer_manager.server_args.enable_dp_attention,
configured_pp_size(),
get_parallel().dp_size,
get_parallel().enable_dp_attention,
),
"loads": loads,
}
@@ -1,6 +1,10 @@
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
@@ -152,7 +156,7 @@ class AiterAttnBackend(AttentionBackend):
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)
@@ -2957,7 +2961,7 @@ class AiterMultiStepDraftBackend:
# Cached variables for generate_draft_decode_kv_indices
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.page_size = model_runner.server_args.page_size
self.page_size = get_schedule().page_size
def common_template(
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,
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
_is_musa = is_musa()
@@ -49,7 +52,7 @@ def create_flashinfer_backend(runner):
)
# Init streams
if runner.server_args.speculative_algorithm == "EAGLE":
if get_spec().speculative_algorithm == "EAGLE":
if (
not hasattr(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):
if not runner.use_mla_backend:
raise ValueError("trtllm_mla backend can only be used with MLA models.")
if (
get_parallel().dcp_enabled
and runner.server_args.speculative_algorithm is not None
):
if get_parallel().dcp_enabled and get_spec().speculative_algorithm is not None:
_, decode_backend = runner.server_args.get_attention_backends()
if decode_backend == "trtllm_mla":
raise ValueError(
@@ -265,7 +265,7 @@ def create_hpc_ops_backend(runner):
raise ValueError(
"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(
"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.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
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.ragged_verify import (
RaggedVerifyMode,
@@ -576,7 +579,7 @@ class DeepseekV4AttnBackend(
self._q8kv8_qpad_buf = None
self._q8kv8_attn_sink_pad = 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"
self.mtp_enabled = self.topk > 0
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.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.ragged_verify import resolve_ragged_verify_layout
from sglang.srt.utils import ceil_align
@@ -456,7 +459,7 @@ class DeepseekV4HipRadixBackend(
self.enable_deepseek_v4_fp4_indexer: bool = (
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"
self.mtp_enabled = self.topk > 0
self.speculative_num_steps = speculative_num_steps
@@ -15,7 +15,11 @@ from typing import (
import torch
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__)
from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
@@ -309,7 +313,7 @@ class DeepseekSparseAttnBackend(
assert isinstance(model_runner.page_size, int)
self.real_page_size = model_runner.page_size
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)
assert self.use_dsa, "DSA backend only supports DeepSeek DSA"
@@ -334,10 +338,8 @@ class DeepseekSparseAttnBackend(
self.use_mha: bool = False
self.supports_mha_one_shot: bool = True
self.dsa_prefill_impl: _DSA_IMPL_T = (
model_runner.server_args.dsa_prefill_backend
)
self.dsa_decode_impl: _DSA_IMPL_T = model_runner.server_args.dsa_decode_backend
self.dsa_prefill_impl: _DSA_IMPL_T = get_exec().kernel.dsa_prefill_backend
self.dsa_decode_impl: _DSA_IMPL_T = get_exec().kernel.dsa_decode_backend
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend(
model_runner.server_args.dsa_topk_backend
)
@@ -390,7 +392,7 @@ class DeepseekSparseAttnBackend(
)
# 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_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_step_id = speculative_step_id
@@ -1,6 +1,9 @@
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.
@@ -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
# More information can be found here: https://github.com/flashinfer-ai/flashinfer/pull/1675
self.enable_deterministic = (
model_runner.server_args.enable_deterministic_inference
get_exec().deterministic.enable_deterministic_inference
)
self.prefill_split_tile_size = None
self.decode_split_tile_size = None
@@ -1,6 +1,11 @@
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.
@@ -1101,7 +1106,7 @@ class FlashInferMLAMultiStepDraftBackend:
# Cached variables for generate_draft_decode_kv_indices
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.page_size = model_runner.server_args.page_size
self.page_size = get_schedule().page_size
def common_template(
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.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import (
get_spec,
)
if TYPE_CHECKING:
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.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token_pool = model_runner.req_to_token_pool
self.spec_attn_is_decode = (
model_runner.server_args.speculative_attention_mode == "decode"
)
self.spec_attn_is_prefill = (
model_runner.server_args.speculative_attention_mode == "prefill"
)
self.spec_attn_is_decode = get_spec().speculative_attention_mode == "decode"
self.spec_attn_is_prefill = get_spec().speculative_attention_mode == "prefill"
# 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:
# 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):
self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens)
if (
self.model_runner.server_args.speculative_algorithm is not None
and self.spec_attn_is_prefill
):
if get_spec().speculative_algorithm is not None and self.spec_attn_is_prefill:
# When speculative decoding is enabled, we need to initialize the backend
# that will be used for target_verify.
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 (
get_exec,
get_memory,
get_spec,
mamba_cache_chunk_size,
)
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
@@ -47,7 +48,7 @@ class MambaAttnBackendBase(AttentionBackend):
super().__init__()
self.pad_slot_id = PAD_SLOT_ID
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.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_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.model_runner import ModelRunner
from sglang.srt.runtime_context import (
get_spec,
)
class KDAKernelDispatcher:
@@ -384,7 +387,7 @@ class KDAAttnBackend(MambaAttnBackendBase):
# traversal). Reject EAGLE tree verify (topk > 1) early at setup, keyed on the
# verify backend (not decode). The kernel keeps a per-call
# 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:
raise ValueError(
"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,
)
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 (
draft_kv_indices_buffer_width,
draft_kv_indices_used_len,
@@ -254,7 +259,7 @@ class TritonAttnBackend(AttentionBackend):
self.static_kv_splits = get_bool_env_var(
"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:
self.max_kv_splits = _mla_decode_kv_splits_cap(
self.max_kv_splits,
@@ -280,11 +285,11 @@ class TritonAttnBackend(AttentionBackend):
cuda_graph_fully_disabled()
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 = (
model_runner.server_args.enable_deterministic_inference
get_exec().deterministic.enable_deterministic_inference
)
if self.enable_deterministic:
@@ -1987,7 +1992,7 @@ class TritonMultiStepDraftBackend:
# Cached variables for generate_draft_decode_kv_indices
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.page_size = model_runner.server_args.page_size
self.page_size = get_schedule().page_size
def common_template(
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.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
if TYPE_CHECKING:
@@ -109,7 +113,7 @@ class WaveAttnBackend(AttentionBackend):
self.static_kv_splits = get_bool_env_var(
"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.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.swa_memory_pool import SWAKVPool
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:
from sglang.srt.layers.radix_attention import RadixAttention
@@ -78,7 +81,7 @@ class XPUAttentionBackend(AttentionBackend):
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
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_draft_tokens = get_spec().speculative_num_draft_tokens
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.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM
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.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
# 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_extend = disagg_mode != "decode"
@@ -13,7 +13,10 @@ from typing import TYPE_CHECKING, Optional
import torch
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
if TYPE_CHECKING:
@@ -88,7 +91,7 @@ def create_kt_config_from_server_args(
cpuinfer_threads=server_args.kt_cpuinfer,
threadpool_count=server_args.kt_threadpool_count,
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,
max_deferred_experts_per_token=server_args.kt_max_deferred_experts_per_token,
num_layers=num_layers,
@@ -31,7 +31,10 @@ from sglang.srt.layers.quantization.base_config import (
QuantizeMethodBase,
)
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 (
cpu_has_amx_support,
get_bool_env_var,
@@ -100,7 +103,9 @@ def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
backend_str = server_args.bf16_gemm_backend
if backend_str == "auto" and is_sm100_supported():
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)
@@ -118,7 +123,7 @@ def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
_hopper_bf16_gemv = hopper_bf16_gemv
_use_hopper_bf16_gemv = use_hopper_bf16_gemv
elif backend.is_cutedsl():
if server_args.enable_deterministic_inference:
if get_exec().deterministic.enable_deterministic_inference:
raise ValueError(
"--bf16-gemm-backend cutedsl is batch-size dependent and cannot "
"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
# explicitly disabled. No-LoRA delegates to the stock Marlin fused path.
lora_enabled = bool(server_args.enable_lora) or (
server_args.enable_lora is None and bool(server_args.lora_paths)
lora_enabled = bool(resolved_args.enable_lora) or (
resolved_args.enable_lora is None and bool(server_args.lora_paths)
)
if not lora_enabled:
return
@@ -61,6 +61,9 @@ from sglang.srt.managers.load_snapshot import (
zmq_reader_owner,
)
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.utils import (
configure_logger,
@@ -667,7 +670,7 @@ class TokenizerWorker(TokenizerManager):
# For PD disaggregation
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
@@ -8,7 +8,10 @@ from typing import TYPE_CHECKING, NamedTuple, Optional
import torch
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
if TYPE_CHECKING:
@@ -112,7 +115,7 @@ class PrefillDelayer:
# env flag is on (or overlap scheduling is disabled), ride the NCCL
# device group on `device` instead of gloo on CPU.
use_nccl = (
server_args.disable_overlap_schedule
get_schedule().disable_overlap_schedule
or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get()
)
if use_nccl:
@@ -140,7 +143,7 @@ class PrefillDelayer:
self.skip_first_delayer = True
assert (
not server_args.disable_overlap_schedule
not get_schedule().disable_overlap_schedule
), "To use PrefillDelayer, disable_overlap_schedule must be False."
def _negotiate_should_allow_prefill(
+6 -4
View File
@@ -27,6 +27,10 @@ from sglang.srt.managers.utils import (
compute_num_reserved_tokens,
msgpack_decode_explained,
)
from sglang.srt.runtime_context import (
get_mm,
get_serving,
)
from sglang.srt.utils.flatten import (
FlatPairColumns,
NestedRowColumns,
@@ -216,9 +220,7 @@ class NativeMmHost:
# `--mm-process-config {"image": {...}}`: only pixel-limit overrides are
# mirrored natively, anything else disables the pipeline.
image_overrides = dict(
(self.server_args.mm_process_config or {}).get("image", {})
)
image_overrides = dict((get_mm().mm_process_config or {}).get("image", {}))
if not set(image_overrides) <= {"min_pixels", "max_pixels"}:
return None
@@ -386,7 +388,7 @@ class RustServer:
# Refuse rather than run: silently dropping it means generating with
# sampling the operator did not configure, and `/get_model_info` would go on
# advertising values no request ever receives.
if server_args.preferred_sampling_params:
if get_serving().preferred_sampling_params:
raise ValueError(
"SGLANG_RUST_SERVER does not yet apply --preferred-sampling-params "
"(the Python TokenizerManager merges it into every request; the rust "
@@ -22,7 +22,13 @@ from sglang.srt.observability.metrics_collector import (
SchedulerStats,
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.scheduler_status_logger import SchedulerStatusLogger
@@ -621,9 +627,8 @@ class SchedulerMetricsReporter:
msg += f"#optimistic-req: {num_optimistic}, "
if (
self.scheduler.server_args.language_only
and self.scheduler.server_args.encoder_transfer_backend
== "zmq_to_scheduler"
get_disagg().language_only
and get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
):
msg += (
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)}, "
if (
self.scheduler.server_args.language_only
and self.scheduler.server_args.encoder_transfer_backend
== "zmq_to_scheduler"
get_disagg().language_only
and get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
):
msg += (
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
@@ -1113,7 +1117,7 @@ class SchedulerMetricsReporter:
active_lora_ids = set()
# 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:
if batch and hasattr(batch, "reqs"):
for req in batch.reqs:
@@ -74,7 +74,11 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromTensorReqOutput,
)
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.utils import (
get_bool_env_var,
@@ -186,7 +190,7 @@ class TokenizerControlMixin:
self: TokenizerManager, obj: AddExternalCorpusReqInput
) -> AddExternalCorpusReqOutput:
self.auto_create_handle_loop()
if self.server_args.speculative_algorithm != "NGRAM":
if get_spec().speculative_algorithm != "NGRAM":
return AddExternalCorpusReqOutput(
success=False,
message="Ngram speculative decoding is not enabled.",
@@ -262,7 +266,7 @@ class TokenizerControlMixin:
self: TokenizerManager, corpus_id: str
) -> RemoveExternalCorpusReqOutput:
self.auto_create_handle_loop()
if self.server_args.speculative_algorithm != "NGRAM":
if get_spec().speculative_algorithm != "NGRAM":
return RemoveExternalCorpusReqOutput(
success=False,
message="Ngram speculative decoding is not enabled.",
@@ -277,7 +281,7 @@ class TokenizerControlMixin:
self: TokenizerManager,
) -> ListExternalCorporaReqOutput:
self.auto_create_handle_loop()
if self.server_args.speculative_algorithm != "NGRAM":
if get_spec().speculative_algorithm != "NGRAM":
return ListExternalCorporaReqOutput(
success=False,
message="Ngram speculative decoding is not enabled.",
@@ -602,7 +606,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop()
try:
if not self.server_args.enable_lora:
if not get_lora().enable_lora:
raise ValueError(
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
)
@@ -680,7 +684,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop()
try:
if not self.server_args.enable_lora:
if not get_lora().enable_lora:
raise ValueError(
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
)
@@ -756,7 +760,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop()
try:
if not self.server_args.enable_lora:
if not get_lora().enable_lora:
raise ValueError(
"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_size=0,
page_size=page_size,
layout=server_args.hicache_mem_layout,
layout=get_memory().hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
pool_label="draft",
)
@@ -70,6 +70,7 @@ from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_memory,
get_mm,
get_parallel,
get_schedule,
get_spec,
@@ -1253,7 +1254,7 @@ class KVCacheConfigurator:
token_to_kv_pool = NPUMiniMaxSparseKVPool(
size=max_total_num_tokens,
page_size=self.server_args.page_size,
page_size=get_schedule().page_size,
dtype=self.kv_cache_dtype,
index_dtype=self.model_dtype,
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(
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
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
# kv cache layout.
if (
server_args.dsa_prefill_backend == "trtllm"
or server_args.dsa_decode_backend == "trtllm"
get_exec().kernel.dsa_prefill_backend == "trtllm"
or get_exec().kernel.dsa_decode_backend == "trtllm"
):
return kv_cache_dim
# On HIP, TileLang and AITER DSA kernels consume the raw MLA KV layout:
# nope(512 fp8) + rope(64 fp8), without extra per-block scales.
if _is_hip and (
server_args.dsa_prefill_backend in ("tilelang", "aiter")
or server_args.dsa_decode_backend in ("tilelang", "aiter")
get_exec().kernel.dsa_prefill_backend in ("tilelang", "aiter")
or get_exec().kernel.dsa_decode_backend in ("tilelang", "aiter")
):
return kv_cache_dim
+1 -1
View File
@@ -214,7 +214,7 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache:
source = "default"
if (
ctx.server_args.enable_hierarchical_cache
get_memory().enable_hierarchical_cache
and ctx.server_args.hicache_host_memory_mode == "buffer_only"
):
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.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 (
empty_context,
log_info_on_rank0,
@@ -585,22 +592,20 @@ class CPUGraphRunner:
self.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 = (
model_runner.server_args.enable_profile_cuda_graph
)
self.tp_size = model_runner.server_args.tp_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_hidden_mode = self.return_hidden_states_mode
# Static capture width: CPU graphs are decode-only.
self.captured_req_width = 1
assert (
not self.model_runner.server_args.enable_lora
), "CPUGraphRunner does not support LoRA yet."
assert not get_lora().enable_lora, "CPUGraphRunner does not support LoRA yet."
assert (
not self.enable_two_batch_overlap
), "CPUGraphRunner does not support two batch overlap yet."
@@ -991,7 +996,7 @@ class CPUGraphRunner:
retrieve_next_sibling=None,
retrieve_cum_len=None,
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,
capture_hidden_mode=CaptureHiddenMode.FULL,
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 (
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 (
is_cpu,
is_cuda,
@@ -922,7 +926,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
if model_runner.lora_manager is not None:
# In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool
# 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.prepare_lora_batch(ret)
@@ -1320,7 +1324,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# graph; larger prefills fall back to eager and keep the
# memory-efficient SUM_LEN. global_num_tokens is identical across ranks
# (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 (
self.can_run_dp_breakable_cuda_graph
and self.is_extend_in_batch
@@ -5,6 +5,9 @@ from typing import TYPE_CHECKING, Dict, Optional
import torch
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.runtime_context import (
get_exec,
)
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
@@ -29,7 +32,7 @@ class GraphSharedOutput:
def create_for_model_runner(
cls, model_runner: ModelRunner
) -> 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:
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.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
if TYPE_CHECKING:
@@ -111,16 +117,15 @@ def capture_cuda_graphs(
if model_runner.is_draft_worker:
moe_runner_backend = (
model_runner.server_args.speculative_moe_runner_backend
or model_runner.server_args.moe_runner_backend
get_spec().speculative_moe_runner_backend
or get_exec().moe.moe_runner_backend
)
moe_a2a_backend = (
model_runner.server_args.speculative_moe_a2a_backend
or model_runner.server_args.moe_a2a_backend
get_spec().speculative_moe_a2a_backend or get_exec().moe.moe_a2a_backend
)
else:
moe_runner_backend = model_runner.server_args.moe_runner_backend
moe_a2a_backend = model_runner.server_args.moe_a2a_backend
moe_runner_backend = get_exec().moe.moe_runner_backend
moe_a2a_backend = get_exec().moe.moe_a2a_backend
uses_deep_gemm_moe_runner = moe_runner_backend == "deep_gemm"
if moe_runner_backend == "auto" and model_runner.model_config.quantization in (
@@ -200,7 +205,7 @@ def capture_cuda_graphs(
prealloc_symmetric_memory_pool(
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,
forward_stream=model_runner.forward_stream,
)
@@ -292,19 +297,17 @@ def capture_prefill_graph(
return result(None)
# 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")
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
context_length = model_runner.model_config.context_len
if prefill_backend == Backend.FULL:
max_capture_requests = prefill_config.full_prefill_max_req
if max_capture_requests is None:
max_capture_requests = max(
model_runner.server_args.chunked_prefill_size // 512, 1
)
max_capture_requests = max(get_schedule().chunked_prefill_size // 512, 1)
max_capture_requests = min(
max_capture_requests, model_runner.req_to_token_pool.size
)
@@ -418,7 +421,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
if (
model_runner.spec_algorithm.is_speculative()
and not model_runner.is_draft_worker
and model_runner.server_args.disaggregation_mode == "prefill"
and get_disagg().disaggregation_mode == "prefill"
):
return no_capture
if not model_runner.is_generation:
@@ -453,7 +456,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
capture_name = f"{role} decode"
num_tokens_per_req = 1
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(
f"Capture {capture_name} {graph_backend[model_runner.device]} begin. "
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.model_executor.cuda_graph_config import Backend
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
if TYPE_CHECKING:
@@ -55,11 +60,11 @@ def compute_post_capture_kv_resize(
headroom_gb = model_runner.pre_model_load_memory * (
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)
running_requests = int(model_runner.max_running_requests or decode_max_bs or 1)
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_max_bs < running_requests
)
@@ -80,7 +85,7 @@ def compute_post_capture_kv_resize(
)
mm_reservation_gb = mm_runtime_reservation_gb(
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 = (
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
), "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._page_size = kvc.page_size
@@ -779,19 +779,19 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
f"local={len(self.compression_ratios)}/{len(cfg.compress_ratios)}"
)
self.swa_page_size = cfg.window_size
self.swa_ratio = kvc.server_args.swa_full_tokens_ratio
self.is_speculative = kvc.server_args.speculative_algorithm is not None
self.swa_ratio = get_schedule().swa_full_tokens_ratio
self.is_speculative = get_spec().speculative_algorithm is not None
self.online_c128_mtp_max_draft_tokens = (
kvc.server_args.max_speculative_num_draft_tokens or 0
)
self.requested_max_running_requests_per_worker = (
kvc.server_args.max_running_requests // kvc.ps.attn_dp_size
if kvc.server_args.max_running_requests is not None
get_schedule().max_running_requests // kvc.ps.attn_dp_size
if get_schedule().max_running_requests is not None
else None
)
self.disaggregation_mode = kvc.server_args.disaggregation_mode
self.disaggregation_mode = get_disagg().disaggregation_mode
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:
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,
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.utils import (
empty_context,
@@ -214,7 +220,7 @@ class BaseRunner(ABC):
self.tp_size = model_runner.server_args.tp_size
# elastic-EP scale-up rewrites dp_size on the published config
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.return_hidden_states_mode = (
CaptureHiddenMode.NULL
@@ -265,7 +271,7 @@ class BaseRunner(ABC):
with custom_all_reduce.register_graph_buffers).
"""
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
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,
dtype=mr.model_config.dtype,
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,
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(),
@@ -410,7 +416,7 @@ class BaseRunner(ABC):
# TARGET_VERIFY dummy forward would trip the linear-attn backend's
# pool-type assert. Warm up in plain DECODE instead.
_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.is_draft_worker:
@@ -514,7 +520,7 @@ class BaseRunner(ABC):
extend_prefix_lens = 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.
pp_hidden_tokens = num_tokens
if (
@@ -638,7 +644,7 @@ class BaseRunner(ABC):
kwargs = {}
if (
mr.server_args.pp_size > 1
configured_pp_size() > 1
and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters
):
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.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.utils import (
empty_context,
@@ -233,7 +238,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.require_mlp_tp_gather or self.require_attn_tp_gather
)
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 = (
model_runner.server_args.enable_two_batch_overlap
@@ -243,7 +248,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
hf_config = model_runner.model_config.hf_config
self.ngram_embedding_n = hf_config.ngram_embedding_n
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 = (
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
# 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(
self.backend, BreakableCudaGraphBackend
), "Breakable CUDA graph is required for --debug-cuda-graph"
@@ -1477,7 +1482,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
retrieve_next_sibling=None,
retrieve_cum_len=None,
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,
capture_hidden_mode=capture_mode,
seq_lens_sum=None,
@@ -116,7 +116,12 @@ from sglang.srt.model_executor.runner_utils.buffers import (
PrefillInputBuffers,
)
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.utils import (
get_available_gpu_memory,
@@ -262,7 +267,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
self.capture_return_pooled_hidden_states = not model_runner.is_generation
# --- 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
# bs in prefill carries the captured shape (token count for
# tc_piecewise) — one shape knob per phase.
@@ -424,7 +429,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
f"{type(attn_backend).__name__} does not support chunked-prefix "
"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_capacity,
@@ -531,7 +536,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
def _is_mamba_track_enabled(self) -> bool:
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):
@@ -802,10 +807,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
model_runner, capture_req_slots: int
) -> tuple[int, int]:
"""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
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:
requested_capacity = prefix_config.max_bs
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 (
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
if TYPE_CHECKING:
@@ -110,7 +113,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
def build_compilation_config(server_args: ServerArgs) -> CompilationConfig:
"""Construct a CompilationConfig from ServerArgs and
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
compiler = prefill.tc_compiler
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 (
TcPiecewiseCudaGraphBackend,
)
from sglang.srt.runtime_context import (
get_exec,
)
if TYPE_CHECKING:
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).
"""
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
enable_memory_saver = model_runner.server_args.enable_memory_saver
@@ -86,7 +89,7 @@ def resolve_decode_backend(
return BreakableCudaGraphBackend(
cuda_graph_runner,
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:
global _TC_PIECEWISE_DECODE_FALLBACK_LOGGED
@@ -106,14 +109,14 @@ def resolve_prefill_backend(
) -> BaseCudaGraphBackend:
"""Pick a backend instance from cuda_graph_config['prefill']['backend']."""
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
if backend_name == Backend.BREAKABLE:
return BreakableCudaGraphBackend(
cuda_graph_runner,
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:
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 (
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.speculative.eagle_info import EagleDraftInput
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.tp_size = model_runner.ps.tp_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.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
@@ -124,7 +128,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
if speculative_num_steps is None
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
# 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 (
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_utils import get_draft_input_from_target_hidden_dim
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.tp_size = model_runner.ps.tp_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.disable_padding = model_runner.server_args.disable_cuda_graph_padding
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 (
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.spec_utils import resolve_num_tokens_per_req
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.tp_size = self.model_runner.ps.tp_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.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.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 (
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_utils import get_draft_input_from_target_hidden_dim
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.tp_size = model_runner.ps.tp_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.disable_padding = model_runner.server_args.disable_cuda_graph_padding
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.speculative_num_steps = get_spec().speculative_num_steps
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 = (
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 (
configured_tp_size,
get_mm,
get_parallel,
)
from sglang.srt.utils.cuda_ipc_transport_utils import (
@@ -933,7 +934,7 @@ class CudaVmmFeatureTransport:
def __init__(self, server_args, mm_processor) -> None:
self.pool: CudaVmmMemoryPool | None = None
if server_args.mm_feature_transport != "cuda_vmm":
if get_mm().mm_feature_transport != "cuda_vmm":
return
if mm_processor is None:
raise RuntimeError(