Gate mamba extra-buffer predicates on uses_mamba_radix_cache (#37474)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
Co-authored-by: Ke Bao <ispobaoke@gmail.com>
Co-authored-by: Sam Shleifer <sam@thinkingmachines.ai>
This commit is contained in:
Sam Shleifer
2026-09-17 14:39:12 +08:00
committed by GitHub
co-authored by Claude Fable 5 Ke Bao Sam Shleifer
parent fd6f96bf96
commit 8ae4a39b50
4 changed files with 30 additions and 0 deletions
@@ -921,6 +921,13 @@ def handle_mamba_radix_cache(server_args: Any, model_arch: str):
run_post_process_pass(server_args, _mamba_radix_cache_resolution)
view = resolved_view(server_args)
if not view.uses_mamba_radix_cache:
# auto is arch-gated, so only an explicit strategy reaches a non-mamba
# arch here, where it would arm the mamba paths and crash at prefill.
if mamba_extra_buffer_of(view):
raise ValueError(
f"--mamba-radix-cache-strategy {view.mamba_radix_cache_strategy} "
f"needs mamba state, got {model_arch}."
)
return
if mamba_extra_buffer_of(view):
@@ -58,6 +58,10 @@ def _inkling_overrides(server_args: Any, hf_config: Any) -> dict:
# spec today, but giving it one would silently stop this pin from firing.
if cfg.mamba_radix_cache_strategy == "auto":
overrides["mamba_radix_cache_strategy"] = "extra_buffer"
# The generic resolution never sets the arch-derived leaf for Inkling, and
# handle_mamba_radix_cache rejects extra_buffer on a model without it.
if not cfg.disable_radix_cache:
overrides["uses_mamba_radix_cache"] = True
# Inkling attention runs only on the fa4 (Blackwell) or triton backends --
# models/inkling_common/attn.py asserts attention_backend in {fa4, triton}.
# The generic resolver would otherwise pick trtllm_mha (SM100) / fa3
@@ -488,6 +488,10 @@ _MAMBA_EXTRA_BUFFER_ARCHS = frozenset(
# KDA backend's track-snapshot writes (decode + extend) so donated
# slots hold real states for prefix-cache restores.
"KimiK3ForConditionalGeneration",
# Inkling asserts enable_mamba_extra_buffer and _inkling_overrides pins it,
# so validate_mamba_extra_buffer runs for these archs and must accept them.
"InklingForConditionalGeneration",
"InklingForConditionalGenerationMTP",
}
)
@@ -591,6 +591,21 @@ class TestGoldenModelOverrides(_IsolatedPublish):
set_global_server_args_for_scheduler(server_args)
return get_server_args()
def test_explicit_extra_buffer_without_mamba_state_fails_fast(self):
with self.assertRaisesRegex(ValueError, "needs mamba state"):
self._construct(
"LlamaForCausalLM", "llama", mamba_radix_cache_strategy="extra_buffer"
)
def test_explicit_extra_buffer_is_harmless_with_radix_cache_disabled(self):
sa = self._construct(
"LlamaForCausalLM",
"llama",
mamba_radix_cache_strategy="extra_buffer",
disable_radix_cache=True,
)
self.assertFalse(self._resolved(sa, "uses_mamba_radix_cache"))
def test_mistral_large3_forces_bfloat16(self):
sa = self._construct("MistralLarge3ForCausalLM", "mistral")
self.assertEqual(