欢迎光临
我们一直在努力

论文精读 - Attention Is All You Need

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(?)

代码复现和实验结果

后续更新

赞(0)
未经允许不得转载:171主机测评 » 论文精读 - Attention Is All You Need
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址