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 头等),这是推理成本的主要来源之一。已有的推理加速路线大致有四类:

  1. 结构改造(如 MQA/GQA、FlashAttention):改模型架构或重写 kernel,需要训练或重训。
  2. 量化(INT8/INT4/FP8):需要校准、可能改权重,部署链路过半。
  3. 剪枝 / 蒸馏:训练时介入,与"无重训部署"目标不兼容。
  4. 投机解码(speculative decoding):依赖 draft model,不在所有场景都划算。

RMM 要解决的问题是:不动权重、不重训的前提下,从单次矩阵乘法内部做"输入自适应"的减法。其核心观察是,矩阵乘法 Y = A·B 沿着缩并维度(k)做求和时,不同 k 切片对最终结果的贡献并不均匀 —— 如果能在线判别哪些 k 切片"信息密度高",只保留它们就能省算力。

核心方法

RMM 的机制可拆为三层:

  1. 沿缩并维度选切片(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,由调用方控制。

  1. 输入自适应打分(input-adaptive scoring):切片分数由当前输入在线算出(不是预先算好的固定 mask),所以同一层在不同 token 上保留的 k 切片集合不同。这是 "input-adaptive" 的字面含义。

  2. 保留率控制 + 平滑权衡(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。

对工程落地的启发

  1. "不重训加速"是当下性价比最高的优化层:相比量化/蒸馏/重训,RMM 的部署成本是"kernel + 校准 r",可作为模型上线后的第一道优化闸门。
  2. 结构非对称 = 调度策略:attention 用激进 r、MLP 用保守 r 是 immediate actionable 的默认策略,可直接进入 serving config。
  3. 长序列场景的最优解:如果你的服务是文档摘要、长上下文 RAG、代码库级 Copilot,RMM 命中 wall-clock 收益区的概率最高。
  4. 与现有优化正交:RMM 与量化、投机解码、KV 缓存压缩都不冲突,理论上可叠加;这是它最容易"塞进现有 stack"的地方。
  5. 打分函数实现是核心 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。在此之前,适合做技术调研和可行性评估,不适合直接上生产。