广播

广播描述了 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,直到所有数组具有相同数量的维度。
  2. 若数组的某维度为 1,而另一数组对应的维度大于 1,则该数组元素的值会进行广播,表现为具有沿该维度最大形状的数组的大小,且各元素值重复多次。

以上的两个规则不太容易理解,这里举例说明。对于如下两个数组 ab,他们的形状分别为 (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]]

这样数组 ab 就能进行运算了。下图为运算过程示意图。

alt text

alt text

关于两个或更多个数组是否能根据以上规则进行运算,可简单地把这些数组的形状放在一起,按照从右到左的顺序比较每一维度,缺失的元素不参与比较,当所有参与比较的元素值都满足如下条件时,就表示他们的维度是匹配的,或者说他们是可广播的:

  1. 两个维度相等(即实现了精确匹配);
  2. 其中一个维度为 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 # 倒数第二个维度不匹配
拷贝和视图高级索引和索引技巧