
理论上,“Attention Is All You Need”。然而在实践中,我们还需要像 FlashAttention 这样经过优化的注意力实现。
尽管这些融合(fused)的注意力实现大幅提升了性能并支持了长上下文,但这种效率是以牺牲灵活性为代价的。你无法再通过编写几行 PyTorch 代码来尝试一种新的注意力变体——你通常需要编写一个新的自定义内核!对于机器学习研究人员来说,这就像一场“软件彩票”:如果你的注意力变体无法适配现有的优化内核,你就注定要面对缓慢的运行速度和 CUDA 显存溢出(OOM)。
一些注意力变体的例子包括:因果掩码(Causal)、相对位置编码、Alibi、滑动窗口注意力、PrefixLM、文档掩码/样本打包/锯齿状张量、Tanh 软上限(Soft-Capping)、PagedAttention 等。更糟糕的是,人们往往需要这些技术的组合!滑动窗口 + 文档掩码 + 因果 + 上下文并行?或者 PagedAttention + 滑动窗口 + Tanh 软上限?
下图左侧展示了当下的现状——一些掩码 + 偏置 + 设置的组合拥有现成的内核实现。但各种选项导致了指数级的设置组合,最终导致支持力度参差不齐。更糟糕的是,研究人员提出的新注意力变体将得到“零”支持。

为了彻底解决这个超立方体问题,我们推出了一个新的 PyTorch API:FlexAttention。
- 我们提供了一个灵活的 API,允许仅用几行地道的 PyTorch 代码实现许多注意力变体(包括目前博文中提到的所有变体)。
- 通过
torch.compile,我们将这些代码下沉为融合的 FlashAttention 内核,生成的内核不会占用额外的显存,且性能与手写内核旗鼓相当。 - 我们还利用 PyTorch 的自动求导机制,自动生成了反向传播过程。
- 最后,我们还能利用注意力掩码的稀疏性,从而相比标准注意力实现获得显著的性能改进。
有了 FlexAttention,我们希望尝试新的注意力变体时,唯一受限的只有你的想象力。
你可以在 Attention Gym 中找到许多 FlexAttention 的示例:https://github.com/pytorch-labs/attention-gym。如果你有任何很酷的应用,欢迎提交示例!
附:我们也觉得这个 API 非常令人兴奋,因为它以一种有趣的方式利用了许多现有的 PyTorch 基础架构——文末会详细说明。
FlexAttention
这是经典的注意力公式

代码形式如下
Q, K, V: Tensor[batch_size, num_heads, sequence_length, head_dim]
score: Tensor[batch_size, num_heads, sequence_length, sequence_length] = (Q @ K) / sqrt(head_dim)
probabilities = softmax(score, dim=-1)
output: Tensor[batch_size, num_heads, sequence_length, head_dim] = probabilities @ V
FlexAttention 允许用户自定义函数 score_mod:

代码形式如下
Q, K, V: Tensor[batch_size, num_heads, sequence_length, head_dim]
score: Tensor[batch_size, num_heads, sequence_length, sequence_length] = (Q @ K) / sqrt(head_dim)
modified_scores: Tensor[batch_size, num_heads, sequence_length, sequence_length] = score_mod(score)
probabilities = softmax(modified_scores, dim=-1)
output: Tensor[batch_size, num_heads, sequence_length, head_dim] = probabilities @ V
该函数允许你在 softmax 之前“修改”注意力分数。令人惊讶的是,这足以应付绝大多数注意力变体(见下文示例)!
具体来说,score_mod 的期望签名比较独特。
def score_mod(score: f32[], b: i32[], h: i32[], q_idx: i32[], kv_idx: i32[])
return score # noop - standard attention
换句话说,score 是一个标量 PyTorch 张量,代表 query token 和 key token 的点积。其余参数告诉你当前正在计算“哪一个”点积——b(批次中的当前元素)、h(当前头)、q_idx(query 中的位置)、kv_idx(key/value 张量中的位置)。
要应用此函数,我们可以这样实现:
for b in range(batch_size):
for h in range(num_heads):
for q_idx in range(sequence_length):
for kv_idx in range(sequence_length):
modified_scores[b, h, q_idx, kv_idx] = score_mod(scores[b, h, q_idx, kv_idx], b, h, q_idx, kv_idx)
当然,这并不是 FlexAttention 底层的实现方式。利用 torch.compile,我们自动将你的函数下沉为一个单一的、融合的 FlexAttention 内核——保证性能,否则退款!
这个 API 的表达能力出人意料地强大。让我们看看一些示例。
Score Mod 示例
全注意力 (Full Attention)
首先来看“全注意力”,即标准的双向注意力。在这种情况下,score_mod 是一个空操作(no-op)——它接收分数作为输入,并按原样返回它们。
def noop(score, b, h, q_idx, kv_idx):
return score
要端到端使用它(包括前向和反向传播):
from torch.nn.attention.flex_attention import flex_attention
flex_attention(query, key, value, score_mod=noop).sum().backward()
相对位置编码
一种常见的注意力变体是“相对位置编码”。它不是在 query 和 key 中编码绝对距离,而是基于 query 和 key 之间的“距离”来调整分数。
def relative_positional(score, b, h, q_idx, kv_idx):
return score + (q_idx - kv_idx)
请注意,与典型实现不同,这不需要实例化一个 SxS 张量。相反,FlexAttention 在内核中“即时”计算偏置值,从而显著改善内存和性能。

