[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
parent a0985dd5e5
commit 160a06cab2
38 changed files with 5396 additions and 169 deletions

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. "