Flash Attention Code Study for Large Language Model(LLM).
从此前GPT-3 on Pytorch的测试结果可以看出Attention主要包含matmul->Dropout->Softmax->Mask->Matmul,其中时间占比并不是matmul最高,Dropout, Mask和Softmax占比也相当明显.

输入矩阵$\mathbf{Q}$,
综合而言,FA的核心思想包含两个: (a) 在前向和后向采用Tiling切分Softmax/score矩阵;(b) 在后向中采用重复计算("以算换存"). 此时的softmax公式在Tiling意义下则等价变换成(以二分块为例):
- Step 1:
$\mathbf{x} = [\mathbf{x^{(1)}} \ \mathbf{x^{(2)}}]$ ,将原始向量$\mathbf{x}$分裂为两个Block; - Step 2:
$m(\mathbf{x}) = \max_{i} (m(\mathbf(x^{(1)})), m(\mathbf(x^{(2)})))$ 等价; - Step 3:
$f(\mathbf{x}) = [ e^{m(\mathbf{x^{(1)}})-m(\mathbf{x)}}f(\mathbf{x^{(1)}}) \qquad e^{m(\mathbf{x^{(2)}})-m(\mathbf{x)}}f(\mathbf{x^{(2)}}) ]$ ; - Step 4:
$l(\mathbf{x}) = l([\mathbf{x^{(1)}} \ \mathbf{x^{(2)}}]) = e^{m(\mathbf{x^{(1)}})-m(\mathbf{x)}}l(\mathbf{x^{(1)}})+e^{m(\mathbf{x^{(2)}})-m(\mathbf{x)}}l(\mathbf{x^{(2)}})$ ; - Step 5: 计算$softmax(\mathbf{x}) = f(\mathbf{x}) / l(\mathbf{x})$.
将以上过程展开,看计算步骤:
假设: Tr = Tc = 4,则循环以上Tiling的计算步骤实际上是先向下行循环,此时第一列的O仅仅涉及本块单元信息.
第一列:
- 例如: 当i=1时
$\mathbf{S}{11}=\mathbf{Q}1{\mathbf{K}1}^T$ , (size: $Br\times Bc$),
$\mathbf{m}{11} = rowmax(\mathbf{S}{11})$ , (size: $Br$),
$\mathbf{P}{11} = exp(\mathbf{S}{11}-\mathbf{m}{11})$, (size:
$Br\times Bc$ ), $\mathbf{L}{11} = rowsum(\mathbf{P}{11})$, (size:$Br$ ) $\mathbf{O}1 = (\mathbf{P}{11} \mathbf{V}1)/\mathbf{L}{11}$, (size:$Br \times d$ ) $\mathbf{L}1 = \mathbf{L}{11}$, $\mathbf{m}1 = \mathbf{m}{11}$, (size:$Br$ )
第二列:
- 例如: 当i=1时
$\mathbf{S}{12}=\mathbf{Q}1{\mathbf{K}2}^T$ , (size: $Br\times Bc$),
$\mathbf{m}{12} = rowmax(\mathbf{S}{12})$ , (size: $Br$),
$\mathbf{P}{12} = exp(\mathbf{S}{12}-\mathbf{m}{12})$, (size:
$Br\times Bc$ ), $\mathbf{L}{12} = rowsum(\mathbf{P}{12})$, (size:$Br$ ) $\mathbf{O}1 = (exp(\mathbf{m}{11}) \mathbf{P}{11} \mathbf{V}1 + exp(\mathbf{m}{12}) \mathbf{P}{12} \mathbf{V}2)/(exp(\mathbf{m}{11})\mathbf{L}{11}+exp(\mathbf{m}{12})\mathbf{L}_{12})$, (size:$Br \times d$ )
或者对于任意块$(i,j0)$,$i \in [1,Tr]$,