ALiBi 偏置

来源:Train Short, Test Long: Attention with Linear Biases Enables Input Length Extrapolation
ALiBi 在 Train Short, Test Long 中被引入,声称在推理时的长度外推方面具有优势。值得注意的是,MosaicML 指出“缺乏内核支持”是他们最终从 ALiBi 转向旋转位置编码(Rotary Embeddings)的主要原因。
Alibi 与相对位置编码类似,只有一个区别——它有一个通常预先计算好的逐头(per-head)因子。
alibi_bias = generate_alibi_bias() # [num_heads]
def alibi(score, b, h, q_idx, kv_idx):
bias = alibi_bias[h] * (kv_idx - q_idx)
return score + bias
这展示了 torch.compile 提供的一个有趣灵活性——即使 alibi_bias 没有被显式地作为输入传入,我们仍然可以从中读取数据!生成的 Triton 内核会计算出从 alibi_bias 张量的正确读取方式并进行融合。请注意,即使你重新生成 alibi_bias,我们也无需重新编译。
软上限 (Soft-capping)
软上限是 Gemma2 和 Grok-1 中使用的一种技术,旨在防止 Logits 变得过大。在 FlexAttention 中,它看起来像这样:
softcap = 20
def soft_cap(score, b, h, q_idx, kv_idx):
score = score / softcap
score = torch.tanh(score)
score = score * softcap
return score
请注意,我们也在这里根据前向传播自动生成了反向传播。此外,尽管该实现语义正确,但在性能上我们通常建议使用 tanh 近似。更多细节请查看 attention-gym。
因果掩码 (Causal Mask)
尽管双向注意力最简单,但最初的 *Attention Is All You Need* 论文和绝大多数大语言模型(LLM)都在仅解码器(decoder-only)设置中使用注意力,其中每个 token 只能关注其之前的 token。人们通常认为这是一个下三角掩码,但使用 score_mod API,它可以表达为:
def causal_mask(score, b, h, q_idx, kv_idx):
return torch.where(q_idx >= kv_idx, score, -float("inf"))
基本上,如果 query token 在 key token “之后”,我们保留分数。否则,我们将其设为 -inf 来屏蔽它,从而确保它不会参与 softmax 计算。
然而,与其他修改相比,掩码比较特殊——如果某些内容被掩码了,我们可以完全跳过其计算!在这种情况下,因果掩码大约有 50% 的稀疏性,如果不利用这种稀疏性,会导致 2 倍的减速。虽然 score_mod 足以“正确”实现因果掩码,但要获得稀疏性的性能优势,需要另一个概念——mask_mod。
Mask Mods
为了利用掩码带来的稀疏性,我们需要做更多工作。具体来说,通过将 mask_mod 传递给 create_block_mask,我们可以创建一个 BlockMask。FlexAttention 随后可以使用 BlockMask 来利用稀疏性!
mask_mod 的签名与 score_mod 非常相似——只是没有 score 参数。特别是:
# returns True if this position should participate in the computation
mask_mod(b, h, q_idx, kv_idx) => bool
请注意,score_mod 在表达能力上严格优于 mask_mod。然而,对于掩码,建议使用 mask_mod 和 create_block_mask,因为性能更好。请参阅关于为何将 score_mod 和 mask_mod 分开的 FAQ。
现在,让我们看看如何使用 mask_mod 实现因果掩码。
因果掩码 (Causal Mask)
from torch.nn.attention.flex_attention import create_block_mask
def causal(b, h, q_idx, kv_idx):
return q_idx >= kv_idx
# Because the sparsity pattern is independent of batch and heads, we'll set them to None (which broadcasts them)
block_mask = create_block_mask(causal, B=None, H=None, Q_LEN=1024, KV_LEN=1024)
# In this case, we don't need a score_mod, so we won't pass any in.
# However, score_mod can still be combined with block_mask if you need the additional flexibility.
flex_attention(query, key, value, block_mask=block_mask)
请注意,create_block_mask 是一个相对昂贵的操作!尽管 FlexAttention 在它发生变化时不需要重新编译,但如果你不注意缓存它,可能会导致明显的性能下降(请查看 FAQ 以获取最佳实践建议)。

