Julia中广播数组求和:无临时矩阵分配的高效实现方案
高效求和方案(无循环、避免大临时矩阵)
原代码的核心问题是逐元素相乘会生成尺寸为(m1,m1,n1,n1,n2,n3)的巨型临时矩阵,完全可以通过数学变形和框架内置的优化操作规避,以下是两种可行方案:
方案一:使用爱因斯坦求和(einsum)
爱因斯坦求和能直接描述计算逻辑,且多数数值计算框架(NumPy、PyTorch、TensorFlow等)会自动优化内存,不会生成中间大数组。
步骤如下:
- 先将
big_mat重塑为更清晰的维度:BM = reshape(big_mat, m1, n1, n2, n3) - 用
einsum直接完成跨维度求和:temp = einsum('ipkl,jqkl->ijpq', BM, BM) D = temp * C
这里einsum表达式'ipkl,jqkl->ijpq'表示:对BM的第3、4维度(对应n2、n3)求和,将BM[i,p,k,l]与BM[j,q,k,l]的乘积累积到结果的[i,j,p,q]位置,最后和C逐元素相乘得到最终结果。
方案二:利用矩阵乘法优化
矩阵乘法通常基于BLAS库实现,效率极高,同样可避免大临时矩阵:
- 将
big_mat重塑为二维矩阵,合并n2、n3维度:BM_flat = reshape(big_mat, m1 * n1, n2 * n3) - 计算矩阵乘积,再调整维度匹配
C的形状:mat_prod = BM_flat @ BM_flat.T # 形状为(m1*n1, m1*n1) mat_prod_reshaped = reshape(mat_prod, m1, n1, m1, n1) temp = transpose(mat_prod_reshaped, (0, 2, 1, 3)) # 调整为(m1, m1, n1, n1) D = temp * C
这种方法通过矩阵乘法一次性完成跨n2、n3维度的求和,后续仅需调整维度即可与C相乘,内存占用远低于原实现。
两种方案都无需循环,内存效率和计算速度均远超原代码,可根据使用的框架选择适配方式。
内容的提问来源于stack exchange,提问作者Rain
相关产品推荐
相关产品推荐

