Support using different attention backend for draft decoding. (#14843)

This commit is contained in:
Charles Chen
2025-12-18 16:49:11 -08:00
committed by GitHub
parent 216067c0cb
commit 9e0ef04e5b
3 changed files with 23 additions and 1 deletions

View File

@@ -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()
)

View File

@@ -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,

View File

@@ -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