如何在PyTorch中实现a_{ij}+b_{kj}=c_{ik}的张量收缩运算?
在PyTorch中实现$a_{ij} + b_{kj} \rightarrow c_{ik}$的张量收缩
首先明确:这里的运算本质是对j维度求和(直接元素级相加维度不匹配),即最终的$c_{ik} = \sum_j (a_{ij} + b_{kj})$。利用PyTorch的广播机制或张量求和可以轻松实现,以下是两种高效方案:
方案1:利用加法分配律优化(推荐)
由于$\sum_j (a_{ij} + b_{kj}) = \sum_j a_{ij} + \sum_j b_{kj}$,我们可以先分别对两个张量的j维度求和,再通过广播实现行向量+列向量的运算,得到目标张量:
import torch # 定义示例张量 I, J, K = 3, 4, 5 a = torch.randn(I, J) # shape: (I, J) b = torch.randn(K, J) # shape: (K, J) # 分别对j维度求和 sum_a = a.sum(dim=1) # shape: (I,) → 每个i对应的a_ij之和 sum_b = b.sum(dim=1) # shape: (K,) → 每个k对应的b_kj之和 # 广播相加得到c_ik c = sum_a.unsqueeze(1) + sum_b.unsqueeze(0) # shape: (I, K)
这种方法计算量更小,尤其当J维度较大时性能更优。
方案2:广播后直接求和(直观易懂)
如果想直观对应“$a_{ij} + b_{kj}$”的形式,可以先通过维度扩展让两个张量广播到相同的三维形状,再相加后对j维度求和:
import torch I, J, K = 3, 4, 5 a = torch.randn(I, J) b = torch.randn(K, J) # 扩展维度:a→(I,1,J),b→(1,K,J),实现广播兼容 a_expanded = a.unsqueeze(1) b_expanded = b.unsqueeze(0) # 相加后对j维度求和 c = (a_expanded + b_expanded).sum(dim=2) # shape: (I, K)
两种方案得到的结果完全一致,可根据需求选择。
内容的提问来源于stack exchange,提问作者olafcx
相关产品推荐
相关产品推荐

