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