config: stop handing the record to code that does not read it (#36252)
This commit is contained in:
@@ -538,7 +538,7 @@ def decode(input_token_ids, batch, 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(
|
||||
batch,
|
||||
model_runner=model_runner,
|
||||
@@ -548,7 +548,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
|
||||
tp_group=model_runner.tp_group,
|
||||
get_idle_batch=None,
|
||||
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,
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
pass
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -307,9 +307,7 @@ def get_mooncake_transfer_engine() -> Optional[MooncakeTransferEngine]:
|
||||
return _mooncake_transfer_engine
|
||||
|
||||
|
||||
def maybe_init_shared_mooncake_transfer_engine(
|
||||
*, server_args: ServerArgs, gpu_id: int
|
||||
) -> None:
|
||||
def maybe_init_shared_mooncake_transfer_engine(*, gpu_id: int) -> None:
|
||||
"""
|
||||
Need MooncakeTransferEngine when:
|
||||
1) PD disaggregation uses mooncake for KV transfer (prefill/decode)
|
||||
|
||||
@@ -111,14 +111,14 @@ class ElasticEPStateManager:
|
||||
|
||||
inst.ep_join_rank_offset = get_parallel().config.ep_join_rank_offset
|
||||
if server_args.is_ep_joiner:
|
||||
cls._init_joiner_state(inst, server_args)
|
||||
cls._init_joiner_state(inst)
|
||||
|
||||
cls._instance = inst
|
||||
|
||||
return cls._instance
|
||||
|
||||
@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()
|
||||
inst.active_ranks.zero_()
|
||||
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(
|
||||
server_args: ServerArgs,
|
||||
url: str,
|
||||
headers: Dict[str, 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 ...")
|
||||
status_codes = asyncio.run(
|
||||
_send_disaggregation_warmup_requests(
|
||||
server_args=server_args,
|
||||
url=url,
|
||||
headers=headers,
|
||||
ssl_verify=ssl_verify,
|
||||
|
||||
@@ -18,11 +18,10 @@ from sglang.srt.eplb.expert_location import (
|
||||
get_global_expert_location_metadata,
|
||||
)
|
||||
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:
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -31,7 +30,6 @@ class EPLBManager:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
server_args: ServerArgs,
|
||||
model_config: ModelConfig,
|
||||
ps: Any,
|
||||
get_model: Callable[[], nn.Module],
|
||||
@@ -43,7 +41,6 @@ class EPLBManager:
|
||||
# These collaborators are set on ModelRunner AFTER EPLBManager is
|
||||
# constructed (model load, expert_backup_client, weight_updater), so
|
||||
# they are read through getters at rebalance time, not captured here.
|
||||
self._server_args = server_args
|
||||
self._model_config = model_config
|
||||
self._ps = ps
|
||||
self._get_model = get_model
|
||||
@@ -51,16 +48,16 @@ class EPLBManager:
|
||||
self._get_expert_backup_client = get_expert_backup_client
|
||||
self._get_weight_updater = get_weight_updater
|
||||
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_logged = False
|
||||
|
||||
# Otherwise, the circular buffer will contain stale data. If the case is needed, it can be implemented.
|
||||
assert (
|
||||
self._server_args.eplb_rebalance_num_iterations
|
||||
>= self._server_args.expert_distribution_recorder_buffer_size
|
||||
get_exec().moe.eplb_rebalance_num_iterations
|
||||
>= get_exec().moe.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:
|
||||
@@ -160,7 +157,7 @@ class EPLBManager:
|
||||
model=self._get_model(),
|
||||
new_expert_location_metadata=expert_location_metadata,
|
||||
update_layer_ids=chunk_layer_ids,
|
||||
nnodes=self._server_args.nnodes,
|
||||
nnodes=get_parallel().config.nnodes,
|
||||
tp_rank=(
|
||||
self._elastic_global_rank()
|
||||
if is_post_scale_rebalance
|
||||
@@ -169,7 +166,7 @@ class EPLBManager:
|
||||
use_flat_topology=is_post_scale_rebalance,
|
||||
expert_backup_client=self._get_expert_backup_client(),
|
||||
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(
|
||||
model_config=self._model_config
|
||||
),
|
||||
@@ -193,7 +190,6 @@ class EPLBManager:
|
||||
) -> ExpertLocationMetadata:
|
||||
if not broadcast_over_world:
|
||||
return ExpertLocationMetadata.init_by_eplb(
|
||||
self._server_args,
|
||||
self._model_config,
|
||||
logical_count,
|
||||
)
|
||||
@@ -204,7 +200,6 @@ class EPLBManager:
|
||||
# the mapping chosen for the expanded world.
|
||||
if dist.get_rank() == 0:
|
||||
computed_metadata = ExpertLocationMetadata.init_by_eplb(
|
||||
self._server_args,
|
||||
self._model_config,
|
||||
logical_count,
|
||||
# Arbitrary append topologies may not preserve node divisibility.
|
||||
@@ -220,14 +215,13 @@ class EPLBManager:
|
||||
|
||||
dist.broadcast(physical_to_logical_map, src=0)
|
||||
return ExpertLocationMetadata.init_by_mapping(
|
||||
self._server_args,
|
||||
self._model_config,
|
||||
physical_to_logical_map,
|
||||
moe_ep_rank=self._elastic_global_rank(),
|
||||
)
|
||||
|
||||
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):
|
||||
if average_utilization_rate_over_window is None:
|
||||
@@ -235,10 +229,10 @@ class EPLBManager:
|
||||
|
||||
if (
|
||||
average_utilization_rate_over_window
|
||||
> self._server_args.eplb_min_rebalancing_utilization_threshold
|
||||
> get_exec().moe.eplb_min_rebalancing_utilization_threshold
|
||||
):
|
||||
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
|
||||
|
||||
|
||||
@@ -671,12 +671,10 @@ class _Accumulator(ABC):
|
||||
expert_location_metadata: ExpertLocationMetadata,
|
||||
rank: int,
|
||||
) -> _Accumulator:
|
||||
return _Accumulator.get_class(server_args)(
|
||||
server_args, expert_location_metadata, rank
|
||||
)
|
||||
return _Accumulator.get_class()(server_args, expert_location_metadata, rank)
|
||||
|
||||
@staticmethod
|
||||
def get_class(server_args: ServerArgs) -> Type[_Accumulator]:
|
||||
def get_class() -> Type[_Accumulator]:
|
||||
return {
|
||||
"stat": _StatAccumulator,
|
||||
"stat_approx": _StatAccumulator,
|
||||
|
||||
@@ -32,7 +32,6 @@ from sglang.srt.runtime_context import (
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -105,11 +104,9 @@ class ExpertLocationMetadata:
|
||||
# -------------------------------- construction ------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def init_trivial(
|
||||
server_args: ServerArgs, model_config: ModelConfig, moe_ep_rank: int
|
||||
):
|
||||
def init_trivial(model_config: ModelConfig, moe_ep_rank: int):
|
||||
"""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:
|
||||
return None
|
||||
@@ -131,7 +128,6 @@ class ExpertLocationMetadata:
|
||||
)
|
||||
|
||||
return ExpertLocationMetadata.init_by_mapping(
|
||||
server_args,
|
||||
model_config,
|
||||
physical_to_logical_map=physical_to_logical_map,
|
||||
moe_ep_rank=moe_ep_rank,
|
||||
@@ -139,7 +135,6 @@ class ExpertLocationMetadata:
|
||||
|
||||
@staticmethod
|
||||
def init_by_mapping(
|
||||
server_args: ServerArgs,
|
||||
model_config: ModelConfig,
|
||||
physical_to_logical_map,
|
||||
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 = 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:
|
||||
return None
|
||||
@@ -179,7 +174,6 @@ class ExpertLocationMetadata:
|
||||
|
||||
@staticmethod
|
||||
def init_by_eplb(
|
||||
server_args: ServerArgs,
|
||||
model_config: ModelConfig,
|
||||
logical_count: torch.Tensor,
|
||||
*,
|
||||
@@ -193,7 +187,7 @@ class ExpertLocationMetadata:
|
||||
|
||||
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:
|
||||
return None
|
||||
@@ -229,7 +223,7 @@ class ExpertLocationMetadata:
|
||||
)
|
||||
|
||||
@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
|
||||
|
||||
model_config_for_expert_location = (
|
||||
@@ -526,9 +520,6 @@ def broadcast_global_expert_location_metadata(
|
||||
src_rank: int = 0,
|
||||
group: Optional[torch.distributed.ProcessGroup] = None,
|
||||
) -> ExpertLocationMetadata:
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
server_args = get_server_args()
|
||||
metadata = get_global_expert_location_metadata()
|
||||
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 = ExpertLocationMetadata.init_by_mapping(
|
||||
server_args,
|
||||
model_config,
|
||||
metadata.physical_to_logical_map,
|
||||
moe_ep_rank=moe_ep_rank,
|
||||
@@ -785,15 +775,12 @@ class ModelConfigForExpertLocation:
|
||||
|
||||
|
||||
def compute_initial_expert_location_metadata(
|
||||
server_args: ServerArgs,
|
||||
model_config: ModelConfig,
|
||||
moe_ep_rank: int,
|
||||
) -> Optional[ExpertLocationMetadata]:
|
||||
data = get_exec().moe.init_expert_location
|
||||
if data == "trivial":
|
||||
return ExpertLocationMetadata.init_trivial(
|
||||
server_args, model_config, moe_ep_rank
|
||||
)
|
||||
return ExpertLocationMetadata.init_trivial(model_config, moe_ep_rank)
|
||||
|
||||
# TODO unify with the utils function
|
||||
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"
|
||||
)
|
||||
return ExpertLocationMetadata.init_by_mapping(
|
||||
server_args,
|
||||
model_config,
|
||||
**data_dict,
|
||||
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"
|
||||
)
|
||||
return ExpertLocationMetadata.init_by_eplb(
|
||||
server_args, model_config, logical_count=data_dict["logical_count"]
|
||||
model_config, logical_count=data_dict["logical_count"]
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
|
||||
@@ -442,9 +442,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
)
|
||||
else:
|
||||
self.workspace_buffer = global_workspace_buffer
|
||||
max_bs = get_cuda_graph_max_batch_size(
|
||||
model_runner.server_args, model_runner.req_to_token_pool.size
|
||||
)
|
||||
max_bs = get_cuda_graph_max_batch_size(model_runner.req_to_token_pool.size)
|
||||
if kv_indptr_buf is None:
|
||||
self.kv_indptr = [
|
||||
torch.zeros(
|
||||
@@ -2254,7 +2252,7 @@ class FlashInferMultiStepDraftBackend:
|
||||
self.page_size = model_runner.page_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(
|
||||
(
|
||||
|
||||
@@ -400,7 +400,6 @@ class GDNAttnBackend(MambaAttnBackendBase):
|
||||
self.verify_intermediate_state_indices = (
|
||||
build_verify_intermediate_state_indices(
|
||||
self.req_to_token_pool.size,
|
||||
model_runner.server_args,
|
||||
model_runner.device,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -408,7 +408,6 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
||||
self.verify_intermediate_state_indices = (
|
||||
build_verify_intermediate_state_indices(
|
||||
self.req_to_token_pool.size,
|
||||
model_runner.server_args,
|
||||
model_runner.device,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -9,7 +9,7 @@ from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils.common import rank0_log
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
pass
|
||||
|
||||
|
||||
class LinearAttnKernelBackend(Enum):
|
||||
@@ -103,9 +103,7 @@ def resolve_linear_attn_backends(
|
||||
return backends
|
||||
|
||||
|
||||
def build_verify_intermediate_state_indices(
|
||||
pool_size: int, server_args: ServerArgs, device
|
||||
):
|
||||
def build_verify_intermediate_state_indices(pool_size: int, device):
|
||||
"""Per-request row index into the speculative intermediate scratch
|
||||
(`intermediate_ssm` / `intermediate_conv_window`) for the MTP /
|
||||
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
|
||||
|
||||
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)
|
||||
if padded_bs > pool_size:
|
||||
indices = torch.cat(
|
||||
|
||||
@@ -477,7 +477,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
|
||||
cp = max(get_cp_padding_align_size(), 1)
|
||||
|
||||
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:
|
||||
# Align to lcm(cp, attn_tp_size) so the CP multiple isn't undone by a
|
||||
|
||||
@@ -94,19 +94,17 @@ class LoRAManager:
|
||||
self.attn_tp_size: int = get_parallel().attn_tp_size
|
||||
self.lora_added_tokens_size: Optional[int] = None
|
||||
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.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._experts_shared_outer_override: Optional[bool] = (
|
||||
server_args.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
|
||||
get_lora().experts_shared_outer_loras
|
||||
)
|
||||
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
|
||||
|
||||
# LoRA backend for running sgemm kernels
|
||||
@@ -1030,7 +1028,6 @@ class LoRAManager:
|
||||
|
||||
def init_lora_cuda_graph_moe_buffers(
|
||||
*,
|
||||
server_args: ServerArgs,
|
||||
model: torch.nn.Module,
|
||||
lora_manager: LoRAManager,
|
||||
dtype: torch.dtype,
|
||||
|
||||
@@ -13,12 +13,9 @@ from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
|
||||
def start_disagg_service(
|
||||
server_args: ServerArgs,
|
||||
):
|
||||
def start_disagg_service():
|
||||
# Start kv bootstrap server on prefill
|
||||
disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend)
|
||||
@@ -32,16 +29,12 @@ def start_disagg_service(
|
||||
host=get_serving().host,
|
||||
port=get_disagg().disaggregation_bootstrap_port,
|
||||
)
|
||||
maybe_create_ascend_config_store(
|
||||
server_args=server_args, transfer_backend=transfer_backend
|
||||
)
|
||||
maybe_create_ascend_config_store(transfer_backend=transfer_backend)
|
||||
|
||||
return bootstrap_server
|
||||
|
||||
|
||||
def maybe_create_ascend_config_store(
|
||||
server_args: ServerArgs, transfer_backend: TransferBackend
|
||||
) -> None:
|
||||
def maybe_create_ascend_config_store(transfer_backend: TransferBackend) -> None:
|
||||
"""Also called directly by the rust-server scheduler: there the KV
|
||||
bootstrap registry is served by the embedded rust server's api listener
|
||||
(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.disaggregation_bootstrap_server = start_disagg_service(self.server_args)
|
||||
self.disaggregation_bootstrap_server = start_disagg_service()
|
||||
|
||||
# Worker IPC names for pause/continue broadcasting
|
||||
self.all_worker_ipcs: set[str] = set()
|
||||
|
||||
@@ -18,14 +18,12 @@ if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
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.ngram_info import NgramVerifyInput
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
|
||||
|
||||
def decide_needs_cpu_seq_lens(
|
||||
server_args: ServerArgs,
|
||||
attn_backends: Sequence[AttentionBackend],
|
||||
) -> bool:
|
||||
"""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 (
|
||||
RaggedVerifyMode,
|
||||
read_ragged_verify_mode,
|
||||
|
||||
@@ -416,7 +416,7 @@ class Scheduler(
|
||||
# init_soft_watchdog starts a daemon thread that reads these on its first tick.
|
||||
self.forward_ct: int = 0
|
||||
self.cur_batch_for_debug: Optional[ScheduleBatch] = None
|
||||
self.init_soft_watchdog(server_args)
|
||||
self.init_soft_watchdog()
|
||||
|
||||
# Parse args
|
||||
self.server_args = server_args
|
||||
@@ -911,7 +911,7 @@ class Scheduler(
|
||||
initialize_bf16_gemm_config(self.server_args)
|
||||
|
||||
# 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):
|
||||
worker_kwargs = dict(
|
||||
@@ -1289,7 +1289,7 @@ class Scheduler(
|
||||
|
||||
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:
|
||||
self.soft_watchdog = create_scheduler_watchdog(
|
||||
self, watchdog_timeout=x, soft=True
|
||||
@@ -1341,7 +1341,6 @@ class Scheduler(
|
||||
and self._hosts_rust_server()
|
||||
):
|
||||
maybe_create_ascend_config_store(
|
||||
server_args=self.server_args,
|
||||
transfer_backend=self.transfer_backend,
|
||||
)
|
||||
|
||||
@@ -1493,8 +1492,8 @@ class Scheduler(
|
||||
)
|
||||
else:
|
||||
attn_backends = (self.tp_worker.model_runner.attn_backend,)
|
||||
needs_cpu_seq_lens = decide_needs_cpu_seq_lens(self.server_args, attn_backends)
|
||||
needs_confidence_relay = decide_needs_confidence_relay(self.server_args)
|
||||
needs_cpu_seq_lens = decide_needs_cpu_seq_lens(attn_backends)
|
||||
needs_confidence_relay = decide_needs_confidence_relay()
|
||||
self.future_map = self.spec_algorithm.create_future_map(
|
||||
self.device,
|
||||
self.req_to_token_pool,
|
||||
@@ -2089,7 +2088,6 @@ class Scheduler(
|
||||
tree_cache=self.tree_cache,
|
||||
offload_tags=self.weight_updater.offload_tags,
|
||||
ps=self.ps,
|
||||
server_args=self.server_args,
|
||||
model_config=self.model_config,
|
||||
enable_overlap=self.enable_overlap,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
@@ -2213,7 +2211,6 @@ class Scheduler(
|
||||
disaggregation_mode=self.disaggregation_mode,
|
||||
enable_overlap=self.enable_overlap,
|
||||
enable_overlap_mlx=self.enable_overlap_mlx,
|
||||
server_args=self.server_args,
|
||||
model_config=self.model_config,
|
||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||
tree_cache=self.tree_cache,
|
||||
|
||||
@@ -70,7 +70,6 @@ if TYPE_CHECKING:
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
|
||||
from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -81,7 +80,6 @@ class SchedulerBatchResultProcessor:
|
||||
disaggregation_mode: DisaggregationMode
|
||||
enable_overlap: bool
|
||||
enable_overlap_mlx: bool
|
||||
server_args: ServerArgs
|
||||
model_config: ModelConfig
|
||||
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
|
||||
tree_cache: BasePrefixCache
|
||||
@@ -280,7 +278,6 @@ class SchedulerBatchResultProcessor:
|
||||
hidden_state_offset = 0
|
||||
prefill_hidden_capture_mode = self._get_prefill_hidden_capture_mode(
|
||||
batch,
|
||||
self.server_args,
|
||||
)
|
||||
|
||||
# Check finish conditions
|
||||
@@ -617,10 +614,7 @@ class SchedulerBatchResultProcessor:
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_prefill_hidden_capture_mode(
|
||||
batch: ScheduleBatch,
|
||||
server_args: ServerArgs,
|
||||
) -> CaptureHiddenMode:
|
||||
def _get_prefill_hidden_capture_mode(batch: ScheduleBatch) -> CaptureHiddenMode:
|
||||
return get_required_capture_hidden_mode(
|
||||
max(
|
||||
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.observability.metrics_collector import DPCooperationInfo
|
||||
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.utils.common import require_mlp_tp_gather
|
||||
|
||||
@@ -401,7 +400,6 @@ class SchedulerDPAttnAdapter:
|
||||
tree_cache: BasePrefixCache
|
||||
offload_tags: set[str]
|
||||
ps: ParallelState
|
||||
server_args: ServerArgs
|
||||
model_config: ModelConfig
|
||||
enable_overlap: bool
|
||||
spec_algorithm: SpeculativeAlgorithm
|
||||
@@ -417,7 +415,7 @@ class SchedulerDPAttnAdapter:
|
||||
tp_group=self.tp_group,
|
||||
get_idle_batch=self.get_idle_batch,
|
||||
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,
|
||||
offload_tags=self.offload_tags,
|
||||
dwdp=get_parallel().config.dwdp_size > 1,
|
||||
|
||||
@@ -81,7 +81,7 @@ from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.server_args import LoRARef, ServerArgs
|
||||
from sglang.srt.server_args import LoRARef
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
normalize_serialized_named_tensor_payloads,
|
||||
@@ -158,7 +158,7 @@ class TokenizerControlMixin:
|
||||
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 = []
|
||||
for spec in _COMMUNICATOR_SPECS:
|
||||
name, resp_type = spec[0], spec[1]
|
||||
|
||||
@@ -663,9 +663,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
# Keep a reference so the bootstrap server is not garbage-collected.
|
||||
self.bootstrap_server = (
|
||||
start_disagg_service(self.server_args)
|
||||
if start_pd_bootstrap_service
|
||||
else None
|
||||
start_disagg_service() if start_pd_bootstrap_service else None
|
||||
)
|
||||
# Single-source counter for auto-assigning fake bootstrap_room.
|
||||
self.fake_bootstrap_room_counter = 0
|
||||
@@ -761,7 +759,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
(ElasticScaleUpdateReq, self.forward_elastic_scale_update),
|
||||
]
|
||||
)
|
||||
self.init_communicators(self.server_args)
|
||||
self.init_communicators()
|
||||
|
||||
self.sampling_params_class = SamplingParams
|
||||
self.signal_handler_class = SignalHandler
|
||||
|
||||
@@ -410,7 +410,6 @@ def build_hybrid_swa_stack(
|
||||
def _deepseek_v4_num_host_pages(
|
||||
*,
|
||||
params: CacheInitParams,
|
||||
server_args: ServerArgs,
|
||||
kvcache: Any,
|
||||
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(
|
||||
params=params,
|
||||
server_args=server_args,
|
||||
kvcache=kvcache,
|
||||
page_size=page_size,
|
||||
swa_page_size=kvcache.swa_page_size,
|
||||
|
||||
@@ -1152,7 +1152,6 @@ class KVCacheConfigurator:
|
||||
kv_cache_dim=calculate_mla_kv_cache_dim(
|
||||
model_config=self.model_config,
|
||||
kv_cache_dtype=self.kv_cache_dtype,
|
||||
server_args=self.server_args,
|
||||
),
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
@@ -1356,7 +1355,6 @@ class KVCacheConfigurator:
|
||||
kv_cache_dim=calculate_mla_kv_cache_dim(
|
||||
model_config=self.model_config,
|
||||
kv_cache_dtype=self.kv_cache_dtype,
|
||||
server_args=self.server_args,
|
||||
),
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
@@ -1396,7 +1394,6 @@ class KVCacheConfigurator:
|
||||
kv_cache_dim=calculate_mla_kv_cache_dim(
|
||||
model_config=self.model_config,
|
||||
kv_cache_dtype=self.kv_cache_dtype,
|
||||
server_args=self.server_args,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -2199,10 +2196,7 @@ class KVCacheConfigurator:
|
||||
|
||||
|
||||
def calculate_mla_kv_cache_dim(
|
||||
*,
|
||||
model_config: ModelConfig,
|
||||
kv_cache_dtype: torch.dtype,
|
||||
server_args: ServerArgs,
|
||||
*, model_config: ModelConfig, kv_cache_dtype: torch.dtype
|
||||
) -> int:
|
||||
is_dsa_model = is_deepseek_dsa(model_config.hf_config)
|
||||
kv_cache_dtype = kv_cache_dtype
|
||||
|
||||
@@ -598,10 +598,10 @@ class CPUGraphRunner:
|
||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
|
||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||
self.require_gathered_buffer = require_gathered_buffer()
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
self.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),
|
||||
).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."""
|
||||
from sglang.srt.utils.common import require_mlp_tp_gather
|
||||
|
||||
dp_rank = get_parallel().attn_dp_rank
|
||||
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]
|
||||
else:
|
||||
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_utils import profile_range
|
||||
from sglang.srt.utils.offloader import (
|
||||
create_offloader_from_server_args,
|
||||
create_offloader,
|
||||
get_offloader,
|
||||
set_offloader,
|
||||
)
|
||||
@@ -270,7 +270,6 @@ class ModelRunnerOutput:
|
||||
def resolve_draft_attention_backend(
|
||||
*,
|
||||
draft_attention_backend: Optional[str],
|
||||
server_args: ServerArgs,
|
||||
is_draft_worker: bool,
|
||||
) -> Optional[str]:
|
||||
"""The attention backend a runner uses because it is a draft runner.
|
||||
@@ -342,7 +341,6 @@ class ModelRunner:
|
||||
self.device = get_device().device
|
||||
self.draft_attention_backend = resolve_draft_attention_backend(
|
||||
draft_attention_backend=draft_attention_backend,
|
||||
server_args=server_args,
|
||||
is_draft_worker=is_draft_worker,
|
||||
)
|
||||
# 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
|
||||
|
||||
# CPU offload
|
||||
set_offloader(
|
||||
create_offloader_from_server_args(server_args, dp_rank=self.ps.dp_rank)
|
||||
)
|
||||
set_offloader(create_offloader(dp_rank=self.ps.dp_rank))
|
||||
|
||||
self._weight_checker = WeightChecker(get_model=lambda: self.model, ps=self.ps)
|
||||
|
||||
@@ -591,7 +587,6 @@ class ModelRunner:
|
||||
model=self.model,
|
||||
model_config=self.model_config,
|
||||
req_to_token_pool=self.req_to_token_pool,
|
||||
server_args=self.server_args,
|
||||
max_running_requests=self.max_running_requests,
|
||||
device=self.device,
|
||||
)
|
||||
@@ -702,7 +697,6 @@ class ModelRunner:
|
||||
)
|
||||
set_global_expert_location_metadata(
|
||||
compute_initial_expert_location_metadata(
|
||||
server_args=self.server_args,
|
||||
model_config=self.model_config,
|
||||
moe_ep_rank=expert_rank,
|
||||
)
|
||||
@@ -727,7 +721,6 @@ class ModelRunner:
|
||||
def maybe_init_eplb_manager(self):
|
||||
self.eplb_manager = (
|
||||
EPLBManager(
|
||||
server_args=self.server_args,
|
||||
model_config=self.model_config,
|
||||
ps=self.ps,
|
||||
get_model=lambda: self.model,
|
||||
@@ -1076,9 +1069,7 @@ class ModelRunner:
|
||||
self.pre_model_load_memory = result.pre_model_load_memory
|
||||
|
||||
def init_shared_mooncake_transfer_engine(self):
|
||||
maybe_init_shared_mooncake_transfer_engine(
|
||||
server_args=self.server_args, gpu_id=self.gpu_id
|
||||
)
|
||||
maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id)
|
||||
|
||||
def load_model(self):
|
||||
tic_total = time.perf_counter()
|
||||
@@ -1091,9 +1082,7 @@ class ModelRunner:
|
||||
if self.device != "cpu":
|
||||
torch.set_num_threads(1)
|
||||
if self.device == "cuda":
|
||||
maybe_downgrade_dtype_for_legacy_gpu(
|
||||
server_args=self.server_args, model_config=self.model_config
|
||||
)
|
||||
maybe_downgrade_dtype_for_legacy_gpu(model_config=self.model_config)
|
||||
|
||||
set_cuda_arch()
|
||||
|
||||
@@ -1113,7 +1102,6 @@ class ModelRunner:
|
||||
# and derive the per-rank daemon socket. Idempotent across reloads.
|
||||
maybe_enable_ipc_weight_cache(
|
||||
load_config=self.load_config,
|
||||
server_args=self.server_args,
|
||||
tp_size=self.ps.tp_size,
|
||||
pp_rank=self.ps.pp_rank,
|
||||
tp_rank=self.ps.tp_rank,
|
||||
@@ -1124,7 +1112,6 @@ class ModelRunner:
|
||||
)
|
||||
|
||||
maybe_trigger_remote_instance_nccl_send_group(
|
||||
server_args=self.server_args,
|
||||
tp_rank=self.ps.tp_rank,
|
||||
load_format=draft_load_format,
|
||||
)
|
||||
@@ -1195,11 +1182,12 @@ class ModelRunner:
|
||||
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(
|
||||
model=self.model,
|
||||
server_args=self.server_args,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
is_draft_worker=self.is_draft_worker,
|
||||
tp_size=self.ps.tp_size,
|
||||
@@ -1281,7 +1269,6 @@ class ModelRunner:
|
||||
)
|
||||
if not cuda_graph_fully_disabled():
|
||||
init_lora_cuda_graph_moe_buffers(
|
||||
server_args=self.server_args,
|
||||
model=self.model,
|
||||
lora_manager=self.lora_manager,
|
||||
dtype=self.dtype,
|
||||
@@ -1471,13 +1458,11 @@ class ModelRunner:
|
||||
if (
|
||||
forward_batch.num_token_non_padded 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_mla_prefill_cp_enabled()
|
||||
):
|
||||
forward_batch.adjust_num_token_non_padded_for_attn_tp(
|
||||
server_args=self.server_args,
|
||||
)
|
||||
forward_batch.adjust_num_token_non_padded_for_attn_tp()
|
||||
|
||||
# Hisparse coordinator — backends now read it from self.model_runner.
|
||||
if self.hisparse_coordinator is not None:
|
||||
@@ -1926,7 +1911,6 @@ class ModelRunner:
|
||||
start=old_num_physical - num_local * initial_ep_size,
|
||||
)
|
||||
new_metadata = ExpertLocationMetadata.init_by_mapping(
|
||||
self.server_args,
|
||||
self.model_config,
|
||||
physical_to_logical_map=expanded_p2l,
|
||||
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
|
||||
|
||||
|
||||
def maybe_downgrade_dtype_for_legacy_gpu(
|
||||
*, server_args: ServerArgs, model_config: ModelConfig
|
||||
) -> None:
|
||||
def maybe_downgrade_dtype_for_legacy_gpu(*, model_config: ModelConfig) -> None:
|
||||
if torch.cuda.get_device_capability()[0] < 8:
|
||||
logger.info(
|
||||
"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(
|
||||
*, server_args: ServerArgs, tp_rank: int, load_format: Optional[str] = None
|
||||
*, tp_rank: int, load_format: Optional[str] = None
|
||||
) -> None:
|
||||
"""``load_format`` is this runner's effective format: a draft loading under
|
||||
``--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
|
||||
|
||||
|
||||
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.
|
||||
quantized_layers = getattr(
|
||||
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(
|
||||
*,
|
||||
model,
|
||||
server_args: ServerArgs,
|
||||
spec_algorithm: SpeculativeAlgorithm,
|
||||
is_draft_worker: bool,
|
||||
tp_size: int,
|
||||
@@ -237,7 +234,6 @@ def build_load_config(
|
||||
def maybe_enable_ipc_weight_cache(
|
||||
*,
|
||||
load_config: LoadConfig,
|
||||
server_args: ServerArgs,
|
||||
tp_size: int,
|
||||
pp_rank: int,
|
||||
tp_rank: int,
|
||||
@@ -311,12 +307,11 @@ def load_model_with_memory_saver(
|
||||
StartupWeightLoadManager,
|
||||
)
|
||||
|
||||
startup_weight_load = StartupWeightLoadManager.create_from_server_args(
|
||||
startup_weight_load = StartupWeightLoadManager.create_from_published_config(
|
||||
loader=loader,
|
||||
model_config=model_config,
|
||||
load_config=load_config,
|
||||
device_config=device_config,
|
||||
server_args=server_args,
|
||||
is_draft_worker=is_draft_worker,
|
||||
)
|
||||
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.mem_cache.memory_pool import ReqToTokenPool
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||
@@ -33,7 +32,6 @@ class NgramEmbeddingManager:
|
||||
model: torch.nn.Module,
|
||||
model_config: ModelConfig,
|
||||
req_to_token_pool: ReqToTokenPool,
|
||||
server_args: ServerArgs,
|
||||
max_running_requests: int,
|
||||
device: str,
|
||||
):
|
||||
|
||||
@@ -203,7 +203,6 @@ def _resolve_dflash_aux_hidden_state(
|
||||
config.dflash_draft_num_layers = int(draft_num_layers)
|
||||
config.dflash_target_layer_ids = target_layer_ids
|
||||
config.dflash_draft_cell_size_per_token = _resolve_dflash_draft_cell_size(
|
||||
server_args=server_args,
|
||||
draft_model_config=draft_model_config,
|
||||
draft_num_layers=int(draft_num_layers),
|
||||
)
|
||||
@@ -211,7 +210,6 @@ def _resolve_dflash_aux_hidden_state(
|
||||
|
||||
def _resolve_dflash_draft_cell_size(
|
||||
*,
|
||||
server_args: ServerArgs,
|
||||
draft_model_config: ModelConfig,
|
||||
draft_num_layers: int,
|
||||
) -> int | None:
|
||||
|
||||
@@ -31,7 +31,6 @@ from sglang.srt.runtime_context import (
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -98,12 +97,16 @@ class StartupWeightLoadOptions:
|
||||
prefetch_num_threads: int
|
||||
|
||||
@classmethod
|
||||
def from_server_args(
|
||||
def from_published_config(
|
||||
cls,
|
||||
*,
|
||||
server_args: ServerArgs,
|
||||
is_draft_worker: bool,
|
||||
) -> 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_enabled = any(
|
||||
getattr(cuda_graph_config, phase).backend != Backend.DISABLED
|
||||
@@ -267,29 +270,27 @@ class StartupWeightLoadManager:
|
||||
self._prefetch_failure_reported = False
|
||||
|
||||
@classmethod
|
||||
def create_from_server_args(
|
||||
def create_from_published_config(
|
||||
cls,
|
||||
*,
|
||||
loader,
|
||||
model_config: ModelConfig,
|
||||
load_config: LoadConfig,
|
||||
device_config: DeviceConfig,
|
||||
server_args: ServerArgs,
|
||||
is_draft_worker: bool,
|
||||
) -> 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
|
||||
knowledge of which server arguments matter, and every support rule,
|
||||
stays in this module.
|
||||
knowledge of which config leaves matter, and every support rule, stays
|
||||
in this module.
|
||||
"""
|
||||
return cls.create(
|
||||
loader=loader,
|
||||
model_config=model_config,
|
||||
load_config=load_config,
|
||||
device_config=device_config,
|
||||
options=StartupWeightLoadOptions.from_server_args(
|
||||
server_args=server_args,
|
||||
options=StartupWeightLoadOptions.from_published_config(
|
||||
is_draft_worker=is_draft_worker,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -259,7 +259,6 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
||||
calculate_mla_kv_cache_dim(
|
||||
model_config=model_config,
|
||||
kv_cache_dtype=kv_cache_dtype,
|
||||
server_args=kvc.server_args,
|
||||
)
|
||||
* effective_num_layers
|
||||
* kv_size
|
||||
@@ -465,7 +464,6 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
||||
calculate_mla_kv_cache_dim(
|
||||
model_config=model_config,
|
||||
kv_cache_dtype=kv_cache_dtype,
|
||||
server_args=kvc.server_args,
|
||||
)
|
||||
* kv_size
|
||||
)
|
||||
|
||||
@@ -70,11 +70,10 @@ def get_batch_sizes_to_capture(
|
||||
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)
|
||||
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
|
||||
# alignment constraint applies per request rather than per token row.
|
||||
alignment_width = captured_req_width
|
||||
@@ -82,7 +81,7 @@ def get_batch_sizes_to_capture(
|
||||
alignment_width = 1
|
||||
|
||||
# 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:
|
||||
# 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.
|
||||
|
||||
@@ -351,7 +351,7 @@ class BaseRunner(ABC):
|
||||
dp_size=get_parallel().config.dp_size,
|
||||
pp_size=get_parallel().config.pp_size,
|
||||
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(),
|
||||
encoder_len_fill_value=(
|
||||
getattr(mr.model_config.hf_config, "max_source_positions", 0)
|
||||
@@ -535,9 +535,9 @@ class BaseRunner(ABC):
|
||||
)
|
||||
|
||||
# TP-gather requirements for global token metadata.
|
||||
require_mlp_tp_gather_ = require_mlp_tp_gather(mr.server_args)
|
||||
require_attn_tp_gather_ = require_attn_tp_gather(mr.server_args)
|
||||
if require_gathered_buffer(mr.server_args):
|
||||
require_mlp_tp_gather_ = require_mlp_tp_gather()
|
||||
require_attn_tp_gather_ = require_attn_tp_gather()
|
||||
if require_gathered_buffer():
|
||||
assert require_mlp_tp_gather_ or require_attn_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.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather(
|
||||
model_runner.server_args
|
||||
) and not self._forward_is_dp_local(model_runner)
|
||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||
self.require_mlp_tp_gather = (
|
||||
require_mlp_tp_gather() and not self._forward_is_dp_local(model_runner)
|
||||
)
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
# Composite predicates derive from the instance values so the dp-local
|
||||
# draft exemption above stays consistent (require_gathered_buffer ==
|
||||
# mlp_tp_gather or attn_tp_gather; require_mlp_sync adds dp attention).
|
||||
@@ -597,7 +597,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
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:
|
||||
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.
|
||||
max_bs *= get_spec().speculative_eagle_topk
|
||||
# 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())
|
||||
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
|
||||
|
||||
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,
|
||||
enable_mamba_track=self.mamba_track_enabled,
|
||||
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=(
|
||||
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.dp_size = get_parallel().config.dp_size
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
|
||||
# --- backend ---------------------------------------------------
|
||||
# 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.fill_(num_tokens)
|
||||
if require_gathered_buffer(self.model_runner.server_args):
|
||||
if require_gathered_buffer():
|
||||
local = compute_local_num_token_non_padded(
|
||||
global_num_token_non_padded=buf,
|
||||
num_tokens_per_dp=num_tokens,
|
||||
|
||||
@@ -121,7 +121,6 @@ from sglang.srt.multimodal.mm_utils import materialize_multimodal_features
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.utils import is_blackwell_supported, is_hip, is_npu, make_layers
|
||||
from sglang.srt.utils.common import (
|
||||
@@ -2146,7 +2145,7 @@ class KimiK3DecoderLayer(nn.Module):
|
||||
self._dp_attention = is_dp_attention_enabled()
|
||||
# mlp-sync (DP attention OR MoE a2a/EP) pads extend batches to
|
||||
# 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
|
||||
# prefix and on the MoE cadence — same predicate the mlp construction
|
||||
# below uses.
|
||||
@@ -2614,7 +2613,7 @@ class KimiK3LinearModel(nn.Module):
|
||||
self.pp_group = get_pp_group()
|
||||
self.dspark_layers_to_capture: Optional[list[int]] = None
|
||||
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:
|
||||
embedding_quant_config = (
|
||||
|
||||
@@ -40,7 +40,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
|
||||
MmItemMemoryPool,
|
||||
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 (
|
||||
CLIENT_MEDIA_EXCEPTIONS,
|
||||
configure_media_url_security,
|
||||
@@ -215,7 +215,7 @@ class BaseMultimodalProcessor(ABC):
|
||||
self.transport_mode = transport_mode
|
||||
configure_media_url_security(
|
||||
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
|
||||
self.mm_feature_transport = (
|
||||
@@ -227,11 +227,11 @@ class BaseMultimodalProcessor(ABC):
|
||||
self.use_ipc_pool_handle_cache = (
|
||||
self.use_cuda_ipc and envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get()
|
||||
)
|
||||
self.image_processor_backend = server_args.image_processor_backend
|
||||
if server_args.disable_fast_image_processor:
|
||||
self.image_processor_backend = get_mm().image_processor_backend
|
||||
if get_mm().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
|
||||
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.
|
||||
# Split the requested service-wide budget so increasing worker count
|
||||
# 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 = (
|
||||
self.auto_mm_preprocess_cache_size_mb
|
||||
if requested_cache_mb is None
|
||||
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
|
||||
self.mm_preprocess_cache = MultimodalPreprocessCache(
|
||||
max_size_bytes=worker_cache_bytes,
|
||||
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
|
||||
# processor state when this processor will never retain artifacts.
|
||||
self.processor_fingerprint = (
|
||||
@@ -288,7 +288,7 @@ class BaseMultimodalProcessor(ABC):
|
||||
# FIXME: not accurate, model and image specific
|
||||
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")
|
||||
if 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,
|
||||
)
|
||||
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 = (
|
||||
1
|
||||
if skip_mm_pool
|
||||
@@ -402,7 +402,7 @@ class BaseMultimodalProcessor(ABC):
|
||||
# SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all
|
||||
# tokenizer workers. Each worker gets an equal share so that adding
|
||||
# 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(
|
||||
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))
|
||||
|
||||
|
||||
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.utils.hf_transformers_utils import get_config
|
||||
|
||||
|
||||
@@ -172,7 +172,7 @@ class DSparkVerifyPlanner:
|
||||
and is_dp_attention_enabled()
|
||||
and get_parallel().attn_tp_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_spec().speculative_skip_dp_mlp_sync
|
||||
and get_disagg().disaggregation_mode == "null"
|
||||
|
||||
@@ -103,7 +103,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
self.page_size = get_schedule().page_size
|
||||
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 = (
|
||||
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.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||
self.require_gathered_buffer = require_gathered_buffer()
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
self.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.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||
self.require_gathered_buffer = require_gathered_buffer()
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
self.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.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||
self.require_gathered_buffer = require_gathered_buffer()
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
self.tp_size = self.model_runner.ps.tp_size
|
||||
self.attn_dp_size = self.model_runner.ps.attn_dp_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.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
|
||||
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
|
||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||
self.require_gathered_buffer = require_gathered_buffer()
|
||||
self.require_mlp_tp_gather = require_mlp_tp_gather()
|
||||
self.require_mlp_sync = require_mlp_sync()
|
||||
self.require_attn_tp_gather = require_attn_tp_gather()
|
||||
self.enable_pdmux = model_runner.server_args.enable_pdmux
|
||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||
|
||||
@@ -846,7 +846,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
||||
self.cuda_graph_runner_for_draft_extend.prune_draft_extend_logits
|
||||
)
|
||||
else:
|
||||
prune_logits = not require_gathered_buffer(self.server_args)
|
||||
prune_logits = not require_gathered_buffer()
|
||||
if prune_logits:
|
||||
forward_batch.spec_info.select_index = select_index
|
||||
# Left unmarked on every platform: each de-tied runner has its own
|
||||
|
||||
@@ -105,7 +105,7 @@ from sglang.srt.runtime_context import (
|
||||
from sglang.srt.utils.video_decoder import _BACKEND, VideoDecoderWrapper
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
pass
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
torch_release = pkg_version.parse(torch.__version__).release
|
||||
@@ -3719,7 +3719,7 @@ class Withable(Generic[T]):
|
||||
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.
|
||||
"""
|
||||
@@ -3763,7 +3763,7 @@ def require_mlp_tp_gather(server_args: ServerArgs):
|
||||
return False
|
||||
|
||||
|
||||
def require_attn_tp_gather(server_args: ServerArgs):
|
||||
def require_attn_tp_gather():
|
||||
"""
|
||||
Check if the input of attention is scattered.
|
||||
"""
|
||||
@@ -3790,35 +3790,33 @@ def require_attn_tp_gather(server_args: ServerArgs):
|
||||
return False
|
||||
|
||||
|
||||
def require_gathered_buffer(server_args: ServerArgs):
|
||||
return require_mlp_tp_gather(server_args) or require_attn_tp_gather(server_args)
|
||||
def require_gathered_buffer():
|
||||
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
|
||||
|
||||
return get_parallel().config.enable_dp_attention or require_gathered_buffer(
|
||||
server_args
|
||||
)
|
||||
return get_parallel().config.enable_dp_attention or require_gathered_buffer()
|
||||
|
||||
|
||||
def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int:
|
||||
def get_cuda_graph_batch_size_alignment() -> int:
|
||||
alignment = 1
|
||||
if get_exec().overlap.enable_two_batch_overlap:
|
||||
alignment *= 2
|
||||
if require_gathered_buffer(server_args):
|
||||
if require_gathered_buffer():
|
||||
alignment *= get_parallel().attn_tp_size
|
||||
if alignment % get_parallel().attn_cp_size != 0:
|
||||
alignment *= get_parallel().attn_cp_size
|
||||
return alignment
|
||||
|
||||
|
||||
def get_cuda_graph_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int:
|
||||
return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment(server_args))
|
||||
def get_cuda_graph_max_batch_size(max_batch_size: int) -> int:
|
||||
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:
|
||||
if not require_mlp_sync(server_args):
|
||||
def get_eager_max_batch_size(max_batch_size: int) -> int:
|
||||
if not require_mlp_sync():
|
||||
return max_batch_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:
|
||||
return get_parallel().config.tp_size // get_parallel().config.dp_size
|
||||
return get_parallel().config.tp_size
|
||||
@@ -947,7 +947,7 @@ class CudaVmmFeatureTransport:
|
||||
memory_size=per_worker_pool_size,
|
||||
recycle_interval=MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
@@ -17,7 +17,6 @@ from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available
|
||||
from sglang.srt.utils.host_shared_memory import (
|
||||
HostSharedMemoryManager,
|
||||
@@ -66,7 +65,7 @@ def set_offloader(instance: BaseOffloader):
|
||||
_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:
|
||||
return OffloaderV1(
|
||||
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):
|
||||
status_codes = await _send_disaggregation_warmup_requests(
|
||||
server_args=server_args,
|
||||
url="http://localhost:30000",
|
||||
headers={"Authorization": "Bearer token"},
|
||||
ssl_verify=False,
|
||||
|
||||
@@ -1200,10 +1200,6 @@ class TestMlxOverlapScheduler(unittest.TestCase):
|
||||
disaggregation_mode=None,
|
||||
enable_overlap=False,
|
||||
enable_overlap_mlx=False,
|
||||
server_args=SimpleNamespace(
|
||||
disaggregation_decode_enable_offload_kvcache=False,
|
||||
enable_hisparse=False,
|
||||
),
|
||||
model_config=None,
|
||||
token_to_kv_pool_allocator=None,
|
||||
tree_cache=tree_cache,
|
||||
|
||||
@@ -31,10 +31,6 @@ def _make_processor(case, server_mode: str = "full") -> SchedulerBatchResultProc
|
||||
disaggregation_mode=None,
|
||||
enable_overlap=False,
|
||||
enable_overlap_mlx=False,
|
||||
server_args=SimpleNamespace(
|
||||
enable_metrics=False,
|
||||
enable_hisparse=False,
|
||||
),
|
||||
model_config=SimpleNamespace(think_end_ids=None),
|
||||
token_to_kv_pool_allocator=Mock(),
|
||||
tree_cache=None,
|
||||
|
||||
@@ -60,7 +60,6 @@ def _make_processor() -> SchedulerBatchResultProcessor:
|
||||
disaggregation_mode=None,
|
||||
enable_overlap=True,
|
||||
enable_overlap_mlx=False,
|
||||
server_args=SimpleNamespace(),
|
||||
model_config=SimpleNamespace(think_end_ids=None),
|
||||
token_to_kv_pool_allocator=MagicMock(),
|
||||
tree_cache=SimpleNamespace(page_size=TRACK_INTERVAL),
|
||||
|
||||
@@ -65,7 +65,6 @@ def _make_processor() -> SchedulerBatchResultProcessor:
|
||||
disaggregation_mode=None,
|
||||
enable_overlap=False,
|
||||
enable_overlap_mlx=False,
|
||||
server_args=SimpleNamespace(enable_metrics=False),
|
||||
model_config=SimpleNamespace(think_end_ids=None),
|
||||
token_to_kv_pool_allocator=None,
|
||||
tree_cache=None,
|
||||
|
||||
@@ -81,21 +81,20 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
|
||||
BaseMultimodalProcessor,
|
||||
)
|
||||
|
||||
# The multimodal config comes from the bags.
|
||||
override = get_context().override_server_args(
|
||||
mm_process_config=mm_process_config,
|
||||
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()
|
||||
self.addCleanup(override.restore)
|
||||
|
||||
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()
|
||||
mock_hf_processor = MagicMock()
|
||||
|
||||
+1
-2
@@ -209,8 +209,7 @@ class TestStartupWeightLoadSelector(CustomTestCase):
|
||||
# The parallel sizes come from the bags, so the config has to be published.
|
||||
publish(server_args, role="test")
|
||||
self.addCleanup(reset_context)
|
||||
options = StartupWeightLoadOptions.from_server_args(
|
||||
server_args=server_args,
|
||||
options = StartupWeightLoadOptions.from_published_config(
|
||||
is_draft_worker=False,
|
||||
)
|
||||
|
||||
|
||||
@@ -648,6 +648,14 @@ def _k3_preprocess_config(
|
||||
)
|
||||
def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls):
|
||||
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",
|
||||
disable_fast_image_processor=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,
|
||||
mm_preprocess_cache_size_mb=0,
|
||||
trust_mm_content_hashes=False,
|
||||
base_gpu_id=0,
|
||||
rl_on_policy_target=None,
|
||||
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(
|
||||
hf_config=SimpleNamespace(media_placeholder_token_id=42),
|
||||
|
||||
@@ -89,14 +89,15 @@ def make_processor(case, config, image_processor_cls=None):
|
||||
allowed_media_domains=[],
|
||||
media_url_max_file_size_mb=64,
|
||||
)
|
||||
# The processor reads its media policy, transport and per-modality limits
|
||||
# from the mm bag, so the fixture publishes before building it.
|
||||
# Left at the default backend, the fast image processor sends the tensor to
|
||||
# `cuda:<base_gpu_id>`, which a CPU-only host cannot do.
|
||||
publish(
|
||||
ServerArgs(
|
||||
model_path="dummy",
|
||||
mm_feature_transport=server_args.mm_feature_transport,
|
||||
mm_process_config=server_args.mm_process_config,
|
||||
allowed_media_domains=server_args.allowed_media_domains,
|
||||
disable_fast_image_processor=server_args.disable_fast_image_processor,
|
||||
),
|
||||
role="tokenizer",
|
||||
)
|
||||
|
||||
@@ -134,14 +134,13 @@ class TestDraftPerRunnerConfig(CustomTestCase):
|
||||
|
||||
def test_an_unresolved_draft_falls_back_to_the_config_field(self):
|
||||
"""The v2 workers pass no backend: --speculative-draft-attention-backend."""
|
||||
server_args = self._seed(
|
||||
self._seed(
|
||||
attention_backend="fa3", speculative_draft_attention_backend="triton"
|
||||
)
|
||||
|
||||
def effective(*, is_draft_worker, passed=None):
|
||||
return resolve_draft_attention_backend(
|
||||
draft_attention_backend=passed,
|
||||
server_args=server_args,
|
||||
is_draft_worker=is_draft_worker,
|
||||
)
|
||||
|
||||
|
||||
@@ -137,7 +137,6 @@ _PASSED = frozenset({"model_path", "device", "random_seed"})
|
||||
_EXPOSED = {
|
||||
("dllm/config.py", "max_running_requests"),
|
||||
("dllm/config.py", "model_path"),
|
||||
("multimodal/processors/base_processor.py", "image_processor_backend"),
|
||||
("speculative/spec_registry.py", "disable_overlap_schedule"),
|
||||
("disaggregation/encoder/server.py", "model_loader_extra_config"),
|
||||
("layers/moe/utils.py", "deepep_mode"),
|
||||
@@ -165,8 +164,6 @@ _EXPOSED = {
|
||||
("entrypoints/engine.py", "enable_symm_mem"),
|
||||
("entrypoints/engine.py", "reasoning_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", "cp_strategy"),
|
||||
("layers/cp/base.py", "enable_prefill_cp"),
|
||||
@@ -178,11 +175,9 @@ _EXPOSED = {
|
||||
("layers/moe/utils.py", "moe_runner_backend"),
|
||||
("layers/moe/utils.py", "quantization"),
|
||||
("layers/moe/utils.py", "speculative_moe_runner_backend"),
|
||||
("lora/lora_manager.py", "enable_lora_overlap_loading"),
|
||||
("lora/marlin_lora_temp/policy.py", "lora_paths"),
|
||||
("model_loader/expert_pack_runtime.py", "model_path"),
|
||||
("model_loader/expert_pack_runtime.py", "tokenizer_path"),
|
||||
("multimodal/processors/base_processor.py", "image_processor_backend"),
|
||||
("parser/template_detection.py", "model_path"),
|
||||
("speculative/adaptive_spec_params.py", "speculative_algorithm"),
|
||||
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
|
||||
|
||||
Reference in New Issue
Block a user