CODA: 通过 GEMM-Epilogue 融合优化 Transformer 训练
Transformer 训练系统从根本上构建在稠密线性代数之上,然而端到端执行时间中的很大一部分往往被内存受限(memory-bound)的算子所消耗。虽然核心矩阵乘法(GEMMs)已经过高度优化,但周围的操作——例如归一化(normalization)、激活(activations)、残差更新(residual updates)和归约(reductions)——经常需要通过全局内存移动大量中间张量。这种重复的数据移动造成了关键的瓶颈,限制了训练栈的整体效率。
为了解决这个问题,研究人员引入了 CODA,这是一种 GPU 内核抽象,旨在将这些计算表达为“GEMM-plus-epilogue”程序。通过重新思考 Transformer 块的执行方式,CODA 旨在消除不必要的内存往返,并最大限度地提高片上数据利用率。
核心问题:Transformer 中的内存墙
在标准的 Transformer 块中,操作序列通常被处理为一系列独立的框架内核(framework kernels)。例如,一个 GEMM 操作会产生一个结果并将其写入全局内存,随后一个激活或归一化内核会再次读取该数据到 GPU 的寄存器或共享内存中,仅为了执行一个简单的算术操作,然后再将其写回。
这种模式本质上是低效的。因为这些“epilogue”操作在移动每字节数据时执行的算术运算非常少,所以它们是内存受限的。随着硬件计算能力的发展速度快于内存带宽,这些非 GEMM 操作在训练时间中占据了不成比例的更大比例。
CODA 方法:GEMM-Epilogue 程序
CODA 基于这样一个观察:许多 Transformer 算子可以通过代数重参数化(reparameterized)来实现。CODA 不再将它们视为独立的步骤,而是将它们集成到 GEMM 过程中。
工作原理
CODA 固定了 GEMM 主循环(mainloop)——即矩阵乘法中计算密集的部分——并暴露了一组可组合的 epilogue 原语(primitives)。这些原语允许进行:
- Scaling: 调整输出分块(tiles)的大小。
- Reductions: 执行求和或其他归约操作。
- Pairwise Transformations: 应用逐元素函数(例如激活函数)。
- Accumulation: 将结果添加到现有张量中(例如残差连接)。
通过在 GEMM 输出分块仍保留在片上时执行这些操作,CODA 避免了昂贵的全局内存读写过程。其结果是一个既能保持专家编写的 GEMM 性能,又能提供高层抽象鲁棒性的内核。
超越手动优化:LLM 的角色
CODA 框架最显著的方面之一是其易用性。研究人员发现,无论是人类编写的还是 LLM 编写的 CODA 内核都实现了高性能。这表明低层 GPU 内核编写方式正在发生转变。
虽然 LLM 在处理低层硬件优化的复杂细节方面(例如管理共享内存库(shared memory banks)或优化 warp-level 原语)往往表现不佳,但它们在高级组合方面表现出色。通过提供一个受限且可组合的 API,CODA 允许 LLM 将专家编写的模块“粘合”成一个功能完备且高效的内核。
正如社区成员在讨论该论文时所指出的:
"设计具有受限且可组合 API 的编译器抽象,以便 LLM 可以轻松地将专家编写的模块粘合在一起,这是一个明智之举。我怀疑随着我们向智能体开发(agentic development)迈进,这最终将成为代码生成(codegens)的常态。"
结论
CODA 代表了一条将框架级生产力与硬件级效率相结合的实用路径。通过将 Transformer 块视为一系列 GEMM-epilogue 程序,它减少了内存瓶颈,并提供了一个结构化的环境,使自动化工具——以及人类——能够快速迭代内核性能。随着更大规模模型的训练不断推向内存带宽的极限,像 CODA 这样的抽象将对于优化下一代 AI 基础设施至关重要。