Skip to content

Commit 7ef887a

Browse files
committed
first draft
1 parent 9eec803 commit 7ef887a

43 files changed

Lines changed: 3799 additions & 706 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

README.md

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -85,6 +85,18 @@ The goal is not to crown a single winner, but to give an honest picture of **wha
8585

8686
-------
8787

88+
# Current local findings
89+
90+
Latest local refresh:
91+
92+
- `py312`: Python 3.12.2, NumPy 1.26.4, Numba 0.59.1, JAX 0.4.25 (CPU).
93+
- `py314t`: Python 3.14.0 free-threaded build, confirmed `sys._is_gil_enabled() == False`.
94+
- k-means at `N=1M`, `k=5`, `d=10`: Numba `0.482 s`, JAX `0.499 s`, NumPy naive `5.22 s`.
95+
- permutation test at `n=10k`, `R=10k`: Numba `0.064 s`, ThreadPool on `py314t` `0.173 s`, NumPy loop `0.856 s`, JAX CPU `37.4 s`.
96+
- See [`experiments/results/README.md`](experiments/results/README.md) for the refreshed figures and commands.
97+
98+
-------
99+
88100
# Notes
89101

90102
## What motivates us to submit this proposal

docs/00-statistician-simulation-process.md

Lines changed: 1466 additions & 0 deletions
Large diffs are not rendered by default.

docs/01-python314-features.md

Lines changed: 11 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,12 @@
1919

2020
### 1.3 实验中的表现(在 GIL 构建上的天花板)
2121

22-
本仓库的置换检验实验在标准 GIL 构建(Python 3.11.6, macOS 15.1, Apple Silicon 8 核心)上运行。即便在 GIL 下,`ThreadPoolExecutor` 仍然取得了 **~1.4× 加速**(n=10k, R=10000,详见 [`perm_scaling.png`](../experiments/results/v2/perm_scaling.png))。原因:NumPy 的 `.sum()``.permutation()` 在 C 层释放 GIL,使线程可以重叠。**这就是 GIL 构建上线程加速的上限。** 在 free-threaded 构建上,同一段 Python 代码应该继续向 8 核心靠拢——这是演讲最具说服力的「升级即收益」论点。
22+
本仓库的置换检验线程实验在同一台 Apple Silicon 8 核机器上比较了 `py312` GIL build 与 `py314t` free-threaded build(n=10k, R=10000,详见 [`perm_threads_py312_py314t.png`](../experiments/results/v2/perm_threads_py312_py314t.png))。同一份 `ThreadPoolExecutor` 代码在 8 workers 下:
23+
24+
- `py312` GIL build:warm median **0.32 s**
25+
- `py314t` free-threaded build:warm median **0.17 s**
26+
27+
标准 GIL 下线程仍然有加速,是因为 NumPy 的 `.sum()``.permutation()` 在 C 层释放 GIL,使线程可以重叠。Free-threaded build 继续降低 Python 层同步成本,把同一段代码的上限往 8 核心推进。
2328

2429
---
2530

