shape 只是折叠方式
上一章说 ndarray 是「一块连续内存 + 一张说明书」。这一章把说明书打开:里面只有三个字段,shape、strides、dtype。搞清楚这三个字段之后,一堆看起来神秘的行为会同时变得显然——为什么转置不花钱,为什么有的 reshape 免费有的要复制,以及为什么「零拷贝」这个词有时候是个陷阱。
一个 4000 × 4000 的 float64 矩阵,占 122 MB。执行 big.T(转置)。
问:这一行大约要多久?
说明书只有三个字段
一个 ndarray 在内存里其实是两样东西:一块裸字节(叫 buffer),和一个很小的头。头里记着:
| 字段 | 意思 | 例子 |
|---|---|---|
dtype | 每一格几个字节,怎么解释 | int64 → 8 字节,补码整数 |
shape | 把这条平的列,折成几维、每维多长 | (3, 4) |
strides | 沿每一维走一步,地址要挪多少字节 | (32, 8) |
strides = (32, 8) 的意思是:往下走一行,地址加 32 字节(也就是 4 个 int64);往右走一列,地址加 8 字节。就这样。a[i, j] 的地址是 基址 + i*32 + j*8,一个乘加就完事。
看懂这张图,转置就没有秘密了。转置只是把 shape 和 strides 同时倒过来:
a shape (3, 4) strides (32, 8) a.T shape (4, 3) strides (8, 32) 底下那块字节:一个都没动。 np.shares_memory(a, a.T) → True
所以那个问题的答案是 63 纳秒——它只是新建了一个头。作为对照,真把它复制一份(big.T.copy())要 35.6 毫秒。两者差 56 万倍。
reshape 什么时候免费,什么时候不免费
判据只有一条:新的形状,能不能用「一组 stride」描述这块现成的字节。能,就改个头;不能,就得复制。
a = np.arange(12).reshape(3, 4) np.shares_memory(a, a.reshape(2, 6)) # True —— 免费 np.shares_memory(a, a.T.reshape(12)) # False —— 复制了
第一个免费,因为 (3,4) 和 (2,6) 在这条平列上都是「从头到尾连着数」,只是断句不同。第二个不行:a.T 的元素顺序是 0 4 8 1 5 9 …,要把它摊成一条连续的十二格,只能真的搬一遍。
顺手说清楚一对经常被搞混的函数:
a.ravel() —— 尽量返回视图;能共享就共享 a.flatten() —— 永远复制 np.shares_memory(a, a.ravel()) True np.shares_memory(a, a.flatten()) False
还有一个参数值得知道,因为它是 numpy 与 MATLAB/Fortran/R 之间的分水岭:
a.ravel() [0 1 2 3 4 5 6 7 8 9 10 11] ← C order,按行摊 a.ravel(order='F') [0 4 8 1 5 9 2 6 10 3 7 11] ← Fortran order,按列摊
零拷贝不等于零成本:账只是被推后了
这是这一章真正想给你的那句话。转置本身免费,但转置出来的那个视图,它的 stride 是「跳着的」。只要后面有哪一步需要一块真正连续的内存,账就在那时候结。
本机拿一个 6000 × 6000(275 MB)的矩阵量了一遍。两个数组数值完全相同,只是一个按 C order 摆、一个按 Fortran order 摆:
| 做的事 | C order | F order | 倍数 |
|---|---|---|---|
.copy()(复制成连续内存) | 4.8 ms | 100.4 ms | 21.0× |
.sum(axis=1) 对 .sum(axis=0) | 4.9 ms | 5.3 ms | 1.07× |
两行数据讲的是两件不同的事,都值得记住:
- 复制差 21 倍。C order 的复制是一次直筒的 memcpy;F order 的复制要按 8 字节一格跳着读,每读一格就跨过 48 KB,缓存行几乎全废。同样多的字节,差一个数量级。
- 求和几乎不差。这一条可能出乎意料。「按列求和一定慢」是个流传很广的说法,但现代 numpy 的归约会分块(一次处理一小块的所有行列),把缓存友好性拿回来了。所以差别只有 1.07 倍。
那「跳着读慢」这条规律还成立吗?成立,只是要看它有没有被 numpy 藏起来。换个不给 numpy 优化机会的写法,它立刻现形:
取 1000 行(每行是连续的一段) 1.1 ms 取 1000 列(每列是跳着的一串) 17.9 ms ← 16.1 倍
所以真正稳的判据不是「按行快按列慢」,是「连着读快,跳着读慢」——而某一步到底连不连着,写在 strides 里,不写在你的直觉里。需要确认的时候有一个现成的开关:
a.flags['C_CONTIGUOUS'] # 这块内存是不是行方向连续 a.flags['F_CONTIGUOUS'] np.ascontiguousarray(a) # 强制变连续(该复制就复制),把账结清
stride 这个词一般翻成「步长」或「跨度」,但它有一个特别容易踩的细节:numpy 的 strides 单位是字节,不是元素个数。所以同一个 (3, 4) 的形状,int64 数组的 strides 是 (32, 8),int32 数组的是 (16, 4)。
另外,「C order」和「Fortran order」跟语言没关系,只是历史留下的名字:C 语言的多维数组按行连续存,Fortran 按列连续存。numpy 默认 C order,MATLAB、R、Julia 默认 Fortran order。跨这几个生态搬代码时,这是最容易静默出错的一处——数还是那些数,摆的顺序不一样。
说明书本身就三行,直接打印出来看:
import numpy as np a = np.arange(12).reshape(3, 4) print(a.shape, a.strides, a.dtype) # (3, 4) (32, 8) int64 print(a.T.shape, a.T.strides) # (4, 3) (8, 32) print(np.shares_memory(a, a.T)) # True —— 转置零拷贝 print(np.shares_memory(a, a.reshape(2, 6))) # True print(np.shares_memory(a, a.T.reshape(12))) # False —— 这次真复制了 print(a.ravel(), a.ravel(order='F'), sep='\n')
再量一遍「零拷贝不等于零成本」。这段要跑几秒钟,会吃掉约 550 MB 内存:
import time, numpy as np
n = 6000
c = np.ascontiguousarray(np.random.default_rng(0).random((n, n)))
f = np.asfortranarray(c)
assert np.array_equal(c, f) # 数值完全一样
def best(fn, r=3):
return min((lambda t: (fn(), time.perf_counter() - t)[1])(time.perf_counter())
for _ in range(r))
print("C 复制 %.1f ms" % (best(lambda: c.copy()) * 1000)) # 4.8
print("F 复制 %.1f ms" % (best(lambda: f.copy()) * 1000)) # 100.4
print("取 1000 行 %.1f ms" % (best(lambda: [c[i].copy() for i in range(1000)]) * 1000))
print("取 1000 列 %.1f ms" % (best(lambda: [c[:, i].copy() for i in range(1000)]) * 1000))
python3 -c "import numpy as np;a=np.arange(12).reshape(3,4);print(a.strides,a.T.strides,np.shares_memory(a,a.T))"
内存不够就把 n 改成 3000,倍数关系不变。
- 图像库里的
stride参数。Android 的Bitmap.copyPixelsToBuffer、OpenCV 的Mat::step、视频解码器输出的 YUV 帧,全都带一个 stride 字段,原因和这一章一模一样:一行像素的字节数常常大于宽度乘以每像素字节数(为了对齐到 16 或 64 字节)。把 stride 当成 width 用,画面就会斜着裂开——这是图形程序员的经典第一个 bug。 - PyTorch 的
.contiguous()。你在深度学习代码里见过的x.permute(0, 2, 1).contiguous(),就是这一章那句「把账结清」。permute免费(改 strides),但接下来的.view()需要连续内存,所以必须先复制一次。忘了写就报错,写多了就白白复制。 - 数据库的行存与列存。Parquet 文件的元数据里有一堆偏移量,作用就是 strides:告诉读取器「第 3 列从第几个字节开始」。第 21 章会把这条线接上。
- 为什么矩阵乘法库要区分
trans参数。BLAS 的dgemm有transa/transb两个开关,而不是让你先转置再传进来——因为库内部可以直接换个方向读,省掉一次 100 毫秒级的复制。
「转置是个重操作,能不转就不转。」
转置本身是全书最便宜的操作之一:63 纳秒,它只改了一个头。真正贵的是转置之后那些需要连续内存的步骤——本机上同一份数据,跳着读的复制比连着读的慢 21 倍。
这个错误还有一个更常见的反向版本:「反正 numpy 的切片和转置都是零拷贝,那我随便切随便转,反正不花钱。」不花钱的是那一行,不是整段代码。判断方法很机械:每写一步,问一句「这一步之后,.flags['C_CONTIGUOUS'] 还是 True 吗?」——如果不是,那笔账迟早有人付。
还有第三个变体,也是最贵的一个:凭直觉断言「按列比按行慢」。本机上求和只差 1.07 倍(numpy 分块优化过),逐列取却差 16.1 倍。这类问题不要推理,要量。
正确答案是 B:大约 63 纳秒。转置只是把 shape 和 strides 倒过来,一个字节都没动。
A 「122 MB 搬一趟」——这是big.T.copy() 的答案,35.6 毫秒。它和 B 相差 56 万倍,而两行代码在屏幕上只差五个字符。这个反差本身就值得记住:在 numpy 里,「看起来在搬数据」和「真的在搬数据」是两回事。
C 「看是不是 C order」——转置的耗时和 order 无关,因为它不读数据。但这个直觉用在下一步就是对的:转置之后的复制,C order 与 F order 差 21 倍。答对了半题,只是答在了错误的那一行上。
D 「缓存全部失效」——描述的是真实存在的现象,只是发生在别的地方。本机实测:求和几乎不受影响(1.07 倍),逐列取影响很大(16.1 倍)。缓存的账要看具体操作有没有被 numpy 分块优化过,不能一概而论。
shape 怎么折、strides 每步跳多远、dtype 每格几个字节。
没人检查。shape 不对不会报错,只会算出一个你没想要的结果——下一章的 axis 和第 4 章的广播,都是这句话的展开。要主动查:a.shape、a.strides、a.flags。
这一章的一句话
ndarray 底下永远是一条平的列,shape 和 strides 只是「怎么折」的说明书;所以改形状常常免费,而免费的那一步只是把账推给了下一步。
下一章处理 numpy 最常被问的那个问题:axis=0 到底是「按行」还是「按列」。答案是两个都不是。而且这个误解会长出一个特别阴险的 bug:在长方形矩阵上它会报错——你运气好;在方阵上它一声不吭,把「按行归一化」做成了「按列归一化」,四行的和分别是 0.8、2.0、0.85、1.25,没有一个是 1。