
LLM 推理越接近逐 token 生成阶段越能看清一个事实系统大部分时间都耗在矩阵乘法上。Reduced Matrix Multiplication 这个方向并不追求发明一种全新的通用矩阵乘法而是想回答另一个问题给定当前输入是否可以把完整矩阵乘法收缩成一个更小的矩阵乘法同时保持输出质量。这个思路在标题里被表达为 Input-Adaptive Matrix-Product Reduction也就是输入自适应的矩阵乘积约简。本文从一名推理系统工程师的视角先拆解原理再给出一个可运行的 NumPy 原型最后说明从原型走向真实推理层时要处理哪些问题。这篇文章会帮助你理解三层内容为什么 LLM 的矩阵乘有“输入动态性”可以利用一个最小可运行的输入自适应约简流程长什么样以及直接用t[mask] W[mask]为什么不能等同于生产级加速。原型部分使用 Python 和 NumPy不依赖特定模型权重也不需要对 PyTorch 有深入理解但需要熟悉ndarray的基本索引。1. 为什么 LLM 推理中的矩阵乘法存在可约简空间1.1 先看 token 生成时的 FFN 计算热点LLM 的 Transformer 块里注意力之后往往是一个大规模 FFN 结构。以常见的门控 FFN 为例单个 token 的前向过程可以拆成gate_pre x W_gate up_pre x W_up t act(gate_pre) * up_pre y t W_down其中x是当前 token 的输入向量维度记作DW_gate和W_up都会把D映射到隐藏维度HW_down再把隐藏维度H映射到输出维度O。假设一次只处理一个 tokenFFN 里的实际乘加次数大约是2 * D * H * 2 2 * H * O前面的2表示乘法和加法分开计数D * H * 2来自W_gate和W_up最后H * O来自W_down。在权重较大时这个数值会非常可观。很多加速思路并不是把某一次矩阵乘法替换成不存在而是想办法少算某些行或某些列。1.2 激活层会制造稀疏性但标准实现仍然按稠密矩阵计算如果 FFN 使用 ReLU 作为激活函数负的gate_pre会被直接置成 0。问题在于标准矩阵乘法并不会因为这个 0 而减少计算gate_pre x W_gate t relu(gate_pre) * up_pre y t W_down标准实现会先完整算出t再完整执行t W_down。但在 w_down 这一步真正参与贡献的是t中非零、或者绝对值很大的那些隐藏单元。对于当前输入完全不激活的隐藏单元W_down对应的行即使参与了乘法结果也会被乘 0。这个“乘 0”就是运算浪费。不同 token 的输入不同t里的零值分布也不同。于是“哪些隐藏单元可以被跳过”就变成由当前输入决定的问题而不是一次训练完就固定的静态剪枝。这种动态性正是 Input-Adaptive Matrix-Product Reduction 的关键。1.3 Reduced Matrix Multiplication 的实际收益模型如果完全不跳过任何行W_down的乘加次数是dense_count H * O如果针对当前 token只取活跃隐藏单元集合S集合大小为R |S|那么实际需要做的乘加次数近似变成reduced_count R * O看起来收益是明显的但它有三个隐藏成本需要额外计算当前输入的活跃集合也就是 mask 生成成本需要把W_down[active_rows]访存出来不能继续按稠密连续权重一次性加载如果输出缓冲区没有清零或者两个 token 的活跃集合不同计算结果可能写错。所以Reduced Matrix Multiplication 并不是简单“少算点矩阵乘法”而是一场动态选择、访存组织和数值误差之间的 trade-off。方案W_down 每 token 乘加次数是否取决于输入优点主要风险稠密矩阵乘法H * O否实现简单硬件利用率高大量无效行也被计算固定剪枝H * O否延迟可预测输入变化后精度劣化输入自适应约简R * O是跟随输入剪掉无效计算需要 mask、gather 和更复杂 kernel块状自适应约简B_count * O是硬件并行性强块内可能仍有无效单元前期可以先验证 R 的分布例如在验证集上统计每一层t ! 0的比例再决定这层是否值得写专门的约简算子。不要在所有层上一刀切。2. 输入自适应的矩阵乘积约简应该怎么设计2.1 一条完整链路选择集合、收集权重、约简乘加把一个完整的约简流程拆开看至少包含四步。第一步是确定活跃集合。如果是 ReLU 这类会产生严格 0 的激活可以直接用t ! 0判断。如果激活函数不会产生严格 0就需要用abs(t) eps或 top-k 的方式选择“值得计算的隐藏单元”。第二步是把权重行收集成紧凑布局。原始W_down是按 H 个隐藏单元保存的但当前输入可能只用到其中 R 行。真实 kernel 里不能依赖太随机的索引否则访存会退化成低效 gather。第三步是执行一个小矩阵乘法y[output_dim] sum_i t[active_i] * W_down[active_i, output_dim]第四步是把结果写进输出。写之前必须把输出缓冲区对应位置初始化为 0因为跳过的行不会再来累加。实际项目里可以直接用 torch.index_select 或 numpy indexing 验证逻辑 但生产推理阶段建议由自定义 CUDA/Triton kernel 完成“读取行号 - 加载权重行 - 乘加”的流程。2.2 选择策略精确零、阈值、top-k 与块状稀疏选择策略直接影响误差和硬件利用率。表里几类策略可以组合使用。策略判定方式误差说明精确零跳过t_i ! 0理论上无误差只消除严格 0 行收益取决于激活稀疏度阈值剪枝abs(t_i) eps有误差eps 越大误差越大但 R 越小top-k取绝对值前 k 个有误差计算量更可控但可能把大权重行错误跳过块稀疏以 32/64 行为一组组内有任一活跃则整组参与取决于块内策略更适合 GPU 并行要注意很多激活函数不会出现严格 0。GELU、SiLU 等激活会在 0 附近有微小但非零的值完整跳过它们会引入误差。为了安全使用阈值策略必须先在验证集上量化质量损失不能只看稀疏率。选择阈值时可以按每个 token 的abs(t)分位数动态调整也可以固定一个常量。固定阈值简单但不同层、不同 batch 的激活分布差异很大容易在某些输入上出现异常误差。top-k 的好处是计算量相对可预测缺点是“最不重要的 k 个不一定真的不重要”因为隐藏单元的重要性不只取决于t_i还取决于W_down对应行的权重范数。2.3 活跃集合的预测器放在哪一层这里有一个常见误区上面算t时已经执行了x W_gate和x W_up两个矩阵乘法。如果仍然先完整算出t再去跳过W_down的行那么优化的只是下游一次大乘法前层的W_gate和W_up并没有被节省。所以真正的“输入自适应”系统还需要一个更早的判定点。常见做法有两种。第一种是用上一层的输出或浅层特征训练一个很小的路由网络预测当前 token 会让哪些隐藏单元接近 0。预测可靠时可以直接跳过W_gate和W_up的对应输出行计算同时跳过W_down的对应输入行。第二种是降低判定成本比如用更小的采样维度、低精度近似或者上一时刻的活跃集合来估计当前 mask。这类方案误差更大但可能减少额外开销。在原型验证阶段先只做W_down的精确零跳过。这能帮你理清 R 的分布和数值误差再逐步把判定前移。2.4 残差、兜底与数值安全当使用近似截断时被剪掉的贡献不是严格 0只是被认为“足够小”。这个误差会在多层堆叠后累积。如果层数很多某一层相对误差 1e-3到最后可能变成肉眼可见的生成质量下降。因此生产实现通常需要设计兜底路径设置一个单 token 最大允许丢弃贡献量若某 token 的剪枝后行权重累计范数超过预设阈值则回退到完整的稠密计算每个模块独立校验输出误差不要让误差跨层无限放大。如果是精确零跳过则不存在空行贡献误差。但即使如此浮点累加顺序和重排也可能带来毫厘级差异仍建议用allclose设置宽松的 rtol 做测试。3. 用 NumPy 写一个可验证的最小原型3.1 环境准备本原型只需要 Python 3 和 NumPy。建议在虚拟环境里执行python -m venv .venv source .venv/bin/activate pip install numpy如果你的电脑没有 GPU也完全可以运行。这里不会使用 PyTorch 层因为最重要的不是 API而是理解约简前后的形状、mask 和误差计算。3.2 生成模拟权重和输入为了演示生成一组随机权重来模拟中型 FFN。隐藏维度 H 设置成 4096用于观察零激活行比例。import numpy as np D 256 H 4096 O 512 seed 0 rng np.random.default_rng(seed) x rng.standard_normal(D).astype(np.float32) W_gate rng.normal(0.0, 0.02, size(D, H)).astype(np.float32) W_up rng.normal(0.0, 0.02, size(D, H)).astype(np.float32) W_down rng.normal(0.0, 0.02, size(H, O)).astype(np.float32)这里没有模拟真实训练好的模型目的是验证约简逻辑本身。真实项目中你需要加载自己模型某一层的权重并且要确认权重和输入的数据类型一致否则x W_gate可能产生额外转换。3.3 实现一个完整矩阵乘法的基准路径先用最直观的方式计算一组参考输出后续所有约简结果都与它对比。def relu(a): return np.maximum(a, 0.0) gate_pre x W_gate up_pre x W_up t relu(gate_pre) * up_pre y_ref t W_down print(y_ref shape:, y_ref.shape) print(total hidden units:, H)t是 FFN 门控后的隐藏激活。这个数组里存在多少严格 0取决于 W_gate 的输入分布。生成的数据中大约会有一半gate_pre小于 0因此t中会有大量严格 0。3.4 实现精确零跳过的约简路径既然t中为 0 的行不会对y t W_down产生任何贡献那么理论上可以只保留非零行计算。active_mask t ! 0.0 active_idx np.flatnonzero(active_mask) y_zero_reduced t[active_idx] W_down[active_idx, :] # 验证与完整计算是否一致 np.testing.assert_allclose(y_zero_reduced, y_ref, rtol1e-6, atol1e-6) print(exact-zero reduced active count:, active_idx.size) print(dense W_down multiply-adds:, H * O) print(reduced W_down multiply-adds:, active_idx.size * O)这段代码的关键在于W_down[active_idx, :]只取出了当前 token 实际用到的权重行。由于跳过的权重行对应的t是 0数学结果不会变所以y_zero_reduced与y_ref完全一致。如果你把active_mask打印出来会发现它依赖t。这就是“输入自适应”的最小体现同一个权重矩阵W_down换一个 token被选中的行可能完全不同。3.5 加入 top-k 近似选择观察误差很多时候t里非零值很多但大部分绝对值很小。为了进一步缩小矩阵乘法可以按权重行范数加权的分数选择最靠前的 k 行。def reduce_by_mask(t, W_down, mask): active_idx np.flatnonzero(mask) if active_idx.size 0: return np.zeros(W_down.shape[1], dtypet.dtype) return t[active_idx] W_down[active_idx, :] # 用 |t_i| * ||W_down[i, :]||_2 作为“该行对输出贡献”的粗略估计 row_l2 np.linalg.norm(W_down, axis1) score np.abs(t) * row_l2 k min(H, int(H * 0.2)) if k 0: top_k_idx np.argpartition(-score, k - 1)[:k] approx_mask np.zeros_like(t, dtypebool) approx_mask[top_k_idx] True y_approx reduce_by_mask(t, W_down, approx_mask) denom np.linalg.norm(y_ref) 1e-6 rel_error float(np.linalg.norm(y_approx - y_ref) / denom) print(top-k k:, k) print(relative error:, rel_error)在这个模拟数据里W_gate和W_up的分布完全随机隐藏单元之间没有强结构所以 top-k 后误差通常不会小到可以忽略。实际模型里只有当激活分布有明显的低秩或稀疏结构时这种近似才有实用价值。这个步骤不是为了追求“切掉 80% 还不掉点”而是为了提醒你任何近似选择都要用误差指标做检查。3.6 约简后的理论收益怎么估算模型推理收益不能只看矩阵乘算力的减少量还要看额外开销。一个相对保守的粗估公式是估算减少的 W_down 乘加比例 (H - R) / H * 100%在精确零跳过的原型里如果 R 为 2048H 为 4096那么 W_down 的乘加减少比例是(4096 - 2048) / 4096 * 100% 50%但端到端延迟未必下降 50%因为W_gate 和 W_up 的完整计算仍然存在活跃行索引的计算和 gather 有额外开销kernel 访存可能由权重行的是否连续决定而不是只由 FLOPs 决定。因此原型阶段的结论应该写成“理论乘加减少了多少”而不是“推理变快了多少”。4. 从 NumPy 原型到真实推理系统真正关键的问题4.1 任意行索引会破坏 GPU 并行与访存连续性在 NumPy 里写W_down[active_idx]很轻松但真实 GPU kernel 里如果每一行是否计算都由当前 token 的 mask 决定线程块会遇到两个问题相邻 token 的活跃行集合不一致难以合并成一个大 GEMM每个 token 需要单独 gather 权重行权重访存很可能变得碎片化。因此生产实现通常不按“任意单行”做跳过而按 block 做跳过。例如以 32 行或 64 行作为最小粒度active_block block_contains_any_nonzero(t)只要这一个 block 里存在任意一个需要计算的行就把整个 block 纳入计算。这样做会多算块内一些接近 0 的行但可以大幅提高访存与并行效率。4.2 输出缓冲区清零必须谨慎如果矩阵乘法使用 beta 累加跳过的行不会写输出那么输出缓冲区必须提前清零否则上一次计算或者垃圾数据会被错误保留下来。一个常见错误是第一次 token 的活跃行是第 1、5、10 行输出缓冲区被写入第二次 token 的活跃行变为第 2、5、11 行但缓冲区没有清零。第二次结果中第 1 行和第 10 行的旧值仍然留在输出里导致推理结果错误。生产 kernel 里应显式管理清零时机if segment_start: zero_output(y) compute_y_with_active_rows(...)在 NumPy 原型里不会出这个问题因为y_approx是由局部