基于MemoryNetwork:计算Softmax注意力时如何忽略零填充列?
解决Memory Networks中Softmax忽略零填充位置的PyTorch方案
我懂你遇到的问题——直接对带零填充的注意力logits做Softmax,会因为填充位的存在(哪怕数值是0),导致概率全被最大的那个有效logits吸走,根本没法聚焦到目标上下文句子上。下面给你几个实用的PyTorch实现方案,核心思路都是把填充位置的logits替换成负无穷,让它们在Softmax计算中彻底失效:
方法一:手动指定有效位置生成掩码
如果你的有效位置是固定规则(比如每个样本都是前4个+最后一个位置有效),可以直接生成掩码标记有效区域:
import torch import torch.nn.functional as F # 模拟你的bmm输出,这里假设是batch_size=1的情况,形状[1, 15] logits = torch.tensor([[109.8601, 77.6376, 68.3927, 199.1673, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 348.0155]]).cuda() # 1. 标记有效位置:前4个索引(0-3)+最后一个索引 seq_len = logits.size(1) valid_indices = [0, 1, 2, 3, seq_len - 1] mask = torch.zeros_like(logits) mask[:, valid_indices] = 1 # 有效位置设为1,填充位为0 # 2. 把填充位的logits替换成负无穷,Softmax时这些位置的概率会趋近于0 logits_masked = logits.masked_fill(mask == 0, -float('inf')) # 3. 计算Softmax,此时只有有效位置参与归一化 attention = F.softmax(logits_masked, dim=1) print(attention)
运行后,填充位的注意力会是0,有效位置的概率会根据各自的logits值正常归一化,不会出现全集中在最大logits的情况。
方法二:基于原始填充标记生成掩码
如果你的填充位是构建batch时就确定的(比如提前记录了原始上下文的padding mask),直接用这个mask更可靠(能避免误判真实logits为0的特殊情况):
# 假设你已经有一个和logits同形状的padding_mask,有效位置为1,填充位为0 padding_mask = torch.tensor([[1,1,1,1,0,0,0,0,0,0,0,0,0,0,1]]).cuda() # 替换填充位为负无穷 logits_masked = logits.masked_fill(padding_mask == 0, -float('inf')) attention = F.softmax(logits_masked, dim=1)
这种方法更通用,尤其是当不同样本的有效上下文长度不一致时,提前记录的padding_mask能准确标记每个样本的有效区域。
原理说明
Softmax的计算逻辑是exp(x_i) / sum(exp(x_j)),当x_i被设为负无穷时,exp(x_i)会趋近于0,完全不会参与分母的求和,这样有效位置的概率就只会基于彼此的logits值进行归一化,完美实现忽略填充位的需求。
内容的提问来源于stack exchange,提问作者jef
相关产品推荐
相关产品推荐

