axis 不是行和列,是哪根轴被吃掉
这是 numpy 被问得最多的一个问题:a.sum(axis=0) 到底是「把每一行加起来」还是「把每一列加起来」?两个答案都是错的——准确说,两个都是记忆,不是理解,所以到了三维以上就全线崩溃。这一章给出唯一那个不会崩的读法,然后演示为什么这个误解在方阵上尤其危险。
一个 4 × 4 的矩阵 sq,你想做「行归一化」——让每一行的和变成 1。你写下:
sq / sq.sum(axis=1)
问:会发生什么?
唯一一个到 N 维都不会崩的读法
先把结论放这儿:
axis=k 的意思是:第 k 根轴被吃掉。
形状 (3, 4) 沿 axis=0 归约,第 0 根轴(长度 3)消失,剩下 (4,)。
沿 axis=1 归约,第 1 根轴(长度 4)消失,剩下 (3,)。
不用记「行」「列」,只需要把那个数字从 shape 里划掉。
验证一下这个读法:
m = [[ 1, 2, 3, 4],
[ 5, 6, 7, 8],
[ 9, 10, 11, 12]] shape (3, 4)
m.sum(axis=0) [15 18 21 24] shape (4,) ← 3 没了
m.sum(axis=1) [10 26 42] shape (3,) ← 4 没了
为什么「按行/按列」这套记忆会崩?因为它只在二维成立,而且中文和英文的「行列」对应关系本身就容易记反。到了三维就完全说不通了:
t = np.arange(24).reshape(2, 3, 4) shape (2, 3, 4) t.sum(axis=0) → (3, 4) 划掉 2 t.sum(axis=1) → (2, 4) 划掉 3 t.sum(axis=2) → (2, 3) 划掉 4 t.sum(axis=(0, 2)) → (3,) 划掉 2 和 4
这里没有任何一根轴叫「行」。而「划掉那个数字」这条规则一次都没失效。
顺带说一个推论:axis=-1 是「最后一根轴」,也就是变化最快、在内存里最连续的那一根。这就是为什么深度学习框架里的 softmax(x, dim=-1)、x.mean(-1) 到处都是——最后一根轴通常是「特征维」,而且它连续。
那为什么开头那行会静默出错
现在回到 sq / sq.sum(axis=1)。左边形状 (4, 4),右边形状 (4,)。numpy 要把它们凑成同一个形状(这是下一章的广播规则),凑法是从右往左对齐:
sq (4, 4)
sq.sum(1) (4,)
─────────────────────
从右往左对齐:
sq (4, 4)
sq.sum(1) (1, 4) ← 左边补一个 1
结果 (4, 4) ← 那个 1 被拉成 4
于是 sq.sum(axis=1) 这四个数,被当成了一整「行」,
横着铺满了整个矩阵 —— 除以的是⟦列方向⟧,不是行方向。
本机跑一遍。sq 的四行和分别是 4、10、4、5:
写错的版本 sq / sq.sum(axis=1) [[0.25 0.1 0.25 0.2 ] [0.25 0.2 0.75 0.8 ] [0. 0. 0.25 0.6 ] [1.25 0. 0. 0. ]] 各行的和 0.8 2.0 0.85 1.25 ★ 没有一个是 1 写对的版本 sq / sq.sum(axis=1, keepdims=True) [[0.25 0.25 0.25 0.25] [0.1 0.2 0.3 0.4 ] [0. 0. 0.25 0.75] [1. 0. 0. 0. ]] 各行的和 1.0 1.0 1.0 1.0
两段代码差 , keepdims=True 十四个字符。第一段不报错、不警告、返回一个形状完全正确的 4 × 4 矩阵——它只是把每一行除错了。如果这四行是四个用户的行为分布、四个类别的概率、四个通道的权重,你会在下游某个地方看到一个说不通的结果,然后回来找三个小时。
为什么方阵最危险
把同一个错误写在长方形矩阵上试试:
>>> m / m.sum(axis=1) # m 是 (3, 4) ValueError: operands could not be broadcast together with shapes (3,4) (3,)
它报错了——这是运气好。(3,) 没法从右往左对上 (3, 4) 的最后一维 4,广播失败,numpy 当场把你拦住。
而方阵上两个维度一样长,所以「对错了也能对上」。这条规律在整本书里会反复出现:形状越整齐,错误越安静。协方差矩阵、邻接矩阵、注意力矩阵、混淆矩阵——数据科学里到处是方阵,而它们恰好是这个 bug 最容易藏身的地方。
所以有一条便宜的自保动作:
# 任何一次归约之后,如果结果还要参与逐元素运算,就写 keepdims=True row_sum = sq.sum(axis=1, keepdims=True) # (4, 1),不是 (4,) sq / row_sum # 现在广播方向一定对 # 或者,把断言写进代码里 assert row_sum.shape == (sq.shape[0], 1)
keepdims=True 保留那根被吃掉的轴,把它压成长度 1。(4,) 变成 (4, 1)——一个竖着的形状,广播时就会往横向铺,正是你要的方向。
归约(reduction):把一根轴上的若干个数合并成一个数的运算。sum、mean、max、argmax、any、all、std 全是归约。它们共享同一套 axis 语义,所以这一章讲的规则对它们全部适用。
注意区分两个容易混的词:归约吃掉一根轴((3,4) → (4,)),广播(下一章)长出一根轴((4,1) + (1,4) → (4,4))。这本书里绝大多数形状事故,都是这两件事凑在一起、方向配反了。
还有一个中文里的坑:axis 通常译作「轴」,但也常被译作「维」。这两个词在 numpy 里指同一样东西(ndim 就是轴的个数),不必分辨。
这一章最值得亲手做的是那个「静默出错」。八行:
import numpy as np
sq = np.array([[1., 1., 1., 1.],
[1., 2., 3., 4.],
[0., 0., 1., 3.],
[5., 0., 0., 0.]])
wrong = sq / sq.sum(axis=1) # 没有报错
right = sq / sq.sum(axis=1, keepdims=True)
print("写错的,各行和:", wrong.sum(axis=1)) # [0.8 2. 0.85 1.25]
print("写对的,各行和:", right.sum(axis=1)) # [1. 1. 1. 1.]
m = np.arange(1, 13.).reshape(3, 4)
print(m / m.sum(axis=1, keepdims=True)) # 长方形上写对了也没事
m / m.sum(axis=1) # 长方形上写错会当场报错
再用一行把「划掉那个数字」这条规则验穿:
t = np.arange(24).reshape(2, 3, 4)
for ax in [0, 1, 2, (0, 2)]:
print(ax, "→", t.sum(axis=ax).shape)
# 0 → (3, 4) 1 → (2, 4) 2 → (2, 3) (0, 2) → (3,)
python3 -c "import numpy as np;t=np.arange(24).reshape(2,3,4);print([t.sum(axis=a).shape for a in (0,1,2)])"
在线跑:jupyter.org/try-jupyter。这一章不需要装任何东西以外的依赖。
- 每一个神经网络的 softmax。
softmax要求「每个样本的各类概率加起来是 1」,写法是exp(x) / exp(x).sum(axis=-1, keepdims=True)。那个keepdims=True不是习惯,是必须——少了它,在 batch size 恰好等于类别数的时候(比如 10 类、batch 10),代码不报错,模型静静地学不好。 - 批归一化和层归一化的区别,只是 axis 不同。BatchNorm 沿 batch 轴归约,LayerNorm 沿特征轴归约,代码几乎一样,改的就是
axis。理解「哪根轴被吃掉」,这两个概念一分钟就能分清。 - Excel 里的「按行汇总/按列汇总」。同一个概念的鼠标版本。Excel 帮你把方向画在了屏幕上,所以很少出错;numpy 把方向藏在一个数字里,所以经常出错。这是 GUI 和代码的一个普遍差异:代码更快,但把「你在对哪个方向操作」这件事变成了不可见的。
- SQL 里没有这个问题。
GROUP BY city明确写出了按什么分组,不存在「方向」的歧义。这是声明式接口的一个真实优势,也是第 9 章groupby比裸 numpy 好用的原因之一。
「axis=0 就是按行操作,axis=1 就是按列操作。记住这句话就够了。」
这句话不但在三维以上完全无效,在二维上还正好容易记反:sum(axis=0) 得到的是「每一列的和」,可它的名字里写着 0,而 0 常被联想成「行」。
换成「第 k 根轴被吃掉」就不会反:(3, 4) 沿 axis=0 归约,把 3 划掉,剩 (4,)——四个数,当然是每列一个。这条规则从一维到 N 维全部适用,也适用于 argmax、cumsum、np.stack、np.concatenate 的 axis 参数。
顺手的判据:写完任何一次归约,先在脑子里划掉那个数字,看剩下的形状是不是你要的;如果结果还要参与逐元素运算,就加 keepdims=True。这十四个字符的成本,远低于在方阵上找三小时。
正确答案是 C:不报错,做的是列归一化。四行的和是 0.8、2.0、0.85、1.25。
A 「正常工作」——如果 numpy 是「按行对齐」的,这就对了。但广播规则是从右往左对齐(下一章会把这条规则完整拆开),所以(4,) 被补成 (1, 4) 而不是 (4, 1),方向正好反了。
B 「报错」——在长方形矩阵上确实会报错,本机原文是 operands could not be broadcast together with shapes (3,4) (3,)。方阵不会。这正是这一章想让你记住的那件事:越整齐的形状,越不会替你把关。
D 「转回来就对了」——不对。它不是转置过的行归一化,而是用错误的除数做的除法:第 0 行的第 1 个元素被 4 除(第 0 行的和),第 1 个元素被 10 除(第 1 行的和)。除数是按列换的,被除数是按行摆的,两者根本没有对应关系,转多少次都救不回来。
keepdims 之后,那根轴被压成长度 1)。
只有形状对不上时 numpy 才拦;形状凑巧对得上的时候,它一声不吭。方阵是重灾区。自己检查:归约后打印 .shape,或者写一句 assert。
这一章的一句话
axis=k 不是「行」也不是「列」,是「第 k 根轴被吃掉」;而当归约的结果还要参与运算时,keepdims=True 是你唯一能主动指定方向的地方——否则方向由广播规则替你决定,它未必替你决定对。
下一章就去看那条替你做决定的规则。广播是 numpy 最省事的设计,也是它最危险的设计:一个 (10000, 1) 加一个 (10000,),你想要的是 76 KB 的一条列,拿到的是 0.75 GB 的一个方阵。它不会问你一句。