[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
@@ -340,6 +340,7 @@ class FlashAttentionBackend(AttentionBackend):
self.full_to_swa_index_mapping = (
model_runner.token_to_kv_pool.full_to_swa_index_mapping
)
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.topk = model_runner.server_args.speculative_eagle_topk or 0
self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens = (
@@ -792,6 +793,15 @@ class FlashAttentionBackend(AttentionBackend):
cu_seqlens_k = swa_spec_metadata.cu_seqlens_k
else:
page_table = metadata.page_table
if self.is_hybrid_swa:
_, is_swa = forward_batch.token_to_kv_pool.layers_mapping[
layer.layer_id
]
if is_swa:
page_table = self.token_to_kv_pool.translate_loc_from_full_to_swa(
page_table
)
window_size = (self.attention_chunk_size, 0)
cu_seqlens_q = metadata.cu_seqlens_q
cache_seqlens = metadata.cache_seqlens_int32
max_seqlen_q = metadata.max_seq_len_q
@@ -807,7 +817,7 @@ class FlashAttentionBackend(AttentionBackend):
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
)
value_cache = value_cache.view(
-1, self.page_size, layer.tp_v_head_num, layer.head_dim
-1, self.page_size, layer.tp_v_head_num, layer.v_head_dim
)
if layer.is_cross_attention:
page_table = metadata.encoder_page_table
@@ -1098,7 +1108,7 @@ class FlashAttentionBackend(AttentionBackend):
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
)
value_cache = value_cache.view(
-1, self.page_size, layer.tp_v_head_num, layer.head_dim
-1, self.page_size, layer.tp_v_head_num, layer.v_head_dim
)
if layer.is_cross_attention:
@@ -1143,6 +1153,17 @@ class FlashAttentionBackend(AttentionBackend):
)
else:
page_table = metadata.page_table
if self.is_hybrid_swa:
_, is_swa = forward_batch.token_to_kv_pool.layers_mapping[
layer.layer_id
]
if is_swa:
page_table = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
page_table
)
)
window_size = (self.attention_chunk_size, 0)
cache_seqlens = metadata.cache_seqlens_int32
cu_seqlens_k = metadata.cu_seqlens_k
max_seqlen_q = metadata.max_seq_len_q
@@ -1743,7 +1764,7 @@ class FlashAttentionBackend(AttentionBackend):
self.target_verify_metadata_topk_swa[bs] = metadata_swa
metadata.swa_spec_metadata = metadata_swa
elif forward_mode.is_draft_extend():
elif forward_mode.is_draft_extend(include_v2=True):
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs
]
@@ -2048,6 +2069,54 @@ class FlashAttentionBackend(AttentionBackend):
]
metadata.page_table[:, :max_seq_pages].copy_(page_indices // self.page_size)
elif forward_mode.is_draft_extend_v2():
metadata = self.draft_extend_metadata[bs]
metadata.cache_seqlens_int32.copy_(seq_lens)
metadata.max_seq_len_k = seq_lens_cpu.max().item()
metadata.cu_seqlens_k[1:].copy_(
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
)
extend_seq_lens_tensor = getattr(spec_info, "extend_seq_lens_tensor", None)
extend_seq_lens_cpu = getattr(spec_info, "extend_seq_lens_cpu", None)
if extend_seq_lens_tensor is not None:
extend_seq_lens = extend_seq_lens_tensor.to(torch.int32)
elif extend_seq_lens_cpu is not None:
extend_seq_lens = torch.as_tensor(
extend_seq_lens_cpu,
dtype=torch.int32,
device=device,
)
else:
default_extend = getattr(
spec_info, "num_tokens_per_batch", self.speculative_num_steps + 1
)
extend_seq_lens = torch.full(
(bs,), default_extend, dtype=torch.int32, device=device
)
extend_seq_lens_cpu = [default_extend] * bs
if extend_seq_lens_cpu:
metadata.max_seq_len_q = int(max(extend_seq_lens_cpu))
else:
metadata.max_seq_len_q = getattr(
spec_info, "num_tokens_per_batch", self.speculative_num_steps + 1
)
metadata.cu_seqlens_q[1:].copy_(
torch.cumsum(extend_seq_lens, dim=0, dtype=torch.int32)
)
max_seq_pages = (
metadata.max_seq_len_k + self.page_size - 1
) // self.page_size
page_indices = self.req_to_token[
req_pool_indices[:, None],
self.draft_extend_metadata["strided_indices"][:max_seq_pages],
]
metadata.page_table[:, :max_seq_pages].copy_(page_indices // self.page_size)
if encoder_lens is not None:
# Only support encoder size 1 for now
metadata.encoder_max_seq_len_k = encoder_lens[0]