1 Transformer结构
Transformer作为编码器-解码器架构的一个实例。正如所见到的,Transformer是由编码器和解码器组成的。与基于Bahdanau注意力实现的序列到序列的模型相比,Transformer的编码器和解码器是基于自注意力的模块叠加而成的,源(输入)序列和目标(输出)序列的嵌入(embedding)表示将加上位置编码(positional encoding),再分别输入到编码器和解码器中。
2 编码器
从宏观角度来看,Transformer的编码器是由多个相同的层叠加而成的,每个层都有两个子层(子层表示为\mathrm{sublayer})。第一个子层是多头自注意力(multi-head self-attention)汇聚;第二个子层是基于位置的前馈网络(positionwise feed-forward network)。具体来说,在计算编码器的自注意力时,查询、键和值都来自前一个编码器层的输出。受残差网络的启发,每个子层都采用了残差连接(residual connection)。在Transformer中,对于序列中任何位置的任何输入\mathbf{x} \in \mathbb{R}^d,都要求满足\mathrm{sublayer}(\mathbf{x}) \in \mathbb{R}^d,以便残差连接满足\mathbf{x} + \mathrm{sublayer}(\mathbf{x}) \in \mathbb{R}^d。在残差连接的加法计算之后,紧接着应用层规范化(layer normalization)。因此,输入序列对应的每个位置,Transformer编码器都将输出一个d维表示向量。
2.1 编码器子层
class EncoderBlock(nn.Module):
"""Transformer编码器块"""
def __init__(self, key_size, query_size, value_size,
num_hiddens, norm_shape, ffn_num_input, ffn_num_hiddens,
num_heads, dropout, use_bias=False, **kwargs):
super(EncoderBlock, self).__init__(**kwargs)
# 多头注意力层 (会在2.1.1中介绍到)
self.attention = MultiHeadAttention(key_size, query_size, value_size, num_hiddens, num_heads, dropout, use_bias)
# 残差连接 + 层归一化 (会在2.1.2中介绍到)
self.addnorm1 = AddNorm(norm_shape, dropout)
# 基于位置的前馈网络(会在2.1.3中介绍到)
self.ffn = PositionWiseFFN(ffn_num_input, ffn_num_hiddens, num_hiddens)
# 残差连接 + 层归一化
self.addnorm2 = AddNorm(norm_shape, dropout)
def forward(self, X, valid_lens):
# 这里的attention传入的Q K V都是X,其实就是自注意力
Y = self.addnorm1(X, self.attention(X, X, X, valid_lens))
return self.addnorm2(Y, self.ffn(Y))
2.1.1 多头注意力 (MultiHeadAttention)
工具函数:DotProductAttention
作用是计算缩放点积注意力(是一种注意力评分函数)
def masked_softmax(X, valid_lens):
"""通过在最后一个轴上掩蔽元素来执行softmax操作"""
# X:3D张量,valid_lens:1D或2D张量
if valid_lens is None:
return nn.functional.softmax(X, dim=-1)
else:
shape = X.shape
if valid_lens.dim() == 1:
valid_lens = torch.repeat_interleave(valid_lens, shape[1])
else:
valid_lens = valid_lens.reshape(-1)
# 最后一轴上被掩蔽的元素使用一个非常大的负值替换,从而其softmax输出为0
X = d2l.sequence_mask(X.reshape(-1, shape[-1]), valid_lens, value=-1e6)
return nn.functional.softmax(X.reshape(shape), dim=-1)
class DotProductAttention(nn.Module):
"""缩放点积注意力"""
def __init__(self, dropout, **kwargs):
super(DotProductAttention, self).__init__(**kwargs)
self.dropout = nn.Dropout(dropout)
# queries的形状:(batch_size,查询的个数,d)
# keys的形状:(batch_size,“键-值”对的个数,d)
# values的形状:(batch_size,“键-值”对的个数,值的维度)
# valid_lens的形状:(batch_size,)或者(batch_size,查询的个数)
def forward(self, queries, keys, values, valid_lens=None):
d = queries.shape[-1]
# 设置transpose_b=True为了交换keys的最后两个维度
scores = torch.bmm(queries, keys.transpose(1,2)) / math.sqrt(d)
self.attention_weights = masked_softmax(scores, valid_lens)
return torch.bmm(self.dropout(self.attention_weights), values)```
工具函数:transpose_qkv、transpose_output
在普通注意力计算中,传入的Q K V单个数据的维度是num_hiddens,就是隐藏层维度。在多头注意力的计算中,对“多头”的划分,实际是对隐藏状态的划分,也就是说将原本的数据形状:
[batch_size,查询或者“键-值”对的个数,num_hiddens]
变化为
[batch_size,查询或者“键-值”对的个数,num_heads, num_hiddens/num_heads]
但是为了方便并行计算,这里把形状继续变换为:
[batch_size*num_heads,查询或者“键-值”对的个数,num_hiddens/num_heads]
所以这里两个函数作用就是,一个将它变换过去,一个将它再变换回来
def transpose_qkv(X, num_heads):
"""为了多注意力头的并行计算而变换形状"""
# X的形状: (batch_size, 查询或者“键-值”对的个数, num_hiddens)
# X的形状变为: (batch_size, 查询或者“键-值”对的个数, num_heads, num_hiddens/num_heads)
X = X.reshape(X.shape[0], X.shape[1], num_heads, -1)
# X的形状变为: (batch_size, num_heads, 查询或者“键-值”对的个数, num_hiddens/num_heads)
X = X.permute(0, 2, 1, 3)
# 最终输出的形状: (batch_size*num_heads, 查询或者“键-值”对的个数, num_hiddens/num_heads)
return X.reshape(-1, X.shape[2], X.shape[3])
def transpose_output(X, num_heads):
"""逆转transpose_qkv函数的操作"""
# X的形状: (batch_size*num_heads, 查询的个数, num_hiddens/num_heads)
# X的形状变为: (batch_size, num_heads, 查询的个数, num_hiddens/num_heads)
X = X.reshape(-1, num_heads, X.shape[1], X.shape[2])
# X的形状变为: (batch_size, 查询的个数, num_heads, num_hiddens/num_heads)
X = X.permute(0, 2, 1, 3)
# 最终输出的形状: (batch_size, 查询的个数, num_hiddens)
return X.reshape(X.shape[0], X.shape[1], -1)```
多头注意力:MultiHeadAttention
有了上面的两个工具函数,接下来是多头注意力,总结就是将原QKV经过对应权重矩阵(也可以说是Linear层)计算,并且转换为可并行的多头的QKV,然后经过缩放点积注意力得分函数计算得到结果,最后通过输出全连接层输出。
class MultiHeadAttention(nn.Module):
"""多头注意力"""
def __init__(self, key_size, query_size, value_size, num_hiddens, num_heads, dropout, bias=False, **kwargs):
super(MultiHeadAttention, self).__init__(**kwargs)
self.num_heads = num_heads
self.attention = DotProductAttention(dropout)
self.W_q = nn.Linear(query_size, num_hiddens, bias=bias)
self.W_k = nn.Linear(key_size, num_hiddens, bias=bias)
self.W_v = nn.Linear(value_size, num_hiddens, bias=bias)
self.W_o = nn.Linear(num_hiddens, num_hiddens, bias=bias)
def forward(self, queries, keys, values, valid_lens):
# queries,keys,values的形状: (batch_size,查询或者“键-值”对的个数,num_hiddens)
# valid_lens 的形状: (batch_size,) 或 (batch_size,查询的个数)
# 经过变换后,输出的queries,keys,values 的形状: (batch_size*num_heads,查询或者“键-值”对的个数,num_hiddens/num_heads)
queries = transpose_qkv(self.W_q(queries), self.num_heads)
keys = transpose_qkv(self.W_k(keys), self.num_heads)
values = transpose_qkv(self.W_v(values), self.num_heads)
if valid_lens is not None:
# 在轴0,将第一项(标量或者矢量)复制num_heads次,
# 然后如此复制第二项,然后诸如此类。
valid_lens = torch.repeat_interleave(valid_lens, repeats=self.num_heads, dim=0)
# output的形状:(batch_size*num_heads,查询的个数,num_hiddens/num_heads)
output = self.attention(queries, keys, values, valid_lens)
# output_concat的形状:(batch_size,查询的个数,num_hiddens)
output_concat = transpose_output(output, self.num_heads)
return self.W_o(output_concat)
2.1.2 残差连接 + 层归一化 (AddNorm)
这里主要做的事是:残差连接,以及将加和后的结果进行层归一化(了解详细请看LayerNorm)
class AddNorm(nn.Module):
"""残差连接后进行层规范化"""
def __init__(self, normalized_shape, dropout, **kwargs):
super(AddNorm, self).__init__(**kwargs)
self.dropout = nn.Dropout(dropout)
self.ln = nn.LayerNorm(normalized_shape)
def forward(self, X, Y):
return self.ln(self.dropout(Y) + X)
2.1.3 基于位置的前馈网络 (PositionWiseFFN)
作用:还记得在编码器输入前还加了一个位置编码(后面会提及如何实现),那么网络如何来利用这个位置信息呢?就是使用的位置前馈网络,这里的ffn_num_input和ffn_num_outputs通常是相等的,一般都是num_hiddens大小,所以可以看出,这个层本质是想提取高级抽象的信息,即位置信息。
class PositionWiseFFN(nn.Module):
"""基于位置的前馈网络"""
def __init__(self, ffn_num_input, ffn_num_hiddens, ffn_num_outputs, **kwargs):
super(PositionWiseFFN, self).__init__(**kwargs)
self.dense1 = nn.Linear(ffn_num_input, ffn_num_hiddens)
self.relu = nn.ReLU()
self.dense2 = nn.Linear(ffn_num_hiddens, ffn_num_outputs)
def forward(self, X):
return self.dense2(self.relu(self.dense1(X)))
2.2 编码器总体结构
2.2.1 位置编码
PositionalEncoding
这里使用的位置编码具有周期性,并且随着行数增大,频率增大。而得到的周期性的作用是可以获取相对位置数据。在一般seq2seq模型中,位置信息一般来自于前一步的隐藏状态,所以不能解决长依赖问题,而此处直接加入了位置数据,就可以针对性地解决长距离依赖问题。
要注意X = X + self.P[:, :X.shape[1], :].to(X.device),加和的左右两部分形状相同,所以位置信息就是直接加和在隐状态值上的,所以上面说到的PositionWiseFFN,说是基于位置的原因也是如此。
class PositionalEncoding(nn.Module):
def __init__(self, num_hiddens, dropout, max_len=1000):
super(PositionalEncoding, self).__init__()
# 创建一个dropout层,用于在训练过程中减少过拟合
self.dropout = nn.Dropout(dropout)
# 创建一个足够长的位置编码矩阵P,其中1表示批次大小,
# max_len表示序列的最大长度,num_hiddens表示嵌入的维度
self.P = torch.zeros((1, max_len, num_hiddens))
# 创建一个长度的向量X,其中包含序列的长度信息
X = (torch.arange(max_len, dtype=torch.float32).reshape(-1, 1)
/ torch.pow(10000, torch.arange(0, num_hiddens, 2, dtype=torch.float32) / num_hiddens))
# 将X的每个维度分割成两部分,奇数维度使用正弦函数,偶数维度使用余弦函数
self.P[:, :, 0::2] = torch.sin(X)
self.P[:, :, 1::2] = torch.cos(X)
def forward(self, X):
# 将位置编码矩阵P与输入X进行拼接,然后将结果送入dropout层
X = X + self.P[:, :X.shape[1], :].to(X.device)
return self.dropout(X)
2.2.2 编码器结构
TransformerEncoder
前面铺垫了那么多,终于到了实现编码器整体结构的时候了。
- 和一般的seq2seq模型类似,先是对输入数据的预处理部分,但是这里引入了位置编码的概念,所以分为两个部分,先对数据进行词向量转换,在加上位置编码(embedding、pos_encoding)。
- 接下来和本文开头图中一样,复制num_layers个编码器子模块。
class TransformerEncoder(nn.Module):
"""Transformer编码器"""
def __init__(self, vocab_size, key_size, query_size, value_size,
num_hiddens, norm_shape, ffn_num_input, ffn_num_hiddens,
num_heads, num_layers, dropout, use_bias=False, **kwargs):
super(TransformerEncoder, self).__init__(**kwargs)
self.num_hiddens = num_hiddens
# 嵌入层
self.embedding = nn.Embedding(vocab_size, num_hiddens)
# 位置编码(加入相对位置信息)
self.pos_encoding = PositionalEncoding(num_hiddens, dropout)
# 编码器模块序列
self.blks = nn.Sequential()
for i in range(num_layers):
self.blks.add_module("block"+str(i),
EncoderBlock(key_size, query_size, value_size,
num_hiddens, norm_shape, ffn_num_input, ffn_num_hiddens,
num_heads, dropout, use_bias))
def forward(self, X, valid_lens, *args):
# 因为位置编码值在-1和1之间,
# 因此嵌入值乘以嵌入维度的平方根进行缩放,
# 然后再与位置编码相加。
X = self.pos_encoding(self.embedding(X) * math.sqrt(self.num_hiddens))
self.attention_weights = [None] * len(self.blks)
for i, blk in enumerate(self.blks):
X = blk(X, valid_lens)
self.attention_weights[i] = blk.attention.attention.attention_weights
return X
3 解码器
Transformer解码器也是由多个相同的层叠加而成的,并且层中使用了残差连接和层规范化。除了编码器中描述的两个子层之外,解码器还在这两个子层之间插入了第三个子层,称为编码器-解码器注意力(encoder-decoder attention)层。在编码器-解码器注意力中,查询来自前一个解码器层的输出,而键和值来自整个编码器的输出。在解码器自注意力中,查询、键和值都来自上一个解码器层的输出。但是,解码器中的每个位置只能考虑该位置之前的所有位置。这种掩蔽(masked)注意力保留了自回归(auto-regressive)属性,确保预测仅依赖于已生成的输出词元。
3.1 解码器子层
模块定义不再赘述,相关的层在编码器那一块已经介绍过。结构与开头结构图中一致。
- 处理key_values
- 如果状态为
None,则直接将当前输入X作为key_values。 - 如果状态不为
None,则将当前输入X与之前的输出state[2][self.i]拼接在一起,形成新的key_values。这样做是为了在解码器中累积之前的输出,以便在后续的注意力机制中使用。
- 如果状态为
- 在训练模式下,计算解码器块的Mask,即
dec_valid_lens,在计算注意力权重时,会将未来位置的权重设置为0,确保解码器在生成每个词时只能依赖于之前的输出。 - 经过两个多头注意力和残差归一化
- 第一个多头注意力是自注意力,得到的是序列内部的共现关系
- 第二个多头注意力是编码器-解码器注意力,得到的是源序列和目标序列之间的关系
- 因此它们共同构成的结构本质是:试图捕捉序列内部和序列之间的复杂关系
- 关于注意力机制,如果不明白还可以看这一篇《注意力机制》
class DecoderBlock(nn.Module):
"""解码器中第i个块"""
def __init__(self, key_size, query_size, value_size, num_hiddens,
norm_shape, ffn_num_input, ffn_num_hiddens, num_heads,
dropout, i, **kwargs):
super(DecoderBlock, self).__init__(**kwargs)
self.i = i
self.attention1 = MultiHeadAttention(key_size, query_size, value_size, num_hiddens, num_heads, dropout)
self.addnorm1 = AddNorm(norm_shape, dropout)
self.attention2 = MultiHeadAttention(key_size, query_size, value_size, num_hiddens, num_heads, dropout)
self.addnorm2 = AddNorm(norm_shape, dropout)
self.ffn = PositionWiseFFN(ffn_num_input, ffn_num_hiddens, num_hiddens)
self.addnorm3 = AddNorm(norm_shape, dropout)
def forward(self, X, state):
enc_outputs, enc_valid_lens = state[0], state[1]
# 训练阶段,输出序列的所有词元都在同一时间处理,
# 因此state[2][self.i]初始化为None。
# 预测阶段,输出序列是通过词元一个接着一个解码的,
# 因此state[2][self.i]包含着直到当前时间步第i个块解码的输出表示
if state[2][self.i] is None:
key_values = X
else:
key_values = torch.cat((state[2][self.i], X), dim=1)
state[2][self.i] = key_values
if self.training:
batch_size, num_steps, _ = X.shape
# dec_valid_lens的开头:(batch_size,num_steps),
# 其中每一行是[1,2,...,num_steps]
dec_valid_lens = torch.arange(1, num_steps + 1, device=X.device).repeat(batch_size, 1)
else:
dec_valid_lens = None
# 自注意力
X2 = self.attention1(X, key_values, key_values, dec_valid_lens)
Y = self.addnorm1(X, X2)
# 编码器-解码器注意力。
# enc_outputs的开头:(batch_size,num_steps,num_hiddens)
Y2 = self.attention2(Y, enc_outputs, enc_outputs, enc_valid_lens)
Z = self.addnorm2(Y, Y2)
return self.addnorm3(Z, self.ffn(Z)), state
3.2 解码器整体结构
类似编码器结构,不同的是,解码器有一个隐藏状态初始化的方法,这是根据编码器输出来决定的;另外解码器有两个注意力模块。
需要注意的是,在前向传播中,state隐藏状态的输入是有层之分的,一般Transformer的编码器和解码器数量是相同的,每一层对应的state。
class TransformerDecoder(nn.Module):
def __init__(self, vocab_size, key_size, query_size, value_size, num_hiddens,
norm_shape, ffn_num_input, ffn_num_hiddens,
num_heads, num_layers, dropout, **kwargs):
super(TransformerDecoder, self).__init__(**kwargs)
self.num_hiddens = num_hiddens
self.num_layers = num_layers
# 嵌入层
self.embedding = nn.Embedding(vocab_size, num_hiddens)
# 位置编码(加入相对位置信息)
self.pos_encoding = PositionalEncoding(num_hiddens, dropout)
# 解码器模块序列
self.blks = nn.Sequential()
for i in range(num_layers):
self.blks.add_module("block"+str(i),
DecoderBlock(key_size, query_size, value_size, num_hiddens,
norm_shape, ffn_num_input, ffn_num_hiddens,
num_heads, dropout, i))
self.dense = nn.Linear(num_hiddens, vocab_size)
def init_state(self, enc_outputs, enc_valid_lens, *args):
return [enc_outputs, enc_valid_lens, [None] * self.num_layers]
def forward(self, X, state):
X = self.pos_encoding(self.embedding(X) * math.sqrt(self.num_hiddens))
self._attention_weights = [[None] * len(self.blks) for _ in range (2)]
for i, blk in enumerate(self.blks):
X, state = blk(X, state)
# 解码器自注意力权重
self._attention_weights[0][i] = blk.attention1.attention.attention_weights
# “编码器-解码器”自注意力权重
self._attention_weights[1][i] = blk.attention2.attention.attention_weights
return self.dense(X), state
@property
def attention_weights(self):
return self._attention_weights
4 组合
class Transformer(nn.Module):
def __init__(self, encoder, decoder, **kwargs):
super(Transformer, self).__init__(**kwargs)
self.encoder = encoder
self.decoder = decoder
def forward(self, enc_X, dec_X, *args):
enc_outputs = self.encoder(enc_X, *args)
dec_state = self.decoder.init_state(enc_outputs, *args)
return self.decoder(dec_X, dec_state)
5 训练
这里提及的权重初始化xavier_init_weights、带遮蔽的softmax交叉熵损失函数MaskedSoftmaxCELoss 以及梯度裁剪grad_clipping在【深度学习】Seq2Seq中都介绍了,包括训练方法其实是一样的。
net.apply(xavier_init_weights)
net.to(device)
net.train()
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
loss = MaskedSoftmaxCELoss()
for epoch in range(num_epochs):
for batch in train_iter:
optimizer.zero_grad()
X, X_valid_len, Y, Y_valid_len = [x.to(device) for x in batch]
bos = torch.tensor([tgt_vocab['<bos>']] * Y.shape[0], device=device).reshape(-1, 1)
dec_input = torch.concat([bos, Y[:, :-1]], 1)
Y_hat, _ = net(X, dec_input, X_valid_len)
l = loss(Y_hat, Y, Y_valid_len)
l.sum().backward()
grad_clipping(net, 1)
optimizer.step()
if (epoch + 1) % 10 == 0:
print( f'epoch {epoch + 1}, ' f'loss {float(l.sum()):.3f}')