虽然 TFlops 大致相同,但 mask_mod 版本的执行时间快了 2 倍!这证明了我们可以在不丢失硬件效率的情况下利用 BlockMask 提供的稀疏性。
滑动窗口 + 因果 (Sliding Window + Causal)

来源:Mistral 7B
由 Mistral 推广,滑动窗口注意力(也称为局部注意力)利用了“最近的 token 最有用”的直觉。特别是,它允许 query token 只关注(例如)最近的 1024 个 token。这通常与因果注意力结合使用。
SLIDING_WINDOW = 1024
def sliding_window_causal(b, h, q_idx, kv_idx):
causal_mask = q_idx >= kv_idx
window_mask = q_idx - kv_idx <= SLIDING_WINDOW
return causal_mask & window_mask
# If you want to be cute...
from torch.nn.attention import and_masks
def sliding_window(b, h, q_idx, kv_idx)
return q_idx - kv_idx <= SLIDING_WINDOW
sliding_window_causal = and_masks(causal_mask, sliding_window)
我们将它与使用滑动窗口掩码的 F.scaled_dot_product_attention 以及使用因果掩码的 FA2(作为性能参考)进行了基准测试。我们不仅比 F.scaled_dot_product_attention 快得多,而且由于该掩码具有显著的稀疏性,我们还比带因果掩码的 FA2 快得多。

PrefixLM

来源:PaliGemma: A versatile 3B VLM for transfer
T5 架构描述了一种注意力变体:对“前缀(prefix)”执行全双向注意力,对其余部分执行因果注意力。我们再次组合两个掩码函数来完成此操作,一个用于因果掩码,另一个基于前缀长度。
prefix_length: [B]
def prefix_mask(b, h, q_idx, kv_idx):
return kv_idx <= prefix_length[b]
prefix_lm_causal = or_masks(prefix_mask, causal_mask)
# In this case, our mask is different per sequence so we set B equal to our batch size
block_mask = create_block_mask(prefix_lm_causal, B=B, H=None, S, S)
就像 score_mod 一样,mask_mod 允许我们引用并非函数显式输入的额外张量!然而,对于 PrefixLM,稀疏模式会随“每个输入”变化。这意味着对于每个新的输入批次,我们需要重新计算 BlockMask。一种常见的模式是在模型开头调用 create_block_mask,并在模型的所有注意力调用中重用该 block_mask。请参阅“重计算 Block Mask 与重编译”。
然而,作为交换,我们不仅能够为 PrefixLM 提供高效的注意力内核,还能够利用输入中存在的任何稀疏性!FlexAttention 会根据 BlockMask 数据动态调整其性能,而无需重新编译内核。
文档掩码/锯齿状序列 (Document Masking/Jagged Sequences)
另一种常见的注意力变体是文档掩码/锯齿状序列。假设你有许多长度不一的序列。你希望将它们一起训练,但不幸的是,大多数运算符只接受矩形张量。
通过 BlockMask,我们也能在 FlexAttention 中高效地支持这一点!
- 首先,我们将所有序列平铺(flatten)为一个包含 sum(序列长度) 个 token 的单一序列。
- 然后,我们计算每个 token 所属的 document_id。
- 最后,在
mask_mod中,我们只需判断 query 和 kv token 是否属于同一个文档!
# The document that each token belongs to.
# e.g. [0, 0, 0, 1, 1, 2, 2, 2, 2, 2, 2] corresponds to sequence lengths 3, 2, and 6.
document_id: [SEQ_LEN]
def document_masking(b, h, q_idx, kv_idx):
return document_id[q_idx] == document_id[kv_idx]
就这样!在这种情况下,我们最终得到一个块对角掩码。

