config: the last runner-side instance reads read the bags
Six reads were left on `self.server_args` outside the per-instance boundary the plan reserves for the tokenizer-manager family, and each had a different reason to be there: - `scheduler.process_input_requests` (`mm_feature_transport`) and `BaseSpecWorker._build_hicache_draft_plan` (`enable_hierarchical_cache`) are plain leaves -> `get_mm()` / `get_memory()`. - `DraftBackendFactory._create_backend` read the split backend through a *runtime-computed name* (`getattr(self.server_args, backend_name)`) and then fell back to the base field by hand -- the census's documented blind spot. The two names it can be handed are exactly the pair `attention_backends()` returns with that fallback already applied, so it reads the pair and indexes it. The draft runner's own stamp still wins when it has one. - `remote_instance_weight_loader_use_transfer_engine` and `pre_capture_activation_reserve_mb` are derived members. Both are computed from published leaves only, so both get a named accessor that derives from the bags (and therefore follows a post-publish override). The first of those two has all its inputs in one bag, so it follows the established shape: one `*_of(cfg)` helper in `arg_groups/overrides.py`, the `ServerArgs` member delegating to it, and the accessor calling it on `get_model()`. `modelexpress_transport_of` splits out the JSON parse both sides need. The second spans four bags plus the configured parallel sizes, so it exists twice like the mamba pair -- and `TestDerivedPredicatesAgreeAcrossTiers` now pins both new pairs equal over their input matrices (92 subtests). `self.server_args.X` outside the tokenizer-manager family: 11 -> 5, and the five that remain are the documented ones (the encode server's own record, the nixl connector's rank arithmetic, `GrammarManager`'s handed instance). The post-capture headroom path calls the same bag-backed `pre_capture_activation_reserve_mb` accessor the configurator uses -- the accessor advertises override-following, and a reserve that reads the record while its sibling reads the bags can disagree after a post-publish override. And the conversions' orphans go with them: `RemoteInstanceWeightTransporter` kept a `server_args` field nothing reads, and `DraftBackendFactory` parked a record it no longer consults -- both drop the parameter, and the four factory call sites stop threading one.
This commit is contained in:
@@ -31,6 +31,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import inspect
|
import inspect
|
||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple
|
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple
|
||||||
|
|
||||||
@@ -290,6 +291,42 @@ def attention_backends_of(cfg: Any) -> tuple:
|
|||||||
return prefill, decode
|
return prefill, decode
|
||||||
|
|
||||||
|
|
||||||
|
def modelexpress_transport_of(cfg: Any) -> str:
|
||||||
|
"""The modelexpress transport a config-shaped object asks for.
|
||||||
|
|
||||||
|
``modelexpress_config`` is a JSON string (or an already-parsed dict) rather
|
||||||
|
than a leaf of its own; this is the shared parse for the transfer-engine
|
||||||
|
gate (`remote_instance_transfer_engine_of`) and any future bag reader.
|
||||||
|
``ServerArgs.modelexpress_transport`` keeps its own instance-cached parse
|
||||||
|
(`_parsed_modelexpress_config`) -- same rule, cached seed-side."""
|
||||||
|
raw = cfg.modelexpress_config
|
||||||
|
if raw is None:
|
||||||
|
parsed = {}
|
||||||
|
elif isinstance(raw, str):
|
||||||
|
parsed = json.loads(raw)
|
||||||
|
else:
|
||||||
|
parsed = raw
|
||||||
|
return parsed.get("transport", "nixl")
|
||||||
|
|
||||||
|
|
||||||
|
def remote_instance_transfer_engine_of(cfg: Any, load_format: Any = None) -> bool:
|
||||||
|
"""Whether remote-instance weight loading runs over the transfer engine.
|
||||||
|
|
||||||
|
``load_format`` overrides the config's: a draft runner loading under
|
||||||
|
``--speculative-draft-load-format`` needs its own transfer engine. Every
|
||||||
|
input is a ``model`` leaf, so this serves both the pre-publish member and
|
||||||
|
the post-publish accessor."""
|
||||||
|
if cfg.remote_instance_weight_loader_start_seed_via_transfer_engine:
|
||||||
|
return True
|
||||||
|
if (load_format or cfg.load_format) != "remote_instance":
|
||||||
|
return False
|
||||||
|
backend = cfg.remote_instance_weight_loader_backend
|
||||||
|
return backend == "transfer_engine" or (
|
||||||
|
backend == "modelexpress"
|
||||||
|
and modelexpress_transport_of(cfg) == "transfer_engine"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def mamba_extra_buffer_of(cfg: Any) -> bool:
|
def mamba_extra_buffer_of(cfg: Any) -> bool:
|
||||||
"""Mid-resolution equivalent of runtime_context.mamba_extra_buffer_enabled:
|
"""Mid-resolution equivalent of runtime_context.mamba_extra_buffer_enabled:
|
||||||
reads the (possibly overlaid) strategy from a config-shaped object.
|
reads the (possibly overlaid) strategy from a config-shaped object.
|
||||||
|
|||||||
@@ -1876,7 +1876,7 @@ class Scheduler(
|
|||||||
def process_input_requests(self, recv_reqs: List):
|
def process_input_requests(self, recv_reqs: List):
|
||||||
now = time.monotonic()
|
now = time.monotonic()
|
||||||
self.session_controller.maybe_reap(now)
|
self.session_controller.maybe_reap(now)
|
||||||
if self.server_args.mm_feature_transport == "cuda_vmm":
|
if get_mm().mm_feature_transport == "cuda_vmm":
|
||||||
for recv_req in recv_reqs:
|
for recv_req in recv_reqs:
|
||||||
self._materialize_cuda_vmm_inputs(recv_req)
|
self._materialize_cuda_vmm_inputs(recv_req)
|
||||||
|
|
||||||
|
|||||||
@@ -75,6 +75,7 @@ from sglang.srt.runtime_context import (
|
|||||||
mamba_extra_buffer_enabled,
|
mamba_extra_buffer_enabled,
|
||||||
mamba_extra_buffer_lazy_enabled,
|
mamba_extra_buffer_lazy_enabled,
|
||||||
max_speculative_num_draft_tokens,
|
max_speculative_num_draft_tokens,
|
||||||
|
pre_capture_activation_reserve_mb,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
@@ -1782,7 +1783,7 @@ class KVCacheConfigurator:
|
|||||||
# Mamba state is a fixed pre-capture allocation, so it can't ride the ~0 post-capture slack.
|
# Mamba state is a fixed pre-capture allocation, so it can't ride the ~0 post-capture slack.
|
||||||
slack_gb = max(
|
slack_gb = max(
|
||||||
slack_gb,
|
slack_gb,
|
||||||
self.server_args.pre_capture_activation_reserve_mb(
|
pre_capture_activation_reserve_mb(
|
||||||
get_device_memory_capacity(self.device)
|
get_device_memory_capacity(self.device)
|
||||||
)
|
)
|
||||||
/ 1024,
|
/ 1024,
|
||||||
|
|||||||
@@ -179,6 +179,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_spec,
|
get_spec,
|
||||||
is_ep_joiner,
|
is_ep_joiner,
|
||||||
is_ep_scale_joiner,
|
is_ep_scale_joiner,
|
||||||
|
remote_instance_transfer_engine_enabled,
|
||||||
set_global_dwdp_manager,
|
set_global_dwdp_manager,
|
||||||
)
|
)
|
||||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
@@ -560,7 +561,6 @@ class ModelRunner:
|
|||||||
|
|
||||||
def init_remote_instance_weight_transporter(self):
|
def init_remote_instance_weight_transporter(self):
|
||||||
self.remote_instance_weight_transporter = RemoteInstanceWeightTransporter(
|
self.remote_instance_weight_transporter = RemoteInstanceWeightTransporter(
|
||||||
server_args=self.server_args,
|
|
||||||
get_model=lambda: self.model,
|
get_model=lambda: self.model,
|
||||||
tp_rank=self.ps.tp_rank,
|
tp_rank=self.ps.tp_rank,
|
||||||
gpu_id=self.gpu_id,
|
gpu_id=self.gpu_id,
|
||||||
@@ -671,9 +671,7 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def maybe_init_remote_instance_transfer_engine(self):
|
def maybe_init_remote_instance_transfer_engine(self):
|
||||||
if self.server_args.remote_instance_weight_loader_use_transfer_engine(
|
if remote_instance_transfer_engine_enabled(load_format=self.draft_load_format):
|
||||||
load_format=self.draft_load_format
|
|
||||||
):
|
|
||||||
self.remote_instance_weight_transporter.init_engine()
|
self.remote_instance_weight_transporter.init_engine()
|
||||||
|
|
||||||
def maybe_init_expert_location_metadata(self):
|
def maybe_init_expert_location_metadata(self):
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from sglang.srt.distributed import get_world_group
|
|||||||
from sglang.srt.mem_cache.kv_cache_configurator import mm_runtime_reservation_gb
|
from sglang.srt.mem_cache.kv_cache_configurator import mm_runtime_reservation_gb
|
||||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
|
from sglang.srt.runtime_context import pre_capture_activation_reserve_mb
|
||||||
from sglang.srt.utils.common import get_available_gpu_memory, get_device_memory_capacity
|
from sglang.srt.utils.common import get_available_gpu_memory, get_device_memory_capacity
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -72,7 +73,7 @@ def compute_post_capture_kv_resize(
|
|||||||
if eager_decode_gap or mambaish_config(model_runner.model_config) is not None:
|
if eager_decode_gap or mambaish_config(model_runner.model_config) is not None:
|
||||||
headroom_gb = max(
|
headroom_gb = max(
|
||||||
headroom_gb,
|
headroom_gb,
|
||||||
model_runner.server_args.pre_capture_activation_reserve_mb(
|
pre_capture_activation_reserve_mb(
|
||||||
get_device_memory_capacity(model_runner.device)
|
get_device_memory_capacity(model_runner.device)
|
||||||
)
|
)
|
||||||
/ 1024,
|
/ 1024,
|
||||||
|
|||||||
+6
-4
@@ -11,8 +11,11 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
|||||||
RemoteInstanceWeightLoaderBackend,
|
RemoteInstanceWeightLoaderBackend,
|
||||||
register_memory_region,
|
register_memory_region,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_model, get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
from sglang.srt.server_args import ServerArgs
|
get_model,
|
||||||
|
get_parallel,
|
||||||
|
remote_instance_transfer_engine_enabled,
|
||||||
|
)
|
||||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -20,7 +23,6 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
@dataclass(slots=True, kw_only=True)
|
@dataclass(slots=True, kw_only=True)
|
||||||
class RemoteInstanceWeightTransporter:
|
class RemoteInstanceWeightTransporter:
|
||||||
server_args: ServerArgs
|
|
||||||
get_model: Callable[[], torch.nn.Module]
|
get_model: Callable[[], torch.nn.Module]
|
||||||
tp_rank: int
|
tp_rank: int
|
||||||
gpu_id: int
|
gpu_id: int
|
||||||
@@ -55,7 +57,7 @@ class RemoteInstanceWeightTransporter:
|
|||||||
|
|
||||||
def maybe_register_and_publish_weight_info(self) -> None:
|
def maybe_register_and_publish_weight_info(self) -> None:
|
||||||
if (
|
if (
|
||||||
self.server_args.remote_instance_weight_loader_use_transfer_engine()
|
remote_instance_transfer_engine_enabled()
|
||||||
# ModelExpress owns TransferEngine memory registration and metadata
|
# ModelExpress owns TransferEngine memory registration and metadata
|
||||||
# publishing for backend=modelexpress. Re-registering here would
|
# publishing for backend=modelexpress. Re-registering here would
|
||||||
# overlap the same weight buffers.
|
# overlap the same weight buffers.
|
||||||
|
|||||||
@@ -1420,6 +1420,53 @@ def mamba_extra_buffer_lazy_enabled() -> bool:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def remote_instance_transfer_engine_enabled(load_format: str | None = None) -> bool:
|
||||||
|
"""Whether remote-instance weight loading runs over the transfer engine.
|
||||||
|
|
||||||
|
Every input is a ``model`` leaf, so this derives from the bags and follows a
|
||||||
|
post-publish override; ``ServerArgs.remote_instance_weight_loader_use_transfer_engine``
|
||||||
|
is the pre-publish equivalent, and both go through the same helper.
|
||||||
|
``load_format`` is the caller's own (a draft runner loading under
|
||||||
|
``--speculative-draft-load-format`` has one the process record does not).
|
||||||
|
"""
|
||||||
|
from sglang.srt.arg_groups.overrides import remote_instance_transfer_engine_of
|
||||||
|
|
||||||
|
return remote_instance_transfer_engine_of(get_model(), load_format)
|
||||||
|
|
||||||
|
|
||||||
|
def pre_capture_activation_reserve_mb(gpu_mem: float | None) -> float:
|
||||||
|
"""The activation working-set reserve held back before cuda-graph capture.
|
||||||
|
|
||||||
|
Derived from published leaves across four bags (``disagg`` / ``schedule`` /
|
||||||
|
``exec.graph`` / ``spec``) plus the configured parallel sizes, so it follows
|
||||||
|
a post-publish override; ``ServerArgs.pre_capture_activation_reserve_mb`` is
|
||||||
|
the pre-publish equivalent and
|
||||||
|
``TestDerivedPredicatesAgreeAcrossTiers`` pins the two equal.
|
||||||
|
"""
|
||||||
|
schedule = get_schedule()
|
||||||
|
if get_disagg().disaggregation_mode == "decode":
|
||||||
|
running_requests = (
|
||||||
|
schedule.max_running_requests
|
||||||
|
or get_exec().graph.cuda_graph_config.decode.max_bs
|
||||||
|
or 1
|
||||||
|
)
|
||||||
|
activation_tokens = max(
|
||||||
|
running_requests * (get_spec().speculative_num_draft_tokens or 1), 2048
|
||||||
|
)
|
||||||
|
elif schedule.chunked_prefill_size > 0:
|
||||||
|
activation_tokens = max(schedule.chunked_prefill_size, 2048)
|
||||||
|
else:
|
||||||
|
activation_tokens = max(schedule.max_prefill_tokens, 2048)
|
||||||
|
reserved_mem = (
|
||||||
|
512
|
||||||
|
+ activation_tokens * 1.5
|
||||||
|
+ _configured_parallel("tp_size") * _configured_parallel("pp_size") / 8 * 1024
|
||||||
|
)
|
||||||
|
if gpu_mem is not None and gpu_mem > 60 * 1024:
|
||||||
|
reserved_mem = max(reserved_mem, 10 * 1024)
|
||||||
|
return reserved_mem
|
||||||
|
|
||||||
|
|
||||||
# --- Derived config accessors ------------------------------------------------
|
# --- Derived config accessors ------------------------------------------------
|
||||||
#
|
#
|
||||||
# A few values are computed from several config fields plus the HF config, so
|
# A few values are computed from several config fields plus the HF config, so
|
||||||
|
|||||||
@@ -44,6 +44,7 @@ from sglang.srt.arg_groups.overrides import (
|
|||||||
attention_backends_of,
|
attention_backends_of,
|
||||||
mamba_extra_buffer_lazy_of,
|
mamba_extra_buffer_lazy_of,
|
||||||
mamba_extra_buffer_of,
|
mamba_extra_buffer_of,
|
||||||
|
remote_instance_transfer_engine_of,
|
||||||
resolved_view,
|
resolved_view,
|
||||||
)
|
)
|
||||||
from sglang.srt.configs.embedding_model_spec import BCGPrefillPolicy
|
from sglang.srt.configs.embedding_model_spec import BCGPrefillPolicy
|
||||||
@@ -9470,20 +9471,7 @@ class ServerArgs:
|
|||||||
def remote_instance_weight_loader_use_transfer_engine(self, load_format=None):
|
def remote_instance_weight_loader_use_transfer_engine(self, load_format=None):
|
||||||
"""``load_format`` overrides the seed's: a draft runner loading under
|
"""``load_format`` overrides the seed's: a draft runner loading under
|
||||||
``--speculative-draft-load-format`` needs its own transfer engine."""
|
``--speculative-draft-load-format`` needs its own transfer engine."""
|
||||||
# Use TransferEngine as seed backend.
|
return remote_instance_transfer_engine_of(self, load_format)
|
||||||
if self.remote_instance_weight_loader_start_seed_via_transfer_engine:
|
|
||||||
return True
|
|
||||||
# Use TransferEngine as client backend.
|
|
||||||
if (load_format or self.load_format) == "remote_instance" and (
|
|
||||||
self.remote_instance_weight_loader_backend == "transfer_engine"
|
|
||||||
or (
|
|
||||||
self.remote_instance_weight_loader_backend == "modelexpress"
|
|
||||||
and self.modelexpress_transport == "transfer_engine"
|
|
||||||
)
|
|
||||||
):
|
|
||||||
return True
|
|
||||||
else:
|
|
||||||
return False
|
|
||||||
|
|
||||||
def describe_kv_events_publisher(self) -> Optional[dict]:
|
def describe_kv_events_publisher(self) -> Optional[dict]:
|
||||||
"""Return a structured description of this server's KV-event
|
"""Return a structured description of this server's KV-event
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ from sglang.srt.model_executor.graph_memory_usage import (
|
|||||||
merge_graph_memory_usage,
|
merge_graph_memory_usage,
|
||||||
merge_graph_time_usage,
|
merge_graph_time_usage,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_exec, get_schedule
|
from sglang.srt.runtime_context import get_exec, get_memory, get_schedule
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
@@ -235,7 +235,7 @@ class BaseSpecWorker(ABC):
|
|||||||
target_model_runner = self.target_worker.model_runner
|
target_model_runner = self.target_worker.model_runner
|
||||||
target_model_runner.mtp_draft_device_pools = ()
|
target_model_runner.mtp_draft_device_pools = ()
|
||||||
spec_algorithm = target_model_runner.spec_algorithm
|
spec_algorithm = target_model_runner.spec_algorithm
|
||||||
if not self.server_args.enable_hierarchical_cache:
|
if not get_memory().enable_hierarchical_cache:
|
||||||
return HiCacheDraftPlan()
|
return HiCacheDraftPlan()
|
||||||
|
|
||||||
draft_runners = self._draft_model_runners()
|
draft_runners = self._draft_model_runners()
|
||||||
|
|||||||
@@ -1,5 +1,4 @@
|
|||||||
from sglang.srt.runtime_context import get_spec
|
from sglang.srt.runtime_context import attention_backends, get_spec
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
is_blackwell,
|
is_blackwell,
|
||||||
@@ -28,13 +27,11 @@ def _assert_draft_needs_no_conv_sidecar(draft_model_runner) -> None:
|
|||||||
class DraftBackendFactory:
|
class DraftBackendFactory:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
|
||||||
draft_model_runner,
|
draft_model_runner,
|
||||||
topk: int,
|
topk: int,
|
||||||
speculative_num_steps: int,
|
speculative_num_steps: int,
|
||||||
seed_dsa_topk_from_draft_extend: bool = False,
|
seed_dsa_topk_from_draft_extend: bool = False,
|
||||||
):
|
):
|
||||||
self.server_args = server_args
|
|
||||||
self.draft_model_runner = draft_model_runner
|
self.draft_model_runner = draft_model_runner
|
||||||
self.topk = topk
|
self.topk = topk
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
@@ -45,13 +42,14 @@ class DraftBackendFactory:
|
|||||||
def _create_backend(
|
def _create_backend(
|
||||||
self, backend_name: str, backend_map: dict, error_template: str
|
self, backend_name: str, backend_map: dict, error_template: str
|
||||||
):
|
):
|
||||||
backend_type = (
|
# The split pair with the base-backend fallback already applied.
|
||||||
self.draft_attn_backend
|
prefill_backend, decode_backend = attention_backends()
|
||||||
if self.draft_attn_backend
|
configured = (
|
||||||
else getattr(self.server_args, backend_name)
|
decode_backend
|
||||||
|
if backend_name == "decode_attention_backend"
|
||||||
|
else prefill_backend
|
||||||
)
|
)
|
||||||
if backend_type is None:
|
backend_type = self.draft_attn_backend or configured
|
||||||
backend_type = self.server_args.attention_backend
|
|
||||||
|
|
||||||
if backend_type not in backend_map:
|
if backend_type not in backend_map:
|
||||||
raise ValueError(error_template.format(backend_type=backend_type))
|
raise ValueError(error_template.format(backend_type=backend_type))
|
||||||
|
|||||||
@@ -320,7 +320,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self.draft_extend_attn_backend = None
|
self.draft_extend_attn_backend = None
|
||||||
|
|
||||||
draft_backend_factory = DraftBackendFactory(
|
draft_backend_factory = DraftBackendFactory(
|
||||||
self.server_args,
|
|
||||||
self.draft_runner,
|
self.draft_runner,
|
||||||
self.topk,
|
self.topk,
|
||||||
self.speculative_num_steps,
|
self.speculative_num_steps,
|
||||||
|
|||||||
@@ -362,7 +362,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
self.draft_extend_attn_backend_list = []
|
self.draft_extend_attn_backend_list = []
|
||||||
for step in range(self.speculative_num_steps):
|
for step in range(self.speculative_num_steps):
|
||||||
draft_backend_factory = DraftBackendFactory(
|
draft_backend_factory = DraftBackendFactory(
|
||||||
self.server_args,
|
|
||||||
self.draft_runner_list[step],
|
self.draft_runner_list[step],
|
||||||
self.topk,
|
self.topk,
|
||||||
self.speculative_num_steps,
|
self.speculative_num_steps,
|
||||||
|
|||||||
-1
@@ -498,7 +498,6 @@ def _build_eagle_draft_extend_fixture(
|
|||||||
speculative_attention_mode="prefill",
|
speculative_attention_mode="prefill",
|
||||||
)
|
)
|
||||||
draft_extend_attn_backend = DraftBackendFactory(
|
draft_extend_attn_backend = DraftBackendFactory(
|
||||||
fixture.runner.server_args,
|
|
||||||
fixture.runner,
|
fixture.runner,
|
||||||
settings.topk,
|
settings.topk,
|
||||||
settings.speculative_num_steps,
|
settings.speculative_num_steps,
|
||||||
|
|||||||
@@ -359,7 +359,6 @@ def _build_eagle_draft_fixture(
|
|||||||
)
|
)
|
||||||
_configure_runner_for_eagle_draft(fixture.runner, case, settings)
|
_configure_runner_for_eagle_draft(fixture.runner, case, settings)
|
||||||
draft_attn_backend = DraftBackendFactory(
|
draft_attn_backend = DraftBackendFactory(
|
||||||
fixture.runner.server_args,
|
|
||||||
fixture.runner,
|
fixture.runner,
|
||||||
settings.topk,
|
settings.topk,
|
||||||
settings.speculative_num_steps,
|
settings.speculative_num_steps,
|
||||||
|
|||||||
@@ -490,6 +490,13 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
|
|
||||||
|
|
||||||
class TestSchedulerMmTransportBoundary(unittest.TestCase):
|
class TestSchedulerMmTransportBoundary(unittest.TestCase):
|
||||||
|
def _publish(self, **fields):
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
|
|
||||||
|
override = get_context().override_server_args(**fields)
|
||||||
|
override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _prepare_scheduler(scheduler):
|
def _prepare_scheduler(scheduler):
|
||||||
scheduler.session_controller = SimpleNamespace(maybe_reap=MagicMock())
|
scheduler.session_controller = SimpleNamespace(maybe_reap=MagicMock())
|
||||||
@@ -501,7 +508,9 @@ class TestSchedulerMmTransportBoundary(unittest.TestCase):
|
|||||||
from sglang.srt.managers import scheduler as scheduler_module
|
from sglang.srt.managers import scheduler as scheduler_module
|
||||||
|
|
||||||
scheduler = object.__new__(scheduler_module.Scheduler)
|
scheduler = object.__new__(scheduler_module.Scheduler)
|
||||||
scheduler.server_args = SimpleNamespace(
|
# The transport gate reads the published bags, so the case publishes
|
||||||
|
# the configuration under test.
|
||||||
|
self._publish(
|
||||||
mm_feature_transport="cuda_vmm",
|
mm_feature_transport="cuda_vmm",
|
||||||
enable_broadcast_mm_inputs_process=True,
|
enable_broadcast_mm_inputs_process=True,
|
||||||
)
|
)
|
||||||
@@ -548,7 +557,7 @@ class TestSchedulerMmTransportBoundary(unittest.TestCase):
|
|||||||
return iter(self.batch)
|
return iter(self.batch)
|
||||||
|
|
||||||
scheduler = object.__new__(scheduler_module.Scheduler)
|
scheduler = object.__new__(scheduler_module.Scheduler)
|
||||||
scheduler.server_args = SimpleNamespace(mm_feature_transport="cuda_vmm")
|
self._publish(mm_feature_transport="cuda_vmm")
|
||||||
self._prepare_scheduler(scheduler)
|
self._prepare_scheduler(scheduler)
|
||||||
raw_inputs = [object(), object()]
|
raw_inputs = [object(), object()]
|
||||||
materialized = [object(), object()]
|
materialized = [object(), object()]
|
||||||
|
|||||||
@@ -391,6 +391,19 @@ class _FakeResolvedArgs:
|
|||||||
speculative_num_draft_tokens: A[int | None, Arg(help="d"), NS("spec")] = None
|
speculative_num_draft_tokens: A[int | None, Arg(help="d"), NS("spec")] = None
|
||||||
speculative_adaptive: A[bool, Arg(help="a"), NS("spec")] = False
|
speculative_adaptive: A[bool, Arg(help="a"), NS("spec")] = False
|
||||||
speculative_adaptive_config: A[str | None, Arg(help="c"), NS("spec")] = None
|
speculative_adaptive_config: A[str | None, Arg(help="c"), NS("spec")] = None
|
||||||
|
load_format: A[str, Arg(help="lf"), NS("model")] = "auto"
|
||||||
|
remote_instance_weight_loader_backend: A[str, Arg(help="rb"), NS("model")] = "nccl"
|
||||||
|
remote_instance_weight_loader_start_seed_via_transfer_engine: A[
|
||||||
|
bool, Arg(help="rs"), NS("model")
|
||||||
|
] = False
|
||||||
|
modelexpress_config: A[str | None, Arg(help="mx"), NS("model")] = None
|
||||||
|
disaggregation_mode: A[str, Arg(help="dm"), NS("disagg")] = "null"
|
||||||
|
max_running_requests: A[int | None, Arg(help="mrr"), NS("schedule")] = None
|
||||||
|
chunked_prefill_size: A[int, Arg(help="cps"), NS("schedule")] = -1
|
||||||
|
max_prefill_tokens: A[int, Arg(help="mpt"), NS("schedule")] = 16384
|
||||||
|
cuda_graph_config: A[object | None, Arg(help="cgc"), NS("exec.graph")] = None
|
||||||
|
tp_size: A[int, Arg(help="tp"), NS("parallel")] = 1
|
||||||
|
pp_size: A[int, Arg(help="pp"), NS("parallel")] = 1
|
||||||
_resolved_overrides: list = dataclasses.field(default_factory=list)
|
_resolved_overrides: list = dataclasses.field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
@@ -1013,6 +1026,74 @@ class TestDerivedPredicatesAgreeAcrossTiers(_IsolatedServerArgs):
|
|||||||
mamba_extra_buffer_lazy_enabled(),
|
mamba_extra_buffer_lazy_enabled(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_activation_reserve_matches_the_member(self):
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import pre_capture_activation_reserve_mb
|
||||||
|
|
||||||
|
graph = SimpleNamespace(decode=SimpleNamespace(max_bs=64))
|
||||||
|
cases = (
|
||||||
|
dict(disaggregation_mode="null", chunked_prefill_size=8192),
|
||||||
|
dict(disaggregation_mode="null", chunked_prefill_size=-1),
|
||||||
|
dict(
|
||||||
|
disaggregation_mode="null",
|
||||||
|
chunked_prefill_size=-1,
|
||||||
|
max_prefill_tokens=1024,
|
||||||
|
),
|
||||||
|
dict(disaggregation_mode="decode", max_running_requests=32),
|
||||||
|
dict(disaggregation_mode="decode", max_running_requests=None),
|
||||||
|
dict(
|
||||||
|
disaggregation_mode="decode",
|
||||||
|
max_running_requests=None,
|
||||||
|
speculative_num_draft_tokens=4,
|
||||||
|
),
|
||||||
|
dict(
|
||||||
|
disaggregation_mode="null",
|
||||||
|
chunked_prefill_size=8192,
|
||||||
|
tp_size=8,
|
||||||
|
pp_size=2,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
for case in cases:
|
||||||
|
for gpu_mem in (None, 20 * 1024, 80 * 1024):
|
||||||
|
with self.subTest(gpu_mem=gpu_mem, **case):
|
||||||
|
args = _FakeResolvedArgs(cuda_graph_config=graph, **case)
|
||||||
|
get_context().set_server_args(args)
|
||||||
|
self.assertEqual(
|
||||||
|
ServerArgs.pre_capture_activation_reserve_mb(args, gpu_mem),
|
||||||
|
pre_capture_activation_reserve_mb(gpu_mem),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_remote_instance_transfer_engine_matches_the_member(self):
|
||||||
|
from sglang.srt.runtime_context import remote_instance_transfer_engine_enabled
|
||||||
|
|
||||||
|
backends = ("nccl", "transfer_engine", "modelexpress")
|
||||||
|
transports = (None, '{"transport": "transfer_engine"}', '{"transport": "nixl"}')
|
||||||
|
for seed_via_te in (False, True):
|
||||||
|
for load_format in ("auto", "remote_instance"):
|
||||||
|
for backend in backends:
|
||||||
|
for mx in transports:
|
||||||
|
with self.subTest(
|
||||||
|
seed=seed_via_te,
|
||||||
|
load_format=load_format,
|
||||||
|
backend=backend,
|
||||||
|
modelexpress=mx,
|
||||||
|
):
|
||||||
|
args = _FakeResolvedArgs(
|
||||||
|
load_format=load_format,
|
||||||
|
remote_instance_weight_loader_backend=backend,
|
||||||
|
remote_instance_weight_loader_start_seed_via_transfer_engine=seed_via_te,
|
||||||
|
modelexpress_config=mx,
|
||||||
|
)
|
||||||
|
get_context().set_server_args(args)
|
||||||
|
for override in (None, "remote_instance", "auto"):
|
||||||
|
self.assertEqual(
|
||||||
|
ServerArgs.remote_instance_weight_loader_use_transfer_engine(
|
||||||
|
args, override
|
||||||
|
),
|
||||||
|
remote_instance_transfer_engine_enabled(override),
|
||||||
|
)
|
||||||
|
|
||||||
def test_attention_backends_match_the_member(self):
|
def test_attention_backends_match_the_member(self):
|
||||||
from sglang.srt.runtime_context import attention_backends
|
from sglang.srt.runtime_context import attention_backends
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user