1 整体结构

可分解注意力模型:Decomposable Attention Model
2 多层感知机 (MLP)
def mlp(num_inputs, num_hiddens, flatten):
net = []
net.append(nn.Dropout(0.2))
net.append(nn.Linear(num_inputs, num_hiddens))
net.append(nn.ReLU())
if flatten:
net.append(nn.Flatten(start_dim=1))
net.append(nn.Dropout(0.2))
net.append(nn.Linear(num_hiddens, num_hiddens))
net.append(nn.ReLU())
if flatten:
net.append(nn.Flatten(start_dim=1))
return nn.Sequential(*net)
3 注意力 (Attend)
class Attend(nn.Module):
def __init__(self, num_inputs, num_hiddens, **kwargs):
super(Attend, self).__init__(**kwargs)
self.f = mlp(num_inputs, num_hiddens, flatten=False)
def forward(self, A, B):
# A/B的形状:(批量大小,序列A/B的词元数,embed_size)
# f_A/f_B的形状:(批量大小,序列A/B的词元数,num_hiddens)
f_A = self.f(A)
f_B = self.f(B)
# e的形状:(批量大小,序列A的词元数,序列B的词元数)
e = torch.bmm(f_A, f_B.permute(0, 2, 1))
# beta的形状:(批量大小,序列A的词元数,embed_size),
# 意味着序列B被软对齐到序列A的每个词元(beta的第1个维度)
beta = torch.bmm(F.softmax(e, dim=-1), B)
# beta的形状:(批量大小,序列B的词元数,embed_size),
# 意味着序列A被软对齐到序列B的每个词元(alpha的第1个维度)
alpha = torch.bmm(F.softmax(e.permute(0, 2, 1), dim=-1), A)
return beta, alpha
这里的注意力有点类似于后来的Transformer的解码器中的“编码器解码器注意力”(可以想象AB分别为编码器输出和解码器输入),并且这里分别对AB都“对称”计算一遍意义就是:
- 这种设计允许模型在处理句子对时,同时考虑两个方向的信息流动,即不仅考虑序列B对序列A的影响(通过
beta),还考虑序列A对序列B的影响(通过alpha)。 - 这种双向的注意力机制有助于模型更全面地捕捉句子对之间的交互信息,从而更好地理解它们之间的关系。
4 比较 (Compare)
class Compare(nn.Module):
def __init__(self, num_inputs, num_hiddens, **kwargs):
super(Compare, self).__init__(**kwargs)
self.g = mlp(num_inputs, num_hiddens, flatten=False)
def forward(self, A, B, beta, alpha):
V_A = self.g(torch.cat([A, beta], dim=2))
V_B = self.g(torch.cat([B, alpha], dim=2))
return V_A, V_B
我们将来自一个序列的词元的连结(运算符[\cdot, \cdot])和来自另一序列的对齐的词元送入函数g(一个多层感知机):
\mathbf{v}_{A,i} = g([\mathbf{a}_i, \boldsymbol{\beta}_i]), i = 1, \ldots, m
\mathbf{v}_{B,j} = g([\mathbf{b}_j, \boldsymbol{\alpha}_j]), j = 1, \ldots, n
- \mathbf{a}_i:序列A(例如,前提)中的第i个词元的嵌入表示。
- \boldsymbol{\beta}_i:序列B(例如,假设)中被软对齐到序列A第i个词元的表示。
- \mathbf{b}_j:序列B中的第j个词元的嵌入表示。
- \boldsymbol{\alpha}_j:序列A中被软对齐到序列B第j个词元的表示。
- g:多层感知机,用于对连结后的表示进行非线性变换。
- \mathbf{v}_{A,i}:序列A第i个词元与序列B中与之软对齐的词元进行比较后的表示。
- \mathbf{v}_{B,j}:序列B第j个词元与序列A中与之软对齐的词元进行比较后的表示。
比较步骤通过将两个序列中的词元与其软对齐的表示进行比较,提取出有助于判断句子对之间关系的特征。这个过程利用了多层感知机进行非线性变换,从而更好地捕捉词元之间的复杂交互关系。在自然语言推理任务中,这些比较后的表示将被用于最终的聚合和关系判断。
5 聚合 (Aggregate)
class Aggregate(nn.Module):
def __init__(self, num_inputs, num_hiddens, num_outputs, **kwargs):
super(Aggregate, self).__init__(**kwargs)
self.h = mlp(num_inputs, num_hiddens, flatten=True)
self.linear = nn.Linear(num_hiddens, num_outputs)
def forward(self, V_A, V_B):
# 对两组比较向量分别求和
V_A = V_A.sum(dim=1)
V_B = V_B.sum(dim=1)
# 将两个求和结果的连结送到多层感知机中
Y_hat = self.linear(self.h(torch.cat([V_A, V_B], dim=1)))
return Y_hat
6 组合
class DecomposableAttention(nn.Module):
def __init__(self, vocab, embed_size, num_hiddens,
num_inputs_attend=100, num_inputs_compare=200, num_inputs_agg=400, **kwargs):
super(DecomposableAttention, self).__init__(**kwargs)
# 嵌入
self.embedding = nn.Embedding(len(vocab), embed_size)
# 注意力
self.attend = Attend(num_inputs_attend, num_hiddens)
# 比较
self.compare = Compare(num_inputs_compare, num_hiddens)
# 有3种可能的输出:蕴涵、矛盾和中性
self.aggregate = Aggregate(num_inputs_agg, num_hiddens, num_outputs=3)
def forward(self, X):
premises, hypotheses = X
A = self.embedding(premises)
B = self.embedding(hypotheses)
beta, alpha = self.attend(A, B)
V_A, V_B = self.compare(A, B, beta, alpha)
Y_hat = self.aggregate(V_A, V_B)
return Y_hat