用 FlexAttention 在纯 Python 里编写自定义注意力内核

主题演讲
上海
15:50 - 16:30
分会场 C(综合楼 421 会议室)
  • Himanshu Sangshetti Mem0 AI 工程师

    Himanshu Sangshetti 是 Mem0 的 AI 工程师,负责面向 AI 智能体的开源记忆层(GitHub 60k+ stars)。

    他的工作处在智能体与基础设施的交叉点,涉及记忆系统、LLM 推理性能,以及围绕这两者的 Python 工具链。他是 AWS Community Builder 与 HashiCorp Ambassador,并创办了免费云学习平台 CloudFluently(10,000+ 学习者)。

    他在印度浦那组织 Build Club AI 社区(1,000+ 成员),同时也是 Cursor Ambassador,负责举办 Cursor 社区活动。他已在全球社区完成 30+ 场技术分享。

    himanshu

摘要

写一个新的注意力变体,过去意味着要写一个新的 CUDA 内核。滑动窗口、ALiBi、文档掩码、PrefixLM,如果你的想法没有对应的现成内核,就只能忍受又慢又占显存的普通 PyTorch 实现。FlexAttention 去掉了这个限制。

详情

取而代之的是两个很小的 Python 函数。一个负责调整注意力分数(score_mod),加位置偏置、做 soft-capping,你的变体需要什么就写什么。另一个只回答一个问题(mask_mod):这个位置到底要不要算。torch.compile 把两者编译成一个融合的 Triton 内核,反向传播由 autograd 自动生成。

最有意思的一点是:为什么需要两个函数,而不是一个。改分数的那个其实已经能表达任何掩码,返回负无穷,这个位置就会在 softmax 中被丢掉。但那是先把所有东西都算出来,再扔掉。把掩码单独声明出来,PyTorch 才能在计算之前就整块跳过,在因果掩码上大约值 2 倍性能;而编译器无法从任意一个 Python 函数里自己还原出这个意图。

十分钟里我们会写一个出来,用 create_block_mask 看看它生成的块结构,以及它的代价:按 PyTorch 自己的基准,前向约为 FlashAttention-2 的 90%,反向 85%。

听完你将能用纯 Python 实现一个没有现成内核的注意力变体,并清楚自己付出了什么代价。