如何利用NumPy的broadcasting机制替代循环实现Kronecker product相关计算以提升效率?
我来帮你搞定这个问题!你想要用NumPy的广播完全替代循环来提速,思路非常对——咱们先拆解原来循环里的计算逻辑,再用广播重构,就能彻底去掉循环啦。
首先看nu的计算:
原来循环里每次对单个g,你用(qu_g[:, None]*f_g[None, :]).reshape(P*H,1)实现了等价于np.kron(qu_g, f_g)的效果,本质是把(P,1)的qu_g和(H,1)的f_g做广播乘积得到(P,H),再拉平成一维。要批量处理所有G个g,咱们只需要调整维度让广播能覆盖所有g:
- 把qu从
(P, G)调整为(P, 1, G),给它加一个中间维度对应H的位置; - 把f从
(H, G)调整为(1, H, G),给它加一个开头维度对应P的位置; - 两者相乘后得到
(P, H, G)的数组,每个切片[:, :, g]就是对应单个g的(P,H)乘积; - 把前两个维度拉平成
(P*H, G),再沿着G轴求和,最后保持二维形状(和原来的nu一致)。
对应的代码是:
nu = (qu[:, None, :] * f[None, :, :]).reshape(P*H, G).sum(axis=1, keepdims=True)
接下来是de的计算:
循环里的(we_g[:, None, :, None]*f_g.dot(f_g.T)[None, :, None, :]).reshape(P*H,P*H),等价于np.kron(we_g, f_g @ f_g.T)。同样用广播批量处理:
- 先计算所有g对应的
f_g @ f_g.T,用广播可以直接得到(H, H, G)的数组:f_outer = f[:, None, :] * f[None, :, :],每个[:, :, g]就是单个g的外积; - 把we从
(P, P, G)调整为(P, 1, P, 1, G),给它插入两个维度对应H的位置; - 把f_outer调整为
(1, H, 1, H, G),给它插入两个维度对应P的位置; - 两者相乘后得到
(P, H, P, H, G)的数组,每个切片[:, :, :, :, g]就是单个g的四维乘积; - 把前两个维度拉平成
(P*H,),中间两个维度拉平成(P*H,),得到(P*H, P*H, G),再沿着G轴求和就是最终的de。
对应的代码是:
f_outer = f[:, None, :] * f[None, :, :] # shape (H, H, G) de = (we[:, None, :, None, :] * f_outer[None, :, None, :, :]).reshape(P*H, P*H, G).sum(axis=2)
为什么你之前的尝试不对?
你用了transpose(2,0,1)把维度顺序改成了(G,P,1)和(G,H,0),这样广播的时候G变成了第一个维度,虽然能相乘,但后续reshape和求和的逻辑和原来的循环累加不匹配——咱们需要保持G作为最后一个维度,让每个g的计算独立,最后再沿G轴求和,这样才能和循环里的累加效果一致。
验证结果
你可以用随机数测试一下,比如运行原来的循环代码和新的广播代码,然后用np.allclose(nu_loop, nu_broadcast)和np.allclose(de_loop, de_broadcast)来验证结果是否一致,误差在浮点精度范围内就没问题。
最后,你原来的nu = nu.sum(axis=0)和de = de.sum(axis=0)其实可以去掉,因为新的广播代码已经直接完成了求和,最后直接计算ga2就可以啦:
ga2 = np.linalg.solve(de, nu).reshape((P, H))
备注:内容来源于stack exchange,提问作者user9875321__

