fix: speculative draft worker clobbering target attention backend (#28559)
This commit is contained in:
@@ -38,6 +38,10 @@ class AttentionBackend(ABC):
|
|||||||
those must migrate to ``init_forward_metadata_out_graph(fb, in_capture)``.
|
those must migrate to ``init_forward_metadata_out_graph(fb, in_capture)``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
# Resolved per-mode backend names, stamped by ModelRunner.init_attention_backend
|
||||||
|
prefill_attention_backend_str: Optional[str] = None
|
||||||
|
decode_attention_backend_str: Optional[str] = None
|
||||||
|
|
||||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
"""Eager entry point. Default = ``_out_graph(fb) + _in_graph(fb)``.
|
"""Eager entry point. Default = ``_out_graph(fb) + _in_graph(fb)``.
|
||||||
|
|
||||||
|
|||||||
@@ -2415,6 +2415,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
else:
|
else:
|
||||||
self.attn_backend = self._get_attention_backend()
|
self.attn_backend = self._get_attention_backend()
|
||||||
|
|
||||||
|
# Record resolved per-mode backends on the backend for model dispatch.
|
||||||
|
self.attn_backend.prefill_attention_backend_str = (
|
||||||
|
self.prefill_attention_backend_str
|
||||||
|
)
|
||||||
|
self.attn_backend.decode_attention_backend_str = (
|
||||||
|
self.decode_attention_backend_str
|
||||||
|
)
|
||||||
|
|
||||||
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
|
draft_attn_backend = self.server_args.speculative_draft_attention_backend
|
||||||
@@ -2422,6 +2430,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
logger.warning(
|
logger.warning(
|
||||||
f"Overriding draft attention backend to {draft_attn_backend}."
|
f"Overriding draft attention backend to {draft_attn_backend}."
|
||||||
)
|
)
|
||||||
|
# Single backend for all draft modes (no prefill/decode split).
|
||||||
|
self.prefill_attention_backend_str = draft_attn_backend
|
||||||
|
self.decode_attention_backend_str = draft_attn_backend
|
||||||
return self._get_attention_backend_from_str(
|
return self._get_attention_backend_from_str(
|
||||||
draft_attn_backend,
|
draft_attn_backend,
|
||||||
init_new_workspace=init_new_workspace,
|
init_new_workspace=init_new_workspace,
|
||||||
@@ -2463,10 +2474,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
init_new_workspace=init_new_workspace,
|
init_new_workspace=init_new_workspace,
|
||||||
)
|
)
|
||||||
|
|
||||||
(
|
|
||||||
get_global_server_args().prefill_attention_backend,
|
|
||||||
get_global_server_args().decode_attention_backend,
|
|
||||||
) = (self.prefill_attention_backend_str, self.decode_attention_backend_str)
|
|
||||||
return attn_backend
|
return attn_backend
|
||||||
|
|
||||||
def _get_attention_backend_from_str(
|
def _get_attention_backend_from_str(
|
||||||
|
|||||||
@@ -132,6 +132,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
check_cuda_graph_backend,
|
check_cuda_graph_backend,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
|
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||||
from sglang.srt.models.deepseek_common.attention_backend_handler import (
|
from sglang.srt.models.deepseek_common.attention_backend_handler import (
|
||||||
AttentionBackendRegistry,
|
AttentionBackendRegistry,
|
||||||
@@ -1740,20 +1741,28 @@ class DeepseekV2AttentionMLA(
|
|||||||
def dispatch_attn_forward_method(
|
def dispatch_attn_forward_method(
|
||||||
self, forward_batch: ForwardBatch
|
self, forward_batch: ForwardBatch
|
||||||
) -> AttnForwardMethod:
|
) -> AttnForwardMethod:
|
||||||
# Determine attention backend used by current forward batch
|
# Determine attention backend name for current forward batch: prefer the
|
||||||
|
# name stamped per-runner on the backend object, else resolve from server args.
|
||||||
|
backend = get_attn_backend()
|
||||||
|
server_args = get_global_server_args()
|
||||||
|
default_prefill_str, default_decode_str = server_args.get_attention_backends()
|
||||||
|
prefill_backend_str = (
|
||||||
|
backend.prefill_attention_backend_str or default_prefill_str
|
||||||
|
)
|
||||||
|
decode_backend_str = backend.decode_attention_backend_str or default_decode_str
|
||||||
if forward_batch.forward_mode.is_decode_or_idle():
|
if forward_batch.forward_mode.is_decode_or_idle():
|
||||||
attention_backend = get_global_server_args().decode_attention_backend
|
attention_backend = decode_backend_str
|
||||||
elif (
|
elif (
|
||||||
forward_batch.forward_mode.is_target_verify()
|
forward_batch.forward_mode.is_target_verify()
|
||||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
):
|
):
|
||||||
# Use the specified backend for speculative operations (both verify and draft extend)
|
# Use the specified backend for speculative operations (both verify and draft extend)
|
||||||
if get_global_server_args().speculative_attention_mode == "decode":
|
if server_args.speculative_attention_mode == "decode":
|
||||||
attention_backend = get_global_server_args().decode_attention_backend
|
attention_backend = decode_backend_str
|
||||||
else: # default to prefill
|
else: # default to prefill
|
||||||
attention_backend = get_global_server_args().prefill_attention_backend
|
attention_backend = prefill_backend_str
|
||||||
else:
|
else:
|
||||||
attention_backend = get_global_server_args().prefill_attention_backend
|
attention_backend = prefill_backend_str
|
||||||
self.current_attention_backend = attention_backend
|
self.current_attention_backend = attention_backend
|
||||||
|
|
||||||
handler = AttentionBackendRegistry.get_handler(attention_backend)
|
handler = AttentionBackendRegistry.get_handler(attention_backend)
|
||||||
|
|||||||
Reference in New Issue
Block a user