🏠 总目录📚 本教程 07 · 把循环改写成数组运算 ← →
📑 本页目录(点开跳转)

07 · 把循环改写成数组运算:六个套路

⏱ 66 分钟 | ⭐ 前六章讲的是「为什么」和「怎么读」,这一章是唯一一章讲「拿到一个循环,手怎么动」


🎯 一句话

大多数循环之所以能改写,是因为它其实在做六件事之一:累积、分组、取前 k 个、开窗、查区间、找相邻变化。 认出是哪一件,对应的函数就只有一个。认不出来的那些,本章最后一节告诉你别硬改。


🧩 一、动手之前,先问三个问题

上一章讲完索引,工具箱就齐了。但不是所有循环都该改写,先用这三个问题筛一遍:

问题 答「是」的话 答「否」的话
① 第 i 次迭代要用到第 i−1 次的结果吗? 只能用专门的累积函数(cumsum / cumprod / np.maximum.accumulate),没有对应函数就别改 随便改
② 改完的中间体有多大? 中间体比原数组大几十倍时,要么分块,要么放弃 随便改
③ 这个循环一次跑几个元素? 几十个:改了也没用,NumPy 每次调用有固定开销 上万个:值得改

⭐ 第 ① 条最容易判错。「用到上一次的结果」不等于「必须写循环」——前缀和明明是顺序依赖,np.cumsum 照样一行搞定。 正确的判据是:这个顺序依赖有没有现成的 ufunc accumulate。有就用,没有(比如「每步的结果决定下一步读哪个下标」)就老实写循环。

⚠️ 第 ③ 条不是玩笑。下面所有倍数都是在 N = 200,000 上测的;N = 100 时数组版经常更慢,因为你花在建临时数组上的时间比省下来的还多。


🧩 二、① 累积:cumsum / cumprod / accumulate

长什么样:循环里有一个「累加变量」,每轮更新它并记下来。

import numpy as np
import timeit

rng = np.random.default_rng(0)
N = 200_000
x = rng.random(N)

def loop_cumsum(x):                      # 循环版:一个累加变量 + 逐个写回
    out = np.empty(len(x))
    s = 0.0
    for i in range(len(x)):
        s += x[i]
        out[i] = s
    return out

print("最大差", np.abs(loop_cumsum(x) - np.cumsum(x)).max())   # ⭐ 先对正确性

t_loop = min(timeit.repeat(lambda: loop_cumsum(x), number=1, repeat=5))
t_arr  = min(timeit.repeat(lambda: np.cumsum(x),   number=1, repeat=5))
print(f"循环 {t_loop*1e3:.2f} ms | np.cumsum {t_arr*1e3:.3f} ms | {t_loop/t_arr:.0f}x")

要点

最大差 0.0

循环 25.59 ms | np.cumsum 1.103 ms | 23x

⭐ cumsum 只是一个特例。任何二元 ufunc 都有 .accumulate:

循环在做 写成
累加 np.cumsum(x) / np.add.accumulate(x)
累乘 np.cumprod(x)
「到目前为止的最大值」 np.maximum.accumulate(x) ⭐ 回撤计算就是这个
「到目前为止是不是一直为真」 np.logical_and.accumulate(m)

⚠️ cumsum 的误差会累积。它是顺序相加的,N 很大且量级差异大时,末尾几位和「先排序再加」不同。要求高时用 math.fsum(慢但精确),或者分块加。


🧩 三、② 分组求和:np.bincount

长什么样:循环里有个 dict,按某个 key 往里 +=。

import numpy as np
import timeit

rng = np.random.default_rng(0)
N, G = 200_000, 1000
g = rng.integers(0, G, N)          # 组号,取值 0..999
x = rng.random(N)

def loop_group(g, x):                    # 循环版:dict 里 +=
    d = {}
    for gi, xi in zip(g, x):
        d[int(gi)] = d.get(int(gi), 0.0) + float(xi)
    return d

def arr_group(g, x):                     # 数组版:一行
    return np.bincount(g, weights=x, minlength=G)

d, a = loop_group(g, x), arr_group(g, x)
print("最大差", max(abs(d[k] - a[k]) for k in d))

