Seq2Seq基础结构可以看上一篇文章【深度学习】Seq2Seq
1 总体结构
Bahdanau注意力也成为加性注意力,其实是一种注意力评分函数
2 编码器结构
- 在编码器部分,和基础的seq2seq结构完全一致
3 解码器结构
- 注意:
- 与基础Seq2Seq结构的主要不同主要是对于解码器结构的部分,对于编码器输出的利用方式的不同。
- 基础的Seq2Seq的解码器使用的初始context(也就是编码器输出的隐藏状态),而带有注意力机制的最大的不同就在于此,它每一步的context都根据最后的隐藏状态计算得出,而这里说的计算就是指的是加性注意力计算函数(3.1有相关说明)。
class Seq2SeqAttentionDecoder(nn.Module):
def __init__(self, vocab_size, embed_size, num_hiddens, num_layers, dropout=0, **kwargs):
super(Seq2SeqAttentionDecoder, self).__init__(**kwargs)
self.attention = AdditiveAttention(num_hiddens, num_hiddens, num_hiddens, dropout)
self.embedding = nn.Embedding(vocab_size, embed_size)
# 这里rnn输入形状是 embed_size + num_hiddens 原因是组合了注意力上下文(num_hiddens)和编码器输入(embed_size)
self.rnn = nn.GRU(embed_size + num_hiddens, num_hiddens, num_layers, dropout=dropout)
self.dense = nn.Linear(num_hiddens, vocab_size)
def init_state(self, enc_outputs, enc_valid_lens, *args):
# outputs的形状为(batch_size,num_steps,num_hiddens)
# hidden_state的形状为(num_layers,batch_size,num_hiddens)
outputs, hidden_state = enc_outputs
return (outputs.permute(1, 0, 2), hidden_state, enc_valid_lens)
def forward(self, X, state):
# enc_outputs的形状为(batch_size,num_steps,num_hiddens)
# hidden_state的形状为(num_layers,batch_size,num_hiddens)
enc_outputs, hidden_state, enc_valid_lens = state
# 输出X的形状为(num_steps,batch_size,embed_size)
X = self.embedding(X).permute(1, 0, 2)
outputs, self._attention_weights = [], []
for x in X:
# query的形状为(batch_size,1,num_hiddens)
query = torch.unsqueeze(hidden_state[-1], dim=1) # 在非注意力解码器中,所谓context就是query
# context的形状为(batch_size,1,num_hiddens)
context = self.attention(query, enc_outputs, enc_outputs, enc_valid_lens) # queries, keys, values, valid_lens
# 在特征维度上连结
x = torch.cat((context, torch.unsqueeze(x, dim=1)), dim=-1)
# 将x变形为(1,batch_size,embed_size+num_hiddens)
out, hidden_state = self.rnn(x.permute(1, 0, 2), hidden_state)
outputs.append(out)
self._attention_weights.append(self.attention.attention_weights)
# 全连接层变换后,outputs的形状为
# (num_steps,batch_size,vocab_size)
outputs = self.dense(torch.cat(outputs, dim=0))
return outputs.permute(1, 0, 2), [enc_outputs, hidden_state, enc_valid_lens]
@property
def attention_weights(self):
return self._attention_weights
3.1 加性注意力评分函数
实质上它是一个MLP结构,是可学习的
- 注意其中的sequence_mask在【深度学习】Seq2Seq的损失函数中写过
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 = sequence_mask(X.reshape(-1, shape[-1]), valid_lens, value=-1e6)
return nn.functional.softmax(X.reshape(shape), dim=-1)
- 以下就是加性注意力评分函数,可以理解为一个简单MLP结构
a(\mathbf q, \mathbf k) = \mathbf w_v^\top \text{tanh}(\mathbf W_q\mathbf q + \mathbf W_k \mathbf k) \in \mathbb{R},
- 主要组件:
W_k:一个线性层,用于将键的维度映射到隐藏层大小。W_q:一个线性层,用于将查询的维度映射到隐藏层大小。w_v:一个线性层,用于将隐藏层的输出映射到一个标量,用于计算注意力得分。dropout:一个dropout层,用于在计算注意力权重之前对特征进行dropout处理。
- 前向传播方法:
- 输入
queries是一个形状为(batch_size, num_queries, query_size)的张量,代表一批查询序列。 - 输入
keys是一个形状为(batch_size, 1, num_keys, key_size)的张量,代表一批键-值对序列。 - 输入
values是一个形状为(batch_size, num_keys, value_size)的张量,代表一批值序列。 - 输入
valid_lens是一个形状为(batch_size,)的张量,表示每个批次中键-值对序列的有效长度。 - 首先,通过
W_q和W_k将查询和键分别映射到隐藏层大小。 - 在特征维度上,将查询和键进行广播加和,然后通过tanh函数进行非线性变换。
- 将变换后的特征通过
w_v映射到标量,得到注意力得分。 - 使用
masked_softmax函数对注意力得分进行处理,得到注意力权重,同时考虑了键-值对的有效长度。 - 最后,通过注意力权重和值进行矩阵乘法,得到加权后的值序列。
- 输入
- 输出:该方法返回一个形状为
(batch_size, num_queries, value_size)的张量,代表加权后的值序列。
class AdditiveAttention(nn.Module):
"""加性注意力"""
def __init__(self, key_size, query_size, num_hiddens, dropout, **kwargs):
super(AdditiveAttention, self).__init__(**kwargs)
self.W_k = nn.Linear(key_size, num_hiddens, bias=False)
self.W_q = nn.Linear(query_size, num_hiddens, bias=False)
self.w_v = nn.Linear(num_hiddens, 1, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, queries, keys, values, valid_lens):
queries, keys = self.W_q(queries), self.W_k(keys)
# 在维度扩展后,
# queries的形状:(batch_size,查询的个数,1,num_hidden)
# key的形状:(batch_size,1,“键-值”对的个数,num_hiddens)
# 使用广播方式进行求和
features = queries.unsqueeze(2) + keys.unsqueeze(1)
features = torch.tanh(features)
# self.w_v仅有一个输出,因此从形状中移除最后那个维度。
# scores的形状:(batch_size,查询的个数,“键-值”对的个数)
scores = self.w_v(features).squeeze(-1)
self.attention_weights = masked_softmax(scores, valid_lens)
# values的形状:(batch_size,“键-值”对的个数,值的维度)
return torch.bmm(self.dropout(self.attention_weights), values)
4 组合
- 和基础的seq2seq结构完全一致
5 损失函数
- 和基础的seq2seq结构完全一致
6 训练(只展示关键部分)
- 和基础的seq2seq结构完全一致 (加入的注意力机制在模型内部,训练的输入和输出结构不影响)
7 评估
- 和基础的seq2seq结构完全一致