【FlashAttention】 FA2与FA1算法区别辨析

看了几篇关于FlashAttention2的文章,对于其中移除冗余的CUDA操作这个算法优化进行了一个综合梳理。

https://zhuanlan.zhihu.com/p/1993815603383902344

https://zhuanlan.zhihu.com/p/668888063

https://zhuanlan.zhihu.com/p/665170554

注意,第10行在部分文章中错写成了diag的逆,应该根据这篇文章的伪代码为准(推测是之前存在笔误,改了之后又重新上传了)。

这里FlashAttention2与FlashAttention1看起来有很大差别,推导如下;

  1. 首先比较重要的一点是,在FA2里,关于m, P的计算都没有mijm_{ij}mij, pijp_{ij}pij的概念,而是直接计算mim_imi和minewm_i^{new}minew,pip_ipi和pinewp_i^{new}pinew。因此此处的mijm_i^jmij就是FA1中的mijm_{ij}mij - minewm_i^{new}minew。另外此处的P也就是FA1中的emij−minew∗Pe^{m_{ij} - m_i^{new}} * Pemij−minew∗P。
  2. 另外第二个点,就是在中间的迭代中不计算L,只在最后一个迭代计算。
相关推荐
论文复现现场2 天前
MiniCPM-o 4.5 需要多少显存?RTX 3090 24GB 的 BF16、AWQ、GGUF 部署边界与验收方法
cuda·多模态大模型·显存优化·rtx3090·minicpm-o
qq_199886873 天前
第8板块·第2节:统一内存的高级特性与性能调优
c++·人工智能·gpu算力·cuda
qq_199886873 天前
第8板块·第3节:设备间拷贝与点对点(P2P)传输
c++·人工智能·gpu算力·cuda
qq_199886873 天前
第8板块·第1节:主机与设备内存管理基础与统一内存
c++·人工智能·gpu算力·cuda
qq_199886875 天前
第7板块·第1节:通用算子分类与设计模式
c++·人工智能·gpu算力·cuda
qq_199886875 天前
第7板块·第3节:CUTLASS 的 GEMM 实现与优化策略
c++·人工智能·gpu算力·cuda
Luchang-Li6 天前
CUDA stream创建和依赖,多流并行
stream·cuda·并行
论文复现现场7 天前
课程作业要跑 PyTorch 训练,学校机房不够用去哪租?云 GPU 选型、环境迁移与防丢数据指南
人工智能·pytorch·深度学习·云计算·gpu·cuda
论文复现现场14 天前
本地 PyTorch 训练 OOM,第一次租 RTX 4090 云 GPU 怎么迁移项目?从环境检查到 100 Step 跑通
人工智能·pytorch·python·深度学习·cuda
2601_9623878216 天前
Python CUDA 编程 - 2 - Numba 简介
python·numpy·cuda·jit编译器·numba