t_loop = min(timeit.repeat(lambda: loop_group(g, x), number=1, repeat=5))
t_arr  = min(timeit.repeat(lambda: arr_group(g, x),  number=1, repeat=5))
print(f"循环 {t_loop*1e3:.2f} ms | np.bincount {t_arr*1e3:.3f} ms | {t_loop/t_arr:.0f}x")

要点

最大差 0.0

循环 35.18 ms | np.bincount 0.283 ms | 124x

⭐ 这是六个套路里倍数最大的一个,因为循环版每轮都在做哈希、装箱、字典查找三件事,而 bincount 只是一趟顺序写内存。

三个必须知道的细节:

细节 说明
minlength 不写的话,输出长度只到「实际出现过的最大组号 + 1」。某一批数据里最后几组没出现,输出就短一截,后面广播直接报错
组号必须是非负整数 是字符串 / 稀疏 ID 的话,先 np.unique(keys, return_inverse=True) 映射成 0..G−1
不带 weights 就是计数 np.bincount(g) = 每组多少个。⭐ 组均值 = bincount(g, weights=x) / bincount(g)

⭐ 和上一章的关系:np.bincount 处理的正是 06 章那个「重复索引」问题。b[[1,1,1]] += 1 只加一次,np.add.at(b, [1,1,1], 1) 才加三次——而 bincount 是同一件事的专用快速通道(np.add.at 通用但慢)。


🧩 四、③ Top-K:np.argpartition

长什么样:np.argsort(-x)[:k]。⚠️ 它不是循环,但它是站内出现频率最高的一个「多做了很多功」的写法——为了拿前 10 个,把 20 万个全排了一遍。

import numpy as np
import timeit

rng = np.random.default_rng(0)
N, k = 200_000, 10
x = rng.random(N)

def topk_sort(x, k):                     # 全排序再切
    return np.argsort(-x)[:k]

def topk_part(x, k):                     # 只保证「前 k 个在前面」,不保证它们有序
    idx = np.argpartition(-x, k)[:k]
    return idx[np.argsort(-x[idx])]      # ⭐ 需要有序的话,再对这 k 个排一次

print("集合一致", set(topk_sort(x, k).tolist()) == set(topk_part(x, k).tolist()))

t_sort = min(timeit.repeat(lambda: topk_sort(x, k), number=1, repeat=5))
t_part = min(timeit.repeat(lambda: topk_part(x, k), number=1, repeat=5))
print(f"argsort {t_sort*1e3:.2f} ms | argpartition {t_part*1e3:.3f} ms | {t_sort/t_part:.1f}x")

要点

集合一致 True

argsort 9.50 ms | argpartition 2.148 ms | 4.4x

⭐ argpartition 返回的前 k 个是无序的,这是它最常被写错的地方。它只保证「第 k 位上的元素是正确的第 k 名,左边都不比它小」。要有序就像上面那样再对这 k 个排一次——k 很小时这一步几乎不要钱。


🧩 五、④ 滑窗:sliding_window_view

长什么样:for i in range(N-w+1): out[i] = f(x[i:i+w])。

import numpy as np
import timeit
from numpy.lib.stride_tricks import sliding_window_view

rng = np.random.default_rng(0)
N, w = 200_000, 50
x = rng.random(N)

def loop_window(x, w):                   # 循环版:切一片算一次均值
    return np.array([x[i:i+w].mean() for i in range(len(x) - w + 1)])

def arr_window(x, w):                    # 数组版:一次性铺开成 (N-w+1, w)
    return sliding_window_view(x, w).mean(-1)

print("最大差", np.abs(loop_window(x, w) - arr_window(x, w)).max())

t_loop = min(timeit.repeat(lambda: loop_window(x, w), number=1, repeat=5))
t_arr  = min(timeit.repeat(lambda: arr_window(x, w),  number=1, repeat=5))
print(f"循环 {t_loop*1e3:.2f} ms | sliding_window_view {t_arr*1e3:.3f} ms | {t_loop/t_arr:.0f}x")

