[Spec] Remove the ServerArgs clone + global save/restore hack from DFlashWorkerV2 (#29995)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)."
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user