diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index efef2e903..b1b107215 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -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 diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 16a028168..ac43738ff 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -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, diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 2692ce002..3db49daf8 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -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): diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 0b4338b23..c8db9208f 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 51f3fadeb..7c38edc35 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -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, diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index 291620c28..2e55c5fd0 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -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)." )