关于文档掩码的一个有趣点是,它很容易与其他掩码组合。例如,我们已经在上一节中定义了 prefixlm_mask。我们现在是否还需要定义一个 prefixlm_document_mask 函数?
在这些情况下,我们发现一种非常有用的模式,即所谓的“更高级的修改”。在这种情况下,我们可以获取现有的 mask_mod 并自动将其转换为适用于锯齿状序列的掩码!
def generate_doc_mask_mod(mask_mod, document_id):
# Get unique document IDs and their counts
_, counts = torch.unique_consecutive(document_id, return_counts=True)
# Create cumulative counts (offsets)
offsets = torch.cat([torch.tensor([0], device=document_id.device), counts.cumsum(0)[:-1]])
def doc_mask_wrapper(b, h, q_idx, kv_idx):
same_doc = document_id[q_idx] == document_id[kv_idx]
q_logical = q_idx - offsets[document_id[q_idx]]
kv_logical = kv_idx - offsets[document_id[kv_idx]]
inner_mask = mask_mod(b, h, q_logical, kv_logical)
return same_doc & inner_mask
return doc_mask_wrapper
例如,给定上面的 prefix_lm_causal 掩码,我们可以将其转换为适用于打包文档的掩码,如下所示:
prefix_length = torch.tensor(2, dtype=torch.int32, device="cuda")
def prefix_mask(b, h, q_idx, kv_idx):
return kv_idx < prefix_length
prefix_lm_causal = or_masks(prefix_mask, causal_mask)
doc_prefix_lm_causal_mask = generate_doc_mask_mod(prefix_lm_causal, document_id)

