改变乘积,保留参数:面向 Transformer 的结合代数层
快速矩阵乘法算法保持乘积不变,只去寻找更廉价的求值方式。我们反过来提问:Transformer 学到的投影,能否直接换用一个更廉价的乘积? 基于一种结合代数构造,我们用更稀疏的交互表替代普通矩阵乘法,同时保留相同的权重分块。该构造在物理块大小固定时,矩阵维度上的算术量为二次,并给出了面向 GPU 执行的有限形状约束。它在其双线性秩意义上由 Alder–Strassen 界可证最优,且可实现为与 causal masking 和 KV-cached decoding 兼容的行类型矩形投影。 实证方面,我们用同一配方与 12.3B token 预算训练了两个约 110M 参数的 decoder-only Transformer LM,二者仅 feed-forward 层不同:一个使用普通稠密矩阵乘法,另一个使用结合代数乘积。在四个 prompt 域上,代数模型端到端生成吞吐提升 6.2–7.8%,但三项下游指标得分均更低。 我们将此视为小规模下对该方法的可行性与可训练性检验,更深入的探究留待未来工作。
论文精读
TL;DR 用结合代数构造的稀疏乘积替代 Transformer 前馈层的稠密矩阵乘法,在固定物理块大小时算术复杂度降为二次方(达到 Alder-Strassen 下界最优),并在 110M 参数 LM 上验证生成吞吐提升 6-8% 而下游质量略降。
问题
问题背景
Transformer 中大量 FLOPs 消耗在投影层与 FFN 的稠密矩阵乘法上,业界长期关注如何降低浮点运算量而不损失模型容量。
现有方法局限
典型路线有两条:一是快速矩阵乘法算法(如 Strassen、Winograd),它们保持乘积语义不变,通过分治减少标量乘法次数,但在中等矩阵尺寸下常数开销大、数值稳定性差、GPU 张量核心利用率低;二是结构化权重层(低秩、Butterfly、Monarch 等),它们约束权重矩阵形态,但往往引入额外近似误差或限制表达能力,且多数仍以普通矩阵乘法为底层运算。二者均未改变“乘积”本身。
为什么这个问题难/重要
直接用更廉价的乘积替换普通矩阵乘法,需同时满足:乘积满足结合律以保证多层堆叠不崩溃;双线性秩低于普通矩阵乘法且可达 Alder–Strassen 下界;权重块形状能映射到 GPU 友好的矩形 tile;前向与反向均能高效实现;还要兼容 causal masking 和 KV-cached decoding。若可行,可在保持参数量的前提下直接提升生成吞吐,而无需重新设计注意力或量化。但 trainability 与下游质量损失仍是风险点。
行业类比
类似将标准卷积替换为深度可分离卷积以降低 FLOPs,但本工作是在代数层面替换乘法运算本身,而非仅做算子分解。
核心洞察
- 将乘积本身视作可学习的架构选择,而非固定不变的计算对象。传统快速矩阵乘法算法在保持乘积结果不变的前提下降低求值复杂度;本文反其道而行之,用关联代数定义的稀疏交互表替换标准矩阵乘法,保留相同权重块,但改变层的前向语义。这一视角将 Transformer 投影层的搜索空间从“如何更快计算矩阵乘法”扩展到“可以用哪些更便宜的乘积替代矩阵乘法”,为设计高效 Transformer 架构提供了新维度。
- 关联代数层可具体实现为兼容因果掩码和 KV-cache 的 GPU 高效操作,并具备理论最优的双线性复杂度。通过行类型矩形投影和有限形状约束,该乘积在解码时维持较低算术开销;构造达到 Alder-Strassen 下界,表明在双线性秩意义上无法更省标量乘法。在 110M 参数语言模型上,该层带来 6.2-7.8% 的生成吞吐量提升,尽管下游指标下降,但验证了训练可行性与实现路径,为后续放大规模提供了工程基线。
方法
输入
该方法作用于 Transformer 前馈层 (FFN):输入为隐藏状态向量 h,其维度为模型宽度 d;原 FFN 包含两次稠密矩阵乘法(升维与降维投影)。本方法保持权重矩阵的物理块大小与参数量不变,仅替换乘积运算。
关键模块
代数乘积构造
将权重矩阵划分为固定大小的块,定义一张稀疏交互表,指定哪些块对参与乘法。乘积结果由这些稀疏块乘加得到,而非全稠密块间运算。该交互表来自一个结合代数,保证乘积满足结合律,从而可安全堆叠多层。复杂度族与最优性
当物理块大小固定时,该乘积的算术复杂度随矩阵维度二次增长(而非普通矩阵乘法的三次增长)。构造达到 Alder–Strassen 下界,即在其双线性秩意义下不可再减少标量乘法次数。有限形状约束
为适配 GPU 执行,推导出可行的块形状与 tile 尺寸约束,使稀疏交互表上的计算可映射为高效 kernel,避免不规则访存。Transformer 兼容性
乘积实现为行类型矩形投影,天然支持因果掩码与 KV 缓存解码:注意力层不变,仅 FFN 中的两次投影被替换为代数乘积。
输出
输出向量与普通 FFN 形状一致,可直接进入残差连接与后续层。实验训练两个约 110M 参数的 decoder-only LM,仅 FFN 层乘积方式不同:代数模型端到端生成吞吐提升 6.2–7.8%,但下游任务分数略有下降,作者定位为小规模可行性验证。
与同类方法的差异
不同于快速矩阵乘法算法(如 Strassen)保持乘积结果固定、仅优化计算路径,本方法直接改变乘积定义,在保持参数块不变的前提下降低算术复杂度,属于架构层面的算子替换。
实验
实验设计
论文训练两个约110M参数的decoder-only Transformer语言模型,使用相同训练配方和12.3B token预算。唯一区别在前馈层:基线使用普通稠密矩阵乘法,实验模型使用结合代数乘积(associative-algebra product)。评估覆盖四个提示领域,测量端到端生成吞吐量和三个下游指标。
关键发现
- 代数模型在生成吞吐量上相对基线提升 6.2%–7.8%。
- 但在三个下游指标上得分均低于基线。
- 作者将此视为小规模可行性与可训练性检验,指出速度提升伴随质量下降。
深度解读
该结果表明,用更稀疏的交互表替代普通矩阵乘法可以在不改变参数量的前提下提高端到端推理速度,但可能削弱模型表达能力。作为概念验证,实验规模较小,下游评估未给出具体任务,无法得出普遍结论。与快速矩阵乘法算法不同,本方法改变乘积本身而非仅优化计算过程。未来需在更大规模、更多任务上验证。
行业影响
落地场景
该技术可直接用于大语言模型推理服务,尤其是生成式 AI 应用中的前馈层(MLP)替换。例如:
- 实时聊天机器人:电商平台的智能客服、内容平台的互动助手,需要高吞吐和低延迟,代数层可将端到端生成吞吐提升 6.2–7.8%,在同等硬件下服务更多并发用户。
- 企业级知识库问答:金融、医疗领域的内部问答系统,推理成本占大头,替换后可降低单位 token 计算成本,同时保持可接受的回答质量(当前下游指标略降,适合质量容忍度稍高的场景)。
商业价值
主要收益来自降本:该构造在参数块固定的前提下,用更少的算术操作实现相同维度的变换,理论上达到最优双线性秩,可减少推理时的 FLOPs,从而降低 GPU 能耗和云服务成本。对模型服务商而言,6–8% 的吞吐提升意味着可扩大服务规模或降价竞争。
此外,延迟的降低能直接改善用户体验:语音助手、代码补全等交互式工具对响应时间敏感,吞吐提升可缩短首 token 延迟,提升留存率。
当前下游指标略低,短期更适合作为成本优化选项,或与普通密集层混合使用(如关键层保留原乘法,次要层替换)。
跟现有产品/工作流的接口
该代数层可设计为drop-in 替换:它使用相同的权重块和参数形状,仅改变投影的乘积方式,并与 causal masking 和 KV-cache 解码兼容。因此,工程师可以将现有的 PyTorch 模型中的 nn.Linear 前馈层替换为自定义的 AssociativeAlgebraLayer,无需改动整体架构或训练框架。
实际集成时:
- 训练阶段:可以联合训练(论文已验证可训练性),或先用普通层训练,再通过权重对齐转换到代数层。
- 推理阶段:需要实现高效的 CUDA kernel 或利用现有稀疏矩阵计算库(如 cuSPARSE、CUTLASS),因为代数乘积表现为稀疏交互表。
- 部署:该层可导出为 ONNX 或 TensorRT 自定义插件,接入现有推理服务栈。
该构造的二次算术复杂度(当块大小固定时)意味着在大矩阵维度下收益更明显,适合长序列或宽 FFN 的模型。
局限
- - 论文明确将实验定位为可行性验证:两个约 **110M 参数** 模型仅在 **12.3B token** 预算下训练,下游指标全部低于基线,作者未提供更大规模或更长训练下的趋势,因此无法判断该代数层是否能在实用模型中保持质量。
- - 该方法仅替换了 Transformer 的前馈层矩阵乘法,注意力等其余部分仍为常规运算,整体端到端吞吐量提升仅 **6.2–7.8%**,且实验用 110M 模型可能因前馈层占比不同而在大模型上有所变化;同时生成吞吐量提升来自自定义 kernel,但 kernel 的通用性和与现有推理框架的集成度尚未讨论。
- - 与保持数值等价的加速方法(如 **FlashAttention**)不同,该代数乘积改变了前馈层的计算语义,可能引入额外表达能力损失或训练不稳定风险;此外,其“在双线性秩下最优”是理论性质,但实际 GPU 执行受限于有限矩形形状约束,未必能转化为 wall-clock 加速;论文未与更成熟的低秩/结构化替代方案(如 **Monarch**, **BTT**)进行系统实验对比。