取而代之的是两个很小的 Python 函数。一个负责调整注意力分数(score_mod),加位置偏置、做 soft-capping,你的变体需要什么就写什么。另一个只回答一个问题(mask_mod):这个位置到底要不要算。torch.compile 把两者编译成一个融合的 Triton 内核,反向传播由 autograd 自动生成。
最有意思的一点是:为什么需要两个函数,而不是一个。改分数的那个其实已经能表达任何掩码,返回负无穷,这个位置就会在 softmax 中被丢掉。但那是先把所有东西都算出来,再扔掉。把掩码单独声明出来,PyTorch 才能在计算之前就整块跳过,在因果掩码上大约值 2 倍性能;而编译器无法从任意一个 Python 函数里自己还原出这个意图。
十分钟里我们会写一个出来,用 create_block_mask 看看它生成的块结构,以及它的代价:按 PyTorch 自己的基准,前向约为 FlashAttention-2 的 90%,反向 85%。
听完你将能用纯 Python 实现一个没有现成内核的注意力变体,并清楚自己付出了什么代价。