🏠 总目录📚 本教程 03 · 广播的三条规则 ← →
📑 本页目录(点开跳转)

03 · 广播的三条规则

⏱ 94 分钟 | ⭐ (5,) - (5,1) 不报错,它给你一张 5×5 的表 —— 你的 MSE 从 0.022 变成 3.822,而程序一声不吭


🎯 一句话

广播只做三件事:左边补 1、逐位比对、把长度 1 的那位拉伸开 —— 全程不看数值,只看形状元组。

这三条规则总共不到三十个字,但它们决定了你写的每一行数组运算是算对、报错、还是静默算错。 第三种最贵,本章第五节就是它。


⚠️ 〇、先消歧:站内的「广播」有两个意思

第 00 章钉过两个同名不同物(「向量化」和「stride」),这里是第三个, 而且它比前两个更容易混,因为两边都在讲「一份数据被很多地方用到」:

站内出处 那里的「广播」是 命中
AI 基础设施 10 · 数据并行与 AllReduce 通信原语:一张卡把结果发给其他所有卡(「机内 NVLink 先归约 → 跨机 → 机内广播」) 3 处(含 ai_md/05)
密码学 14 · RSA Håstad 广播攻击:同一条消息用 e=3 发给 3 个不同的人 3 处
⭐ ML 基础 18 · 手搓 mini-torch NumPy 义:形状不同的两个数组怎么对齐相加 9 处
ML 基础 附录 C 同上(# (n,k) 广播算距离) 1 处

⭐ 注意最后两行的分布:NumPy 义的 10 处里有 9 处集中在 ml_md/18, 而那一章要你做的事是 —— 实现广播的反向传播:

ml_md/18:63:T2 (120min) 实现 + * @ 三个运算的前向和反向。⭐ 难点:广播的反向(形状不匹配时梯度要 sum 回去)

ml_md/18:156(专属坑表):广播的反向 ⭐ 最常错 —— (32,64) + (64,) 的反向要把梯度 sum 到 (64,)。不处理会形状报错或静默算错

💀 也就是说:全站唯一一处认真对待广播的地方,是让你写它的逆运算 —— 而正向那三条规则从头到尾没有一行字写出来过。 本章补的就是这个。 第八节会把 ml_md/18 里那个 unbroadcast() 函数逐行摆回三条规则上,你会看到它就是规则的镜像。

⚠️ 本章往下所有的「广播」一律是 NumPy 义,不再重复声明。


📏 一、三条规则

拿到两个形状不同的数组做逐元素运算(+ - * / 和一切 ufunc),NumPy 只做这三步:

# 规则 说的是什么
① 补齐 维数少的那个,在形状元组左边补 1,补到两边一样长
② 比对 对齐之后逐位比:要么两个数相等,要么其中一个是 1。有一位不满足就 ValueError
③ 拉伸 是 1 的那一位,被拉伸到对方的长度。⭐ 拉伸不复制内存(第三节证明)

⚠️ 规则 ① 只往左边补,永远不往右边补。 这一条是后面所有坑的根源 —— (3,4) 和 (4,) 能算((4,) 补成 (1,4),对上了), (3,4) 和 (3,) 就不能((3,) 补成 (1,3),末位 4 对 3,炸)。

import numpy as np

a = np.ones((3, 4))          # 3 个样本 x 4 个特征
b = np.array([10., 20., 30., 40.])   # 4 个特征各自的偏移

print("a.shape =", a.shape, "  b.shape =", b.shape)
print("(a + b).shape =", (a + b).shape)
print(a + b)

c = np.array([10., 20., 30.])        # 3 个样本各自的偏移
print("\nc.shape =", c.shape)
try:
    a + c
except ValueError as e:
    print("a + c  ->  ValueError:", e)

print("\n给 c 补一根轴,让它变成 (3,1):")
print("c[:, None].shape =", c[:, None].shape)
print((a + c[:, None]))

实跑输出:

关键信息

a.shape = (3, 4) b.shape = (4,)
(a + b).shape = (3, 4)
[[11. 21. 31. 41.]
[11. 21. 31. 41.]
[11. 21. 31. 41.]]
c.shape = (3,)
a + c -> ValueError: operands could not be broadcast together with shapes (3,4) (3,)
给 c 补一根轴,让它变成 (3,1):
c[:, None].shape = (3, 1)
[[11. 11. 11. 11.]
[21. 21. 21. 21.]
[31. 31. 31. 31.]]

⭐ 这段输出就是全章的缩影: 「每个特征加一个偏移」((4,))直接就能写,因为特征轴在最右边,规则 ① 自动帮你补上了; 「每个样本加一个偏移」((3,))必须自己写 [:, None],因为样本轴在最左边,规则 ① 帮不上忙。

⭐ 一句话记法:右边的轴是免费的,左边的轴要你自己动手。 第 02 章说「写库的人几乎一律用 axis=-1」是同一个理由的另一面 —— 最右边那根轴最不需要照顾。

⚠️ 广播不看数值,只看形状。 a + c 报错不是因为 10、20、30 有什么问题, 是因为 (3,4) 和 (1,3) 的末位对不上。报错信息里也只有形状,这是个好消息 —— 你永远可以只盯着形状元组调试,不用管数据内容。


🧮 二、手算:一张对照表

规则说完了,但真正的功夫是看到两个形状能立刻在脑子里跑一遍。 np.broadcast_shapes 可以让你不造数据就问结果:

import numpy as np

pairs = [
    ((3, 4), (4,)),
    ((3, 4), (3,)),
    ((3, 4), (3, 1)),
    ((3, 1), (1, 4)),
    ((2, 3, 4), (4,)),
    ((2, 3, 4), (3, 1)),
    ((2, 3, 4), (2, 1, 4)),
    ((2, 3, 4), (2, 3)),
    ((5,), (5, 1)),
]

for s1, s2 in pairs:
    try:
        out = np.broadcast_shapes(s1, s2)
        print(f"{str(s1):>12} + {str(s2):>10}  ->  {out}")
    except ValueError as e:
        print(f"{str(s1):>12} + {str(s2):>10}  ->  ValueError: {e}")

实跑输出:

关键信息

(3, 4) + (4,) -> (3, 4)
(3, 4) + (3,) -> ValueError: shape mismatch: objects cannot be broadcast to a single shape. Mismatch is between arg 0 with shape (3, 4) and arg 1 with shape (3,).
(3, 4) + (3, 1) -> (3, 4)
(3, 1) + (1, 4) -> (3, 4)
(2, 3, 4) + (4,) -> (2, 3, 4)
(2, 3, 4) + (3, 1) -> (2, 3, 4)
(2, 3, 4) + (2, 1, 4) -> (2, 3, 4)
(2, 3, 4) + (2, 3) -> ValueError: shape mismatch: objects cannot be broadcast to a single shape. Mismatch is between arg 0 with shape (2, 3, 4) and arg 1 with shape (2, 3).
(5,) + (5, 1) -> (5, 5)

手算的时候把两个形状右对齐写成两行,比对着看:

例子 右对齐之后 结论
(3,4) + (4,) 3 4 / 1 4 首位 3 对 1 → 拉伸;末位 4=4 → ✅ (3,4)
(3,4) + (3,) 3 4 / 1 3 末位 4 对 3,都不是 1 → ❌
(3,1) + (1,4) 3 1 / 1 4 两位各有一个 1 → 两边都被拉伸 → ✅ (3,4)
(2,3,4) + (3,1) 2 3 4 / 1 3 1 全部满足 → ✅ (2,3,4)
(2,3,4) + (2,3) 2 3 4 / 1 2 3 末位 4 对 3 → ❌(⚠️ (2,3) 看着像前两位,实际被补到了后两位)

⚠️ 倒数第二行是最反直觉的一条:(2,3,4) 和 (2,3) 看起来「前两维完全一样,应该能算」, 但规则 ① 是往左补,(2,3) 变成 (1,2,3) 而不是 (2,3,1)。 要表达「前两维对齐」得自己写成 (2,3,1),也就是 y[..., None]。

⭐ np.broadcast_shapes() 值得记住 —— 调试形状问题时它比造两个数组再相加快得多, 而且不分配任何内存。


🧠 三、为什么拉伸是免费的

规则 ③ 说「拉伸不复制内存」。这不是修辞,是可以直接验的:

import numpy as np

v = np.arange(5, dtype=np.float64)          # 5 个数,40 字节
big = np.broadcast_to(v, (1_000_000, 5))    # 假装有 100 万行

print("v.shape       =", v.shape,   " v.nbytes =", v.nbytes)
print("big.shape     =", big.shape, " big.nbytes(名义) =", big.nbytes)
print("共享内存吗   ->", np.shares_memory(v, big))
print("big 自己有数据吗 -> big.base is v :", big.base is v)
print("能写吗       -> big.flags.writeable :", big.flags.writeable)

# 真正实体化一份要多少
real = np.ascontiguousarray(big)
print("\nnp.ascontiguousarray(big).nbytes =", real.nbytes, "= %.1f MB" % (real.nbytes / 1024 / 1024))
print("real 共享内存吗 ->", np.shares_memory(v, real))

实跑输出:

算一算

v.shape = (5,) v.nbytes = 40

big.shape = (1000000, 5) big.nbytes(名义) = 40000000

共享内存吗 -> True

big 自己有数据吗 -> big.base is v : True

能写吗 -> big.flags.writeable : False

np.ascontiguousarray(big).nbytes = 40000000 = 38.1 MB

real 共享内存吗 -> False

⭐ 一个 40 字节的数组,被「拉伸」成号称 38.1 MB 的形状,实际占用还是那 40 字节。 shares_memory 是 True,big.base is v 也是 True —— 它压根没有自己的数据, 只是给同一块内存换了一套读法。真要实体化(ascontiguousarray)才会掏出那 38.1 MB。

⚠️ 代价是它只读(writeable 是 False)。理由很直白: big[0, 0] 和 big[999999, 0] 是同一个字节,往里写一个数会让一百万行同时变, NumPy 干脆禁掉这个歧义。这一条在第六节还会以另一副面孔出现。

「换一套读法」到底怎么换的,是 第 05 章 的正题 (答案是那根轴的步长被设成 0)。这一章你只需要接受「它是免费的」这个事实。 第 07 章 里 sliding_window_view 展开成 (199951, 50) 却 shares_memory=True,用的是同一个机制。


🚀 四、最值钱的模式:一次拿到所有两两组合

第 02 章末尾埋了 v[:, None] - v[None, :] 这个组合,现在兑现它:

import numpy as np
import time

rng = np.random.default_rng(0)
X = rng.random((400, 2))

# 双重循环版
t0 = time.perf_counter()
D_loop = np.empty((400, 400))
for i in range(400):
    for j in range(400):
        d = X[i] - X[j]
        D_loop[i, j] = np.sqrt(d[0] * d[0] + d[1] * d[1])
t_loop = time.perf_counter() - t0

# 广播版
t0 = time.perf_counter()
diff = X[:, None, :] - X[None, :, :]      # (400,1,2) 和 (1,400,2) -> (400,400,2)
D_bc = np.sqrt((diff ** 2).sum(axis=-1))
t_bc = time.perf_counter() - t0

print("diff.shape =", diff.shape)
print("D_bc.shape =", D_bc.shape)
print("两版结果最大差 =", np.abs(D_loop - D_bc).max())
print("循环版 %.4f s   广播版 %.4f s   快 %.0f 倍" % (t_loop, t_bc, t_loop / t_bc))

# 点数一涨,中间量会怎样(只算不分配)
for n in (400, 4000, 20000):
    b = n * n * 2 * 8
    print("N=%-6d 时中间量 diff 要 %10.1f MB" % (n, b / 1024 ** 2))

实跑输出:

算一算

diff.shape = (400, 400, 2)

D_bc.shape = (400, 400)

两版结果最大差 = 0.0

循环版 0.5425 s 广播版 0.0122 s 快 44 倍

N=400 时中间量 diff 要 2.4 MB

N=4000 时中间量 diff 要 244.1 MB

N=20000 时中间量 diff 要 6103.5 MB

⚠️ 倍数看量级不看单点 —— 这台机器有热降频,同一段代码重跑三次是 44x / 57x / 68x, 第 01 章也是给的区间。「快几十倍」是结论,「44」不是。 两版结果最大差是 0.0(完全相同,不是「误差很小」),说明改写没有引入任何数值差异。

怎么读 X[:, None, :] - X[None, :, :]:

写法 形状 含义
X (400, 2) 400 个点,每个 2 维
X[:, None, :] (400, 1, 2) 「我是谁」放 0 号轴
X[None, :, :] (1, 400, 2) 「他是谁」放 1 号轴
相减(规则 ③ 两边都拉伸) (400, 400, 2) [i, j] 就是第 i 个点减第 j 个点
.sum(axis=-1) 消掉坐标轴 (400, 400) 完整的距离矩阵

⭐ 这个模式的通用形状是「把要配对的两个东西分别塞进相邻的两根轴,中间用 None 隔开」。 ML 基础 附录 C 第 399 行的 K-means d = ((X[:, None, :] - C[None, :, :]) ** 2).sum(-1) 就是它(那里配对的是样本和聚类中心, 注释只写了一句「(n,k) 广播算距离」,为什么是 (n,k) 没有解释过)。

⚠️ 代价写在最后三行:中间量 diff 的大小是 N² × D × 8 字节, 400 个点才 2.4 MB,20000 个点就要 6103.5 MB —— 这是第七节的正题。


🛑 读到这里可以停 —— 前半章讲完了(约 34 分钟)。 后半章还有:最贵的坑:(n,) 和 (n,1) 不报错 · 原地运算时,广播是单向的 · 内存爆炸和它的解药 · 反过来看:unbroadcast 就是三条规则的镜像 · 一页速记 回来的时候不用重读,直接从下一节接着看就行。


💀 五、最贵的坑:(n,) 和 (n,1) 不报错

前面所有例子里,形状对不上都会 ValueError。真正会让你损失时间的是不报错的那一类。

回头看第二节最后一行:(5,) + (5,1) -> (5,5)。它完全符合三条规则 —— (5,) 补成 (1,5),和 (5,1) 逐位一比,两位各有一个 1,两边都拉伸,得 (5,5)。 规则很讲道理,结果很致命:

import numpy as np

y_true = np.array([1.0, 2.0, 3.0, 4.0, 5.0])          # (5,)
y_pred = np.array([[1.1], [2.1], [2.9], [4.2], [4.8]])  # (5,1)  <- 从 model.predict 出来常是这个形状

print("y_true.shape =", y_true.shape, "  y_pred.shape =", y_pred.shape)

err = y_true - y_pred
print("(y_true - y_pred).shape =", err.shape, "  <- 不报错,变成了 5x5 的差值表")
print("mse 算出来 =", (err ** 2).mean())

right = y_true - y_pred.ravel()
print("\n压平之后 shape =", right.shape)
print("mse 正确值 =", (right ** 2).mean())

print("\n用 np.broadcast_shapes 提前查:")
print("np.broadcast_shapes((5,), (5,1)) =", np.broadcast_shapes((5,), (5, 1)))

实跑输出:

对照

y_true.shape = (5,) y_pred.shape = (5, 1)

(y_true - y_pred).shape = (5, 5) <- 不报错,变成了 5x5 的差值表

mse 算出来 = 3.822

压平之后 shape = (5,)

mse 正确值 = 0.022000000000000037

💀 正确的 MSE 是 0.022,算出来的是 3.822 —— 大了 173 倍,而程序一个字都没吭。 你的模型明明学得不错,指标却糟得离谱,于是你会去查数据、查学习率、查特征, 唯独不会怀疑那行减法,因为它跑通了。

⚠️ 为什么这个坑高频:(n,1) 是很多东西的天然输出形状 —— sklearn 某些 predict 的返回、keepdims=True 的归约结果、 数据库读出来的单列、reshape(-1, 1) 之后的特征。 而 (n,) 是标签的天然形状。两者相遇,就是一张 n×n 的表。

⭐ 三条防线,按可靠性排序:

防线 写法 说明
⭐ 断言形状 assert y_true.shape == y_pred.shape, (y_true.shape, y_pred.shape) 最硬,失败时把两个形状一起打出来
入口统一压平 y_pred = np.asarray(y_pred).ravel() 在函数入口做一次,后面全干净
事前查 np.broadcast_shapes(a.shape, b.shape) 调试时用,看到 (5,5) 就知道错了

⚠️ 不要指望「结果形状不对我会看见」 —— .mean() 之后就是一个标量, 形状信息在报出指标之前就已经被 mean() 抹掉了。


🧯 六、原地运算时,广播是单向的

a + b 和 a += b 在广播这件事上不是同一回事:

import numpy as np

a = np.zeros((3, 4))
b = np.array([1., 2., 3., 4.])
a += b                      # (3,4) += (4,)  结果还是 (3,4),能原地写
print("a += b 之后 a.shape =", a.shape)
print(a)

small = np.zeros(4)
big = np.zeros((3, 4))
try:
    small += big            # 结果是 (3,4),塞不回 (4,) 的容器
except ValueError as e:
    print("\nsmall += big  ->  ValueError:", e)

print("\n非原地写法就没事:", (small + big).shape)

# 广播出来的那份是只读的
v = np.arange(3)
bt = np.broadcast_to(v, (2, 3))
try:
    bt[0, 0] = 99
except ValueError as e:
    print("\n往 broadcast_to 的结果里写  ->  ValueError:", e)

实跑输出:

关键信息

a += b 之后 a.shape = (3, 4)
[[1. 2. 3. 4.]
[1. 2. 3. 4.]
[1. 2. 3. 4.]]
small += big -> ValueError: non-broadcastable output operand with shape (4,) doesn't match the broadcast shape (3,4)
非原地写法就没事: (3, 4)
往 broadcast_to 的结果里写 -> ValueError: assignment destination is read-only

⭐ 规则很简单:a += b 要求广播结果的形状和 a 完全一样。 a 是那个「容器」,它不许在原地长大。所以广播只能是单向的 —— b 可以被拉伸去迎合 a, a 不能被拉伸去迎合 b。

⚠️ non-broadcastable output operand 这句报错专属于原地运算。 看到它不要去查两个形状能不能广播(它们能,small + big 就成功了), 要去看等号左边那个数组的形状。同一族的还有 Incompatible shape for in-place modification,第 00 章 的症状路线表把这两条都指向了 第 05 章。

⭐ 顺带一提:a += b 和 a = a + b 对 NumPy 数组是两件不同的事 —— 前者原地改那块内存(别的视图看得见),后者造一个新数组换个名字(别人看不见)。 这句话的完整版是 第 04 章 的正题。


🔬 七、内存爆炸和它的解药

第四节最后那三行数字是本节的动机:广播本身免费,但广播出来的结果是实打实的内存。 X[:, None, :] - X[None, :, :] 里被拉伸的两个操作数不占内存,它们相减产生的那个 (N,N,D) 占。

经典解药是把平方展开成矩阵乘法:‖a−b‖² = ‖a‖² − 2a·b + ‖b‖², 右边三项的中间量最大只有 (N,N),把 D 那一维彻底消掉:

import numpy as np
import time

rng = np.random.default_rng(0)
A = rng.random((3000, 16))
B = rng.random((3000, 16))

# 广播版:中间量 (3000, 3000, 16)
t0 = time.perf_counter()
D1 = np.sqrt(((A[:, None, :] - B[None, :, :]) ** 2).sum(-1))
t1 = time.perf_counter() - t0
mid = 3000 * 3000 * 16 * 8

# 展开版:|a-b|^2 = |a|^2 - 2a·b + |b|^2,中间量只有 (3000, 3000)
t0 = time.perf_counter()
D2 = np.sqrt(np.maximum(
    (A ** 2).sum(1)[:, None] - 2 * (A @ B.T) + (B ** 2).sum(1)[None, :], 0))
t2 = time.perf_counter() - t0

print("两版最大差 =", np.abs(D1 - D2).max())
print("广播版 %.3f s,中间量 %.0f MB" % (t1, mid / 1024 ** 2))
print("展开版 %.3f s,中间量 %.0f MB" % (t2, 3000 * 3000 * 8 / 1024 ** 2))
print("快 %.0f 倍" % (t1 / t2))

实跑输出(重跑两次都是 8 倍):

要点

两版最大差 = 3.4416913763379853e-15

广播版 1.661 s,中间量 1099 MB

展开版 0.219 s,中间量 69 MB

快 8 倍

⭐ 中间量从 1099 MB 降到 69 MB,顺带快了 8 倍。 快的原因不是「少算了」—— 两版算的是同一件事(最大差 3.4e-15,就是 float64 的舍入噪声), 而是 A @ B.T 走的是 BLAS 的矩阵乘法内核,比逐元素减法+平方+求和的访存量小一个量级。

⚠️ np.maximum(..., 0) 那一层不是装饰。 展开式在数值上可能算出 -1e-16 这样的负数 (本该是 0 的对角线附近),直接开根号会得到 nan。这是展开法唯一的代价,别省。

⭐ 什么时候该换写法:一个粗判据是估一下中间量。 N × M × D × 8 字节超过你内存的一半就别广播了。 第四节那张表给了三个锚点:N=400 是 2.4 MB,N=4000 是 244.1 MB,N=20000 是 6103.5 MB。 ⭐ 另一条路是分块:把 N 切成每块 1000 行,块内照样用广播 —— 中间量降到 1/N 块数,代码几乎不用改。


🔁 八、反过来看:unbroadcast 就是三条规则的镜像

现在回到第〇节那个证据。ml_md/18 让你实现广播的反向,它给的参考实现是这样的 (ml_md/18:140-152,逻辑一字未改(docstring 换了措辞)):

import numpy as np


def unbroadcast(grad, shape):
    """把广播后的梯度还原回原形状(和 ml_md/18 里那段是同一个)"""
    while grad.ndim > len(shape):     # ← 撤销规则一:前面补出来的轴,sum 掉
        grad = grad.sum(axis=0)
    for i, s in enumerate(shape):     # ← 撤销规则三:被拉伸的长度 1 轴,sum 回长度 1
        if s == 1:
            grad = grad.sum(axis=i, keepdims=True)
    return grad


W = np.zeros((32, 64))
bias = np.zeros((64,))
print("前向: (32,64) + (64,) ->", (W + bias).shape)

g = np.ones((32, 64))                        # 假装上游传下来的梯度全是 1
gW, gb = unbroadcast(g, W.shape), unbroadcast(g, bias.shape)
print("还给 W    的梯度:", gW.shape, " 前 3 个 =", gW[0, :3])
print("还给 bias 的梯度:", gb.shape, "  前 3 个 =", gb[:3], "<- 每个 bias 被 32 行共用")

col, row, g2 = np.zeros((3, 1)), np.zeros((1, 4)), np.ones((3, 4))
print("\n前向: (3,1) + (1,4) ->", (col + row).shape)
print("还给 (3,1):", unbroadcast(g2, col.shape).shape,
      " 值 =", unbroadcast(g2, col.shape).ravel(), "<- 每个被 4 列共用")
print("还给 (1,4):", unbroadcast(g2, row.shape).shape,
      " 值 =", unbroadcast(g2, row.shape).ravel(), "<- 每个被 3 行共用")

实跑输出:

关键信息

前向: (32,64) + (64,) -> (32, 64)
还给 W 的梯度: (32, 64) 前 3 个 = [1. 1. 1.]
还给 bias 的梯度: (64,) 前 3 个 = [32. 32. 32.] <- 每个 bias 被 32 行共用
前向: (3,1) + (1,4) -> (3, 4)
还给 (3,1): (3, 1) 值 = [4. 4. 4.] <- 每个被 4 列共用
还给 (1,4): (1, 4) 值 = [3. 3. 3. 3.] <- 每个被 3 行共用

⭐ 两行代码,两条规则,严格一一对应:

unbroadcast 那一行 撤销的是 为什么是 sum
while grad.ndim > len(shape): grad = grad.sum(axis=0) 规则 ①(左边补出来的轴) 补出来的轴上,原数组是同一份数据被复用了 N 次,梯度要加起来
if s == 1: grad = grad.sum(axis=i, keepdims=True) 规则 ③(长度 1 被拉伸) 同理,拉伸多少份就加多少份;keepdims=True 保住那根长度 1 的轴

⭐ 规则 ② 不需要撤销 —— 它只是个检查,不改变任何形状。

为什么梯度是 sum 不是 mean:广播是复用同一份数据, bias[0] 这一个数真的参与了 32 行的计算,32 条路径的梯度按链式法则要相加。 实跑出来正好是 32.,(3,1) 那个是 4.(被 4 列共用),(1,4) 是 3.(被 3 行共用)—— 这些数字本身就是「被复用了多少次」。

⭐ 这也解释了 ml_md/18:156 那句「不处理会形状报错或静默算错」: 少了第一个 while,梯度形状 (32,64) 塞不进 (64,) 的 bias.grad,会报错; 少了第二个 for,(3,1) 那种情况梯度形状碰巧还能加上去(+= 触发广播), 不报错但值错了 4 倍 —— 这就是「静默算错」。


🛑 读到这里可以停 —— 已经读了约 65 分钟。 最后一段还有(约 25 分钟):一页速记 · 检查点与走神救援 回来的时候不用重读,直接从下一节接着看就行。


📋 九、一页速记

想干的事 / 遇到的事 怎么办
每个特征加/减一个数((N,D) 和 (D,)) 直接写,右边的轴免费
每个样本加/减一个数((N,D) 和 (N,)) 写 v[:, None],或者归约时加 keepdims=True
不造数据就想知道结果形状 np.broadcast_shapes(s1, s2)
一次拿到所有两两组合 A[:, None, :] - B[None, :, :]
想显式看到拉伸后的样子 np.broadcast_to(v, shape)(⚠️ 只读)
could not be broadcast together 右对齐两个形状,找末位对不上的那一位
non-broadcastable output operand 看等号左边那个数组的形状,原地运算不许长大
结果形状莫名其妙变成方阵 十有八九是 (n,) 撞上了 (n,1),assert a.shape == b.shape
中间量太大爆内存 估 N×M×D×8;换展开式(‖a‖²−2ab+‖b‖²)或者分块
广播的反向传播 sum 回去:补出来的轴 sum(axis=0),长度 1 的轴 sum(axis=i, keepdims=True)

⚠️ 一条容易漏的:广播只管形状,不管 dtype。 int8 数组和 int8 数组广播完还是 int8,该溢出照样溢出 —— 那是 第 08 章 的题。


🔗 这一章连到哪里

相关的地方 为什么
02 · shape 和 axis 到底怎么数 ⭐ 那一章的 keepdims=True 和 v[:, None] 只说了「留一根长度 1 的轴」,这一章的规则 ③ 才是它们管用的理由
04 · 视图还是拷贝 ⭐ 本章第六节说 a += b 和 a = a + b 是两件事,那一章讲清「两件事」到底差在哪、以及为什么会波及别人
05 · 内存布局与 stride ⭐ 本章第三节只说「拉伸是免费的」,为什么免费(那根轴的步长被设成 0)在那里;两条原地运算的报错也归那一章
07 · 把循环改写成数组运算 那一章的 sliding_window_view 展开成 (199951, 50) 却 shares_memory=True,和本章第三节是同一个机制
08 · 整数 dtype 的真相 ⚠️ 广播只对齐形状,不改 dtype。两个 int8 广播完还是 int8,溢出照旧
⭐ ML 基础 18 · 手搓 mini-torch ⭐ 本章的立项证据,也是它的兑现处:那一章的 T2 要你实现「广播的反向」并给了 unbroadcast(),却假定你已经知道正向的三条规则
ML 基础 附录 C · 手撕代码速查 第 399 行 K-means 的 (X[:, None, :] - C[None, :, :]).sum(-1) 是第四节那个模式的现场,那里的注释只有「(n,k) 广播算距离」六个字
数学原理 05b · PCA 的推导 那里的 Xc = X - X.mean(axis=0) 是本章第一节「右边的轴免费」的最常见现场:(N,D) 减 (D,)
AI 基础设施 10 · 数据并行与 AllReduce ⚠️ 同名不同物:那里的「广播」是通信原语(一张卡发给所有卡),和本章毫无关系
密码学 14 · RSA ⚠️ 同名不同物:那里的「Håstad 广播攻击」是把同一条消息发给多个人

✅ 检查点

  1. 广播的三条规则分别是什么?哪一条是「只检查不改形状」的?
  2. (3,4) 和 (4,) 能算,(3,4) 和 (3,) 不能算。用规则 ① 解释为什么。
  3. (2,3,4) 和 (2,3) 能不能算?如果我确实想让它们前两维对齐,该怎么写?
  4. ⭐ (5,) 减 (5,1) 得到什么形状?为什么这是本章最贵的坑?实跑里 MSE 从多少变成了多少?
  5. np.broadcast_to(v, (1000000, 5)) 里 v 只有 5 个数,结果占多少内存?为什么它是只读的?
  6. a += b 和 a + b 在广播上有什么区别?看到 non-broadcastable output operand 该去查什么?
  7. X[:, None, :] - X[None, :, :] 得到什么形状?中间量多大?点数从 400 涨到 20000 时它变成多少?
  8. ⭐ 为什么广播的反向传播是 sum 而不是 mean?(32,64) + (64,) 的反向里,bias 的梯度实跑是多少?
  9. unbroadcast() 里那两行分别撤销的是哪条规则?为什么没有一行对应规则 ②?
  10. 距离矩阵爆内存时,展开成 ‖a‖²−2ab+‖b‖² 之后中间量从多少降到多少?为什么必须加 np.maximum(..., 0)?
👀 答案
  1. ① 补齐:维数少的在形状元组左边补 1;② 比对:逐位比,要么相等要么其中一个是 1,否则 ValueError;③ 拉伸:是 1 的那一位拉伸到对方长度,且不复制内存。规则 ② 只检查不改形状 —— 所以第八节的 unbroadcast() 里没有一行对应它。
  2. 因为规则 ① 只往左边补,永远不往右边补。(4,) 补成 (1,4),末位 4=4 对上了;(3,) 补成 (1,3),末位是 4 对 3,两个都不是 1,直接报 operands could not be broadcast together with shapes (3,4) (3,)。一句话:右边的轴是免费的,左边的轴要自己动手。
  3. 不能。 实跑报 Mismatch is between arg 0 with shape (2, 3, 4) and arg 1 with shape (2, 3) —— (2,3) 被补成 (1,2,3) 而不是 (2,3,1),末位 4 对 3。想让前两维对齐要自己写成 (2,3,1),也就是 y[..., None]。
  4. 得 (5,5)((5,)→(1,5),和 (5,1) 两位各有一个 1,两边都被拉伸)。💀 贵在它不报错:实跑里正确的 MSE 是 0.022,算出来是 3.822,大了 173 倍而程序一个字没吭,于是你会去查数据、查学习率,唯独不怀疑那行减法。⚠️ (n,1) 是 predict 返回、keepdims=True 结果、reshape(-1,1) 的天然形状,(n,) 是标签的天然形状,两者相遇就是一张 n×n 的表。防线按可靠性排:assert y_true.shape == y_pred.shape > 入口 .ravel() > np.broadcast_shapes 事前查。
  5. 还是 40 字节(v.nbytes = 40)。名义上 big.nbytes 是 40000000(38.1 MB),但 np.shares_memory(v, big) 是 True、big.base is v 是 True —— 它没有自己的数据。⚠️ 只读是因为 big[0,0] 和 big[999999,0] 是同一个字节,写一个数会让一百万行同时变,NumPy 干脆禁掉(writeable 是 False)。真要实体化 np.ascontiguousarray(big) 才掏出那 38.1 MB。
  6. a += b 要求广播结果的形状和 a 完全一样 —— a 是容器,不许原地长大,所以广播是单向的。实跑:small(4,) += big(3,4) 报 non-broadcastable output operand with shape (4,) doesn't match the broadcast shape (3,4),而 small + big 好好的给出 (3,4)。⚠️ 看到这条报错不要去查两个形状能不能广播(它们能),要去看等号左边那个数组的形状。
  7. (400,1,2) 和 (1,400,2) 相减得 (400,400,2),.sum(axis=-1) 之后是 (400,400)。中间量 = N² × D × 8 字节:N=400 是 2.4 MB,N=4000 是 244.1 MB,N=20000 是 6103.5 MB。实跑广播版比双重循环快几十倍(重跑三次 44x/57x/68x,⚠️ 这台机器有热降频,看量级别看单点),两版最大差是 0.0。
  8. 因为广播是复用同一份数据:bias[0] 这一个数真的参与了 32 行的计算,32 条路径的梯度按链式法则要相加。实跑 unbroadcast(ones((32,64)), (64,)) 出来正好是 [32. 32. 32. ...];(3,1) 那个是 [4. 4. 4.](被 4 列共用),(1,4) 是 [3. 3. 3. 3.](被 3 行共用)—— 这些数字本身就是「被复用了多少次」。
  9. while grad.ndim > len(shape): grad = grad.sum(axis=0) 撤销 规则 ①(左边补出来的轴);if s == 1: grad = grad.sum(axis=i, keepdims=True) 撤销 规则 ③(长度 1 被拉伸)。规则 ② 只是个检查、不改变形状,所以没有对应行。⚠️ 少了第一个 while 会形状报错;少了第二个 for,(3,1) 那种情况 += 会触发广播,不报错但值错 4 倍(ml_md/18:156 说的「静默算错」)。
  10. 中间量从 1099 MB 降到 69 MB,顺带快了 8 倍(重跑两次都是 8 倍)。快不是因为少算了 —— 两版最大差 3.4e-15,就是 float64 舍入噪声;是因为 A @ B.T 走 BLAS 矩阵乘法内核,访存量小一个量级。⚠️ np.maximum(..., 0) 必须加:展开式在数值上会算出 -1e-16 这样本该是 0 的负数,直接开根号得 nan。

🛑 可以停在这里

⚡ 走神救援

⭐ 广播只做三件事:① 补齐 —— 维数少的那个在形状元组左边补 1(永远不往右边补,这是后面所有坑的根源);② 比对 —— 逐位比,要么相等、要么其中一个是 1,否则 ValueError;③ 拉伸 —— 是 1 的那一位被拉到对方长度,且不复制内存。

记法:右边的轴是免费的,左边的轴要你自己动手。(3,4)+(4,) 直接能算;(3,4)+(3,) 报错,必须写成 c[:, None]。最反直觉的一条:(2,3,4)+(2,3) 不能算,因为 (2,3) 被补成 (1,2,3) 而不是 (2,3,1),要写 y[..., None]。调试形状用 np.broadcast_shapes(),它不分配内存。

💀 最贵的坑是不报错的那一类:(5,) 减 (5,1) 完全合规地给你 (5,5),实跑 MSE 从 0.022 变成 3.822(大 173 倍)。而 (n,1) 正是 predict 返回、keepdims=True 结果、reshape(-1,1) 的天然形状,(n,) 正是标签的形状。⭐ 防线按可靠性排:assert y_true.shape == y_pred.shape > 入口 .ravel() > 事前 broadcast_shapes。

最值钱的模式是 X[:, None, :] - X[None, :, :]:一次拿到所有两两距离,比双重循环快一个量级。⚠️ 代价是中间量 N²×D×8 —— N=400 是 2.4 MB,N=4000 就是 244.1 MB。解药是展开成 ‖a‖²−2ab+‖b‖²(走 BLAS),但必须套一层 np.maximum(..., 0),否则本该是 0 的位置算出 -1e-16、开根号得 nan。

⚠️ 原地运算的广播是单向的:a += b 要求结果形状和 a 一样,看到 non-broadcastable output operand 就去查等号左边那个数组。

下一节 👉 04-视图还是拷贝.md

打卡记录保存在你的浏览器里,首页能看到总进度