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:
Cheng Wan
2026-08-15 00:37:06 -07:00
committed by GitHub
parent ab810e4052
commit 97279980cf
16 changed files with 201 additions and 43 deletions
+37
View File
@@ -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.
+1 -1
View File
@@ -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,
@@ -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.
+47
View File
@@ -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
+2 -14
View File
@@ -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()
+8 -10
View File
@@ -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,
@@ -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