refactor: wrap split backends once on full-attention backends (#31439)

This commit is contained in:
Mick
2026-07-17 19:15:04 +08:00
committed by GitHub
parent 24a8944e15
commit 681c223570
2 changed files with 102 additions and 11 deletions
@@ -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)
@@ -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"]))