[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}."
|
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:
|
if server_args.max_running_requests is None:
|
||||||
server_args.max_running_requests = 48
|
server_args.max_running_requests = 48
|
||||||
logger.warning(
|
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:
|
def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None:
|
||||||
if server_args.max_running_requests is None:
|
if server_args.max_running_requests is None:
|
||||||
server_args.max_running_requests = 48
|
server_args.max_running_requests = 48
|
||||||
|
|||||||
@@ -483,6 +483,7 @@ class ModelConfig:
|
|||||||
model_path: str = None,
|
model_path: str = None,
|
||||||
model_revision: str = None,
|
model_revision: str = None,
|
||||||
is_draft_model: bool = False,
|
is_draft_model: bool = False,
|
||||||
|
context_length: Optional[int] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
quantization = (
|
quantization = (
|
||||||
@@ -499,7 +500,11 @@ class ModelConfig:
|
|||||||
model_path=model_path or server_args.model_path,
|
model_path=model_path or server_args.model_path,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=model_revision or server_args.revision,
|
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,
|
model_override_args=server_args.json_model_override_args,
|
||||||
is_embedding=server_args.is_embedding,
|
is_embedding=server_args.is_embedding,
|
||||||
enable_multimodal=server_args.enable_multimodal,
|
enable_multimodal=server_args.enable_multimodal,
|
||||||
|
|||||||
@@ -241,6 +241,7 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
|
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
|
||||||
memory_pool_config: Optional[MemoryPoolConfig] = None,
|
memory_pool_config: Optional[MemoryPoolConfig] = None,
|
||||||
is_multi_layer_eagle: bool = False,
|
is_multi_layer_eagle: bool = False,
|
||||||
|
context_length: Optional[int] = None,
|
||||||
):
|
):
|
||||||
# Parse args
|
# Parse args
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
@@ -261,6 +262,9 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
self.moe_dp_rank = moe_dp_rank
|
self.moe_dp_rank = moe_dp_rank
|
||||||
# Draft worker: target's resolved MemoryPoolConfig (forwarded to ModelRunner).
|
# Draft worker: target's resolved MemoryPoolConfig (forwarded to ModelRunner).
|
||||||
self.memory_pool_config = memory_pool_config
|
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
|
# MTP model runners
|
||||||
self.model_runner_list: List[ModelRunner] = []
|
self.model_runner_list: List[ModelRunner] = []
|
||||||
@@ -273,7 +277,9 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
|
|
||||||
self._init_dllm_algorithm()
|
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
|
self.tokenizer = self.processor = None
|
||||||
else:
|
else:
|
||||||
if self.model_config.is_multimodal:
|
if self.model_config.is_multimodal:
|
||||||
@@ -370,6 +376,7 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
else self.server_args.speculative_draft_model_revision
|
else self.server_args.speculative_draft_model_revision
|
||||||
),
|
),
|
||||||
is_draft_model=self.is_draft_worker,
|
is_draft_model=self.is_draft_worker,
|
||||||
|
context_length=self.context_length,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_model_runner(self):
|
def _init_model_runner(self):
|
||||||
|
|||||||
@@ -528,12 +528,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
# Model-specific adjustment
|
# Model-specific adjustment
|
||||||
self.model_specific_adjustment()
|
self.model_specific_adjustment()
|
||||||
|
|
||||||
# Set the global server_args in the scheduler process
|
# 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)
|
set_global_server_args_for_scheduler(server_args)
|
||||||
global_server_args = get_global_server_args()
|
|
||||||
|
|
||||||
# FIXME: hacky set `use_mla_backend`
|
# FIXME: hacky set `use_mla_backend`
|
||||||
global_server_args.use_mla_backend = self.use_mla_backend
|
get_global_server_args().use_mla_backend = self.use_mla_backend
|
||||||
|
|
||||||
# Init OpenMP threads binding for CPU
|
# Init OpenMP threads binding for CPU
|
||||||
if self.device == "cpu":
|
if self.device == "cpu":
|
||||||
@@ -1128,6 +1128,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def model_specific_adjustment(self):
|
def model_specific_adjustment(self):
|
||||||
|
if self.is_draft_worker:
|
||||||
|
return
|
||||||
|
|
||||||
server_args = self.server_args
|
server_args = self.server_args
|
||||||
|
|
||||||
# HRM-Text needs bidirectional prompt attention (prefill), which only the
|
# HRM-Text needs bidirectional prompt attention (prefill), which only the
|
||||||
|
|||||||
@@ -1,6 +1,5 @@
|
|||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
from copy import deepcopy
|
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -15,11 +14,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardMode,
|
ForwardMode,
|
||||||
compute_position,
|
compute_position,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import ServerArgs
|
||||||
ServerArgs,
|
|
||||||
get_global_server_args,
|
|
||||||
set_global_server_args_for_scheduler,
|
|
||||||
)
|
|
||||||
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
||||||
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
|
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
|
||||||
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
||||||
@@ -134,52 +129,8 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
self._logged_first_verify = False
|
self._logged_first_verify = False
|
||||||
|
|
||||||
# Draft runner (separate KV cache + attention backend).
|
# 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(
|
self._draft_worker = TpModelWorker(
|
||||||
server_args=draft_server_args,
|
server_args=server_args,
|
||||||
gpu_id=gpu_id,
|
gpu_id=gpu_id,
|
||||||
tp_rank=tp_rank,
|
tp_rank=tp_rank,
|
||||||
moe_ep_rank=moe_ep_rank,
|
moe_ep_rank=moe_ep_rank,
|
||||||
@@ -189,8 +140,8 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
dp_rank=dp_rank,
|
dp_rank=dp_rank,
|
||||||
nccl_port=nccl_port,
|
nccl_port=nccl_port,
|
||||||
is_draft_worker=True,
|
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_model_runner = self._draft_worker.model_runner
|
||||||
self._draft_sampler = None
|
self._draft_sampler = None
|
||||||
# Keep the same alias that other spec-v2 workers expose.
|
# Keep the same alias that other spec-v2 workers expose.
|
||||||
@@ -226,7 +177,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
if self.tp_rank == 0:
|
if self.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s",
|
"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.draft_model.__class__.__name__,
|
||||||
self.block_size,
|
self.block_size,
|
||||||
self.draft_window_size,
|
self.draft_window_size,
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
import logging
|
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
|
from sglang.srt.utils.common import is_blackwell, is_hip, is_musa, is_npu
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -118,7 +118,7 @@ class DraftBackendFactory:
|
|||||||
return DeepseekSparseAttnBackend(self.draft_model_runner, skip_prefill=False)
|
return DeepseekSparseAttnBackend(self.draft_model_runner, skip_prefill=False)
|
||||||
|
|
||||||
def _create_flashinfer_decode_backend(self):
|
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 (
|
from sglang.srt.layers.attention.flashinfer_backend import (
|
||||||
FlashInferMultiStepDraftBackend,
|
FlashInferMultiStepDraftBackend,
|
||||||
)
|
)
|
||||||
@@ -193,7 +193,7 @@ class DraftBackendFactory:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _create_trtllm_mla_decode_backend(self, backend: str = "trtllm-gen"):
|
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(
|
raise ValueError(
|
||||||
"trtllm_mla backend requires MLA model (use_mla_backend=True)."
|
"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")
|
return self._create_trtllm_mla_decode_backend(backend="cute-dsl")
|
||||||
|
|
||||||
def _create_tokenspeed_mla_decode_backend(self):
|
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(
|
raise ValueError(
|
||||||
"tokenspeed_mla backend requires MLA model (use_mla_backend=True)."
|
"tokenspeed_mla backend requires MLA model (use_mla_backend=True)."
|
||||||
)
|
)
|
||||||
@@ -259,7 +259,7 @@ class DraftBackendFactory:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _create_flashinfer_prefill_backend(self):
|
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 (
|
from sglang.srt.layers.attention.flashinfer_backend import (
|
||||||
FlashInferAttnBackend,
|
FlashInferAttnBackend,
|
||||||
)
|
)
|
||||||
@@ -307,7 +307,7 @@ class DraftBackendFactory:
|
|||||||
return TRTLLMHAAttnBackend(self.draft_model_runner, skip_prefill=False)
|
return TRTLLMHAAttnBackend(self.draft_model_runner, skip_prefill=False)
|
||||||
|
|
||||||
def _create_trtllm_mla_prefill_backend(self):
|
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(
|
raise ValueError(
|
||||||
"trtllm_mla backend requires MLA model (use_mla_backend=True)."
|
"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)
|
return TRTLLMMLABackend(self.draft_model_runner, skip_prefill=False)
|
||||||
|
|
||||||
def _create_tokenspeed_mla_prefill_backend(self):
|
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(
|
raise ValueError(
|
||||||
"tokenspeed_mla backend requires MLA model (use_mla_backend=True)."
|
"tokenspeed_mla backend requires MLA model (use_mla_backend=True)."
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user