[refactor] Move model-capability adjustments into the resolution pipeline (#30299)
This commit is contained in:
@@ -175,12 +175,15 @@ def refine_server_args(server_args: ServerArgs, compile_args: CompileArgs):
|
||||
# legacy disable_cuda_graph field, so flip both phases directly.
|
||||
server_args.cuda_graph_config[Phase.DECODE].backend = Backend.DISABLED
|
||||
server_args.cuda_graph_config[Phase.PREFILL].backend = Backend.DISABLED
|
||||
server_args.enable_torch_compile = False
|
||||
print(f"Disable CUDA Graph and Torch Compile to save time...")
|
||||
|
||||
# Set watchdog timeout to compile_args.timeout because compilation will take a long time
|
||||
server_args.watchdog_timeout = compile_args.timeout
|
||||
server_args.warmups = "compile-deep-gemm"
|
||||
# Watchdog timeout follows compile_args.timeout because compilation takes long.
|
||||
server_args.override(
|
||||
"compile_deep_gemm.refine_server_args",
|
||||
enable_torch_compile=False,
|
||||
watchdog_timeout=compile_args.timeout,
|
||||
warmups="compile-deep-gemm",
|
||||
)
|
||||
|
||||
|
||||
def run_compile(server_args: ServerArgs, compile_args: CompileArgs):
|
||||
|
||||
@@ -390,7 +390,7 @@ class Runtime:
|
||||
for port in range(self.server_args.port, 40000):
|
||||
if is_port_available(port):
|
||||
break
|
||||
self.server_args.port = port
|
||||
self.server_args.override("runtime_endpoint.port_alloc", port=port)
|
||||
|
||||
self.url = self.server_args.url()
|
||||
self.generate_url = self.url + "/generate"
|
||||
|
||||
@@ -197,8 +197,18 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None:
|
||||
stash.append(entry)
|
||||
validate_declarations(server_args, [entry])
|
||||
if getattr(server_args, "_declarations_materialized", False):
|
||||
for field, value in declared.items():
|
||||
setattr(server_args, field, value)
|
||||
_apply_fields(server_args, declared)
|
||||
|
||||
|
||||
def _apply_fields(server_args: Any, fields: Dict[str, Any]) -> None:
|
||||
"""Write fields on behalf of the pipeline (bypasses the strict bare-
|
||||
assignment guard that protects post-resolution mutation)."""
|
||||
object.__setattr__(server_args, "_in_override", True)
|
||||
try:
|
||||
for field, value in fields.items():
|
||||
setattr(server_args, field, value)
|
||||
finally:
|
||||
object.__setattr__(server_args, "_in_override", False)
|
||||
|
||||
|
||||
def materialize_declarations(server_args: Any) -> None:
|
||||
@@ -257,8 +267,7 @@ def declare_load_time_override(source: str, declared: Dict[str, Any]) -> None:
|
||||
ctx = get_context()
|
||||
entry = (source, dict(declared))
|
||||
validate_declarations(ctx.server_args, [entry])
|
||||
for field, value in declared.items():
|
||||
setattr(ctx.server_args, field, value)
|
||||
_apply_fields(ctx.server_args, declared)
|
||||
ctx.record_runtime_overrides([entry])
|
||||
|
||||
|
||||
@@ -2146,25 +2155,19 @@ def validate_declarations(
|
||||
)
|
||||
|
||||
|
||||
def refresh_declared_fields(server_args: Any, fields: Iterable[str]) -> None:
|
||||
"""Helper for legacy code that overwrites a resolved field AFTER
|
||||
materialization (e.g. ``ModelRunner.model_specific_adjustment`` forcing
|
||||
``attention_backend`` for HRM-Text). Redeclares the live value so the
|
||||
publish parity holds and the flags tier materializes the adjusted end
|
||||
state.
|
||||
"""
|
||||
_missing = object()
|
||||
declarations = server_args._resolved_overrides
|
||||
for field in fields:
|
||||
effective = _missing
|
||||
for _source, decl in declarations:
|
||||
if field in decl:
|
||||
effective = decl[field]
|
||||
if effective is _missing:
|
||||
continue
|
||||
live = getattr(server_args, field)
|
||||
if effective != live:
|
||||
declarations.append((f"runtime_adjustment[{field}]", {field: live}))
|
||||
def _hrm_text_attention_force(view: Any) -> dict:
|
||||
"""HRM-Text's bidirectional prefix attention only works on the Triton
|
||||
backend. Invoked as the last attention declaration of the resolution
|
||||
(mirroring the legacy runner-side force, which ran after the whole
|
||||
pipeline)."""
|
||||
if view.attention_backend not in (None, "triton"):
|
||||
logger.warning(
|
||||
f"Overriding --attention-backend "
|
||||
f"{view.attention_backend!r} -> 'triton': only the "
|
||||
"Triton backend supports HRM-Text's bidirectional prefix "
|
||||
"attention."
|
||||
)
|
||||
return {"attention_backend": "triton"}
|
||||
|
||||
|
||||
def assert_flag_parity(
|
||||
|
||||
@@ -44,7 +44,9 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None:
|
||||
"with speculative decoding "
|
||||
f"(--speculative-algorithm {server_args.speculative_algorithm})"
|
||||
)
|
||||
if server_args.enable_dp_attention:
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
if resolved_view(server_args).enable_dp_attention:
|
||||
logger.warning(
|
||||
"EXPERIMENTAL: Decode radix cache with DP attention. "
|
||||
"Requires prefix-aware DP rank routing for optimal cache hits."
|
||||
|
||||
@@ -272,7 +272,7 @@ def create_grammar_backend(
|
||||
"Falling back to grammar_backend='none'. "
|
||||
"Structured outputs (JSON schema, regex, EBNF) will not be available."
|
||||
)
|
||||
server_args.grammar_backend = "none"
|
||||
server_args.override("grammar.import_fallback", grammar_backend="none")
|
||||
return None
|
||||
elif name == "llguidance":
|
||||
from sglang.srt.constrained.llguidance_backend import GuidanceBackend
|
||||
|
||||
@@ -1355,7 +1355,9 @@ async def update_weight_version(
|
||||
# since weight_version update is a simple operation that doesn't affect model weights
|
||||
try:
|
||||
# Update the weight version in server args (the single source of truth)
|
||||
_global_state.tokenizer_manager.server_args.weight_version = obj.new_version
|
||||
_global_state.tokenizer_manager.server_args.override(
|
||||
"http.update_weight_version", weight_version=obj.new_version
|
||||
)
|
||||
|
||||
return ORJSONResponse(
|
||||
{
|
||||
|
||||
@@ -188,6 +188,10 @@ class ToolStrictLevel(IntEnum):
|
||||
|
||||
class Envs:
|
||||
|
||||
# Raise on bare server_args field assignments after resolution; mutation
|
||||
# must go through ServerArgs.override() (enabled by the test harness).
|
||||
SGLANG_STRICT_CONFIG_MUTATION = EnvBool(False)
|
||||
|
||||
# Model & File Download
|
||||
SGLANG_USE_MODELSCOPE = EnvBool(False)
|
||||
# Controls weight-file ordering for load-time I/O optimization.
|
||||
|
||||
@@ -601,14 +601,19 @@ class TokenizerWorker(TokenizerManager):
|
||||
setproctitle.setproctitle(f"sglang::tokenizer_worker:{os.getpid()}")
|
||||
# prevent init prefill bootstrapserver again
|
||||
disaggregation_mode = server_args.disaggregation_mode
|
||||
server_args.disaggregation_mode = "null"
|
||||
server_args.override(
|
||||
"tokenizer_worker.suppress_bootstrap", disaggregation_mode="null"
|
||||
)
|
||||
super().__init__(server_args, port_args)
|
||||
|
||||
self.worker_id = os.getpid()
|
||||
self.tokenizer_ipc_name = port_args.tokenizer_ipc_name
|
||||
|
||||
# For PD disaggregtion
|
||||
self.server_args.disaggregation_mode = disaggregation_mode
|
||||
self.server_args.override(
|
||||
"tokenizer_worker.restore_disaggregation_mode",
|
||||
disaggregation_mode=disaggregation_mode,
|
||||
)
|
||||
self.disaggregation_mode = DisaggregationMode(
|
||||
self.server_args.disaggregation_mode
|
||||
)
|
||||
|
||||
@@ -586,15 +586,6 @@ class Scheduler(
|
||||
if self.server_args.dllm_algorithm is not None
|
||||
else None
|
||||
)
|
||||
if self.dllm_config:
|
||||
if self.dllm_config.block_size < self.page_size:
|
||||
logger.warning(
|
||||
"WARNING: "
|
||||
f"The page size {self.page_size} should not be larger than dllm block size {self.dllm_config.block_size}."
|
||||
f"Page size now falls back to {self.dllm_config.block_size}"
|
||||
)
|
||||
self.page_size = self.dllm_config.block_size
|
||||
self.server_args.page_size = self.dllm_config.block_size
|
||||
|
||||
def init_metrics_collector(
|
||||
self, tp_rank: int, pp_rank: int, dp_rank: Optional[int]
|
||||
@@ -806,8 +797,9 @@ class Scheduler(
|
||||
)
|
||||
|
||||
if self.server_args.speculative_draft_load_format is not None:
|
||||
self.server_args.load_format = (
|
||||
self.server_args.speculative_draft_load_format
|
||||
self.server_args.override(
|
||||
"scheduler.draft_load_format",
|
||||
load_format=self.server_args.speculative_draft_load_format,
|
||||
)
|
||||
logger.info(
|
||||
f"Using draft model load_format: '{self.server_args.speculative_draft_load_format}'"
|
||||
@@ -910,8 +902,11 @@ class Scheduler(
|
||||
min_free_slots=min_free_slots
|
||||
)
|
||||
if not get_global_server_args().pp_max_micro_batch_size:
|
||||
get_global_server_args().pp_max_micro_batch_size = max(
|
||||
self.max_running_requests // self.ps.pp_size, 1
|
||||
get_global_server_args().override(
|
||||
"scheduler.pp_max_micro_batch_size_default",
|
||||
pp_max_micro_batch_size=max(
|
||||
self.max_running_requests // self.ps.pp_size, 1
|
||||
),
|
||||
)
|
||||
|
||||
self.tp_group = get_tp_group()
|
||||
@@ -3683,17 +3678,20 @@ class Scheduler(
|
||||
return AttachHiCacheStorageReqOutput(success=False, message=str(e))
|
||||
if ok:
|
||||
self.enable_hicache_storage = True
|
||||
self.server_args.hicache_storage_backend = recv_req.hicache_storage_backend
|
||||
hicache_fields = {
|
||||
"hicache_storage_backend": recv_req.hicache_storage_backend
|
||||
}
|
||||
if recv_req.hicache_storage_backend_extra_config_json is not None:
|
||||
self.server_args.hicache_storage_backend_extra_config = (
|
||||
hicache_fields["hicache_storage_backend_extra_config"] = (
|
||||
recv_req.hicache_storage_backend_extra_config_json
|
||||
)
|
||||
if recv_req.hicache_storage_prefetch_policy is not None:
|
||||
self.server_args.hicache_storage_prefetch_policy = (
|
||||
hicache_fields["hicache_storage_prefetch_policy"] = (
|
||||
recv_req.hicache_storage_prefetch_policy
|
||||
)
|
||||
if recv_req.hicache_write_policy is not None:
|
||||
self.server_args.hicache_write_policy = recv_req.hicache_write_policy
|
||||
hicache_fields["hicache_write_policy"] = recv_req.hicache_write_policy
|
||||
self.server_args.override("scheduler.attach_hicache", **hicache_fields)
|
||||
logger.info(
|
||||
f"Attached HiCache storage backend: {recv_req.hicache_storage_backend}"
|
||||
)
|
||||
@@ -3734,8 +3732,11 @@ class Scheduler(
|
||||
if ok or (not self.enable_hicache_storage):
|
||||
# Treat "already disabled / nothing to do" as success for idempotence.
|
||||
self.enable_hicache_storage = False
|
||||
self.server_args.hicache_storage_backend = None
|
||||
self.server_args.hicache_storage_backend_extra_config = None
|
||||
self.server_args.override(
|
||||
"scheduler.detach_hicache",
|
||||
hicache_storage_backend=None,
|
||||
hicache_storage_backend_extra_config=None,
|
||||
)
|
||||
logger.info("Detached HiCache storage backend.")
|
||||
return DetachHiCacheStorageReqOutput(
|
||||
success=True, message=msg or "HiCache storage backend is detached."
|
||||
|
||||
@@ -215,7 +215,10 @@ class SchedulerMetricsReporter:
|
||||
if base_endpoint is None:
|
||||
ipc_path = tempfile.NamedTemporaryFile(delete=False).name
|
||||
base_endpoint = f"ipc://{ipc_path}"
|
||||
self.scheduler.server_args.forward_pass_metrics_ipc_name = base_endpoint
|
||||
self.scheduler.server_args.override(
|
||||
"metrics_reporter.ipc_endpoint",
|
||||
forward_pass_metrics_ipc_name=base_endpoint,
|
||||
)
|
||||
endpoint = f"{base_endpoint}.{self.scheduler._fpm_dp_rank}"
|
||||
self.scheduler._fpm_publisher = _FpmPublisherThread(
|
||||
endpoint,
|
||||
|
||||
@@ -292,17 +292,18 @@ class TokenizerControlMixin:
|
||||
# TODO: partial rollback if failed
|
||||
if all_success:
|
||||
# Keep tokenizer side server_info consistent with scheduler side.
|
||||
self.server_args.hicache_storage_backend = hicache_storage_backend
|
||||
hicache_fields = {"hicache_storage_backend": hicache_storage_backend}
|
||||
if hicache_storage_backend_extra_config_json is not None:
|
||||
self.server_args.hicache_storage_backend_extra_config = (
|
||||
hicache_fields["hicache_storage_backend_extra_config"] = (
|
||||
hicache_storage_backend_extra_config_json
|
||||
)
|
||||
if hicache_storage_prefetch_policy is not None:
|
||||
self.server_args.hicache_storage_prefetch_policy = (
|
||||
hicache_fields["hicache_storage_prefetch_policy"] = (
|
||||
hicache_storage_prefetch_policy
|
||||
)
|
||||
if hicache_write_policy is not None:
|
||||
self.server_args.hicache_write_policy = hicache_write_policy
|
||||
hicache_fields["hicache_write_policy"] = hicache_write_policy
|
||||
self.server_args.override("tokenizer.attach_hicache", **hicache_fields)
|
||||
return out
|
||||
|
||||
async def detach_hicache_storage(
|
||||
@@ -318,8 +319,11 @@ class TokenizerControlMixin:
|
||||
out = DetachHiCacheStorageReqOutput(success=all_success, message=all_message)
|
||||
# TODO: partial rollback if failed
|
||||
if all_success:
|
||||
self.server_args.hicache_storage_backend = None
|
||||
self.server_args.hicache_storage_backend_extra_config = None
|
||||
self.server_args.override(
|
||||
"tokenizer.detach_hicache",
|
||||
hicache_storage_backend=None,
|
||||
hicache_storage_backend_extra_config=None,
|
||||
)
|
||||
return out
|
||||
|
||||
async def start_profile(
|
||||
@@ -869,4 +873,6 @@ class TokenizerControlMixin:
|
||||
) -> None:
|
||||
"""Update weight version if provided."""
|
||||
if weight_version is not None:
|
||||
self.server_args.weight_version = weight_version
|
||||
self.server_args.override(
|
||||
"tokenizer.weight_version", weight_version=weight_version
|
||||
)
|
||||
|
||||
@@ -1749,8 +1749,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
|
||||
def _update_model_path_info(self, model_path: str, load_format: str):
|
||||
self.served_model_name = model_path
|
||||
self.server_args.model_path = model_path
|
||||
self.server_args.load_format = load_format
|
||||
self.server_args.override(
|
||||
"tokenizer.update_weights", model_path=model_path, load_format=load_format
|
||||
)
|
||||
self.model_path = model_path
|
||||
|
||||
async def _wait_for_model_update_from_disk(
|
||||
|
||||
@@ -101,7 +101,9 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
self._enable_metrics_flag = params.enable_metrics
|
||||
if server_args.hicache_io_backend == "direct":
|
||||
if server_args.hicache_mem_layout == "page_first":
|
||||
server_args.hicache_mem_layout = "page_first_direct"
|
||||
server_args.override(
|
||||
"hicache.mem_layout_force", hicache_mem_layout="page_first_direct"
|
||||
)
|
||||
logger.warning(
|
||||
"Page first layout is not supported with direct IO backend, "
|
||||
"switching to page first direct layout"
|
||||
|
||||
@@ -496,7 +496,9 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||
# Direct IO layout fixup (must happen before pool creation)
|
||||
if server_args.hicache_io_backend == "direct":
|
||||
if server_args.hicache_mem_layout == "page_first":
|
||||
server_args.hicache_mem_layout = "page_first_direct"
|
||||
server_args.override(
|
||||
"hicache.mem_layout_force", hicache_mem_layout="page_first_direct"
|
||||
)
|
||||
logger.warning(
|
||||
"Page first layout is not supported with direct IO backend, "
|
||||
"switching to page first direct layout"
|
||||
|
||||
@@ -184,8 +184,10 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.server_args import (
|
||||
from sglang.srt.server_args import ( # noqa: F401 (re-export)
|
||||
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
|
||||
ServerArgs,
|
||||
add_chunked_prefix_cache_attention_backend,
|
||||
get_global_server_args,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
@@ -271,16 +273,6 @@ MLA_ATTENTION_BACKENDS = [
|
||||
"intel_xpu",
|
||||
]
|
||||
|
||||
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS = [
|
||||
"flashinfer",
|
||||
"fa3",
|
||||
"fa4",
|
||||
"flashmla",
|
||||
"cutedsl_mla",
|
||||
"cutlass_mla",
|
||||
"trtllm_mla",
|
||||
"tokenspeed_mla",
|
||||
]
|
||||
|
||||
TORCH_DTYPE_TO_KV_CACHE_STR = {
|
||||
torch.float8_e4m3fn: "fp8_e4m3",
|
||||
@@ -296,14 +288,6 @@ def add_mla_attention_backend(backend_name):
|
||||
logger.info(f"Added {backend_name} to MLA_ATTENTION_BACKENDS.")
|
||||
|
||||
|
||||
def add_chunked_prefix_cache_attention_backend(backend_name):
|
||||
if backend_name not in CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS:
|
||||
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS.append(backend_name)
|
||||
logger.info(
|
||||
f"Added {backend_name} to CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS."
|
||||
)
|
||||
|
||||
|
||||
# Detect stragger ranks in model loading
|
||||
UNBALANCED_MODEL_LOADING_TIMEOUT_S = 480 # leave more time for post data processing
|
||||
|
||||
@@ -528,8 +512,24 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if server_args.show_time_cost:
|
||||
enable_show_time_cost()
|
||||
|
||||
# Model-specific adjustment
|
||||
self.model_specific_adjustment()
|
||||
# Chunked prefix caching requires an MLA model on a backend whose
|
||||
# kernels read that layout. This is a load-time gate, not a
|
||||
# resolution-time one: out-of-tree platforms register their supported
|
||||
# backends in init_backend(), which runs when this module is imported
|
||||
# — after ServerArgs.__post_init__. Target runner only: a draft
|
||||
# model's (often non-MLA) config must not flip the shared setting.
|
||||
if not self.is_draft_worker and (
|
||||
not self.use_mla_backend
|
||||
or server_args.attention_backend
|
||||
not in CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS
|
||||
):
|
||||
if not server_args.disable_chunked_prefix_cache:
|
||||
server_args.override(
|
||||
"model_runner.chunked_prefix_cache_gate",
|
||||
disable_chunked_prefix_cache=True,
|
||||
)
|
||||
if not self.is_draft_worker and not server_args.disable_chunked_prefix_cache:
|
||||
logger.info("Chunked prefix cache is turned on.")
|
||||
|
||||
# Set the global server_args in the scheduler process (target worker
|
||||
# only, so a draft init cannot clobber target-derived global state).
|
||||
@@ -1128,70 +1128,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
f"Failed to register transfer engine info for tp_rank={self.tp_rank}: {e}"
|
||||
)
|
||||
|
||||
def model_specific_adjustment(self):
|
||||
if self.is_draft_worker:
|
||||
return
|
||||
|
||||
server_args = self.server_args
|
||||
|
||||
# HRM-Text needs bidirectional prompt attention (prefill), which only the
|
||||
# Triton backend honors and only with cuda graph / chunked prefill off
|
||||
# (TritonAttnBackend.allow_bidirectional_attention_in_extend). Radix cache
|
||||
# is also unsafe: the recurrent forward writes direction-dependent KV
|
||||
# across many slots.
|
||||
hf_config = self.model_config.hf_config
|
||||
is_hrm_text = getattr(
|
||||
hf_config, "model_type", None
|
||||
) == "hrm_text" or "HrmTextForCausalLM" in getattr(
|
||||
hf_config, "architectures", []
|
||||
)
|
||||
# prefix_lm defaults to True upstream; defaulting False would skip the
|
||||
# bidirectional-attention forcing and silently produce junk output.
|
||||
is_prefix_lm_recurrent = is_hrm_text and getattr(hf_config, "prefix_lm", True)
|
||||
if is_prefix_lm_recurrent:
|
||||
if server_args.attention_backend not in (None, "triton"):
|
||||
logger.warning(
|
||||
f"Overriding --attention-backend "
|
||||
f"{server_args.attention_backend!r} -> 'triton': only the "
|
||||
"Triton backend supports HRM-Text's bidirectional prefix "
|
||||
"attention."
|
||||
)
|
||||
server_args.attention_backend = "triton"
|
||||
server_args.chunked_prefill_size = -1
|
||||
server_args.disable_radix_cache = True
|
||||
server_args.disable_cuda_graph = True
|
||||
logger.warning(
|
||||
"HRM-Text (prefix_lm) detected: forcing --attention-backend "
|
||||
"triton, --chunked-prefill-size -1, --disable-radix-cache, and "
|
||||
"--disable-cuda-graph for correctness of the bidirectional "
|
||||
"prompt attention."
|
||||
)
|
||||
|
||||
if self.is_multimodal:
|
||||
if not self.is_multimodal_chunked_prefill_supported:
|
||||
server_args.chunked_prefill_size = -1
|
||||
logger.info(
|
||||
f"Automatically turn off --chunked-prefill-size as it is not supported for "
|
||||
f"{self.model_config.hf_config.model_type}"
|
||||
)
|
||||
|
||||
if (
|
||||
not self.use_mla_backend
|
||||
or server_args.attention_backend
|
||||
not in CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS
|
||||
):
|
||||
server_args.disable_chunked_prefix_cache = True
|
||||
|
||||
if not server_args.disable_chunked_prefix_cache:
|
||||
log_info_on_rank0(logger, "Chunked prefix cache is turned on.")
|
||||
|
||||
# The imperative adjustments above may overwrite fields the resolution passes
|
||||
# already declared (HRM-Text forces attention_backend); redeclare the
|
||||
# adjusted values so publish parity holds.
|
||||
from sglang.srt.arg_groups.overrides import refresh_declared_fields
|
||||
|
||||
refresh_declared_fields(server_args, ("attention_backend",))
|
||||
|
||||
def check_quantized_moe_compatibility(self):
|
||||
if (
|
||||
quantization_config := getattr(
|
||||
@@ -1436,7 +1372,13 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
logger.info(
|
||||
"Compute capability below sm80. Use float16 due to lack of bfloat16 support."
|
||||
)
|
||||
self.server_args.dtype = "float16"
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
declare_load_time_override,
|
||||
)
|
||||
|
||||
declare_load_time_override(
|
||||
"ModelRunner._sm80_dtype_fallback", {"dtype": "float16"}
|
||||
)
|
||||
self.model_config.dtype = torch.float16
|
||||
if torch.cuda.get_device_capability()[1] < 5:
|
||||
raise RuntimeError("SGLang only supports sm75 and above.")
|
||||
@@ -1894,8 +1836,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
return False, message
|
||||
|
||||
self.model = model
|
||||
self.server_args.model_path = model_path
|
||||
self.server_args.load_format = load_format
|
||||
self.server_args.override(
|
||||
"model_runner.update_weights",
|
||||
model_path=model_path,
|
||||
load_format=load_format,
|
||||
)
|
||||
self.load_config = load_config
|
||||
|
||||
if recapture_cuda_graph and (
|
||||
@@ -2443,7 +2388,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
{"kv_cache_dtype": resolved},
|
||||
)
|
||||
else:
|
||||
self.server_args.kv_cache_dtype = resolved
|
||||
self.server_args.override(
|
||||
"ModelRunner.configure_kv_cache_dtype", kv_cache_dtype=resolved
|
||||
)
|
||||
|
||||
def configure_kv_cache_dtype(self):
|
||||
if self.server_args.kv_cache_dtype == "auto":
|
||||
|
||||
@@ -159,8 +159,10 @@ class ModelRunnerKVCacheMixin:
|
||||
|
||||
if server_args.max_mamba_cache_size is not None:
|
||||
# Use explicitly set max_mamba_cache_size
|
||||
server_args.max_mamba_cache_size = server_args.max_mamba_cache_size // (
|
||||
server_args.dp_size if server_args.enable_dp_attention else 1
|
||||
server_args.override(
|
||||
"mamba_pool.per_dp_shard",
|
||||
max_mamba_cache_size=server_args.max_mamba_cache_size
|
||||
// (server_args.dp_size if server_args.enable_dp_attention else 1),
|
||||
)
|
||||
# Reserve intermediate memory based on capped max_num_reqs
|
||||
if has_spec_dec:
|
||||
@@ -181,8 +183,10 @@ class ModelRunnerKVCacheMixin:
|
||||
and server_args.max_running_requests is not None
|
||||
):
|
||||
# Use explicitly set max_running_requests when radix cache is disabled
|
||||
server_args.max_mamba_cache_size = server_args.max_running_requests // (
|
||||
server_args.dp_size if server_args.enable_dp_attention else 1
|
||||
server_args.override(
|
||||
"mamba_pool.from_max_running_requests",
|
||||
max_mamba_cache_size=server_args.max_running_requests
|
||||
// (server_args.dp_size if server_args.enable_dp_attention else 1),
|
||||
)
|
||||
# Reserve intermediate memory based on capped max_num_reqs
|
||||
if has_spec_dec:
|
||||
@@ -213,8 +217,11 @@ class ModelRunnerKVCacheMixin:
|
||||
ratio = self._calculate_mamba_ratio()
|
||||
D = server_args.speculative_num_draft_tokens
|
||||
# Joint solve: main_state + intermediate = mamba_budget
|
||||
server_args.max_mamba_cache_size = int(
|
||||
mamba_budget_bytes // (per_req * (1 + D / ratio))
|
||||
server_args.override(
|
||||
"mamba_pool.memory_budget_spec",
|
||||
max_mamba_cache_size=int(
|
||||
mamba_budget_bytes // (per_req * (1 + D / ratio))
|
||||
),
|
||||
)
|
||||
# Intermediate memory is included in mamba_budget, subtract it
|
||||
# so the return value only has main_state subtracted from total
|
||||
@@ -226,7 +233,10 @@ class ModelRunnerKVCacheMixin:
|
||||
intermediate_size = per_req * capped_reqs * D
|
||||
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
|
||||
else:
|
||||
server_args.max_mamba_cache_size = int(mamba_budget_bytes // per_req)
|
||||
server_args.override(
|
||||
"mamba_pool.memory_budget",
|
||||
max_mamba_cache_size=int(mamba_budget_bytes // per_req),
|
||||
)
|
||||
|
||||
# Validate: max_mamba_cache_size must be positive after memory allocation.
|
||||
# A non-positive value means GPU memory is insufficient for the requested
|
||||
|
||||
@@ -234,7 +234,7 @@ class RayEngine(Engine):
|
||||
if "log_level" not in kwargs:
|
||||
kwargs["log_level"] = "error"
|
||||
server_args = ServerArgs(**kwargs)
|
||||
server_args.placement_group = placement_group
|
||||
server_args.override("ray.placement_group", placement_group=placement_group)
|
||||
super().__init__(server_args=server_args)
|
||||
|
||||
def shutdown(self):
|
||||
@@ -463,7 +463,9 @@ class RayEngine(Engine):
|
||||
)
|
||||
# dataclasses.replace only copies declared fields; placement_group is
|
||||
# a dynamic attribute that must be manually appended after the rebuild.
|
||||
dp_server_args.placement_group = server_args.placement_group
|
||||
dp_server_args.override(
|
||||
"ray.placement_group", placement_group=server_args.placement_group
|
||||
)
|
||||
|
||||
# Create the DP controller in-process. This blocks until all actors
|
||||
# are initialized and their event loops have started.
|
||||
|
||||
@@ -44,7 +44,7 @@ def launch_server(
|
||||
if execute_warmup_func is None:
|
||||
execute_warmup_func = _execute_server_warmup
|
||||
|
||||
server_args.placement_group = None
|
||||
server_args.override("ray.http_server.clear_placement_group", placement_group=None)
|
||||
|
||||
(
|
||||
tokenizer_manager,
|
||||
|
||||
@@ -125,6 +125,30 @@ LOAD_FORMAT_CHOICES = [
|
||||
|
||||
# TODO: this list should likely contain only methods that support online quantization, or that support using custom quantization classes compatible with a given `quant_method` in config.json.
|
||||
# Some of the choices here do NOT support online quantization.
|
||||
# Attention backends whose kernels read the chunked prefix-cache layout.
|
||||
# Out-of-tree platforms may extend this list (via
|
||||
# add_chunked_prefix_cache_attention_backend) before ServerArgs construction;
|
||||
# the chunked-prefix gate is evaluated during resolution.
|
||||
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS = [
|
||||
"flashinfer",
|
||||
"fa3",
|
||||
"fa4",
|
||||
"flashmla",
|
||||
"cutedsl_mla",
|
||||
"cutlass_mla",
|
||||
"trtllm_mla",
|
||||
"tokenspeed_mla",
|
||||
]
|
||||
|
||||
|
||||
def add_chunked_prefix_cache_attention_backend(backend_name):
|
||||
if backend_name not in CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS:
|
||||
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS.append(backend_name)
|
||||
logger.info(
|
||||
f"Added {backend_name} to CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS."
|
||||
)
|
||||
|
||||
|
||||
QUANTIZATION_CHOICES = [
|
||||
"awq",
|
||||
"fp8", # MOE + linear online quantization.
|
||||
@@ -2845,6 +2869,10 @@ class ServerArgs:
|
||||
# Handle any other necessary validations.
|
||||
self._handle_other_validations()
|
||||
|
||||
# Model-capability adjustments that legacy code applied at model-load
|
||||
# time; last declarations of the resolution, mirroring that order.
|
||||
self._handle_model_capability_adjustments()
|
||||
|
||||
# End of resolution: apply the accumulated declarations onto the
|
||||
# fields once (gate order). From here on server_args carries the
|
||||
# resolved configuration — post-init readers, in any process, read
|
||||
@@ -2853,6 +2881,54 @@ class ServerArgs:
|
||||
|
||||
materialize_declarations(self)
|
||||
|
||||
def _handle_model_capability_adjustments(self):
|
||||
if parse_connector_type(self.model_path) == ConnectorType.INSTANCE:
|
||||
return
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
_hrm_text_attention_force,
|
||||
run_post_process_pass,
|
||||
)
|
||||
|
||||
model_config = self.get_model_config()
|
||||
hf_config = model_config.hf_config
|
||||
|
||||
# HRM-Text needs bidirectional prompt attention (prefill), which only
|
||||
# the Triton backend honors at the kernel level. Radix/prefix reuse is
|
||||
# also unsafe: the recurrent forward writes direction-dependent KV
|
||||
# across many slots.
|
||||
is_hrm_text = getattr(
|
||||
hf_config, "model_type", None
|
||||
) == "hrm_text" or "HrmTextForCausalLM" in getattr(
|
||||
hf_config, "architectures", []
|
||||
)
|
||||
# prefix_lm defaults to True upstream; defaulting False would skip the
|
||||
# bidirectional-attention forcing and silently produce junk output.
|
||||
if is_hrm_text and getattr(hf_config, "prefix_lm", True):
|
||||
run_post_process_pass(self, _hrm_text_attention_force)
|
||||
self.chunked_prefill_size = -1
|
||||
self.disable_radix_cache = True
|
||||
self.disable_cuda_graph = True
|
||||
# cuda_graph_config was already parsed from the legacy boolean, so
|
||||
# flipping the boolean alone would not stop graph capture.
|
||||
self.cuda_graph_config.decode.backend = Backend.DISABLED
|
||||
self.cuda_graph_config.prefill.backend = Backend.DISABLED
|
||||
logger.warning(
|
||||
"HRM-Text (prefix_lm) detected: forcing --attention-backend "
|
||||
"triton, --chunked-prefill-size -1, --disable-radix-cache, and "
|
||||
"--disable-cuda-graph for correctness of the bidirectional "
|
||||
"prompt attention."
|
||||
)
|
||||
|
||||
if (
|
||||
model_config.is_multimodal
|
||||
and not model_config.is_multimodal_chunked_prefill_supported
|
||||
):
|
||||
self.chunked_prefill_size = -1
|
||||
logger.info(
|
||||
f"Automatically turn off --chunked-prefill-size as it is not supported for "
|
||||
f"{hf_config.model_type}"
|
||||
)
|
||||
|
||||
def _handle_model_source_paths(self):
|
||||
"""Resolve model/tokenizer paths backed by remote object stores."""
|
||||
if is_runai_obj_uri(self.model_path):
|
||||
@@ -3938,7 +4014,7 @@ class ServerArgs:
|
||||
|
||||
run_post_process_pass(self, _dsa_kv_cache_dtype_default)
|
||||
|
||||
def _set_default_dsa_backends(self, kv_cache_dtype: str, major: int) -> None:
|
||||
def _set_default_dsa_backends(self, major: int) -> None:
|
||||
# Moved to the resolution pipeline (arg_groups/overrides.py:
|
||||
# _dsa_split_backend_resolution), invoked here at its legacy slot.
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
@@ -4069,7 +4145,7 @@ class ServerArgs:
|
||||
self._set_default_dsa_kv_cache_dtype(
|
||||
major, resolved_view(self).quantization
|
||||
)
|
||||
self._set_default_dsa_backends(self.kv_cache_dtype, major)
|
||||
self._set_default_dsa_backends(major)
|
||||
|
||||
if self.enable_prefill_cp:
|
||||
assert (
|
||||
@@ -6607,6 +6683,58 @@ class ServerArgs:
|
||||
|
||||
return resolved_view(self)
|
||||
|
||||
def override(self, source: str, **fields) -> None:
|
||||
"""The single post-resolution mutation point.
|
||||
|
||||
After ``__post_init__`` the configuration is resolved; the audited
|
||||
runtime adjustments (load-resolved values, control-plane
|
||||
reconfiguration, deployment wiring) go through here instead of
|
||||
assigning fields. Whitelisted resolvable fields also join the
|
||||
declaration stash, so a republish resolves the same values;
|
||||
everything is recorded with its ``source`` for provenance.
|
||||
"""
|
||||
from sglang.srt.arg_groups.arg_utils import resolvable_fields
|
||||
|
||||
whitelist = resolvable_fields(type(self))
|
||||
declared = {k: v for k, v in fields.items() if k in whitelist}
|
||||
rest = {k: v for k, v in fields.items() if k not in whitelist}
|
||||
if declared:
|
||||
stash = getattr(self, "_resolved_overrides", None)
|
||||
if stash is None:
|
||||
stash = []
|
||||
object.__setattr__(self, "_resolved_overrides", stash)
|
||||
stash.append((source, dict(declared)))
|
||||
if rest:
|
||||
log = getattr(self, "_runtime_mutations", None)
|
||||
if log is None:
|
||||
log = []
|
||||
object.__setattr__(self, "_runtime_mutations", log)
|
||||
log.append((source, dict(rest)))
|
||||
object.__setattr__(self, "_in_override", True)
|
||||
try:
|
||||
for field, value in fields.items():
|
||||
setattr(self, field, value)
|
||||
finally:
|
||||
object.__setattr__(self, "_in_override", False)
|
||||
|
||||
def __setattr__(self, name, value):
|
||||
# After materialization the fields are the resolved configuration:
|
||||
# under the strict test harness, a bare assignment outside
|
||||
# ServerArgs.override() (and the resolution pipeline itself) raises.
|
||||
if (
|
||||
not name.startswith("_")
|
||||
and getattr(self, "_declarations_materialized", False)
|
||||
and not getattr(self, "_in_override", False)
|
||||
):
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
if envs.SGLANG_STRICT_CONFIG_MUTATION.get():
|
||||
raise AttributeError(
|
||||
f"server_args.{name} assigned after resolution; use "
|
||||
"server_args.override(source, ...) instead."
|
||||
)
|
||||
object.__setattr__(self, name, value)
|
||||
|
||||
def _resolved_attention_backends(self):
|
||||
"""Mid-resolution (prefill, decode) backends: reads through the pass
|
||||
view so declared fields resolve from the declaration stash."""
|
||||
@@ -6901,7 +7029,7 @@ class ServerArgs:
|
||||
# Enable LoRA if any LoRA paths are provided for backward compatibility.
|
||||
if self.lora_paths:
|
||||
if self.enable_lora is None:
|
||||
self.enable_lora = True
|
||||
self.override("check_lora_server_args", enable_lora=True)
|
||||
logger.warning(
|
||||
"--enable-lora is set to True because --lora-paths is provided."
|
||||
)
|
||||
@@ -6912,7 +7040,9 @@ class ServerArgs:
|
||||
|
||||
if self.enable_lora:
|
||||
if self.enable_lora_overlap_loading is None:
|
||||
self.enable_lora_overlap_loading = False
|
||||
self.override(
|
||||
"check_lora_server_args", enable_lora_overlap_loading=False
|
||||
)
|
||||
|
||||
if self.enable_lora_overlap_loading:
|
||||
# TODO (glenliu21): use some sort of buffer with eviction instead of enforcing a limit
|
||||
@@ -6933,9 +7063,8 @@ class ServerArgs:
|
||||
|
||||
# Parse lora_paths
|
||||
if isinstance(self.lora_paths, list):
|
||||
lora_paths = self.lora_paths
|
||||
self.lora_paths = []
|
||||
for lora_path in lora_paths:
|
||||
parsed_lora_paths = []
|
||||
for lora_path in self.lora_paths:
|
||||
if isinstance(lora_path, str):
|
||||
if "=" in lora_path:
|
||||
name, path = lora_path.split("=", 1)
|
||||
@@ -6969,19 +7098,23 @@ class ServerArgs:
|
||||
f"Invalid type for item in --lora-paths list: {type(lora_path)}. "
|
||||
"Expected a string or a dictionary."
|
||||
)
|
||||
self.lora_paths.append(lora_ref)
|
||||
parsed_lora_paths.append(lora_ref)
|
||||
self.override("check_lora_server_args", lora_paths=parsed_lora_paths)
|
||||
elif isinstance(self.lora_paths, dict):
|
||||
self.lora_paths = [
|
||||
LoRARef(
|
||||
lora_id=LoRARef.deterministic_id(k, v),
|
||||
lora_name=k,
|
||||
lora_path=v,
|
||||
pinned=False,
|
||||
)
|
||||
for k, v in self.lora_paths.items()
|
||||
]
|
||||
self.override(
|
||||
"check_lora_server_args",
|
||||
lora_paths=[
|
||||
LoRARef(
|
||||
lora_id=LoRARef.deterministic_id(k, v),
|
||||
lora_name=k,
|
||||
lora_path=v,
|
||||
pinned=False,
|
||||
)
|
||||
for k, v in self.lora_paths.items()
|
||||
],
|
||||
)
|
||||
elif self.lora_paths is None:
|
||||
self.lora_paths = []
|
||||
self.override("check_lora_server_args", lora_paths=[])
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid type for --lora-paths: {type(self.lora_paths)}. "
|
||||
@@ -6991,7 +7124,10 @@ class ServerArgs:
|
||||
# Normalize target modules to a set; keep {"all"} as a sentinel
|
||||
# that gets resolved model-awarely in lora_manager.init_lora_shapes().
|
||||
if self.lora_target_modules:
|
||||
self.lora_target_modules = set(self.lora_target_modules)
|
||||
self.override(
|
||||
"check_lora_server_args",
|
||||
lora_target_modules=set(self.lora_target_modules),
|
||||
)
|
||||
if "all" in self.lora_target_modules:
|
||||
assert (
|
||||
len(self.lora_target_modules) == 1
|
||||
|
||||
@@ -320,8 +320,11 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
||||
self.hot_token_id = None
|
||||
elif self.server_args.speculative_token_map is not None:
|
||||
self.hot_token_id = load_token_map(self.server_args.speculative_token_map)
|
||||
self.server_args.json_model_override_args = (
|
||||
f'{{"hot_vocab_size": {len(self.hot_token_id)}}}'
|
||||
self.server_args.override(
|
||||
"eagle_worker.hot_token_map",
|
||||
json_model_override_args=(
|
||||
f'{{"hot_vocab_size": {len(self.hot_token_id)}}}'
|
||||
),
|
||||
)
|
||||
else:
|
||||
self.hot_token_id = None
|
||||
@@ -1060,7 +1063,10 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
||||
)
|
||||
|
||||
# Override the context length of the draft model to be the same as the target model.
|
||||
server_args.context_length = target_worker.model_runner.model_config.context_len
|
||||
server_args.override(
|
||||
"spec_worker.match_target_context_length",
|
||||
context_length=target_worker.model_runner.model_config.context_len,
|
||||
)
|
||||
|
||||
self._draft_worker = EagleDraftWorker(
|
||||
server_args,
|
||||
@@ -1462,9 +1468,10 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
||||
)
|
||||
|
||||
# Sync server_args
|
||||
self.server_args.speculative_num_steps = state.speculative_num_steps
|
||||
self.server_args.speculative_num_draft_tokens = (
|
||||
state.speculative_num_draft_tokens
|
||||
self.server_args.override(
|
||||
"adaptive_spec.restore",
|
||||
speculative_num_steps=state.speculative_num_steps,
|
||||
speculative_num_draft_tokens=state.speculative_num_draft_tokens,
|
||||
)
|
||||
|
||||
@contextlib.contextmanager
|
||||
@@ -1498,16 +1505,21 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
||||
self.speculative_num_draft_tokens = speculative_num_draft_tokens
|
||||
dw.speculative_num_steps = speculative_num_steps
|
||||
dw.speculative_num_draft_tokens = speculative_num_draft_tokens
|
||||
sa.speculative_num_steps = speculative_num_steps
|
||||
sa.speculative_num_draft_tokens = speculative_num_draft_tokens
|
||||
sa.override(
|
||||
"adaptive_spec.capture_override",
|
||||
speculative_num_steps=speculative_num_steps,
|
||||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||||
)
|
||||
if cuda_graph_bs is not None:
|
||||
sa.cuda_graph_bs_decode = cuda_graph_bs
|
||||
# BS-aware adaptive spec may prune cuda_graph_bs to an empty list
|
||||
# for steps that no BS range uses (e.g. step=1). Disable graph
|
||||
# capture for those steps; restore in finally so subsequent steps
|
||||
# are not affected.
|
||||
if not cuda_graph_bs:
|
||||
sa.disable_cuda_graph = True
|
||||
sa.override(
|
||||
"adaptive_spec.capture_override",
|
||||
cuda_graph_bs_decode=cuda_graph_bs,
|
||||
**({"disable_cuda_graph": True} if not cuda_graph_bs else {}),
|
||||
)
|
||||
dw._rebuild_topk1_chain_buffers()
|
||||
|
||||
try:
|
||||
@@ -1524,11 +1536,14 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
||||
dw.draft_runner.attn_backend,
|
||||
dw.cuda_graph_runner,
|
||||
dw.cuda_graph_runner_for_draft_extend,
|
||||
sa.speculative_num_steps,
|
||||
sa.speculative_num_draft_tokens,
|
||||
sa.cuda_graph_bs_decode,
|
||||
sa.disable_cuda_graph,
|
||||
) = backup
|
||||
) = backup[:10]
|
||||
sa.override(
|
||||
"adaptive_spec.capture_restore",
|
||||
speculative_num_steps=backup[10],
|
||||
speculative_num_draft_tokens=backup[11],
|
||||
cuda_graph_bs_decode=backup[12],
|
||||
disable_cuda_graph=backup[13],
|
||||
)
|
||||
dw._rebuild_topk1_chain_buffers()
|
||||
|
||||
def verify(self, batch: ScheduleBatch):
|
||||
|
||||
@@ -676,7 +676,10 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
|
||||
target_worker.get_memory_pool()
|
||||
)
|
||||
# Match the draft context length to the target (assistant reads target KV).
|
||||
server_args.context_length = target_worker.model_runner.model_config.context_len
|
||||
server_args.override(
|
||||
"spec_worker.match_target_context_length",
|
||||
context_length=target_worker.model_runner.model_config.context_len,
|
||||
)
|
||||
|
||||
self._draft_worker = FrozenKVMTPDraftWorker(
|
||||
server_args,
|
||||
|
||||
@@ -678,7 +678,10 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
||||
)
|
||||
|
||||
# Override the context length of the draft model to be the same as the target model.
|
||||
server_args.context_length = target_worker.model_runner.model_config.context_len
|
||||
server_args.override(
|
||||
"spec_worker.match_target_context_length",
|
||||
context_length=target_worker.model_runner.model_config.context_len,
|
||||
)
|
||||
|
||||
self._draft_worker = MultiLayerEagleDraftWorker(
|
||||
server_args,
|
||||
|
||||
@@ -181,7 +181,10 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
|
||||
)
|
||||
|
||||
# Override the context length of the draft model to be the same as the target model.
|
||||
server_args.context_length = target_worker.model_runner.model_config.context_len
|
||||
server_args.override(
|
||||
"spec_worker.match_target_context_length",
|
||||
context_length=target_worker.model_runner.model_config.context_len,
|
||||
)
|
||||
|
||||
# Create our custom draft worker that doesn't share embeddings/lm_head
|
||||
self._draft_worker = StandaloneDraftWorker(
|
||||
|
||||
@@ -390,7 +390,9 @@ def _build_frozen_kv_mtp_fixture(
|
||||
runner_batch_size=settings.capture_batch_size,
|
||||
)
|
||||
_configure_runner_for_eagle_draft(fixture.runner, case, settings)
|
||||
fixture.runner.server_args.speculative_algorithm = "FROZEN_KV_MTP"
|
||||
fixture.runner.server_args.override(
|
||||
"attention_unittest.frozen_kv_draft", speculative_algorithm="FROZEN_KV_MTP"
|
||||
)
|
||||
fixture.runner.spec_algorithm = SpeculativeAlgorithm.FROZEN_KV_MTP
|
||||
fixture.runner.draft_attn_backend = fixture.backend
|
||||
fixture.runner.attn_backend = fixture.backend
|
||||
|
||||
@@ -8,6 +8,10 @@ import inspect
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
||||
# Registered tests run with the strict config-mutation guard: bare
|
||||
# server_args assignments after resolution raise (use ServerArgs.override).
|
||||
os.environ.setdefault("SGLANG_STRICT_CONFIG_MUTATION", "1")
|
||||
import random
|
||||
import re
|
||||
import shlex
|
||||
|
||||
Reference in New Issue
Block a user