NumPy中将一维数组广播至指定维度的最简实现方法
嗨,我完全懂你想找更简洁写法的心情——你现在用的两种方法确实都能完美实现需求,不过我这里有几个更短的实现方式,能帮你少写点冗余代码~
先再明确下你的需求:把一维数组a转换成dim_total维数组,其中只有第dim_a个维度保留原数组的长度,其余维度都设为1,这样后续就能方便地进行广播运算对吧?
方法1:更紧凑的reshape写法
你原来用列表推导式生成目标形状,其实可以直接用元组拼接的方式构造形状参数,代码更短也更直观:
import numpy as np a = np.arange(10) dim_a = 2 dim_total = 4 # 最简reshape实现 a.reshape((1,) * dim_a + (-1,) + (1,) * (dim_total - dim_a - 1))
这里(1,)*dim_a会生成前dim_a个维度全为1的元组,(-1,)对应我们要保留原数组长度的dim_a维度,最后的(1,)*(dim_total - dim_a -1)补全剩下的维度为1,全程不用循环和条件判断,比列表推导式清爽多了。
方法2:简化版np.expand_dims调用
你原来的np.expand_dims写法需要先构造轴列表再删除指定项,其实可以直接生成所有需要扩展的轴的元组,一步传参:
np.expand_dims(a, axis=tuple(range(dim_a)) + tuple(range(dim_a+1, dim_total)))
range(dim_a)是dim_a之前的所有轴,range(dim_a+1, dim_total)是dim_a之后的所有轴,把它们拼接成元组传给axis参数,就会在这些轴上都扩展出长度为1的维度,刚好符合我们的要求。
方法3:利用索引语法的小技巧
NumPy里None和np.newaxis是等价的,都能用来扩展维度,我们还可以直接用索引语法实现,连函数都不用调用:
a[(None,) * dim_a + (slice(None),) + (None,) * (dim_total - dim_a - 1)]
这里slice(None)就等价于索引里的:,用来取原数组的所有元素,前后的(None,)*n则用来扩展出n个长度为1的维度,写法非常简洁。
以上这几种方法和你原来的代码效果完全一致,但代码量都更短,你可以根据自己的编程习惯选择合适的写法~
备注:内容来源于stack exchange,提问作者Axel Donath

