📑 本页目录(点开跳转)
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]))
实跑输出:
关键信息
⭐ 这段输出就是全章的缩影:
「每个特征加一个偏移」((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 / 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 完全一样。
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 行共用")
实跑输出:
关键信息
⭐ 两行代码,两条规则,严格一一对应:
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 广播攻击」是把同一条消息发给多个人 |
✅ 检查点
- 广播的三条规则分别是什么?哪一条是「只检查不改形状」的?
(3,4)和(4,)能算,(3,4)和(3,)不能算。用规则 ① 解释为什么。(2,3,4)和(2,3)能不能算?如果我确实想让它们前两维对齐,该怎么写?- ⭐
(5,)减(5,1)得到什么形状?为什么这是本章最贵的坑?实跑里 MSE 从多少变成了多少? np.broadcast_to(v, (1000000, 5))里v只有 5 个数,结果占多少内存?为什么它是只读的?a += b和a + b在广播上有什么区别?看到non-broadcastable output operand该去查什么?X[:, None, :] - X[None, :, :]得到什么形状?中间量多大?点数从 400 涨到 20000 时它变成多少?- ⭐ 为什么广播的反向传播是
sum而不是mean?(32,64) + (64,)的反向里,bias的梯度实跑是多少? unbroadcast()里那两行分别撤销的是哪条规则?为什么没有一行对应规则 ②?- 距离矩阵爆内存时,展开成
‖a‖²−2ab+‖b‖²之后中间量从多少降到多少?为什么必须加np.maximum(..., 0)?
👀 答案
- ① 补齐:维数少的在形状元组左边补 1;② 比对:逐位比,要么相等要么其中一个是 1,否则
ValueError;③ 拉伸:是 1 的那一位拉伸到对方长度,且不复制内存。规则 ② 只检查不改形状 —— 所以第八节的unbroadcast()里没有一行对应它。 - 因为规则 ① 只往左边补,永远不往右边补。
(4,)补成(1,4),末位 4=4 对上了;(3,)补成(1,3),末位是 4 对 3,两个都不是 1,直接报operands could not be broadcast together with shapes (3,4) (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]。 - 得
(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事前查。 - 还是 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。 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)。⚠️ 看到这条报错不要去查两个形状能不能广播(它们能),要去看等号左边那个数组的形状。(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。- 因为广播是复用同一份数据:
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 行共用)—— 这些数字本身就是「被复用了多少次」。 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说的「静默算错」)。- 中间量从 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