关于self-attention模块中torch.einsum语句的解读及替代实现咨询
解读Self-Attention中的torch.einsum语句及替代实现
维度含义说明
先明确张量维度在Self-Attention场景下的典型意义,方便理解:
b: 批量大小(batch size)h: 注意力头数量(number of attention heads)l: 查询(query)序列长度r: 键/值(key/value)序列长度(Self-Attention中通常与l相等)d: 单个注意力头的维度(dimension per head)
第一个einsum语句:'bhld,lrd->bhlr'
计算逻辑
输入两个张量:
- 4D张量
bhld:形状为(b, h, l, d),对应每个batch、每个注意力头下的查询序列特征 - 3D张量
lrd:形状为(l, r, d),对应每个查询位置、键位置下的键特征
einsum的作用是对维度d做求和,最终输出形状为(b, h, l, r)的张量。具体计算式为:
对于每个b, h, l, r,计算 sum_d(bhld[b,h,l,d] * lrd[l,r,d])
这个结果对应注意力分数计算的核心逻辑——查询与键的点积操作。
替代实现(无需einsum)
利用PyTorch的torch.matmul实现,核心是调整第二个张量的维度顺序,让可相乘的维度对齐:
# 原einsum代码 x = torch.randn(b, h, l, d) # bhld y = torch.randn(l, r, d) # lrd result_einsum = torch.einsum('bhld,lrd->bhlr', x, y) # 替代实现 y_transposed = y.transpose(1, 2) # 将y从(l,r,d)转为(l,d,r) result_matmul = torch.matmul(x, y_transposed) # 输出形状(b,h,l,r),与einsum结果一致
第二个einsum语句:'bhrd,lrd->bhlr'
计算逻辑
输入两个张量:
- 4D张量
bhrd:形状为(b, h, r, d),对应每个batch、每个注意力头下的值序列特征 - 3D张量
lrd:形状为(l, r, d),对应每个查询位置、键位置下的键特征
einsum对维度d做求和,最终输出形状为(b, h, l, r)的张量。具体计算式为:
对于每个b, h, l, r,计算 sum_d(bhrd[b,h,r,d] * lrd[l,r,d])
替代实现(无需einsum)
需要先调整第二个张量的维度顺序,做完矩阵乘法后再调整输出维度:
# 原einsum代码 x = torch.randn(b, h, r, d) # bhrd y = torch.randn(l, r, d) # lrd result_einsum = torch.einsum('bhrd,lrd->bhlr', x, y) # 替代实现 y_permuted = y.permute(1, 2, 0) # 将y从(l,r,d)转为(r,d,l) result_matmul = torch.matmul(x, y_permuted) # 输出形状(b,h,r,l) result_final = result_matmul.transpose(2, 3) # 转置后得到(b,h,l,r),与einsum结果一致
内容的提问来源于stack exchange,提问作者clueless
相关产品推荐
相关产品推荐

