广播省掉的循环,和它悄悄给你的 (n, n)
广播(broadcasting)是这套工具里最像魔法的一环:形状不一样的两个数组,直接写加号就能算。它省掉的循环是实打实的——本机上一个距离矩阵,双重循环 134.3 毫秒,广播 2.28 毫秒,结果逐位相同。但同一条规则也是「内存炸了」这类事故的头号来源。这一章把规则完整摊开,它只有三行。
你有一万个数,想让每个数减去这一万个数的第一个(做一次基线对齐)。手一滑写成了这样:
col = x.reshape(-1, 1) # 形状 (10000, 1) row = x # 形状 (10000,) result = col - row
问:result 有多大?
规则只有三行
广播的完整规则如下。把两个形状右对齐写好,然后逐位比较:
- 维数不够的,左边补 1。
(4,)面对(3, 4)时,先变成(1, 4)。 - 每一位上,两个数要么相等,要么其中一个是 1。是 1 的那个被拉伸到另一个的长度。
- 否则报错。
关键在于第 1 条:补的是左边。所以一个一维数组永远被当成「横着的一行」,不是「竖着的一列」。上一章那个静默出错,根子就在这里。
拉伸不是真的复制。numpy 的做法是把那一维的 stride 设成 0——「往这个方向走一步,地址不动」。所以广播本身不花内存,花内存的是结果。
拿几个例子过一遍规则:
(3, 4) 与 (4,) → (1, 4) 拉成 (3, 4) ✓ (3, 4) (3, 4) 与 (3,) → (1, 3) 对不上 4 ✗ ⟦报错⟧ (3, 4) 与 (3, 1) → 第 1 维 1 拉成 4 ✓ (3, 4) (4, 1) 与 (1, 4) → 两边各拉一次 ✓ (4, 4) ★ (2, 3, 4) 与 (3, 1) → (1, 3, 1) 拉两次 ✓ (2, 3, 4)
标星的那一行是这一章的主角。(n, 1) 和 (1, n) 相遇,两边各拉一次,结果是 (n, n)。形状从 n 变成 n²——n = 10000 时,从 76 KB 变成 0.75 GB,正好 10000 倍。
而 (n,) 会被补成 (1, n)。所以只要你手上有一个 (n, 1)(很常见:任何一次 keepdims=True 的归约、任何一次 reshape(-1, 1)),再和一个原始的一维数组相遇,方阵就出现了,而且没有任何提示。
它省下来的东西是真的
说完危险,说说它为什么值得。经典例子:算 400 个三维点两两之间的距离。
# 双重循环
out = np.empty((n, n))
for i in range(n):
for j in range(n):
d = pts[i] - pts[j]
out[i, j] = math.sqrt(d[0]**2 + d[1]**2 + d[2]**2)
# 广播
np.sqrt(((pts[:, None, :] - pts[None, :, :]) ** 2).sum(-1))
本机实测:
双重循环 134.3 ms 广播 2.28 ms ★ 快 59 倍 两者最大差值 0.0 (逐位相同,不是「近似相等」)
那个 pts[:, None, :] 是广播的标准写法:None(也可以写 np.newaxis)在那个位置插一根长度为 1 的新轴。于是 (400, 3) 变成 (400, 1, 3) 和 (1, 400, 3),一广播就是 (400, 400, 3)——正好是「每一对点的坐标差」。
这段代码值得多看两眼,因为它演示了这套工具的思维方式:你不再描述「怎么遍历」,你描述「结果的形状是什么,每个格子里装什么」。循环消失了,不是因为它被优化掉了,是因为它被表达式吸收了。
但注意:这次省下的时间,是拿内存换的
上面那个 400 点的例子,中间结果 (400, 400, 3) 是 3.7 MB,无所谓。换成 40000 个点呢?(40000, 40000, 3) 是 34 TB。
这是广播的真实代价,也是它和「循环」的本质区别:循环是一次算一个,广播是一次把所有中间结果都摆出来。所以工程上有三条现成的出路:
| 做法 | 什么时候用 | 代价 |
|---|---|---|
| 分块(每次广播一小批) | 最常用。外层一个几十次的循环,内层广播 | 代码多几行 |
| 换一个不产生中间结果的公式 | 距离矩阵可以用 |a|²+|b|²−2a·b,走 BLAS | 数值精度略降 |
| 用专门的库 | scipy.spatial.cdist、sklearn 的 pairwise_distances | 多一个依赖 |
一条实用的自查:写完一个广播表达式,先在脑子里把结果形状算出来,再乘 8 字节。如果这个数超过你机器内存的十分之一,就该换写法了——而不是等它跑起来。
广播(broadcasting)这个词容易让人想到「广播电台」,其实它说的是把一个小的东西沿某个方向铺开,去匹配一个大的东西。更贴切的中文是「摊开」或者「铺满」。
要点是:铺开是虚的。numpy 把那一维的 stride 设成 0,读的时候反复读同一个地址,不占内存。所以「广播一个 (1, 10000) 的数组」不花钱,花钱的是它参与运算之后产生的那个结果。「广播免费,结果收费」——这句话是这一章的收据。
另外,同一套规则在 PyTorch、TensorFlow、JAX、Julia 里是一模一样的(它们都遵循 NumPy 的广播语义)。学一次,用四个生态。
第一件事:亲手看一次那个 (n, n)。用小的 n,别真开 0.75 GB。
import numpy as np x = np.arange(5.) print((x.reshape(-1, 1) - x).shape) # (5, 5) ← 你想要的是 (5,) print((x - x[0]).shape) # (5,) ← 这才是「减基线」 # 规则本身可以直接问 numpy,不用真的分配内存 print(np.broadcast_shapes((10000, 1), (10000,))) # (10000, 10000) print(np.broadcast_shapes((3, 4), (4,))) # (3, 4) np.broadcast_shapes((3, 4), (3,)) # 直接报错,省得真跑
第二件事:那个 59 倍。
import time, math
pts = np.random.default_rng(1).random((400, 3))
def loop():
out = np.empty((400, 400))
for i in range(400):
for j in range(400):
d = pts[i] - pts[j]
out[i, j] = math.sqrt(d[0]**2 + d[1]**2 + d[2]**2)
return out
def bcast():
return np.sqrt(((pts[:, None, :] - pts[None, :, :]) ** 2).sum(-1))
t = time.perf_counter(); a = loop(); t1 = time.perf_counter() - t
t = time.perf_counter(); b = bcast(); t2 = time.perf_counter() - t
print("%.1f ms vs %.2f ms 最大差 %g" % (t1*1e3, t2*1e3, np.abs(a-b).max()))
# 134.3 ms vs 2.28 ms 最大差 0
python3 -c "import numpy as np;print(np.broadcast_shapes((10000,1),(10000,)))"
np.broadcast_shapes 是这一章最值得记住的一个函数:它只算形状,不分配内存,所以可以拿它当广播规则的口算器用。
- 「CUDA out of memory」十次里有三次是这个。深度学习里最常见的显存爆炸不是模型太大,是某个中间张量的形状被广播成了
(batch, seq, seq, hidden)。注意力机制本身就是一个(n, n),所以序列一长显存就是平方增长——这也是各种「线性注意力」论文要解决的那个 n²。 - 图像处理里的每一次通道运算。给一张
(H, W, 3)的图做白平衡,写的是img * gains,其中gains是(3,)。广播规则从右往左对齐,正好对上通道维——这不是巧合,numpy 把通道放最后一维就是为了这个。如果你的图是(3, H, W)(PyTorch 的习惯),同一行代码就错了,要写成gains[:, None, None]。 - SQL 里的隐式笛卡尔积。
SELECT * FROM a, b忘了写WHERE,返回的行数是两表相乘。这和广播出一个(n, n)是同一种事故:一个本该是「配对」的操作,退化成了「所有组合」。第 10 章的 merge 会第三次遇到它。 - Excel 的绝对引用与相对引用。
$A1与A$1的区别,就是「哪个方向被固定、哪个方向被铺开」。拖动公式时方向拖反了,得到的也是一张全错但看起来很正常的表。
「形状不一样,numpy 会报错,所以广播出问题的时候我一定会知道。」
恰恰相反:广播的设计目的就是让形状不一样也能算。它只在完全对不上时报错((3, 4) 与 (3,)),而在对得上但不是你想要的那种对法时,它一声不吭地给你一个更大的数组。
最危险的组合是 (n, 1) 遇上 (n,):两个都长 n,看起来天经地义,结果是 (n, n)。而 (n, 1) 的来源特别多——keepdims=True、reshape(-1, 1)、pandas 的 df[['col']](双层方括号)、sklearn 要求的二维 X。
判据:任何一次形状不同的逐元素运算,先用 np.broadcast_shapes 问一遍结果形状。它不分配内存,一秒钟的事。真出了事故的另一个症状很好认:程序没报错,但内存占用突然涨到几个 GB,或者进程被系统杀掉——那多半就是某处广播出了一个方阵。
正确答案是 C:0.75 GB,一个 10000 × 10000 的方阵。
A 「76 KB」——这是你想要的大小。想要它的写法是x - x[0](右边是一个标量),或者 x - row[0:1]。多出来的那个 .reshape(-1, 1) 就是全部的差别。
B 「报错」——广播的存在意义就是不报错。(10000, 1) 和 (1, 10000)(补 1 之后)在每一位上都满足「相等或其中一个是 1」,规则完全通过,所以它开开心心地给了你一个方阵。
D 「152 KB,两个都留着」——广播不复制输入。输入那两个数组还是原来那么大(各 76 KB),被撑大的是输出。这个区分很重要:广播免费,广播的结果收费。
再补一句:0.75 GB 这个数还算「幸运」的,因为它至少能分配出来,你会看到进程变慢、风扇转起来。n 再大一点(比如 5 万),分配直接失败,抛 MemoryError——那反而是好事,起码它告诉你了。
(n, 1) 遇 (n,) 属于「对得上」,它会给你 (n, n)。自己检查:np.broadcast_shapes(a.shape, b.shape)。
这一章的一句话
广播用三行规则换掉了你所有的逐元素循环,代价是它永远不会问你要不要——形状对得上它就算,哪怕算出来的是一个比输入大一万倍的方阵。
下一章是卷 I 的最后一章,也是这套工具里最容易造成「改了这个那个跟着变」的一处:a[2:6] 和 a[[2,3,4,5]] 打印出来一模一样,但一个和 a 共用内存,一个不共用。改前者会改掉 a,改后者不会。而这条区别一路管到 pandas——第 7 章那个著名的 SettingWithCopyWarning,讲的是同一件事。