【RustyML入门】6.3. 并行归约

6.3. 并行归约

rustyml::math::reduction 不提供 sum()mean() 函数。它提供 2 个泛型折叠组合子,

det_reduce

det_reduce_range

外加 1 个常量 DET_REDUCE_BLOCK

这个 crate 里所有的并行归约都建立在这 2 个函数之上。例子包括线性回归里的误差平方和、

clip-by-global-norm 用到的全局梯度范数,以及标准化里的一遍式 Welford 矩。其他例子还包括

k-means 的惯性和逻辑回归的对数损失。这两个函数要解决的是普通并行求和解决不了的一个问题:

得到一个不依赖线程数的结果。

6.3.1. 模块暴露了什么

公开接口一共有 3 个条目。这个模块躲在 math feature 后面。其他每个 feature

machine_learningneural_networkutilsmetrics)都会拉入 math feature。所以只要

RustyML 能编译,这 3 个条目就总是可用(见

1.2. 安装与Feature配置)。

条目 签名(省略约束) 作用
DET_REDUCE_BLOCK pub const DET_REDUCE_BLOCK: usize = 16_384 固定块大小(以元素为单位),决定分组方式
det_reduce fn det_reduce<T, A, F, M>(slice: &[T], parallel: bool, fold_block: F, merge: M, identity: A) -> A 按固定块折叠一个切片
det_reduce_range fn det_reduce_range<A, F, M>(n: usize, parallel: bool, fold_block: F, merge: M, identity: A) -> A 按固定块折叠索引区间 0..n

完整的约束如下所示。围绕它们的编译器报错信息,如果不了解这些约束,会很难读懂。

rust,ignore 复制代码
pub fn det_reduce<T, A, F, M>(slice: &[T], parallel: bool, fold_block: F, merge: M, identity: A) -> A
where
    T: Sync,
    A: Send,
    F: Fn(&[T]) -> A + Sync + Send,   // 对 1 个块做串行折叠
    M: Fn(A, A) -> A,                 // 合并 2 个部分结果
{ /* ... */ }

fold_block 把 1 个块归约成累加器类型 A 的一个部分结果。merge 合并 2 个部分结果。

identity 是空输入时返回的值,同时也是最终合并的种子。

fold_block 必须是 Sync + Send,因为 rayon 可能在任意 worker 上调用它。哪怕传入的是

parallel = false,这条约束依然成立,因为约束加在类型上,而不是那个 flag 上。merge 不需要

这两条约束中的任何一条,因为它永远在 1 个线程上按块顺序运行。

A 可以是任何 Send 的类型:一个标量、一个像 (sum, sum_of_squares) 这样的元组、一个

Welford 三元组,或者一个按桶分组的和数组。

det_reduce_range 在索引区间而非切片上运行同一套算法。当归约需要一次读多个数组、或者按行

索引矩阵时,用它。它的 fold_block 收到的是 Range<usize>,而不是 &[T]

这个模块特意没有提供 sum 包装函数。在并行阈值之下,一行 slice.iter().sum() 已经是对的

工具。在阈值之上,调用方几乎总想把一个 map 融进同一遍扫描里。常见的例子包括平方、exp

或者某种距离函数。融合胜过先构建一个中间数组。把折叠交给你,而不是给一个写死的归约,正是

为了把这种融合留在调用点。

6.3.2. 为什么朴素的并行求和不确定

浮点加法不满足结合律。(a + b) + ca + (b + c) 可能舍入到不同的 f64 值。这不是硬件的

bug。这是每次运算后都舍入到 53 位这一定义本身。只要求和在 1 个线程上从左到右进行,次序就是

固定的,结果也就可复现。并行执行会去掉这个固定的次序。

一个裸的 slice.par_iter().sum::<f64>(),或者 fold().reduce(),会自适应地切分工作。rayon

的工作窃取调度器决定哪个 worker 折叠哪段子区间。它还决定这些部分和以什么次序合并。

在一台有 4 个空闲核心的机器上跑一次,会得到 1 种分组。同样的输入用 RAYON_NUM_THREADS=1

跑,会得到另一种分组。在一台繁忙的 16 核机器上跑 2 次,两次结果也可能对不上,因为某个线程

在不同的时刻被抢占了。这些结果中的每一个都是同一组数字的正确求和。它们只是舍入方式不同,

通常差在最后几个 ULP 上。

对很多数值代码来说,这种抖动无伤大雅。但对一个机器学习库来说,它是腐蚀性的。一个低位来回

晃动的损失值,会让早停检查在不同运行里于不同迭代触发。一个依赖线程数的梯度范数,会让

clip-by-global-norm 在每次运行时剪裁得略有出入。它还可能让同一份拟合的 2 次运行得到 2 个