v = sliding_window_view(x, w)
print("展开后形状", v.shape, "| 与原数组共享内存", np.shares_memory(v, x))
print("如果真复制要占", round(v.nbytes/1024**2, 1), "MiB;实际 x 只占",
      round(x.nbytes/1024**2, 1), "MiB")

要点

最大差 0.0

循环 850.47 ms | sliding_window_view 5.950 ms | 143x

展开后形状 (199951, 50) | 与原数组共享内存 True

如果真复制要占 76.3 MiB;实际 x 只占 1.5 MiB

⭐ 最后两行是这一节真正的重点:(199951, 50) 这个形状看着像把数据复制了 50 份,但 shares_memory 是 True ——它没有复制任何东西,只是给同一块内存换了一套读法。76.3 MiB 是「如果真复制」的账,实际开销是 0。

为什么能这样,是 05 · 内存布局与 stride 那一章的正题;这里只需要知道它免费。

⚠️ 两个坑:


🛑 读到这里可以停 —— 前半章讲完了(约 26 分钟)。 后半章还有:⑤ 分桶:np.searchsorted · ⑥ 找相邻变化:diff + flatnonzero · 一张表:六个套路的倍数 · ⭐ 倍数是规模的函数,不是常数 · 什么时候不该改写 回来的时候不用重读,直接从下一节接着看就行。


🧩 六、⑤ 分桶:np.searchsorted

长什么样:循环里对每个元素做二分查找,或者一串 if / elif 判断落在哪个区间。

import numpy as np
import timeit
from bisect import bisect_right

rng = np.random.default_rng(0)
N = 200_000
edges = np.array([0.0, 0.1, 0.25, 0.5, 0.75, 0.9, 1.0])   # ⭐ 必须已排序
x = rng.random(N)

def loop_bucket(x, edges):               # 循环版:一个一个 bisect
    e = edges.tolist()
    return np.array([bisect_right(e, xi) for xi in x])

def arr_bucket(x, edges):                # 数组版:一次全查完
    return np.searchsorted(edges, x, side="right")

print("完全一致", np.array_equal(loop_bucket(x, edges), arr_bucket(x, edges)))

t_loop = min(timeit.repeat(lambda: loop_bucket(x, edges), number=1, repeat=5))
t_arr  = min(timeit.repeat(lambda: arr_bucket(x, edges),  number=1, repeat=5))
print(f"循环 {t_loop*1e3:.2f} ms | np.searchsorted {t_arr*1e3:.3f} ms | {t_loop/t_arr:.0f}x")

要点

完全一致 True

循环 61.14 ms | np.searchsorted 6.436 ms | 9x

⚠️ side 那个参数决定边界归左还是归右:值正好等于某个边界时,"left" 把它算进左边那个桶(和 0.49 同桶),"right" 算进右边(和 0.51 同桶)。等值样本多的时候(比如大量 0),这一个字母会改变一整个桶的大小。

⭐ searchsorted 不止能分桶。它的本质是「在一个已排序数组里,问这些值该插在哪」,所以还能做:把一批 ID 映射到它们在排序表中的位置、两个有序数组求交、找最近邻的那个刻度。


🧩 七、⑥ 找相邻变化:diff + flatnonzero

长什么样:循环比较 x[i] 和 x[i+1],记下不一样的位置。游程编码、找变号点、找状态切换点都是这一类。

import numpy as np
import timeit

rng = np.random.default_rng(0)
N = 200_000
y = rng.normal(size=N)                   # 有正有负的序列

def loop_sign(y):                        # 循环版:比较相邻两个
    out = []
    for i in range(len(y) - 1):
        if (y[i] < 0) != (y[i+1] < 0):
            out.append(i)
    return np.array(out)

def arr_sign(y):                         # 数组版:signbit -> diff -> 非零位置
    return np.flatnonzero(np.diff(np.signbit(y)))

print("完全一致", np.array_equal(loop_sign(y), arr_sign(y)), "| 变号次数", len(arr_sign(y)))

