[Spec] Remove the ServerArgs clone + global save/restore hack from DFlashWorkerV2 (#29995)

This commit is contained in:
Cheng Wan
2026-07-02 23:37:04 -07:00
committed by GitHub
parent d364cd8ead
commit 76f7f7c006
6 changed files with 75 additions and 68 deletions
@@ -229,6 +229,8 @@ def _handle_dflash(server_args: ServerArgs) -> None:
f"window_size={server_args.speculative_draft_window_size}, block_size={draft_tokens}."
)
_resolve_dflash_draft_attention_backend(server_args)
if server_args.max_running_requests is None:
server_args.max_running_requests = 48
logger.warning(
@@ -242,6 +244,45 @@ def _handle_dflash(server_args: ServerArgs) -> None:
)
def _resolve_dflash_draft_attention_backend(server_args: ServerArgs) -> None:
"""Resolve `speculative_draft_attention_backend` to a final, supported value.
Consumed by ModelRunner's `is_draft_worker` override (one backend for all
draft modes).
"""
from sglang.srt.utils import is_hip
supported_draft_backends = ("flashinfer", "fa3", "fa4", "triton", "ascend")
# Use triton on ROCm (no FlashInfer), flashinfer on CUDA.
fallback_backend = "triton" if is_hip() else "flashinfer"
draft_backend = server_args.speculative_draft_attention_backend
if draft_backend is None:
draft_backend, _ = server_args.get_attention_backends()
if draft_backend is None:
draft_backend = fallback_backend
elif draft_backend == "trtllm_mha":
logger.warning(
"DFLASH draft worker does not support 'trtllm_mha' because the "
"draft path requires per-layer DFlash attention. Falling back to "
"'%s'.",
fallback_backend,
)
draft_backend = fallback_backend
elif draft_backend not in supported_draft_backends:
logger.warning(
"DFLASH draft worker only supports attention_backend in %s for now, "
"but got %r. Falling back to '%s'.",
supported_draft_backends,
draft_backend,
fallback_backend,
)
draft_backend = fallback_backend
# FIXME: avoid overriding server args directly; pass the resolved draft
# backend to the draft worker explicitly instead.
server_args.speculative_draft_attention_backend = draft_backend
def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None:
if server_args.max_running_requests is None:
server_args.max_running_requests = 48
+6 -1
View File
@@ -483,6 +483,7 @@ class ModelConfig:
model_path: str = None,
model_revision: str = None,
is_draft_model: bool = False,
context_length: Optional[int] = None,
**kwargs,
):
quantization = (
@@ -499,7 +500,11 @@ class ModelConfig:
model_path=model_path or server_args.model_path,
trust_remote_code=server_args.trust_remote_code,
revision=model_revision or server_args.revision,
context_length=server_args.context_length,
context_length=(
context_length
if context_length is not None
else server_args.context_length
),
model_override_args=server_args.json_model_override_args,
is_embedding=server_args.is_embedding,
enable_multimodal=server_args.enable_multimodal,
+8 -1
View File
@@ -241,6 +241,7 @@ class TpModelWorker(BaseTpWorker):
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
memory_pool_config: Optional[MemoryPoolConfig] = None,
is_multi_layer_eagle: bool = False,
context_length: Optional[int] = None,
):
# Parse args
self.server_args = server_args
@@ -261,6 +262,9 @@ class TpModelWorker(BaseTpWorker):
self.moe_dp_rank = moe_dp_rank
# Draft worker: target's resolved MemoryPoolConfig (forwarded to ModelRunner).
self.memory_pool_config = memory_pool_config
# Draft worker: target's effective context length; the draft runs at
# absolute target positions. None keeps server_args.context_length.
self.context_length = context_length
# MTP model runners
self.model_runner_list: List[ModelRunner] = []
@@ -273,7 +277,9 @@ class TpModelWorker(BaseTpWorker):
self._init_dllm_algorithm()
if server_args.skip_tokenizer_init:
if server_args.skip_tokenizer_init or self.is_draft_worker:
# A draft worker's tokenizer would only duplicate the target's:
# tokenizer_path always points at the target model.
self.tokenizer = self.processor = None
else:
if self.model_config.is_multimodal:
@@ -370,6 +376,7 @@ class TpModelWorker(BaseTpWorker):
else self.server_args.speculative_draft_model_revision
),
is_draft_model=self.is_draft_worker,
context_length=self.context_length,
)
def _init_model_runner(self):
@@ -528,12 +528,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
# Model-specific adjustment
self.model_specific_adjustment()
# Set the global server_args in the scheduler process
set_global_server_args_for_scheduler(server_args)
global_server_args = get_global_server_args()
# FIXME: hacky set `use_mla_backend`
global_server_args.use_mla_backend = self.use_mla_backend
# Set the global server_args in the scheduler process (target worker
# only, so a draft init cannot clobber target-derived global state).
if not self.is_draft_worker:
set_global_server_args_for_scheduler(server_args)
# FIXME: hacky set `use_mla_backend`
get_global_server_args().use_mla_backend = self.use_mla_backend
# Init OpenMP threads binding for CPU
if self.device == "cpu":
@@ -1128,6 +1128,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
)
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
@@ -1,6 +1,5 @@
import logging
import math
from copy import deepcopy
from typing import List, Optional
import torch
@@ -15,11 +14,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode,
compute_position,
)
from sglang.srt.server_args import (
ServerArgs,
get_global_server_args,
set_global_server_args_for_scheduler,
)
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
@@ -134,52 +129,8 @@ class DFlashWorkerV2(BaseSpecWorker):
self._logged_first_verify = False
# Draft runner (separate KV cache + attention backend).
draft_server_args = deepcopy(server_args)
draft_server_args.skip_tokenizer_init = True
draft_backend = draft_server_args.speculative_draft_attention_backend
supported_draft_backends = ("flashinfer", "fa3", "fa4", "triton", "ascend")
if draft_backend is None:
draft_backend, _ = draft_server_args.get_attention_backends()
if draft_backend is None:
# Use triton on ROCm (no FlashInfer), flashinfer on CUDA
import torch as _torch
draft_backend = "triton" if _torch.version.hip else "flashinfer"
elif draft_backend == "trtllm_mha":
import torch as _torch
_fb = "triton" if _torch.version.hip else "flashinfer"
logger.warning(
"DFLASH draft worker does not support 'trtllm_mha' because the "
"draft path requires per-layer DFlash attention. Falling back to "
"'%s'.",
_fb,
)
draft_backend = _fb
elif draft_backend not in supported_draft_backends:
import torch as _torch
_fb = "triton" if _torch.version.hip else "flashinfer"
logger.warning(
"DFLASH draft worker only supports attention_backend in %s for now, "
"but got %r. Falling back to '%s'.",
supported_draft_backends,
draft_backend,
_fb,
)
draft_backend = _fb
# Make the draft worker backend explicit and self-contained (no further overrides).
draft_server_args.speculative_draft_attention_backend = None
draft_server_args.prefill_attention_backend = None
draft_server_args.decode_attention_backend = None
draft_server_args.attention_backend = draft_backend
# Keep draft context length aligned with the target.
draft_server_args.context_length = (
target_worker.model_runner.model_config.context_len
)
saved_server_args = get_global_server_args()
self._draft_worker = TpModelWorker(
server_args=draft_server_args,
server_args=server_args,
gpu_id=gpu_id,
tp_rank=tp_rank,
moe_ep_rank=moe_ep_rank,
@@ -189,8 +140,8 @@ class DFlashWorkerV2(BaseSpecWorker):
dp_rank=dp_rank,
nccl_port=nccl_port,
is_draft_worker=True,
context_length=target_worker.model_runner.model_config.context_len,
)
set_global_server_args_for_scheduler(saved_server_args)
self.draft_model_runner = self._draft_worker.model_runner
self._draft_sampler = None
# Keep the same alias that other spec-v2 workers expose.
@@ -226,7 +177,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if self.tp_rank == 0:
logger.info(
"Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s",
getattr(draft_server_args, "attention_backend", None),
server_args.speculative_draft_attention_backend,
self.draft_model.__class__.__name__,
self.block_size,
self.draft_window_size,
+7 -7
View File
@@ -1,6 +1,6 @@
import logging
from sglang.srt.server_args import ServerArgs, get_global_server_args
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import is_blackwell, is_hip, is_musa, is_npu
logger = logging.getLogger(__name__)
@@ -118,7 +118,7 @@ class DraftBackendFactory:
return DeepseekSparseAttnBackend(self.draft_model_runner, skip_prefill=False)
def _create_flashinfer_decode_backend(self):
if not get_global_server_args().use_mla_backend:
if not self.draft_model_runner.use_mla_backend:
from sglang.srt.layers.attention.flashinfer_backend import (
FlashInferMultiStepDraftBackend,
)
@@ -193,7 +193,7 @@ class DraftBackendFactory:
)
def _create_trtllm_mla_decode_backend(self, backend: str = "trtllm-gen"):
if not get_global_server_args().use_mla_backend:
if not self.draft_model_runner.use_mla_backend:
raise ValueError(
"trtllm_mla backend requires MLA model (use_mla_backend=True)."
)
@@ -213,7 +213,7 @@ class DraftBackendFactory:
return self._create_trtllm_mla_decode_backend(backend="cute-dsl")
def _create_tokenspeed_mla_decode_backend(self):
if not get_global_server_args().use_mla_backend:
if not self.draft_model_runner.use_mla_backend:
raise ValueError(
"tokenspeed_mla backend requires MLA model (use_mla_backend=True)."
)
@@ -259,7 +259,7 @@ class DraftBackendFactory:
)
def _create_flashinfer_prefill_backend(self):
if not get_global_server_args().use_mla_backend:
if not self.draft_model_runner.use_mla_backend:
from sglang.srt.layers.attention.flashinfer_backend import (
FlashInferAttnBackend,
)
@@ -307,7 +307,7 @@ class DraftBackendFactory:
return TRTLLMHAAttnBackend(self.draft_model_runner, skip_prefill=False)
def _create_trtllm_mla_prefill_backend(self):
if not get_global_server_args().use_mla_backend:
if not self.draft_model_runner.use_mla_backend:
raise ValueError(
"trtllm_mla backend requires MLA model (use_mla_backend=True)."
)
@@ -317,7 +317,7 @@ class DraftBackendFactory:
return TRTLLMMLABackend(self.draft_model_runner, skip_prefill=False)
def _create_tokenspeed_mla_prefill_backend(self):
if not get_global_server_args().use_mla_backend:
if not self.draft_model_runner.use_mla_backend:
raise ValueError(
"tokenspeed_mla backend requires MLA model (use_mla_backend=True)."
)