FlashRNN:针对现代硬件的传统RNN的I/O感知优化
机器学习
2025-03-14 v3 人工智能
摘要
虽然Transformer和其他可并行化的序列神经网络架构似乎在序列建模中占据当前最佳地位,但它们特别缺乏状态跟踪能力,这在时间序列任务和逻辑推理中至关重要。传统RNN如LSTM和GRU,以及现代变体如sLSTM,都具备这些能力,代价是严格的串行处理。虽然这常被视为强大的限制,我们展示了在Triton和CUDA中的FlashRNN硬件优化下,这些网络可以多快。我们在现代GPU上针对寄存器级别优化内核,对传统RNN进行了并行化扩展,通过并行处理多个小隐藏状态的RNN(类似于Transformer中的头级处理),实现了并行化。为实现不同GPU变体上的灵活性,我们引入了一种新的硬件内部缓存大小、内存和计算处理的优化框架。它使用多面体式约束,包括可整除性的概念。这加快了ConstrINT库中通用整数约束满足问题(integer CSPs)的求解速度。我们展示,我们的内核相对于vanilla PyTorch实现可实现50倍加速,允许比我们Triton实现的40倍更大的隐藏层大小。我们发布的开源内核和优化库旨在推动以状态跟踪为特征的RNN和序列建模的研究:https://github.com/NX-AI/flashrnn
引用
@article{arxiv.2412.07752,
title = {FlashRNN: I/O-Aware Optimization of Traditional RNNs on modern hardware},
author = {Korbinian Pöppel and Maximilian Beck and Sepp Hochreiter},
journal= {arXiv preprint arXiv:2412.07752},
year = {2025}
}