diff --git a/python/sglang/srt/model_executor/model_runner_components/attention_backend_setup.py b/python/sglang/srt/model_executor/model_runner_components/attention_backend_setup.py index 0aec1107a..e9a863690 100644 --- a/python/sglang/srt/model_executor/model_runner_components/attention_backend_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/attention_backend_setup.py @@ -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) diff --git a/test/registered/unit/model_executor/model_runner_components/test_attention_backend_setup.py b/test/registered/unit/model_executor/model_runner_components/test_attention_backend_setup.py new file mode 100644 index 000000000..2f9f9e5dc --- /dev/null +++ b/test/registered/unit/model_executor/model_runner_components/test_attention_backend_setup.py @@ -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"]))