Support using different attention backend for draft decoding. (#14843)
This commit is contained in:
@@ -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()
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user