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):
|
def _get_attention_backend(self, init_new_workspace: bool = False):
|
||||||
"""Init attention kernel backend."""
|
"""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.prefill_attention_backend_str, self.decode_attention_backend_str = (
|
||||||
self.server_args.get_attention_backends()
|
self.server_args.get_attention_backends()
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -424,6 +424,7 @@ class ServerArgs:
|
|||||||
speculative_accept_threshold_acc: float = 1.0
|
speculative_accept_threshold_acc: float = 1.0
|
||||||
speculative_token_map: Optional[str] = None
|
speculative_token_map: Optional[str] = None
|
||||||
speculative_attention_mode: str = "prefill"
|
speculative_attention_mode: str = "prefill"
|
||||||
|
speculative_draft_attention_backend: Optional[str] = None
|
||||||
speculative_moe_runner_backend: Optional[str] = None
|
speculative_moe_runner_backend: Optional[str] = None
|
||||||
speculative_moe_a2a_backend: Optional[str] = None
|
speculative_moe_a2a_backend: Optional[str] = None
|
||||||
speculative_draft_model_quantization: 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'.",
|
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,
|
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(
|
parser.add_argument(
|
||||||
"--speculative-moe-runner-backend",
|
"--speculative-moe-runner-backend",
|
||||||
type=str,
|
type=str,
|
||||||
|
|||||||
@@ -18,11 +18,16 @@ class DraftBackendFactory:
|
|||||||
self.draft_model_runner = draft_model_runner
|
self.draft_model_runner = draft_model_runner
|
||||||
self.topk = topk
|
self.topk = topk
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
|
self.draft_attn_backend = server_args.speculative_draft_attention_backend
|
||||||
|
|
||||||
def _create_backend(
|
def _create_backend(
|
||||||
self, backend_name: str, backend_map: dict, error_template: str
|
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:
|
if backend_type is None:
|
||||||
backend_type = self.server_args.attention_backend
|
backend_type = self.server_args.attention_backend
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user