RMM:把 Transformer 矩阵乘法按"输入自适应"砍一刀
- 关联论文:2608.13426
- 作者:flyP
- 更新:2026-08-15
一句话结论:RMM(Reduced Matrix Multiplication)是一个训练免、改权免、输入自适应的推理加速方法 —— 它在矩阵乘法的缩并维度(contraction dimension)上按 token 动态挑选"信息量大的切片"参与计算,其余切片丢掉;在不修改模型权重的条件下,1B–70B 模型上获得"精度-效率"可预测的平滑权衡,长序列场景下落到 A100 自定义 kernel 上能拿到实际 wall-clock 收益。
解决什么真问题
Transformer 类语言模型在推理时反复跑高维矩阵乘法(Q/K/V 投影、MLP 中间层、logits 头等),这是推理成本的主要来源之一。已有的推理加速路线大致有四类:
- 结构改造(如 MQA/GQA、FlashAttention):改模型架构或重写 kernel,需要训练或重训。
- 量化(INT8/INT4/FP8):需要校准、可能改权重,部署链路过半。
- 剪枝 / 蒸馏:训练时介入,与"无重训部署"目标不兼容。
- 投机解码(speculative decoding):依赖 draft model,不在所有场景都划算。
RMM 要解决的问题是:不动权重、不重训的前提下,从单次矩阵乘法内部做"输入自适应"的减法。其核心观察是,矩阵乘法 Y = A·B 沿着缩并维度(k)做求和时,不同 k 切片对最终结果的贡献并不均匀 —— 如果能在线判别哪些 k 切片"信息密度高",只保留它们就能省算力。
核心方法
RMM 的机制可拆为三层:
- 沿缩并维度选切片(slice selection along contraction dim):对 Y = A·B 中的缩并维 k,按某种重要性分数只保留 top-r% 的 k 切片;其余切片置零。形式上:
Y_full = sum_k a[:, k] * b[k, :]
Y_rmm = sum_{k in S} a[:, k] * b[k, :], S ⊂ {0..K-1}, |S|/K = r
r 是 retention ratio,由调用方控制。
-
输入自适应打分(input-adaptive scoring):切片分数由当前输入在线算出(不是预先算好的固定 mask),所以同一层在不同 token 上保留的 k 切片集合不同。这是 "input-adaptive" 的字面含义。
-
保留率控制 + 平滑权衡(retention-ratio control):r 是单一旋钮 —— 调小 r → 算力下降、可能掉点;调大 r → 接近原模型精度。论文强调"a smooth and predictable accuracy-efficiency trade-off",意味着在某个 r 区间内精度几乎不掉,超过阈值才掉。
结构非对称发现(mechanistic ablation):
"attention-side computations are substantially more reducible than MLP components."
即 attention 路径的矩阵乘可以砍掉更多切片而不掉点,MLP 路径更敏感。这意味着实际部署时对 attention 层用更激进的 r,对 MLP 层用更保守的 r,可作为默认调度策略。
伪代码(简化):
def rmm_matmul(A, B, r=0.8):
# A: [M, K], B: [K, N]
scores = importance_per_k(A, B) # 输入自适应打分
keep = topk(scores, int(K * r)) # 保留 r 比例的 k 切片
return A[:, keep] @ B[keep, :]
⚠️ 关键工程细节原文未明确:importance_per_k 的具体形式(基于 A 的范数?基于 A·B 子乘积的 L2?基于 attention 分数?)在 abstract 层级未给出;落地时需查正文 §3。
关键实验与数据
- 规模跨度:1B 到 70B 参数的语言模型,覆盖多个模型族。
- 任务覆盖:判别式任务、自回归生成长文本、长上下文设置、多模态 VLM 推理。
- 核心趋势:在中等 reduction(⚠️ 原文未给出单一具体 r 数值 —— 建议查正文表 1 取代表性数字)下,RMM 在四类设置上都保持鲁棒。reduction tolerance 随模型规模上升(更大的模型更耐砍)。
- 消融:attention 路径可减比例显著高于 MLP(结构性非对称)。
- Wall-clock:在 NVIDIA A100 上用自定义 kernel 验证,长序列下能转化为实际运行时间收益。短序列收益小或可忽略。
- ⚠️ 未量化数字:具体加速比(如 1.4× / 1.7× / 2.1×)、各模型族最优 r、A100 长序列下"长"的具体阈值,abstract 层面未公开。落地前必须查正文表 1 与 §5。
亮点与局限
亮点
- 完全无重训:不改权重、不量化、不蒸馏,对已部署模型零侵入。
- 单一旋钮 r:调一个超参就能在精度-效率曲线上滑动,工程友好。
- 结构非对称发现:attention 路径的冗余度比 MLP 路径高得多,给后续工作留了清晰的优化方向。
- 多模态外延:同一原则在 VLM 推理上仍生效,跨模态泛化有据可查。
- 长序列 winner:在 A100 + 长序列下能转化为 wall-clock 收益,正好打在 LLM 推理最痛处。
局限 / ⚠️ 风险边界
- 依赖自定义 kernel:要拿到实际 wall-clock 收益,需要在推理框架里实现切片选择 + 不规则矩阵乘 kernel;标准 cuBLAS 直接用不上。
- 打分开销:input-adaptive 的代价是每 token 都要算一次 importance;短序列或 batch 很大时,这部分开销可能吃掉收益。
- 最优 r 与模型强耦合:从 abstract 表述看,r 与"模型族 + 任务 + 组件 + 保留率"四元组都相关,没有"通用 r";每个部署都要重新校准。
- A100 之外未证:abstract 显式说"on an NVIDIA A100",H100/B200/RTX 50 系上的等效收益未在 abstract 给出(⚠️ 原文未明确)。
- 极端 reduction 必掉点:本质是有损方法,关键场景(高风险决策、医疗/法律)不建议开激进 r。
对工程落地的启发
- "不重训加速"是当下性价比最高的优化层:相比量化/蒸馏/重训,RMM 的部署成本是"kernel + 校准 r",可作为模型上线后的第一道优化闸门。
- 结构非对称 = 调度策略:attention 用激进 r、MLP 用保守 r 是 immediate actionable 的默认策略,可直接进入 serving config。
- 长序列场景的最优解:如果你的服务是文档摘要、长上下文 RAG、代码库级 Copilot,RMM 命中 wall-clock 收益区的概率最高。
- 与现有优化正交:RMM 与量化、投机解码、KV 缓存压缩都不冲突,理论上可叠加;这是它最容易"塞进现有 stack"的地方。
- 打分函数实现是核心 IP:谁把 importance scoring 做便宜做准,谁就拿到 RMM 的工程红利。
与同方向工作的关系
- 稀疏注意力 / 滑动窗口注意力(Longformer、BigBird、StreamingLLM):与 RMM 同属"砍掉部分计算"家族,但 RMM 砍的是缩并维度(k 维)而非 token 维度,应用面更广(不限于 attention)。
- 激活稀疏化(Activation Outliers 系列、DejaVu):与 RMM "MLP 路径冗余更少" 的发现同源,但 DejaVu 关注"剪掉激活通道",RMM 关注"剪掉 k 切片",粒度更细。
- 训练免推理优化(Speculative Decoding、Early Exit、SkipDecode):同一阵营的不同切入;RMM 是"输入粒度的算力重分配"。
- 结构化矩阵低秩(LoRA-as-inference、SVD-pruning):与 RMM 不同 —— RMM 不改权重,结构化低秩改权重;RMM 是"推理时算力重分配",低秩是"权重静态降维"。
适合谁读
- LLM serving 工程师:拿 RMM 作为上线后第一道推理加速闸门,关注长序列场景。
- kernel 开发者:实现 r-controlled 不规则 matmul kernel 是直接落地点。
- 推理框架维护者(vLLM / TGI / TensorRT-LLM 社区):评估是否值得把 RMM 作为官方 plugin。
- 多模态 VLM 团队:方法在 VLM 上同样生效,跨模态泛化有据可查。
- 学术研究者:RMM 揭示的"attention 比 MLP 更可砍"是后续工作的明确入口(为什么 attention 更冗余?这是新问题)。
⚠️ 诚实标注:本篇解读基于 arXiv abstract (2608.13426v1) + 关联 paper_card。具体加速比数字、importance scoring 形式、A100 以外硬件表现、各模型族最优 r 等定量细节abstract 未给出,落地前请查正文表 1、§3(机制)与 §5(wall-clock 基准)。
工程落地与核查(Jay)
事实核查结果
| 核查项 | 结论 | 存疑级别 |
|---|---|---|
| arXiv 2608.13426 存在 | ✅ 核实,标题吻合 | — |
| 1B–70B 规模验证 | ✅ abstract 原文有描述;未经独立复现 | — |
| NVIDIA A100 wall-clock 收益 | ✅ abstract 明确说 A100;⚠️ 其他硬件未验证 | 中 |
| attention 比 MLP 更可砍 | ⚠️ abstract 引述为结论性语句;具体 r 值需正文消融表 | 中 |
| 具体加速比 | ⚠️ abstract 未给出任何数字;需正文表 1 | 高 |
| importance_per_k 函数形式 | ⚠️ 全文未披露;kernel 实现无参考 | 高 |
| GitHub 代码仓库 | ⚠️ 未 fetch 核实;abstract 无链接 | 高 |
| vLLM / TGI 集成状态 | ⚠️ abstract 未提;无官方 plugin 声明 | 高 |
实际系统怎么用
⚠️ 前置条件:必须写自定义 CUDA kernel
RMM 的切片选择 + 不规则矩阵乘法在 A100 上拿到 wall-clock 收益的前提是跳过 cuBLAS 的 GEMM 路径,直接实现 A[:, keep] @ B[keep, :] 的原子操作。标准 vLLM / TGI / TensorRT-LLM 不用自定义 kernel,直接跑不到收益。
概念验证级 Python 实现(不可用于生产):
import torch
import torch.nn.functional as F
def importance_per_k_default(A: torch.Tensor, B: torch.Tensor) -> torch.Tensor:
"""
⚠️ 原文未披露实际实现;以下为合理猜测(基于激活范数思路)。
真实实现需读正文 §3 确认。
"""
# 猜测方案A:基于 A 的逐列 L2 范数
# cols_norm = torch.norm(A, dim=0) # [K]
# 猜测方案B:基于 A·B 子乘积的 L2
# sub = A[:, :512] @ B[:512, :] # 这会引入额外计算,不划算
# 猜测方案C:基于 attention score 的统计量
# return torch.mean(torch.abs(A), dim=0) * torch.mean(torch.abs(B), dim=1)
raise NotImplementedError("importance_per_k() form undisclosed in abstract; see paper §3")
def rmm_matmul(A, B, r=0.8):
K = A.shape[-1]
keep_k = max(1, int(K * r))
# scores = importance_per_k_default(A, B) # TODO: confirm scoring form
# keep_idx = torch.topk(scores, keep_k).indices
# return A[:, keep_idx] @ B[keep_idx, :]
raise NotImplementedError("Must implement custom kernel; cuBLAS path will not yield speedup")
生产级 kernel 实现路径(CUDA C++/PTX):
// RMM GEMM kernel 伪代码(概念级)
__global__ void rmm_gemm_kernel(
const float* __restrict__ A, // [M, K]
const float* __restrict__ B, // [K, N]
float* __restrict__ C, // [M, N]
const int K, const float r,
const int* __restrict__ keep_idx // 预计算的 top-r% 索引
) {
int m = blockIdx.x * blockDim.x + threadIdx.x;
int n = blockIdx.y * blockDim.y + threadIdx.y;
if (m >= M || n >= N) return;
float sum = 0.0f;
for (int ki = 0; ki < keep_k; ++ki) {
int k = keep_idx[ki]; // 不连续内存访问,需要专门优化
sum += A[m * K + k] * B[k * N + n];
}
C[m * N + n] = sum;
}
⚠️ 坑点 #1:不连续内存访问抵消切片收益
A[:, keep_idx] 和 B[keep_idx, :] 产生非连续内存访问(strided/gather)。若 keep_k=0.8K,则每行仍有 20% 的列被跳过,但访问模式变成 gather → 内存带宽可能不降反升。需要 kernel 级 swizzle / tiling 优化才能真正省带宽。
⚠️ 坑点 #2:打分开销可能吃光收益
importance_per_k 若需要一次完整的 A·B 子乘积(大小 M×K × K×N),则每 token 要跑一次额外矩阵乘,等于 overhead = O(M·K·N) 而节省 = O(M·K·N·(1-r))。若 r=0.8,节省 20% 但开销 100%,得不偿失。必须确保打分函数的复杂度 << O(M·K·N)。
⚠️ 坑点 #3:attention vs MLP 分层调度需精确实现 原文说 attention 层比 MLP 层更可砍,但实际操作时需要逐层识别当前在跑的是 attention 投影还是 MLP。在 vLLM 的 fused kernel 里两者通常是 fused 在一起的,分离需要修改 kernel 发射逻辑,不是简单的 r 值调整。
⚠️ 坑点 #4:H100/B200 上收益未知 A100 的 tensor core 架构 vs H100 的 Hopper 架构差异显著。Hopper 对 irregular sparse GEMM 的支持更好(Transformer Engine 支持 dynamic sparsity),但 abstract 未报告 H100 数据,不能假设收益可移植。
⚠️ 坑点 #5:短序列几乎无收益 若序列长度 < 512 tokens,wall-clock 收益可能 < 5%,而打分开销固定。RMM 只适合长序列场景(文档摘要/代码库/RAG+长 context)。
与推理框架集成
vLLM / TGI / TensorRT-LLM
↓ 原生 GEMM 路径
Custom RMM GEMM Plugin (CUDA kernel)
├→ importance scoring (per token)
├→ topk keep indices
└→ irregular gather + GEMM
↓
Output (speedup if long seq + A100)
当前主流推理框架(vLLM 0.6.x / TGI 2.x / TensorRT-LLM)均无 RMM 官方支持。集成需要: 1. 等官方插件(若论文团队与框架 maintainer 合作) 2. 自写 vLLM custom operator(需 fork vLLM + 写 CUDA kernel + 通过 benchmark CI) 3. 实际上线时间:最快 6–12 个月(包括 debug + benchmark 验证 + upstream PR merge)
总结
RMM 工程可行性中等偏低 —— 核心障碍是无开源 kernel 实现 + 依赖自定义 CUDA 代码。单一旋钮 r + 无重训 + attention/MLP 非对称调度是极具吸引力的工程价值,但 importance_per_k 形式未公开、A100 以外硬件无数据、主流推理框架无集成,使其短期内难以直接落地。最快路径:等待正文 §3 披露 scoring 函数 + 等 vLLM/TGI 官方 plugin。在此之前,适合做技术调研和可行性评估,不适合直接上生产。