大家好,好久不见,今天想给大家分享一些关于Mask的一些理解和小技巧。
来源丨433的3号同学
作者丨贾彦
如果你要对transformer进行改动,就必然要去重新构建你的mask,那么本文会对transformer中的mask机制进行详细解读,教你如何根据输入的改变来修改你的mask。
重新来看Self-attention


figure 1
figure 1中,Q矩阵的形状为(max_langth=5,dim=4),对于每个时刻的Querie,我们用q来表示。那么图中代表5个时刻的维度为4的Querie,构成Q矩阵。同样的道理适用于K,只不过K需要经过转置,那么转置后的形状为(dim=4,max_langth=5)
根据矩阵乘法规则,有如下运算:

figure 2

figure 3

figure 4


figure 5

figure 6

figure 7
最终的result可以被认为是V根据score矩阵进行加权相加后的结果,比如:r1可以被认为是v1~v5根据s1的权重相乘相加的结果。r4可以被认为是v1~v5根据s4的权重相乘相加的结果。
Mask工作原理


figure 8
上图为self-attention的一种常见Mask,其中1代表参与计算,0代表不参与计算。当我们需要self-attention的输入每次只计算历史和自身的信息,不可以看到未来信息,我们可以设计Mask为下三角矩阵。
那么,为什么Mask会是长这样呢?


figure 9

figure 10


figure 11


figure 12


figure 13
因为vector只是额外信息,产生的attention输出并不参与到后续的计算,我们可以设置Mask前五行的值为任意。而对于Mask后五行真正的输出,我们根据之前描述的Mask行列的含义,在原来的Mask(绿色)增加了部分Mask(黑色)。拿其中几行举例:当我们q1的输出需要k1和a1,我们将Mask第六行第1列(代表a1)、第六行第6列(代表k1)设为1,第六行其他列都设为0。当我们q3的输出需要k1、k2、k3、a3,我们将Mask第八行第3列(代表a3)、第八行第6-8列(代表k1-k3)设为1,第八行其他列的都设为0。


figure 14
这样,我们就完成了整个计算。
利用相似的设计规则,理论上我们可以构造任何序列的self-attention的特殊计算。大家只用搞清楚score矩阵和Mask矩阵行列的具体含义,设计变化灵活的self-attention也就是顺理成章的事了。祝大家学习快乐。