t_loop = min(timeit.repeat(lambda: loop_sign(y), number=1, repeat=5))
t_arr  = min(timeit.repeat(lambda: arr_sign(y),  number=1, repeat=5))
print(f"循环 {t_loop*1e3:.2f} ms | flatnonzero+diff {t_arr*1e3:.3f} ms | {t_loop/t_arr:.0f}x")

要点

完全一致 True | 变号次数 100527

循环 296.44 ms | flatnonzero+diff 0.706 ms | 420x

⭐ 这个套路的心法是「把状态变成布尔数组,再看它哪里跳变」:

想找 写法
变号位置 np.flatnonzero(np.diff(np.signbit(y)))
值发生变化的位置(游程边界) np.flatnonzero(np.diff(a) != 0)
每段游程的长度 在上面的边界数组两头补 -1 和 len(a)-1,再 np.diff
布尔序列里连续 True 的段 np.diff(m.astype(np.int8)),+1 是起点、-1 是终点

⚠️ flatnonzero(d) 给的是 d 的下标,不是原数组的下标。d = np.diff(y) 比 y 短 1,d[i] 描述的是 y[i] 和 y[i+1] 之间——所以「变化发生在 y 的第 i+1 个」。差这个 1 是本套路唯一的高频错误。


🧩 八、一张表:六个套路的倍数

⚠️ 看倍数,不要看绝对耗时。 下面是同一台笔记本(Windows 11 · Python 3.13.14 · NumPy 2.4.6 · 无 GPU 参与)上 N = 200,000、timeit.repeat(number=1, repeat=5) 取 min、三次独立进程跑出来的区间。这台机器会热降频,所以给区间不给单点。

套路 循环版 数组版 循环 ms 数组 ms 倍数(三次)
① 累积 累加变量 + 逐个写回 np.cumsum 23.6 – 39.9 1.10 – 1.37 20 – 29x
② 分组求和 dict 里 += np.bincount(g, weights=x) 34.7 – 43.9 0.28 – 0.32 120 – 137x
③ Top-K(k=10) np.argsort(-x)[:10] np.argpartition(-x,10)[:10] 7.0 – 14.3 1.98 – 2.50 3.5 – 5.7x
④ 滑窗均值(w=50) for + 切片 .mean() sliding_window_view(x,50).mean(-1) 683 – 940 5.5 – 7.7 89 – 170x
⑤ 分桶 for + bisect np.searchsorted 39.0 – 61.1 4.85 – 6.44 8 – 10x
⑥ 找变号位置 for + 比较相邻 np.flatnonzero(np.diff(np.signbit(y))) 296 – 306 0.68 – 0.93 320 – 450x

⭐ 你的绝对耗时一定和上表不同,看量级和倍数就行。 同一台机器上三次跑,④ 的倍数就在 89x 和 170x 之间摆——这不是测量出错,是笔记本降频。真要给出可比的数字,得用固定频率的机器多轮取分位数,那套口径归《Python》板块的性能剖析章。

⚠️⚠️ ③ 只有 3.5–5.7x,这个数没有粉饰。 六个套路里它最不划算,原因见下一节——它恰好是最有教学价值的一个。


🧩 九、⭐ 倍数是规模的函数,不是常数

上表所有数字都只对 N = 200,000 成立。把 Top-K 这一条放到不同规模上看:

import numpy as np
import timeit

k = 10
for N in (200_000, 5_000_000):
    x = np.random.default_rng(0).random(N)
    t_sort = min(timeit.repeat(lambda: np.argsort(-x)[:k],         number=1, repeat=5))
    t_part = min(timeit.repeat(lambda: np.argpartition(-x, k)[:k], number=1, repeat=5))
    print(f"N = {N:>9,} : argsort {t_sort*1e3:7.2f} ms | "
          f"argpartition {t_part*1e3:6.2f} ms -> {t_sort/t_part:.1f}x")

算一算

N = 200,000 : argsort 8.91 ms | argpartition 2.66 ms -> 3.4x

N = 5,000,000 : argsort 603.00 ms | argpartition 56.09 ms -> 10.8x

(三次独立跑:N=200,000 时 3.4x / 3.2x / 2.8x;N=5,000,000 时 10.8x / 8.6x / 14.7x。倍数抖,但趋势不抖。)

