MiMo-V2.5 推理优化解读:Hybrid SWA 的工程落地

MiMo-V2.5 推理优化解读:Hybrid SWA 的工程落地

小米 MiMo 团队的《Full-Pipeline Inference Optimization for MiMo-V2.5 Series》(arXiv:2607.13095)核心论点很直白:Hybrid SWA 理论上能把 KVCache 和 attention 计算量压到 Full Attention 的 1/7,但如果 KVCache 系统不跟着改造,实际反而会变慢

本文不写综述,重点拆解几个关键工程问题:SWA 的 KV cache 怎么存、怎么 offload、共享前缀命中时的重算问题、numa_balancing 为什么 +10%、以及 Prefill 为什么也要开 MTP


1. 论文速览

MiMo-V2.5-Pro 的架构组合拳:

1
2
Hybrid SWA + 稀疏 MoE + 多模态编码器
70 层 = 10 Full Attention + 60 SWA(窗口 W=128)

理论上 attention FLOPs 和 KVCache 存储都降到 1/7,但论文的核心观点是:这些理论红利不会自动兑现。Hybrid SWA 在 KVCache 管理、前缀匹配、Full 和 SWA 层语义一致性上都引入了新问题。

论文涉及的优化面很广:KVCache 分层(L1 GPU / L2 Host / L3 GCache)、调度(LLM-Router)、Prefill/Decode pipeline、多模态等。其中 LLM-Router 和 GCache 属于常见的分布式缓存和调度实现,本文不展开。重点聚焦在 SWA 相关 KVCache 这条主线上。


2. SWA 的 KV Cache 怎么存

核心设计:物理分离 + 逻辑统一

物理层:两个独立池子

每层存储(L1 GPU / L2 Host / L3 GCache)都并排维护两个 KV Pool:

1
2
3
4
5
6
7
8
9
┌─────────────────────────────────────┐
│ Full KV Pool │
│ ├─ 大小: O(seq_len × 10 层) │
│ └─ 淘汰: 按整个序列做 LRU │
├─────────────────────────────────────┤
│ SWA KV Pool │
│ ├─ 大小: O(W × 60 层) ← 严格 O(W) │
│ └─ 淘汰: 独立按 window 做 eviction │
└─────────────────────────────────────┘

关键点:SWA pool 物理上就只有 W 大小,写入新 token 时按窗口 evict 旧的。不是「存完整历史然后只用尾部 128 个」,而是存进来就已经是窗口截断后的样子

直观对比(128K token 序列):

Full Attention Hybrid SWA
60 个 SWA 层总量 60 × 128K = 7.68M 60 × 128 = 7.68K
10 个 Full 层总量 10 × 128K = 1.28M 10 × 128K = 1.28M
合计 8.96M ≈ 1.29M(省 7×)

Hybrid SWA 存储量几乎等于只算 Full 层的量,SWA 层几乎不占空间。

逻辑层:单序列视图 + 双索引前缀树

上层(前缀树、调度器)只看到一个逻辑序列,底下用 Full→SWA mapping 做透明分层。每个前缀树节点存两套元数据:

1
2
3
4
5
PrefixTreeNode {
tokens: [t0, t1, ..., tN]
full_seg_idx: [Full KV pool 位置] # Full 层复用
swa_seg_map: [SWA KV pool 位置 or 空] # 判断 window 安全性
}

用途

  1. 命中判定:token 相等 + tail W 个 token 的 swa_seg_map 都非空 → 才算真命中
  2. 独立淘汰:window 外的 SWA 段可以先扔,Full 段还留着

一句话总结:SWA 存的就是「每 SWA 层永远只存 W 个 token 的 KV」,SWA pool 和 Full pool 物理分离,上层用双索引前缀树抽象成单序列视图。


3. SWA Cache 怎么 Offload 到 CPU 内存

存什么:只存「窗口内有效 slot」

对一条 100K token 的序列:

1
2
Full 层 (10 层):  每层 100K 个 slot  ← 全部下沉到 CPU
SWA 层 (60 层): 每层只有 W=128 个 slot ← 只存窗口内

SWA 下沉到 CPU 的就是窗口截断后的形态,不是完整历史。

怎么存:Host 侧独立的 SWA Pool

L2 结构完全镜像 L1:

1
2
3
Host DRAM (L2):
Full KV Pool: [10 层, seq_len, head_dim], paged + pinned memory
SWA KV Pool: [60 层, W=128, head_dim], 独立 eviction

工程要点:

  • Pin memory:Host 池子用固定内存,D2H/H2D 用异步 DMA
  • Paged 分配:SWA pool 拆成 block,滑出窗口的 block 归还池子
  • 和 Full pool 完全独立:分开的地址空间、eviction 策略

怎么搬:按 SWA mask 传输

论文 §3.1.1 的关键一句:

“Cross-tier transfers are performed based solely on the SWA mask, ensuring only valid window data is moved.”

D2H / H2D 时不搬窗口外的数据

1
2
3
D2H writeback: 只把 mask=1 的 slot 打包 DMA 到 host
H2D prefetch: 只拉窗口内的 W 个 token 的 KV
→ 数据量小,layerwise prefetch 能 overlap 计算

这里省的是带宽。按 mask 走,60 层每层就 128 个 token,总量小到可以忽略。

D2H 时机

  • Eviction 前的备份:GPU SWA pool 满了,先 D2H 到 host 再释放
  • 一致性修复(§3.1.4):前缀树节点合并 / prefill chunk 完成时,检查 device 和 host 的 SWA 占用差 → host 补齐槽位 → 异步 D2H
  • 请求结束 / 会话切换:高价值会话主动 D2H 到 host,甚至下沉到 L3

一句话总结:SWA 下沉到 CPU 的就是窗口截断后的 128 个 slot,Host 侧用独立 pool + pinned memory 存,搬运只按 SWA mask 走有效位置。


4. 共享前缀命中时,SWA 要不要重算

这是 Hybrid SWA 落地生产最难受的一个点

问题本质

典型场景:共享 system prompt(2000 tokens),下面挂了很多用户会话:

1
2
3
4
PrefixTree:
system_prompt (2000 tokens) ← 非叶子节点
├─ user_A_session (往下延展了 50K tokens)
└─ user_B_session (往下延展了 30K tokens)

关键观察:user_A 会话延展了 50K tokens,当前窗口在 [52000-128, 52000] 附近。system_prompt 段(位置 0-2000)早就滑出窗口了,SWA KV 被 evict 掉了。

这就是「伪命中」

user_D 带着同一个 system prompt 进来:

1
2
3
4
5
命中判定:
1. Token equality 检查 → 匹配到 2000 tokens ✅
2. Full KV 检查 → 都在 → Full 命中 2000 tokens ✅
3. Window-safe 检查 → tail 128 个 SWA slot 在不在?
└─► ❌ 已被 evict

按「window-safe length」规则,SWA 匹配长度被 clip。Full 层能复用 2000 tokens,但 SWA 层不行——SWA 层要看窗口内 KV,位置 2001 的 attention 需要 [1873, 2001] 的 KV,这段没了 → 只能重算 SWA 层的 prefill。

代价

  • Full 层:省了(Full pool 还留着)
  • SWA 层:60 层全部重跑一次 prefill(占 6/7)
  • 60 层重算,几乎等于白干

小米的应对:SWA 保留策略(§3.1.4 第 4 条)

论文直接点名了这个场景:

“Medium/short sequence SWA retention strategy. Based on user request patterns, we retain relatively dense SWA KV Cache at fixed length positions for medium/short sequences… particularly beneficial for long agent sessions, multi-user shared system prompts, and repeated tool calls to the same codebase.

翻译:对高频复用的前缀(比如 system prompt),SWA pool 不严格 O(W),而是在关键位置留几份 SWA 快照

推测的实现方式:

1
2
3
4
5
6
7
普通序列的 SWA pool:
只留 tail W 个 slot ← 严格 O(W)

高频前缀(如 system prompt):
在几个 anchor 位置多留 W 个 slot
比如: [0, W], [500-W, 500], [1000-W, 1000], [2000-W, 2000]
← 相当于"检查点",每个都能作为 SWA 计算的起点

这样 user_D 命中时:匹配到 system_prompt 末尾 tail W=128 的 SWA slot 确实存在(是保留下来的 anchor),Full + SWA 都能完整复用 → 不需要重算 SWA 层

Trade-off 账

严格 O(W) 高频前缀密集保留
SWA 存储占比 极低(1/7) 略高(几倍 W)
单节点并发能力 稍低
共享前缀命中率 差(伪命中) 好(真命中)

核心思路用一点点 SWA 存储换取显著提升的共享前缀命中率。因为 system prompt 被成千上万用户共享,每次省掉 60 层 SWA prefill 的收益,远远大于多存几份 128-slot 快照的代价。

一句话总结:非叶子共享前缀命中时,默认严格 O(W) 实现会因为 tail 被 evict 导致 SWA 层必须重算。小米的解法是对高频共享前缀主动保留多个 SWA 快照,让 Full + SWA 都能完整复用,真正兑现 Hybrid SWA 在多用户共享场景下的效率红利。


5. numa_balancing 为什么影响 +10%

论文里就一句话:「关掉 numa_balancing,端到端 +10%」。背后是操作系统调度和 GPU 推理的碰撞

numa_balancing 是干嘛的

现代服务器都是 NUMA 架构,CPU 访问本 NUMA 节点内存 << 访问远端节点内存(延迟差 1.5-2×)。kernel.numa_balancing=1 是 Linux 的「自动 NUMA 均衡」特性:

1
2
3
4
5
6
7
内核周期性做的事:
1. 扫描进程页表 → 把页面标记为 PROT_NONE
2. 进程下次访问 → 触发缺页异常
3. 内核在 fault handler 里统计 NUMA locality
4. 如果 CPU 和 page 不在同一 node:
- 迁移页面到 CPU 所在 node
- 或者迁移线程到 page 所在 node

但在 GPU 推理场景下变成灾难

SGLang 有自己的 --numa-node 配置,会主动把 rank 绑到指定 NUMA node(PIN 死)。冲突来了

  • SGLang 说:这个 rank 就吃 Node 0,别乱跑
  • 内核 numa_balancing 说:我来「优化」一下

内核会做几件破坏性的事:

  1. 周期性把页面标 PROT_NONE:即使你已经手动 pin 好了,下一次访问必触发 page fault
  2. Page fault 引起大 stall:GPU 计算 → 触发 page fault → 内核处理 → GPU 继续(延迟数百 μs 到 ms)
  3. 随机页面迁移:pinned memory 被迁移 → GPU 之前记录的 DMA 地址失效

为什么这个 bug 特别难缠

论文里这句话说到点子上了:

“In multi-node multi-GPU deployments, these gaps appear at random positions across ranks, and each inter-rank synchronization is bottlenecked by the slowest rank.”

关键组合拳

  • 随机性:numa_balancing 是周期性异步扫描,不同 rank、不同时刻会随机中招
  • 全局同步放大伤害:LLM 推理每层 attention 之后都要 all_reduce最慢那个 rank 拖累所有人
1
2
3
所有 rank 计算 ──► 集合通信同步 ──► 下一层

最慢那个 rank 拖累所有人

如果 8 个 rank 里有 1 个刚好被 page fault 命中,那 8 个 rank 全都得等它——尾延迟灾难。对一次前向 60+ 层,只要偶尔一层有一个 rank 中招,整个 latency 就翻倍。

为啥关掉直接 +10%

  1. SGLang 已经手动 pin 好了 NUMA,调度器安排得明明白白
  2. numa_balancing 是「帮倒忙」,它假设进程没做优化,主动帮你搬,反而破坏了原有布局
  3. 代价是 hard cost:page fault、TLB shootdown、页面迁移都是不可避免的 CPU 侧开销

关掉后:内核不再主动折腾 → 页表稳定 → 没有意外 page fault → GPU kernel 之间没有随机 gap → 集合通信不再被拖尾 → 直接 +10%。

一句话总结:SGLang 已经手动做了 NUMA pin,numa_balancing 在这基础上做二次干扰,通过 page fault + 页面迁移引入随机 stall,被 GPU 集合通信同步放大成尾延迟灾难,关掉就直接 +10%。


6. Prefill 为什么也要开 MTP

默认直觉是「MTP 是加速 decode 的,跟 prefill 有啥关系」,但正是这个直觉害惨了 agentic 场景。

MTP 是什么

MTP (Multi-Token Prediction) = 多 token 预测。MiMo-V2.5 系列原生支持 3 层 MTP

1
2
3
4
主模型: 预测第 t+1 个 token
MTP Layer 1: 同时预测第 t+2 个 token
MTP Layer 2: 同时预测第 t+3 个 token
MTP Layer 3: 同时预测第 t+4 个 token

核心思想:一次 forward 出 4 个 token,然后主模型验证,接受的部分直接输出。理想情况下 decode 速度接近 4×。

关键前提:MTP 层本身有参数、有 KV cache,需要**跟着上下文一起「预热」**才能预测准。

原始问题:Prefill 不跑 MTP 的后果

默认实现:Prefill 阶段只跑主模型,跳过 MTP 层。

问题:MTP 层需要看历史 context 才能预测。如果 prefill 跳过 MTP:

1
2
3
4
5
6
7
Prefill 结束时:
主模型 KV cache: [完整 prompt 的 KV] ✅
MTP 层 KV cache: [空 / invalid] ❌

Decode 第 1 个 token:
主模型: ✅
MTP 层: context 是空的,瞎猜 → 接受率极低 ❌

MTP 需要多久「热身」:论文说是 128 个 token。因为 MTP 层需要逐个 decode 的 token 慢慢累积 KV。

这段时间 MTP 基本白搭:主模型每步生成 1 个 token,MTP 提议的 3 个 token 因为 context 不足几乎全被拒,有效加速 ≈ 1×。

为什么 agentic 场景特别惨

论文这句话点破了关键:

Since agentic scenarios involve mostly short output sequences, this limitation significantly limited MTP’s effective speedup.”

Agentic 场景的输出特征:

1
2
3
一次工具调用响应: ~20-50 tokens
一次 function call: ~30-80 tokens
一次思考步骤: ~50-150 tokens

几乎所有输出都在 128 token 之内——这正好是 MTP 需要热身的窗口。结果:

1
2
理论加速比: 3× (3层 MTP)
Agentic 场景实际加速比: ~1× (刚热身完就结束了)

等于 MTP 层白装了

解法:Prefill 阶段也跑 MTP

小米的改动很直接:prefill 时也让 MTP 层参与前向

1
2
3
4
5
6
7
Prefill 结束时:
主模型 KV cache: [完整 prompt 的 KV] ✅
MTP 层 KV cache: [完整 prompt 的 MTP KV] ✅ ← 新增

Decode 第 1 个 token:
主模型: ✅
MTP 层: 已经用整个 prompt 预热过,接受率立即很高 ✅

效果(论文数据)

  • 0-128 token 加速: 2.3× ← 从近 1× 提升到 2.3×,agentic 场景直接受益
  • 128-256 token 加速: 1.5×

工程适配

论文说:

“By introducing MTP support during prefill with dedicated adaptations and optimizations for HiCache L2/L3…”

关键点:MTP 层有自己的 KV cache,之前 HiCache 那套 offloading 系统只管主模型的 KV。要 prefill 阶段跑 MTP,就得:

  1. MTP KV 也要进 HiCache:L1 (GPU) ↔ L2 (Host) ↔ L3 (GCache)
  2. Prefix cache 得覆盖 MTP KV:前缀树节点要标 MTP 层的状态
  3. 传输和存储成本:MTP 有 3 层,KVCache 总量变多,L2/L3 得扛住

Trade-off 账

维度 不开 prefill MTP 开 prefill MTP
Prefill 计算 70 层 70 + 3 = 73 层(多 4%)
Decode 0-128 tokens 加速 ~1× 2.3×
Agentic 场景总吞吐 显著提升

核心逻辑用 prefill 的少量额外开销,换 decode 从第 1 个 token 就享受 MTP 加速。Agentic 场景短输出多,这笔账非常划算。

具体场景:Agent 调用工具,输入 5K token,输出 60 token

1
2
3
没开 prefill MTP: 100ms (prefill) + 1200ms (decode) = 1300ms
开了 prefill MTP: 104ms (prefill) + 520ms (decode) = 624ms
省 52%

一句话总结:Prefill 开 MTP 是为了给 MTP 层的 KV cache 做「上下文预热」,避免 decode 初期因为 MTP 无历史导致接受率极低。Agentic 场景输出短,正好都落在这个「未热身窗口」里,所以收益特别大——0-128 token 加速直接从 ~1× 提升到 2.3×。


总结

这篇论文的核心贡献不是某个单点优化,而是把 Hybrid SWA + MoE + 多模态这套组合架构从理论拉到生产的完整工程实践。本文聚焦在 SWA 相关 KVCache 这条主线,拆解了几个关键问题:

  1. SWA 存储:物理分离 Full/SWA pool,SWA 严格 O(W),前缀树双索引做透明分层
  2. SWA offload:Host 侧独立 SWA pool + pinned memory,搬运只按 mask 走有效位置
  3. 共享前缀命中:默认严格 O(W) 会伪命中,小米用「SWA 快照保留」策略换真实命中率
  4. numa_balancing:操作系统自动优化 vs 手动 NUMA pin 的冲突,关掉 +10%
  5. Prefill MTP:给 MTP 层 KV cache 做上下文预热,agentic 短输出场景 0-128 token 从 ~1× 提升到 2.3×

可迁移的启示

  • 架构红利不会自动兑现,工程系统必须跟着架构一起改造
  • Hybrid SWA 的价值在多用户共享场景,但需要专门的缓存策略才能兑现
  • 推理系统的 pipeline 各阶段不是孤立的,跨阶段依赖的组件需要协同优化
  • **系统级坑(numa_balancing、THP、CPU governor)**对已经手动优化过的高性能场景是纯干扰,该关就关

参考

  • 论文:MiMo Team, Xiaomi. Full-Pipeline Inference Optimization for MiMo-V2.5 Series: Pushing Hybrid SWA Efficiency to the Limit. arXiv:2607.13095, 2026.
  • 模型:Xiaomi MiMo Team. MiMo-V2.5 / MiMo-V2.5-Pro. HuggingFace, 2026.

本文基于 MiMo-V2.5 论文、个人笔记《MiMo-V2.5 SWA offload 方法》整理。