[Feature] Xiaomi MiMo-V2-Flash day0 support (#15207)

Co-authored-by: 谢学扬 <xiexueyang@xiaomi.com>
Co-authored-by: tz <tangzhen3@xiaomi.com>
Co-authored-by: 李家乐 <lijiale10@xiaomi.com>
Co-authored-by: 张晨 <zhangchen50@xiaomi.com>
Co-authored-by: Shaohui Liu <liushaohui3@xiaomi.com>
Co-authored-by: 王晨 <wangchen77@xiaomi.com>
Co-authored-by: jiangzihan <jiangzihan@xiaomi.com>
Co-authored-by: xiexueyang <xyxie_wangyi@163.com>
Co-authored-by: Linghao Zhang <zhanglinghao@xiaomi.com>
Co-authored-by: ispobock <ispobaoke@gmail.com>
Co-authored-by: Liangsheng Yin <lsyincs@gmail.com>
Co-authored-by: JoyFuture <35593546+JoyFuture@users.noreply.github.com>
Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
Co-authored-by: Qiaolin Yu <liin1211@outlook.com>
Co-authored-by: root <root@bj9-ml-g8h20e-k8s-slave106-20251106.alicn.idc.xiaomi.com>
This commit is contained in:
Yingchun Lai
2025-12-19 11:40:07 +08:00
committed by GitHub
co-authored by 谢学扬 tz 李家乐 张晨 Shaohui Liu 王晨 jiangzihan xiexueyang Linghao Zhang ispobock Liangsheng Yin JoyFuture Liangsheng Yin Qiaolin Yu root
parent a0985dd5e5
commit 160a06cab2
38 changed files with 5396 additions and 169 deletions
+100 -45
View File
@@ -296,6 +296,7 @@ class ModelRunner:
is_draft_worker: bool = False,
req_to_token_pool: Optional[ReqToTokenPool] = None,
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
draft_model_idx: Optional[int] = None,
):
# Parse args
self.mem_fraction_static = mem_fraction_static
@@ -324,10 +325,13 @@ class ModelRunner:
self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
self.is_hybrid_swa = model_config.is_hybrid_swa
self.is_hybrid_swa_compress = model_config.is_hybrid_swa_compress
self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
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
self.remote_instance_transfer_engine_session_id = ""
@@ -481,6 +485,8 @@ class ModelRunner:
self.model_config.num_attention_layers,
)
)
if self.model_config.hf_config.architectures[0] == "MiMoV2MTP":
model_num_layers = 1
self.start_layer = getattr(self.model, "start_layer", 0)
self.end_layer = getattr(self.model, "end_layer", model_num_layers)
self.num_effective_layers = self.end_layer - self.start_layer
@@ -493,6 +499,23 @@ class ModelRunner:
)
), "PP is not compatible with MTP models."
# Consider PP, so use start_layer and end_layer.
full_attention_layer_ids = [
layer_idx
for layer_idx in range(self.start_layer, self.end_layer + 1)
if hasattr(self.model_config, "full_attention_layer_ids")
and layer_idx in self.model_config.full_attention_layer_ids
]
swa_attention_layer_ids = [
layer_idx
for layer_idx in range(self.start_layer, self.end_layer + 1)
if hasattr(self.model_config, "swa_attention_layer_ids")
and layer_idx in self.model_config.swa_attention_layer_ids
]
# Update back to model_config.
self.model_config.swa_attention_layer_ids = swa_attention_layer_ids
self.model_config.full_attention_layer_ids = full_attention_layer_ids
# Apply torchao quantization
torchao_applied = getattr(self.model, "torchao_applied", False)
# In layered loading, torchao may have been applied
@@ -811,6 +834,7 @@ class ModelRunner:
remote_instance_weight_loader_transfer_engine=self.remote_instance_transfer_engine,
modelopt_config=modelopt_config,
rl_quant_profile=self.server_args.rl_quant_profile,
draft_model_idx=self.draft_model_idx,
)
if self.device == "cpu":
self.model_config = adjust_config_with_unaligned_cpu_tp(
@@ -1431,6 +1455,8 @@ class ModelRunner:
)
elif config := self.mambaish_config:
num_layers = len(config.full_attention_layer_ids)
elif self.model_config.full_attention_layer_ids:
num_layers = len(self.model_config.full_attention_layer_ids)
else:
num_layers = self.num_effective_layers
if self.use_mla_backend:
@@ -1468,9 +1494,8 @@ class ModelRunner:
else:
cell_size = (
self.model_config.get_num_kv_heads(get_attention_tp_size())
* self.model_config.head_dim
* (self.model_config.head_dim + self.model_config.v_head_dim)
* num_layers
* 2
* torch._utils._element_size(self.kv_cache_dtype)
)
@@ -1491,12 +1516,24 @@ class ModelRunner:
// scale_block_size
)
if self.model_config.hf_config.architectures[0] == "MiMoV2FlashForCausalLM":
cell_size += (
self.model_config.get_swa_num_kv_heads(get_attention_tp_size())
* (
self.model_config.hf_text_config.swa_head_dim
+ self.model_config.hf_text_config.swa_v_head_dim
)
* len(self.model_config.swa_attention_layer_ids)
* torch._utils._element_size(self.kv_cache_dtype)
)
rest_memory = available_gpu_memory - total_gpu_memory * (
1 - self.mem_fraction_static
)
if self.mambaish_config is not None:
rest_memory = self.handle_max_mamba_cache(rest_memory)
max_num_token = int(rest_memory * (1 << 30) // cell_size)
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
def handle_max_mamba_cache(self, total_rest_memory):
@@ -1578,6 +1615,14 @@ class ModelRunner:
return config.llm_config
return None
@property
def max_token_pool_size(self):
"""Return the max token pool size considering hybrid swa settings."""
if self.is_hybrid_swa:
return min(self.swa_max_total_num_tokens, self.max_total_num_tokens)
else:
return self.max_total_num_tokens
@property
def kimi_linear_config(self):
config = self.model_config.hf_config
@@ -1590,6 +1635,7 @@ class ModelRunner:
return self.mamba2_config or self.hybrid_gdn_config or self.kimi_linear_config
def set_num_token_hybrid(self):
page_size = self.server_args.page_size
if (
"Llama4ForConditionalGeneration"
in self.model_config.hf_config.architectures
@@ -1607,44 +1653,33 @@ class ModelRunner:
4 * self.max_total_num_tokens
- 12 * self.max_total_num_tokens * temp_ratio // (3 * temp_ratio + 1)
)
self.swa_max_total_num_tokens = int(
self.swa_max_total_num_tokens
// self.server_args.page_size
* self.server_args.page_size
self.swa_max_total_num_tokens = (
self.swa_max_total_num_tokens // page_size * page_size
)
self.full_max_total_num_tokens = int(
self.full_max_total_num_tokens
// self.server_args.page_size
* self.server_args.page_size
self.full_max_total_num_tokens = (
self.full_max_total_num_tokens // page_size * page_size
)
self.max_total_num_tokens = self.full_max_total_num_tokens
elif "MiMoV2MTP" in self.model_config.hf_config.architectures:
assert self.is_draft_worker
# MiMoV2MTP uses SWA, so set full KV cache to 0
self.full_max_total_num_tokens = 0
self.swa_max_total_num_tokens = (
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_attention_layer_ids = []
swa_attention_layer_ids = []
try:
layers = self.model.model.layers
except:
try:
layers = self.model.language_model.model.layers
except:
try:
layers = self.model.language_model.layers
except:
self.is_hybrid_swa = False
return
for layer in layers:
if (
layer.self_attn.attn.sliding_window_size is None
or layer.self_attn.attn.sliding_window_size == -1
):
full_attention_layer_ids.append(layer.layer_id)
else:
swa_attention_layer_ids.append(layer.layer_id)
self.model_config.swa_attention_layer_ids = swa_attention_layer_ids
self.model_config.full_attention_layer_ids = full_attention_layer_ids
full_layers_num = len(self.model_config.full_attention_layer_ids)
swa_layers_num = len(self.model_config.swa_attention_layer_ids)
# Algorithm:
# Existing max_total_num_tokens is per layer and assume all layers have the same number of tokens.
@@ -1653,8 +1688,6 @@ class ModelRunner:
total_tokens = (
self.max_total_num_tokens * self.model_config.num_hidden_layers
)
full_layers_num = len(full_attention_layer_ids)
swa_layers_num = len(swa_attention_layer_ids)
swa_full_tokens_ratio = self.server_args.swa_full_tokens_ratio
# Solve the equations:
@@ -1667,9 +1700,9 @@ class ModelRunner:
)
self.max_total_num_tokens = self.full_max_total_num_tokens
logger.info(
f"Use Sliding window memory pool. full_layer_tokens={self.full_max_total_num_tokens}, swa_layer_tokens={self.swa_max_total_num_tokens}"
)
logger.info(
f"Use sliding window memory pool. full_layer_tokens={self.full_max_total_num_tokens}, swa_layer_tokens={self.swa_max_total_num_tokens}"
)
def can_run_piecewise_cuda_graph(self):
if self.server_args.enable_torch_compile:
@@ -1778,10 +1811,9 @@ class ModelRunner:
else:
# We are sharing the `token_to_kv_pool`, and both verify and draft tokens
# can be concurrently allocated, so we should give a headroom for it.
self.server_args.draft_runner_cache_size = (
self.max_total_num_tokens
extra_tokens = (
# draft
+ max_num_reqs
max_num_reqs
* self.server_args.speculative_num_steps
* self.server_args.speculative_eagle_topk
# verify
@@ -1791,7 +1823,9 @@ class ModelRunner:
)
# Target worker and draft worker shares the same indices for the
# token_to_kv_pool, so we should make sure to match max_total_num_tokens.
self.max_total_num_tokens = self.server_args.draft_runner_cache_size
self.max_total_num_tokens += extra_tokens
self.server_args.draft_runner_cache_size = self.max_total_num_tokens
self.server_args.max_num_reqs = max_num_reqs
if max_total_tokens is not None:
@@ -1988,6 +2022,18 @@ class ModelRunner:
)
else:
if self.is_hybrid_swa:
kwargs = {}
if self.is_hybrid_swa_compress:
kwargs = {
"swa_head_num": max(
1,
self.model_config.hf_text_config.swa_num_key_value_heads
// get_attention_tp_size(),
),
"swa_head_dim": self.model_config.hf_text_config.swa_head_dim,
"swa_v_head_dim": self.model_config.hf_text_config.swa_v_head_dim,
"v_head_dim": self.model_config.hf_text_config.v_head_dim,
}
self.token_to_kv_pool = SWAKVPool(
size=self.full_max_total_num_tokens,
size_swa=self.swa_max_total_num_tokens,
@@ -2000,6 +2046,7 @@ class ModelRunner:
full_attention_layer_ids=self.model_config.full_attention_layer_ids,
enable_kvcache_transpose=False,
device=self.device,
**kwargs,
)
elif config := self.mambaish_config:
extra_args = {}
@@ -2117,6 +2164,14 @@ class ModelRunner:
)
else:
assert self.is_draft_worker
if self.is_hybrid_swa:
assert (
self.token_to_kv_pool_allocator.__class__
== SWATokenToKVPoolAllocator
)
self.token_to_kv_pool.full_to_swa_index_mapping = (
self.token_to_kv_pool_allocator.full_to_swa_index_mapping
)
logger.info(
f"Memory pool end. "