[MiMoV2Flash] fix: respect --swa-full-tokens-ratio arg (#15488)

This commit is contained in:
Yingchun Lai
2025-12-25 21:02:56 +08:00
committed by GitHub
parent de03b0cd30
commit bb9e6cdf9e
2 changed files with 16 additions and 16 deletions
@@ -334,7 +334,6 @@ class ModelRunner:
self.attention_chunk_size = model_config.attention_chunk_size self.attention_chunk_size = model_config.attention_chunk_size
self.forward_pass_id = 0 self.forward_pass_id = 0
self.init_new_workspace = False self.init_new_workspace = False
self.kv_cache_memory = 0
self.draft_model_idx = draft_model_idx self.draft_model_idx = draft_model_idx
self.remote_instance_transfer_engine = None self.remote_instance_transfer_engine = None
@@ -1582,10 +1581,9 @@ class ModelRunner:
) )
if self.mambaish_config is not None: if self.mambaish_config is not None:
rest_memory = self.handle_max_mamba_cache(rest_memory) 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.") 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): def handle_max_mamba_cache(self, total_rest_memory):
config = self.mambaish_config config = self.mambaish_config
@@ -1719,14 +1717,6 @@ class ModelRunner:
self.max_total_num_tokens // page_size * page_size self.max_total_num_tokens // page_size * page_size
) )
self.max_total_num_tokens = self.swa_max_total_num_tokens 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: else:
assert self.sliding_window_size is not None and self.sliding_window_size > 0 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) 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.swa_max_total_num_tokens = int(
self.full_max_total_num_tokens * swa_full_tokens_ratio 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 self.max_total_num_tokens = self.full_max_total_num_tokens
logger.info( logger.info(
+6 -4
View File
@@ -1203,11 +1203,11 @@ class ServerArgs:
"Spec v2 is enabled for multi-layer EAGLE speculative decoding." "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: 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 self.disable_hybrid_swa_memory = True
logger.warning( logger.warning(
"Disable hybrid SWA memory for MiMoV2FlashForCausalLM model with hierarchical cache" "Disable hybrid SWA memory for MiMoV2FlashForCausalLM model with hierarchical cache"
@@ -2263,6 +2263,8 @@ class ServerArgs:
raise ValueError( raise ValueError(
"Spec v2 and decode offload kv cache are incompatible and cannot be enabled together." "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): def _handle_deterministic_inference(self):
if self.rl_on_policy_target is not None: if self.rl_on_policy_target is not None: