From 9e0ef04e5bb2b26f8b67944a25b6b7e19cb27a0a Mon Sep 17 00:00:00 2001 From: Charles Chen Date: Thu, 18 Dec 2025 16:49:11 -0800 Subject: [PATCH] Support using different attention backend for draft decoding. (#14843) --- python/sglang/srt/model_executor/model_runner.py | 10 ++++++++++ python/sglang/srt/server_args.py | 7 +++++++ python/sglang/srt/speculative/draft_utils.py | 7 ++++++- 3 files changed, 23 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 479723d74..ebc418731 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2147,6 +2147,16 @@ class ModelRunner: def _get_attention_backend(self, init_new_workspace: bool = False): """Init attention kernel backend.""" + draft_attn_backend = self.server_args.speculative_draft_attention_backend + if self.is_draft_worker and draft_attn_backend: + logger.warning( + f"Overriding draft attention backend to {draft_attn_backend}." + ) + return self._get_attention_backend_from_str( + draft_attn_backend, + init_new_workspace=init_new_workspace, + ) + self.prefill_attention_backend_str, self.decode_attention_backend_str = ( self.server_args.get_attention_backends() ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 27fc37880..a7c48fad4 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -424,6 +424,7 @@ class ServerArgs: speculative_accept_threshold_acc: float = 1.0 speculative_token_map: Optional[str] = None speculative_attention_mode: str = "prefill" + speculative_draft_attention_backend: Optional[str] = None speculative_moe_runner_backend: Optional[str] = None speculative_moe_a2a_backend: Optional[str] = None speculative_draft_model_quantization: Optional[str] = None @@ -3331,6 +3332,12 @@ class ServerArgs: help="Attention backend for speculative decoding operations (both target verify and draft extend). Can be one of 'prefill' (default) or 'decode'.", default=ServerArgs.speculative_attention_mode, ) + parser.add_argument( + "--speculative-draft-attention-backend", + type=str, + help="Attention backend for speculative decoding drafting.", + default=ServerArgs.speculative_draft_attention_backend, + ) parser.add_argument( "--speculative-moe-runner-backend", type=str, diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index d3246d30b..9c630da72 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -18,11 +18,16 @@ class DraftBackendFactory: self.draft_model_runner = draft_model_runner self.topk = topk self.speculative_num_steps = speculative_num_steps + self.draft_attn_backend = server_args.speculative_draft_attention_backend def _create_backend( self, backend_name: str, backend_map: dict, error_template: str ): - backend_type = getattr(self.server_args, backend_name) + backend_type = ( + self.draft_attn_backend + if self.draft_attn_backend + else getattr(self.server_args, backend_name) + ) if backend_type is None: backend_type = self.server_args.attention_backend