广播描述了 NumPy 如何在算术运算时处理具有不同形状数组的方式,从而使通用函数以直白的方式处理形状不完全相同的输入。在受规则约束的前提下,较小的数组在较大的数组上“广播”,以便它们具有兼容的形状。
NumPy 操作通常是在一对数组上逐元素进行。在最简单的情况下,两个数组必须具有完全相同的形状,如下所示:
import numpy as np
a = np.array([1., 2., 3.])
b = np.array([2., 2., 2.])
print(a * b) # [2. 4. 6.]
当数组的形状满足某些条件时,NumPy 的广播规则放宽了形状约束。如一个数组和一个标量值的操作组合,会发生最简单的广播示例:
import numpy as np
a = np.array([1., 2., 3.])
b = 2.0
print(a * b) # [2. 4. 6.]
这里如同标量 b 被拉伸为和 a 形状相同的数组(NumPy 足够聪明,实际并没有发生拉伸)。
当对两个数组进行运算时,理想情况是这两个数组的形状完全相同(维度数相同,各维度也相同;或者说形状元组的元素个数相同,各元素值也相同)。广播机制使得即便达不到理想情况,也能进行运算。针对维度数不同和各维度不同,分别有两条规则:
- 当两个数组的维度数不同时,对于较小维度数的数组,会在形状元组的左侧反复为缺失的维度补充 1,直到所有数组具有相同数量的维度。
- 若数组的某维度为 1,而另一数组对应的维度大于 1,则该数组元素的值会进行广播,表现为具有沿该维度最大形状的数组的大小,且各元素值重复多次。
以上的两个规则不太容易理解,这里举例说明。对于如下两个数组 a 和 b,他们的形状分别为 (4, 3) 和 (3)。
a = np.array([[ 0.0, 0.0, 0.0],
[10.0, 10.0, 10.0],
[20.0, 20.0, 20.0],
[30.0, 30.0, 30.0]])
b = np.array([1.0, 2.0, 3.0])
print(a + b)
# [[ 1. 2. 3.]
# [11. 12. 13.]
# [21. 22. 23.]
# [31. 32. 33.]]
根据第 1 条规则,会把数组 b 的形状看作是 (1, 3),即其值由 [1.0, 2.0, 3.0] 变为 [[1.0 2.0 3.0]],这样两个数组将具有相同的维度数。
根据第 2 条规则,会进一步把形状 (1, 3) 中的维度 1 拉伸为 4,即将该维度对应的唯一元素 [1.0 2.0 3.0] 重复 4 次,相当于数组 b 变为:
[[1.0 2.0 3.0]
[1.0 2.0 3.0]
[1.0 2.0 3.0]
[1.0 2.0 3.0]]
这样数组 a 和 b 就能进行运算了。下图为运算过程示意图。
alt text
关于两个或更多个数组是否能根据以上规则进行运算,可简单地把这些数组的形状放在一起,按照从右到左的顺序比较每一维度,缺失的元素不参与比较,当所有参与比较的元素值都满足如下条件时,就表示他们的维度是匹配的,或者说他们是可广播的:
- 两个维度相等(即实现了精确匹配);
- 其中一个维度为 1(即维度 1 可以和任何维度相匹配)。
若不满足这两个条件,在运算时会抛出一个特定的异常,表明这参与运算的数组的形状不匹配。
根据此方法,可快速得出以下数组是可广播的:
A (2d array): 5 x 4
B (1d array): 1
Result (2d array): 5 x 4
A (2d array): 5 x 4
B (1d array): 4
Result (2d array): 5 x 4
A (3d array): 15 x 3 x 5
B (3d array): 15 x 1 x 5
Result (3d array): 15 x 3 x 5
A (3d array): 15 x 3 x 5
B (2d array): 3 x 5
Result (3d array): 15 x 3 x 5
A (3d array): 15 x 3 x 5
B (2d array): 3 x 1
Result (3d array): 15 x 3 x 5
以下数组是不可广播的:
A (1d array): 3
B (1d array): 4 # 最后的维度不匹配
A (2d array): 2 x 1
B (3d array): 8 x 4 x 3 # 倒数第二个维度不匹配