不同的模型。

可复现性在 RustyML 里是头等承诺(见

7.1. 可复现性与随机种子)。一个依赖调度器的归约

会打破这个承诺,无论你把 RNG 的种子设置得多小心。

6.3.3. 分块算法

解法是把分组从调度器手里拿走,固定到一个常量上。det_reduce 把输入切成固定

DET_REDUCE_BLOCK 个元素的块。它用你的 fold_block 串行折叠每个块。它按块顺序收集每块的

部分结果。然后它用你的 merge 从左到右合并这些结果。并行和串行这 2 条路径,唯一的差别在于

块怎么运行:

rust,ignore 复制代码
if parallel {
    let parts: Vec<A> = slice.par_chunks(DET_REDUCE_BLOCK).map(fold_block).collect();
    parts.into_iter().fold(identity, merge)          // 按块顺序合并
} else {
    slice.chunks(DET_REDUCE_BLOCK).map(fold_block).fold(identity, merge)
}

关键细节在于,rayon 的 par_chunks(...).collect::<Vec<_>>() 是一个有序(indexed)并行

迭代器。无论工作窃取把这些块怎样分派到各个线程上,收集回来的 Vec 都会按原本的块顺序

返回。所以这棵归约树,是输入长度和 DET_REDUCE_BLOCK 的纯函数。这棵树决定了哪些元素落进

哪个块、块又以什么次序合并。这棵树不依赖线程数,不依赖调度,也不依赖 parallel 这个

flag。两条路径折叠的都是同样的 16 384 个元素的块,次序相同,合并方式也相同。

det_reduce_rangen.div_ceil(DET_REDUCE_BLOCK) 个索引块上做完全相同的事。第 b

覆盖的区间是 b * BLOCK .. ((b + 1) * BLOCK).min(n)

这就让 parallel 参数变成一个纯粹的性能提示。它绝不改变哪些数字以什么次序相加。它只决定

这些块是跑在 rayon 线程池上,还是跑在一个普通的顺序循环里。crate 自己的测试用逐位 ==

比较(不是某个 epsilon)来检查这一点,覆盖了空输入、不足一块、恰好一块,以及参差不齐的

多块长度。所以在给定的一次构建里,这 2 条路径逐位相同。改动 RAYON_NUM_THREADS 无法改变

结果。

模块文档仍然指出,结果并不总是逐位可复现的。那条附注针对的是跨机器和跨构建的差异。比如

不同的 libm sin、被不同 target CPU 拨动的 FMA 收缩,或者你自己 fold_block 里不同的

舍入方式。它不涉及线程数,线程数已经被分块彻底钉死。

DET_REDUCE_BLOCK 取 16 384,是因为这个大小落在实测吞吐量平台期附近。在一个 420 万元素的

f64 平方和基准上,实测的加速比从 1024 元素块的约 14 倍开始上升。它在 32 768 元素块处达到

峰值,约 18 倍。随后在 65 536 元素块处回落到约 15 倍。在 262 144 元素块处进一步跌到约 11

倍。此时剩下的块太少,无法在核心之间均衡负载。

等价的 f32 基准(累加器同样是 f64)峰值更高,在 65 536 元素块处约为 21 倍。16 384 同时

接近这两个峰值:比 f64 峰值低约 3%,比 f32 峰值低约 8%。它对两种元素类型都表现不错,

不需要为每种类型单独设一个常量。

这个常量数的是元素,不是字节,并且被所有元素类型共用。一个 16 384 个 f32 值的块是

64 KB。一个 16 384 个 f64 值的块是 128 KB。这两个大小都稳稳落在各自元素类型的平台期上。

块大小定义了分组方式,所以它是可复现性接口的一部分。改动它,会在低位上改变这个(仍然确定

的)结果。正因如此,DET_REDUCE_BLOCK 是一个 const,而不是一个运行时旋钮。

6.3.4. 精度是副产品,不是目标

选择分块是为了确定性。它顺带也提升了精度,算是一个副产品。这一点对两条路径都成立。串行

路径同样会切块,所以即便 parallel = false,它也不是对整个数组的朴素左折叠。

从左到右把 n 个浮点数相加,最坏情况的舍入误差随 n 线性增长。它大致遵循

(n - 1) * eps * S,其中 eps 是机器精度(machine epsilon),S 是所有输入绝对值之和。

det_reduce 用的是一个两级方案。每个块串行折叠 b = 16 384 项。然后来自

ceil(n / b) 个块的部分结果再串行折叠。误差上界因此大致变成 (b + n / b) * eps * S

对一个 420 万元素的求和,这个上界大约是 (16 384 + 256) * eps。同一个求和的朴素上界大约是

4.2 million * eps(约 420 万乘以 eps)。分块后的上界在最坏情况下大约收紧了 250 倍。无论

这些块是并行跑的还是串行跑的,这个改进都一样。

这个方案是扁平分块加一次串行合并,不是完整的成对(log n)求和树。实践中,累加器的位宽比

树的形状更要紧。global_grad_normf32 梯度归约进 fold_block 内部的一个 f64

累加器,所以这个平方梯度和从头到尾都留在 f64 里。这个累加器的选择,对精度的提升比分块

本身更大。分块设定的是确定性的底线。当精度本身才是关心的问题时,宽累加器才是正确的选择。

6.3.5. 并行路径何时开启

det_reduce 不决定 parallel。由调用方传入这个 flag。在这个 crate 内部,一道校准过的尺寸

阈值产生这个布尔值。在大约 1 块以下,没什么可并行的。一个短于 16 384 元素的输入就是单独

一块,把它 fork 到 rayon 上只会平添 join 开销。

这些阈值住在 rustyml::tuning::reduction 里。每道阈值按成本类别共享,而不是按调用点各自

定义:

阈值(tuning::reduction 中的 getter) 默认值 守护对象
get_sum_f64 262 144 f64 求和类归约(SSE、Welford 矩、k-means 惯性)
get_sq_sum_f32 65 536 clip-by-global-norm 用到的 f32f64 平方和
get_scan_f64 262 144 f64 逐行扫描(arg-min、距离扫描)
get_exp_reduce 32 768 exp 密集的逻辑回归对数损失归约

每个调用点都遵循同样的写法:拿一个工作量指标去跟 gate() 比较,把比较结果当作 flag 传进去。

大多数调用点用 slice.len() 作为这个指标。k-means 的质心累加则用 n_samples * n_features

因为这个乘积才是它那个分块折叠真正遍历的元素数。阈值移动的是切换点,但从不触碰正确性,因为

分块折叠在切换点两侧给的是同一个答案。

每道阈值都配有一个对应的 setter,比如 set_sum_f64set_sq_sum_f32,用来为不同硬件调整

切换点。7.3. 性能调优与并行 讲了具体机制和校准过程。

exp 归约那道阈值设得最低,是 32 768,因为那里每个元素都要付一次 exp 和一次 ln。并行在那里

比普通加法更早摊平成本。

6.3.6. 在你自己的代码里使用它们

最小的调用,是在一个 Vec<f64> 上算融合的平方和,并保持串行:

rust 复制代码
use rustyml::math::reduction::det_reduce;

fn main() {
    let data: Vec<f64> = (0..1_000).map(|i| (i as f64).sin()).collect();
    let sum_sq = det_reduce(
        &data,
        false, // 性能提示:输入小,保持串行
        |block| block.iter().map(|&x| x * x).sum::<f64>(),
        |a, b| a + b,
        0.0,
    );
    println!("sum of squares = {sum_sq}");
}

这里有 2 处需要留意。第一,det_reduce 接收的是 &[T],所以数据必须是一段连续的切片。

一个 ndarray 数组只能通过 as_slice() 给出连续切片,而对于非标准布局的视图,这个方法会

返回 None。crate 自己的惯用写法,是在连续的快路径上用 det_reduce 归约,其余情况回退到

ndarray 的串行内核。这套写法把 flag 按尺寸类别来门控:

rust 复制代码
use ndarray::Array1;
use rustyml::math::reduction::det_reduce;
use rustyml::tuning::reduction::get_sum_f64;

fn main() {
    let v: Array1<f64> = (0..10_000).map(|i| i as f64).collect();
    let sum = match v.as_slice() {
        Some(slice) => det_reduce(
            slice,
            slice.len() >= get_sum_f64(),
            |block| block.iter().sum::<f64>(),
            |a, b| a + b,
            0.0,
        ),
        None => v.sum(), // 非连续:走 ndarray 的串行折叠
    };
    println!("sum = {sum}");
}

第二,累加器不必是标量。这正是这个模块把折叠交给你、而不是给一个写死的归约的原因。一遍

扫描就能同时返回和与平方和,足够算出均值和方差。用一个元组累加器配一个元组 merge

rust 复制代码
use rustyml::math::reduction::det_reduce;

fn main() {
    let data: Vec<f64> = (0..10_000).map(|i| (i as f64).sin()).collect();
    let (sum, sum_sq) = det_reduce(
        &data,
        false,
        |block| block.iter().fold((0.0f64, 0.0f64), |(s, sq), &x| (s + x, sq + x * x)),
        |(sa, sqa), (sb, sqb)| (sa + sb, sqa + sqb),
        (0.0, 0.0),
    );
    let n = data.len() as f64;
    let mean = sum / n;
    let variance = sum_sq / n - mean * mean;
    println!("mean = {mean}, variance = {variance}");
}

