Seq2Seq基础结构可以看上一篇文章【深度学习】Seq2Seq

1 总体结构

Bahdanau注意力也成为加性注意力,其实是一种注意力评分函数

seq2seq-attention-details.svg

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结构,是可学习的

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_qW_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结构完全一致