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 = []
object.__setattr__(server_args, "_runtime_mutations", log)
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)
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:
"""Apply the accumulated declarations onto ``server_args`` once, at 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
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:
"""Read-only view of the resolving configuration for mid-resolution code
that is not a pass (``__post_init__`` handlers and hooks). Internal to
@@ -5,7 +5,10 @@ import logging
import os
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:
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.
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:
assert server_args.speculative_algorithm == "EAGLE", (
@@ -163,7 +170,13 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
_init_adaptive_speculative_params(server_args)
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:
@@ -36,7 +36,12 @@ from sglang.srt.layers.dp_attention import (
get_attention_dp_rank,
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.utils.network import (
NetworkAddress,
@@ -164,9 +169,9 @@ class CommonKVManager(BaseKVManager):
envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get()
)
# for p/d multi node infer
self.bootstrap_host = server_args.host
self.bootstrap_port = server_args.disaggregation_bootstrap_port
self.dist_init_addr = server_args.dist_init_addr
self.bootstrap_host = get_serving().host
self.bootstrap_port = get_disagg().disaggregation_bootstrap_port
self.dist_init_addr = get_parallel().dist_init_addr
parallel = get_parallel()
self.attn_tp_size = parallel.attn_tp_size
self.attn_tp_rank = parallel.attn_tp_rank
@@ -182,7 +187,7 @@ class CommonKVManager(BaseKVManager):
self.system_dp_rank = (
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.local_ip = get_local_ip_auto()
cp_sharded_prefill = self.attn_cp_size > 1 and (
@@ -21,7 +21,11 @@ from sglang.srt.mem_cache.memory_pool import (
MLATokenToKVPool,
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.utils.common import ceil_align
@@ -68,10 +72,10 @@ class DecodeKVCacheOffloadManager:
self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group)
hicache_storage_backend_extra_config = {}
if server_args.hicache_storage_backend_extra_config:
if get_memory().hicache_storage_backend_extra_config:
try:
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:
raise ValueError(
@@ -83,10 +87,10 @@ class DecodeKVCacheOffloadManager:
mem_pool_host=self.decode_host_mem_pool,
page_size=self.page_size,
tp_group=tp_group,
io_backend=server_args.hicache_io_backend,
io_backend=get_memory().hicache_io_backend,
load_cache_event=threading.Event(),
storage_backend=server_args.hicache_storage_backend,
model_name=server_args.served_model_name,
storage_backend=get_memory().hicache_storage_backend,
model_name=get_serving().served_model_name,
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 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
if TYPE_CHECKING:
@@ -313,27 +318,27 @@ def maybe_init_shared_mooncake_transfer_engine(
"""
use_mooncake_te = (
(
server_args.disaggregation_mode != "null"
and server_args.disaggregation_transfer_backend == "mooncake"
get_disagg().disaggregation_mode != "null"
and get_disagg().disaggregation_transfer_backend == "mooncake"
)
or (
server_args.enable_hierarchical_cache
and server_args.hicache_storage_backend == "mooncake"
get_memory().enable_hierarchical_cache
and get_memory().hicache_storage_backend == "mooncake"
and envs.SGLANG_HICACHE_MOONCAKE_REUSE_TE.get()
)
or (
server_args.encoder_only
and server_args.encoder_transfer_backend == "mooncake"
get_disagg().encoder_only
and get_disagg().encoder_transfer_backend == "mooncake"
)
or (
server_args.language_only
and server_args.encoder_transfer_backend == "mooncake"
get_disagg().language_only
and get_disagg().encoder_transfer_backend == "mooncake"
)
or (
server_args.enable_elastic_expert_backup
and server_args.elastic_ep_backend is not None
get_exec().moe.enable_elastic_expert_backup
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:
@@ -341,11 +346,12 @@ def maybe_init_shared_mooncake_transfer_engine(
hostname=get_local_ip_auto(),
gpu_id=gpu_id,
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:
from mooncake.pg import set_transfer_engine
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.eplb.expert_location import broadcast_global_expert_location_metadata
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
if TYPE_CHECKING:
@@ -87,9 +91,9 @@ class ElasticEPStateManager:
if cls._instance is not None:
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()
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, (
f"--max-ep-size ({active_rank_capacity}) must be >= "
f"world_size ({world_size})."
@@ -103,10 +107,10 @@ class ElasticEPStateManager:
inst.snapshot_active_to_last()
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
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:
cls._init_joiner_state(inst, server_args)
@@ -122,12 +126,13 @@ class ElasticEPStateManager:
inst.snapshot_active_to_last()
inst.sync_active_to_cpu()
if server_args.ep_join_mode == "scale":
if get_exec().moe.ep_join_mode == "scale":
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 = (
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
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.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 (
PortArgs,
ServerArgs,
@@ -38,14 +45,14 @@ def extract_expert_id(param_name):
class ExpertBackupManager:
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.continuous_buffer = None
self.weight_pointer_map = {}
self.transfer_engine = None
self.session_id = None
self.engine_num = server_args.nnodes
self.engine_rank = server_args.node_rank
self.engine_num = get_parallel().nnodes
self.engine_rank = get_parallel().node_rank
self.expert_num = self.model_config.hf_config.n_routed_experts
self.idmn = (self.expert_num // self.engine_num) * self.engine_rank
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.
self.recv_from_expert_backup_client = context.socket(zmq.PULL)
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.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.start_transfer_server()
@@ -66,7 +73,7 @@ class ExpertBackupManager:
# losing the initial PUB message due to slow joiners.
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)
num_ready_clients += 1
@@ -168,7 +175,7 @@ def run_expert_backup_manager_process(
hostname=get_local_ip_auto(),
gpu_id=0,
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)
+19 -15
View File
@@ -45,7 +45,11 @@ from sglang.srt.observability.metrics_collector import (
ExpertDispatchCollector,
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.utils import Withable, get_device, get_int_env_var
@@ -84,7 +88,7 @@ class ExpertDistributionRecorder(ABC):
expert_location_metadata: ExpertLocationMetadata,
rank: int,
):
if server_args.expert_distribution_recorder_mode is not None:
if get_exec().moe.expert_distribution_recorder_mode is not None:
assert (
expert_location_metadata is not None
), "ExpertLocationMetadata is required for expert distribution recording. One possible"
@@ -178,7 +182,7 @@ class _ExpertDistributionRecorderReal(ExpertDistributionRecorder):
if server_args.should_report_expert_balancedness():
logger.info(
"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()
@@ -328,30 +332,30 @@ class _SinglePassGatherer(ABC):
expert_location_metadata: ExpertLocationMetadata,
rank: int,
) -> _SinglePassGatherer:
if server_args.expert_distribution_recorder_mode == "per_token":
if get_exec().moe.expert_distribution_recorder_mode == "per_token":
return _DetailSinglePassGatherer(
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)
if server_args.expert_distribution_recorder_mode == "stat_approx":
if server_args.moe_a2a_backend != "none" and (
server_args.deepep_mode == "normal"
if get_exec().moe.expert_distribution_recorder_mode == "stat_approx":
if get_exec().moe.moe_a2a_backend != "none" and (
get_exec().moe.deepep_mode == "normal"
):
return _DeepepNormalSinglePassGatherer(expert_location_metadata, rank)
else:
raise NotImplementedError
if server_args.moe_a2a_backend == "deepep":
if server_args.deepep_mode == "normal":
if get_exec().moe.moe_a2a_backend == "deepep":
if get_exec().moe.deepep_mode == "normal":
return _SelectExpertsSinglePassGatherer(expert_location_metadata, rank)
elif server_args.deepep_mode == "low_latency":
elif get_exec().moe.deepep_mode == "low_latency":
return _DeepepLowLatencySinglePassGatherer(
expert_location_metadata,
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:
raise NotImplementedError
@@ -412,11 +416,11 @@ class _DetailSinglePassGatherer(_SinglePassGatherer):
self._TOP_K_NUM,
),
dtype=torch.int32,
device=server_args.device,
device=get_device_namespace().device,
)
self._misc_objects: List[Dict[str, Any]] = []
assert (
not server_args.enable_two_batch_overlap
not get_exec().overlap.enable_two_batch_overlap
), "DetailSinglePassGatherer does not support TBO yet"
# TODO assert shared experts fusion is disabled, o/w data is wrong
@@ -678,7 +682,7 @@ class _Accumulator(ABC):
"stat_approx": _StatAccumulator,
"per_pass": _DetailAccumulator,
"per_token": _DetailAccumulator,
}[server_args.expert_distribution_recorder_mode]
}[get_exec().moe.expert_distribution_recorder_mode]
def __init__(
self,
+13 -7
View File
@@ -25,6 +25,12 @@ import torch
import torch.distributed
import torch.nn.functional as F
from sglang.srt.runtime_context import (
configured_tp_size,
get_device,
get_exec,
)
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
@@ -141,7 +147,7 @@ class ExpertLocationMetadata:
):
if not isinstance(physical_to_logical_map, torch.Tensor):
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)
@@ -184,7 +190,7 @@ class ExpertLocationMetadata:
logical_count = torch.tensor(logical_count)
if len(logical_count.shape) == 2:
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
@@ -208,7 +214,7 @@ class ExpertLocationMetadata:
num_groups=num_groups,
num_nodes=num_nodes,
algorithm=eplb_algorithms.compute_algorithm(
raw_algorithm=server_args.eplb_algorithm,
raw_algorithm=get_exec().moe.eplb_algorithm,
num_groups=num_groups,
num_nodes=num_nodes,
),
@@ -217,9 +223,9 @@ class ExpertLocationMetadata:
return ExpertLocationMetadata._init_raw(
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(
server_args.device
get_device().device
),
)
@@ -246,7 +252,7 @@ class ExpertLocationMetadata:
if get_exec().moe.ep_join_mode == "scale":
ep_size = max(
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 = (
_compute_elastic_expert_layout(
@@ -781,7 +787,7 @@ def compute_initial_expert_location_metadata(
model_config: ModelConfig,
moe_ep_rank: int,
) -> Optional[ExpertLocationMetadata]:
data = server_args.init_expert_location
data = get_exec().moe.init_expert_location
if data == "trivial":
return ExpertLocationMetadata.init_trivial(
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,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import (
get_disagg,
get_spec,
)
if TYPE_CHECKING:
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
)
launch_capacities = CanaryLaunchCapacities.from_args(
server_args=model_runner.server_args,
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],
pool_slot_count=model_runner.max_total_num_tokens,
)
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(
config=config,
perturb_config=perturb_config,
@@ -88,7 +91,7 @@ def install_canary(
"install_canary: disaggregation_mode=%s config=%s perturb_config=%s "
"launch_capacities=%s n_buffer_groups=%d buffer_group_kinds=%s "
"swa_window_size=%d speculative_num_steps=%d",
server_args.disaggregation_mode,
get_disagg().disaggregation_mode,
config,
perturb_config,
launch_capacities,
+9 -8
View File
@@ -2,10 +2,12 @@ from __future__ import annotations
import math
from dataclasses import dataclass
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
from sglang.srt.runtime_context import (
get_exec,
get_schedule,
get_spec,
)
@dataclass(frozen=True, slots=True, kw_only=True)
@@ -43,7 +45,6 @@ class CanaryLaunchCapacities:
def from_args(
cls,
*,
server_args: ServerArgs,
req_to_token_pool_size: int,
max_seq_len_per_req: int,
pool_slot_count: int,
@@ -63,7 +64,7 @@ class CanaryLaunchCapacities:
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_config.decode.max_bs if cuda_graph_config is not None else 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}"
)
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:
spec_num_draft_tokens = 0
if spec_num_draft_tokens < 0:
@@ -81,7 +82,7 @@ class CanaryLaunchCapacities:
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:
raise ValueError(
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)
chunked_prefill_size = server_args.chunked_prefill_size
chunked_prefill_size = get_schedule().chunked_prefill_size
chunked_limit = (
chunked_prefill_size
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,
get_kv_class,
)
from sglang.srt.runtime_context import (
get_disagg,
get_parallel,
get_serving,
)
from sglang.srt.server_args import ServerArgs
@@ -15,8 +20,8 @@ def start_disagg_service(
server_args: ServerArgs,
):
# Start kv bootstrap server on prefill
disagg_mode = DisaggregationMode(server_args.disaggregation_mode)
transfer_backend = TransferBackend(server_args.disaggregation_transfer_backend)
disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode)
transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend)
if disagg_mode == DisaggregationMode.PREFILL:
# only start bootstrap server on prefill tm
@@ -24,8 +29,8 @@ def start_disagg_service(
transfer_backend, KVClassType.BOOTSTRAP_SERVER
)
bootstrap_server = kv_bootstrap_server_class(
host=server_args.host,
port=server_args.disaggregation_bootstrap_port,
host=get_serving().host,
port=get_disagg().disaggregation_bootstrap_port,
)
maybe_create_ascend_config_store(
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``,
which the rust registry ports verbatim), leaving this store as the only
``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
try:
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.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
if TYPE_CHECKING:
@@ -34,10 +38,10 @@ def decide_needs_cpu_seq_lens(
# importable everywhere; spec_info pulls in the spec/schedule_batch graph.
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
return True
algo = SpeculativeAlgorithm.from_string(server_args.speculative_algorithm)
algo = SpeculativeAlgorithm.from_string(get_spec().speculative_algorithm)
if algo.is_ngram():
# 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).
@@ -56,7 +60,7 @@ def decide_needs_confidence_relay(server_args: ServerArgs) -> bool:
)
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():
return False
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)
# pass offload_kv=False to skip the wasteful device->host copy.
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(
req,
tree_cache,
@@ -2826,7 +2826,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self, server_args: ServerArgs
) -> Tuple[List[Req], float, List[Req]]:
"""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 = []
reqs_to_abort: List[Req] = []
@@ -2886,9 +2886,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
return retracted_reqs, new_estimate_ratio, reqs_to_abort
@staticmethod
def _get_decode_retraction_order(
reqs: List[Req], server_args: ServerArgs
) -> List[int]:
def _get_decode_retraction_order(reqs: List[Req]) -> List[int]:
"""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
@@ -2901,15 +2899,17 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
def length_key(req: Req) -> Tuple[int, int]:
return (len(req.output_ids), -len(req.origin_input_ids))
if server_args.retraction_policy == "priority":
priority_sign = 1 if server_args.schedule_low_priority_values_first else -1
if get_schedule().retraction_policy == "priority":
priority_sign = (
1 if get_schedule().schedule_low_priority_values_first else -1
)
def retraction_key(req: Req) -> Tuple[int, int, int]:
priority = req.priority
if priority is None:
priority = (
sys.maxsize
if server_args.schedule_low_priority_values_first
if get_schedule().schedule_low_priority_values_first
else -sys.maxsize - 1
)
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 (
attention_backends,
configured_attn_cp_size,
configured_dcp_size,
configured_moe_dp_size,
configured_pp_size,
configured_tp_size,
@@ -417,71 +418,75 @@ class Scheduler(
# Parse args
self.server_args = server_args
self.nccl_port = port_args.nccl_port
self.schedule_policy = server_args.schedule_policy
self.enable_priority_scheduling = server_args.enable_priority_scheduling
self.schedule_policy = get_schedule().schedule_policy
self.enable_priority_scheduling = get_schedule().enable_priority_scheduling
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 = (
server_args.schedule_low_priority_values_first
get_schedule().schedule_low_priority_values_first
)
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_overlap_loading = server_args.enable_lora_overlap_loading
self.max_loras_per_batch = server_args.max_loras_per_batch
self.enable_overlap = not server_args.disable_overlap_schedule and not use_mlx()
self.enable_overlap_mlx = not server_args.disable_overlap_schedule and use_mlx()
self.enable_pdmux = server_args.enable_pdmux
self.enable_lora = get_lora().enable_lora
self.enable_lora_overlap_loading = get_lora().enable_lora_overlap_loading
self.max_loras_per_batch = get_lora().max_loras_per_batch
self.enable_overlap = (
not get_schedule().disable_overlap_schedule and not use_mlx()
)
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.stream_interval = server_args.stream_interval
self.stream_interval = get_serving().stream_interval
self.spec_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
get_spec().speculative_algorithm
)
self.page_size = get_schedule().page_size
self.enable_hierarchical_cache = server_args.enable_hierarchical_cache
self.enable_session_radix_cache = server_args.enable_session_radix_cache
self.enable_hicache_storage = server_args.hicache_storage_backend is not None
self.enable_hierarchical_cache = get_memory().enable_hierarchical_cache
self.enable_session_radix_cache = get_memory().enable_session_radix_cache
self.enable_hicache_storage = get_memory().hicache_storage_backend is not None
self.enable_decode_hicache = (
server_args.disaggregation_decode_enable_radix_cache
get_disagg().disaggregation_decode_enable_radix_cache
and self.enable_hierarchical_cache
)
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.enable_hisparse = server_args.enable_hisparse
self.enable_hisparse = get_memory().enable_hisparse
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
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
compute_dp_attention_world_info(
get_parallel().enable_dp_attention,
tp_rank,
server_args.tp_size,
configured_tp_size(),
get_parallel().dp_size,
server_args.attn_cp_size,
configured_attn_cp_size(),
)
)
self.ps = ParallelState(
tp_rank=tp_rank,
tp_size=server_args.tp_size,
tp_size=configured_tp_size(),
pp_rank=pp_rank,
pp_size=server_args.pp_size,
pp_size=configured_pp_size(),
dp_rank=dp_rank,
dp_size=get_parallel().dp_size,
attn_tp_rank=attn_tp_rank,
attn_tp_size=attn_tp_size,
attn_cp_rank=attn_cp_rank,
attn_cp_size=server_args.attn_cp_size,
attn_dcp_rank=tp_rank % server_args.dcp_size,
attn_dcp_size=server_args.dcp_size,
attn_cp_size=configured_attn_cp_size(),
attn_dcp_rank=tp_rank % configured_dcp_size(),
attn_dcp_size=configured_dcp_size(),
attn_dp_rank=attn_dp_rank,
attn_dp_size=attn_dp_size,
moe_ep_rank=moe_ep_rank,
moe_ep_size=get_parallel().ep_size,
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,
)
@@ -839,7 +844,7 @@ class Scheduler(
if (
self.model_config.is_multimodal
and self.processor is not None
and not server_args.language_model_only
and not get_disagg().language_model_only
):
try:
import_processors("sglang.srt.multimodal.processors")
@@ -1282,7 +1287,7 @@ class Scheduler(
self.new_token_ratio_tracker = NewTokenRatioTracker.from_config()
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, watchdog_timeout=x, soft=True
)
@@ -5172,16 +5177,16 @@ def run_scheduler_process(
parent_process = psutil.Process().parent()
# Set up tracing
if server_args.enable_trace:
if get_observability().enable_trace:
process_tracing_init(
server_args.otlp_traces_endpoint,
get_observability().otlp_traces_endpoint,
"sglang",
trace_modules=server_args.trace_modules,
trace_modules=get_observability().trace_modules,
)
thread_label = "Scheduler"
if server_args.disaggregation_mode == "prefill":
if get_disagg().disaggregation_mode == "prefill":
thread_label = "Prefill Scheduler"
elif server_args.disaggregation_mode == "decode":
elif get_disagg().disaggregation_mode == "decode":
thread_label = "Decode Scheduler"
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_mm,
get_model,
get_observability,
get_parallel,
get_schedule,
get_serving,
get_spec,
)
@@ -411,14 +413,14 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.elastic_pending_ep_size = None
self.elastic_scale_phase = "idle"
self.elastic_last_error = None
self.enable_metrics = server_args.enable_metrics
self.incremental_streaming_output = server_args.incremental_streaming_output
self.enable_metrics = get_observability().enable_metrics
self.incremental_streaming_output = get_serving().incremental_streaming_output
self.enable_lora = get_lora().enable_lora
self.enable_trace = server_args.enable_trace
self.allow_auto_truncate = server_args.allow_auto_truncate
self.skip_tokenizer_init = server_args.skip_tokenizer_init
self.enable_trace = get_observability().enable_trace
self.allow_auto_truncate = get_serving().allow_auto_truncate
self.skip_tokenizer_init = get_serving().skip_tokenizer_init
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
self.init_model_config()
@@ -462,15 +464,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
model_config_class = getattr(self, "model_config_class", ModelConfig)
# Read model args
self.model_path = server_args.model_path
self.served_model_name = server_args.served_model_name
self.model_path = get_model().model_path
self.served_model_name = get_serving().served_model_name
self.model_config = model_config_class.from_server_args(server_args)
self.is_generation = self.model_config.is_generation
self.context_len = self.model_config.context_len
self.image_token_id = self.model_config.image_token_id
self.max_req_input_len = None # Will be set later in engine.py
self.enable_priority_scheduling = server_args.enable_priority_scheduling
self.default_priority_value = server_args.default_priority_value
self.enable_priority_scheduling = get_schedule().enable_priority_scheduling
self.default_priority_value = get_schedule().default_priority_value
self.num_reserved_tokens = compute_num_reserved_tokens()
self.validate_total_tokens = True
@@ -478,7 +480,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
server_args = self.server_args
# 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")
if mm_process_pkg := envs.SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE.get():
import_processors(mm_process_pkg, overwrite=True)
@@ -496,7 +498,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
model_config=self.model_config,
)
if server_args.skip_tokenizer_init:
if get_serving().skip_tokenizer_init:
self.tokenizer = self.processor = None
else:
self.processor = _processor
@@ -505,26 +507,26 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
else:
self.mm_processor = self.processor = None
if server_args.skip_tokenizer_init:
if get_serving().skip_tokenizer_init:
self.tokenizer = None
else:
self.tokenizer = get_tokenizer(
get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
tokenizer_backend=server_args.tokenizer_backend,
tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=get_model().trust_remote_code,
revision=get_model().revision,
tokenizer_backend=get_serving().tokenizer_backend,
)
# Initialize async dynamic batch tokenizer if enabled (common for both multimodal and non-multimodal)
if (
server_args.enable_dynamic_batch_tokenizer
and not server_args.skip_tokenizer_init
get_serving().enable_dynamic_batch_tokenizer
and not get_serving().skip_tokenizer_init
):
self.async_dynamic_batch_tokenizer = AsyncDynamicbatchTokenizer(
self.tokenizer,
max_batch_size=server_args.dynamic_batch_tokenizer_batch_size,
batch_wait_timeout_s=server_args.dynamic_batch_tokenizer_batch_timeout,
max_batch_size=get_serving().dynamic_batch_tokenizer_batch_size,
batch_wait_timeout_s=get_serving().dynamic_batch_tokenizer_batch_timeout,
)
else:
self.async_dynamic_batch_tokenizer = None
@@ -547,7 +549,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.recv_from_detokenizer = get_zmq_socket(
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(
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.
# Request logging
self.request_logger = RequestLogger(
log_requests=self.server_args.log_requests,
log_requests_level=self.server_args.log_requests_level,
log_requests_format=self.server_args.log_requests_format,
log_requests_target=self.server_args.log_requests_target,
log_requests=get_observability().log_requests,
log_requests_level=get_observability().log_requests_level,
log_requests_format=get_observability().log_requests_format,
log_requests_target=get_observability().log_requests_target,
)
# Dumping
@@ -621,7 +623,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def init_weight_update(self):
# Initial weights status
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
# Weight updates
@@ -668,7 +670,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Encoder Disaggregation
self.encoder_bootstrap_server = None
if self.server_args.language_only:
if get_disagg().language_only:
from sglang.srt.disaggregation.encoder.receiver import (
EncoderBootstrapServer,
)
@@ -677,10 +679,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# entries as encoders register, the receiver reads from the same
# list. Pre-populated with static --encoder-urls so the legacy
# 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(
host=self.server_args.host,
port=self.server_args.encoder_bootstrap_port,
host=get_serving().host,
port=get_disagg().encoder_bootstrap_port,
urls=self.encoder_urls,
)
self.mm_receiver = create_mm_receiver(
@@ -698,16 +700,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
)
labels = {
"model_name": self.server_args.served_model_name,
"model_name": get_serving().served_model_name,
"engine_type": engine_type,
}
if self.enable_priority_scheduling:
labels["priority"] = ""
if self.server_args.tokenizer_metrics_allowed_custom_labels:
for label in self.server_args.tokenizer_metrics_allowed_custom_labels:
if get_observability().tokenizer_metrics_allowed_custom_labels:
for (
label
) in get_observability().tokenizer_metrics_allowed_custom_labels:
labels[label] = ""
if self.server_args.extra_metric_labels:
labels.update(self.server_args.extra_metric_labels)
if get_observability().extra_metric_labels:
labels.update(get_observability().extra_metric_labels)
tokenizer_collector_cls = resolve_collector_class(
self.server_args,
STAT_LOGGER_ROLE_TOKENIZER,
@@ -716,15 +720,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.metrics_collector = tokenizer_collector_cls(
server_args=self.server_args,
labels=labels,
bucket_time_to_first_token=self.server_args.bucket_time_to_first_token,
bucket_e2e_request_latency=self.server_args.bucket_e2e_request_latency,
bucket_inter_token_latency=self.server_args.bucket_inter_token_latency,
bucket_time_to_first_token=get_observability().bucket_time_to_first_token,
bucket_e2e_request_latency=get_observability().bucket_e2e_request_latency,
bucket_inter_token_latency=get_observability().bucket_inter_token_latency,
)
start_cpu_monitor_thread("tokenizer")
if self.server_args.gc_warning_threshold_secs > 0.0:
configure_gc_warning(self.server_args.gc_warning_threshold_secs)
if get_observability().gc_warning_threshold_secs > 0.0:
configure_gc_warning(get_observability().gc_warning_threshold_secs)
self.soft_watchdog = Watchdog.create(
debug_name="TokenizerManager",
watchdog_timeout=get_device().soft_watchdog_timeout,
@@ -773,7 +777,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if (
isinstance(obj, GenerateReqInput)
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(
"max_thinking_tokens requires the server to be launched with "
@@ -793,7 +797,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self._init_req_state(obj, request)
try:
if self.server_args.language_only:
if get_disagg().language_only:
self._handle_epd_disaggregation_encode_request(obj)
# Log the request
@@ -993,7 +997,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
)
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(
"Multimodal inputs are not supported when --language-model-only "
"is set; the encoder is not loaded. Restart without the flag."
@@ -1028,10 +1032,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
)
if (
not self.server_args.language_only
not get_disagg().language_only
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(
request_obj=obj,
mm_processor=self.mm_processor,
@@ -1053,7 +1057,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
max_req_input_len=self.max_req_input_len,
)
elif (
self.server_args.language_only
get_disagg().language_only
and get_disagg().encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
and not obj.need_wait_for_mm_inputs
@@ -1238,7 +1242,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
)
if (
obj.custom_logit_processor
and not self.server_args.enable_custom_logit_processor
and not get_exec().features.enable_custom_logit_processor
):
raise ValueError(
"The server is not configured to enable custom logit processor. "
@@ -1945,7 +1949,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
return
if (
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
):
return
@@ -2175,7 +2179,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
else:
customized_info = None
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):
state = self.rid_to_state.get(rid, None)
if state is None:
@@ -3183,7 +3187,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state.obj.top_logprobs_num,
state.obj.token_ids_logprob,
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
@@ -3304,12 +3308,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
unique_lora_paths = set(obj.lora_path)
if (
self.server_args.max_loaded_loras is not None
and len(unique_lora_paths) > self.server_args.max_loaded_loras
get_lora().max_loaded_loras is not None
and len(unique_lora_paths) > get_lora().max_loaded_loras
):
raise ValueError(
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
@@ -3442,7 +3446,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if isinstance(obj, GenerateReqInput) and obj.contains_mm_input():
# dispatch to encoder by default
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)
# 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:
is_cross_node = server_args.dist_init_addr
is_cross_node = get_parallel().dist_init_addr
if is_cross_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,
)
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.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed
from sglang.srt.utils.hf_transformers_utils import (
@@ -341,28 +348,28 @@ class TpModelWorker(BaseTpWorker):
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:
# tokenizer_path always points at the target model.
self.tokenizer = self.processor = None
else:
if self.model_config.is_multimodal:
self.processor = get_processor(
server_args.tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
tokenizer_backend=server_args.tokenizer_backend,
model_name=server_args.model_path,
get_serving().tokenizer_path,
tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=get_model().trust_remote_code,
revision=get_model().revision,
tokenizer_backend=get_serving().tokenizer_backend,
model_name=get_model().model_path,
)
self.tokenizer = get_tokenizer_from_processor(self.processor)
else:
self.tokenizer = get_tokenizer(
server_args.tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
tokenizer_backend=server_args.tokenizer_backend,
get_serving().tokenizer_path,
tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=get_model().trust_remote_code,
revision=get_model().revision,
tokenizer_backend=get_serving().tokenizer_backend,
)
self.device = self.model_runner.device
@@ -373,18 +380,18 @@ class TpModelWorker(BaseTpWorker):
# Sync random seed across TP workers.
# Elastic joiners cannot enter the launch-time WORLD broadcast.
if server_args.is_ep_joiner:
self.random_seed = server_args.random_seed
self.random_seed = get_device().random_seed
else:
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.world_group.cpu_group,
src=self.world_group.ranks[0],
)[0]
set_random_seed(self.random_seed)
self.enable_overlap = not server_args.disable_overlap_schedule
self.enable_spec = server_args.speculative_algorithm is not None
self.enable_overlap = not get_schedule().disable_overlap_schedule
self.enable_spec = get_spec().speculative_algorithm is not None
self.hicache_layer_transfer_counter = None
def alloc_memory_pool(
+19 -15
View File
@@ -65,7 +65,11 @@ from sglang.srt.observability.metrics_collector import (
StorageMetricsCollector,
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:
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.kv_cache,
get_memory().hicache_ratio,
server_args.hicache_size,
get_memory().hicache_size,
self.page_size,
server_args.hicache_mem_layout,
get_memory().hicache_mem_layout,
allocator_type=allocator_type,
)
elif isinstance(self.kv_cache, DSATokenToKVPool):
@@ -106,9 +110,9 @@ class HiRadixCache(RadixCache):
self.token_to_kv_pool_host = MLATokenToKVPoolHost(
self.kv_cache,
get_memory().hicache_ratio,
server_args.hicache_size,
get_memory().hicache_size,
self.page_size,
server_args.hicache_mem_layout,
get_memory().hicache_mem_layout,
allocator_type=allocator_type,
dcp_size=_parallel.attn_dcp_size,
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.pp_rank = params.pp_rank
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.extra_metric_labels = server_args.extra_metric_labels
self.extra_metric_labels = get_observability().extra_metric_labels
(
extra_config,
@@ -133,11 +137,11 @@ class HiRadixCache(RadixCache):
prefetch_timeout_config,
hicache_storage_pass_prefix_keys,
) = 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
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()
if isinstance(self.kv_cache, DSATokenToKVPool):
@@ -174,16 +178,16 @@ class HiRadixCache(RadixCache):
attn_cp_group=self.attn_cp_group,
attn_tp_group=self.attn_tp_group,
pp_group=self.pp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=server_args.hicache_storage_backend,
write_policy=get_memory().hicache_write_policy,
io_backend=get_memory().hicache_io_backend,
storage_backend=get_memory().hicache_storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=server_args.served_model_name,
model_name=get_serving().served_model_name,
storage_backend_extra_config=extra_config,
enable_storage_metrics=self.enable_storage_metrics,
)
self._apply_storage_runtime_config(
storage_backend=server_args.hicache_storage_backend,
storage_backend=get_memory().hicache_storage_backend,
prefetch_threshold=prefetch_threshold,
prefetch_timeout_config=prefetch_timeout_config,
hicache_storage_pass_prefix_keys=hicache_storage_pass_prefix_keys,
@@ -205,7 +209,7 @@ class HiRadixCache(RadixCache):
self.work_list: List[torch.distributed.Work] = []
# todo: dynamically adjust the 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
# 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.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:
import torch
@@ -100,9 +104,9 @@ def build_kv_host_pool(
return kv_host_pool_cls(
kv_pool,
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,
server_args.hicache_mem_layout,
get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args),
pool_label=pool_label,
**kwargs,
@@ -299,8 +303,8 @@ def build_kv_only_stack(
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
write_policy=get_memory().hicache_write_policy,
io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
@@ -340,9 +344,9 @@ def build_hybrid_swa_stack(
)
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(
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(
@@ -370,8 +374,8 @@ def build_hybrid_swa_stack(
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
write_policy=get_memory().hicache_write_policy,
io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
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
if server_args.hicache_size > 0:
if get_memory().hicache_size > 0:
raise ValueError(
"DeepSeek V4 HiCache currently does not support --hicache-size; "
"use --hicache-ratio instead."
@@ -496,7 +500,7 @@ def build_deepseek_v4_hicache_stack(
)
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 = [
build_pool_entry(
@@ -519,7 +523,7 @@ def build_deepseek_v4_hicache_stack(
item_bytes=kvcache.swa_kv_pool.bytes_per_page_padded,
num_host_pages=swa_num_host_pages,
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),
)
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,
num_host_pages=num_host_pages,
slot_page_size=page_size,
layout=server_args.hicache_mem_layout,
layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args),
)
c4_indexer_host_pool = DeepSeekV4PagedHostPool(
@@ -557,7 +561,7 @@ def build_deepseek_v4_hicache_stack(
),
num_host_pages=num_host_pages,
slot_page_size=page_size,
layout=server_args.hicache_mem_layout,
layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args),
)
entries.extend(
@@ -588,7 +592,7 @@ def build_deepseek_v4_hicache_stack(
],
num_host_pages=swa_num_host_pages,
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),
)
c4_indexer_state_host_pool = DeepSeekV4StateHostPool(
@@ -599,7 +603,7 @@ def build_deepseek_v4_hicache_stack(
],
num_host_pages=swa_num_host_pages,
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),
)
entries.extend(
@@ -631,7 +635,7 @@ def build_deepseek_v4_hicache_stack(
item_bytes=c128_item_bytes,
num_host_pages=num_host_pages,
slot_page_size=page_size,
layout=server_args.hicache_mem_layout,
layout=get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args),
)
# 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_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
write_policy=get_memory().hicache_write_policy,
io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
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
)
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(
server_args.hicache_size, (kv_pool, mamba_pool)
get_memory().hicache_size, (kv_pool, mamba_pool)
)
kv_host_pool = build_kv_host_pool(
kv_pool=kv_pool,
@@ -721,7 +725,7 @@ def build_hybrid_mamba_stack(
get_memory().hicache_ratio,
mamba_host_size,
allocator_type=_get_allocator_type(server_args),
layout=server_args.hicache_mem_layout,
layout=get_memory().hicache_mem_layout,
)
entries = [
build_pool_entry(
@@ -754,8 +758,8 @@ def build_hybrid_mamba_stack(
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
write_policy=get_memory().hicache_write_policy,
io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
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
mamba_allocator = params.req_to_token_pool.mamba_allocator
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(
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_pool=full_kv_pool,
@@ -825,8 +829,8 @@ def build_hybrid_mamba_swa_stack(
mamba_pool,
get_memory().hicache_ratio,
mamba_host_size,
allocator_type=server_args.hicache_storage_backend,
layout=server_args.hicache_mem_layout,
allocator_type=get_memory().hicache_storage_backend,
layout=get_memory().hicache_mem_layout,
)
entries = [
build_pool_entry(
@@ -870,8 +874,8 @@ def build_hybrid_mamba_swa_stack(
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
write_policy=get_memory().hicache_write_policy,
io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
@@ -948,8 +952,8 @@ def build_anchor_sidecar_stack(
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
write_policy=get_memory().hicache_write_policy,
io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
@@ -1019,7 +1023,7 @@ def build_full_draft_pools(
pool=pool,
host_to_device_ratio=host_pool_group.logical_size / pool.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),
pool_label="draft",
)
@@ -1045,7 +1049,7 @@ def build_full_draft_pools(
indexer_host_pool = DSAIndexerPoolHost(
pool,
draft_host_pool,
server_args.hicache_mem_layout,
get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args),
)
specs.append(
@@ -1502,7 +1506,7 @@ class _DsaStrategy(StackStrategy):
sidecar_host_pool_factory=lambda kv_host_pool: DSAIndexerPoolHost(
full_kv_pool,
kv_host_pool,
server_args.hicache_mem_layout,
get_memory().hicache_mem_layout,
allocator_type=_get_allocator_type(server_args),
),
prefetch_threshold=prefetch_threshold,
@@ -1732,7 +1736,7 @@ def attach_hybrid_pool_to_unified_cache(
storage_backend=storage_backend,
storage_backend_extra_config=storage_extra_config,
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,
)
_apply_stack_result(cache, kvcache, params, result)
@@ -1796,8 +1800,8 @@ def build_minimax_sparse_hicache_stack(
index_host_pool = MHATokenToKOnlyPoolHost(
index_k_pool,
kv_host_pool,
server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
get_memory().hicache_mem_layout,
allocator_type=get_memory().hicache_storage_backend,
)
entries.append(
build_pool_entry(
@@ -1821,8 +1825,8 @@ def build_minimax_sparse_hicache_stack(
load_cache_event=load_cache_event,
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
write_policy=get_memory().hicache_write_policy,
io_backend=get_memory().hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
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)
},
load_cache_event=load_cache_event,
storage_backend=server_args.hicache_storage_backend,
storage_backend=get_memory().hicache_storage_backend,
use_mla=False,
prefetch_threshold=prefetch_threshold,
model_name=server_args.served_model_name,
model_name=get_serving().served_model_name,
storage_backend_extra_config=extra_config,
enable_storage_metrics=enable_storage_metrics,
)
@@ -1885,9 +1889,9 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
server_args=server_args,
sparse_pool=sparse_pool,
load_cache_event=load_cache_event,
storage_backend=server_args.hicache_storage_backend,
storage_backend=get_memory().hicache_storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=server_args.served_model_name,
model_name=get_serving().served_model_name,
storage_backend_extra_config=extra_config,
enable_storage_metrics=enable_storage_metrics,
)
@@ -1933,17 +1937,17 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
sidecar_pool_name=PoolName.INDEXER,
full_layer_mapping=layer_mapping,
load_cache_event=load_cache_event,
storage_backend=server_args.hicache_storage_backend,
storage_backend=get_memory().hicache_storage_backend,
use_mla=True,
override_kv_cache_dim=kv.kv_cache_dim,
prefetch_threshold=prefetch_threshold,
sidecar_host_pool_factory=lambda kv_host_pool: DSAIndexerPoolHost(
kv,
kv_host_pool,
server_args.hicache_mem_layout,
get_memory().hicache_mem_layout,
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,
enable_storage_metrics=enable_storage_metrics,
)
@@ -22,6 +22,9 @@ from sglang.srt.mem_cache.cpp_radix_tree.radix_tree import (
TreeNodeCpp,
)
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.runtime_context import (
get_memory,
)
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
@@ -59,7 +62,7 @@ class RadixCacheCpp(BasePrefixCache):
self.ongoing_load_back: Set[IOHandle] = set()
# todo: dynamically adjust the 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.device = self.token_to_kv_pool_allocator.device
@@ -72,7 +75,7 @@ class RadixCacheCpp(BasePrefixCache):
if params.enable_metrics:
self.init_metrics_collector()
if not server_args.enable_hierarchical_cache:
if not get_memory().enable_hierarchical_cache:
self.tree = RadixTreeCpp(
disabled=self.disable,
page_size=self.page_size,
@@ -170,10 +170,13 @@ from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import (
ensure_published,
get_context,
get_device,
get_exec,
get_global_dwdp_manager,
get_lora,
get_memory,
get_model,
get_observability,
get_parallel,
get_schedule,
get_spec,
@@ -279,7 +282,7 @@ def resolve_draft_attention_backend(
"""
if not is_draft_worker:
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:
@@ -323,7 +326,6 @@ class ModelRunner:
# workers so they reuse target's resolved sizes (replaces legacy
# `server_args._draft_pool_config` mutation hack).
self.memory_pool_config = memory_pool_config
self.device = server_args.device
self.gpu_id = gpu_id
self.ps = ps
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
# draft runners, which serve adapters' target model unadapted.
self.lora_manager: Optional[LoRAManager] = None
self.device = get_device().device
self.draft_attention_backend = resolve_draft_attention_backend(
draft_attention_backend=draft_attention_backend,
server_args=server_args,
@@ -356,7 +359,7 @@ class ModelRunner:
model_config.is_multimodal_chunked_prefill_supported
)
self.spec_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
get_spec().speculative_algorithm
)
self.capture_tail_hooks = []
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.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
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._pending_elastic_scale_update = None
self.init_new_workspace = False
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.init_startup_observability()
@@ -385,7 +388,7 @@ class ModelRunner:
self.init_spec_aux_hidden_state()
# Apply the rank zero filter to logger
if server_args.show_time_cost:
if get_observability().show_time_cost:
enable_show_time_cost()
misc_utils.maybe_disable_chunked_prefix_cache(
@@ -1348,7 +1351,7 @@ class ModelRunner:
else False
),
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
@@ -25,6 +25,11 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
trigger_init_weights_send_group_for_remote_instance_request,
)
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.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
target's format cannot answer for it."""
if (
(load_format or server_args.load_format) == LoadFormat.REMOTE_INSTANCE
and server_args.remote_instance_weight_loader_backend
== RemoteInstanceWeightLoaderBackend.NCCL
):
load_format or get_model().load_format
) == LoadFormat.REMOTE_INSTANCE and get_model().remote_instance_weight_loader_backend == RemoteInstanceWeightLoaderBackend.NCCL:
if tp_rank == 0:
instance_ip = NetworkAddress.resolve_host(socket.gethostname())
t = threading.Thread(
target=trigger_init_weights_send_group_for_remote_instance_request,
args=(
server_args.remote_instance_weight_loader_seed_instance_ip,
server_args.remote_instance_weight_loader_seed_instance_service_port,
server_args.remote_instance_weight_loader_send_weights_group_ports,
get_model().remote_instance_weight_loader_seed_instance_ip,
get_model().remote_instance_weight_loader_seed_instance_service_port,
get_model().remote_instance_weight_loader_send_weights_group_ports,
instance_ip,
),
)
@@ -111,12 +114,12 @@ def load_kv_cache_scales(
defaulted: a fallback to ``server_args`` would be a hidden global read for
any future caller that forgets to pass one."""
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)):
model.load_kv_cache_scales(server_args.quantization_param_path)
model.load_kv_cache_scales(get_model().quantization_param_path)
logger.info(
"Loaded KV cache scaling factors from %s",
server_args.quantization_param_path,
get_model().quantization_param_path,
)
else:
raise RuntimeError(
@@ -154,13 +157,13 @@ def report_online_quantization(*, model, server_args: ServerArgs) -> None:
getattr(model, "quant_config", None), "quantized_layers", None
)
if (
server_args.quantization is not None
get_model().quantization is not None
and isinstance(quantized_layers, tuple)
and len(quantized_layers) == 2
):
layer_types, quantized_layers_count = quantized_layers
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,
pp_rank: int,
) -> None:
if server_args.debug_tensor_dump_output_folder is not None:
dump_folder = server_args.debug_tensor_dump_output_folder
if get_observability().debug_tensor_dump_output_folder is not None:
dump_folder = get_observability().debug_tensor_dump_output_folder
if spec_algorithm.is_eagle():
role = "draft" if is_draft_worker else "target"
dump_folder = os.path.join(dump_folder, role)
register_forward_hook_for_model(
model,
dump_folder,
server_args.debug_tensor_dump_layers,
get_observability().debug_tensor_dump_layers,
tp_size,
tp_rank,
pp_rank,
@@ -203,28 +206,28 @@ def build_load_config(
from sglang.srt.configs.modelopt_config import ModelOptConfig
modelopt_config = ModelOptConfig(
quant=server_args.modelopt_quant,
checkpoint_restore_path=server_args.modelopt_checkpoint_restore_path,
checkpoint_save_path=server_args.modelopt_checkpoint_save_path,
export_path=server_args.modelopt_export_path,
quantize_and_serve=server_args.quantize_and_serve,
quant=get_model().modelopt_quant,
checkpoint_restore_path=get_model().modelopt_checkpoint_restore_path,
checkpoint_save_path=get_model().modelopt_checkpoint_save_path,
export_path=get_model().modelopt_export_path,
quantize_and_serve=get_model().quantize_and_serve,
)
return LoadConfig(
load_format=load_format or server_args.load_format,
download_dir=server_args.download_dir,
model_loader_extra_config=server_args.model_loader_extra_config,
load_format=load_format or get_model().load_format,
download_dir=get_model().download_dir,
model_loader_extra_config=get_model().model_loader_extra_config,
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_service_port=server_args.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_backend=server_args.remote_instance_weight_loader_backend,
remote_instance_weight_loader_seed_instance_ip=get_model().remote_instance_weight_loader_seed_instance_ip,
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=get_model().remote_instance_weight_loader_send_weights_group_ports,
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_session_id=remote_instance_weight_transporter_session_id,
modelexpress_url=server_args.modelexpress_url,
modelexpress_transport=server_args.modelexpress_transport,
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,
weight_cache_mode=weight_cache_mode,
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
weight reload) can't overwrite the captured fallback format.
"""
if server_args.weight_cache_mode == "off":
if get_model().weight_cache_mode == "off":
return
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
monkey_patch_vllm_parallel_state()
enable_cpu_backup = server_args.enable_weights_cpu_backup or (
is_draft_worker and server_args.enable_draft_weights_cpu_backup
enable_cpu_backup = get_exec().features.enable_weights_cpu_backup or (
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
# 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:
logger.warning(
"[ModelRunner] Disabling weights CPU backup in zero-copy IPC mode — "
@@ -6,6 +6,11 @@ from typing import TYPE_CHECKING, Any, Optional
import msgspec
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.runtime_context import (
configured_tp_size,
get_model,
get_spec,
)
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
@@ -75,13 +80,13 @@ def _resolve_eagle_aux_hidden_state(
if (
(spec_algorithm.is_eagle() or spec_algorithm.is_standalone())
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
draft_model_config = ModelConfig.from_server_args(
server_args,
model_path=server_args.speculative_draft_model_path,
model_revision=server_args.speculative_draft_model_revision,
model_path=get_spec().speculative_draft_model_path,
model_revision=get_spec().speculative_draft_model_revision,
is_draft_model=True,
)
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.
draft_model_config = ModelConfig.from_server_args(
server_args,
model_path=(server_args.speculative_draft_model_path),
model_revision=server_args.speculative_draft_model_revision,
model_path=(get_spec().speculative_draft_model_path),
model_revision=get_spec().speculative_draft_model_revision,
is_draft_model=True,
)
dflash_draft_config = parse_dflash_draft_config(
@@ -221,23 +226,23 @@ def _resolve_dflash_draft_cell_size(
try:
_, 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=(
server_args.speculative_draft_kv_cache_dtype
get_spec().speculative_draft_kv_cache_dtype
),
model=None,
model_dtype=draft_model_config.dtype,
is_draft_worker=True,
is_dflash=True,
speculative_draft_attention_backend=(
server_args.speculative_draft_attention_backend
get_spec().speculative_draft_attention_backend
),
)
return dflash_draft_cell_size_per_token(
draft_model_config=draft_model_config,
draft_num_layers=draft_num_layers,
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
logger.warning(
@@ -20,7 +20,18 @@ from sglang.srt.model_loader.weight_utils import (
CheckpointFilePrefetchHandle,
)
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:
from sglang.srt.configs.model_config import ModelConfig
@@ -97,37 +108,37 @@ class StartupWeightLoadOptions:
server_args: ServerArgs,
is_draft_worker: bool,
) -> StartupWeightLoadOptions:
cuda_graph_config = server_args.cuda_graph_config
cuda_graph_config = get_exec().graph.cuda_graph_config
cuda_graph_enabled = any(
getattr(cuda_graph_config, phase).backend != Backend.DISABLED
for phase in Phase.ALL
)
return cls(
device=server_args.device,
device=get_device().device,
is_cuda_platform=current_platform.is_cuda(),
cuda_graph_enabled=cuda_graph_enabled,
prefill_cuda_graph_backend=cuda_graph_config.prefill.backend,
is_draft_worker=is_draft_worker,
speculative_algorithm=server_args.speculative_algorithm,
tp_size=server_args.tp_size,
attn_cp_size=server_args.attn_cp_size,
dcp_size=server_args.dcp_size,
pp_size=server_args.pp_size,
speculative_algorithm=get_spec().speculative_algorithm,
tp_size=configured_tp_size(),
attn_cp_size=configured_attn_cp_size(),
dcp_size=configured_dcp_size(),
pp_size=configured_pp_size(),
dp_size=get_parallel().dp_size,
ep_size=get_parallel().ep_size,
cpu_offload_gb=server_args.cpu_offload_gb,
offload_group_size=server_args.offload_group_size,
enable_memory_saver=server_args.enable_memory_saver,
enable_weights_cpu_backup=server_args.enable_weights_cpu_backup,
enable_lora=server_args.enable_lora,
has_lora_paths=bool(server_args.lora_paths),
weight_loader_disable_mmap=server_args.weight_loader_disable_mmap,
cpu_offload_gb=get_exec().offload.cpu_offload_gb,
offload_group_size=get_exec().offload.offload_group_size,
enable_memory_saver=get_exec().features.enable_memory_saver,
enable_weights_cpu_backup=get_exec().features.enable_weights_cpu_backup,
enable_lora=get_lora().enable_lora,
has_lora_paths=bool(get_lora().lora_paths),
weight_loader_disable_mmap=get_model().weight_loader_disable_mmap,
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),
enable_torch_compile=server_args.enable_torch_compile,
prefetch_num_threads=server_args.weight_loader_prefetch_num_threads,
has_custom_weight_loader=bool(get_model().custom_weight_loader),
enable_torch_compile=get_exec().graph.enable_torch_compile,
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.model_executor.forward_batch_info import ForwardMode
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.utils import get_bool_env_var
from sglang.srt.utils.gauge_histogram import GaugeHistogram
@@ -957,7 +963,7 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
# =================================================================
# 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(
name="sglang:prefill_delayer_wait_forward_passes",
documentation="Histogram of forward passes waited by prefill delayer.",
@@ -966,7 +972,7 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
set(
x
for x in (
server_args.prefill_delayer_forward_passes_buckets
get_schedule().prefill_delayer_forward_passes_buckets
or [5, 20, 50, 100, 200]
)
if x < max_delay
@@ -981,7 +987,7 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
labelnames=labels.keys(),
buckets=sorted(
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]
)
# Need bucket "<=0" for zero-delay cases
@@ -1077,13 +1083,14 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
enable_lora: bool,
enable_hierarchical_cache: bool,
) -> SchedulerMetricsCollectorContext:
enable_metrics = server_args.enable_metrics
enable_metrics = get_observability().enable_metrics
is_stats_logging_rank = ps.attn_tp_rank == 0
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(
server_args.kv_events_config
get_observability().kv_events_config
and ps.pp_rank == 0
and ps.attn_tp_rank == 0
and ps.attn_cp_rank == 0
@@ -1091,10 +1098,10 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
collector: Optional[SchedulerMetricsCollector] = None
if enable_metrics:
engine_type = DisaggregationMode.to_engine_type(
server_args.disaggregation_mode
get_disagg().disaggregation_mode
)
labels = {
"model_name": server_args.served_model_name,
"model_name": get_serving().served_model_name,
"engine_type": engine_type,
"tp_rank": tp_rank,
"pp_rank": pp_rank,
@@ -1104,8 +1111,8 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
labels["priority"] = ""
if dp_rank is not None:
labels["dp_rank"] = dp_rank
if server_args.extra_metric_labels:
labels.update(server_args.extra_metric_labels)
if get_observability().extra_metric_labels:
labels.update(get_observability().extra_metric_labels)
scheduler_collector_cls = resolve_collector_class(
server_args, STAT_LOGGER_ROLE_SCHEDULER, cls
)
@@ -1113,7 +1120,7 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
labels=labels,
enable_lora=enable_lora,
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,
)
return SchedulerMetricsCollectorContext(
@@ -1568,7 +1575,7 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
documentation="Histogram of prompt token length.",
labelnames=labels.keys(),
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(
@@ -1576,7 +1583,7 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
documentation="Histogram of uncached (compute) prompt token length.",
labelnames=labels.keys(),
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(
@@ -1584,7 +1591,7 @@ class TokenizerMetricsCollector(_StatLoggerDIMixin):
documentation="Histogram of generation token length.",
labelnames=labels.keys(),
buckets=generate_buckets(
server_args.generation_tokens_buckets,
get_observability().generation_tokens_buckets,
default_bucket_prompt_tokens,
),
)
+27 -5
View File
@@ -678,18 +678,22 @@ class _ConfigBag:
def _build_config_bags(server_args: Any) -> dict:
"""Snapshot resolved ``server_args`` into the namespace bag tree, driven by
the ``NS(...)`` metadata on the dataclass fields. Returns
"""Snapshot the resolution result into the namespace bag tree, driven by
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.``).
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
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.overrides import resolution_result
_MISSING = object()
tops: dict = {}
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:
# Every NS-declared field is a dataclass field, so a resolved config
# 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_parallel_config = ctx.parallel._config
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")
# Underscore names seed private property caches (the strict guard
@@ -1040,7 +1047,18 @@ class _ServerArgsOverride:
raise ValueError(
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
# strict mutation guard unarmed on the published object — mark it
# 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")
def configured_dcp_size() -> int:
return _configured_parallel("dcp_size")
def is_ep_joiner() -> bool:
"""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 (
attention_backends_of,
declare_direct_writes,
mamba_extra_buffer_lazy_of,
mamba_extra_buffer_of,
remote_instance_transfer_engine_of,
@@ -3677,6 +3678,12 @@ class ServerArgs:
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
# short-circuit (none/dummy model paths) so run_post_process_pass and
# direct handler invocations can rely on it even when
@@ -3742,8 +3749,13 @@ class ServerArgs:
self._handle_mps_backends()
self._handle_xpu_backends()
# Allow OOT platform plugins to apply server args defaults.
current_platform.apply_server_args_defaults(self)
# OOT platform plugins set fields directly (an interface this tree
# 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.
gpu_mem = get_device_memory_capacity(self.device)
@@ -33,7 +33,12 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode,
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.speculative.base_spec_worker import BaseSpecWorker
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
@@ -279,9 +284,7 @@ class DFlashWorkerV2(BaseSpecWorker):
self._need_mamba_verify_commit = False
self.page_size = get_schedule().page_size
# Normalized in arg_groups.speculative_hook.handle_speculative_decoding.
self.draft_window_size: Optional[int] = (
server_args.speculative_draft_window_size
)
self.draft_window_size: Optional[int] = get_spec().speculative_draft_window_size
self.use_compact_draft_cache = self.draft_window_size is not None
self.device = target_worker.device
@@ -305,11 +308,11 @@ class DFlashWorkerV2(BaseSpecWorker):
draft_config = parse_dflash_draft_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.
self.block_size = int(draft_config.resolve_block_size(default=16))
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
if model_block_size is 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
from sglang.srt.runtime_context import (
get_model,
get_spec,
)
from sglang.srt.speculative.dflash_utils import parse_dflash_draft_config
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.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:
return False
draft_hf_config = get_config(
draft_model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.speculative_draft_model_revision,
model_override_args=json.loads(server_args.json_model_override_args),
model_config_parser=server_args.model_config_parser,
trust_remote_code=get_model().trust_remote_code,
revision=get_spec().speculative_draft_model_revision,
model_override_args=json.loads(get_model().json_model_override_args),
model_config_parser=get_model().model_config_parser,
)
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]:
"""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
resolving = resolved_view(server_args)
draft_hf_config = get_config(
server_args.speculative_draft_model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.speculative_draft_model_revision,
model_override_args=json.loads(server_args.json_model_override_args),
resolving.speculative_draft_model_path,
trust_remote_code=resolving.trust_remote_code,
revision=resolving.speculative_draft_model_revision,
model_override_args=json.loads(resolving.json_model_override_args),
)
return parse_dspark_draft_config(draft_hf_config=draft_hf_config).resolve_gamma(
default=None
@@ -21,6 +21,7 @@ from sglang.srt.model_executor.forward_batch_info import (
compute_position,
)
from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_parallel,
get_schedule,
@@ -106,7 +107,7 @@ class DSparkWorkerV2(BaseSpecWorker):
self._draft_dp_context_enabled = (
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 = (
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
runtime_config = resolve_runtime_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),
)
self.gamma = runtime_config.gamma
@@ -47,6 +47,7 @@ from sglang.srt.model_executor.runner import (
)
from sglang.srt.runtime_context import (
get_context,
get_device,
get_exec,
get_model,
get_parallel,
@@ -145,14 +146,14 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.target_worker = target_worker
# Args for easy access
self.device = server_args.device
self.topk = server_args.speculative_eagle_topk
self.device = get_device().device
self.topk = get_spec().speculative_eagle_topk
if get_spec().speculative_use_rejection_sampling:
assert self.topk == 1, "Chain speculative sampling supports only topk=1"
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
get_spec().speculative_algorithm
)
self._rebuild_topk1_chain_buffers()
@@ -1060,16 +1061,16 @@ class EAGLEWorkerV2(BaseSpecWorker):
# Parse arguments
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id
self.device = server_args.device
self.device = get_device().device
self._target_worker = target_worker
self.page_size = get_schedule().page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
get_spec().speculative_algorithm
)
self._draft_worker = EagleDraftWorker(
@@ -1082,10 +1083,10 @@ class EAGLEWorkerV2(BaseSpecWorker):
# Adaptive speculative
self.adaptive_controller: Optional[AdaptiveController] = None
if server_args.speculative_adaptive:
if get_spec().speculative_adaptive:
self.adaptive_controller = AdaptiveController(
self,
config_path=server_args.speculative_adaptive_config,
config_path=get_spec().speculative_adaptive_config,
)
# 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.runtime_context import (
attention_backends,
get_device,
get_parallel,
get_schedule,
get_spec,
@@ -107,16 +108,16 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
EagleDraftWorkerBase.__init__(self)
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id
self.device = server_args.device
self.device = get_device().device
self.target_worker = target_worker
self.page_size = get_schedule().page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
get_spec().speculative_algorithm
)
assert self.speculative_algorithm.is_frozen_kv_mtp(), (
"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
# no KV, so we mirror the relevant setup and build a FrozenKVMTPDraftWorker.
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id
self.device = server_args.device
self.device = get_device().device
self._target_worker = target_worker
self.page_size = get_schedule().page_size
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 = (
@@ -720,7 +721,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
# Frozen MTP does not wire the adaptive controller yet.
assert (
not server_args.speculative_adaptive
not get_spec().speculative_adaptive
), "Frozen-KV MTP does not support adaptive speculative decoding yet."
self.adaptive_controller = None
@@ -43,7 +43,12 @@ from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
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.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
from sglang.srt.speculative.draft_utils import DraftBackendFactory
@@ -128,22 +133,22 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.model_config = target_worker.model_config
# Args for easy access
self.device = server_args.device
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.device = get_device().device
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
# 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.
# Single-CG runner samples in-graph (_sample_draft_proposal); per-step
# 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, (
"multi-layer EAGLE requires speculative_num_draft_tokens == "
"speculative_num_steps + 1, "
f"got {self.speculative_num_draft_tokens} and {self.speculative_num_steps}"
)
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
get_spec().speculative_algorithm
)
self._rebuild_topk1_chain_buffers()
@@ -928,15 +933,15 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
# Parse arguments
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.gpu_id = gpu_id
self.device = server_args.device
self.device = get_device().device
self._target_worker = target_worker
self.page_size = get_schedule().page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
get_spec().speculative_algorithm
)
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.model_executor.forward_batch_info import ForwardMode
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.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus
@@ -88,19 +92,19 @@ class NGRAMWorker(BaseSpecWorker):
super().__init__()
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.model_runner = target_worker.model_runner
self.tp_rank = ps.tp_rank
self.page_size = get_schedule().page_size
self.draft_token_num: int = server_args.speculative_num_draft_tokens
self.max_trie_depth: int = server_args.speculative_ngram_max_trie_depth
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.draft_token_num: int = get_spec().speculative_num_draft_tokens
self.max_trie_depth: int = get_spec().speculative_ngram_max_trie_depth
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
# req_to_token_pool / token_to_kv_pool_allocator are set in
# alloc_memory_pool(), after the target pools are allocated.
self.device = server_args.device
self.device = get_device().device
self.adaptive_controller = None
# 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.ngram_corpus = NgramCorpus(
min_bfs_breadth=server_args.speculative_ngram_min_bfs_breadth,
max_bfs_breadth=server_args.speculative_ngram_max_bfs_breadth,
match_type=server_args.speculative_ngram_match_type,
capacity=server_args.speculative_ngram_capacity,
max_trie_depth=server_args.speculative_ngram_max_trie_depth,
draft_token_num=server_args.speculative_num_draft_tokens,
external_sam_budget=server_args.speculative_ngram_external_sam_budget,
external_corpus_max_tokens=server_args.speculative_ngram_external_corpus_max_tokens,
min_bfs_breadth=get_spec().speculative_ngram_min_bfs_breadth,
max_bfs_breadth=get_spec().speculative_ngram_max_bfs_breadth,
match_type=get_spec().speculative_ngram_match_type,
capacity=get_spec().speculative_ngram_capacity,
max_trie_depth=get_spec().speculative_ngram_max_trie_depth,
draft_token_num=get_spec().speculative_num_draft_tokens,
external_sam_budget=get_spec().speculative_ngram_external_sam_budget,
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 (
iter_external_corpus_chunks,
)
corpus_path = server_args.speculative_ngram_external_corpus_path
corpus_path = get_spec().speculative_ngram_external_corpus_path
chunks = list(
iter_external_corpus_chunks(
corpus_path,
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)
@@ -10,7 +10,12 @@ from sglang.srt.layers.moe.utils import (
speculative_moe_backend_context,
)
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.speculative.adaptive_runtime_state import (
AdaptiveController,
@@ -53,12 +58,12 @@ class StandaloneDraftWorker(EagleDraftWorker):
self.target_worker = target_worker
# Args for easy access
self.device = server_args.device
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.device = get_device().device
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
get_spec().speculative_algorithm
)
self._rebuild_topk1_chain_buffers()
@@ -158,15 +163,15 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
# Parse arguments
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.gpu_id = gpu_id
self.device = server_args.device
self.device = get_device().device
self._target_worker = target_worker
self.page_size = get_schedule().page_size
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
+14
View File
@@ -2105,6 +2105,20 @@ def server_args_variant(server_args, **fields):
}
if 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():
object.__setattr__(variant, name, value)
return variant