nn.TransformerEncoderLayer中的src_mask,src_key_padding_mask解析
创始人
2024-06-03 08:08:33
0

注意,不同版本的pytorch,对nn.TransformerEncdoerLayer部分代码差别很大,比如1.8.0版本中没有batch_first参数,而1.10.1版本中就增加了这个参数,笔者这里使用pytorch1.10.1版本实验。

attention mask

要搞清楚src_mask和src_key_padding_mask的区别,关键在于搞清楚在self-attention中attention mask的作用是啥。
attetnionscore=softmax(QKTdk)Vattetnion \ score = softmax({QK^{T} \over \sqrt d_{k} })V attetnion score=softmax(d​k​QKT​)V
上式中,并没有体现出pad的token,认为所有token都是有用的,但是实际写代码时使用batch进行训练,所以要将所有token序列pad到相同的长度。
attention mask的作用就是,在计算注意力分数的时候,告诉模型,哪些token是pad的,不应该分配注意力分数。

针对一条长度为LLL的token序列,其attention mask的矩阵应该是L∗LL*LL∗L,下图是一个attention mask,蓝色的表示不是pad的token,灰色的表示pad的token。

在这里插入图片描述
但是针对attention mask中蓝色位置和灰色位置中的值,目前有两种做法:

  • 在huggingface的transformers中实现是,将蓝色位置填1 ,灰色位置填0,也就是1表示真实序列,不需要被mask,而0表示pad序列,需要被mask。但是为了用户操作,huggingface并没有要求用户输入一个B∗L∗LB*L*LB∗L∗L的mask矩阵,而是输入B∗LB*LB∗L的矩阵即可,然后在forward函数中使用get_extended_attention_mask方法将其扩展为B∗L∗LB*L*LB∗L∗L的mask矩阵。
  • 在pytorch的transformers中的实现是,蓝色的位置填0,灰色的位置填float(“-inf”),但是在实现时,又分为了src_mask和src_key_padding_mask,而最终的attention mask矩阵,是通过这个两个矩阵得到的。
    其中:

src_mask: 必须是2D或者3D的矩阵,形状为[L,S][L,S][L,S]或者[B∗num_heads,L,S][B*num\_heads, L, S][B∗num_heads,L,S],LLL是目标序列长度,SSS是源序列长度(只有涉及到机器翻译这种encoder-decoder框架目标序列和源序列才有意义,如果只是用transformer encoder做编码,则L=SL=SL=S),BBB是batch size,numheadnum\ headnum head表示头数。另外src_mask的取值有三种,

  1. 可以是binary mask,True的位置表示需要被mask,
  2. 可以是byte mask,非零的位置表示需要被mask,
  3. 可以float mask,这时float(“-inf”)的位置需要被mask。

src_key_padding_mask:是一个2D的矩阵,形状为[B,S][B, S][B,S],取值有两种,

  1. 可以是binary mask,True的位置表示key矩阵需要被mask,
  2. 可以是byte mask,非零的位置表示key矩阵需要被mask,

这里的key矩阵应该也是为了涵盖encoder-decoder这样的情况,对于只用transformer encoder的情况,src_key_padding_mask则更像是huggingface 中的attention mask。

其实在pytorch官方代码中,是通过src_mask和src_key_padding_mask二者综合得到最终的attention_mask。对于绝大多数情况,我们只需要使用src_key_padding_mask即可。

相关内容

热门资讯

【NI Multisim 14...   目录 序言 一、工具栏 🍊1.“标准”工具栏 🍊 2.视图工具...
银河麒麟V10SP1高级服务器... 银河麒麟高级服务器操作系统简介: 银河麒麟高级服务器操作系统V10是针对企业级关键业务...
不能访问光猫的的管理页面 光猫是现代家庭宽带网络的重要组成部分,它可以提供高速稳定的网络连接。但是,有时候我们会遇到不能访问光...
AWSECS:访问外部网络时出... 如果您在AWS ECS中部署了应用程序,并且该应用程序需要访问外部网络,但是无法正常访问,可能是因为...
Android|无法访问或保存... 这个问题可能是由于权限设置不正确导致的。您需要在应用程序清单文件中添加以下代码来请求适当的权限:此外...
北信源内网安全管理卸载 北信源内网安全管理是一款网络安全管理软件,主要用于保护内网安全。在日常使用过程中,卸载该软件是一种常...
AWSElasticBeans... 在Dockerfile中手动配置nginx反向代理。例如,在Dockerfile中添加以下代码:FR...
AsusVivobook无法开... 首先,我们可以尝试重置BIOS(Basic Input/Output System)来解决这个问题。...
ASM贪吃蛇游戏-解决错误的问... 要解决ASM贪吃蛇游戏中的错误问题,你可以按照以下步骤进行:首先,确定错误的具体表现和问题所在。在贪...
月入8000+的steam搬砖... 大家好,我是阿阳 今天要给大家介绍的是 steam 游戏搬砖项目,目前...