@@ -34,11 +34,11 @@ C_sq = np.einsum("ij,ij->i", C, C)[None, :] # (1, K)
3434dists_sq = X_sq + C_sq - 2.0 * (X @ C.T) # (N, K)
3535```
3636
37- ** 内存** :只需 (N, K) 矩阵,N=1M 时约 40 MB,** 比朴素广播小约 10 ×** 。
37+ ** 内存** :只需 (N, K) 矩阵,N=1M 时约 40 MB,当前本机 tracemalloc 峰值约 207 MiB, ** 比朴素广播小约 4 ×** 。
3838
3939** 【但是:实战性能取决于 K×d】**
40- - 在 K=5、d=10 这种「极瘦」配置下,矩阵乘法的收益会被两次 norm 计算吃掉,在我们的 M1 上甚至 ** 比朴素版慢 ~ 50% ** 。
41- - 如果改成 K=50、d=100(更接近真实生物统计模型的 EM 步骤),` X @ C.T ` 会显著快于广播。演讲中可以把这一条作为 ** 「加速窍门不是万能药」** 的例证:你得知道自己的 (N, K, d) 是否撑得起一次 BLAS 调用。
40+ - 在 K=5、d=10 这种「极瘦」配置下,矩阵乘法的收益很小;当前本机 N=1M 时 matmul 与 naive 接近(4.96 s vs 5.22 s),小 N 下还可能更慢 。
41+ - 如果改成 K=50、d=100(更接近真实生物统计模型的 EM 步骤),` X @ C.T ` 会显著快于广播。当前本机 N=5k 时 matmul 为 0.054 s,对照 naive 0.300 s,约 ** 5.5× ** 。 演讲中可以把这一条作为 ** 「加速窍门不是万能药」** 的例证:你得知道自己的 (N, K, d) 是否撑得起一次 BLAS 调用。
4242
4343## 3. 纯 Python 双层循环版本(用于 JIT 演示)
4444
@@ -52,23 +52,23 @@ for i in range(n):
5252```
5353
5454** 【专家深度点评:JIT 的用武之地】**
55- 这种循环在传统 CPython 中极慢,因为涉及数以亿计的字节码分发(Dispatch)和对象装箱/拆箱(Boxing/Unboxing)。在 N=2000, d=10, K=5 上一次完整运行需 ** 0.84 s** ([ ` kmeans_sweep.csv ` ] ( ../experiments/results/v2/kmeans_sweep.csv ) ),而 Numba 在 N=100k、更大 50× 的数据上只要 ** 0.010 s** ——两者相差 4 个数量级 。
55+ 这种循环在传统 CPython 中极慢,因为涉及数以亿计的字节码分发(Dispatch)和对象装箱/拆箱(Boxing/Unboxing)。在当前本机 N=2000, d=10, K=5 上一次完整运行需 ** 0.414 s** ([ ` kmeans_sweep.csv ` ] ( ../experiments/results/v2/kmeans_sweep.csv ) ),而 Numba 在 N=100k、更大 50× 的数据上只要 ** 0.006 s** 。
5656
57- Python 3.14 实验性 JIT(Copy-and-patch)理论上能把这 0.84 s 砍 2–3×,但 ** 永远达不到 Numba 的水准 ** 。这是演讲的重要边界 :JIT 是在既有解释器开销上打折,不是换算法或换运行时。
57+ Python 3.14 实验性 JIT(Copy-and-patch)理论上能给这种纯 Python 循环打折,但本地 ` py314 ` 的 ` sys._jit.is_available() ` 为 ` False ` ,所以本轮不伪造 JIT 数字。演讲里的边界应当讲清楚 :JIT 是在既有解释器开销上打折,不是换算法或换运行时。
5858
5959## 4. Numba 版本:最纯粹的高效
6060
6161- 将 ** 内层循环** 、距离与部分归约用 ` @njit ` 编译。
6262- ** 内存优势** :由于 Numba 能够在每次外层循环处理一个点时就立即计算最近距离并归约,它** 完全不需要生成 ` (N, K, d) ` 的中间张量** !它的空间复杂度从 NumPy 朴素的 \( O(NKd)\) 断崖式下降到 \( O(N + Kd)\) (只需 labels 与 centroids)。
63- - ** 实测** :N=1M 峰值分配 ** 22 MB ** ,对比 NumPy 朴素的 ** 810 MB ** — 38 × 的差距。
64- - ** 实测运行时** :N=1M 稳态 ** 1.21 s** ,对比 NumPy 朴素的 12.6 s — 10× 的速度差。
63+ - ** 实测** :N=1M 峰值分配 ** 22 MiB ** ,对比 NumPy 朴素的 ** 810 MiB ** — 36 × 的差距。
64+ - ** 实测运行时** :N=1M 稳态 ** 0.482 s** ,对比 NumPy 朴素的 5.22 s — 10.8 × 的速度差。
6565
6666## 5. JAX 版本:XLA 的威力
6767
6868- 整体 ` jax.jit ` ;使用 ` jax.lax.scan ` 固定 ` max_iter ` ;质心作为 ` carry ` 。
69- - ** 内存视角** :利用 ` tracemalloc ` 测试 JAX 的内存时,Python 解释器层面几乎没有任何峰值分配(N=1M 只有 2.2 MB )。这是因为 ` jax.jit ` 将计算完全下放到 XLA 运行时中,** 从 Python 层看不到** 。但 XLA 的设备内存仍然被占用——如果要报 GPU VRAM,需要用 JAX 自己的 profiler。
70- - ** 实测运行时** :N=1M 稳态 ** 1.35 s** ,与 Numba 差距仅约 12%。JAX 追上 Numba 的转折点恰好在 N=1M 附近——小 N 时 Numba 领先,大 N 时 BLAS-backed XLA 追上来 。
71- - ** 实测冷启动** :N=10k 的 ` cold_s = 0.68 s ` , warm 中位数 0.015 s。** 首次调用 45 × 慢于稳态** 。对「一次性脚本」这是实打实的代价。
69+ - ** 内存视角** :利用 ` tracemalloc ` 测试 JAX 的内存时,Python 解释器层面几乎没有任何峰值分配(N=1M 约 0.9 MiB )。这是因为 ` jax.jit ` 将计算完全下放到 XLA 运行时中,** 从 Python 层看不到** 。但 XLA 的设备内存仍然被占用——如果要报 GPU VRAM,需要用 JAX 自己的 profiler。
70+ - ** 实测运行时** :N=1M 稳态 ** 0.499 s** ,与 Numba 的 0.482 s 非常接近 。
71+ - ** 实测冷启动** :N=10k 的 ` cold_s = 0.21 s ` , warm 中位数 0.008 s。** 首次调用约 26 × 慢于稳态** 。对「一次性脚本」这是实打实的代价。
7272
7373## 6. 正确性检查
7474
0 commit comments