返回题库
#5困难

实现 FlashAttention-2 前向

FlashAttentionAttentionIO-aware

题目描述

实现 FlashAttention-2 的前向计算(单头,简化版)。给定 Qseq_len × d)、KV,用分块 + 在线 softmax计算输出,不实例化完整的 seq_len × seq_len 注意力矩阵。

要求

  • 实现右侧代码模板中给出的函数,签名以模板为准
  • 使用 online softmax,维护 running max 和 running sum
  • 缩放因子为

提示

FlashAttention-2 相比 v1 减少了非矩阵乘的 rescale 操作,并把并行维度放在 seq_len 上。

示例

示例 1

输入:
seq_len=1, d=1, Q=[1], K=[1], V=[5]
输出:
[5.0]

说明:单 token,softmax 后权重为 1

示例 2

输入:
seq_len=2, d=1, Q=[0,0], K=[0,0], V=[1,3]
输出:
[2.0, 2.0]

说明:scores 全 0,softmax 均匀,输出为 V 均值

讨论区(0)

还没有评论,来做第一个发言的人吧

Python · solution.py
编辑器加载中…

运行结果

点击「调试」(只跑样例)或「提交」(跑全部用例)查看结果