Extract attention-backend setup into a module (#31167)
This commit is contained in:
@@ -74,12 +74,7 @@ from sglang.srt.kv_canary.api import install_canary
|
||||
from sglang.srt.kv_canary.runner.canary_manager import context_tuple
|
||||
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
|
||||
from sglang.srt.layers import deep_gemm_wrapper, model_parallel
|
||||
from sglang.srt.layers.attention.attention_registry import (
|
||||
ATTENTION_BACKENDS,
|
||||
attn_backend_wrapper,
|
||||
)
|
||||
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
||||
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
||||
from sglang.srt.layers.cp.utils import (
|
||||
get_cp_strategy,
|
||||
)
|
||||
@@ -115,6 +110,11 @@ from sglang.srt.model_executor.forward_context import (
|
||||
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
|
||||
from sglang.srt.model_executor.hook_manager import register_forward_hooks
|
||||
from sglang.srt.model_executor.model_runner_components import misc_utils
|
||||
from sglang.srt.model_executor.model_runner_components.attention_backend_setup import (
|
||||
build_attention_backends,
|
||||
configure_aux_hidden_state_capture,
|
||||
get_attention_backend,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.kv_pool_runtime import (
|
||||
compute_post_capture_kv_resize,
|
||||
is_post_capture_kv_active,
|
||||
@@ -192,7 +192,6 @@ from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
enable_show_time_cost,
|
||||
get_available_gpu_memory,
|
||||
init_cublas,
|
||||
is_host_cpu_arm64,
|
||||
is_npu,
|
||||
log_info_on_rank0,
|
||||
@@ -728,31 +727,22 @@ class ModelRunner:
|
||||
|
||||
def init_attention_backends(self):
|
||||
"""Initialize attention backends only (no cuda graph capture)."""
|
||||
# TODO: Refactor device-specific init branches into platform interface (separate PR).
|
||||
# Must be called BEFORE init_decode_cuda_graph() so CUDA graph capture
|
||||
# runs with aux hidden state capture enabled.
|
||||
self.init_aux_hidden_state_capture()
|
||||
|
||||
if self.device == "cuda" or self.device == "musa":
|
||||
init_cublas()
|
||||
self.init_attention_backend()
|
||||
elif self.device in ["cpu", "xpu"]:
|
||||
self.init_attention_backend()
|
||||
elif self.device == "npu":
|
||||
self.init_attention_backend()
|
||||
# lazy init for zbal with mix mode (before graph capture when enable_cuda_graph)
|
||||
if envs.SGLANG_ZBAL_LOCAL_MEM_SIZE.get() > 0 and not self.is_draft_worker:
|
||||
from sglang.srt.hardware_backend.npu.utils import lazy_init_zbal_gva_mem
|
||||
|
||||
lazy_init_zbal_gva_mem(
|
||||
self.device,
|
||||
self.gpu_id,
|
||||
get_world_group().rank_in_group,
|
||||
get_world_group().world_size,
|
||||
get_world_group().cpu_group,
|
||||
)
|
||||
else:
|
||||
self.init_attention_backend()
|
||||
configure_aux_hidden_state_capture(
|
||||
model=self.model,
|
||||
eagle_use_aux_hidden_state=self.spec_aux_config.eagle_use_aux_hidden_state,
|
||||
eagle_aux_hidden_state_layer_ids=self.spec_aux_config.eagle_aux_hidden_state_layer_ids,
|
||||
dflash_use_aux_hidden_state=self.spec_aux_config.dflash_use_aux_hidden_state,
|
||||
dflash_target_layer_ids=self.spec_aux_config.dflash_target_layer_ids,
|
||||
is_dspark=self.spec_algorithm.is_dspark(),
|
||||
)
|
||||
backends = build_attention_backends(model_runner=self)
|
||||
self.attn_backend = backends.attn_backend
|
||||
self.decode_attn_backend = backends.decode_attn_backend
|
||||
self.decode_attn_backend_group = backends.decode_attn_backend_group
|
||||
self.prefill_attention_backend_str = backends.prefill_attention_backend_str
|
||||
self.decode_attention_backend_str = backends.decode_attention_backend_str
|
||||
|
||||
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
|
||||
"""Capture cuda graphs. Requires init_attention_backends() to have run.
|
||||
@@ -835,34 +825,6 @@ class ModelRunner:
|
||||
)
|
||||
)
|
||||
|
||||
def init_aux_hidden_state_capture(self):
|
||||
"""Configure auxiliary hidden state capture for speculative decoding.
|
||||
|
||||
Must be called before CUDA graph capture so the captured graphs
|
||||
include aux hidden state output paths.
|
||||
"""
|
||||
if self.spec_aux_config.eagle_use_aux_hidden_state:
|
||||
self.model.set_eagle3_layers_to_capture(
|
||||
self.spec_aux_config.eagle_aux_hidden_state_layer_ids
|
||||
)
|
||||
if self.spec_aux_config.dflash_use_aux_hidden_state:
|
||||
if self.spec_algorithm.is_dspark() and hasattr(
|
||||
self.model, "set_dspark_layers_to_capture"
|
||||
):
|
||||
self.model.set_dspark_layers_to_capture(
|
||||
self.spec_aux_config.dflash_target_layer_ids
|
||||
)
|
||||
elif hasattr(self.model, "set_dflash_layers_to_capture"):
|
||||
self.model.set_dflash_layers_to_capture(
|
||||
self.spec_aux_config.dflash_target_layer_ids
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Model {self.model.__class__.__name__} implements neither "
|
||||
"set_dspark_layers_to_capture nor set_dflash_layers_to_capture, "
|
||||
"one of which is required for DFLASH/DSPARK."
|
||||
)
|
||||
|
||||
def check_quantized_moe_compatibility(self):
|
||||
check_quantized_moe_compatibility(
|
||||
model_config=self.model_config,
|
||||
@@ -1125,88 +1087,10 @@ class ModelRunner:
|
||||
if resolved_kv_cache_dtype is not None:
|
||||
self._record_kv_cache_dtype(resolved_kv_cache_dtype)
|
||||
|
||||
def init_attention_backend(self):
|
||||
"""Init attention kernel backend."""
|
||||
if self.server_args.enable_pdmux:
|
||||
self.attn_backend = self._get_attention_backend(init_new_workspace=True)
|
||||
self.decode_attn_backend_group = []
|
||||
for _ in range(self.server_args.sm_group_num):
|
||||
self.decode_attn_backend_group.append(self._get_attention_backend())
|
||||
self.decode_attn_backend = self.decode_attn_backend_group[0]
|
||||
elif self.server_args.enable_two_batch_overlap and not self.is_draft_worker:
|
||||
self.attn_backend = TboAttnBackend.init_new(self._get_attention_backend)
|
||||
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
|
||||
if self.is_draft_worker and draft_attn_backend:
|
||||
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,
|
||||
)
|
||||
|
||||
(
|
||||
self.prefill_attention_backend_str,
|
||||
self.decode_attention_backend_str,
|
||||
) = self.server_args.get_attention_backends()
|
||||
|
||||
if self.decode_attention_backend_str != self.prefill_attention_backend_str:
|
||||
from sglang.srt.layers.attention.hybrid_attn_backend import (
|
||||
HybridAttnBackend,
|
||||
)
|
||||
|
||||
attn_backend = HybridAttnBackend(
|
||||
self,
|
||||
decode_backend=self._get_attention_backend_from_str(
|
||||
self.decode_attention_backend_str,
|
||||
init_new_workspace=init_new_workspace,
|
||||
),
|
||||
prefill_backend=self._get_attention_backend_from_str(
|
||||
self.prefill_attention_backend_str,
|
||||
init_new_workspace=init_new_workspace,
|
||||
),
|
||||
)
|
||||
logger.info(
|
||||
f"Using hybrid attention backend for decode and prefill: "
|
||||
f"decode_backend={self.decode_attention_backend_str}, "
|
||||
f"prefill_backend={self.prefill_attention_backend_str}."
|
||||
)
|
||||
logger.warning(
|
||||
"Warning: Attention backend specified by --attention-backend or default backend might be overridden."
|
||||
"The feature of hybrid attention backend is experimental and unstable. Please raise an issue if you encounter any problem."
|
||||
)
|
||||
else:
|
||||
attn_backend = self._get_attention_backend_from_str(
|
||||
self.server_args.attention_backend,
|
||||
init_new_workspace=init_new_workspace,
|
||||
)
|
||||
|
||||
return attn_backend
|
||||
|
||||
def _get_attention_backend_from_str(
|
||||
self, backend_str: str, init_new_workspace: bool = False
|
||||
):
|
||||
if backend_str not in ATTENTION_BACKENDS:
|
||||
raise ValueError(f"Invalid attention backend: {backend_str}")
|
||||
self.init_new_workspace = init_new_workspace
|
||||
full_attention_backend = ATTENTION_BACKENDS[backend_str](self)
|
||||
return attn_backend_wrapper(self, full_attention_backend)
|
||||
return get_attention_backend(
|
||||
model_runner=self, init_new_workspace=init_new_workspace
|
||||
)
|
||||
|
||||
def init_decode_cuda_graph(self):
|
||||
"""Capture device graphs."""
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import msgspec
|
||||
|
||||
from sglang.srt.distributed import get_world_group
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.attention_registry import (
|
||||
ATTENTION_BACKENDS,
|
||||
attn_backend_wrapper,
|
||||
)
|
||||
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
||||
from sglang.srt.utils import init_cublas
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ResolvedAttentionBackendStr(msgspec.Struct, frozen=True, kw_only=True):
|
||||
prefill: str
|
||||
decode: str
|
||||
is_draft_override: bool = False
|
||||
|
||||
|
||||
class AttentionBackends(msgspec.Struct, frozen=True, kw_only=True):
|
||||
attn_backend: AttentionBackend
|
||||
decode_attn_backend: Optional[AttentionBackend]
|
||||
decode_attn_backend_group: list[AttentionBackend]
|
||||
prefill_attention_backend_str: str
|
||||
decode_attention_backend_str: str
|
||||
|
||||
|
||||
def configure_aux_hidden_state_capture(
|
||||
*,
|
||||
model,
|
||||
eagle_use_aux_hidden_state: bool,
|
||||
eagle_aux_hidden_state_layer_ids,
|
||||
dflash_use_aux_hidden_state: bool,
|
||||
dflash_target_layer_ids,
|
||||
is_dspark: bool,
|
||||
) -> None:
|
||||
"""Configure auxiliary hidden state capture for speculative decoding.
|
||||
|
||||
Must be called before CUDA graph capture so the captured graphs
|
||||
include aux hidden state output paths.
|
||||
"""
|
||||
if eagle_use_aux_hidden_state:
|
||||
model.set_eagle3_layers_to_capture(eagle_aux_hidden_state_layer_ids)
|
||||
if dflash_use_aux_hidden_state:
|
||||
if is_dspark and hasattr(model, "set_dspark_layers_to_capture"):
|
||||
model.set_dspark_layers_to_capture(dflash_target_layer_ids)
|
||||
elif hasattr(model, "set_dflash_layers_to_capture"):
|
||||
model.set_dflash_layers_to_capture(dflash_target_layer_ids)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Model {model.__class__.__name__} implements neither "
|
||||
"set_dspark_layers_to_capture nor set_dflash_layers_to_capture, "
|
||||
"one of which is required for DFLASH/DSPARK."
|
||||
)
|
||||
|
||||
|
||||
def build_attention_backends(*, model_runner: ModelRunner) -> AttentionBackends:
|
||||
"""Init attention kernel backend."""
|
||||
server_args = model_runner.server_args
|
||||
|
||||
# TODO: Refactor device-specific init branches into platform interface (separate PR).
|
||||
if model_runner.device in ("cuda", "musa"):
|
||||
init_cublas()
|
||||
|
||||
resolved = _resolve_attention_backend_strs(
|
||||
server_args=server_args, is_draft_worker=model_runner.is_draft_worker
|
||||
)
|
||||
|
||||
if server_args.enable_pdmux:
|
||||
attn_backend = _build_resolved_backend(
|
||||
model_runner=model_runner, resolved=resolved, init_new_workspace=True
|
||||
)
|
||||
decode_attn_backend_group = [
|
||||
_build_resolved_backend(
|
||||
model_runner=model_runner,
|
||||
resolved=resolved,
|
||||
init_new_workspace=False,
|
||||
)
|
||||
for _ in range(server_args.sm_group_num)
|
||||
]
|
||||
decode_attn_backend = decode_attn_backend_group[0]
|
||||
elif server_args.enable_two_batch_overlap and not model_runner.is_draft_worker:
|
||||
attn_backend = TboAttnBackend.init_new(
|
||||
lambda: _build_resolved_backend(
|
||||
model_runner=model_runner,
|
||||
resolved=resolved,
|
||||
init_new_workspace=False,
|
||||
)
|
||||
)
|
||||
decode_attn_backend = None
|
||||
decode_attn_backend_group = []
|
||||
else:
|
||||
attn_backend = _build_resolved_backend(
|
||||
model_runner=model_runner, resolved=resolved, init_new_workspace=False
|
||||
)
|
||||
decode_attn_backend = None
|
||||
decode_attn_backend_group = []
|
||||
|
||||
if (
|
||||
model_runner.device == "npu"
|
||||
and envs.SGLANG_ZBAL_LOCAL_MEM_SIZE.get() > 0
|
||||
and not model_runner.is_draft_worker
|
||||
):
|
||||
# lazy init for zbal with mix mode (before graph capture when enable_cuda_graph)
|
||||
from sglang.srt.hardware_backend.npu.utils import lazy_init_zbal_gva_mem
|
||||
|
||||
lazy_init_zbal_gva_mem(
|
||||
model_runner.device,
|
||||
model_runner.gpu_id,
|
||||
get_world_group().rank_in_group,
|
||||
get_world_group().world_size,
|
||||
get_world_group().cpu_group,
|
||||
)
|
||||
|
||||
# Record resolved per-mode backends on the backend for model dispatch.
|
||||
attn_backend.prefill_attention_backend_str = resolved.prefill
|
||||
attn_backend.decode_attention_backend_str = resolved.decode
|
||||
|
||||
return AttentionBackends(
|
||||
attn_backend=attn_backend,
|
||||
decode_attn_backend=decode_attn_backend,
|
||||
decode_attn_backend_group=decode_attn_backend_group,
|
||||
prefill_attention_backend_str=resolved.prefill,
|
||||
decode_attention_backend_str=resolved.decode,
|
||||
)
|
||||
|
||||
|
||||
def get_attention_backend(
|
||||
*, model_runner: ModelRunner, init_new_workspace: bool = False
|
||||
) -> AttentionBackend:
|
||||
"""Init attention kernel backend."""
|
||||
resolved = _resolve_attention_backend_strs(
|
||||
server_args=model_runner.server_args,
|
||||
is_draft_worker=model_runner.is_draft_worker,
|
||||
)
|
||||
return _build_resolved_backend(
|
||||
model_runner=model_runner,
|
||||
resolved=resolved,
|
||||
init_new_workspace=init_new_workspace,
|
||||
)
|
||||
|
||||
|
||||
def _resolve_attention_backend_strs(
|
||||
*, server_args: ServerArgs, is_draft_worker: bool
|
||||
) -> ResolvedAttentionBackendStr:
|
||||
draft_attn_backend = server_args.speculative_draft_attention_backend
|
||||
if is_draft_worker and draft_attn_backend:
|
||||
logger.warning(f"Overriding draft attention backend to {draft_attn_backend}.")
|
||||
# Single backend for all draft modes (no prefill/decode split).
|
||||
return ResolvedAttentionBackendStr(
|
||||
prefill=draft_attn_backend,
|
||||
decode=draft_attn_backend,
|
||||
is_draft_override=True,
|
||||
)
|
||||
prefill, decode = server_args.get_attention_backends()
|
||||
return ResolvedAttentionBackendStr(prefill=prefill, decode=decode)
|
||||
|
||||
|
||||
def _build_resolved_backend(
|
||||
*,
|
||||
model_runner: ModelRunner,
|
||||
resolved: ResolvedAttentionBackendStr,
|
||||
init_new_workspace: bool,
|
||||
) -> AttentionBackend:
|
||||
if resolved.is_draft_override:
|
||||
attn_backend = _build_backend_from_str(
|
||||
model_runner=model_runner,
|
||||
backend_str=resolved.prefill,
|
||||
init_new_workspace=init_new_workspace,
|
||||
)
|
||||
elif resolved.decode != resolved.prefill:
|
||||
from sglang.srt.layers.attention.hybrid_attn_backend import (
|
||||
HybridAttnBackend,
|
||||
)
|
||||
|
||||
attn_backend = HybridAttnBackend(
|
||||
model_runner,
|
||||
decode_backend=_build_backend_from_str(
|
||||
model_runner=model_runner,
|
||||
backend_str=resolved.decode,
|
||||
init_new_workspace=init_new_workspace,
|
||||
),
|
||||
prefill_backend=_build_backend_from_str(
|
||||
model_runner=model_runner,
|
||||
backend_str=resolved.prefill,
|
||||
init_new_workspace=init_new_workspace,
|
||||
),
|
||||
)
|
||||
logger.info(
|
||||
f"Using hybrid attention backend for decode and prefill: "
|
||||
f"decode_backend={resolved.decode}, "
|
||||
f"prefill_backend={resolved.prefill}."
|
||||
)
|
||||
logger.warning(
|
||||
"Warning: Attention backend specified by --attention-backend or default backend might be overridden."
|
||||
"The feature of hybrid attention backend is experimental and unstable. Please raise an issue if you encounter any problem."
|
||||
)
|
||||
else:
|
||||
attn_backend = _build_backend_from_str(
|
||||
model_runner=model_runner,
|
||||
backend_str=model_runner.server_args.attention_backend,
|
||||
init_new_workspace=init_new_workspace,
|
||||
)
|
||||
return attn_backend
|
||||
|
||||
|
||||
def _build_backend_from_str(
|
||||
*, model_runner: ModelRunner, backend_str: str, init_new_workspace: bool
|
||||
) -> AttentionBackend:
|
||||
if backend_str not in ATTENTION_BACKENDS:
|
||||
raise ValueError(f"Invalid attention backend: {backend_str}")
|
||||
model_runner.init_new_workspace = init_new_workspace
|
||||
full_attention_backend = ATTENTION_BACKENDS[backend_str](model_runner)
|
||||
return attn_backend_wrapper(model_runner, full_attention_backend)
|
||||
Reference in New Issue
Block a user