config: stop handing the record to code that does not read it (#36252)

This commit is contained in:
Cheng Wan
2026-08-26 05:02:17 -07:00
committed by GitHub
parent d7b144f64e
commit ae5feb4b9c
60 changed files with 168 additions and 280 deletions
+2 -2
View File
@@ -538,7 +538,7 @@ def decode(input_token_ids, batch, model_runner):
def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner): def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
if require_mlp_sync(model_runner.server_args): if require_mlp_sync():
prepare_mlp_sync_batch_raw( prepare_mlp_sync_batch_raw(
batch, batch,
model_runner=model_runner, model_runner=model_runner,
@@ -548,7 +548,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
tp_group=model_runner.tp_group, tp_group=model_runner.tp_group,
get_idle_batch=None, get_idle_batch=None,
disable_cuda_graph=cuda_graph_fully_disabled(), disable_cuda_graph=cuda_graph_fully_disabled(),
require_mlp_tp_gather=require_mlp_tp_gather(model_runner.server_args), require_mlp_tp_gather=require_mlp_tp_gather(),
disable_overlap_schedule=get_schedule().disable_overlap_schedule, disable_overlap_schedule=get_schedule().disable_overlap_schedule,
offload_tags=set(), offload_tags=set(),
) )
@@ -14,7 +14,7 @@ from sglang.srt.runtime_context import (
from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs pass
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -307,9 +307,7 @@ def get_mooncake_transfer_engine() -> Optional[MooncakeTransferEngine]:
return _mooncake_transfer_engine return _mooncake_transfer_engine
def maybe_init_shared_mooncake_transfer_engine( def maybe_init_shared_mooncake_transfer_engine(*, gpu_id: int) -> None:
*, server_args: ServerArgs, gpu_id: int
) -> None:
""" """
Need MooncakeTransferEngine when: Need MooncakeTransferEngine when:
1) PD disaggregation uses mooncake for KV transfer (prefill/decode) 1) PD disaggregation uses mooncake for KV transfer (prefill/decode)
+2 -2
View File
@@ -111,14 +111,14 @@ class ElasticEPStateManager:
inst.ep_join_rank_offset = get_parallel().config.ep_join_rank_offset inst.ep_join_rank_offset = get_parallel().config.ep_join_rank_offset
if server_args.is_ep_joiner: if server_args.is_ep_joiner:
cls._init_joiner_state(inst, server_args) cls._init_joiner_state(inst)
cls._instance = inst cls._instance = inst
return cls._instance return cls._instance
@classmethod @classmethod
def _init_joiner_state(cls, inst: ElasticEPState, server_args: ServerArgs) -> None: def _init_joiner_state(cls, inst: ElasticEPState) -> None:
global_rank = torch.distributed.get_rank() global_rank = torch.distributed.get_rank()
inst.active_ranks.zero_() inst.active_ranks.zero_()
inst.active_ranks[global_rank] = 1 inst.active_ranks[global_rank] = 1
@@ -2142,7 +2142,6 @@ def _get_vlm_warmup_image_base64(model_info: dict) -> str:
async def _send_disaggregation_warmup_requests( async def _send_disaggregation_warmup_requests(
server_args: ServerArgs,
url: str, url: str,
headers: Dict[str, str], headers: Dict[str, str],
ssl_verify: Union[bool, str], ssl_verify: Union[bool, str],
@@ -2321,7 +2320,6 @@ def _execute_server_warmup(server_args: ServerArgs):
logger.info(f"Start of pd disaggregation warmup ...") logger.info(f"Start of pd disaggregation warmup ...")
status_codes = asyncio.run( status_codes = asyncio.run(
_send_disaggregation_warmup_requests( _send_disaggregation_warmup_requests(
server_args=server_args,
url=url, url=url,
headers=headers, headers=headers,
ssl_verify=ssl_verify, ssl_verify=ssl_verify,
+10 -16
View File
@@ -18,11 +18,10 @@ from sglang.srt.eplb.expert_location import (
get_global_expert_location_metadata, get_global_expert_location_metadata,
) )
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
from sglang.srt.runtime_context import get_model from sglang.srt.runtime_context import get_exec, get_model, get_parallel
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -31,7 +30,6 @@ class EPLBManager:
def __init__( def __init__(
self, self,
*, *,
server_args: ServerArgs,
model_config: ModelConfig, model_config: ModelConfig,
ps: Any, ps: Any,
get_model: Callable[[], nn.Module], get_model: Callable[[], nn.Module],
@@ -43,7 +41,6 @@ class EPLBManager:
# These collaborators are set on ModelRunner AFTER EPLBManager is # These collaborators are set on ModelRunner AFTER EPLBManager is
# constructed (model load, expert_backup_client, weight_updater), so # constructed (model load, expert_backup_client, weight_updater), so
# they are read through getters at rebalance time, not captured here. # they are read through getters at rebalance time, not captured here.
self._server_args = server_args
self._model_config = model_config self._model_config = model_config
self._ps = ps self._ps = ps
self._get_model = get_model self._get_model = get_model
@@ -51,16 +48,16 @@ class EPLBManager:
self._get_expert_backup_client = get_expert_backup_client self._get_expert_backup_client = get_expert_backup_client
self._get_weight_updater = get_weight_updater self._get_weight_updater = get_weight_updater
self._rebalance_layers_per_chunk = ( self._rebalance_layers_per_chunk = (
self._server_args.eplb_rebalance_layers_per_chunk get_exec().moe.eplb_rebalance_layers_per_chunk
) )
self._rebalance_num_iterations = self._server_args.eplb_rebalance_num_iterations self._rebalance_num_iterations = get_exec().moe.eplb_rebalance_num_iterations
self._rebalance_disabled_reason = None self._rebalance_disabled_reason = None
self._rebalance_disabled_logged = False self._rebalance_disabled_logged = False
# Otherwise, the circular buffer will contain stale data. If the case is needed, it can be implemented. # Otherwise, the circular buffer will contain stale data. If the case is needed, it can be implemented.
assert ( assert (
self._server_args.eplb_rebalance_num_iterations get_exec().moe.eplb_rebalance_num_iterations
>= self._server_args.expert_distribution_recorder_buffer_size >= get_exec().moe.expert_distribution_recorder_buffer_size
), "eplb_rebalance_num_iterations must be greater than expert_distribution_recorder_buffer_size" ), "eplb_rebalance_num_iterations must be greater than expert_distribution_recorder_buffer_size"
if not get_global_expert_distribution_recorder().recording: if not get_global_expert_distribution_recorder().recording:
@@ -160,7 +157,7 @@ class EPLBManager:
model=self._get_model(), model=self._get_model(),
new_expert_location_metadata=expert_location_metadata, new_expert_location_metadata=expert_location_metadata,
update_layer_ids=chunk_layer_ids, update_layer_ids=chunk_layer_ids,
nnodes=self._server_args.nnodes, nnodes=get_parallel().config.nnodes,
tp_rank=( tp_rank=(
self._elastic_global_rank() self._elastic_global_rank()
if is_post_scale_rebalance if is_post_scale_rebalance
@@ -169,7 +166,7 @@ class EPLBManager:
use_flat_topology=is_post_scale_rebalance, use_flat_topology=is_post_scale_rebalance,
expert_backup_client=self._get_expert_backup_client(), expert_backup_client=self._get_expert_backup_client(),
update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk, update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk,
ep_dispatch_algorithm=self._server_args.ep_dispatch_algorithm, ep_dispatch_algorithm=get_exec().moe.ep_dispatch_algorithm,
init_lplb_solvers_callable=lambda: init_lplb_solvers( init_lplb_solvers_callable=lambda: init_lplb_solvers(
model_config=self._model_config model_config=self._model_config
), ),
@@ -193,7 +190,6 @@ class EPLBManager:
) -> ExpertLocationMetadata: ) -> ExpertLocationMetadata:
if not broadcast_over_world: if not broadcast_over_world:
return ExpertLocationMetadata.init_by_eplb( return ExpertLocationMetadata.init_by_eplb(
self._server_args,
self._model_config, self._model_config,
logical_count, logical_count,
) )
@@ -204,7 +200,6 @@ class EPLBManager:
# the mapping chosen for the expanded world. # the mapping chosen for the expanded world.
if dist.get_rank() == 0: if dist.get_rank() == 0:
computed_metadata = ExpertLocationMetadata.init_by_eplb( computed_metadata = ExpertLocationMetadata.init_by_eplb(
self._server_args,
self._model_config, self._model_config,
logical_count, logical_count,
# Arbitrary append topologies may not preserve node divisibility. # Arbitrary append topologies may not preserve node divisibility.
@@ -220,14 +215,13 @@ class EPLBManager:
dist.broadcast(physical_to_logical_map, src=0) dist.broadcast(physical_to_logical_map, src=0)
return ExpertLocationMetadata.init_by_mapping( return ExpertLocationMetadata.init_by_mapping(
self._server_args,
self._model_config, self._model_config,
physical_to_logical_map, physical_to_logical_map,
moe_ep_rank=self._elastic_global_rank(), moe_ep_rank=self._elastic_global_rank(),
) )
def _elastic_global_rank(self) -> int: def _elastic_global_rank(self) -> int:
return self._ps.tp_rank + self._server_args.ep_join_rank_offset return self._ps.tp_rank + get_parallel().config.ep_join_rank_offset
def _check_rebalance_needed(self, average_utilization_rate_over_window): def _check_rebalance_needed(self, average_utilization_rate_over_window):
if average_utilization_rate_over_window is None: if average_utilization_rate_over_window is None:
@@ -235,10 +229,10 @@ class EPLBManager:
if ( if (
average_utilization_rate_over_window average_utilization_rate_over_window
> self._server_args.eplb_min_rebalancing_utilization_threshold > get_exec().moe.eplb_min_rebalancing_utilization_threshold
): ):
logger.info( logger.info(
f"[EPLBManager] Skipped ep rebalancing: current GPU utilization {average_utilization_rate_over_window:.2f} > minimum rebalance threshold {self._server_args.eplb_min_rebalancing_utilization_threshold:.2f}" f"[EPLBManager] Skipped ep rebalancing: current GPU utilization {average_utilization_rate_over_window:.2f} > minimum rebalance threshold {get_exec().moe.eplb_min_rebalancing_utilization_threshold:.2f}"
) )
return False return False
@@ -671,12 +671,10 @@ class _Accumulator(ABC):
expert_location_metadata: ExpertLocationMetadata, expert_location_metadata: ExpertLocationMetadata,
rank: int, rank: int,
) -> _Accumulator: ) -> _Accumulator:
return _Accumulator.get_class(server_args)( return _Accumulator.get_class()(server_args, expert_location_metadata, rank)
server_args, expert_location_metadata, rank
)
@staticmethod @staticmethod
def get_class(server_args: ServerArgs) -> Type[_Accumulator]: def get_class() -> Type[_Accumulator]:
return { return {
"stat": _StatAccumulator, "stat": _StatAccumulator,
"stat_approx": _StatAccumulator, "stat_approx": _StatAccumulator,
+7 -21
View File
@@ -32,7 +32,6 @@ from sglang.srt.runtime_context import (
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -105,11 +104,9 @@ class ExpertLocationMetadata:
# -------------------------------- construction ------------------------------------ # -------------------------------- construction ------------------------------------
@staticmethod @staticmethod
def init_trivial( def init_trivial(model_config: ModelConfig, moe_ep_rank: int):
server_args: ServerArgs, model_config: ModelConfig, moe_ep_rank: int
):
"""Trivial location - logical expert i corresponds to physical expert i""" """Trivial location - logical expert i corresponds to physical expert i"""
common = ExpertLocationMetadata._init_common(server_args, model_config) common = ExpertLocationMetadata._init_common(model_config)
if common is None: if common is None:
return None return None
@@ -131,7 +128,6 @@ class ExpertLocationMetadata:
) )
return ExpertLocationMetadata.init_by_mapping( return ExpertLocationMetadata.init_by_mapping(
server_args,
model_config, model_config,
physical_to_logical_map=physical_to_logical_map, physical_to_logical_map=physical_to_logical_map,
moe_ep_rank=moe_ep_rank, moe_ep_rank=moe_ep_rank,
@@ -139,7 +135,6 @@ class ExpertLocationMetadata:
@staticmethod @staticmethod
def init_by_mapping( def init_by_mapping(
server_args: ServerArgs,
model_config: ModelConfig, model_config: ModelConfig,
physical_to_logical_map, physical_to_logical_map,
moe_ep_rank: int = None, moe_ep_rank: int = None,
@@ -148,7 +143,7 @@ class ExpertLocationMetadata:
physical_to_logical_map = torch.tensor(physical_to_logical_map) physical_to_logical_map = torch.tensor(physical_to_logical_map)
physical_to_logical_map = physical_to_logical_map.to(get_device().device) physical_to_logical_map = physical_to_logical_map.to(get_device().device)
common = ExpertLocationMetadata._init_common(server_args, model_config) common = ExpertLocationMetadata._init_common(model_config)
if common is None: if common is None:
return None return None
@@ -179,7 +174,6 @@ class ExpertLocationMetadata:
@staticmethod @staticmethod
def init_by_eplb( def init_by_eplb(
server_args: ServerArgs,
model_config: ModelConfig, model_config: ModelConfig,
logical_count: torch.Tensor, logical_count: torch.Tensor,
*, *,
@@ -193,7 +187,7 @@ class ExpertLocationMetadata:
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
common = ExpertLocationMetadata._init_common(server_args, model_config) common = ExpertLocationMetadata._init_common(model_config)
if common is None: if common is None:
return None return None
@@ -229,7 +223,7 @@ class ExpertLocationMetadata:
) )
@staticmethod @staticmethod
def _init_common(server_args: ServerArgs, model_config: ModelConfig): def _init_common(model_config: ModelConfig):
from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.runtime_context import get_exec, get_parallel
model_config_for_expert_location = ( model_config_for_expert_location = (
@@ -526,9 +520,6 @@ def broadcast_global_expert_location_metadata(
src_rank: int = 0, src_rank: int = 0,
group: Optional[torch.distributed.ProcessGroup] = None, group: Optional[torch.distributed.ProcessGroup] = None,
) -> ExpertLocationMetadata: ) -> ExpertLocationMetadata:
from sglang.srt.runtime_context import get_server_args
server_args = get_server_args()
metadata = get_global_expert_location_metadata() metadata = get_global_expert_location_metadata()
assert metadata is not None assert metadata is not None
@@ -537,7 +528,6 @@ def broadcast_global_expert_location_metadata(
metadata.physical_to_logical_map, src=src_rank, group=group metadata.physical_to_logical_map, src=src_rank, group=group
) )
metadata = ExpertLocationMetadata.init_by_mapping( metadata = ExpertLocationMetadata.init_by_mapping(
server_args,
model_config, model_config,
metadata.physical_to_logical_map, metadata.physical_to_logical_map,
moe_ep_rank=moe_ep_rank, moe_ep_rank=moe_ep_rank,
@@ -785,15 +775,12 @@ class ModelConfigForExpertLocation:
def compute_initial_expert_location_metadata( def compute_initial_expert_location_metadata(
server_args: ServerArgs,
model_config: ModelConfig, model_config: ModelConfig,
moe_ep_rank: int, moe_ep_rank: int,
) -> Optional[ExpertLocationMetadata]: ) -> Optional[ExpertLocationMetadata]:
data = get_exec().moe.init_expert_location data = get_exec().moe.init_expert_location
if data == "trivial": if data == "trivial":
return ExpertLocationMetadata.init_trivial( return ExpertLocationMetadata.init_trivial(model_config, moe_ep_rank)
server_args, model_config, moe_ep_rank
)
# TODO unify with the utils function # TODO unify with the utils function
if data.endswith(".pt"): if data.endswith(".pt"):
@@ -808,7 +795,6 @@ def compute_initial_expert_location_metadata(
"init_expert_location from init_by_mapping using ServerArgs.init_expert_location" "init_expert_location from init_by_mapping using ServerArgs.init_expert_location"
) )
return ExpertLocationMetadata.init_by_mapping( return ExpertLocationMetadata.init_by_mapping(
server_args,
model_config, model_config,
**data_dict, **data_dict,
moe_ep_rank=moe_ep_rank, moe_ep_rank=moe_ep_rank,
@@ -818,7 +804,7 @@ def compute_initial_expert_location_metadata(
"init_expert_location from init_by_eplb using ServerArgs.init_expert_location" "init_expert_location from init_by_eplb using ServerArgs.init_expert_location"
) )
return ExpertLocationMetadata.init_by_eplb( return ExpertLocationMetadata.init_by_eplb(
server_args, model_config, logical_count=data_dict["logical_count"] model_config, logical_count=data_dict["logical_count"]
) )
else: else:
raise NotImplementedError( raise NotImplementedError(
@@ -442,9 +442,7 @@ class FlashInferAttnBackend(AttentionBackend):
) )
else: else:
self.workspace_buffer = global_workspace_buffer self.workspace_buffer = global_workspace_buffer
max_bs = get_cuda_graph_max_batch_size( max_bs = get_cuda_graph_max_batch_size(model_runner.req_to_token_pool.size)
model_runner.server_args, model_runner.req_to_token_pool.size
)
if kv_indptr_buf is None: if kv_indptr_buf is None:
self.kv_indptr = [ self.kv_indptr = [
torch.zeros( torch.zeros(
@@ -2254,7 +2252,7 @@ class FlashInferMultiStepDraftBackend:
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
max_bs = get_cuda_graph_max_batch_size( max_bs = get_cuda_graph_max_batch_size(
model_runner.server_args, model_runner.req_to_token_pool.size * self.topk model_runner.req_to_token_pool.size * self.topk
) )
self.kv_indptr = torch.zeros( self.kv_indptr = torch.zeros(
( (
@@ -400,7 +400,6 @@ class GDNAttnBackend(MambaAttnBackendBase):
self.verify_intermediate_state_indices = ( self.verify_intermediate_state_indices = (
build_verify_intermediate_state_indices( build_verify_intermediate_state_indices(
self.req_to_token_pool.size, self.req_to_token_pool.size,
model_runner.server_args,
model_runner.device, model_runner.device,
) )
) )
@@ -408,7 +408,6 @@ class KDAAttnBackend(MambaAttnBackendBase):
self.verify_intermediate_state_indices = ( self.verify_intermediate_state_indices = (
build_verify_intermediate_state_indices( build_verify_intermediate_state_indices(
self.req_to_token_pool.size, self.req_to_token_pool.size,
model_runner.server_args,
model_runner.device, model_runner.device,
) )
) )
@@ -9,7 +9,7 @@ from sglang.srt.runtime_context import get_exec
from sglang.srt.utils.common import rank0_log from sglang.srt.utils.common import rank0_log
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs pass
class LinearAttnKernelBackend(Enum): class LinearAttnKernelBackend(Enum):
@@ -103,9 +103,7 @@ def resolve_linear_attn_backends(
return backends return backends
def build_verify_intermediate_state_indices( def build_verify_intermediate_state_indices(pool_size: int, device):
pool_size: int, server_args: ServerArgs, device
):
"""Per-request row index into the speculative intermediate scratch """Per-request row index into the speculative intermediate scratch
(`intermediate_ssm` / `intermediate_conv_window`) for the MTP / (`intermediate_ssm` / `intermediate_conv_window`) for the MTP /
target_verify path: request slot i owns scratch row i. target_verify path: request slot i owns scratch row i.
@@ -123,7 +121,7 @@ def build_verify_intermediate_state_indices(
from sglang.srt.utils.common import get_eager_max_batch_size from sglang.srt.utils.common import get_eager_max_batch_size
padded_bs = max(get_eager_max_batch_size(server_args, pool_size), pool_size) padded_bs = max(get_eager_max_batch_size(pool_size), pool_size)
indices = torch.arange(pool_size, dtype=torch.int32, device=device) indices = torch.arange(pool_size, dtype=torch.int32, device=device)
if padded_bs > pool_size: if padded_bs > pool_size:
indices = torch.cat( indices = torch.cat(
@@ -477,7 +477,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
cp = max(get_cp_padding_align_size(), 1) cp = max(get_cp_padding_align_size(), 1)
attn_tp_size = get_parallel().attn_tp_size attn_tp_size = get_parallel().attn_tp_size
mlp_sync = require_mlp_sync(model_runner.server_args) mlp_sync = require_mlp_sync()
def _align(bs: int) -> int: def _align(bs: int) -> int:
# Align to lcm(cp, attn_tp_size) so the CP multiple isn't undone by a # Align to lcm(cp, attn_tp_size) so the CP multiple isn't undone by a
+5 -8
View File
@@ -94,19 +94,17 @@ class LoRAManager:
self.attn_tp_size: int = get_parallel().attn_tp_size self.attn_tp_size: int = get_parallel().attn_tp_size
self.lora_added_tokens_size: Optional[int] = None self.lora_added_tokens_size: Optional[int] = None
self.enable_lora_overlap_loading: Optional[bool] = ( self.enable_lora_overlap_loading: Optional[bool] = (
server_args.enable_lora_overlap_loading get_lora().enable_lora_overlap_loading
) )
self.pending_lora_load_events = {} self.pending_lora_load_events = {}
self.eviction_policy = server_args.lora_eviction_policy self.eviction_policy = get_lora().lora_eviction_policy
self.enable_dp_attention: bool = get_parallel().config.enable_dp_attention self.enable_dp_attention: bool = get_parallel().config.enable_dp_attention
self._experts_shared_outer_override: Optional[bool] = ( self._experts_shared_outer_override: Optional[bool] = (
server_args.experts_shared_outer_loras get_lora().experts_shared_outer_loras
)
self.lora_use_virtual_experts: bool = server_args.lora_use_virtual_experts
self.lora_strict_loading: bool = getattr(
server_args, "lora_strict_loading", False
) )
self.lora_use_virtual_experts: bool = get_lora().lora_use_virtual_experts
self.lora_strict_loading: bool = get_lora().lora_strict_loading
self.speculative_algorithm: Optional[str] = get_spec().speculative_algorithm self.speculative_algorithm: Optional[str] = get_spec().speculative_algorithm
# LoRA backend for running sgemm kernels # LoRA backend for running sgemm kernels
@@ -1030,7 +1028,6 @@ class LoRAManager:
def init_lora_cuda_graph_moe_buffers( def init_lora_cuda_graph_moe_buffers(
*, *,
server_args: ServerArgs,
model: torch.nn.Module, model: torch.nn.Module,
lora_manager: LoRAManager, lora_manager: LoRAManager,
dtype: torch.dtype, dtype: torch.dtype,
+3 -10
View File
@@ -13,12 +13,9 @@ from sglang.srt.runtime_context import (
get_parallel, get_parallel,
get_serving, get_serving,
) )
from sglang.srt.server_args import ServerArgs
def start_disagg_service( def start_disagg_service():
server_args: ServerArgs,
):
# Start kv bootstrap server on prefill # Start kv bootstrap server on prefill
disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode) disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode)
transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend) transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend)
@@ -32,16 +29,12 @@ def start_disagg_service(
host=get_serving().host, host=get_serving().host,
port=get_disagg().disaggregation_bootstrap_port, port=get_disagg().disaggregation_bootstrap_port,
) )
maybe_create_ascend_config_store( maybe_create_ascend_config_store(transfer_backend=transfer_backend)
server_args=server_args, transfer_backend=transfer_backend
)
return bootstrap_server return bootstrap_server
def maybe_create_ascend_config_store( def maybe_create_ascend_config_store(transfer_backend: TransferBackend) -> None:
server_args: ServerArgs, transfer_backend: TransferBackend
) -> None:
"""Also called directly by the rust-server scheduler: there the KV """Also called directly by the rust-server scheduler: there the KV
bootstrap registry is served by the embedded rust server's api listener bootstrap registry is served by the embedded rust server's api listener
(one rust implementation covers every transfer backend — their (one rust implementation covers every transfer backend — their
@@ -478,7 +478,7 @@ class MultiTokenizerRouter:
) )
self._loop.call_soon_threadsafe(self._register_load_snapshot_reader) self._loop.call_soon_threadsafe(self._register_load_snapshot_reader)
self.disaggregation_bootstrap_server = start_disagg_service(self.server_args) self.disaggregation_bootstrap_server = start_disagg_service()
# Worker IPC names for pause/continue broadcasting # Worker IPC names for pause/continue broadcasting
self.all_worker_ipcs: set[str] = set() self.all_worker_ipcs: set[str] = set()
+1 -3
View File
@@ -18,14 +18,12 @@ if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.ngram_info import NgramVerifyInput from sglang.srt.speculative.ngram_info import NgramVerifyInput
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
def decide_needs_cpu_seq_lens( def decide_needs_cpu_seq_lens(
server_args: ServerArgs,
attn_backends: Sequence[AttentionBackend], attn_backends: Sequence[AttentionBackend],
) -> bool: ) -> bool:
"""Whether FutureMap must publish seq_lens_cpu / sum. """Whether FutureMap must publish seq_lens_cpu / sum.
@@ -53,7 +51,7 @@ def decide_needs_cpu_seq_lens(
) )
def decide_needs_confidence_relay(server_args: ServerArgs) -> bool: def decide_needs_confidence_relay() -> bool:
from sglang.srt.speculative.ragged_verify import ( from sglang.srt.speculative.ragged_verify import (
RaggedVerifyMode, RaggedVerifyMode,
read_ragged_verify_mode, read_ragged_verify_mode,
+5 -8
View File
@@ -416,7 +416,7 @@ class Scheduler(
# init_soft_watchdog starts a daemon thread that reads these on its first tick. # init_soft_watchdog starts a daemon thread that reads these on its first tick.
self.forward_ct: int = 0 self.forward_ct: int = 0
self.cur_batch_for_debug: Optional[ScheduleBatch] = None self.cur_batch_for_debug: Optional[ScheduleBatch] = None
self.init_soft_watchdog(server_args) self.init_soft_watchdog()
# Parse args # Parse args
self.server_args = server_args self.server_args = server_args
@@ -911,7 +911,7 @@ class Scheduler(
initialize_bf16_gemm_config(self.server_args) initialize_bf16_gemm_config(self.server_args)
# This must be called after initialize_moe_config # This must be called after initialize_moe_config
self.require_mlp_sync = require_mlp_sync(self.server_args) self.require_mlp_sync = require_mlp_sync()
def init_tp_model_worker(self): def init_tp_model_worker(self):
worker_kwargs = dict( worker_kwargs = dict(
@@ -1289,7 +1289,7 @@ class Scheduler(
self.new_token_ratio_tracker = NewTokenRatioTracker.from_config() self.new_token_ratio_tracker = NewTokenRatioTracker.from_config()
def init_soft_watchdog(self, server_args: ServerArgs): def init_soft_watchdog(self):
if (x := get_device().soft_watchdog_timeout) is not None: if (x := get_device().soft_watchdog_timeout) is not None:
self.soft_watchdog = create_scheduler_watchdog( self.soft_watchdog = create_scheduler_watchdog(
self, watchdog_timeout=x, soft=True self, watchdog_timeout=x, soft=True
@@ -1341,7 +1341,6 @@ class Scheduler(
and self._hosts_rust_server() and self._hosts_rust_server()
): ):
maybe_create_ascend_config_store( maybe_create_ascend_config_store(
server_args=self.server_args,
transfer_backend=self.transfer_backend, transfer_backend=self.transfer_backend,
) )
@@ -1493,8 +1492,8 @@ class Scheduler(
) )
else: else:
attn_backends = (self.tp_worker.model_runner.attn_backend,) attn_backends = (self.tp_worker.model_runner.attn_backend,)
needs_cpu_seq_lens = decide_needs_cpu_seq_lens(self.server_args, attn_backends) needs_cpu_seq_lens = decide_needs_cpu_seq_lens(attn_backends)
needs_confidence_relay = decide_needs_confidence_relay(self.server_args) needs_confidence_relay = decide_needs_confidence_relay()
self.future_map = self.spec_algorithm.create_future_map( self.future_map = self.spec_algorithm.create_future_map(
self.device, self.device,
self.req_to_token_pool, self.req_to_token_pool,
@@ -2089,7 +2088,6 @@ class Scheduler(
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
offload_tags=self.weight_updater.offload_tags, offload_tags=self.weight_updater.offload_tags,
ps=self.ps, ps=self.ps,
server_args=self.server_args,
model_config=self.model_config, model_config=self.model_config,
enable_overlap=self.enable_overlap, enable_overlap=self.enable_overlap,
spec_algorithm=self.spec_algorithm, spec_algorithm=self.spec_algorithm,
@@ -2213,7 +2211,6 @@ class Scheduler(
disaggregation_mode=self.disaggregation_mode, disaggregation_mode=self.disaggregation_mode,
enable_overlap=self.enable_overlap, enable_overlap=self.enable_overlap,
enable_overlap_mlx=self.enable_overlap_mlx, enable_overlap_mlx=self.enable_overlap_mlx,
server_args=self.server_args,
model_config=self.model_config, model_config=self.model_config,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
@@ -70,7 +70,6 @@ if TYPE_CHECKING:
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -81,7 +80,6 @@ class SchedulerBatchResultProcessor:
disaggregation_mode: DisaggregationMode disaggregation_mode: DisaggregationMode
enable_overlap: bool enable_overlap: bool
enable_overlap_mlx: bool enable_overlap_mlx: bool
server_args: ServerArgs
model_config: ModelConfig model_config: ModelConfig
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
tree_cache: BasePrefixCache tree_cache: BasePrefixCache
@@ -280,7 +278,6 @@ class SchedulerBatchResultProcessor:
hidden_state_offset = 0 hidden_state_offset = 0
prefill_hidden_capture_mode = self._get_prefill_hidden_capture_mode( prefill_hidden_capture_mode = self._get_prefill_hidden_capture_mode(
batch, batch,
self.server_args,
) )
# Check finish conditions # Check finish conditions
@@ -617,10 +614,7 @@ class SchedulerBatchResultProcessor:
) )
@staticmethod @staticmethod
def _get_prefill_hidden_capture_mode( def _get_prefill_hidden_capture_mode(batch: ScheduleBatch) -> CaptureHiddenMode:
batch: ScheduleBatch,
server_args: ServerArgs,
) -> CaptureHiddenMode:
return get_required_capture_hidden_mode( return get_required_capture_hidden_mode(
max( max(
batch.return_hidden_states_mode, batch.return_hidden_states_mode,
@@ -27,7 +27,6 @@ from sglang.srt.model_executor.cuda_graph_config import (
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.metrics_collector import DPCooperationInfo from sglang.srt.observability.metrics_collector import DPCooperationInfo
from sglang.srt.runtime_context import get_parallel, get_schedule from sglang.srt.runtime_context import get_parallel, get_schedule
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.common import require_mlp_tp_gather from sglang.srt.utils.common import require_mlp_tp_gather
@@ -401,7 +400,6 @@ class SchedulerDPAttnAdapter:
tree_cache: BasePrefixCache tree_cache: BasePrefixCache
offload_tags: set[str] offload_tags: set[str]
ps: ParallelState ps: ParallelState
server_args: ServerArgs
model_config: ModelConfig model_config: ModelConfig
enable_overlap: bool enable_overlap: bool
spec_algorithm: SpeculativeAlgorithm spec_algorithm: SpeculativeAlgorithm
@@ -417,7 +415,7 @@ class SchedulerDPAttnAdapter:
tp_group=self.tp_group, tp_group=self.tp_group,
get_idle_batch=self.get_idle_batch, get_idle_batch=self.get_idle_batch,
disable_cuda_graph=cuda_graph_fully_disabled(), disable_cuda_graph=cuda_graph_fully_disabled(),
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args), require_mlp_tp_gather=require_mlp_tp_gather(),
disable_overlap_schedule=get_schedule().disable_overlap_schedule, disable_overlap_schedule=get_schedule().disable_overlap_schedule,
offload_tags=self.offload_tags, offload_tags=self.offload_tags,
dwdp=get_parallel().config.dwdp_size > 1, dwdp=get_parallel().config.dwdp_size > 1,
@@ -81,7 +81,7 @@ from sglang.srt.runtime_context import (
get_parallel, get_parallel,
get_spec, get_spec,
) )
from sglang.srt.server_args import LoRARef, ServerArgs from sglang.srt.server_args import LoRARef
from sglang.srt.utils import ( from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
normalize_serialized_named_tensor_payloads, normalize_serialized_named_tensor_payloads,
@@ -158,7 +158,7 @@ class TokenizerControlMixin:
FanOutCommunicator, as opposed to data-plane inference requests multiplexed by rid. FanOutCommunicator, as opposed to data-plane inference requests multiplexed by rid.
""" """
def init_communicators(self: TokenizerManager, server_args: ServerArgs): def init_communicators(self: TokenizerManager):
dispatch_pairs = [] dispatch_pairs = []
for spec in _COMMUNICATOR_SPECS: for spec in _COMMUNICATOR_SPECS:
name, resp_type = spec[0], spec[1] name, resp_type = spec[0], spec[1]
@@ -663,9 +663,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode) self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
# Keep a reference so the bootstrap server is not garbage-collected. # Keep a reference so the bootstrap server is not garbage-collected.
self.bootstrap_server = ( self.bootstrap_server = (
start_disagg_service(self.server_args) start_disagg_service() if start_pd_bootstrap_service else None
if start_pd_bootstrap_service
else None
) )
# Single-source counter for auto-assigning fake bootstrap_room. # Single-source counter for auto-assigning fake bootstrap_room.
self.fake_bootstrap_room_counter = 0 self.fake_bootstrap_room_counter = 0
@@ -761,7 +759,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
(ElasticScaleUpdateReq, self.forward_elastic_scale_update), (ElasticScaleUpdateReq, self.forward_elastic_scale_update),
] ]
) )
self.init_communicators(self.server_args) self.init_communicators()
self.sampling_params_class = SamplingParams self.sampling_params_class = SamplingParams
self.signal_handler_class = SignalHandler self.signal_handler_class = SignalHandler
@@ -410,7 +410,6 @@ def build_hybrid_swa_stack(
def _deepseek_v4_num_host_pages( def _deepseek_v4_num_host_pages(
*, *,
params: CacheInitParams, params: CacheInitParams,
server_args: ServerArgs,
kvcache: Any, kvcache: Any,
page_size: int, page_size: int,
swa_page_size: int, swa_page_size: int,
@@ -511,7 +510,6 @@ def build_deepseek_v4_hicache_stack(
} }
num_host_pages, swa_num_host_pages = _deepseek_v4_num_host_pages( num_host_pages, swa_num_host_pages = _deepseek_v4_num_host_pages(
params=params, params=params,
server_args=server_args,
kvcache=kvcache, kvcache=kvcache,
page_size=page_size, page_size=page_size,
swa_page_size=kvcache.swa_page_size, swa_page_size=kvcache.swa_page_size,
@@ -1152,7 +1152,6 @@ class KVCacheConfigurator:
kv_cache_dim=calculate_mla_kv_cache_dim( kv_cache_dim=calculate_mla_kv_cache_dim(
model_config=self.model_config, model_config=self.model_config,
kv_cache_dtype=self.kv_cache_dtype, kv_cache_dtype=self.kv_cache_dtype,
server_args=self.server_args,
), ),
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=get_exec().features.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
@@ -1356,7 +1355,6 @@ class KVCacheConfigurator:
kv_cache_dim=calculate_mla_kv_cache_dim( kv_cache_dim=calculate_mla_kv_cache_dim(
model_config=self.model_config, model_config=self.model_config,
kv_cache_dtype=self.kv_cache_dtype, kv_cache_dtype=self.kv_cache_dtype,
server_args=self.server_args,
), ),
enable_memory_saver=get_exec().features.enable_memory_saver, enable_memory_saver=get_exec().features.enable_memory_saver,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
@@ -1396,7 +1394,6 @@ class KVCacheConfigurator:
kv_cache_dim=calculate_mla_kv_cache_dim( kv_cache_dim=calculate_mla_kv_cache_dim(
model_config=self.model_config, model_config=self.model_config,
kv_cache_dtype=self.kv_cache_dtype, kv_cache_dtype=self.kv_cache_dtype,
server_args=self.server_args,
), ),
) )
@@ -2199,10 +2196,7 @@ class KVCacheConfigurator:
def calculate_mla_kv_cache_dim( def calculate_mla_kv_cache_dim(
*, *, model_config: ModelConfig, kv_cache_dtype: torch.dtype
model_config: ModelConfig,
kv_cache_dtype: torch.dtype,
server_args: ServerArgs,
) -> int: ) -> int:
is_dsa_model = is_deepseek_dsa(model_config.hf_config) is_dsa_model = is_deepseek_dsa(model_config.hf_config)
kv_cache_dtype = kv_cache_dtype kv_cache_dtype = kv_cache_dtype
@@ -598,10 +598,10 @@ class CPUGraphRunner:
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync(model_runner.server_args) self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_two_batch_overlap = ( self.enable_two_batch_overlap = (
model_runner.server_args.enable_two_batch_overlap model_runner.server_args.enable_two_batch_overlap
) )
@@ -986,14 +986,14 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
pin_memory=is_pin_memory_available(batch.device), pin_memory=is_pin_memory_available(batch.device),
).to(batch.device, non_blocking=True) ).to(batch.device, non_blocking=True)
def adjust_num_token_non_padded_for_attn_tp(self, server_args) -> None: def adjust_num_token_non_padded_for_attn_tp(self) -> None:
"""Make num_token_non_padded local to this attention-TP rank.""" """Make num_token_non_padded local to this attention-TP rank."""
from sglang.srt.utils.common import require_mlp_tp_gather from sglang.srt.utils.common import require_mlp_tp_gather
dp_rank = get_parallel().attn_dp_rank dp_rank = get_parallel().attn_dp_rank
assert self.global_num_tokens_cpu is not None assert self.global_num_tokens_cpu is not None
if require_mlp_tp_gather(server_args): if require_mlp_tp_gather():
num_tokens_per_dp = self.global_num_tokens_cpu[dp_rank] num_tokens_per_dp = self.global_num_tokens_cpu[dp_rank]
else: else:
num_tokens_per_dp = self.global_num_tokens_cpu[0] num_tokens_per_dp = self.global_num_tokens_cpu[0]
@@ -223,7 +223,7 @@ from sglang.srt.utils.device_timer import device_timer_ctx
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
from sglang.srt.utils.nvtx_utils import profile_range from sglang.srt.utils.nvtx_utils import profile_range
from sglang.srt.utils.offloader import ( from sglang.srt.utils.offloader import (
create_offloader_from_server_args, create_offloader,
get_offloader, get_offloader,
set_offloader, set_offloader,
) )
@@ -270,7 +270,6 @@ class ModelRunnerOutput:
def resolve_draft_attention_backend( def resolve_draft_attention_backend(
*, *,
draft_attention_backend: Optional[str], draft_attention_backend: Optional[str],
server_args: ServerArgs,
is_draft_worker: bool, is_draft_worker: bool,
) -> Optional[str]: ) -> Optional[str]:
"""The attention backend a runner uses because it is a draft runner. """The attention backend a runner uses because it is a draft runner.
@@ -342,7 +341,6 @@ class ModelRunner:
self.device = get_device().device self.device = get_device().device
self.draft_attention_backend = resolve_draft_attention_backend( self.draft_attention_backend = resolve_draft_attention_backend(
draft_attention_backend=draft_attention_backend, draft_attention_backend=draft_attention_backend,
server_args=server_args,
is_draft_worker=is_draft_worker, is_draft_worker=is_draft_worker,
) )
# This runner's own load format, resolved before anything keys off it: # This runner's own load format, resolved before anything keys off it:
@@ -430,9 +428,7 @@ class ModelRunner:
self.shared_read_done_event: Optional[torch.cuda.Event] = None self.shared_read_done_event: Optional[torch.cuda.Event] = None
# CPU offload # CPU offload
set_offloader( set_offloader(create_offloader(dp_rank=self.ps.dp_rank))
create_offloader_from_server_args(server_args, dp_rank=self.ps.dp_rank)
)
self._weight_checker = WeightChecker(get_model=lambda: self.model, ps=self.ps) self._weight_checker = WeightChecker(get_model=lambda: self.model, ps=self.ps)
@@ -591,7 +587,6 @@ class ModelRunner:
model=self.model, model=self.model,
model_config=self.model_config, model_config=self.model_config,
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
server_args=self.server_args,
max_running_requests=self.max_running_requests, max_running_requests=self.max_running_requests,
device=self.device, device=self.device,
) )
@@ -702,7 +697,6 @@ class ModelRunner:
) )
set_global_expert_location_metadata( set_global_expert_location_metadata(
compute_initial_expert_location_metadata( compute_initial_expert_location_metadata(
server_args=self.server_args,
model_config=self.model_config, model_config=self.model_config,
moe_ep_rank=expert_rank, moe_ep_rank=expert_rank,
) )
@@ -727,7 +721,6 @@ class ModelRunner:
def maybe_init_eplb_manager(self): def maybe_init_eplb_manager(self):
self.eplb_manager = ( self.eplb_manager = (
EPLBManager( EPLBManager(
server_args=self.server_args,
model_config=self.model_config, model_config=self.model_config,
ps=self.ps, ps=self.ps,
get_model=lambda: self.model, get_model=lambda: self.model,
@@ -1076,9 +1069,7 @@ class ModelRunner:
self.pre_model_load_memory = result.pre_model_load_memory self.pre_model_load_memory = result.pre_model_load_memory
def init_shared_mooncake_transfer_engine(self): def init_shared_mooncake_transfer_engine(self):
maybe_init_shared_mooncake_transfer_engine( maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id)
server_args=self.server_args, gpu_id=self.gpu_id
)
def load_model(self): def load_model(self):
tic_total = time.perf_counter() tic_total = time.perf_counter()
@@ -1091,9 +1082,7 @@ class ModelRunner:
if self.device != "cpu": if self.device != "cpu":
torch.set_num_threads(1) torch.set_num_threads(1)
if self.device == "cuda": if self.device == "cuda":
maybe_downgrade_dtype_for_legacy_gpu( maybe_downgrade_dtype_for_legacy_gpu(model_config=self.model_config)
server_args=self.server_args, model_config=self.model_config
)
set_cuda_arch() set_cuda_arch()
@@ -1113,7 +1102,6 @@ class ModelRunner:
# and derive the per-rank daemon socket. Idempotent across reloads. # and derive the per-rank daemon socket. Idempotent across reloads.
maybe_enable_ipc_weight_cache( maybe_enable_ipc_weight_cache(
load_config=self.load_config, load_config=self.load_config,
server_args=self.server_args,
tp_size=self.ps.tp_size, tp_size=self.ps.tp_size,
pp_rank=self.ps.pp_rank, pp_rank=self.ps.pp_rank,
tp_rank=self.ps.tp_rank, tp_rank=self.ps.tp_rank,
@@ -1124,7 +1112,6 @@ class ModelRunner:
) )
maybe_trigger_remote_instance_nccl_send_group( maybe_trigger_remote_instance_nccl_send_group(
server_args=self.server_args,
tp_rank=self.ps.tp_rank, tp_rank=self.ps.tp_rank,
load_format=draft_load_format, load_format=draft_load_format,
) )
@@ -1195,11 +1182,12 @@ class ModelRunner:
f"mem usage={self.weight_load_mem_usage:.2f} GB." f"mem usage={self.weight_load_mem_usage:.2f} GB."
) )
report_online_quantization(model=self.model, server_args=self.server_args) report_online_quantization(
model=self.model,
)
maybe_register_debug_tensor_dump_hook( maybe_register_debug_tensor_dump_hook(
model=self.model, model=self.model,
server_args=self.server_args,
spec_algorithm=self.spec_algorithm, spec_algorithm=self.spec_algorithm,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
tp_size=self.ps.tp_size, tp_size=self.ps.tp_size,
@@ -1281,7 +1269,6 @@ class ModelRunner:
) )
if not cuda_graph_fully_disabled(): if not cuda_graph_fully_disabled():
init_lora_cuda_graph_moe_buffers( init_lora_cuda_graph_moe_buffers(
server_args=self.server_args,
model=self.model, model=self.model,
lora_manager=self.lora_manager, lora_manager=self.lora_manager,
dtype=self.dtype, dtype=self.dtype,
@@ -1471,13 +1458,11 @@ class ModelRunner:
if ( if (
forward_batch.num_token_non_padded is not None forward_batch.num_token_non_padded is not None
and forward_batch.global_num_tokens_gpu is not None and forward_batch.global_num_tokens_gpu is not None
and require_gathered_buffer(self.server_args) and require_gathered_buffer()
and not is_dsa_enable_prefill_cp() and not is_dsa_enable_prefill_cp()
and not is_mla_prefill_cp_enabled() and not is_mla_prefill_cp_enabled()
): ):
forward_batch.adjust_num_token_non_padded_for_attn_tp( forward_batch.adjust_num_token_non_padded_for_attn_tp()
server_args=self.server_args,
)
# Hisparse coordinator — backends now read it from self.model_runner. # Hisparse coordinator — backends now read it from self.model_runner.
if self.hisparse_coordinator is not None: if self.hisparse_coordinator is not None:
@@ -1926,7 +1911,6 @@ class ModelRunner:
start=old_num_physical - num_local * initial_ep_size, start=old_num_physical - num_local * initial_ep_size,
) )
new_metadata = ExpertLocationMetadata.init_by_mapping( new_metadata = ExpertLocationMetadata.init_by_mapping(
self.server_args,
self.model_config, self.model_config,
physical_to_logical_map=expanded_p2l, physical_to_logical_map=expanded_p2l,
moe_ep_rank=self._elastic_global_rank(), moe_ep_rank=self._elastic_global_rank(),
@@ -67,9 +67,7 @@ class LoadedModel(msgspec.Struct, frozen=True, kw_only=True):
startup_weight_load: Optional[Any] = None startup_weight_load: Optional[Any] = None
def maybe_downgrade_dtype_for_legacy_gpu( def maybe_downgrade_dtype_for_legacy_gpu(*, model_config: ModelConfig) -> None:
*, server_args: ServerArgs, model_config: ModelConfig
) -> None:
if torch.cuda.get_device_capability()[0] < 8: if torch.cuda.get_device_capability()[0] < 8:
logger.info( logger.info(
"Compute capability below sm80. Use float16 due to lack of bfloat16 support." "Compute capability below sm80. Use float16 due to lack of bfloat16 support."
@@ -85,7 +83,7 @@ def maybe_downgrade_dtype_for_legacy_gpu(
def maybe_trigger_remote_instance_nccl_send_group( def maybe_trigger_remote_instance_nccl_send_group(
*, server_args: ServerArgs, tp_rank: int, load_format: Optional[str] = None *, tp_rank: int, load_format: Optional[str] = None
) -> None: ) -> None:
"""``load_format`` is this runner's effective format: a draft loading under """``load_format`` is this runner's effective format: a draft loading under
``--speculative-draft-draft-load-format`` needs its own send group, and the ``--speculative-draft-draft-load-format`` needs its own send group, and the
@@ -151,7 +149,7 @@ def resolve_sliding_window_size(model, model_config: ModelConfig) -> Optional[in
return sliding_window_size return sliding_window_size
def report_online_quantization(*, model, server_args: ServerArgs) -> None: def report_online_quantization(*, model) -> None:
# TODO: Make sure all models have `quant_config` attribute, and all online quantization methods register which layers they actually quantize. # TODO: Make sure all models have `quant_config` attribute, and all online quantization methods register which layers they actually quantize.
quantized_layers = getattr( quantized_layers = getattr(
getattr(model, "quant_config", None), "quantized_layers", None getattr(model, "quant_config", None), "quantized_layers", None
@@ -170,7 +168,6 @@ def report_online_quantization(*, model, server_args: ServerArgs) -> None:
def maybe_register_debug_tensor_dump_hook( def maybe_register_debug_tensor_dump_hook(
*, *,
model, model,
server_args: ServerArgs,
spec_algorithm: SpeculativeAlgorithm, spec_algorithm: SpeculativeAlgorithm,
is_draft_worker: bool, is_draft_worker: bool,
tp_size: int, tp_size: int,
@@ -237,7 +234,6 @@ def build_load_config(
def maybe_enable_ipc_weight_cache( def maybe_enable_ipc_weight_cache(
*, *,
load_config: LoadConfig, load_config: LoadConfig,
server_args: ServerArgs,
tp_size: int, tp_size: int,
pp_rank: int, pp_rank: int,
tp_rank: int, tp_rank: int,
@@ -311,12 +307,11 @@ def load_model_with_memory_saver(
StartupWeightLoadManager, StartupWeightLoadManager,
) )
startup_weight_load = StartupWeightLoadManager.create_from_server_args( startup_weight_load = StartupWeightLoadManager.create_from_published_config(
loader=loader, loader=loader,
model_config=model_config, model_config=model_config,
load_config=load_config, load_config=load_config,
device_config=device_config, device_config=device_config,
server_args=server_args,
is_draft_worker=is_draft_worker, is_draft_worker=is_draft_worker,
) )
model = startup_weight_load.prepare() model = startup_weight_load.prepare()
@@ -12,7 +12,6 @@ from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.managers.schedule_batch import ForwardMode from sglang.srt.managers.schedule_batch import ForwardMode
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.runtime_context import get_schedule from sglang.srt.runtime_context import get_schedule
from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
@@ -33,7 +32,6 @@ class NgramEmbeddingManager:
model: torch.nn.Module, model: torch.nn.Module,
model_config: ModelConfig, model_config: ModelConfig,
req_to_token_pool: ReqToTokenPool, req_to_token_pool: ReqToTokenPool,
server_args: ServerArgs,
max_running_requests: int, max_running_requests: int,
device: str, device: str,
): ):
@@ -203,7 +203,6 @@ def _resolve_dflash_aux_hidden_state(
config.dflash_draft_num_layers = int(draft_num_layers) config.dflash_draft_num_layers = int(draft_num_layers)
config.dflash_target_layer_ids = target_layer_ids config.dflash_target_layer_ids = target_layer_ids
config.dflash_draft_cell_size_per_token = _resolve_dflash_draft_cell_size( config.dflash_draft_cell_size_per_token = _resolve_dflash_draft_cell_size(
server_args=server_args,
draft_model_config=draft_model_config, draft_model_config=draft_model_config,
draft_num_layers=int(draft_num_layers), draft_num_layers=int(draft_num_layers),
) )
@@ -211,7 +210,6 @@ def _resolve_dflash_aux_hidden_state(
def _resolve_dflash_draft_cell_size( def _resolve_dflash_draft_cell_size(
*, *,
server_args: ServerArgs,
draft_model_config: ModelConfig, draft_model_config: ModelConfig,
draft_num_layers: int, draft_num_layers: int,
) -> int | None: ) -> int | None:
@@ -31,7 +31,6 @@ from sglang.srt.runtime_context import (
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -98,12 +97,16 @@ class StartupWeightLoadOptions:
prefetch_num_threads: int prefetch_num_threads: int
@classmethod @classmethod
def from_server_args( def from_published_config(
cls, cls,
*, *,
server_args: ServerArgs,
is_draft_worker: bool, is_draft_worker: bool,
) -> StartupWeightLoadOptions: ) -> StartupWeightLoadOptions:
"""Everything this needs is a published leaf; nothing comes off a record.
`is_draft_worker` is the exception and travels as an argument: it is
this runner's role, not the process's configuration.
"""
cuda_graph_config = get_exec().graph.cuda_graph_config cuda_graph_config = get_exec().graph.cuda_graph_config
cuda_graph_enabled = any( cuda_graph_enabled = any(
getattr(cuda_graph_config, phase).backend != Backend.DISABLED getattr(cuda_graph_config, phase).backend != Backend.DISABLED
@@ -267,29 +270,27 @@ class StartupWeightLoadManager:
self._prefetch_failure_reported = False self._prefetch_failure_reported = False
@classmethod @classmethod
def create_from_server_args( def create_from_published_config(
cls, cls,
*, *,
loader, loader,
model_config: ModelConfig, model_config: ModelConfig,
load_config: LoadConfig, load_config: LoadConfig,
device_config: DeviceConfig, device_config: DeviceConfig,
server_args: ServerArgs,
is_draft_worker: bool, is_draft_worker: bool,
) -> StartupWeightLoadManager: ) -> StartupWeightLoadManager:
"""Build a manager straight from ``ServerArgs``. """Build a manager from the published configuration.
Callers on the model-loading path only decide *whether* to overlap; the Callers on the model-loading path only decide *whether* to overlap; the
knowledge of which server arguments matter, and every support rule, knowledge of which config leaves matter, and every support rule, stays
stays in this module. in this module.
""" """
return cls.create( return cls.create(
loader=loader, loader=loader,
model_config=model_config, model_config=model_config,
load_config=load_config, load_config=load_config,
device_config=device_config, device_config=device_config,
options=StartupWeightLoadOptions.from_server_args( options=StartupWeightLoadOptions.from_published_config(
server_args=server_args,
is_draft_worker=is_draft_worker, is_draft_worker=is_draft_worker,
), ),
) )
@@ -259,7 +259,6 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
calculate_mla_kv_cache_dim( calculate_mla_kv_cache_dim(
model_config=model_config, model_config=model_config,
kv_cache_dtype=kv_cache_dtype, kv_cache_dtype=kv_cache_dtype,
server_args=kvc.server_args,
) )
* effective_num_layers * effective_num_layers
* kv_size * kv_size
@@ -465,7 +464,6 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
calculate_mla_kv_cache_dim( calculate_mla_kv_cache_dim(
model_config=model_config, model_config=model_config,
kv_cache_dtype=kv_cache_dtype, kv_cache_dtype=kv_cache_dtype,
server_args=kvc.server_args,
) )
* kv_size * kv_size
) )
@@ -70,11 +70,10 @@ def get_batch_sizes_to_capture(
constraints and clamps to req_to_token_pool.size. constraints and clamps to req_to_token_pool.size.
""" """
server_args = model_runner.server_args
capture_bs = list(get_exec().graph.cuda_graph_config.decode.bs) capture_bs = list(get_exec().graph.cuda_graph_config.decode.bs)
num_max_requests = model_runner.req_to_token_pool.size num_max_requests = model_runner.req_to_token_pool.size
mul_base = get_cuda_graph_batch_size_alignment(server_args) mul_base = get_cuda_graph_batch_size_alignment()
# TBO splits each request's rows across two micro-batches, so the # TBO splits each request's rows across two micro-batches, so the
# alignment constraint applies per request rather than per token row. # alignment constraint applies per request rather than per token row.
alignment_width = captured_req_width alignment_width = captured_req_width
@@ -82,7 +81,7 @@ def get_batch_sizes_to_capture(
alignment_width = 1 alignment_width = 1
# pad `num_max_requests` to avoid being filtered out # pad `num_max_requests` to avoid being filtered out
num_max_requests = get_cuda_graph_max_batch_size(server_args, num_max_requests) num_max_requests = get_cuda_graph_max_batch_size(num_max_requests)
if max(capture_bs) > num_max_requests: if max(capture_bs) > num_max_requests:
# In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests # In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests
# is very small. We add more values here to make sure we capture the maximum bs. # is very small. We add more values here to make sure we capture the maximum bs.
@@ -351,7 +351,7 @@ class BaseRunner(ABC):
dp_size=get_parallel().config.dp_size, dp_size=get_parallel().config.dp_size,
pp_size=get_parallel().config.pp_size, pp_size=get_parallel().config.pp_size,
is_encoder_decoder=mr.model_config.is_encoder_decoder, is_encoder_decoder=mr.model_config.is_encoder_decoder,
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args), require_mlp_tp_gather=require_mlp_tp_gather(),
seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(), seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(),
encoder_len_fill_value=( encoder_len_fill_value=(
getattr(mr.model_config.hf_config, "max_source_positions", 0) getattr(mr.model_config.hf_config, "max_source_positions", 0)
@@ -535,9 +535,9 @@ class BaseRunner(ABC):
) )
# TP-gather requirements for global token metadata. # TP-gather requirements for global token metadata.
require_mlp_tp_gather_ = require_mlp_tp_gather(mr.server_args) require_mlp_tp_gather_ = require_mlp_tp_gather()
require_attn_tp_gather_ = require_attn_tp_gather(mr.server_args) require_attn_tp_gather_ = require_attn_tp_gather()
if require_gathered_buffer(mr.server_args): if require_gathered_buffer():
assert require_mlp_tp_gather_ or require_attn_tp_gather_ assert require_mlp_tp_gather_ or require_attn_tp_gather_
if require_mlp_tp_gather_: if require_mlp_tp_gather_:
@@ -227,10 +227,10 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
self.require_mlp_tp_gather = require_mlp_tp_gather( self.require_mlp_tp_gather = (
model_runner.server_args require_mlp_tp_gather() and not self._forward_is_dp_local(model_runner)
) and not self._forward_is_dp_local(model_runner) )
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.require_attn_tp_gather = require_attn_tp_gather()
# Composite predicates derive from the instance values so the dp-local # Composite predicates derive from the instance values so the dp-local
# draft exemption above stays consistent (require_gathered_buffer == # draft exemption above stays consistent (require_gathered_buffer ==
# mlp_tp_gather or attn_tp_gather; require_mlp_sync adds dp attention). # mlp_tp_gather or attn_tp_gather; require_mlp_sync adds dp attention).
@@ -597,7 +597,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
draft_is_deepseek_v4, draft_is_deepseek_v4,
) )
return not draft_is_deepseek_v4(server_args=model_runner.server_args) return not draft_is_deepseek_v4()
def _ragged_capture_slots(self, num_tokens: int) -> int: def _ragged_capture_slots(self, num_tokens: int) -> int:
if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get(): if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get():
@@ -113,10 +113,10 @@ class EagerRunner(BaseRunner):
# (expand_for_topk_draft) before the eager fallback. # (expand_for_topk_draft) before the eager fallback.
max_bs *= get_spec().speculative_eagle_topk max_bs *= get_spec().speculative_eagle_topk
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies. # Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
max_bs = get_eager_max_batch_size(sa, max_bs) max_bs = get_eager_max_batch_size(max_bs)
prefill_ceiling = max(mr.max_total_num_tokens, max_prefill_buffer_tokens()) prefill_ceiling = max(mr.max_total_num_tokens, max_prefill_buffer_tokens())
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req) max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
if require_mlp_sync(sa): if require_mlp_sync():
from sglang.srt.layers.cp.padding import get_cp_padding_align_size from sglang.srt.layers.cp.padding import get_cp_padding_align_size
max_num_token = ceil_align(max_num_token, self.attn_tp_size) max_num_token = ceil_align(max_num_token, self.attn_tp_size)
@@ -332,7 +332,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
embed_dtype=self.model_runner.dtype, embed_dtype=self.model_runner.dtype,
enable_mamba_track=self.mamba_track_enabled, enable_mamba_track=self.mamba_track_enabled,
enable_num_token_non_padded=enable_num_token_non_padded(), enable_num_token_non_padded=enable_num_token_non_padded(),
require_gathered_buffer=require_gathered_buffer(model_runner.server_args), require_gathered_buffer=require_gathered_buffer(),
enable_prefill_cp=( enable_prefill_cp=(
is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled() is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled()
), ),
@@ -349,8 +349,8 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None) self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None)
self.dp_size = get_parallel().config.dp_size self.dp_size = get_parallel().config.dp_size
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.require_attn_tp_gather = require_attn_tp_gather()
# --- backend --------------------------------------------------- # --- backend ---------------------------------------------------
# TcPiecewise resolves by running a compile pass that calls back into # TcPiecewise resolves by running a compile pass that calls back into
@@ -602,7 +602,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
buf = self.buffer_registry.get_slot("num_token_non_padded").buffer buf = self.buffer_registry.get_slot("num_token_non_padded").buffer
buf.fill_(num_tokens) buf.fill_(num_tokens)
if require_gathered_buffer(self.model_runner.server_args): if require_gathered_buffer():
local = compute_local_num_token_non_padded( local = compute_local_num_token_non_padded(
global_num_token_non_padded=buf, global_num_token_non_padded=buf,
num_tokens_per_dp=num_tokens, num_tokens_per_dp=num_tokens,
+2 -3
View File
@@ -121,7 +121,6 @@ from sglang.srt.multimodal.mm_utils import materialize_multimodal_features
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec, get_exec,
get_parallel, get_parallel,
get_server_args,
) )
from sglang.srt.utils import is_blackwell_supported, is_hip, is_npu, make_layers from sglang.srt.utils import is_blackwell_supported, is_hip, is_npu, make_layers
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
@@ -2146,7 +2145,7 @@ class KimiK3DecoderLayer(nn.Module):
self._dp_attention = is_dp_attention_enabled() self._dp_attention = is_dp_attention_enabled()
# mlp-sync (DP attention OR MoE a2a/EP) pads extend batches to # mlp-sync (DP attention OR MoE a2a/EP) pads extend batches to
# attn_tp multiples; attention must then run on the real rows only. # attn_tp multiples; attention must then run on the real rows only.
self._trim_padded_attn = require_mlp_sync(get_server_args()) self._trim_padded_attn = require_mlp_sync()
# A layer runs MoE (vs a plain dense MLP) iff it is past the dense # A layer runs MoE (vs a plain dense MLP) iff it is past the dense
# prefix and on the MoE cadence — same predicate the mlp construction # prefix and on the MoE cadence — same predicate the mlp construction
# below uses. # below uses.
@@ -2614,7 +2613,7 @@ class KimiK3LinearModel(nn.Module):
self.pp_group = get_pp_group() self.pp_group = get_pp_group()
self.dspark_layers_to_capture: Optional[list[int]] = None self.dspark_layers_to_capture: Optional[list[int]] = None
self._dp_attention = is_dp_attention_enabled() self._dp_attention = is_dp_attention_enabled()
self._trim_padded_attn = require_mlp_sync(get_server_args()) self._trim_padded_attn = require_mlp_sync()
if self.pp_group.is_first_rank: if self.pp_group.is_first_rank:
embedding_quant_config = ( embedding_quant_config = (
@@ -40,7 +40,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
MmItemMemoryPool, MmItemMemoryPool,
get_mm_feature_pool_size_per_worker, get_mm_feature_pool_size_per_worker,
) )
from sglang.srt.runtime_context import get_mm from sglang.srt.runtime_context import get_mm, get_serving
from sglang.srt.utils import ( from sglang.srt.utils import (
CLIENT_MEDIA_EXCEPTIONS, CLIENT_MEDIA_EXCEPTIONS,
configure_media_url_security, configure_media_url_security,
@@ -215,7 +215,7 @@ class BaseMultimodalProcessor(ABC):
self.transport_mode = transport_mode self.transport_mode = transport_mode
configure_media_url_security( configure_media_url_security(
get_mm().allowed_media_domains, get_mm().allowed_media_domains,
server_args.media_url_max_file_size_mb, get_mm().media_url_max_file_size_mb,
) )
configured_mm_feature_transport = get_mm().mm_feature_transport configured_mm_feature_transport = get_mm().mm_feature_transport
self.mm_feature_transport = ( self.mm_feature_transport = (
@@ -227,11 +227,11 @@ class BaseMultimodalProcessor(ABC):
self.use_ipc_pool_handle_cache = ( self.use_ipc_pool_handle_cache = (
self.use_cuda_ipc and envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get() self.use_cuda_ipc and envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get()
) )
self.image_processor_backend = server_args.image_processor_backend self.image_processor_backend = get_mm().image_processor_backend
if server_args.disable_fast_image_processor: if get_mm().disable_fast_image_processor:
self.image_processor_backend = "pil" self.image_processor_backend = "pil"
self.disable_fast_image_processor = self.image_processor_backend == "pil" self.disable_fast_image_processor = self.image_processor_backend == "pil"
self.skip_tokenizer_init = server_args.skip_tokenizer_init self.skip_tokenizer_init = get_serving().skip_tokenizer_init
mm_process_config = get_mm().mm_process_config mm_process_config = get_mm().mm_process_config
self.image_config = mm_process_config.get("image", {}) self.image_config = mm_process_config.get("image", {})
@@ -241,19 +241,19 @@ class BaseMultimodalProcessor(ABC):
# Each tokenizer worker is a separate process with its own CPU cache. # Each tokenizer worker is a separate process with its own CPU cache.
# Split the requested service-wide budget so increasing worker count # Split the requested service-wide budget so increasing worker count
# does not silently multiply host-memory usage. # does not silently multiply host-memory usage.
requested_cache_mb = self.server_args.mm_preprocess_cache_size_mb requested_cache_mb = get_mm().mm_preprocess_cache_size_mb
total_cache_mb = ( total_cache_mb = (
self.auto_mm_preprocess_cache_size_mb self.auto_mm_preprocess_cache_size_mb
if requested_cache_mb is None if requested_cache_mb is None
else requested_cache_mb else requested_cache_mb
) )
tokenizer_worker_num = max(int(self.server_args.tokenizer_worker_num), 1) tokenizer_worker_num = max(int(get_serving().tokenizer_worker_num), 1)
worker_cache_bytes = total_cache_mb * 1024 * 1024 // tokenizer_worker_num worker_cache_bytes = total_cache_mb * 1024 * 1024 // tokenizer_worker_num
self.mm_preprocess_cache = MultimodalPreprocessCache( self.mm_preprocess_cache = MultimodalPreprocessCache(
max_size_bytes=worker_cache_bytes, max_size_bytes=worker_cache_bytes,
max_entries=8192, max_entries=8192,
) )
self.trust_mm_content_hashes = bool(self.server_args.trust_mm_content_hashes) self.trust_mm_content_hashes = bool(get_mm().trust_mm_content_hashes)
# The fingerprint is needed only to build artifact keys. Avoid inspecting # The fingerprint is needed only to build artifact keys. Avoid inspecting
# processor state when this processor will never retain artifacts. # processor state when this processor will never retain artifacts.
self.processor_fingerprint = ( self.processor_fingerprint = (
@@ -288,7 +288,7 @@ class BaseMultimodalProcessor(ABC):
# FIXME: not accurate, model and image specific # FIXME: not accurate, model and image specific
self.NUM_TOKEN_PER_FRAME = 330 self.NUM_TOKEN_PER_FRAME = 330
requested_mm_io_worker_num = self.server_args.mm_io_worker_num requested_mm_io_worker_num = get_mm().mm_io_worker_num
env_mm_io_worker_num = os.environ.get("SGLANG_IO_WORKERS") env_mm_io_worker_num = os.environ.get("SGLANG_IO_WORKERS")
if requested_mm_io_worker_num: if requested_mm_io_worker_num:
self.mm_io_worker_num = requested_mm_io_worker_num self.mm_io_worker_num = requested_mm_io_worker_num
@@ -310,7 +310,7 @@ class BaseMultimodalProcessor(ABC):
io_worker_source, io_worker_source,
) )
skip_mm_pool = kwargs.get("skip_mm_pool", False) skip_mm_pool = kwargs.get("skip_mm_pool", False)
requested_mm_processor_worker_num = self.server_args.mm_processor_worker_num requested_mm_processor_worker_num = get_mm().mm_processor_worker_num
self.mm_processor_worker_num = ( self.mm_processor_worker_num = (
1 1
if skip_mm_pool if skip_mm_pool
@@ -402,7 +402,7 @@ class BaseMultimodalProcessor(ABC):
# SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all # SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all
# tokenizer workers. Each worker gets an equal share so that adding # tokenizer workers. Each worker gets an equal share so that adding
# workers doesn't multiply the GPU-side footprint. # workers doesn't multiply the GPU-side footprint.
worker_num = self.server_args.tokenizer_worker_num worker_num = get_serving().tokenizer_worker_num
per_worker_pool_size = get_mm_feature_pool_size_per_worker( per_worker_pool_size = get_mm_feature_pool_size_per_worker(
MM_FEATURE_CACHE_SIZE, worker_num MM_FEATURE_CACHE_SIZE, worker_num
) )
@@ -30,7 +30,7 @@ def get_dspark_sample_from_anchor(draft_hf_config: Any) -> bool:
return bool(_cfg_get(draft_hf_config, "sample_from_anchor", True)) return bool(_cfg_get(draft_hf_config, "sample_from_anchor", True))
def draft_is_deepseek_v4(*, server_args: ServerArgs) -> bool: def draft_is_deepseek_v4() -> bool:
from sglang.srt.configs.model_config import is_deepseek_v4 from sglang.srt.configs.model_config import is_deepseek_v4
from sglang.srt.utils.hf_transformers_utils import get_config from sglang.srt.utils.hf_transformers_utils import get_config
@@ -172,7 +172,7 @@ class DSparkVerifyPlanner:
and is_dp_attention_enabled() and is_dp_attention_enabled()
and get_parallel().attn_tp_size == 1 and get_parallel().attn_tp_size == 1
and get_parallel().attn_cp_size == 1 and get_parallel().attn_cp_size == 1
and require_mlp_tp_gather(self.server_args) and require_mlp_tp_gather()
and not get_schedule().disable_overlap_schedule and not get_schedule().disable_overlap_schedule
and not get_spec().speculative_skip_dp_mlp_sync and not get_spec().speculative_skip_dp_mlp_sync
and get_disagg().disaggregation_mode == "null" and get_disagg().disaggregation_mode == "null"
@@ -103,7 +103,7 @@ class DSparkWorkerV2(BaseSpecWorker):
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
self.device = target_worker.device self.device = target_worker.device
self._draft_is_moe = draft_is_deepseek_v4(server_args=server_args) self._draft_is_moe = draft_is_deepseek_v4()
self._draft_dp_context_enabled = ( self._draft_dp_context_enabled = (
get_parallel().config.enable_dp_attention and not self._draft_is_moe get_parallel().config.enable_dp_attention and not self._draft_is_moe
) )
@@ -116,10 +116,10 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
self.pp_size = get_parallel().config.pp_size self.pp_size = get_parallel().config.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync(model_runner.server_args) self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph model_runner.server_args.enable_profile_cuda_graph
) )
@@ -102,10 +102,10 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.pp_size = get_parallel().config.pp_size self.pp_size = get_parallel().config.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync(model_runner.server_args) self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_profile_cuda_graph = ( self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph model_runner.server_args.enable_profile_cuda_graph
) )
@@ -93,10 +93,10 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync(model_runner.server_args) self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.require_attn_tp_gather = require_attn_tp_gather()
self.tp_size = self.model_runner.ps.tp_size self.tp_size = self.model_runner.ps.tp_size
self.attn_dp_size = self.model_runner.ps.attn_dp_size self.attn_dp_size = self.model_runner.ps.attn_dp_size
self.pp_size = get_parallel().config.pp_size self.pp_size = get_parallel().config.pp_size
@@ -158,10 +158,10 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.pp_size = get_parallel().config.pp_size self.pp_size = get_parallel().config.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args) self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync(model_runner.server_args) self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_pdmux = model_runner.server_args.enable_pdmux self.enable_pdmux = model_runner.server_args.enable_pdmux
self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
@@ -846,7 +846,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.cuda_graph_runner_for_draft_extend.prune_draft_extend_logits self.cuda_graph_runner_for_draft_extend.prune_draft_extend_logits
) )
else: else:
prune_logits = not require_gathered_buffer(self.server_args) prune_logits = not require_gathered_buffer()
if prune_logits: if prune_logits:
forward_batch.spec_info.select_index = select_index forward_batch.spec_info.select_index = select_index
# Left unmarked on every platform: each de-tied runner has its own # Left unmarked on every platform: each de-tied runner has its own
+13 -15
View File
@@ -105,7 +105,7 @@ from sglang.srt.runtime_context import (
from sglang.srt.utils.video_decoder import _BACKEND, VideoDecoderWrapper from sglang.srt.utils.video_decoder import _BACKEND, VideoDecoderWrapper
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs pass
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
torch_release = pkg_version.parse(torch.__version__).release torch_release = pkg_version.parse(torch.__version__).release
@@ -3719,7 +3719,7 @@ class Withable(Generic[T]):
self._value = None self._value = None
def require_mlp_tp_gather(server_args: ServerArgs): def require_mlp_tp_gather():
""" """
Check if the input of MLP is obtained by all-gather rather than all-reduce. This only happens when each MLP TP group contains multiple attention DP groups. Check if the input of MLP is obtained by all-gather rather than all-reduce. This only happens when each MLP TP group contains multiple attention DP groups.
""" """
@@ -3763,7 +3763,7 @@ def require_mlp_tp_gather(server_args: ServerArgs):
return False return False
def require_attn_tp_gather(server_args: ServerArgs): def require_attn_tp_gather():
""" """
Check if the input of attention is scattered. Check if the input of attention is scattered.
""" """
@@ -3790,35 +3790,33 @@ def require_attn_tp_gather(server_args: ServerArgs):
return False return False
def require_gathered_buffer(server_args: ServerArgs): def require_gathered_buffer():
return require_mlp_tp_gather(server_args) or require_attn_tp_gather(server_args) return require_mlp_tp_gather() or require_attn_tp_gather()
def require_mlp_sync(server_args: ServerArgs): def require_mlp_sync():
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
return get_parallel().config.enable_dp_attention or require_gathered_buffer( return get_parallel().config.enable_dp_attention or require_gathered_buffer()
server_args
)
def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int: def get_cuda_graph_batch_size_alignment() -> int:
alignment = 1 alignment = 1
if get_exec().overlap.enable_two_batch_overlap: if get_exec().overlap.enable_two_batch_overlap:
alignment *= 2 alignment *= 2
if require_gathered_buffer(server_args): if require_gathered_buffer():
alignment *= get_parallel().attn_tp_size alignment *= get_parallel().attn_tp_size
if alignment % get_parallel().attn_cp_size != 0: if alignment % get_parallel().attn_cp_size != 0:
alignment *= get_parallel().attn_cp_size alignment *= get_parallel().attn_cp_size
return alignment return alignment
def get_cuda_graph_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int: def get_cuda_graph_max_batch_size(max_batch_size: int) -> int:
return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment(server_args)) return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment())
def get_eager_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int: def get_eager_max_batch_size(max_batch_size: int) -> int:
if not require_mlp_sync(server_args): if not require_mlp_sync():
return max_batch_size return max_batch_size
from sglang.srt.layers.cp.padding import get_cp_padding_align_size from sglang.srt.layers.cp.padding import get_cp_padding_align_size
@@ -161,7 +161,7 @@ def _contains_tensor_container(value) -> bool:
) )
def get_vmm_feature_consumer_count(server_args) -> int: def get_vmm_feature_consumer_count() -> int:
if get_parallel().config.enable_dp_attention: if get_parallel().config.enable_dp_attention:
return get_parallel().config.tp_size // get_parallel().config.dp_size return get_parallel().config.tp_size // get_parallel().config.dp_size
return get_parallel().config.tp_size return get_parallel().config.tp_size
@@ -947,7 +947,7 @@ class CudaVmmFeatureTransport:
memory_size=per_worker_pool_size, memory_size=per_worker_pool_size,
recycle_interval=MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL, recycle_interval=MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
base_gpu_id=server_args.base_gpu_id, base_gpu_id=server_args.base_gpu_id,
consumer_count=get_vmm_feature_consumer_count(server_args), consumer_count=get_vmm_feature_consumer_count(),
allow_posix_fallback=server_args.nnodes == 1, allow_posix_fallback=server_args.nnodes == 1,
) )
+1 -2
View File
@@ -17,7 +17,6 @@ from sglang.srt.runtime_context import (
get_parallel, get_parallel,
get_stream, get_stream,
) )
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available
from sglang.srt.utils.host_shared_memory import ( from sglang.srt.utils.host_shared_memory import (
HostSharedMemoryManager, HostSharedMemoryManager,
@@ -66,7 +65,7 @@ def set_offloader(instance: BaseOffloader):
_instance = instance _instance = instance
def create_offloader_from_server_args(server_args: ServerArgs, dp_rank: int): def create_offloader(dp_rank: int):
if get_exec().offload.cpu_offload_gb > 0: if get_exec().offload.cpu_offload_gb > 0:
return OffloaderV1( return OffloaderV1(
cpu_offload_max_bytes=int(get_exec().offload.cpu_offload_gb * 1024**3) cpu_offload_max_bytes=int(get_exec().offload.cpu_offload_gb * 1024**3)
@@ -57,7 +57,6 @@ class TestDisaggregationServerWarmup(unittest.IsolatedAsyncioTestCase):
with patch("sglang.srt.entrypoints.http_server.aiohttp.ClientSession", Session): with patch("sglang.srt.entrypoints.http_server.aiohttp.ClientSession", Session):
status_codes = await _send_disaggregation_warmup_requests( status_codes = await _send_disaggregation_warmup_requests(
server_args=server_args,
url="http://localhost:30000", url="http://localhost:30000",
headers={"Authorization": "Bearer token"}, headers={"Authorization": "Bearer token"},
ssl_verify=False, ssl_verify=False,
@@ -1200,10 +1200,6 @@ class TestMlxOverlapScheduler(unittest.TestCase):
disaggregation_mode=None, disaggregation_mode=None,
enable_overlap=False, enable_overlap=False,
enable_overlap_mlx=False, enable_overlap_mlx=False,
server_args=SimpleNamespace(
disaggregation_decode_enable_offload_kvcache=False,
enable_hisparse=False,
),
model_config=None, model_config=None,
token_to_kv_pool_allocator=None, token_to_kv_pool_allocator=None,
tree_cache=tree_cache, tree_cache=tree_cache,
@@ -31,10 +31,6 @@ def _make_processor(case, server_mode: str = "full") -> SchedulerBatchResultProc
disaggregation_mode=None, disaggregation_mode=None,
enable_overlap=False, enable_overlap=False,
enable_overlap_mlx=False, enable_overlap_mlx=False,
server_args=SimpleNamespace(
enable_metrics=False,
enable_hisparse=False,
),
model_config=SimpleNamespace(think_end_ids=None), model_config=SimpleNamespace(think_end_ids=None),
token_to_kv_pool_allocator=Mock(), token_to_kv_pool_allocator=Mock(),
tree_cache=None, tree_cache=None,
@@ -60,7 +60,6 @@ def _make_processor() -> SchedulerBatchResultProcessor:
disaggregation_mode=None, disaggregation_mode=None,
enable_overlap=True, enable_overlap=True,
enable_overlap_mlx=False, enable_overlap_mlx=False,
server_args=SimpleNamespace(),
model_config=SimpleNamespace(think_end_ids=None), model_config=SimpleNamespace(think_end_ids=None),
token_to_kv_pool_allocator=MagicMock(), token_to_kv_pool_allocator=MagicMock(),
tree_cache=SimpleNamespace(page_size=TRACK_INTERVAL), tree_cache=SimpleNamespace(page_size=TRACK_INTERVAL),
@@ -65,7 +65,6 @@ def _make_processor() -> SchedulerBatchResultProcessor:
disaggregation_mode=None, disaggregation_mode=None,
enable_overlap=False, enable_overlap=False,
enable_overlap_mlx=False, enable_overlap_mlx=False,
server_args=SimpleNamespace(enable_metrics=False),
model_config=SimpleNamespace(think_end_ids=None), model_config=SimpleNamespace(think_end_ids=None),
token_to_kv_pool_allocator=None, token_to_kv_pool_allocator=None,
tree_cache=None, tree_cache=None,
@@ -81,21 +81,20 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
BaseMultimodalProcessor, BaseMultimodalProcessor,
) )
# The multimodal config comes from the bags.
override = get_context().override_server_args( override = get_context().override_server_args(
mm_process_config=mm_process_config, mm_process_config=mm_process_config,
allowed_media_domains=[], allowed_media_domains=[],
mm_processor_worker_num=mm_processor_worker_num,
mm_io_worker_num=mm_io_worker_num,
mm_preprocess_cache_size_mb=None,
tokenizer_worker_num=1,
trust_mm_content_hashes=False,
media_url_max_file_size_mb=64,
) )
override.install() override.install()
self.addCleanup(override.restore) self.addCleanup(override.restore)
server_args = MagicMock() server_args = MagicMock()
server_args.mm_processor_worker_num = mm_processor_worker_num
server_args.mm_io_worker_num = mm_io_worker_num
server_args.mm_preprocess_cache_size_mb = None
server_args.tokenizer_worker_num = 1
server_args.trust_mm_content_hashes = False
server_args.media_url_max_file_size_mb = 64
hf_config = MagicMock() hf_config = MagicMock()
mock_hf_processor = MagicMock() mock_hf_processor = MagicMock()
@@ -209,8 +209,7 @@ class TestStartupWeightLoadSelector(CustomTestCase):
# The parallel sizes come from the bags, so the config has to be published. # The parallel sizes come from the bags, so the config has to be published.
publish(server_args, role="test") publish(server_args, role="test")
self.addCleanup(reset_context) self.addCleanup(reset_context)
options = StartupWeightLoadOptions.from_server_args( options = StartupWeightLoadOptions.from_published_config(
server_args=server_args,
is_draft_worker=False, is_draft_worker=False,
) )
+8 -6
View File
@@ -648,6 +648,14 @@ def _k3_preprocess_config(
) )
def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls): def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls):
server_args = SimpleNamespace( server_args = SimpleNamespace(
base_gpu_id=0,
rl_on_policy_target=None,
tp_size=1,
)
with get_context().override_server_args(
mm_feature_transport="cpu",
mm_process_config={},
allowed_media_domains=[],
image_processor_backend="auto", image_processor_backend="auto",
disable_fast_image_processor=False, disable_fast_image_processor=False,
skip_tokenizer_init=False, skip_tokenizer_init=False,
@@ -656,13 +664,7 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls
tokenizer_worker_num=1, tokenizer_worker_num=1,
mm_preprocess_cache_size_mb=0, mm_preprocess_cache_size_mb=0,
trust_mm_content_hashes=False, trust_mm_content_hashes=False,
base_gpu_id=0,
rl_on_policy_target=None,
media_url_max_file_size_mb=64, media_url_max_file_size_mb=64,
)
# The multimodal config comes from the bags.
with get_context().override_server_args(
mm_feature_transport="cpu", mm_process_config={}, allowed_media_domains=[]
): ):
processor = processor_cls( processor = processor_cls(
hf_config=SimpleNamespace(media_placeholder_token_id=42), hf_config=SimpleNamespace(media_placeholder_token_id=42),
@@ -89,14 +89,15 @@ def make_processor(case, config, image_processor_cls=None):
allowed_media_domains=[], allowed_media_domains=[],
media_url_max_file_size_mb=64, media_url_max_file_size_mb=64,
) )
# The processor reads its media policy, transport and per-modality limits # Left at the default backend, the fast image processor sends the tensor to
# from the mm bag, so the fixture publishes before building it. # `cuda:<base_gpu_id>`, which a CPU-only host cannot do.
publish( publish(
ServerArgs( ServerArgs(
model_path="dummy", model_path="dummy",
mm_feature_transport=server_args.mm_feature_transport, mm_feature_transport=server_args.mm_feature_transport,
mm_process_config=server_args.mm_process_config, mm_process_config=server_args.mm_process_config,
allowed_media_domains=server_args.allowed_media_domains, allowed_media_domains=server_args.allowed_media_domains,
disable_fast_image_processor=server_args.disable_fast_image_processor,
), ),
role="tokenizer", role="tokenizer",
) )
@@ -134,14 +134,13 @@ class TestDraftPerRunnerConfig(CustomTestCase):
def test_an_unresolved_draft_falls_back_to_the_config_field(self): def test_an_unresolved_draft_falls_back_to_the_config_field(self):
"""The v2 workers pass no backend: --speculative-draft-attention-backend.""" """The v2 workers pass no backend: --speculative-draft-attention-backend."""
server_args = self._seed( self._seed(
attention_backend="fa3", speculative_draft_attention_backend="triton" attention_backend="fa3", speculative_draft_attention_backend="triton"
) )
def effective(*, is_draft_worker, passed=None): def effective(*, is_draft_worker, passed=None):
return resolve_draft_attention_backend( return resolve_draft_attention_backend(
draft_attention_backend=passed, draft_attention_backend=passed,
server_args=server_args,
is_draft_worker=is_draft_worker, is_draft_worker=is_draft_worker,
) )
@@ -137,7 +137,6 @@ _PASSED = frozenset({"model_path", "device", "random_seed"})
_EXPOSED = { _EXPOSED = {
("dllm/config.py", "max_running_requests"), ("dllm/config.py", "max_running_requests"),
("dllm/config.py", "model_path"), ("dllm/config.py", "model_path"),
("multimodal/processors/base_processor.py", "image_processor_backend"),
("speculative/spec_registry.py", "disable_overlap_schedule"), ("speculative/spec_registry.py", "disable_overlap_schedule"),
("disaggregation/encoder/server.py", "model_loader_extra_config"), ("disaggregation/encoder/server.py", "model_loader_extra_config"),
("layers/moe/utils.py", "deepep_mode"), ("layers/moe/utils.py", "deepep_mode"),
@@ -165,8 +164,6 @@ _EXPOSED = {
("entrypoints/engine.py", "enable_symm_mem"), ("entrypoints/engine.py", "enable_symm_mem"),
("entrypoints/engine.py", "reasoning_parser"), ("entrypoints/engine.py", "reasoning_parser"),
("entrypoints/engine.py", "tool_call_parser"), ("entrypoints/engine.py", "tool_call_parser"),
("eplb/eplb_manager.py", "ep_dispatch_algorithm"),
("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"),
("layers/cp/base.py", "attn_cp_size"), ("layers/cp/base.py", "attn_cp_size"),
("layers/cp/base.py", "cp_strategy"), ("layers/cp/base.py", "cp_strategy"),
("layers/cp/base.py", "enable_prefill_cp"), ("layers/cp/base.py", "enable_prefill_cp"),
@@ -178,11 +175,9 @@ _EXPOSED = {
("layers/moe/utils.py", "moe_runner_backend"), ("layers/moe/utils.py", "moe_runner_backend"),
("layers/moe/utils.py", "quantization"), ("layers/moe/utils.py", "quantization"),
("layers/moe/utils.py", "speculative_moe_runner_backend"), ("layers/moe/utils.py", "speculative_moe_runner_backend"),
("lora/lora_manager.py", "enable_lora_overlap_loading"),
("lora/marlin_lora_temp/policy.py", "lora_paths"), ("lora/marlin_lora_temp/policy.py", "lora_paths"),
("model_loader/expert_pack_runtime.py", "model_path"), ("model_loader/expert_pack_runtime.py", "model_path"),
("model_loader/expert_pack_runtime.py", "tokenizer_path"), ("model_loader/expert_pack_runtime.py", "tokenizer_path"),
("multimodal/processors/base_processor.py", "image_processor_backend"),
("parser/template_detection.py", "model_path"), ("parser/template_detection.py", "model_path"),
("speculative/adaptive_spec_params.py", "speculative_algorithm"), ("speculative/adaptive_spec_params.py", "speculative_algorithm"),
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"), ("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),