[1706.03762] Attention Is All You Need

核心创新点:Transformer 架构和 Self-Attention 自注意力机制

Transformer 之前,主流序列转导模型通常具有以下特征:

  1. 主体结构多基于 RNN,CNN 作为替代路线
  2. 通常采用 encoder-decoder 架构
  3. 使用注意力机制增强

Transformer 结构的创新:

  1. 用 Self-Attention 取代 RNN/CNN 作为序列建模核心
  2. 引入 Multi-Head Attention,让模型从多个子空间看关系
  3. 用 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$ 步根据:

  1. encoder 给出的上下文向量 $c$
  2. 上一步生成的 token $y_{t-1}$
  3. 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 的所有隐藏状态

解决的问题:

  • 解决模型处理长序列时的“遗忘”问题

  • 解决不同时间步输入对当前时刻输出的“重要性”问题

    decoder 每生成一个词,都计算当前解码状态和所有 encoder 隐藏状态的相关性,然后加权汇总

    $$ c_t = \sum_{i=1}^{n} \alpha_{t,i} h_i $$

位置编码

因为没有 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有这么做过,直接用整数位置有几个问题:

  1. 标量位置太粗糙,难以和高维词向量融合
  2. 数值大小会随序列长度变大,尺度不稳定
  3. 不容易表达相对位置关系,比如“前一个词”“后两个词”

正弦/余弦编码的好处是:每个位置被表示成一个稳定的高维模式,而且不同频率可以覆盖不同尺度的位置关系

论文中给出的一个重要动机是:对于固定偏移 $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) # [B, L, 512]
K = k_proj(x) # [B, L, 512]
V = v_proj(x) # [B, L, 512]
# 分头
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

1
head_i: [B, L, 64]

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$ 越大,点积结果越容易变得很大或很小

为了把点积分数标准化到相对稳定的尺度,需要除以标准差

如果不除会导致

  1. attention 权重过早接近 one-hot,模型探索性变差
  2. 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(距离, 输入内容, 中间上下文)

可以实现根据内容选择性压缩历史