卷 I · 一条列CH 03深度 3/23

axis 不是行和列,是哪根轴被吃掉

这是 numpy 被问得最多的一个问题:a.sum(axis=0) 到底是「把每一行加起来」还是「把每一列加起来」?两个答案都是错的——准确说,两个都是记忆,不是理解,所以到了三维以上就全线崩溃。这一章给出唯一那个不会崩的读法,然后演示为什么这个误解在方阵上尤其危险。

被吃掉的那根轴keepdims方阵上不报错

▷ 先猜一下

一个 4 × 4 的矩阵 sq,你想做「行归一化」——让每一行的和变成 1。你写下:

sq / sq.sum(axis=1)

问:会发生什么?

A 正常工作。每行和都变成 1.0 B 报错。形状对不上,numpy 会拦住你 C 不报错,但做的是列归一化。四行的和分别是 0.8、2.0、0.85、1.25 D 不报错,结果是转置过的行归一化,转回来就对了

唯一一个到 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):把一根轴上的若干个数合并成一个数的运算。summeanmaxargmaxanyallstd 全是归约。它们共享同一套 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 维全部适用,也适用于 argmaxcumsumnp.stacknp.concatenateaxis 参数。

顺手的判据:写完任何一次归约,先在脑子里划掉那个数字,看剩下的形状是不是你要的;如果结果还要参与逐元素运算,就加 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 的和)。除数是按列换的,被除数是按行摆的,两者根本没有对应关系,转多少次都救不回来。
⌗ 交接单
一个 N 维数组,加上你想沿哪根轴汇总。 少了一根轴的数组(或者加了 keepdims 之后,那根轴被压成长度 1)。 只有形状对不上时 numpy 才拦;形状凑巧对得上的时候,它一声不吭。方阵是重灾区。自己检查:归约后打印 .shape,或者写一句 assert

这一章的一句话

axis=k 不是「行」也不是「列」,是「第 k 根轴被吃掉」;而当归约的结果还要参与运算时,keepdims=True 是你唯一能主动指定方向的地方——否则方向由广播规则替你决定,它未必替你决定对。

下一章就去看那条替你做决定的规则。广播是 numpy 最省事的设计,也是它最危险的设计:一个 (10000, 1) 加一个 (10000,),你想要的是 76 KB 的一条列,拿到的是 0.75 GB 的一个方阵。它不会问你一句。