[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:
@@ -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. "
|
||||
|
||||
Reference in New Issue
Block a user