如何修正numpy.einsum向量化图像旋转代码的下标错误?
修复NumPy einsum向量化旋转代码的问题
我来帮你搞定这个einsum的报错问题!首先得说,你想用向量化替代三重循环的思路非常对,毕竟循环在NumPy里确实慢得让人头疼。不过你之前的einsum写法下标没对应对,才导致了那个"subscripts too many"的错误。
错误原因分析
你原来的代码里,把旋转矩阵的每一列单独拿出来和数组的每个通道做einsum,但rot[:3,0]是个一维数组(形状(3,)),而start[:,:,0]是二维数组(形状(429,1024)),这时候你写'ij,j'就完全不匹配——前者只有1个维度,后者有2个,下标数量对应不上,自然报错了。
修正后的向量化代码
其实不用拆分通道逐个计算,用einsum可以一次性完成整个向量的矩阵乘法。这里的核心是匹配维度:你的start和norm都是(H,W,3)的三维数组,每个(H,W)位置上都是一个3维向量;旋转矩阵rot是(3,3),我们要做的是把每个3维向量和旋转矩阵相乘,得到新的3维向量。
正确的代码应该是这样:
import numpy as np # 计算旋转矩阵(和你原来的逻辑一致) s = np.sin(np.pi * 30 / 180) c = np.cos(np.pi * 30 / 180) rot = np.array([[1.0, 0.0, 0.0], [0.0, c, s], [0.0, -s, c]]) # 用einsum完成向量化旋转 start_rotated = np.einsum('hwj,jk->hwk', start, rot) norm_rotated = np.einsum('hwj,jk->hwk', norm, rot) # 如果需要原地修改原数组的话 start[:] = start_rotated norm[:] = norm_rotated
代码细节解释
'hwj,jk->hwk'这个下标规则是什么意思?hwj对应start的三个维度:高度(H)、宽度(W)、向量的3个分量(j)jk对应旋转矩阵rot的两个维度:输入向量分量(j)、输出向量分量(k)->hwk表示输出的维度是H、W、新的向量分量(k)- 简单说就是:对每个H和W位置上的j分量,和rot的j行k列相乘求和,得到该位置的k分量,完美对应你原来三重循环里的矩阵乘法逻辑!
如果你想让代码更通用(不管前面有多少维度都能处理),还可以用省略号...代替hw:
start[:] = np.einsum('...j,jk->...k', start, rot) norm[:] = np.einsum('...j,jk->...k', norm, rot)
这样写出来的代码不仅不会报错,运行速度还会比原来的三重循环快很多,完全发挥NumPy向量化的优势!
内容的提问来源于stack exchange,提问作者Ketan Chaudhari
相关产品推荐
相关产品推荐

