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.runner.canary_manager import context_tuple
|
||||||
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
|
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 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.dsa.utils import is_dsa_enable_prefill_cp
|
||||||
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
|
||||||
from sglang.srt.layers.cp.utils import (
|
from sglang.srt.layers.cp.utils import (
|
||||||
get_cp_strategy,
|
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.graph_shared_output import GraphSharedOutput
|
||||||
from sglang.srt.model_executor.hook_manager import register_forward_hooks
|
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 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 (
|
from sglang.srt.model_executor.model_runner_components.kv_pool_runtime import (
|
||||||
compute_post_capture_kv_resize,
|
compute_post_capture_kv_resize,
|
||||||
is_post_capture_kv_active,
|
is_post_capture_kv_active,
|
||||||
@@ -192,7 +192,6 @@ from sglang.srt.utils import (
|
|||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
enable_show_time_cost,
|
enable_show_time_cost,
|
||||||
get_available_gpu_memory,
|
get_available_gpu_memory,
|
||||||
init_cublas,
|
|
||||||
is_host_cpu_arm64,
|
is_host_cpu_arm64,
|
||||||
is_npu,
|
is_npu,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
@@ -728,31 +727,22 @@ class ModelRunner:
|
|||||||
|
|
||||||
def init_attention_backends(self):
|
def init_attention_backends(self):
|
||||||
"""Initialize attention backends only (no cuda graph capture)."""
|
"""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
|
# Must be called BEFORE init_decode_cuda_graph() so CUDA graph capture
|
||||||
# runs with aux hidden state capture enabled.
|
# runs with aux hidden state capture enabled.
|
||||||
self.init_aux_hidden_state_capture()
|
configure_aux_hidden_state_capture(
|
||||||
|
model=self.model,
|
||||||
if self.device == "cuda" or self.device == "musa":
|
eagle_use_aux_hidden_state=self.spec_aux_config.eagle_use_aux_hidden_state,
|
||||||
init_cublas()
|
eagle_aux_hidden_state_layer_ids=self.spec_aux_config.eagle_aux_hidden_state_layer_ids,
|
||||||
self.init_attention_backend()
|
dflash_use_aux_hidden_state=self.spec_aux_config.dflash_use_aux_hidden_state,
|
||||||
elif self.device in ["cpu", "xpu"]:
|
dflash_target_layer_ids=self.spec_aux_config.dflash_target_layer_ids,
|
||||||
self.init_attention_backend()
|
is_dspark=self.spec_algorithm.is_dspark(),
|
||||||
elif self.device == "npu":
|
)
|
||||||
self.init_attention_backend()
|
backends = build_attention_backends(model_runner=self)
|
||||||
# lazy init for zbal with mix mode (before graph capture when enable_cuda_graph)
|
self.attn_backend = backends.attn_backend
|
||||||
if envs.SGLANG_ZBAL_LOCAL_MEM_SIZE.get() > 0 and not self.is_draft_worker:
|
self.decode_attn_backend = backends.decode_attn_backend
|
||||||
from sglang.srt.hardware_backend.npu.utils import lazy_init_zbal_gva_mem
|
self.decode_attn_backend_group = backends.decode_attn_backend_group
|
||||||
|
self.prefill_attention_backend_str = backends.prefill_attention_backend_str
|
||||||
lazy_init_zbal_gva_mem(
|
self.decode_attention_backend_str = backends.decode_attention_backend_str
|
||||||
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()
|
|
||||||
|
|
||||||
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
|
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
|
||||||
"""Capture cuda graphs. Requires init_attention_backends() to have run.
|
"""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):
|
def check_quantized_moe_compatibility(self):
|
||||||
check_quantized_moe_compatibility(
|
check_quantized_moe_compatibility(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
@@ -1125,88 +1087,10 @@ class ModelRunner:
|
|||||||
if resolved_kv_cache_dtype is not None:
|
if resolved_kv_cache_dtype is not None:
|
||||||
self._record_kv_cache_dtype(resolved_kv_cache_dtype)
|
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):
|
def _get_attention_backend(self, init_new_workspace: bool = False):
|
||||||
"""Init attention kernel backend."""
|
return get_attention_backend(
|
||||||
draft_attn_backend = self.server_args.speculative_draft_attention_backend
|
model_runner=self, init_new_workspace=init_new_workspace
|
||||||
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)
|
|
||||||
|
|
||||||
def init_decode_cuda_graph(self):
|
def init_decode_cuda_graph(self):
|
||||||
"""Capture device graphs."""
|
"""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