⭐ 原因是复杂度不同:argsort 是 O(n log n),argpartition 是 O(n)。它们的比值本身就含一个 log n,所以规模涨 25 倍,优势从 3 倍涨到 9–15 倍。

💀 这条比六个漂亮的倍数都更值钱:任何「X 比 Y 快 N 倍」的说法,不写清楚 N 是多少就没有意义。你在 20 万行上量出的 3 倍,上线跑 500 万行时是 10 倍;反过来,你在 20 万行上量出的 100 倍,缩到 100 行时可能是 0.5 倍(更慢)。


🚦 十、什么时候不该改写

改写不是免费的。这四种情况请直接停手:

情况 为什么 怎么办
真正的顺序依赖 第 i 步要根据第 i−1 步的结果决定读哪个下标(比如并查集、A* 的开放列表),没有 ufunc 能表达 写循环。真嫌慢就换语言层,见 numba / C 扩展
中间体爆炸 03 章那个 (400,400,2) 的距离中间体是 2.4 MB,放大到 (3000,3000,16) 就是 1.1 GB —— N 翻 10 倍它翻 100 倍 分块(chunk)循环 + 块内向量化。⭐ 外层留一个循环是完全正常的写法
n 很小 每次 NumPy 调用有几微秒的固定开销,n=50 时纯 Python 反而赢 别改
可读性崩了 一行 np.flatnonzero(np.diff(np.signbit(np.where(...)))) 三个月后没人看得懂 ⭐ 拆成两三行并留注释,倍数几乎不变

⭐ 一条经验:先把循环改对,再测。本章每段代码都是「先 print 最大差 / 是否完全一致,再打印耗时」——顺序不能反。改快了但改错了,是最难查的一类 bug,因为它不报错。


🔗 这一章连到哪里

相关的地方 为什么
06 · 索引的四种形态 ② 用的 bincount 是那一章「重复索引只写一次」的专用快速通道。先看懂 np.add.at 为什么存在,再回来用 bincount,否则你不知道自己躲开了什么坑
05 · 内存布局与 stride ④ 的 sliding_window_view 铺出 (199951, 50) 却不复制数据,为什么能这样,答案全在那一章的 stride 三元组。这里只用结论
03 · 广播的三条规则 「中间体爆炸」那一条的具体账(2.4 MB → 1.1 GB)在那一章算过。⭐ 决定要不要分块之前,先去那里学会估中间体大小
01 · 为什么写循环是错的 那一章讲为什么数组版更快(解释器开销、连续内存、SIMD),这一章只讲手怎么动。倍数看不懂就回去补那三层
ML 基础 · 附录 C · 手撕代码速查 那里的 conv2d 用 Ho×Wo 双重 Python 循环写前向——⭐ 那是故意的,白板题要的是「拆开给你看」。但你在自己项目里写卷积时,套路 ④ 是正确答案
Kaggle 竞赛方法论 · 01 · 数据策略与特征工程 套路 ② 和 ⑤(分组统计、分桶)就是特征工程里最高频的两个动作。那边讲造什么特征,这里讲怎么算得快
08 · 整数 dtype 的真相 ⚠️ 套路 ② 的组号是整数数组。组号很多时用 int16 存看起来省内存,但它会静默回绕——下一章的正题

✅ 检查点

  1. 判断一个循环该不该改写,本章给了哪三个问题?
  2. 「第 i 次要用第 i−1 次的结果」是不是就一定不能改写?举一个反例。
  3. np.bincount 不写 minlength 会出什么问题?
  4. np.argpartition(-x, k)[:k] 拿到的前 k 个是有序的吗?要有序该怎么补?
  5. sliding_window_view(x, 50) 在 N=200,000 上展开成 (199951, 50),这一步占多少内存?为什么?
  6. 本章六个套路里,倍数最小的是哪个?大概多少?为什么它最小?
  7. Top-K 的倍数从 N=200,000 到 N=5,000,000 是怎么变的?原因是什么?
  8. np.flatnonzero(np.diff(y) != 0) 返回的下标,对应原数组 y 的第几个元素?
  9. 本章列了四种「不该改写」的情况,哪一种明确说了「外层留一个循环是正常的」?
