From 0cb37c018c2227c4c41685587b840a3bb06dbc6b Mon Sep 17 00:00:00 2001 From: William Hu <67835823+willhu-jpg@users.noreply.github.com> Date: Mon, 21 Sep 2026 13:34:08 -0400 Subject: [PATCH] [KDA] Enable ReplaySSM for GLM-5.3 Flash (#40517) --- python/sglang/srt/configs/hybrid_arch.py | 7 +++++-- .../srt/mem_cache/kv_cache_configurator.py | 17 +++++++++-------- 2 files changed, 14 insertions(+), 10 deletions(-) diff --git a/python/sglang/srt/configs/hybrid_arch.py b/python/sglang/srt/configs/hybrid_arch.py index 5b364d0b8..d9d0620cc 100644 --- a/python/sglang/srt/configs/hybrid_arch.py +++ b/python/sglang/srt/configs/hybrid_arch.py @@ -133,6 +133,10 @@ def glm5_next_config(model_config: ModelConfig): return None +def hybrid_kda_config(model_config: ModelConfig): + return kimi_linear_config(model_config) or glm5_next_config(model_config) + + def linear_attn_model_spec(model_config: ModelConfig): result = _get_linear_attn_registry_result(model_config) return result[0] if result else None @@ -142,8 +146,7 @@ def mambaish_config(model_config: ModelConfig): existing = ( mamba2_config(model_config) or hybrid_gdn_config(model_config) - or kimi_linear_config(model_config) - or glm5_next_config(model_config) + or hybrid_kda_config(model_config) or hybrid_lightning_config(model_config) ) if existing: diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 2bd16e700..49e552c61 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -12,7 +12,7 @@ import torch from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.configs.hybrid_arch import ( hybrid_gdn_config, - kimi_linear_config, + hybrid_kda_config, mambaish_config, ) from sglang.srt.configs.model_config import ( @@ -285,12 +285,14 @@ class KVCacheConfigurator: kv_cache_dtype_str: Optional[str] = None mambaish_config: Optional[Any] = field(init=False) hybrid_gdn_config: Optional[Any] = field(init=False) + hybrid_kda_config: Optional[Any] = field(init=False) is_hybrid_swa_mtp_draft: bool = field(init=False) draft_swa_full_capacity: bool = field(init=False) def __post_init__(self) -> None: self.mambaish_config = mambaish_config(self.model_config) self.hybrid_gdn_config = hybrid_gdn_config(self.model_config) + self.hybrid_kda_config = hybrid_kda_config(self.model_config) self.is_hybrid_swa_mtp_draft = ( self.is_draft_worker and self.draft_model_idx is not None @@ -1091,7 +1093,7 @@ class KVCacheConfigurator: get_exec().mamba.enable_linear_replayssm_spec and ( self.hybrid_gdn_config is not None - or kimi_linear_config(self.model_config) is not None + or self.hybrid_kda_config is not None ) ), ) @@ -1132,11 +1134,11 @@ class KVCacheConfigurator: if ( get_exec().mamba.enable_linear_replayssm_spec and _algo in ("DSPARK", "DFLASH") - and kimi_linear_config(self.model_config) is None + and self.hybrid_kda_config is None ): raise ValueError( "--enable-linear-replayssm-spec with DSPARK/DFLASH requires a KDA " - "(kimi_linear) model; got a non-KDA model." + "model; got a non-KDA model." ) req_to_token_pool = HybridReqToTokenPool( size=max_num_reqs, @@ -1172,7 +1174,7 @@ class KVCacheConfigurator: get_exec().mamba.enable_linear_replayssm_spec and ( self.hybrid_gdn_config is not None - or kimi_linear_config(self.model_config) is not None + or self.hybrid_kda_config is not None ) ), ) @@ -2458,8 +2460,7 @@ class KVCacheConfigurator: # The ring is not part of mamba_cache_per_req. GDN replay is fixed-size # request scratch; KDA replay remains attached to each mamba slot. replayssm_active = get_exec().mamba.enable_linear_replayssm_spec and ( - self.hybrid_gdn_config is not None - or kimi_linear_config(self.model_config) is not None + self.hybrid_gdn_config is not None or self.hybrid_kda_config is not None ) if replayssm_active: record_len = get_exec().mamba.linear_replayssm_cache_len @@ -2471,7 +2472,7 @@ class KVCacheConfigurator: else: replayssm_ring_per_req = 0 replayssm_ring_per_req = int(replayssm_ring_per_req * pp_layer_scale) - if replayssm_active and kimi_linear_config(self.model_config) is None: + if replayssm_active and self.hybrid_kda_config is None: replay_req_slots = ( get_schedule().max_running_requests // self.ps.attn_dp_size + 1 )