diff --git a/python/sglang/srt/layers/attention/base_attn_backend.py b/python/sglang/srt/layers/attention/base_attn_backend.py index 8694299b8..901200f79 100644 --- a/python/sglang/srt/layers/attention/base_attn_backend.py +++ b/python/sglang/srt/layers/attention/base_attn_backend.py @@ -38,6 +38,10 @@ class AttentionBackend(ABC): 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): """Eager entry point. Default = ``_out_graph(fb) + _in_graph(fb)``. diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 441ef5a06..c32a1c06f 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2415,6 +2415,14 @@ class ModelRunner(ModelRunnerKVCacheMixin): else: 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): """Init attention kernel backend.""" draft_attn_backend = self.server_args.speculative_draft_attention_backend @@ -2422,6 +2430,9 @@ class ModelRunner(ModelRunnerKVCacheMixin): logger.warning( 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( draft_attn_backend, init_new_workspace=init_new_workspace, @@ -2463,10 +2474,6 @@ class ModelRunner(ModelRunnerKVCacheMixin): 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 def _get_attention_backend_from_str( diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 3d6dbb376..4468090ef 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -132,6 +132,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( check_cuda_graph_backend, ) 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.models.deepseek_common.attention_backend_handler import ( AttentionBackendRegistry, @@ -1740,20 +1741,28 @@ class DeepseekV2AttentionMLA( def dispatch_attn_forward_method( self, forward_batch: ForwardBatch ) -> 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(): - attention_backend = get_global_server_args().decode_attention_backend + attention_backend = decode_backend_str elif ( forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_draft_extend_v2() ): # Use the specified backend for speculative operations (both verify and draft extend) - if get_global_server_args().speculative_attention_mode == "decode": - attention_backend = get_global_server_args().decode_attention_backend + if server_args.speculative_attention_mode == "decode": + attention_backend = decode_backend_str else: # default to prefill - attention_backend = get_global_server_args().prefill_attention_backend + attention_backend = prefill_backend_str else: - attention_backend = get_global_server_args().prefill_attention_backend + attention_backend = prefill_backend_str self.current_attention_backend = attention_backend handler = AttentionBackendRegistry.get_handler(attention_backend)