PyTorch中masked_fill函数在该张量操作代码场景下的作用是什么?
代码作用说明
前提张量信息
- 张量
k1:形状为(3, 1, 1, 9) - 张量
p1:形状为(3, 7, 9, 9)
代码执行逻辑
这段是PyTorch框架下的张量操作代码,完整代码为:
p1 = p1.masked_fill(k1 == 0, float("-1e30"))
具体作用拆解如下:
- 首先生成布尔掩码:运行
k1 == 0会得到和k1形状完全一致的布尔张量,k1里值为0的位置对应掩码值为True,其余位置为False。 - 维度自动广播:PyTorch的广播机制会把形状为
(3, 1, 1, 9)的布尔掩码自动扩展到和p1匹配的(3, 7, 9, 9)形状:掩码中两个长度为1的维度会分别复制扩展到对应p1的第2维(长度7)、第3维(长度9),确保掩码和p1的每个位置一一对应。 - 按掩码填充数值:
masked_fill方法会把p1中所有对应掩码值为True的位置,统一替换为指定的极小值-1e30。
典型使用场景
这个操作通常用在注意力机制的掩码逻辑中:后续对p1做Softmax运算时,被替换为-1e30的位置输出会无限趋近于0,相当于直接屏蔽这些位置的注意力权重,避免模型关注到无效、填充或者不需要计算的位置。
内容的提问来源于stack exchange,提问作者Farhang Amaji
相关产品推荐
相关产品推荐