@@ -38,7 +43,7 @@ Python 3.14 搭载的实验性 JIT 并非传统的追踪式 JIT(如 PyPy 的 T
3843

3944
- JIT 对于 **纯 Python 密集循环**、条件分支(控制流)有着显著加速。
4045
- **NumPy/C 扩展盲区**:如果在 CPython JIT 下运行高度 NumPy 向量化的代码,JIT **毫无作用**。因为运行时大部分时间在 `numpy` 的 C 库中。
41-
- 在我们的 k-means 实验中,`kmeans_loops.py`(N=2000, d=10, k=5)在标准 CPython 3.11 下单次完整跑需 **0.84 s**;对照 Numba 版本的 **0.010 s**(N=100k,80 倍大)。这一巨大差距正是 Python 循环解释开销的量化,也是 CPython JIT 值得去优化的地方——但即便 3.14 JIT 把纯循环加速 3×,距 Numba 也还有一个数量级**这是向观众传达的真实界限**:JIT 加速的是 Python 解释器层面的开销,而不是改变算法的渐进复杂度或 C 核心的执行速度。
46+
- 在我们的 k-means 实验中,`kmeans_loops.py`(N=2000, d=10, k=5)在标准 CPython 3.12 下 warm median **0.414 s**;对照 Numba 版本在 **N=100k** 时 warm median **0.0064 s**。这一巨大差距正是 Python 循环解释开销的量化,也是 CPython JIT 值得去优化的地方——但本地 `py314` 只暴露 `sys._jit``is_available()``False`,所以本轮只把 JIT 作为诚实限制呈现**这是向观众传达的真实界限**:JIT 加速的是 Python 解释器层面的开销,而不是改变算法的渐进复杂度或 C 核心的执行速度。
4247

4348
---
4449

@@ -66,8 +71,8 @@ Python 3.14 搭载的实验性 JIT 并非传统的追踪式 JIT(如 PyPy 的 T
6671

6772
在准备这篇演讲时,应该让观众明确:Python 3.14 的到来并不意味着我们可以盲目写纯 Python 循环。它降低了"偶尔写出慢代码"的惩罚,并为构建高性能的多线程 Python 库(如在单进程内共享巨大矩阵的统计工具)提供了基础设施。
6873

69-
**实测数据要点(Apple Silicon 8 核, Python 3.11, NumPy 1.24, Numba 0.57, JAX 0.4)**
74+
**实测数据要点(Apple Silicon 8 核, Python 3.12.2, NumPy 1.26.4, Numba 0.59.1, JAX 0.4.25;另测 Python 3.14t free-threaded**
7075

71-
- 即使在 GIL 上,`ThreadPoolExecutor` 在 NumPy-heavy 置换检验里仍能取得 1.4× 加速(见 [`perm_scaling.png`](../experiments/results/v2/perm_scaling.png))。无 GIL 构建应进一步扩大此差距
72-
- `multiprocessing` 在 R=10000 时的子进程 RSS 加总 **757 MB**(8 worker × 10k-length float64 array + 解释器基础内存),对照同规模线程池的 **~1.4 MB**[`perm_memory.png`](../experiments/results/v2/perm_memory.png) 把这一差距做成了一张直观的对数刻度条形图。
73-
- 在 CPU 上运行基于 `jax.vmap` 的置换检验(R=10000)耗时 **~75 s**——比纯 NumPy 慢 37×。JAX 在加速器上才能发挥,它在 CPU 上并非默认选项。
76+
- 同一份 ThreadPool 置换检验代码在 8 workers 下,`py312` GIL build 为 **0.32 s**`py314t` free-threaded build 为 **0.17 s**
77+
- `multiprocessing` 在 R=10000 时的子进程 RSS 加总 **833 MiB**,对照同规模线程池的 Python-level 峰值约 **1-2 MiB**[`perm_memory.png`](../experiments/results/v2/perm_memory.png) 把这一差距做成了一张直观的对数刻度条形图。
78+
- 在 CPU 上运行基于 `jax.vmap` 的置换检验(R=10000)耗时 **~37 s**——比纯 NumPy 慢约 44×。JAX 在加速器上才能发挥,它在 CPU 上并非默认选项。

docs/02-numba-guide.md

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66

77
- **Numba**:通过 LLVM 将装饰的函数编译为机器码,适合 **NumPy 数组 + 数值循环**
88
- **`@njit`**(即 `@jit(nopython=True)`):**nopython 模式**,避免回退到对象模式,性能最可预期。
9-
- **冷启动**:首次编译某函数通常有 **数百毫秒到数秒** 级延迟。例如本仓库 k-means Numba 核(d=10, K=5)首次编译 ≈ 0.9 s([`kmeans_sweep.csv`](../experiments/results/v2/kmeans_sweep.csv)`cold_s` 列)。必须靠 `warmup``cache=True` 抵消。
9+
- **冷启动**:首次编译某函数通常有 **数百毫秒到数秒** 级延迟。例如本仓库 k-means Numba 核(d=10, K=5)在当前本机 `py312` 下首次调用约 0.36 s([`kmeans_sweep.csv`](../experiments/results/v2/kmeans_sweep.csv)`cold_s` 列)。必须靠 `warmup``cache=True` 抵消。
1010

1111
## 2. `@njit`:单线程加速
1212

@@ -41,13 +41,13 @@ def parallel_kernel(...):
4141
- **只读大数组**:通常所有线程共享只读视图,避免在并行区写同一元素造成数据竞争。
4242
- Numba 有自己的线程池,即便 CPython 是 **GIL 构建**,它也能让外层 `prange` 并行执行——这是置换检验里它全场最快的原因。
4343

44-
## 4. 本仓库实测(Apple Silicon 8 核, Python 3.11.6
44+
## 4. 本仓库实测(Apple Silicon 8 核, Python 3.12.2
4545

4646
| 实验 | Numba 版 warm 中位数 | 对照项 | 加速比 |
4747
|------|---------------------|--------|--------|
48-
| k-means N=1M, max_iter=30 | **1.21 s** | NumPy 朴素广播 12.6 s | **10×** |
49-
| Permutation R=10000, n=10k | **0.158 s** | NumPy 朴素循环 2.04 s | **13×** |
50-
| Permutation R=10000, n=10k | **0.158 s** | `multiprocessing` (8 workers) 3.34 s | **21×** |
48+
| k-means N=1M, max_iter=30 | **0.482 s** | NumPy 朴素广播 5.22 s | **10.8×** |
49+
| Permutation R=10000, n=10k | **0.064 s** | NumPy 朴素循环 0.856 s | **13.4×** |
50+
| Permutation R=10000, n=10k | **0.064 s** | `multiprocessing` (8 workers) 1.71 s | **26.8×** |
5151

5252
[图表链接](../experiments/results/v2/perm_speedup.png) · [数据链接](../experiments/results/v2/perm_sweep.csv)
5353

docs/03-jax-guide.md

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414

1515
- 将 Python 可追踪函数编译为 XLA 程序。
1616
- 首次调用会 **追踪 + 编译**,耗时明显;基准需 **warmup** 并多次运行取中位数。
17-
- 本仓库 k-means JAX 核 (N=10k, d=10, K=5) 冷启动 ≈ 0.68 s,稳态 0.015 s(45× 差距)。在 N=1M 稳态下,编译成本相对于一次运行已微不足道。
17+
- 本仓库 k-means JAX 核 (N=10k, d=10, K=5) 在当前本机 `py312` 下冷启动约 0.21 s,稳态约 0.008 s(约 26× 差距)。在 N=1M 稳态下,编译成本相对于一次运行已微不足道。
1818

1919
### `jax.lax.scan`
2020

@@ -45,17 +45,17 @@
4545

4646
| R | JAX vmap (perm) | Numba prange | NumPy naive | JAX 对 NumPy 倍率 |
4747
|---|-----------------|--------------|-------------|-------------------|
48-
| 500 | 4.0 s | 0.018 s | 0.091 s | 44× **slower** |
49-
| 2 000 | 16.4 s | 0.054 s | 0.42 s | 39× slower |
50-
| 10 000 | 75.5 s | 0.158 s | 2.04 s | 37× slower |
48+
| 500 | 1.71 s | 0.0029 s | 0.044 s | 39× **slower** |
49+
| 2 000 | 7.44 s | 0.0155 s | 0.175 s | 42× slower |
50+
| 10 000 | 37.8 s | 0.064 s | 0.856 s | 44× slower |
5151

5252
(来源:[`perm_sweep.csv`](../experiments/results/v2/perm_sweep.csv)
5353

5454
我们还测试了「算法小聪明」版本——用 `jax.random.choice(replace=False)` 代替 `permutation`,期望少做一半的 shuffle 工作,结果 **几乎没差别**。因为 `choice(replace=False)` 在 JAX 内部也是基于 permutation 实现的。
5555

5656
**演讲要点***「如果你的代码只跑在 CPU 上,JAX 不是默认选项;`jax.vmap` 是为加速器批处理写的,不是为 CPU 外层并行写的。」*
5757

58-
反过来,k-means 这种 **BLAS-heavy 内核**(矩阵乘法在 XLA 上能良好下降到 BLAS),CPU JAX 已经可以在 N=1M 上跑到 1.35 s,与 Numba 只相差约 10%
58+
反过来,k-means 这种 **BLAS-heavy 内核**(矩阵乘法在 XLA 上能良好下降到 BLAS),CPU JAX 已经可以在 N=1M 上跑到约 0.50 s,与 Numba 的 0.48 s 非常接近
5959

6060
## 5. GPU 注意事项
6161

docs/04-kmeans-algorithms.md

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -34,11 +34,11 @@ C_sq = np.einsum("ij,ij->i", C, C)[None, :] # (1, K)
3434
dists_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

Comments
 (0)