新PyTorch API:几行代码实现不同注意力变体,兼具FlashAttention性能和PyTorch灵活性

新PyTorch API:几行代码实现不同注意力变体,兼具FlashAttention性能和PyTorch灵活性

💡 原文中文,约2400字,阅读约需6分钟。
📝

内容提要

PyTorch团队引入了FlexAttention,一个灵活的API,允许用户使用几行PyTorch代码实现多个注意力变体。通过torch.compile将其降低到一个融合的FlashAttention内核中,生成了一个不会占用额外内存且性能可与手写内核相媲美的FlashAttention内核。FlexAttention具有令人惊讶的表达能力,可以满足大多数用户对注意力变体的需求。

🔎

延伸解读

FlexAttention 的定位与价值

FlexAttention 旨在解决现有优化注意力内核灵活性不足的问题。传统方法如 FlashAttention 虽然性能高,但只支持特定变体,用户若需要新变体则面临性能下降和内存不足。FlexAttention 通过允许用户用几行 PyTorch 代码定义 score_mod 函数,在保持高性能的同时提供了灵活性,使研究人员能轻松实现和组合多种注意力变体。

score_mod 的表达能力与实现

score_mod 是一个用户定义的函数,在 softmax 之前修改注意力分数。它接受分数和索引参数,返回修改后的分数。通过示例可见,score_mod 能实现全注意力、相对位置编码、Soft-capping、因果掩码、滑动窗口等多种变体,甚至组合使用。这种设计无需具体化大型张量,动态计算偏差,提升了内存和性能。

性能表现与权衡

FlexAttention 的性能接近手写的 Triton 内核,前向传播达到 FlashAttention2 的 90%,反向传播达到 85%。由于通用性,会有轻微性能损失和额外延迟。目前使用确定性反向算法,比 FAv2 重计算更多中间体,但团队计划改进以缩小差距。对于需要灵活性的用户,这一性能损失可能是可接受的。

❓

Q&A

FlexAttention 是什么?

FlexAttention 是一个灵活的 PyTorch API,允许用户用几行代码实现多个注意力变体。

FlexAttention 如何提高性能?

通过 torch.compile,FlexAttention 被降低到一个融合的 FlashAttention 内核,性能可与手写内核相媲美且不占用额外内存。

FlexAttention 支持哪些注意力变体?

FlexAttention 支持因果注意力、相对位置嵌入、滑动窗口注意力等多种注意力变体。

使用 FlexAttention 的好处是什么?

使用 FlexAttention,用户可以灵活定义注意力变体,避免了运行缓慢和 CUDA 内存不足的问题。

FlexAttention 的性能与手写内核相比如何?

FlexAttention 的性能接近手写的 Triton 内核,前向传播实现了 FlashAttention2 性能的 90%,反向传播实现了 85%。

未来对 FlexAttention 的改进计划是什么?

研究者计划改进 FlexAttention 的反向算法,以缩小与 FlashAttention2 的性能差距。

🏷️

标签

➡️

继续阅读