👀 答案
  1. ① 有没有顺序依赖 ② 中间体有多大 ③ 一次跑几个元素(几十个就别改,本章倍数全部在 N = 200,000 上测)。
  2. 不一定。前缀和就是顺序依赖,但 np.cumsum 一行搞定。真正的判据是「这个顺序依赖有没有现成的 ufunc accumulate」——np.add.accumulate / np.maximum.accumulate / np.logical_and.accumulate 都算。没有对应函数(比如每步结果决定下一步读哪个下标)才必须写循环。
  3. 输出长度只到「实际出现过的最大组号 + 1」。某批数据里最后几组恰好没出现,输出就短一截,后面广播直接报错。
  4. 无序。argpartition 只保证第 k 位是正确的第 k 名、左边都不比它小。要有序就 idx[np.argsort(-x[idx])] 再对这 k 个排一次,k 小时几乎不花钱。
  5. 不占(0 额外内存)。实测 np.shares_memory(v, x) 是 True ——它没复制,只是给同一块内存换了套 stride 读法。「如果真复制」是 76.3 MiB,而 x 本身只有 1.5 MiB。为什么能这样见 05 章。
  6. ③ Top-K,只有 3.5–5.7x。因为 argsort 和 argpartition 都是 C 实现的数组操作,两边都没有 Python 循环开销,差的只是 O(n log n) 和 O(n) 这个算法复杂度;其余五个是「Python 循环 vs C 循环」,差的是一整层解释器开销。
  7. 从 3.4x 涨到 10.8x(三次跑:200,000 时 3.4/3.2/2.8x,5,000,000 时 10.8/8.6/14.7x)。因为两者比值里含一个 log n,规模涨 25 倍,优势就涨了一档。结论:倍数是规模的函数,不写清 N 的倍数没有意义。
  8. 对应 y 的第 i+1 个。np.diff(y) 比 y 短 1,d[i] 描述的是 y[i] 和 y[i+1] 之间的关系。差这个 1 是本套路唯一的高频错误。
  9. 中间体爆炸那一条:分块循环 + 块内向量化,外层留一个循环是完全正常的写法。

🛑 可以停在这里

⚡ 走神救援

这一章只做一件事:给「拿到一个循环,手怎么动」列一张对照表。 动手前先过三个问题——有没有顺序依赖、中间体多大、一次跑几个元素(只有几十个就别改)。

⭐ 第一个问题最容易判错:「要用上一次的结果」不等于「必须写循环」——前缀和就是顺序依赖,而一个 cumsum 就够。真正的判据是这个依赖有没有现成的 ufunc accumulate。

六个套路里最值得记的四条:分组求和用 bincount 最猛(循环版每轮都在做哈希、装箱、字典查找),⚠️ 必须写 minlength,否则某批数据末尾几组没出现,输出就短一截;Top-K 用 argpartition 而不是全排序,⚠️ 结果无序;滑窗用 sliding_window_view——⭐ 它展开成一个大得多的形状却和原数组共享内存,「如果真复制」会是几十倍的内存;分桶用 searchsorted,⚠️ side 决定等值归左还是归右。

⭐⭐ 全章最该带走的一条不是任何一个函数,是「倍数是规模的函数」:同一个改写在小规模上只快几倍、在大规模上快十几倍,因为比值里含一个 log。所以「X 比 Y 快 N 倍」不写清 N 就没有意义。 ⚠️ 同理,同一台笔记本三次独立跑,同一个倍数能差近一倍——那是热降频不是测错,看量级不看绝对值。

最后四种情况直接停手不要改:真正的顺序依赖、中间体爆炸(⭐ 这时外层留一个循环是正常写法)、n 很小(每次调用有固定开销)、可读性崩了。

⭐ 顺序永远是先改对再测——本章每段代码都先打印「最大差 / 是否完全一致」再打印耗时。改快了但改错了不报错,是最难查的一类 bug。

下一节 👉 08-整数dtype的真相.md

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