换基而不失速度:DNN 中 MatMul 的 GPU 高效替代方案
机器学习
2025-10-22 v3 人工智能
数据结构与算法
摘要
现代 AI 依赖巨大的矩阵乘法(MatMul),其计算对推理和训练构成了可扩展性挑战。我们提出了一种替代方案——一种 GPU 原生的双线性算子,用于替代神经网络中的 MatMul,它在速度、准确性和参数数量之间提供了三方权衡。具体而言,该算子在求值时需要显著更少的 FLOPs(),但与 MatMul 相比增加了参数数量()。我们将该算子称为 Strassen-Tile(STL)。STL 的核心思想是在权重矩阵和激活矩阵的块(tiles)上应用局部可学习的基变换,然后在块之间进行逐元素乘积,同时通过 MatMul 实现。我们研究的关键技术问题是如何优化给定层的基变换,这是一个高度非凸问题。我们表明,基于理论的初始化(受快速矩阵和多项式乘法启发)比随机 SGD 初始化带来了显著更好的准确性。这一现象激发了对 STL 在 DNN 中优化的进一步算法研究。我们的实验表明,STL 可以在减少 2.66 倍 FLOPs 的同时近似 4×4 的块 MatMul,并能在降低 FLOPs 的情况下提升 SoTA T2T-ViT-7(4.3M 参数)在 ImageNet-1K 上的准确率。即使使用未经 CUDA 优化的 PyTorch 代码,STL 在计算密集型场景下也能实现实际的时钟速度提升。这些结果连同其理论基础,表明 STL 是可扩展且经济高效的 AI 的一个有前景的构建模块。
引用
@article{arxiv.2503.12211,
title = {Changing Base Without Losing Pace: A GPU-Efficient Alternative to MatMul in DNNs},
author = {Nir Ailon and Akhiad Bercovich and Yahel Uffenheimer and Omri Weinstein},
journal= {arXiv preprint arXiv:2503.12211},
year = {2025}
}