如何用NumPy einsum实现含嵌套索引的双重循环计算?
用NumPy einsum替代双重循环的实现方案
原循环代码
for j in range(100): p[j] = 0 for i in range(100): if i!=j: p[j] += S[i,j]*B[T[i,j], i]
数组维度说明:
p.shape = (1,100)S.shape = (100,100)B.shape = (N,100)T.shape = (100,100),且N大于T中的最大值
可以用einsum实现,步骤如下:
处理i≠j的条件:构造单位矩阵掩码,将S中i=j的位置置为0,求和时自动忽略这些无效项:
import numpy as np mask = np.eye(100, dtype=bool) S_masked = S.copy() S_masked[mask] = 0提取B中对应T的元素:利用NumPy高级索引,直接取出
B[T[i,j], i]对应的所有元素,生成和S同维度的数组:# np.arange(100)对应每个i的列索引,T的每个元素是B的行索引 B_selected = B[T, np.arange(100)]用einsum计算求和:
p = np.einsum('ij,ij->j', S_masked, B_selected).reshape(1, 100)表达式
'ij,ij->j'的含义:对两个形状为(100,100)的数组对应位置相乘,然后沿i轴(第一个维度)累加,最终得到长度为100的结果,再reshape为(1,100)匹配p的维度。
其他优化方法(无需einsum)
如果不想使用einsum,直接用向量运算也能达到同样的优化效果,代码更直观:
p = (S_masked * B_selected).sum(axis=0).reshape(1, 100)
这种方式和einsum逻辑一致,都是利用NumPy的向量化操作替代显式循环,效率远高于原双重循环。
内容的提问来源于stack exchange,提问作者FatPanda01
相关产品推荐
相关产品推荐

