config: project the config bags from the resolution result (#35906)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-08-23 01:18:24 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 0e22777572
commit 4bc79a1b49
44 changed files with 1448 additions and 661 deletions
+73
View File
@@ -292,9 +292,61 @@ def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> Non
log = [] log = []
object.__setattr__(server_args, "_runtime_mutations", log) object.__setattr__(server_args, "_runtime_mutations", log)
log.append((source, dict(fields))) log.append((source, dict(fields)))
stash = getattr(server_args, "_resolved_overrides", None)
if stash is None:
stash = []
object.__setattr__(server_args, "_resolved_overrides", stash)
stash.append((source, dict(fields)))
_apply_fields(server_args, fields) _apply_fields(server_args, fields)
def declare_direct_writes(
server_args: Any, source: str, resolve: Callable[[Any], Any]
) -> Any:
"""Run a resolver that writes the fields directly, and declare what it moved.
Returns whatever the resolver returned, so a provider with a return value
can go through the same capture.
Out-of-tree platform plugins are handed the record and set fields on it.
Their implementations live outside this tree, so they cannot be converted
by editing the resolver; and the raw snapshot is taken before the pipeline
starts, so a plugin's default is neither declared nor raw.
Rebinding is what the diff sees, and rebinding is all it needs to see: a
plugin that mutates a value in place reaches the projection anyway, because
the raw snapshot and the stash entries hold the same object it mutated.
A stand-in record (tests drive the hooks with a plain namespace) has no
fields to diff and no projection to feed, so the resolver runs uncaptured.
"""
if not dataclasses.is_dataclass(server_args):
return resolve(server_args)
before = {
field.name: getattr(server_args, field.name)
for field in dataclasses.fields(server_args)
}
already = len(getattr(server_args, "_resolved_overrides", None) or ())
result = resolve(server_args)
stash = getattr(server_args, "_resolved_overrides", None)
if stash is None:
stash = []
object.__setattr__(server_args, "_resolved_overrides", stash)
# A resolver reached this way can also declare properly -- the in-tree
# implementations of these hooks do. Those fields are already explained, and
# recording them again would attribute them to the wrapper and bury an
# actual direct write among the echoes.
declared = {name for _source, fields in stash[already:] for name in fields}
changed = {
name: getattr(server_args, name)
for name, previous in before.items()
if name not in declared and getattr(server_args, name) is not previous
}
if changed:
stash.append((source, changed))
return result
def materialize_declarations(server_args: Any) -> None: def materialize_declarations(server_args: Any) -> None:
"""Apply the accumulated declarations onto ``server_args`` once, at the """Apply the accumulated declarations onto ``server_args`` once, at the
end of ``__post_init__`` (gate order: last writer wins). After this the end of ``__post_init__`` (gate order: last writer wins). After this the
@@ -307,6 +359,27 @@ def materialize_declarations(server_args: Any) -> None:
server_args._declarations_materialized = True server_args._declarations_materialized = True
def resolution_result(server_args: Any, field: str, default: Any = None) -> Any:
"""What resolution decided for ``field``: the declaration if there is one,
otherwise what the caller supplied.
This is what the config projection reads. Reading the field instead would
work only for as long as declarations materialize onto the record -- and
the point of declaring is that they will not, so the projection must not
depend on it. A config that never ran the pipeline (a mock, a partial
fixture) carries no raw snapshot; its fields are all it has.
"""
for _source, declared in reversed(
getattr(server_args, "_resolved_overrides", None) or ()
):
if field in declared:
return declared[field]
raw = getattr(server_args, "_raw_input", None)
if raw is not None and field in raw:
return raw[field]
return getattr(server_args, field, default)
def resolved_view(server_args: Any) -> ResolvedView: def resolved_view(server_args: Any) -> ResolvedView:
"""Read-only view of the resolving configuration for mid-resolution code """Read-only view of the resolving configuration for mid-resolution code
that is not a pass (``__post_init__`` handlers and hooks). Internal to that is not a pass (``__post_init__`` handlers and hooks). Internal to
@@ -5,7 +5,10 @@ import logging
import os import os
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING, Optional
from sglang.srt.arg_groups.overrides import declare_resolution from sglang.srt.arg_groups.overrides import (
declare_direct_writes,
declare_resolution,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -149,7 +152,11 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
# TODO: move the per-algorithm validation below into spec module hooks. # TODO: move the per-algorithm validation below into spec module hooks.
if isinstance(algo, CustomSpecAlgo) and algo.validate_server_args is not None: if isinstance(algo, CustomSpecAlgo) and algo.validate_server_args is not None:
algo.validate_server_args(server_args) declare_direct_writes(
server_args,
"handle_speculative_decoding.custom_validate",
algo.validate_server_args,
)
if server_args.speculative_skip_dp_mlp_sync: if server_args.speculative_skip_dp_mlp_sync:
assert server_args.speculative_algorithm == "EAGLE", ( assert server_args.speculative_algorithm == "EAGLE", (
@@ -163,7 +170,13 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
_init_adaptive_speculative_params(server_args) _init_adaptive_speculative_params(server_args)
if algo is not None: if algo is not None:
algo.handle_server_args(server_args) # A registered algorithm's callback lives outside this tree and sets
# fields on the record, so the writes are captured around the call.
declare_direct_writes(
server_args,
"handle_speculative_decoding.custom_algo",
algo.handle_server_args,
)
def _handle_dflash(server_args: ServerArgs) -> None: def _handle_dflash(server_args: ServerArgs) -> None:
@@ -36,7 +36,12 @@ from sglang.srt.layers.dp_attention import (
get_attention_dp_rank, get_attention_dp_rank,
get_attention_dp_size, get_attention_dp_size,
) )
from sglang.srt.runtime_context import get_parallel, get_serving from sglang.srt.runtime_context import (
configured_pp_size,
get_disagg,
get_parallel,
get_serving,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import ( from sglang.srt.utils.network import (
NetworkAddress, NetworkAddress,
@@ -164,9 +169,9 @@ class CommonKVManager(BaseKVManager):
envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get() envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get()
) )
# for p/d multi node infer # for p/d multi node infer
self.bootstrap_host = server_args.host self.bootstrap_host = get_serving().host
self.bootstrap_port = server_args.disaggregation_bootstrap_port self.bootstrap_port = get_disagg().disaggregation_bootstrap_port
self.dist_init_addr = server_args.dist_init_addr self.dist_init_addr = get_parallel().dist_init_addr
parallel = get_parallel() parallel = get_parallel()
self.attn_tp_size = parallel.attn_tp_size self.attn_tp_size = parallel.attn_tp_size
self.attn_tp_rank = parallel.attn_tp_rank self.attn_tp_rank = parallel.attn_tp_rank
@@ -182,7 +187,7 @@ class CommonKVManager(BaseKVManager):
self.system_dp_rank = ( self.system_dp_rank = (
self.kv_args.system_dp_rank if self.kv_args.system_dp_rank else 0 self.kv_args.system_dp_rank if self.kv_args.system_dp_rank else 0
) )
self.pp_size = server_args.pp_size self.pp_size = configured_pp_size()
self.pp_rank = self.kv_args.pp_rank self.pp_rank = self.kv_args.pp_rank
self.local_ip = get_local_ip_auto() self.local_ip = get_local_ip_auto()
cp_sharded_prefill = self.attn_cp_size > 1 and ( cp_sharded_prefill = self.attn_cp_size > 1 and (
@@ -21,7 +21,11 @@ from sglang.srt.mem_cache.memory_pool import (
MLATokenToKVPool, MLATokenToKVPool,
ReqToTokenPool, ReqToTokenPool,
) )
from sglang.srt.runtime_context import get_schedule from sglang.srt.runtime_context import (
get_memory,
get_schedule,
get_serving,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import ceil_align from sglang.srt.utils.common import ceil_align
@@ -68,10 +72,10 @@ class DecodeKVCacheOffloadManager:
self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group) self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group)
hicache_storage_backend_extra_config = {} hicache_storage_backend_extra_config = {}
if server_args.hicache_storage_backend_extra_config: if get_memory().hicache_storage_backend_extra_config:
try: try:
hicache_storage_backend_extra_config = json.loads( hicache_storage_backend_extra_config = json.loads(
server_args.hicache_storage_backend_extra_config get_memory().hicache_storage_backend_extra_config
) )
except json.JSONDecodeError as e: except json.JSONDecodeError as e:
raise ValueError( raise ValueError(
@@ -83,10 +87,10 @@ class DecodeKVCacheOffloadManager:
mem_pool_host=self.decode_host_mem_pool, mem_pool_host=self.decode_host_mem_pool,
page_size=self.page_size, page_size=self.page_size,
tp_group=tp_group, tp_group=tp_group,
io_backend=server_args.hicache_io_backend, io_backend=get_memory().hicache_io_backend,
load_cache_event=threading.Event(), load_cache_event=threading.Event(),
storage_backend=server_args.hicache_storage_backend, storage_backend=get_memory().hicache_storage_backend,
model_name=server_args.served_model_name, model_name=get_serving().served_model_name,
storage_backend_extra_config=hicache_storage_backend_extra_config, storage_backend_extra_config=hicache_storage_backend_extra_config,
) )
@@ -6,6 +6,11 @@ import os
from typing import TYPE_CHECKING, Dict, List, Optional, Union from typing import TYPE_CHECKING, Dict, List, Optional, Union
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_memory,
)
from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -313,27 +318,27 @@ def maybe_init_shared_mooncake_transfer_engine(
""" """
use_mooncake_te = ( use_mooncake_te = (
( (
server_args.disaggregation_mode != "null" get_disagg().disaggregation_mode != "null"
and server_args.disaggregation_transfer_backend == "mooncake" and get_disagg().disaggregation_transfer_backend == "mooncake"
) )
or ( or (
server_args.enable_hierarchical_cache get_memory().enable_hierarchical_cache
and server_args.hicache_storage_backend == "mooncake" and get_memory().hicache_storage_backend == "mooncake"
and envs.SGLANG_HICACHE_MOONCAKE_REUSE_TE.get() and envs.SGLANG_HICACHE_MOONCAKE_REUSE_TE.get()
) )
or ( or (
server_args.encoder_only get_disagg().encoder_only
and server_args.encoder_transfer_backend == "mooncake" and get_disagg().encoder_transfer_backend == "mooncake"
) )
or ( or (
server_args.language_only get_disagg().language_only
and server_args.encoder_transfer_backend == "mooncake" and get_disagg().encoder_transfer_backend == "mooncake"
) )
or ( or (
server_args.enable_elastic_expert_backup get_exec().moe.enable_elastic_expert_backup
and server_args.elastic_ep_backend is not None and get_exec().moe.elastic_ep_backend is not None
) )
or server_args.elastic_ep_backend == "mooncake" or get_exec().moe.elastic_ep_backend == "mooncake"
) )
if use_mooncake_te: if use_mooncake_te:
@@ -341,11 +346,12 @@ def maybe_init_shared_mooncake_transfer_engine(
hostname=get_local_ip_auto(), hostname=get_local_ip_auto(),
gpu_id=gpu_id, gpu_id=gpu_id,
ib_device=( ib_device=(
server_args.disaggregation_ib_device or server_args.mooncake_ib_device get_disagg().disaggregation_ib_device
or get_exec().moe.mooncake_ib_device
), ),
) )
if server_args.elastic_ep_backend == "mooncake": if get_exec().moe.elastic_ep_backend == "mooncake":
try: try:
from mooncake.pg import set_transfer_engine from mooncake.pg import set_transfer_engine
except ImportError as e: except ImportError as e:
+13 -8
View File
@@ -11,7 +11,11 @@ from sglang.srt.distributed import get_world_group, parallel_state
from sglang.srt.distributed.utils import get_global_tcp_store from sglang.srt.distributed.utils import get_global_tcp_store
from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata
from sglang.srt.managers.schedule_batch import ServerArgs from sglang.srt.managers.schedule_batch import ServerArgs
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import (
configured_tp_size,
get_exec,
get_parallel,
)
from sglang.srt.utils import is_cpu, is_cuda from sglang.srt.utils import is_cpu, is_cuda
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -87,9 +91,9 @@ class ElasticEPStateManager:
if cls._instance is not None: if cls._instance is not None:
return cls._instance return cls._instance
if server_args.elastic_ep_backend is not None: if get_exec().moe.elastic_ep_backend is not None:
world_size = torch.distributed.get_world_size() world_size = torch.distributed.get_world_size()
active_rank_capacity = server_args.max_ep_size or world_size active_rank_capacity = get_parallel().max_ep_size or world_size
assert active_rank_capacity >= world_size, ( assert active_rank_capacity >= world_size, (
f"--max-ep-size ({active_rank_capacity}) must be >= " f"--max-ep-size ({active_rank_capacity}) must be >= "
f"world_size ({world_size})." f"world_size ({world_size})."
@@ -103,10 +107,10 @@ class ElasticEPStateManager:
inst.snapshot_active_to_last() inst.snapshot_active_to_last()
inst.sync_active_to_cpu() inst.sync_active_to_cpu()
if server_args.moe_a2a_backend == "nixl": if get_exec().moe.moe_a2a_backend == "nixl":
cls._on_scale = cls._on_scale_nixl cls._on_scale = cls._on_scale_nixl
inst.ep_join_rank_offset = server_args.ep_join_rank_offset inst.ep_join_rank_offset = get_parallel().ep_join_rank_offset
if server_args.is_ep_joiner: if server_args.is_ep_joiner:
cls._init_joiner_state(inst, server_args) cls._init_joiner_state(inst, server_args)
@@ -122,12 +126,13 @@ class ElasticEPStateManager:
inst.snapshot_active_to_last() inst.snapshot_active_to_last()
inst.sync_active_to_cpu() inst.sync_active_to_cpu()
if server_args.ep_join_mode == "scale": if get_exec().moe.ep_join_mode == "scale":
inst.effective_ep_size = ( inst.effective_ep_size = (
server_args.ep_join_rank_offset + server_args.tp_size get_parallel().ep_join_rank_offset + configured_tp_size()
) )
inst.original_ep_size = ( inst.original_ep_size = (
server_args.elastic_ep_initial_size or server_args.ep_join_rank_offset get_parallel().elastic_ep_initial_size
or get_parallel().ep_join_rank_offset
) )
inst.has_scaled = True inst.has_scaled = True
else: else:
@@ -17,7 +17,14 @@ from sglang.srt.managers.io_struct import (
) )
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
from sglang.srt.model_loader.utils import set_default_torch_dtype from sglang.srt.model_loader.utils import set_default_torch_dtype
from sglang.srt.runtime_context import publish from sglang.srt.runtime_context import (
configured_tp_size,
get_disagg,
get_exec,
get_model,
get_parallel,
publish,
)
from sglang.srt.server_args import ( from sglang.srt.server_args import (
PortArgs, PortArgs,
ServerArgs, ServerArgs,
@@ -38,14 +45,14 @@ def extract_expert_id(param_name):
class ExpertBackupManager: class ExpertBackupManager:
def __init__(self, server_args: ServerArgs, port_args: PortArgs): def __init__(self, server_args: ServerArgs, port_args: PortArgs):
self.load_format = server_args.load_format self.load_format = get_model().load_format
self.model_config = ModelConfig.from_server_args(server_args) self.model_config = ModelConfig.from_server_args(server_args)
self.continuous_buffer = None self.continuous_buffer = None
self.weight_pointer_map = {} self.weight_pointer_map = {}
self.transfer_engine = None self.transfer_engine = None
self.session_id = None self.session_id = None
self.engine_num = server_args.nnodes self.engine_num = get_parallel().nnodes
self.engine_rank = server_args.node_rank self.engine_rank = get_parallel().node_rank
self.expert_num = self.model_config.hf_config.n_routed_experts self.expert_num = self.model_config.hf_config.n_routed_experts
self.idmn = (self.expert_num // self.engine_num) * self.engine_rank self.idmn = (self.expert_num // self.engine_num) * self.engine_rank
self.idmx = (self.expert_num // self.engine_num) * (self.engine_rank + 1) self.idmx = (self.expert_num // self.engine_num) * (self.engine_rank + 1)
@@ -53,11 +60,11 @@ class ExpertBackupManager:
# Synchronization socket to avoid PUB/SUB slow joiner issues. # Synchronization socket to avoid PUB/SUB slow joiner issues.
self.recv_from_expert_backup_client = context.socket(zmq.PULL) self.recv_from_expert_backup_client = context.socket(zmq.PULL)
self.recv_from_expert_backup_client.bind( self.recv_from_expert_backup_client.bind(
f"tcp://{get_local_ip_auto()}:{PORT_BASE + server_args.node_rank * 2}" f"tcp://{get_local_ip_auto()}:{PORT_BASE + get_parallel().node_rank * 2}"
) )
self.send_to_expert_backup_client = context.socket(zmq.PUB) self.send_to_expert_backup_client = context.socket(zmq.PUB)
self.send_to_expert_backup_client.bind( self.send_to_expert_backup_client.bind(
f"tcp://{get_local_ip_auto()}:{PORT_BASE + server_args.node_rank * 2 + 1}" f"tcp://{get_local_ip_auto()}:{PORT_BASE + get_parallel().node_rank * 2 + 1}"
) )
self.backup_weights_from_disk() self.backup_weights_from_disk()
self.start_transfer_server() self.start_transfer_server()
@@ -66,7 +73,7 @@ class ExpertBackupManager:
# losing the initial PUB message due to slow joiners. # losing the initial PUB message due to slow joiners.
num_ready_clients = 0 num_ready_clients = 0
while num_ready_clients < server_args.tp_size: while num_ready_clients < configured_tp_size():
sock_recv(self.recv_from_expert_backup_client) sock_recv(self.recv_from_expert_backup_client)
num_ready_clients += 1 num_ready_clients += 1
@@ -168,7 +175,7 @@ def run_expert_backup_manager_process(
hostname=get_local_ip_auto(), hostname=get_local_ip_auto(),
gpu_id=0, gpu_id=0,
ib_device=( ib_device=(
server_args.disaggregation_ib_device or server_args.mooncake_ib_device get_disagg().disaggregation_ib_device or get_exec().moe.mooncake_ib_device
), ),
) )
manager = ExpertBackupManager(server_args, port_args) manager = ExpertBackupManager(server_args, port_args)
+19 -15
View File
@@ -45,7 +45,11 @@ from sglang.srt.observability.metrics_collector import (
ExpertDispatchCollector, ExpertDispatchCollector,
resolve_collector_class, resolve_collector_class,
) )
from sglang.srt.runtime_context import get_schedule from sglang.srt.runtime_context import get_device as get_device_namespace
from sglang.srt.runtime_context import (
get_exec,
get_schedule,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import Withable, get_device, get_int_env_var from sglang.srt.utils import Withable, get_device, get_int_env_var
@@ -84,7 +88,7 @@ class ExpertDistributionRecorder(ABC):
expert_location_metadata: ExpertLocationMetadata, expert_location_metadata: ExpertLocationMetadata,
rank: int, rank: int,
): ):
if server_args.expert_distribution_recorder_mode is not None: if get_exec().moe.expert_distribution_recorder_mode is not None:
assert ( assert (
expert_location_metadata is not None expert_location_metadata is not None
), "ExpertLocationMetadata is required for expert distribution recording. One possible" ), "ExpertLocationMetadata is required for expert distribution recording. One possible"
@@ -178,7 +182,7 @@ class _ExpertDistributionRecorderReal(ExpertDistributionRecorder):
if server_args.should_report_expert_balancedness(): if server_args.should_report_expert_balancedness():
logger.info( logger.info(
"ExpertDistributionRecorder auto start record since " "ExpertDistributionRecorder auto start record since "
f"expert_balancedness_report_mode={server_args.expert_balancedness_report_mode}" f"expert_balancedness_report_mode={get_exec().moe.expert_balancedness_report_mode}"
) )
self.start_record() self.start_record()
@@ -328,30 +332,30 @@ class _SinglePassGatherer(ABC):
expert_location_metadata: ExpertLocationMetadata, expert_location_metadata: ExpertLocationMetadata,
rank: int, rank: int,
) -> _SinglePassGatherer: ) -> _SinglePassGatherer:
if server_args.expert_distribution_recorder_mode == "per_token": if get_exec().moe.expert_distribution_recorder_mode == "per_token":
return _DetailSinglePassGatherer( return _DetailSinglePassGatherer(
server_args, expert_location_metadata, rank server_args, expert_location_metadata, rank
) )
if server_args.moe_a2a_backend == "mori": if get_exec().moe.moe_a2a_backend == "mori":
return _DeepepLowLatencySinglePassGatherer(expert_location_metadata, rank) return _DeepepLowLatencySinglePassGatherer(expert_location_metadata, rank)
if server_args.expert_distribution_recorder_mode == "stat_approx": if get_exec().moe.expert_distribution_recorder_mode == "stat_approx":
if server_args.moe_a2a_backend != "none" and ( if get_exec().moe.moe_a2a_backend != "none" and (
server_args.deepep_mode == "normal" get_exec().moe.deepep_mode == "normal"
): ):
return _DeepepNormalSinglePassGatherer(expert_location_metadata, rank) return _DeepepNormalSinglePassGatherer(expert_location_metadata, rank)
else: else:
raise NotImplementedError raise NotImplementedError
if server_args.moe_a2a_backend == "deepep": if get_exec().moe.moe_a2a_backend == "deepep":
if server_args.deepep_mode == "normal": if get_exec().moe.deepep_mode == "normal":
return _SelectExpertsSinglePassGatherer(expert_location_metadata, rank) return _SelectExpertsSinglePassGatherer(expert_location_metadata, rank)
elif server_args.deepep_mode == "low_latency": elif get_exec().moe.deepep_mode == "low_latency":
return _DeepepLowLatencySinglePassGatherer( return _DeepepLowLatencySinglePassGatherer(
expert_location_metadata, expert_location_metadata,
rank, rank,
elastic_ep_enabled=server_args.elastic_ep_backend is not None, elastic_ep_enabled=get_exec().moe.elastic_ep_backend is not None,
) )
else: else:
raise NotImplementedError raise NotImplementedError
@@ -412,11 +416,11 @@ class _DetailSinglePassGatherer(_SinglePassGatherer):
self._TOP_K_NUM, self._TOP_K_NUM,
), ),
dtype=torch.int32, dtype=torch.int32,
device=server_args.device, device=get_device_namespace().device,
) )
self._misc_objects: List[Dict[str, Any]] = [] self._misc_objects: List[Dict[str, Any]] = []
assert ( assert (
not server_args.enable_two_batch_overlap not get_exec().overlap.enable_two_batch_overlap
), "DetailSinglePassGatherer does not support TBO yet" ), "DetailSinglePassGatherer does not support TBO yet"
# TODO assert shared experts fusion is disabled, o/w data is wrong # TODO assert shared experts fusion is disabled, o/w data is wrong
@@ -678,7 +682,7 @@ class _Accumulator(ABC):
"stat_approx": _StatAccumulator, "stat_approx": _StatAccumulator,
"per_pass": _DetailAccumulator, "per_pass": _DetailAccumulator,
"per_token": _DetailAccumulator, "per_token": _DetailAccumulator,
}[server_args.expert_distribution_recorder_mode] }[get_exec().moe.expert_distribution_recorder_mode]
def __init__( def __init__(
self, self,
+13 -7
View File
@@ -25,6 +25,12 @@ import torch
import torch.distributed import torch.distributed
import torch.nn.functional as F import torch.nn.functional as F
from sglang.srt.runtime_context import (
configured_tp_size,
get_device,
get_exec,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -141,7 +147,7 @@ class ExpertLocationMetadata:
): ):
if not isinstance(physical_to_logical_map, torch.Tensor): if not isinstance(physical_to_logical_map, torch.Tensor):
physical_to_logical_map = torch.tensor(physical_to_logical_map) physical_to_logical_map = torch.tensor(physical_to_logical_map)
physical_to_logical_map = physical_to_logical_map.to(server_args.device) physical_to_logical_map = physical_to_logical_map.to(get_device().device)
common = ExpertLocationMetadata._init_common(server_args, model_config) common = ExpertLocationMetadata._init_common(server_args, model_config)
@@ -184,7 +190,7 @@ class ExpertLocationMetadata:
logical_count = torch.tensor(logical_count) logical_count = torch.tensor(logical_count)
if len(logical_count.shape) == 2: if len(logical_count.shape) == 2:
logical_count = logical_count.unsqueeze(0) logical_count = logical_count.unsqueeze(0)
logical_count = logical_count.to(server_args.device) logical_count = logical_count.to(get_device().device)
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
@@ -208,7 +214,7 @@ class ExpertLocationMetadata:
num_groups=num_groups, num_groups=num_groups,
num_nodes=num_nodes, num_nodes=num_nodes,
algorithm=eplb_algorithms.compute_algorithm( algorithm=eplb_algorithms.compute_algorithm(
raw_algorithm=server_args.eplb_algorithm, raw_algorithm=get_exec().moe.eplb_algorithm,
num_groups=num_groups, num_groups=num_groups,
num_nodes=num_nodes, num_nodes=num_nodes,
), ),
@@ -217,9 +223,9 @@ class ExpertLocationMetadata:
return ExpertLocationMetadata._init_raw( return ExpertLocationMetadata._init_raw(
ep_size=common["ep_size"], ep_size=common["ep_size"],
physical_to_logical_map=physical_to_logical_map.to(server_args.device), physical_to_logical_map=physical_to_logical_map.to(get_device().device),
logical_to_all_physical_map=logical_to_all_physical_map.to( logical_to_all_physical_map=logical_to_all_physical_map.to(
server_args.device get_device().device
), ),
) )
@@ -246,7 +252,7 @@ class ExpertLocationMetadata:
if get_exec().moe.ep_join_mode == "scale": if get_exec().moe.ep_join_mode == "scale":
ep_size = max( ep_size = max(
ep_size, ep_size,
get_parallel().ep_join_rank_offset + server_args.tp_size, get_parallel().ep_join_rank_offset + configured_tp_size(),
) )
num_physical_experts, num_local_physical_experts = ( num_physical_experts, num_local_physical_experts = (
_compute_elastic_expert_layout( _compute_elastic_expert_layout(
@@ -781,7 +787,7 @@ def compute_initial_expert_location_metadata(
model_config: ModelConfig, model_config: ModelConfig,
moe_ep_rank: int, moe_ep_rank: int,
) -> Optional[ExpertLocationMetadata]: ) -> Optional[ExpertLocationMetadata]:
data = server_args.init_expert_location data = get_exec().moe.init_expert_location
if data == "trivial": if data == "trivial":
return ExpertLocationMetadata.init_trivial( return ExpertLocationMetadata.init_trivial(
server_args, model_config, moe_ep_rank server_args, model_config, moe_ep_rank
+6 -3
View File
@@ -18,6 +18,10 @@ from sglang.srt.model_executor.cuda_graph_config import (
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import (
get_disagg,
get_spec,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
@@ -59,13 +63,12 @@ def install_canary(
allocator if isinstance(allocator, SWATokenToKVPoolAllocator) else None allocator if isinstance(allocator, SWATokenToKVPoolAllocator) else None
) )
launch_capacities = CanaryLaunchCapacities.from_args( launch_capacities = CanaryLaunchCapacities.from_args(
server_args=model_runner.server_args,
req_to_token_pool_size=model_runner.req_to_token_pool.size, req_to_token_pool_size=model_runner.req_to_token_pool.size,
max_seq_len_per_req=model_runner.req_to_token_pool.req_to_token.shape[1], max_seq_len_per_req=model_runner.req_to_token_pool.req_to_token.shape[1],
pool_slot_count=model_runner.max_total_num_tokens, pool_slot_count=model_runner.max_total_num_tokens,
) )
swa_window_size = model_runner.sliding_window_size or 0 swa_window_size = model_runner.sliding_window_size or 0
speculative_num_steps = int(server_args.speculative_num_steps or 1) speculative_num_steps = int(get_spec().speculative_num_steps or 1)
manager = CanaryManager( manager = CanaryManager(
config=config, config=config,
perturb_config=perturb_config, perturb_config=perturb_config,
@@ -88,7 +91,7 @@ def install_canary(
"install_canary: disaggregation_mode=%s config=%s perturb_config=%s " "install_canary: disaggregation_mode=%s config=%s perturb_config=%s "
"launch_capacities=%s n_buffer_groups=%d buffer_group_kinds=%s " "launch_capacities=%s n_buffer_groups=%d buffer_group_kinds=%s "
"swa_window_size=%d speculative_num_steps=%d", "swa_window_size=%d speculative_num_steps=%d",
server_args.disaggregation_mode, get_disagg().disaggregation_mode,
config, config,
perturb_config, perturb_config,
launch_capacities, launch_capacities,
+9 -8
View File
@@ -2,10 +2,12 @@ from __future__ import annotations
import math import math
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING
if TYPE_CHECKING: from sglang.srt.runtime_context import (
from sglang.srt.server_args import ServerArgs get_exec,
get_schedule,
get_spec,
)
@dataclass(frozen=True, slots=True, kw_only=True) @dataclass(frozen=True, slots=True, kw_only=True)
@@ -43,7 +45,6 @@ class CanaryLaunchCapacities:
def from_args( def from_args(
cls, cls,
*, *,
server_args: ServerArgs,
req_to_token_pool_size: int, req_to_token_pool_size: int,
max_seq_len_per_req: int, max_seq_len_per_req: int,
pool_slot_count: int, pool_slot_count: int,
@@ -63,7 +64,7 @@ class CanaryLaunchCapacities:
f"kv-canary: pool_slot_count must be positive, got {pool_slot_count}" f"kv-canary: pool_slot_count must be positive, got {pool_slot_count}"
) )
cuda_graph_config = server_args.cuda_graph_config cuda_graph_config = get_exec().graph.cuda_graph_config
cuda_graph_max_bs = ( cuda_graph_max_bs = (
cuda_graph_config.decode.max_bs if cuda_graph_config is not None else 0 cuda_graph_config.decode.max_bs if cuda_graph_config is not None else 0
) or 0 ) or 0
@@ -72,7 +73,7 @@ class CanaryLaunchCapacities:
f"kv-canary: cuda_graph_max_bs must be non-negative, got {cuda_graph_max_bs}" f"kv-canary: cuda_graph_max_bs must be non-negative, got {cuda_graph_max_bs}"
) )
spec_num_draft_tokens = server_args.speculative_num_draft_tokens spec_num_draft_tokens = get_spec().speculative_num_draft_tokens
if spec_num_draft_tokens is None: if spec_num_draft_tokens is None:
spec_num_draft_tokens = 0 spec_num_draft_tokens = 0
if spec_num_draft_tokens < 0: if spec_num_draft_tokens < 0:
@@ -81,7 +82,7 @@ class CanaryLaunchCapacities:
f"got {spec_num_draft_tokens}" f"got {spec_num_draft_tokens}"
) )
max_prefill_tokens = server_args.max_prefill_tokens max_prefill_tokens = get_schedule().max_prefill_tokens
if max_prefill_tokens <= 0: if max_prefill_tokens <= 0:
raise ValueError( raise ValueError(
f"kv-canary: max_prefill_tokens must be positive, got {max_prefill_tokens}" f"kv-canary: max_prefill_tokens must be positive, got {max_prefill_tokens}"
@@ -93,7 +94,7 @@ class CanaryLaunchCapacities:
max_bs = max(cuda_graph_max_bs, req_to_token_pool_size) max_bs = max(cuda_graph_max_bs, req_to_token_pool_size)
chunked_prefill_size = server_args.chunked_prefill_size chunked_prefill_size = get_schedule().chunked_prefill_size
chunked_limit = ( chunked_limit = (
chunked_prefill_size chunked_prefill_size
if chunked_prefill_size is not None and chunked_prefill_size >= 0 if chunked_prefill_size is not None and chunked_prefill_size >= 0
+12 -5
View File
@@ -8,6 +8,11 @@ from sglang.srt.disaggregation.utils import (
TransferBackend, TransferBackend,
get_kv_class, get_kv_class,
) )
from sglang.srt.runtime_context import (
get_disagg,
get_parallel,
get_serving,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -15,8 +20,8 @@ def start_disagg_service(
server_args: ServerArgs, server_args: ServerArgs,
): ):
# Start kv bootstrap server on prefill # Start kv bootstrap server on prefill
disagg_mode = DisaggregationMode(server_args.disaggregation_mode) disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode)
transfer_backend = TransferBackend(server_args.disaggregation_transfer_backend) transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend)
if disagg_mode == DisaggregationMode.PREFILL: if disagg_mode == DisaggregationMode.PREFILL:
# only start bootstrap server on prefill tm # only start bootstrap server on prefill tm
@@ -24,8 +29,8 @@ def start_disagg_service(
transfer_backend, KVClassType.BOOTSTRAP_SERVER transfer_backend, KVClassType.BOOTSTRAP_SERVER
) )
bootstrap_server = kv_bootstrap_server_class( bootstrap_server = kv_bootstrap_server_class(
host=server_args.host, host=get_serving().host,
port=server_args.disaggregation_bootstrap_port, port=get_disagg().disaggregation_bootstrap_port,
) )
maybe_create_ascend_config_store( maybe_create_ascend_config_store(
server_args=server_args, transfer_backend=transfer_backend server_args=server_args, transfer_backend=transfer_backend
@@ -43,7 +48,9 @@ def maybe_create_ascend_config_store(
bootstrap-server subclasses are all plain ``CommonKVBootstrapServer``, bootstrap-server subclasses are all plain ``CommonKVBootstrapServer``,
which the rust registry ports verbatim), leaving this store as the only which the rust registry ports verbatim), leaving this store as the only
``start_disagg_service`` duty left to perform.""" ``start_disagg_service`` duty left to perform."""
if not (server_args.node_rank == 0 and transfer_backend == TransferBackend.ASCEND): if not (
get_parallel().node_rank == 0 and transfer_backend == TransferBackend.ASCEND
):
return return
try: try:
from memfabric_hybrid import create_config_store from memfabric_hybrid import create_config_store
+7 -3
View File
@@ -8,6 +8,10 @@ import torch
from sglang.kernels.ops.speculative.gather_spec_extras import gather_spec_extras from sglang.kernels.ops.speculative.gather_spec_extras import gather_spec_extras
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.runtime_context import (
get_exec,
get_spec,
)
from sglang.srt.utils import is_cuda, is_hip, is_npu from sglang.srt.utils import is_cuda, is_hip, is_npu
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -34,10 +38,10 @@ def decide_needs_cpu_seq_lens(
# importable everywhere; spec_info pulls in the spec/schedule_batch graph. # importable everywhere; spec_info pulls in the spec/schedule_batch graph.
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
if server_args.enable_two_batch_overlap: if get_exec().overlap.enable_two_batch_overlap:
# FIXME: support TBO without seq lens cpu value # FIXME: support TBO without seq lens cpu value
return True return True
algo = SpeculativeAlgorithm.from_string(server_args.speculative_algorithm) algo = SpeculativeAlgorithm.from_string(get_spec().speculative_algorithm)
if algo.is_ngram(): if algo.is_ngram():
# ngram's USE_FULL_MASK verify path reads seq_lens_cpu per req to size # ngram's USE_FULL_MASK verify path reads seq_lens_cpu per req to size
# the tree mask, regardless of the attn backend (e.g. Triton opts out). # the tree mask, regardless of the attn backend (e.g. Triton opts out).
@@ -56,7 +60,7 @@ def decide_needs_confidence_relay(server_args: ServerArgs) -> bool:
) )
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
algo = SpeculativeAlgorithm.from_string(server_args.speculative_algorithm) algo = SpeculativeAlgorithm.from_string(get_spec().speculative_algorithm)
if not algo.is_dspark(): if not algo.is_dspark():
return False return False
return read_ragged_verify_mode() is not RaggedVerifyMode.STATIC return read_ragged_verify_mode() is not RaggedVerifyMode.STATIC
+8 -8
View File
@@ -1930,7 +1930,7 @@ def release_req(
# Callers that will recompute the KV instead (PD true-retraction rebootstrap) # Callers that will recompute the KV instead (PD true-retraction rebootstrap)
# pass offload_kv=False to skip the wasteful device->host copy. # pass offload_kv=False to skip the wasteful device->host copy.
backup_saved = True backup_saved = True
if server_args.disaggregation_mode == "decode" and offload_kv: if get_disagg().disaggregation_mode == "decode" and offload_kv:
backup_saved = retraction_backup( backup_saved = retraction_backup(
req, req,
tree_cache, tree_cache,
@@ -2826,7 +2826,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self, server_args: ServerArgs self, server_args: ServerArgs
) -> Tuple[List[Req], float, List[Req]]: ) -> Tuple[List[Req], float, List[Req]]:
"""Retract the decoding requests when there is not enough memory.""" """Retract the decoding requests when there is not enough memory."""
sorted_indices = self._get_decode_retraction_order(self.reqs, server_args) sorted_indices = self._get_decode_retraction_order(self.reqs)
retracted_reqs = [] retracted_reqs = []
reqs_to_abort: List[Req] = [] reqs_to_abort: List[Req] = []
@@ -2886,9 +2886,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
return retracted_reqs, new_estimate_ratio, reqs_to_abort return retracted_reqs, new_estimate_ratio, reqs_to_abort
@staticmethod @staticmethod
def _get_decode_retraction_order( def _get_decode_retraction_order(reqs: List[Req]) -> List[int]:
reqs: List[Req], server_args: ServerArgs
) -> List[int]:
"""Return indices ordered from most-preferred to least-preferred to keep. """Return indices ordered from most-preferred to least-preferred to keep.
The retraction loop pops from the end of this list, so the least-preferred The retraction loop pops from the end of this list, so the least-preferred
@@ -2901,15 +2899,17 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
def length_key(req: Req) -> Tuple[int, int]: def length_key(req: Req) -> Tuple[int, int]:
return (len(req.output_ids), -len(req.origin_input_ids)) return (len(req.output_ids), -len(req.origin_input_ids))
if server_args.retraction_policy == "priority": if get_schedule().retraction_policy == "priority":
priority_sign = 1 if server_args.schedule_low_priority_values_first else -1 priority_sign = (
1 if get_schedule().schedule_low_priority_values_first else -1
)
def retraction_key(req: Req) -> Tuple[int, int, int]: def retraction_key(req: Req) -> Tuple[int, int, int]:
priority = req.priority priority = req.priority
if priority is None: if priority is None:
priority = ( priority = (
sys.maxsize sys.maxsize
if server_args.schedule_low_priority_values_first if get_schedule().schedule_low_priority_values_first
else -sys.maxsize - 1 else -sys.maxsize - 1
) )
return (priority * (-priority_sign), *length_key(req)) return (priority * (-priority_sign), *length_key(req))
+39 -34
View File
@@ -30,6 +30,7 @@ from typing import TYPE_CHECKING, Any, Deque, Dict, List, Optional, Set, Tuple,
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
attention_backends, attention_backends,
configured_attn_cp_size, configured_attn_cp_size,
configured_dcp_size,
configured_moe_dp_size, configured_moe_dp_size,
configured_pp_size, configured_pp_size,
configured_tp_size, configured_tp_size,
@@ -417,71 +418,75 @@ class Scheduler(
# Parse args # Parse args
self.server_args = server_args self.server_args = server_args
self.nccl_port = port_args.nccl_port self.nccl_port = port_args.nccl_port
self.schedule_policy = server_args.schedule_policy self.schedule_policy = get_schedule().schedule_policy
self.enable_priority_scheduling = server_args.enable_priority_scheduling self.enable_priority_scheduling = get_schedule().enable_priority_scheduling
self.abort_on_priority_when_disabled = ( self.abort_on_priority_when_disabled = (
server_args.abort_on_priority_when_disabled get_schedule().abort_on_priority_when_disabled
) )
self.schedule_low_priority_values_first = ( self.schedule_low_priority_values_first = (
server_args.schedule_low_priority_values_first get_schedule().schedule_low_priority_values_first
) )
self.priority_scheduling_preemption_threshold = ( self.priority_scheduling_preemption_threshold = (
server_args.priority_scheduling_preemption_threshold get_schedule().priority_scheduling_preemption_threshold
) )
self.enable_lora = server_args.enable_lora self.enable_lora = get_lora().enable_lora
self.enable_lora_overlap_loading = server_args.enable_lora_overlap_loading self.enable_lora_overlap_loading = get_lora().enable_lora_overlap_loading
self.max_loras_per_batch = server_args.max_loras_per_batch self.max_loras_per_batch = get_lora().max_loras_per_batch
self.enable_overlap = not server_args.disable_overlap_schedule and not use_mlx() self.enable_overlap = (
self.enable_overlap_mlx = not server_args.disable_overlap_schedule and use_mlx() not get_schedule().disable_overlap_schedule and not use_mlx()
self.enable_pdmux = server_args.enable_pdmux )
self.enable_overlap_mlx = (
not get_schedule().disable_overlap_schedule and use_mlx()
)
self.enable_pdmux = get_disagg().enable_pdmux
self.skip_tokenizer_init = get_serving().skip_tokenizer_init self.skip_tokenizer_init = get_serving().skip_tokenizer_init
self.stream_interval = server_args.stream_interval self.stream_interval = get_serving().stream_interval
self.spec_algorithm = SpeculativeAlgorithm.from_string( self.spec_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
self.enable_hierarchical_cache = server_args.enable_hierarchical_cache self.enable_hierarchical_cache = get_memory().enable_hierarchical_cache
self.enable_session_radix_cache = server_args.enable_session_radix_cache self.enable_session_radix_cache = get_memory().enable_session_radix_cache
self.enable_hicache_storage = server_args.hicache_storage_backend is not None self.enable_hicache_storage = get_memory().hicache_storage_backend is not None
self.enable_decode_hicache = ( self.enable_decode_hicache = (
server_args.disaggregation_decode_enable_radix_cache get_disagg().disaggregation_decode_enable_radix_cache
and self.enable_hierarchical_cache and self.enable_hierarchical_cache
) )
self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get() self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get()
self.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get() self.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get()
self.enable_hisparse = server_args.enable_hisparse self.enable_hisparse = get_memory().enable_hisparse
self.enable_dp_attention = get_parallel().enable_dp_attention self.enable_dp_attention = get_parallel().enable_dp_attention
self.enable_unified_memory = server_args.enable_unified_memory self.enable_unified_memory = get_memory().enable_unified_memory
# Distributed rank info # Distributed rank info
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = ( attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
compute_dp_attention_world_info( compute_dp_attention_world_info(
get_parallel().enable_dp_attention, get_parallel().enable_dp_attention,
tp_rank, tp_rank,
server_args.tp_size, configured_tp_size(),
get_parallel().dp_size, get_parallel().dp_size,
server_args.attn_cp_size, configured_attn_cp_size(),
) )
) )
self.ps = ParallelState( self.ps = ParallelState(
tp_rank=tp_rank, tp_rank=tp_rank,
tp_size=server_args.tp_size, tp_size=configured_tp_size(),
pp_rank=pp_rank, pp_rank=pp_rank,
pp_size=server_args.pp_size, pp_size=configured_pp_size(),
dp_rank=dp_rank, dp_rank=dp_rank,
dp_size=get_parallel().dp_size, dp_size=get_parallel().dp_size,
attn_tp_rank=attn_tp_rank, attn_tp_rank=attn_tp_rank,
attn_tp_size=attn_tp_size, attn_tp_size=attn_tp_size,
attn_cp_rank=attn_cp_rank, attn_cp_rank=attn_cp_rank,
attn_cp_size=server_args.attn_cp_size, attn_cp_size=configured_attn_cp_size(),
attn_dcp_rank=tp_rank % server_args.dcp_size, attn_dcp_rank=tp_rank % configured_dcp_size(),
attn_dcp_size=server_args.dcp_size, attn_dcp_size=configured_dcp_size(),
attn_dp_rank=attn_dp_rank, attn_dp_rank=attn_dp_rank,
attn_dp_size=attn_dp_size, attn_dp_size=attn_dp_size,
moe_ep_rank=moe_ep_rank, moe_ep_rank=moe_ep_rank,
moe_ep_size=get_parallel().ep_size, moe_ep_size=get_parallel().ep_size,
moe_dp_rank=moe_dp_rank, moe_dp_rank=moe_dp_rank,
moe_dp_size=server_args.moe_dp_size, moe_dp_size=configured_moe_dp_size(),
gpu_id=gpu_id, gpu_id=gpu_id,
) )
@@ -839,7 +844,7 @@ class Scheduler(
if ( if (
self.model_config.is_multimodal self.model_config.is_multimodal
and self.processor is not None and self.processor is not None
and not server_args.language_model_only and not get_disagg().language_model_only
): ):
try: try:
import_processors("sglang.srt.multimodal.processors") import_processors("sglang.srt.multimodal.processors")
@@ -1282,7 +1287,7 @@ class Scheduler(
self.new_token_ratio_tracker = NewTokenRatioTracker.from_config() self.new_token_ratio_tracker = NewTokenRatioTracker.from_config()
def init_soft_watchdog(self, server_args: ServerArgs): def init_soft_watchdog(self, server_args: ServerArgs):
if (x := server_args.soft_watchdog_timeout) is not None: if (x := get_device().soft_watchdog_timeout) is not None:
self.soft_watchdog = create_scheduler_watchdog( self.soft_watchdog = create_scheduler_watchdog(
self, watchdog_timeout=x, soft=True self, watchdog_timeout=x, soft=True
) )
@@ -5172,16 +5177,16 @@ def run_scheduler_process(
parent_process = psutil.Process().parent() parent_process = psutil.Process().parent()
# Set up tracing # Set up tracing
if server_args.enable_trace: if get_observability().enable_trace:
process_tracing_init( process_tracing_init(
server_args.otlp_traces_endpoint, get_observability().otlp_traces_endpoint,
"sglang", "sglang",
trace_modules=server_args.trace_modules, trace_modules=get_observability().trace_modules,
) )
thread_label = "Scheduler" thread_label = "Scheduler"
if server_args.disaggregation_mode == "prefill": if get_disagg().disaggregation_mode == "prefill":
thread_label = "Prefill Scheduler" thread_label = "Prefill Scheduler"
elif server_args.disaggregation_mode == "decode": elif get_disagg().disaggregation_mode == "decode":
thread_label = "Decode Scheduler" thread_label = "Decode Scheduler"
trace_set_thread_info(thread_label, tp_rank, dp_rank, pp_rank) trace_set_thread_info(thread_label, tp_rank, dp_rank, pp_rank)
+60 -56
View File
@@ -131,7 +131,9 @@ from sglang.srt.runtime_context import (
get_memory, get_memory,
get_mm, get_mm,
get_model, get_model,
get_observability,
get_parallel, get_parallel,
get_schedule,
get_serving, get_serving,
get_spec, get_spec,
) )
@@ -411,14 +413,14 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.elastic_pending_ep_size = None self.elastic_pending_ep_size = None
self.elastic_scale_phase = "idle" self.elastic_scale_phase = "idle"
self.elastic_last_error = None self.elastic_last_error = None
self.enable_metrics = server_args.enable_metrics self.enable_metrics = get_observability().enable_metrics
self.incremental_streaming_output = server_args.incremental_streaming_output self.incremental_streaming_output = get_serving().incremental_streaming_output
self.enable_lora = get_lora().enable_lora self.enable_lora = get_lora().enable_lora
self.enable_trace = server_args.enable_trace self.enable_trace = get_observability().enable_trace
self.allow_auto_truncate = server_args.allow_auto_truncate self.allow_auto_truncate = get_serving().allow_auto_truncate
self.skip_tokenizer_init = server_args.skip_tokenizer_init self.skip_tokenizer_init = get_serving().skip_tokenizer_init
self.preferred_sampling_params = get_serving().preferred_sampling_params self.preferred_sampling_params = get_serving().preferred_sampling_params
self.crash_dump_folder = server_args.crash_dump_folder self.crash_dump_folder = get_observability().crash_dump_folder
# Init model config # Init model config
self.init_model_config() self.init_model_config()
@@ -462,15 +464,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
model_config_class = getattr(self, "model_config_class", ModelConfig) model_config_class = getattr(self, "model_config_class", ModelConfig)
# Read model args # Read model args
self.model_path = server_args.model_path self.model_path = get_model().model_path
self.served_model_name = server_args.served_model_name self.served_model_name = get_serving().served_model_name
self.model_config = model_config_class.from_server_args(server_args) self.model_config = model_config_class.from_server_args(server_args)
self.is_generation = self.model_config.is_generation self.is_generation = self.model_config.is_generation
self.context_len = self.model_config.context_len self.context_len = self.model_config.context_len
self.image_token_id = self.model_config.image_token_id self.image_token_id = self.model_config.image_token_id
self.max_req_input_len = None # Will be set later in engine.py self.max_req_input_len = None # Will be set later in engine.py
self.enable_priority_scheduling = server_args.enable_priority_scheduling self.enable_priority_scheduling = get_schedule().enable_priority_scheduling
self.default_priority_value = server_args.default_priority_value self.default_priority_value = get_schedule().default_priority_value
self.num_reserved_tokens = compute_num_reserved_tokens() self.num_reserved_tokens = compute_num_reserved_tokens()
self.validate_total_tokens = True self.validate_total_tokens = True
@@ -478,7 +480,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
server_args = self.server_args server_args = self.server_args
# Initialize tokenizer and processor # Initialize tokenizer and processor
if self.model_config.is_multimodal and not server_args.language_model_only: if self.model_config.is_multimodal and not get_disagg().language_model_only:
import_processors("sglang.srt.multimodal.processors") import_processors("sglang.srt.multimodal.processors")
if mm_process_pkg := envs.SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE.get(): if mm_process_pkg := envs.SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE.get():
import_processors(mm_process_pkg, overwrite=True) import_processors(mm_process_pkg, overwrite=True)
@@ -496,7 +498,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
model_config=self.model_config, model_config=self.model_config,
) )
if server_args.skip_tokenizer_init: if get_serving().skip_tokenizer_init:
self.tokenizer = self.processor = None self.tokenizer = self.processor = None
else: else:
self.processor = _processor self.processor = _processor
@@ -505,26 +507,26 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
else: else:
self.mm_processor = self.processor = None self.mm_processor = self.processor = None
if server_args.skip_tokenizer_init: if get_serving().skip_tokenizer_init:
self.tokenizer = None self.tokenizer = None
else: else:
self.tokenizer = get_tokenizer( self.tokenizer = get_tokenizer(
get_serving().tokenizer_path, get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode, tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=get_model().trust_remote_code,
revision=server_args.revision, revision=get_model().revision,
tokenizer_backend=server_args.tokenizer_backend, tokenizer_backend=get_serving().tokenizer_backend,
) )
# Initialize async dynamic batch tokenizer if enabled (common for both multimodal and non-multimodal) # Initialize async dynamic batch tokenizer if enabled (common for both multimodal and non-multimodal)
if ( if (
server_args.enable_dynamic_batch_tokenizer get_serving().enable_dynamic_batch_tokenizer
and not server_args.skip_tokenizer_init and not get_serving().skip_tokenizer_init
): ):
self.async_dynamic_batch_tokenizer = AsyncDynamicbatchTokenizer( self.async_dynamic_batch_tokenizer = AsyncDynamicbatchTokenizer(
self.tokenizer, self.tokenizer,
max_batch_size=server_args.dynamic_batch_tokenizer_batch_size, max_batch_size=get_serving().dynamic_batch_tokenizer_batch_size,
batch_wait_timeout_s=server_args.dynamic_batch_tokenizer_batch_timeout, batch_wait_timeout_s=get_serving().dynamic_batch_tokenizer_batch_timeout,
) )
else: else:
self.async_dynamic_batch_tokenizer = None self.async_dynamic_batch_tokenizer = None
@@ -547,7 +549,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.recv_from_detokenizer = get_zmq_socket( self.recv_from_detokenizer = get_zmq_socket(
context, zmq.PULL, port_args.tokenizer_ipc_name, True context, zmq.PULL, port_args.tokenizer_ipc_name, True
) )
if self.server_args.tokenizer_worker_num == 1: if get_serving().tokenizer_worker_num == 1:
self.send_to_scheduler = get_zmq_socket( self.send_to_scheduler = get_zmq_socket(
context, zmq.PUSH, port_args.scheduler_input_ipc_name, True context, zmq.PUSH, port_args.scheduler_input_ipc_name, True
) )
@@ -595,10 +597,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# TODO: Refactor and organize the log export code. # TODO: Refactor and organize the log export code.
# Request logging # Request logging
self.request_logger = RequestLogger( self.request_logger = RequestLogger(
log_requests=self.server_args.log_requests, log_requests=get_observability().log_requests,
log_requests_level=self.server_args.log_requests_level, log_requests_level=get_observability().log_requests_level,
log_requests_format=self.server_args.log_requests_format, log_requests_format=get_observability().log_requests_format,
log_requests_target=self.server_args.log_requests_target, log_requests_target=get_observability().log_requests_target,
) )
# Dumping # Dumping
@@ -621,7 +623,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def init_weight_update(self): def init_weight_update(self):
# Initial weights status # Initial weights status
self.initial_weights_loaded = True self.initial_weights_loaded = True
if self.server_args.checkpoint_engine_wait_weights_before_ready: if get_model().checkpoint_engine_wait_weights_before_ready:
self.initial_weights_loaded = False self.initial_weights_loaded = False
# Weight updates # Weight updates
@@ -668,7 +670,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Encoder Disaggregation # Encoder Disaggregation
self.encoder_bootstrap_server = None self.encoder_bootstrap_server = None
if self.server_args.language_only: if get_disagg().language_only:
from sglang.srt.disaggregation.encoder.receiver import ( from sglang.srt.disaggregation.encoder.receiver import (
EncoderBootstrapServer, EncoderBootstrapServer,
) )
@@ -677,10 +679,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# entries as encoders register, the receiver reads from the same # entries as encoders register, the receiver reads from the same
# list. Pre-populated with static --encoder-urls so the legacy # list. Pre-populated with static --encoder-urls so the legacy
# CLI flag still works (alongside dynamic registrations). # CLI flag still works (alongside dynamic registrations).
self.encoder_urls: List[str] = list(self.server_args.encoder_urls) self.encoder_urls: List[str] = list(get_disagg().encoder_urls)
self.encoder_bootstrap_server = EncoderBootstrapServer( self.encoder_bootstrap_server = EncoderBootstrapServer(
host=self.server_args.host, host=get_serving().host,
port=self.server_args.encoder_bootstrap_port, port=get_disagg().encoder_bootstrap_port,
urls=self.encoder_urls, urls=self.encoder_urls,
) )
self.mm_receiver = create_mm_receiver( self.mm_receiver = create_mm_receiver(
@@ -698,16 +700,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
) )
labels = { labels = {
"model_name": self.server_args.served_model_name, "model_name": get_serving().served_model_name,
"engine_type": engine_type, "engine_type": engine_type,
} }
if self.enable_priority_scheduling: if self.enable_priority_scheduling:
labels["priority"] = "" labels["priority"] = ""
if self.server_args.tokenizer_metrics_allowed_custom_labels: if get_observability().tokenizer_metrics_allowed_custom_labels:
for label in self.server_args.tokenizer_metrics_allowed_custom_labels: for (
label
) in get_observability().tokenizer_metrics_allowed_custom_labels:
labels[label] = "" labels[label] = ""
if self.server_args.extra_metric_labels: if get_observability().extra_metric_labels:
labels.update(self.server_args.extra_metric_labels) labels.update(get_observability().extra_metric_labels)
tokenizer_collector_cls = resolve_collector_class( tokenizer_collector_cls = resolve_collector_class(
self.server_args, self.server_args,
STAT_LOGGER_ROLE_TOKENIZER, STAT_LOGGER_ROLE_TOKENIZER,
@@ -716,15 +720,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.metrics_collector = tokenizer_collector_cls( self.metrics_collector = tokenizer_collector_cls(
server_args=self.server_args, server_args=self.server_args,
labels=labels, labels=labels,
bucket_time_to_first_token=self.server_args.bucket_time_to_first_token, bucket_time_to_first_token=get_observability().bucket_time_to_first_token,
bucket_e2e_request_latency=self.server_args.bucket_e2e_request_latency, bucket_e2e_request_latency=get_observability().bucket_e2e_request_latency,
bucket_inter_token_latency=self.server_args.bucket_inter_token_latency, bucket_inter_token_latency=get_observability().bucket_inter_token_latency,
) )
start_cpu_monitor_thread("tokenizer") start_cpu_monitor_thread("tokenizer")
if self.server_args.gc_warning_threshold_secs > 0.0: if get_observability().gc_warning_threshold_secs > 0.0:
configure_gc_warning(self.server_args.gc_warning_threshold_secs) configure_gc_warning(get_observability().gc_warning_threshold_secs)
self.soft_watchdog = Watchdog.create( self.soft_watchdog = Watchdog.create(
debug_name="TokenizerManager", debug_name="TokenizerManager",
watchdog_timeout=get_device().soft_watchdog_timeout, watchdog_timeout=get_device().soft_watchdog_timeout,
@@ -773,7 +777,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if ( if (
isinstance(obj, GenerateReqInput) isinstance(obj, GenerateReqInput)
and obj.max_thinking_tokens is not None and obj.max_thinking_tokens is not None
and not self.server_args.enable_strict_thinking and not get_serving().enable_strict_thinking
): ):
raise ValueError( raise ValueError(
"max_thinking_tokens requires the server to be launched with " "max_thinking_tokens requires the server to be launched with "
@@ -793,7 +797,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self._init_req_state(obj, request) self._init_req_state(obj, request)
try: try:
if self.server_args.language_only: if get_disagg().language_only:
self._handle_epd_disaggregation_encode_request(obj) self._handle_epd_disaggregation_encode_request(obj)
# Log the request # Log the request
@@ -993,7 +997,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
) )
contains_mm_input = obj.contains_mm_input() contains_mm_input = obj.contains_mm_input()
if contains_mm_input and self.server_args.language_model_only: if contains_mm_input and get_disagg().language_model_only:
raise ValueError( raise ValueError(
"Multimodal inputs are not supported when --language-model-only " "Multimodal inputs are not supported when --language-model-only "
"is set; the encoder is not loaded. Restart without the flag." "is set; the encoder is not loaded. Restart without the flag."
@@ -1028,10 +1032,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
) )
if ( if (
not self.server_args.language_only not get_disagg().language_only
or get_disagg().encoder_transfer_backend == "zmq_to_tokenizer" or get_disagg().encoder_transfer_backend == "zmq_to_tokenizer"
): ):
if self.server_args.language_only: if get_disagg().language_only:
mm_inputs = await self.mm_receiver.recv_mm_data( mm_inputs = await self.mm_receiver.recv_mm_data(
request_obj=obj, request_obj=obj,
mm_processor=self.mm_processor, mm_processor=self.mm_processor,
@@ -1053,7 +1057,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
max_req_input_len=self.max_req_input_len, max_req_input_len=self.max_req_input_len,
) )
elif ( elif (
self.server_args.language_only get_disagg().language_only
and get_disagg().encoder_transfer_backend and get_disagg().encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"] in ["zmq_to_scheduler", "mooncake"]
and not obj.need_wait_for_mm_inputs and not obj.need_wait_for_mm_inputs
@@ -1238,7 +1242,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
) )
if ( if (
obj.custom_logit_processor obj.custom_logit_processor
and not self.server_args.enable_custom_logit_processor and not get_exec().features.enable_custom_logit_processor
): ):
raise ValueError( raise ValueError(
"The server is not configured to enable custom logit processor. " "The server is not configured to enable custom logit processor. "
@@ -1945,7 +1949,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
return return
if ( if (
not abort_all not abort_all
and self.server_args.tokenizer_worker_num == 1 and get_serving().tokenizer_worker_num == 1
and rid not in self.rid_to_state and rid not in self.rid_to_state
): ):
return return
@@ -2175,7 +2179,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
else: else:
customized_info = None customized_info = None
pending_notify: dict[str, ReqState] = {} pending_notify: dict[str, ReqState] = {}
batch_notify_size = self.server_args.batch_notify_size batch_notify_size = get_serving().batch_notify_size
for i, rid in enumerate(recv_obj.rids): for i, rid in enumerate(recv_obj.rids):
state = self.rid_to_state.get(rid, None) state = self.rid_to_state.get(rid, None)
if state is None: if state is None:
@@ -3183,7 +3187,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state.obj.top_logprobs_num, state.obj.top_logprobs_num,
state.obj.token_ids_logprob, state.obj.token_ids_logprob,
state.obj.return_text_in_logprobs state.obj.return_text_in_logprobs
and not self.server_args.skip_tokenizer_init, and not get_serving().skip_tokenizer_init,
) )
output_ids = state.output_ids output_ids = state.output_ids
@@ -3304,12 +3308,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
unique_lora_paths = set(obj.lora_path) unique_lora_paths = set(obj.lora_path)
if ( if (
self.server_args.max_loaded_loras is not None get_lora().max_loaded_loras is not None
and len(unique_lora_paths) > self.server_args.max_loaded_loras and len(unique_lora_paths) > get_lora().max_loaded_loras
): ):
raise ValueError( raise ValueError(
f"Received request with {len(unique_lora_paths)} unique loras requested " f"Received request with {len(unique_lora_paths)} unique loras requested "
f"but max loaded loras is {self.server_args.max_loaded_loras}" f"but max loaded loras is {get_lora().max_loaded_loras}"
) )
# Reload all existing LoRA adapters that have been dynamically unloaded # Reload all existing LoRA adapters that have been dynamically unloaded
@@ -3442,7 +3446,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if isinstance(obj, GenerateReqInput) and obj.contains_mm_input(): if isinstance(obj, GenerateReqInput) and obj.contains_mm_input():
# dispatch to encoder by default # dispatch to encoder by default
should_dispatch = True should_dispatch = True
if self.server_args.enable_adaptive_dispatch_to_encoder: if get_disagg().enable_adaptive_dispatch_to_encoder:
should_dispatch = self._should_dispatch_to_encoder(obj) should_dispatch = self._should_dispatch_to_encoder(obj)
# Set need_wait_for_mm_inputs flag based on whether we dispatch to encoder # Set need_wait_for_mm_inputs flag based on whether we dispatch to encoder
@@ -3579,7 +3583,7 @@ def get_processor_wrapper(server_args):
def determine_tensor_transport_mode(server_args: ServerArgs) -> TensorTransportMode: def determine_tensor_transport_mode(server_args: ServerArgs) -> TensorTransportMode:
is_cross_node = server_args.dist_init_addr is_cross_node = get_parallel().dist_init_addr
if is_cross_node: if is_cross_node:
# Fallback to default CPU transport for multi-node # Fallback to default CPU transport for multi-node
+24 -17
View File
@@ -52,7 +52,14 @@ from sglang.srt.model_executor.graph_memory_usage import (
merge_graph_time_usage, merge_graph_time_usage,
) )
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.runtime_context import get_exec, get_model, get_schedule, get_spec from sglang.srt.runtime_context import (
get_device,
get_exec,
get_model,
get_schedule,
get_serving,
get_spec,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed
from sglang.srt.utils.hf_transformers_utils import ( from sglang.srt.utils.hf_transformers_utils import (
@@ -341,28 +348,28 @@ class TpModelWorker(BaseTpWorker):
self._init_dllm_algorithm() self._init_dllm_algorithm()
if server_args.skip_tokenizer_init or self.is_draft_worker: if get_serving().skip_tokenizer_init or self.is_draft_worker:
# A draft worker's tokenizer would only duplicate the target's: # A draft worker's tokenizer would only duplicate the target's:
# tokenizer_path always points at the target model. # tokenizer_path always points at the target model.
self.tokenizer = self.processor = None self.tokenizer = self.processor = None
else: else:
if self.model_config.is_multimodal: if self.model_config.is_multimodal:
self.processor = get_processor( self.processor = get_processor(
server_args.tokenizer_path, get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode, tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=get_model().trust_remote_code,
revision=server_args.revision, revision=get_model().revision,
tokenizer_backend=server_args.tokenizer_backend, tokenizer_backend=get_serving().tokenizer_backend,
model_name=server_args.model_path, model_name=get_model().model_path,
) )
self.tokenizer = get_tokenizer_from_processor(self.processor) self.tokenizer = get_tokenizer_from_processor(self.processor)
else: else:
self.tokenizer = get_tokenizer( self.tokenizer = get_tokenizer(
server_args.tokenizer_path, get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode, tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=get_model().trust_remote_code,
revision=server_args.revision, revision=get_model().revision,
tokenizer_backend=server_args.tokenizer_backend, tokenizer_backend=get_serving().tokenizer_backend,
) )
self.device = self.model_runner.device self.device = self.model_runner.device
@@ -373,18 +380,18 @@ class TpModelWorker(BaseTpWorker):
# Sync random seed across TP workers. # Sync random seed across TP workers.
# Elastic joiners cannot enter the launch-time WORLD broadcast. # Elastic joiners cannot enter the launch-time WORLD broadcast.
if server_args.is_ep_joiner: if server_args.is_ep_joiner:
self.random_seed = server_args.random_seed self.random_seed = get_device().random_seed
else: else:
self.random_seed = broadcast_pyobj( self.random_seed = broadcast_pyobj(
[server_args.random_seed], [get_device().random_seed],
self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank, self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank,
self.world_group.cpu_group, self.world_group.cpu_group,
src=self.world_group.ranks[0], src=self.world_group.ranks[0],
)[0] )[0]
set_random_seed(self.random_seed) set_random_seed(self.random_seed)
self.enable_overlap = not server_args.disable_overlap_schedule self.enable_overlap = not get_schedule().disable_overlap_schedule
self.enable_spec = server_args.speculative_algorithm is not None self.enable_spec = get_spec().speculative_algorithm is not None
self.hicache_layer_transfer_counter = None self.hicache_layer_transfer_counter = None
def alloc_memory_pool( def alloc_memory_pool(
+19 -15
View File
@@ -65,7 +65,11 @@ from sglang.srt.observability.metrics_collector import (
StorageMetricsCollector, StorageMetricsCollector,
resolve_collector_class, resolve_collector_class,
) )
from sglang.srt.runtime_context import get_memory from sglang.srt.runtime_context import (
get_memory,
get_observability,
get_serving,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.mem_cache.cache_init_params import CacheInitParams from sglang.srt.mem_cache.cache_init_params import CacheInitParams
@@ -88,9 +92,9 @@ class HiRadixCache(RadixCache):
self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)( self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)(
self.kv_cache, self.kv_cache,
get_memory().hicache_ratio, get_memory().hicache_ratio,
server_args.hicache_size, get_memory().hicache_size,
self.page_size, self.page_size,
server_args.hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=allocator_type, allocator_type=allocator_type,
) )
elif isinstance(self.kv_cache, DSATokenToKVPool): elif isinstance(self.kv_cache, DSATokenToKVPool):
@@ -106,9 +110,9 @@ class HiRadixCache(RadixCache):
self.token_to_kv_pool_host = MLATokenToKVPoolHost( self.token_to_kv_pool_host = MLATokenToKVPoolHost(
self.kv_cache, self.kv_cache,
get_memory().hicache_ratio, get_memory().hicache_ratio,
server_args.hicache_size, get_memory().hicache_size,
self.page_size, self.page_size,
server_args.hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=allocator_type, allocator_type=allocator_type,
dcp_size=_parallel.attn_dcp_size, dcp_size=_parallel.attn_dcp_size,
dcp_rank=_parallel.attn_dcp_rank, dcp_rank=_parallel.attn_dcp_rank,
@@ -123,9 +127,9 @@ class HiRadixCache(RadixCache):
self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group) self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group)
self.pp_rank = params.pp_rank self.pp_rank = params.pp_rank
self.pp_size = params.pp_size self.pp_size = params.pp_size
self.enable_storage = server_args.hicache_storage_backend is not None self.enable_storage = get_memory().hicache_storage_backend is not None
self.enable_storage_metrics = self.enable_storage and params.enable_metrics self.enable_storage_metrics = self.enable_storage and params.enable_metrics
self.extra_metric_labels = server_args.extra_metric_labels self.extra_metric_labels = get_observability().extra_metric_labels
( (
extra_config, extra_config,
@@ -133,11 +137,11 @@ class HiRadixCache(RadixCache):
prefetch_timeout_config, prefetch_timeout_config,
hicache_storage_pass_prefix_keys, hicache_storage_pass_prefix_keys,
) = self._parse_storage_backend_extra_config( ) = self._parse_storage_backend_extra_config(
server_args.hicache_storage_backend_extra_config get_memory().hicache_storage_backend_extra_config
) )
# TODO: support more timeout check functions # TODO: support more timeout check functions
self.is_prefetch_timeout = self._prefetch_timeout_check_linear_func self.is_prefetch_timeout = self._prefetch_timeout_check_linear_func
self.prefetch_stop_policy = server_args.hicache_storage_prefetch_policy self.prefetch_stop_policy = get_memory().hicache_storage_prefetch_policy
self.load_cache_event = threading.Event() self.load_cache_event = threading.Event()
if isinstance(self.kv_cache, DSATokenToKVPool): if isinstance(self.kv_cache, DSATokenToKVPool):
@@ -174,16 +178,16 @@ class HiRadixCache(RadixCache):
attn_cp_group=self.attn_cp_group, attn_cp_group=self.attn_cp_group,
attn_tp_group=self.attn_tp_group, attn_tp_group=self.attn_tp_group,
pp_group=self.pp_group, pp_group=self.pp_group,
write_policy=server_args.hicache_write_policy, write_policy=get_memory().hicache_write_policy,
io_backend=server_args.hicache_io_backend, io_backend=get_memory().hicache_io_backend,
storage_backend=server_args.hicache_storage_backend, storage_backend=get_memory().hicache_storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=server_args.served_model_name, model_name=get_serving().served_model_name,
storage_backend_extra_config=extra_config, storage_backend_extra_config=extra_config,
enable_storage_metrics=self.enable_storage_metrics, enable_storage_metrics=self.enable_storage_metrics,
) )
self._apply_storage_runtime_config( self._apply_storage_runtime_config(
storage_backend=server_args.hicache_storage_backend, storage_backend=get_memory().hicache_storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
prefetch_timeout_config=prefetch_timeout_config, prefetch_timeout_config=prefetch_timeout_config,
hicache_storage_pass_prefix_keys=hicache_storage_pass_prefix_keys, hicache_storage_pass_prefix_keys=hicache_storage_pass_prefix_keys,
@@ -205,7 +209,7 @@ class HiRadixCache(RadixCache):
self.work_list: List[torch.distributed.Work] = [] self.work_list: List[torch.distributed.Work] = []
# todo: dynamically adjust the threshold # todo: dynamically adjust the threshold
self.write_through_threshold = ( self.write_through_threshold = (
1 if server_args.hicache_write_policy == "write_through" else 2 1 if get_memory().hicache_write_policy == "write_through" else 2
) )
self.load_back_threshold = 10 self.load_back_threshold = 10
# Detach storage backend automatically on process shutdown # Detach storage backend automatically on process shutdown
@@ -28,7 +28,11 @@ from sglang.srt.mem_cache.pool_host.mha import (
) )
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType from sglang.srt.mem_cache.unified_cache.component_type import ComponentType
from sglang.srt.runtime_context import get_memory, get_parallel from sglang.srt.runtime_context import (
get_memory,
get_parallel,
get_serving,
)
if TYPE_CHECKING: if TYPE_CHECKING:
import torch import torch
@@ -100,9 +104,9 @@ def build_kv_host_pool(
return kv_host_pool_cls( return kv_host_pool_cls(
kv_pool, kv_pool,
get_memory().hicache_ratio, get_memory().hicache_ratio,
server_args.hicache_size if host_size is None else host_size, get_memory().hicache_size if host_size is None else host_size,
page_size, page_size,
server_args.hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
pool_label=pool_label, pool_label=pool_label,
**kwargs, **kwargs,
@@ -299,8 +303,8 @@ def build_kv_only_stack(
attn_cp_group=params.attn_cp_cache_group, attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group, attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group, pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy, write_policy=get_memory().hicache_write_policy,
io_backend=server_args.hicache_io_backend, io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend, storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=model_name, model_name=model_name,
@@ -340,9 +344,9 @@ def build_hybrid_swa_stack(
) )
kv_host_size = swa_host_size = None kv_host_size = swa_host_size = None
if server_args.hicache_size > 0: if get_memory().hicache_size > 0:
kv_host_size, swa_host_size = _split_hicache_size( kv_host_size, swa_host_size = _split_hicache_size(
server_args.hicache_size, (full_kv_pool, swa_kv_pool) get_memory().hicache_size, (full_kv_pool, swa_kv_pool)
) )
host_pool_group = build_hybrid_swa_group( host_pool_group = build_hybrid_swa_group(
@@ -370,8 +374,8 @@ def build_hybrid_swa_stack(
attn_cp_group=params.attn_cp_cache_group, attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group, attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group, pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy, write_policy=get_memory().hicache_write_policy,
io_backend=server_args.hicache_io_backend, io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend, storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=model_name, model_name=model_name,
@@ -399,7 +403,7 @@ def _deepseek_v4_num_host_pages(
device_swa_pages = (kvcache.swa_size + swa_page_size - 1) // swa_page_size device_swa_pages = (kvcache.swa_size + swa_page_size - 1) // swa_page_size
if server_args.hicache_size > 0: if get_memory().hicache_size > 0:
raise ValueError( raise ValueError(
"DeepSeek V4 HiCache currently does not support --hicache-size; " "DeepSeek V4 HiCache currently does not support --hicache-size; "
"use --hicache-ratio instead." "use --hicache-ratio instead."
@@ -496,7 +500,7 @@ def build_deepseek_v4_hicache_stack(
) )
logical_host_pool = LogicalHostPool( logical_host_pool = LogicalHostPool(
num_host_pages * page_size, page_size, layout=server_args.hicache_mem_layout num_host_pages * page_size, page_size, layout=get_memory().hicache_mem_layout
) )
entries = [ entries = [
build_pool_entry( build_pool_entry(
@@ -519,7 +523,7 @@ def build_deepseek_v4_hicache_stack(
item_bytes=kvcache.swa_kv_pool.bytes_per_page_padded, item_bytes=kvcache.swa_kv_pool.bytes_per_page_padded,
num_host_pages=swa_num_host_pages, num_host_pages=swa_num_host_pages,
slot_page_size=kvcache.swa_page_size, slot_page_size=kvcache.swa_page_size,
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
) )
swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
@@ -545,7 +549,7 @@ def build_deepseek_v4_hicache_stack(
item_bytes=c4_item_bytes, item_bytes=c4_item_bytes,
num_host_pages=num_host_pages, num_host_pages=num_host_pages,
slot_page_size=page_size, slot_page_size=page_size,
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
) )
c4_indexer_host_pool = DeepSeekV4PagedHostPool( c4_indexer_host_pool = DeepSeekV4PagedHostPool(
@@ -557,7 +561,7 @@ def build_deepseek_v4_hicache_stack(
), ),
num_host_pages=num_host_pages, num_host_pages=num_host_pages,
slot_page_size=page_size, slot_page_size=page_size,
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
) )
entries.extend( entries.extend(
@@ -588,7 +592,7 @@ def build_deepseek_v4_hicache_stack(
], ],
num_host_pages=swa_num_host_pages, num_host_pages=swa_num_host_pages,
swa_page_size=kvcache.swa_page_size, swa_page_size=kvcache.swa_page_size,
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
) )
c4_indexer_state_host_pool = DeepSeekV4StateHostPool( c4_indexer_state_host_pool = DeepSeekV4StateHostPool(
@@ -599,7 +603,7 @@ def build_deepseek_v4_hicache_stack(
], ],
num_host_pages=swa_num_host_pages, num_host_pages=swa_num_host_pages,
swa_page_size=kvcache.swa_page_size, swa_page_size=kvcache.swa_page_size,
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
) )
entries.extend( entries.extend(
@@ -631,7 +635,7 @@ def build_deepseek_v4_hicache_stack(
item_bytes=c128_item_bytes, item_bytes=c128_item_bytes,
num_host_pages=num_host_pages, num_host_pages=num_host_pages,
slot_page_size=page_size, slot_page_size=page_size,
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
) )
# C128 state pool is intentionally not registered with hicache. # C128 state pool is intentionally not registered with hicache.
@@ -658,8 +662,8 @@ def build_deepseek_v4_hicache_stack(
attn_cp_group=params.attn_cp_cache_group, attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group, attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group, pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy, write_policy=get_memory().hicache_write_policy,
io_backend=server_args.hicache_io_backend, io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend, storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=model_name, model_name=model_name,
@@ -697,9 +701,9 @@ def build_hybrid_mamba_stack(
pool.full_kv_pool for pool in params.mtp_draft_device_pools pool.full_kv_pool for pool in params.mtp_draft_device_pools
) )
kv_host_size, mamba_host_size = None, 0 kv_host_size, mamba_host_size = None, 0
if server_args.hicache_size > 0: if get_memory().hicache_size > 0:
kv_host_size, mamba_host_size = _split_hicache_size( kv_host_size, mamba_host_size = _split_hicache_size(
server_args.hicache_size, (kv_pool, mamba_pool) get_memory().hicache_size, (kv_pool, mamba_pool)
) )
kv_host_pool = build_kv_host_pool( kv_host_pool = build_kv_host_pool(
kv_pool=kv_pool, kv_pool=kv_pool,
@@ -721,7 +725,7 @@ def build_hybrid_mamba_stack(
get_memory().hicache_ratio, get_memory().hicache_ratio,
mamba_host_size, mamba_host_size,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
) )
entries = [ entries = [
build_pool_entry( build_pool_entry(
@@ -754,8 +758,8 @@ def build_hybrid_mamba_stack(
attn_cp_group=params.attn_cp_cache_group, attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group, attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group, pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy, write_policy=get_memory().hicache_write_policy,
io_backend=server_args.hicache_io_backend, io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend, storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=model_name, model_name=model_name,
@@ -801,9 +805,9 @@ def build_hybrid_mamba_swa_stack(
swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
mamba_allocator = params.req_to_token_pool.mamba_allocator mamba_allocator = params.req_to_token_pool.mamba_allocator
kv_host_size, swa_host_size, mamba_host_size = None, None, 0 kv_host_size, swa_host_size, mamba_host_size = None, None, 0
if server_args.hicache_size > 0: if get_memory().hicache_size > 0:
kv_host_size, swa_host_size, mamba_host_size = _split_hicache_size( kv_host_size, swa_host_size, mamba_host_size = _split_hicache_size(
server_args.hicache_size, (full_kv_pool, swa_kv_pool, mamba_pool) get_memory().hicache_size, (full_kv_pool, swa_kv_pool, mamba_pool)
) )
kv_host_pool = build_kv_host_pool( kv_host_pool = build_kv_host_pool(
kv_pool=full_kv_pool, kv_pool=full_kv_pool,
@@ -825,8 +829,8 @@ def build_hybrid_mamba_swa_stack(
mamba_pool, mamba_pool,
get_memory().hicache_ratio, get_memory().hicache_ratio,
mamba_host_size, mamba_host_size,
allocator_type=server_args.hicache_storage_backend, allocator_type=get_memory().hicache_storage_backend,
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
) )
entries = [ entries = [
build_pool_entry( build_pool_entry(
@@ -870,8 +874,8 @@ def build_hybrid_mamba_swa_stack(
attn_cp_group=attn_cp_group, attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group, attn_tp_group=attn_tp_group,
pp_group=pp_group, pp_group=pp_group,
write_policy=server_args.hicache_write_policy, write_policy=get_memory().hicache_write_policy,
io_backend=server_args.hicache_io_backend, io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend, storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=model_name, model_name=model_name,
@@ -948,8 +952,8 @@ def build_anchor_sidecar_stack(
attn_cp_group=params.attn_cp_cache_group, attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group, attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group, pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy, write_policy=get_memory().hicache_write_policy,
io_backend=server_args.hicache_io_backend, io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend, storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=model_name, model_name=model_name,
@@ -1019,7 +1023,7 @@ def build_full_draft_pools(
pool=pool, pool=pool,
host_to_device_ratio=host_pool_group.logical_size / pool.size, host_to_device_ratio=host_pool_group.logical_size / pool.size,
page_size=controller.page_size, page_size=controller.page_size,
layout=server_args.hicache_mem_layout, layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
pool_label="draft", pool_label="draft",
) )
@@ -1045,7 +1049,7 @@ def build_full_draft_pools(
indexer_host_pool = DSAIndexerPoolHost( indexer_host_pool = DSAIndexerPoolHost(
pool, pool,
draft_host_pool, draft_host_pool,
server_args.hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
) )
specs.append( specs.append(
@@ -1502,7 +1506,7 @@ class _DsaStrategy(StackStrategy):
sidecar_host_pool_factory=lambda kv_host_pool: DSAIndexerPoolHost( sidecar_host_pool_factory=lambda kv_host_pool: DSAIndexerPoolHost(
full_kv_pool, full_kv_pool,
kv_host_pool, kv_host_pool,
server_args.hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
), ),
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
@@ -1732,7 +1736,7 @@ def attach_hybrid_pool_to_unified_cache(
storage_backend=storage_backend, storage_backend=storage_backend,
storage_backend_extra_config=storage_extra_config, storage_backend_extra_config=storage_extra_config,
prefetch_threshold=storage_prefetch_threshold, prefetch_threshold=storage_prefetch_threshold,
model_name=server_args.served_model_name, model_name=get_serving().served_model_name,
enable_storage_metrics=cache._enable_metrics_flag, enable_storage_metrics=cache._enable_metrics_flag,
) )
_apply_stack_result(cache, kvcache, params, result) _apply_stack_result(cache, kvcache, params, result)
@@ -1796,8 +1800,8 @@ def build_minimax_sparse_hicache_stack(
index_host_pool = MHATokenToKOnlyPoolHost( index_host_pool = MHATokenToKOnlyPoolHost(
index_k_pool, index_k_pool,
kv_host_pool, kv_host_pool,
server_args.hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend, allocator_type=get_memory().hicache_storage_backend,
) )
entries.append( entries.append(
build_pool_entry( build_pool_entry(
@@ -1821,8 +1825,8 @@ def build_minimax_sparse_hicache_stack(
load_cache_event=load_cache_event, load_cache_event=load_cache_event,
attn_cp_group=params.attn_cp_cache_group, attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group, attn_tp_group=params.attn_tp_cache_group,
write_policy=server_args.hicache_write_policy, write_policy=get_memory().hicache_write_policy,
io_backend=server_args.hicache_io_backend, io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend, storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=model_name, model_name=model_name,
@@ -1871,10 +1875,10 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
layer_id: layer_id for layer_id in range(main_pool.layer_num) layer_id: layer_id for layer_id in range(main_pool.layer_num)
}, },
load_cache_event=load_cache_event, load_cache_event=load_cache_event,
storage_backend=server_args.hicache_storage_backend, storage_backend=get_memory().hicache_storage_backend,
use_mla=False, use_mla=False,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=server_args.served_model_name, model_name=get_serving().served_model_name,
storage_backend_extra_config=extra_config, storage_backend_extra_config=extra_config,
enable_storage_metrics=enable_storage_metrics, enable_storage_metrics=enable_storage_metrics,
) )
@@ -1885,9 +1889,9 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
server_args=server_args, server_args=server_args,
sparse_pool=sparse_pool, sparse_pool=sparse_pool,
load_cache_event=load_cache_event, load_cache_event=load_cache_event,
storage_backend=server_args.hicache_storage_backend, storage_backend=get_memory().hicache_storage_backend,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
model_name=server_args.served_model_name, model_name=get_serving().served_model_name,
storage_backend_extra_config=extra_config, storage_backend_extra_config=extra_config,
enable_storage_metrics=enable_storage_metrics, enable_storage_metrics=enable_storage_metrics,
) )
@@ -1933,17 +1937,17 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
sidecar_pool_name=PoolName.INDEXER, sidecar_pool_name=PoolName.INDEXER,
full_layer_mapping=layer_mapping, full_layer_mapping=layer_mapping,
load_cache_event=load_cache_event, load_cache_event=load_cache_event,
storage_backend=server_args.hicache_storage_backend, storage_backend=get_memory().hicache_storage_backend,
use_mla=True, use_mla=True,
override_kv_cache_dim=kv.kv_cache_dim, override_kv_cache_dim=kv.kv_cache_dim,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
sidecar_host_pool_factory=lambda kv_host_pool: DSAIndexerPoolHost( sidecar_host_pool_factory=lambda kv_host_pool: DSAIndexerPoolHost(
kv, kv,
kv_host_pool, kv_host_pool,
server_args.hicache_mem_layout, get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
), ),
model_name=server_args.served_model_name, model_name=get_serving().served_model_name,
storage_backend_extra_config=extra_config, storage_backend_extra_config=extra_config,
enable_storage_metrics=enable_storage_metrics, enable_storage_metrics=enable_storage_metrics,
) )
@@ -22,6 +22,9 @@ from sglang.srt.mem_cache.cpp_radix_tree.radix_tree import (
TreeNodeCpp, TreeNodeCpp,
) )
from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.runtime_context import (
get_memory,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
@@ -59,7 +62,7 @@ class RadixCacheCpp(BasePrefixCache):
self.ongoing_load_back: Set[IOHandle] = set() self.ongoing_load_back: Set[IOHandle] = set()
# todo: dynamically adjust the threshold # todo: dynamically adjust the threshold
self.write_through_threshold = ( self.write_through_threshold = (
1 if server_args.hicache_write_policy == "write_through" else 2 1 if get_memory().hicache_write_policy == "write_through" else 2
) )
self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator
self.device = self.token_to_kv_pool_allocator.device self.device = self.token_to_kv_pool_allocator.device
@@ -72,7 +75,7 @@ class RadixCacheCpp(BasePrefixCache):
if params.enable_metrics: if params.enable_metrics:
self.init_metrics_collector() self.init_metrics_collector()
if not server_args.enable_hierarchical_cache: if not get_memory().enable_hierarchical_cache:
self.tree = RadixTreeCpp( self.tree = RadixTreeCpp(
disabled=self.disable, disabled=self.disable,
page_size=self.page_size, page_size=self.page_size,
@@ -170,10 +170,13 @@ from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
ensure_published, ensure_published,
get_context, get_context,
get_device,
get_exec, get_exec,
get_global_dwdp_manager, get_global_dwdp_manager,
get_lora, get_lora,
get_memory,
get_model, get_model,
get_observability,
get_parallel, get_parallel,
get_schedule, get_schedule,
get_spec, get_spec,
@@ -279,7 +282,7 @@ def resolve_draft_attention_backend(
""" """
if not is_draft_worker: if not is_draft_worker:
return None return None
return draft_attention_backend or server_args.speculative_draft_attention_backend return draft_attention_backend or get_spec().speculative_draft_attention_backend
class ModelRunner: class ModelRunner:
@@ -323,7 +326,6 @@ class ModelRunner:
# workers so they reuse target's resolved sizes (replaces legacy # workers so they reuse target's resolved sizes (replaces legacy
# `server_args._draft_pool_config` mutation hack). # `server_args._draft_pool_config` mutation hack).
self.memory_pool_config = memory_pool_config self.memory_pool_config = memory_pool_config
self.device = server_args.device
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps self.ps = ps
self.model_config = model_config self.model_config = model_config
@@ -340,6 +342,7 @@ class ModelRunner:
# Set by maybe_init_lora_manager; stays None when LoRA is off and on # Set by maybe_init_lora_manager; stays None when LoRA is off and on
# draft runners, which serve adapters' target model unadapted. # draft runners, which serve adapters' target model unadapted.
self.lora_manager: Optional[LoRAManager] = None self.lora_manager: Optional[LoRAManager] = None
self.device = get_device().device
self.draft_attention_backend = resolve_draft_attention_backend( self.draft_attention_backend = resolve_draft_attention_backend(
draft_attention_backend=draft_attention_backend, draft_attention_backend=draft_attention_backend,
server_args=server_args, server_args=server_args,
@@ -356,7 +359,7 @@ class ModelRunner:
model_config.is_multimodal_chunked_prefill_supported model_config.is_multimodal_chunked_prefill_supported
) )
self.spec_algorithm = SpeculativeAlgorithm.from_string( self.spec_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
self.capture_tail_hooks = [] self.capture_tail_hooks = []
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
@@ -367,12 +370,12 @@ class ModelRunner:
self.is_hybrid_swa_compress = model_config.is_hybrid_swa_compress self.is_hybrid_swa_compress = model_config.is_hybrid_swa_compress
self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
self.attention_chunk_size = model_config.attention_chunk_size self.attention_chunk_size = model_config.attention_chunk_size
self.enable_elastic_ep = server_args.elastic_ep_backend is not None self.enable_elastic_ep = get_exec().moe.elastic_ep_backend is not None
self.forward_pass_id = 0 self.forward_pass_id = 0
self._pending_elastic_scale_update = None self._pending_elastic_scale_update = None
self.init_new_workspace = False self.init_new_workspace = False
self.draft_model_idx = draft_model_idx self.draft_model_idx = draft_model_idx
self.enable_hisparse = server_args.enable_hisparse self.enable_hisparse = get_memory().enable_hisparse
self._sampling_observer: Optional[SamplingObserver] = None self._sampling_observer: Optional[SamplingObserver] = None
self.init_startup_observability() self.init_startup_observability()
@@ -385,7 +388,7 @@ class ModelRunner:
self.init_spec_aux_hidden_state() self.init_spec_aux_hidden_state()
# Apply the rank zero filter to logger # Apply the rank zero filter to logger
if server_args.show_time_cost: if get_observability().show_time_cost:
enable_show_time_cost() enable_show_time_cost()
misc_utils.maybe_disable_chunked_prefix_cache( misc_utils.maybe_disable_chunked_prefix_cache(
@@ -1348,7 +1351,7 @@ class ModelRunner:
else False else False
), ),
speculative_draft_attention_backend=self.draft_attention_backend, speculative_draft_attention_backend=self.draft_attention_backend,
speculative_draft_kv_cache_dtype=self.server_args.speculative_draft_kv_cache_dtype, speculative_draft_kv_cache_dtype=get_spec().speculative_draft_kv_cache_dtype,
) )
) )
# This runner's OWN resolved dtype string (target or draft). Attention # This runner's OWN resolved dtype string (target or draft). Attention
@@ -25,6 +25,11 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
trigger_init_weights_send_group_for_remote_instance_request, trigger_init_weights_send_group_for_remote_instance_request,
) )
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import (
get_exec,
get_model,
get_observability,
)
from sglang.srt.utils.common import is_npu from sglang.srt.utils.common import is_npu
from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.network import NetworkAddress
@@ -86,18 +91,16 @@ def maybe_trigger_remote_instance_nccl_send_group(
``--speculative-draft-draft-load-format`` needs its own send group, and the ``--speculative-draft-draft-load-format`` needs its own send group, and the
target's format cannot answer for it.""" target's format cannot answer for it."""
if ( if (
(load_format or server_args.load_format) == LoadFormat.REMOTE_INSTANCE load_format or get_model().load_format
and server_args.remote_instance_weight_loader_backend ) == LoadFormat.REMOTE_INSTANCE and get_model().remote_instance_weight_loader_backend == RemoteInstanceWeightLoaderBackend.NCCL:
== RemoteInstanceWeightLoaderBackend.NCCL
):
if tp_rank == 0: if tp_rank == 0:
instance_ip = NetworkAddress.resolve_host(socket.gethostname()) instance_ip = NetworkAddress.resolve_host(socket.gethostname())
t = threading.Thread( t = threading.Thread(
target=trigger_init_weights_send_group_for_remote_instance_request, target=trigger_init_weights_send_group_for_remote_instance_request,
args=( args=(
server_args.remote_instance_weight_loader_seed_instance_ip, get_model().remote_instance_weight_loader_seed_instance_ip,
server_args.remote_instance_weight_loader_seed_instance_service_port, get_model().remote_instance_weight_loader_seed_instance_service_port,
server_args.remote_instance_weight_loader_send_weights_group_ports, get_model().remote_instance_weight_loader_send_weights_group_ports,
instance_ip, instance_ip,
), ),
) )
@@ -111,12 +114,12 @@ def load_kv_cache_scales(
defaulted: a fallback to ``server_args`` would be a hidden global read for defaulted: a fallback to ``server_args`` would be a hidden global read for
any future caller that forgets to pass one.""" any future caller that forgets to pass one."""
if kv_cache_dtype == "fp8_e4m3": if kv_cache_dtype == "fp8_e4m3":
if server_args.quantization_param_path is not None: if get_model().quantization_param_path is not None:
if callable(getattr(model, "load_kv_cache_scales", None)): if callable(getattr(model, "load_kv_cache_scales", None)):
model.load_kv_cache_scales(server_args.quantization_param_path) model.load_kv_cache_scales(get_model().quantization_param_path)
logger.info( logger.info(
"Loaded KV cache scaling factors from %s", "Loaded KV cache scaling factors from %s",
server_args.quantization_param_path, get_model().quantization_param_path,
) )
else: else:
raise RuntimeError( raise RuntimeError(
@@ -154,13 +157,13 @@ def report_online_quantization(*, model, server_args: ServerArgs) -> None:
getattr(model, "quant_config", None), "quantized_layers", None getattr(model, "quant_config", None), "quantized_layers", None
) )
if ( if (
server_args.quantization is not None get_model().quantization is not None
and isinstance(quantized_layers, tuple) and isinstance(quantized_layers, tuple)
and len(quantized_layers) == 2 and len(quantized_layers) == 2
): ):
layer_types, quantized_layers_count = quantized_layers layer_types, quantized_layers_count = quantized_layers
logger.info( logger.info(
f"Online {server_args.quantization} quantization: quantized {quantized_layers_count} layers of types: {layer_types}" f"Online {get_model().quantization} quantization: quantized {quantized_layers_count} layers of types: {layer_types}"
) )
@@ -174,15 +177,15 @@ def maybe_register_debug_tensor_dump_hook(
tp_rank: int, tp_rank: int,
pp_rank: int, pp_rank: int,
) -> None: ) -> None:
if server_args.debug_tensor_dump_output_folder is not None: if get_observability().debug_tensor_dump_output_folder is not None:
dump_folder = server_args.debug_tensor_dump_output_folder dump_folder = get_observability().debug_tensor_dump_output_folder
if spec_algorithm.is_eagle(): if spec_algorithm.is_eagle():
role = "draft" if is_draft_worker else "target" role = "draft" if is_draft_worker else "target"
dump_folder = os.path.join(dump_folder, role) dump_folder = os.path.join(dump_folder, role)
register_forward_hook_for_model( register_forward_hook_for_model(
model, model,
dump_folder, dump_folder,
server_args.debug_tensor_dump_layers, get_observability().debug_tensor_dump_layers,
tp_size, tp_size,
tp_rank, tp_rank,
pp_rank, pp_rank,
@@ -203,28 +206,28 @@ def build_load_config(
from sglang.srt.configs.modelopt_config import ModelOptConfig from sglang.srt.configs.modelopt_config import ModelOptConfig
modelopt_config = ModelOptConfig( modelopt_config = ModelOptConfig(
quant=server_args.modelopt_quant, quant=get_model().modelopt_quant,
checkpoint_restore_path=server_args.modelopt_checkpoint_restore_path, checkpoint_restore_path=get_model().modelopt_checkpoint_restore_path,
checkpoint_save_path=server_args.modelopt_checkpoint_save_path, checkpoint_save_path=get_model().modelopt_checkpoint_save_path,
export_path=server_args.modelopt_export_path, export_path=get_model().modelopt_export_path,
quantize_and_serve=server_args.quantize_and_serve, quantize_and_serve=get_model().quantize_and_serve,
) )
return LoadConfig( return LoadConfig(
load_format=load_format or server_args.load_format, load_format=load_format or get_model().load_format,
download_dir=server_args.download_dir, download_dir=get_model().download_dir,
model_loader_extra_config=server_args.model_loader_extra_config, model_loader_extra_config=get_model().model_loader_extra_config,
tp_rank=tp_rank, tp_rank=tp_rank,
remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip, remote_instance_weight_loader_seed_instance_ip=get_model().remote_instance_weight_loader_seed_instance_ip,
remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port, remote_instance_weight_loader_seed_instance_service_port=get_model().remote_instance_weight_loader_seed_instance_service_port,
remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports, remote_instance_weight_loader_send_weights_group_ports=get_model().remote_instance_weight_loader_send_weights_group_ports,
remote_instance_weight_loader_backend=server_args.remote_instance_weight_loader_backend, remote_instance_weight_loader_backend=get_model().remote_instance_weight_loader_backend,
remote_instance_weight_loader_transfer_engine=remote_instance_weight_transporter_engine, remote_instance_weight_loader_transfer_engine=remote_instance_weight_transporter_engine,
remote_instance_weight_loader_transfer_engine_session_id=remote_instance_weight_transporter_session_id, remote_instance_weight_loader_transfer_engine_session_id=remote_instance_weight_transporter_session_id,
modelexpress_url=server_args.modelexpress_url, modelexpress_url=server_args.modelexpress_url,
modelexpress_transport=server_args.modelexpress_transport, modelexpress_transport=server_args.modelexpress_transport,
modelopt_config=modelopt_config, modelopt_config=modelopt_config,
rl_quant_profile=server_args.rl_quant_profile, rl_quant_profile=get_model().rl_quant_profile,
draft_model_idx=draft_model_idx, draft_model_idx=draft_model_idx,
weight_cache_mode=weight_cache_mode, weight_cache_mode=weight_cache_mode,
weight_cache_socket=weight_cache_socket, weight_cache_socket=weight_cache_socket,
@@ -246,7 +249,7 @@ def maybe_enable_ipc_weight_cache(
the format swap is guarded on ``!= IPC_CACHE`` so a second call (e.g. a the format swap is guarded on ``!= IPC_CACHE`` so a second call (e.g. a
weight reload) can't overwrite the captured fallback format. weight reload) can't overwrite the captured fallback format.
""" """
if server_args.weight_cache_mode == "off": if get_model().weight_cache_mode == "off":
return return
if load_config.load_format != LoadFormat.IPC_CACHE: if load_config.load_format != LoadFormat.IPC_CACHE:
@@ -278,13 +281,13 @@ def load_model_with_memory_saver(
# Remove monkey_patch when linear.py quant remove dependencies with vllm # Remove monkey_patch when linear.py quant remove dependencies with vllm
monkey_patch_vllm_parallel_state() monkey_patch_vllm_parallel_state()
enable_cpu_backup = server_args.enable_weights_cpu_backup or ( enable_cpu_backup = get_exec().features.enable_weights_cpu_backup or (
is_draft_worker and server_args.enable_draft_weights_cpu_backup is_draft_worker and get_exec().features.enable_draft_weights_cpu_backup
) )
# In zero-copy IPC mode, the weights are shared with the daemon via # In zero-copy IPC mode, the weights are shared with the daemon via
# CUDA IPC and must not be offloaded/reloaded by the memory saver. # CUDA IPC and must not be offloaded/reloaded by the memory saver.
is_ipc_zero_copy = server_args.weight_cache_mode != "off" is_ipc_zero_copy = get_model().weight_cache_mode != "off"
if is_ipc_zero_copy and enable_cpu_backup: if is_ipc_zero_copy and enable_cpu_backup:
logger.warning( logger.warning(
"[ModelRunner] Disabling weights CPU backup in zero-copy IPC mode — " "[ModelRunner] Disabling weights CPU backup in zero-copy IPC mode — "
@@ -6,6 +6,11 @@ from typing import TYPE_CHECKING, Any, Optional
import msgspec import msgspec
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.runtime_context import (
configured_tp_size,
get_model,
get_spec,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -75,13 +80,13 @@ def _resolve_eagle_aux_hidden_state(
if ( if (
(spec_algorithm.is_eagle() or spec_algorithm.is_standalone()) (spec_algorithm.is_eagle() or spec_algorithm.is_standalone())
and not is_draft_worker and not is_draft_worker
and server_args.speculative_draft_model_path and get_spec().speculative_draft_model_path
): ):
# Load draft config to get layer count for KV cache sizing # Load draft config to get layer count for KV cache sizing
draft_model_config = ModelConfig.from_server_args( draft_model_config = ModelConfig.from_server_args(
server_args, server_args,
model_path=server_args.speculative_draft_model_path, model_path=get_spec().speculative_draft_model_path,
model_revision=server_args.speculative_draft_model_revision, model_revision=get_spec().speculative_draft_model_revision,
is_draft_model=True, is_draft_model=True,
) )
num_nextn_predict_layers = draft_model_config.num_nextn_predict_layers num_nextn_predict_layers = draft_model_config.num_nextn_predict_layers
@@ -134,8 +139,8 @@ def _resolve_dflash_aux_hidden_state(
# Select target layers to capture for building draft context features. # Select target layers to capture for building draft context features.
draft_model_config = ModelConfig.from_server_args( draft_model_config = ModelConfig.from_server_args(
server_args, server_args,
model_path=(server_args.speculative_draft_model_path), model_path=(get_spec().speculative_draft_model_path),
model_revision=server_args.speculative_draft_model_revision, model_revision=get_spec().speculative_draft_model_revision,
is_draft_model=True, is_draft_model=True,
) )
dflash_draft_config = parse_dflash_draft_config( dflash_draft_config = parse_dflash_draft_config(
@@ -221,23 +226,23 @@ def _resolve_dflash_draft_cell_size(
try: try:
_, draft_kv_cache_dtype = configure_kv_cache_dtype( _, draft_kv_cache_dtype = configure_kv_cache_dtype(
server_args_kv_cache_dtype=server_args.kv_cache_dtype, server_args_kv_cache_dtype=get_model().kv_cache_dtype,
speculative_draft_kv_cache_dtype=( speculative_draft_kv_cache_dtype=(
server_args.speculative_draft_kv_cache_dtype get_spec().speculative_draft_kv_cache_dtype
), ),
model=None, model=None,
model_dtype=draft_model_config.dtype, model_dtype=draft_model_config.dtype,
is_draft_worker=True, is_draft_worker=True,
is_dflash=True, is_dflash=True,
speculative_draft_attention_backend=( speculative_draft_attention_backend=(
server_args.speculative_draft_attention_backend get_spec().speculative_draft_attention_backend
), ),
) )
return dflash_draft_cell_size_per_token( return dflash_draft_cell_size_per_token(
draft_model_config=draft_model_config, draft_model_config=draft_model_config,
draft_num_layers=draft_num_layers, draft_num_layers=draft_num_layers,
draft_kv_cache_dtype=draft_kv_cache_dtype, draft_kv_cache_dtype=draft_kv_cache_dtype,
tp_size=server_args.tp_size, tp_size=configured_tp_size(),
) )
except Exception as e: # noqa: BLE001 except Exception as e: # noqa: BLE001
logger.warning( logger.warning(
@@ -20,7 +20,18 @@ from sglang.srt.model_loader.weight_utils import (
CheckpointFilePrefetchHandle, CheckpointFilePrefetchHandle,
) )
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 (
configured_attn_cp_size,
configured_dcp_size,
configured_pp_size,
configured_tp_size,
get_device,
get_exec,
get_lora,
get_model,
get_parallel,
get_spec,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
@@ -97,37 +108,37 @@ class StartupWeightLoadOptions:
server_args: ServerArgs, server_args: ServerArgs,
is_draft_worker: bool, is_draft_worker: bool,
) -> StartupWeightLoadOptions: ) -> StartupWeightLoadOptions:
cuda_graph_config = server_args.cuda_graph_config cuda_graph_config = get_exec().graph.cuda_graph_config
cuda_graph_enabled = any( cuda_graph_enabled = any(
getattr(cuda_graph_config, phase).backend != Backend.DISABLED getattr(cuda_graph_config, phase).backend != Backend.DISABLED
for phase in Phase.ALL for phase in Phase.ALL
) )
return cls( return cls(
device=server_args.device, device=get_device().device,
is_cuda_platform=current_platform.is_cuda(), is_cuda_platform=current_platform.is_cuda(),
cuda_graph_enabled=cuda_graph_enabled, cuda_graph_enabled=cuda_graph_enabled,
prefill_cuda_graph_backend=cuda_graph_config.prefill.backend, prefill_cuda_graph_backend=cuda_graph_config.prefill.backend,
is_draft_worker=is_draft_worker, is_draft_worker=is_draft_worker,
speculative_algorithm=server_args.speculative_algorithm, speculative_algorithm=get_spec().speculative_algorithm,
tp_size=server_args.tp_size, tp_size=configured_tp_size(),
attn_cp_size=server_args.attn_cp_size, attn_cp_size=configured_attn_cp_size(),
dcp_size=server_args.dcp_size, dcp_size=configured_dcp_size(),
pp_size=server_args.pp_size, pp_size=configured_pp_size(),
dp_size=get_parallel().dp_size, dp_size=get_parallel().dp_size,
ep_size=get_parallel().ep_size, ep_size=get_parallel().ep_size,
cpu_offload_gb=server_args.cpu_offload_gb, cpu_offload_gb=get_exec().offload.cpu_offload_gb,
offload_group_size=server_args.offload_group_size, offload_group_size=get_exec().offload.offload_group_size,
enable_memory_saver=server_args.enable_memory_saver, enable_memory_saver=get_exec().features.enable_memory_saver,
enable_weights_cpu_backup=server_args.enable_weights_cpu_backup, enable_weights_cpu_backup=get_exec().features.enable_weights_cpu_backup,
enable_lora=server_args.enable_lora, enable_lora=get_lora().enable_lora,
has_lora_paths=bool(server_args.lora_paths), has_lora_paths=bool(get_lora().lora_paths),
weight_loader_disable_mmap=server_args.weight_loader_disable_mmap, weight_loader_disable_mmap=get_model().weight_loader_disable_mmap,
weight_loader_drop_cache_after_load=( weight_loader_drop_cache_after_load=(
server_args.weight_loader_drop_cache_after_load get_model().weight_loader_drop_cache_after_load
), ),
has_custom_weight_loader=bool(server_args.custom_weight_loader), has_custom_weight_loader=bool(get_model().custom_weight_loader),
enable_torch_compile=server_args.enable_torch_compile, enable_torch_compile=get_exec().graph.enable_torch_compile,
prefetch_num_threads=server_args.weight_loader_prefetch_num_threads, prefetch_num_threads=get_model().weight_loader_prefetch_num_threads,
) )
@@ -26,6 +26,12 @@ from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional, Set, Union
from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.utils import exponential_buckets, generate_buckets from sglang.srt.observability.utils import exponential_buckets, generate_buckets
from sglang.srt.runtime_context import (
get_disagg,
get_observability,
get_schedule,
get_serving,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import get_bool_env_var from sglang.srt.utils import get_bool_env_var
from sglang.srt.utils.gauge_histogram import GaugeHistogram from sglang.srt.utils.gauge_histogram import GaugeHistogram
@@ -957,7 +963,7 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
# ================================================================= # =================================================================
# Prefill delayer # Prefill delayer
# ================================================================= # =================================================================
max_delay = server_args.prefill_delayer_max_delay_passes max_delay = get_schedule().prefill_delayer_max_delay_passes
self.prefill_delayer_wait_forward_passes = Histogram( self.prefill_delayer_wait_forward_passes = Histogram(
name="sglang:prefill_delayer_wait_forward_passes", name="sglang:prefill_delayer_wait_forward_passes",
documentation="Histogram of forward passes waited by prefill delayer.", documentation="Histogram of forward passes waited by prefill delayer.",
@@ -966,7 +972,7 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
set( set(
x x
for x in ( for x in (
server_args.prefill_delayer_forward_passes_buckets get_schedule().prefill_delayer_forward_passes_buckets
or [5, 20, 50, 100, 200] or [5, 20, 50, 100, 200]
) )
if x < max_delay if x < max_delay
@@ -981,7 +987,7 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
labelnames=labels.keys(), labelnames=labels.keys(),
buckets=sorted( buckets=sorted(
set( set(
server_args.prefill_delayer_wait_seconds_buckets get_schedule().prefill_delayer_wait_seconds_buckets
or [1, 2, 5, 10, 20, 50, 100, 200, 500] or [1, 2, 5, 10, 20, 50, 100, 200, 500]
) )
# Need bucket "<=0" for zero-delay cases # Need bucket "<=0" for zero-delay cases
@@ -1077,13 +1083,14 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
enable_lora: bool, enable_lora: bool,
enable_hierarchical_cache: bool, enable_hierarchical_cache: bool,
) -> SchedulerMetricsCollectorContext: ) -> SchedulerMetricsCollectorContext:
enable_metrics = server_args.enable_metrics enable_metrics = get_observability().enable_metrics
is_stats_logging_rank = ps.attn_tp_rank == 0 is_stats_logging_rank = ps.attn_tp_rank == 0
current_scheduler_metrics_enabled = enable_metrics and ( current_scheduler_metrics_enabled = enable_metrics and (
is_stats_logging_rank or server_args.enable_metrics_for_all_schedulers is_stats_logging_rank
or get_observability().enable_metrics_for_all_schedulers
) )
enable_kv_cache_events = bool( enable_kv_cache_events = bool(
server_args.kv_events_config get_observability().kv_events_config
and ps.pp_rank == 0 and ps.pp_rank == 0
and ps.attn_tp_rank == 0 and ps.attn_tp_rank == 0
and ps.attn_cp_rank == 0 and ps.attn_cp_rank == 0
@@ -1091,10 +1098,10 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
collector: Optional[SchedulerMetricsCollector] = None collector: Optional[SchedulerMetricsCollector] = None
if enable_metrics: if enable_metrics:
engine_type = DisaggregationMode.to_engine_type( engine_type = DisaggregationMode.to_engine_type(
server_args.disaggregation_mode get_disagg().disaggregation_mode
) )
labels = { labels = {
"model_name": server_args.served_model_name, "model_name": get_serving().served_model_name,
"engine_type": engine_type, "engine_type": engine_type,
"tp_rank": tp_rank, "tp_rank": tp_rank,
"pp_rank": pp_rank, "pp_rank": pp_rank,
@@ -1104,8 +1111,8 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
labels["priority"] = "" labels["priority"] = ""
if dp_rank is not None: if dp_rank is not None:
labels["dp_rank"] = dp_rank labels["dp_rank"] = dp_rank
if server_args.extra_metric_labels: if get_observability().extra_metric_labels:
labels.update(server_args.extra_metric_labels) labels.update(get_observability().extra_metric_labels)
scheduler_collector_cls = resolve_collector_class( scheduler_collector_cls = resolve_collector_class(
server_args, STAT_LOGGER_ROLE_SCHEDULER, cls server_args, STAT_LOGGER_ROLE_SCHEDULER, cls
) )
@@ -1113,7 +1120,7 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
labels=labels, labels=labels,
enable_lora=enable_lora, enable_lora=enable_lora,
enable_hierarchical_cache=enable_hierarchical_cache, enable_hierarchical_cache=enable_hierarchical_cache,
enable_streaming_session=server_args.enable_streaming_session, enable_streaming_session=get_serving().enable_streaming_session,
server_args=server_args, server_args=server_args,
) )
return SchedulerMetricsCollectorContext( return SchedulerMetricsCollectorContext(
@@ -1568,7 +1575,7 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
documentation="Histogram of prompt token length.", documentation="Histogram of prompt token length.",
labelnames=labels.keys(), labelnames=labels.keys(),
buckets=generate_buckets( buckets=generate_buckets(
server_args.prompt_tokens_buckets, default_bucket_prompt_tokens get_observability().prompt_tokens_buckets, default_bucket_prompt_tokens
), ),
) )
self.uncached_prompt_tokens_histogram = Histogram( self.uncached_prompt_tokens_histogram = Histogram(
@@ -1576,7 +1583,7 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
documentation="Histogram of uncached (compute) prompt token length.", documentation="Histogram of uncached (compute) prompt token length.",
labelnames=labels.keys(), labelnames=labels.keys(),
buckets=generate_buckets( buckets=generate_buckets(
server_args.prompt_tokens_buckets, default_bucket_prompt_tokens get_observability().prompt_tokens_buckets, default_bucket_prompt_tokens
), ),
) )
self.generation_tokens_histogram = Histogram( self.generation_tokens_histogram = Histogram(
@@ -1584,7 +1591,7 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
documentation="Histogram of generation token length.", documentation="Histogram of generation token length.",
labelnames=labels.keys(), labelnames=labels.keys(),
buckets=generate_buckets( buckets=generate_buckets(
server_args.generation_tokens_buckets, get_observability().generation_tokens_buckets,
default_bucket_prompt_tokens, default_bucket_prompt_tokens,
), ),
) )
+27 -5
View File
@@ -678,18 +678,22 @@ class _ConfigBag:
def _build_config_bags(server_args: Any) -> dict: def _build_config_bags(server_args: Any) -> dict:
"""Snapshot resolved ``server_args`` into the namespace bag tree, driven by """Snapshot the resolution result into the namespace bag tree, driven by
the ``NS(...)`` metadata on the dataclass fields. Returns the ``NS(...)`` metadata on the dataclass fields. Each leaf comes from
``resolution_result`` -- the declaration if resolution made one, else what
the caller supplied -- rather than from the field, which carries the same
value only while declarations still materialize. Returns
``{top_level_name: _ConfigBag}``, arbitrarily nested (``exec.moe.eplb.…``). ``{top_level_name: _ConfigBag}``, arbitrarily nested (``exec.moe.eplb.…``).
Only dataclass fields carry ``NS`` markers, so derived properties/methods are Only dataclass fields carry ``NS`` markers, so derived properties/methods are
naturally excluded (they stay on the bag). A name used as both a leaf and a naturally excluded (they stay on the bag). A name used as both a leaf and a
subgroup at the same level is a hard error — no silent shadowing.""" subgroup at the same level is a hard error — no silent shadowing."""
from sglang.srt.arg_groups.arg_utils import namespace_of from sglang.srt.arg_groups.arg_utils import namespace_of
from sglang.srt.arg_groups.overrides import resolution_result
_MISSING = object() _MISSING = object()
tops: dict = {} tops: dict = {}
for field, path in namespace_of(type(server_args)).items(): for field, path in namespace_of(type(server_args)).items():
value = getattr(server_args, field, _MISSING) value = resolution_result(server_args, field, _MISSING)
if value is _MISSING: if value is _MISSING:
# Every NS-declared field is a dataclass field, so a resolved config # Every NS-declared field is a dataclass field, so a resolved config
# always carries it; a miss means a malformed/partial config object # always carries it; a miss means a malformed/partial config object
@@ -1028,7 +1032,10 @@ class _ServerArgsOverride:
self._prev_publish_role = ctx._publish_role self._prev_publish_role = ctx._publish_role
self._prev_parallel_config = ctx.parallel._config self._prev_parallel_config = ctx.parallel._config
self._prev_capture = ctx.flags.capture.enable_torch_compile self._prev_capture = ctx.flags.capture.enable_torch_compile
from sglang.srt.arg_groups.overrides import _apply_fields from sglang.srt.arg_groups.overrides import (
_apply_fields,
declare_resolution,
)
server_args = ServerArgs(model_path="dummy") server_args = ServerArgs(model_path="dummy")
# Underscore names seed private property caches (the strict guard # Underscore names seed private property caches (the strict guard
@@ -1040,7 +1047,18 @@ class _ServerArgsOverride:
raise ValueError( raise ValueError(
f"override_server_args: unknown ServerArgs field(s): {sorted(unknown)}" f"override_server_args: unknown ServerArgs field(s): {sorted(unknown)}"
) )
_apply_fields(server_args, self._fields) # Declared so the projection sees it.
# Underscore names are not fields at all (they seed private property
# caches), so they stay a direct write.
declared = {
name: value for name, value in self._fields.items() if name[0] != "_"
}
if declared:
declare_resolution(server_args, "override_server_args", **declared)
_apply_fields(
server_args,
{name: value for name, value in self._fields.items() if name[0] == "_"},
)
# The dummy boundary skips materialization, which would leave the # The dummy boundary skips materialization, which would leave the
# strict mutation guard unarmed on the published object — mark it # strict mutation guard unarmed on the published object — mark it
# materialized so bare post-publish writes raise like they do on a # materialized so bare post-publish writes raise like they do on a
@@ -1662,6 +1680,10 @@ def configured_attn_cp_size() -> int:
return _configured_parallel("attn_cp_size") return _configured_parallel("attn_cp_size")
def configured_dcp_size() -> int:
return _configured_parallel("dcp_size")
def is_ep_joiner() -> bool: def is_ep_joiner() -> bool:
"""True in a process launched as an elastic-EP joiner (scale or recover). """True in a process launched as an elastic-EP joiner (scale or recover).
+14 -2
View File
@@ -42,6 +42,7 @@ from sglang.srt.arg_groups.argparse_actions import (
) )
from sglang.srt.arg_groups.overrides import ( from sglang.srt.arg_groups.overrides import (
attention_backends_of, attention_backends_of,
declare_direct_writes,
mamba_extra_buffer_lazy_of, mamba_extra_buffer_lazy_of,
mamba_extra_buffer_of, mamba_extra_buffer_of,
remote_instance_transfer_engine_of, remote_instance_transfer_engine_of,
@@ -3677,6 +3678,12 @@ class ServerArgs:
belong in the helper or signal that the helper should be split. belong in the helper or signal that the helper should be split.
""" """
# What the caller asked for, before any handler runs; this plus the
# stash is the resolution result the projection reads.
self._raw_input = {
field.name: getattr(self, field.name) for field in dataclasses.fields(self)
}
# Declaration stash for the override/post-process passes. Set before any # Declaration stash for the override/post-process passes. Set before any
# short-circuit (none/dummy model paths) so run_post_process_pass and # short-circuit (none/dummy model paths) so run_post_process_pass and
# direct handler invocations can rely on it even when # direct handler invocations can rely on it even when
@@ -3742,8 +3749,13 @@ class ServerArgs:
self._handle_mps_backends() self._handle_mps_backends()
self._handle_xpu_backends() self._handle_xpu_backends()
# Allow OOT platform plugins to apply server args defaults. # OOT platform plugins set fields directly (an interface this tree
current_platform.apply_server_args_defaults(self) # does not own); the diff records what they applied.
declare_direct_writes(
self,
f"platform:{current_platform.device_name}",
current_platform.apply_server_args_defaults,
)
# Get GPU memory capacity, which is a common dependency for several configuration steps. # Get GPU memory capacity, which is a common dependency for several configuration steps.
gpu_mem = get_device_memory_capacity(self.device) gpu_mem = get_device_memory_capacity(self.device)
@@ -33,7 +33,12 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
compute_position, compute_position,
) )
from sglang.srt.runtime_context import get_exec, get_schedule, mamba_track_grid from sglang.srt.runtime_context import (
get_exec,
get_schedule,
get_spec,
mamba_track_grid,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
from sglang.srt.speculative.dflash_info import DFlashVerifyInput from sglang.srt.speculative.dflash_info import DFlashVerifyInput
@@ -279,9 +284,7 @@ class DFlashWorkerV2(BaseSpecWorker):
self._need_mamba_verify_commit = False self._need_mamba_verify_commit = False
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
# Normalized in arg_groups.speculative_hook.handle_speculative_decoding. # Normalized in arg_groups.speculative_hook.handle_speculative_decoding.
self.draft_window_size: Optional[int] = ( self.draft_window_size: Optional[int] = get_spec().speculative_draft_window_size
server_args.speculative_draft_window_size
)
self.use_compact_draft_cache = self.draft_window_size is not None self.use_compact_draft_cache = self.draft_window_size is not None
self.device = target_worker.device self.device = target_worker.device
@@ -305,11 +308,11 @@ class DFlashWorkerV2(BaseSpecWorker):
draft_config = parse_dflash_draft_config( draft_config = parse_dflash_draft_config(
draft_hf_config=self.draft_model_runner.model_config.hf_config draft_hf_config=self.draft_model_runner.model_config.hf_config
) )
if server_args.speculative_num_draft_tokens is None: if get_spec().speculative_num_draft_tokens is None:
# Should not happen (ServerArgs should have inferred it), but keep a fallback. # Should not happen (ServerArgs should have inferred it), but keep a fallback.
self.block_size = int(draft_config.resolve_block_size(default=16)) self.block_size = int(draft_config.resolve_block_size(default=16))
else: else:
self.block_size = int(server_args.speculative_num_draft_tokens) self.block_size = int(get_spec().speculative_num_draft_tokens)
model_block_size = draft_config.block_size model_block_size = draft_config.block_size
if model_block_size is None: if model_block_size is None:
model_block_size = getattr(self.draft_model, "block_size", None) model_block_size = getattr(self.draft_model, "block_size", None)
@@ -6,6 +6,10 @@ from typing import TYPE_CHECKING, Any, List, Optional
import msgspec import msgspec
from sglang.srt.runtime_context import (
get_model,
get_spec,
)
from sglang.srt.speculative.dflash_utils import parse_dflash_draft_config from sglang.srt.speculative.dflash_utils import parse_dflash_draft_config
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -25,15 +29,15 @@ def draft_is_deepseek_v4(*, server_args: ServerArgs) -> bool:
from sglang.srt.configs.model_config import is_deepseek_v4 from sglang.srt.configs.model_config import is_deepseek_v4
from sglang.srt.utils.hf_transformers_utils import get_config from sglang.srt.utils.hf_transformers_utils import get_config
draft_model_path = server_args.speculative_draft_model_path draft_model_path = get_spec().speculative_draft_model_path
if not draft_model_path: if not draft_model_path:
return False return False
draft_hf_config = get_config( draft_hf_config = get_config(
draft_model_path, draft_model_path,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=get_model().trust_remote_code,
revision=server_args.speculative_draft_model_revision, revision=get_spec().speculative_draft_model_revision,
model_override_args=json.loads(server_args.json_model_override_args), model_override_args=json.loads(get_model().json_model_override_args),
model_config_parser=server_args.model_config_parser, model_config_parser=get_model().model_config_parser,
) )
return draft_hf_config is not None and is_deepseek_v4(draft_hf_config) return draft_hf_config is not None and is_deepseek_v4(draft_hf_config)
@@ -126,14 +130,23 @@ def resolve_runtime_config(
def read_draft_checkpoint_gamma(*, server_args: ServerArgs) -> Optional[int]: def read_draft_checkpoint_gamma(*, server_args: ServerArgs) -> Optional[int]:
"""Load the draft checkpoint's hf config and read its DSpark gamma """Load the draft checkpoint's hf config and read its DSpark gamma
(block_size). Raises on config-load failure; callers pick the fallback.""" (block_size). Raises on config-load failure; callers pick the fallback.
Reads the *resolving* configuration, not the bags: the speculative hook
calls this from inside resolution, where no bag exists yet -- and the
caller swallows exceptions, so a bag read here does not fail loudly, it
silently drops the checkpoint's gamma and the cross-check with
`--speculative-num-draft-tokens` along with it.
"""
from sglang.srt.arg_groups.overrides import resolved_view
from sglang.srt.utils.hf_transformers_utils import get_config from sglang.srt.utils.hf_transformers_utils import get_config
resolving = resolved_view(server_args)
draft_hf_config = get_config( draft_hf_config = get_config(
server_args.speculative_draft_model_path, resolving.speculative_draft_model_path,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=resolving.trust_remote_code,
revision=server_args.speculative_draft_model_revision, revision=resolving.speculative_draft_model_revision,
model_override_args=json.loads(server_args.json_model_override_args), model_override_args=json.loads(resolving.json_model_override_args),
) )
return parse_dspark_draft_config(draft_hf_config=draft_hf_config).resolve_gamma( return parse_dspark_draft_config(draft_hf_config=draft_hf_config).resolve_gamma(
default=None default=None
@@ -21,6 +21,7 @@ from sglang.srt.model_executor.forward_batch_info import (
compute_position, compute_position,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_disagg,
get_exec, get_exec,
get_parallel, get_parallel,
get_schedule, get_schedule,
@@ -106,7 +107,7 @@ class DSparkWorkerV2(BaseSpecWorker):
self._draft_dp_context_enabled = ( self._draft_dp_context_enabled = (
get_parallel().enable_dp_attention and not self._draft_is_moe get_parallel().enable_dp_attention and not self._draft_is_moe
) )
self._is_pd_prefill = server_args.disaggregation_mode == "prefill" self._is_pd_prefill = get_disagg().disaggregation_mode == "prefill"
self._decode_graph_allowed = ( self._decode_graph_allowed = (
not get_exec().graph.disable_cuda_graph and not self._is_pd_prefill not get_exec().graph.disable_cuda_graph and not self._is_pd_prefill
) )
@@ -156,7 +157,7 @@ class DSparkWorkerV2(BaseSpecWorker):
self._target_is_mambaish = mambaish_config(target_model_config) is not None self._target_is_mambaish = mambaish_config(target_model_config) is not None
runtime_config = resolve_runtime_config( runtime_config = resolve_runtime_config(
draft_hf_config=self.draft_model_runner.model_config.hf_config, draft_hf_config=self.draft_model_runner.model_config.hf_config,
speculative_num_draft_tokens=server_args.speculative_num_draft_tokens, speculative_num_draft_tokens=get_spec().speculative_num_draft_tokens,
target_vocab_size=int(target_embed_rows), target_vocab_size=int(target_embed_rows),
) )
self.gamma = runtime_config.gamma self.gamma = runtime_config.gamma
@@ -47,6 +47,7 @@ from sglang.srt.model_executor.runner import (
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_context, get_context,
get_device,
get_exec, get_exec,
get_model, get_model,
get_parallel, get_parallel,
@@ -145,14 +146,14 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.target_worker = target_worker self.target_worker = target_worker
# Args for easy access # Args for easy access
self.device = server_args.device self.device = get_device().device
self.topk = server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
if get_spec().speculative_use_rejection_sampling: if get_spec().speculative_use_rejection_sampling:
assert self.topk == 1, "Chain speculative sampling supports only topk=1" assert self.topk == 1, "Chain speculative sampling supports only topk=1"
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_algorithm = SpeculativeAlgorithm.from_string( self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
self._rebuild_topk1_chain_buffers() self._rebuild_topk1_chain_buffers()
@@ -1060,16 +1061,16 @@ class EAGLEWorkerV2(BaseSpecWorker):
# Parse arguments # Parse arguments
self.server_args = server_args self.server_args = server_args
self.topk = server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = server_args.device self.device = get_device().device
self._target_worker = target_worker self._target_worker = target_worker
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string( self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
self._draft_worker = EagleDraftWorker( self._draft_worker = EagleDraftWorker(
@@ -1082,10 +1083,10 @@ class EAGLEWorkerV2(BaseSpecWorker):
# Adaptive speculative # Adaptive speculative
self.adaptive_controller: Optional[AdaptiveController] = None self.adaptive_controller: Optional[AdaptiveController] = None
if server_args.speculative_adaptive: if get_spec().speculative_adaptive:
self.adaptive_controller = AdaptiveController( self.adaptive_controller = AdaptiveController(
self, self,
config_path=server_args.speculative_adaptive_config, config_path=get_spec().speculative_adaptive_config,
) )
# Some dummy tensors # Some dummy tensors
@@ -46,6 +46,7 @@ from sglang.srt.model_executor.forward_context import ForwardContext, forward_co
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
attention_backends, attention_backends,
get_device,
get_parallel, get_parallel,
get_schedule, get_schedule,
get_spec, get_spec,
@@ -107,16 +108,16 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
EagleDraftWorkerBase.__init__(self) EagleDraftWorkerBase.__init__(self)
self.server_args = server_args self.server_args = server_args
self.topk = server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = server_args.device self.device = get_device().device
self.target_worker = target_worker self.target_worker = target_worker
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string( self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
assert self.speculative_algorithm.is_frozen_kv_mtp(), ( assert self.speculative_algorithm.is_frozen_kv_mtp(), (
"FrozenKVMTPDraftWorker should only be instantiated for " "FrozenKVMTPDraftWorker should only be instantiated for "
@@ -695,16 +696,16 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
# an EagleDraftWorker (with its own draft KV pool). The frozen draft owns # an EagleDraftWorker (with its own draft KV pool). The frozen draft owns
# no KV, so we mirror the relevant setup and build a FrozenKVMTPDraftWorker. # no KV, so we mirror the relevant setup and build a FrozenKVMTPDraftWorker.
self.server_args = server_args self.server_args = server_args
self.topk = server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = server_args.device self.device = get_device().device
self._target_worker = target_worker self._target_worker = target_worker
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string( self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
self.req_to_token_pool, self.token_to_kv_pool_allocator = ( self.req_to_token_pool, self.token_to_kv_pool_allocator = (
@@ -720,7 +721,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
# Frozen MTP does not wire the adaptive controller yet. # Frozen MTP does not wire the adaptive controller yet.
assert ( assert (
not server_args.speculative_adaptive not get_spec().speculative_adaptive
), "Frozen-KV MTP does not support adaptive speculative decoding yet." ), "Frozen-KV MTP does not support adaptive speculative decoding yet."
self.adaptive_controller = None self.adaptive_controller = None
@@ -43,7 +43,12 @@ from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode, CaptureHiddenMode,
ForwardBatch, ForwardBatch,
) )
from sglang.srt.runtime_context import get_parallel, get_schedule from sglang.srt.runtime_context import (
get_device,
get_parallel,
get_schedule,
get_spec,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
from sglang.srt.speculative.draft_utils import DraftBackendFactory from sglang.srt.speculative.draft_utils import DraftBackendFactory
@@ -128,22 +133,22 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.model_config = target_worker.model_config self.model_config = target_worker.model_config
# Args for easy access # Args for easy access
self.device = server_args.device self.device = get_device().device
self.topk = server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
# Leviathan/Chen rejection sampling (temp>0): the draft samples X ~ q and # Leviathan/Chen rejection sampling (temp>0): the draft samples X ~ q and
# provides q so the verify accepts iff coin*q < p and resamples the residual. # provides q so the verify accepts iff coin*q < p and resamples the residual.
# Single-CG runner samples in-graph (_sample_draft_proposal); per-step # Single-CG runner samples in-graph (_sample_draft_proposal); per-step
# runner samples worker-side between replays. # runner samples worker-side between replays.
self.use_rejection_sampling = server_args.speculative_use_rejection_sampling self.use_rejection_sampling = get_spec().speculative_use_rejection_sampling
assert self.speculative_num_draft_tokens == self.speculative_num_steps + 1, ( assert self.speculative_num_draft_tokens == self.speculative_num_steps + 1, (
"multi-layer EAGLE requires speculative_num_draft_tokens == " "multi-layer EAGLE requires speculative_num_draft_tokens == "
"speculative_num_steps + 1, " "speculative_num_steps + 1, "
f"got {self.speculative_num_draft_tokens} and {self.speculative_num_steps}" f"got {self.speculative_num_draft_tokens} and {self.speculative_num_steps}"
) )
self.speculative_algorithm = SpeculativeAlgorithm.from_string( self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
self._rebuild_topk1_chain_buffers() self._rebuild_topk1_chain_buffers()
@@ -928,15 +933,15 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
# Parse arguments # Parse arguments
self.server_args = server_args self.server_args = server_args
self.topk = server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = server_args.device self.device = get_device().device
self._target_worker = target_worker self._target_worker = target_worker
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string( self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
self._draft_worker = MultiLayerEagleDraftWorker( self._draft_worker = MultiLayerEagleDraftWorker(
+23 -19
View File
@@ -15,7 +15,11 @@ from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.req_time_stats import set_time_batch from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.runtime_context import get_schedule from sglang.srt.runtime_context import (
get_device,
get_schedule,
get_spec,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus
@@ -88,19 +92,19 @@ class NGRAMWorker(BaseSpecWorker):
super().__init__() super().__init__()
self.server_args = server_args self.server_args = server_args
self.enable_overlap = not server_args.disable_overlap_schedule self.enable_overlap = not get_schedule().disable_overlap_schedule
self._target_worker = target_worker self._target_worker = target_worker
self.model_runner = target_worker.model_runner self.model_runner = target_worker.model_runner
self.tp_rank = ps.tp_rank self.tp_rank = ps.tp_rank
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
self.draft_token_num: int = server_args.speculative_num_draft_tokens self.draft_token_num: int = get_spec().speculative_num_draft_tokens
self.max_trie_depth: int = server_args.speculative_ngram_max_trie_depth self.max_trie_depth: int = get_spec().speculative_ngram_max_trie_depth
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.topk = server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
# req_to_token_pool / token_to_kv_pool_allocator are set in # req_to_token_pool / token_to_kv_pool_allocator are set in
# alloc_memory_pool(), after the target pools are allocated. # alloc_memory_pool(), after the target pools are allocated.
self.device = server_args.device self.device = get_device().device
self.adaptive_controller = None self.adaptive_controller = None
# rids of the last decode batch; used to erase corpus match state for # rids of the last decode batch; used to erase corpus match state for
@@ -109,26 +113,26 @@ class NGRAMWorker(BaseSpecWorker):
self.grammar_tree_host: Optional[tuple] = None self.grammar_tree_host: Optional[tuple] = None
self.ngram_corpus = NgramCorpus( self.ngram_corpus = NgramCorpus(
min_bfs_breadth=server_args.speculative_ngram_min_bfs_breadth, min_bfs_breadth=get_spec().speculative_ngram_min_bfs_breadth,
max_bfs_breadth=server_args.speculative_ngram_max_bfs_breadth, max_bfs_breadth=get_spec().speculative_ngram_max_bfs_breadth,
match_type=server_args.speculative_ngram_match_type, match_type=get_spec().speculative_ngram_match_type,
capacity=server_args.speculative_ngram_capacity, capacity=get_spec().speculative_ngram_capacity,
max_trie_depth=server_args.speculative_ngram_max_trie_depth, max_trie_depth=get_spec().speculative_ngram_max_trie_depth,
draft_token_num=server_args.speculative_num_draft_tokens, draft_token_num=get_spec().speculative_num_draft_tokens,
external_sam_budget=server_args.speculative_ngram_external_sam_budget, external_sam_budget=get_spec().speculative_ngram_external_sam_budget,
external_corpus_max_tokens=server_args.speculative_ngram_external_corpus_max_tokens, external_corpus_max_tokens=get_spec().speculative_ngram_external_corpus_max_tokens,
) )
if server_args.speculative_ngram_external_corpus_path is not None: if get_spec().speculative_ngram_external_corpus_path is not None:
from sglang.srt.speculative.cpp_ngram.external_corpus import ( from sglang.srt.speculative.cpp_ngram.external_corpus import (
iter_external_corpus_chunks, iter_external_corpus_chunks,
) )
corpus_path = server_args.speculative_ngram_external_corpus_path corpus_path = get_spec().speculative_ngram_external_corpus_path
chunks = list( chunks = list(
iter_external_corpus_chunks( iter_external_corpus_chunks(
corpus_path, corpus_path,
target_worker.tokenizer, target_worker.tokenizer,
server_args.speculative_ngram_external_corpus_max_tokens, get_spec().speculative_ngram_external_corpus_max_tokens,
) )
) )
loaded = self.add_external_corpus(corpus_path, chunks) loaded = self.add_external_corpus(corpus_path, chunks)
@@ -10,7 +10,12 @@ from sglang.srt.layers.moe.utils import (
speculative_moe_backend_context, speculative_moe_backend_context,
) )
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.runtime_context import get_parallel, get_schedule from sglang.srt.runtime_context import (
get_device,
get_parallel,
get_schedule,
get_spec,
)
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.adaptive_runtime_state import ( from sglang.srt.speculative.adaptive_runtime_state import (
AdaptiveController, AdaptiveController,
@@ -53,12 +58,12 @@ class StandaloneDraftWorker(EagleDraftWorker):
self.target_worker = target_worker self.target_worker = target_worker
# Args for easy access # Args for easy access
self.device = server_args.device self.device = get_device().device
self.topk = server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_algorithm = SpeculativeAlgorithm.from_string( self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
self._rebuild_topk1_chain_buffers() self._rebuild_topk1_chain_buffers()
@@ -158,15 +163,15 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
# Parse arguments # Parse arguments
self.server_args = server_args self.server_args = server_args
self.topk = server_args.speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = server_args.device self.device = get_device().device
self._target_worker = target_worker self._target_worker = target_worker
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string( self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm get_spec().speculative_algorithm
) )
# Create our custom draft worker that doesn't share embeddings/lm_head # Create our custom draft worker that doesn't share embeddings/lm_head
+14
View File
@@ -2105,6 +2105,20 @@ def server_args_variant(server_args, **fields):
} }
if unknown: if unknown:
raise ValueError(f"unknown ServerArgs field(s): {sorted(unknown)}") raise ValueError(f"unknown ServerArgs field(s): {sorted(unknown)}")
# Reach the stash as well as the fields (the bags project from raw input
# + declarations); through `object` because the copy keeps its read-only
# guard.
stash = getattr(variant, "_resolved_overrides", None)
if stash is None:
stash = []
object.__setattr__(variant, "_resolved_overrides", stash)
declared = {
name: value
for name, value in fields.items()
if name in cls.__dataclass_fields__
}
if declared:
stash.append(("server_args_variant", dict(declared)))
for name, value in fields.items(): for name, value in fields.items():
object.__setattr__(variant, name, value) object.__setattr__(variant, name, value)
return variant return variant
@@ -1,7 +1,6 @@
from __future__ import annotations from __future__ import annotations
import unittest import unittest
from types import SimpleNamespace
from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities
from sglang.srt.model_executor.cuda_graph_config import ( from sglang.srt.model_executor.cuda_graph_config import (
@@ -9,6 +8,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
CudaGraphConfig, CudaGraphConfig,
PhaseConfig, PhaseConfig,
) )
from sglang.srt.runtime_context import get_context
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
@@ -16,28 +16,33 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestComputeLaunchCapacities(CustomTestCase): class TestComputeLaunchCapacities(CustomTestCase):
@staticmethod
def _make_server_args(*, max_bs: int) -> SimpleNamespace:
return SimpleNamespace(
cuda_graph_config=CudaGraphConfig(
decode=PhaseConfig(backend=Backend.FULL, max_bs=max_bs)
),
speculative_num_draft_tokens=0,
chunked_prefill_size=None,
max_prefill_tokens=128,
)
@staticmethod
def _from_args( def _from_args(
self,
*, *,
max_bs: int, max_bs: int,
max_seq_len: int, max_seq_len: int,
max_total_num_tokens: int | None = None, max_total_num_tokens: int | None = None,
speculative_num_draft_tokens: int | None = 0,
) -> CanaryLaunchCapacities: ) -> CanaryLaunchCapacities:
"""`from_args` reads the published configuration, so publish one.
Handing it a stand-in object stopped meaning anything when the reads
moved to the config bags: the parameter was ignored and the values
under test came from whatever the process had published.
"""
if max_total_num_tokens is None: if max_total_num_tokens is None:
max_total_num_tokens = max_bs * max_seq_len max_total_num_tokens = max_bs * max_seq_len
override = get_context().override_server_args(
cuda_graph_config=CudaGraphConfig(
decode=PhaseConfig(backend=Backend.FULL, max_bs=max_bs)
),
speculative_num_draft_tokens=speculative_num_draft_tokens,
chunked_prefill_size=None,
max_prefill_tokens=128,
)
override.install()
self.addCleanup(override.restore)
return CanaryLaunchCapacities.from_args( return CanaryLaunchCapacities.from_args(
server_args=TestComputeLaunchCapacities._make_server_args(max_bs=max_bs),
req_to_token_pool_size=max_bs, req_to_token_pool_size=max_bs,
max_seq_len_per_req=max_seq_len, max_seq_len_per_req=max_seq_len,
pool_slot_count=max_total_num_tokens, pool_slot_count=max_total_num_tokens,
@@ -60,14 +65,11 @@ class TestComputeLaunchCapacities(CustomTestCase):
def test_from_args_treats_missing_speculative_draft_tokens_as_zero(self) -> None: def test_from_args_treats_missing_speculative_draft_tokens_as_zero(self) -> None:
"""per_forward_write_entry_capacity is floored by max_prefill_tokens when batch * tokens_per_req is smaller.""" """per_forward_write_entry_capacity is floored by max_prefill_tokens when batch * tokens_per_req is smaller."""
server_args = self._make_server_args(max_bs=2) capacities = self._from_args(
server_args.speculative_num_draft_tokens = None max_bs=2,
max_seq_len=32,
capacities = CanaryLaunchCapacities.from_args( max_total_num_tokens=64,
server_args=server_args, speculative_num_draft_tokens=None,
req_to_token_pool_size=2,
max_seq_len_per_req=32,
pool_slot_count=64,
) )
self.assertEqual(capacities.per_forward_write_entry_capacity, 128) self.assertEqual(capacities.per_forward_write_entry_capacity, 128)
@@ -84,12 +86,7 @@ class TestComputeLaunchCapacities(CustomTestCase):
def test_from_args_rejects_empty_pool_capacity(self) -> None: def test_from_args_rejects_empty_pool_capacity(self) -> None:
"""Verify derived launch capacities reject invalid pool sizing.""" """Verify derived launch capacities reject invalid pool sizing."""
with self.assertRaisesRegex(ValueError, "pool_slot_count"): with self.assertRaisesRegex(ValueError, "pool_slot_count"):
CanaryLaunchCapacities.from_args( self._from_args(max_bs=1, max_seq_len=1, max_total_num_tokens=0)
server_args=self._make_server_args(max_bs=1),
req_to_token_pool_size=1,
max_seq_len_per_req=1,
pool_slot_count=0,
)
if __name__ == "__main__": if __name__ == "__main__":
@@ -2,6 +2,8 @@ import unittest
from types import SimpleNamespace from types import SimpleNamespace
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs
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
@@ -17,14 +19,21 @@ def _req(output_len: int, input_len: int = 8, priority=None):
def _args(policy: str = "length", low_first: bool = False): def _args(policy: str = "length", low_first: bool = False):
return SimpleNamespace( """The retraction order reads the schedule bag, so the policy has to be
published rather than handed in."""
return ServerArgs(
model_path="dummy",
retraction_policy=policy, retraction_policy=policy,
schedule_low_priority_values_first=low_first, schedule_low_priority_values_first=low_first,
) )
def _order(reqs, args): def _order(reqs, args):
return ScheduleBatch._get_decode_retraction_order(reqs, args) publish(args, role="test")
try:
return ScheduleBatch._get_decode_retraction_order(reqs)
finally:
reset_context()
class TestRetractionOrder(CustomTestCase): class TestRetractionOrder(CustomTestCase):
@@ -20,6 +20,7 @@ from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.managers.scheduler_components.pool_stats_observer import PoolStats from sglang.srt.managers.scheduler_components.pool_stats_observer import PoolStats
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 publish, reset_context
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
register_cpu_ci(est_time=15, suite="base-a-test-cpu") register_cpu_ci(est_time=15, suite="base-a-test-cpu")
@@ -27,6 +28,15 @@ register_cpu_ci(est_time=9, suite="base-c-test-cpu")
class TestSchedulerPauseGeneration(unittest.TestCase): class TestSchedulerPauseGeneration(unittest.TestCase):
def setUp(self):
# The scheduler runs after its process publishes; retraction reads the
# disaggregation and schedule bags rather than the record it is handed.
from sglang.srt.server_args import ServerArgs
super().setUp()
publish(ServerArgs(model_path="dummy"), role="test")
self.addCleanup(reset_context)
def _new_scheduler(self) -> Scheduler: def _new_scheduler(self) -> Scheduler:
scheduler = Scheduler.__new__(Scheduler) scheduler = Scheduler.__new__(Scheduler)
scheduler._engine_paused = False scheduler._engine_paused = False
@@ -15,6 +15,7 @@ see.
""" """
import ast import ast
import copy
import dataclasses import dataclasses
import json import json
import os import os
@@ -22,8 +23,11 @@ import pathlib
import shutil import shutil
import tempfile import tempfile
import unittest import unittest
import unittest.mock
import sglang import sglang
from sglang.srt import server_args as server_args_module
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
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
@@ -203,6 +207,36 @@ def _stash_overlay(server_args):
return overlay return overlay
def _live_topology_leaves():
"""Names `ParallelContext` serves from the live topology, not the config.
Read out of the class: each shadowed name arrives as `self._v("<name>",
<getter>)`. Inferring them from "did the read raise" is wrong -- it only
raises while the process groups are missing, so in a process where an
earlier test built them the property answers the *live* size and a leaf
check reads it as a config mismatch (`parallel.tp_size: bag=1
resolution=2`). Whether they are shadowed is a property of the class, not
of the process.
"""
tree = ast.parse((_SRT / "runtime_context.py").read_text(encoding="utf-8-sig"))
parallel = next(
node
for node in ast.walk(tree)
if isinstance(node, ast.ClassDef) and node.name == "ParallelContext"
)
names = set()
for node in ast.walk(parallel):
if (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "_v"
and node.args
and isinstance(node.args[0], ast.Constant)
):
names.add(node.args[0].value)
return frozenset(names)
class TestResolutionDeclarations(CustomTestCase): class TestResolutionDeclarations(CustomTestCase):
def setUp(self): def setUp(self):
# Resolution writes environment variables, and those outlive the # Resolution writes environment variables, and those outlive the
@@ -275,6 +309,224 @@ class TestResolutionDeclarations(CustomTestCase):
+ "\n ".join(unexplained), + "\n ".join(unexplained),
) )
def test_the_projection_input_is_the_resolved_configuration(self):
"""What the bags are built from equals what the record ends up holding.
The projection reads `raw input + declarations` rather than the
fields, so that it keeps working when the declarations stop
materializing. While they still do, the two have to agree leaf for
leaf -- a difference means the projection would publish something the
record does not say, which is the failure this whole transition is
meant to avoid.
"""
from sglang.srt.arg_groups.arg_utils import namespace_of
from sglang.srt.arg_groups.overrides import resolution_result
differences = []
for shape in _SHAPES:
server_args = self._resolve(shape)
for field in namespace_of(type(server_args)):
projected = resolution_result(server_args, field)
on_record = getattr(server_args, field)
if projected != on_record:
differences.append(
f"{shape} -> {field}: projection={projected!r} "
f"record={on_record!r}"
)
self.assertEqual(
differences,
[],
"the projection and the record disagree about a config leaf:\n "
+ "\n ".join(differences),
)
def test_every_published_leaf_is_what_resolution_decided(self):
"""One hop further than the check above: the leaf a reader reads.
The projection's *input* agreeing with the record says nothing about
the last hop: whether the leaf is reachable through the path the
metadata declares, and whether it carries the resolved value once it
is. Both sides here come from that metadata, so this cannot tell that
a field is assigned to the *wrong* group -- the readers are the
independent source for that, and
`test_server_args_namespaces.py::test_the_readers_agree_with_the_namespace_metadata`
is where the two are compared.
"""
import sglang.srt.runtime_context as runtime_context
from sglang.srt.arg_groups.arg_utils import namespace_of
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.runtime_context import publish, reset_context
mapping = namespace_of(ServerArgs)
self.assertGreater(len(mapping), 400, "the namespace mapping collapsed")
shadowed = _live_topology_leaves()
self.assertGreaterEqual(
shadowed
& {
"tp_size",
"pp_size",
"moe_dp_size",
"attn_cp_size",
"dcp_size",
},
{"tp_size", "pp_size", "moe_dp_size", "attn_cp_size", "dcp_size"},
"a parallel size stopped being served from the live topology; if it "
"is a plain config leaf now, it belongs in the comparison below",
)
compared = 0
unreachable, mismatched = [], []
for shape in _SHAPES:
self.addCleanup(reset_context)
server_args = self._resolve(shape)
publish(server_args, role="scheduler")
for field, path in mapping.items():
if field in shadowed:
# Served from the process groups by design; `configured_*()`
# is what answers with the configured value, and
# test_launch_path_reads_configured_sizes pins that.
continue
groups = path.split(".")
accessor = getattr(runtime_context, f"get_{groups[0]}", None)
if accessor is None:
unreachable.append(f"no get_{groups[0]}() for {path}.{field}")
continue
node = accessor()
try:
for group in groups[1:]:
node = getattr(node, group)
leaf = getattr(node, field)
except Exception as exc:
unreachable.append(f"{path}.{field}: {type(exc).__name__}: {exc}")
continue
decided = resolution_result(server_args, field)
compared += 1
if leaf is not decided and leaf != decided:
mismatched.append(
f"{shape} -> {path}.{field}: bag={leaf!r} resolution={decided!r}"
)
reset_context()
self.assertEqual(
unreachable,
[],
"these leaves are mapped to a namespace that cannot serve them, so "
"a reader following the mapping raises:\n " + "\n ".join(unreachable),
)
self.assertEqual(
mismatched,
[],
"the published leaf and the resolution result disagree:\n "
+ "\n ".join(mismatched),
)
self.assertGreater(
compared, 2000, f"only {compared} leaves were compared; the walk broke"
)
def test_a_child_that_received_the_record_publishes_the_same_bags(self):
"""A forked worker gets the record by pickle, and re-projects from it.
Every process publishes, so a child's bags are only right if the
declarations travelled with the object -- and the gate has to hold on
the far side, or the child re-runs handlers over their own output. The
parent's bags are the reference: this is the multi-process half of the
projection, and nothing else exercises it.
"""
import pickle
import sglang.srt.runtime_context as runtime_context
from sglang.srt.arg_groups.arg_utils import namespace_of
from sglang.srt.runtime_context import publish, reset_context
mapping = namespace_of(ServerArgs)
def leaves():
out = {}
for field, path in mapping.items():
groups = path.split(".")
accessor = getattr(runtime_context, f"get_{groups[0]}", None)
if accessor is None:
continue
node = accessor()
try:
for group in groups[1:]:
node = getattr(node, group)
out[f"{path}.{field}"] = repr(getattr(node, field))
except Exception:
continue
return out
for shape in _SHAPES:
self.addCleanup(reset_context)
parent = self._resolve(shape)
publish(parent, role="scheduler")
expected = leaves()
blob = pickle.dumps(parent)
reset_context()
child = pickle.loads(blob)
entered = []
original = ServerArgs._run_resolution_pipeline
def counted(self, _original=original):
entered.append(1)
return _original(self)
with unittest.mock.patch.object(
ServerArgs, "_run_resolution_pipeline", counted
):
publish(child, role="scheduler")
self.assertEqual(
entered,
[],
f"{shape}: the child resolved again, so its handlers ran over "
"the parent's output",
)
differences = {
key: (expected[key], value)
for key, value in leaves().items()
if expected.get(key) != value
}
self.assertEqual(
differences,
{},
f"{shape}: the child published different values than the "
f"parent: {differences}",
)
reset_context()
def test_late_resolution_reaches_the_projection(self):
"""Resolution staged after `__post_init__` is still resolution.
The parser detection and the LoRA normalization run at launcher stage --
they need a tokenizer, a chat template, an adapter directory -- and they
write through `declare_late_resolution`. If those writes only reached
the fields, the bags would describe the *unresolved* value: a server
launched with `--reasoning-parser auto` would advertise and apply
`auto` after detection had already replaced it.
A real model path, not the dummy one: a dummy record never materializes,
so its `resolve_once` re-runs and re-snapshots the raw input from
already-late-resolved fields, which hides exactly this.
"""
from sglang.srt.arg_groups.overrides import declare_late_resolution
from sglang.srt.runtime_context import get_serving, publish, reset_context
server_args = self._resolve({"reasoning_parser": "auto"})
self.addCleanup(reset_context)
declare_late_resolution(
server_args, "template-detection", reasoning_parser="qwen3"
)
self.assertEqual(
resolution_result(server_args, "reasoning_parser"),
"qwen3",
"the projection still reports what the caller asked for, so the "
"bags would publish an unresolved parser",
)
publish(server_args, role="tokenizer")
self.assertEqual(get_serving().reasoning_parser, "qwen3")
self.assertEqual(server_args.reasoning_parser, get_serving().reasoning_parser)
def test_the_stash_agrees_with_the_fields_it_declared(self): def test_the_stash_agrees_with_the_fields_it_declared(self):
mismatches = [] mismatches = []
for shape in _SHAPES: for shape in _SHAPES:
@@ -362,6 +614,90 @@ class TestResolutionDeclarations(CustomTestCase):
+ "\n ".join(inversions), + "\n ".join(inversions),
) )
def test_a_nested_resolution_decision_reaches_the_bags(self):
"""Resolution also decides *inside* a declared object.
The graph sizing writes `cuda_graph_config.decode.max_bs` through the
object the parse step declared -- no field is assigned, so nothing
records it. It reaches the bags because the stash holds that same
object; a copy taken when it was declared would publish the `None` the
parse step declared while the process runs with a real batch size.
"""
from sglang.srt.runtime_context import get_exec, publish, reset_context
server_args = self._resolve({"disaggregation_mode": "prefill"})
self.addCleanup(reset_context)
# Snapshot before publishing: the bag serves the very object the record
# holds, so comparing them after the fact compares an object with
# itself and passes however the projection behaves.
expected = copy.deepcopy(server_args.cuda_graph_config)
publish(server_args, role="scheduler")
published = get_exec().graph.cuda_graph_config
resolved = expected
self.assertIsNotNone(
published.decode.max_bs,
"the published graph config carries the batch size the parse step "
"declared, not the one the sizing handler decided",
)
self.assertEqual(
(
published.decode.max_bs,
published.decode.backend,
published.prefill.max_bs,
published.prefill.backend,
),
(
resolved.decode.max_bs,
resolved.decode.backend,
resolved.prefill.max_bs,
resolved.prefill.backend,
),
"the bags and the record disagree about the graph configuration, "
"so a decision made inside the declared object was dropped",
)
def test_every_platform_hook_that_takes_the_record_is_captured(self):
"""A second out-of-tree config hook must not arrive uncaptured.
`apply_server_args_defaults` is the one method on the platform
interface that is handed the record, and its implementations live in
other distributions -- no source scan of this tree can see what they
write, so the pipeline diffs the record across the call instead. A new
hook of the same shape would be invisible again, and this is what
notices. Derived from the interface rather than listed: a rename keeps
working, an addition fails.
"""
interface = _SRT / "platforms" / "interface.py"
tree = ast.parse(interface.read_text(encoding="utf-8-sig"))
taking_the_record = set()
for node in ast.walk(tree):
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
arguments = node.args
names = [
arg.arg
for arg in arguments.posonlyargs + arguments.args + arguments.kwonlyargs
]
if any(name == "server_args" or name.endswith("_args") for name in names):
taking_the_record.add(node.name)
self.assertEqual(
taking_the_record,
{"apply_server_args_defaults"},
"the platform interface hands the startup record to a method this "
"test does not know about; either it only reads, or its writes need "
"capturing like apply_server_args_defaults",
)
pipeline = (_SRT / "server_args.py").read_text(encoding="utf-8-sig")
for hook in sorted(taking_the_record):
self.assertIn(
f"current_platform.{hook},",
pipeline,
f"{hook} is called directly instead of through the write "
"capture, so an out-of-tree plugin's defaults would be dropped "
"by the projection",
)
def test_the_shapes_reach_the_fields_they_are_meant_to(self): def test_the_shapes_reach_the_fields_they_are_meant_to(self):
"""A green agreement check over an empty stash would prove nothing.""" """A green agreement check over an empty stash would prove nothing."""
declared = set() declared = set()
@@ -376,6 +712,41 @@ class TestResolutionDeclarations(CustomTestCase):
+ "\n ".join(missing), + "\n ".join(missing),
) )
def test_a_platform_plugin_default_reaches_the_projection(self):
"""An out-of-tree platform writes the fields; the diff declares them.
The plugin interface is not ours to convert -- implementations live in
other distributions -- so its writes are captured rather than declared.
Without the capture the projection falls through to the raw snapshot,
which was taken before the plugin ran, and publishes the value the
plugin overrode.
"""
# The pipeline asks the platform other questions on the way through
# (whether it is out of tree, whether it supports piecewise capture),
# and which of those it reaches depends on the host.
class _Plugin(type(server_args_module.current_platform)):
device_name = "oot"
def apply_server_args_defaults(self, server_args):
server_args.attention_backend = "triton"
server_args.schedule_conservativeness = 0.5
with unittest.mock.patch.object(
server_args_module, "current_platform", _Plugin()
):
server_args = self._resolve({})
self.assertEqual(
(
resolution_result(server_args, "attention_backend"),
resolution_result(server_args, "schedule_conservativeness"),
),
("triton", 0.5),
"the platform plugin's defaults did not reach the resolution "
"result, so the projection publishes what the operator passed "
"instead of what the platform decided",
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -104,6 +104,54 @@ _CONFIGURED_SIZE_CALL_SITES = {
("srt/managers/scheduler.py", "configured_attn_cp_size"): ( ("srt/managers/scheduler.py", "configured_attn_cp_size"): (
"same pre-distributed-init arithmetic in configure_scheduler_process" "same pre-distributed-init arithmetic in configure_scheduler_process"
), ),
("srt/managers/scheduler.py", "configured_dcp_size"): (
"same pre-distributed-init arithmetic in configure_scheduler_process"
),
("srt/disaggregation/common/conn.py", "configured_pp_size"): (
"the bootstrap connection is built by the KV manager on the transfer "
"path, which the CPU-only conn tests exercise without ever starting "
"torch.distributed"
),
("srt/elastic_ep/elastic_ep.py", "configured_tp_size"): (
"the joiner's rank window is computed against the size the process was "
"configured with, not the size of the group it is about to join"
),
("srt/elastic_ep/expert_backup_manager.py", "configured_tp_size"): (
"the backup server counts the clients it expects to report in, which "
"is how many the launch configured -- the live group is what they are "
"still joining"
),
(
"srt/model_executor/model_runner_components/startup_weight_load.py",
"configured_tp_size",
): (
"the load options are assembled in ModelRunner.__init__ for a runner "
"that may be a draft, whose groups are the target's; the configured "
"sizes are what the record answered before"
),
(
"srt/model_executor/model_runner_components/startup_weight_load.py",
"configured_pp_size",
): ("same options object, same reason"),
(
"srt/model_executor/model_runner_components/startup_weight_load.py",
"configured_attn_cp_size",
): ("same options object, same reason"),
(
"srt/model_executor/model_runner_components/startup_weight_load.py",
"configured_dcp_size",
): ("same options object, same reason"),
(
"srt/model_executor/model_runner_components/spec_aux_hidden_state.py",
"configured_tp_size",
): (
"the draft KV bytes/token estimate sizes the memory pool before the "
"draft runner exists, so its shard count is configuration"
),
("srt/eplb/expert_location.py", "configured_tp_size"): (
"the elastic-EP joiner window, used to size the expert layout: the "
"size the process was configured with, not the group it is joining"
),
("srt/utils/cuda_vmm_transport_utils.py", "configured_tp_size"): ( ("srt/utils/cuda_vmm_transport_utils.py", "configured_tp_size"): (
"the consumer count is configured fan-out arithmetic (tp_size // " "the consumer count is configured fan-out arithmetic (tp_size // "
"dp_size), which is what the record answered before" "dp_size), which is what the record answered before"
@@ -28,7 +28,7 @@ _LIVE_SHADOWED = {
"pp_size": "configured_pp_size()", "pp_size": "configured_pp_size()",
"moe_dp_size": "configured_moe_dp_size()", "moe_dp_size": "configured_moe_dp_size()",
"attn_cp_size": "configured_attn_cp_size()", "attn_cp_size": "configured_attn_cp_size()",
"dcp_size": "a configured accessor (none exists yet; add one beside configured_pp_size)", "dcp_size": "configured_dcp_size()",
} }
# Launch paths that decide how many children to spawn are derived below # Launch paths that decide how many children to spawn are derived below
@@ -234,6 +234,119 @@ def _launch_paths():
class TestLaunchPathsReadConfiguredSizes(CustomTestCase): class TestLaunchPathsReadConfiguredSizes(CustomTestCase):
def test_configured_sizes_hold_when_the_live_topology_disagrees(self):
"""The other direction: groups exist and answer something else.
The check above proves nobody reads a live size too early. It says
nothing about what `configured_*()` returns once the groups *are* up
and answering a different number -- which is not hypothetical: elastic
EP scales the live topology away from what the operator configured, and
that divergence is the entire reason these five helpers exist. With
only the early-read direction covered, a helper that quietly delegated
to the live property would look correct.
"""
import json
import os
import tempfile
from unittest.mock import patch
from sglang.srt.runtime_context import (
configured_attn_cp_size,
configured_dcp_size,
configured_moe_dp_size,
configured_pp_size,
configured_tp_size,
get_parallel,
publish,
reset_context,
)
from sglang.srt.server_args import ServerArgs
directory = tempfile.mkdtemp(prefix="configured_sizes_")
with open(os.path.join(directory, "config.json"), "w") as handle:
json.dump(
{
"architectures": ["LlamaForCausalLM"],
"model_type": "llama",
"hidden_size": 16,
"intermediate_size": 32,
"num_attention_heads": 2,
"num_key_value_heads": 2,
"num_hidden_layers": 2,
"vocab_size": 128,
"max_position_embeddings": 2048,
},
handle,
)
# No resolve_once() here: `tp_size` is raw input, so the configured
# value is 2 either way.
server_args = ServerArgs(model_path=directory, device="cuda", tp_size=2)
self.addCleanup(reset_context)
publish(server_args, role="scheduler")
# The live getter behind each property, read out of ParallelContext
# rather than listed here.
context_source = ast.parse(
(_PACKAGE_ROOT / "srt" / "runtime_context.py").read_text(
encoding="utf-8-sig"
)
)
parallel_class = next(
node
for node in ast.walk(context_source)
if isinstance(node, ast.ClassDef) and node.name == "ParallelContext"
)
live_getter = {}
for method in parallel_class.body:
if not isinstance(method, ast.FunctionDef):
continue
for call in ast.walk(method):
if not (
isinstance(call, ast.Call)
and isinstance(call.func, ast.Attribute)
and call.func.attr == "_v"
and call.args
and isinstance(call.args[0], ast.Constant)
):
continue
getter = call.args[1]
if isinstance(getter, ast.Attribute):
live_getter[call.args[0].value] = getter.attr
state = "sglang.srt.distributed.parallel_state"
helpers = {
"tp_size": configured_tp_size,
"pp_size": configured_pp_size,
"moe_dp_size": configured_moe_dp_size,
"attn_cp_size": configured_attn_cp_size,
"dcp_size": configured_dcp_size,
}
missing = sorted(set(helpers) - set(live_getter))
self.assertEqual(
missing,
[],
f"these sizes no longer have a live property to diverge from: {missing}",
)
cases = tuple(
(name, helper, f"{state}.{live_getter[name]}")
for name, helper in helpers.items()
)
for name, helper, target in cases:
with self.subTest(size=name):
configured = helper()
with patch(target, return_value=configured + 41):
self.assertEqual(
get_parallel().__getattribute__(name),
configured + 41,
f"{name} no longer follows the live topology",
)
self.assertEqual(
helper(),
configured,
f"configured_{name}() followed the live topology instead "
"of the published configuration",
)
reset_context()
def test_no_live_topology_read_before_distributed_init(self): def test_no_live_topology_read_before_distributed_init(self):
offenders = [] offenders = []
for rel, tree in _launch_paths(): for rel, tree in _launch_paths():
@@ -47,6 +47,156 @@ def _field_names():
class TestServerArgsNamespaces(CustomTestCase): class TestServerArgsNamespaces(CustomTestCase):
def test_no_module_shadows_a_bag_accessor(self):
"""An accessor name bound twice in one module is a silent wrong read.
This has happened twice. Once a module imported `get_model` from the
context and a same-named helper from elsewhere, and once `get_device`
-- which names three different things in this tree: the bag accessor,
the device-string utility, and a platform method. The second import
wins, the converted line calls the wrong callable, and the failure is
an AttributeError on whichever branch reaches it, which for a
per-pass recorder or a specific accelerator can be none of the ones a
CPU suite runs. Nothing else notices; a name scan looks fine.
"""
import ast
import collections
import pathlib as _pathlib
import sglang
srt = _pathlib.Path(sglang.__file__).resolve().parent / "srt"
context_module = ast.parse(
(srt / "runtime_context.py").read_text(encoding="utf-8-sig")
)
accessors = {
node.name
for node in context_module.body
if isinstance(node, ast.FunctionDef)
and (node.name.startswith("get_") or node.name.startswith("configured_"))
}
self.assertGreater(len(accessors), 20, "the accessor derivation broke")
shadowed = []
for path in sorted(srt.rglob("*.py")):
if path.name == "runtime_context.py":
continue
try:
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
except SyntaxError:
self.fail(f"unparsable module in the census: {path}")
bindings = collections.defaultdict(set)
for node in ast.walk(tree):
if isinstance(node, ast.ImportFrom):
origin = node.module or ""
kind = (
"context"
if origin.endswith("runtime_context")
else f"{origin or '.'}"
)
for alias in node.names:
bindings[alias.asname or alias.name].add((kind, node.lineno))
elif isinstance(node, ast.Import):
for alias in node.names:
bindings[(alias.asname or alias.name).split(".")[0]].add(
("import", node.lineno)
)
elif isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
bindings[node.name].add(("def", node.lineno))
elif isinstance(node, ast.Assign):
for target in node.targets:
if isinstance(target, ast.Name):
bindings[target.id].add(("assign", node.lineno))
for name, where in bindings.items():
if name not in accessors:
continue
kinds = {kind for kind, _ in where}
# The same accessor imported from the context more than once
# (module level plus a lazy import inside a function) is one
# object under one name; a *different* origin is the hazard.
if "context" in kinds and kinds - {"context"}:
shadowed.append(
f"{path.relative_to(srt)}: {name} <- "
+ ", ".join(
f"{k}@{l}" for k, l in sorted(where, key=lambda w: w[1])
)
)
self.assertEqual(
shadowed,
[],
"a bag accessor shares its name with another binding in the same "
"module, so the converted reads call whichever import came last; "
"alias one of them:\n " + "\n ".join(shadowed),
)
def test_the_readers_agree_with_the_namespace_metadata(self):
"""Two independent sources say where a leaf lives; they must match.
The metadata is one source and the ~2000 hand-written reads
(`get_schedule().chunked_prefill_size`) are the other. Checking the
projection against the metadata cannot catch a field assigned to the
wrong group -- both sides come from the same marker, so the check is
true by construction. The readers are written by hand, so a
disagreement means one of the two is wrong, and every reader on the
losing side raises `has no leaf/subgroup` at runtime on whichever
branch reaches it first.
"""
import ast
import pathlib as _pathlib
import sglang
srt = _pathlib.Path(sglang.__file__).resolve().parent / "srt"
mapping = namespace_of(ServerArgs)
accessors = {
f"get_{group}" for group in {p.split(".")[0] for p in mapping.values()}
}
sites = 0
disagreements = []
for path in sorted(srt.rglob("*.py")):
try:
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
except SyntaxError:
self.fail(f"unparsable module in the census: {path}")
for node in ast.walk(tree):
if not isinstance(node, ast.Attribute):
continue
chain, cursor = [], node
while isinstance(cursor, ast.Attribute):
chain.append(cursor.attr)
cursor = cursor.value
if not (
isinstance(cursor, ast.Call)
and isinstance(cursor.func, ast.Name)
and cursor.func.id in accessors
):
continue
chain.reverse()
field = chain[-1]
if field not in mapping:
continue
sites += 1
read = [cursor.func.id[len("get_") :]] + chain[:-1]
if mapping[field].split(".") != read:
disagreements.append(
f"{path.relative_to(srt)}:{node.lineno} reads "
f"{'.'.join(read)}.{field}, metadata says "
f"{mapping[field]}.{field}"
)
self.assertEqual(
disagreements,
[],
"a reader and the namespace metadata disagree about where a leaf "
"lives; one of them is wrong:\n " + "\n ".join(disagreements),
)
self.assertGreater(
sites,
1500,
f"only {sites} bag reads were matched; the scan broke and this "
"check stopped covering anything",
)
def test_every_field_has_a_namespace(self): def test_every_field_has_a_namespace(self):
nsmap = namespace_of(ServerArgs) nsmap = namespace_of(ServerArgs)
missing = sorted(_field_names() - set(nsmap)) missing = sorted(_field_names() - set(nsmap))
@@ -129,6 +129,15 @@ _ENV_MATRIX = (({}, {"SGLANG_IS_IN_CI": "true"}),)
_PASSED = frozenset({"model_path", "device", "random_seed"}) _PASSED = frozenset({"model_path", "device", "random_seed"})
_EXPOSED = { _EXPOSED = {
("dllm/config.py", "max_running_requests"),
("dllm/config.py", "model_path"),
("multimodal/processors/base_processor.py", "image_processor_backend"),
("speculative/spec_registry.py", "disable_overlap_schedule"),
("layers/moe/utils.py", "deepep_mode"),
("layers/moe/utils.py", "moe_a2a_backend"),
("layers/moe/utils.py", "moe_runner_backend"),
("layers/moe/utils.py", "quantization"),
("layers/moe/utils.py", "speculative_moe_runner_backend"),
("entrypoints/sidecar.py", "grpc_port"), ("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"),
@@ -144,10 +153,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/common/conn.py", "disaggregation_bootstrap_port"),
("disaggregation/common/conn.py", "pp_size"),
("disaggregation/decode_kvcache_offload_manager.py", "hicache_io_backend"),
("disaggregation/decode_kvcache_offload_manager.py", "served_model_name"),
("disaggregation/utils.py", "disaggregation_transfer_backend"), ("disaggregation/utils.py", "disaggregation_transfer_backend"),
("distributed/bootstrap.py", "disable_custom_all_reduce"), ("distributed/bootstrap.py", "disable_custom_all_reduce"),
("distributed/bootstrap.py", "enable_symm_mem"), ("distributed/bootstrap.py", "enable_symm_mem"),
@@ -155,38 +160,6 @@ _EXPOSED = {
("distributed/bootstrap.py", "flashinfer_allreduce_fusion_backend"), ("distributed/bootstrap.py", "flashinfer_allreduce_fusion_backend"),
("distributed/bootstrap.py", "moe_a2a_backend"), ("distributed/bootstrap.py", "moe_a2a_backend"),
("distributed/bootstrap.py", "pre_warm_nccl"), ("distributed/bootstrap.py", "pre_warm_nccl"),
(
"distributed/device_communicators/mooncake_transfer_engine.py",
"disaggregation_ib_device",
),
(
"distributed/device_communicators/mooncake_transfer_engine.py",
"disaggregation_mode",
),
(
"distributed/device_communicators/mooncake_transfer_engine.py",
"disaggregation_transfer_backend",
),
(
"distributed/device_communicators/mooncake_transfer_engine.py",
"enable_hierarchical_cache",
),
(
"distributed/device_communicators/mooncake_transfer_engine.py",
"encoder_transfer_backend",
),
(
"distributed/device_communicators/mooncake_transfer_engine.py",
"mooncake_ib_device",
),
("dllm/config.py", "max_running_requests"),
("dllm/config.py", "model_path"),
("elastic_ep/elastic_ep.py", "elastic_ep_initial_size"),
("elastic_ep/elastic_ep.py", "ep_join_mode"),
("elastic_ep/elastic_ep.py", "moe_a2a_backend"),
("elastic_ep/expert_backup_manager.py", "disaggregation_ib_device"),
("elastic_ep/expert_backup_manager.py", "load_format"),
("elastic_ep/expert_backup_manager.py", "mooncake_ib_device"),
("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"),
@@ -198,17 +171,6 @@ _EXPOSED = {
("entrypoints/engine.py", "tool_call_parser"), ("entrypoints/engine.py", "tool_call_parser"),
("eplb/eplb_manager.py", "ep_dispatch_algorithm"), ("eplb/eplb_manager.py", "ep_dispatch_algorithm"),
("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"), ("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"),
("eplb/expert_distribution.py", "deepep_mode"),
("eplb/expert_distribution.py", "device"),
("eplb/expert_distribution.py", "expert_distribution_recorder_mode"),
("eplb/expert_distribution.py", "moe_a2a_backend"),
("eplb/expert_location.py", "device"),
("eplb/expert_location.py", "eplb_algorithm"),
("kv_canary/api.py", "disaggregation_mode"),
("kv_canary/api.py", "speculative_num_steps"),
("kv_canary/capacities.py", "chunked_prefill_size"),
("kv_canary/capacities.py", "cuda_graph_config"),
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
("layers/cp/base.py", "attn_cp_size"), ("layers/cp/base.py", "attn_cp_size"),
("layers/cp/base.py", "cp_strategy"), ("layers/cp/base.py", "cp_strategy"),
("layers/cp/base.py", "enable_prefill_cp"), ("layers/cp/base.py", "enable_prefill_cp"),
@@ -216,11 +178,6 @@ _EXPOSED = {
("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/moe/kt_ep_wrapper.py", "chunked_prefill_size"),
("layers/moe/utils.py", "deepep_mode"),
("layers/moe/utils.py", "moe_a2a_backend"),
("layers/moe/utils.py", "moe_runner_backend"),
("layers/moe/utils.py", "quantization"),
("layers/moe/utils.py", "speculative_moe_runner_backend"),
("layers/quantization/unquant.py", "enable_deterministic_inference"), ("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", "enable_lora"),
@@ -231,124 +188,18 @@ _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/disagg_service.py", "disaggregation_bootstrap_port"),
("managers/disagg_service.py", "disaggregation_mode"),
("managers/disagg_service.py", "disaggregation_transfer_backend"),
("managers/overlap_utils.py", "speculative_algorithm"),
("managers/prefill_delayer.py", "disable_overlap_schedule"), ("managers/prefill_delayer.py", "disable_overlap_schedule"),
("managers/rust_server.py", "mm_process_config"), ("managers/rust_server.py", "mm_process_config"),
("managers/schedule_batch.py", "disaggregation_mode"),
("managers/scheduler.py", "attn_cp_size"),
("managers/scheduler.py", "disable_overlap_schedule"),
("managers/scheduler.py", "disaggregation_mode"),
("managers/scheduler.py", "enable_hierarchical_cache"),
("managers/scheduler.py", "enable_lora"),
("managers/scheduler.py", "enable_lora_overlap_loading"),
("managers/scheduler.py", "moe_dp_size"),
("managers/scheduler.py", "pp_size"),
("managers/scheduler.py", "soft_watchdog_timeout"),
("managers/scheduler.py", "speculative_algorithm"),
("managers/tokenizer_manager.py", "served_model_name"),
("managers/tp_worker.py", "disable_overlap_schedule"),
("managers/tp_worker.py", "model_path"),
("managers/tp_worker.py", "random_seed"),
("managers/tp_worker.py", "speculative_algorithm"),
("managers/tp_worker.py", "tokenizer_path"),
("mem_cache/hiradix_cache.py", "hicache_io_backend"),
("mem_cache/hiradix_cache.py", "hicache_mem_layout"),
("mem_cache/hiradix_cache.py", "served_model_name"),
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "hicache_io_backend"),
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "hicache_mem_layout"),
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "served_model_name"),
("mem_cache/kv_cache_builder.py", "hicache_mem_layout"), ("mem_cache/kv_cache_builder.py", "hicache_mem_layout"),
("mem_cache/radix_cache_cpp.py", "enable_hierarchical_cache"),
("model_executor/model_runner.py", "device"),
("model_executor/model_runner.py", "speculative_algorithm"),
("model_executor/model_runner.py", "speculative_draft_attention_backend"),
("model_executor/model_runner_components/load_model_utils.py", "load_format"),
("model_executor/model_runner_components/load_model_utils.py", "quantization"),
(
"model_executor/model_runner_components/spec_aux_hidden_state.py",
"speculative_draft_attention_backend",
),
(
"model_executor/model_runner_components/spec_aux_hidden_state.py",
"speculative_draft_model_path",
),
(
"model_executor/model_runner_components/spec_aux_hidden_state.py",
"speculative_draft_model_revision",
),
("model_executor/model_runner_components/startup_weight_load.py", "attn_cp_size"),
(
"model_executor/model_runner_components/startup_weight_load.py",
"cuda_graph_config",
),
(
"model_executor/model_runner_components/startup_weight_load.py",
"custom_weight_loader",
),
("model_executor/model_runner_components/startup_weight_load.py", "device"),
("model_executor/model_runner_components/startup_weight_load.py", "enable_lora"),
("model_executor/model_runner_components/startup_weight_load.py", "lora_paths"),
("model_executor/model_runner_components/startup_weight_load.py", "pp_size"),
(
"model_executor/model_runner_components/startup_weight_load.py",
"speculative_algorithm",
),
( (
"model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py", "model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py",
"cuda_graph_config", "cuda_graph_config",
), ),
("multimodal/processors/base_processor.py", "image_processor_backend"),
("observability/metrics_collector.py", "disaggregation_mode"),
("observability/metrics_collector.py", "prefill_delayer_max_delay_passes"),
("observability/metrics_collector.py", "served_model_name"),
("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"),
("speculative/dflash_worker_v2.py", "speculative_draft_window_size"),
("speculative/dflash_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/draft_worker_common.py", "speculative_draft_attention_backend"), ("speculative/draft_worker_common.py", "speculative_draft_attention_backend"),
("speculative/dspark_components/dspark_config.py", "speculative_draft_model_path"),
(
"speculative/dspark_components/dspark_config.py",
"speculative_draft_model_revision",
),
("speculative/dspark_components/dspark_worker_v2.py", "disaggregation_mode"),
(
"speculative/dspark_components/dspark_worker_v2.py",
"speculative_num_draft_tokens",
),
("speculative/eagle_worker_v2.py", "device"),
("speculative/eagle_worker_v2.py", "speculative_adaptive"),
("speculative/eagle_worker_v2.py", "speculative_algorithm"),
("speculative/eagle_worker_v2.py", "speculative_eagle_topk"),
("speculative/eagle_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/eagle_worker_v2.py", "speculative_num_steps"),
("speculative/frozen_kv_mtp_worker_v2.py", "device"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_adaptive"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_algorithm"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_eagle_topk"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_num_steps"),
("speculative/multi_layer_eagle_worker_v2.py", "device"),
("speculative/multi_layer_eagle_worker_v2.py", "speculative_algorithm"),
("speculative/multi_layer_eagle_worker_v2.py", "speculative_eagle_topk"),
("speculative/multi_layer_eagle_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/multi_layer_eagle_worker_v2.py", "speculative_num_steps"),
("speculative/ngram_worker.py", "device"),
("speculative/ngram_worker.py", "disable_overlap_schedule"),
("speculative/ngram_worker.py", "speculative_eagle_topk"),
("speculative/ngram_worker.py", "speculative_num_draft_tokens"),
("speculative/ngram_worker.py", "speculative_num_steps"),
("speculative/spec_info.py", "enable_multi_layer_eagle"), ("speculative/spec_info.py", "enable_multi_layer_eagle"),
("speculative/spec_registry.py", "disable_overlap_schedule"),
("speculative/standalone_worker_v2.py", "device"),
("speculative/standalone_worker_v2.py", "speculative_algorithm"),
("speculative/standalone_worker_v2.py", "speculative_eagle_topk"),
("speculative/standalone_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/standalone_worker_v2.py", "speculative_num_steps"),
("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/cuda_vmm_transport_utils.py", "mm_feature_transport"),
@@ -367,56 +218,19 @@ _EXPOSED_CUDA_ONLY: frozenset = frozenset()
# some code overrides post-publish. Each needs an ordering judgment, not a blanket # some code overrides post-publish. Each needs an ordering judgment, not a blanket
# conversion; the list exists so a new one is a decision made when it is written. # conversion; the list exists so a new one is a decision made when it is written.
_OVERRIDDEN_AND_READ = { _OVERRIDDEN_AND_READ = {
("dllm/config.py", "model_path"),
("entrypoints/engine.py", "reasoning_parser"), ("entrypoints/engine.py", "reasoning_parser"),
("entrypoints/engine.py", "tool_call_parser"), ("entrypoints/engine.py", "tool_call_parser"),
("configs/model_config.py", "dtype"), ("configs/model_config.py", "dtype"),
("configs/model_config.py", "model_path"), ("configs/model_config.py", "model_path"),
("disaggregation/decode_kvcache_offload_manager.py", "hicache_storage_backend"),
(
"disaggregation/decode_kvcache_offload_manager.py",
"hicache_storage_backend_extra_config",
),
(
"distributed/device_communicators/mooncake_transfer_engine.py",
"hicache_storage_backend",
),
("dllm/config.py", "model_path"),
("elastic_ep/expert_backup_manager.py", "load_format"),
("kv_canary/api.py", "speculative_num_steps"),
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
("managers/scheduler.py", "hicache_storage_backend"),
("managers/tp_worker.py", "model_path"),
("mem_cache/hiradix_cache.py", "hicache_storage_backend"),
("mem_cache/hiradix_cache.py", "hicache_storage_backend_extra_config"),
("mem_cache/hiradix_cache.py", "hicache_storage_prefetch_policy"),
("mem_cache/hiradix_cache.py", "hicache_write_policy"),
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "hicache_storage_backend"),
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "hicache_write_policy"),
("mem_cache/kv_cache_builder.py", "hicache_storage_backend"), ("mem_cache/kv_cache_builder.py", "hicache_storage_backend"),
("mem_cache/pool_host/common.py", "hicache_storage_backend"), ("mem_cache/pool_host/common.py", "hicache_storage_backend"),
("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"), ("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"),
("mem_cache/radix_cache_cpp.py", "hicache_write_policy"),
("mem_cache/unified_radix_cache.py", "hicache_storage_backend"), ("mem_cache/unified_radix_cache.py", "hicache_storage_backend"),
("mem_cache/unified_radix_cache.py", "hicache_storage_backend_extra_config"), ("mem_cache/unified_radix_cache.py", "hicache_storage_backend_extra_config"),
("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"), ("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"),
("mem_cache/unified_radix_cache.py", "hicache_write_policy"), ("mem_cache/unified_radix_cache.py", "hicache_write_policy"),
("model_executor/model_runner_components/load_model_utils.py", "load_format"),
("parser/template_detection.py", "model_path"), ("parser/template_detection.py", "model_path"),
("speculative/dflash_worker_v2.py", "speculative_num_draft_tokens"),
(
"speculative/dspark_components/dspark_worker_v2.py",
"speculative_num_draft_tokens",
),
("speculative/eagle_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/eagle_worker_v2.py", "speculative_num_steps"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/frozen_kv_mtp_worker_v2.py", "speculative_num_steps"),
("speculative/multi_layer_eagle_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/multi_layer_eagle_worker_v2.py", "speculative_num_steps"),
("speculative/ngram_worker.py", "speculative_num_draft_tokens"),
("speculative/ngram_worker.py", "speculative_num_steps"),
("speculative/standalone_worker_v2.py", "speculative_num_draft_tokens"),
("speculative/standalone_worker_v2.py", "speculative_num_steps"),
("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"),
} }