在PyTorch中如何计算注意力分数与编码器输出的加权平均值
上下文向量计算方案(PyTorch实现)
计算逻辑前提
注意力加权平均的本质是对每个batch样本,将19个时间步对应的编码器隐状态分别乘对应时间步的注意力分数,再对时间步维度求和,最终得到每个样本维度为256的上下文向量。
注意:输入的注意力分数需要提前在seq_len(也就是第1维,索引为1)维度做过softmax归一化,保证每个样本的19个注意力分数加和为1,这样加权和才符合加权平均的定义
常用实现方法
方法1:广播乘法+求和(直观易理解)
你的注意力分数维度已经是[64,19,1],最后一维的单值刚好可以和编码器隐状态的256维做广播匹配,直接相乘后在时间步维度求和即可:
import torch # 模拟你的输入张量 attn_scores = torch.randn(64, 19, 1) encoder_out = torch.randn(64, 19, 256) # 广播乘法输出维度:[64, 19, 256] weighted_hidden = attn_scores * encoder_out # 时间步维度求和,输出上下文向量维度:[64, 256] context_vec = weighted_hidden.sum(dim=1)
方法2:批量矩阵乘法(运算效率更高)
使用PyTorch内置的批量矩阵乘法接口torch.bmm,仅需要调整注意力分数的维度顺序即可:
# 调整注意力分数维度为 [batch_size, 1, seq_len] → [64, 1, 19] attn_scores_trans = attn_scores.transpose(1, 2) # 批量矩阵乘法运算:[64,1,19] × [64,19,256] = [64,1,256], squeeze后得到[64,256] context_vec = torch.bmm(attn_scores_trans, encoder_out).squeeze(1)
高维带batch运算逻辑梳理
理清这类运算只需要记住三个核心规则:
- 最左侧的batch维度永远保持不变,运算只会在后面的样本内维度执行,每个样本的计算完全独立,不需要手动写循环遍历batch
- 广播乘法、批量矩阵乘法都会自动对齐batch维度,无需额外处理
- 运算前可以先明确每个维度的物理含义,推导输出维度后可以先构造小尺寸的测试张量验证形状,再代入正式数据运行
内容的提问来源于stack exchange,提问作者Cole
相关产品推荐
相关产品推荐

