[refactor] Migrate the page_size resolution chain (stack 12/15) (#30074)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
abbb41a214
commit
276fbfe880
@@ -311,7 +311,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
|
||||
args.model_config.hf_config.dual_chunk_attention_config = None
|
||||
return args
|
||||
|
||||
@patch("sglang.srt.server_args.is_sm100_supported", return_value=True)
|
||||
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
|
||||
@patch("sglang.srt.server_args.ServerArgs.use_mla_backend", return_value=False)
|
||||
def test_combined_attention_backend_fa4_forces_page_size_128(
|
||||
self, _mock_mla, _mock_sm100
|
||||
@@ -323,7 +323,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
|
||||
|
||||
self.assertEqual(args.page_size, 128)
|
||||
|
||||
@patch("sglang.srt.server_args.is_sm100_supported", return_value=True)
|
||||
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
|
||||
@patch("sglang.srt.server_args.ServerArgs.use_mla_backend", return_value=False)
|
||||
def test_explicit_prefill_fa4_forces_page_size_128(self, _mock_mla, _mock_sm100):
|
||||
# `--prefill-attention-backend fa4`: the previously-covered path.
|
||||
|
||||
@@ -69,6 +69,7 @@ class TestModelOverridableWhitelist(CustomTestCase):
|
||||
"disable_hybrid_swa_memory",
|
||||
"sampling_backend",
|
||||
"attention_backend",
|
||||
"page_size",
|
||||
}
|
||||
),
|
||||
)
|
||||
@@ -742,6 +743,189 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
_dllm_attention_backend(_view(dllm_algorithm=None)), {}
|
||||
)
|
||||
|
||||
def test_page_size_default_pass(self):
|
||||
from sglang.srt.arg_groups.overrides import ResolvedView, _page_size_default
|
||||
|
||||
# user-set page_size: nothing to declare
|
||||
self.assertEqual(
|
||||
_page_size_default(ResolvedView(SimpleNamespace(page_size=64))), {}
|
||||
)
|
||||
# default fill on non-HIP/non-MUSA platforms is 1
|
||||
with patch.object(overrides_module, "is_hip", return_value=False):
|
||||
with patch.object(overrides_module, "is_musa", return_value=False):
|
||||
self.assertEqual(
|
||||
_page_size_default(ResolvedView(SimpleNamespace(page_size=None))),
|
||||
{"page_size": 1},
|
||||
)
|
||||
with patch.object(overrides_module, "is_musa", return_value=True):
|
||||
self.assertEqual(
|
||||
_page_size_default(ResolvedView(SimpleNamespace(page_size=None))),
|
||||
{"page_size": 64},
|
||||
)
|
||||
|
||||
def test_dllm_page_size_pass(self):
|
||||
from sglang.srt.arg_groups.overrides import ResolvedView, _dllm_page_size
|
||||
|
||||
def _view(**kw):
|
||||
defaults = dict(
|
||||
dllm_algorithm="LowConfidence", disable_radix_cache=False, page_size=1
|
||||
)
|
||||
defaults.update(kw)
|
||||
return ResolvedView(SimpleNamespace(**defaults))
|
||||
|
||||
with patch(
|
||||
"sglang.srt.dllm.config.DllmConfig.from_server_args",
|
||||
return_value=SimpleNamespace(block_size=32),
|
||||
):
|
||||
self.assertEqual(_view() and _dllm_page_size(_view()), {"page_size": 32})
|
||||
self.assertEqual(_dllm_page_size(_view(page_size=64)), {}) # aligned
|
||||
self.assertEqual(_dllm_page_size(_view(dllm_algorithm=None)), {})
|
||||
self.assertEqual(_dllm_page_size(_view(disable_radix_cache=True)), {})
|
||||
|
||||
def test_page_size_leaf_materializes_end_state(self):
|
||||
sa = self._construct("LlamaForCausalLM", "llama")
|
||||
declared = {f for _s, d in sa._resolved_overrides for f in d}
|
||||
self.assertIn("page_size", declared) # default fill declared
|
||||
self.assertEqual(self._publish(sa).page_size, sa.page_size)
|
||||
|
||||
def test_qwen3_5_hybrid_coupled_declaration(self):
|
||||
from sglang.srt.arg_groups.overrides import _qwen3_5_hybrid_overrides
|
||||
|
||||
def _args(default_backend, **kw):
|
||||
defaults = dict(
|
||||
attention_backend=None,
|
||||
_get_default_attn_backend=lambda **_: default_backend,
|
||||
use_mla_backend=lambda: False,
|
||||
get_model_config=lambda: None,
|
||||
enable_mamba_extra_buffer=lambda: False,
|
||||
disable_radix_cache=False,
|
||||
speculative_algorithm=None,
|
||||
)
|
||||
defaults.update(kw)
|
||||
return SimpleNamespace(**defaults)
|
||||
|
||||
with patch.object(overrides_module, "is_sm100_supported", return_value=True):
|
||||
# radix on + no extra buffer + no spec -> page_size=1 path
|
||||
self.assertEqual(
|
||||
_qwen3_5_hybrid_overrides(_args("trtllm_mha"), None),
|
||||
{"attention_backend": "triton", "page_size": 1},
|
||||
)
|
||||
# spec decoding present -> trtllm_mha + page 64 (coupled)
|
||||
self.assertEqual(
|
||||
_qwen3_5_hybrid_overrides(
|
||||
_args("trtllm_mha", speculative_algorithm="EAGLE"), None
|
||||
),
|
||||
{"attention_backend": "trtllm_mha", "page_size": 64},
|
||||
)
|
||||
# user-set backend: nothing declared
|
||||
self.assertEqual(
|
||||
_qwen3_5_hybrid_overrides(
|
||||
_args("trtllm_mha", attention_backend="fa3"), None
|
||||
),
|
||||
{},
|
||||
)
|
||||
with patch.object(overrides_module, "is_sm100_supported", return_value=False):
|
||||
self.assertEqual(_qwen3_5_hybrid_overrides(_args("fa3"), None), {})
|
||||
|
||||
def test_qwen3vl_page_size(self):
|
||||
from sglang.srt.arg_groups.overrides import _qwen3vl_overrides
|
||||
|
||||
with patch.object(overrides_module, "is_hip", return_value=True):
|
||||
with patch("sglang.srt.environ.envs.SGLANG_USE_AITER_UNIFIED_ATTN") as e:
|
||||
e.get.return_value = True
|
||||
self.assertEqual(
|
||||
_qwen3vl_overrides(SimpleNamespace(page_size=None), None),
|
||||
{"page_size": 16},
|
||||
)
|
||||
self.assertEqual(
|
||||
_qwen3vl_overrides(SimpleNamespace(page_size=64), None), {}
|
||||
)
|
||||
|
||||
def test_page_constraint_passes_at_callable_level(self):
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
ResolvedView,
|
||||
_fa4_page_constraint,
|
||||
_intel_xpu_page_constraint,
|
||||
_mla_backend_page_constraints,
|
||||
)
|
||||
|
||||
def _view(**kw):
|
||||
defaults = dict(
|
||||
attention_backend=None,
|
||||
decode_attention_backend=None,
|
||||
prefill_attention_backend=None,
|
||||
page_size=1,
|
||||
)
|
||||
defaults.update(kw)
|
||||
return ResolvedView(SimpleNamespace(**defaults))
|
||||
|
||||
# flashmla snaps to 64 (unconditional within the backend match)
|
||||
self.assertEqual(
|
||||
_mla_backend_page_constraints(_view(attention_backend="flashmla")),
|
||||
{"page_size": 64},
|
||||
)
|
||||
# trtllm_mla with already-valid page: no declaration
|
||||
self.assertEqual(
|
||||
_mla_backend_page_constraints(
|
||||
_view(attention_backend="trtllm_mla", page_size=32)
|
||||
),
|
||||
{},
|
||||
)
|
||||
# chained: flashmla via decode -> 64, then trtllm_mha accepts 64
|
||||
self.assertEqual(
|
||||
_mla_backend_page_constraints(
|
||||
_view(
|
||||
decode_attention_backend="flashmla",
|
||||
prefill_attention_backend="trtllm_mha",
|
||||
)
|
||||
),
|
||||
{"page_size": 64},
|
||||
)
|
||||
# no matching backend: nothing declared
|
||||
self.assertEqual(_mla_backend_page_constraints(_view()), {})
|
||||
|
||||
with patch.object(overrides_module, "is_sm100_supported", return_value=True):
|
||||
self.assertEqual(
|
||||
_fa4_page_constraint(
|
||||
_view(
|
||||
attention_backend="fa4",
|
||||
use_mla_backend=lambda: False,
|
||||
speculative_eagle_topk=None,
|
||||
)
|
||||
),
|
||||
{"page_size": 128},
|
||||
)
|
||||
self.assertEqual(
|
||||
_fa4_page_constraint(
|
||||
_view(
|
||||
attention_backend="fa4",
|
||||
use_mla_backend=lambda: False,
|
||||
speculative_eagle_topk=2, # EAGLE topk>1 keeps default
|
||||
)
|
||||
),
|
||||
{},
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
_intel_xpu_page_constraint(
|
||||
_view(
|
||||
get_attention_backends=lambda: (None, "intel_xpu"),
|
||||
use_mla_backend=lambda: False,
|
||||
)
|
||||
),
|
||||
{"page_size": 128},
|
||||
)
|
||||
self.assertEqual(
|
||||
_intel_xpu_page_constraint(
|
||||
_view(
|
||||
get_attention_backends=lambda: (None, "intel_xpu"),
|
||||
use_mla_backend=lambda: True,
|
||||
page_size=16, # MLA decode accepts 16
|
||||
)
|
||||
),
|
||||
{},
|
||||
)
|
||||
|
||||
def test_monolith_attention_families_at_callable_level(self):
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
_falcon_h1_jet_overrides,
|
||||
|
||||
Reference in New Issue
Block a user