[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
+3
View File
@@ -89,6 +89,9 @@ class LoadConfig:
None # Path to rollout quantization profile (e.g., /root/profile.7b.pt)
)
# For multi-layer MTP
draft_model_idx: Optional[int] = None
def __post_init__(self):
model_loader_extra_config = self.model_loader_extra_config or {}
if isinstance(model_loader_extra_config, str):
+66 -13
View File
@@ -100,6 +100,7 @@ class ModelConfig:
model_impl: Union[str, ModelImpl] = ModelImpl.AUTO,
sampling_defaults: str = "openai",
quantize_and_serve: bool = False,
is_mtp: bool = False,
encoder_only: bool = False,
language_only: bool = False,
) -> None:
@@ -111,6 +112,7 @@ class ModelConfig:
self.model_impl = model_impl
self.sampling_defaults = sampling_defaults
self.quantize_and_serve = quantize_and_serve
self.is_mtp = is_mtp
# Validate quantize_and_serve configuration
self._validate_quantize_and_serve_config()
@@ -158,18 +160,6 @@ class ModelConfig:
self.attention_chunk_size = getattr(
self.hf_text_config, "attention_chunk_size", None
)
self.is_hybrid_swa = is_hybrid_model(
self.hf_config.architectures,
hybrid_kvcache_ratio=hybrid_kvcache_ratio,
context_length=context_length,
attention_chunk_size=self.attention_chunk_size,
)
if self.is_hybrid_swa is not None:
self.swa_attention_layer_ids, self.full_attention_layer_ids = (
get_hybrid_layer_ids(
self.hf_config.architectures, self.hf_text_config.num_hidden_layers
)
)
self.is_generation = is_generation_model(
self.hf_config.architectures, is_embedding
)
@@ -204,6 +194,9 @@ class ModelConfig:
self._derive_context_length(context_length)
self._derive_model_shapes()
# Update hybrid model
self._derive_hybrid_model(hybrid_kvcache_ratio)
# Verify quantization
self._verify_quantization()
@@ -259,6 +252,7 @@ class ModelConfig:
sampling_defaults=server_args.sampling_defaults,
quantize_and_serve=server_args.quantize_and_serve,
override_config_file=server_args.decrypted_config_file,
is_mtp=server_args.enable_mtp,
language_only=server_args.language_only,
encoder_only=server_args.encoder_only,
is_draft_model=is_draft_model,
@@ -286,6 +280,11 @@ class ModelConfig:
if is_draft_model and self.hf_config.architectures[0] == "MiMoForCausalLM":
self.hf_config.architectures[0] = "MiMoMTP"
if (
is_draft_model
and self.hf_config.architectures[0] == "MiMoV2FlashForCausalLM"
):
self.hf_config.architectures[0] = "MiMoV2MTP"
if is_draft_model and self.hf_config.architectures[0] in [
"BailingMoeV2ForCausalLM",
"BailingMoeForCausalLM",
@@ -301,6 +300,28 @@ class ModelConfig:
self.hf_config.architectures[0] = "Qwen3NextForCausalLMMTP"
self.hf_config.num_nextn_predict_layers = 1
def _derive_hybrid_model(self, hybrid_kvcache_ratio: Optional[float] = None):
# Use self.context_len after it has been initialized to prevent using context_len which may be None.
self.is_hybrid_swa = is_hybrid_model(
self.hf_config.architectures,
hybrid_kvcache_ratio=hybrid_kvcache_ratio,
context_length=self.context_len,
attention_chunk_size=self.attention_chunk_size,
)
if self.is_hybrid_swa is not None:
self.swa_attention_layer_ids, self.full_attention_layer_ids = (
get_hybrid_layer_ids(
self.hf_config.architectures,
self.hf_text_config.num_hidden_layers,
getattr(self.hf_text_config, "hybrid_layer_pattern", None),
)
)
self.is_hybrid_swa_compress = self.hf_config.architectures[0] in [
"MiMoV2FlashForCausalLM",
"MiMoV2MTP",
]
def _derive_context_length(self, context_length: int):
is_draft_model = self.is_draft_model
derived_context_len = get_context_length(self.hf_text_config)
@@ -342,6 +363,11 @@ class ModelConfig:
"head_dim",
self.hf_text_config.hidden_size // self.hf_text_config.num_attention_heads,
)
self.v_head_dim = getattr(
self.hf_text_config,
"v_head_dim",
self.head_dim,
)
# FIXME: temporary special judge for MLA architecture
if (
@@ -526,6 +552,15 @@ class ModelConfig:
# parallel size so each GPU has at least one KV head.
return max(1, total_num_kv_heads // tensor_parallel_size)
def get_swa_num_kv_heads(self, tensor_parallel_size) -> int:
"""Similar to get_num_kv_heads(), but for SWA."""
if not self.is_hybrid_swa_compress:
return 0
# For MiMoV2FlashForCausalLM models
total_num_kv_heads = self.hf_text_config.swa_num_key_value_heads
return max(1, total_num_kv_heads // tensor_parallel_size)
# adapted from https://github.com/vllm-project/vllm/blob/v0.6.4.post1/vllm/config.py
def _parse_quant_hf_config(self):
quant_cfg = getattr(self.hf_config, "quantization_config", None)
@@ -1106,6 +1141,11 @@ def is_hybrid_model(
context_length: Optional[int],
attention_chunk_size: Optional[int],
):
if model_architectures[0] in [
"MiMoV2FlashForCausalLM",
"MiMoV2MTP",
]:
return 1
if hybrid_kvcache_ratio is None:
return None
elif (
@@ -1118,7 +1158,11 @@ def is_hybrid_model(
return None
def get_hybrid_layer_ids(model_architectures: List[str], num_hidden_layers: int):
def get_hybrid_layer_ids(
model_architectures: List[str],
num_hidden_layers: int,
hybrid_layer_pattern: Optional[List[int]] = None,
):
if "Llama4ForConditionalGeneration" in model_architectures:
swa_attention_layer_ids = [
i for i in range(num_hidden_layers) if (i + 1) % 4 != 0
@@ -1126,6 +1170,15 @@ def get_hybrid_layer_ids(model_architectures: List[str], num_hidden_layers: int)
full_attention_layer_ids = [
i for i in range(num_hidden_layers) if (i + 1) % 4 == 0
]
elif "MiMoV2FlashForCausalLM" in model_architectures:
swa_attention_layer_ids = [
i for i in range(num_hidden_layers) if hybrid_layer_pattern[i] == 1
]
full_attention_layer_ids = [
i for i in range(num_hidden_layers) if hybrid_layer_pattern[i] == 0
]
elif "MiMoV2MTP" in model_architectures:
return [0], []
else:
swa_attention_layer_ids = None
full_attention_layer_ids = None