[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

View File

@@ -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(

View File

@@ -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: