📑 目录
一句话结论:RMM 在矩阵乘的收缩维度上按当前激活的列 L2 范数做 TopK 选择,只算保留的切片——不训练、不动权重,一个 retention-ratio 旋钮给出可预测的精度-效率权衡。实测:70B 在保留 80% 时几乎无损、Llama3.1 8B 长序列端到端 1.40× 加速、4096 序列下 70B 模型免 OOM;机制上注意力侧比 MLP 可裁剪得多(Q 投影 RR=0.5 只掉 2pp,整 MLP 掉 29.5pp)。
背景与动机
Transformer 推理算力大头是高维矩阵乘(QK^T、PV、FFN 三投影),但计算大量冗余:注意力得分稀疏、FFN 激活高维稀疏。现有方案两难:
| 路线 | 代表 | 缺陷 |
|---|---|---|
| 训练式稀疏 | SparseGPT/Wanda/SliceGPT | 要改权重、成本高 |
| 静态剪枝 | Magnitude 等 | 与输入无关,跨分布急剧退化 |
RMM 补的空档:不改权重 + 随输入动态裁剪。
核心思路(公式级)
收缩维 TopK 选择
矩阵乘 Y = A·B(A∈ℝ^{n×d} 激活、B∈ℝ^{d×m}),沿收缩维 d 选索引集 ℐ(|ℐ| = ⌈ρd⌉):
RMM_ρ(A,B) = A[:,ℐ] · B[ℐ,:]
重要性度量 = 激活列 L2 范数:s_j = ||A[:,j]||₂,取 TopK 最大的 ⌈ρd⌉ 个。
性质:
- 确定性:同输入同选择;输入自适应:逐层/逐头/逐 token 变化
- minimax 最优(Theorem 1):TopK 列范数在给定预算下最小化任意 B 的最坏近似误差
- 误差界:
||AB − A[:,ℐ]B[ℐ,:]||_F ≤ Σ_{j∉ℐ} ||A[:,j]||₂·||B[j,:]||₂ - 复杂度:O(n·ρd·m)(稠密 O(n·d·m)),列范数 O(n·d) + TopK 开销小
组件应用:QK^T 按头特征维选择(评分 Q 列范数)、PV 按 token 位置选择(可选)、MLP/线性投影按激活隐层维选择;GQA 下在 Q 上按头选择、K/V 收集对应维度。
retention-ratio(ρ)旋钮
ρ∈(0,1] 直接控制保留维度数 ⌈ρd⌉——平滑可预测的权衡。组件差异化:注意力侧可激进(RR 低至 0.5),MLP 需保守且按投影类型区分。无标注数据时可用 ~100 个无标签样本做一致性扫描(Llama-3.1-8B RR=0.7 下 87/100 Wikipedia 段落与稠密序列级完全一致)。
实验数据
规模规律(8 任务 × RR 0.9→0.5)
| 模型 | RR=0.8 表现 | RR=0.5 表现 |
|---|---|---|
| Llama3.1 70B | 接近满性能(MMLU 75.0→72.6) | 仍可用(GSM8K 53.7→19.9 掉但多数任务平缓) |
| Qwen3 32B | 几乎无损(MMLU 80.8→78.6) | 平缓退化 |
| Llama3.1 8B | 轻微下降 | GSM8K 26.2→5.9 明显掉 |
| Qwen3.1 7B | 下降明显 | GSM8K 39.9→1.7 崩 |
规律:模型越大冗余越多、容忍度越高;小模型在 RR=0.7 已现拐点(WikiText 困惑度:Llama3.2-1B RR=0.7 时 20.04→31.29)。
与静态剪枝对比(RR=0.5,Llama3.1 8B,5 QA 平均)
| 方法 | 平均 |
|---|---|
| 全模型 | 69.8 |
| RMM | 59.8 |
| SparseGPT | 56.1 |
| Wanda | 52.7 |
| Magnitude | 39.3 |
| SliceGPT | 37.0 |
注意力 vs MLP:结构性不对称(表 16,8B,5 QA 平均)
| 裁剪目标 | RR=0.9 | RR=0.7 | RR=0.5 |
|---|---|---|---|
| Q 投影 | 69.60 | 70.01 | 67.80(几乎不掉) |
| QKV 投影 | 68.92 | 67.35 | 59.79 |
| 注意力内部(QK^T+PV) | 69.45 | 66.98 | 59.56 |
| 整 MLP | 63.06 | 55.93 | 40.28(暴跌) |
| MLP Up | 65.69 | 59.88 | 52.44 |
| MLP Down | 67.43 | 65.75 | 61.36 |
补充(ARC-Easy RR=0.7 归一化对比):注意力侧掉 3.52 点(保留能量 89.69%)、MLP Up 掉 16.32(82.24%)、MLP Down 掉 3.51(99.02%)、整 MLP 掉 18.78(87.85%)——Down 投影最鲁棒、Up 最敏感、误差会跨投影累积。
长上下文(Ruler,RR=0.5 仍持平)
CWE 5K/15K/30K:98.0/94.0/28.9 vs 基线 98.2/94.0/29.6——裁剪不放大长上下文退化。
A100 实测(ρ=0.8,batch=1)
| 序列长 | QK^T | AV | 端到端(8B) | 端到端(70B) |
|---|---|---|---|---|
| 1024 | 1.36× | 1.67× | 1.05× | 1.03× |
| 2048 | 1.29× | 1.81× | 1.27× | 1.41× |
| 4096 | 1.56× | 1.89× | 1.40× | OOM→可跑 |
规律:序列越长加速越明显(短序列选择开销占比大);70B 在 4096 序列从 OOM 变为可推理——省内存与省延迟双赢。
兼容性与泛化
- 与 INT8 正交:INT8 + RMM(注意力侧 RR=0.8) COPA 81.40→77.40——降精度 × 减算力可叠加
- VLM 泛化:Qwen2.5-VL-7B RR=0.8 几乎无损(POPE 83.7→82.0);InternVL3-8B 到 RR=0.5 仍 92.33 持平
- vs TEAL(激活稀疏):TEAL 只对投影输入激活稀疏、无法裁剪 QK^T/PV 内部矩阵乘;RMM 矩阵乘积视角覆盖面更广
工程落地要点
- 接入:包注意力/FFN 算子,PyTorch 可原型验证;生产需自定义 kernel 才有真实加速
- 配置:注意力侧激进(RR 0.5~0.7)、MLP 保守(0.8+,Down 可更低);预填充裁 FFN、解码裁注意力分开调
- 踩坑:短序列收益小;GSM8K 类强推理任务对裁剪最敏感(掉得最快),数学场景降 RR 要谨慎
- 验证:用 100 个无标签样本一致性扫描快速选 ρ,无需标注集
适用边界与取舍
- 适合:长上下文、批量生成、已上线模型降本、4096+ 序列的内存受限部署;与量化叠加
- 不适合:短序列高并发小 batch(收益被 GEMM 库摊薄);精度严格敏感业务
- 取舍:vs 静态稀疏(动态鲁棒但需 kernel);vs 量化(正交可叠);vs 激活稀疏 TEAL(覆盖矩阵乘范围更广)
复现要点
- arXiv: 2608.13426(8-13,24 页);作者 Zixuan Lan 等,未注明开源仓库
- 复现路径:实现列范数 TopK 切片算子 → 8B 模型跑 RR 曲线 → 长序列 A100 基准
- 注意力侧与 MLP 侧分开配 RR 是关键工程决策点
参与讨论
GitHub Discussions 驱动评论由 GitHub Discussions 驱动,数据存储于 hackcv/blog 仓库;需要 GitHub 账号登录后参与,支持 Markdown 与表情回应。