Skip to content

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

4 Commits
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

FlashAttention

Flash Attention Code Study for Large Language Model(LLM).

大模型算子

1. Attention

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

(1) Standard Attention

 输入矩阵$\mathbf{Q}$, $\mathbf{K}$, $\mathbf{V}$,大小均为$N \times d$,且初始存储在HBM上,其中$N$代表sequence length,$d$代表head dimension.  Step 1: 从HBM上Load $\mathbf{Q}$, $\mathbf{K}$(按照block方式),在运算单元上完成$\mathbf{S}=\mathbf{Q}\mathbf{K}^T$,进而将矩阵$\mathbf{S}$写入HBM中存储;  Step 2: 从HBM读取矩阵$\mathbf{S}$,进而在运算单元计算$\mathbf{P}=softmax(\mathbf{S})$,将矩阵$\mathbf{P}$写入HBM;  Step 3: 从HBM中Load矩阵$\mathbf{P}$和$\mathbf{V}$,然后完成计算$\mathbf{O}=\mathbf{PV}$,进而将矩阵$\mathbf{O}$写入HBM中,返回Attention矩阵$\mathbf{O}$.  从上面流程可以看到,矩阵$\mathbf{S}$和$\mathbf{P}$矩阵均为中间生成矩阵,相比于原生输入矩阵在HBM中本身占有的存储,这两个矩阵涉及跨越内存层级搬运,且总矩阵大小为$N*N$.  其中softmax的计算公式如下(这里向量$\mathbf{x}$大小是$B \times 1$): $$m(\mathbf{x}) = \max_{i} \ x_i, \qquad f(\mathbf{x}) = e^{\mathbf{x}-m(\mathbf{x})}, \qquad l(\mathbf{x})=\sum_i {f(x_i)}, \qquad softmax(\mathbf{x}) = f(\mathbf{x})/l(\mathbf{x})$$

(2) Flash Attention 1.0

 综合而言,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, ..., Tr$, $j=1$ $\mathbf{S}{i1}=\mathbf{Q}i{\mathbf{K}1}^T$ ,    (size: $Br\times Bc$), $\mathbf{m}{i1} = rowmax(\mathbf{S}{i1})$ ,    (size: $Br$), $\mathbf{P}{i1} = exp(\mathbf{S}{i1}-\mathbf{m}{i1})$,    (size: $Br\times Bc$), $\mathbf{L}{i1} = rowsum(\mathbf{P}{i1})$,    (size: $Br$) $\mathbf{O}i = (\mathbf{P}{i1} \mathbf{V}1)/\mathbf{L}{i1}$,    (size: $Br \times d$) $\mathbf{L}i = \mathbf{L}{i1}$, $\mathbf{m}i = \mathbf{m}{i1}$,    (size: $Br$)

  • 例如: 当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, ..., Tr$, $j=1$ $\mathbf{S}{i2}=\mathbf{Q}i{\mathbf{K}2}^T$ ,    (size: $Br\times Bc$), $\mathbf{m}{i2} = rowmax(\mathbf{S}{i2})$ ,    (size: $Br$), $\mathbf{P}{i2} = exp(\mathbf{S}{i2}-\mathbf{m}{i2})$,    (size: $Br\times Bc$), $\mathbf{L}{i2} = rowsum(\mathbf{P}{i2})$,    (size: $Br$) $\mathbf{O}i = (exp(\mathbf{m}{i1}) \mathbf{P}{i1} \mathbf{V}1 + exp(\mathbf{m}{i2}) \mathbf{P}{i2} \mathbf{V}2)/(exp(\mathbf{m}{i1})\mathbf{L}{i1}+exp(\mathbf{m}{i2})\mathbf{L}_{i2})$,    (size: $Br \times d$)

  • 例如: 当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]$, $j0 \in [1,Tc]$其Attention Matrix$((Br*Tr)*d)$的输出结果: $$O_i = \frac{\sum_{j=1}^{j=j0} {(e^{\mathbf{S}{ij}}\mathbf{V}j)}} {rowsum({\sum{j=1}^{j=j0} {e^{\mathbf{S}{ij}}}})}$$

About

Flash Attention Code Study for Large Language Model(LLM).

Resources

Stars

Watchers

Forks

Releases

Packages

Contributors

Languages