From bb9e6cdf9e3c1868550c57439bd1bd6fdd0da679 Mon Sep 17 00:00:00 2001 From: Yingchun Lai Date: Thu, 25 Dec 2025 21:02:56 +0800 Subject: [PATCH] [MiMoV2Flash] fix: respect --swa-full-tokens-ratio arg (#15488) --- .../sglang/srt/model_executor/model_runner.py | 22 +++++++++---------- python/sglang/srt/server_args.py | 10 +++++---- 2 files changed, 16 insertions(+), 16 deletions(-) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 5218e1649..370cd4eb8 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -334,7 +334,6 @@ class ModelRunner: self.attention_chunk_size = model_config.attention_chunk_size self.forward_pass_id = 0 self.init_new_workspace = False - self.kv_cache_memory = 0 self.draft_model_idx = draft_model_idx self.remote_instance_transfer_engine = None @@ -1582,10 +1581,9 @@ class ModelRunner: ) if self.mambaish_config is not None: rest_memory = self.handle_max_mamba_cache(rest_memory) - self.kv_cache_memory = int(rest_memory * (1 << 30)) - max_num_token = int(self.kv_cache_memory // cell_size) + logger.info(f"The available memory for KV cache is {rest_memory:.2f} GB.") - return max_num_token + return int(rest_memory * (1 << 30)) // cell_size def handle_max_mamba_cache(self, total_rest_memory): config = self.mambaish_config @@ -1719,14 +1717,6 @@ class ModelRunner: self.max_total_num_tokens // page_size * page_size ) self.max_total_num_tokens = self.swa_max_total_num_tokens - elif self.model_config.hf_config.architectures[0] == "MiMoV2FlashForCausalLM": - self.full_max_total_num_tokens = ( - self.max_total_num_tokens // page_size * page_size - ) - self.swa_max_total_num_tokens = ( - self.max_total_num_tokens // page_size * page_size - ) - self.max_total_num_tokens = self.full_max_total_num_tokens else: assert self.sliding_window_size is not None and self.sliding_window_size > 0 full_layers_num = len(self.model_config.full_attention_layer_ids) @@ -1749,6 +1739,14 @@ class ModelRunner: self.swa_max_total_num_tokens = int( self.full_max_total_num_tokens * swa_full_tokens_ratio ) + + self.full_max_total_num_tokens = ( + self.full_max_total_num_tokens // page_size * page_size + ) + self.swa_max_total_num_tokens = ( + self.swa_max_total_num_tokens // page_size * page_size + ) + self.max_total_num_tokens = self.full_max_total_num_tokens logger.info( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 57bba8230..b7dc490e7 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1203,11 +1203,11 @@ class ServerArgs: "Spec v2 is enabled for multi-layer EAGLE speculative decoding." ) - self.swa_full_tokens_ratio = 1.0 - logger.warning( - "Reset swa_full_tokens_ratio to 1.0 for MiMoV2FlashForCausalLM model" - ) if self.enable_hierarchical_cache: + self.swa_full_tokens_ratio = 1.0 + logger.warning( + "Reset swa_full_tokens_ratio to 1.0 for MiMoV2FlashForCausalLM model with hierarchical cache" + ) self.disable_hybrid_swa_memory = True logger.warning( "Disable hybrid SWA memory for MiMoV2FlashForCausalLM model with hierarchical cache" @@ -2263,6 +2263,8 @@ class ServerArgs: raise ValueError( "Spec v2 and decode offload kv cache are incompatible and cannot be enabled together." ) + if not (0 < self.swa_full_tokens_ratio <= 1.0): + raise ValueError("--swa-full-tokens-ratio should be in range (0, 1.0].") def _handle_deterministic_inference(self): if self.rl_on_policy_target is not None: