借助 Rust 新 API 实现更快的浮点运算
摘要
Rust 1.98 引入了一个新 API,通过允许更激进的编译器优化来加速浮点运算,同时仍让开发者控制舍入误差。
<p><a href="https://lobste.rs/s/jnznnu/faster_floating_point_math_with_rust_s_new">评论</a></p>
查看缓存全文
缓存时间: 2026/08/03 01:33
# 使用 Rust 新 API 实现更快的浮点数学运算
来源:https://pythonspeed.com/articles/faster-float-math-rust
浮点数学运算通常比整数数学运算慢,因为编译器在优化你的代码时采取了保守策略。虽然其他一些编程语言早有某种解决方案,但直到现在,Rust 还没有一种稳定的好办法来应对这一限制。不过,从 1.98 版本开始,Rust 将允许告诉编译器可以进一步优化你的代码——但带有额外的控制,以便你仍然可以编写舍入误差极小的数值算法。
在本文中,你将了解到:
- 为什么默认情况下编译器对浮点数学运算的优化不如对整数数学运算那样激进。
- Rust 用来解决这一限制的新 API。
- 使用这个新 API 的例子、它带来的速度影响,以及如何控制它的使用范围。
## 对整数求和很快
我先从一个整数示例开始,作为某种可能达到的性能基线。
为了获得最快速的代码生成,我告诉 Rust:现在已经不是 2004 年了(https://en.wikipedia.org/wiki/X86-64#Microarchitecture_levels),它可以生成需要现代硬件支持的 CPU 指令,也就是大约最近十年的 x86-64 机器。具体来说,本文中的所有代码都使用 `RUSTFLAGS="-C target-cpu=x86-64-v3"` 编译。(为了获得最大兼容性,在实际使用中,你可以为较旧的计算机提供备用实现。)
下面是一个 Rust 函数,用于对 int64 数字切片求和:
``
fn naive_sum_i64(values: &[i64]) -> i64 {
let mut total = 0;
for value in values {
total += value;
}
total
}
``
> 我将省略把这个函数暴露给 Python 的代码,但它是之前文章中(https://pythonspeed.com/articles/branchless-binary-search/)的 Rust/Python 代码的一个变体。
为了进行基准测试,我将在 NumPy 中创建一个整数数组:
``
import numpy as np
DATA_INT = np.ones((1_000_000,), dtype=np.int64)
assert naive_sum_i64(DATA_INT) == 1_000_000
``
现在我可以测量对这个数组求和的速度:
代码➘ 耗时(微秒)➘ 每个值的 CPU 指令数`naive\_sum\_i64\(DATA\_INT\)`168\.10\.5➘ 数字越低越好
每个值仅需 0.5 条 CPU 指令!这怎么可能?
编译器很可能使用了专门的单指令多数据(SIMD)CPU 指令,这些指令可以同时对多个值进行批量操作。我这里使用的 i7-12700K CPU 具有 256 位 SIMD 指令,这意味着它可以一次对四个 64 位整数执行某些特定操作。如果有一条专门的 SIMD 求和 CPU 指令,CPU 只需要循环 250,000 次,然后在每次迭代中求和 4 个整数。
实际上:
代码➘ 耗时(微秒)➘ CPU 指令256 位 SIMD 整数指令`naive\_sum\_i64\(DATA\_INT\)`156\.8521,280250,003➘ 数字越低越好
简而言之,通过使用专门的 SIMD 指令,我的 CPU 可以非常快速地求和整数。
## 对浮点数求和很慢?!
那么浮点数呢——它们也很快吗?
同样,我创建一百万个浮点数值:
``
# Array of 1M float64 values between 0 and 1.
DATA = np.random.random((1_000_000,))
``
我将实现一个简单的浮点求和函数:
``
fn naive_sum(values: &[f64]) -> f64 {
let mut total = 0.0;
for value in values {
total += value;
}
total
}
``
并比较对整数和浮点数求和时的性能:
代码➘ 耗时(微秒)➘ CPU 指令256 位 SIMD 整数指令256 位 SIMD 浮点指令`naive\_sum\_i64\(DATA\_INT\)`151\.9521,214250,0030`naive\_sum\(DATA\)`595\.21,458,26900➘ 数字越低越好
浮点求和比整数求和慢得多,而且编译器没有使用 SIMD 浮点操作。为什么会这样?
### 浮点运算不满足结合律
与大多数编译器一样,当 Rust 通过发布模式编译你的代码时,它会优化你的代码,以各种方式转换它,(希望)使其更快。但编译器在这样做时有一个承诺:优化后的代码的行为将与未优化代码*完全一致*。
如果我把三个整数 `a`、`b` 和 `c` 相加,`a + (b + c) == (a + b) + c`。这给了编译器充足的余地来优化代码的执行方式,例如使用可能稍微改变加法顺序的 SIMD 操作。
浮点数则不同。例如,由于浮点数跨越了从极小到极大的数值范围,将一个足够大的数与一个足够小的数相加,结果仍然是那个大的数:
``
print(
"Does adding a small number do nothing?",
1e16 + 1.0 == 1e16
)
``
``
Does adding a small number do nothing? True
``
更广泛地说,对于浮点数,`a + (b + c)` 并不总是等同于 `(a + b) + c`,至少当你连续相加多个数字时是这样。假设我有一个以 `1e16` 开头、后面跟着许多 `1.0` 值的数组,另一个数组则相反。对这些数组求和会得到不同的结果:
``
import math
HIGH_VALUE_FIRST = np.ones((1_000_000,), dtype=np.float64)
HIGH_VALUE_FIRST[0] = 1e16
HIGH_VALUE_LAST = np.ones((1_000_000,), dtype=np.float64)
HIGH_VALUE_LAST[-1] = 1e16
print(
"Is the sum the same?",
naive_sum(HIGH_VALUE_FIRST) == naive_sum(HIGH_VALUE_LAST)
)
``
``
Is the sum the same? False
``
因为求和的顺序会影响结果,编译器理所应当地认为,我要求这个特定顺序是有原因的。因此,编译器*不会重新排列这些操作*。它也不会应用任何其他可能改变结果的优化,即使生成的代码会更慢。
## Rust 新的代数运算符:告诉编译器何时可以灵活处理
虽然编译器采取保守策略是正确的默认行为,但有时你作为程序员知道重新排列操作不是问题。在这种情况下,能够告诉编译器:在*这里*它不应该改变操作顺序,而在*那里*实际上没问题,这会很好。
从 Rust 1.98 开始,有一个新功能可以实现这一点。除了通常的浮点数算术运算之外,还有一组新的所谓“代数”算术运算符(https://doc.rust-lang.org/std/primitive.f32.html#algebraic-operators),根据文档,它们“允许编译器利用实数的所有常见代数性质来优化浮点运算”,包括改变操作顺序。
> Rust 1.98 的发布日期是 2026 年 8 月 20 日。由于我是在那之前写的这篇文章,我使用了 `"beta"` 渠道来运行本文中的代码,该渠道包含相同的功能。
### 示例:优化的成对求和
让我们看看这些运算符的实际效果,以及它们如何让代码变得更快。
由于浮点求和可能出现意外行为——还记得 `1e16 + 1.0 == 1e16`——如果你想要对大量浮点数求和,有多种算法可以用来最小化由此产生的舍入误差。它们之间的权衡通常在于速度和精度之间的取舍。`numpy.sum()`(https://numpy.org/doc/stable/reference/generated/numpy.sum.html)主要使用一种称为成对求和(https://en.wikipedia.org/wiki/Pairwise_summation)的算法,它在速度和减少累积误差之间达到了良好的平衡。当然,相比我上面用 `naive_sum()` 那样按顺序一个接一个地相加浮点数,它的误差累积要小得多。
成对求和的基本思想是将数组分成两半,使用相同的算法递归地对每一半求和,然后将得到的两个浮点数相加。当数组大小低于某个阈值时——在 NumPy 中是 128——就按普通方式求和……而且可以任意顺序进行。
让我们在 Rust 中实现这个算法。在相加最顶层的两个浮点数时,我将使用普通加法,这样编译器就无法重排操作或做任何改变输出的操作。一旦函数到达进行普通求和的阈值,我将切换到代数加法,因为此时我不关心加法的顺序,只想要速度。
``
fn pairwise_sum(values: &[f64]) -> f64 {
let n = values.len();
if n > 128 {
// Precise addition of two recursive applications
// of the algorithm:
let half = n / 2;
pairwise_sum(&values[0..half])
+ pairwise_sum(&values[half..n])
} else {
// 😎 Normal addition, where the order doesn't matter
// as far as the algorithm is concerned. Thus using a
// algebraic add is fine—and that gives the compiler
// permission to optimize aggressively.
let mut total: f64 = 0.0;
for value in values {
total = total.algebraic_add(*value);
}
total
}
}
``
这种成对求和的实现速度很快,实际上比 NumPy 的实现还要快:
代码➘ 耗时(微秒)➘ CPU 指令256 位 SIMD 浮点指令`naive\_sum\(DATA\)`563\.11,458,2790`np\.sum\(DATA\)`190\.72,191,7670`pairwise\_sum\(DATA\)`144\.5🏆1,298,028270,336➘ 数字越低越好
而且它与 NumPy 的实现具有相同的(大致)精度,因为它实现了相同的算法:
``
from math import fsum
# fsum() uses an accurate summation algorithm that gets rid of any
# avoidable floating-point error.
assert fsum(HIGH_VALUE_FIRST) == fsum(HIGH_VALUE_LAST)
correct_sum = fsum(HIGH_VALUE_FIRST)
print(
"naive_sum() error: ",
naive_sum(HIGH_VALUE_FIRST) - correct_sum
)
print(
"np.sum() error: ",
np.sum(HIGH_VALUE_FIRST) - correct_sum
)
print(
"pairwise_sum() error:",
pairwise_sum(HIGH_VALUE_FIRST) - correct_sum
)
``
``
naive_sum() error: -1000000.0
np.sum() error: -14.0
pairwise_sum() error: -6.0
``
> `np.sum()` 和 `pairwise_sum()` 之间的误差差异并不具有实际意义,只是运气问题;关键点在于,与 `naive_sum()` 相比,它们的误差在数量级上同样都很小。
额外的好处是,`pairwise_sum()` 用来鼓励生成 SIMD 代码的结构比 NumPy 的实现(https://github.com/numpy/numpy/blob/841147dd0bde1ec13ceb89fd77a9265d8b9b5918/numpy/_core/src/umath/loops_utils.h.src#L80-L145)要简洁得多。
## 另一个示例:平方差之和
除了用代数运算符做加法,你还可以做更多事情。在下面的例子中,我将计算两个数组之间的平方差之和。
我需要两个数组:
``
DATA1 = np.random.random((1_000_000,))
DATA2 = np.random.random((1_000_000,))
``
这是使用普通算术运算符的实现:
``
fn ssd_normal(arr1: &[f64], arr2: &[f64]) -> f64 {
assert_eq!(arr1.len(), arr2.len());
let mut total = 0.0;
for (val1, val2) in arr1.iter().zip(arr2) {
total += (val1 - val2).powi(2);
}
total
}
``
这是使用代数运算符的版本;我不太确定 `f64.powi()` 是否会使用代数运算符,所以我显式地使用了它们:
``
fn ssd_optimized(arr1: &[f64], arr2: &[f64]) -> f64 {
assert_eq!(arr1.len(), arr2.len());
let mut total: f64 = 0.0;
for (val1, val2) in arr1.iter().zip(arr2) {
// 😎 Algebraic operations to allow more compiler
// optimizations:
let diff = val1.algebraic_sub(*val2);
let squared_diff = diff.algebraic_mul(diff);
total = total.algebraic_add(squared_diff);
}
total
}
``
首先,快速测试一下以确保结果相似:
``
print(ssd_normal(DATA1, DATA2))
print(ssd_optimized(DATA1, DATA2))
``
``
166770.0055951995
166770.00559520238
``
接下来,我测量性能:
代码➘ 耗时(微秒)➘ 每个值的 CPU 指令数`ssd\_normal\(DATA1, DATA2\)`628\.74\.5`ssd\_optimized\(DATA1, DATA2\)`371\.1🏆1\.0➘ 数字越低越好
使用代数运算使得编译器能够生成在我的电脑上运行速度快一倍的代码。
## 快去加速你的数值代码吧!
成对求和算法是一个很好的例子,说明你为什么需要两种运算符:严格的和宽松的。
- 如果你只使用严格的按序操作,代码会更慢。
- 如果你只使用宽松的“随意优化”代数操作,编译器可能会将算法优化到不存在,从而失去算法所追求的精度。
我上面展示的实现,在算法的不同部分同时使用普通加法——为了精度——和代数加法——为了速度,从而获益。
如果你正在使用 Rust 编写数值代码,你的代码也可能从中受益——一旦你可以使用 Rust 1.98,就去试一试吧。如果你还没有使用 Rust,这又是一个切换到 Rust 的好理由。
相似文章
Rust 中的安全 SIMD,即使内部也安全
Rust 的 SIMD 抽象现在允许在不使用 unsafe 代码的情况下安全使用,这得益于 Rust 1.87 引入的 CPU 特性令牌,从而实现了简洁且可移植的向量操作。
Rust Decimal 库的比较与基准测试
一篇详细的技术文章,比较和基准测试了多种 Rust Decimal 库,涵盖了定点数与浮点数、固定精度与任意精度设计。
RISC-V 与浮点运算
关于 RISC-V 架构浮点功能及更新的报告。
中间浮点精度
本文探讨了C++代码中的中间浮点精度如何依赖于编译器设置、CPU标志和架构,尤其是在x87 FPU上,以及这如何影响性能和计算结果。
使用 Rust 表达式插件扩展 Polars
本文解释了 fenic 为何以及如何使用 Rust 表达式插件来扩展 Polars,以在引擎原生执行文本操作(分块、提示模板化、模糊匹配等),从而避免 Python UDF 的性能和组合问题。