有些归约需要一次读多个数组,比如点积、距离累加,或者逐行挑选。遇到这些情况,用

det_reduce_range,在块内部做索引:

rust 复制代码
use rustyml::math::reduction::det_reduce_range;

fn main() {
    let xs: Vec<f64> = (0..5_000).map(|i| i as f64).collect();
    let ys: Vec<f64> = (0..5_000).map(|i| (i as f64).cos()).collect();
    let dot = det_reduce_range(
        xs.len(),
        false,
        |range| range.map(|i| xs[i] * ys[i]).sum::<f64>(),
        |a, b| a + b,
        0.0,
    );
    println!("dot = {dot}");
}

det_reduce 对比 ndarray 的 .sum()

ndarray 的 .sum().dot().mean() 是串行、单线程的。它们内部的分组方式自成一套,

一般不会和 det_reduce 的分块逐位吻合。对小数组来说,ndarray 是对的选择:写起来更短,

不需要闭包,而且在阈值之下 det_reduce 反正也是串行跑,还要多写一些准备代码。

det_reduce 需要同时满足 3 个条件。缓冲区又大又连续。这个归约需要并行运行。并行结果

需要可复现。ndarray 不提供这种组合。即便开了 ndarray 的 rayon feature,一个裸的并行求和

依然依赖调度器。

当你想把一个 map 融进归约、或者累加比标量更丰富的东西时,det_reduce 同样是趁手的工具。

默认用 ndarray,图个方便,应付小数据。等到你正要写下一个 par_iter().sum()、又需要答案

保持稳定的那个点,就换成 det_reduce。互操作的细节见

1.3. 使用ndarray准备数据

4.2. 标准化与归一化 里有一个正是架在这套折叠之上的

真实 Welford 归约。

6.3.7. 跨线程数验证确定性

这个论断可以从 crate 外部验证。下面这个程序在 rayon 路径上归约 200 万个值,并以完整精度

打印出和:

rust 复制代码
use rustyml::math::reduction::det_reduce;

fn main() {
    let data: Vec<f64> = (0..2_000_000).map(|i| (i as f64 * 0.001).sin()).collect();
    let sum = det_reduce(
        &data,
        true, // 强制走并行路径
        |block| block.iter().sum::<f64>(),
        |a, b| a + b,
        0.0,
    );
    // 完整精度打印,这样低位的差异也会显形
    println!("{sum:.17e}");
}

编译一次,然后通过设置 RAYON_NUM_THREADS 在不同线程数下运行它,这个变量给 rayon 的全局

线程池设了上限:

bash 复制代码
RAYON_NUM_THREADS=1 ./target/release/demo
RAYON_NUM_THREADS=2 ./target/release/demo
RAYON_NUM_THREADS=8 ./target/release/demo

3 次运行都打印出完全相同的 17 位尾数。每次运行折叠的都是同样的 16 384 个元素的块,并按

同样的块顺序合并,跟当时有多少个 worker 扛着这些块无关。把 flag 设成 false,输出同样

不会变。把函数体换成 data.par_iter().sum::<f64>() 的版本,表现则不同。在不同的

RAYON_NUM_THREADS 取值下,只要输入够大,那个版本打印出的和就会在最后几位数字上出现

差异。这正是 det_reduce 要堵住的那个失效模式。

相关的数值原语共享同一套确定性纪律。见 6.1. 距离度量

6.2. 矩阵乘法,以及更宏观的 6.0. 数学工具 概览。

相关推荐
richard_first1 小时前
Transformer 与大语言模型:第9章 LayerNorm 归一化层
人工智能·深度学习·机器学习·transformer
Profile排查笔记2 小时前
指纹浏览器哪个好?从 Profile、代理、权限和自动化能力判断是否适合
前端·人工智能·后端·自动化
用户938515635072 小时前
用 AI 结对编程从 0 搭一个"单词后台管理系统":Next.js + Supabase + Drizzle + shadcn/ui 全记录
后端·postgresql·next.js
mqiqe2 小时前
AgentScope Java 2.0 协议集成全景解析:A2A、AG-UI、Agent Protocol 三大开放协议实战指南
java·开发语言·ui
lzhdim2 小时前
12、JavaScript常见的内存泄露问题 - JavaScript学习系列文章
开发语言·前端·javascript·学习·ecmascript
脉动数据行情2 小时前
Java 实现台股 TWSE/TPEx 行情采集(个股 + K 线)
java·开发语言·twse·tpex·台股
爱学习的小邓同学2 小时前
Golang --- (1)第一个Golang程序
开发语言·后端·golang
红尘散仙3 小时前
TypeScript WIT Guest:类型约束与 ComponentizeJS 工程化
rust·typescript·webassembly