Skip to content

NumPy - 广播

广播 (Broadcasting) 描述了 NumPy 如何处理不同形状的数组之间的算术运算。通常,运算是逐元素 (element-wise) 进行的,这要求数组具有完全相同的形状 (shape)。

两个相同形状数组之间的逐元素乘法。

import numpy as np
a = np.array([1, 2, 3, 4])
b = np.array([10, 20, 30, 40])
c = a * b # 逐元素乘法
print(f"Array a: {a}")
print(f"Array b: {b}")
print(f"Result c = a * b: {c}")

输出:

Array a: [1 2 3 4]
Array b: [10 20 30 40]
Result c = a * b: [ 10 40 90 160]

然而,NumPy 还可以通过一个强大的机制,称为 广播 (broadcasting),对兼容但形状不同的数组执行运算。如果满足某些规则,较小的数组在概念上会被“广播”(拉伸或平铺),以匹配较大数组的形状,而无需实际复制数据,这使得运算非常节省内存。

广播规则:

如果对于每个维度(从末尾维度开始),以下任一条件成立,则两个数组兼容广播:

  • 它们的维度大小相等,或者
  • 其中一个维度的大小是 1。

如果数组具有不同数量的维度,维度较少的数组的形状在概念上会在其左侧用 1 进行填充,直到维度数量匹配。

在运算过程中,大小为 1 的维度会被“拉伸”或“平铺”,以匹配另一个数组中相应的维度大小。结果数组的形状是输入数组在该维度大小中的最大值。

常见场景:数组和标量 (scalar)(单个数字)之间的运算。标量会被广播以匹配数组的形状。

import numpy as np
arr = np.array([1, 2, 3])
scalar = 10
result = arr + scalar # 标量被广播
print(result) # Output: [11 12 13]

示例 2:将一维数组广播到二维数组

Section titled “示例 2:将一维数组广播到二维数组”

演示将一个一维数组添加到二维数组的每一行。

import numpy as np
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]
]) # 形状 (4, 3)
b = np.array([1.0, 2.0, 3.0]) # 形状 (3,)
print('First array (a, shape {}):'.format(a.shape))
print(a)
print('\n')
print('Second array (b, shape {}):'.format(b.shape))
print(b)
print('\n')
# 将 a 和 b 相加。这里发生了广播。
# b 的形状 (3,) 与 a 的末尾维度 (3) 兼容。
# 在概念上,b 沿着第一维度拉伸,以匹配 a 的形状 (4, 3)。
c = a + b
print('Result of a + b (shape {}):'.format(c.shape))
print(c)

输出:

First array (a, shape (4, 3)):
[[ 0. 0. 0.]
[10. 10. 10.]
[20. 20. 20.]
[30. 30. 30.]]
Second array (b, shape (3,)):
[1. 2. 3.]
Result of a + b (shape (4, 3)):
[[ 1. 2. 3.]
[11. 12. 13.]
[21. 22. 23.]
[31. 32. 33.]]

广播的可视化:

要执行 a + b,其中 a 的形状为 (4, 3),b 的形状为 (3,):

  1. NumPy 从右到左比较形状。
  2. 末尾维度:a 有 3,b 有 3。它们匹配。
  3. 下一个维度:a 有 4,b 在这里没有维度。NumPy 在概念上为 b 添加一个大小为 1 的维度,使其形状变为 (1, 3)。
  4. 再次比较:(4, 3) 和 (1, 3)。
  5. 末尾维度(3 对 3):匹配。
  6. 第一维度(4 对 1):其中一个是 1。兼容。
  7. 结果形状是每个维度中输入数组大小的最大值:(max(4,1), max(3,3)) = (4, 3)。
  8. 在运算过程中,b 中大小为 1 的维度(现在在概念上是 (1, 3))沿着第一轴(行)拉伸,以匹配 a 的大小 4。对于加法运算,它实际上就像 [[1., 2., 3.], [1., 2., 3.], [1., 2., 3.], [1., 2., 3.]] 一样,但无需在内存中创建这个临时数组。

广播是一个强大的功能,可以编写简洁高效的代码,但请注意遵循规则,以避免意外的形状不匹配或非预期的操作。