[1706.03762] Attention Is All You Need
核心创新点:Transformer 架构和 Self-Attention 自注意力机制
Transformer 之前,主流序列转导模型通常具有以下特征:
- 主体结构多基于 RNN,CNN 作为替代路线
- 通常采用 encoder-decoder 架构
- 使用注意力机制增强
Transformer 结构的创新:
- 用 Self-Attention 取代 RNN/CNN 作为序列建模核心
- 引入 Multi-Head Attention,让模型从多个子空间看关系
- 用 Positional Encoding 补足序列位置信息
Transformer是首个完全基于自注意力、不使用RNN或CNN的序列建模模型
Encoder-Decoder 架构
早期 seq2seq 的核心思想:encoder 先压缩,decoder 再生成
- 编码器负责“读懂输入序列”,把它压缩成一个向量
- 解码器负责“根据这个向量逐步生成输出序列
编码器解码器结构解决了输入输出不等长的问题
早期 encoder 通常是一个 RNN / LSTM / GRU,编码器按顺序读入 token
每读一个 token,更新一次隐藏状态
$$
h_t = f(h_{t-1}, x_t)
$$
最后,encoder 会得到一个整体表示,通常记作$c = h_n$ (context vector),也可以理解成“源句子的压缩表示”
解码器 decoder 在第 $t$ 步根据:
- encoder 给出的上下文向量 $c$
- 上一步生成的 token $y_{t-1}$
- decoder 自己的隐藏状态 $s_{t-1}$
预测当前 token $P(y_t \mid y_{<t}, c)$
早期模型最大的问题是:不管输入句子多长,都要压缩进一个固定长度向量 $c$
这会导致:
- 长句信息容易丢失;
- 远距离依赖难以保存;
- decoder 生成后半句时,可能已经“忘了”输入前面的内容;
- 所有源句信息都挤在一个固定向量里,形成 bottleneck
Attention 机制
加入 attention 之后
1 2 3 4 5 6 7
| source sequence -> Encoder -> h1, h2, ..., hn | Attention | Decoder | target sequence
|
decoder 不再只看一个固定向量 $c$,而是在生成每个词时动态查看 encoder 的所有隐藏状态
解决的问题:
位置编码
因为没有 RNN 的时间递推,也没有 CNN 的局部窗口顺序结构,Self-Attention 本身不知道 token 的顺序
在 encoder 和 decoder 底部给输入 embedding 加入 positional encodings,用来注入 token 的位置信息
$$
X_{pos} = E_{pos} + PE_{pos}
$$
原始 Transformer 用的正弦/余弦位置编码
$$
PE_{(pos,2i)}=\sin
\left(\frac{pos}{10000^{2i/d_{\text{model}}}}\right)\\
PE_{(pos,2i+1)}=
\cos\left(\frac{pos}{10000^{2i/d_{\text{model}}}}\right)
$$
假设模型维度是 $d_{\text{model}} = 4$,那么一个位置 $pos$ 的位置编码大概长这样
$$
PE_{pos}=[\sin(\omega_1 pos),\cos(\omega_1 pos),\sin(\omega_2 pos),\cos(\omega_2 pos)]
$$
不同维度使用不同频率,
- 低维度:变化快,区分近距离位置
- 高维度:变化慢,表达长距离趋势
为什么不用简单的 1, 2, 3, 4?
Bert有这么做过,直接用整数位置有几个问题:
- 标量位置太粗糙,难以和高维词向量融合
- 数值大小会随序列长度变大,尺度不稳定
- 不容易表达相对位置关系,比如“前一个词”“后两个词”
正弦/余弦编码的好处是:每个位置被表示成一个稳定的高维模式,而且不同频率可以覆盖不同尺度的位置关系
论文中给出的一个重要动机是:对于固定偏移 $k$,$PE_{pos+k}$ 可以表示为 $PE_{pos}$ 的线性函数,因此模型可能更容易学习相对位置关系
多头注意力机制
多头注意力 Multi-Head Attention 的核心思想是:把同一批 token 映射到多个不同的表示子空间中,并行做多次 attention,然后把结果拼接起来
Transformer 使用的基础注意力是 Scaled Dot-Product Attention
$Q,K,V$ 都来自同一个输入序列,让输入序列内部的 token 互相建模关系
$$
\mathrm{Attention}(Q,K,V)=\mathrm{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V
$$
$$
\mathrm{Softmax}(z_i)=\frac{e^{z_i}}{\sum_{j=1}^{k} e^{z_j}}
$$
QK^T:计算每个 token 对其他 token 的关注程度
- softmax:把关注分数归一化成权重
- 乘 V:根据权重加权汇聚信息
经过这个Attention以后的向量信息会融合其它token的信息
如果只有一个 attention head,那么模型只有一套 $Q,K,V$ 投影,那么token只能用一种方式去理解自己
但是语言中token的关系从来不是一维的,单头使得被迫将不同层面的关系压缩进一组权重,注意力权重变成一个折中的分布,一个注意力分布无法承载多种独立的关系模式,表达能力被严重限制
多头注意力的设计就是:允许模型在多个子空间中并行学习不同的相关性模式
多头注意力不是直接对原始 $Q,K,V$ 做很多次一样的 attention,而是先用不同的线性变换得到不同 head 的 $Q_i,K_i,V_i$
1 2 3 4 5 6 7 8 9 10 11
| q_proj = nn.Linear(d_model, d_model) k_proj = nn.Linear(d_model, d_model) v_proj = nn.Linear(d_model, d_model)
Q = q_proj(x) K = k_proj(x) V = v_proj(x)
Q = Q.view(B, L, num_heads, d_head).transpose(1, 2) K = K.view(B, L, num_heads, d_head).transpose(1, 2) V = V.view(B, L, num_heads, d_head).transpose(1, 2)
|
第 $i$ 个 head 是
$$
\mathrm{head}_i=\mathrm{Attention}\left(QW_i^Q,KW_i^K,VW_i^V\right)
$$
然后把所有 head 的输出拼接起来
$$
\mathrm{MultiHead}(Q,K,V)=\mathrm{Concat}(head_1,\dots,head_h)W^O
$$
再经过一个线性变换,融合多头信息
$$
\mathrm{FFN}(x)=W_2\,\sigma(W_1x+b_1)+b_2
$$
> 相当于八个低清晰度的表情包你给拼一块了说它是高清的,那肯定不行,还得再处理一下
这种模式使得每个头只需要捕获一种或少量关系模式
这里也会有一个问题,拆成多个头降低了维度,可能会损失表达能力,但损失的远小于收益
拆分成多头在同等计算量上获取到的信息大于单头,并且低维子空间相当于隐含的正则化,保证每种学习关系独立,也防止了每个头在高维空间过拟合,多头保证了多种关系可以并行捕捉
head数量需要根据实际情况决定,head 太少,表达能力可能不足;head 太多,每个 head 的维度会变小,单个 head 的表示能力下降,而且计算和工程开销增加
通常 $h$ 是一个超参数,需要和 $d_{\text{model}}$ 配合
Mask
mask 不是加在 $Q$、$K$、$V$ 上,而是加在 attention logits 上
$$
\mathrm{MaskedAttention}(Q,K,V)=
\mathrm{softmax}\left(\frac{QK^\top}{\sqrt{d_k}} + M\right)V
$$
其中 $M$ 是 mask 矩阵
$$
M_{ij}=
\begin{cases}
0, & j \le i \\
-\infty, & j > i
\end{cases}
$$
这样可以获得每个 token 根据注意力权重,从它能看到的 token 的 value 向量中汇聚出来的新表示
加Mask是为了避免关注到后面的信息,只注意自己和之前的信息
张量形状
每个 token 原始 hidden state 的维度是
$$
X \in \mathbb{R}^{B \times L \times d_{\text{model}}}
$$
通过线性层得到
1 2 3
| Q: [B, L, d_model] K: [B, L, d_model] V: [B, L, d_model]
|
然后拆成 $h$ 个 head:
$$
d_{\text{head}} = \frac{d_{\text{model}}}{h}
$$
如果 $d_{\text{model}}=512$,$h=8$,那么
1 2 3
| Q: [B, h, L, d_head] = [B, 8, L, 64] K: [B, h, L, d_head] = [B, 8, L, 64] V: [B, h, L, d_head] = [B, 8, L, 64]
|
每个 head 单独做 attention
8 个 head 拼接后
1
| Concat(heads): [B, L, 512]
|
最后再经过输出投影 $W^O$
分母 $\sqrt{d_k}$
注意力分数是:
$$
q \cdot k = \sum_{i=1}^{d_k} q_i k_i
$$
假设$q_i, k_i$都是均值为 $0$、方差为 $1$ 的独立随机变量
那么单项乘积 $q_i k_i$ 的均值大约是 $0$,方差大约是 $1$
点积是 $d_k$ 个这样的项相加,所以方差大约是
$$
\mathrm{Var}(q \cdot k) \approx d_k
$$
$d_k$ 越大,点积结果越容易变得很大或很小
为了把点积分数标准化到相对稳定的尺度,需要除以标准差
如果不除会导致
- attention 权重过早接近 one-hot,模型探索性变差
- softmax 进入饱和区,梯度很小,训练变慢或不稳定
本质上类似于初始化中的方差归一化思想
Post-LN和Pre-LN的区别
Post-LayerNorm是原文设计
1 2 3 4 5 6 7 8 9
| x │ ├── Sublayer(x) (Attention 或 FFN) │ ├── + residual (x) │ └── LayerNorm │ y
|
Attention block:
$$
x_1 = \mathrm{LayerNorm}(x+\mathrm{Attention}(x))
$$
FFN block:
$$
x_2 = \mathrm{LayerNorm}(x_1+\mathrm{FFN}(x_1))
$$
可以这么理解
$$
y = \mathrm{LN}(x+F(x))\quad z= x+F(x)
$$
损失对输入的梯度:
$$
\frac{\partial L}{\partial x}=\frac{\partial L}{\partial y} \cdot \frac{\partial y}{\partial x} = \frac{\partial L}{\partial y} \cdot \frac{\partial \mathrm{LN}}{\partial z}(I
+\frac{\partial F}{\partial x})
$$
问题在于LN在在残差连接之后,${\partial \mathrm{LN}}/{\partial z}$梯度值近似为$1/\sqrt d_k$,使得残差传播的恒等路径梯度不再是1,每层都要缩放一次
所以梯度随网络深度呈指数衰减,导致低层(靠近输入的层)梯度几乎消失,梯度消失会导致Adam等优化器的更新变得不稳定
Layer Norm
对每个 token 的 hidden dimension 做归一化
- 不跨 batch 归一化
- 不跨 sequence length 归一化
- 只在每个 token 自己的特征维度上归一化
$$
\mu = \frac{1}{d}\sum_{i=1}^{d} x_i\qquad \sigma^2 = \frac{1}{d}\sum_{i=1}^{d}(x_i - \mu)^2
$$
归一化:
$$
\hat{x}_i = \frac{x_i - \mu}{\sqrt{\sigma^2 + \epsilon}}
$$
不是像 Batch Norm 那样依赖 mini-batch 统计
Transformer 需要 Layer Norm是因为每一层都有复杂变换
如果不做归一化,层数加深后 hidden states 的尺度可能变得不稳定
Layer Norm 的作用是把每个 token 的 hidden state 拉回到相对稳定的尺度
Post-LayerNorm
Post-LN 是原始 Transformer 的写法;Pre-LN 是后来更常用的稳定训练写法
1 2 3 4 5 6 7 8 9
| x │ ├── Sublayer(x) (Attention 或 FFN) │ ├── + residual (x) │ └── LayerNorm │ y
|
Attention block:
$$
x_1 =\mathrm{LayerNorm}(x+\mathrm{Attention}(x))
$$
FFN block:
$$
x_2 = \mathrm{LayerNorm}(x_1+\mathrm{FFN}(x_1))
$$
| 优点 | 缺点 |
| ------------------------ | ------------------------------------------------------------ |
| 每一层输出都被及时归一化 | 训练初期更容易不稳定
深层模型难训练(论文原文中提及)
对超参数(学习率、初始化、warm-up)更敏感 |
Pre-LayerNorm
把 LayerNorm 移到子层前面
1 2 3 4 5 6 7 8 9
| x │ ├── LayerNorm │ ├── Sublayer │ ├── + residual │ └── y
|
Attention block:
$$
x_1 =x+\mathrm{Attention}( \mathrm{LayerNorm}(x))
$$
FFN block:
$$
x_2 = x_1+\mathrm{FFN}(\mathrm{LayerNorm}(x_1))
$$
最关键的在于这样使得残差路径变成恒等映射
$$
y = x+F(\mathrm{LN}(x))\quad u = \mathrm{LN}(x)
$$
损失对输入的梯度:
$$
\frac{\partial L}{\partial x}=\frac{\partial L}{\partial y} \cdot \frac{\partial y}{\partial x} = \frac{\partial L}{\partial y}(I+\frac{\partial F}{\partial u}\cdot \frac{\partial u}{\partial x})
$$
| 优点 | 缺点 |
| ------------------------------------------------------------ | -------------------------- |
| 深层模型更容易训练
对超参数(学习率、初始化、warm-up)更鲁棒
残差连接保留信息更直接,梯度传播更稳定 | 深层更新可能被残差主干稀释 |
Mamba
对于transformer,自注意力需要计算所有token的关系,那么 $QK^T$ 的大小为 $L\times L$,标准注意力的复杂度为 $O(L^2d)$
当上下文非常长时,例如数十万甚至上百万 token,计算量和显存占用都会明显增加
Mamba 的目标是把序列复杂度降低到 $O(Ld)$
基础是状态空间模型(State Space Model, SSM)
[2111.00396] Efficiently Modeling Long Sequences with Structured State Spaces 别名 S4
SSM 的基本结构
1 2 3 4 5 6 7 8 9 10 11
| +-------------+ | h(t-1) | +------+------+ | | A: state transition v +-------------+ B: write +-------------+ C: read +-------------+ | x(t) |------------->| h(t) |------------->| y(t) | +------+------+ +-------------+ +------^------+ | | +--------------------- D: skip connection ------------------+
|
$$
\frac{\mathrm d h(t)}{\mathrm d t}=Ah(t)+Bx(t),
\qquad
y(t)=Ch(t).
$$
S4 使用的是双线性离散化
$$
\overline A=\left(I-\frac{\Delta}{2}A\right)^{-1}\left(I+\frac{\Delta}{2}A\right) \qquad
\overline B=\left(I-\frac{\Delta}{2}A\right)^{-1}\Delta B\qquad
\overline C = C
$$
参数一般不随序列位置变化
$$
\overline A_t=\overline A,\qquad \overline B_t=\overline B \qquad C_t=C
$$
Mamba 采用零阶保持法(ZOH)离散化
$$
\overline{A}_t=\exp(\Delta_t A) \qquad \overline{B}_t=(\Delta_tA)^{-1}\left(\exp(\Delta_tA)-I\right)\Delta_tB_t
$$
然后得到真正执行的离散递推
$$
h_t=\overline{A}_t h_{t-1}+\overline{B}_t x_t\qquad y_t=C_t h_t
$$
对于S4 $A, B, C ,\Delta $对于所有时间步都是固定的,$x_i$ 对输出的影响只取决于二者的距离,相当于固定卷积核
对于Mamba
| 参数 |
功能 |
直观含义 |
| $(\Delta_t)$ |
控制状态更新和遗忘速度 |
什么时候保留、什么时候刷新 |
| $(B_t)$ |
控制输入如何写入状态 |
当前 token 写入哪些状态维度 |
| $(C_t)$ |
控制状态如何被读出 |
当前输出需要读取哪些信息 |
| $(A)$ |
定义基础状态动力学 |
信息在状态中的衰减和传播规律 |
$A$ 仍然不随 token 改变,$B_t, C_t, \Delta _t$ 依赖当前位置的输入特征,形成了一个动态卷积核
1 2 3 4 5
| 普通 S4: 权重 = f(距离)
Mamba: 权重 = f(距离, 输入内容, 中间上下文)
|
可以实现根据内容选择性压缩历史