卷 I · 一条列CH 04深度 4/23

广播省掉的循环,和它悄悄给你的 (n, n)

广播(broadcasting)是这套工具里最像魔法的一环:形状不一样的两个数组,直接写加号就能算。它省掉的循环是实打实的——本机上一个距离矩阵,双重循环 134.3 毫秒,广播 2.28 毫秒,结果逐位相同。但同一条规则也是「内存炸了」这类事故的头号来源。这一章把规则完整摊开,它只有三行。

规则三行快 59 倍76 KB → 0.75 GB

▷ 先猜一下

你有一万个数,想让每个数减去这一万个数的第一个(做一次基线对齐)。手一滑写成了这样:

col = x.reshape(-1, 1)      # 形状 (10000, 1)
row = x                     # 形状 (10000,)
result = col - row

问:result 有多大?

A 76 KB。一万个 float64,本来就该这么大 B 报错。(10000, 1) 和 (10000,) 形状不同,不能相减 C 0.75 GB。结果是一个 10000 × 10000 的方阵 D 152 KB。numpy 会把两个都留着

规则只有三行

广播的完整规则如下。把两个形状右对齐写好,然后逐位比较:

◆ 主线
  1. 维数不够的,左边补 1。(4,) 面对 (3, 4) 时,先变成 (1, 4)
  2. 每一位上,两个数要么相等,要么其中一个是 1。是 1 的那个被拉伸到另一个的长度。
  3. 否则报错。

关键在于第 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 = 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.cdistsklearnpairwise_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 的绝对引用与相对引用。$A1A$1 的区别,就是「哪个方向被固定、哪个方向被铺开」。拖动公式时方向拖反了,得到的也是一张全错但看起来很正常的表。
✗ 这个直觉是错的

「形状不一样,numpy 会报错,所以广播出问题的时候我一定会知道。」

恰恰相反:广播的设计目的就是让形状不一样也能算。它只在完全对不上时报错((3, 4)(3,)),而在对得上但不是你想要的那种对法时,它一声不吭地给你一个更大的数组。

最危险的组合是 (n, 1) 遇上 (n,):两个都长 n,看起来天经地义,结果是 (n, n)。而 (n, 1) 的来源特别多——keepdims=Truereshape(-1, 1)、pandas 的 df[['col']](双层方括号)、sklearn 要求的二维 X。

判据:任何一次形状不同的逐元素运算,先用 np.broadcast_shapes 问一遍结果形状。它不分配内存,一秒钟的事。真出了事故的另一个症状很好认:程序没报错,但内存占用突然涨到几个 GB,或者进程被系统杀掉——那多半就是某处广播出了一个方阵。

◇ 揭晓

正确答案是 C0.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——那反而是好事,起码它告诉你了。
⌗ 交接单
两个形状不同的数组。 一个形状是「两者右对齐后逐位取大」的新数组。它可能比两个输入加起来还大好几个数量级。 只有完全对不上时 numpy 才报错。(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,讲的是同一件事。