将 N 维数组广播到匹配的 (N+1) 维数组的最直接方法是使用np.broadcast_to():
import numpy as np
arr = np.random.randint(0, 100, (2, 3))
mask = np.random.randint(0, 2, (2, 3, 4), dtype=bool)
b_arr = np.broadcast_to(arr[..., None], mask.shape)
print(mask.shape == b_arr.shape)
# True
但是,正如@hpaulj 已经指出的那样,您不能使用mask 对b_arr 进行切片而不丢失尺寸。
鉴于您只想将元素加在一起并将零相加“不会造成伤害”,您可以简单地将数组和掩码按元素相乘,以保持正确的维度,但 False 中的元素掩码与相应数组元素的后续sum 无关:
result = np.sum(b_arr * mask, axis=tuple(range(mask.ndim - 1)))
或者,因为* 会自动进行广播:
result = np.sum(arr[..., None] * mask, axis=tuple(range(mask.ndim - 1)))
首先不需要使用np.broadcast_to()(但您仍然需要匹配维度的数量,即使用arr[..., None] 而不仅仅是arr)。
作为@PaulPanzer already pointed out,由于您想要总结除一维之外的所有维度,因此可以使用np.matmul()/@ 进一步简化:
result2 = arr.ravel() @ mask.reshape(-1, mask.shape[-1])
print(np.all(result == result2))
# True
对于涉及求和的更高级的操作,请查看np.einsum()。
编辑
广播的问题是它会在评估表达式期间创建临时数组。
对于您似乎正在处理的数字,当我遇到MemoryError 时,我根本无法使用广播数组,但从时间上看,元素乘法可能仍然是比您最初建议的更好的方法。
或者,如果您追求速度,您可以通过 Cython 或 Numba 中的显式循环在较低级别上执行此操作。
您可以在下面找到几个基于 Numba 的解决方案(处理 ravel()-ed 数据):
-
_vector_matrix_product(): 不使用任何临时数组
-
_vector_matrix_product_mp(): 部分同上,但使用并行执行
-
_vector_matrix_product_sum():使用np.sum() 和并行执行
import numpy as np
import numba as nb
@nb.jit(nopython=True)
def _vector_matrix_product(
vect_arr,
mat_arr,
result_arr):
rows, cols = mat_arr.shape
if vect_arr.shape == result_arr.shape:
for i in range(rows):
for j in range(cols):
result_arr[i] += vect_arr[j] * mat_arr[i, j]
else:
for i in range(rows):
for j in range(cols):
result_arr[j] += vect_arr[i] * mat_arr[i, j]
@nb.jit(nopython=True, parallel=True)
def _vector_matrix_product_mp(
vect_arr,
mat_arr,
result_arr):
rows, cols = mat_arr.shape
if vect_arr.shape == result_arr.shape:
for i in nb.prange(rows):
for j in nb.prange(cols):
result_arr[i] += vect_arr[j] * mat_arr[i, j]
else:
for i in nb.prange(rows):
for j in nb.prange(cols):
result_arr[j] += vect_arr[i] * mat_arr[i, j]
@nb.jit(nopython=True, parallel=True)
def _vector_matrix_product_sum(
vect_arr,
mat_arr,
result_arr):
rows, cols = mat_arr.shape
if vect_arr.shape == result_arr.shape:
for i in nb.prange(rows):
result_arr[i] = np.sum(vect_arr * mat_arr[i, :])
else:
for j in nb.prange(cols):
result_arr[j] = np.sum(vect_arr * mat_arr[:, j])
def vector_matrix_product(
vect_arr,
mat_arr,
swap=False,
dtype=None,
mode=None):
rows, cols = mat_arr.shape
if not dtype:
dtype = (vect_arr[0] * mat_arr[0, 0]).dtype
if not swap:
result_arr = np.zeros(cols, dtype=dtype)
else:
result_arr = np.zeros(rows, dtype=dtype)
if mode == 'sum':
_vector_matrix_product_sum(vect_arr, mat_arr, result_arr)
elif mode == 'mp':
_vector_matrix_product_mp(vect_arr, mat_arr, result_arr)
else:
_vector_matrix_product(vect_arr, mat_arr, result_arr)
return result_arr
np.random.seed(0)
arr = np.random.randint(0, 100, (2, 3, 4))
mask = np.random.randint(0, 2, (2, 3, 4, 5), dtype=bool)
target = arr.ravel() @ mask.reshape(-1, mask.shape[-1])
print(target)
# [820 723 861 486 408]
result1 = vector_matrix_product(arr.ravel(), mask.reshape(-1, mask.shape[-1]))
print(result1)
# [820 723 861 486 408]
result2 = vector_matrix_product(arr.ravel(), mask.reshape(-1, mask.shape[-1]), mode='mp')
print(result2)
# [820 723 861 486 408]
result3 = vector_matrix_product(arr.ravel(), mask.reshape(-1, mask.shape[-1]), mode='sum')
print(result3)
# [820 723 861 486 408]
与任何基于 list-comprehension 的解决方案相比,时间有所改进:
arr = np.random.randint(0, 100, (256, 256, 256))
mask = np.random.randint(0, 2, (256, 256, 256, 128), dtype=bool)
%timeit np.sum(arr[..., None] * mask, axis=tuple(range(mask.ndim - 1)))
# MemoryError
%timeit arr.ravel() @ mask.reshape(-1, mask.shape[-1])
# MemoryError
%timeit np.array([np.sum(arr * mask[..., i], axis=tuple(range(mask.ndim - 1))) for i in range(mask.shape[-1])])
# 24.1 s ± 105 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
%timeit np.array([np.sum(arr[mask[..., i]]) for i in range(mask.shape[-1])])
# 46 s ± 119 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
%timeit vector_matrix_product(arr.ravel(), mask.reshape(-1, mask.shape[-1]))
# 408 ms ± 2.12 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
%timeit vector_matrix_product(arr.ravel(), mask.reshape(-1, mask.shape[-1]), mode='mp')
# 1.63 s ± 3.58 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
%timeit vector_matrix_product(arr.ravel(), mask.reshape(-1, mask.shape[-1]), mode='sum')
# 7.17 s ± 258 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
正如预期的那样,JIT 加速版本是最快的,并且对代码强制执行并行性并不会提高速度。
另请注意,逐元素乘法的方法比切片更快(这些基准的速度大约是切片的两倍)。
编辑 2
按照@max9111 的建议,首先按行循环,然后按列循环会导致最耗时的循环在连续数据上运行,从而显着提高速度。
如果没有这个技巧,_vector_matrix_product_sum() 和 _vector_matrix_product_mp() 将以基本相同的速度运行。