6倍更快的二分查找:从编译代码到机械共鸣
摘要
本文详细介绍了Rust中二分查找的一系列底层优化,通过利用分支预测和SIMD等CPU架构特性,实现了6倍的加速,并将其应用于scikit-learn梯度提升用例中。
<p><a href="https://lobste.rs/s/czbhmr/6x_faster_binary_search_from_compiled">评论</a></p>
查看缓存全文
缓存时间: 2026/07/14 12:16
# 6倍更快的二分搜索:从编译代码到机械同理心
来源:https://pythonspeed.com/articles/branchless-binary-search
如何加速计算密集型 Python 代码?一个常见且有用的起点是:
1. 选择一个好的算法。
2. 用编译语言编写 Python 扩展。
3. 或许加上并行,以便使用多个 CPU 核心。
但如果你需要更快的速度呢?考虑以下实际场景——scikit-learn 梯度直方图提升算法中的一个步骤:
- 你有一个大的浮点数数组。
- 你想将它们均匀分配到 0-254 的整数范围内。
scikit-learn 的实现方法是将浮点值的整个范围分成 255 个桶,创建一个有序的桶边界数组,然后对每个值使用二分搜索来选择合适的桶。二分搜索是用编译语言实现的,并且可以并行在多个核心上运行。
最近,作为我在 [Quansight](https://labs.quansight.org/) 工作的一部分,并受 [Paul Khuong 的两篇文章](https://pvk.ca/Blog/2012/07/03/binary-search-star-eliminates-star-branch-mispredictions/) [启发](https://pvk.ca/Blog/2015/11/29/retrospective-on-binary-search-and-on-compression-slash-compilation/),我显著加速了这一实现。如何做到的?通过确保代码不“对抗”CPU。在本文中,我将通过一个简化的例子带你了解这种加速过程。然后,我将演示一系列额外的优化,最终版本比原始版本快 6 倍。
值得注意的是,我会快速掠过许多底层硬件相关的话题:指令级并行、分支(误)预测、内存缓存、SIMD 等等。这只是一篇文章,只能简要介绍可能性,不能作为深入教程。因此,我将在文章末尾讨论如何进一步学习这些主题。
## 起点:标准二分搜索
原始的 scikit-learn 代码是用 Cython 实现的,但本文我将使用 Rust。以下是一个相当标准的二分搜索实现(基于 NumPy 中的版本),专为给定边界数组查找桶而设计:
```rust
use std::cmp::Ordering;
/// Rust 不允许用普通的 < 运算符比较浮点数
/// (因为 NaN 会导致比较结果不一致),
/// 所以实现一个自定义函数来完成。
fn less_than(a: f64, b: &f64) -> bool {
a.total_cmp(b) == Ordering::Less
}
/// 将浮点数值转换为整数值,
/// 根据桶边界确定它们属于哪个桶。
fn bucketize_classic_impl(
arr: &[f64],
boundaries: &[f64],
) -> Vec<usize> {
// Vec 或 vector 是 Rust 中 Python 列表的等价物。
// 这里我创建一个空的 Vec,分配足够存储 `arr.len()` 个值的内存:
let mut result = Vec::with_capacity(arr.len());
for value in arr {
// 标准二分搜索算法:
let mut min_idx = 0;
let mut max_idx = boundaries.len();
while min_idx < max_idx {
let middle = min_idx + ((max_idx - min_idx) / 2);
if less_than(boundaries[middle], value) {
min_idx = middle + 1;
} else {
max_idx = middle;
}
}
// 这等同于 Python 中的 a_list.append():
result.push(min_idx);
}
// 返回结果:
result
}
```
为完整起见,下面是如何将其连接到 Python,使其接收并返回 NumPy 数组;这是样板代码,所以对于后面函数我不会再展示。如果熟悉 Rust 和 PyO3 可以跳过,或者不关心也没关系;它和文章其余部分无关。
点击查看代码
```rust
use pyo3::prelude::*;
use numpy::ndarray::Array;
use numpy::{PyArray1, PyReadonlyArray1};
#[pyfunction]
fn bucketize_classic<'py>(
py: Python<'py>,
// 这两个参数是一维浮点数数组:
arr: PyReadonlyArray1<f64>,
boundaries: PyReadonlyArray1<f64>,
) -> PyResult<&'py PyArray1<usize>> {
let result = bucketize_classic_impl(
arr.as_slice().unwrap(),
boundaries.as_slice().unwrap(),
);
// 将 Rust vector 转换为一维 NumPy 数组,以便返回给 Python:
let result = PyArray1::from_owned_array(py, Array::from_vec(result));
Ok(result)
}
```
## 分支误预测会拖慢你的代码
如何加速这个二分搜索实现?它已经使用了可扩展的算法和编译语言。并行化当然是一个选项,但我会采用不同的方法:机械同理心——更好地理解 CPU 的工作方式。
我将从快速回顾现代 CPU 如何在单核内部并行运行代码开始。对于 Python 代码,一个合理的心理模型是代码一次执行一条指令。算术操作翻倍,代码运行速度就会减半。一旦切换到编译语言,操作有时映射为一两条 CPU 指令,这个心理模型就不再正确。现代 CPU 有时可以在单个 CPU 核心上同时运行多条独立的 CPU 指令(“指令级并行”),从而加快执行速度。
```rust
fn two_adds(a: i64, b: i64, c: i64, d: i64) -> i64 {
// 你的 CPU 很可能可以自动在单个核心上并行执行这两次加法:
let t1 = a + b;
let t2 = c + d;
// 返回结果:
t1 + t2
}
```
然而,由 `if`/`while`/`for` 表达式产生的分支会带来问题:你的代码可能会走某一条路径,也可能走另一条。面对两个选择,CPU 应该尝试并行执行哪一组未来可能的指令?
```rust
fn maybe_add(
a: i64, b: i64, c: i64, d: i64,
add: bool
) -> i64 {
let t1 = a + b;
// CPU 应该与 `a + b` 并行执行这两个分支中的哪一个?
let t2 = if add { c + d } else { c * d };
t1 + t2
}
```
为了确保快速执行,CPU 有一个分支预测器,它通过启发方式选择哪个分支并行执行。如果猜对了,代码运行更快。如果猜错了,CPU 最终会注意到,撤销错误的工作,然后执行正确的分支……这意味着代码变慢。在某些情况下,慢非常多。
不幸的是,上面二分搜索算法对于给定的输入数据非常不可预测。回忆一下,桶边界被选中使得输入值均匀分布在所有桶中。这意味着在二分搜索中选择向左还是向右是完全不可预测的:
```rust
// CPU 无法可靠地猜测要走哪条路径:
if less_than(boundaries[middle], value) {
min_idx = middle + 1;
} else {
max_idx = middle;
}
```
同样,循环次数也可能变化:
```rust
// 这个循环会继续多少次?根据所使用的特定数据,无法知道。
while min_idx < max_idx {
// ...
}
```
为了验证这个假设,我可以使用 CPU 的硬件计数器(通过 Python 的 [py-perf-event](https://github.com/pythonspeed/py-perf-event/) 暴露)来测量运行代码执行了多少分支,以及其中有多少被错误预测。以下是我将使用的输入:
```python
import numpy as np
from numba import jit
# 0 到 1 之间的值:
DATA = np.random.random(1_000_000)
# 均匀间隔的桶边界:
BOUNDARIES = np.linspace(0.0, 1.0, 255)[1:-1]
```
以下是运行代码的结果:
| 代码 | 耗时(微秒) | 分支指令数 | 分支误预测百分比 |
|------|-------------|-----------|-----------------|
| `bucketize_classic(DATA, BOUNDARIES)` | 45,870.2 | 26,997,038 | 16.6% |
16% 的分支被错误预测并不理想,而且每个值对应的分支数量也相当多。`DATA` 有 1,000,000 个值,所以总分支 2700 万意味着每个值 27 个分支。
## 转向无分支执行
我将消除这两种不可预测分支的来源。对于 `while` 循环迭代次数,我改为固定迭代次数,即桶数量取对数底 2。对于某些值,如果之前可以在一次或两次迭代中找到桶,这可能会多做一点工作,但避免了分支误预测,速度上的节省会弥补这一点。
对于 `if` 表达式,当前代码有时设置一个变量,有时设置另一个。我将用总是设置同一个变量的代码替换它,即使它没有被改变。然后,我将使用 Rust 的 `std::hint::select_unpredictable()`,它告诉编译器尽可能避免发出分支。通常 CPU 有特殊的指令可以在没有分支情况下条件性地选择两个值之一。
以下是新的无分支版本:
```rust
use std::hint::select_unpredictable;
fn bucketize_branchless_impl(
arr: &[f64],
boundaries: &[f64],
) -> Vec<usize> {
let size = boundaries.len();
let n_iterations = (size as f64).log2().ceil() as usize;
let mut result = Vec::with_capacity(arr.len());
for value in arr {
let mut left = 0;
let mut remaining_size = size;
// 无分支二分搜索:不断将搜索区域减半。
for _ in 0..n_iterations {
let half = remaining_size / 2;
let middle = left + half;
// 条件操作:如果 bool 为真,则取 a,否则取 b。
// 希望编译器生成无实际分支的指令。
left = select_unpredictable(
less_than(boundaries[middle], value),
middle,
left,
);
remaining_size -= half;
}
// 修复与原始算法相差 1 的问题。
left = select_unpredictable(
less_than(boundaries[left], value),
left + 1,
left,
);
result.push(left);
}
result
}
```
以下是两个版本性能的比较:
| 代码 | 耗时(微秒) | CPU 指令数 | 分支指令数 | 分支误预测百分比 | IPC |
|------|-------------|-----------|-----------|-----------------|-----|
| `bucketize_classic(DATA, BOUNDARIES)` | 45,810.2 | 184,909,902 | 26,997,064 | 16.6% | 1.1 |
| `bucketize_branchless(DATA, BOUNDARIES)` | 13,197.8 🏆 | 188,101,254 | 19,020,569 | 0.0% | 4.0 |
注意:
1. 新版本有更少的分支:每个值 19 个,而不是 27 个。
2. 分支误预测完全消失。
3. IPC(每周期指令数)大大提高。这衡量 CPU 并行执行多条指令的能力,越高越好:由于分支更少更可预测,CPU 现在可以利用更多的指令并行性。
结果是一个快得多的实现,尽管它使用了稍微多一点的 CPU 指令。
## 消除分支和多余工作
改进了代码的“机械同理心”后,我现在有机会回到其他速度来源:调整编译,并通过做更少的工作使代码在算法上更高效。
首先,无分支代码在输入数组的每个值上仍然有 19 个分支。其中 7 或 8 个来自 `for _ in 0..n_iterations`:判断 `for` 循环是否继续或结束需要一个分支,而 `n_iterations` 大约是 255 的对数。但其余的分支来自哪里?
这些额外的分支是因为 Rust 编译器会对 vector 和数组访问添加边界检查。每次索引读取和写入都会检查索引是否越界,以确保内存安全:对长度为 10 的 vector 写入索引 100 可能导致程序崩溃或内存损坏。边界检查是有好处的,因为它们能捕获 bug,但不好之处在于它们增加了计算量。
因此,我将使用 Rust 的 `unsafe` 代码调用像 `get_unchecked()` 这样的 API,跳过边界检查,这意味着现在由我负责确保代码仍然正确。
其次,注意到 `half`/`remaining_size` 的计算在每次二分搜索中完全相同。这是不必要的计算。所以我将尝试预先计算这些值,减少计算工作量。
以下的结果代码:
```rust
fn bucketize_branchless2_impl(
arr: &[f64],
boundaries: &[f64],
) -> Vec<usize> {
let size = boundaries.len();
let n_iterations = (size as f64).log2().ceil() as usize;
// 预先计算减半值。创建一个大小为 n_iterations 的 Vec,初始化为 0:
let mut halves = vec![0usize; n_iterations];
let mut remaining_size = size;
for i in 0..n_iterations {
let half = remaining_size / 2;
halves[i] = half;
remaining_size -= half;
}
// 告诉 Rust(以及读者)halves 不再可变:
let halves = halves;
// 安全性检查:后续不安全操作依赖于 halves 之和小于 boundaries 的长度。
assert!(
halves.iter().copied().sum::<usize>() < boundaries.len()
);
let mut result = Vec::with_capacity(arr.len());
for value in arr {
let mut left = 0;
for half in &halves {
let middle = left + half;
left = select_unpredictable(
// 安全性:`left` 和 `middle` 最多为 `halves` 之和,
// 而上述代码断言该和小于 `boundaries` 的长度。
less_than(
unsafe { *boundaries.get_unchecked(middle) },
value,
),
middle,
left,
);
}
left = select_unpredictable(
// 安全性:见上文注释。
less_than(
unsafe { *boundaries.get_unchecked(left) },
value,
),
left + 1,
left,
);
result.push(left);
}
result
}
```
以下是其性能:
| 代码 | 耗时(微秒) | CPU 指令数 | 分支指令数 | 分支误预测百分比 | IPC |
|------|-------------|-----------|-----------|-----------------|-----|
| `bucketize_branchless(DATA, BOUNDARIES)` | 13,236.3 | 188,101,364 | 19,020,585 | 0.0% | 4.0 |
| `bucketize_branchless2(DATA, BOUNDARIES)` | 10,042.8 🏆 | 121,101,769 | 8,020,646 | 0.0% | 3.4 |
正如预期,现在分支更少了。总的来说,指令数更少(好),IPC 略有下降(坏),最重要的是代码更快。
## 自动向量化与 SIMD
最后,我将回到“机械同理心”。值得注意的是,输入数组中的不同值基本上执行相同的操作。在我们的原始版本中,内部迭代次数不同,但现在每个值执行相同次数的迭代。如果内存中的多个值执行相同的操作,专门的 SIMD CPU 指令可以发挥作用:全称是“单指令多数据”(SIMD),它们专为并行处理数据而设计。你可以手动使用它们,也可以让 Rust 编译器“自动向量化”你的代码,即如果它认为有用,就自动生成 SIMD 指令。
为了实现这一点,我将做两件事。首先,我告诉 Rust [它不再是 2004 年了](https://en.wikipedia.org/wiki/X86-64#Microarchitecture_levels),具体来说,它可以生成为过去约 10 年的现代硬件设计的 CPU 指令(x86-64 机器)。这意味着可以访问更多种类的 SIMD CPU 指令,代价是生成的机器码不能在旧计算机上运行。具体做法是:
```bash
$ export RUSTFLAGS="-C target-cpu=x86-64-v3"
```
然后,我重构代码,使得对 `halves` 的迭代成为**外层**循环,而对 `arr` 中值的迭代成为**内层**循环;之前是相反的。这样一来,相邻数据上显然有更多相同的操作,希望有助于编译器找出将这段代码转化为 SIMD 指令的方法:
```rust
fn bucketize_branchless3_impl(
arr: &[f64],
boundaries: &[f64],
) -> Vec<usize> {
let size = boundaries.len();
let n_iterations = (size as f64).log2().ceil() as usize;
let mut halves = vec![0usize; n_iterations];
let mut remaining_size = size;
for i in 0..n_iterations {
let half = remaining_size >> 1;
halves[i] = half;
remaining_size -= half;
}
let halves = halves;
assert!(
halves.iter().copied().sum::<usize>() < boundaries.len()
);
let mut result = Vec::with_capacity(arr.len());
// 按 16 个一组迭代:
for chunk in arr.chunks(16) {
// 创建一个长度为 16 的数组(类似 Vec 但固定长度且在栈上),初始化为 0:
let mut lefts = [0; 16];
// 不再...
相似文章
在4 GB的数据堆中找针:Go语言性能从0.75 GB/s提升到49 GB/s
一位开发者详细介绍了将Go文件搜索从0.75 GB/s优化到49 GB/s的过程,利用SIMD等技术并理解内存层次结构,包括Go 1.26新推出的`simd/archsimd`包。
Rust 中的安全 SIMD,即使内部也安全
Rust 的 SIMD 抽象现在允许在不使用 unsafe 代码的情况下安全使用,这得益于 Rust 1.87 引入的 CPU 特性令牌,从而实现了简洁且可移植的向量操作。
比较向量搜索库
对向量搜索库(Faiss、Scann、Usearch)进行基准测试,涵盖从500到100万样本的数据集大小,评估速度、内存使用和精确度,并提供结果和代码。
优化模型以快速进行代码生成(8分钟阅读)
Morph LLC描述了三种关键技术——基于编码输出训练投机模型、在廉价GPU上自动搜索内核、以及编写自定义互连——以大幅加速像Qwen和DeepSeek这样的开放模型在编码代理工作负载上的运行,实现了最高3倍的投机解码加速,并在7000美元的GPU上达到97-162 tok/s。
优化 #[sqlx::test] 的重建时间
一位 Rust 开发者对 SQLx 测试的增量重建时间进行了性能分析和优化,识别了调试信息生成和过程宏开销等瓶颈,并提出了加速测试编译的改进方案。