如何优雅识别NumPy广播过程中被复制的轴?
确定NumPy广播中被复制的轴的优化实现
你的原实现思路是正确的:先将原形状左侧补1对齐目标形状的长度,再对比对应轴的维度差异来找到被复制的轴。这里提供几种更简洁/高效的实现方式:
方式一:一行式列表推导(最简洁)
直接在列表推导中完成形状扩展与对比,省去单独定义extended_shape的变量:
shape = (3, 1, 1) broadcast_to = (4, 3, 1, 6) axis = [i for i, (orig_dim, target_dim) in enumerate( zip((1,) * (len(broadcast_to) - len(shape)) + shape, broadcast_to) ) if orig_dim != target_dim] # 输出: [0, 3]
方式二:利用NumPy向量化操作(适合复杂/大规模形状)
如果需要处理更复杂的形状或追求更高性能,可以用NumPy的数组操作来实现:
import numpy as np shape = (3, 1, 1) broadcast_to = (4, 3, 1, 6) # 将原形状左侧补1,对齐目标形状长度 extended_shape = np.pad( np.array(shape), pad_width=(len(broadcast_to)-len(shape), 0), mode='constant', constant_values=1 ) # 找到维度不匹配的轴索引 axis = np.where(extended_shape != np.array(broadcast_to))[0].tolist() # 输出: [0, 3]
方式三:借助广播机制验证(更直观)
可以通过创建空数组并执行广播,直接对比原数组与广播后数组的形状扩展关系:
import numpy as np shape = (3, 1, 1) broadcast_to = (4, 3, 1, 6) # 创建原形状空数组,广播到目标形状 orig_arr = np.empty(shape) broadcast_arr = np.broadcast_to(orig_arr, broadcast_to) # 扩展原形状到广播后形状的长度 extended_shape = (1,) * (broadcast_arr.ndim - orig_arr.ndim) + orig_arr.shape # 找到被复制的轴 axis = [i for i in range(broadcast_arr.ndim) if broadcast_arr.shape[i] != extended_shape[i]] # 输出: [0, 3]
这些方式都保留了原逻辑的正确性,同时在简洁性或性能上有所优化,可以根据实际场景选择。
内容的提问来源于stack exchange,提问作者Manuel Schmidt
相关产品推荐
相关产品推荐

