diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index f9de19ff1..66c83d463 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -1235,14 +1235,17 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict: "minicpm_flashattn": ("fa4" if is_blackwell_supported() else "fa3"), "minicpm_flashinfer": "flashinfer", } - for backend_field in ( - "attention_backend", - "prefill_attention_backend", - "decode_attention_backend", - ): - dense_backend = dense_backends.get(getattr(server_args, backend_field)) - if dense_backend is not None: - overrides[backend_field] = dense_backend + # Literal keys keep the written-field set statically derivable; a loop + # variable hides it from the census in test_chain_read_ratchet.py. + dense_attention = dense_backends.get(server_args.attention_backend) + if dense_attention is not None: + overrides["attention_backend"] = dense_attention + dense_prefill = dense_backends.get(server_args.prefill_attention_backend) + if dense_prefill is not None: + overrides["prefill_attention_backend"] = dense_prefill + dense_decode = dense_backends.get(server_args.decode_attention_backend) + if dense_decode is not None: + overrides["decode_attention_backend"] = dense_decode elif has_sparse_attention: uses_sparse_backend = server_args.is_attention_backend_not_set() or any( backend in ("minicpm_flashattn", "minicpm_flashinfer") diff --git a/python/sglang/srt/layers/attention/minicpm/backend.py b/python/sglang/srt/layers/attention/minicpm/backend.py index 657c41f97..b3dfea7e7 100644 --- a/python/sglang/srt/layers/attention/minicpm/backend.py +++ b/python/sglang/srt/layers/attention/minicpm/backend.py @@ -17,7 +17,7 @@ from sglang.srt.layers.attention.minicpm.attention_adapter import ( ) from sglang.srt.layers.attention.minicpm.cache import attach_compressed_cache from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_schedule from sglang.srt.utils import is_blackwell_supported, next_power_of_2 if TYPE_CHECKING: @@ -273,7 +273,7 @@ class MiniCPMSparseBackend(AttentionBackend): "local_blocks": self.local_blocks, "dtype_str": dtype_str, } - chunked_prefill_size = model_runner.server_args.chunked_prefill_size + chunked_prefill_size = get_schedule().chunked_prefill_size if self.minicpm_fuse_topk and chunked_prefill_size <= 0: raise ValueError( "MiniCPM fused top-k requires a positive --chunked-prefill-size." diff --git a/test/registered/unit/layers/test_minicpm_sparse_metadata.py b/test/registered/unit/layers/test_minicpm_sparse_metadata.py index ab0ea4b0a..6e501b1c9 100644 --- a/test/registered/unit/layers/test_minicpm_sparse_metadata.py +++ b/test/registered/unit/layers/test_minicpm_sparse_metadata.py @@ -34,6 +34,7 @@ with patch.dict( from sglang.srt.layers.attention.minicpm.sparse_utils import ( CompressionLevelMetadata, ) + from sglang.srt.runtime_context import get_context, get_schedule register_cpu_ci(est_time=5, suite="base-a-test-cpu") @@ -73,7 +74,6 @@ def _construct_sparse_backend( token_to_kv_pool_allocator=SimpleNamespace(), server_args=SimpleNamespace( enable_memory_saver=False, - chunked_prefill_size=chunked_prefill_size, ), model_config=SimpleNamespace( hf_config=SimpleNamespace( @@ -94,6 +94,7 @@ def _construct_sparse_backend( ), ) with ( + get_schedule().override(chunked_prefill_size=chunked_prefill_size), patch.object(backend_module, "MiniCPMHybridConfig", SimpleNamespace), patch.object(backend_module, "is_blackwell_supported", return_value=blackwell), patch.object( @@ -151,6 +152,15 @@ class _SingleTensorConversion: class TestMiniCPMSparseMetadata(CustomTestCase): + def setUp(self): + super().setUp() + # The backend reads chunked_prefill_size off the schedule bag, so the + # context has to be published before any construction; the helper + # scopes a different value on top of this one where a case needs it. + override = get_context().override_server_args(chunked_prefill_size=64) + override.install() + self.addCleanup(override.restore) + def test_sparse_backend_rejects_context_too_short_for_layout(self): with self.assertRaisesRegex( ValueError, @@ -369,7 +379,6 @@ class TestMiniCPMSparseMetadata(CustomTestCase): attention_backend="minicpm_flashattn", disable_cuda_graph=False, enable_memory_saver=False, - chunked_prefill_size=64, ), model_config=SimpleNamespace( hf_config=hf_config,