https://huggingface.co/papers/1706.03762
这篇文章主要贡献是提出了Transformer架构,用attention mechanisms替代了传统的RNN和CNN,从而实现更强的并行能力
Transformer整体还是采用了encoder-decoder的结构,是因为机器翻译本质是序列到序列任务,需要将输入序列编码为语义表示,再基于该表示进行有条件的逐步生成,Transformer的创立动机就是因为传统的RNN的输入状态取决于上一轮的输出状态,使得RNN的计算必须串行执行,这也导致了RNN的训练非常缓慢,因此创立Transformer大幅提升了模型训练的并行度
而attention可以让输入输出关联起来,不会受到序列距离的限制,实现了token之间的动态关系建模
因此基本的架构图如下

Attention
attention可以抽象为在数据库中做查询,也就是通常由query,key,value组成,在一般我们写程序中,我们可以用判断语句来比较每个key是否符合要求,再进行下一步,但是在有时候问题抽象到我们无法用简单的判断语句来比较,这时候我们的方法是把query和key各建模成一个向量,然后对query和key之间计算一个相似度,以这个相似度为权重,计算value的加权和。attention指的就是这里的权重
所以我们的实现只需要对query和key做矩阵乘法然后用softmax把权重归一化最后对value进行加权求和就可以
Scaled Dot-Product Attention
刚才说的计算方法就是Transformer中的attention,这种计算在论文里称为放缩点乘注意力,在实现中,为了提高效率,我们不会逐个计算查询,而是将所有查询打包成矩阵Q,键和值分别组成矩阵K和V,然后通过一次矩阵乘法就可以同时完成所有位置的注意力计算,整体形式如下
$$ \\mathrm{Attention}(Q, K, V) = \\mathrm{softmax}\\left(\\frac{QK^T}{\\sqrt{d_k}}\\right) V $$
这里的 $d_k$ 就是query和key向量的长度, 由于query和key向量要做点乘,这两种向量的长度必须一致
我们现在都知道$QK^T$是计算相似度,在点积很大时,softmax对于大数很敏感,会导致梯度消失,因此用 $\\sqrt{d_k}$ 来缩放这个量,提升了模型训练的稳定性

class SingleHeadAttention(nn.Module):
def __init__(self, d_model, d_k, d_v, dropout=0.1):
super().__init__()
self.W_q = nn.Linear(d_model, d_k, bias=False)
self.W_k = nn.Linear(d_model, d_k, bias=False)
self.W_v = nn.Linear(d_model, d_v, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
"""
x: [batch_size, seq_len, d_model]
mask: [batch_size, seq_len, seq_len]
"""
Q = self.W_q(x)
K = self.W_k(x)
V = self.W_v(x)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(K.size(-1))
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
output = torch.matmul(attn_weights, V)
return output, attn_weights
Multi-Head Attention
单头注意力只能在一个表示空间里计算谁该关注谁,当有好几个类型的信息时候效果会很差,而多头注意力可以让每个头去关注不同类别的关系,在单头注意力中,所有信息都在一个 $d_{model}$ 维空间计算,而多头注意力其实就是将所有输出的head连接到一起,如下
$$ \\mathrm{MultiHead}(Q, K, V) = \\mathrm{Concat}(head_1, head_2, \\ldots, head_n) W^O $$
$$ head_i = \\mathrm{Attention}(QW_i^Q, KW_i^K, VW_i^V) $$

class MultiHeadAttention(nn.Module):
def __init__(self, d_model, num_heads, dropout=0.1):
super().__init__()
assert d_model % num_heads == 0
self.h = num_heads
self.d_k = d_model // num_heads
self.W_q = nn.Linear(d_model, d_model, bias=False)
self.W_k = nn.Linear(d_model, d_model, bias=False)
self.W_v = nn.Linear(d_model, d_model, bias=False)
self.W_o = nn.Linear(d_model, d_model, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
B, T, _ = x.shape
Q, K, V = [
proj(x).view(B, T, self.h, self.d_k).transpose(1, 2)
for proj in (self.W_q, self.W_k, self.W_v)
]
scores = (Q @ K.transpose(-2, -1)) / math.sqrt(self.d_k)
attn = self.dropout(F.softmax(scores, dim=-1))
out = (attn @ V).transpose(1, 2).contiguous().view(B, T, -1)
return self.W_o(out), attn
Transformer 架构
下面就来看下Transformer的具体每一层是什么样的,首先先看下Encoder和Decoder是怎么组织起来的
残差连接
残差连接的核心思想就是在每个子层中,将输入x和子层的输出Sublayer(x)相加,然后通过LayerNorm来进行归一化,从而让模型只需要学习对原始输入的“增量修正”,而不是重新学习整个映射
-
LayerNorm:比如给一个向量$[x_1,x_2,x_3,…,x_d]$,LayerNorm先对每个向量求均值和方差,将其标准化为均值为0、方差为1的分布,然后再通过变换,从而在稳定数值分布的同时保留模型的表达能力
前馈网络
前馈网络(Feed Forward)就是在每个token上独立应用的两层全连接网络,中间用ReLU作为激活函数
$$ \\mathrm{FFN}(x) = \\max(0, xW_1 + b_1) W_2 + b_2 $$
class FeedForward(nn.Module):
def __init__(self, d_model, d_ff=2048, dropout=0.1):
super().__init__()
self.net = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.ReLU(),
nn.Linear(d_ff, d_model),
nn.Dropout(dropout)
)
def forward(self, x):
return self.net(x)
掩码多头注意力
掩码多头注意力(Masked Multi-Head Attention就是在多头注意力中,通过mask让每个位置只能看到自己和之前的位置,不能看到未来信息,在attention里有一步 $\\frac{QK^T}{\\sqrt{d_k}}$,简单说,掩码注意力就是把未来位置没发生的分数设为-INF,比如原始的attention score是[2, 3, 5],在mask之后就变成了[2, 3, – INF]
位置编码
因为我们的模型中不包含RNN和CNN,而Transformer的主干网络不能利用到序列顺序信息,因此引入了“位置编码”这样的机制,能够向模型中注入关于token的相对或绝对位置信息
在这篇文章中,利用正弦函数和余弦函数来构造位置编码,公式如下
$$ PE(pos, 2i) = \\sin\\left(\\frac{pos}{10000^{\\frac{2i}{d_{\\text{model}}}}}\\right) $$
$$ PE(pos, 2i+1) = \\cos\\left(\\frac{pos}{10000^{\\frac{2i}{d_{\\text{model}}}}}\\right) $$
总结
Transformer用attention机制无疑在翻译方面远胜过RNN和CNN,但是我认为在未来还是有许多优化的空间,例如self attention的复杂度是 $O(n^2)$,没有天然的顺序偏置,需要靠位置编码,以及对于FNN的改进,例如可能用GELU替代ReLU(?)
代码复现和实验结果
后续更新





