refactor: wrap split backends once on full-attention backends (#31439)
This commit is contained in:
+32
-11
@@ -184,17 +184,26 @@ def _build_resolved_backend(
|
||||
HybridAttnBackend,
|
||||
)
|
||||
|
||||
attn_backend = HybridAttnBackend(
|
||||
# Compose the two full-attention backends first, then apply model-level
|
||||
# wrappers once. Wrapping each child independently duplicates the
|
||||
# linear/sparse side backend for hybrid models (for example, two GDN
|
||||
# dispatchers for Qwen3.5 when prefill and decode use different MHA
|
||||
# backends), duplicating initialization and associated state while only
|
||||
# one side backend can be active in a forward pass.
|
||||
attn_backend = attn_backend_wrapper(
|
||||
model_runner,
|
||||
decode_backend=_build_backend_from_str(
|
||||
HybridAttnBackend(
|
||||
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,
|
||||
decode_backend=_build_full_attention_backend_from_str(
|
||||
model_runner=model_runner,
|
||||
backend_str=resolved.decode,
|
||||
init_new_workspace=init_new_workspace,
|
||||
),
|
||||
prefill_backend=_build_full_attention_backend_from_str(
|
||||
model_runner=model_runner,
|
||||
backend_str=resolved.prefill,
|
||||
init_new_workspace=init_new_workspace,
|
||||
),
|
||||
),
|
||||
)
|
||||
logger.info(
|
||||
@@ -217,9 +226,21 @@ def _build_resolved_backend(
|
||||
|
||||
def _build_backend_from_str(
|
||||
*, model_runner: ModelRunner, backend_str: str, init_new_workspace: bool
|
||||
) -> AttentionBackend:
|
||||
return attn_backend_wrapper(
|
||||
model_runner,
|
||||
_build_full_attention_backend_from_str(
|
||||
model_runner=model_runner,
|
||||
backend_str=backend_str,
|
||||
init_new_workspace=init_new_workspace,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _build_full_attention_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)
|
||||
return ATTENTION_BACKENDS[backend_str](model_runner)
|
||||
|
||||
+70
@@ -0,0 +1,70 @@
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend
|
||||
from sglang.srt.model_executor.model_runner_components import (
|
||||
attention_backend_setup,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.attention_backend_setup import (
|
||||
ResolvedAttentionBackendStr,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class _FakeBackend:
|
||||
def __init__(self, name):
|
||||
self.name = name
|
||||
|
||||
|
||||
def test_split_full_attention_applies_model_wrapper_once():
|
||||
runner = SimpleNamespace(
|
||||
server_args=SimpleNamespace(speculative_attention_mode="prefill"),
|
||||
kv_cache_dtype=None,
|
||||
token_to_kv_pool=object(),
|
||||
req_to_token_pool=object(),
|
||||
init_new_workspace=None,
|
||||
)
|
||||
wrapper_inputs = []
|
||||
wrapped_backend = object()
|
||||
|
||||
def wrap_once(model_runner, backend):
|
||||
assert model_runner is runner
|
||||
wrapper_inputs.append(backend)
|
||||
return wrapped_backend
|
||||
|
||||
constructors = {
|
||||
"decode-test": lambda model_runner: _FakeBackend("decode"),
|
||||
"prefill-test": lambda model_runner: _FakeBackend("prefill"),
|
||||
}
|
||||
resolved = ResolvedAttentionBackendStr(decode="decode-test", prefill="prefill-test")
|
||||
|
||||
with (
|
||||
patch.dict(attention_backend_setup.ATTENTION_BACKENDS, constructors),
|
||||
patch.object(
|
||||
attention_backend_setup,
|
||||
"attn_backend_wrapper",
|
||||
side_effect=wrap_once,
|
||||
),
|
||||
):
|
||||
result = attention_backend_setup._build_resolved_backend(
|
||||
model_runner=runner,
|
||||
resolved=resolved,
|
||||
init_new_workspace=True,
|
||||
)
|
||||
|
||||
assert result is wrapped_backend
|
||||
assert len(wrapper_inputs) == 1
|
||||
split_backend = wrapper_inputs[0]
|
||||
assert isinstance(split_backend, HybridAttnBackend)
|
||||
assert split_backend.decode_backend.name == "decode"
|
||||
assert split_backend.prefill_backend.name == "prefill"
|
||||
assert runner.init_new_workspace is True
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
Reference in New Issue
Block a user