📑 目录

一句话结论: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 与表情回应。