config: an out-of-tree replacement point for every resolution-pipeline step (#39134)

This commit is contained in:
Cheng Wan
2026-09-12 17:26:54 -07:00
committed by GitHub
parent a8b5616303
commit 6804eeaabe
8 changed files with 593 additions and 92 deletions
+3 -1
View File
@@ -18,6 +18,7 @@ from sglang.srt.arg_groups.overrides import (
from sglang.srt.environ import envs
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase
from sglang.srt.runtime_context import get_platform
from sglang.srt.utils.common import get_device_memory_capacity
logger = logging.getLogger(__name__)
@@ -37,7 +38,7 @@ def handle_offload_compatibility(server_args: Any) -> None:
)
def handle_gpu_memory_settings(server_args: Any, gpu_mem):
def handle_gpu_memory_settings(server_args: Any):
"""
Configure GPU memory-dependent settings including
chunked_prefill_size, cuda_graph_config[decode].max_bs, and mem_fraction_static.
@@ -68,6 +69,7 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
)
cfg = resolving_view(server_args)
gpu_mem = get_device_memory_capacity(cfg.device)
# A copy, so an earlier declaration keeps the value it recorded.
cuda_graph_config = copy.deepcopy(cfg.cuda_graph_config)
decode_cuda_graph_config = cuda_graph_config.decode
+6 -1
View File
@@ -28,6 +28,7 @@ from sglang.srt.arg_groups.overrides import (
use_mla_backend,
validate_declarations,
)
from sglang.srt.arg_groups.resolution_hooks import run_hook
from sglang.srt.configs.embedding_model_spec import BCGPrefillPolicy
from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch
from sglang.srt.connector import ConnectorType
@@ -786,7 +787,11 @@ def handle_model_capability_adjustments(server_args: Any):
"_handle_model_capability_adjustments",
prefill_only_disable_kv_cache=True,
)
validate_prefill_only_disable_kv_cache_args(server_args)
# Through the registry, not a bare call: an out-of-tree
# replacement registered at this validator's own pipeline
# position must also win here, at this later re-validation after
# the Hopper/Blackwell no-KV-pool default declares itself.
run_hook(validate_prefill_only_disable_kv_cache_args, server_args)
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
@@ -19,6 +19,7 @@ from sglang.srt.arg_groups.overrides import (
run_post_process_pass,
should_report_expert_balancedness,
)
from sglang.srt.arg_groups.resolution_hooks import run_hook
from sglang.srt.connector import ConnectorType
from sglang.srt.environ import envs
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase
@@ -29,7 +30,12 @@ logger = logging.getLogger(__name__)
def handle_context_parallelism(server_args: Any):
validate_prefill_cp_platform(server_args)
# Through the registry, not a bare call: an out-of-tree replacement of
# `validate_prefill_cp_platform` registered at its own (earlier) pipeline
# position must also win here, or a package permitting prefill CP on its
# own qualified HIP/NPU/MUSA build would still hit the original rejection
# at this later, nested call.
run_hook(validate_prefill_cp_platform, server_args)
cfg = resolving_view(server_args)
if parse_connector_type(cfg.model_path) != ConnectorType.INSTANCE:
+83 -80
View File
@@ -18,7 +18,7 @@ from sglang.srt.arg_groups.overrides import (
resolving_view,
run_post_process_pass,
)
from sglang.srt.utils.common import get_device_memory_capacity
from sglang.srt.arg_groups.resolution_hooks import run_hook
def run_resolution_pipeline(server_args: Any) -> None:
@@ -43,6 +43,11 @@ def run_resolution_pipeline(server_args: Any) -> None:
5. Give each handler one clear contract: what state it expects, what it
may mutate, and whether it validates only. Long ordering comments
belong in the helper or signal that the helper should be split.
6. Call each step through ``run_hook(handle_x, server_args, ...)``, not
``handle_x(server_args, ...)`` directly -- see
``arg_groups/resolution_hooks.py``. This is every step's fixed
position in the pipeline either way; registering an override changes
what runs here, never when.
"""
# What the caller asked for, before any handler runs; this plus the
@@ -62,7 +67,7 @@ def run_resolution_pipeline(server_args: Any) -> None:
from sglang.srt.arg_groups.mega_moe_hook import handle_mega_moe
handle_mega_moe(server_args)
run_hook(handle_mega_moe, server_args)
from sglang.srt.arg_groups.serving_hook import (
handle_asr_validation,
handle_crash_dump_env,
@@ -81,17 +86,17 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_tokenizer_batching,
)
handle_return_hidden_states_mode(server_args)
handle_media_url_security(server_args)
run_hook(handle_return_hidden_states_mode, server_args)
run_hook(handle_media_url_security, server_args)
from sglang.srt.arg_groups.hicache_hook import (
handle_hicache,
handle_hicache_ratio_default,
)
handle_hicache_ratio_default(server_args)
run_hook(handle_hicache_ratio_default, server_args)
from sglang.srt.arg_groups.memory_hook import handle_offload_compatibility
handle_offload_compatibility(server_args)
run_hook(handle_offload_compatibility, server_args)
from sglang.srt.arg_groups.validation_hook import (
default_unset_prefill_decode_interval,
validate_experimental_sgl_marlin,
@@ -99,8 +104,8 @@ def run_resolution_pipeline(server_args: Any) -> None:
validate_sampling_mask_max_tokens,
)
validate_prefill_decode_interval(server_args)
validate_sampling_mask_max_tokens(server_args)
run_hook(validate_prefill_decode_interval, server_args)
run_hook(validate_sampling_mask_max_tokens, server_args)
# Reject an explicitly enabled but incompatible hardware runtime before
# model path resolution, downloads, or the dummy-model short circuit.
@@ -109,8 +114,8 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_hardware_runtime_validation,
)
validate_prefill_cp_platform(server_args)
handle_hardware_runtime_validation()
run_hook(validate_prefill_cp_platform, server_args)
run_hook(handle_hardware_runtime_validation, server_args)
if cfg.model_path.lower() in ["none", "dummy"]:
return
@@ -119,30 +124,30 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_model_source_paths,
)
handle_model_source_paths(server_args)
run_hook(handle_model_source_paths, server_args)
# Validate mm_process_config.
handle_multimodal(server_args)
run_hook(handle_multimodal, server_args)
# Validate SSL arguments early.
handle_ssl_validation(server_args)
run_hook(handle_ssl_validation, server_args)
# Validate transcription/ASR-specific server args.
handle_asr_validation(server_args)
run_hook(handle_asr_validation, server_args)
# Handle deprecated arguments.
handle_deprecated_args(server_args)
run_hook(handle_deprecated_args, server_args)
# Handle deprecated environment variables for prefill delayer.
handle_prefill_delayer_env_compat(server_args)
run_hook(handle_prefill_delayer_env_compat, server_args)
# Set missing default values.
handle_missing_default_values(server_args)
run_hook(handle_missing_default_values, server_args)
# expert_pack may replace a raw GGUF input with its generated local
# model metadata before any model-specific handler calls model_config_of.
# It also establishes eager-only invariants before CUDA graph parsing.
from sglang.srt.arg_groups.expert_pack_hook import handle_expert_pack
handle_expert_pack(server_args)
run_hook(handle_expert_pack, server_args)
# Validate PD disaggregation flags before CUDA graph config.
from sglang.srt.arg_groups.pd_disaggregation_hook import (
@@ -150,7 +155,7 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_pd_disaggregation,
)
handle_pd_disaggregation(server_args)
run_hook(handle_pd_disaggregation, server_args)
from sglang.srt.arg_groups.kv_cache_hook import (
handle_cache_compatibility,
@@ -171,8 +176,8 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_expert_distribution_metrics,
)
validate_prefill_only_disable_kv_cache_args(server_args)
handle_decode_context_parallelism(server_args)
run_hook(validate_prefill_only_disable_kv_cache_args, server_args)
run_hook(handle_decode_context_parallelism, server_args)
# Model-arch prefill CUDA-graph default must land before cuda-graph
# resolution (the declarative registry materializes too late to affect
@@ -186,16 +191,16 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_cuda_graph_config,
)
apply_inkling_prefill_cuda_graph_default(server_args)
apply_muse_glimmer_prefill_cuda_graph_max_bs_default(server_args)
run_hook(apply_inkling_prefill_cuda_graph_default, server_args)
run_hook(apply_muse_glimmer_prefill_cuda_graph_max_bs_default, server_args)
# must run before _handle_cuda_graph_config and _handle_data_parallelism
handle_dwdp(server_args)
run_hook(handle_dwdp, server_args)
handle_cuda_graph_config(server_args)
run_hook(handle_cuda_graph_config, server_args)
# Requires the parsed backend and explicit-input locks, and must precede
# handle_gpu_memory_settings so the chunk size feeds memory budgeting.
apply_glm5_chunked_prefill_default(server_args)
run_hook(apply_glm5_chunked_prefill_default, server_args)
# Handle device-specific backends.
from sglang.srt.arg_groups.platform_hook import (
@@ -210,23 +215,21 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_xpu_backends,
)
handle_hpu_backends(server_args)
handle_cpu_backends(server_args)
handle_npu_backends(server_args)
handle_mps_backends(server_args)
handle_xpu_backends(server_args)
run_hook(handle_hpu_backends, server_args)
run_hook(handle_cpu_backends, server_args)
run_hook(handle_npu_backends, server_args)
run_hook(handle_mps_backends, server_args)
run_hook(handle_xpu_backends, server_args)
# Must precede handle_gpu_memory_settings: its symm-mem prealloc default
# keys off enable_symm_mem.
handle_symm_mem_device_support(server_args)
run_hook(handle_symm_mem_device_support, server_args)
handle_platform_defaults(server_args)
gpu_mem = get_device_memory_capacity(cfg.device)
run_hook(handle_platform_defaults, server_args)
# Handle memory-related, chunked prefill, and CUDA graph batch size configurations.
from sglang.srt.arg_groups.memory_hook import handle_gpu_memory_settings
handle_gpu_memory_settings(server_args, gpu_mem)
run_hook(handle_gpu_memory_settings, server_args)
# Apply model-specific adjustments.
from sglang.srt.arg_groups.model_hook import (
@@ -234,10 +237,10 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_model_specific_adjustments,
)
handle_model_specific_adjustments(server_args)
default_unset_prefill_decode_interval(server_args)
run_hook(handle_model_specific_adjustments, server_args)
run_hook(default_unset_prefill_decode_interval, server_args)
# After the model overrides: Qwen4-Exp declares the PLE offload default there.
handle_offload_compatibility(server_args)
run_hook(handle_offload_compatibility, server_args)
# Set kernel backends.
run_post_process_pass(server_args, _sampling_backend_default)
@@ -250,49 +253,49 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_multi_item_scoring,
)
handle_deterministic_inference(server_args)
handle_attention_backend_compatibility(server_args)
run_hook(handle_deterministic_inference, server_args)
run_hook(handle_attention_backend_compatibility, server_args)
# Must run after the attention backend is resolved so the trtllm_mla
# default (auto-selected for DeepseekV3ForCausalLM on sm100) is visible.
disable_prefill_cuda_graph_for_deepseek_trtllm_mla(server_args)
run_hook(disable_prefill_cuda_graph_for_deepseek_trtllm_mla, server_args)
from sglang.srt.arg_groups.mamba_hook import (
handle_int8_mamba_checkpoint,
handle_mamba_backend,
)
handle_mamba_backend(server_args)
handle_int8_mamba_checkpoint(server_args)
handle_linear_attn_backend(server_args)
apply_glm5_prefill_cuda_graph_policy(server_args)
handle_kv4_compatibility(server_args)
handle_mxfp8_kv_cache_compatibility(server_args)
run_hook(handle_mamba_backend, server_args)
run_hook(handle_int8_mamba_checkpoint, server_args)
run_hook(handle_linear_attn_backend, server_args)
run_hook(apply_glm5_prefill_cuda_graph_policy, server_args)
run_hook(handle_kv4_compatibility, server_args)
run_hook(handle_mxfp8_kv_cache_compatibility, server_args)
run_post_process_pass(server_args, _page_size_default)
handle_amd_specifics(server_args)
handle_nccl_pre_warm(server_args)
handle_grammar_backend(server_args)
run_hook(handle_amd_specifics, server_args)
run_hook(handle_nccl_pre_warm, server_args)
run_hook(handle_grammar_backend, server_args)
# Handle multi-item scoring constraints. Must run after the above so
# the final attention backend and chunked_prefill_size are in effect.
handle_multi_item_scoring(server_args)
run_hook(handle_multi_item_scoring, server_args)
# Backend-dependent half of --prefill-only-disable-kv-cache validation.
# Must stay after _handle_attention_backend_compatibility() (above) and
# _handle_multi_item_scoring() so the resolved prefill backend is final;
# the flag/precondition half runs earlier in
# _validate_prefill_only_disable_kv_cache_args().
handle_prefill_only_disable_kv_cache(server_args)
run_hook(handle_prefill_only_disable_kv_cache, server_args)
# Handle Hicache settings.
handle_hicache(server_args)
run_hook(handle_hicache, server_args)
# Handle data parallelism.
handle_data_parallelism(server_args)
run_hook(handle_data_parallelism, server_args)
# Normalize load balancing defaults.
handle_load_balance_method(server_args)
run_hook(handle_load_balance_method, server_args)
# Handle context parallelism.
handle_context_parallelism(server_args)
run_hook(handle_context_parallelism, server_args)
# Handle MoE configurations.
from sglang.srt.arg_groups.moe_hook import (
@@ -303,12 +306,12 @@ def run_resolution_pipeline(server_args: Any) -> None:
validate_deepep_v2_speculative_draft,
)
handle_moe_kernel_config(server_args)
handle_a2a_moe(server_args)
handle_eplb_and_dispatch(server_args)
handle_expert_distribution_metrics(server_args)
handle_elastic_ep(server_args)
validate_experimental_sgl_marlin(server_args)
run_hook(handle_moe_kernel_config, server_args)
run_hook(handle_a2a_moe, server_args)
run_hook(handle_eplb_and_dispatch, server_args)
run_hook(handle_expert_distribution_metrics, server_args)
run_hook(handle_elastic_ep, server_args)
run_hook(validate_experimental_sgl_marlin, server_args)
# Handle pipeline parallelism.
run_post_process_pass(server_args, _pipeline_parallel_overlap_disable)
@@ -317,55 +320,55 @@ def run_resolution_pipeline(server_args: Any) -> None:
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
handle_speculative_decoding(server_args)
run_hook(handle_speculative_decoding, server_args)
# After the speculative hook so speculative_algorithm is final.
from sglang.srt.arg_groups.layernorm_sp_hook import handle_layernorm_sp
handle_layernorm_sp(server_args)
run_hook(handle_layernorm_sp, server_args)
# Validate the CuteDSL A2A token budget now that num_tokens_per_req is final.
validate_cutedsl_a2a_token_budget(server_args)
run_hook(validate_cutedsl_a2a_token_budget, server_args)
# Handle model loading format.
handle_load_format(server_args)
run_hook(handle_load_format, server_args)
# Handle Encoder disaggregation.
handle_encoder_disaggregation(server_args)
run_hook(handle_encoder_disaggregation, server_args)
# Validate tokenizer settings.
handle_tokenizer_batching(server_args)
run_hook(handle_tokenizer_batching, server_args)
# Propagate environment variables.
handle_environment_variables(server_args)
run_hook(handle_environment_variables, server_args)
# Validate cache settings.
handle_cache_compatibility(server_args)
run_hook(handle_cache_compatibility, server_args)
handle_page_major_kv_layout(server_args)
run_hook(handle_page_major_kv_layout, server_args)
handle_unified_memory_pool(server_args)
run_hook(handle_unified_memory_pool, server_args)
# Handle diffusion LLM inference.
from sglang.srt.arg_groups.dllm_hook import handle_dllm_inference
handle_dllm_inference(server_args)
run_hook(handle_dllm_inference, server_args)
# Handle crash dump environment variables (must run before CUDA init).
handle_crash_dump_env(server_args)
run_hook(handle_crash_dump_env, server_args)
# Handle debug utilities.
handle_debug_utils(server_args)
run_hook(handle_debug_utils, server_args)
# Handle any other necessary validations.
handle_other_validations(server_args)
run_hook(handle_other_validations, server_args)
# Model-capability adjustments that legacy code applied at model-load
# time; last declarations of the resolution, mirroring that order.
handle_model_capability_adjustments(server_args)
run_hook(handle_model_capability_adjustments, server_args)
# Validate after all batch-size declarations are visible.
validate_deepep_v2_speculative_draft(server_args)
validate_deepep_v2_dispatch_token_budget(server_args)
run_hook(validate_deepep_v2_speculative_draft, server_args)
run_hook(validate_deepep_v2_dispatch_token_budget, server_args)
server_args._resolution_finished = True
@@ -20,11 +20,13 @@ from sglang.srt.utils.common import is_host_cpu_arm64
logger = logging.getLogger(__name__)
def handle_hardware_runtime_validation():
# This is intentionally independent of `server_args.device`: setting
# SGLANG_USE_MLX opts into the MLX backend and must fail immediately if
# the environment cannot honor that request. With the flag unset,
# use_mlx() remains lazy and does not import MLX.
def handle_hardware_runtime_validation(server_args: Any):
# `server_args` is accepted, not read: every resolution-hook step takes
# it, uniformly, so `run_hook` never has to special-case an arity. The
# check below is intentionally independent of `server_args.device`:
# setting SGLANG_USE_MLX opts into the MLX backend and must fail
# immediately if the environment cannot honor that request. With the
# flag unset, use_mlx() remains lazy and does not import MLX.
use_mlx()
@@ -0,0 +1,204 @@
"""Out-of-tree replacement for a named step of the resolution pipeline.
`run_resolution_pipeline` calls its steps by name, hardcoded, with no
per-step dispatch through `self` -- there is nothing on `ServerArgs` left to
subclass in order to change how one step decides. This is the replacement
for that: a decorator that wraps whatever currently runs under a given name,
and a dispatcher the pipeline calls instead of the bare function. The name
itself is never spelled out at the call site -- `run_hook` reads it off the
function it was handed -- so there is exactly one place a step's name is
written by hand: the whitelist below, and whatever a downstream registrant
passes to the decorator.
Whitelisted names only, the same discipline `Arg(resolvable=True)` uses for
declarable fields: this project has already been bitten once by an
unqualified name collision in this exact pipeline (`_parse_cuda_graph_config`
and `_handle_cuda_graph_config` merged under one rename and the dispatcher
called itself). A name that is not on the list fails loudly at import time,
not silently at the call site three modules away.
Every step takes exactly `(server_args)`, `handle_hardware_runtime_validation`
included -- it does not read `server_args` (see the comment at its
definition), but it takes the parameter anyway so `run_hook` never has to
special-case an arity. An override's own signature always matches: `def
mine(server_args, previous)`.
"""
from __future__ import annotations
from typing import Any, Callable, Dict, FrozenSet, List
# Every step this pipeline runs by name that a downstream package may replace.
# Add a name here only alongside the call site's own switch to `run_hook` --
# an entry with no matching `run_hook(...)` call is a name nothing will ever
# look up. `test_every_whitelisted_hook_has_a_call_site` in
# test_resolution_hook_registry.py holds the two directions of this equal by
# construction, so this list cannot drift from `pipeline.py` silently.
#
# `handle_offload_compatibility` runs at two different points in the pipeline
# (once before model-specific adjustments, once after -- see the comments at
# its call sites). An override registered for it applies at both, identically;
# there is no way to target "just the second call" through this mechanism,
# because the name is all `run_hook` has to key on.
_OVERRIDABLE_HOOKS: FrozenSet[str] = frozenset(
{
"handle_mega_moe",
"handle_return_hidden_states_mode",
"handle_media_url_security",
"handle_hicache_ratio_default",
"handle_offload_compatibility",
"validate_prefill_decode_interval",
"default_unset_prefill_decode_interval",
"validate_sampling_mask_max_tokens",
"validate_prefill_cp_platform",
"handle_hardware_runtime_validation",
"handle_model_source_paths",
"handle_multimodal",
"handle_ssl_validation",
"handle_asr_validation",
"handle_deprecated_args",
"handle_prefill_delayer_env_compat",
"handle_missing_default_values",
"handle_expert_pack",
"handle_pd_disaggregation",
"validate_prefill_only_disable_kv_cache_args",
"handle_decode_context_parallelism",
"apply_inkling_prefill_cuda_graph_default",
"apply_muse_glimmer_prefill_cuda_graph_max_bs_default",
"handle_dwdp",
"handle_cuda_graph_config",
"apply_glm5_chunked_prefill_default",
"handle_hpu_backends",
"handle_cpu_backends",
"handle_npu_backends",
"handle_mps_backends",
"handle_xpu_backends",
"handle_symm_mem_device_support",
"handle_platform_defaults",
"handle_gpu_memory_settings",
"handle_model_specific_adjustments",
"handle_deterministic_inference",
"handle_attention_backend_compatibility",
"disable_prefill_cuda_graph_for_deepseek_trtllm_mla",
"handle_mamba_backend",
"handle_int8_mamba_checkpoint",
"handle_linear_attn_backend",
"apply_glm5_prefill_cuda_graph_policy",
"handle_kv4_compatibility",
"handle_mxfp8_kv_cache_compatibility",
"handle_amd_specifics",
"handle_nccl_pre_warm",
"handle_grammar_backend",
"handle_multi_item_scoring",
"handle_prefill_only_disable_kv_cache",
"handle_hicache",
"handle_data_parallelism",
"handle_load_balance_method",
"handle_context_parallelism",
"handle_moe_kernel_config",
"handle_a2a_moe",
"handle_eplb_and_dispatch",
"handle_expert_distribution_metrics",
"handle_elastic_ep",
"validate_experimental_sgl_marlin",
"handle_speculative_decoding",
"handle_layernorm_sp",
"validate_cutedsl_a2a_token_budget",
"handle_load_format",
"handle_encoder_disaggregation",
"handle_tokenizer_batching",
"handle_environment_variables",
"handle_cache_compatibility",
"handle_page_major_kv_layout",
"handle_unified_memory_pool",
"handle_dllm_inference",
"handle_crash_dump_env",
"handle_debug_utils",
"handle_other_validations",
"handle_model_capability_adjustments",
"validate_deepep_v2_speculative_draft",
"validate_deepep_v2_dispatch_token_budget",
}
)
# name -> registered overrides, oldest first. Each takes `(server_args,
# previous)`, where `previous` is the callable it wraps -- the built-in on
# the first registration, the previous registrant's own wrapper on every one
# after. Process-global as a `dict[str, list]` so a test isolates it the way
# test_model_overrides.py isolates `_MODEL_OVERRIDE_FNS`:
# `patch.dict(..., clear=True)`.
_HOOKS: Dict[str, List[Callable[[Any, Callable[[Any], None]], None]]] = {}
def register_resolution_hook(name: str):
"""Replace (or wrap) the pipeline step named ``name``.
The decorated function is called as ``fn(server_args, previous)``.
``previous`` is a plain ``server_args -> None`` callable: the built-in
step on the first registration for this name, or the previous
registrant's own wrapper on every registration after that. Call it to
run what would have run without this override -- the `super().handle_x()`
shape, expressed as an explicit argument instead of a method-resolution
lookup, because there is no class hierarchy here for `super()` to walk.
Not calling it is a full replacement.
Registering twice for the same name does not replace the first
registration; it wraps it. The **last** registration is outermost --
runs first, and decides whether/when its `previous` (everything
registered before it, down to the built-in) runs at all. Two downstream
packages that both target the same name compose in whichever order they
happened to import in; if that order matters to you, make one of them
import the other first.
This changes *what* runs at the step's existing position in the
pipeline, never *when*: the call site in `pipeline.py` is unmoved, so
every other step keeps the order it already had. A wrapped step's own
declarations reach `resolution_result` the same way any declaration
does -- only readers from this position onward see them; a step that
already ran and read the old value before this one's `previous` (or the
built-in) declared its replacement has already made its decision on it.
"""
if name not in _OVERRIDABLE_HOOKS:
raise ValueError(
f"{name!r} is not an overridable resolution hook; the "
f"overridable set is {sorted(_OVERRIDABLE_HOOKS)}. A new entry "
"needs a matching `run_hook(...)` call at the step's site in "
"pipeline.py, not just a name here."
)
def decorator(fn):
_HOOKS.setdefault(name, []).append(fn)
return fn
return decorator
def run_hook(builtin: Callable[[Any], None], server_args: Any) -> None:
"""Run ``builtin`` -- the registered chain if anything overrode it under
its name, ``builtin`` directly otherwise.
The name is ``builtin.__name__``, not a second argument: the call site
already has the function in scope (a function-local import, same as
every other step), and spelling the name out again next to it is the
exact "two copies that can silently disagree" shape this project keeps
removing elsewhere. ``builtin`` must therefore be a plain, named
function -- every real call site is -- not a lambda or a bound method.
Called from the step's fixed position in `run_resolution_pipeline`. The
chain is rebuilt from the registry on every call rather than cached at
registration time, because at registration time (import time, before any
`ServerArgs` exists) there is no `server_args` yet and the first
registrant's `previous` cannot be bound to anything real until a call
actually happens.
"""
step = builtin
for fn in _HOOKS.get(builtin.__name__, ()):
step = _bind(fn, step)
step(server_args)
def _bind(fn, previous):
def wrapped(server_args):
fn(server_args, previous)
return wrapped
@@ -0,0 +1,268 @@
"""Unit tests for the out-of-tree resolution-hook registry: whitelist
enforcement, wrap-vs-replace semantics, multi-registrant composition, and the
end-to-end proof that an override reaches a real `resolve_once()`."""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
import ast
import glob
import json
import os
import shutil
import tempfile
import unittest
from unittest.mock import patch
from sglang.srt.arg_groups import pipeline as pipeline_module
from sglang.srt.arg_groups import resolution_hooks as hooks_module
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.arg_groups.resolution_hooks import (
_OVERRIDABLE_HOOKS,
register_resolution_hook,
run_hook,
)
from sglang.srt.server_args import ServerArgs
from sglang.test.test_utils import CustomTestCase
# `model_path="dummy"` short-circuits the pipeline before this step ever
# runs (the fixture the rest of this file uses on purpose, since it does not
# need the step to have run). The end-to-end proof below needs the real
# pipeline, so it needs a real, tiny HF config on disk instead.
_MINI_CONFIG = {
"architectures": ["LlamaForCausalLM"],
"model_type": "llama",
"hidden_size": 16,
"intermediate_size": 32,
"num_attention_heads": 2,
"num_key_value_heads": 2,
"num_hidden_layers": 2,
"vocab_size": 128,
"max_position_embeddings": 2048,
}
class _IsolatedRegistry(CustomTestCase):
"""Run each test against an empty registry (it is process-global)."""
def setUp(self):
super().setUp()
self._patch = patch.dict(hooks_module._HOOKS, clear=True)
self._patch.start()
self.addCleanup(self._patch.stop)
class TestWhitelist(_IsolatedRegistry):
def test_an_unlisted_name_is_refused_at_registration(self):
with self.assertRaisesRegex(ValueError, "not an overridable resolution hook"):
@register_resolution_hook("handle_something_nobody_whitelisted")
def _fn(server_args, previous):
pass
def test_the_whitelisted_name_registers(self):
@register_resolution_hook("handle_cuda_graph_config")
def _fn(server_args, previous):
pass
self.assertIn(_fn, hooks_module._HOOKS["handle_cuda_graph_config"])
class TestWhitelistMatchesThePipeline(CustomTestCase):
"""The whitelist and `pipeline.py` are two hand-written lists that have to
agree; nothing keeps them in sync on its own. This derives both sides
fresh and asserts they are the same set, so a name added to one without
the other fails here instead of surfacing as "my override never runs" or
"this step nobody can ever replace" months later."""
def _run_hook_call_targets(self) -> set:
"""The first argument of every `run_hook(...)` call in pipeline.py,
statically -- the set of steps the pipeline actually dispatches
through the registry."""
tree = ast.parse(open(pipeline_module.__file__).read())
names = set()
for node in ast.walk(tree):
if (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Name)
and node.func.id == "run_hook"
and node.args
and isinstance(node.args[0], ast.Name)
):
names.add(node.args[0].id)
return names
def test_every_whitelisted_hook_has_a_call_site(self):
called = self._run_hook_call_targets()
self.assertEqual(
(sorted(_OVERRIDABLE_HOOKS - called), sorted(called - _OVERRIDABLE_HOOKS)),
([], []),
"(whitelisted but never called, called but not whitelisted) -- "
"both should be empty",
)
def test_no_bare_calls_to_a_whitelisted_hook_anywhere_in_arg_groups(self):
"""A step called through `run_hook` at its own pipeline position can
still be called *again*, directly, from inside another hook's body --
a re-validation after a later declaration, say. That nested call
bypasses the registry: an out-of-tree replacement registered for the
name wins at the pipeline position but not here, which is exactly the
kind of gap `test_every_whitelisted_hook_has_a_call_site` cannot see,
since it only looks at `run_hook(...)` call targets, not at every
other way a whitelisted name's bare function can be invoked.
Two real instances of this existed (`validate_prefill_cp_platform`
called directly inside `handle_context_parallelism`,
`validate_prefill_only_disable_kv_cache_args` called directly inside
`handle_model_capability_adjustments`) before both were routed
through `run_hook` too. This asserts the count stays at zero rather
than grandfathering it, since a bare call to a whitelisted name from
inside `arg_groups/` is never correct -- it always means the same
function is reachable two ways, only one of which a downstream
override can see.
"""
hook_dir = os.path.dirname(pipeline_module.__file__)
bypasses = []
for path in sorted(
glob.glob(os.path.join(hook_dir, "**", "*.py"), recursive=True)
):
if os.path.basename(path) in ("pipeline.py", "resolution_hooks.py"):
continue
tree = ast.parse(open(path).read())
for node in ast.walk(tree):
if (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Name)
and node.func.id in _OVERRIDABLE_HOOKS
):
bypasses.append(
f"{os.path.relpath(path, hook_dir)}:{node.lineno} "
f"calls {node.func.id}(...) directly"
)
self.assertEqual(bypasses, [])
class TestRunHook(_IsolatedRegistry):
"""``run_hook`` reads the name off the function it is handed
(``builtin.__name__``), so every ``builtin`` stand-in below that needs to
line up with a registration is a real named function called
``handle_cuda_graph_config`` -- a lambda's ``__name__`` is ``'<lambda>'``
and would silently miss the registry entirely."""
def test_nothing_registered_runs_the_builtin_directly(self):
calls = []
run_hook(calls.append, "sa")
self.assertEqual(calls, ["sa"])
def test_an_override_that_calls_previous_wraps_the_builtin(self):
order = []
@register_resolution_hook("handle_cuda_graph_config")
def _wraps(server_args, previous):
order.append(("before", server_args))
previous(server_args)
order.append(("after", server_args))
def handle_cuda_graph_config(server_args):
order.append(("builtin", server_args))
run_hook(handle_cuda_graph_config, "sa")
self.assertEqual(
order,
[("before", "sa"), ("builtin", "sa"), ("after", "sa")],
)
def test_an_override_that_never_calls_previous_replaces_the_builtin(self):
builtin_ran = []
@register_resolution_hook("handle_cuda_graph_config")
def _replaces(server_args, previous):
pass # deliberately does not call `previous`
def handle_cuda_graph_config(server_args):
builtin_ran.append(server_args)
run_hook(handle_cuda_graph_config, "sa")
self.assertEqual(builtin_ran, [], "the replaced builtin must not have run")
def test_two_registrants_compose_last_registered_outermost(self):
order = []
@register_resolution_hook("handle_cuda_graph_config")
def _first(server_args, previous):
order.append("first-before")
previous(server_args)
order.append("first-after")
@register_resolution_hook("handle_cuda_graph_config")
def _second(server_args, previous):
order.append("second-before")
previous(server_args)
order.append("second-after")
def handle_cuda_graph_config(server_args):
order.append("builtin")
run_hook(handle_cuda_graph_config, "sa")
self.assertEqual(
order,
[
"second-before", # last registered runs first (outermost)
"first-before",
"builtin",
"first-after",
"second-after",
],
)
class TestEndToEnd(_IsolatedRegistry):
"""The proof that matters: a real `resolve_once()` picks up the override,
at the same position the built-in occupied, without disturbing the
neighboring steps documented at that call site."""
def setUp(self):
super().setUp()
self._config_dir = tempfile.mkdtemp(prefix="resolution_hook_registry_")
self.addCleanup(shutil.rmtree, self._config_dir, ignore_errors=True)
with open(os.path.join(self._config_dir, "config.json"), "w") as handle:
json.dump(_MINI_CONFIG, handle)
def test_an_override_reaches_resolution_result(self):
@register_resolution_hook("handle_cuda_graph_config")
def _mark_it(server_args, previous):
previous(server_args)
from sglang.srt.arg_groups.overrides import declare_resolution
declare_resolution(server_args, "test_plugin", random_seed=999)
server_args._test_plugin_ran = True
sa = ServerArgs(model_path=self._config_dir, device="cuda")
sa.resolve_once()
self.assertTrue(getattr(sa, "_test_plugin_ran", False))
self.assertEqual(resolution_result(sa, "random_seed"), 999)
# And the neighboring step (must run right after, per the comment at
# the call site) still ran and still saw a real config to chunk.
self.assertIsNotNone(
resolution_result(sa, "chunked_prefill_size"),
"apply_glm5_chunked_prefill_default's neighbor did not run",
)
def test_with_nothing_registered_resolution_is_unchanged(self):
baseline = ServerArgs(model_path=self._config_dir, device="cuda")
baseline.resolve_once()
sa = ServerArgs(model_path=self._config_dir, device="cuda")
sa.resolve_once()
self.assertEqual(
resolution_result(sa, "chunked_prefill_size"),
resolution_result(baseline, "chunked_prefill_size"),
)
self.assertEqual(
resolution_result(sa, "cuda_graph_config"),
resolution_result(baseline, "cuda_graph_config"),
)
if __name__ == "__main__":
unittest.main()
@@ -2153,11 +2153,22 @@ class TestPipelineParallelPrefillCudaGraphPolicy(CustomTestCase):
args._cuda_graph_config_locked = {(Phase.PREFILL, "backend")} | (
{(Phase.PREFILL, "max_bs")} if max_bs is not None else set()
)
with patch(
"sglang.srt.arg_groups.memory_hook.use_mla_backend",
return_value=False,
with (
patch(
"sglang.srt.arg_groups.memory_hook.use_mla_backend",
return_value=False,
),
patch(
# `handle_gpu_memory_settings` computes `gpu_mem` itself
# now (`get_device_memory_capacity(cfg.device)`),
# imported at module scope into `memory_hook` -- patch
# the name where it is looked up, not its origin
# module.
"sglang.srt.arg_groups.memory_hook.get_device_memory_capacity",
return_value=None,
),
):
handle_gpu_memory_settings(args, gpu_mem=None)
handle_gpu_memory_settings(args)
prefill = resolution_result(args, "cuda_graph_config").prefill
self.assertEqual((prefill.max_bs, prefill.bs[-1]), (expected, expected))