现在,这个掩码的形状是“块-前缀LM-对角线”的。🙂
以上就是所有的示例!注意力变体比我们能列出的要多得多,所以请查看 Attention Gym 获取更多示例。我们希望社区也能贡献他们最喜欢的 FlexAttention 应用。
FAQ
Q:FlexAttention 何时需要重新编译?
由于 FlexAttention 利用 torch.compile 进行图捕获,它实际上可以在广泛的情况下避免重新编译。值得注意的是,即使捕获的张量值发生变化,它也不需要重新编译!
flex_attention = torch.compile(flex_attention)
def create_bias_mod(bias)
def bias_mod(score, b, h, q_idx, kv_idx):
return score + bias
return bias_mod
bias_mod1 = create_bias_mod(torch.tensor(0))
flex_attention(..., score_mod=bias_mod1) # Compiles the kernel here
bias_mod2 = create_bias_mod(torch.tensor(2))
flex_attention(..., score_mod=bias_mod2) # Doesn't need to recompile!
即使改变块稀疏度也不需要重新编译。但是,如果块稀疏度发生了变化,我们确实需要“重新计算” BlockMask。
Q:我们何时应该重新计算 BlockMask?
每当块稀疏度发生变化时,我们都需要重新计算 BlockMask。尽管计算 BlockMask 比重新编译要便宜得多(大约几百微秒,而不是几秒),但你仍应注意不要过度计算 BlockMask。
以下是一些常见模式以及关于如何处理它们的建议。
掩码从不改变(例如因果掩码)
在这种情况下,你可以简单地预计算块掩码并将其全局缓存,在所有注意力调用中重用它。
block_mask = create_block_mask(causal_mask, 1, 1, S,S)
causal_attention = functools.partial(flex_attention, block_mask=block_mask)
掩码每批次改变(例如文档掩码)
在这种情况下,我们建议在模型开头计算 BlockMask 并将其传入模型——在所有层中重用 BlockMask。
def forward(self, x, doc_mask):
# Compute block mask at beginning of forwards
block_mask = create_block_mask(doc_mask, None, None, S, S)
x = self.layer1(x, block_mask)
x = self.layer2(x, block_mask)
...
# amortize block mask construction cost across all layers
x = self.layer3(x, block_mask)
return x
掩码每层改变(例如数据依赖的稀疏性)
这是最难的情况,因为我们无法在多个 FlexAttention 调用中摊销块掩码计算。虽然 FlexAttention 在这种情况下肯定仍有益处,但 BlockMask 的实际效益取决于你的注意力掩码有多稀疏,以及我们构建 BlockMask 的速度有多快。这就引出了……
Q:我们如何更快地计算 BlockMask?
不幸的是,create_block_mask 在内存和计算方面都相当昂贵,因为确定一个块是否完全稀疏需要评估块中每个点的 mask_mod。有几种方法可以解决这个问题:
- 如果你的掩码在批次大小或头维度上是相同的,请确保你对它们进行了广播(即在
create_block_mask中将它们设置为None)。 - 编译
create_block_mask。不幸的是,目前由于一些局限性,torch.compile不能直接作用于create_block_mask。但是,你可以设置_compile=True,这将显著降低峰值内存和运行时间(在我们的测试中通常能降低一个数量级)。 - 为 BlockMask 编写自定义构造函数。BlockMask 的元数据非常简单(请查看文档)。它本质上是两个张量:a.
num_blocks:为每个查询块计算的 KV 块数量。
b.indices:为每个查询块计算的 KV 块的位置。例如,这是causal_mask的自定义 BlockMask 构造函数。
def create_causal_mask(S):
BLOCK_SIZE = 128
# The first query block computes one block, the second query block computes 2 blocks, etc.
num_blocks = torch.arange(S // BLOCK_SIZE, device="cuda") + 1
# Since we're always computing from the left to the right,
# we can use the indices [0, 1, 2, ...] for every query block.
indices = torch.arange(S // BLOCK_SIZE, device="cuda").expand(
S // BLOCK_SIZE, S // BLOCK_SIZE
)
num_blocks = num_blocks[None, None, :]
indices = indices[None, None, :]
return BlockMask(num_blocks, indices, BLOCK_SIZE=BLOCK_SIZE, mask_mod=causal_mask)
Q:为什么 score_mod 和 mask_mod 不同?难道 mask_mod 不是 score_mod 的特例吗?
这是一个非常敏锐的问题!事实上,任何 mask_mod 都可以轻松转换为 score_mod(我们不建议在实践中使用此函数!)。
def mask_mod_as_score_mod(b, h, q_idx, kv_idx):
return torch.where(mask_mod(b, h, q_idx, kv_idx), score, -float("inf"))
那么,如果 score_mod 可以实现 mask_mod 能做的所有事情,为什么还要 mask_mod 呢?
一个直接的挑战是:score_mod 需要实际的 score 值作为输入,但在我们预计算 BlockMask 时,我们没有实际的 score 值。我们可以通过传入全零来模拟这些值,如果 score_mod 返回 -inf,我们就认为它被屏蔽了(事实上,我们最初就是这样做的!)。
然而,有两个问题。首先,这很 hacky——如果用户的 score_mod 在输入为 0 时返回 -inf 怎么办?或者如果用户的 score_mod 用一个很大的负值而不是 -inf 来屏蔽怎么办?这似乎是在强行把方木楔入圆孔。然而,将 mask_mod 从 score_mod 中分离出来还有一个更重要的原因——它本质上效率更高!
事实证明,对计算出的每一个元素应用掩码实际上非常昂贵——我们的基准测试显示性能下降了约 15-20%!因此,虽然我们可以通过跳过一半计算获得显著加速,但我们需要为屏蔽每个元素付出代价!
幸运的是,如果我们可视化因果掩码,我们会注意到绝大多数块根本不需要“因果掩码”——它们是完全计算的!只有对角线上的块,因为是部分计算和部分掩码的,才需要应用掩码。

BlockMask 之前告诉我们哪些块需要计算,哪些块可以跳过。现在,我们进一步扩充了这个数据结构,告诉我们哪些块是“完全计算的”(即可以跳过掩码)与“部分计算的”(即需要应用掩码)。但请注意,尽管在“完全计算”的块上可以跳过掩码,但其他 score_mod(如相对位置编码)仍然需要应用。
仅给定一个 score_mod,我们没有可靠的方法来判断其中的哪些部分是“屏蔽”。因此,用户必须亲自将这些分离到 mask_mod 中。
Q:BlockMask 需要多少额外内存?
BlockMask 元数据的大小为 [BATCH_SIZE, NUM_HEADS, QUERY_LEN//BLOCK_SIZE, KV_LEN//BLOCK_SIZE]。如果掩码在批次或头维度上相同,它可以跨该维度广播以节省内存。
在默认的 BLOCK_SIZE 为 128 时,我们预计在大多数用例中内存使用量微不足道。例如,对于 100 万的序列长度,BlockMask 只会使用额外的 60MB 内存。如果这是一个问题,你可以增加块大小:create_block_mask(..., BLOCK_SIZE=1024)。例如,将 BLOCK_SIZE 增加到 1024 会使该元数据降至不到 1MB。
Q:数值计算结果如何比较?
尽管结果不是按位完全相同的,但我们确信 FlexAttention 在数值上与 FlashAttention 一样准确。我们在多种输入下对因果和非因果注意力变体进行了比较,FlashAttention 与 FlexAttention 的误差分布几乎完全一致。

性能
总的来说,FlexAttention 的性能几乎与手写的 Triton 内核一样,因为我们大量利用了手写的 Triton 内核。然而,由于其通用性,我们确实承担了小部分的性能损失。例如,我们必须增加一些延迟来确定接下来要计算哪个块。在某些情况下,我们提供了一些内核选项来在更改内核行为的同时影响其性能。它们可以在这里找到:性能旋钮。
作为案例研究,让我们探索旋钮如何影响因果注意力的性能。我们将比较在 A100 上 Triton 内核与 FlashAttentionv2 的性能。脚本可以在这里找到。
FlexAttention 在前向传播中达到了 FlashAttention2 性能的 90%,在反向传播中达到了 85%。FlexAttention 目前使用的是一种确定性算法,比 FAv2 重计算了更多的中间量,但我们计划改进 FlexAttention 的反向算法,并希望缩小这一差距!


结论
我们希望你们使用 FlexAttention 也能和我们开发它一样开心!在进行此项工作时,我们发现该 API 的应用远超我们的预期。我们已经看到它将 torchtune 的样本打包吞吐量提高了 71%,取代了研究人员花费一周多时间编写自定义 Triton 内核的需求,并提供了与自定义手写注意力变体相当的性能。
让 FlexAttention 的实现变得有趣的一点是,我们能够以有趣的方式利用许多现有的 PyTorch 基础架构。例如,TorchDynamo(torch.compile 的前端)的独特之处在于,它不需要在编译函数中使用的张量被显式地作为输入传入。这使我们能够编译像文档掩码这样的模组,它们需要访问全局变量,而全局变量是需要改变的!
bias = torch.randn(1024, 1024)
def score_mod(score, b, h, q_idx, kv_idx):
return score + bias[q_idx][kv_idx] # The bias tensor can change!
此外,torch.compile 作为一种通用的图捕获机制,也支持更“高级”的转换,例如将任何 mask_mod 转换为适用于锯齿状张量的高阶转换。
我们还利用了 TorchInductor(torch.compile 的后端)基础设施来实现 Triton 模板。这不仅使支持 FlexAttention 代码生成变得容易,还自动赋予了我们对动态形状以及尾部融合(即在注意力末尾融合算子)的支持!未来,我们计划扩展此支持,以允许注意力量化版本或诸如 RadixAttention 之类的技术。
此外,我们还利用了高阶算子、PyTorch 的自动求导来自动生成反向传播,以及 vmap 来自动应用 score_mod 以创建 BlockMask。
当然,如果没有 Triton 和 TorchInductor 生成 Triton 代码的能力,这个项目是不可能实现的。
我们期待在未来将这里使用的方法应用于更多的应用场景!
局限性与未来工作
- FlexAttention 目前已在 PyTorch 夜间版本中可用,我们计划在 2.5.0 版本中将其作为原型功能发布。
- 我们没有在这里涵盖如何使用 FlexAttention 进行推理(或如何实现 PagedAttention)——我们将在以后的文章中介绍。
- 我们正在努力改进 FlexAttention 的性能,以在 H100 GPU 上匹配 FlashAttention3。
- FlexAttention 要求所有序列长度必须是 128 的倍数——这个问题很快会得到解决。
- 我们计划很快添加对 GQA 的支持——目前,你可以简单地复制 kv 头。
致谢
我们要强调一些启发了 FlexAttention 的先前工作(和相关人员)。
- Tri Dao 在 FlashAttention 上的工作
- Francisco Massa 和 Xformers 团队在 Triton 中实现的 BlockSparseAttention
- Jax 团队在 SplashAttention 上的工作
- Philippe Tillet 和 Keren Zhou 在 Triton 方面给予我们的帮助
- Ali Hassani 关于邻域注意力的讨论
- 所有抱怨注意力内核不支持他们最喜欢的注意力变体的人 🙂